mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
feat: add Go backend compatibility service scaffold
This commit is contained in:
@@ -28,6 +28,35 @@ services:
|
|||||||
retries: 5
|
retries: 5
|
||||||
start_period: 60s
|
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:
|
frontend:
|
||||||
image: ghcr.io/sagit-chu/vite-frontend:2.0.7-beta
|
image: ghcr.io/sagit-chu/vite-frontend:2.0.7-beta
|
||||||
container_name: vite-frontend
|
container_name: vite-frontend
|
||||||
|
|||||||
@@ -28,6 +28,35 @@ services:
|
|||||||
retries: 5
|
retries: 5
|
||||||
start_period: 60s
|
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:
|
frontend:
|
||||||
image: ghcr.io/sagit-chu/vite-frontend:2.0.7-beta
|
image: ghcr.io/sagit-chu/vite-frontend:2.0.7-beta
|
||||||
container_name: vite-frontend
|
container_name: vite-frontend
|
||||||
|
|||||||
@@ -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"]
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
)
|
||||||
@@ -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=
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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))
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -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))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user