feat: add sqlite-backed admin management
Some checks failed
Go / build (.exe, 386, windows, windows-386) (push) Has been cancelled
Go / build (.exe, amd64, windows, windows-amd64) (push) Has been cancelled
Go / build (.exe, arm64, windows, windows-arm64) (push) Has been cancelled
Go / build (386, freebsd, freebsd-386) (push) Has been cancelled
Go / build (386, linux, linux-386) (push) Has been cancelled
Go / build (386, netbsd, netbsd-386) (push) Has been cancelled
Go / build (386, openbsd, openbsd-386) (push) Has been cancelled
Go / build (386, plan9, plan9-386) (push) Has been cancelled
Go / build (amd64, darwin, darwin-amd64) (push) Has been cancelled
Go / build (amd64, dragonfly, dragonfly-amd64) (push) Has been cancelled
Go / build (amd64, freebsd, freebsd-amd64) (push) Has been cancelled
Go / build (amd64, illumos, illumos-amd64) (push) Has been cancelled
Go / build (amd64, linux, linux-amd64) (push) Has been cancelled
Go / build (amd64, netbsd, netbsd-amd64) (push) Has been cancelled
Go / build (amd64, openbsd, openbsd-amd64) (push) Has been cancelled
Go / build (amd64, plan9, plan9-amd64) (push) Has been cancelled
Go / build (amd64, solaris, solaris-amd64) (push) Has been cancelled
Go / build (arm, 6, linux, linux-armv6) (push) Has been cancelled
Go / build (arm, 7, linux, linux-armv7) (push) Has been cancelled
Go / build (arm, freebsd, freebsd-arm) (push) Has been cancelled
Go / build (arm, netbsd, netbsd-arm) (push) Has been cancelled
Go / build (arm, openbsd, openbsd-arm) (push) Has been cancelled
Go / build (arm, plan9, plan9-arm) (push) Has been cancelled
Go / build (arm64, darwin, darwin-arm64) (push) Has been cancelled
Go / build (arm64, freebsd, freebsd-arm64) (push) Has been cancelled
Go / build (arm64, linux, linux-arm64) (push) Has been cancelled
Go / build (arm64, netbsd, netbsd-arm64) (push) Has been cancelled
Go / build (arm64, openbsd, openbsd-arm64) (push) Has been cancelled
Go / build (loong64, linux, linux-loong64) (push) Has been cancelled
Go / build (mips, linux, linux-mips) (push) Has been cancelled
Go / build (mips64, linux, linux-mips64) (push) Has been cancelled
Go / build (mips64le, linux, linux-mips64le) (push) Has been cancelled
Go / build (mipsle, linux, linux-mipsle) (push) Has been cancelled
Go / build (ppc64, aix, aix-ppc64) (push) Has been cancelled
Go / build (ppc64, linux, linux-ppc64) (push) Has been cancelled
Go / build (ppc64, openbsd, openbsd-ppc64) (push) Has been cancelled
Go / build (ppc64le, linux, linux-ppc64le) (push) Has been cancelled
Go / build (riscv64, freebsd, freebsd-riscv64) (push) Has been cancelled
Go / build (riscv64, linux, linux-riscv64) (push) Has been cancelled
Go / build (riscv64, openbsd, openbsd-riscv64) (push) Has been cancelled
Go / build (s390x, linux, linux-s390x) (push) Has been cancelled
Go / merge-artifacts (push) Has been cancelled

This commit is contained in:
2026-06-25 08:53:44 +08:00
parent ada051341d
commit 13ab1f6964
73 changed files with 6926 additions and 782 deletions

129
README.md
View File

@@ -1,52 +1,70 @@
# mc-gateway
一个简易的 Minecraft 网关,通过 host 将客户端的流量转发到对应的后端 Minecraft 服务器。
一个简易的 Minecraft 网关,通过客户端握手里的 host 将流量转发到对应的后端 Minecraft 服务器。
## 配置
## 启动
目前 mc-gateway 只支持读取当前目录的 `config.toml` 作为配置。在 `config.toml` 被修改时,可以自动加载并更新部分配置,以达到不停机修改配置的效果。支持热加载的配置有
mc-gateway 默认不依赖配置文件。直接启动后会在 `25565` 端口同时提供 Minecraft TCP 转发入口和后台管理入口
- hosts
- log
- KCP 的 data_shards 和 parity_Shards
- QUIC 的 application_protocols
- pid_file
- Admin 页面:`/admin/`
- Admin API`/admin/api`
- SQLite 数据库:`mc-gateway.sqlite3`
### 顶层配置
首次启动时,如果用户表为空,可以通过 Admin 页面初始化管理员账号。也可以用 `MC_GATEWAY_ADMIN_PASSWORD` 在启动时创建默认管理员用户 `admin`
| 配置 | 类型 | 备注 |
| -------- | ------ | -------- |
| pid_file | string | pid 文件 |
### 启动期环境变量
> pid_file 在非 windows 平台默认会写入 /var/run/mc-gateway.pid
> 在 windows 平台默认不会写入任何文件
| 环境变量 | 默认值 | 说明 |
| --- | --- | --- |
| `MC_GATEWAY_TCP_ADMIN_PORT` | `25565` | TCP/Admin 共享监听端口 |
| `MC_GATEWAY_ADMIN_PATH` | `/admin/` | Admin 页面路径 |
| `MC_GATEWAY_ADMIN_API_PREFIX` | `/admin/api` | Admin API 前缀 |
| `MC_GATEWAY_DB` | `mc-gateway.sqlite3` | SQLite 数据库路径 |
| `MC_GATEWAY_ADMIN_PASSWORD` | 空 | 首次启动时创建默认管理员密码 |
### hosts
服务启停、KCP/QUIC/WebSocket 参数、用户、权限和路由都通过后台管理写入 SQLite不再使用 `config.toml` 作为启动配置或路由来源。
hosts 使用期望的 host 做 key转发的目的地址为 value。参考`config.example.toml`。默认的 fallback host 配置 key 为 `default`
## 权限
#### upstream
后台管理内置三类角色:
host 支持多种协议的上游服务器,包括 tcp、kcp、quic、websocket 等。只需要在原始地址前加对应的协议名称即可,如 `quic://127.0.0.1:8080`。支持的列表如下:
| 角色 | 能力 |
| --- | --- |
| 管理员 | 管理用户、服务、路由和审计日志 |
| 成员 | 查看状态,管理路由 |
| 游客 | 查看当前路由 |
| 协议名称 | 前缀 | 备注 |
| -------- | ---------- | ----------------------------- |
| tcp | 无 | 原始的tcp连接 |
| kcp | kcp:// | 使用 kcp 协议连接到服务器 |
| quic | quic:// | 使用 quic 协议连接到服务器 |
| haproxy | haproxy:// | 使用 HAProxy 协议连接到服务器 |
## 路由
> HAProxy 协议头会保存客户端的真实 ip大部分支持 HAProxy 的 mod或插件都支持从协议头获取真实 ip
> 这样服务端就能够获取到真实的客户端 ip了以此兼容现有的 ban ip 或者统计等插件。
路由记录存储在 SQLite 中,后台修改后会刷新内存快照。默认 fallback 路由的 host 为 `default`
### log
### upstream
| 配置 | 类型 | 备注 |
| ----- | ------ | -------- |
| level | Level | 日志等级 |
| file | string | 日志文件 |
路由上游支持多种协议。TCP 上游直接填写地址,其他协议在地址前加协议前缀:
> 日志适配 logrotate可以使用 logrotate 进行日志分片、压缩等日常运维操作,参考配置:
| 协议 | 前缀 | 说明 |
| --- | --- | --- |
| tcp | 无 | 原始 TCP 连接 |
| kcp | `kcp://` | 使用 KCP 协议连接到服务器 |
| quic | `quic://` | 使用 QUIC 协议连接到服务器 |
| haproxy | `haproxy://` | 使用 HAProxy 协议连接到服务器 |
HAProxy 协议头会保存客户端真实 IP适合需要在后端服务端获取真实客户端 IP 的场景。
## 可选服务
TCP/Admin listener 是基础入口默认启用。KCP、QUIC、WebSocket 默认禁用,可以在后台管理中启用并配置端口和参数;配置修改后第一版按重启后生效处理。
| 服务 | 默认端口 | 说明 |
| --- | --- | --- |
| TCP/Admin | `25565` | Minecraft TCP 转发和后台管理共享入口 |
| KCP | `25565` | 可选 KCP 入口 |
| QUIC | `25565` | 可选 QUIC 入口 |
| WebSocket | `25566` | 可选 WebSocket 入口 |
## 日志
默认日志级别为 `info`,输出到标准输出。日志文件重开逻辑支持 logrotate 场景:
```logrotate
/var/log/mc-gateway.log {
@@ -63,52 +81,9 @@ host 支持多种协议的上游服务器,包括 tcp、kcp、quic、websocket
delaycompress
dateext
postrotate
# 向程序发送 SIGHUP 信号
# mc-gateway 默认会将当前进程的 pid 写入 /var/run/mc-gateway.pid
if [ -f /var/run/mc-gateway.pid ]; then
kill -SIGHUP $(cat /var/run/mc-gateway.pid)
if [ -f /dev/shm/mc-gateway.pid ]; then
kill -SIGHUP $(cat /dev/shm/mc-gateway.pid)
fi
endscript
}
```
### tcp
| 配置 | 类型 | 备注 |
| ------ | ---- | -------- |
| enable | bool | 是否启用 |
| port | int | 端口 |
> 默认端口为 25565与 Minecraft 服务端保持一致
### kcp
| 配置 | 类型 | 备注 |
| ------------- | ---- | -------- |
| enable | bool | 是否启用 |
| port | int | 端口 |
| data_shards | int | 数据分片 |
| parity_Shards | int | 校验分片 |
### quic
| 配置 | 类型 | 备注 |
| --------------------- | -------- | ------------ |
| enable | bool | 是否启用 |
| port | int | 端口 |
| application_protocols | []string | 应用协议列表 |
> application_protocols 只要客户端与服务端有一个能够对应上就可以成功连接
> 默认值为 ["minecraft", "quic", "raw", "h3"]
### websocket
| 配置 | 类型 | 备注 |
| ------ | ---- | -------- |
| enable | bool | 是否启用 |
| port | int | 端口 |
| path | str | 接口路径 |
> path 默认为 "/",会对所有路径的请求进行处理
>
> port 默认为 25566

32
cmd/gateway/admin_api.go Normal file
View File

@@ -0,0 +1,32 @@
package main
import (
"net/http"
"github.com/tursom/mc-gateway/internal/adminhttp"
)
func newAdminAPIHandler() http.HandlerFunc {
return adminhttp.NewAPIHandler(adminStartup.AdminAPIPrefix, adminhttp.APIHandlers{
SetupStatus: handleAdminSetupStatus,
Setup: handleAdminSetup,
Login: handleAdminLogin,
Logout: handleAdminLogout,
Me: handleAdminMe,
Status: handleAdminStatus,
RoutesList: handleAdminRoutesList,
RouteItem: handleAdminRouteItem,
ServicesList: handleAdminServicesList,
ServiceItem: handleAdminServiceItem,
Metrics: handleAdminMetrics,
UsersList: handleAdminUsersList,
UsersCreate: handleAdminUsersCreate,
UserItem: handleAdminUserItem,
AuditLogs: handleAdminAuditLogs,
})
}

View File

@@ -0,0 +1,281 @@
package main
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
)
func TestAdminSetupLoginAndPermissions(t *testing.T) {
handler := newAdminTestHandler(t)
resp := adminTestRequest(t, handler, http.MethodGet, "/admin/api/setup", "", nil)
if resp.Code != http.StatusOK {
t.Fatalf("GET setup status = %d, want %d", resp.Code, http.StatusOK)
}
if got := adminTestJSON(t, resp)["required"]; got != true {
t.Fatalf("setup required = %v, want true", got)
}
resp = adminTestRequest(t, handler, http.MethodPost, "/admin/api/setup", "", map[string]any{
"username": "admin",
"password": "secret",
})
if resp.Code != http.StatusCreated {
t.Fatalf("POST setup status = %d, body=%s", resp.Code, resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodPost, "/admin/api/setup", "", map[string]any{
"username": "admin2",
"password": "secret",
})
if resp.Code != http.StatusConflict {
t.Fatalf("repeat setup status = %d, want %d", resp.Code, http.StatusConflict)
}
adminToken := adminTestLogin(t, handler, "admin", "secret")
resp = adminTestRequest(t, handler, http.MethodGet, "/admin/api/status", adminToken, nil)
if resp.Code != http.StatusOK {
t.Fatalf("admin status = %d, body=%s", resp.Code, resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodPost, "/admin/api/users", adminToken, map[string]any{
"username": "guest",
"role": "guest",
"password": "guest-secret",
})
if resp.Code != http.StatusCreated {
t.Fatalf("create guest status = %d, body=%s", resp.Code, resp.Body.String())
}
guestToken := adminTestLogin(t, handler, "guest", "guest-secret")
resp = adminTestRequest(t, handler, http.MethodGet, "/admin/api/routes", guestToken, nil)
if resp.Code != http.StatusOK {
t.Fatalf("guest routes status = %d, body=%s", resp.Code, resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodPut, "/admin/api/routes/play.example", guestToken, map[string]any{
"upstream": "127.0.0.1:25565",
"enabled": true,
})
if resp.Code != http.StatusForbidden {
t.Fatalf("guest route write status = %d, want %d", resp.Code, http.StatusForbidden)
}
resp = adminTestRequest(t, handler, http.MethodGet, "/admin/api/users", guestToken, nil)
if resp.Code != http.StatusForbidden {
t.Fatalf("guest users status = %d, want %d", resp.Code, http.StatusForbidden)
}
}
func TestAdminRoutesRefreshSnapshot(t *testing.T) {
handler := newAdminTestHandlerWithAdmin(t)
token := adminTestLogin(t, handler, "admin", "secret")
resp := adminTestRequest(t, handler, http.MethodPut, "/admin/api/routes/play.example", token, map[string]any{
"upstream": "127.0.0.1:25565",
"enabled": true,
"note": "primary",
})
if resp.Code != http.StatusOK {
t.Fatalf("route upsert status = %d, body=%s", resp.Code, resp.Body.String())
}
upstream, ok := lookupRoute("play.example")
if !ok || upstream != "127.0.0.1:25565" {
t.Fatalf("lookupRoute() = %q, %v; want route", upstream, ok)
}
resp = adminTestRequest(t, handler, http.MethodDelete, "/admin/api/routes/play.example", token, nil)
if resp.Code != http.StatusOK {
t.Fatalf("route delete status = %d, body=%s", resp.Code, resp.Body.String())
}
if upstream, ok := lookupRoute("play.example"); ok || upstream != "" {
t.Fatalf("lookupRoute() after delete = %q, %v; want miss", upstream, ok)
}
}
func TestAdminCustomPathAndAPIPrefixFromEnv(t *testing.T) {
t.Cleanup(saveGatewayState(t))
t.Setenv(adminEnvDB, filepath.Join(t.TempDir(), "gateway.sqlite3"))
t.Setenv(adminEnvPath, "/ops")
t.Setenv(adminEnvAPIPrefix, "/ops/api")
if err := loadConfig(); err != nil {
t.Fatalf("loadConfig() error = %v", err)
}
handler := newGatewayHTTPHandler()
resp := adminTestRequest(t, handler, http.MethodGet, "/ops", "", nil)
if resp.Code != http.StatusMovedPermanently {
t.Fatalf("admin path redirect status = %d, want %d", resp.Code, http.StatusMovedPermanently)
}
if got := resp.Header().Get("Location"); got != "/ops/" {
t.Fatalf("admin path redirect location = %q, want /ops/", got)
}
resp = adminTestRequest(t, handler, http.MethodGet, "/ops/", "", nil)
if resp.Code != http.StatusOK {
t.Fatalf("custom admin page status = %d, body=%s", resp.Code, resp.Body.String())
}
if !strings.Contains(resp.Body.String(), `data-api-prefix="/ops/api"`) {
t.Fatalf("custom admin page does not contain API prefix: %s", resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodGet, "/ops/api/setup", "", nil)
if resp.Code != http.StatusOK {
t.Fatalf("custom setup status = %d, body=%s", resp.Code, resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodGet, "/admin/api/setup", "", nil)
if resp.Code != http.StatusNotFound {
t.Fatalf("default setup status = %d, want %d", resp.Code, http.StatusNotFound)
}
}
func TestAdminServiceUpdateMarksRestartRequired(t *testing.T) {
handler := newAdminTestHandlerWithAdmin(t)
token := adminTestLogin(t, handler, "admin", "secret")
resp := adminTestRequest(t, handler, http.MethodPut, "/admin/api/services/kcp", token, map[string]any{
"enabled": true,
"port": 25570,
"options": map[string]any{
"data_shards": 12,
"parity_shards": 4,
},
})
if resp.Code != http.StatusOK {
t.Fatalf("service update status = %d, body=%s", resp.Code, resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodGet, "/admin/api/services", token, nil)
if resp.Code != http.StatusOK {
t.Fatalf("services status = %d, body=%s", resp.Code, resp.Body.String())
}
body := adminTestJSON(t, resp)
services := body["services"].([]any)
var found map[string]any
for _, item := range services {
service := item.(map[string]any)
if service["name"] == "kcp" {
found = service
break
}
}
if found == nil {
t.Fatal("kcp service not found")
}
if found["enabled"] != true || found["restart_required"] != true || found["running"] != false {
t.Fatalf("kcp service = %#v, want enabled restart_required and not running", found)
}
}
func TestAdminUserPatchInvalidatesExistingSession(t *testing.T) {
handler := newAdminTestHandlerWithAdmin(t)
adminToken := adminTestLogin(t, handler, "admin", "secret")
resp := adminTestRequest(t, handler, http.MethodPost, "/admin/api/users", adminToken, map[string]any{
"username": "member",
"role": "member",
"password": "member-secret",
})
if resp.Code != http.StatusCreated {
t.Fatalf("create member status = %d, body=%s", resp.Code, resp.Body.String())
}
memberToken := adminTestLogin(t, handler, "member", "member-secret")
resp = adminTestRequest(t, handler, http.MethodGet, "/admin/api/status", memberToken, nil)
if resp.Code != http.StatusOK {
t.Fatalf("member status before patch = %d, body=%s", resp.Code, resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodPatch, "/admin/api/users/member", adminToken, map[string]any{
"role": "guest",
})
if resp.Code != http.StatusOK {
t.Fatalf("patch member status = %d, body=%s", resp.Code, resp.Body.String())
}
resp = adminTestRequest(t, handler, http.MethodGet, "/admin/api/status", memberToken, nil)
if resp.Code != http.StatusUnauthorized {
t.Fatalf("old member token status = %d, want %d", resp.Code, http.StatusUnauthorized)
}
}
func newAdminTestHandler(t *testing.T) http.Handler {
t.Helper()
t.Cleanup(saveGatewayState(t))
t.Setenv(adminEnvDB, filepath.Join(t.TempDir(), "gateway.sqlite3"))
if err := loadConfig(); err != nil {
t.Fatalf("loadConfig() error = %v", err)
}
return newGatewayHTTPHandler()
}
func newAdminTestHandlerWithAdmin(t *testing.T) http.Handler {
t.Helper()
handler := newAdminTestHandler(t)
resp := adminTestRequest(t, handler, http.MethodPost, "/admin/api/setup", "", map[string]any{
"username": "admin",
"password": "secret",
})
if resp.Code != http.StatusCreated {
t.Fatalf("setup status = %d, body=%s", resp.Code, resp.Body.String())
}
return handler
}
func adminTestLogin(t *testing.T, handler http.Handler, username, password string) string {
t.Helper()
resp := adminTestRequest(t, handler, http.MethodPost, "/admin/api/auth/login", "", map[string]any{
"username": username,
"password": password,
})
if resp.Code != http.StatusOK {
t.Fatalf("login status = %d, body=%s", resp.Code, resp.Body.String())
}
body := adminTestJSON(t, resp)
token, ok := body["token"].(string)
if !ok || token == "" {
t.Fatalf("login token = %#v", body["token"])
}
return token
}
func adminTestRequest(t *testing.T, handler http.Handler, method, target, token string, body any) *httptest.ResponseRecorder {
t.Helper()
var payload bytes.Buffer
if body != nil {
if err := json.NewEncoder(&payload).Encode(body); err != nil {
t.Fatalf("Encode() error = %v", err)
}
}
req := httptest.NewRequest(method, target, &payload)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, req)
return resp
}
func adminTestJSON(t *testing.T, resp *httptest.ResponseRecorder) map[string]any {
t.Helper()
var body map[string]any
if err := json.Unmarshal(resp.Body.Bytes(), &body); err != nil {
t.Fatalf("Unmarshal(%q) error = %v", resp.Body.String(), err)
}
return body
}

View File

@@ -0,0 +1,15 @@
package main
import (
"context"
"github.com/tursom/mc-gateway/internal/adminaudit"
)
func recordAudit(ctx context.Context, actor, sourceIP, action, targetType, targetID string, success bool, message string) {
_ = adminaudit.NewRepository(adminDB).Record(ctx, actor, sourceIP, action, targetType, targetID, success, message)
}
func listAuditLogs(ctx context.Context) ([]adminaudit.Record, error) {
return adminaudit.NewRepository(adminDB).List(ctx, adminaudit.DefaultListLimit)
}

View File

@@ -0,0 +1,113 @@
package main
import (
"net/http"
"os"
"strings"
"time"
"github.com/tursom/mc-gateway/internal/adminhttp"
"github.com/tursom/mc-gateway/internal/adminuser"
)
func handleAdminSetupStatus(w http.ResponseWriter, r *http.Request) {
empty, err := usersTableEmpty(r.Context(), adminDB)
if err != nil {
adminhttp.WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"required": empty})
}
func handleAdminSetup(w http.ResponseWriter, r *http.Request) {
var req adminhttp.SetupRequest
if !adminhttp.DecodeJSONRequest(w, r, &req) {
return
}
if strings.TrimSpace(req.Username) == "" {
req.Username = "admin"
}
err := createInitialAdmin(r.Context(), req.Username, req.Password)
if err != nil {
recordAudit(r.Context(), "setup", adminhttp.RequestSourceIP(r), "setup", "user", req.Username, false, err.Error())
status := http.StatusBadRequest
if strings.Contains(err.Error(), "already") {
status = http.StatusConflict
}
adminhttp.WriteAPIError(w, status, err.Error())
return
}
recordAudit(r.Context(), "setup", adminhttp.RequestSourceIP(r), "setup", "user", req.Username, true, "created initial admin")
adminhttp.WriteJSON(w, http.StatusCreated, map[string]any{"ok": true})
}
func handleAdminLogin(w http.ResponseWriter, r *http.Request) {
var req adminhttp.LoginRequest
if !adminhttp.DecodeJSONRequest(w, r, &req) {
return
}
user, err := authenticateUser(r.Context(), req.Username, req.Password)
if err != nil {
recordAudit(r.Context(), req.Username, adminhttp.RequestSourceIP(r), "login", "user", req.Username, false, err.Error())
adminhttp.WriteAPIError(w, http.StatusUnauthorized, err.Error())
return
}
session, err := createSession(user.Username, user.Role)
if err != nil {
adminhttp.WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
recordAudit(r.Context(), user.Username, adminhttp.RequestSourceIP(r), "login", "user", user.Username, true, "login success")
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{
"token": session.Token,
"expires_at": session.ExpiresAt.Unix(),
"user": user,
})
}
func handleAdminLogout(w http.ResponseWriter, r *http.Request) {
session, ok := requireSession(w, r)
if !ok {
return
}
deleteSession(session.Token)
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "logout", "session", session.Username, true, "logout success")
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"ok": true})
}
func handleAdminMe(w http.ResponseWriter, r *http.Request) {
session, ok := requireSession(w, r)
if !ok {
return
}
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{
"username": session.Username,
"role": session.Role,
"expires_at": session.ExpiresAt.Unix(),
"permissions": adminuser.Permissions(session.Role),
})
}
func handleAdminStatus(w http.ResponseWriter, r *http.Request) {
if _, ok := requireRole(w, r, adminRoleMember); !ok {
return
}
services, err := listServiceConfigs(r.Context(), adminDB)
if err != nil {
adminhttp.WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{
"pid": os.Getpid(),
"uptime_seconds": int64(time.Since(processStartAt).Seconds()),
"db_path": adminDBPath,
"tcp_admin_port": adminStartup.TCPAdminPort,
"admin_path": adminStartup.AdminPath,
"admin_api_prefix": adminStartup.AdminAPIPrefix,
"services": services,
})
}

View File

@@ -0,0 +1,17 @@
package main
import (
"net/http"
"github.com/tursom/mc-gateway/internal/adminhttp"
"github.com/tursom/mc-gateway/internal/gatewaymetrics"
)
var gatewayMetrics = gatewaymetrics.New()
func handleAdminMetrics(w http.ResponseWriter, r *http.Request) {
if _, ok := requireRole(w, r, adminRoleMember); !ok {
return
}
adminhttp.WriteJSON(w, http.StatusOK, gatewayMetrics.Snapshot())
}

View File

@@ -0,0 +1,62 @@
package main
import (
"net/http"
"github.com/tursom/mc-gateway/internal/adminhttp"
)
func handleAdminRoutesList(w http.ResponseWriter, r *http.Request) {
if _, ok := requireRole(w, r, adminRoleGuest); !ok {
return
}
routes, err := listRoutes(r.Context(), r.URL.Query().Get("q"))
if err != nil {
adminhttp.WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"routes": routes})
}
func handleAdminRouteItem(w http.ResponseWriter, r *http.Request, rawHost string) {
session, ok := requireRole(w, r, adminRoleMember)
if !ok {
return
}
host, err := adminhttp.PathSegment(rawHost)
if err != nil {
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
switch r.Method {
case http.MethodPut:
var req adminhttp.RouteRequest
if !adminhttp.DecodeJSONRequest(w, r, &req) {
return
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
err := upsertRoute(r.Context(), session.Username, host, req.Upstream, enabled, req.Note)
if err != nil {
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "route_upsert", "route", host, false, err.Error())
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "route_upsert", "route", host, true, "route saved")
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"ok": true})
case http.MethodDelete:
err := deleteRoute(r.Context(), session.Username, host)
if err != nil {
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "route_delete", "route", host, false, err.Error())
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "route_delete", "route", host, true, "route deleted")
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"ok": true})
default:
adminhttp.WriteAPIError(w, http.StatusMethodNotAllowed, "method not allowed")
}
}

View File

@@ -0,0 +1,60 @@
package main
import (
"context"
"sync"
"github.com/tursom/mc-gateway/internal/adminroute"
)
var (
routeSnapshot = adminroute.NewSnapshot()
routeWriteLock sync.Mutex
)
func refreshRouteSnapshot(ctx context.Context) error {
if adminDB == nil {
publishRouteSnapshot(map[string]string{})
return nil
}
routes, err := adminroute.NewRepository(adminDB).EnabledMap(ctx)
if err != nil {
return err
}
publishRouteSnapshot(routes)
return nil
}
func publishRouteSnapshot(routes map[string]string) {
routeSnapshot.Store(routes)
}
func lookupRoute(host string) (string, bool) {
return routeSnapshot.Lookup(host)
}
func listRoutes(ctx context.Context, query string) ([]adminroute.Record, error) {
return adminroute.NewRepository(adminDB).List(ctx, query)
}
func upsertRoute(ctx context.Context, actor, host, upstream string, enabled bool, note string) error {
routeWriteLock.Lock()
defer routeWriteLock.Unlock()
if err := adminroute.NewRepository(adminDB).Upsert(ctx, actor, host, upstream, enabled, note); err != nil {
return err
}
return refreshRouteSnapshot(ctx)
}
func deleteRoute(ctx context.Context, actor, host string) error {
routeWriteLock.Lock()
defer routeWriteLock.Unlock()
if err := adminroute.NewRepository(adminDB).Delete(ctx, host); err != nil {
return err
}
_ = actor
return refreshRouteSnapshot(ctx)
}

View File

@@ -0,0 +1,92 @@
package main
import (
"context"
"database/sql"
"os"
"time"
"github.com/tursom/mc-gateway/internal/adminconfig"
"github.com/tursom/mc-gateway/internal/admindb"
"github.com/tursom/mc-gateway/internal/adminservice"
)
const (
defaultAdminDBPath = adminconfig.DefaultDBPath
defaultAdminPath = adminconfig.DefaultAdminPath
defaultAdminAPIPrefix = adminconfig.DefaultAdminAPIPrefix
defaultAdminSessionTTL = adminconfig.DefaultSessionTTL
defaultKCPPort = adminservice.DefaultKCPPort
defaultKCPDataShards = adminservice.DefaultKCPDataShards
defaultKCPParityShards = adminservice.DefaultKCPParityShards
defaultQUICPort = adminservice.DefaultQUICPort
defaultWebSocketPort = adminservice.DefaultWebSocketPort
defaultWebSocketPath = adminservice.DefaultWebSocketPath
serviceNameTCPAdmin = adminservice.NameTCPAdmin
serviceNameKCP = adminservice.NameKCP
serviceNameQUIC = adminservice.NameQUIC
serviceNameWebSocket = adminservice.NameWebSocket
adminEnvDB = adminconfig.EnvDB
adminEnvTCPAdminPort = adminconfig.EnvTCPAdminPort
adminEnvPath = adminconfig.EnvPath
adminEnvAPIPrefix = adminconfig.EnvAPIPrefix
adminEnvInitialPassword = adminconfig.EnvInitialPassword
)
var (
adminStartup = adminconfig.Config{
DBPath: defaultAdminDBPath,
TCPAdminPort: defaultTCPPort,
AdminPath: defaultAdminPath,
AdminAPIPrefix: defaultAdminAPIPrefix,
SessionTTL: defaultAdminSessionTTL,
}
adminDB *sql.DB
adminDBPath string
processStartAt = time.Now()
)
func initializeGatewayRuntime() error {
startup, err := parseStartupConfig(os.Getenv)
if err != nil {
return err
}
adminStartup = startup
db, err := admindb.Open(startup.DBPath)
if err != nil {
return err
}
if adminDB != nil && adminDB != db {
_ = adminDB.Close()
}
adminDB = db
adminDBPath = startup.DBPath
if err := admindb.Migrate(db); err != nil {
return err
}
if err := ensureDefaultServices(context.Background(), db, startup.TCPAdminPort); err != nil {
return err
}
if err := applyServiceConfig(context.Background(), db); err != nil {
return err
}
if err := ensureInitialAdminFromEnv(context.Background(), db, os.Getenv(adminEnvInitialPassword)); err != nil {
return err
}
return refreshRouteSnapshot(context.Background())
}
func closeGatewayRuntime() {
if adminDB != nil {
_ = adminDB.Close()
adminDB = nil
}
}
func parseStartupConfig(getenv func(string) string) (adminconfig.Config, error) {
return adminconfig.Parse(getenv)
}

View File

@@ -0,0 +1,74 @@
package main
import (
"net/http"
"strings"
"github.com/tursom/mc-gateway/internal/adminhttp"
"github.com/tursom/mc-gateway/internal/adminservice"
)
func handleAdminServicesList(w http.ResponseWriter, r *http.Request) {
if _, ok := requireRole(w, r, adminRoleMember); !ok {
return
}
services, err := listServiceConfigs(r.Context(), adminDB)
if err != nil {
adminhttp.WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"services": services})
}
func handleAdminServiceItem(w http.ResponseWriter, r *http.Request, rawName string) {
session, ok := requireRole(w, r, adminRoleAdmin)
if !ok {
return
}
if strings.HasSuffix(rawName, "/restart") {
name := strings.TrimSuffix(rawName, "/restart")
if r.Method != http.MethodPost {
adminhttp.WriteAPIError(w, http.StatusMethodNotAllowed, "method not allowed")
return
}
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "service_restart", "service", name, false, "restart is not implemented")
adminhttp.WriteAPIError(w, http.StatusNotImplemented, "service restart is not implemented; restart the gateway process")
return
}
name, err := adminhttp.PathSegment(rawName)
if err != nil {
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
if r.Method != http.MethodPut {
adminhttp.WriteAPIError(w, http.StatusMethodNotAllowed, "method not allowed")
return
}
var req adminhttp.ServiceRequest
if !adminhttp.DecodeJSONRequest(w, r, &req) {
return
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
if req.Port == 0 {
req.Port = defaultServicePort(name)
}
err = updateServiceConfig(r.Context(), session.Username, name, enabled, req.Port, req.Options)
if err != nil {
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "service_update", "service", name, false, err.Error())
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "service_update", "service", name, true, "service saved")
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"ok": true, "restart_required": true})
}
func defaultServicePort(name string) int {
return adminservice.DefaultPort(name, adminStartup.TCPAdminPort)
}

View File

@@ -0,0 +1,85 @@
package main
import (
"context"
"database/sql"
"github.com/tursom/mc-gateway/internal/adminservice"
)
func ensureDefaultServices(ctx context.Context, db *sql.DB, tcpAdminPort int) error {
return adminservice.NewRepository(db).EnsureDefaults(ctx, tcpAdminPort)
}
func applyServiceConfig(ctx context.Context, db *sql.DB) error {
services, err := listServiceConfigs(ctx, db)
if err != nil {
return err
}
config.Tcp.Enable = true
config.Tcp.Port = adminStartup.TCPAdminPort
for _, service := range services {
switch service.Name {
case serviceNameKCP:
config.Kcp.Enable = service.Enabled
config.Kcp.Port = service.Port
config.Kcp.DataShards = adminservice.IntOption(service.Options, "data_shards", defaultKCPDataShards)
config.Kcp.ParityShards = adminservice.IntOption(service.Options, "parity_shards", defaultKCPParityShards)
case serviceNameQUIC:
config.Quic.Enable = service.Enabled
config.Quic.Port = service.Port
config.Quic.ApplicationProtocols = adminservice.StringSliceOption(service.Options, "application_protocols")
case serviceNameWebSocket:
config.WebSocket.Enable = service.Enabled
config.WebSocket.Port = service.Port
config.WebSocket.Path = adminservice.StringOption(service.Options, "path", defaultWebSocketPath)
}
}
if config.Kcp.Port == 0 {
config.Kcp.Port = defaultKCPPort
}
if config.Quic.Port == 0 {
config.Quic.Port = defaultQUICPort
}
if config.WebSocket.Port == 0 {
config.WebSocket.Port = defaultWebSocketPort
}
if config.WebSocket.Path == "" {
config.WebSocket.Path = defaultWebSocketPath
}
return nil
}
func listServiceConfigs(ctx context.Context, db *sql.DB) ([]adminservice.Record, error) {
services, err := adminservice.NewRepository(db).List(ctx)
if err != nil {
return nil, err
}
for i := range services {
services[i].Running = serviceIsRunning(services[i])
}
return services, nil
}
func serviceIsRunning(service adminservice.Record) bool {
switch service.Name {
case serviceNameTCPAdmin:
return true
case serviceNameKCP:
return config.Kcp.Enable
case serviceNameQUIC:
return config.Quic.Enable
case serviceNameWebSocket:
return config.WebSocket.Enable
default:
return false
}
}
func updateServiceConfig(ctx context.Context, actor, name string, enabled bool, port int, options map[string]any) error {
return adminservice.NewRepository(adminDB).Update(ctx, actor, name, enabled, port, options)
}

View File

@@ -0,0 +1,58 @@
package main
import (
"net/http"
"strings"
"github.com/tursom/mc-gateway/internal/adminhttp"
"github.com/tursom/mc-gateway/internal/adminsession"
"github.com/tursom/mc-gateway/internal/adminuser"
)
var adminSessionManager = adminsession.NewManager()
func createSession(username, role string) (adminsession.Session, error) {
return adminSessionManager.Create(username, role, adminStartup.SessionTTL)
}
func getSession(token string) (adminsession.Session, bool) {
return adminSessionManager.Get(token)
}
func deleteSession(token string) {
adminSessionManager.Delete(token)
}
func removeSessionsForUser(username string) {
adminSessionManager.RemoveUser(username)
}
func sessionFromRequest(r *http.Request) (adminsession.Session, bool) {
auth := r.Header.Get("Authorization")
token, ok := strings.CutPrefix(auth, "Bearer ")
if !ok || strings.TrimSpace(token) == "" {
return adminsession.Session{}, false
}
return getSession(strings.TrimSpace(token))
}
func requireSession(w http.ResponseWriter, r *http.Request) (adminsession.Session, bool) {
session, ok := sessionFromRequest(r)
if !ok {
adminhttp.WriteAPIError(w, http.StatusUnauthorized, "login required")
return adminsession.Session{}, false
}
return session, true
}
func requireRole(w http.ResponseWriter, r *http.Request, role string) (adminsession.Session, bool) {
session, ok := requireSession(w, r)
if !ok {
return adminsession.Session{}, false
}
if !adminuser.HasRole(session.Role, role) {
adminhttp.WriteAPIError(w, http.StatusForbidden, "permission denied")
return adminsession.Session{}, false
}
return session, true
}

View File

@@ -0,0 +1,23 @@
package main
import (
"embed"
"net/http"
"github.com/tursom/mc-gateway/internal/adminhttp"
)
//go:embed admin_static/index.html admin_static/app.css admin_static/app.js
var adminStaticFS embed.FS
func newGatewayHTTPHandler() http.Handler {
return adminhttp.NewGatewayHandler(adminhttp.GatewayHandlerOptions{
AdminPath: adminStartup.AdminPath,
AdminAPIPrefix: adminStartup.AdminAPIPrefix,
Assets: adminStaticFS,
APIHandler: newAdminAPIHandler(),
WebSocketEnabled: config.WebSocket.Enable,
WebSocketPath: normalizedWebSocketPath(),
WebSocketHandler: handleWebSocket,
})
}

View File

@@ -0,0 +1,377 @@
:root {
color-scheme: light;
--bg: #f6f7f4;
--panel: #ffffff;
--text: #1c2623;
--muted: #66716c;
--line: #d9ded8;
--accent: #1f7a5d;
--accent-dark: #145944;
--warn: #a66321;
--danger: #b42318;
--shadow: 0 16px 40px rgba(21, 32, 28, 0.08);
}
* {
box-sizing: border-box;
}
body {
margin: 0;
background: var(--bg);
color: var(--text);
font: 14px/1.5 system-ui, -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;
}
button,
input,
select {
font: inherit;
}
button {
min-height: 36px;
border: 0;
border-radius: 6px;
padding: 0 14px;
background: var(--accent);
color: #fff;
cursor: pointer;
}
button:hover {
background: var(--accent-dark);
}
button.ghost,
button.secondary {
border: 1px solid var(--line);
background: #fff;
color: var(--text);
}
button.ghost:hover,
button.secondary:hover {
border-color: #aeb8b2;
background: #eef2ef;
}
button.danger {
background: var(--danger);
}
button.danger:hover {
background: #8f1d15;
}
input,
select {
width: 100%;
min-height: 38px;
border: 1px solid var(--line);
border-radius: 6px;
padding: 7px 10px;
background: #fff;
color: var(--text);
}
label {
display: grid;
gap: 6px;
color: var(--muted);
font-size: 13px;
}
label.inline {
display: flex;
align-items: center;
gap: 8px;
color: var(--text);
}
label.inline input {
width: auto;
min-height: 0;
}
.shell {
width: min(1180px, calc(100vw - 32px));
margin: 0 auto;
padding: 24px 0 40px;
}
.topbar {
display: flex;
align-items: center;
justify-content: space-between;
gap: 16px;
margin-bottom: 18px;
}
.topbar h1 {
margin: 0;
font-size: 24px;
font-weight: 700;
}
.topbar p {
margin: 2px 0 0;
color: var(--muted);
}
.session {
display: flex;
align-items: center;
gap: 12px;
color: var(--muted);
}
.auth-view {
min-height: calc(100vh - 120px);
display: grid;
place-items: center;
}
.panel {
background: var(--panel);
border: 1px solid var(--line);
border-radius: 8px;
box-shadow: var(--shadow);
padding: 18px;
}
.panel.compact {
width: min(420px, 100%);
display: grid;
gap: 14px;
}
.panel h2 {
margin: 0 0 8px;
font-size: 18px;
}
.alert {
border: 1px solid #f0c36a;
border-radius: 6px;
background: #fff7df;
color: #5f410c;
padding: 10px 12px;
margin-bottom: 14px;
}
.hidden {
display: none !important;
}
.status-grid,
.service-grid {
display: grid;
grid-template-columns: repeat(4, minmax(0, 1fr));
gap: 12px;
margin-bottom: 16px;
}
.stat,
.service {
min-width: 0;
background: #fff;
border: 1px solid var(--line);
border-radius: 8px;
padding: 14px;
}
.stat span,
.service span {
display: block;
color: var(--muted);
font-size: 12px;
}
.stat strong,
.service strong {
display: block;
margin-top: 5px;
overflow-wrap: anywhere;
font-size: 17px;
}
.tabs {
display: flex;
gap: 4px;
border-bottom: 1px solid var(--line);
margin: 8px 0 16px;
}
.tabs button {
border-radius: 6px 6px 0 0;
background: transparent;
color: var(--muted);
}
.tabs button.active {
background: #fff;
color: var(--text);
border: 1px solid var(--line);
border-bottom-color: #fff;
}
.toolbar {
display: flex;
gap: 10px;
align-items: center;
justify-content: space-between;
margin-bottom: 12px;
}
.toolbar input {
max-width: 420px;
}
.table-wrap {
overflow-x: auto;
border: 1px solid var(--line);
border-radius: 8px;
background: #fff;
}
table {
width: 100%;
border-collapse: collapse;
min-width: 720px;
}
th,
td {
padding: 10px 12px;
border-bottom: 1px solid var(--line);
text-align: left;
vertical-align: top;
}
th {
color: var(--muted);
font-size: 12px;
font-weight: 600;
background: #fafbf9;
}
td {
overflow-wrap: anywhere;
}
tr:last-child td {
border-bottom: 0;
}
.actions {
width: 180px;
}
.row-actions {
display: flex;
gap: 8px;
flex-wrap: wrap;
}
.badge {
display: inline-flex;
align-items: center;
min-height: 24px;
border-radius: 999px;
padding: 0 9px;
background: #e8f3ee;
color: var(--accent-dark);
font-size: 12px;
}
.badge.off {
background: #f4ece6;
color: var(--warn);
}
.chips {
display: flex;
gap: 8px;
flex-wrap: wrap;
}
.chip {
border: 1px solid var(--line);
border-radius: 999px;
padding: 5px 9px;
background: #fff;
}
.service {
display: grid;
gap: 10px;
}
.service form {
display: grid;
gap: 10px;
}
dialog {
width: min(460px, calc(100vw - 28px));
border: 1px solid var(--line);
border-radius: 8px;
padding: 0;
box-shadow: var(--shadow);
}
dialog::backdrop {
background: rgba(20, 27, 24, 0.36);
}
dialog form {
display: grid;
gap: 12px;
padding: 18px;
}
dialog h2 {
margin: 0;
font-size: 18px;
}
.dialog-actions {
display: flex;
justify-content: flex-end;
gap: 10px;
margin-top: 4px;
}
@media (max-width: 860px) {
.status-grid,
.service-grid {
grid-template-columns: repeat(2, minmax(0, 1fr));
}
.topbar,
.toolbar {
align-items: stretch;
flex-direction: column;
}
.toolbar input {
max-width: none;
}
}
@media (max-width: 560px) {
.shell {
width: min(100vw - 20px, 1180px);
padding-top: 14px;
}
.status-grid,
.service-grid {
grid-template-columns: 1fr;
}
.tabs {
overflow-x: auto;
}
}

View File

@@ -0,0 +1,579 @@
const apiBase = document.body.dataset.apiPrefix || "/admin/api";
const state = {
token: sessionStorage.getItem("mcGatewayAdminToken") || "",
user: null,
routes: [],
services: [],
users: [],
};
const el = (id) => document.getElementById(id);
function showAlert(message) {
const box = el("alert");
box.textContent = message;
box.classList.toggle("hidden", !message);
}
function setView(name) {
for (const id of ["setupView", "loginView", "appView"]) {
el(id).classList.toggle("hidden", id !== name);
}
}
async function api(path, options = {}) {
const headers = { "Accept": "application/json" };
if (options.body !== undefined) {
headers["Content-Type"] = "application/json";
}
if (state.token) {
headers.Authorization = `Bearer ${state.token}`;
}
const res = await fetch(apiBase + path, {
method: options.method || "GET",
headers,
body: options.body === undefined ? undefined : JSON.stringify(options.body),
});
let data = {};
const text = await res.text();
if (text) {
try {
data = JSON.parse(text);
} catch {
data = { error: text };
}
}
if (!res.ok) {
throw new Error(data.error || res.statusText);
}
return data;
}
async function boot() {
bindEvents();
try {
const setup = await api("/setup");
if (setup.required) {
setView("setupView");
el("subtitle").textContent = "Setup";
return;
}
} catch (err) {
showAlert(err.message);
}
if (!state.token) {
setView("loginView");
el("subtitle").textContent = "Login";
return;
}
try {
state.user = await api("/me");
await showApp();
} catch {
sessionStorage.removeItem("mcGatewayAdminToken");
state.token = "";
setView("loginView");
el("subtitle").textContent = "Login";
}
}
function bindEvents() {
el("setupForm").addEventListener("submit", submitSetup);
el("loginForm").addEventListener("submit", submitLogin);
el("logoutBtn").addEventListener("click", logout);
el("routeSearch").addEventListener("input", debounce(loadRoutes, 180));
el("newRouteBtn").addEventListener("click", () => openRouteDialog());
el("newUserBtn").addEventListener("click", () => openUserDialog());
el("routeForm").addEventListener("submit", saveRoute);
el("userForm").addEventListener("submit", saveUser);
for (const button of document.querySelectorAll("[data-close]")) {
button.addEventListener("click", () => button.closest("dialog").close());
}
for (const button of document.querySelectorAll(".tabs button")) {
button.addEventListener("click", () => selectTab(button.dataset.tab));
}
}
async function submitSetup(event) {
event.preventDefault();
const form = new FormData(event.currentTarget);
try {
await api("/setup", {
method: "POST",
body: {
username: form.get("username"),
password: form.get("password"),
},
});
showAlert("");
setView("loginView");
} catch (err) {
showAlert(err.message);
}
}
async function submitLogin(event) {
event.preventDefault();
const form = new FormData(event.currentTarget);
try {
const data = await api("/auth/login", {
method: "POST",
body: {
username: form.get("username"),
password: form.get("password"),
},
});
state.token = data.token;
state.user = data.user;
sessionStorage.setItem("mcGatewayAdminToken", state.token);
showAlert("");
await showApp();
} catch (err) {
showAlert(err.message);
}
}
async function logout() {
try {
await api("/auth/logout", { method: "POST", body: {} });
} catch {
}
sessionStorage.removeItem("mcGatewayAdminToken");
state.token = "";
state.user = null;
setView("loginView");
}
async function showApp() {
setView("appView");
el("subtitle").textContent = "Admin";
el("sessionUser").textContent = `${state.user.username} (${state.user.role})`;
el("logoutBtn").classList.remove("hidden");
applyRoleVisibility();
await loadRoutes();
if (isMember()) {
await loadStatus();
await loadServices();
await loadMetrics();
}
if (isAdmin()) {
await loadUsers();
await loadAudit();
}
}
function applyRoleVisibility() {
const member = isMember();
const admin = isAdmin();
el("statusGrid").classList.toggle("hidden", !member);
el("newRouteBtn").classList.toggle("hidden", !member);
toggleTab("services", member);
toggleTab("metrics", member);
toggleTab("users", admin);
toggleTab("audit", admin);
selectTab("routes");
}
function toggleTab(name, visible) {
document.querySelector(`[data-tab="${name}"]`).classList.toggle("hidden", !visible);
}
function selectTab(name) {
for (const button of document.querySelectorAll(".tabs button")) {
button.classList.toggle("active", button.dataset.tab === name);
}
for (const panel of document.querySelectorAll(".tab-panel")) {
panel.classList.add("hidden");
}
el(`${name}Tab`).classList.remove("hidden");
}
async function loadStatus() {
try {
const status = await api("/status");
el("statusGrid").innerHTML = [
stat("PID", status.pid),
stat("Uptime", `${status.uptime_seconds}s`),
stat("SQLite", status.db_path),
stat("TCP/Admin", status.tcp_admin_port),
].join("");
} catch (err) {
showAlert(err.message);
}
}
async function loadRoutes() {
try {
const q = encodeURIComponent(el("routeSearch").value || "");
const data = await api(q ? `/routes?q=${q}` : "/routes");
state.routes = data.routes || [];
renderRoutes();
} catch (err) {
showAlert(err.message);
}
}
function renderRoutes() {
const canWrite = isMember();
el("routesBody").innerHTML = state.routes.map((route) => `
<tr>
<td>${escapeHTML(route.host)}</td>
<td>${escapeHTML(route.upstream)}</td>
<td>${badge(route.enabled ? "Enabled" : "Disabled", !route.enabled)}</td>
<td>${escapeHTML(route.note || "")}</td>
<td class="actions">${canWrite ? routeActions(route) : ""}</td>
</tr>
`).join("");
for (const button of document.querySelectorAll("[data-edit-route]")) {
button.addEventListener("click", () => {
const route = state.routes.find((item) => item.host === button.dataset.editRoute);
openRouteDialog(route);
});
}
for (const button of document.querySelectorAll("[data-delete-route]")) {
button.addEventListener("click", () => removeRoute(button.dataset.deleteRoute));
}
}
function routeActions(route) {
return `
<div class="row-actions">
<button class="secondary" type="button" data-edit-route="${escapeAttr(route.host)}">Edit</button>
<button class="danger" type="button" data-delete-route="${escapeAttr(route.host)}">Delete</button>
</div>
`;
}
function openRouteDialog(route = null) {
const form = el("routeForm");
form.reset();
form.dataset.originalHost = route ? route.host : "";
form.elements.host.disabled = Boolean(route);
if (route) {
form.elements.host.value = route.host;
form.elements.upstream.value = route.upstream;
form.elements.enabled.checked = route.enabled;
form.elements.note.value = route.note || "";
} else {
form.elements.enabled.checked = true;
}
el("routeDialog").showModal();
}
async function saveRoute(event) {
event.preventDefault();
const form = event.currentTarget;
const host = form.dataset.originalHost || form.elements.host.value;
try {
await api(`/routes/${encodeURIComponent(host)}`, {
method: "PUT",
body: {
upstream: form.elements.upstream.value,
enabled: form.elements.enabled.checked,
note: form.elements.note.value,
},
});
el("routeDialog").close();
await loadRoutes();
showAlert("");
} catch (err) {
showAlert(err.message);
}
}
async function removeRoute(host) {
if (host === "default" && !confirm("Delete default route?")) {
return;
}
try {
await api(`/routes/${encodeURIComponent(host)}`, { method: "DELETE" });
await loadRoutes();
} catch (err) {
showAlert(err.message);
}
}
async function loadServices() {
try {
const data = await api("/services");
state.services = data.services || [];
renderServices();
} catch (err) {
showAlert(err.message);
}
}
function renderServices() {
el("servicesGrid").innerHTML = state.services.map((service) => `
<article class="service">
<div>
<span>${escapeHTML(service.name)}</span>
<strong>${service.enabled ? "Enabled" : "Disabled"}${service.restart_required ? " / restart required" : ""}</strong>
</div>
${isAdmin() ? serviceForm(service) : serviceSummary(service)}
</article>
`).join("");
for (const form of document.querySelectorAll("[data-service-form]")) {
form.addEventListener("submit", saveService);
}
for (const button of document.querySelectorAll("[data-restart-service]")) {
button.addEventListener("click", () => restartService(button.dataset.restartService));
}
}
function serviceSummary(service) {
return `<div><span>Port</span><strong>${service.port}</strong></div>`;
}
function serviceForm(service) {
const disabled = service.name === "tcp_admin" ? "disabled" : "";
const optionFields = serviceOptionFields(service);
return `
<form data-service-form="${escapeAttr(service.name)}">
<label class="inline">
<input name="enabled" type="checkbox" ${service.enabled ? "checked" : ""} ${disabled}>
Enabled
</label>
<label>
Port
<input name="port" type="number" min="1" max="65535" value="${service.port}">
</label>
${optionFields}
<div class="row-actions">
<button type="submit">Save</button>
<button class="secondary" type="button" data-restart-service="${escapeAttr(service.name)}">Restart</button>
</div>
</form>
`;
}
function serviceOptionFields(service) {
const options = service.options || {};
if (service.name === "kcp") {
return `
<label>Data shards<input name="data_shards" type="number" min="1" value="${options.data_shards || 10}"></label>
<label>Parity shards<input name="parity_shards" type="number" min="1" value="${options.parity_shards || 3}"></label>
`;
}
if (service.name === "quic") {
const protocols = Array.isArray(options.application_protocols) ? options.application_protocols.join(",") : "";
return `<label>Protocols<input name="application_protocols" value="${escapeAttr(protocols)}"></label>`;
}
if (service.name === "websocket") {
return `<label>Path<input name="path" value="${escapeAttr(options.path || "/")}"></label>`;
}
return "";
}
async function saveService(event) {
event.preventDefault();
const form = event.currentTarget;
const name = form.dataset.serviceForm;
const options = {};
if (name === "kcp") {
options.data_shards = Number(form.elements.data_shards.value);
options.parity_shards = Number(form.elements.parity_shards.value);
} else if (name === "quic") {
options.application_protocols = form.elements.application_protocols.value.split(",").map((item) => item.trim()).filter(Boolean);
} else if (name === "websocket") {
options.path = form.elements.path.value;
}
try {
await api(`/services/${encodeURIComponent(name)}`, {
method: "PUT",
body: {
enabled: name === "tcp_admin" ? true : form.elements.enabled.checked,
port: Number(form.elements.port.value),
options,
},
});
await loadServices();
} catch (err) {
showAlert(err.message);
}
}
async function restartService(name) {
try {
await api(`/services/${encodeURIComponent(name)}/restart`, { method: "POST", body: {} });
await loadServices();
} catch (err) {
showAlert(err.message);
}
}
async function loadUsers() {
try {
const data = await api("/users");
state.users = data.users || [];
renderUsers();
} catch (err) {
showAlert(err.message);
}
}
function renderUsers() {
el("usersBody").innerHTML = state.users.map((user) => `
<tr>
<td>${escapeHTML(user.username)}</td>
<td>${escapeHTML(user.role)}</td>
<td>${badge(user.disabled ? "Disabled" : "Active", user.disabled)}</td>
<td class="actions">
<div class="row-actions">
<button class="secondary" type="button" data-edit-user="${escapeAttr(user.username)}">Edit</button>
<button class="danger" type="button" data-delete-user="${escapeAttr(user.username)}">Delete</button>
</div>
</td>
</tr>
`).join("");
for (const button of document.querySelectorAll("[data-edit-user]")) {
button.addEventListener("click", () => {
const user = state.users.find((item) => item.username === button.dataset.editUser);
openUserDialog(user);
});
}
for (const button of document.querySelectorAll("[data-delete-user]")) {
button.addEventListener("click", () => removeUser(button.dataset.deleteUser));
}
}
function openUserDialog(user = null) {
const form = el("userForm");
form.reset();
form.dataset.originalUsername = user ? user.username : "";
form.elements.username.disabled = Boolean(user);
form.elements.password.required = !user;
if (user) {
form.elements.username.value = user.username;
form.elements.role.value = user.role;
form.elements.disabled.checked = user.disabled;
} else {
form.elements.role.value = "member";
}
el("userDialog").showModal();
}
async function saveUser(event) {
event.preventDefault();
const form = event.currentTarget;
const username = form.dataset.originalUsername || form.elements.username.value;
const body = {
role: form.elements.role.value,
disabled: form.elements.disabled.checked,
};
if (form.elements.password.value) {
body.password = form.elements.password.value;
}
try {
if (form.dataset.originalUsername) {
await api(`/users/${encodeURIComponent(username)}`, { method: "PATCH", body });
} else {
body.username = username;
await api("/users", { method: "POST", body });
}
el("userDialog").close();
await loadUsers();
} catch (err) {
showAlert(err.message);
}
}
async function removeUser(username) {
if (!confirm(`Delete user ${username}?`)) {
return;
}
try {
await api(`/users/${encodeURIComponent(username)}`, { method: "DELETE" });
await loadUsers();
} catch (err) {
showAlert(err.message);
}
}
async function loadMetrics() {
try {
const data = await api("/metrics");
el("metricsGrid").innerHTML = [
stat("Total", data.total_connections),
stat("Active", data.active_connections),
stat("TCP", data.tcp_connections),
stat("WebSocket", data.websocket_connections),
stat("Misses", data.route_misses),
stat("Dial errors", data.upstream_dial_errors),
].join("");
const hits = data.route_hits || {};
el("routeHits").innerHTML = Object.keys(hits).length
? Object.entries(hits).map(([host, count]) => `<span class="chip">${escapeHTML(host)}: ${count}</span>`).join("")
: `<span class="chip">No hits</span>`;
} catch (err) {
showAlert(err.message);
}
}
async function loadAudit() {
try {
const data = await api("/audit-logs");
el("auditBody").innerHTML = (data.audit_logs || []).map((item) => `
<tr>
<td>${new Date(item.created_at * 1000).toLocaleString()}</td>
<td>${escapeHTML(item.actor)}</td>
<td>${escapeHTML(item.action)}</td>
<td>${escapeHTML(item.target_type)}:${escapeHTML(item.target_id)}</td>
<td>${badge(item.success ? "Success" : "Failed", !item.success)}</td>
<td>${escapeHTML(item.message || "")}</td>
</tr>
`).join("");
} catch (err) {
showAlert(err.message);
}
}
function stat(label, value) {
return `<div class="stat"><span>${escapeHTML(label)}</span><strong>${escapeHTML(String(value ?? ""))}</strong></div>`;
}
function badge(text, off = false) {
return `<span class="badge ${off ? "off" : ""}">${escapeHTML(text)}</span>`;
}
function isAdmin() {
return state.user && state.user.role === "admin";
}
function isMember() {
return state.user && (state.user.role === "admin" || state.user.role === "member");
}
function debounce(fn, wait) {
let id = 0;
return (...args) => {
clearTimeout(id);
id = setTimeout(() => fn(...args), wait);
};
}
function escapeHTML(value) {
return String(value).replace(/[&<>"']/g, (ch) => ({
"&": "&amp;",
"<": "&lt;",
">": "&gt;",
"\"": "&quot;",
"'": "&#39;",
}[ch]));
}
function escapeAttr(value) {
return escapeHTML(value).replace(/`/g, "&#96;");
}
boot();

View File

@@ -0,0 +1,195 @@
<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>mc-gateway admin</title>
<link rel="stylesheet" href="app.css">
</head>
<body data-api-prefix="__ADMIN_API_PREFIX__">
<main class="shell">
<header class="topbar">
<div>
<h1>mc-gateway</h1>
<p id="subtitle">Admin</p>
</div>
<div class="session">
<span id="sessionUser"></span>
<button id="logoutBtn" class="ghost hidden" type="button">Logout</button>
</div>
</header>
<section id="alert" class="alert hidden"></section>
<section id="setupView" class="auth-view hidden">
<form id="setupForm" class="panel compact">
<h2>Initial admin</h2>
<label>
Username
<input name="username" autocomplete="username" value="admin">
</label>
<label>
Password
<input name="password" autocomplete="new-password" type="password" required>
</label>
<button type="submit">Create admin</button>
</form>
</section>
<section id="loginView" class="auth-view hidden">
<form id="loginForm" class="panel compact">
<h2>Login</h2>
<label>
Username
<input name="username" autocomplete="username" required>
</label>
<label>
Password
<input name="password" autocomplete="current-password" type="password" required>
</label>
<button type="submit">Login</button>
</form>
</section>
<section id="appView" class="hidden">
<section id="statusGrid" class="status-grid"></section>
<nav class="tabs">
<button data-tab="routes" class="active" type="button">Routes</button>
<button data-tab="services" type="button">Services</button>
<button data-tab="users" type="button">Users</button>
<button data-tab="metrics" type="button">Metrics</button>
<button data-tab="audit" type="button">Audit</button>
</nav>
<section id="routesTab" class="tab-panel">
<div class="toolbar">
<input id="routeSearch" placeholder="Search host, upstream, note">
<button id="newRouteBtn" type="button">New route</button>
</div>
<div class="table-wrap">
<table>
<thead>
<tr>
<th>Host</th>
<th>Upstream</th>
<th>Enabled</th>
<th>Note</th>
<th class="actions">Actions</th>
</tr>
</thead>
<tbody id="routesBody"></tbody>
</table>
</div>
</section>
<section id="servicesTab" class="tab-panel hidden">
<div id="servicesGrid" class="service-grid"></div>
</section>
<section id="usersTab" class="tab-panel hidden">
<div class="toolbar">
<button id="newUserBtn" type="button">New user</button>
</div>
<div class="table-wrap">
<table>
<thead>
<tr>
<th>Username</th>
<th>Role</th>
<th>Disabled</th>
<th class="actions">Actions</th>
</tr>
</thead>
<tbody id="usersBody"></tbody>
</table>
</div>
</section>
<section id="metricsTab" class="tab-panel hidden">
<div id="metricsGrid" class="status-grid"></div>
<div class="panel">
<h2>Route hits</h2>
<div id="routeHits" class="chips"></div>
</div>
</section>
<section id="auditTab" class="tab-panel hidden">
<div class="table-wrap">
<table>
<thead>
<tr>
<th>Time</th>
<th>Actor</th>
<th>Action</th>
<th>Target</th>
<th>Result</th>
<th>Message</th>
</tr>
</thead>
<tbody id="auditBody"></tbody>
</table>
</div>
</section>
</section>
</main>
<dialog id="routeDialog">
<form id="routeForm" method="dialog">
<h2>Route</h2>
<label>
Host
<input name="host" required>
</label>
<label>
Upstream
<input name="upstream" required placeholder="127.0.0.1:25565">
</label>
<label class="inline">
<input name="enabled" type="checkbox" checked>
Enabled
</label>
<label>
Note
<input name="note">
</label>
<div class="dialog-actions">
<button type="button" data-close>Cancel</button>
<button type="submit">Save</button>
</div>
</form>
</dialog>
<dialog id="userDialog">
<form id="userForm" method="dialog">
<h2>User</h2>
<label>
Username
<input name="username" required>
</label>
<label>
Role
<select name="role">
<option value="admin">Admin</option>
<option value="member">Member</option>
<option value="guest">Guest</option>
</select>
</label>
<label>
Password
<input name="password" type="password">
</label>
<label class="inline">
<input name="disabled" type="checkbox">
Disabled
</label>
<div class="dialog-actions">
<button type="button" data-close>Cancel</button>
<button type="submit">Save</button>
</div>
</form>
</dialog>
<script src="app.js"></script>
</body>
</html>

View File

@@ -0,0 +1,90 @@
package main
import (
"net/http"
"github.com/tursom/mc-gateway/internal/adminhttp"
)
func handleAdminUsersList(w http.ResponseWriter, r *http.Request) {
if _, ok := requireRole(w, r, adminRoleAdmin); !ok {
return
}
users, err := listUsers(r.Context())
if err != nil {
adminhttp.WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"users": users})
}
func handleAdminUsersCreate(w http.ResponseWriter, r *http.Request) {
session, ok := requireRole(w, r, adminRoleAdmin)
if !ok {
return
}
var req adminhttp.CreateUserRequest
if !adminhttp.DecodeJSONRequest(w, r, &req) {
return
}
err := createUser(r.Context(), session.Username, req.Username, req.Role, req.Password, req.Disabled)
if err != nil {
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "user_create", "user", req.Username, false, err.Error())
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "user_create", "user", req.Username, true, "user created")
adminhttp.WriteJSON(w, http.StatusCreated, map[string]any{"ok": true})
}
func handleAdminUserItem(w http.ResponseWriter, r *http.Request, rawUsername string) {
session, ok := requireRole(w, r, adminRoleAdmin)
if !ok {
return
}
username, err := adminhttp.PathSegment(rawUsername)
if err != nil {
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
switch r.Method {
case http.MethodPatch:
var req adminhttp.PatchUserRequest
if !adminhttp.DecodeJSONRequest(w, r, &req) {
return
}
err := patchUser(r.Context(), session.Username, username, req.Role, req.Disabled, req.Password)
if err != nil {
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "user_patch", "user", username, false, err.Error())
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "user_patch", "user", username, true, "user updated")
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"ok": true})
case http.MethodDelete:
err := deleteUser(r.Context(), username)
if err != nil {
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "user_delete", "user", username, false, err.Error())
adminhttp.WriteAPIError(w, http.StatusBadRequest, err.Error())
return
}
recordAudit(r.Context(), session.Username, adminhttp.RequestSourceIP(r), "user_delete", "user", username, true, "user deleted")
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"ok": true})
default:
adminhttp.WriteAPIError(w, http.StatusMethodNotAllowed, "method not allowed")
}
}
func handleAdminAuditLogs(w http.ResponseWriter, r *http.Request) {
if _, ok := requireRole(w, r, adminRoleAdmin); !ok {
return
}
logs, err := listAuditLogs(r.Context())
if err != nil {
adminhttp.WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
adminhttp.WriteJSON(w, http.StatusOK, map[string]any{"audit_logs": logs})
}

139
cmd/gateway/admin_users.go Normal file
View File

@@ -0,0 +1,139 @@
package main
import (
"context"
"database/sql"
"errors"
"strings"
"github.com/tursom/mc-gateway/internal/adminuser"
"golang.org/x/crypto/bcrypt"
)
const (
adminRoleAdmin = adminuser.RoleAdmin
adminRoleMember = adminuser.RoleMember
adminRoleGuest = adminuser.RoleGuest
)
func ensureInitialAdminFromEnv(ctx context.Context, db *sql.DB, password string) error {
if strings.TrimSpace(password) == "" {
return nil
}
empty, err := usersTableEmpty(ctx, db)
if err != nil {
return err
}
if !empty {
return nil
}
return createUser(ctx, "system", "admin", adminRoleAdmin, password, false)
}
func usersTableEmpty(ctx context.Context, db *sql.DB) (bool, error) {
return adminuser.NewRepository(db).TableEmpty(ctx)
}
func createInitialAdmin(ctx context.Context, username, password string) error {
empty, err := usersTableEmpty(ctx, adminDB)
if err != nil {
return err
}
if !empty {
return errors.New("initial admin has already been created")
}
return createUser(ctx, "setup", username, adminRoleAdmin, password, false)
}
func createUser(ctx context.Context, actor, username, role, password string, disabled bool) error {
username = strings.TrimSpace(username)
if err := adminuser.ValidateUsername(username); err != nil {
return err
}
if err := adminuser.ValidateRole(role); err != nil {
return err
}
if err := adminuser.ValidatePassword(password); err != nil {
return err
}
hash, err := hashPassword(password)
if err != nil {
return err
}
_ = actor
return adminuser.NewRepository(adminDB).Create(ctx, username, role, hash, disabled)
}
func authenticateUser(ctx context.Context, username, password string) (adminuser.User, error) {
user, hash, err := getUserWithHash(ctx, username)
if err != nil {
return adminuser.User{}, err
}
if user.Disabled {
return adminuser.User{}, errors.New("user is disabled")
}
if err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)); err != nil {
return adminuser.User{}, errors.New("invalid username or password")
}
return user, nil
}
func getUserWithHash(ctx context.Context, username string) (adminuser.User, string, error) {
return adminuser.NewRepository(adminDB).GetWithHash(ctx, username)
}
func listUsers(ctx context.Context) ([]adminuser.User, error) {
return adminuser.NewRepository(adminDB).List(ctx)
}
func patchUser(ctx context.Context, actor, username string, role *string, disabled *bool, password *string) error {
username = strings.TrimSpace(username)
if err := adminuser.ValidateUsername(username); err != nil {
return err
}
var passwordHashProvider func() (string, error)
if password != nil {
passwordValue := *password
passwordHashProvider = func() (string, error) {
if err := adminuser.ValidatePassword(passwordValue); err != nil {
return "", err
}
return hashPassword(passwordValue)
}
}
result, err := adminuser.NewRepository(adminDB).Patch(ctx, username, adminuser.Patch{
Role: role,
Disabled: disabled,
PasswordHashProvider: passwordHashProvider,
})
if err != nil {
return err
}
if result.InvalidateSessions {
removeSessionsForUser(username)
}
_ = actor
return nil
}
func deleteUser(ctx context.Context, username string) error {
username = strings.TrimSpace(username)
if err := adminuser.ValidateUsername(username); err != nil {
return err
}
if err := adminuser.NewRepository(adminDB).Delete(ctx, username); err != nil {
return err
}
removeSessionsForUser(username)
return nil
}
func hashPassword(password string) (string, error) {
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
return string(hash), err
}

View File

@@ -1,120 +1,24 @@
package main
import (
"io"
"os"
"strings"
"sync"
"time"
"github.com/BurntSushi/toml"
"github.com/fsnotify/fsnotify"
"github.com/mitchellh/mapstructure"
"github.com/rs/zerolog/log"
"github.com/tursom/mc-gateway/internal/gatewayconfig"
)
var (
configFile = "config.toml"
config Config
config gatewayconfig.Config
configLoadLock sync.Mutex
)
type (
Config struct {
Tcp ProtocolConfig `toml:"tcp"`
Quic QuicConfig `toml:"quic"`
Kcp KcpConfig `toml:"kcp"`
WebSocket WebSocketConfig `toml:"websocket"`
Hosts map[string]string `toml:"hosts"`
Log LogConfig `toml:"log"`
PidFile string `toml:"pid_file"`
Plugin map[string]map[string]any `toml:"plugin"`
}
ProtocolConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
}
KcpConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
DataShards int `toml:"data_shards"`
ParityShards int `toml:"parity_Shards"`
}
QuicConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
ApplicationProtocols []string `toml:"application_protocols"`
}
WebSocketConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
Path string `toml:"path"`
}
// LogConfig 定义了日志的配置,包括日志级别和日志文件路径
// 日志级别可以是 "trade", "debug", "info", "warn", "error", "fatal", "disabled"
// 默认日志级别为 "info"
// 日志文件路径指定了日志输出的位置
// 如果日志文件路径为空,则日志将只输出到标准输出
LogConfig struct {
Level string `toml:"level"`
File string `toml:"file"`
}
// serviceConfig 定义了一个服务的配置,包括是否启用和运行函数
// 运行函数接收一个 WaitGroup用于在服务运行时进行同步
// 这样可以确保所有服务在主函数退出前都能正确关闭
// 运行函数通常会在 goroutine 中执行,以便并发处理多个服务
serviceConfig struct {
enable *bool
run func(wg *sync.WaitGroup)
}
)
var services = []serviceConfig{
{
enable: &config.Tcp.Enable,
run: runTcp,
},
{
enable: &config.Kcp.Enable,
run: runKcp,
},
{
enable: &config.Quic.Enable,
run: runQuic,
},
{
enable: &config.WebSocket.Enable,
run: runWebSocket,
},
}
func loadConfig() error {
configLoadLock.Lock()
defer configLoadLock.Unlock()
file, err := os.Open(configFile)
if err != nil {
if err := initializeGatewayRuntime(); err != nil {
return err
}
defer file.Close()
byteValue, err := io.ReadAll(file)
if err != nil {
return err
}
if err := toml.Unmarshal(byteValue, &config); err != nil {
return err
}
writePIDFile()
if err := loadLogger(); err != nil {
return err
@@ -125,69 +29,10 @@ func loadConfig() error {
return nil
}
func watchConfig() *fsnotify.Watcher {
go func() {
for {
time.Sleep(time.Minute)
if err := loadConfig(); err != nil {
log.Error().Err(err).Msg("Failed to reload config")
}
}
}()
watcher, err := fsnotify.NewWatcher()
if err != nil {
log.Fatal().Err(err).Msg("Failed to create config watcher")
}
go func() {
for {
select {
case event, ok := <-watcher.Events:
if !ok {
log.Error().Msg("watcher.Events channel closed")
return
}
if !strings.HasSuffix(event.Name, configFile) {
continue
}
case err, ok := <-watcher.Errors:
if !ok {
log.Error().Msg("watcher.Errors channel closed")
return
}
log.Error().Err(err).Msg("watcher error")
continue
}
log.Info().Msg("reload config")
if err := loadConfig(); err != nil {
log.Error().Err(err).Msg("Failed to reload config")
}
}
}()
if err = watcher.Add("."); err != nil {
log.Fatal().Err(err).Msg("Failed to watch config file")
}
return watcher
}
func loadPluginConfig(cfg map[string]any, pluginCfg any) error {
log.Info().
Any("config", cfg).
Msg("Loading plugin config")
decoder, err := mapstructure.NewDecoder(&mapstructure.DecoderConfig{
Result: pluginCfg,
TagName: "toml",
})
if err != nil {
return err
}
if err := decoder.Decode(cfg); err != nil {
return err
}
return nil
return gatewayconfig.DecodePluginConfig(cfg, pluginCfg)
}

View File

@@ -1,10 +1,7 @@
package main
import (
"fmt"
"os"
"path/filepath"
"strings"
"testing"
)
@@ -46,111 +43,153 @@ func TestLoadPluginConfigReturnsDecodeError(t *testing.T) {
}
}
func TestLoadConfigReadsTomlAndAppliesSideEffects(t *testing.T) {
func TestParseStartupConfigDefaultsAndEnv(t *testing.T) {
cfg, err := parseStartupConfig(func(string) string { return "" })
if err != nil {
t.Fatalf("parseStartupConfig() error = %v", err)
}
if cfg.DBPath != defaultAdminDBPath {
t.Fatalf("DBPath = %q, want %q", cfg.DBPath, defaultAdminDBPath)
}
if cfg.TCPAdminPort != defaultTCPPort {
t.Fatalf("TCPAdminPort = %d, want %d", cfg.TCPAdminPort, defaultTCPPort)
}
if cfg.AdminPath != defaultAdminPath {
t.Fatalf("AdminPath = %q, want %q", cfg.AdminPath, defaultAdminPath)
}
if cfg.AdminAPIPrefix != defaultAdminAPIPrefix {
t.Fatalf("AdminAPIPrefix = %q, want %q", cfg.AdminAPIPrefix, defaultAdminAPIPrefix)
}
env := map[string]string{
adminEnvDB: "/tmp/mc.db",
adminEnvTCPAdminPort: "25575",
adminEnvPath: "/ops",
adminEnvAPIPrefix: "/ops/api/",
}
cfg, err = parseStartupConfig(func(key string) string { return env[key] })
if err != nil {
t.Fatalf("parseStartupConfig(env) error = %v", err)
}
if cfg.DBPath != "/tmp/mc.db" || cfg.TCPAdminPort != 25575 || cfg.AdminPath != "/ops/" || cfg.AdminAPIPrefix != "/ops/api" {
t.Fatalf("startup config = %+v", cfg)
}
}
func TestParseStartupConfigReturnsErrors(t *testing.T) {
tests := []struct {
name string
env map[string]string
}{
{
name: "invalid port",
env: map[string]string{adminEnvTCPAdminPort: "70000"},
},
{
name: "invalid admin path",
env: map[string]string{adminEnvPath: "admin"},
},
{
name: "invalid api prefix",
env: map[string]string{adminEnvAPIPrefix: "api"},
},
{
name: "api prefix equals admin path",
env: map[string]string{
adminEnvPath: "/admin",
adminEnvAPIPrefix: "/admin",
},
},
{
name: "api prefix under asset path",
env: map[string]string{adminEnvAPIPrefix: "/admin/app.js/api"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := parseStartupConfig(func(key string) string { return tt.env[key] })
if err == nil {
t.Fatal("parseStartupConfig() error = nil, want error")
}
})
}
}
func TestLoadConfigInitializesSQLiteDefaults(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)
}
dbPath := filepath.Join(tmpDir, "gateway.sqlite3")
t.Setenv(adminEnvDB, dbPath)
t.Setenv(adminEnvTCPAdminPort, "25575")
t.Setenv(adminEnvPath, "/ops")
t.Setenv(adminEnvAPIPrefix, "/ops/api")
if err := loadConfig(); err != nil {
t.Fatalf("loadConfig() error = %v", err)
}
if !config.Tcp.Enable || config.Tcp.Port != 25565 {
if !config.Tcp.Enable || config.Tcp.Port != 25575 {
t.Fatalf("tcp config = %+v", config.Tcp)
}
if !config.Quic.Enable || config.Quic.Port != 25566 {
t.Fatalf("quic config = %+v", config.Quic)
if config.Quic.Enable || config.Kcp.Enable || config.WebSocket.Enable {
t.Fatalf("optional services should be disabled by default: quic=%+v kcp=%+v websocket=%+v", config.Quic, config.Kcp, config.WebSocket)
}
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 {
if config.Kcp.Port != defaultKCPPort || config.Kcp.DataShards != defaultKCPDataShards || config.Kcp.ParityShards != defaultKCPParityShards {
t.Fatalf("kcp config = %+v", config.Kcp)
}
if !config.WebSocket.Enable || config.WebSocket.Path != "/gateway" {
t.Fatalf("websocket config = %+v", config.WebSocket)
if adminStartup.DBPath != dbPath || adminStartup.AdminPath != "/ops/" || adminStartup.AdminAPIPrefix != "/ops/api" {
t.Fatalf("adminStartup = %+v", adminStartup)
}
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)
if adminDB == nil {
t.Fatal("adminDB = nil")
}
pidBytes, err := os.ReadFile(pidPath)
var enabled, port int
if err := adminDB.QueryRow(`SELECT enabled, port FROM services WHERE name = ?`, serviceNameTCPAdmin).Scan(&enabled, &port); err != nil {
t.Fatalf("query tcp_admin service: %v", err)
}
if enabled != 1 || port != 25575 {
t.Fatalf("tcp_admin service enabled=%d port=%d, want enabled=1 port=25575", enabled, port)
}
}
func TestLoadConfigCreatesInitialAdminFromEnv(t *testing.T) {
defer saveGatewayState(t)()
t.Setenv(adminEnvDB, filepath.Join(t.TempDir(), "gateway.sqlite3"))
t.Setenv(adminEnvInitialPassword, "secret")
if err := loadConfig(); err != nil {
t.Fatalf("loadConfig() error = %v", err)
}
user, err := authenticateUser(t.Context(), "admin", "secret")
if err != nil {
t.Fatalf("ReadFile(pid) error = %v", err)
t.Fatalf("authenticateUser() 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)
if user.Role != adminRoleAdmin {
t.Fatalf("admin role = %q, want %q", user.Role, adminRoleAdmin)
}
}
func TestLoadConfigReturnsErrors(t *testing.T) {
t.Run("missing file", func(t *testing.T) {
func TestLoadConfigReturnsEnvErrors(t *testing.T) {
t.Run("invalid port", func(t *testing.T) {
defer saveGatewayState(t)()
configFile = filepath.Join(t.TempDir(), "missing.toml")
t.Setenv(adminEnvDB, filepath.Join(t.TempDir(), "gateway.sqlite3"))
t.Setenv(adminEnvTCPAdminPort, "0")
if err := loadConfig(); err == nil {
t.Fatal("loadConfig() error = nil, want error")
}
})
t.Run("invalid toml", func(t *testing.T) {
t.Run("invalid admin path", 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)
}
t.Setenv(adminEnvDB, filepath.Join(t.TempDir(), "gateway.sqlite3"))
t.Setenv(adminEnvPath, "admin")
if err := loadConfig(); err == nil {
t.Fatal("loadConfig() error = nil, want error")
}

View File

@@ -14,9 +14,9 @@ func TestHandleRequestProxiesAndClosesConnections(t *testing.T) {
packet := gatewayTestPacket("play.example")
source := newGatewayTestConn(packet)
upstream := newGatewayTestConn([]byte("reply"))
config.Hosts = map[string]string{
setGatewayTestRoutes(map[string]string{
"play.example": "backend.example:25565",
}
})
registerGatewayUpstreamHook(
t,

View File

@@ -10,12 +10,14 @@ import (
func haProxyUpstream(source net.Conn, host string) net.Conn {
target, err := net.ResolveTCPAddr("tcp", host)
if err != nil {
gatewayMetrics.UpstreamDialError()
log.Err(err).Msg("failed to resolve TCP address")
return nil
}
conn, err := tcpDialer.Dial("tcp", target.String())
if err != nil {
gatewayMetrics.UpstreamDialError()
log.Err(err).Msg("failed to dial TCP")
return nil
}

View File

@@ -44,6 +44,7 @@ func runKcp(wg *sync.WaitGroup) {
func upstreamKcp(host string) net.Conn {
conn, err := kcp.DialWithOptions(host, nil, config.Kcp.DataShards, config.Kcp.ParityShards)
if err != nil {
gatewayMetrics.UpstreamDialError()
log.Error().Err(err).
Msg("Failed to dial KCP server")
return nil

View File

@@ -2,10 +2,10 @@ package main
import (
"net"
"strings"
"sync"
"github.com/rs/zerolog/log"
"github.com/tursom/mc-gateway/internal/upstreamtarget"
"github.com/tursom/mc-gateway/plugin/api"
"github.com/tursom/mc-gateway/protocol"
)
@@ -19,9 +19,7 @@ func main() {
log.Err(err).Msg("Failed to write PID file")
}
defer removePIDFile()
watcher := watchConfig()
defer watcher.Close()
defer closeGatewayRuntime()
go handleLogRotate()
@@ -31,23 +29,16 @@ func main() {
}
func startEnabledServices() {
if tcpWebPortReuseEnabled() {
startService(runTcpWebPortReuse)
if config.Kcp.Enable {
startService(runKcp)
}
if config.Quic.Enable {
startService(runQuic)
}
return
startService(runTcpWebPortReuse)
if config.Kcp.Enable {
startService(runKcp)
}
for _, service := range services {
if !*service.enable {
continue
}
startService(service.run)
if config.Quic.Enable {
startService(runQuic)
}
if config.WebSocket.Enable && normalizedWebSocketPort() != normalizedTCPPort() {
startService(runWebSocket)
}
}
@@ -57,6 +48,9 @@ func startService(run func(wg *sync.WaitGroup)) {
}
func handleRequest(conn net.Conn) {
gatewayMetrics.ConnectionStarted()
defer gatewayMetrics.ConnectionFinished()
defer func() {
rec := recover()
if rec == nil {
@@ -104,29 +98,30 @@ func mapToHost(conn net.Conn) net.Conn {
return nil
}
mc_host := protocol.GetMcHost(buf[:n])
if mc_host == "" {
mcHost := protocol.GetMcHost(buf[:n])
if mcHost == "" {
log.Err(errEmptyBuffer).
Str("client", conn.RemoteAddr().String()).
Msg("failed to parse mc host from buffer")
return nil
}
host, ok := config.Hosts[mc_host]
if !ok {
host = config.Hosts["default"]
}
host, ok := lookupRoute(mcHost)
if host == "" {
gatewayMetrics.RouteMiss()
log.Err(errEmptyBuffer).
Str("client", conn.RemoteAddr().String()).
Str("host", mc_host).
Str("host", mcHost).
Msg("failed to route host")
return nil
}
if ok {
gatewayMetrics.RouteHit(mcHost)
}
log.Debug().
Str("client", conn.RemoteAddr().String()).
Str("host", mc_host).
Str("host", mcHost).
Str("mc", host).
Msg("map to host")
@@ -143,14 +138,16 @@ func mapToHost(conn net.Conn) net.Conn {
}
if !ok {
if host, ok := strings.CutPrefix(host, "quic://"); ok {
client = upstreamQuic(host)
} else if host, ok := strings.CutPrefix(host, "kcp://"); ok {
client = upstreamKcp(host)
} else if host, ok := strings.CutPrefix(host, "haproxy://"); ok {
client = haProxyUpstream(conn, host)
} else {
client = upstreamTcp(host)
target := upstreamtarget.Parse(host)
switch target.Protocol {
case upstreamtarget.ProtocolQUIC:
client = upstreamQuic(target.Address)
case upstreamtarget.ProtocolKCP:
client = upstreamKcp(target.Address)
case upstreamtarget.ProtocolHAProxy:
client = haProxyUpstream(conn, target.Address)
default:
client = upstreamTcp(target.Address)
}
}
if client == nil {
@@ -160,7 +157,7 @@ func mapToHost(conn net.Conn) net.Conn {
if err := writeAll(client, buf[:n]); err != nil {
log.Err(err).
Str("client", conn.RemoteAddr().String()).
Str("host", mc_host).
Str("host", mcHost).
Str("mc", host).
Msg("failed to write initial packet to upstream")
client.Close()

View File

@@ -13,9 +13,9 @@ func TestMapToHostRoutesThroughHookAndForwardsInitialPacket(t *testing.T) {
packet := gatewayTestPacket("play.example", 0x63, 0x00)
source := newGatewayTestConn(packet)
upstream := newGatewayTestConn(nil)
config.Hosts = map[string]string{
setGatewayTestRoutes(map[string]string{
"play.example": "backend.example:25565",
}
})
var gotSource net.Conn
var gotHost string
@@ -52,9 +52,9 @@ func TestMapToHostUsesDefaultRoute(t *testing.T) {
packet := gatewayTestPacket("unknown.example")
source := newGatewayTestConn(packet)
upstream := newGatewayTestConn(nil)
config.Hosts = map[string]string{
setGatewayTestRoutes(map[string]string{
"default": "fallback.example:25565",
}
})
registerGatewayUpstreamHook(
t,
@@ -105,7 +105,7 @@ func TestMapToHostRejectsInvalidOrUnroutedPackets(t *testing.T) {
if tt.packet == nil {
source.readErr = errors.New("read failed")
}
config.Hosts = tt.hosts
setGatewayTestRoutes(tt.hosts)
if got := mapToHost(source); got != nil {
t.Fatalf("mapToHost() = %v, want nil", got)
@@ -118,9 +118,9 @@ func TestMapToHostReturnsNilWhenHookFails(t *testing.T) {
defer saveGatewayState(t)()
source := newGatewayTestConn(gatewayTestPacket("play.example"))
config.Hosts = map[string]string{
setGatewayTestRoutes(map[string]string{
"play.example": "backend.example:25565",
}
})
wantErr := errors.New("hook failed")
registerGatewayUpstreamHook(
@@ -142,9 +142,9 @@ func TestMapToHostClosesUpstreamWhenInitialWriteFails(t *testing.T) {
source := newGatewayTestConn(gatewayTestPacket("play.example"))
upstream := newGatewayTestConn(nil)
upstream.writeErr = errors.New("write failed")
config.Hosts = map[string]string{
setGatewayTestRoutes(map[string]string{
"play.example": "backend.example:25565",
}
})
registerGatewayUpstreamHook(
t,

View File

@@ -72,12 +72,14 @@ func upstreamQuic(host string) net.Conn {
conn, err := quic.DialAddr(ctx, host, tlsConf, nil)
if err != nil {
gatewayMetrics.UpstreamDialError()
log.Err(err).Str("host", host).Msg("Failed to dial QUIC")
return nil
}
stream, err := conn.OpenStream()
if err != nil {
gatewayMetrics.UpstreamDialError()
log.Err(err).Str("host", host).Msg("Failed to open stream")
conn.CloseWithError(0, "failed to open stream")
return nil

View File

@@ -72,6 +72,8 @@ func BenchmarkMapToHostInitialPacket(b *testing.B) {
packet := benchmarkHandshakePacket(hostName)
source := newBenchmarkConn(benchmarkAddr("client:25565"))
upstream := newBenchmarkConn(benchmarkAddr("upstream:25565"))
publishRouteSnapshot(map[string]string{hostName: upstreamHost})
defer publishRouteSnapshot(nil)
restore := installBenchmarkUpstreamHook(b, hostName, upstreamHost, upstream)
defer restore()
@@ -191,10 +193,6 @@ func installBenchmarkUpstreamHook(b *testing.B, hostName, upstreamHost string, u
previousConfig := config
previousHooks := hooks
config.Hosts = map[string]string{
hostName: upstreamHost,
}
pluginLock.Lock()
hooks = map[string]map[string]any{
"benchmark": {},

View File

@@ -35,6 +35,7 @@ func runTcp(wg *sync.WaitGroup) {
}
setSocketOptions(conn)
// 处理连接
gatewayMetrics.TCPConnectionStarted()
go handleRequest(conn)
}
}
@@ -42,6 +43,7 @@ func runTcp(wg *sync.WaitGroup) {
func upstreamTcp(host string) net.Conn {
conn, err := tcpDialer.Dial("tcp", host)
if err != nil {
gatewayMetrics.UpstreamDialError()
log.Err(err).Str("host", host).Msg("Error dialing upstream")
return nil
}

View File

@@ -1,54 +1,21 @@
package main
import (
"bytes"
"errors"
"fmt"
"io"
"net"
"net/http"
"sync"
"time"
"github.com/rs/zerolog/log"
"github.com/tursom/mc-gateway/internal/tcphttpmux"
)
const (
defaultTCPPort = 25565
defaultWebSocketPort = 25566
tcpWebInitialPacketTimeout = time.Second
httpConnBacklog = 128
maxHTTPMethodPrefixLen = len("OPTIONS ")
tcpWebInitialPacketTimeout = tcphttpmux.DefaultInitialPacketTimeout
)
var httpMethodPrefixes = [][]byte{
[]byte("GET "),
[]byte("POST "),
[]byte("HEAD "),
[]byte("PUT "),
[]byte("PATCH "),
[]byte("DELETE "),
[]byte("OPTIONS "),
[]byte("CONNECT "),
[]byte("TRACE "),
}
type replayConn struct {
net.Conn
reader io.Reader
}
func (c *replayConn) Read(p []byte) (int, error) {
return c.reader.Read(p)
}
type chanListener struct {
conns chan net.Conn
closed chan struct{}
closeOnce sync.Once
addr net.Addr
}
func normalizedTCPPort() int {
if config.Tcp.Port == 0 {
return defaultTCPPort
@@ -91,189 +58,38 @@ func runTcpWebPortReuse(wg *sync.WaitGroup) {
log.Info().
Int("port", port).
Str("path", normalizedWebSocketPath()).
Msg("Listening for shared TCP and WebSocket connections")
Str("admin_path", adminStartup.AdminPath).
Msg("Listening for shared TCP and Admin connections")
if err := serveTcpWebPortReuse(listener, newWebSocketHandler(), handleRequest); err != nil && !errors.Is(err, net.ErrClosed) {
if err := serveTcpWebPortReuse(listener, newGatewayHTTPHandler(), handleRequest); err != nil && !errors.Is(err, net.ErrClosed) {
log.Fatal().Err(err).
Int("port", port).
Msg("Shared TCP/WebSocket server stopped")
Msg("Shared TCP/Admin server stopped")
}
}
func serveTcpWebPortReuse(listener net.Listener, handler http.Handler, tcpHandler func(net.Conn)) error {
defer listener.Close()
webListener := newChanListener(listener.Addr())
webServer := &http.Server{Handler: handler}
webServerDone := make(chan error, 1)
go func() {
err := webServer.Serve(webListener)
if err != nil && !errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
webServerDone <- err
return
}
webServerDone <- nil
}()
defer func() {
_ = webListener.Close()
_ = webServer.Close()
<-webServerDone
}()
for {
conn, err := listener.Accept()
if err != nil {
if errors.Is(err, net.ErrClosed) {
return err
}
return tcphttpmux.Serve(listener, handler, tcpHandler, tcphttpmux.Options{
InitialPacketTimeout: tcpWebInitialPacketTimeout,
SetSocketOptions: setSocketOptions,
OnTCPConnection: gatewayMetrics.TCPConnectionStarted,
OnAcceptError: func(err error) {
log.Err(err).Msg("Error accepting shared TCP/WebSocket connection")
continue
}
setSocketOptions(conn)
go handleTcpWebPortReuseConn(conn, webListener, tcpHandler, tcpWebInitialPacketTimeout)
}
}
func handleTcpWebPortReuseConn(conn net.Conn, webListener *chanListener, tcpHandler func(net.Conn), timeout time.Duration) {
peeked, err := readInitialPacket(conn, timeout)
if err != nil {
log.Debug().Err(err).
Str("client", conn.RemoteAddr().String()).
Msg("failed to read initial packet")
conn.Close()
return
}
if len(peeked) == 0 {
log.Debug().
Str("client", conn.RemoteAddr().String()).
Msg("initial packet is empty")
conn.Close()
return
}
replayed := newReplayConn(conn, peeked)
if isHTTPInitialPacket(peeked) {
if !webListener.deliver(replayed) {
},
OnInitialPacketError: func(conn net.Conn, err error) {
log.Debug().Err(err).
Str("client", conn.RemoteAddr().String()).
Msg("failed to read initial packet")
},
OnEmptyInitialPacket: func(conn net.Conn) {
log.Debug().
Str("client", conn.RemoteAddr().String()).
Msg("initial packet is empty")
},
OnHTTPDeliveryFailed: func(conn net.Conn) {
log.Debug().
Str("client", conn.RemoteAddr().String()).
Msg("failed to deliver HTTP connection")
conn.Close()
}
return
}
tcpHandler(replayed)
}
func readInitialPacket(conn net.Conn, timeout time.Duration) ([]byte, error) {
if timeout > 0 {
if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil {
return nil, err
}
defer conn.SetReadDeadline(time.Time{})
}
buf := getProxyBuffer()
defer putProxyBuffer(buf)
var peeked []byte
for {
n, err := conn.Read(buf)
if n > 0 {
peeked = append(peeked, buf[:n]...)
if isHTTPInitialPacket(peeked) ||
!isPotentialHTTPInitialPacket(peeked) ||
len(peeked) >= maxHTTPMethodPrefixLen {
return peeked, nil
}
}
if err != nil {
if len(peeked) > 0 && errors.Is(err, io.EOF) {
return peeked, nil
}
return peeked, err
}
if n == 0 {
return peeked, io.ErrNoProgress
}
}
}
func newReplayConn(conn net.Conn, peeked []byte) net.Conn {
return &replayConn{
Conn: conn,
reader: io.MultiReader(bytes.NewReader(peeked), conn),
}
}
func isHTTPInitialPacket(buf []byte) bool {
for _, prefix := range httpMethodPrefixes {
if bytes.HasPrefix(buf, prefix) {
return true
}
}
return false
}
func isPotentialHTTPInitialPacket(buf []byte) bool {
if len(buf) == 0 {
return true
}
for _, prefix := range httpMethodPrefixes {
if len(buf) <= len(prefix) && bytes.HasPrefix(prefix, buf) {
return true
}
}
return false
}
func newChanListener(addr net.Addr) *chanListener {
return &chanListener{
conns: make(chan net.Conn, httpConnBacklog),
closed: make(chan struct{}),
addr: addr,
}
}
func (l *chanListener) Accept() (net.Conn, error) {
select {
case conn := <-l.conns:
return conn, nil
case <-l.closed:
return nil, net.ErrClosed
}
}
func (l *chanListener) Close() error {
l.closeOnce.Do(func() {
close(l.closed)
},
})
return nil
}
func (l *chanListener) Addr() net.Addr {
return l.addr
}
func (l *chanListener) deliver(conn net.Conn) bool {
select {
case <-l.closed:
return false
default:
}
select {
case l.conns <- conn:
return true
case <-l.closed:
return false
default:
return false
}
}

View File

@@ -1,60 +1,58 @@
package main
import (
"bytes"
"errors"
"io"
"net"
"net/http"
"strings"
"testing"
"time"
"github.com/gorilla/websocket"
"github.com/tursom/mc-gateway/internal/gatewayconfig"
"github.com/tursom/mc-gateway/protocol"
)
func TestTCPWebPortReuseEnabledUsesNormalizedPorts(t *testing.T) {
tests := []struct {
name string
tcp ProtocolConfig
websocket WebSocketConfig
tcp gatewayconfig.ProtocolConfig
websocket gatewayconfig.WebSocketConfig
want bool
}{
{
name: "explicit same port",
tcp: ProtocolConfig{Enable: true, Port: 25565},
websocket: WebSocketConfig{Enable: true, Port: 25565},
tcp: gatewayconfig.ProtocolConfig{Enable: true, Port: 25565},
websocket: gatewayconfig.WebSocketConfig{Enable: true, Port: 25565},
want: true,
},
{
name: "websocket default port",
tcp: ProtocolConfig{Enable: true, Port: defaultWebSocketPort},
websocket: WebSocketConfig{Enable: true},
tcp: gatewayconfig.ProtocolConfig{Enable: true, Port: defaultWebSocketPort},
websocket: gatewayconfig.WebSocketConfig{Enable: true},
want: true,
},
{
name: "tcp default port",
tcp: ProtocolConfig{Enable: true},
websocket: WebSocketConfig{Enable: true, Port: defaultTCPPort},
tcp: gatewayconfig.ProtocolConfig{Enable: true},
websocket: gatewayconfig.WebSocketConfig{Enable: true, Port: defaultTCPPort},
want: true,
},
{
name: "different ports",
tcp: ProtocolConfig{Enable: true, Port: 25565},
websocket: WebSocketConfig{Enable: true, Port: 25566},
tcp: gatewayconfig.ProtocolConfig{Enable: true, Port: 25565},
websocket: gatewayconfig.WebSocketConfig{Enable: true, Port: 25566},
want: false,
},
{
name: "tcp disabled",
tcp: ProtocolConfig{Enable: false, Port: 25565},
websocket: WebSocketConfig{Enable: true, Port: 25565},
tcp: gatewayconfig.ProtocolConfig{Enable: false, Port: 25565},
websocket: gatewayconfig.WebSocketConfig{Enable: true, Port: 25565},
want: false,
},
{
name: "websocket disabled",
tcp: ProtocolConfig{Enable: true, Port: 25565},
websocket: WebSocketConfig{Enable: false, Port: 25565},
tcp: gatewayconfig.ProtocolConfig{Enable: true, Port: 25565},
websocket: gatewayconfig.WebSocketConfig{Enable: false, Port: 25565},
want: false,
},
}
@@ -73,156 +71,6 @@ func TestTCPWebPortReuseEnabledUsesNormalizedPorts(t *testing.T) {
}
}
func TestHTTPInitialPacketRecognition(t *testing.T) {
for _, method := range []string{
"GET ",
"POST ",
"HEAD ",
"PUT ",
"PATCH ",
"DELETE ",
"OPTIONS ",
"CONNECT ",
"TRACE ",
} {
t.Run(strings.TrimSpace(method), func(t *testing.T) {
if !isHTTPInitialPacket([]byte(method + "/gateway HTTP/1.1\r\n")) {
t.Fatalf("isHTTPInitialPacket(%q) = false, want true", method)
}
})
}
if isHTTPInitialPacket(gatewayTestPacket("play.example")) {
t.Fatal("Minecraft handshake was recognized as HTTP")
}
if isHTTPInitialPacket([]byte("GE")) {
t.Fatal("partial HTTP method was recognized as complete HTTP")
}
if !isPotentialHTTPInitialPacket([]byte("GE")) {
t.Fatal("partial HTTP method was not recognized as a possible HTTP prefix")
}
if isPotentialHTTPInitialPacket([]byte("GOT ")) {
t.Fatal("invalid HTTP method was recognized as a possible HTTP prefix")
}
}
func TestReplayConnReadsPeekedBytesBeforeUnderlyingConn(t *testing.T) {
base := newGatewayTestConn([]byte("rest"))
conn := newReplayConn(base, []byte("peek-"))
got, err := io.ReadAll(conn)
if err != nil {
t.Fatalf("ReadAll() error = %v", err)
}
if string(got) != "peek-rest" {
t.Fatalf("replayed data = %q, want peek-rest", got)
}
}
func TestChanListenerAcceptCloseAndDeliver(t *testing.T) {
listener := newChanListener(benchmarkAddr("listener"))
conn := newGatewayTestConn(nil)
if !listener.deliver(conn) {
t.Fatal("deliver() = false, want true")
}
got, err := listener.Accept()
if err != nil {
t.Fatalf("Accept() error = %v", err)
}
if got != conn {
t.Fatalf("Accept() = %v, want delivered conn", got)
}
if err := listener.Close(); err != nil {
t.Fatalf("Close() error = %v", err)
}
if listener.deliver(newGatewayTestConn(nil)) {
t.Fatal("deliver() after Close = true, want false")
}
_, err = listener.Accept()
if !errors.Is(err, net.ErrClosed) {
t.Fatalf("Accept() error = %v, want %v", err, net.ErrClosed)
}
}
func TestHandleTcpWebPortReuseConnRoutesHTTP(t *testing.T) {
defer saveGatewayState(t)()
listener := newChanListener(benchmarkAddr("listener"))
source := newGatewayTestConn([]byte("GET / HTTP/1.1\r\n\r\n"))
tcpCalled := false
handleTcpWebPortReuseConn(source, listener, func(net.Conn) {
tcpCalled = true
}, time.Second)
if tcpCalled {
t.Fatal("TCP handler was called for HTTP request")
}
conn, err := listener.Accept()
if err != nil {
t.Fatalf("Accept() error = %v", err)
}
got, err := io.ReadAll(conn)
if err != nil {
t.Fatalf("ReadAll() error = %v", err)
}
if string(got) != "GET / HTTP/1.1\r\n\r\n" {
t.Fatalf("HTTP replay = %q", got)
}
}
func TestHandleTcpWebPortReuseConnRoutesMinecraft(t *testing.T) {
defer saveGatewayState(t)()
packet := gatewayTestPacket("play.example")
listener := newChanListener(benchmarkAddr("listener"))
source := newGatewayTestConn(packet)
var got []byte
handleTcpWebPortReuseConn(source, listener, func(conn net.Conn) {
var err error
got, err = io.ReadAll(conn)
if err != nil {
t.Fatalf("ReadAll() error = %v", err)
}
}, time.Second)
if !bytes.Equal(got, packet) {
t.Fatalf("Minecraft replay = %v, want %v", got, packet)
}
}
func TestHandleTcpWebPortReuseConnTimeoutClosesConn(t *testing.T) {
defer saveGatewayState(t)()
client, server := net.Pipe()
defer client.Close()
listener := newChanListener(benchmarkAddr("listener"))
done := make(chan struct{})
go func() {
handleTcpWebPortReuseConn(server, listener, func(net.Conn) {
t.Error("TCP handler was called after timeout")
}, 10*time.Millisecond)
close(done)
}()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for initial packet timeout")
}
if _, err := client.Write([]byte("x")); err == nil {
t.Fatal("client Write() error = nil, want closed connection error")
}
}
func TestServeTcpWebPortReuseServesHTTPOnSharedPort(t *testing.T) {
defer saveGatewayState(t)()

View File

@@ -11,6 +11,10 @@ import (
"github.com/rs/zerolog"
"github.com/rs/zerolog/log"
"github.com/tursom/mc-gateway/internal/adminconfig"
"github.com/tursom/mc-gateway/internal/adminsession"
"github.com/tursom/mc-gateway/internal/gatewayconfig"
"github.com/tursom/mc-gateway/internal/gatewaymetrics"
"github.com/tursom/mc-gateway/plugin/api"
)
@@ -18,10 +22,15 @@ func saveGatewayState(t *testing.T) func() {
t.Helper()
oldConfig := config
oldConfigFile := configFile
oldCurrentPidFile := currentPidFile
oldCurrentLogFile := currentLogFile
oldLogger := log.Logger
oldAdminStartup := adminStartup
oldAdminDB := adminDB
oldAdminDBPath := adminDBPath
oldAdminSessionManager := adminSessionManager
oldRouteSnapshot := routeSnapshot.Clone()
oldGatewayMetrics := gatewayMetrics
pluginLock.Lock()
oldPlugins := plugins
@@ -30,21 +39,40 @@ func saveGatewayState(t *testing.T) func() {
hooks = make(map[string]map[string]any)
pluginLock.Unlock()
config = Config{}
configFile = "config.toml"
config = gatewayconfig.Config{}
currentPidFile = ""
currentLogFile = ""
adminStartup = adminconfig.Config{
DBPath: defaultAdminDBPath,
TCPAdminPort: defaultTCPPort,
AdminPath: defaultAdminPath,
AdminAPIPrefix: defaultAdminAPIPrefix,
SessionTTL: defaultAdminSessionTTL,
}
adminDB = nil
adminDBPath = ""
adminSessionManager = adminsession.NewManager()
publishRouteSnapshot(nil)
gatewayMetrics = gatewaymetrics.New()
log.Logger = zerolog.New(io.Discard)
return func() {
if adminDB != nil && adminDB != oldAdminDB {
_ = adminDB.Close()
}
if currentPidFile != "" && currentPidFile != oldCurrentPidFile {
_ = os.Remove(currentPidFile)
}
config = oldConfig
configFile = oldConfigFile
currentPidFile = oldCurrentPidFile
currentLogFile = oldCurrentLogFile
adminStartup = oldAdminStartup
adminDB = oldAdminDB
adminDBPath = oldAdminDBPath
adminSessionManager = oldAdminSessionManager
publishRouteSnapshot(oldRouteSnapshot)
gatewayMetrics = oldGatewayMetrics
log.Logger = oldLogger
pluginLock.Lock()
@@ -54,6 +82,10 @@ func saveGatewayState(t *testing.T) func() {
}
}
func setGatewayTestRoutes(routes map[string]string) {
publishRouteSnapshot(routes)
}
func gatewayTestPacket(host string, tail ...byte) []byte {
packet := []byte{
byte(4 + 1 + len(host) + len(tail)),

View File

@@ -37,6 +37,7 @@ func handleWebSocket(w http.ResponseWriter, r *http.Request) {
}
defer conn.Close()
gatewayMetrics.WebSocketConnectionStarted()
handleRequest(&webSocketConn{Conn: conn})
}

432
docs/admin-page-design.md Normal file
View File

@@ -0,0 +1,432 @@
# 管理页面设计
## 背景
当前 gateway 已经具备 TCP 与 HTTP/WebSocket 的端口复用能力:
- Minecraft 原始 TCP 连接进入 `handleRequest`,再按握手包里的 host 路由到上游。
- WebSocket 连接先进入 HTTP handler再 upgrade 为 WebSocket最后同样复用 `handleRequest`
- `tcp_web_port_reuse.go` 已经提供首包识别、连接回放和 `http.Server.Serve(chanListener)` 这条分流路径。
新的管理方案不再以 `config.toml` 为中心。进程默认启动即可工作:
- 默认在 `25565` 端口启动 Minecraft TCP 监听。
- 同一个 `25565` 端口同时承载后台管理页面和管理 API。
- 用户、权限、路由、其他协议服务配置都持久化到 SQLite。
- 其他服务和路由通过后台管理配置,不再依赖配置文件。
## 功能需求对齐
第一版管理页面按以下边界设计。
### 必须支持
- 进程启动不依赖配置文件;缺省状态下直接监听 `25565`
- TCP 转发和后台管理默认共用同一个 TCP listener。
- 管理入口必须有认证保护。复用公网 Minecraft 端口时,不能出现无认证管理 API。
- 管理页面内置用户和权限系统,角色分为管理员、成员和游客。
- 所有持久化状态进入 SQLite包括用户、路由、服务配置和审计日志。
- 页面提供一个实际可用的后台界面,而不是只返回 JSON
- 概览:进程 PID、运行时长、SQLite 路径、TCP/Admin 端口、各协议服务状态。
- 路由管理:展示、搜索、新增、修改、删除 SQLite 中的路由记录。
- 服务管理:配置 KCP、QUIC、WebSocket 的启用状态、端口和协议参数。
- 用户管理:管理员可以新增、禁用、修改角色和重置密码。
- 连接统计:展示总连接数、当前活跃连接数、路由命中次数、上游连接失败次数。
- 游客登录后可以查看当前 SQLite 路由记录,但不能修改路由、配置服务或查看用户列表。
- 路由记录必须持久化到 SQLite进程重启后仍然生效。
- 服务配置必须持久化到 SQLite进程重启后按后台配置恢复。
- 所有写操作需要记录审计日志,至少包含操作者、来源 IP、操作类型、目标对象和结果。
### 暂不支持
- 不再支持 `config.toml` 作为启动配置、热加载配置或路由来源。
- 不做多租户、组织空间和自定义细粒度权限。
- 不做 HTTPS/TLS 终止;需要 HTTPS 时由外部反向代理或负载均衡器处理。
- 不做 Minecraft 玩家在线列表、踢人、封禁等游戏服管理能力。
- 不做插件管理、二进制升级、进程重启。
- 不把每条连接的完整客户端地址、目标地址长期保存在内存里。
## 默认启动设计
无配置文件时使用以下默认值:
| 项 | 默认值 | 环境变量 |
| --- | --- | --- |
| SQLite 数据库 | `mc-gateway.sqlite3` | `MC_GATEWAY_DB` |
| TCP/Admin 监听端口 | `25565` | `MC_GATEWAY_TCP_ADMIN_PORT` |
| Admin 页面路径 | `/admin/` | `MC_GATEWAY_ADMIN_PATH` |
| Admin API 前缀 | `/admin/api` | `MC_GATEWAY_ADMIN_API_PREFIX` |
| KCP | 默认禁用 | 后台配置 |
| QUIC | 默认禁用 | 后台配置 |
| WebSocket | 默认禁用 | 后台配置 |
| 路由表 | 默认空表 | 后台配置 |
| 会话有效期 | 8 小时 | 后台配置 |
启动期环境变量规则:
- `MC_GATEWAY_TCP_ADMIN_PORT` 只在启动时读取,必须是 `1-65535` 的整数;为空时使用 `25565`
- `MC_GATEWAY_ADMIN_PATH` 只在启动时读取,必须以 `/` 开头,规范化为以 `/` 结尾;为空时使用 `/admin/`
- `MC_GATEWAY_ADMIN_API_PREFIX` 只在启动时读取,必须以 `/` 开头,规范化为不以 `/` 结尾;为空时使用 `/admin/api`
- `MC_GATEWAY_ADMIN_API_PREFIX` 不能等于 `MC_GATEWAY_ADMIN_PATH`,也不能落在静态资源路径下。
- 环境变量覆盖的是本次进程的 Admin 入口;第一次创建 SQLite 默认服务配置时,应把解析后的 TCP/Admin 端口写入 `services.tcp_admin.port`
启动流程:
1. 读取并校验启动期环境变量,得到 SQLite 路径、TCP/Admin 端口、Admin 页面路径、Admin API 前缀。
2. 打开 SQLite 数据库;不存在时自动创建。
3. 执行 schema migration。
4. 确保默认服务配置存在TCP/Admin listener 启用,端口为启动期解析后的端口。
5. 如果用户表为空,进入首次初始化模式。
6. 加载启用的路由记录为内存只读快照。
7. 启动 TCP/Admin 共享 listener。
8. 按 SQLite 中的服务配置启动 KCP、QUIC、WebSocket。
首次初始化:
- 当用户表为空时,启动期解析后的 Admin 页面路径显示初始化管理员页面。
- 初始化接口只允许创建第一个管理员账号。
- 第一个管理员创建成功后,初始化接口永久关闭。
- 也可以通过环境变量 `MC_GATEWAY_ADMIN_PASSWORD` 配合默认用户名 `admin` 在启动时创建初始管理员;这不是配置文件,只是无交互部署入口。
SQLite 路径:
- 默认使用当前工作目录下的 `mc-gateway.sqlite3`
- 如需改路径,第一版只接受启动参数或环境变量,例如 `MC_GATEWAY_DB`;不引入配置文件。
## 运行模型
TCP/Admin listener 是基础入口:
```text
net.Listen(:startup env port or db service setting)
|
v
Accept
|
v
读取首批字节
|
+-- HTTP/Admin/WebSocket
| -> HTTP channel listener
| -> http.Server.Serve
|
+-- Minecraft TCP
-> 回放首批字节
-> handleRequest
-> mapToHost
-> route snapshot
-> proxyConnections
```
核心规则:
- TCP/Admin listener 默认永远启用,避免后台管理入口丢失。
- TCP/Admin 端口可以由启动期环境变量覆盖,也可以保存在 SQLite最终监听端口以启动期解析结果优先。
- TCP/Admin 端口修改后第一版按重启后生效处理。
- Admin 页面路径和 API 前缀只由启动期环境变量控制,不通过后台页面修改,避免运行中替换管理入口导致当前会话失效。
- KCP、QUIC、WebSocket 由后台管理配置启用状态和端口。
- 如果某个服务的端口或协议参数无法热更新,页面必须标记为“重启后生效”或提供明确的重启服务操作。
- 路由变更不需要重启,也不需要 reloadSQLite 写入成功后刷新路由快照即可影响新连接。
## HTTP 路由与页面
页面使用 Go `embed` 打包到单个二进制,不引入 Node 构建链。
建议目录:
```text
cmd/gateway/admin.go
cmd/gateway/admin_api.go
cmd/gateway/admin_auth.go
cmd/gateway/admin_db.go
cmd/gateway/admin_static.go
cmd/gateway/admin_static/
index.html
app.css
app.js
```
API 路由:
| 方法 | 路径 | 权限 | 说明 |
| --- | --- | --- | --- |
| GET | `/admin/` | 公开页面 | 管理页面入口或首次初始化页面 |
| GET | `/admin/app.css` | 公开页面 | 页面样式 |
| GET | `/admin/app.js` | 公开页面 | 页面脚本 |
| POST | `/admin/api/setup` | 仅用户表为空 | 创建第一个管理员 |
| POST | `/admin/api/auth/login` | 未登录 | 用户名密码登录 |
| POST | `/admin/api/auth/logout` | 已登录 | 注销当前会话 |
| GET | `/admin/api/me` | 已登录 | 当前用户、角色和权限 |
| GET | `/admin/api/status` | 成员、管理员 | 进程、端口、数据库、服务状态 |
| GET | `/admin/api/routes` | 游客、成员、管理员 | 当前 SQLite 路由列表 |
| PUT | `/admin/api/routes/{host}` | 成员、管理员 | 新增或更新一条路由 |
| DELETE | `/admin/api/routes/{host}` | 成员、管理员 | 删除一条路由 |
| GET | `/admin/api/services` | 成员、管理员 | 服务配置和运行状态 |
| PUT | `/admin/api/services/{name}` | 管理员 | 修改 KCP、QUIC、WebSocket 配置 |
| POST | `/admin/api/services/{name}/restart` | 管理员 | 重启指定服务 |
| GET | `/admin/api/metrics` | 成员、管理员 | 连接和路由统计 |
| GET | `/admin/api/users` | 管理员 | 用户列表 |
| POST | `/admin/api/users` | 管理员 | 新增用户 |
| PATCH | `/admin/api/users/{username}` | 管理员 | 修改用户角色、状态或密码 |
| DELETE | `/admin/api/users/{username}` | 管理员 | 删除用户 |
| GET | `/admin/api/audit-logs` | 管理员 | 审计日志 |
页面形态:
- 用户表为空时只显示初始化管理员界面。
- 未登录且已初始化时只显示登录界面。
- 成员和管理员顶部固定显示 gateway 状态、SQLite 路径和服务状态。
- 路由表支持按 host、upstream、启用状态搜索。
- 游客登录后只展示当前路由表,不展示新增、编辑、删除、服务配置、用户管理入口。
- 新增、编辑、删除路由使用弹窗或行内表单,不单独跳转页面。
- 删除 `default` 路由需要二次确认,因为它是 fallback 路由。
- API 错误直接展示服务端返回的错误信息,便于运维判断问题。
## 用户与权限设计
第一版采用 SQLite 用户表和内存会话,不开放匿名管理 API。这里的“游客”是已登录用户的只读角色不是未登录访问。
角色权限:
| 角色 | 权限 |
| --- | --- |
| 管理员 `admin` | 所有管理能力,包括用户管理、服务配置、路由管理、状态和指标查看 |
| 成员 `member` | 查看状态、指标和服务状态;新增、修改、删除路由 |
| 游客 `guest` | 只能登录、注销、查看自己的信息、查看当前 SQLite 路由 |
认证流程:
- `POST /admin/api/setup` 只在用户表为空时可用,用于创建第一个管理员。
- `POST /admin/api/auth/login` 使用用户名和密码登录。
- 登录成功后服务端生成随机会话 token响应给前端。
- 前端把会话 token 保存在 `sessionStorage`,后续 API 使用 `Authorization: Bearer <session_token>`
- 服务端在内存中保存会话,超过会话有效期后失效;进程重启后所有会话失效。
- 未登录返回 `401`,已登录但权限不足返回 `403`
用户规则:
- 密码只保存哈希,建议使用 bcrypt 或 argon2id不保存明文密码。
- 用户字段至少包含:`username``role``password_hash``disabled``created_at``updated_at`
- 管理员不能删除或禁用最后一个可用管理员账号。
- 修改用户、重置密码、禁用用户都要记录审计日志。
## SQLite 存储设计
建议使用一个 SQLite 数据库保存所有后台状态。优先选择不依赖 CGO 的 SQLite driver降低交叉编译和容器部署成本。
基础表:
```sql
CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
applied_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS users (
username TEXT PRIMARY KEY,
role TEXT NOT NULL,
password_hash TEXT NOT NULL,
disabled INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS routes (
host TEXT PRIMARY KEY,
upstream TEXT NOT NULL,
enabled INTEGER NOT NULL DEFAULT 1,
note TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
updated_by TEXT NOT NULL DEFAULT ''
);
CREATE TABLE IF NOT EXISTS services (
name TEXT PRIMARY KEY,
enabled INTEGER NOT NULL,
port INTEGER NOT NULL DEFAULT 0,
options_json TEXT NOT NULL DEFAULT '{}',
restart_required INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
updated_by TEXT NOT NULL DEFAULT ''
);
CREATE TABLE IF NOT EXISTS audit_logs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
actor TEXT NOT NULL,
source_ip TEXT NOT NULL,
action TEXT NOT NULL,
target_type TEXT NOT NULL,
target_id TEXT NOT NULL,
success INTEGER NOT NULL,
message TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_routes_enabled ON routes(enabled);
CREATE INDEX IF NOT EXISTS idx_audit_logs_created_at ON audit_logs(created_at);
```
服务配置记录:
| name | 默认 enabled | 默认 port | 说明 |
| --- | --- | --- | --- |
| `tcp_admin` | 1 | 25565 | Minecraft TCP 和 Admin 共用入口 |
| `kcp` | 0 | 25566 | KCP 服务 |
| `quic` | 0 | 25565 | QUIC 服务,具体协议参数放在 `options_json` |
| `websocket` | 0 | 25566 | WebSocket 服务path 放在 `options_json` |
注意:
- `tcp_admin` 是基础入口,不允许在后台禁用。
- `tcp_admin.port` 修改后第一版按重启后生效处理。
- KCP、QUIC、WebSocket 配置修改后,如果当前实现无法安全热更新,则标记 `restart_required = 1` 并在页面提示。
- 后续可以把服务运行状态放在内存中,不需要写入 SQLite。
SQLite 运行参数:
- 启动后设置 `PRAGMA journal_mode=WAL`
- 设置合理的 `busy_timeout`,避免后台写入和读取短暂冲突直接失败。
- 所有写操作使用事务。
- 路由和服务配置写入成功后,刷新对应的内存快照。
## 路由存储与运行时读取
路由管理以 SQLite 为唯一持久化来源。
运行时读取:
- 启动时从 SQLite 加载所有 `enabled = 1` 的路由,发布为内存只读快照。
- `mapToHost` 从路由快照读取 host 到 upstream 的映射,不再读取 `config.Hosts`
- 路由 API 写入 SQLite 成功后,重新加载路由快照;不需要 reload。
- 如果快照中没有目标 host则尝试 `default`;没有 `default` 时按现有未命中逻辑关闭连接并记录日志。
写入流程:
1. 校验 host、upstream、enabled 和 note。
2. 在 SQLite 事务中执行 insert、update、delete 或 enable/disable。
3. 事务提交成功后重新加载路由快照。
4. 快照发布成功后返回 API 成功。
5. 如果快照发布失败,返回错误并保留 SQLite 已提交数据;旧路由快照继续服务已有逻辑。
并发与一致性:
- 后台写操作串行执行,避免并发修改同一 host 时互相覆盖。
- 路由快照使用 copy-on-write 发布,转发热路径只做只读 map 查询。
- 没有 `config.toml` 热加载,也不存在配置文件覆盖后台路由的问题。
## 路由校验规则
`host` 校验:
- 不能为空。
- 允许 `default`
- 不允许包含空白字符。
- 不允许包含 `/`,避免和 URL 路由混淆。
`upstream` 校验:
- 不能为空。
- 支持现有前缀:无前缀 TCP、`kcp://``quic://``haproxy://`
- 去掉协议前缀后,必须能解析为 `host:port`
- 不在第一版新增 websocket upstream 前缀,因为当前 README 未列出该前缀的实际实现。
## 指标设计
第一版只做轻量统计,避免影响 TCP 转发热路径。
建议新增全局 `gatewayMetrics`,内部使用 `sync/atomic`
- `total_connections`
- `active_connections`
- `tcp_connections`
- `websocket_connections`
- `route_hits{host}`
- `route_misses`
- `upstream_dial_errors`
实现原则:
- 只在连接进入、退出、路由成功或失败时更新计数。
- 不在每次 `copyForward` 读写时计数。
- route 维度只记录 host 计数,不记录完整客户端 IP。
## 测试计划
单元测试:
- 无配置文件时默认创建 SQLite 并写入默认服务配置。
- 默认启动配置包含 `tcp_admin` enabled 和端口 `25565`
- Admin path、api_path 默认分别为 `/admin/``/admin/api`
- `MC_GATEWAY_TCP_ADMIN_PORT``MC_GATEWAY_ADMIN_PATH``MC_GATEWAY_ADMIN_API_PREFIX` 能覆盖启动期入口配置。
- 启动期环境变量非法时返回明确错误,不创建含错误值的默认服务配置。
- 首次初始化管理员、重复初始化拒绝、管理员登录和禁用用户拒绝登录。
- 登录、注销、会话过期。
- 权限 middleware 的 200、401、403 场景。
- 管理员用户 API 的新增、改角色、重置密码、禁用和删除。
- 不能删除或禁用最后一个可用管理员。
- 路由 API 的新增、修改、删除和校验失败。
- SQLite schema migration、空库初始化、WAL 和 busy timeout 配置。
- 路由写入 SQLite 后能刷新内存路由快照。
- 服务配置 API 能更新 KCP、QUIC、WebSocket 配置并标记是否需要重启。
- 游客可以读取 routes但不能写 routes、配置 services、读取 metrics 或 users。
集成测试:
- 不提供 `config.toml` 时,进程默认通过 `25565` 提供 TCP/Admin 入口。
- 设置 `MC_GATEWAY_TCP_ADMIN_PORT` 和 Admin path/API prefix 后,进程通过指定端口和路径提供后台入口。
- 设置自定义 Admin API prefix 后,登录、状态、路由等 API 都挂载到新的 prefix 下。
- 用户表为空时可访问初始化页面并创建第一个管理员。
- 管理员或成员登录后可访问 `/admin/api/status`
- 游客账号登录后可以通过 `25565` 读取 `/admin/api/routes`
- 成员通过后台新增或修改路由后,新的 Minecraft host 立即按 SQLite 记录转发。
- 同一端口下 Minecraft 握手仍然进入 TCP 分支,并能完成 host 路由。
- WebSocket 启用后,按后台配置的端口和 path 提供服务。
性能验证:
- 管理页面变更后继续运行现有 TCP 转发 benchmark。
- 重点确认路由快照查询和指标计数没有进入已建立连接后的 `copyForward` 热路径。
建议命令:
```bash
go test ./...
go test -race ./...
go test -run TestTcpWebPortReuse ./cmd/gateway
go test -bench 'Benchmark.*tcp' ./cmd/gateway
```
## 实施步骤
1. 移除 `config.toml` 启动依赖,改为无配置默认值启动。
2. 新增 SQLite 打开、schema migration、WAL 和默认服务配置初始化。
3. 增加启动期环境变量解析和校验,覆盖 TCP/Admin 端口、Admin 页面路径和 Admin API 前缀。
4. 将 TCP/Admin 共享 listener 固定为基础入口,默认端口 `25565`,并支持启动期端口覆盖。
5. 把 WebSocket handler 创建逻辑扩展为统一 `newGatewayHTTPHandler()`
6. 增加用户表、首次初始化管理员、密码哈希和会话管理。
7. 增加 Admin 静态页面、初始化页面、登录页面和权限 middleware。
8. 增加 auth、me、users、status、routes、services、metrics、audit-logs API。
9.`mapToHost` 改为读取 SQLite 路由快照,不再读取 `config.Hosts`
10. 实现 routes 写入 SQLite、事务提交和路由快照刷新。
11. 实现 services 写入 SQLite并接入 KCP、QUIC、WebSocket 启动配置。
12. 实现 users 写入 SQLite 和权限校验。
13. 移除或废弃配置文件 watcher 和基于文件的 reload 逻辑。
14. 补充测试和 benchmark 验证。
## 验收标准
- 没有 `config.toml`gateway 默认在 `25565` 启动 TCP 转发和后台管理页面。
- `http://<host>:25565/admin/` 可访问初始化或登录页面。
- 设置 `MC_GATEWAY_TCP_ADMIN_PORT``MC_GATEWAY_ADMIN_PATH``MC_GATEWAY_ADMIN_API_PREFIX` 后,后台入口使用环境变量指定的端口和路径。
- 同一端口下 Minecraft 客户端连接和 host 路由行为正常。
- 用户表为空时只能创建第一个管理员;创建后初始化接口不可再次使用。
- 未登录的管理 API 请求返回 `401`,已登录但权限不足返回 `403`
- 管理员可以新增用户、禁用用户、修改角色和重置密码。
- 管理员可以配置 KCP、QUIC、WebSocket 服务。
- 成员可以查看、搜索、新增、修改、删除 SQLite 路由记录,变更后立即影响新连接路由。
- 游客可以查看和搜索当前 SQLite 路由记录,但看不到写操作入口,直接调用写 API 返回 `403`
- 删除或修改旧 `config.toml` 不影响后台管理中的用户、服务配置和路由。
- SQLite 写入失败或路由快照刷新失败时,页面展示明确错误,旧路由快照仍可继续工作。
- `go test ./...``go test -race ./...` 和 TCP benchmark 完成后无新增失败。

29
go.mod
View File

@@ -1,40 +1,47 @@
module github.com/tursom/mc-gateway
go 1.24
go 1.24.0
toolchain go1.24.4
require (
github.com/BurntSushi/toml v1.5.0
github.com/fsnotify/fsnotify v1.7.0
github.com/gorilla/websocket v1.5.3
github.com/mitchellh/mapstructure v1.5.0
github.com/pires/go-proxyproto v0.8.1
github.com/quic-go/quic-go v0.52.0
github.com/rs/zerolog v1.33.0
github.com/xtaci/kcp-go v5.4.20+incompatible
golang.org/x/crypto v0.43.0
modernc.org/sqlite v1.45.0
)
require (
github.com/dustin/go-humanize v1.0.1 // indirect
github.com/go-task/slim-sprig v0.0.0-20230315185526-52ccab3ef572 // indirect
github.com/google/go-cmp v0.7.0 // indirect
github.com/google/pprof v0.0.0-20210407192527-94a9f03dee38 // indirect
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/klauspost/cpuid/v2 v2.2.8 // indirect
github.com/klauspost/reedsolomon v1.12.4 // indirect
github.com/mattn/go-colorable v0.1.13 // indirect
github.com/mattn/go-isatty v0.0.19 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
github.com/ncruces/go-strftime v1.0.0 // indirect
github.com/onsi/ginkgo/v2 v2.9.5 // indirect
github.com/pkg/errors v0.9.1 // indirect
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
github.com/stretchr/testify v1.10.0 // indirect
github.com/templexxx/cpufeat v0.0.0-20180724012125-cef66df7f161 // indirect
github.com/templexxx/xor v0.0.0-20191217153810-f85b25db303b // indirect
github.com/tjfoc/gmsm v1.4.1 // indirect
github.com/xtaci/lossyconn v0.0.0-20200209145036-adba10fffc37 // indirect
go.uber.org/mock v0.5.0 // indirect
golang.org/x/crypto v0.38.0 // indirect
golang.org/x/mod v0.24.0 // indirect
golang.org/x/net v0.40.0 // indirect
golang.org/x/sync v0.14.0 // indirect
golang.org/x/sys v0.33.0 // indirect
golang.org/x/tools v0.33.0 // indirect
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect
golang.org/x/mod v0.29.0 // indirect
golang.org/x/net v0.46.0 // indirect
golang.org/x/sync v0.17.0 // indirect
golang.org/x/sys v0.37.0 // indirect
golang.org/x/tools v0.38.0 // indirect
modernc.org/libc v1.67.6 // indirect
modernc.org/mathutil v1.7.1 // indirect
modernc.org/memory v1.11.0 // indirect
)

84
go.sum
View File

@@ -1,22 +1,17 @@
cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw=
github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU=
github.com/BurntSushi/toml v1.5.0 h1:W5quZX/G/csjUnuI8SUYlsHs9M38FC7znL0lIO+DvMg=
github.com/BurntSushi/toml v1.5.0/go.mod h1:ukJfTF/6rtPPRCnwkur4qwRxa8vTRFBF0uk2lLoLwho=
github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU=
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc=
github.com/coreos/go-systemd/v22 v22.5.0/go.mod h1:Y58oyj3AT4RCenI/lSvhwexgC+NSVTIJ3seZv2GcEnc=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
github.com/fsnotify/fsnotify v1.7.0 h1:8JEhPFa5W2WU7YfeZzPNqzMP6Lwt7L2715Ggo0nosvA=
github.com/fsnotify/fsnotify v1.7.0/go.mod h1:40Bi/Hjc2AVfZrqy+aj+yEI+/bRxZnMJyTJwOpGvigM=
github.com/go-logr/logr v1.2.4 h1:g01GSCwiDw2xSZfjJ2/T9M+S6pFdcNtFYsp+Y43HYDQ=
github.com/go-logr/logr v1.2.4/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A=
github.com/go-task/slim-sprig v0.0.0-20230315185526-52ccab3ef572 h1:tfuBGBXKqDEevZMzYi5KSi8KkcZtzBcTgAUUtapy0OI=
@@ -41,11 +36,14 @@ github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMyw
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/pprof v0.0.0-20210407192527-94a9f03dee38 h1:yAJXTCF9TqKcTiHJAE8dj7HMvPfh66eeA2JYW7eFpSE=
github.com/google/pprof v0.0.0-20210407192527-94a9f03dee38/go.mod h1:kpwsk12EmLew5upagYY7GY0pfYCcupk39gWOCRROcvE=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/ianlancetaylor/demangle v0.0.0-20200824232613-28f6c0f3b639/go.mod h1:aSSvb/t6k1mPoxDqO4vJh6VOCGPwU4O0C2/Eqndh1Sc=
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
github.com/klauspost/cpuid/v2 v2.2.8 h1:+StwCXwm9PdpiEkPyzBXIy+M9KUb4ODm0Zarf1kS5BM=
github.com/klauspost/cpuid/v2 v2.2.8/go.mod h1:Lcz8mBdAVJIBVzewtcLocK12l3Y+JytZYpaMropDUws=
github.com/klauspost/reedsolomon v1.12.4 h1:5aDr3ZGoJbgu/8+j45KtUJxzYm8k08JGtB9Wx1VQ4OA=
@@ -53,10 +51,13 @@ github.com/klauspost/reedsolomon v1.12.4/go.mod h1:d3CzOMOt0JXGIFZm1StgkyF14EYr3
github.com/mattn/go-colorable v0.1.13 h1:fFA4WZxdEF4tXPZVKMLwD8oUnCTTo08duU7wxecdEvA=
github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg=
github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM=
github.com/mattn/go-isatty v0.0.19 h1:JITubQf0MOLdlGRuRq+jtsDlekdYPia9ZFsB8h/APPA=
github.com/mattn/go-isatty v0.0.19/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mitchellh/mapstructure v1.5.0 h1:jeMsZIYE/09sWLaz43PL7Gy6RuMjD2eJVyuac5Z2hdY=
github.com/mitchellh/mapstructure v1.5.0/go.mod h1:bFUtVrKA4DC2yAKiSyO/QUcy7e+RRV2QTWOzhPopBRo=
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
github.com/onsi/ginkgo/v2 v2.9.5 h1:+6Hr4uxzP4XIUyAkg61dWBw8lb/gc4/X5luuxN/EC+Q=
github.com/onsi/ginkgo/v2 v2.9.5/go.mod h1:tvAoo1QUJwNEU2ITftXTpR7R1RbCzoZUOs3RonqW57k=
github.com/onsi/gomega v1.27.6 h1:ENqfyGeS5AX/rlXDd/ETokDz93u0YufY1Pgxuy/PvWE=
@@ -70,6 +71,8 @@ github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZN
github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA=
github.com/quic-go/quic-go v0.52.0 h1:/SlHrCRElyaU6MaEPKqKr9z83sBg2v4FLLvWM+Z47pA=
github.com/quic-go/quic-go v0.52.0/go.mod h1:MFlGGpcpJqRAfmYi6NC2cptDPSxRWTOGNuP4wqrWmzQ=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
github.com/rs/xid v1.5.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
github.com/rs/zerolog v1.33.0 h1:1cU2KZkvPxNyfgEmhHAz/1A9Bz+llsdYzklWFzgp0r8=
github.com/rs/zerolog v1.33.0/go.mod h1:/7mN4D5sKwJLZQ2b/znpjC3/GQWY/xaDXUM0kKWRHss=
@@ -92,50 +95,51 @@ go.uber.org/mock v0.5.0/go.mod h1:ge71pBPLYDk7QIi1LupWxdAykm7KIEFchiOqd6z7qMM=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.38.0 h1:jt+WWG8IZlBnVbomuhg2Mdq0+BBQaHbtqHEFEigjUV8=
golang.org/x/crypto v0.38.0/go.mod h1:MvrbAqul58NNYPKnOra203SB9vpuZW0e+RRZV+Ggqjw=
golang.org/x/crypto v0.43.0 h1:dduJYIi3A3KOfdGOHX8AVZ/jGiyPa3IbBozJ5kNuE04=
golang.org/x/crypto v0.43.0/go.mod h1:BFbav4mRNlXJL4wNeejLpWxB7wMbc79PdRGhWKncxR0=
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY=
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70=
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU=
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
golang.org/x/mod v0.24.0 h1:ZfthKaKaT4NrhGVZHO1/WDTwGES4De8KtWO0SIbNJMU=
golang.org/x/mod v0.24.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww=
golang.org/x/mod v0.29.0 h1:HV8lRxZC4l2cr3Zq1LvtOsi/ThTgWnUk/y64QSs8GwA=
golang.org/x/mod v0.29.0/go.mod h1:NyhrlYXJ2H4eJiRy/WDBO6HMqZQ6q9nk4JzS3NuCK+w=
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.40.0 h1:79Xs7wF06Gbdcg4kdCCIQArK11Z1hr5POQ6+fIYHNuY=
golang.org/x/net v0.40.0/go.mod h1:y0hY0exeL2Pku80/zKK7tpntoX23cqL3Oa6njdgRtds=
golang.org/x/net v0.46.0 h1:giFlY12I07fugqwPuWJi68oOnpfqFnJIJzaIIm2JVV4=
golang.org/x/net v0.46.0/go.mod h1:Q9BGdFy1y4nkUwiLvT5qtyhAnEHgnQ/zd8PfU6nc210=
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.14.0 h1:woo0S4Yywslg6hp4eUFjTVOyKt0RookbpAHG4c1HmhQ=
golang.org/x/sync v0.14.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20191204072324-ce4227a45e2e/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/sys v0.37.0 h1:fdNQudmxPjkdUTPnLn5mdQv7Zwvbvpaxqs831goi9kQ=
golang.org/x/sys v0.37.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.25.0 h1:qVyWApTSYLk/drJRO5mDlNYskwQznZmkpV2c8q9zls4=
golang.org/x/text v0.25.0/go.mod h1:WEdwpYrmk1qmdHvhkSTNPm3app7v4rsT8F2UD6+VHIA=
golang.org/x/text v0.30.0 h1:yznKA/E9zq54KzlzBEAWn1NXSQ8DIp/NYMy88xJjl4k=
golang.org/x/text v0.30.0/go.mod h1:yDdHFIX9t+tORqspjENWgzaCVXgk0yYnYuSZ8UzzBVM=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
golang.org/x/tools v0.33.0 h1:4qz2S3zmRxbGIhDIAgjxvFutSvH5EfnsYrRBj0UI0bc=
golang.org/x/tools v0.33.0/go.mod h1:CIJMaWEY88juyUfo7UbgPqbC8rU2OqfAV1h2Qp0oMYI=
golang.org/x/tools v0.38.0 h1:Hx2Xv8hISq8Lm16jvBZ2VQf+RLmbd7wVUsALibYI/IQ=
golang.org/x/tools v0.38.0/go.mod h1:yEsQ/d/YK8cjh0L6rZlY8tgtlKiBNTL14pGDJPJpYQs=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM=
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
@@ -159,3 +163,31 @@ gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE=
modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
modernc.org/libc v1.67.6 h1:eVOQvpModVLKOdT+LvBPjdQqfrZq+pC39BygcT+E7OI=
modernc.org/libc v1.67.6/go.mod h1:JAhxUVlolfYDErnwiqaLvUqc8nfb2r6S6slAgZOnaiE=
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
modernc.org/sqlite v1.45.0 h1:r51cSGzKpbptxnby+EIIz5fop4VuE4qFoVEjNvWoObs=
modernc.org/sqlite v1.45.0/go.mod h1:CzbrU2lSB1DKUusvwGz7rqEKIq+NUd8GWuBBZDs9/nA=
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM=

View File

@@ -0,0 +1,86 @@
package adminaudit
import (
"context"
"database/sql"
"time"
)
const DefaultListLimit = 200
type Record struct {
ID int64 `json:"id"`
Actor string `json:"actor"`
SourceIP string `json:"source_ip"`
Action string `json:"action"`
TargetType string `json:"target_type"`
TargetID string `json:"target_id"`
Success bool `json:"success"`
Message string `json:"message"`
CreatedAt int64 `json:"created_at"`
}
type Repository struct {
db *sql.DB
now func() time.Time
}
func NewRepository(db *sql.DB) Repository {
return Repository{
db: db,
now: time.Now,
}
}
func NewRepositoryWithClock(db *sql.DB, now func() time.Time) Repository {
repo := NewRepository(db)
if now != nil {
repo.now = now
}
return repo
}
func (r Repository) Record(ctx context.Context, actor, sourceIP, action, targetType, targetID string, success bool, message string) error {
if r.db == nil {
return nil
}
_, err := r.db.ExecContext(ctx, `
INSERT INTO audit_logs(actor, source_ip, action, target_type, target_id, success, message, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)`,
actor, sourceIP, action, targetType, targetID, boolToInt(success), message, r.now().Unix())
return err
}
func (r Repository) List(ctx context.Context, limit int) ([]Record, error) {
if limit <= 0 {
limit = DefaultListLimit
}
rows, err := r.db.QueryContext(ctx, `
SELECT id, actor, source_ip, action, target_type, target_id, success, message, created_at
FROM audit_logs
ORDER BY id DESC
LIMIT ?`, limit)
if err != nil {
return nil, err
}
defer rows.Close()
var logs []Record
for rows.Next() {
var item Record
var success int
if err := rows.Scan(&item.ID, &item.Actor, &item.SourceIP, &item.Action, &item.TargetType, &item.TargetID, &success, &item.Message, &item.CreatedAt); err != nil {
return nil, err
}
item.Success = success != 0
logs = append(logs, item)
}
return logs, rows.Err()
}
func boolToInt(value bool) int {
if value {
return 1
}
return 0
}

View File

@@ -0,0 +1,64 @@
package adminaudit
import (
"context"
"path/filepath"
"testing"
"time"
"github.com/tursom/mc-gateway/internal/admindb"
)
func TestRepositoryRecordAndList(t *testing.T) {
db, err := admindb.Open(filepath.Join(t.TempDir(), "gateway.sqlite3"))
if err != nil {
t.Fatalf("Open() error = %v", err)
}
defer db.Close()
if err := admindb.Migrate(db); err != nil {
t.Fatalf("Migrate() error = %v", err)
}
now := time.Unix(123, 0)
repo := NewRepositoryWithClock(db, func() time.Time { return now })
if err := repo.Record(context.Background(), "admin", "127.0.0.1", "route_upsert", "route", "play.example", true, "route saved"); err != nil {
t.Fatalf("Record(success) error = %v", err)
}
now = time.Unix(124, 0)
if err := repo.Record(context.Background(), "admin", "127.0.0.1", "route_delete", "route", "play.example", false, "route missing"); err != nil {
t.Fatalf("Record(failure) error = %v", err)
}
logs, err := repo.List(context.Background(), DefaultListLimit)
if err != nil {
t.Fatalf("List() error = %v", err)
}
if len(logs) != 2 {
t.Fatalf("List() len = %d, want 2", len(logs))
}
if logs[0].Action != "route_delete" || logs[0].Success {
t.Fatalf("logs[0] = %+v, want latest failure", logs[0])
}
if logs[0].CreatedAt != 124 {
t.Fatalf("logs[0].CreatedAt = %d, want 124", logs[0].CreatedAt)
}
if logs[1].Action != "route_upsert" || !logs[1].Success {
t.Fatalf("logs[1] = %+v, want earlier success", logs[1])
}
limited, err := repo.List(context.Background(), 1)
if err != nil {
t.Fatalf("List(limit) error = %v", err)
}
if len(limited) != 1 || limited[0].Action != "route_delete" {
t.Fatalf("List(limit) = %+v", limited)
}
}
func TestRepositoryRecordIgnoresNilDB(t *testing.T) {
repo := NewRepository(nil)
if err := repo.Record(context.Background(), "admin", "127.0.0.1", "login", "user", "admin", true, "ok"); err != nil {
t.Fatalf("Record(nil DB) error = %v", err)
}
}

View File

@@ -0,0 +1,142 @@
package adminconfig
import (
"errors"
"fmt"
"path"
"strconv"
"strings"
"time"
)
const (
DefaultDBPath = "mc-gateway.sqlite3"
DefaultTCPAdminPort = 25565
DefaultAdminPath = "/admin/"
DefaultAdminAPIPrefix = "/admin/api"
DefaultSessionTTL = 8 * time.Hour
EnvDB = "MC_GATEWAY_DB"
EnvTCPAdminPort = "MC_GATEWAY_TCP_ADMIN_PORT"
EnvPath = "MC_GATEWAY_ADMIN_PATH"
EnvAPIPrefix = "MC_GATEWAY_ADMIN_API_PREFIX"
EnvInitialPassword = "MC_GATEWAY_ADMIN_PASSWORD"
)
type Config struct {
DBPath string
TCPAdminPort int
AdminPath string
AdminAPIPrefix string
SessionTTL time.Duration
}
func Default() Config {
return Config{
DBPath: DefaultDBPath,
TCPAdminPort: DefaultTCPAdminPort,
AdminPath: DefaultAdminPath,
AdminAPIPrefix: DefaultAdminAPIPrefix,
SessionTTL: DefaultSessionTTL,
}
}
func Parse(getenv func(string) string) (Config, error) {
cfg := Default()
if dbPath := strings.TrimSpace(getenv(EnvDB)); dbPath != "" {
cfg.DBPath = dbPath
}
port, err := parseOptionalPort(getenv(EnvTCPAdminPort), DefaultTCPAdminPort, EnvTCPAdminPort)
if err != nil {
return Config{}, err
}
cfg.TCPAdminPort = port
adminPath, err := normalizeAdminPath(getenv(EnvPath))
if err != nil {
return Config{}, err
}
cfg.AdminPath = adminPath
apiPrefix, err := normalizeAdminAPIPrefix(getenv(EnvAPIPrefix))
if err != nil {
return Config{}, err
}
cfg.AdminAPIPrefix = apiPrefix
if err := validateAdminPaths(cfg.AdminPath, cfg.AdminAPIPrefix); err != nil {
return Config{}, err
}
return cfg, nil
}
func parseOptionalPort(value string, fallback int, name string) (int, error) {
value = strings.TrimSpace(value)
if value == "" {
return fallback, nil
}
port, err := strconv.Atoi(value)
if err != nil || port < 1 || port > 65535 {
return 0, fmt.Errorf("%s must be an integer from 1 to 65535", name)
}
return port, nil
}
func normalizeAdminPath(value string) (string, error) {
value = strings.TrimSpace(value)
if value == "" {
value = DefaultAdminPath
}
if !strings.HasPrefix(value, "/") {
return "", fmt.Errorf("%s must start with /", EnvPath)
}
cleaned := path.Clean(value)
if cleaned == "." {
cleaned = "/"
}
if !strings.HasPrefix(cleaned, "/") {
cleaned = "/" + cleaned
}
if !strings.HasSuffix(cleaned, "/") {
cleaned += "/"
}
return cleaned, nil
}
func normalizeAdminAPIPrefix(value string) (string, error) {
value = strings.TrimSpace(value)
if value == "" {
value = DefaultAdminAPIPrefix
}
if !strings.HasPrefix(value, "/") {
return "", fmt.Errorf("%s must start with /", EnvAPIPrefix)
}
cleaned := path.Clean(value)
if cleaned == "." || cleaned == "/" {
return "", fmt.Errorf("%s must not be /", EnvAPIPrefix)
}
if !strings.HasPrefix(cleaned, "/") {
cleaned = "/" + cleaned
}
return strings.TrimRight(cleaned, "/"), nil
}
func validateAdminPaths(adminPath, apiPrefix string) error {
adminRoot := strings.TrimRight(adminPath, "/")
if apiPrefix == adminRoot {
return errors.New("admin API prefix cannot equal admin page path")
}
for _, asset := range []string{"app.css", "app.js"} {
assetPath := strings.TrimRight(adminPath, "/") + "/" + asset
if apiPrefix == assetPath || strings.HasPrefix(apiPrefix+"/", assetPath+"/") {
return errors.New("admin API prefix cannot be under static asset path")
}
}
return nil
}

View File

@@ -0,0 +1,76 @@
package adminconfig
import "testing"
func TestParseDefaultsAndEnv(t *testing.T) {
cfg, err := Parse(func(string) string { return "" })
if err != nil {
t.Fatalf("Parse() error = %v", err)
}
if cfg.DBPath != DefaultDBPath {
t.Fatalf("DBPath = %q, want %q", cfg.DBPath, DefaultDBPath)
}
if cfg.TCPAdminPort != DefaultTCPAdminPort {
t.Fatalf("TCPAdminPort = %d, want %d", cfg.TCPAdminPort, DefaultTCPAdminPort)
}
if cfg.AdminPath != DefaultAdminPath {
t.Fatalf("AdminPath = %q, want %q", cfg.AdminPath, DefaultAdminPath)
}
if cfg.AdminAPIPrefix != DefaultAdminAPIPrefix {
t.Fatalf("AdminAPIPrefix = %q, want %q", cfg.AdminAPIPrefix, DefaultAdminAPIPrefix)
}
env := map[string]string{
EnvDB: "/tmp/mc.db",
EnvTCPAdminPort: "25575",
EnvPath: "/ops",
EnvAPIPrefix: "/ops/api/",
}
cfg, err = Parse(func(key string) string { return env[key] })
if err != nil {
t.Fatalf("Parse(env) error = %v", err)
}
if cfg.DBPath != "/tmp/mc.db" || cfg.TCPAdminPort != 25575 || cfg.AdminPath != "/ops/" || cfg.AdminAPIPrefix != "/ops/api" {
t.Fatalf("config = %+v", cfg)
}
}
func TestParseReturnsErrors(t *testing.T) {
tests := []struct {
name string
env map[string]string
}{
{
name: "invalid port",
env: map[string]string{EnvTCPAdminPort: "70000"},
},
{
name: "invalid admin path",
env: map[string]string{EnvPath: "admin"},
},
{
name: "invalid api prefix",
env: map[string]string{EnvAPIPrefix: "api"},
},
{
name: "api prefix equals admin path",
env: map[string]string{
EnvPath: "/admin",
EnvAPIPrefix: "/admin",
},
},
{
name: "api prefix under asset path",
env: map[string]string{EnvAPIPrefix: "/admin/app.js/api"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := Parse(func(key string) string { return tt.env[key] })
if err == nil {
t.Fatal("Parse() error = nil, want error")
}
})
}
}

91
internal/admindb/db.go Normal file
View File

@@ -0,0 +1,91 @@
package admindb
import (
"database/sql"
"os"
"path/filepath"
_ "modernc.org/sqlite"
)
func Open(dbPath string) (*sql.DB, error) {
if dir := filepath.Dir(dbPath); dir != "." && dir != "" {
if err := os.MkdirAll(dir, 0755); err != nil {
return nil, err
}
}
db, err := sql.Open("sqlite", dbPath)
if err != nil {
return nil, err
}
db.SetMaxOpenConns(1)
if _, err := db.Exec(`PRAGMA journal_mode=WAL`); err != nil {
db.Close()
return nil, err
}
if _, err := db.Exec(`PRAGMA busy_timeout=5000`); err != nil {
db.Close()
return nil, err
}
return db, nil
}
func Migrate(db *sql.DB) error {
const schema = `
CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
applied_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS users (
username TEXT PRIMARY KEY,
role TEXT NOT NULL,
password_hash TEXT NOT NULL,
disabled INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS routes (
host TEXT PRIMARY KEY,
upstream TEXT NOT NULL,
enabled INTEGER NOT NULL DEFAULT 1,
note TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
updated_by TEXT NOT NULL DEFAULT ''
);
CREATE TABLE IF NOT EXISTS services (
name TEXT PRIMARY KEY,
enabled INTEGER NOT NULL,
port INTEGER NOT NULL DEFAULT 0,
options_json TEXT NOT NULL DEFAULT '{}',
restart_required INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
updated_by TEXT NOT NULL DEFAULT ''
);
CREATE TABLE IF NOT EXISTS audit_logs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
actor TEXT NOT NULL,
source_ip TEXT NOT NULL,
action TEXT NOT NULL,
target_type TEXT NOT NULL,
target_id TEXT NOT NULL,
success INTEGER NOT NULL,
message TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_routes_enabled ON routes(enabled);
CREATE INDEX IF NOT EXISTS idx_audit_logs_created_at ON audit_logs(created_at);
INSERT OR IGNORE INTO schema_migrations(version, applied_at) VALUES (1, strftime('%s','now'));
`
_, err := db.Exec(schema)
return err
}

View File

@@ -0,0 +1,80 @@
package admindb
import (
"database/sql"
"path/filepath"
"testing"
)
func TestOpenCreatesDirectoryAndConfiguresSQLite(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "nested", "gateway.sqlite3")
db, err := Open(dbPath)
if err != nil {
t.Fatalf("Open() error = %v", err)
}
defer db.Close()
var journalMode string
if err := db.QueryRow(`PRAGMA journal_mode`).Scan(&journalMode); err != nil {
t.Fatalf("PRAGMA journal_mode error = %v", err)
}
if journalMode != "wal" {
t.Fatalf("journal_mode = %q, want wal", journalMode)
}
var busyTimeout int
if err := db.QueryRow(`PRAGMA busy_timeout`).Scan(&busyTimeout); err != nil {
t.Fatalf("PRAGMA busy_timeout error = %v", err)
}
if busyTimeout != 5000 {
t.Fatalf("busy_timeout = %d, want 5000", busyTimeout)
}
}
func TestMigrateCreatesSchema(t *testing.T) {
db, err := Open(filepath.Join(t.TempDir(), "gateway.sqlite3"))
if err != nil {
t.Fatalf("Open() error = %v", err)
}
defer db.Close()
if err := Migrate(db); err != nil {
t.Fatalf("Migrate() error = %v", err)
}
if err := Migrate(db); err != nil {
t.Fatalf("Migrate() second run error = %v", err)
}
for _, name := range []string{"schema_migrations", "users", "routes", "services", "audit_logs"} {
t.Run("table "+name, func(t *testing.T) {
if !sqliteObjectExists(t, db, "table", name) {
t.Fatalf("table %q was not created", name)
}
})
}
for _, name := range []string{"idx_routes_enabled", "idx_audit_logs_created_at"} {
t.Run("index "+name, func(t *testing.T) {
if !sqliteObjectExists(t, db, "index", name) {
t.Fatalf("index %q was not created", name)
}
})
}
var version int
if err := db.QueryRow(`SELECT version FROM schema_migrations WHERE version = 1`).Scan(&version); err != nil {
t.Fatalf("schema migration version query error = %v", err)
}
if version != 1 {
t.Fatalf("schema migration version = %d, want 1", version)
}
}
func sqliteObjectExists(t *testing.T, db *sql.DB, objectType, name string) bool {
t.Helper()
var count int
if err := db.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE type = ? AND name = ?`, objectType, name).Scan(&count); err != nil {
t.Fatalf("sqlite_master query error = %v", err)
}
return count == 1
}

92
internal/adminhttp/api.go Normal file
View File

@@ -0,0 +1,92 @@
package adminhttp
import (
"net/http"
"strings"
)
type SegmentHandlerFunc func(http.ResponseWriter, *http.Request, string)
type APIHandlers struct {
SetupStatus http.HandlerFunc
Setup http.HandlerFunc
Login http.HandlerFunc
Logout http.HandlerFunc
Me http.HandlerFunc
Status http.HandlerFunc
RoutesList http.HandlerFunc
RouteItem SegmentHandlerFunc
ServicesList http.HandlerFunc
ServiceItem SegmentHandlerFunc
Metrics http.HandlerFunc
UsersList http.HandlerFunc
UsersCreate http.HandlerFunc
UserItem SegmentHandlerFunc
AuditLogs http.HandlerFunc
}
func NewAPIHandler(prefix string, handlers APIHandlers) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Cache-Control", "no-store")
path := strings.TrimPrefix(r.URL.Path, prefix)
if path == "" {
path = "/"
}
switch {
case path == "/setup" && r.Method == http.MethodGet:
callHandler(w, r, handlers.SetupStatus)
case path == "/setup" && r.Method == http.MethodPost:
callHandler(w, r, handlers.Setup)
case path == "/auth/login" && r.Method == http.MethodPost:
callHandler(w, r, handlers.Login)
case path == "/auth/logout" && r.Method == http.MethodPost:
callHandler(w, r, handlers.Logout)
case path == "/me" && r.Method == http.MethodGet:
callHandler(w, r, handlers.Me)
case path == "/status" && r.Method == http.MethodGet:
callHandler(w, r, handlers.Status)
case path == "/routes" && r.Method == http.MethodGet:
callHandler(w, r, handlers.RoutesList)
case strings.HasPrefix(path, "/routes/"):
callSegmentHandler(w, r, handlers.RouteItem, strings.TrimPrefix(path, "/routes/"))
case path == "/services" && r.Method == http.MethodGet:
callHandler(w, r, handlers.ServicesList)
case strings.HasPrefix(path, "/services/"):
callSegmentHandler(w, r, handlers.ServiceItem, strings.TrimPrefix(path, "/services/"))
case path == "/metrics" && r.Method == http.MethodGet:
callHandler(w, r, handlers.Metrics)
case path == "/users" && r.Method == http.MethodGet:
callHandler(w, r, handlers.UsersList)
case path == "/users" && r.Method == http.MethodPost:
callHandler(w, r, handlers.UsersCreate)
case strings.HasPrefix(path, "/users/"):
callSegmentHandler(w, r, handlers.UserItem, strings.TrimPrefix(path, "/users/"))
case path == "/audit-logs" && r.Method == http.MethodGet:
callHandler(w, r, handlers.AuditLogs)
default:
WriteAPIError(w, http.StatusNotFound, "not found")
}
}
}
func callHandler(w http.ResponseWriter, r *http.Request, handler http.HandlerFunc) {
if handler == nil {
WriteAPIError(w, http.StatusNotFound, "not found")
return
}
handler(w, r)
}
func callSegmentHandler(w http.ResponseWriter, r *http.Request, handler SegmentHandlerFunc, segment string) {
if handler == nil {
WriteAPIError(w, http.StatusNotFound, "not found")
return
}
handler(w, r, segment)
}

View File

@@ -0,0 +1,114 @@
package adminhttp
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestNewAPIHandlerRoutesRequests(t *testing.T) {
tests := []struct {
name string
method string
path string
wantCall string
wantSegment string
}{
{name: "setup status", method: http.MethodGet, path: "/admin/api/setup", wantCall: "setup_status"},
{name: "setup", method: http.MethodPost, path: "/admin/api/setup", wantCall: "setup"},
{name: "login", method: http.MethodPost, path: "/admin/api/auth/login", wantCall: "login"},
{name: "logout", method: http.MethodPost, path: "/admin/api/auth/logout", wantCall: "logout"},
{name: "me", method: http.MethodGet, path: "/admin/api/me", wantCall: "me"},
{name: "status", method: http.MethodGet, path: "/admin/api/status", wantCall: "status"},
{name: "routes list", method: http.MethodGet, path: "/admin/api/routes", wantCall: "routes_list"},
{name: "route item", method: http.MethodPut, path: "/admin/api/routes/play.example", wantCall: "route_item", wantSegment: "play.example"},
{name: "services list", method: http.MethodGet, path: "/admin/api/services", wantCall: "services_list"},
{name: "service restart", method: http.MethodPost, path: "/admin/api/services/kcp/restart", wantCall: "service_item", wantSegment: "kcp/restart"},
{name: "metrics", method: http.MethodGet, path: "/admin/api/metrics", wantCall: "metrics"},
{name: "users list", method: http.MethodGet, path: "/admin/api/users", wantCall: "users_list"},
{name: "users create", method: http.MethodPost, path: "/admin/api/users", wantCall: "users_create"},
{name: "user item", method: http.MethodPatch, path: "/admin/api/users/member", wantCall: "user_item", wantSegment: "member"},
{name: "audit logs", method: http.MethodGet, path: "/admin/api/audit-logs", wantCall: "audit_logs"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var gotCall, gotSegment string
handler := NewAPIHandler("/admin/api", APIHandlers{
SetupStatus: recordCall(&gotCall, "setup_status"),
Setup: recordCall(&gotCall, "setup"),
Login: recordCall(&gotCall, "login"),
Logout: recordCall(&gotCall, "logout"),
Me: recordCall(&gotCall, "me"),
Status: recordCall(&gotCall, "status"),
RoutesList: recordCall(&gotCall, "routes_list"),
RouteItem: recordSegmentCall(&gotCall, &gotSegment, "route_item"),
ServicesList: recordCall(&gotCall, "services_list"),
ServiceItem: recordSegmentCall(&gotCall, &gotSegment, "service_item"),
Metrics: recordCall(&gotCall, "metrics"),
UsersList: recordCall(&gotCall, "users_list"),
UsersCreate: recordCall(&gotCall, "users_create"),
UserItem: recordSegmentCall(&gotCall, &gotSegment, "user_item"),
AuditLogs: recordCall(&gotCall, "audit_logs"),
})
resp := httptest.NewRecorder()
req := httptest.NewRequest(tt.method, tt.path, nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusNoContent {
t.Fatalf("status = %d, want %d; body=%s", resp.Code, http.StatusNoContent, resp.Body.String())
}
if gotCall != tt.wantCall {
t.Fatalf("call = %q, want %q", gotCall, tt.wantCall)
}
if gotSegment != tt.wantSegment {
t.Fatalf("segment = %q, want %q", gotSegment, tt.wantSegment)
}
if got := resp.Header().Get("Cache-Control"); got != "no-store" {
t.Fatalf("Cache-Control = %q, want no-store", got)
}
})
}
}
func TestNewAPIHandlerWritesNotFound(t *testing.T) {
handler := NewAPIHandler("/admin/api", APIHandlers{})
resp := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/admin/api/auth/login", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusNotFound {
t.Fatalf("wrong method status = %d, want %d", resp.Code, http.StatusNotFound)
}
if strings.TrimSpace(resp.Body.String()) != `{"error":"not found"}` {
t.Fatalf("wrong method body = %q", resp.Body.String())
}
resp = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/admin/api/missing", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusNotFound {
t.Fatalf("missing status = %d, want %d", resp.Code, http.StatusNotFound)
}
}
func recordCall(got *string, call string) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
*got = call
w.WriteHeader(http.StatusNoContent)
}
}
func recordSegmentCall(gotCall, gotSegment *string, call string) SegmentHandlerFunc {
return func(w http.ResponseWriter, r *http.Request, segment string) {
*gotCall = call
*gotSegment = segment
w.WriteHeader(http.StatusNoContent)
}
}

View File

@@ -0,0 +1,107 @@
package adminhttp
import (
"io/fs"
"net/http"
"strings"
)
const (
staticIndexFile = "admin_static/index.html"
staticCSSFile = "admin_static/app.css"
staticJSFile = "admin_static/app.js"
)
type GatewayHandlerOptions struct {
AdminPath string
AdminAPIPrefix string
Assets fs.FS
APIHandler http.HandlerFunc
WebSocketEnabled bool
WebSocketPath string
WebSocketHandler http.HandlerFunc
}
func NewGatewayHandler(opts GatewayHandlerOptions) http.Handler {
mux := http.NewServeMux()
registerAdminHandlers(mux, opts)
if opts.WebSocketEnabled &&
opts.WebSocketHandler != nil &&
!WebSocketPathConflictsWithAdmin(opts.WebSocketPath, opts.AdminPath, opts.AdminAPIPrefix) {
mux.HandleFunc(opts.WebSocketPath, opts.WebSocketHandler)
}
return mux
}
func WebSocketPathConflictsWithAdmin(webSocketPath, adminPath, adminAPIPrefix string) bool {
adminRoot := strings.TrimRight(adminPath, "/")
if webSocketPath == adminPath || webSocketPath == adminRoot {
return true
}
if webSocketPath == adminAPIPrefix || strings.HasPrefix(adminAPIPrefix+"/", webSocketPath+"/") {
return webSocketPath != "/"
}
return false
}
func registerAdminHandlers(mux *http.ServeMux, opts GatewayHandlerOptions) {
adminRoot := strings.TrimRight(opts.AdminPath, "/")
mux.HandleFunc(opts.AdminAPIPrefix, opts.APIHandler)
mux.HandleFunc(opts.AdminAPIPrefix+"/", opts.APIHandler)
if adminRoot != "" {
mux.HandleFunc(adminRoot, func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, opts.AdminPath, http.StatusMovedPermanently)
})
}
mux.HandleFunc(opts.AdminPath, func(w http.ResponseWriter, r *http.Request) {
serveAdminStatic(w, r, opts)
})
}
func serveAdminStatic(w http.ResponseWriter, r *http.Request, opts GatewayHandlerOptions) {
if r.Method != http.MethodGet && r.Method != http.MethodHead {
WriteAPIError(w, http.StatusMethodNotAllowed, "method not allowed")
return
}
if r.URL.Path == opts.AdminPath {
serveAdminIndex(w, opts)
return
}
rel := strings.TrimPrefix(r.URL.Path, opts.AdminPath)
switch rel {
case "app.css":
serveAdminFile(w, r, opts.Assets, staticCSSFile, "text/css; charset=utf-8")
case "app.js":
serveAdminFile(w, r, opts.Assets, staticJSFile, "application/javascript; charset=utf-8")
default:
http.NotFound(w, r)
}
}
func serveAdminIndex(w http.ResponseWriter, opts GatewayHandlerOptions) {
data, err := fs.ReadFile(opts.Assets, staticIndexFile)
if err != nil {
WriteAPIError(w, http.StatusInternalServerError, err.Error())
return
}
html := strings.ReplaceAll(string(data), "__ADMIN_API_PREFIX__", opts.AdminAPIPrefix)
w.Header().Set("Content-Type", "text/html; charset=utf-8")
w.Header().Set("Cache-Control", "no-store")
_, _ = w.Write([]byte(html))
}
func serveAdminFile(w http.ResponseWriter, r *http.Request, assets fs.FS, name, contentType string) {
data, err := fs.ReadFile(assets, name)
if err != nil {
http.NotFound(w, r)
return
}
w.Header().Set("Content-Type", contentType)
w.Header().Set("Cache-Control", "no-store")
_, _ = w.Write(data)
}

View File

@@ -0,0 +1,58 @@
package adminhttp
import (
"encoding/json"
"errors"
"net"
"net/http"
"net/url"
"strings"
)
func DecodeJSONRequest(w http.ResponseWriter, r *http.Request, dst any) bool {
defer r.Body.Close()
decoder := json.NewDecoder(r.Body)
decoder.UseNumber()
if err := decoder.Decode(dst); err != nil {
WriteAPIError(w, http.StatusBadRequest, "invalid JSON: "+err.Error())
return false
}
return true
}
func WriteJSON(w http.ResponseWriter, status int, value any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(value)
}
func WriteAPIError(w http.ResponseWriter, status int, message string) {
WriteJSON(w, status, map[string]any{
"error": message,
})
}
func PathSegment(raw string) (string, error) {
if raw == "" || strings.Contains(raw, "/") {
return "", errors.New("invalid path segment")
}
value, err := url.PathUnescape(raw)
if err != nil {
return "", err
}
if value == "" || strings.Contains(value, "/") {
return "", errors.New("invalid path segment")
}
return value, nil
}
func RequestSourceIP(r *http.Request) string {
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err == nil {
return host
}
if r.RemoteAddr != "" {
return r.RemoteAddr
}
return "unknown"
}

View File

@@ -0,0 +1,230 @@
package adminhttp
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"testing/fstest"
)
func TestDecodeJSONRequest(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/api", strings.NewReader(`{"port":25565}`))
resp := httptest.NewRecorder()
var body map[string]any
if !DecodeJSONRequest(resp, req, &body) {
t.Fatal("DecodeJSONRequest() = false, want true")
}
port, ok := body["port"].(json.Number)
if !ok {
t.Fatalf("port = %#v, want json.Number", body["port"])
}
if port.String() != "25565" {
t.Fatalf("port = %q, want 25565", port.String())
}
}
func TestDecodeJSONRequestWritesError(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/api", strings.NewReader(`{`))
resp := httptest.NewRecorder()
var body map[string]any
if DecodeJSONRequest(resp, req, &body) {
t.Fatal("DecodeJSONRequest(invalid) = true, want false")
}
if resp.Code != http.StatusBadRequest {
t.Fatalf("status = %d, want %d", resp.Code, http.StatusBadRequest)
}
var errorBody map[string]string
if err := json.Unmarshal(resp.Body.Bytes(), &errorBody); err != nil {
t.Fatalf("Unmarshal(%q) error = %v", resp.Body.String(), err)
}
if !strings.HasPrefix(errorBody["error"], "invalid JSON: ") {
t.Fatalf("error = %q, want invalid JSON prefix", errorBody["error"])
}
}
func TestWriteJSONAndAPIError(t *testing.T) {
resp := httptest.NewRecorder()
WriteJSON(resp, http.StatusCreated, map[string]any{"ok": true})
if resp.Code != http.StatusCreated {
t.Fatalf("status = %d, want %d", resp.Code, http.StatusCreated)
}
if got := resp.Header().Get("Content-Type"); got != "application/json; charset=utf-8" {
t.Fatalf("Content-Type = %q", got)
}
if strings.TrimSpace(resp.Body.String()) != `{"ok":true}` {
t.Fatalf("body = %q", resp.Body.String())
}
resp = httptest.NewRecorder()
WriteAPIError(resp, http.StatusForbidden, "permission denied")
if resp.Code != http.StatusForbidden {
t.Fatalf("error status = %d, want %d", resp.Code, http.StatusForbidden)
}
if strings.TrimSpace(resp.Body.String()) != `{"error":"permission denied"}` {
t.Fatalf("error body = %q", resp.Body.String())
}
}
func TestPathSegment(t *testing.T) {
tests := map[string]string{
"play.example": "play.example",
"play%2Etest": "play.test",
}
for raw, want := range tests {
t.Run(raw, func(t *testing.T) {
got, err := PathSegment(raw)
if err != nil {
t.Fatalf("PathSegment(%q) error = %v", raw, err)
}
if got != want {
t.Fatalf("PathSegment(%q) = %q, want %q", raw, got, want)
}
})
}
for _, raw := range []string{"", "a/b", "%2F", "%zz"} {
t.Run("invalid "+raw, func(t *testing.T) {
if _, err := PathSegment(raw); err == nil {
t.Fatalf("PathSegment(%q) error = nil, want error", raw)
}
})
}
}
func TestRequestSourceIP(t *testing.T) {
tests := []struct {
remoteAddr string
want string
}{
{"127.0.0.1:1234", "127.0.0.1"},
{"[::1]:1234", "::1"},
{"unix-socket", "unix-socket"},
{"", "unknown"},
}
for _, tt := range tests {
t.Run(tt.remoteAddr, func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = tt.remoteAddr
if got := RequestSourceIP(req); got != tt.want {
t.Fatalf("RequestSourceIP(%q) = %q, want %q", tt.remoteAddr, got, tt.want)
}
})
}
}
func TestNewGatewayHandlerServesAdminAndAPI(t *testing.T) {
assets := fstest.MapFS{
staticIndexFile: {Data: []byte(`<html data-api-prefix="__ADMIN_API_PREFIX__"></html>`)},
staticCSSFile: {Data: []byte(`body{color:red}`)},
staticJSFile: {Data: []byte(`console.log("admin")`)},
}
apiCalled := false
handler := NewGatewayHandler(GatewayHandlerOptions{
AdminPath: "/ops/",
AdminAPIPrefix: "/ops/api",
Assets: assets,
APIHandler: func(w http.ResponseWriter, r *http.Request) {
apiCalled = true
w.WriteHeader(http.StatusNoContent)
},
})
resp := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/ops", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusMovedPermanently {
t.Fatalf("redirect status = %d, want %d", resp.Code, http.StatusMovedPermanently)
}
if got := resp.Header().Get("Location"); got != "/ops/" {
t.Fatalf("redirect location = %q, want /ops/", got)
}
resp = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/ops/", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("admin page status = %d, body=%s", resp.Code, resp.Body.String())
}
if !strings.Contains(resp.Body.String(), `data-api-prefix="/ops/api"`) {
t.Fatalf("admin page = %q, want API prefix", resp.Body.String())
}
if got := resp.Header().Get("Cache-Control"); got != "no-store" {
t.Fatalf("admin page Cache-Control = %q, want no-store", got)
}
resp = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/ops/app.css", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusOK || strings.TrimSpace(resp.Body.String()) != `body{color:red}` {
t.Fatalf("css response status=%d body=%q", resp.Code, resp.Body.String())
}
resp = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodPost, "/ops/", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusMethodNotAllowed {
t.Fatalf("POST admin status = %d, want %d", resp.Code, http.StatusMethodNotAllowed)
}
resp = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/ops/api/setup", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusNoContent {
t.Fatalf("api status = %d, want %d", resp.Code, http.StatusNoContent)
}
if !apiCalled {
t.Fatal("API handler was not called")
}
}
func TestNewGatewayHandlerRegistersWebSocketWhenPathDoesNotConflict(t *testing.T) {
assets := fstest.MapFS{
staticIndexFile: {Data: []byte(``)},
}
handler := NewGatewayHandler(GatewayHandlerOptions{
AdminPath: "/admin/",
AdminAPIPrefix: "/admin/api",
Assets: assets,
APIHandler: func(w http.ResponseWriter, r *http.Request) {},
WebSocketEnabled: true,
WebSocketPath: "/ws",
WebSocketHandler: func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusAccepted)
},
})
resp := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/ws", nil)
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusAccepted {
t.Fatalf("websocket path status = %d, want %d", resp.Code, http.StatusAccepted)
}
}
func TestWebSocketPathConflictsWithAdmin(t *testing.T) {
tests := []struct {
name string
webSocketPath string
want bool
}{
{name: "admin path", webSocketPath: "/admin/", want: true},
{name: "admin root", webSocketPath: "/admin", want: true},
{name: "api prefix", webSocketPath: "/admin/api", want: true},
{name: "api child", webSocketPath: "/admin/api/ws", want: false},
{name: "root", webSocketPath: "/", want: false},
{name: "separate path", webSocketPath: "/ws", want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := WebSocketPathConflictsWithAdmin(tt.webSocketPath, "/admin/", "/admin/api")
if got != tt.want {
t.Fatalf("WebSocketPathConflictsWithAdmin(%q) = %v, want %v", tt.webSocketPath, got, tt.want)
}
})
}
}

View File

@@ -0,0 +1,36 @@
package adminhttp
type LoginRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
type SetupRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
type RouteRequest struct {
Upstream string `json:"upstream"`
Enabled *bool `json:"enabled"`
Note string `json:"note"`
}
type ServiceRequest struct {
Enabled *bool `json:"enabled"`
Port int `json:"port"`
Options map[string]any `json:"options"`
}
type CreateUserRequest struct {
Username string `json:"username"`
Role string `json:"role"`
Password string `json:"password"`
Disabled bool `json:"disabled"`
}
type PatchUserRequest struct {
Role *string `json:"role"`
Password *string `json:"password"`
Disabled *bool `json:"disabled"`
}

View File

@@ -0,0 +1,55 @@
package adminhttp
import (
"encoding/json"
"testing"
)
func TestRouteRequestKeepsOptionalEnabled(t *testing.T) {
var req RouteRequest
if err := json.Unmarshal([]byte(`{"upstream":"127.0.0.1:25565","note":"primary"}`), &req); err != nil {
t.Fatalf("Unmarshal() error = %v", err)
}
if req.Enabled != nil {
t.Fatalf("Enabled = %v, want nil when omitted", *req.Enabled)
}
if err := json.Unmarshal([]byte(`{"enabled":false}`), &req); err != nil {
t.Fatalf("Unmarshal(enabled) error = %v", err)
}
if req.Enabled == nil || *req.Enabled {
t.Fatalf("Enabled = %v, want false pointer", req.Enabled)
}
}
func TestPatchUserRequestKeepsOmittedFieldsNil(t *testing.T) {
var req PatchUserRequest
if err := json.Unmarshal([]byte(`{"role":"guest"}`), &req); err != nil {
t.Fatalf("Unmarshal() error = %v", err)
}
if req.Role == nil || *req.Role != "guest" {
t.Fatalf("Role = %v, want guest pointer", req.Role)
}
if req.Password != nil {
t.Fatal("Password pointer should be nil when omitted")
}
if req.Disabled != nil {
t.Fatal("Disabled pointer should be nil when omitted")
}
}
func TestServiceRequestDecodesOptions(t *testing.T) {
var req ServiceRequest
if err := json.Unmarshal([]byte(`{"enabled":true,"port":25570,"options":{"path":"/ws"}}`), &req); err != nil {
t.Fatalf("Unmarshal() error = %v", err)
}
if req.Enabled == nil || !*req.Enabled {
t.Fatalf("Enabled = %v, want true pointer", req.Enabled)
}
if req.Port != 25570 {
t.Fatalf("Port = %d, want 25570", req.Port)
}
if req.Options["path"] != "/ws" {
t.Fatalf("path option = %#v, want /ws", req.Options["path"])
}
}

View File

@@ -0,0 +1,141 @@
package adminroute
import (
"context"
"database/sql"
"strings"
"time"
)
type Record struct {
Host string `json:"host"`
Upstream string `json:"upstream"`
Enabled bool `json:"enabled"`
Note string `json:"note"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
UpdatedBy string `json:"updated_by"`
}
type Repository struct {
db *sql.DB
now func() time.Time
}
func NewRepository(db *sql.DB) Repository {
return Repository{
db: db,
now: time.Now,
}
}
func NewRepositoryWithClock(db *sql.DB, now func() time.Time) Repository {
repo := NewRepository(db)
if now != nil {
repo.now = now
}
return repo
}
func (r Repository) EnabledMap(ctx context.Context) (map[string]string, error) {
rows, err := r.db.QueryContext(ctx, `SELECT host, upstream FROM routes WHERE enabled = 1`)
if err != nil {
return nil, err
}
defer rows.Close()
routes := make(map[string]string)
for rows.Next() {
var host, upstream string
if err := rows.Scan(&host, &upstream); err != nil {
return nil, err
}
routes[host] = upstream
}
return routes, rows.Err()
}
func (r Repository) List(ctx context.Context, query string) ([]Record, error) {
sqlQuery := `
SELECT host, upstream, enabled, note, created_at, updated_at, updated_by
FROM routes`
var args []any
if query = strings.TrimSpace(query); query != "" {
sqlQuery += ` WHERE host LIKE ? OR upstream LIKE ? OR note LIKE ?`
like := "%" + query + "%"
args = append(args, like, like, like)
}
sqlQuery += ` ORDER BY host`
rows, err := r.db.QueryContext(ctx, sqlQuery, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var routes []Record
for rows.Next() {
var route Record
var enabled int
if err := rows.Scan(&route.Host, &route.Upstream, &enabled, &route.Note, &route.CreatedAt, &route.UpdatedAt, &route.UpdatedBy); err != nil {
return nil, err
}
route.Enabled = enabled != 0
routes = append(routes, route)
}
return routes, rows.Err()
}
func (r Repository) Upsert(ctx context.Context, actor, host, upstream string, enabled bool, note string) error {
if err := ValidateHost(host); err != nil {
return err
}
if err := ValidateUpstream(upstream); err != nil {
return err
}
now := r.now().Unix()
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if _, err := tx.ExecContext(ctx, `
INSERT INTO routes(host, upstream, enabled, note, created_at, updated_at, updated_by)
VALUES (?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(host) DO UPDATE SET
upstream = excluded.upstream,
enabled = excluded.enabled,
note = excluded.note,
updated_at = excluded.updated_at,
updated_by = excluded.updated_by`,
host, upstream, boolToInt(enabled), note, now, now, actor); err != nil {
return err
}
return tx.Commit()
}
func (r Repository) Delete(ctx context.Context, host string) error {
if err := ValidateHost(host); err != nil {
return err
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if _, err := tx.ExecContext(ctx, `DELETE FROM routes WHERE host = ?`, host); err != nil {
return err
}
return tx.Commit()
}
func boolToInt(value bool) int {
if value {
return 1
}
return 0
}

View File

@@ -0,0 +1,127 @@
package adminroute
import (
"context"
"path/filepath"
"reflect"
"testing"
"time"
"github.com/tursom/mc-gateway/internal/admindb"
)
func TestRepositoryUpsertListEnabledMapAndDelete(t *testing.T) {
repo, closeDB := newRouteTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.Upsert(ctx, "admin", "play.example", "127.0.0.1:25565", true, "primary"); err != nil {
t.Fatalf("Upsert(play) error = %v", err)
}
if err := repo.Upsert(ctx, "admin", "dev.example", "127.0.0.1:25566", false, "disabled"); err != nil {
t.Fatalf("Upsert(dev) error = %v", err)
}
routes, err := repo.List(ctx, "")
if err != nil {
t.Fatalf("List() error = %v", err)
}
if len(routes) != 2 {
t.Fatalf("List() len = %d, want 2", len(routes))
}
if routes[0].Host != "dev.example" || routes[0].Enabled {
t.Fatalf("routes[0] = %+v, want disabled dev route", routes[0])
}
if routes[1].Host != "play.example" || !routes[1].Enabled || routes[1].CreatedAt != 100 || routes[1].UpdatedAt != 100 || routes[1].UpdatedBy != "admin" {
t.Fatalf("routes[1] = %+v, want enabled play route with timestamps", routes[1])
}
enabled, err := repo.EnabledMap(ctx)
if err != nil {
t.Fatalf("EnabledMap() error = %v", err)
}
if want := map[string]string{"play.example": "127.0.0.1:25565"}; !reflect.DeepEqual(enabled, want) {
t.Fatalf("EnabledMap() = %#v, want %#v", enabled, want)
}
filtered, err := repo.List(ctx, "primary")
if err != nil {
t.Fatalf("List(query) error = %v", err)
}
if len(filtered) != 1 || filtered[0].Host != "play.example" {
t.Fatalf("List(query) = %+v", filtered)
}
if err := repo.Delete(ctx, "play.example"); err != nil {
t.Fatalf("Delete() error = %v", err)
}
enabled, err = repo.EnabledMap(ctx)
if err != nil {
t.Fatalf("EnabledMap(after delete) error = %v", err)
}
if len(enabled) != 0 {
t.Fatalf("EnabledMap(after delete) = %#v, want empty", enabled)
}
}
func TestRepositoryUpsertUpdatesExistingRoute(t *testing.T) {
now := time.Unix(100, 0)
repo, closeDB := newRouteTestRepository(t, now)
defer closeDB()
ctx := context.Background()
if err := repo.Upsert(ctx, "admin", "play.example", "127.0.0.1:25565", true, "primary"); err != nil {
t.Fatalf("Upsert(create) error = %v", err)
}
repo.now = func() time.Time { return time.Unix(200, 0) }
if err := repo.Upsert(ctx, "member", "play.example", "127.0.0.1:25566", false, "updated"); err != nil {
t.Fatalf("Upsert(update) error = %v", err)
}
routes, err := repo.List(ctx, "")
if err != nil {
t.Fatalf("List() error = %v", err)
}
if len(routes) != 1 {
t.Fatalf("List() len = %d, want 1", len(routes))
}
route := routes[0]
if route.CreatedAt != 100 || route.UpdatedAt != 200 || route.UpdatedBy != "member" || route.Enabled {
t.Fatalf("updated route = %+v", route)
}
if route.Upstream != "127.0.0.1:25566" || route.Note != "updated" {
t.Fatalf("updated route = %+v", route)
}
}
func TestRepositoryValidationErrors(t *testing.T) {
repo, closeDB := newRouteTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.Upsert(ctx, "admin", "bad host", "127.0.0.1:25565", true, ""); err == nil {
t.Fatal("Upsert(invalid host) error = nil, want error")
}
if err := repo.Upsert(ctx, "admin", "play.example", "127.0.0.1", true, ""); err == nil {
t.Fatal("Upsert(invalid upstream) error = nil, want error")
}
if err := repo.Delete(ctx, "bad/host"); err == nil {
t.Fatal("Delete(invalid host) error = nil, want error")
}
}
func newRouteTestRepository(t *testing.T, now time.Time) (Repository, func()) {
t.Helper()
db, err := admindb.Open(filepath.Join(t.TempDir(), "gateway.sqlite3"))
if err != nil {
t.Fatalf("Open() error = %v", err)
}
if err := admindb.Migrate(db); err != nil {
db.Close()
t.Fatalf("Migrate() error = %v", err)
}
return NewRepositoryWithClock(db, func() time.Time { return now }), func() {
_ = db.Close()
}
}

View File

@@ -0,0 +1,55 @@
package adminroute
import "sync/atomic"
type Snapshot struct {
value atomic.Value
}
func NewSnapshot() *Snapshot {
return &Snapshot{}
}
func (s *Snapshot) Store(routes map[string]string) {
if routes == nil {
routes = map[string]string{}
}
copied := make(map[string]string, len(routes))
for host, upstream := range routes {
copied[host] = upstream
}
s.value.Store(copied)
}
func (s *Snapshot) Clone() map[string]string {
value := s.value.Load()
if value == nil {
return nil
}
routes, ok := value.(map[string]string)
if !ok {
return nil
}
copied := make(map[string]string, len(routes))
for host, upstream := range routes {
copied[host] = upstream
}
return copied
}
func (s *Snapshot) Lookup(host string) (string, bool) {
value := s.value.Load()
if value == nil {
return "", false
}
routes, ok := value.(map[string]string)
if !ok {
return "", false
}
upstream, ok := routes[host]
if ok {
return upstream, true
}
upstream, ok = routes["default"]
return upstream, ok
}

View File

@@ -0,0 +1,48 @@
package adminroute
import "testing"
func TestSnapshotLookup(t *testing.T) {
snapshot := NewSnapshot()
if upstream, ok := snapshot.Lookup("play.example"); ok || upstream != "" {
t.Fatalf("Lookup(empty) = %q, %v; want miss", upstream, ok)
}
snapshot.Store(map[string]string{
"play.example": "127.0.0.1:25565",
"default": "127.0.0.1:25566",
})
if upstream, ok := snapshot.Lookup("play.example"); !ok || upstream != "127.0.0.1:25565" {
t.Fatalf("Lookup(play.example) = %q, %v", upstream, ok)
}
if upstream, ok := snapshot.Lookup("unknown.example"); !ok || upstream != "127.0.0.1:25566" {
t.Fatalf("Lookup(default fallback) = %q, %v", upstream, ok)
}
}
func TestSnapshotStoreCopiesRoutes(t *testing.T) {
snapshot := NewSnapshot()
routes := map[string]string{"play.example": "127.0.0.1:25565"}
snapshot.Store(routes)
routes["play.example"] = "changed"
if upstream, ok := snapshot.Lookup("play.example"); !ok || upstream != "127.0.0.1:25565" {
t.Fatalf("Lookup(after source mutation) = %q, %v", upstream, ok)
}
}
func TestSnapshotCloneCopiesRoutes(t *testing.T) {
snapshot := NewSnapshot()
if clone := snapshot.Clone(); clone != nil {
t.Fatalf("Clone(empty) = %#v, want nil", clone)
}
snapshot.Store(map[string]string{"play.example": "127.0.0.1:25565"})
clone := snapshot.Clone()
clone["play.example"] = "changed"
if upstream, ok := snapshot.Lookup("play.example"); !ok || upstream != "127.0.0.1:25565" {
t.Fatalf("Lookup(after clone mutation) = %q, %v", upstream, ok)
}
}

View File

@@ -0,0 +1,45 @@
package adminroute
import (
"errors"
"net"
"strconv"
"strings"
)
func ValidateHost(host string) error {
host = strings.TrimSpace(host)
if host == "" {
return errors.New("host is required")
}
if strings.ContainsAny(host, " \t\r\n") {
return errors.New("host must not contain whitespace")
}
if strings.Contains(host, "/") {
return errors.New("host must not contain /")
}
return nil
}
func ValidateUpstream(upstream string) error {
upstream = strings.TrimSpace(upstream)
if upstream == "" {
return errors.New("upstream is required")
}
for _, prefix := range []string{"kcp://", "quic://", "haproxy://"} {
upstream = strings.TrimPrefix(upstream, prefix)
}
host, portValue, err := net.SplitHostPort(upstream)
if err != nil {
return errors.New("upstream must be host:port")
}
if strings.TrimSpace(host) == "" {
return errors.New("upstream host is required")
}
port, err := strconv.Atoi(portValue)
if err != nil || port < 1 || port > 65535 {
return errors.New("upstream port must be an integer from 1 to 65535")
}
return nil
}

View File

@@ -0,0 +1,49 @@
package adminroute
import "testing"
func TestValidateHost(t *testing.T) {
valid := []string{"default", "play.example", "dev.lan"}
for _, host := range valid {
t.Run(host, func(t *testing.T) {
if err := ValidateHost(host); err != nil {
t.Fatalf("ValidateHost(%q) error = %v", host, err)
}
})
}
invalid := []string{"", " ", "play example", "play/example"}
for _, host := range invalid {
t.Run("invalid "+host, func(t *testing.T) {
if err := ValidateHost(host); err == nil {
t.Fatalf("ValidateHost(%q) error = nil, want error", host)
}
})
}
}
func TestValidateUpstream(t *testing.T) {
valid := []string{
"127.0.0.1:25565",
"localhost:25565",
"kcp://127.0.0.1:25565",
"quic://127.0.0.1:25565",
"haproxy://127.0.0.1:25565",
}
for _, upstream := range valid {
t.Run(upstream, func(t *testing.T) {
if err := ValidateUpstream(upstream); err != nil {
t.Fatalf("ValidateUpstream(%q) error = %v", upstream, err)
}
})
}
invalid := []string{"", "127.0.0.1", ":25565", "127.0.0.1:0", "127.0.0.1:70000", "127.0.0.1:not-a-port"}
for _, upstream := range invalid {
t.Run("invalid "+upstream, func(t *testing.T) {
if err := ValidateUpstream(upstream); err == nil {
t.Fatalf("ValidateUpstream(%q) error = nil, want error", upstream)
}
})
}
}

View File

@@ -0,0 +1,104 @@
package adminservice
import (
"context"
"database/sql"
"encoding/json"
"time"
)
type Repository struct {
db *sql.DB
now func() time.Time
}
func NewRepository(db *sql.DB) Repository {
return Repository{
db: db,
now: time.Now,
}
}
func NewRepositoryWithClock(db *sql.DB, now func() time.Time) Repository {
repo := NewRepository(db)
if now != nil {
repo.now = now
}
return repo
}
func (r Repository) EnsureDefaults(ctx context.Context, tcpAdminPort int) error {
now := r.now().Unix()
for _, service := range DefaultRecords(tcpAdminPort) {
options, err := json.Marshal(service.Options)
if err != nil {
return err
}
if _, err := r.db.ExecContext(ctx, `
INSERT INTO services(name, enabled, port, options_json, restart_required, created_at, updated_at)
VALUES (?, ?, ?, ?, 0, ?, ?)
ON CONFLICT(name) DO NOTHING`,
service.Name, boolToInt(service.Enabled), service.Port, string(options), now, now); err != nil {
return err
}
}
return nil
}
func (r Repository) List(ctx context.Context) ([]Record, error) {
rows, err := r.db.QueryContext(ctx, `
SELECT name, enabled, port, options_json, restart_required, created_at, updated_at, updated_by
FROM services
ORDER BY CASE name
WHEN 'tcp_admin' THEN 0
WHEN 'kcp' THEN 1
WHEN 'quic' THEN 2
WHEN 'websocket' THEN 3
ELSE 4
END, name`)
if err != nil {
return nil, err
}
defer rows.Close()
var services []Record
for rows.Next() {
var service Record
var enabled, restartRequired int
var optionsJSON string
if err := rows.Scan(&service.Name, &enabled, &service.Port, &optionsJSON, &restartRequired, &service.CreatedAt, &service.UpdatedAt, &service.UpdatedBy); err != nil {
return nil, err
}
service.Enabled = enabled != 0
service.RestartRequired = restartRequired != 0
service.Options = DecodeOptions(optionsJSON)
services = append(services, service)
}
return services, rows.Err()
}
func (r Repository) Update(ctx context.Context, actor, name string, enabled bool, port int, options map[string]any) error {
if err := ValidateUpdate(name, enabled, port); err != nil {
return err
}
optionsJSON, err := json.Marshal(NormalizeOptions(name, options))
if err != nil {
return err
}
_, err = r.db.ExecContext(ctx, `
UPDATE services
SET enabled = ?, port = ?, options_json = ?, restart_required = 1, updated_at = ?, updated_by = ?
WHERE name = ?`,
boolToInt(enabled), port, string(optionsJSON), r.now().Unix(), actor, name)
return err
}
func boolToInt(value bool) int {
if value {
return 1
}
return 0
}

View File

@@ -0,0 +1,109 @@
package adminservice
import (
"context"
"path/filepath"
"reflect"
"testing"
"time"
"github.com/tursom/mc-gateway/internal/admindb"
)
func TestRepositoryEnsureDefaultsAndList(t *testing.T) {
repo, closeDB := newServiceTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.EnsureDefaults(ctx, 25575); err != nil {
t.Fatalf("EnsureDefaults() error = %v", err)
}
if err := repo.EnsureDefaults(ctx, 25576); err != nil {
t.Fatalf("EnsureDefaults(second) error = %v", err)
}
services, err := repo.List(ctx)
if err != nil {
t.Fatalf("List() error = %v", err)
}
if len(services) != 4 {
t.Fatalf("List() len = %d, want 4", len(services))
}
names := []string{services[0].Name, services[1].Name, services[2].Name, services[3].Name}
if want := []string{NameTCPAdmin, NameKCP, NameQUIC, NameWebSocket}; !reflect.DeepEqual(names, want) {
t.Fatalf("service order = %#v, want %#v", names, want)
}
if services[0].Port != 25575 || !services[0].Enabled || services[0].CreatedAt != 100 || services[0].UpdatedAt != 100 {
t.Fatalf("tcp_admin service = %+v", services[0])
}
if services[1].Options["data_shards"] != float64(DefaultKCPDataShards) {
t.Fatalf("kcp options = %#v", services[1].Options)
}
if got := StringSliceOption(services[2].Options, "application_protocols"); !reflect.DeepEqual(got, []string{"minecraft", "quic", "raw", "h3"}) {
t.Fatalf("quic protocols = %#v", got)
}
}
func TestRepositoryUpdate(t *testing.T) {
repo, closeDB := newServiceTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.EnsureDefaults(ctx, 25565); err != nil {
t.Fatalf("EnsureDefaults() error = %v", err)
}
repo.now = func() time.Time { return time.Unix(200, 0) }
if err := repo.Update(ctx, "admin", NameWebSocket, true, 25580, map[string]any{"path": "gateway"}); err != nil {
t.Fatalf("Update() error = %v", err)
}
services, err := repo.List(ctx)
if err != nil {
t.Fatalf("List() error = %v", err)
}
var websocket Record
for _, service := range services {
if service.Name == NameWebSocket {
websocket = service
break
}
}
if !websocket.Enabled || websocket.Port != 25580 || !websocket.RestartRequired || websocket.UpdatedAt != 200 || websocket.UpdatedBy != "admin" {
t.Fatalf("websocket service = %+v", websocket)
}
if websocket.Options["path"] != "/gateway" {
t.Fatalf("websocket options = %#v", websocket.Options)
}
}
func TestRepositoryUpdateValidationErrors(t *testing.T) {
repo, closeDB := newServiceTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.Update(ctx, "admin", "unknown", true, 25565, nil); err == nil {
t.Fatal("Update(unknown) error = nil, want error")
}
if err := repo.Update(ctx, "admin", NameKCP, true, 0, nil); err == nil {
t.Fatal("Update(port 0) error = nil, want error")
}
if err := repo.Update(ctx, "admin", NameTCPAdmin, false, 25565, nil); err == nil {
t.Fatal("Update(disable tcp_admin) error = nil, want error")
}
}
func newServiceTestRepository(t *testing.T, now time.Time) (Repository, func()) {
t.Helper()
db, err := admindb.Open(filepath.Join(t.TempDir(), "gateway.sqlite3"))
if err != nil {
t.Fatalf("Open() error = %v", err)
}
if err := admindb.Migrate(db); err != nil {
db.Close()
t.Fatalf("Migrate() error = %v", err)
}
return NewRepositoryWithClock(db, func() time.Time { return now }), func() {
_ = db.Close()
}
}

View File

@@ -0,0 +1,190 @@
package adminservice
import (
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
)
const (
NameTCPAdmin = "tcp_admin"
NameKCP = "kcp"
NameQUIC = "quic"
NameWebSocket = "websocket"
DefaultKCPPort = 25566
DefaultKCPDataShards = 10
DefaultKCPParityShards = 3
DefaultQUICPort = 25565
DefaultWebSocketPort = 25566
DefaultWebSocketPath = "/"
)
var defaultQUICApplicationProtocols = []string{"minecraft", "quic", "raw", "h3"}
type Record struct {
Name string `json:"name"`
Enabled bool `json:"enabled"`
Port int `json:"port"`
Options map[string]any `json:"options"`
RestartRequired bool `json:"restart_required"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
UpdatedBy string `json:"updated_by"`
Running bool `json:"running"`
}
func DefaultRecords(tcpAdminPort int) []Record {
return []Record{
{Name: NameTCPAdmin, Enabled: true, Port: tcpAdminPort, Options: map[string]any{}},
{Name: NameKCP, Enabled: false, Port: DefaultKCPPort, Options: map[string]any{
"data_shards": DefaultKCPDataShards,
"parity_shards": DefaultKCPParityShards,
}},
{Name: NameQUIC, Enabled: false, Port: DefaultQUICPort, Options: map[string]any{
"application_protocols": append([]string(nil), defaultQUICApplicationProtocols...),
}},
{Name: NameWebSocket, Enabled: false, Port: DefaultWebSocketPort, Options: map[string]any{
"path": DefaultWebSocketPath,
}},
}
}
func DefaultPort(name string, tcpAdminPort int) int {
switch name {
case NameTCPAdmin:
return tcpAdminPort
case NameKCP:
return DefaultKCPPort
case NameQUIC:
return DefaultQUICPort
case NameWebSocket:
return DefaultWebSocketPort
default:
return 0
}
}
func ValidateUpdate(name string, enabled bool, port int) error {
if !IsKnown(name) {
return fmt.Errorf("unknown service %q", name)
}
if port < 1 || port > 65535 {
return errors.New("port must be an integer from 1 to 65535")
}
if name == NameTCPAdmin && !enabled {
return errors.New("tcp_admin cannot be disabled")
}
return nil
}
func IsKnown(name string) bool {
switch name {
case NameTCPAdmin, NameKCP, NameQUIC, NameWebSocket:
return true
default:
return false
}
}
func DecodeOptions(optionsJSON string) map[string]any {
options := map[string]any{}
if strings.TrimSpace(optionsJSON) == "" {
return options
}
if err := json.Unmarshal([]byte(optionsJSON), &options); err != nil {
return map[string]any{}
}
return options
}
func NormalizeOptions(name string, options map[string]any) map[string]any {
if options == nil {
options = map[string]any{}
}
normalized := map[string]any{}
for key, value := range options {
normalized[key] = value
}
switch name {
case NameKCP:
if IntOption(normalized, "data_shards", 0) <= 0 {
normalized["data_shards"] = DefaultKCPDataShards
}
if IntOption(normalized, "parity_shards", 0) <= 0 {
normalized["parity_shards"] = DefaultKCPParityShards
}
case NameQUIC:
if len(StringSliceOption(normalized, "application_protocols")) == 0 {
normalized["application_protocols"] = append([]string(nil), defaultQUICApplicationProtocols...)
}
case NameWebSocket:
p := StringOption(normalized, "path", DefaultWebSocketPath)
if !strings.HasPrefix(p, "/") {
p = "/" + p
}
normalized["path"] = p
}
return normalized
}
func IntOption(options map[string]any, key string, fallback int) int {
value, ok := options[key]
if !ok {
return fallback
}
switch v := value.(type) {
case int:
return v
case int64:
return int(v)
case float64:
return int(v)
case json.Number:
i, err := v.Int64()
if err == nil {
return int(i)
}
case string:
i, err := strconv.Atoi(v)
if err == nil {
return i
}
}
return fallback
}
func StringOption(options map[string]any, key, fallback string) string {
value, ok := options[key]
if !ok {
return fallback
}
if s, ok := value.(string); ok && s != "" {
return s
}
return fallback
}
func StringSliceOption(options map[string]any, key string) []string {
value, ok := options[key]
if !ok {
return nil
}
switch v := value.(type) {
case []string:
return v
case []any:
out := make([]string, 0, len(v))
for _, item := range v {
if s, ok := item.(string); ok && s != "" {
out = append(out, s)
}
}
return out
}
return nil
}

View File

@@ -0,0 +1,171 @@
package adminservice
import (
"encoding/json"
"reflect"
"strings"
"testing"
)
func TestDefaultRecords(t *testing.T) {
records := DefaultRecords(25575)
if len(records) != 4 {
t.Fatalf("DefaultRecords() len = %d, want 4", len(records))
}
want := []Record{
{Name: NameTCPAdmin, Enabled: true, Port: 25575, Options: map[string]any{}},
{Name: NameKCP, Enabled: false, Port: DefaultKCPPort, Options: map[string]any{
"data_shards": DefaultKCPDataShards,
"parity_shards": DefaultKCPParityShards,
}},
{Name: NameQUIC, Enabled: false, Port: DefaultQUICPort, Options: map[string]any{
"application_protocols": []string{"minecraft", "quic", "raw", "h3"},
}},
{Name: NameWebSocket, Enabled: false, Port: DefaultWebSocketPort, Options: map[string]any{
"path": DefaultWebSocketPath,
}},
}
if !reflect.DeepEqual(records, want) {
t.Fatalf("DefaultRecords() = %#v, want %#v", records, want)
}
records[2].Options["application_protocols"].([]string)[0] = "changed"
if got := DefaultRecords(25575)[2].Options["application_protocols"].([]string)[0]; got != "minecraft" {
t.Fatalf("DefaultRecords() reused mutable protocol defaults, got %q", got)
}
}
func TestDefaultPort(t *testing.T) {
tests := []struct {
name string
want int
}{
{NameTCPAdmin, 25575},
{NameKCP, DefaultKCPPort},
{NameQUIC, DefaultQUICPort},
{NameWebSocket, DefaultWebSocketPort},
{"unknown", 0},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := DefaultPort(tt.name, 25575); got != tt.want {
t.Fatalf("DefaultPort(%q) = %d, want %d", tt.name, got, tt.want)
}
})
}
}
func TestValidateUpdate(t *testing.T) {
valid := []struct {
name string
enabled bool
port int
}{
{NameTCPAdmin, true, 25565},
{NameKCP, false, 25566},
{NameQUIC, true, 25565},
{NameWebSocket, true, 25566},
}
for _, tt := range valid {
t.Run(tt.name, func(t *testing.T) {
if err := ValidateUpdate(tt.name, tt.enabled, tt.port); err != nil {
t.Fatalf("ValidateUpdate() error = %v", err)
}
})
}
invalid := []struct {
name string
enabled bool
port int
}{
{"unknown", true, 25565},
{NameKCP, true, 0},
{NameKCP, true, 70000},
{NameTCPAdmin, false, 25565},
}
for _, tt := range invalid {
t.Run(tt.name, func(t *testing.T) {
if err := ValidateUpdate(tt.name, tt.enabled, tt.port); err == nil {
t.Fatal("ValidateUpdate() error = nil, want error")
}
})
}
}
func TestDecodeAndReadOptions(t *testing.T) {
options := DecodeOptions(`{
"int": 12,
"json_number": 13,
"string_int": "14",
"text": "value",
"items": ["minecraft", "", "raw"]
}`)
if got := IntOption(options, "int", 0); got != 12 {
t.Fatalf("IntOption(float64) = %d, want 12", got)
}
decoder := json.NewDecoder(strings.NewReader(`{"json_number":15}`))
decoder.UseNumber()
var withNumber map[string]any
if err := decoder.Decode(&withNumber); err != nil {
t.Fatalf("Decode() error = %v", err)
}
if got := IntOption(withNumber, "json_number", 0); got != 15 {
t.Fatalf("IntOption(json.Number) = %d, want 15", got)
}
if got := IntOption(options, "string_int", 0); got != 14 {
t.Fatalf("IntOption(string) = %d, want 14", got)
}
if got := StringOption(options, "text", "fallback"); got != "value" {
t.Fatalf("StringOption() = %q, want value", got)
}
if got := StringSliceOption(options, "items"); !reflect.DeepEqual(got, []string{"minecraft", "raw"}) {
t.Fatalf("StringSliceOption() = %#v", got)
}
if got := DecodeOptions("not-json"); len(got) != 0 {
t.Fatalf("DecodeOptions(invalid) = %#v, want empty", got)
}
}
func TestNormalizeOptions(t *testing.T) {
tests := []struct {
name string
options map[string]any
want map[string]any
}{
{
name: NameKCP,
options: map[string]any{"data_shards": 0, "parity_shards": "4"},
want: map[string]any{
"data_shards": DefaultKCPDataShards,
"parity_shards": "4",
},
},
{
name: NameQUIC,
options: map[string]any{"application_protocols": []any{}},
want: map[string]any{
"application_protocols": []string{"minecraft", "quic", "raw", "h3"},
},
},
{
name: NameWebSocket,
options: map[string]any{"path": "gateway"},
want: map[string]any{"path": "/gateway"},
},
{
name: NameWebSocket,
options: map[string]any{"path": "/gateway"},
want: map[string]any{"path": "/gateway"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := NormalizeOptions(tt.name, tt.options)
if !reflect.DeepEqual(got, tt.want) {
t.Fatalf("NormalizeOptions() = %#v, want %#v", got, tt.want)
}
})
}
}

View File

@@ -0,0 +1,92 @@
package adminsession
import (
"crypto/rand"
"encoding/base64"
"sync"
"time"
)
type Session struct {
Token string
Username string
Role string
ExpiresAt time.Time
}
type Manager struct {
mu sync.Mutex
sessions map[string]Session
now func() time.Time
}
func NewManager() *Manager {
return &Manager{
sessions: make(map[string]Session),
now: time.Now,
}
}
func NewManagerWithClock(now func() time.Time) *Manager {
manager := NewManager()
if now != nil {
manager.now = now
}
return manager
}
func (m *Manager) Create(username, role string, ttl time.Duration) (Session, error) {
tokenBytes := make([]byte, 32)
if _, err := rand.Read(tokenBytes); err != nil {
return Session{}, err
}
session := Session{
Token: base64.RawURLEncoding.EncodeToString(tokenBytes),
Username: username,
Role: role,
ExpiresAt: m.now().Add(ttl),
}
m.mu.Lock()
m.sessions[session.Token] = session
m.mu.Unlock()
return session, nil
}
func (m *Manager) Get(token string) (Session, bool) {
m.mu.Lock()
defer m.mu.Unlock()
session, ok := m.sessions[token]
if !ok {
return Session{}, false
}
if m.now().After(session.ExpiresAt) {
delete(m.sessions, token)
return Session{}, false
}
return session, true
}
func (m *Manager) Delete(token string) {
m.mu.Lock()
delete(m.sessions, token)
m.mu.Unlock()
}
func (m *Manager) RemoveUser(username string) {
m.mu.Lock()
defer m.mu.Unlock()
for token, session := range m.sessions {
if session.Username == username {
delete(m.sessions, token)
}
}
}
func (m *Manager) Clear() {
m.mu.Lock()
m.sessions = make(map[string]Session)
m.mu.Unlock()
}

View File

@@ -0,0 +1,86 @@
package adminsession
import (
"testing"
"time"
)
func TestManagerCreateAndGet(t *testing.T) {
now := time.Unix(100, 0)
manager := NewManagerWithClock(func() time.Time { return now })
session, err := manager.Create("admin", "admin", time.Hour)
if err != nil {
t.Fatalf("Create() error = %v", err)
}
if session.Token == "" {
t.Fatal("Create() token is empty")
}
if len(session.Token) != 43 {
t.Fatalf("Create() token len = %d, want 43", len(session.Token))
}
if !session.ExpiresAt.Equal(now.Add(time.Hour)) {
t.Fatalf("ExpiresAt = %v, want %v", session.ExpiresAt, now.Add(time.Hour))
}
got, ok := manager.Get(session.Token)
if !ok {
t.Fatal("Get() ok = false, want true")
}
if got.Username != "admin" || got.Role != "admin" {
t.Fatalf("Get() = %+v", got)
}
}
func TestManagerExpiresSessions(t *testing.T) {
now := time.Unix(100, 0)
manager := NewManagerWithClock(func() time.Time { return now })
session, err := manager.Create("admin", "admin", time.Hour)
if err != nil {
t.Fatalf("Create() error = %v", err)
}
now = now.Add(time.Hour + time.Second)
if _, ok := manager.Get(session.Token); ok {
t.Fatal("Get(expired) ok = true, want false")
}
if _, ok := manager.Get(session.Token); ok {
t.Fatal("Get(expired deleted) ok = true, want false")
}
}
func TestManagerDeleteAndRemoveUser(t *testing.T) {
manager := NewManager()
admin, err := manager.Create("admin", "admin", time.Hour)
if err != nil {
t.Fatalf("Create(admin) error = %v", err)
}
member, err := manager.Create("member", "member", time.Hour)
if err != nil {
t.Fatalf("Create(member) error = %v", err)
}
guest, err := manager.Create("guest", "guest", time.Hour)
if err != nil {
t.Fatalf("Create(guest) error = %v", err)
}
manager.Delete(guest.Token)
if _, ok := manager.Get(guest.Token); ok {
t.Fatal("guest token still exists after Delete")
}
manager.RemoveUser("admin")
if _, ok := manager.Get(admin.Token); ok {
t.Fatal("admin token still exists after RemoveUser")
}
if _, ok := manager.Get(member.Token); !ok {
t.Fatal("member token was removed by RemoveUser(admin)")
}
manager.Clear()
if _, ok := manager.Get(member.Token); ok {
t.Fatal("member token still exists after Clear")
}
}

View File

@@ -0,0 +1,233 @@
package adminuser
import (
"context"
"database/sql"
"errors"
"strings"
"time"
)
type Repository struct {
db *sql.DB
now func() time.Time
}
type Patch struct {
Role *string
Disabled *bool
PasswordHashProvider func() (string, error)
}
type PatchResult struct {
InvalidateSessions bool
}
func NewRepository(db *sql.DB) Repository {
return Repository{
db: db,
now: time.Now,
}
}
func NewRepositoryWithClock(db *sql.DB, now func() time.Time) Repository {
repo := NewRepository(db)
if now != nil {
repo.now = now
}
return repo
}
func (r Repository) TableEmpty(ctx context.Context) (bool, error) {
var count int
if err := r.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users`).Scan(&count); err != nil {
return false, err
}
return count == 0, nil
}
func (r Repository) Create(ctx context.Context, username, role, passwordHash string, disabled bool) error {
username = strings.TrimSpace(username)
if err := ValidateUsername(username); err != nil {
return err
}
if err := ValidateRole(role); err != nil {
return err
}
if strings.TrimSpace(passwordHash) == "" {
return errors.New("password hash is required")
}
now := r.now().Unix()
_, err := r.db.ExecContext(ctx, `
INSERT INTO users(username, role, password_hash, disabled, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?)`,
username, role, passwordHash, boolToInt(disabled), now, now)
return err
}
func (r Repository) GetWithHash(ctx context.Context, username string) (User, string, error) {
user, hash, err := r.getWithHash(ctx, r.db, username, "invalid username or password")
return user, hash, err
}
func (r Repository) List(ctx context.Context) ([]User, error) {
rows, err := r.db.QueryContext(ctx, `
SELECT username, role, disabled, created_at, updated_at
FROM users
ORDER BY username`)
if err != nil {
return nil, err
}
defer rows.Close()
var users []User
for rows.Next() {
var user User
var disabled int
if err := rows.Scan(&user.Username, &user.Role, &disabled, &user.CreatedAt, &user.UpdatedAt); err != nil {
return nil, err
}
user.Disabled = disabled != 0
users = append(users, user)
}
return users, rows.Err()
}
func (r Repository) Patch(ctx context.Context, username string, patch Patch) (PatchResult, error) {
username = strings.TrimSpace(username)
if err := ValidateUsername(username); err != nil {
return PatchResult{}, err
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return PatchResult{}, err
}
defer tx.Rollback()
current, _, err := r.getWithHash(ctx, tx, username, "user not found")
if err != nil {
return PatchResult{}, err
}
nextRole := current.Role
if patch.Role != nil {
if err := ValidateRole(*patch.Role); err != nil {
return PatchResult{}, err
}
nextRole = *patch.Role
}
nextDisabled := current.Disabled
if patch.Disabled != nil {
nextDisabled = *patch.Disabled
}
if current.Role == RoleAdmin && (nextRole != RoleAdmin || nextDisabled) {
count, err := r.enabledAdminCount(ctx, tx, username)
if err != nil {
return PatchResult{}, err
}
if count == 0 {
return PatchResult{}, errors.New("cannot remove the last enabled admin")
}
}
sets := []string{"role = ?", "disabled = ?", "updated_at = ?"}
args := []any{nextRole, boolToInt(nextDisabled), r.now().Unix()}
if patch.PasswordHashProvider != nil {
passwordHash, err := patch.PasswordHashProvider()
if err != nil {
return PatchResult{}, err
}
if strings.TrimSpace(passwordHash) == "" {
return PatchResult{}, errors.New("password hash is required")
}
sets = append(sets, "password_hash = ?")
args = append(args, passwordHash)
}
args = append(args, username)
if _, err := tx.ExecContext(ctx, "UPDATE users SET "+strings.Join(sets, ", ")+" WHERE username = ?", args...); err != nil {
return PatchResult{}, err
}
if err := tx.Commit(); err != nil {
return PatchResult{}, err
}
return PatchResult{
InvalidateSessions: nextDisabled || nextRole != current.Role || patch.PasswordHashProvider != nil,
}, nil
}
func (r Repository) Delete(ctx context.Context, username string) error {
username = strings.TrimSpace(username)
if err := ValidateUsername(username); err != nil {
return err
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
current, _, err := r.getWithHash(ctx, tx, username, "user not found")
if err != nil {
return err
}
if current.Role == RoleAdmin {
count, err := r.enabledAdminCount(ctx, tx, username)
if err != nil {
return err
}
if count == 0 {
return errors.New("cannot delete the last enabled admin")
}
}
if _, err := tx.ExecContext(ctx, `DELETE FROM users WHERE username = ?`, username); err != nil {
return err
}
return tx.Commit()
}
type userQueryer interface {
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
}
func (r Repository) getWithHash(ctx context.Context, q userQueryer, username, notFoundMessage string) (User, string, error) {
var user User
var hash string
var disabled int
err := q.QueryRowContext(ctx, `
SELECT username, role, password_hash, disabled, created_at, updated_at
FROM users
WHERE username = ?`, username).
Scan(&user.Username, &user.Role, &hash, &disabled, &user.CreatedAt, &user.UpdatedAt)
if errors.Is(err, sql.ErrNoRows) {
return User{}, "", errors.New(notFoundMessage)
}
if err != nil {
return User{}, "", err
}
user.Disabled = disabled != 0
return user, hash, nil
}
func (r Repository) enabledAdminCount(ctx context.Context, tx *sql.Tx, excludingUsername string) (int, error) {
var count int
err := tx.QueryRowContext(ctx, `
SELECT COUNT(*)
FROM users
WHERE role = ? AND disabled = 0 AND username <> ?`, RoleAdmin, excludingUsername).Scan(&count)
return count, err
}
func boolToInt(value bool) int {
if value {
return 1
}
return 0
}

View File

@@ -0,0 +1,185 @@
package adminuser
import (
"context"
"path/filepath"
"reflect"
"strings"
"testing"
"time"
"github.com/tursom/mc-gateway/internal/admindb"
)
func TestRepositoryCreateGetListAndTableEmpty(t *testing.T) {
repo, closeDB := newUserTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
empty, err := repo.TableEmpty(ctx)
if err != nil {
t.Fatalf("TableEmpty() error = %v", err)
}
if !empty {
t.Fatal("TableEmpty() = false, want true")
}
if err := repo.Create(ctx, "admin", RoleAdmin, "admin-hash", false); err != nil {
t.Fatalf("Create(admin) error = %v", err)
}
if err := repo.Create(ctx, "guest", RoleGuest, "guest-hash", true); err != nil {
t.Fatalf("Create(guest) error = %v", err)
}
empty, err = repo.TableEmpty(ctx)
if err != nil {
t.Fatalf("TableEmpty(after create) error = %v", err)
}
if empty {
t.Fatal("TableEmpty(after create) = true, want false")
}
user, hash, err := repo.GetWithHash(ctx, "admin")
if err != nil {
t.Fatalf("GetWithHash() error = %v", err)
}
if user.Username != "admin" || user.Role != RoleAdmin || user.Disabled || user.CreatedAt != 100 || user.UpdatedAt != 100 || hash != "admin-hash" {
t.Fatalf("GetWithHash() user=%+v hash=%q", user, hash)
}
users, err := repo.List(ctx)
if err != nil {
t.Fatalf("List() error = %v", err)
}
want := []User{
{Username: "admin", Role: RoleAdmin, CreatedAt: 100, UpdatedAt: 100},
{Username: "guest", Role: RoleGuest, Disabled: true, CreatedAt: 100, UpdatedAt: 100},
}
if !reflect.DeepEqual(users, want) {
t.Fatalf("List() = %#v, want %#v", users, want)
}
}
func TestRepositoryPatch(t *testing.T) {
now := time.Unix(100, 0)
repo, closeDB := newUserTestRepository(t, now)
defer closeDB()
ctx := context.Background()
if err := repo.Create(ctx, "admin", RoleAdmin, "admin-hash", false); err != nil {
t.Fatalf("Create(admin) error = %v", err)
}
if err := repo.Create(ctx, "member", RoleMember, "member-hash", false); err != nil {
t.Fatalf("Create(member) error = %v", err)
}
nextRole := RoleGuest
repo.now = func() time.Time { return time.Unix(200, 0) }
result, err := repo.Patch(ctx, "member", Patch{
Role: &nextRole,
PasswordHashProvider: func() (string, error) {
return "new-member-hash", nil
},
})
if err != nil {
t.Fatalf("Patch() error = %v", err)
}
if !result.InvalidateSessions {
t.Fatal("Patch() InvalidateSessions = false, want true")
}
user, hash, err := repo.GetWithHash(ctx, "member")
if err != nil {
t.Fatalf("GetWithHash(member) error = %v", err)
}
if user.Role != RoleGuest || user.Disabled || user.UpdatedAt != 200 || hash != "new-member-hash" {
t.Fatalf("patched user=%+v hash=%q", user, hash)
}
result, err = repo.Patch(ctx, "member", Patch{})
if err != nil {
t.Fatalf("Patch(no-op) error = %v", err)
}
if result.InvalidateSessions {
t.Fatal("Patch(no-op) InvalidateSessions = true, want false")
}
}
func TestRepositoryPreventsRemovingLastEnabledAdmin(t *testing.T) {
repo, closeDB := newUserTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.Create(ctx, "admin", RoleAdmin, "admin-hash", false); err != nil {
t.Fatalf("Create(admin) error = %v", err)
}
guestRole := RoleGuest
if _, err := repo.Patch(ctx, "admin", Patch{Role: &guestRole}); err == nil || !strings.Contains(err.Error(), "last enabled admin") {
t.Fatalf("Patch(last admin role) error = %v, want last enabled admin", err)
}
disabled := true
if _, err := repo.Patch(ctx, "admin", Patch{Disabled: &disabled}); err == nil || !strings.Contains(err.Error(), "last enabled admin") {
t.Fatalf("Patch(last admin disabled) error = %v, want last enabled admin", err)
}
if err := repo.Delete(ctx, "admin"); err == nil || !strings.Contains(err.Error(), "last enabled admin") {
t.Fatalf("Delete(last admin) error = %v, want last enabled admin", err)
}
}
func TestRepositoryDeleteAdminWhenAnotherAdminExists(t *testing.T) {
repo, closeDB := newUserTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.Create(ctx, "admin1", RoleAdmin, "hash-1", false); err != nil {
t.Fatalf("Create(admin1) error = %v", err)
}
if err := repo.Create(ctx, "admin2", RoleAdmin, "hash-2", false); err != nil {
t.Fatalf("Create(admin2) error = %v", err)
}
if err := repo.Delete(ctx, "admin1"); err != nil {
t.Fatalf("Delete(admin1) error = %v", err)
}
if _, _, err := repo.GetWithHash(ctx, "admin1"); err == nil || err.Error() != "invalid username or password" {
t.Fatalf("GetWithHash(deleted) error = %v, want invalid username or password", err)
}
}
func TestRepositoryValidationErrors(t *testing.T) {
repo, closeDB := newUserTestRepository(t, time.Unix(100, 0))
defer closeDB()
ctx := context.Background()
if err := repo.Create(ctx, "bad user", RoleAdmin, "hash", false); err == nil {
t.Fatal("Create(invalid username) error = nil, want error")
}
if err := repo.Create(ctx, "admin", "owner", "hash", false); err == nil {
t.Fatal("Create(invalid role) error = nil, want error")
}
if err := repo.Create(ctx, "admin", RoleAdmin, "", false); err == nil {
t.Fatal("Create(empty hash) error = nil, want error")
}
if _, err := repo.Patch(ctx, "missing", Patch{}); err == nil || err.Error() != "user not found" {
t.Fatalf("Patch(missing) error = %v, want user not found", err)
}
if err := repo.Delete(ctx, "missing"); err == nil || err.Error() != "user not found" {
t.Fatalf("Delete(missing) error = %v, want user not found", err)
}
}
func newUserTestRepository(t *testing.T, now time.Time) (Repository, func()) {
t.Helper()
db, err := admindb.Open(filepath.Join(t.TempDir(), "gateway.sqlite3"))
if err != nil {
t.Fatalf("Open() error = %v", err)
}
if err := admindb.Migrate(db); err != nil {
db.Close()
t.Fatalf("Migrate() error = %v", err)
}
return NewRepositoryWithClock(db, func() time.Time { return now }), func() {
_ = db.Close()
}
}

View File

@@ -0,0 +1,74 @@
package adminuser
import (
"errors"
"fmt"
"strings"
)
const (
RoleAdmin = "admin"
RoleMember = "member"
RoleGuest = "guest"
)
type User struct {
Username string `json:"username"`
Role string `json:"role"`
Disabled bool `json:"disabled"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
}
func ValidateUsername(username string) error {
if username == "" {
return errors.New("username is required")
}
if strings.ContainsAny(username, " \t\r\n/") {
return errors.New("username must not contain whitespace or /")
}
return nil
}
func ValidateRole(role string) error {
switch role {
case RoleAdmin, RoleMember, RoleGuest:
return nil
default:
return fmt.Errorf("invalid role %q", role)
}
}
func ValidatePassword(password string) error {
if strings.TrimSpace(password) == "" {
return errors.New("password is required")
}
return nil
}
func RoleRank(role string) int {
switch role {
case RoleAdmin:
return 3
case RoleMember:
return 2
case RoleGuest:
return 1
default:
return 0
}
}
func HasRole(actual, required string) bool {
return RoleRank(actual) >= RoleRank(required)
}
func Permissions(role string) map[string]bool {
return map[string]bool{
"read_routes": HasRole(role, RoleGuest),
"write_routes": HasRole(role, RoleMember),
"read_status": HasRole(role, RoleMember),
"manage_users": HasRole(role, RoleAdmin),
"manage_services": HasRole(role, RoleAdmin),
}
}

View File

@@ -0,0 +1,102 @@
package adminuser
import (
"reflect"
"testing"
)
func TestValidateUsername(t *testing.T) {
valid := []string{"admin", "member-1", "guest.example"}
for _, username := range valid {
t.Run(username, func(t *testing.T) {
if err := ValidateUsername(username); err != nil {
t.Fatalf("ValidateUsername(%q) error = %v", username, err)
}
})
}
invalid := []string{"", "admin user", "admin/user", "admin\tuser", "admin\nuser"}
for _, username := range invalid {
t.Run("invalid "+username, func(t *testing.T) {
if err := ValidateUsername(username); err == nil {
t.Fatalf("ValidateUsername(%q) error = nil, want error", username)
}
})
}
}
func TestValidateRole(t *testing.T) {
for _, role := range []string{RoleAdmin, RoleMember, RoleGuest} {
t.Run(role, func(t *testing.T) {
if err := ValidateRole(role); err != nil {
t.Fatalf("ValidateRole(%q) error = %v", role, err)
}
})
}
if err := ValidateRole("owner"); err == nil {
t.Fatal("ValidateRole(owner) error = nil, want error")
}
}
func TestValidatePassword(t *testing.T) {
if err := ValidatePassword("secret"); err != nil {
t.Fatalf("ValidatePassword(secret) error = %v", err)
}
for _, password := range []string{"", " ", "\t"} {
t.Run("empty", func(t *testing.T) {
if err := ValidatePassword(password); err == nil {
t.Fatal("ValidatePassword() error = nil, want error")
}
})
}
}
func TestRoleRankAndHasRole(t *testing.T) {
tests := []struct {
role string
want int
}{
{RoleAdmin, 3},
{RoleMember, 2},
{RoleGuest, 1},
{"unknown", 0},
}
for _, tt := range tests {
t.Run(tt.role, func(t *testing.T) {
if got := RoleRank(tt.role); got != tt.want {
t.Fatalf("RoleRank(%q) = %d, want %d", tt.role, got, tt.want)
}
})
}
if !HasRole(RoleAdmin, RoleGuest) {
t.Fatal("admin should satisfy guest")
}
if HasRole(RoleGuest, RoleMember) {
t.Fatal("guest should not satisfy member")
}
}
func TestPermissions(t *testing.T) {
wantGuest := map[string]bool{
"read_routes": true,
"write_routes": false,
"read_status": false,
"manage_users": false,
"manage_services": false,
}
if got := Permissions(RoleGuest); !reflect.DeepEqual(got, wantGuest) {
t.Fatalf("Permissions(guest) = %#v, want %#v", got, wantGuest)
}
wantAdmin := map[string]bool{
"read_routes": true,
"write_routes": true,
"read_status": true,
"manage_users": true,
"manage_services": true,
}
if got := Permissions(RoleAdmin); !reflect.DeepEqual(got, wantAdmin) {
t.Fatalf("Permissions(admin) = %#v, want %#v", got, wantAdmin)
}
}

View File

@@ -0,0 +1,40 @@
package gatewayconfig
type Config struct {
Tcp ProtocolConfig `toml:"tcp"`
Quic QuicConfig `toml:"quic"`
Kcp KcpConfig `toml:"kcp"`
WebSocket WebSocketConfig `toml:"websocket"`
Log LogConfig `toml:"log"`
PidFile string `toml:"pid_file"`
Plugin map[string]map[string]any `toml:"plugin"`
}
type ProtocolConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
}
type KcpConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
DataShards int `toml:"data_shards"`
ParityShards int `toml:"parity_Shards"`
}
type QuicConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
ApplicationProtocols []string `toml:"application_protocols"`
}
type WebSocketConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
Path string `toml:"path"`
}
type LogConfig struct {
Level string `toml:"level"`
File string `toml:"file"`
}

View File

@@ -0,0 +1,15 @@
package gatewayconfig
import "github.com/mitchellh/mapstructure"
func DecodePluginConfig(cfg map[string]any, pluginCfg any) error {
decoder, err := mapstructure.NewDecoder(&mapstructure.DecoderConfig{
Result: pluginCfg,
TagName: "toml",
})
if err != nil {
return err
}
return decoder.Decode(cfg)
}

View File

@@ -0,0 +1,37 @@
package gatewayconfig
import "testing"
func TestDecodePluginConfig(t *testing.T) {
type pluginConfig struct {
Enable bool `toml:"enable"`
Name string `toml:"name"`
Count int `toml:"count"`
}
var got pluginConfig
err := DecodePluginConfig(map[string]any{
"enable": true,
"name": "plugin-a",
"count": 7,
}, &got)
if err != nil {
t.Fatalf("DecodePluginConfig() error = %v", err)
}
want := pluginConfig{Enable: true, Name: "plugin-a", Count: 7}
if got != want {
t.Fatalf("DecodePluginConfig() = %+v, want %+v", got, want)
}
}
func TestDecodePluginConfigReturnsDecodeError(t *testing.T) {
type pluginConfig struct {
Count int `toml:"count"`
}
var got pluginConfig
if err := DecodePluginConfig(map[string]any{"count": "not-an-int"}, &got); err == nil {
t.Fatal("DecodePluginConfig() error = nil, want error")
}
}

View File

@@ -0,0 +1,74 @@
package gatewaymetrics
import (
"sync"
"sync/atomic"
)
type Counters struct {
totalConnections atomic.Uint64
activeConnections atomic.Int64
tcpConnections atomic.Uint64
webSocketConns atomic.Uint64
routeMisses atomic.Uint64
upstreamDialErrs atomic.Uint64
routeHitsMu sync.Mutex
routeHits map[string]uint64
}
func New() *Counters {
return &Counters{
routeHits: make(map[string]uint64),
}
}
func (m *Counters) ConnectionStarted() {
m.totalConnections.Add(1)
m.activeConnections.Add(1)
}
func (m *Counters) ConnectionFinished() {
m.activeConnections.Add(-1)
}
func (m *Counters) TCPConnectionStarted() {
m.tcpConnections.Add(1)
}
func (m *Counters) WebSocketConnectionStarted() {
m.webSocketConns.Add(1)
}
func (m *Counters) RouteHit(host string) {
m.routeHitsMu.Lock()
m.routeHits[host]++
m.routeHitsMu.Unlock()
}
func (m *Counters) RouteMiss() {
m.routeMisses.Add(1)
}
func (m *Counters) UpstreamDialError() {
m.upstreamDialErrs.Add(1)
}
func (m *Counters) Snapshot() map[string]any {
m.routeHitsMu.Lock()
routeHits := make(map[string]uint64, len(m.routeHits))
for host, count := range m.routeHits {
routeHits[host] = count
}
m.routeHitsMu.Unlock()
return map[string]any{
"total_connections": m.totalConnections.Load(),
"active_connections": m.activeConnections.Load(),
"tcp_connections": m.tcpConnections.Load(),
"websocket_connections": m.webSocketConns.Load(),
"route_hits": routeHits,
"route_misses": m.routeMisses.Load(),
"upstream_dial_errors": m.upstreamDialErrs.Load(),
}
}

View File

@@ -0,0 +1,60 @@
package gatewaymetrics
import "testing"
func TestCountersSnapshot(t *testing.T) {
metrics := New()
metrics.ConnectionStarted()
metrics.ConnectionStarted()
metrics.ConnectionFinished()
metrics.TCPConnectionStarted()
metrics.WebSocketConnectionStarted()
metrics.RouteHit("play.example")
metrics.RouteHit("play.example")
metrics.RouteHit("dev.example")
metrics.RouteMiss()
metrics.UpstreamDialError()
snapshot := metrics.Snapshot()
if got := snapshot["total_connections"]; got != uint64(2) {
t.Fatalf("total_connections = %#v, want 2", got)
}
if got := snapshot["active_connections"]; got != int64(1) {
t.Fatalf("active_connections = %#v, want 1", got)
}
if got := snapshot["tcp_connections"]; got != uint64(1) {
t.Fatalf("tcp_connections = %#v, want 1", got)
}
if got := snapshot["websocket_connections"]; got != uint64(1) {
t.Fatalf("websocket_connections = %#v, want 1", got)
}
if got := snapshot["route_misses"]; got != uint64(1) {
t.Fatalf("route_misses = %#v, want 1", got)
}
if got := snapshot["upstream_dial_errors"]; got != uint64(1) {
t.Fatalf("upstream_dial_errors = %#v, want 1", got)
}
routeHits, ok := snapshot["route_hits"].(map[string]uint64)
if !ok {
t.Fatalf("route_hits = %#v, want map[string]uint64", snapshot["route_hits"])
}
if routeHits["play.example"] != 2 || routeHits["dev.example"] != 1 {
t.Fatalf("route_hits = %#v", routeHits)
}
}
func TestSnapshotCopiesRouteHits(t *testing.T) {
metrics := New()
metrics.RouteHit("play.example")
snapshot := metrics.Snapshot()
routeHits := snapshot["route_hits"].(map[string]uint64)
routeHits["play.example"] = 100
next := metrics.Snapshot()["route_hits"].(map[string]uint64)
if next["play.example"] != 1 {
t.Fatalf("route_hits was not copied, got %#v", next)
}
}

276
internal/tcphttpmux/mux.go Normal file
View File

@@ -0,0 +1,276 @@
package tcphttpmux
import (
"bytes"
"errors"
"io"
"net"
"net/http"
"sync"
"time"
)
const (
DefaultInitialPacketTimeout = time.Second
DefaultHTTPConnBacklog = 128
DefaultReadBufferSize = 64 * 1024
maxHTTPMethodPrefixLen = len("OPTIONS ")
)
var httpMethodPrefixes = [][]byte{
[]byte("GET "),
[]byte("POST "),
[]byte("HEAD "),
[]byte("PUT "),
[]byte("PATCH "),
[]byte("DELETE "),
[]byte("OPTIONS "),
[]byte("CONNECT "),
[]byte("TRACE "),
}
type Options struct {
InitialPacketTimeout time.Duration
HTTPConnBacklog int
SetSocketOptions func(net.Conn)
OnTCPConnection func()
OnAcceptError func(error)
OnInitialPacketError func(net.Conn, error)
OnEmptyInitialPacket func(net.Conn)
OnHTTPDeliveryFailed func(net.Conn)
}
type replayConn struct {
net.Conn
reader io.Reader
}
func (c *replayConn) Read(p []byte) (int, error) {
return c.reader.Read(p)
}
type ChanListener struct {
conns chan net.Conn
closed chan struct{}
closeOnce sync.Once
addr net.Addr
}
var readBufferPool = sync.Pool{
New: func() any {
buf := make([]byte, DefaultReadBufferSize)
return &buf
},
}
func Serve(listener net.Listener, handler http.Handler, tcpHandler func(net.Conn), opts Options) error {
defer listener.Close()
webListener := NewChanListener(listener.Addr(), opts.normalizedHTTPConnBacklog())
webServer := &http.Server{Handler: handler}
webServerDone := make(chan error, 1)
go func() {
err := webServer.Serve(webListener)
if err != nil && !errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
webServerDone <- err
return
}
webServerDone <- nil
}()
defer func() {
_ = webListener.Close()
_ = webServer.Close()
<-webServerDone
}()
for {
conn, err := listener.Accept()
if err != nil {
if errors.Is(err, net.ErrClosed) {
return err
}
if opts.OnAcceptError != nil {
opts.OnAcceptError(err)
}
continue
}
if opts.SetSocketOptions != nil {
opts.SetSocketOptions(conn)
}
go HandleConn(conn, webListener, tcpHandler, opts)
}
}
func HandleConn(conn net.Conn, webListener *ChanListener, tcpHandler func(net.Conn), opts Options) {
peeked, err := ReadInitialPacket(conn, opts.normalizedInitialPacketTimeout())
if err != nil {
if opts.OnInitialPacketError != nil {
opts.OnInitialPacketError(conn, err)
}
conn.Close()
return
}
if len(peeked) == 0 {
if opts.OnEmptyInitialPacket != nil {
opts.OnEmptyInitialPacket(conn)
}
conn.Close()
return
}
replayed := NewReplayConn(conn, peeked)
if IsHTTPInitialPacket(peeked) {
if !webListener.Deliver(replayed) {
if opts.OnHTTPDeliveryFailed != nil {
opts.OnHTTPDeliveryFailed(conn)
}
conn.Close()
}
return
}
if opts.OnTCPConnection != nil {
opts.OnTCPConnection()
}
tcpHandler(replayed)
}
func ReadInitialPacket(conn net.Conn, timeout time.Duration) ([]byte, error) {
if timeout > 0 {
if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil {
return nil, err
}
defer conn.SetReadDeadline(time.Time{})
}
buf := getReadBuffer()
defer putReadBuffer(buf)
var peeked []byte
for {
n, err := conn.Read(buf)
if n > 0 {
peeked = append(peeked, buf[:n]...)
if IsHTTPInitialPacket(peeked) ||
!IsPotentialHTTPInitialPacket(peeked) ||
len(peeked) >= maxHTTPMethodPrefixLen {
return peeked, nil
}
}
if err != nil {
if len(peeked) > 0 && errors.Is(err, io.EOF) {
return peeked, nil
}
return peeked, err
}
if n == 0 {
return peeked, io.ErrNoProgress
}
}
}
func NewReplayConn(conn net.Conn, peeked []byte) net.Conn {
return &replayConn{
Conn: conn,
reader: io.MultiReader(bytes.NewReader(peeked), conn),
}
}
func IsHTTPInitialPacket(buf []byte) bool {
for _, prefix := range httpMethodPrefixes {
if bytes.HasPrefix(buf, prefix) {
return true
}
}
return false
}
func IsPotentialHTTPInitialPacket(buf []byte) bool {
if len(buf) == 0 {
return true
}
for _, prefix := range httpMethodPrefixes {
if len(buf) <= len(prefix) && bytes.HasPrefix(prefix, buf) {
return true
}
}
return false
}
func getReadBuffer() []byte {
return *readBufferPool.Get().(*[]byte)
}
func putReadBuffer(buf []byte) {
if cap(buf) != DefaultReadBufferSize {
return
}
buf = buf[:DefaultReadBufferSize]
readBufferPool.Put(&buf)
}
func NewChanListener(addr net.Addr, backlog int) *ChanListener {
if backlog <= 0 {
backlog = DefaultHTTPConnBacklog
}
return &ChanListener{
conns: make(chan net.Conn, backlog),
closed: make(chan struct{}),
addr: addr,
}
}
func (l *ChanListener) Accept() (net.Conn, error) {
select {
case conn := <-l.conns:
return conn, nil
case <-l.closed:
return nil, net.ErrClosed
}
}
func (l *ChanListener) Close() error {
l.closeOnce.Do(func() {
close(l.closed)
})
return nil
}
func (l *ChanListener) Addr() net.Addr {
return l.addr
}
func (l *ChanListener) Deliver(conn net.Conn) bool {
select {
case <-l.closed:
return false
default:
}
select {
case l.conns <- conn:
return true
case <-l.closed:
return false
default:
return false
}
}
func (opts Options) normalizedInitialPacketTimeout() time.Duration {
if opts.InitialPacketTimeout == 0 {
return DefaultInitialPacketTimeout
}
return opts.InitialPacketTimeout
}
func (opts Options) normalizedHTTPConnBacklog() int {
if opts.HTTPConnBacklog <= 0 {
return DefaultHTTPConnBacklog
}
return opts.HTTPConnBacklog
}

View File

@@ -0,0 +1,222 @@
package tcphttpmux
import (
"bytes"
"errors"
"io"
"net"
"strings"
"testing"
"time"
)
func TestHTTPInitialPacketRecognition(t *testing.T) {
for _, method := range []string{
"GET ",
"POST ",
"HEAD ",
"PUT ",
"PATCH ",
"DELETE ",
"OPTIONS ",
"CONNECT ",
"TRACE ",
} {
t.Run(strings.TrimSpace(method), func(t *testing.T) {
if !IsHTTPInitialPacket([]byte(method + "/gateway HTTP/1.1\r\n")) {
t.Fatalf("IsHTTPInitialPacket(%q) = false, want true", method)
}
})
}
if IsHTTPInitialPacket([]byte{0x00, 0x01, 0x02}) {
t.Fatal("non-HTTP packet was recognized as HTTP")
}
if IsHTTPInitialPacket([]byte("GE")) {
t.Fatal("partial HTTP method was recognized as complete HTTP")
}
if !IsPotentialHTTPInitialPacket([]byte("GE")) {
t.Fatal("partial HTTP method was not recognized as a possible HTTP prefix")
}
if IsPotentialHTTPInitialPacket([]byte("GOT ")) {
t.Fatal("invalid HTTP method was recognized as a possible HTTP prefix")
}
}
func TestReplayConnReadsPeekedBytesBeforeUnderlyingConn(t *testing.T) {
base := newMuxTestConn([]byte("rest"))
conn := NewReplayConn(base, []byte("peek-"))
got, err := io.ReadAll(conn)
if err != nil {
t.Fatalf("ReadAll() error = %v", err)
}
if string(got) != "peek-rest" {
t.Fatalf("replayed data = %q, want peek-rest", got)
}
}
func TestChanListenerAcceptCloseAndDeliver(t *testing.T) {
listener := NewChanListener(muxTestAddr("listener"), DefaultHTTPConnBacklog)
conn := newMuxTestConn(nil)
if !listener.Deliver(conn) {
t.Fatal("Deliver() = false, want true")
}
got, err := listener.Accept()
if err != nil {
t.Fatalf("Accept() error = %v", err)
}
if got != conn {
t.Fatalf("Accept() = %v, want delivered conn", got)
}
if err := listener.Close(); err != nil {
t.Fatalf("Close() error = %v", err)
}
if listener.Deliver(newMuxTestConn(nil)) {
t.Fatal("Deliver() after Close = true, want false")
}
_, err = listener.Accept()
if !errors.Is(err, net.ErrClosed) {
t.Fatalf("Accept() error = %v, want %v", err, net.ErrClosed)
}
}
func TestHandleConnRoutesHTTP(t *testing.T) {
listener := NewChanListener(muxTestAddr("listener"), DefaultHTTPConnBacklog)
source := newMuxTestConn([]byte("GET / HTTP/1.1\r\n\r\n"))
tcpCalled := false
HandleConn(source, listener, func(net.Conn) {
tcpCalled = true
}, Options{InitialPacketTimeout: time.Second})
if tcpCalled {
t.Fatal("TCP handler was called for HTTP request")
}
conn, err := listener.Accept()
if err != nil {
t.Fatalf("Accept() error = %v", err)
}
got, err := io.ReadAll(conn)
if err != nil {
t.Fatalf("ReadAll() error = %v", err)
}
if string(got) != "GET / HTTP/1.1\r\n\r\n" {
t.Fatalf("HTTP replay = %q", got)
}
}
func TestHandleConnRoutesTCP(t *testing.T) {
packet := []byte{0x00, 0x01, 0x02, 'm', 'c'}
listener := NewChanListener(muxTestAddr("listener"), DefaultHTTPConnBacklog)
source := newMuxTestConn(packet)
tcpStarted := false
var got []byte
HandleConn(source, listener, func(conn net.Conn) {
var err error
got, err = io.ReadAll(conn)
if err != nil {
t.Fatalf("ReadAll() error = %v", err)
}
}, Options{
InitialPacketTimeout: time.Second,
OnTCPConnection: func() {
tcpStarted = true
},
})
if !tcpStarted {
t.Fatal("OnTCPConnection was not called")
}
if !bytes.Equal(got, packet) {
t.Fatalf("TCP replay = %v, want %v", got, packet)
}
}
func TestHandleConnTimeoutClosesConn(t *testing.T) {
client, server := net.Pipe()
defer client.Close()
listener := NewChanListener(muxTestAddr("listener"), DefaultHTTPConnBacklog)
done := make(chan struct{})
go func() {
HandleConn(server, listener, func(net.Conn) {
t.Error("TCP handler was called after timeout")
}, Options{InitialPacketTimeout: 10 * time.Millisecond})
close(done)
}()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for initial packet timeout")
}
if _, err := client.Write([]byte("x")); err == nil {
t.Fatal("client Write() error = nil, want closed connection error")
}
}
type muxTestConn struct {
reader *bytes.Reader
closed bool
}
func newMuxTestConn(data []byte) *muxTestConn {
return &muxTestConn{reader: bytes.NewReader(data)}
}
func (c *muxTestConn) Read(p []byte) (int, error) {
if c.closed {
return 0, net.ErrClosed
}
return c.reader.Read(p)
}
func (c *muxTestConn) Write(p []byte) (int, error) {
if c.closed {
return 0, net.ErrClosed
}
return len(p), nil
}
func (c *muxTestConn) Close() error {
c.closed = true
return nil
}
func (c *muxTestConn) LocalAddr() net.Addr {
return muxTestAddr("local")
}
func (c *muxTestConn) RemoteAddr() net.Addr {
return muxTestAddr("remote")
}
func (c *muxTestConn) SetDeadline(time.Time) error {
return nil
}
func (c *muxTestConn) SetReadDeadline(time.Time) error {
return nil
}
func (c *muxTestConn) SetWriteDeadline(time.Time) error {
return nil
}
type muxTestAddr string
func (a muxTestAddr) Network() string {
return "test"
}
func (a muxTestAddr) String() string {
return string(a)
}

View File

@@ -0,0 +1,30 @@
package upstreamtarget
import "strings"
type Protocol string
const (
ProtocolTCP Protocol = "tcp"
ProtocolQUIC Protocol = "quic"
ProtocolKCP Protocol = "kcp"
ProtocolHAProxy Protocol = "haproxy"
)
type Target struct {
Protocol Protocol
Address string
}
func Parse(raw string) Target {
if address, ok := strings.CutPrefix(raw, "quic://"); ok {
return Target{Protocol: ProtocolQUIC, Address: address}
}
if address, ok := strings.CutPrefix(raw, "kcp://"); ok {
return Target{Protocol: ProtocolKCP, Address: address}
}
if address, ok := strings.CutPrefix(raw, "haproxy://"); ok {
return Target{Protocol: ProtocolHAProxy, Address: address}
}
return Target{Protocol: ProtocolTCP, Address: raw}
}

View File

@@ -0,0 +1,27 @@
package upstreamtarget
import "testing"
func TestParse(t *testing.T) {
tests := []struct {
raw string
protocol Protocol
address string
}{
{raw: "127.0.0.1:25565", protocol: ProtocolTCP, address: "127.0.0.1:25565"},
{raw: "quic://127.0.0.1:25565", protocol: ProtocolQUIC, address: "127.0.0.1:25565"},
{raw: "kcp://127.0.0.1:25565", protocol: ProtocolKCP, address: "127.0.0.1:25565"},
{raw: "haproxy://127.0.0.1:25565", protocol: ProtocolHAProxy, address: "127.0.0.1:25565"},
{raw: "quic://", protocol: ProtocolQUIC, address: ""},
{raw: "http://127.0.0.1:25565", protocol: ProtocolTCP, address: "http://127.0.0.1:25565"},
}
for _, tt := range tests {
t.Run(tt.raw, func(t *testing.T) {
got := Parse(tt.raw)
if got.Protocol != tt.protocol || got.Address != tt.address {
t.Fatalf("Parse(%q) = %+v, want protocol=%q address=%q", tt.raw, got, tt.protocol, tt.address)
}
})
}
}