diff --git a/docker-compose-v4.yml b/docker-compose-v4.yml index 81fff43..f234503 100644 --- a/docker-compose-v4.yml +++ b/docker-compose-v4.yml @@ -28,6 +28,35 @@ services: retries: 5 start_period: 60s + backend-go: + profiles: ["go-backend"] + build: + context: ./go-backend + container_name: go-backend + restart: unless-stopped + logging: + driver: json-file + options: + max-size: "20m" + environment: + DB_PATH: /app/data/gost.db + JWT_SECRET: ${JWT_SECRET} + LOG_DIR: /app/logs + SERVER_ADDR: :6365 + ports: + - "${GO_BACKEND_PORT:-6366}:6365" + volumes: + - backend_logs:/app/logs + - sqlite_data:/app/data + networks: + - gost-network + healthcheck: + test: ["CMD", "sh", "-c", "wget --no-verbose --tries=1 --spider http://localhost:6365/flow/test || exit 1"] + interval: 30s + timeout: 10s + retries: 5 + start_period: 30s + frontend: image: ghcr.io/sagit-chu/vite-frontend:2.0.7-beta container_name: vite-frontend diff --git a/docker-compose-v6.yml b/docker-compose-v6.yml index 8694a25..fc8795d 100644 --- a/docker-compose-v6.yml +++ b/docker-compose-v6.yml @@ -28,6 +28,35 @@ services: retries: 5 start_period: 60s + backend-go: + profiles: ["go-backend"] + build: + context: ./go-backend + container_name: go-backend + restart: unless-stopped + logging: + driver: json-file + options: + max-size: "20m" + environment: + DB_PATH: /app/data/gost.db + JWT_SECRET: ${JWT_SECRET} + LOG_DIR: /app/logs + SERVER_ADDR: :6365 + ports: + - "${GO_BACKEND_PORT:-6366}:6365" + volumes: + - backend_logs:/app/logs + - sqlite_data:/app/data + networks: + - gost-network + healthcheck: + test: ["CMD", "sh", "-c", "wget --no-verbose --tries=1 --spider http://localhost:6365/flow/test || exit 1"] + interval: 30s + timeout: 10s + retries: 5 + start_period: 30s + frontend: image: ghcr.io/sagit-chu/vite-frontend:2.0.7-beta container_name: vite-frontend diff --git a/go-backend/Dockerfile b/go-backend/Dockerfile new file mode 100644 index 0000000..fcc41d4 --- /dev/null +++ b/go-backend/Dockerfile @@ -0,0 +1,17 @@ +FROM golang:1.23-bookworm AS builder +WORKDIR /src + +COPY go.mod ./ +RUN go mod download + +COPY . . +RUN CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -o /out/paneld ./cmd/paneld + +FROM debian:bookworm-slim +WORKDIR /app +RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates wget && rm -rf /var/lib/apt/lists/* +COPY --from=builder /out/paneld /app/paneld + +ENV SERVER_ADDR=:6365 +EXPOSE 6365 +ENTRYPOINT ["/app/paneld"] diff --git a/go-backend/Makefile b/go-backend/Makefile new file mode 100644 index 0000000..5ab2068 --- /dev/null +++ b/go-backend/Makefile @@ -0,0 +1,12 @@ +GO ?= go + +.PHONY: test build run + +test: + $(GO) test ./... + +build: + $(GO) build ./cmd/paneld + +run: + SERVER_ADDR=:6365 $(GO) run ./cmd/paneld diff --git a/go-backend/cmd/paneld/main.go b/go-backend/cmd/paneld/main.go new file mode 100644 index 0000000..b74d984 --- /dev/null +++ b/go-backend/cmd/paneld/main.go @@ -0,0 +1,50 @@ +package main + +import ( + "context" + "errors" + "log" + "net/http" + "os" + "os/signal" + "syscall" + "time" + + "go-backend/internal/app" + "go-backend/internal/config" +) + +func main() { + cfg := config.FromEnv() + if cfg.JWTSecret == "" { + log.Println("warning: JWT_SECRET is empty") + } + + a, err := app.New(cfg) + if err != nil { + log.Fatalf("failed to create app: %v", err) + } + + errCh := make(chan error, 1) + go func() { + errCh <- a.Run() + }() + + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + + select { + case sig := <-sigCh: + log.Printf("received signal %s, shutting down", sig) + case runErr := <-errCh: + if runErr != nil && !errors.Is(runErr, http.ErrServerClosed) { + log.Fatalf("server stopped unexpectedly: %v", runErr) + } + } + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if err := a.Shutdown(ctx); err != nil { + log.Fatalf("shutdown failed: %v", err) + } +} diff --git a/go-backend/go.mod b/go-backend/go.mod new file mode 100644 index 0000000..2c26ac9 --- /dev/null +++ b/go-backend/go.mod @@ -0,0 +1,23 @@ +module go-backend + +go 1.23.0 + +toolchain go1.24.4 + +require ( + github.com/gorilla/websocket v1.5.3 + modernc.org/sqlite v1.37.1 +) + +require ( + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/ncruces/go-strftime v0.1.9 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect + golang.org/x/sys v0.33.0 // indirect + modernc.org/libc v1.65.7 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect +) diff --git a/go-backend/go.sum b/go-backend/go.sum new file mode 100644 index 0000000..fa6b48b --- /dev/null +++ b/go-backend/go.sum @@ -0,0 +1,49 @@ +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/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/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/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= +github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM= +golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8= +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/sync v0.14.0 h1:woo0S4Yywslg6hp4eUFjTVOyKt0RookbpAHG4c1HmhQ= +golang.org/x/sync v0.14.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA= +golang.org/x/sys v0.6.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/tools v0.33.0 h1:4qz2S3zmRxbGIhDIAgjxvFutSvH5EfnsYrRBj0UI0bc= +golang.org/x/tools v0.33.0/go.mod h1:CIJMaWEY88juyUfo7UbgPqbC8rU2OqfAV1h2Qp0oMYI= +modernc.org/cc/v4 v4.26.1 h1:+X5NtzVBn0KgsBCBe+xkDC7twLb/jNVj9FPgiwSQO3s= +modernc.org/cc/v4 v4.26.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= +modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU= +modernc.org/ccgo/v4 v4.28.0/go.mod h1:JygV3+9AV6SmPhDasu4JgquwU81XAKLd3OKTUDNOiKE= +modernc.org/fileutil v1.3.1 h1:8vq5fe7jdtEvoCf3Zf9Nm0Q05sH6kGx0Op2CPx1wTC8= +modernc.org/fileutil v1.3.1/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/libc v1.65.7 h1:Ia9Z4yzZtWNtUIuiPuQ7Qf7kxYrxP1/jeHZzG8bFu00= +modernc.org/libc v1.65.7/go.mod h1:011EQibzzio/VX3ygj1qGFt5kMjP0lHb0qCW5/D/pQU= +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.37.1 h1:EgHJK/FPoqC+q2YBXg7fUmES37pCHFc97sI7zSayBEs= +modernc.org/sqlite v1.37.1/go.mod h1:XwdRtsE1MpiBcL54+MbKcaDvcuej+IYSMfLN6gSKV8g= +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= diff --git a/go-backend/internal/app/app.go b/go-backend/internal/app/app.go new file mode 100644 index 0000000..42be18e --- /dev/null +++ b/go-backend/internal/app/app.go @@ -0,0 +1,53 @@ +package app + +import ( + "context" + "fmt" + "net/http" + "time" + + "go-backend/internal/config" + httpserver "go-backend/internal/http" + "go-backend/internal/http/handler" + "go-backend/internal/store/sqlite" +) + +type App struct { + cfg config.Config + server *http.Server + repo *sqlite.Repository +} + +func New(cfg config.Config) (*App, error) { + repo, err := sqlite.Open(cfg.DBPath) + if err != nil { + return nil, fmt.Errorf("open sqlite: %w", err) + } + + h := handler.New(repo, cfg.JWTSecret) + router := httpserver.NewRouter(h, cfg.JWTSecret) + + s := &http.Server{ + Addr: cfg.Addr, + Handler: router, + ReadTimeout: 30 * time.Second, + ReadHeaderTimeout: 5 * time.Second, + WriteTimeout: 30 * time.Second, + IdleTimeout: 60 * time.Second, + } + + return &App{cfg: cfg, server: s, repo: repo}, nil +} + +func (a *App) Run() error { + return a.server.ListenAndServe() +} + +func (a *App) Shutdown(ctx context.Context) error { + shutdownErr := a.server.Shutdown(ctx) + closeErr := a.repo.Close() + if shutdownErr != nil { + return shutdownErr + } + return closeErr +} diff --git a/go-backend/internal/auth/jwt.go b/go-backend/internal/auth/jwt.go new file mode 100644 index 0000000..3bc6d9f --- /dev/null +++ b/go-backend/internal/auth/jwt.go @@ -0,0 +1,121 @@ +package auth + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" + "strconv" + "time" +) + +const ( + algorithm = "HmacSHA256" + expireTime = 90 * 24 * time.Hour +) + +type Claims struct { + Sub string `json:"sub"` + Iat int64 `json:"iat"` + Exp int64 `json:"exp"` + User string `json:"user"` + Name string `json:"name"` + RoleID int `json:"role_id"` +} + +type tokenHeader struct { + Alg string `json:"alg"` + Typ string `json:"typ"` +} + +func GenerateToken(userID int64, username string, roleID int, secret string) (string, error) { + now := time.Now() + header := tokenHeader{Alg: algorithm, Typ: "JWT"} + claims := Claims{ + Sub: strconv.FormatInt(userID, 10), + Iat: now.Unix(), + Exp: now.Add(expireTime).Unix(), + User: username, + Name: username, + RoleID: roleID, + } + + headerPart, err := encodeJSON(header) + if err != nil { + return "", err + } + payloadPart, err := encodeJSON(claims) + if err != nil { + return "", err + } + sig := sign(headerPart+"."+payloadPart, secret) + + return headerPart + "." + payloadPart + "." + sig, nil +} + +func ValidateToken(token, secret string) (Claims, bool) { + claims, err := ParseClaims(token, secret) + if err != nil { + return Claims{}, false + } + return claims, true +} + +func ParseClaims(token, secret string) (Claims, error) { + parts := splitToken(token) + if len(parts) != 3 { + return Claims{}, errors.New("invalid token") + } + + signedContent := parts[0] + "." + parts[1] + expected := sign(signedContent, secret) + if !hmac.Equal([]byte(expected), []byte(parts[2])) { + return Claims{}, errors.New("invalid signature") + } + + payloadBytes, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return Claims{}, err + } + + var claims Claims + if err := json.Unmarshal(payloadBytes, &claims); err != nil { + return Claims{}, err + } + + if claims.Exp <= time.Now().Unix() { + return Claims{}, errors.New("token expired") + } + + return claims, nil +} + +func splitToken(token string) []string { + parts := make([]string, 0, 3) + current := "" + for i := 0; i < len(token); i++ { + if token[i] == '.' { + parts = append(parts, current) + current = "" + continue + } + current += string(token[i]) + } + parts = append(parts, current) + return parts +} + +func encodeJSON(v interface{}) (string, error) { + raw, err := json.Marshal(v) + if err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(raw), nil +} + +func sign(content, secret string) string { + h := hmac.New(sha256.New, []byte(secret)) + h.Write([]byte(content)) + return base64.RawURLEncoding.EncodeToString(h.Sum(nil)) +} diff --git a/go-backend/internal/config/config.go b/go-backend/internal/config/config.go new file mode 100644 index 0000000..043730c --- /dev/null +++ b/go-backend/internal/config/config.go @@ -0,0 +1,28 @@ +package config + +import "os" + +type Config struct { + Addr string + DBPath string + JWTSecret string + LogDir string +} + +func FromEnv() Config { + cfg := Config{ + Addr: getEnv("SERVER_ADDR", ":6365"), + DBPath: getEnv("DB_PATH", "/app/data/gost.db"), + JWTSecret: getEnv("JWT_SECRET", ""), + LogDir: getEnv("LOG_DIR", "/app/logs"), + } + + return cfg +} + +func getEnv(key, fallback string) string { + if v := os.Getenv(key); v != "" { + return v + } + return fallback +} diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go new file mode 100644 index 0000000..c2d6209 --- /dev/null +++ b/go-backend/internal/http/handler/handler.go @@ -0,0 +1,574 @@ +package handler + +import ( + "database/sql" + "encoding/json" + "io" + "net/http" + "sort" + "strconv" + "strings" + "time" + + "go-backend/internal/auth" + "go-backend/internal/http/middleware" + "go-backend/internal/http/response" + "go-backend/internal/security" + "go-backend/internal/store/sqlite" + "go-backend/internal/ws" +) + +type Handler struct { + repo *sqlite.Repository + jwtSecret string + wsServer *ws.Server +} + +type loginRequest struct { + Username string `json:"username"` + Password string `json:"password"` + CaptchaID string `json:"captchaId"` +} + +type nameRequest struct { + Name string `json:"name"` +} + +type configSingleRequest struct { + Name string `json:"name"` + Value string `json:"value"` +} + +type changePasswordRequest struct { + NewUsername string `json:"newUsername"` + CurrentPassword string `json:"currentPassword"` + NewPassword string `json:"newPassword"` + ConfirmPassword string `json:"confirmPassword"` +} + +type flowItem struct { + N string `json:"n"` + U int64 `json:"u"` + D int64 `json:"d"` +} + +func New(repo *sqlite.Repository, jwtSecret string) *Handler { + return &Handler{repo: repo, jwtSecret: jwtSecret, wsServer: ws.NewServer(repo, jwtSecret)} +} + +func (h *Handler) WebSocketHandler() http.Handler { + return h.wsServer +} + +func (h *Handler) Register(mux *http.ServeMux) { + mux.HandleFunc("/api/v1/user/login", h.login) + mux.HandleFunc("/api/v1/config/get", h.getConfigByName) + mux.HandleFunc("/api/v1/config/list", h.getConfigs) + mux.HandleFunc("/api/v1/config/update", h.updateConfigs) + mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig) + mux.HandleFunc("/api/v1/captcha/check", h.checkCaptcha) + mux.HandleFunc("/api/v1/user/package", h.userPackage) + mux.HandleFunc("/api/v1/user/updatePassword", h.updatePassword) + + mux.HandleFunc("/flow/test", h.flowTest) + mux.HandleFunc("/flow/config", h.flowConfig) + mux.HandleFunc("/flow/upload", h.flowUpload) +} + +func (h *Handler) login(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var req loginRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.Err(500, "请求参数错误")) + return + } + + if strings.TrimSpace(req.Username) == "" { + response.WriteJSON(w, response.Err(500, "用户名不能为空")) + return + } + if strings.TrimSpace(req.Password) == "" { + response.WriteJSON(w, response.Err(500, "密码不能为空")) + return + } + + captchaEnabled, err := h.captchaEnabled() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if captchaEnabled && strings.TrimSpace(req.CaptchaID) == "" { + response.WriteJSON(w, response.ErrDefault("验证码校验失败")) + return + } + + user, err := h.repo.GetUserByUsername(req.Username) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if user == nil { + response.WriteJSON(w, response.ErrDefault("账号或密码错误")) + return + } + if user.Pwd != security.MD5(req.Password) { + response.WriteJSON(w, response.ErrDefault("账号或密码错误")) + return + } + if user.Status == 0 { + response.WriteJSON(w, response.ErrDefault("账号被停用")) + return + } + + token, err := auth.GenerateToken(user.ID, user.User, user.RoleID, h.jwtSecret) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + requirePasswordChange := req.Username == "admin_user" || req.Password == "admin_user" + response.WriteJSON(w, response.OK(map[string]interface{}{ + "token": token, + "name": user.User, + "role_id": user.RoleID, + "requirePasswordChange": requirePasswordChange, + })) +} + +func (h *Handler) getConfigByName(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var req nameRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("配置名称不能为空")) + return + } + if strings.TrimSpace(req.Name) == "" { + response.WriteJSON(w, response.ErrDefault("配置名称不能为空")) + return + } + + cfg, err := h.repo.GetConfigByName(req.Name) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if cfg == nil { + response.WriteJSON(w, response.ErrDefault("配置不存在")) + return + } + + response.WriteJSON(w, response.OK(cfg)) +} + +func (h *Handler) getConfigs(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + cfgMap, err := h.repo.ListConfigs() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + response.WriteJSON(w, response.OK(cfgMap)) +} + +func (h *Handler) checkCaptcha(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + enabled, err := h.captchaEnabled() + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if enabled { + response.WriteJSON(w, response.OK(1)) + return + } + response.WriteJSON(w, response.OK(0)) +} + +func (h *Handler) flowTest(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + _, _ = w.Write([]byte("test")) +} + +func (h *Handler) flowConfig(w http.ResponseWriter, r *http.Request) { + secret := r.URL.Query().Get("secret") + if ok, _ := h.repo.NodeExistsBySecret(secret); !ok { + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + _, _ = w.Write([]byte("ok")) + return + } + + _, _ = readAndDecryptFlowBody(r.Body, secret) + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + _, _ = w.Write([]byte("ok")) +} + +func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) { + secret := r.URL.Query().Get("secret") + if ok, _ := h.repo.NodeExistsBySecret(secret); !ok { + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + _, _ = w.Write([]byte("ok")) + return + } + + raw, err := readAndDecryptFlowBody(r.Body, secret) + if err == nil && strings.TrimSpace(raw) != "" { + var items []flowItem + if json.Unmarshal([]byte(raw), &items) == nil { + for _, item := range items { + parts := strings.Split(item.N, "_") + if len(parts) < 3 || item.N == "web_api" { + continue + } + forwardID, err1 := strconv.ParseInt(parts[0], 10, 64) + userID, err2 := strconv.ParseInt(parts[1], 10, 64) + userTunnelID, err3 := strconv.ParseInt(parts[2], 10, 64) + if err1 != nil || err2 != nil || err3 != nil { + continue + } + _ = h.repo.AddFlow(forwardID, userID, userTunnelID, item.D, item.U) + } + } + } + + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + _, _ = w.Write([]byte("ok")) +} + +func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var payload map[string]string + if err := decodeJSON(r.Body, &payload); err != nil { + response.WriteJSON(w, response.ErrDefault("配置数据不能为空")) + return + } + if len(payload) == 0 { + response.WriteJSON(w, response.ErrDefault("配置数据不能为空")) + return + } + + now := time.Now().UnixMilli() + for k, v := range payload { + key := strings.TrimSpace(k) + if key == "" { + continue + } + if err := h.repo.UpsertConfig(key, v, now); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + } + + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + var req configSingleRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("配置名称不能为空")) + return + } + if strings.TrimSpace(req.Name) == "" { + response.WriteJSON(w, response.ErrDefault("配置名称不能为空")) + return + } + if strings.TrimSpace(req.Value) == "" { + response.WriteJSON(w, response.ErrDefault("配置值不能为空")) + return + } + + if err := h.repo.UpsertConfig(strings.TrimSpace(req.Name), req.Value, time.Now().UnixMilli()); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + claims, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims) + if !ok { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + + userID, err := parseUserID(claims.Sub) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + + user, err := h.repo.GetUserByID(userID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if user == nil { + response.WriteJSON(w, response.ErrDefault("用户不存在")) + return + } + + tunnels, err := h.repo.GetUserPackageTunnels(userID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + forwards, err := h.repo.GetUserPackageForwards(userID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + stats, err := h.repo.GetStatisticsFlows(userID, 24) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + sort.Slice(stats, func(i, j int) bool { return stats[i].ID < stats[j].ID }) + + tunnelOut := make([]map[string]interface{}, 0, len(tunnels)) + for _, t := range tunnels { + item := map[string]interface{}{ + "id": t.ID, + "userId": t.UserID, + "tunnelId": t.TunnelID, + "tunnelName": t.TunnelName, + "tunnelFlow": t.TunnelFlow, + "flow": t.Flow, + "inFlow": t.InFlow, + "outFlow": t.OutFlow, + "num": t.Num, + "flowResetTime": t.FlowResetTime, + "expTime": t.ExpTime, + "speedId": nil, + "speedLimitName": nil, + "speed": nil, + } + if t.SpeedID.Valid { + item["speedId"] = t.SpeedID.Int64 + } + if t.SpeedLimit.Valid { + item["speedLimitName"] = t.SpeedLimit.String + } + if t.Speed.Valid { + item["speed"] = t.Speed.Int64 + } + tunnelOut = append(tunnelOut, item) + } + + forwardOut := make([]map[string]interface{}, 0, len(forwards)) + for _, f := range forwards { + item := map[string]interface{}{ + "id": f.ID, + "name": f.Name, + "tunnelId": f.TunnelID, + "tunnelName": f.TunnelName, + "inIp": f.InIP, + "inPort": nil, + "remoteAddr": f.RemoteAddr, + "inFlow": f.InFlow, + "outFlow": f.OutFlow, + "status": f.Status, + "createdTime": f.CreatedAt, + } + if f.InPort.Valid { + item["inPort"] = f.InPort.Int64 + } + forwardOut = append(forwardOut, item) + } + + payload := map[string]interface{}{ + "userInfo": map[string]interface{}{ + "id": user.ID, + "name": user.User, + "user": user.User, + "status": user.Status, + "flow": user.Flow, + "inFlow": user.InFlow, + "outFlow": user.OutFlow, + "num": user.Num, + "expTime": user.ExpTime, + "flowResetTime": user.FlowResetTime, + "createdTime": user.CreatedTime, + "updatedTime": nullableNullInt64(user.UpdatedTime), + }, + "tunnelPermissions": tunnelOut, + "forwards": forwardOut, + "statisticsFlows": stats, + } + + response.WriteJSON(w, response.OK(payload)) +} + +func (h *Handler) updatePassword(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + response.WriteJSON(w, response.ErrDefault("请求失败")) + return + } + + claims, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims) + if !ok { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + + userID, err := parseUserID(claims.Sub) + if err != nil { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + + var req changePasswordRequest + if err := decodeJSON(r.Body, &req); err != nil { + response.WriteJSON(w, response.ErrDefault("修改账号密码时发生错误")) + return + } + + if strings.TrimSpace(req.NewUsername) == "" { + response.WriteJSON(w, response.ErrDefault("新用户名不能为空")) + return + } + if strings.TrimSpace(req.CurrentPassword) == "" { + response.WriteJSON(w, response.ErrDefault("当前密码不能为空")) + return + } + if strings.TrimSpace(req.NewPassword) == "" { + response.WriteJSON(w, response.ErrDefault("新密码不能为空")) + return + } + if strings.TrimSpace(req.ConfirmPassword) == "" { + response.WriteJSON(w, response.ErrDefault("确认密码不能为空")) + return + } + if req.NewPassword != req.ConfirmPassword { + response.WriteJSON(w, response.ErrDefault("新密码和确认密码不匹配")) + return + } + + user, err := h.repo.GetUserByID(userID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if user == nil { + response.WriteJSON(w, response.ErrDefault("用户不存在")) + return + } + + if user.Pwd != security.MD5(req.CurrentPassword) { + response.WriteJSON(w, response.ErrDefault("当前密码错误")) + return + } + + exists, err := h.repo.UsernameExistsExceptID(req.NewUsername, userID) + if err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + if exists { + response.WriteJSON(w, response.ErrDefault("用户名已存在")) + return + } + + if err := h.repo.UpdateUserNameAndPassword(userID, req.NewUsername, security.MD5(req.NewPassword), time.Now().UnixMilli()); err != nil { + response.WriteJSON(w, response.Err(-2, err.Error())) + return + } + + response.WriteJSON(w, response.OKEmpty()) +} + +func (h *Handler) captchaEnabled() (bool, error) { + cfg, err := h.repo.GetConfigByName("captcha_enabled") + if err != nil { + return false, err + } + if cfg == nil { + return false, nil + } + return strings.EqualFold(cfg.Value, "true"), nil +} + +func decodeJSON(body io.ReadCloser, out interface{}) error { + defer body.Close() + decoder := json.NewDecoder(body) + decoder.DisallowUnknownFields() + return decoder.Decode(out) +} + +func parseUserID(sub string) (int64, error) { + id, err := strconv.ParseInt(sub, 10, 64) + if err != nil || id <= 0 { + return 0, strconv.ErrSyntax + } + return id, nil +} + +func nullableNullInt64(v sql.NullInt64) interface{} { + if v.Valid { + return v.Int64 + } + return nil +} + +func readAndDecryptFlowBody(body io.ReadCloser, secret string) (string, error) { + defer body.Close() + raw, err := io.ReadAll(body) + if err != nil { + return "", err + } + text := strings.TrimSpace(string(raw)) + if text == "" { + return "", nil + } + + var wrap struct { + Encrypted bool `json:"encrypted"` + Data string `json:"data"` + Timestamp int64 `json:"timestamp"` + } + if err := json.Unmarshal(raw, &wrap); err != nil || !wrap.Encrypted || strings.TrimSpace(wrap.Data) == "" { + return text, nil + } + + crypto, err := security.NewAESCrypto(secret) + if err != nil { + return text, nil + } + plain, err := crypto.Decrypt(wrap.Data) + if err != nil { + return text, nil + } + return string(plain), nil +} diff --git a/go-backend/internal/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go new file mode 100644 index 0000000..6493afa --- /dev/null +++ b/go-backend/internal/http/middleware/auth.go @@ -0,0 +1,117 @@ +package middleware + +import ( + "context" + "net/http" + "strings" + + "go-backend/internal/auth" + "go-backend/internal/http/response" +) + +type contextKey string + +const ClaimsContextKey contextKey = "claims" + +type AuthOptions struct { + JWTSecret string +} + +func JWT(opts AuthOptions) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if shouldSkip(r.URL.Path) { + next.ServeHTTP(w, r) + return + } + + if !strings.HasPrefix(r.URL.Path, "/api/") { + next.ServeHTTP(w, r) + return + } + + token := strings.TrimSpace(r.Header.Get("Authorization")) + if token == "" { + response.WriteJSON(w, response.Err(401, "未登录或token已过期")) + return + } + + claims, ok := auth.ValidateToken(token, opts.JWTSecret) + if !ok { + response.WriteJSON(w, response.Err(401, "无效的token或token已过期")) + return + } + + if requiresAdmin(r.URL.Path) && claims.RoleID != 0 { + response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作")) + return + } + + ctx := context.WithValue(r.Context(), ClaimsContextKey, claims) + next.ServeHTTP(w, r.WithContext(ctx)) + }) + } +} + +func RequireAdmin(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + raw := r.Context().Value(ClaimsContextKey) + claims, ok := raw.(auth.Claims) + if !ok { + response.WriteJSON(w, response.Err(401, "无法获取用户权限信息")) + return + } + if claims.RoleID != 0 { + response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作")) + return + } + next.ServeHTTP(w, r) + }) +} + +func shouldSkip(path string) bool { + switch { + case strings.HasPrefix(path, "/flow/"): + return true + case strings.HasPrefix(path, "/api/v1/open_api/"): + return true + case strings.HasPrefix(path, "/api/v1/captcha/"): + return true + case path == "/api/v1/config/get": + return true + case path == "/api/v1/user/login": + return true + default: + return false + } +} + +func requiresAdmin(path string) bool { + if strings.HasPrefix(path, "/api/v1/group/") { + return true + } + + if strings.HasPrefix(path, "/api/v1/node/") { + return true + } + + if strings.HasPrefix(path, "/api/v1/speed-limit/") { + return true + } + + if strings.HasPrefix(path, "/api/v1/tunnel/") { + if strings.HasPrefix(path, "/api/v1/tunnel/user/tunnel") { + return false + } + return true + } + + switch path { + case "/api/v1/user/create", "/api/v1/user/list", "/api/v1/user/update", "/api/v1/user/delete", "/api/v1/user/reset": + return true + case "/api/v1/config/update", "/api/v1/config/update-single": + return true + default: + return false + } +} diff --git a/go-backend/internal/http/middleware/cors.go b/go-backend/internal/http/middleware/cors.go new file mode 100644 index 0000000..b33100c --- /dev/null +++ b/go-backend/internal/http/middleware/cors.go @@ -0,0 +1,17 @@ +package middleware + +import "net/http" + +func CORS(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Access-Control-Allow-Origin", "*") + w.Header().Set("Access-Control-Allow-Headers", "*") + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, DELETE, PUT, OPTIONS") + w.Header().Set("Access-Control-Expose-Headers", "Authorization") + if r.Method == http.MethodOptions { + w.WriteHeader(http.StatusNoContent) + return + } + next.ServeHTTP(w, r) + }) +} diff --git a/go-backend/internal/http/middleware/recover.go b/go-backend/internal/http/middleware/recover.go new file mode 100644 index 0000000..0728ac5 --- /dev/null +++ b/go-backend/internal/http/middleware/recover.go @@ -0,0 +1,19 @@ +package middleware + +import ( + "fmt" + "net/http" + + "go-backend/internal/http/response" +) + +func Recover(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer func() { + if rec := recover(); rec != nil { + response.WriteJSON(w, response.Err(-2, fmt.Sprint(rec))) + } + }() + next.ServeHTTP(w, r) + }) +} diff --git a/go-backend/internal/http/response/r.go b/go-backend/internal/http/response/r.go new file mode 100644 index 0000000..be21806 --- /dev/null +++ b/go-backend/internal/http/response/r.go @@ -0,0 +1,48 @@ +package response + +import ( + "encoding/json" + "net/http" + "time" +) + +type R struct { + Code int `json:"code"` + Msg string `json:"msg"` + TS int64 `json:"ts"` + Data interface{} `json:"data,omitempty"` +} + +func OK(data interface{}) R { + return R{ + Code: 0, + Msg: "操作成功", + TS: time.Now().UnixMilli(), + Data: data, + } +} + +func OKEmpty() R { + return R{ + Code: 0, + Msg: "操作成功", + TS: time.Now().UnixMilli(), + } +} + +func Err(code int, msg string) R { + return R{ + Code: code, + Msg: msg, + TS: time.Now().UnixMilli(), + } +} + +func ErrDefault(msg string) R { + return Err(-1, msg) +} + +func WriteJSON(w http.ResponseWriter, payload R) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + _ = json.NewEncoder(w).Encode(payload) +} diff --git a/go-backend/internal/http/router.go b/go-backend/internal/http/router.go new file mode 100644 index 0000000..7bc2b8b --- /dev/null +++ b/go-backend/internal/http/router.go @@ -0,0 +1,19 @@ +package httpserver + +import ( + "net/http" + + "go-backend/internal/http/handler" + "go-backend/internal/http/middleware" +) + +func NewRouter(h *handler.Handler, jwtSecret string) http.Handler { + mux := http.NewServeMux() + h.Register(mux) + mux.Handle("/system-info", h.WebSocketHandler()) + + wrapped := middleware.Recover(mux) + wrapped = middleware.JWT(middleware.AuthOptions{JWTSecret: jwtSecret})(wrapped) + wrapped = middleware.CORS(wrapped) + return wrapped +} diff --git a/go-backend/internal/security/aes.go b/go-backend/internal/security/aes.go new file mode 100644 index 0000000..56f1ec4 --- /dev/null +++ b/go-backend/internal/security/aes.go @@ -0,0 +1,65 @@ +package security + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "fmt" +) + +type AESCrypto struct { + key []byte +} + +func NewAESCrypto(secret string) (*AESCrypto, error) { + if secret == "" { + return nil, fmt.Errorf("secret is empty") + } + hash := sha256.Sum256([]byte(secret)) + return &AESCrypto{key: hash[:]}, nil +} + +func (a *AESCrypto) Encrypt(plain []byte) (string, error) { + if len(plain) == 0 { + return "", fmt.Errorf("empty plaintext") + } + block, err := aes.NewCipher(a.key) + if err != nil { + return "", err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return "", err + } + nonce := make([]byte, gcm.NonceSize()) + if _, err := rand.Read(nonce); err != nil { + return "", err + } + sealed := gcm.Seal(nil, nonce, plain, nil) + data := append(nonce, sealed...) + return base64.StdEncoding.EncodeToString(data), nil +} + +func (a *AESCrypto) Decrypt(cipherText string) ([]byte, error) { + raw, err := base64.StdEncoding.DecodeString(cipherText) + if err != nil { + return nil, err + } + block, err := aes.NewCipher(a.key) + if err != nil { + return nil, err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return nil, err + } + nonceSize := gcm.NonceSize() + if len(raw) < nonceSize { + return nil, fmt.Errorf("ciphertext too short") + } + nonce := raw[:nonceSize] + data := raw[nonceSize:] + return gcm.Open(nil, nonce, data, nil) +} diff --git a/go-backend/internal/security/md5.go b/go-backend/internal/security/md5.go new file mode 100644 index 0000000..17263cc --- /dev/null +++ b/go-backend/internal/security/md5.go @@ -0,0 +1,11 @@ +package security + +import ( + "crypto/md5" + "fmt" +) + +func MD5(input string) string { + hash := md5.Sum([]byte(input)) + return fmt.Sprintf("%x", hash) +} diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go new file mode 100644 index 0000000..1cade0a --- /dev/null +++ b/go-backend/internal/store/sqlite/repository.go @@ -0,0 +1,420 @@ +package sqlite + +import ( + "database/sql" + "errors" + "time" + + _ "modernc.org/sqlite" +) + +type Repository struct { + db *sql.DB +} + +type User struct { + ID int64 + User string + Pwd string + RoleID int + ExpTime int64 + Flow int64 + InFlow int64 + OutFlow int64 + FlowResetTime int64 + Num int + CreatedTime int64 + UpdatedTime sql.NullInt64 + Status int +} + +type ViteConfig struct { + ID int64 `json:"id"` + Name string `json:"name"` + Value string `json:"value"` + Time int64 `json:"time"` +} + +type UserTunnelDetail struct { + ID int64 + UserID int64 + TunnelID int64 + TunnelName string + TunnelFlow int + Flow int64 + InFlow int64 + OutFlow int64 + Num int + FlowResetTime int64 + ExpTime int64 + SpeedID sql.NullInt64 + SpeedLimit sql.NullString + Speed sql.NullInt64 +} + +type UserForwardDetail struct { + ID int64 + Name string + TunnelID int64 + TunnelName string + InIP string + InPort sql.NullInt64 + RemoteAddr string + InFlow int64 + OutFlow int64 + Status int + CreatedAt int64 +} + +type StatisticsFlow struct { + ID int64 `json:"id"` + UserID int64 `json:"userId"` + Flow int64 `json:"flow"` + TotalFlow int64 `json:"totalFlow"` + Time string `json:"time"` +} + +type Node struct { + ID int64 + Secret string + Version sql.NullString + HTTP int + TLS int + Socks int + Status int +} + +func Open(path string) (*Repository, error) { + db, err := sql.Open("sqlite", path) + if err != nil { + return nil, err + } + + if err := db.Ping(); err != nil { + _ = db.Close() + return nil, err + } + + return &Repository{db: db}, nil +} + +func (r *Repository) Close() error { + if r == nil || r.db == nil { + return nil + } + return r.db.Close() +} + +func (r *Repository) GetUserByUsername(username string) (*User, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + row := r.db.QueryRow(` + SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status + FROM user WHERE user = ? LIMIT 1 + `, username) + user := &User{} + if err := row.Scan( + &user.ID, &user.User, &user.Pwd, &user.RoleID, &user.ExpTime, + &user.Flow, &user.InFlow, &user.OutFlow, &user.FlowResetTime, + &user.Num, &user.CreatedTime, &user.UpdatedTime, &user.Status, + ); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return user, nil +} + +func (r *Repository) GetConfigByName(name string) (*ViteConfig, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + row := r.db.QueryRow(`SELECT id, name, value, time FROM vite_config WHERE name = ? LIMIT 1`, name) + cfg := &ViteConfig{} + if err := row.Scan(&cfg.ID, &cfg.Name, &cfg.Value, &cfg.Time); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return cfg, nil +} + +func (r *Repository) ListConfigs() (map[string]string, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + rows, err := r.db.Query(`SELECT name, value FROM vite_config`) + if err != nil { + return nil, err + } + defer rows.Close() + + result := make(map[string]string) + for rows.Next() { + var name, value string + if err := rows.Scan(&name, &value); err != nil { + return nil, err + } + result[name] = value + } + if err := rows.Err(); err != nil { + return nil, err + } + return result, nil +} + +func (r *Repository) UpsertConfig(name, value string, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + + _, err := r.db.Exec(` + INSERT INTO vite_config(name, value, time) + VALUES(?, ?, ?) + ON CONFLICT(name) DO UPDATE SET value=excluded.value, time=excluded.time + `, name, value, now) + return err +} + +func (r *Repository) GetUserByID(id int64) (*User, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + row := r.db.QueryRow(` + SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status + FROM user WHERE id = ? LIMIT 1 + `, id) + user := &User{} + if err := row.Scan( + &user.ID, &user.User, &user.Pwd, &user.RoleID, &user.ExpTime, + &user.Flow, &user.InFlow, &user.OutFlow, &user.FlowResetTime, + &user.Num, &user.CreatedTime, &user.UpdatedTime, &user.Status, + ); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return user, nil +} + +func (r *Repository) UsernameExistsExceptID(username string, exceptID int64) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + + row := r.db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, exceptID) + var count int + if err := row.Scan(&count); err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) UpdateUserNameAndPassword(userID int64, username, passwordMD5 string, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + _, err := r.db.Exec(`UPDATE user SET user = ?, pwd = ?, updated_time = ? WHERE id = ?`, username, passwordMD5, now, userID) + return err +} + +func (r *Repository) GetUserPackageTunnels(userID int64) ([]UserTunnelDetail, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + rows, err := r.db.Query(` + SELECT ut.id, ut.user_id, ut.tunnel_id, t.name, t.flow, ut.flow, ut.in_flow, ut.out_flow, + ut.num, ut.flow_reset_time, ut.exp_time, ut.speed_id, sl.name, sl.speed + FROM user_tunnel ut + LEFT JOIN tunnel t ON t.id = ut.tunnel_id + LEFT JOIN speed_limit sl ON sl.id = ut.speed_id + WHERE ut.user_id = ? + ORDER BY ut.id ASC + `, userID) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]UserTunnelDetail, 0) + for rows.Next() { + var item UserTunnelDetail + if err := rows.Scan( + &item.ID, &item.UserID, &item.TunnelID, &item.TunnelName, &item.TunnelFlow, + &item.Flow, &item.InFlow, &item.OutFlow, &item.Num, &item.FlowResetTime, + &item.ExpTime, &item.SpeedID, &item.SpeedLimit, &item.Speed, + ); err != nil { + return nil, err + } + items = append(items, item) + } + + if err := rows.Err(); err != nil { + return nil, err + } + + return items, nil +} + +func (r *Repository) GetUserPackageForwards(userID int64) ([]UserForwardDetail, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + rows, err := r.db.Query(` + SELECT f.id, f.name, f.tunnel_id, t.name, f.remote_addr, f.in_flow, f.out_flow, f.status, f.created_time, + GROUP_CONCAT(n.server_ip || ':' || fp.port), MIN(fp.port) + FROM forward f + LEFT JOIN tunnel t ON t.id = f.tunnel_id + LEFT JOIN forward_port fp ON fp.forward_id = f.id + LEFT JOIN node n ON n.id = fp.node_id + WHERE f.user_id = ? + GROUP BY f.id + ORDER BY f.id ASC + `, userID) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]UserForwardDetail, 0) + for rows.Next() { + var item UserForwardDetail + if err := rows.Scan( + &item.ID, &item.Name, &item.TunnelID, &item.TunnelName, &item.RemoteAddr, + &item.InFlow, &item.OutFlow, &item.Status, &item.CreatedAt, &item.InIP, &item.InPort, + ); err != nil { + return nil, err + } + items = append(items, item) + } + + if err := rows.Err(); err != nil { + return nil, err + } + + return items, nil +} + +func (r *Repository) GetStatisticsFlows(userID int64, limit int) ([]StatisticsFlow, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + rows, err := r.db.Query(` + SELECT id, user_id, flow, total_flow, time + FROM statistics_flow + WHERE user_id = ? + ORDER BY id DESC + LIMIT ? + `, userID, limit) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]StatisticsFlow, 0) + for rows.Next() { + var item StatisticsFlow + if err := rows.Scan(&item.ID, &item.UserID, &item.Flow, &item.TotalFlow, &item.Time); err != nil { + return nil, err + } + items = append(items, item) + } + + if err := rows.Err(); err != nil { + return nil, err + } + + return items, nil +} + +func (r *Repository) NodeExistsBySecret(secret string) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + + row := r.db.QueryRow(`SELECT COUNT(1) FROM node WHERE secret = ?`, secret) + var count int + if err := row.Scan(&count); err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) GetNodeBySecret(secret string) (*Node, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + row := r.db.QueryRow(`SELECT id, secret, version, http, tls, socks, status FROM node WHERE secret = ? LIMIT 1`, secret) + var n Node + if err := row.Scan(&n.ID, &n.Secret, &n.Version, &n.HTTP, &n.TLS, &n.Socks, &n.Status); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return &n, nil +} + +func (r *Repository) UpdateNodeOnline(nodeID int64, status int, version string, httpVal, tlsVal, socksVal int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + _, err := r.db.Exec(`UPDATE node SET status = ?, version = ?, http = ?, tls = ?, socks = ?, updated_time = ? WHERE id = ?`, + status, version, httpVal, tlsVal, socksVal, unixMilliNow(), nodeID) + return err +} + +func (r *Repository) UpdateNodeStatus(nodeID int64, status int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + _, err := r.db.Exec(`UPDATE node SET status = ?, updated_time = ? WHERE id = ?`, status, unixMilliNow(), nodeID) + return err +} + +func (r *Repository) AddFlow(forwardID, userID int64, userTunnelID int64, inFlow, outFlow int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + + tx, err := r.db.Begin() + if err != nil { + return err + } + defer func() { + if err != nil { + _ = tx.Rollback() + } + }() + + if _, err = tx.Exec(`UPDATE forward SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, forwardID); err != nil { + return err + } + if _, err = tx.Exec(`UPDATE user SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, userID); err != nil { + return err + } + if userTunnelID > 0 { + if _, err = tx.Exec(`UPDATE user_tunnel SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, userTunnelID); err != nil { + return err + } + } + + err = tx.Commit() + return err +} + +func unixMilliNow() int64 { + return time.Now().UnixMilli() +} diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go new file mode 100644 index 0000000..a12d3db --- /dev/null +++ b/go-backend/internal/ws/server.go @@ -0,0 +1,228 @@ +package ws + +import ( + "encoding/json" + "log" + "net/http" + "strconv" + "strings" + "sync" + + "github.com/gorilla/websocket" + + "go-backend/internal/auth" + "go-backend/internal/security" + "go-backend/internal/store/sqlite" +) + +type encryptedMessage struct { + Encrypted bool `json:"encrypted"` + Data string `json:"data"` + Timestamp int64 `json:"timestamp"` +} + +type broadcastMessage struct { + ID int64 `json:"id"` + Type string `json:"type"` + Data string `json:"data"` +} + +type connWrap struct { + conn *websocket.Conn + mu sync.Mutex +} + +type nodeSession struct { + nodeID int64 + secret string + conn *connWrap +} + +type Server struct { + repo *sqlite.Repository + jwtSecret string + upgrader websocket.Upgrader + + mu sync.RWMutex + admins map[*connWrap]struct{} + nodes map[int64]*nodeSession + byConn map[*websocket.Conn]*nodeSession +} + +func NewServer(repo *sqlite.Repository, jwtSecret string) *Server { + return &Server{ + repo: repo, + jwtSecret: jwtSecret, + upgrader: websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { return true }, + }, + admins: make(map[*connWrap]struct{}), + nodes: make(map[int64]*nodeSession), + byConn: make(map[*websocket.Conn]*nodeSession), + } +} + +func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + query := r.URL.Query() + typeVal := query.Get("type") + secret := query.Get("secret") + + if typeVal == "1" { + node, err := s.repo.GetNodeBySecret(secret) + if err != nil || node == nil { + http.Error(w, "forbidden", http.StatusForbidden) + return + } + s.handleNode(w, r, node.ID, secret) + return + } + + if typeVal == "0" { + if _, ok := auth.ValidateToken(secret, s.jwtSecret); !ok { + http.Error(w, "forbidden", http.StatusForbidden) + return + } + s.handleAdmin(w, r) + return + } + + http.Error(w, "bad request", http.StatusBadRequest) +} + +func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) { + conn, err := s.upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + cw := &connWrap{conn: conn} + + s.mu.Lock() + s.admins[cw] = struct{}{} + s.mu.Unlock() + + defer func() { + s.mu.Lock() + delete(s.admins, cw) + s.mu.Unlock() + _ = conn.Close() + }() + + for { + if _, _, err := conn.ReadMessage(); err != nil { + return + } + } +} + +func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64, secret string) { + conn, err := s.upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + cw := &connWrap{conn: conn} + + version := r.URL.Query().Get("version") + httpVal := parseIntDefault(r.URL.Query().Get("http"), 0) + tlsVal := parseIntDefault(r.URL.Query().Get("tls"), 0) + socksVal := parseIntDefault(r.URL.Query().Get("socks"), 0) + + s.mu.Lock() + if old, ok := s.nodes[nodeID]; ok { + _ = old.conn.conn.Close() + delete(s.byConn, old.conn.conn) + } + ns := &nodeSession{nodeID: nodeID, secret: secret, conn: cw} + s.nodes[nodeID] = ns + s.byConn[conn] = ns + s.mu.Unlock() + + _ = s.repo.UpdateNodeOnline(nodeID, 1, version, httpVal, tlsVal, socksVal) + s.broadcastStatus(nodeID, 1) + + defer func() { + needOfflineBroadcast := false + s.mu.Lock() + current, ok := s.nodes[nodeID] + if ok && current.conn.conn == conn { + delete(s.nodes, nodeID) + needOfflineBroadcast = true + } + delete(s.byConn, conn) + s.mu.Unlock() + if needOfflineBroadcast { + _ = s.repo.UpdateNodeStatus(nodeID, 0) + s.broadcastStatus(nodeID, 0) + } + _ = conn.Close() + }() + + for { + _, payload, err := conn.ReadMessage() + if err != nil { + return + } + + msg := decryptIfNeeded(payload, secret) + s.broadcastInfo(nodeID, msg) + } +} + +func (s *Server) broadcastStatus(nodeID int64, status int) { + payload := map[string]interface{}{ + "id": strconv.FormatInt(nodeID, 10), + "type": "status", + "data": status, + } + raw, _ := json.Marshal(payload) + s.broadcastToAdmins(string(raw)) +} + +func (s *Server) broadcastInfo(nodeID int64, data string) { + payload := broadcastMessage{ID: nodeID, Type: "info", Data: data} + raw, _ := json.Marshal(payload) + s.broadcastToAdmins(string(raw)) +} + +func (s *Server) broadcastToAdmins(message string) { + s.mu.RLock() + admins := make([]*connWrap, 0, len(s.admins)) + for c := range s.admins { + admins = append(admins, c) + } + s.mu.RUnlock() + + for _, c := range admins { + c.mu.Lock() + err := c.conn.WriteMessage(websocket.TextMessage, []byte(message)) + c.mu.Unlock() + if err != nil { + log.Printf("websocket broadcast failed: %v", err) + } + } +} + +func decryptIfNeeded(payload []byte, secret string) string { + text := string(payload) + var wrap encryptedMessage + if err := json.Unmarshal(payload, &wrap); err != nil || !wrap.Encrypted || strings.TrimSpace(wrap.Data) == "" { + return text + } + + crypto, err := security.NewAESCrypto(secret) + if err != nil { + return text + } + plain, err := crypto.Decrypt(wrap.Data) + if err != nil { + return text + } + return string(plain) +} + +func parseIntDefault(v string, fallback int) int { + x, err := strconv.Atoi(v) + if err != nil { + return fallback + } + return x +} diff --git a/go-backend/tests/contract/auth_contract_test.go b/go-backend/tests/contract/auth_contract_test.go new file mode 100644 index 0000000..43d01c8 --- /dev/null +++ b/go-backend/tests/contract/auth_contract_test.go @@ -0,0 +1,90 @@ +package contract_test + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "go-backend/internal/auth" + "go-backend/internal/http/middleware" + "go-backend/internal/http/response" +) + +func TestJWTMiddlewareContracts(t *testing.T) { + secret := "unit-test-secret" + + next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + response.WriteJSON(w, response.OK("pass")) + }) + + wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret})(next) + + t.Run("login path is excluded", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", nil) + res := httptest.NewRecorder() + wrapped.ServeHTTP(res, req) + assertCode(t, res, 0) + }) + + t.Run("missing token returns 401 contract message", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil) + res := httptest.NewRecorder() + wrapped.ServeHTTP(res, req) + assertCodeMsg(t, res, 401, "未登录或token已过期") + }) + + t.Run("invalid token returns 401 contract message", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil) + req.Header.Set("Authorization", "invalid.token.value") + res := httptest.NewRecorder() + wrapped.ServeHTTP(res, req) + assertCodeMsg(t, res, 401, "无效的token或token已过期") + }) + + t.Run("valid token reaches next", func(t *testing.T) { + token, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil) + req.Header.Set("Authorization", token) + res := httptest.NewRecorder() + wrapped.ServeHTTP(res, req) + assertCode(t, res, 0) + }) + + t.Run("non-admin blocked on admin path", func(t *testing.T) { + token, err := auth.GenerateToken(2, "normal_user", 1, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", nil) + req.Header.Set("Authorization", token) + res := httptest.NewRecorder() + wrapped.ServeHTTP(res, req) + assertCodeMsg(t, res, 403, "权限不足,仅管理员可操作") + }) +} + +func assertCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) { + t.Helper() + var out response.R + if err := json.NewDecoder(rec.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != expected { + t.Fatalf("expected code %d, got %d", expected, out.Code) + } +} + +func assertCodeMsg(t *testing.T, rec *httptest.ResponseRecorder, expectedCode int, expectedMsg string) { + t.Helper() + var out response.R + if err := json.NewDecoder(rec.Body).Decode(&out); err != nil { + t.Fatalf("decode response: %v", err) + } + if out.Code != expectedCode || out.Msg != expectedMsg { + t.Fatalf("expected (%d,%q), got (%d,%q)", expectedCode, expectedMsg, out.Code, out.Msg) + } +} diff --git a/go-backend/tests/contract/flow_contract_test.go b/go-backend/tests/contract/flow_contract_test.go new file mode 100644 index 0000000..f45f247 --- /dev/null +++ b/go-backend/tests/contract/flow_contract_test.go @@ -0,0 +1,44 @@ +package contract_test + +import ( + "io" + "net/http" + "net/http/httptest" + "testing" + + "go-backend/internal/http/handler" +) + +func TestFlowEndpointsStringResponses(t *testing.T) { + h := handler.New(nil, "secret") + mux := http.NewServeMux() + h.Register(mux) + + tests := []struct { + name string + method string + path string + expected string + }{ + {name: "flow test", method: http.MethodGet, path: "/flow/test", expected: "test"}, + {name: "flow config", method: http.MethodPost, path: "/flow/config?secret=abc", expected: "ok"}, + {name: "flow upload", method: http.MethodPost, path: "/flow/upload?secret=abc", expected: "ok"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + req := httptest.NewRequest(tc.method, tc.path, nil) + res := httptest.NewRecorder() + mux.ServeHTTP(res, req) + + body, err := io.ReadAll(res.Body) + if err != nil { + t.Fatalf("read body: %v", err) + } + + if string(body) != tc.expected { + t.Fatalf("expected %q, got %q", tc.expected, string(body)) + } + }) + } +}