feat: add Go backend compatibility service scaffold

This commit is contained in:
sagit
2026-02-06 10:42:48 +00:00
parent f5a40bf530
commit 0f37017760
22 changed files with 2063 additions and 0 deletions
+29
View File
@@ -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
+29
View File
@@ -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
+17
View File
@@ -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"]
+12
View File
@@ -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
+50
View File
@@ -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)
}
}
+23
View File
@@ -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
)
+49
View File
@@ -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=
+53
View File
@@ -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
}
+121
View File
@@ -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))
}
+28
View File
@@ -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
}
+574
View File
@@ -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
}
+117
View File
@@ -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
}
}
@@ -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)
})
}
@@ -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)
})
}
+48
View File
@@ -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)
}
+19
View File
@@ -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
}
+65
View File
@@ -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)
}
+11
View File
@@ -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)
}
@@ -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()
}
+228
View File
@@ -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
}
@@ -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)
}
}
@@ -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))
}
})
}
}