test: 补全单元测试
Some checks failed
Go / build (.exe, 386, windows, windows-386) (push) Has been cancelled
Go / build (.exe, amd64, windows, windows-amd64) (push) Has been cancelled
Go / build (.exe, arm64, windows, windows-arm64) (push) Has been cancelled
Go / build (386, freebsd, freebsd-386) (push) Has been cancelled
Go / build (386, linux, linux-386) (push) Has been cancelled
Go / build (386, netbsd, netbsd-386) (push) Has been cancelled
Go / build (386, openbsd, openbsd-386) (push) Has been cancelled
Go / build (386, plan9, plan9-386) (push) Has been cancelled
Go / build (amd64, darwin, darwin-amd64) (push) Has been cancelled
Go / build (amd64, dragonfly, dragonfly-amd64) (push) Has been cancelled
Go / build (amd64, freebsd, freebsd-amd64) (push) Has been cancelled
Go / build (amd64, illumos, illumos-amd64) (push) Has been cancelled
Go / build (amd64, linux, linux-amd64) (push) Has been cancelled
Go / build (amd64, netbsd, netbsd-amd64) (push) Has been cancelled
Go / build (amd64, openbsd, openbsd-amd64) (push) Has been cancelled
Go / build (amd64, plan9, plan9-amd64) (push) Has been cancelled
Go / build (amd64, solaris, solaris-amd64) (push) Has been cancelled
Go / build (arm, 6, linux, linux-armv6) (push) Has been cancelled
Go / build (arm, 7, linux, linux-armv7) (push) Has been cancelled
Go / build (arm, freebsd, freebsd-arm) (push) Has been cancelled
Go / build (arm, netbsd, netbsd-arm) (push) Has been cancelled
Go / build (arm, openbsd, openbsd-arm) (push) Has been cancelled
Go / build (arm, plan9, plan9-arm) (push) Has been cancelled
Go / build (arm64, darwin, darwin-arm64) (push) Has been cancelled
Go / build (arm64, freebsd, freebsd-arm64) (push) Has been cancelled
Go / build (arm64, linux, linux-arm64) (push) Has been cancelled
Go / build (arm64, netbsd, netbsd-arm64) (push) Has been cancelled
Go / build (arm64, openbsd, openbsd-arm64) (push) Has been cancelled
Go / build (loong64, linux, linux-loong64) (push) Has been cancelled
Go / build (mips, linux, linux-mips) (push) Has been cancelled
Go / build (mips64, linux, linux-mips64) (push) Has been cancelled
Go / build (mips64le, linux, linux-mips64le) (push) Has been cancelled
Go / build (mipsle, linux, linux-mipsle) (push) Has been cancelled
Go / build (ppc64, aix, aix-ppc64) (push) Has been cancelled
Go / build (ppc64, linux, linux-ppc64) (push) Has been cancelled
Go / build (ppc64, openbsd, openbsd-ppc64) (push) Has been cancelled
Go / build (ppc64le, linux, linux-ppc64le) (push) Has been cancelled
Go / build (riscv64, freebsd, freebsd-riscv64) (push) Has been cancelled
Go / build (riscv64, linux, linux-riscv64) (push) Has been cancelled
Go / build (riscv64, openbsd, openbsd-riscv64) (push) Has been cancelled
Go / build (s390x, linux, linux-s390x) (push) Has been cancelled
Go / merge-artifacts (push) Has been cancelled
Some checks failed
Go / build (.exe, 386, windows, windows-386) (push) Has been cancelled
Go / build (.exe, amd64, windows, windows-amd64) (push) Has been cancelled
Go / build (.exe, arm64, windows, windows-arm64) (push) Has been cancelled
Go / build (386, freebsd, freebsd-386) (push) Has been cancelled
Go / build (386, linux, linux-386) (push) Has been cancelled
Go / build (386, netbsd, netbsd-386) (push) Has been cancelled
Go / build (386, openbsd, openbsd-386) (push) Has been cancelled
Go / build (386, plan9, plan9-386) (push) Has been cancelled
Go / build (amd64, darwin, darwin-amd64) (push) Has been cancelled
Go / build (amd64, dragonfly, dragonfly-amd64) (push) Has been cancelled
Go / build (amd64, freebsd, freebsd-amd64) (push) Has been cancelled
Go / build (amd64, illumos, illumos-amd64) (push) Has been cancelled
Go / build (amd64, linux, linux-amd64) (push) Has been cancelled
Go / build (amd64, netbsd, netbsd-amd64) (push) Has been cancelled
Go / build (amd64, openbsd, openbsd-amd64) (push) Has been cancelled
Go / build (amd64, plan9, plan9-amd64) (push) Has been cancelled
Go / build (amd64, solaris, solaris-amd64) (push) Has been cancelled
Go / build (arm, 6, linux, linux-armv6) (push) Has been cancelled
Go / build (arm, 7, linux, linux-armv7) (push) Has been cancelled
Go / build (arm, freebsd, freebsd-arm) (push) Has been cancelled
Go / build (arm, netbsd, netbsd-arm) (push) Has been cancelled
Go / build (arm, openbsd, openbsd-arm) (push) Has been cancelled
Go / build (arm, plan9, plan9-arm) (push) Has been cancelled
Go / build (arm64, darwin, darwin-arm64) (push) Has been cancelled
Go / build (arm64, freebsd, freebsd-arm64) (push) Has been cancelled
Go / build (arm64, linux, linux-arm64) (push) Has been cancelled
Go / build (arm64, netbsd, netbsd-arm64) (push) Has been cancelled
Go / build (arm64, openbsd, openbsd-arm64) (push) Has been cancelled
Go / build (loong64, linux, linux-loong64) (push) Has been cancelled
Go / build (mips, linux, linux-mips) (push) Has been cancelled
Go / build (mips64, linux, linux-mips64) (push) Has been cancelled
Go / build (mips64le, linux, linux-mips64le) (push) Has been cancelled
Go / build (mipsle, linux, linux-mipsle) (push) Has been cancelled
Go / build (ppc64, aix, aix-ppc64) (push) Has been cancelled
Go / build (ppc64, linux, linux-ppc64) (push) Has been cancelled
Go / build (ppc64, openbsd, openbsd-ppc64) (push) Has been cancelled
Go / build (ppc64le, linux, linux-ppc64le) (push) Has been cancelled
Go / build (riscv64, freebsd, freebsd-riscv64) (push) Has been cancelled
Go / build (riscv64, linux, linux-riscv64) (push) Has been cancelled
Go / build (riscv64, openbsd, openbsd-riscv64) (push) Has been cancelled
Go / build (s390x, linux, linux-s390x) (push) Has been cancelled
Go / merge-artifacts (push) Has been cancelled
This commit is contained in:
158
cmd/gateway/config_test.go
Normal file
158
cmd/gateway/config_test.go
Normal file
@@ -0,0 +1,158 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestLoadPluginConfig(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
type pluginConfig struct {
|
||||
Enable bool `toml:"enable"`
|
||||
Name string `toml:"name"`
|
||||
Count int `toml:"count"`
|
||||
}
|
||||
|
||||
var got pluginConfig
|
||||
err := loadPluginConfig(map[string]any{
|
||||
"enable": true,
|
||||
"name": "plugin-a",
|
||||
"count": 7,
|
||||
}, &got)
|
||||
if err != nil {
|
||||
t.Fatalf("loadPluginConfig() error = %v", err)
|
||||
}
|
||||
|
||||
want := pluginConfig{Enable: true, Name: "plugin-a", Count: 7}
|
||||
if got != want {
|
||||
t.Fatalf("loadPluginConfig() = %+v, want %+v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadPluginConfigReturnsDecodeError(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
type pluginConfig struct {
|
||||
Count int `toml:"count"`
|
||||
}
|
||||
|
||||
var got pluginConfig
|
||||
if err := loadPluginConfig(map[string]any{"count": "not-an-int"}, &got); err == nil {
|
||||
t.Fatal("loadPluginConfig() error = nil, want error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadConfigReadsTomlAndAppliesSideEffects(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
pidPath := filepath.Join(tmpDir, "gateway.pid")
|
||||
logPath := filepath.Join(tmpDir, "logs", "gateway.log")
|
||||
configFile = filepath.Join(tmpDir, "config.toml")
|
||||
|
||||
toml := fmt.Sprintf(`
|
||||
pid_file = %q
|
||||
|
||||
[log]
|
||||
level = "debug"
|
||||
file = %q
|
||||
|
||||
[tcp]
|
||||
enable = true
|
||||
port = 25565
|
||||
|
||||
[quic]
|
||||
enable = true
|
||||
port = 25566
|
||||
application_protocols = ["minecraft", "raw"]
|
||||
|
||||
[kcp]
|
||||
enable = true
|
||||
port = 25567
|
||||
data_shards = 10
|
||||
parity_Shards = 3
|
||||
|
||||
[websocket]
|
||||
enable = true
|
||||
port = 25568
|
||||
path = "/gateway"
|
||||
|
||||
[hosts]
|
||||
"play.example" = "backend.example:25565"
|
||||
default = "fallback.example:25565"
|
||||
|
||||
[plugin.disabled]
|
||||
enable = false
|
||||
`, pidPath, logPath)
|
||||
|
||||
if err := os.WriteFile(configFile, []byte(toml), 0644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
if err := loadConfig(); err != nil {
|
||||
t.Fatalf("loadConfig() error = %v", err)
|
||||
}
|
||||
|
||||
if !config.Tcp.Enable || config.Tcp.Port != 25565 {
|
||||
t.Fatalf("tcp config = %+v", config.Tcp)
|
||||
}
|
||||
if !config.Quic.Enable || config.Quic.Port != 25566 {
|
||||
t.Fatalf("quic config = %+v", config.Quic)
|
||||
}
|
||||
if got := strings.Join(config.Quic.ApplicationProtocols, ","); got != "minecraft,raw" {
|
||||
t.Fatalf("application protocols = %q, want minecraft,raw", got)
|
||||
}
|
||||
if !config.Kcp.Enable || config.Kcp.DataShards != 10 || config.Kcp.ParityShards != 3 {
|
||||
t.Fatalf("kcp config = %+v", config.Kcp)
|
||||
}
|
||||
if !config.WebSocket.Enable || config.WebSocket.Path != "/gateway" {
|
||||
t.Fatalf("websocket config = %+v", config.WebSocket)
|
||||
}
|
||||
if got := config.Hosts["play.example"]; got != "backend.example:25565" {
|
||||
t.Fatalf("host route = %q, want backend.example:25565", got)
|
||||
}
|
||||
if currentPidFile != pidPath {
|
||||
t.Fatalf("currentPidFile = %q, want %q", currentPidFile, pidPath)
|
||||
}
|
||||
if currentLogFile != logPath {
|
||||
t.Fatalf("currentLogFile = %q, want %q", currentLogFile, logPath)
|
||||
}
|
||||
|
||||
pidBytes, err := os.ReadFile(pidPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(pid) error = %v", err)
|
||||
}
|
||||
if wantPID := fmt.Sprintf("%d\n", os.Getpid()); string(pidBytes) != wantPID {
|
||||
t.Fatalf("pid file = %q, want %q", string(pidBytes), wantPID)
|
||||
}
|
||||
if _, err := os.Stat(logPath); err != nil {
|
||||
t.Fatalf("Stat(log file) error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadConfigReturnsErrors(t *testing.T) {
|
||||
t.Run("missing file", func(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
configFile = filepath.Join(t.TempDir(), "missing.toml")
|
||||
if err := loadConfig(); err == nil {
|
||||
t.Fatal("loadConfig() error = nil, want error")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid toml", func(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
configFile = filepath.Join(t.TempDir(), "config.toml")
|
||||
if err := os.WriteFile(configFile, []byte("[tcp\n"), 0644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
if err := loadConfig(); err == nil {
|
||||
t.Fatal("loadConfig() error = nil, want error")
|
||||
}
|
||||
})
|
||||
}
|
||||
97
cmd/gateway/handle_request_test.go
Normal file
97
cmd/gateway/handle_request_test.go
Normal file
@@ -0,0 +1,97 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestHandleRequestProxiesAndClosesConnections(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
packet := gatewayTestPacket("play.example")
|
||||
source := newGatewayTestConn(packet)
|
||||
upstream := newGatewayTestConn([]byte("reply"))
|
||||
config.Hosts = map[string]string{
|
||||
"play.example": "backend.example:25565",
|
||||
}
|
||||
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(net.Conn, string) bool { return true },
|
||||
func(net.Conn, string) (net.Conn, error) {
|
||||
return upstream, nil
|
||||
},
|
||||
)
|
||||
|
||||
handleRequest(source)
|
||||
|
||||
if !source.closed {
|
||||
t.Fatal("source connection was not closed")
|
||||
}
|
||||
if !upstream.closed {
|
||||
t.Fatal("upstream connection was not closed")
|
||||
}
|
||||
if !bytes.Equal(upstream.writeBuf.Bytes(), packet) {
|
||||
t.Fatalf("upstream initial packet = %v, want %v", upstream.writeBuf.Bytes(), packet)
|
||||
}
|
||||
if got := source.writeBuf.String(); got != "reply" {
|
||||
t.Fatalf("proxied reply = %q, want reply", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleRequestRecoversAndClosesConnection(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
source := &panicReadGatewayConn{gatewayTestConn: newGatewayTestConn(nil)}
|
||||
|
||||
handleRequest(source)
|
||||
|
||||
if !source.closed {
|
||||
t.Fatal("source connection was not closed after panic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayHandleConnStartsRequestGoroutine(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
source := newGatewayTestConn(nil)
|
||||
source.readErr = errors.New("read failed")
|
||||
|
||||
(&Gateway{}).HandleConn(source)
|
||||
|
||||
deadline := time.After(2 * time.Second)
|
||||
ticker := time.NewTicker(time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-deadline:
|
||||
t.Fatal("timed out waiting for HandleConn goroutine")
|
||||
case <-ticker.C:
|
||||
if source.isClosed() {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayTestOpPanics(t *testing.T) {
|
||||
defer func() {
|
||||
if recover() == nil {
|
||||
t.Fatal("TestOp() did not panic")
|
||||
}
|
||||
}()
|
||||
|
||||
(&Gateway{}).TestOp()
|
||||
}
|
||||
|
||||
type panicReadGatewayConn struct {
|
||||
*gatewayTestConn
|
||||
}
|
||||
|
||||
func (c *panicReadGatewayConn) Read([]byte) (int, error) {
|
||||
panic("read panic")
|
||||
}
|
||||
68
cmd/gateway/haproxy_test.go
Normal file
68
cmd/gateway/haproxy_test.go
Normal file
@@ -0,0 +1,68 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"net"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestHaProxyUpstreamWritesProxyHeader(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("Listen() error = %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
headerCh := make(chan string, 1)
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
conn, err := listener.Accept()
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
header, err := bufio.NewReader(conn).ReadString('\n')
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
headerCh <- header
|
||||
}()
|
||||
|
||||
source := newGatewayTestConn(nil)
|
||||
source.remote = &net.TCPAddr{IP: net.ParseIP("127.0.0.2"), Port: 45678}
|
||||
|
||||
conn := haProxyUpstream(source, listener.Addr().String())
|
||||
if conn == nil {
|
||||
t.Fatal("haProxyUpstream() = nil, want connection")
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
select {
|
||||
case header := <-headerCh:
|
||||
if !strings.HasPrefix(header, "PROXY TCP4 127.0.0.2 127.0.0.1 45678 ") {
|
||||
t.Fatalf("proxy header = %q", header)
|
||||
}
|
||||
if !strings.HasSuffix(header, "\r\n") {
|
||||
t.Fatalf("proxy header missing CRLF: %q", header)
|
||||
}
|
||||
case err := <-errCh:
|
||||
t.Fatalf("accept/read header error = %v", err)
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for proxy header")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHaProxyUpstreamReturnsNilForInvalidHost(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
if got := haProxyUpstream(newGatewayTestConn(nil), "not a tcp address"); got != nil {
|
||||
t.Fatalf("haProxyUpstream() = %v, want nil", got)
|
||||
}
|
||||
}
|
||||
137
cmd/gateway/log_pid_test.go
Normal file
137
cmd/gateway/log_pid_test.go
Normal file
@@ -0,0 +1,137 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/rs/zerolog"
|
||||
"github.com/rs/zerolog/log"
|
||||
)
|
||||
|
||||
func TestWritePIDFileCreatesAndSwitchesFiles(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
firstPID := filepath.Join(tmpDir, "first.pid")
|
||||
secondPID := filepath.Join(tmpDir, "second.pid")
|
||||
|
||||
config.PidFile = firstPID
|
||||
if err := writePIDFile(); err != nil {
|
||||
t.Fatalf("writePIDFile(first) error = %v", err)
|
||||
}
|
||||
assertPIDFile(t, firstPID)
|
||||
|
||||
config.PidFile = secondPID
|
||||
if err := writePIDFile(); err != nil {
|
||||
t.Fatalf("writePIDFile(second) error = %v", err)
|
||||
}
|
||||
if _, err := os.Stat(firstPID); !os.IsNotExist(err) {
|
||||
t.Fatalf("first pid file still exists or stat failed: %v", err)
|
||||
}
|
||||
assertPIDFile(t, secondPID)
|
||||
|
||||
removePIDFile()
|
||||
if _, err := os.Stat(secondPID); !os.IsNotExist(err) {
|
||||
t.Fatalf("second pid file still exists or stat failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWritePIDFileNoopsWhenPathUnchanged(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
pidPath := filepath.Join(t.TempDir(), "gateway.pid")
|
||||
config.PidFile = pidPath
|
||||
if err := writePIDFile(); err != nil {
|
||||
t.Fatalf("writePIDFile() error = %v", err)
|
||||
}
|
||||
firstStat, err := os.Stat(pidPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat(first) error = %v", err)
|
||||
}
|
||||
|
||||
if err := writePIDFile(); err != nil {
|
||||
t.Fatalf("writePIDFile() second error = %v", err)
|
||||
}
|
||||
secondStat, err := os.Stat(pidPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat(second) error = %v", err)
|
||||
}
|
||||
if !firstStat.ModTime().Equal(secondStat.ModTime()) {
|
||||
t.Fatalf("pid file mod time changed: %v -> %v", firstStat.ModTime(), secondStat.ModTime())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadLoggerDefaultLevelAndInvalidLevel(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
if err := loadLogger(); err != nil {
|
||||
t.Fatalf("loadLogger() error = %v", err)
|
||||
}
|
||||
if config.Log.Level != "info" {
|
||||
t.Fatalf("default log level = %q, want info", config.Log.Level)
|
||||
}
|
||||
if got := log.Logger.GetLevel(); got != zerolog.InfoLevel {
|
||||
t.Fatalf("logger level = %v, want %v", got, zerolog.InfoLevel)
|
||||
}
|
||||
|
||||
config.Log.Level = "not-a-level"
|
||||
if err := loadLogger(); err == nil {
|
||||
t.Fatal("loadLogger() error = nil, want invalid level error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadLoggerCreatesLogFile(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
logPath := filepath.Join(t.TempDir(), "nested", "gateway.log")
|
||||
config.Log.Level = "debug"
|
||||
config.Log.File = logPath
|
||||
|
||||
if err := loadLogger(); err != nil {
|
||||
t.Fatalf("loadLogger() error = %v", err)
|
||||
}
|
||||
if currentLogFile != logPath {
|
||||
t.Fatalf("currentLogFile = %q, want %q", currentLogFile, logPath)
|
||||
}
|
||||
if _, err := os.Stat(logPath); err != nil {
|
||||
t.Fatalf("Stat(log file) error = %v", err)
|
||||
}
|
||||
if got := log.Logger.GetLevel(); got != zerolog.DebugLevel {
|
||||
t.Fatalf("logger level = %v, want %v", got, zerolog.DebugLevel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReopenLogFile(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
if err := reopenLogFile(); err != nil {
|
||||
t.Fatalf("reopenLogFile(empty) error = %v", err)
|
||||
}
|
||||
|
||||
logPath := filepath.Join(t.TempDir(), "gateway.log")
|
||||
config.Log.File = logPath
|
||||
if err := reopenLogFile(); err != nil {
|
||||
t.Fatalf("reopenLogFile() error = %v", err)
|
||||
}
|
||||
if currentLogFile != logPath {
|
||||
t.Fatalf("currentLogFile = %q, want %q", currentLogFile, logPath)
|
||||
}
|
||||
if _, err := os.Stat(logPath); err != nil {
|
||||
t.Fatalf("Stat(log file) error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertPIDFile(t *testing.T, path string) {
|
||||
t.Helper()
|
||||
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", path, err)
|
||||
}
|
||||
want := fmt.Sprintf("%d\n", os.Getpid())
|
||||
if string(got) != want {
|
||||
t.Fatalf("pid file %s = %q, want %q", path, string(got), want)
|
||||
}
|
||||
}
|
||||
163
cmd/gateway/main_test.go
Normal file
163
cmd/gateway/main_test.go
Normal file
@@ -0,0 +1,163 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMapToHostRoutesThroughHookAndForwardsInitialPacket(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
packet := gatewayTestPacket("play.example", 0x63, 0x00)
|
||||
source := newGatewayTestConn(packet)
|
||||
upstream := newGatewayTestConn(nil)
|
||||
config.Hosts = map[string]string{
|
||||
"play.example": "backend.example:25565",
|
||||
}
|
||||
|
||||
var gotSource net.Conn
|
||||
var gotHost string
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(source net.Conn, host string) bool {
|
||||
gotSource = source
|
||||
gotHost = host
|
||||
return host == "backend.example:25565"
|
||||
},
|
||||
func(source net.Conn, host string) (net.Conn, error) {
|
||||
return upstream, nil
|
||||
},
|
||||
)
|
||||
|
||||
got := mapToHost(source)
|
||||
if got != upstream {
|
||||
t.Fatalf("mapToHost() = %v, want upstream conn", got)
|
||||
}
|
||||
if gotSource != source {
|
||||
t.Fatalf("hook source = %v, want original source", gotSource)
|
||||
}
|
||||
if gotHost != "backend.example:25565" {
|
||||
t.Fatalf("hook host = %q, want backend.example:25565", gotHost)
|
||||
}
|
||||
if !bytes.Equal(upstream.writeBuf.Bytes(), packet) {
|
||||
t.Fatalf("upstream initial packet = %v, want %v", upstream.writeBuf.Bytes(), packet)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapToHostUsesDefaultRoute(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
packet := gatewayTestPacket("unknown.example")
|
||||
source := newGatewayTestConn(packet)
|
||||
upstream := newGatewayTestConn(nil)
|
||||
config.Hosts = map[string]string{
|
||||
"default": "fallback.example:25565",
|
||||
}
|
||||
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(_ net.Conn, host string) bool {
|
||||
return host == "fallback.example:25565"
|
||||
},
|
||||
func(net.Conn, string) (net.Conn, error) {
|
||||
return upstream, nil
|
||||
},
|
||||
)
|
||||
|
||||
if got := mapToHost(source); got != upstream {
|
||||
t.Fatalf("mapToHost() = %v, want fallback upstream", got)
|
||||
}
|
||||
if !bytes.Equal(upstream.writeBuf.Bytes(), packet) {
|
||||
t.Fatalf("upstream initial packet = %v, want %v", upstream.writeBuf.Bytes(), packet)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapToHostRejectsInvalidOrUnroutedPackets(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
packet []byte
|
||||
hosts map[string]string
|
||||
}{
|
||||
{
|
||||
name: "read error",
|
||||
packet: nil,
|
||||
hosts: map[string]string{"default": "fallback.example:25565"},
|
||||
},
|
||||
{
|
||||
name: "malformed packet",
|
||||
packet: []byte{0x01, 0x02, 0x03, 0x04, 0x08, 'a'},
|
||||
hosts: map[string]string{"default": "fallback.example:25565"},
|
||||
},
|
||||
{
|
||||
name: "missing route",
|
||||
packet: gatewayTestPacket("unknown.example"),
|
||||
hosts: map[string]string{},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
source := newGatewayTestConn(tt.packet)
|
||||
if tt.packet == nil {
|
||||
source.readErr = errors.New("read failed")
|
||||
}
|
||||
config.Hosts = tt.hosts
|
||||
|
||||
if got := mapToHost(source); got != nil {
|
||||
t.Fatalf("mapToHost() = %v, want nil", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapToHostReturnsNilWhenHookFails(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
source := newGatewayTestConn(gatewayTestPacket("play.example"))
|
||||
config.Hosts = map[string]string{
|
||||
"play.example": "backend.example:25565",
|
||||
}
|
||||
wantErr := errors.New("hook failed")
|
||||
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(net.Conn, string) bool { return true },
|
||||
func(net.Conn, string) (net.Conn, error) {
|
||||
return nil, wantErr
|
||||
},
|
||||
)
|
||||
|
||||
if got := mapToHost(source); got != nil {
|
||||
t.Fatalf("mapToHost() = %v, want nil", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapToHostClosesUpstreamWhenInitialWriteFails(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
source := newGatewayTestConn(gatewayTestPacket("play.example"))
|
||||
upstream := newGatewayTestConn(nil)
|
||||
upstream.writeErr = errors.New("write failed")
|
||||
config.Hosts = map[string]string{
|
||||
"play.example": "backend.example:25565",
|
||||
}
|
||||
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(net.Conn, string) bool { return true },
|
||||
func(net.Conn, string) (net.Conn, error) {
|
||||
return upstream, nil
|
||||
},
|
||||
)
|
||||
|
||||
if got := mapToHost(source); got != nil {
|
||||
t.Fatalf("mapToHost() = %v, want nil", got)
|
||||
}
|
||||
if !upstream.closed {
|
||||
t.Fatal("upstream was not closed after write failure")
|
||||
}
|
||||
}
|
||||
13
cmd/gateway/pid_unix_test.go
Normal file
13
cmd/gateway/pid_unix_test.go
Normal file
@@ -0,0 +1,13 @@
|
||||
//go:build unix || plan9
|
||||
|
||||
package main
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestGetPidFileFromConfigDefaultOnUnix(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
if got := getPidFileFromConfig(); got != "/dev/shm/mc-gateway.pid" {
|
||||
t.Fatalf("getPidFileFromConfig() = %q, want /dev/shm/mc-gateway.pid", got)
|
||||
}
|
||||
}
|
||||
282
cmd/gateway/plugin_test.go
Normal file
282
cmd/gateway/plugin_test.go
Normal file
@@ -0,0 +1,282 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/tursom/mc-gateway/plugin/api"
|
||||
)
|
||||
|
||||
func TestGatewayHookStoresHandler(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
pluginLock.Lock()
|
||||
hooks["plugin-a"] = make(map[string]any)
|
||||
pluginLock.Unlock()
|
||||
|
||||
handler := "handler"
|
||||
if err := (&Gateway{pluginId: "plugin-a"}).Hook("hook-a", handler); err != nil {
|
||||
t.Fatalf("Hook() error = %v", err)
|
||||
}
|
||||
if got := hooks["plugin-a"]["hook-a"]; got != handler {
|
||||
t.Fatalf("stored hook = %v, want %v", got, handler)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerAdaptors(t *testing.T) {
|
||||
if got := Handler1[int, bool](7)(func(v int) bool { return v == 7 }); !got {
|
||||
t.Fatal("Handler1 did not pass its argument")
|
||||
}
|
||||
if got := Handler2[int, string, bool](7, "x")(func(v int, s string) bool {
|
||||
return v == 7 && s == "x"
|
||||
}); !got {
|
||||
t.Fatal("Handler2 did not pass its arguments")
|
||||
}
|
||||
gotString, gotBool := Handler1R2[int, string, bool](7)(func(v int) (string, bool) {
|
||||
return "ok", v == 7
|
||||
})
|
||||
if gotString != "ok" || !gotBool {
|
||||
t.Fatalf("Handler1R2() = (%q, %v), want (ok, true)", gotString, gotBool)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvokeFirstHookHandler(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
source := newGatewayTestConn(nil)
|
||||
upstream := newGatewayTestConn(nil)
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(gotSource net.Conn, host string) bool {
|
||||
return gotSource == source && host == "backend.example:25565"
|
||||
},
|
||||
func(net.Conn, string) (net.Conn, error) {
|
||||
return upstream, nil
|
||||
},
|
||||
)
|
||||
|
||||
var handled bool
|
||||
ok, err := invokeFirstHookHandler(
|
||||
api.HookUpstream,
|
||||
Handler2[net.Conn, string, bool](source, "backend.example:25565"),
|
||||
func(handler func(net.Conn, string) (net.Conn, error)) error {
|
||||
got, err := handler(source, "backend.example:25565")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
handled = got == upstream
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("invokeFirstHookHandler() error = %v", err)
|
||||
}
|
||||
if !ok {
|
||||
t.Fatal("invokeFirstHookHandler() ok = false, want true")
|
||||
}
|
||||
if !handled {
|
||||
t.Fatal("first matching handler was not invoked")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvokeFirstHookHandlerNoMatch(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(net.Conn, string) bool { return false },
|
||||
func(net.Conn, string) (net.Conn, error) {
|
||||
t.Fatal("handler should not be invoked")
|
||||
return nil, nil
|
||||
},
|
||||
)
|
||||
|
||||
ok, err := invokeFirstHookHandler(
|
||||
api.HookUpstream,
|
||||
Handler2[net.Conn, string, bool](newGatewayTestConn(nil), "backend.example:25565"),
|
||||
func(func(net.Conn, string) (net.Conn, error)) error {
|
||||
t.Fatal("callback should not be invoked")
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("invokeFirstHookHandler() error = %v", err)
|
||||
}
|
||||
if ok {
|
||||
t.Fatal("invokeFirstHookHandler() ok = true, want false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvokeFirstHookHandlerReturnsCallbackError(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
wantErr := errors.New("callback failed")
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(net.Conn, string) bool { return true },
|
||||
func(net.Conn, string) (net.Conn, error) { return nil, nil },
|
||||
)
|
||||
|
||||
ok, err := invokeFirstHookHandler(
|
||||
api.HookUpstream,
|
||||
Handler2[net.Conn, string, bool](newGatewayTestConn(nil), "backend.example:25565"),
|
||||
func(func(net.Conn, string) (net.Conn, error)) error {
|
||||
return wantErr
|
||||
},
|
||||
)
|
||||
if !ok {
|
||||
t.Fatal("invokeFirstHookHandler() ok = false, want true")
|
||||
}
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("invokeFirstHookHandler() error = %v, want %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvokeAllHookHandler(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
for _, pluginID := range []string{"plugin-a", "plugin-b"} {
|
||||
pluginLock.Lock()
|
||||
hooks[pluginID] = make(map[string]any)
|
||||
pluginLock.Unlock()
|
||||
if err := api.RegisterHookHandler(
|
||||
&Gateway{pluginId: pluginID},
|
||||
api.HookUpstream,
|
||||
func(net.Conn, string) bool { return true },
|
||||
func(net.Conn, string) (net.Conn, error) { return nil, nil },
|
||||
); err != nil {
|
||||
t.Fatalf("RegisterHookHandler(%s) error = %v", pluginID, err)
|
||||
}
|
||||
}
|
||||
|
||||
calls := 0
|
||||
err := invokeAllHookHandler(
|
||||
api.HookUpstream,
|
||||
Handler2[net.Conn, string, bool](newGatewayTestConn(nil), "backend.example:25565"),
|
||||
func(func(net.Conn, string) (net.Conn, error)) error {
|
||||
calls++
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("invokeAllHookHandler() error = %v", err)
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("invokeAllHookHandler() calls = %d, want 2", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvokeAllHookHandlerReturnsCallbackError(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
wantErr := errors.New("callback failed")
|
||||
registerGatewayUpstreamHook(
|
||||
t,
|
||||
func(net.Conn, string) bool { return true },
|
||||
func(net.Conn, string) (net.Conn, error) { return nil, nil },
|
||||
)
|
||||
|
||||
err := invokeAllHookHandler(
|
||||
api.HookUpstream,
|
||||
Handler2[net.Conn, string, bool](newGatewayTestConn(nil), "backend.example:25565"),
|
||||
func(func(net.Conn, string) (net.Conn, error)) error {
|
||||
return wantErr
|
||||
},
|
||||
)
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("invokeAllHookHandler() error = %v, want %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadPluginsDisablesExistingPlugin(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
plugin := &gatewayPluginStub{}
|
||||
plugins["plugin-a"] = plugin
|
||||
hooks["plugin-a"] = map[string]any{"hook": "handler"}
|
||||
config.Plugin = map[string]map[string]any{
|
||||
"plugin-a": {"enable": false},
|
||||
}
|
||||
|
||||
loadPlugins()
|
||||
|
||||
if plugin.destroyCalls != 1 {
|
||||
t.Fatalf("Destroy() calls = %d, want 1", plugin.destroyCalls)
|
||||
}
|
||||
if _, ok := plugins["plugin-a"]; ok {
|
||||
t.Fatal("disabled plugin was not removed")
|
||||
}
|
||||
if _, ok := hooks["plugin-a"]; ok {
|
||||
t.Fatal("disabled plugin hooks were not removed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadPluginsSkipsAlreadyLoadedEnabledPlugin(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
plugin := &gatewayPluginStub{}
|
||||
plugins["plugin-a"] = plugin
|
||||
hooks["plugin-a"] = make(map[string]any)
|
||||
config.Plugin = map[string]map[string]any{
|
||||
"plugin-a": {"enable": true, "file": "missing-plugin-file"},
|
||||
}
|
||||
|
||||
loadPlugins()
|
||||
|
||||
if plugins["plugin-a"] != plugin {
|
||||
t.Fatal("existing enabled plugin was replaced")
|
||||
}
|
||||
if plugin.destroyCalls != 0 {
|
||||
t.Fatalf("Destroy() calls = %d, want 0", plugin.destroyCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadPluginsIgnoresMissingPluginFile(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
config.Plugin = map[string]map[string]any{
|
||||
"plugin-a": {"enable": true, "file": "missing-plugin-file"},
|
||||
}
|
||||
|
||||
loadPlugins()
|
||||
|
||||
if _, ok := plugins["plugin-a"]; ok {
|
||||
t.Fatal("missing plugin file should not be registered")
|
||||
}
|
||||
if _, ok := hooks["plugin-a"]; ok {
|
||||
t.Fatal("missing plugin file should not create hooks")
|
||||
}
|
||||
}
|
||||
|
||||
type gatewayPluginStub struct {
|
||||
destroyCalls int
|
||||
}
|
||||
|
||||
func (p *gatewayPluginStub) Init(api.Gateway) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *gatewayPluginStub) Destroy() error {
|
||||
p.destroyCalls++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *gatewayPluginStub) NewConfigObj() any {
|
||||
return &struct{}{}
|
||||
}
|
||||
|
||||
func (p *gatewayPluginStub) ReloadConfig(any) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestGatewayExitWaitGroup(t *testing.T) {
|
||||
if got := (&Gateway{}).ExitWaitGroup(); got != &exitWaitGroup {
|
||||
t.Fatalf("ExitWaitGroup() = %p, want %p", got, &exitWaitGroup)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayPluginStubSatisfiesInterface(t *testing.T) {
|
||||
var _ api.Plugin = (*gatewayPluginStub)(nil)
|
||||
var _ api.Gateway = (*Gateway)(nil)
|
||||
}
|
||||
51
cmd/gateway/quic_test.go
Normal file
51
cmd/gateway/quic_test.go
Normal file
@@ -0,0 +1,51 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGetQuicNextProtos(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
wantDefault := []string{"minecraft", "quic", "raw", "h3"}
|
||||
if got := getQuicNextProtos(); !reflect.DeepEqual(got, wantDefault) {
|
||||
t.Fatalf("getQuicNextProtos() = %v, want %v", got, wantDefault)
|
||||
}
|
||||
|
||||
config.Quic.ApplicationProtocols = []string{"minecraft", "custom"}
|
||||
if got := getQuicNextProtos(); !reflect.DeepEqual(got, config.Quic.ApplicationProtocols) {
|
||||
t.Fatalf("getQuicNextProtos() = %v, want %v", got, config.Quic.ApplicationProtocols)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateTLSConfig(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
config.Quic.ApplicationProtocols = []string{"minecraft", "custom"}
|
||||
tlsConfig, err := generateTLSConfig()
|
||||
if err != nil {
|
||||
t.Fatalf("generateTLSConfig() error = %v", err)
|
||||
}
|
||||
if len(tlsConfig.Certificates) != 1 {
|
||||
t.Fatalf("certificates len = %d, want 1", len(tlsConfig.Certificates))
|
||||
}
|
||||
if !reflect.DeepEqual(tlsConfig.NextProtos, config.Quic.ApplicationProtocols) {
|
||||
t.Fatalf("NextProtos = %v, want %v", tlsConfig.NextProtos, config.Quic.ApplicationProtocols)
|
||||
}
|
||||
|
||||
cert, err := x509.ParseCertificate(tlsConfig.Certificates[0].Certificate[0])
|
||||
if err != nil {
|
||||
t.Fatalf("ParseCertificate() error = %v", err)
|
||||
}
|
||||
if !cert.IsCA && !cert.BasicConstraintsValid {
|
||||
t.Fatal("generated certificate has invalid basic constraints")
|
||||
}
|
||||
if len(cert.ExtKeyUsage) != 1 || cert.ExtKeyUsage[0] != x509.ExtKeyUsageServerAuth {
|
||||
t.Fatalf("ExtKeyUsage = %v, want server auth", cert.ExtKeyUsage)
|
||||
}
|
||||
if !cert.NotAfter.After(cert.NotBefore) {
|
||||
t.Fatalf("certificate validity range is invalid: %v - %v", cert.NotBefore, cert.NotAfter)
|
||||
}
|
||||
}
|
||||
207
cmd/gateway/relay_test.go
Normal file
207
cmd/gateway/relay_test.go
Normal file
@@ -0,0 +1,207 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWriteAll(t *testing.T) {
|
||||
t.Run("partial writes", func(t *testing.T) {
|
||||
writer := &partialGatewayWriter{chunkSize: 2}
|
||||
if err := writeAll(writer, []byte("abcdef")); err != nil {
|
||||
t.Fatalf("writeAll() error = %v", err)
|
||||
}
|
||||
if got := writer.buf.String(); got != "abcdef" {
|
||||
t.Fatalf("written data = %q, want abcdef", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("short write", func(t *testing.T) {
|
||||
if err := writeAll(shortGatewayWriter{}, []byte("abcdef")); !errors.Is(err, io.ErrShortWrite) {
|
||||
t.Fatalf("writeAll() error = %v, want %v", err, io.ErrShortWrite)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("writer error", func(t *testing.T) {
|
||||
wantErr := errors.New("write failed")
|
||||
writer := &partialGatewayWriter{chunkSize: 2, err: wantErr}
|
||||
if err := writeAll(writer, []byte("abcdef")); !errors.Is(err, wantErr) {
|
||||
t.Fatalf("writeAll() error = %v, want %v", err, wantErr)
|
||||
}
|
||||
if got := writer.buf.String(); got != "ab" {
|
||||
t.Fatalf("written data = %q, want ab", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestCopyForwardCopiesData(t *testing.T) {
|
||||
var dst bytes.Buffer
|
||||
n, err := copyForward(&dst, strings.NewReader("payload"))
|
||||
if err != nil {
|
||||
t.Fatalf("copyForward() error = %v", err)
|
||||
}
|
||||
if n != int64(len("payload")) {
|
||||
t.Fatalf("copyForward() n = %d, want %d", n, len("payload"))
|
||||
}
|
||||
if got := dst.String(); got != "payload" {
|
||||
t.Fatalf("copyForward() data = %q, want payload", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyForwardWithPlainReaderWriter(t *testing.T) {
|
||||
dst := &plainGatewayWriter{}
|
||||
n, err := copyForward(dst, &plainGatewayReader{data: []byte("plain")})
|
||||
if err != nil {
|
||||
t.Fatalf("copyForward() error = %v", err)
|
||||
}
|
||||
if n != int64(len("plain")) {
|
||||
t.Fatalf("copyForward() n = %d, want %d", n, len("plain"))
|
||||
}
|
||||
if got := dst.buf.String(); got != "plain" {
|
||||
t.Fatalf("copyForward() data = %q, want plain", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyCopyClosesSides(t *testing.T) {
|
||||
src := &closingGatewayReader{}
|
||||
dst := &closingGatewayWriter{}
|
||||
|
||||
proxyCopy(dst, src)
|
||||
|
||||
if !src.closeReadCalled {
|
||||
t.Fatal("CloseRead was not called")
|
||||
}
|
||||
if !dst.closeWriteCalled {
|
||||
t.Fatal("CloseWrite was not called")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyCopyRecoversAndClosesSides(t *testing.T) {
|
||||
src := &panicGatewayReader{}
|
||||
dst := &closingGatewayWriter{}
|
||||
|
||||
proxyCopy(dst, src)
|
||||
|
||||
if !src.closeReadCalled {
|
||||
t.Fatal("CloseRead was not called after panic")
|
||||
}
|
||||
if !dst.closeWriteCalled {
|
||||
t.Fatal("CloseWrite was not called after panic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloseWriteFallsBackToClose(t *testing.T) {
|
||||
closer := &gatewayCloser{}
|
||||
closeWrite(closer)
|
||||
if !closer.closed {
|
||||
t.Fatal("Close was not called")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAndPutProxyBuffer(t *testing.T) {
|
||||
buf := getProxyBuffer()
|
||||
if len(buf) != proxyBufferSize {
|
||||
t.Fatalf("buffer len = %d, want %d", len(buf), proxyBufferSize)
|
||||
}
|
||||
if cap(buf) != proxyBufferSize {
|
||||
t.Fatalf("buffer cap = %d, want %d", cap(buf), proxyBufferSize)
|
||||
}
|
||||
|
||||
putProxyBuffer(buf[:1])
|
||||
putProxyBuffer(make([]byte, 1))
|
||||
}
|
||||
|
||||
type partialGatewayWriter struct {
|
||||
chunkSize int
|
||||
err error
|
||||
buf bytes.Buffer
|
||||
}
|
||||
|
||||
func (w *partialGatewayWriter) Write(p []byte) (int, error) {
|
||||
if len(p) > w.chunkSize {
|
||||
p = p[:w.chunkSize]
|
||||
}
|
||||
n, _ := w.buf.Write(p)
|
||||
if w.err != nil {
|
||||
return n, w.err
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
type shortGatewayWriter struct{}
|
||||
|
||||
func (shortGatewayWriter) Write([]byte) (int, error) {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
type plainGatewayReader struct {
|
||||
data []byte
|
||||
}
|
||||
|
||||
func (r *plainGatewayReader) Read(p []byte) (int, error) {
|
||||
if len(r.data) == 0 {
|
||||
return 0, io.EOF
|
||||
}
|
||||
n := copy(p, r.data)
|
||||
r.data = r.data[n:]
|
||||
return n, nil
|
||||
}
|
||||
|
||||
type plainGatewayWriter struct {
|
||||
buf bytes.Buffer
|
||||
}
|
||||
|
||||
func (w *plainGatewayWriter) Write(p []byte) (int, error) {
|
||||
return w.buf.Write(p)
|
||||
}
|
||||
|
||||
type closingGatewayReader struct {
|
||||
closeReadCalled bool
|
||||
}
|
||||
|
||||
func (r *closingGatewayReader) Read([]byte) (int, error) {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
func (r *closingGatewayReader) CloseRead() error {
|
||||
r.closeReadCalled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
type panicGatewayReader struct {
|
||||
closeReadCalled bool
|
||||
}
|
||||
|
||||
func (r *panicGatewayReader) Read([]byte) (int, error) {
|
||||
panic("read panic")
|
||||
}
|
||||
|
||||
func (r *panicGatewayReader) CloseRead() error {
|
||||
r.closeReadCalled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
type closingGatewayWriter struct {
|
||||
closeWriteCalled bool
|
||||
}
|
||||
|
||||
func (w *closingGatewayWriter) Write(p []byte) (int, error) {
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (w *closingGatewayWriter) CloseWrite() error {
|
||||
w.closeWriteCalled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
type gatewayCloser struct {
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (c *gatewayCloser) Close() error {
|
||||
c.closed = true
|
||||
return nil
|
||||
}
|
||||
59
cmd/gateway/tcp_test.go
Normal file
59
cmd/gateway/tcp_test.go
Normal file
@@ -0,0 +1,59 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestSetSocketOptions(t *testing.T) {
|
||||
client, server := net.Pipe()
|
||||
defer client.Close()
|
||||
defer server.Close()
|
||||
|
||||
setSocketOptions(client)
|
||||
}
|
||||
|
||||
func TestUpstreamTcp(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("Listen() error = %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan net.Conn, 1)
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
conn, err := listener.Accept()
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
accepted <- conn
|
||||
}()
|
||||
|
||||
conn := upstreamTcp(listener.Addr().String())
|
||||
if conn == nil {
|
||||
t.Fatal("upstreamTcp() = nil, want connection")
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
select {
|
||||
case serverConn := <-accepted:
|
||||
serverConn.Close()
|
||||
case err := <-errCh:
|
||||
t.Fatalf("Accept() error = %v", err)
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for upstream TCP connection")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpstreamTcpReturnsNilOnDialError(t *testing.T) {
|
||||
defer saveGatewayState(t)()
|
||||
|
||||
if got := upstreamTcp("not a tcp address"); got != nil {
|
||||
t.Fatalf("upstreamTcp() = %v, want nil", got)
|
||||
}
|
||||
}
|
||||
164
cmd/gateway/test_helpers_test.go
Normal file
164
cmd/gateway/test_helpers_test.go
Normal file
@@ -0,0 +1,164 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rs/zerolog"
|
||||
"github.com/rs/zerolog/log"
|
||||
"github.com/tursom/mc-gateway/plugin/api"
|
||||
)
|
||||
|
||||
func saveGatewayState(t *testing.T) func() {
|
||||
t.Helper()
|
||||
|
||||
oldConfig := config
|
||||
oldConfigFile := configFile
|
||||
oldCurrentPidFile := currentPidFile
|
||||
oldCurrentLogFile := currentLogFile
|
||||
oldLogger := log.Logger
|
||||
|
||||
pluginLock.Lock()
|
||||
oldPlugins := plugins
|
||||
oldHooks := hooks
|
||||
plugins = make(map[string]api.Plugin)
|
||||
hooks = make(map[string]map[string]any)
|
||||
pluginLock.Unlock()
|
||||
|
||||
config = Config{}
|
||||
configFile = "config.toml"
|
||||
currentPidFile = ""
|
||||
currentLogFile = ""
|
||||
log.Logger = zerolog.New(io.Discard)
|
||||
|
||||
return func() {
|
||||
if currentPidFile != "" && currentPidFile != oldCurrentPidFile {
|
||||
_ = os.Remove(currentPidFile)
|
||||
}
|
||||
|
||||
config = oldConfig
|
||||
configFile = oldConfigFile
|
||||
currentPidFile = oldCurrentPidFile
|
||||
currentLogFile = oldCurrentLogFile
|
||||
log.Logger = oldLogger
|
||||
|
||||
pluginLock.Lock()
|
||||
plugins = oldPlugins
|
||||
hooks = oldHooks
|
||||
pluginLock.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func gatewayTestPacket(host string, tail ...byte) []byte {
|
||||
packet := []byte{
|
||||
byte(4 + 1 + len(host) + len(tail)),
|
||||
0x00,
|
||||
0x00,
|
||||
0x00,
|
||||
byte(len(host)),
|
||||
}
|
||||
packet = append(packet, host...)
|
||||
packet = append(packet, tail...)
|
||||
return packet
|
||||
}
|
||||
|
||||
type gatewayTestConn struct {
|
||||
mu sync.Mutex
|
||||
readBuf []byte
|
||||
readErr error
|
||||
writeBuf bytes.Buffer
|
||||
writeErr error
|
||||
|
||||
closed bool
|
||||
local net.Addr
|
||||
remote net.Addr
|
||||
}
|
||||
|
||||
func newGatewayTestConn(readBuf []byte) *gatewayTestConn {
|
||||
return &gatewayTestConn{
|
||||
readBuf: append([]byte(nil), readBuf...),
|
||||
local: &net.TCPAddr{IP: net.ParseIP("127.0.0.1"), Port: 25565},
|
||||
remote: &net.TCPAddr{IP: net.ParseIP("127.0.0.2"), Port: 45678},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) Read(p []byte) (int, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.readErr != nil {
|
||||
return 0, c.readErr
|
||||
}
|
||||
if len(c.readBuf) == 0 {
|
||||
return 0, io.EOF
|
||||
}
|
||||
n := copy(p, c.readBuf)
|
||||
c.readBuf = c.readBuf[n:]
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) Write(p []byte) (int, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.writeErr != nil {
|
||||
return 0, c.writeErr
|
||||
}
|
||||
return c.writeBuf.Write(p)
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.closed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) isClosed() bool {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
return c.closed
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) LocalAddr() net.Addr {
|
||||
return c.local
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) RemoteAddr() net.Addr {
|
||||
return c.remote
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) SetDeadline(time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) SetReadDeadline(time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *gatewayTestConn) SetWriteDeadline(time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func registerGatewayUpstreamHook(
|
||||
t *testing.T,
|
||||
acceptor func(net.Conn, string) bool,
|
||||
handler func(net.Conn, string) (net.Conn, error),
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
pluginLock.Lock()
|
||||
hooks["test-plugin"] = make(map[string]any)
|
||||
pluginLock.Unlock()
|
||||
|
||||
if err := api.RegisterHookHandler(&Gateway{pluginId: "test-plugin"}, api.HookUpstream, acceptor, handler); err != nil {
|
||||
t.Fatalf("RegisterHookHandler() error = %v", err)
|
||||
}
|
||||
}
|
||||
135
cmd/gateway/websocket_test.go
Normal file
135
cmd/gateway/websocket_test.go
Normal file
@@ -0,0 +1,135 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
func TestWebSocketConnReadWriteAndDeadline(t *testing.T) {
|
||||
serverConn, clientConn, cleanup := newGatewayWebSocketPair(t)
|
||||
defer cleanup()
|
||||
|
||||
conn := &webSocketConn{Conn: serverConn}
|
||||
|
||||
if err := clientConn.WriteMessage(websocket.TextMessage, []byte("hello")); err != nil {
|
||||
t.Fatalf("client WriteMessage() error = %v", err)
|
||||
}
|
||||
|
||||
buf := make([]byte, len("hello"))
|
||||
n, err := conn.Read(buf)
|
||||
if err != nil {
|
||||
t.Fatalf("Read() error = %v", err)
|
||||
}
|
||||
if got := string(buf[:n]); got != "hello" {
|
||||
t.Fatalf("Read() = %q, want hello", got)
|
||||
}
|
||||
|
||||
if err := conn.SetDeadline(time.Now().Add(time.Second)); err != nil {
|
||||
t.Fatalf("SetDeadline() error = %v", err)
|
||||
}
|
||||
|
||||
payload := []byte("response")
|
||||
n, err = conn.Write(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("Write() error = %v", err)
|
||||
}
|
||||
if n != len(payload) {
|
||||
t.Fatalf("Write() n = %d, want %d", n, len(payload))
|
||||
}
|
||||
|
||||
messageType, got, err := clientConn.ReadMessage()
|
||||
if err != nil {
|
||||
t.Fatalf("client ReadMessage() error = %v", err)
|
||||
}
|
||||
if messageType != websocket.BinaryMessage {
|
||||
t.Fatalf("message type = %d, want %d", messageType, websocket.BinaryMessage)
|
||||
}
|
||||
if string(got) != string(payload) {
|
||||
t.Fatalf("message payload = %q, want %q", got, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebSocketConnReadContinuesAcrossFrames(t *testing.T) {
|
||||
serverConn, clientConn, cleanup := newGatewayWebSocketPair(t)
|
||||
defer cleanup()
|
||||
|
||||
conn := &webSocketConn{Conn: serverConn}
|
||||
|
||||
if err := clientConn.WriteMessage(websocket.BinaryMessage, []byte("abc")); err != nil {
|
||||
t.Fatalf("client WriteMessage(first) error = %v", err)
|
||||
}
|
||||
if err := clientConn.WriteMessage(websocket.BinaryMessage, []byte("def")); err != nil {
|
||||
t.Fatalf("client WriteMessage(second) error = %v", err)
|
||||
}
|
||||
|
||||
buf := make([]byte, 2)
|
||||
n, err := conn.Read(buf)
|
||||
if err != nil {
|
||||
t.Fatalf("Read(first) error = %v", err)
|
||||
}
|
||||
if got := string(buf[:n]); got != "ab" {
|
||||
t.Fatalf("Read(first) = %q, want ab", got)
|
||||
}
|
||||
n, err = conn.Read(buf)
|
||||
if err != nil {
|
||||
t.Fatalf("Read(second) error = %v", err)
|
||||
}
|
||||
if got := string(buf[:n]); got != "c" {
|
||||
t.Fatalf("Read(second) = %q, want c", got)
|
||||
}
|
||||
n, err = conn.Read(buf)
|
||||
if err != nil {
|
||||
t.Fatalf("Read(third) error = %v", err)
|
||||
}
|
||||
if got := string(buf[:n]); got != "de" {
|
||||
t.Fatalf("Read(third) = %q, want de", got)
|
||||
}
|
||||
}
|
||||
|
||||
func newGatewayWebSocketPair(t *testing.T) (*websocket.Conn, *websocket.Conn, func()) {
|
||||
t.Helper()
|
||||
|
||||
serverConnCh := make(chan *websocket.Conn, 1)
|
||||
errCh := make(chan error, 1)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
serverConnCh <- conn
|
||||
}))
|
||||
|
||||
url := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
clientConn, _, err := websocket.DefaultDialer.Dial(url, nil)
|
||||
if err != nil {
|
||||
server.Close()
|
||||
t.Fatalf("Dial() error = %v", err)
|
||||
}
|
||||
|
||||
var serverConn *websocket.Conn
|
||||
select {
|
||||
case serverConn = <-serverConnCh:
|
||||
case err := <-errCh:
|
||||
clientConn.Close()
|
||||
server.Close()
|
||||
t.Fatalf("Upgrade() error = %v", err)
|
||||
case <-time.After(2 * time.Second):
|
||||
clientConn.Close()
|
||||
server.Close()
|
||||
t.Fatal("timed out waiting for websocket upgrade")
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
_ = clientConn.Close()
|
||||
_ = serverConn.Close()
|
||||
server.Close()
|
||||
}
|
||||
|
||||
return serverConn, clientConn, cleanup
|
||||
}
|
||||
94
cmd/kcp/main_test.go
Normal file
94
cmd/kcp/main_test.go
Normal file
@@ -0,0 +1,94 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCopyData(t *testing.T) {
|
||||
var dst bytes.Buffer
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
|
||||
copyData(strings.NewReader("payload"), &dst, &wg)
|
||||
wg.Wait()
|
||||
|
||||
if got := dst.String(); got != "payload" {
|
||||
t.Fatalf("copyData() wrote %q, want payload", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyDataIgnoresEOFAndLogsOtherErrors(t *testing.T) {
|
||||
copyData(strings.NewReader(""), io.Discard, nil)
|
||||
copyData(errorReader{err: errors.New("read failed")}, io.Discard, nil)
|
||||
}
|
||||
|
||||
func TestSetSocketOptions(t *testing.T) {
|
||||
client, server := net.Pipe()
|
||||
defer client.Close()
|
||||
defer server.Close()
|
||||
|
||||
setSocketOptions(client)
|
||||
}
|
||||
|
||||
func TestSetSocketOptionsTCPConn(t *testing.T) {
|
||||
client, server := newKcpTestTCPConnPair(t)
|
||||
defer client.Close()
|
||||
defer server.Close()
|
||||
|
||||
setSocketOptions(client)
|
||||
setSocketOptions(server)
|
||||
}
|
||||
|
||||
type errorReader struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (r errorReader) Read([]byte) (int, error) {
|
||||
return 0, r.err
|
||||
}
|
||||
|
||||
func newKcpTestTCPConnPair(t *testing.T) (*net.TCPConn, *net.TCPConn) {
|
||||
t.Helper()
|
||||
|
||||
listener, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1")})
|
||||
if err != nil {
|
||||
t.Fatalf("ListenTCP() error = %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan *net.TCPConn, 1)
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
conn, err := listener.AcceptTCP()
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
accepted <- conn
|
||||
}()
|
||||
|
||||
client, err := net.DialTCP("tcp", nil, listener.Addr().(*net.TCPAddr))
|
||||
if err != nil {
|
||||
t.Fatalf("DialTCP() error = %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case server := <-accepted:
|
||||
return client, server
|
||||
case err := <-errCh:
|
||||
client.Close()
|
||||
t.Fatalf("AcceptTCP() error = %v", err)
|
||||
case <-time.After(2 * time.Second):
|
||||
client.Close()
|
||||
t.Fatal("timed out waiting for TCP accept")
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
94
cmd/quic/main_test.go
Normal file
94
cmd/quic/main_test.go
Normal file
@@ -0,0 +1,94 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCopyData(t *testing.T) {
|
||||
var dst bytes.Buffer
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
|
||||
copyData(strings.NewReader("payload"), &dst, &wg)
|
||||
wg.Wait()
|
||||
|
||||
if got := dst.String(); got != "payload" {
|
||||
t.Fatalf("copyData() wrote %q, want payload", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCopyDataIgnoresEOFAndLogsOtherErrors(t *testing.T) {
|
||||
copyData(strings.NewReader(""), io.Discard, nil)
|
||||
copyData(errorReader{err: errors.New("read failed")}, io.Discard, nil)
|
||||
}
|
||||
|
||||
func TestSetSocketOptions(t *testing.T) {
|
||||
client, server := net.Pipe()
|
||||
defer client.Close()
|
||||
defer server.Close()
|
||||
|
||||
setSocketOptions(client)
|
||||
}
|
||||
|
||||
func TestSetSocketOptionsTCPConn(t *testing.T) {
|
||||
client, server := newQuicTestTCPConnPair(t)
|
||||
defer client.Close()
|
||||
defer server.Close()
|
||||
|
||||
setSocketOptions(client)
|
||||
setSocketOptions(server)
|
||||
}
|
||||
|
||||
type errorReader struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (r errorReader) Read([]byte) (int, error) {
|
||||
return 0, r.err
|
||||
}
|
||||
|
||||
func newQuicTestTCPConnPair(t *testing.T) (*net.TCPConn, *net.TCPConn) {
|
||||
t.Helper()
|
||||
|
||||
listener, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1")})
|
||||
if err != nil {
|
||||
t.Fatalf("ListenTCP() error = %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan *net.TCPConn, 1)
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
conn, err := listener.AcceptTCP()
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
accepted <- conn
|
||||
}()
|
||||
|
||||
client, err := net.DialTCP("tcp", nil, listener.Addr().(*net.TCPAddr))
|
||||
if err != nil {
|
||||
t.Fatalf("DialTCP() error = %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case server := <-accepted:
|
||||
return client, server
|
||||
case err := <-errCh:
|
||||
client.Close()
|
||||
t.Fatalf("AcceptTCP() error = %v", err)
|
||||
case <-time.After(2 * time.Second):
|
||||
client.Close()
|
||||
t.Fatal("timed out waiting for TCP accept")
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
91
plugin/api/api_test.go
Normal file
91
plugin/api/api_test.go
Normal file
@@ -0,0 +1,91 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAbstractPluginDefaults(t *testing.T) {
|
||||
var plugin AbstractPlugin
|
||||
|
||||
if err := plugin.Init(nil); err != nil {
|
||||
t.Fatalf("Init() error = %v", err)
|
||||
}
|
||||
if err := plugin.Destroy(); err != nil {
|
||||
t.Fatalf("Destroy() error = %v", err)
|
||||
}
|
||||
if err := plugin.ReloadConfig(struct{}{}); err != nil {
|
||||
t.Fatalf("ReloadConfig() error = %v", err)
|
||||
}
|
||||
if cfg := plugin.NewConfigObj(); cfg == nil {
|
||||
t.Fatal("NewConfigObj() returned nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHookTypesAndHandlers(t *testing.T) {
|
||||
if got := HookUpstream.Key(); got != "upstream" {
|
||||
t.Fatalf("HookUpstream.Key() = %q, want upstream", got)
|
||||
}
|
||||
if got := HookUpstream.AsAny().Key(); got != HookUpstream.Key() {
|
||||
t.Fatalf("AsAny().Key() = %q, want %q", got, HookUpstream.Key())
|
||||
}
|
||||
|
||||
acceptor := func(net.Conn, string) bool { return true }
|
||||
handler := func(net.Conn, string) (net.Conn, error) { return nil, nil }
|
||||
hookHandler := HookHandler[
|
||||
func(net.Conn, string) bool,
|
||||
func(net.Conn, string) (net.Conn, error),
|
||||
]{acceptor: acceptor, handler: handler}
|
||||
|
||||
if hookHandler.Acceptor() == nil {
|
||||
t.Fatal("Acceptor() returned nil")
|
||||
}
|
||||
if hookHandler.Handler() == nil {
|
||||
t.Fatal("Handler() returned nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHookHandler(t *testing.T) {
|
||||
gateway := &recordingGateway{hooks: make(map[string]any)}
|
||||
acceptor := func(net.Conn, string) bool { return true }
|
||||
handler := func(net.Conn, string) (net.Conn, error) { return nil, nil }
|
||||
|
||||
if err := RegisterHookHandler(gateway, HookUpstream, acceptor, handler); err != nil {
|
||||
t.Fatalf("RegisterHookHandler() error = %v", err)
|
||||
}
|
||||
|
||||
rawHandler, ok := gateway.hooks[HookUpstream.Key()]
|
||||
if !ok {
|
||||
t.Fatalf("hook %q was not registered", HookUpstream.Key())
|
||||
}
|
||||
registered, ok := rawHandler.(HookHandler[
|
||||
func(net.Conn, string) bool,
|
||||
func(net.Conn, string) (net.Conn, error),
|
||||
])
|
||||
if !ok {
|
||||
t.Fatalf("registered hook type = %T", rawHandler)
|
||||
}
|
||||
if registered.Acceptor() == nil {
|
||||
t.Fatal("registered acceptor is nil")
|
||||
}
|
||||
if registered.Handler() == nil {
|
||||
t.Fatal("registered handler is nil")
|
||||
}
|
||||
}
|
||||
|
||||
type recordingGateway struct {
|
||||
hooks map[string]any
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func (g *recordingGateway) HandleConn(net.Conn) {}
|
||||
|
||||
func (g *recordingGateway) ExitWaitGroup() *sync.WaitGroup {
|
||||
return &g.wg
|
||||
}
|
||||
|
||||
func (g *recordingGateway) Hook(hook string, handler any) error {
|
||||
g.hooks[hook] = handler
|
||||
return nil
|
||||
}
|
||||
105
protocol/mc_test.go
Normal file
105
protocol/mc_test.go
Normal file
@@ -0,0 +1,105 @@
|
||||
package protocol
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGetMcHost(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
buf []byte
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "short packet",
|
||||
buf: []byte{0x01, 0x02, 0x03, 0x04},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "host length exceeds packet",
|
||||
buf: []byte{0x01, 0x02, 0x03, 0x04, 0x05, 'a'},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "plain host",
|
||||
buf: mcTestPacket("play.example", 0x63, 0x00),
|
||||
want: "play.example",
|
||||
},
|
||||
{
|
||||
name: "host with null suffix",
|
||||
buf: mcTestPacket("play.example\x00FML\x00", 0x63, 0x00),
|
||||
want: "play.example",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := GetMcHost(tt.buf); got != tt.want {
|
||||
t.Fatalf("GetMcHost() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplaceMcHost(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
buf []byte
|
||||
host string
|
||||
want []byte
|
||||
}{
|
||||
{
|
||||
name: "short packet",
|
||||
buf: []byte{0x01, 0x02, 0x03, 0x04},
|
||||
host: "new.example",
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "host length exceeds packet",
|
||||
buf: []byte{0x01, 0x02, 0x03, 0x04, 0x05, 'a'},
|
||||
host: "new.example",
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "replace plain host",
|
||||
buf: mcTestPacket("old.example", 0x63, 0x00),
|
||||
host: "new.example",
|
||||
want: mcTestPacket("new.example", 0x63, 0x00),
|
||||
},
|
||||
{
|
||||
name: "preserve null suffix",
|
||||
buf: mcTestPacket("old.example\x00FML\x00", 0x63, 0x00),
|
||||
host: "new.example",
|
||||
want: mcTestPacket("new.example\x00FML\x00", 0x63, 0x00),
|
||||
},
|
||||
{
|
||||
name: "replace with shorter host",
|
||||
buf: mcTestPacket("long.example", 0x01, 0x02, 0x03),
|
||||
host: "x",
|
||||
want: mcTestPacket("x", 0x01, 0x02, 0x03),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := ReplaceMcHost(append([]byte(nil), tt.buf...), tt.host)
|
||||
if !bytes.Equal(got, tt.want) {
|
||||
t.Fatalf("ReplaceMcHost() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func mcTestPacket(host string, tail ...byte) []byte {
|
||||
packet := []byte{
|
||||
byte(4 + 1 + len(host) + len(tail)),
|
||||
0x00,
|
||||
0x00,
|
||||
0x00,
|
||||
byte(len(host)),
|
||||
}
|
||||
packet = append(packet, host...)
|
||||
packet = append(packet, tail...)
|
||||
return packet
|
||||
}
|
||||
Reference in New Issue
Block a user