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

This commit is contained in:
2026-06-23 13:06:31 +08:00
parent ca9241595e
commit 935fe4b808
16 changed files with 1918 additions and 0 deletions

158
cmd/gateway/config_test.go Normal file
View 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")
}
})
}

View 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")
}

View 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
View 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
View 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")
}
}

View 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
View 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
View 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
View 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
View 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)
}
}

View 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)
}
}

View 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
View 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
View 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
View 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
View 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
}