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