From 935fe4b8083d1ebba4890ab2709a75c1d3058573 Mon Sep 17 00:00:00 2001 From: tursom Date: Tue, 23 Jun 2026 13:06:31 +0800 Subject: [PATCH] =?UTF-8?q?test:=20=E8=A1=A5=E5=85=A8=E5=8D=95=E5=85=83?= =?UTF-8?q?=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- cmd/gateway/config_test.go | 158 ++++++++++++++++ cmd/gateway/handle_request_test.go | 97 ++++++++++ cmd/gateway/haproxy_test.go | 68 +++++++ cmd/gateway/log_pid_test.go | 137 ++++++++++++++ cmd/gateway/main_test.go | 163 +++++++++++++++++ cmd/gateway/pid_unix_test.go | 13 ++ cmd/gateway/plugin_test.go | 282 +++++++++++++++++++++++++++++ cmd/gateway/quic_test.go | 51 ++++++ cmd/gateway/relay_test.go | 207 +++++++++++++++++++++ cmd/gateway/tcp_test.go | 59 ++++++ cmd/gateway/test_helpers_test.go | 164 +++++++++++++++++ cmd/gateway/websocket_test.go | 135 ++++++++++++++ cmd/kcp/main_test.go | 94 ++++++++++ cmd/quic/main_test.go | 94 ++++++++++ plugin/api/api_test.go | 91 ++++++++++ protocol/mc_test.go | 105 +++++++++++ 16 files changed, 1918 insertions(+) create mode 100644 cmd/gateway/config_test.go create mode 100644 cmd/gateway/handle_request_test.go create mode 100644 cmd/gateway/haproxy_test.go create mode 100644 cmd/gateway/log_pid_test.go create mode 100644 cmd/gateway/main_test.go create mode 100644 cmd/gateway/pid_unix_test.go create mode 100644 cmd/gateway/plugin_test.go create mode 100644 cmd/gateway/quic_test.go create mode 100644 cmd/gateway/relay_test.go create mode 100644 cmd/gateway/tcp_test.go create mode 100644 cmd/gateway/test_helpers_test.go create mode 100644 cmd/gateway/websocket_test.go create mode 100644 cmd/kcp/main_test.go create mode 100644 cmd/quic/main_test.go create mode 100644 plugin/api/api_test.go create mode 100644 protocol/mc_test.go diff --git a/cmd/gateway/config_test.go b/cmd/gateway/config_test.go new file mode 100644 index 0000000..d455d4e --- /dev/null +++ b/cmd/gateway/config_test.go @@ -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") + } + }) +} diff --git a/cmd/gateway/handle_request_test.go b/cmd/gateway/handle_request_test.go new file mode 100644 index 0000000..bf66410 --- /dev/null +++ b/cmd/gateway/handle_request_test.go @@ -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") +} diff --git a/cmd/gateway/haproxy_test.go b/cmd/gateway/haproxy_test.go new file mode 100644 index 0000000..23ce397 --- /dev/null +++ b/cmd/gateway/haproxy_test.go @@ -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) + } +} diff --git a/cmd/gateway/log_pid_test.go b/cmd/gateway/log_pid_test.go new file mode 100644 index 0000000..84be2f7 --- /dev/null +++ b/cmd/gateway/log_pid_test.go @@ -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) + } +} diff --git a/cmd/gateway/main_test.go b/cmd/gateway/main_test.go new file mode 100644 index 0000000..d83b1a3 --- /dev/null +++ b/cmd/gateway/main_test.go @@ -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") + } +} diff --git a/cmd/gateway/pid_unix_test.go b/cmd/gateway/pid_unix_test.go new file mode 100644 index 0000000..79b132e --- /dev/null +++ b/cmd/gateway/pid_unix_test.go @@ -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) + } +} diff --git a/cmd/gateway/plugin_test.go b/cmd/gateway/plugin_test.go new file mode 100644 index 0000000..39d5309 --- /dev/null +++ b/cmd/gateway/plugin_test.go @@ -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) +} diff --git a/cmd/gateway/quic_test.go b/cmd/gateway/quic_test.go new file mode 100644 index 0000000..c4b6fee --- /dev/null +++ b/cmd/gateway/quic_test.go @@ -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) + } +} diff --git a/cmd/gateway/relay_test.go b/cmd/gateway/relay_test.go new file mode 100644 index 0000000..bf54210 --- /dev/null +++ b/cmd/gateway/relay_test.go @@ -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 +} diff --git a/cmd/gateway/tcp_test.go b/cmd/gateway/tcp_test.go new file mode 100644 index 0000000..3bd7d0c --- /dev/null +++ b/cmd/gateway/tcp_test.go @@ -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) + } +} diff --git a/cmd/gateway/test_helpers_test.go b/cmd/gateway/test_helpers_test.go new file mode 100644 index 0000000..4acc66d --- /dev/null +++ b/cmd/gateway/test_helpers_test.go @@ -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) + } +} diff --git a/cmd/gateway/websocket_test.go b/cmd/gateway/websocket_test.go new file mode 100644 index 0000000..1b2c0e2 --- /dev/null +++ b/cmd/gateway/websocket_test.go @@ -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 +} diff --git a/cmd/kcp/main_test.go b/cmd/kcp/main_test.go new file mode 100644 index 0000000..f0232a0 --- /dev/null +++ b/cmd/kcp/main_test.go @@ -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 +} diff --git a/cmd/quic/main_test.go b/cmd/quic/main_test.go new file mode 100644 index 0000000..865c9a1 --- /dev/null +++ b/cmd/quic/main_test.go @@ -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 +} diff --git a/plugin/api/api_test.go b/plugin/api/api_test.go new file mode 100644 index 0000000..ba16e57 --- /dev/null +++ b/plugin/api/api_test.go @@ -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 +} diff --git a/protocol/mc_test.go b/protocol/mc_test.go new file mode 100644 index 0000000..3443eb4 --- /dev/null +++ b/protocol/mc_test.go @@ -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 +}