diff --git a/.github/workflows/ci-build.yml b/.github/workflows/ci-build.yml
index 4d11043..89efde9 100644
--- a/.github/workflows/ci-build.yml
+++ b/.github/workflows/ci-build.yml
@@ -28,23 +28,25 @@ jobs:
run: npm run build
backend:
- name: Build Backend
+ name: Build Go Backend
runs-on: ubuntu-latest
defaults:
run:
- working-directory: springboot-backend
+ working-directory: go-backend
steps:
- uses: actions/checkout@v4
- - name: Setup Java 21
- uses: actions/setup-java@v4
+ - name: Setup Go
+ uses: actions/setup-go@v5
with:
- java-version: '21'
- distribution: 'temurin'
- cache: 'maven'
+ go-version: '1.23'
+ cache-dependency-path: go-backend/go.sum
- - name: Build with Maven
- run: mvn clean package -DskipTests
+ - name: Download dependencies
+ run: go mod download
+
+ - name: Build
+ run: go build -v ./...
agent:
name: Build Agent
diff --git a/.github/workflows/docker-build.yml b/.github/workflows/docker-build.yml
index 995ca65..6b205ef 100644
--- a/.github/workflows/docker-build.yml
+++ b/.github/workflows/docker-build.yml
@@ -165,8 +165,8 @@ jobs:
-t ${{ env.REGISTRY }}/${OWNER}/vite-frontend:${VERSION} \
./vite-frontend
- build-java:
- name: Build & Push Spring Boot Backend
+ build-go-backend:
+ name: Build & Push Go Backend
needs: check-version
if: needs.check-version.outputs.should_build == 'true'
runs-on: ubuntu-latest
@@ -176,22 +176,23 @@ jobs:
steps:
- uses: actions/checkout@v4
- - name: Set up JDK and Maven
- uses: actions/setup-java@v4
+ - name: Set up Go
+ uses: actions/setup-go@v5
with:
- java-version: 21
- distribution: 'temurin'
+ go-version: '1.23'
- - name: Cache Maven dependencies
+ - name: Cache Go dependencies
uses: actions/cache@v4
with:
- path: ~/.m2
- key: ${{ runner.os }}-m2-${{ hashFiles('**/pom.xml') }}
- restore-keys: ${{ runner.os }}-m2
+ path: |
+ ~/.cache/go-build
+ ~/go/pkg/mod
+ key: ${{ runner.os }}-go-backend-${{ hashFiles('go-backend/go.sum') }}
+ restore-keys: ${{ runner.os }}-go-backend-
- - name: Build Java JAR
- working-directory: ./springboot-backend
- run: mvn clean package -DskipTests
+ - name: Download dependencies
+ working-directory: ./go-backend
+ run: go mod download
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
@@ -203,7 +204,7 @@ jobs:
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- - name: Build and push Java Docker images
+ - name: Build and push Go backend Docker images
run: |
VERSION="${{ needs.check-version.outputs.version }}"
OWNER="${{ needs.check-version.outputs.image_owner }}"
@@ -211,13 +212,13 @@ jobs:
docker buildx build \
--platform linux/amd64,linux/arm64 \
--push \
- -t ${{ env.REGISTRY }}/${OWNER}/springboot-backend:latest \
- -t ${{ env.REGISTRY }}/${OWNER}/springboot-backend:${VERSION} \
- ./springboot-backend
+ -t ${{ env.REGISTRY }}/${OWNER}/go-backend:latest \
+ -t ${{ env.REGISTRY }}/${OWNER}/go-backend:${VERSION} \
+ ./go-backend
create-release:
name: Create Release (Tag Only)
- needs: [check-version, build-gost, build-vite, build-java]
+ needs: [check-version, build-gost, build-vite, build-go-backend]
if: needs.check-version.outputs.is_tag == 'true'
runs-on: ubuntu-latest
permissions:
@@ -252,10 +253,10 @@ jobs:
cp docker-compose-v6.yml ./artifacts/docker-compose-v6.yml
# 替换镜像地址为 GHCR
- sed -i "s|bqlpfy/springboot-backend:[^[:space:]]*|${{ env.REGISTRY }}/${OWNER}/springboot-backend:${VERSION}|g" ./artifacts/docker-compose-v4.yml
- sed -i "s|bqlpfy/vite-frontend:[^[:space:]]*|${{ env.REGISTRY }}/${OWNER}/vite-frontend:${VERSION}|g" ./artifacts/docker-compose-v4.yml
- sed -i "s|bqlpfy/springboot-backend:[^[:space:]]*|${{ env.REGISTRY }}/${OWNER}/springboot-backend:${VERSION}|g" ./artifacts/docker-compose-v6.yml
- sed -i "s|bqlpfy/vite-frontend:[^[:space:]]*|${{ env.REGISTRY }}/${OWNER}/vite-frontend:${VERSION}|g" ./artifacts/docker-compose-v6.yml
+ sed -i "s|image: .*go-backend:[^[:space:]]*|image: ${{ env.REGISTRY }}/${OWNER}/go-backend:${VERSION}|g" ./artifacts/docker-compose-v4.yml
+ sed -i "s|image: .*vite-frontend:[^[:space:]]*|image: ${{ env.REGISTRY }}/${OWNER}/vite-frontend:${VERSION}|g" ./artifacts/docker-compose-v4.yml
+ sed -i "s|image: .*go-backend:[^[:space:]]*|image: ${{ env.REGISTRY }}/${OWNER}/go-backend:${VERSION}|g" ./artifacts/docker-compose-v6.yml
+ sed -i "s|image: .*vite-frontend:[^[:space:]]*|image: ${{ env.REGISTRY }}/${OWNER}/vite-frontend:${VERSION}|g" ./artifacts/docker-compose-v6.yml
# 复制并修改安装脚本
cp install.sh ./artifacts/install.sh
@@ -294,7 +295,7 @@ jobs:
\`\`\`bash
# Backend
- docker pull ${{ env.REGISTRY }}/${OWNER}/springboot-backend:${VERSION}
+ docker pull ${{ env.REGISTRY }}/${OWNER}/go-backend:${VERSION}
# Frontend
docker pull ${{ env.REGISTRY }}/${OWNER}/vite-frontend:${VERSION}
diff --git a/.gitignore b/.gitignore
index b21deae..31b9b70 100644
--- a/.gitignore
+++ b/.gitignore
@@ -257,4 +257,7 @@ gitee/
doraemon.jks
device.id
commit.sh
-sql/
\ No newline at end of file
+sql/
+!go-backend/internal/store/sqlite/sql/
+!go-backend/internal/store/sqlite/sql/schema.sql
+!go-backend/internal/store/sqlite/sql/data.sql
diff --git a/3x_base.go b/3x_base.go
new file mode 100644
index 0000000..7bc61b6
--- /dev/null
+++ b/3x_base.go
@@ -0,0 +1,42 @@
+// Package controller provides HTTP request handlers and controllers for the 3x-ui web management panel.
+// It handles routing, authentication, and API endpoints for managing Xray inbounds, settings, and more.
+package controller
+
+import (
+ "net/http"
+
+ "github.com/mhsanaei/3x-ui/v2/logger"
+ "github.com/mhsanaei/3x-ui/v2/web/locale"
+ "github.com/mhsanaei/3x-ui/v2/web/session"
+
+ "github.com/gin-gonic/gin"
+)
+
+// BaseController provides common functionality for all controllers, including authentication checks.
+type BaseController struct{}
+
+// checkLogin is a middleware that verifies user authentication and handles unauthorized access.
+func (a *BaseController) checkLogin(c *gin.Context) {
+ if !session.IsLogin(c) {
+ if isAjax(c) {
+ pureJsonMsg(c, http.StatusUnauthorized, false, I18nWeb(c, "pages.login.loginAgain"))
+ } else {
+ c.Redirect(http.StatusTemporaryRedirect, c.GetString("base_path"))
+ }
+ c.Abort()
+ } else {
+ c.Next()
+ }
+}
+
+// I18nWeb retrieves an internationalized message for the web interface based on the current locale.
+func I18nWeb(c *gin.Context, name string, params ...string) string {
+ anyfunc, funcExists := c.Get("I18n")
+ if !funcExists {
+ logger.Warning("I18n function not exists in gin context!")
+ return ""
+ }
+ i18nFunc, _ := anyfunc.(func(i18nType locale.I18nType, key string, keyParams ...string) string)
+ msg := i18nFunc(locale.Web, name, params...)
+ return msg
+}
diff --git a/3x_inbound.go b/3x_inbound.go
new file mode 100644
index 0000000..8317de3
--- /dev/null
+++ b/3x_inbound.go
@@ -0,0 +1,424 @@
+package controller
+
+import (
+ "encoding/json"
+ "fmt"
+ "strconv"
+
+ "github.com/mhsanaei/3x-ui/v2/database/model"
+ "github.com/mhsanaei/3x-ui/v2/web/service"
+ "github.com/mhsanaei/3x-ui/v2/web/session"
+ "github.com/mhsanaei/3x-ui/v2/web/websocket"
+
+ "github.com/gin-gonic/gin"
+)
+
+// InboundController handles HTTP requests related to Xray inbounds management.
+type InboundController struct {
+ inboundService service.InboundService
+ xrayService service.XrayService
+}
+
+// NewInboundController creates a new InboundController and sets up its routes.
+func NewInboundController(g *gin.RouterGroup) *InboundController {
+ a := &InboundController{}
+ a.initRouter(g)
+ return a
+}
+
+// initRouter initializes the routes for inbound-related operations.
+func (a *InboundController) initRouter(g *gin.RouterGroup) {
+
+ g.GET("/list", a.getInbounds)
+ g.GET("/get/:id", a.getInbound)
+ g.GET("/getClientTraffics/:email", a.getClientTraffics)
+ g.GET("/getClientTrafficsById/:id", a.getClientTrafficsById)
+
+ g.POST("/add", a.addInbound)
+ g.POST("/del/:id", a.delInbound)
+ g.POST("/update/:id", a.updateInbound)
+ g.POST("/clientIps/:email", a.getClientIps)
+ g.POST("/clearClientIps/:email", a.clearClientIps)
+ g.POST("/addClient", a.addInboundClient)
+ g.POST("/:id/delClient/:clientId", a.delInboundClient)
+ g.POST("/updateClient/:clientId", a.updateInboundClient)
+ g.POST("/:id/resetClientTraffic/:email", a.resetClientTraffic)
+ g.POST("/resetAllTraffics", a.resetAllTraffics)
+ g.POST("/resetAllClientTraffics/:id", a.resetAllClientTraffics)
+ g.POST("/delDepletedClients/:id", a.delDepletedClients)
+ g.POST("/import", a.importInbound)
+ g.POST("/onlines", a.onlines)
+ g.POST("/lastOnline", a.lastOnline)
+ g.POST("/updateClientTraffic/:email", a.updateClientTraffic)
+ g.POST("/:id/delClientByEmail/:email", a.delInboundClientByEmail)
+}
+
+// getInbounds retrieves the list of inbounds for the logged-in user.
+func (a *InboundController) getInbounds(c *gin.Context) {
+ user := session.GetLoginUser(c)
+ inbounds, err := a.inboundService.GetInbounds(user.Id)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.obtain"), err)
+ return
+ }
+ jsonObj(c, inbounds, nil)
+}
+
+// getInbound retrieves a specific inbound by its ID.
+func (a *InboundController) getInbound(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "get"), err)
+ return
+ }
+ inbound, err := a.inboundService.GetInbound(id)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.obtain"), err)
+ return
+ }
+ jsonObj(c, inbound, nil)
+}
+
+// getClientTraffics retrieves client traffic information by email.
+func (a *InboundController) getClientTraffics(c *gin.Context) {
+ email := c.Param("email")
+ clientTraffics, err := a.inboundService.GetClientTrafficByEmail(email)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.trafficGetError"), err)
+ return
+ }
+ jsonObj(c, clientTraffics, nil)
+}
+
+// getClientTrafficsById retrieves client traffic information by inbound ID.
+func (a *InboundController) getClientTrafficsById(c *gin.Context) {
+ id := c.Param("id")
+ clientTraffics, err := a.inboundService.GetClientTrafficByID(id)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.trafficGetError"), err)
+ return
+ }
+ jsonObj(c, clientTraffics, nil)
+}
+
+// addInbound creates a new inbound configuration.
+func (a *InboundController) addInbound(c *gin.Context) {
+ inbound := &model.Inbound{}
+ err := c.ShouldBind(inbound)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundCreateSuccess"), err)
+ return
+ }
+ user := session.GetLoginUser(c)
+ inbound.UserId = user.Id
+ if inbound.Listen == "" || inbound.Listen == "0.0.0.0" || inbound.Listen == "::" || inbound.Listen == "::0" {
+ inbound.Tag = fmt.Sprintf("inbound-%v", inbound.Port)
+ } else {
+ inbound.Tag = fmt.Sprintf("inbound-%v:%v", inbound.Listen, inbound.Port)
+ }
+
+ inbound, needRestart, err := a.inboundService.AddInbound(inbound)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsgObj(c, I18nWeb(c, "pages.inbounds.toasts.inboundCreateSuccess"), inbound, nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+ // Broadcast inbounds update via WebSocket
+ inbounds, _ := a.inboundService.GetInbounds(user.Id)
+ websocket.BroadcastInbounds(inbounds)
+}
+
+// delInbound deletes an inbound configuration by its ID.
+func (a *InboundController) delInbound(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundDeleteSuccess"), err)
+ return
+ }
+ needRestart, err := a.inboundService.DelInbound(id)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsgObj(c, I18nWeb(c, "pages.inbounds.toasts.inboundDeleteSuccess"), id, nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+ // Broadcast inbounds update via WebSocket
+ user := session.GetLoginUser(c)
+ inbounds, _ := a.inboundService.GetInbounds(user.Id)
+ websocket.BroadcastInbounds(inbounds)
+}
+
+// updateInbound updates an existing inbound configuration.
+func (a *InboundController) updateInbound(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+ inbound := &model.Inbound{
+ Id: id,
+ }
+ err = c.ShouldBind(inbound)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+ inbound, needRestart, err := a.inboundService.UpdateInbound(inbound)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsgObj(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), inbound, nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+ // Broadcast inbounds update via WebSocket
+ user := session.GetLoginUser(c)
+ inbounds, _ := a.inboundService.GetInbounds(user.Id)
+ websocket.BroadcastInbounds(inbounds)
+}
+
+// getClientIps retrieves the IP addresses associated with a client by email.
+func (a *InboundController) getClientIps(c *gin.Context) {
+ email := c.Param("email")
+
+ ips, err := a.inboundService.GetInboundClientIps(email)
+ if err != nil || ips == "" {
+ jsonObj(c, "No IP Record", nil)
+ return
+ }
+
+ jsonObj(c, ips, nil)
+}
+
+// clearClientIps clears the IP addresses for a client by email.
+func (a *InboundController) clearClientIps(c *gin.Context) {
+ email := c.Param("email")
+
+ err := a.inboundService.ClearClientIps(email)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.updateSuccess"), err)
+ return
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.logCleanSuccess"), nil)
+}
+
+// addInboundClient adds a new client to an existing inbound.
+func (a *InboundController) addInboundClient(c *gin.Context) {
+ data := &model.Inbound{}
+ err := c.ShouldBind(data)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+
+ needRestart, err := a.inboundService.AddInboundClient(data)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundClientAddSuccess"), nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+}
+
+// delInboundClient deletes a client from an inbound by inbound ID and client ID.
+func (a *InboundController) delInboundClient(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+ clientId := c.Param("clientId")
+
+ needRestart, err := a.inboundService.DelInboundClient(id, clientId)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundClientDeleteSuccess"), nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+}
+
+// updateInboundClient updates a client's configuration in an inbound.
+func (a *InboundController) updateInboundClient(c *gin.Context) {
+ clientId := c.Param("clientId")
+
+ inbound := &model.Inbound{}
+ err := c.ShouldBind(inbound)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+
+ needRestart, err := a.inboundService.UpdateInboundClient(inbound, clientId)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundClientUpdateSuccess"), nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+}
+
+// resetClientTraffic resets the traffic counter for a specific client in an inbound.
+func (a *InboundController) resetClientTraffic(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+ email := c.Param("email")
+
+ needRestart, err := a.inboundService.ResetClientTraffic(id, email)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.resetInboundClientTrafficSuccess"), nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+}
+
+// resetAllTraffics resets all traffic counters across all inbounds.
+func (a *InboundController) resetAllTraffics(c *gin.Context) {
+ err := a.inboundService.ResetAllTraffics()
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ } else {
+ a.xrayService.SetToNeedRestart()
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.resetAllTrafficSuccess"), nil)
+}
+
+// resetAllClientTraffics resets traffic counters for all clients in a specific inbound.
+func (a *InboundController) resetAllClientTraffics(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+
+ err = a.inboundService.ResetAllClientTraffics(id)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ } else {
+ a.xrayService.SetToNeedRestart()
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.resetAllClientTrafficSuccess"), nil)
+}
+
+// importInbound imports an inbound configuration from provided data.
+func (a *InboundController) importInbound(c *gin.Context) {
+ inbound := &model.Inbound{}
+ err := json.Unmarshal([]byte(c.PostForm("data")), inbound)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ user := session.GetLoginUser(c)
+ inbound.Id = 0
+ inbound.UserId = user.Id
+ if inbound.Listen == "" || inbound.Listen == "0.0.0.0" || inbound.Listen == "::" || inbound.Listen == "::0" {
+ inbound.Tag = fmt.Sprintf("inbound-%v", inbound.Port)
+ } else {
+ inbound.Tag = fmt.Sprintf("inbound-%v:%v", inbound.Listen, inbound.Port)
+ }
+
+ for index := range inbound.ClientStats {
+ inbound.ClientStats[index].Id = 0
+ inbound.ClientStats[index].Enable = true
+ }
+
+ needRestart := false
+ inbound, needRestart, err = a.inboundService.AddInbound(inbound)
+ jsonMsgObj(c, I18nWeb(c, "pages.inbounds.toasts.inboundCreateSuccess"), inbound, err)
+ if err == nil && needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+}
+
+// delDepletedClients deletes clients in an inbound who have exhausted their traffic limits.
+func (a *InboundController) delDepletedClients(c *gin.Context) {
+ id, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+ err = a.inboundService.DelDepletedClients(id)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.delDepletedClientsSuccess"), nil)
+}
+
+// onlines retrieves the list of currently online clients.
+func (a *InboundController) onlines(c *gin.Context) {
+ jsonObj(c, a.inboundService.GetOnlineClients(), nil)
+}
+
+// lastOnline retrieves the last online timestamps for clients.
+func (a *InboundController) lastOnline(c *gin.Context) {
+ data, err := a.inboundService.GetClientsLastOnline()
+ jsonObj(c, data, err)
+}
+
+// updateClientTraffic updates the traffic statistics for a client by email.
+func (a *InboundController) updateClientTraffic(c *gin.Context) {
+ email := c.Param("email")
+
+ // Define the request structure for traffic update
+ type TrafficUpdateRequest struct {
+ Upload int64 `json:"upload"`
+ Download int64 `json:"download"`
+ }
+
+ var request TrafficUpdateRequest
+ err := c.ShouldBindJSON(&request)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundUpdateSuccess"), err)
+ return
+ }
+
+ err = a.inboundService.UpdateClientTrafficByEmail(email, request.Upload, request.Download)
+ if err != nil {
+ jsonMsg(c, I18nWeb(c, "somethingWentWrong"), err)
+ return
+ }
+
+ jsonMsg(c, I18nWeb(c, "pages.inbounds.toasts.inboundClientUpdateSuccess"), nil)
+}
+
+// delInboundClientByEmail deletes a client from an inbound by email address.
+func (a *InboundController) delInboundClientByEmail(c *gin.Context) {
+ inboundId, err := strconv.Atoi(c.Param("id"))
+ if err != nil {
+ jsonMsg(c, "Invalid inbound ID", err)
+ return
+ }
+
+ email := c.Param("email")
+ needRestart, err := a.inboundService.DelInboundClientByEmail(inboundId, email)
+ if err != nil {
+ jsonMsg(c, "Failed to delete client by email", err)
+ return
+ }
+
+ jsonMsg(c, "Client deleted successfully", nil)
+ if needRestart {
+ a.xrayService.SetToNeedRestart()
+ }
+}
diff --git a/3x_xui.go b/3x_xui.go
new file mode 100644
index 0000000..5150290
--- /dev/null
+++ b/3x_xui.go
@@ -0,0 +1,54 @@
+package controller
+
+import (
+ "github.com/gin-gonic/gin"
+)
+
+// XUIController is the main controller for the X-UI panel, managing sub-controllers.
+type XUIController struct {
+ BaseController
+
+ settingController *SettingController
+ xraySettingController *XraySettingController
+}
+
+// NewXUIController creates a new XUIController and initializes its routes.
+func NewXUIController(g *gin.RouterGroup) *XUIController {
+ a := &XUIController{}
+ a.initRouter(g)
+ return a
+}
+
+// initRouter sets up the main panel routes and initializes sub-controllers.
+func (a *XUIController) initRouter(g *gin.RouterGroup) {
+ g = g.Group("/panel")
+ g.Use(a.checkLogin)
+
+ g.GET("/", a.index)
+ g.GET("/inbounds", a.inbounds)
+ g.GET("/settings", a.settings)
+ g.GET("/xray", a.xraySettings)
+
+ a.settingController = NewSettingController(g)
+ a.xraySettingController = NewXraySettingController(g)
+}
+
+// index renders the main panel index page.
+func (a *XUIController) index(c *gin.Context) {
+ html(c, "index.html", "pages.index.title", nil)
+}
+
+// inbounds renders the inbounds management page.
+func (a *XUIController) inbounds(c *gin.Context) {
+ html(c, "inbounds.html", "pages.inbounds.title", nil)
+}
+
+// settings renders the settings management page.
+func (a *XUIController) settings(c *gin.Context) {
+ html(c, "settings.html", "pages.settings.title", nil)
+}
+
+// xraySettings renders the Xray settings page.
+func (a *XUIController) xraySettings(c *gin.Context) {
+ html(c, "xray.html", "pages.xray.title", nil)
+}
diff --git a/docker-compose-v4.yml b/docker-compose-v4.yml
index c0d3527..7449ab2 100644
--- a/docker-compose-v4.yml
+++ b/docker-compose-v4.yml
@@ -1,7 +1,7 @@
services:
backend:
- image: ghcr.io/sagit-chu/springboot-backend:${FLUX_VERSION:-latest}
- container_name: springboot-backend
+ image: ghcr.io/sagit-chu/go-backend:${FLUX_VERSION:-latest}
+ container_name: go-backend
restart: unless-stopped
logging:
driver: json-file
@@ -11,7 +11,7 @@ services:
DB_PATH: /app/data/gost.db
JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
- JAVA_OPTS: "-Xms256m -Xmx512m -Dfile.encoding=UTF-8 -Duser.timezone=Asia/Shanghai"
+ SERVER_ADDR: :6365
ports:
- "${BACKEND_PORT}:6365"
volumes:
@@ -26,7 +26,7 @@ services:
interval: 30s
timeout: 10s
retries: 5
- start_period: 60s
+ start_period: 30s
frontend:
image: ghcr.io/sagit-chu/vite-frontend:${FLUX_VERSION:-latest}
diff --git a/docker-compose-v6.yml b/docker-compose-v6.yml
index ec791e9..91d707c 100644
--- a/docker-compose-v6.yml
+++ b/docker-compose-v6.yml
@@ -1,7 +1,7 @@
services:
backend:
- image: ghcr.io/sagit-chu/springboot-backend:${FLUX_VERSION:-latest}
- container_name: springboot-backend
+ image: ghcr.io/sagit-chu/go-backend:${FLUX_VERSION:-latest}
+ container_name: go-backend
restart: unless-stopped
logging:
driver: json-file
@@ -11,7 +11,7 @@ services:
DB_PATH: /app/data/gost.db
JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
- JAVA_OPTS: "-Xms256m -Xmx512m -Dfile.encoding=UTF-8 -Duser.timezone=Asia/Shanghai"
+ SERVER_ADDR: :6365
ports:
- "${BACKEND_PORT}:6365"
volumes:
@@ -26,7 +26,7 @@ services:
interval: 30s
timeout: 10s
retries: 5
- start_period: 60s
+ start_period: 30s
frontend:
image: ghcr.io/sagit-chu/vite-frontend:${FLUX_VERSION:-latest}
diff --git a/go-backend/Dockerfile b/go-backend/Dockerfile
new file mode 100644
index 0000000..fcc41d4
--- /dev/null
+++ b/go-backend/Dockerfile
@@ -0,0 +1,17 @@
+FROM golang:1.23-bookworm AS builder
+WORKDIR /src
+
+COPY go.mod ./
+RUN go mod download
+
+COPY . .
+RUN CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -o /out/paneld ./cmd/paneld
+
+FROM debian:bookworm-slim
+WORKDIR /app
+RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates wget && rm -rf /var/lib/apt/lists/*
+COPY --from=builder /out/paneld /app/paneld
+
+ENV SERVER_ADDR=:6365
+EXPOSE 6365
+ENTRYPOINT ["/app/paneld"]
diff --git a/go-backend/Makefile b/go-backend/Makefile
new file mode 100644
index 0000000..5ab2068
--- /dev/null
+++ b/go-backend/Makefile
@@ -0,0 +1,12 @@
+GO ?= go
+
+.PHONY: test build run
+
+test:
+ $(GO) test ./...
+
+build:
+ $(GO) build ./cmd/paneld
+
+run:
+ SERVER_ADDR=:6365 $(GO) run ./cmd/paneld
diff --git a/go-backend/cmd/paneld/main.go b/go-backend/cmd/paneld/main.go
new file mode 100644
index 0000000..832b2dd
--- /dev/null
+++ b/go-backend/cmd/paneld/main.go
@@ -0,0 +1,51 @@
+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")
+ }
+ log.Printf("starting go-backend on %s (db=%s)", cfg.Addr, cfg.DBPath)
+
+ a, err := app.New(cfg)
+ if err != nil {
+ log.Fatalf("failed to create app: %v", err)
+ }
+
+ errCh := make(chan error, 1)
+ go func() {
+ errCh <- a.Run()
+ }()
+
+ sigCh := make(chan os.Signal, 1)
+ signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
+
+ select {
+ case sig := <-sigCh:
+ log.Printf("received signal %s, shutting down", sig)
+ case runErr := <-errCh:
+ if runErr != nil && !errors.Is(runErr, http.ErrServerClosed) {
+ log.Fatalf("server stopped unexpectedly: %v", runErr)
+ }
+ }
+
+ ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
+ defer cancel()
+ if err := a.Shutdown(ctx); err != nil {
+ log.Fatalf("shutdown failed: %v", err)
+ }
+}
diff --git a/go-backend/go.mod b/go-backend/go.mod
new file mode 100644
index 0000000..2c26ac9
--- /dev/null
+++ b/go-backend/go.mod
@@ -0,0 +1,23 @@
+module go-backend
+
+go 1.23.0
+
+toolchain go1.24.4
+
+require (
+ github.com/gorilla/websocket v1.5.3
+ modernc.org/sqlite v1.37.1
+)
+
+require (
+ github.com/dustin/go-humanize v1.0.1 // indirect
+ github.com/google/uuid v1.6.0 // indirect
+ github.com/mattn/go-isatty v0.0.20 // indirect
+ github.com/ncruces/go-strftime v0.1.9 // indirect
+ github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
+ golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect
+ golang.org/x/sys v0.33.0 // indirect
+ modernc.org/libc v1.65.7 // indirect
+ modernc.org/mathutil v1.7.1 // indirect
+ modernc.org/memory v1.11.0 // indirect
+)
diff --git a/go-backend/go.sum b/go-backend/go.sum
new file mode 100644
index 0000000..fa6b48b
--- /dev/null
+++ b/go-backend/go.sum
@@ -0,0 +1,49 @@
+github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
+github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
+github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
+github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
+github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
+github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
+github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
+github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
+github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
+github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
+github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4=
+github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
+github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
+github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
+golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM=
+golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8=
+golang.org/x/mod v0.24.0 h1:ZfthKaKaT4NrhGVZHO1/WDTwGES4De8KtWO0SIbNJMU=
+golang.org/x/mod v0.24.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww=
+golang.org/x/sync v0.14.0 h1:woo0S4Yywslg6hp4eUFjTVOyKt0RookbpAHG4c1HmhQ=
+golang.org/x/sync v0.14.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
+golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
+golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
+golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
+golang.org/x/tools v0.33.0 h1:4qz2S3zmRxbGIhDIAgjxvFutSvH5EfnsYrRBj0UI0bc=
+golang.org/x/tools v0.33.0/go.mod h1:CIJMaWEY88juyUfo7UbgPqbC8rU2OqfAV1h2Qp0oMYI=
+modernc.org/cc/v4 v4.26.1 h1:+X5NtzVBn0KgsBCBe+xkDC7twLb/jNVj9FPgiwSQO3s=
+modernc.org/cc/v4 v4.26.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
+modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU=
+modernc.org/ccgo/v4 v4.28.0/go.mod h1:JygV3+9AV6SmPhDasu4JgquwU81XAKLd3OKTUDNOiKE=
+modernc.org/fileutil v1.3.1 h1:8vq5fe7jdtEvoCf3Zf9Nm0Q05sH6kGx0Op2CPx1wTC8=
+modernc.org/fileutil v1.3.1/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
+modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
+modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
+modernc.org/libc v1.65.7 h1:Ia9Z4yzZtWNtUIuiPuQ7Qf7kxYrxP1/jeHZzG8bFu00=
+modernc.org/libc v1.65.7/go.mod h1:011EQibzzio/VX3ygj1qGFt5kMjP0lHb0qCW5/D/pQU=
+modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
+modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
+modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
+modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
+modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
+modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
+modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
+modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
+modernc.org/sqlite v1.37.1 h1:EgHJK/FPoqC+q2YBXg7fUmES37pCHFc97sI7zSayBEs=
+modernc.org/sqlite v1.37.1/go.mod h1:XwdRtsE1MpiBcL54+MbKcaDvcuej+IYSMfLN6gSKV8g=
+modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
+modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
+modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
+modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM=
diff --git a/go-backend/internal/app/app.go b/go-backend/internal/app/app.go
new file mode 100644
index 0000000..425b05a
--- /dev/null
+++ b/go-backend/internal/app/app.go
@@ -0,0 +1,60 @@
+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
+ h *handler.Handler
+}
+
+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, h: h}, nil
+}
+
+func (a *App) Run() error {
+ if a.h != nil {
+ a.h.StartBackgroundJobs()
+ }
+ return a.server.ListenAndServe()
+}
+
+func (a *App) Shutdown(ctx context.Context) error {
+ if a.h != nil {
+ a.h.StopBackgroundJobs()
+ }
+ shutdownErr := a.server.Shutdown(ctx)
+ closeErr := a.repo.Close()
+ if shutdownErr != nil {
+ return shutdownErr
+ }
+ return closeErr
+}
diff --git a/go-backend/internal/auth/jwt.go b/go-backend/internal/auth/jwt.go
new file mode 100644
index 0000000..3bc6d9f
--- /dev/null
+++ b/go-backend/internal/auth/jwt.go
@@ -0,0 +1,121 @@
+package auth
+
+import (
+ "crypto/hmac"
+ "crypto/sha256"
+ "encoding/base64"
+ "encoding/json"
+ "errors"
+ "strconv"
+ "time"
+)
+
+const (
+ algorithm = "HmacSHA256"
+ expireTime = 90 * 24 * time.Hour
+)
+
+type Claims struct {
+ Sub string `json:"sub"`
+ Iat int64 `json:"iat"`
+ Exp int64 `json:"exp"`
+ User string `json:"user"`
+ Name string `json:"name"`
+ RoleID int `json:"role_id"`
+}
+
+type tokenHeader struct {
+ Alg string `json:"alg"`
+ Typ string `json:"typ"`
+}
+
+func GenerateToken(userID int64, username string, roleID int, secret string) (string, error) {
+ now := time.Now()
+ header := tokenHeader{Alg: algorithm, Typ: "JWT"}
+ claims := Claims{
+ Sub: strconv.FormatInt(userID, 10),
+ Iat: now.Unix(),
+ Exp: now.Add(expireTime).Unix(),
+ User: username,
+ Name: username,
+ RoleID: roleID,
+ }
+
+ headerPart, err := encodeJSON(header)
+ if err != nil {
+ return "", err
+ }
+ payloadPart, err := encodeJSON(claims)
+ if err != nil {
+ return "", err
+ }
+ sig := sign(headerPart+"."+payloadPart, secret)
+
+ return headerPart + "." + payloadPart + "." + sig, nil
+}
+
+func ValidateToken(token, secret string) (Claims, bool) {
+ claims, err := ParseClaims(token, secret)
+ if err != nil {
+ return Claims{}, false
+ }
+ return claims, true
+}
+
+func ParseClaims(token, secret string) (Claims, error) {
+ parts := splitToken(token)
+ if len(parts) != 3 {
+ return Claims{}, errors.New("invalid token")
+ }
+
+ signedContent := parts[0] + "." + parts[1]
+ expected := sign(signedContent, secret)
+ if !hmac.Equal([]byte(expected), []byte(parts[2])) {
+ return Claims{}, errors.New("invalid signature")
+ }
+
+ payloadBytes, err := base64.RawURLEncoding.DecodeString(parts[1])
+ if err != nil {
+ return Claims{}, err
+ }
+
+ var claims Claims
+ if err := json.Unmarshal(payloadBytes, &claims); err != nil {
+ return Claims{}, err
+ }
+
+ if claims.Exp <= time.Now().Unix() {
+ return Claims{}, errors.New("token expired")
+ }
+
+ return claims, nil
+}
+
+func splitToken(token string) []string {
+ parts := make([]string, 0, 3)
+ current := ""
+ for i := 0; i < len(token); i++ {
+ if token[i] == '.' {
+ parts = append(parts, current)
+ current = ""
+ continue
+ }
+ current += string(token[i])
+ }
+ parts = append(parts, current)
+ return parts
+}
+
+func encodeJSON(v interface{}) (string, error) {
+ raw, err := json.Marshal(v)
+ if err != nil {
+ return "", err
+ }
+ return base64.RawURLEncoding.EncodeToString(raw), nil
+}
+
+func sign(content, secret string) string {
+ h := hmac.New(sha256.New, []byte(secret))
+ h.Write([]byte(content))
+ return base64.RawURLEncoding.EncodeToString(h.Sum(nil))
+}
diff --git a/go-backend/internal/config/config.go b/go-backend/internal/config/config.go
new file mode 100644
index 0000000..043730c
--- /dev/null
+++ b/go-backend/internal/config/config.go
@@ -0,0 +1,28 @@
+package config
+
+import "os"
+
+type Config struct {
+ Addr string
+ DBPath string
+ JWTSecret string
+ LogDir string
+}
+
+func FromEnv() Config {
+ cfg := Config{
+ Addr: getEnv("SERVER_ADDR", ":6365"),
+ DBPath: getEnv("DB_PATH", "/app/data/gost.db"),
+ JWTSecret: getEnv("JWT_SECRET", ""),
+ LogDir: getEnv("LOG_DIR", "/app/logs"),
+ }
+
+ return cfg
+}
+
+func getEnv(key, fallback string) string {
+ if v := os.Getenv(key); v != "" {
+ return v
+ }
+ return fallback
+}
diff --git a/go-backend/internal/http/handler/captcha_state.go b/go-backend/internal/http/handler/captcha_state.go
new file mode 100644
index 0000000..d0b03b4
--- /dev/null
+++ b/go-backend/internal/http/handler/captcha_state.go
@@ -0,0 +1,67 @@
+package handler
+
+import "time"
+
+const captchaTokenTTL = 5 * time.Minute
+
+func (h *Handler) storeCaptchaToken(token string) {
+ if h == nil {
+ return
+ }
+ token = normalizeCaptchaToken(token)
+ if token == "" {
+ return
+ }
+
+ h.captchaMu.Lock()
+ defer h.captchaMu.Unlock()
+
+ now := time.Now().UnixMilli()
+ h.pruneExpiredCaptchaTokensLocked(now)
+ h.captchaTokens[token] = now + int64(captchaTokenTTL/time.Millisecond)
+}
+
+func (h *Handler) consumeCaptchaToken(token string) bool {
+ if h == nil {
+ return false
+ }
+ token = normalizeCaptchaToken(token)
+ if token == "" {
+ return false
+ }
+
+ h.captchaMu.Lock()
+ defer h.captchaMu.Unlock()
+
+ now := time.Now().UnixMilli()
+ h.pruneExpiredCaptchaTokensLocked(now)
+ expiresAt, ok := h.captchaTokens[token]
+ if !ok || expiresAt <= now {
+ delete(h.captchaTokens, token)
+ return false
+ }
+ delete(h.captchaTokens, token)
+ return true
+}
+
+func (h *Handler) pruneExpiredCaptchaTokensLocked(now int64) {
+ for token, expiresAt := range h.captchaTokens {
+ if expiresAt <= now {
+ delete(h.captchaTokens, token)
+ }
+ }
+}
+
+func normalizeCaptchaToken(token string) string {
+ return trimToken(token)
+}
+
+func trimToken(token string) string {
+ for len(token) > 0 && (token[0] == ' ' || token[0] == '\t' || token[0] == '\n' || token[0] == '\r') {
+ token = token[1:]
+ }
+ for len(token) > 0 && (token[len(token)-1] == ' ' || token[len(token)-1] == '\t' || token[len(token)-1] == '\n' || token[len(token)-1] == '\r') {
+ token = token[:len(token)-1]
+ }
+ return token
+}
diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go
new file mode 100644
index 0000000..2a91cb3
--- /dev/null
+++ b/go-backend/internal/http/handler/control_plane.go
@@ -0,0 +1,967 @@
+package handler
+
+import (
+ "database/sql"
+ "errors"
+ "fmt"
+ "net"
+ "net/http"
+ "sort"
+ "strconv"
+ "strings"
+ "time"
+
+ "go-backend/internal/ws"
+)
+
+var errForwardNotFound = errors.New("forward not found")
+
+type forwardRecord struct {
+ ID int64
+ UserID int64
+ UserName string
+ Name string
+ TunnelID int64
+ RemoteAddr string
+ Strategy string
+ Status int
+}
+
+type tunnelRecord struct {
+ ID int64
+ Type int
+ Status int
+ Flow int64
+ TrafficRatio float64
+}
+
+type forwardPortRecord struct {
+ NodeID int64
+ Port int
+}
+
+type nodeRecord struct {
+ ID int64
+ Name string
+ ServerIP string
+ ServerIPv4 string
+ ServerIPv6 string
+ Status int
+ PortRange string
+ TCPListenAddr string
+ UDPListenAddr string
+ InterfaceName string
+}
+
+type chainNodeRecord struct {
+ ChainType int
+ Inx int64
+ NodeID int64
+ Port int
+ NodeName string
+}
+
+type diagnosisTarget struct {
+ Address string
+ IP string
+ Port int
+}
+
+func (h *Handler) resolveForwardAccess(r *http.Request, forwardID int64) (*forwardRecord, int64, int, error) {
+ userID, roleID, err := userRoleFromRequest(r)
+ if err != nil {
+ return nil, 0, 0, err
+ }
+ forward, err := h.ensureForwardAccessByActor(userID, roleID, forwardID)
+ if err != nil {
+ return nil, userID, roleID, err
+ }
+ return forward, userID, roleID, nil
+}
+
+func (h *Handler) ensureForwardAccessByActor(actorUserID int64, actorRole int, forwardID int64) (*forwardRecord, error) {
+ forward, err := h.getForwardRecord(forwardID)
+ if err != nil {
+ return nil, err
+ }
+ if actorRole != 0 && forward.UserID != actorUserID {
+ return nil, errForwardNotFound
+ }
+ return forward, nil
+}
+
+func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int64) error {
+ if roleID == 0 {
+ return nil
+ }
+ var count int
+ err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? AND status = 1`, userID, tunnelID).Scan(&count)
+ if err != nil {
+ return err
+ }
+ if count <= 0 {
+ return errors.New("你没有该隧道的权限")
+ }
+ return nil
+}
+
+func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) {
+ row := h.repo.DB().QueryRow(`
+ SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
+ FROM forward WHERE id = ? LIMIT 1
+ `, forwardID)
+ var fr forwardRecord
+ err := row.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, errForwardNotFound
+ }
+ return nil, err
+ }
+ if strings.TrimSpace(fr.Strategy) == "" {
+ fr.Strategy = "fifo"
+ }
+ return &fr, nil
+}
+
+func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) {
+ row := h.repo.DB().QueryRow(`SELECT id, type, status, flow, traffic_ratio FROM tunnel WHERE id = ? LIMIT 1`, tunnelID)
+ var tr tunnelRecord
+ err := row.Scan(&tr.ID, &tr.Type, &tr.Status, &tr.Flow, &tr.TrafficRatio)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, errors.New("隧道不存在")
+ }
+ return nil, err
+ }
+ if tr.Flow <= 0 {
+ tr.Flow = 1
+ }
+ if tr.TrafficRatio <= 0 {
+ tr.TrafficRatio = 1
+ }
+ return &tr, nil
+}
+
+func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) {
+ rows, err := h.repo.DB().Query(`
+ SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
+ FROM forward
+ WHERE tunnel_id = ?
+ ORDER BY id ASC
+ `, tunnelID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ result := make([]forwardRecord, 0)
+ for rows.Next() {
+ var fr forwardRecord
+ if err := rows.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status); err != nil {
+ return nil, err
+ }
+ if strings.TrimSpace(fr.Strategy) == "" {
+ fr.Strategy = "fifo"
+ }
+ result = append(result, fr)
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return result, nil
+}
+
+func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error) {
+ rows, err := h.repo.DB().Query(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ result := make([]forwardPortRecord, 0)
+ for rows.Next() {
+ var item forwardPortRecord
+ if err := rows.Scan(&item.NodeID, &item.Port); err != nil {
+ return nil, err
+ }
+ result = append(result, item)
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return result, nil
+}
+
+func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
+ row := h.repo.DB().QueryRow(`
+ SELECT id, name, server_ip, server_ip_v4, server_ip_v6, status, port, tcp_listen_addr, udp_listen_addr, interface_name
+ FROM node
+ WHERE id = ?
+ LIMIT 1
+ `, nodeID)
+ var n nodeRecord
+ var serverIPv4 sql.NullString
+ var serverIPv6 sql.NullString
+ var portRange sql.NullString
+ var tcpListen sql.NullString
+ var udpListen sql.NullString
+ var iface sql.NullString
+ err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Status, &portRange, &tcpListen, &udpListen, &iface)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, errors.New("节点不存在")
+ }
+ return nil, err
+ }
+ n.ServerIPv4 = strings.TrimSpace(serverIPv4.String)
+ n.ServerIPv6 = strings.TrimSpace(serverIPv6.String)
+ n.PortRange = strings.TrimSpace(portRange.String)
+ n.TCPListenAddr = strings.TrimSpace(tcpListen.String)
+ n.UDPListenAddr = strings.TrimSpace(udpListen.String)
+ n.InterfaceName = strings.TrimSpace(iface.String)
+ if n.TCPListenAddr == "" {
+ n.TCPListenAddr = "[::]"
+ }
+ if n.UDPListenAddr == "" {
+ n.UDPListenAddr = "[::]"
+ }
+ if strings.TrimSpace(n.Name) == "" {
+ n.Name = fmt.Sprintf("node_%d", n.ID)
+ }
+ return &n, nil
+}
+
+func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int, error) {
+ row := h.repo.DB().QueryRow(`
+ SELECT ut.id, sl.speed
+ FROM user_tunnel ut
+ LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
+ WHERE ut.user_id = ? AND ut.tunnel_id = ?
+ LIMIT 1
+ `, userID, tunnelID)
+ var userTunnelID int64
+ var speed sql.NullInt64
+ err := row.Scan(&userTunnelID, &speed)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return 0, nil, nil
+ }
+ return 0, nil, err
+ }
+ if !speed.Valid || speed.Int64 <= 0 {
+ return userTunnelID, nil, nil
+ }
+ v := int(speed.Int64)
+ return userTunnelID, &v, nil
+}
+
+func (h *Handler) syncForwardServices(forward *forwardRecord, method string, allowFallbackAdd bool) error {
+ if h == nil || forward == nil {
+ return errors.New("invalid forward sync context")
+ }
+
+ tunnel, err := h.getTunnelRecord(forward.TunnelID)
+ if err != nil {
+ return err
+ }
+ ports, err := h.listForwardPorts(forward.ID)
+ if err != nil {
+ return err
+ }
+ if len(ports) == 0 {
+ return errors.New("转发入口端口不存在")
+ }
+
+ userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
+ if err != nil {
+ return err
+ }
+ serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
+
+ for _, fp := range ports {
+ node, err := h.getNodeRecord(fp.NodeID)
+ if err != nil {
+ return err
+ }
+ services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiter)
+ _, err = h.sendNodeCommand(node.ID, method, services, true, false)
+ if err != nil && allowFallbackAdd && method == "UpdateService" {
+ _, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
+ }
+ if err != nil {
+ return fmt.Errorf("节点 %s 下发失败: %w", node.Name, err)
+ }
+ }
+ return nil
+}
+
+func (h *Handler) controlForwardServices(forward *forwardRecord, commandType string, tolerateNotFound bool) error {
+ if h == nil || forward == nil {
+ return errors.New("invalid forward control context")
+ }
+ ports, err := h.listForwardPorts(forward.ID)
+ if err != nil {
+ return err
+ }
+ if len(ports) == 0 {
+ return nil
+ }
+ userTunnelID, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
+ if err != nil {
+ return err
+ }
+ base := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
+ payload := map[string]interface{}{
+ "services": buildForwardControlServiceNames(base, commandType),
+ }
+ seen := map[int64]struct{}{}
+ for _, fp := range ports {
+ if _, ok := seen[fp.NodeID]; ok {
+ continue
+ }
+ seen[fp.NodeID] = struct{}{}
+ _, err := h.sendNodeCommand(fp.NodeID, commandType, payload, false, tolerateNotFound)
+ if err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+func (h *Handler) applyNodeProtocolChange(nodeID int64, httpVal, tlsVal, socksVal int) error {
+ _, err := h.sendNodeCommand(nodeID, "SetProtocol", map[string]interface{}{
+ "http": httpVal,
+ "tls": tlsVal,
+ "socks": socksVal,
+ }, false, false)
+ return err
+}
+
+func (h *Handler) sendNodeCommand(nodeID int64, commandType string, data interface{}, tolerateExists bool, tolerateNotFound bool) (ws.CommandResult, error) {
+ result, err := h.wsServer.SendCommand(nodeID, commandType, data, 12*time.Second)
+ if err == nil {
+ return result, nil
+ }
+ msg := strings.ToLower(strings.TrimSpace(err.Error()))
+ if tolerateExists {
+ if strings.Contains(msg, "exists") || strings.Contains(msg, "already") || strings.Contains(msg, "已存在") {
+ return result, nil
+ }
+ }
+ if tolerateNotFound {
+ if strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") {
+ return result, nil
+ }
+ }
+ return result, err
+}
+
+func (h *Handler) diagnoseForwardRuntime(forward *forwardRecord) (map[string]interface{}, error) {
+ if forward == nil {
+ return nil, errForwardNotFound
+ }
+ targets, err := resolveDiagnosisTargets(forward.RemoteAddr)
+ if err != nil {
+ return nil, err
+ }
+
+ tunnel, err := h.getTunnelRecord(forward.TunnelID)
+ if err != nil {
+ return nil, err
+ }
+
+ chainRows, err := h.listChainNodesForTunnel(forward.TunnelID)
+ if err != nil {
+ return nil, err
+ }
+ if len(chainRows) == 0 {
+ return nil, errors.New("隧道配置不完整")
+ }
+
+ inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
+ results := make([]map[string]interface{}, 0, len(chainRows)*2+len(targets))
+ nodeCache := map[int64]*nodeRecord{}
+
+ switch tunnel.Type {
+ case 1:
+ for _, inNode := range inNodes {
+ for _, target := range targets {
+ description := fmt.Sprintf("入口(%s)->目标(%s)", inNode.NodeName, target.Address)
+ h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, target.IP, target.Port, description, map[string]interface{}{
+ "fromChainType": 1,
+ })
+ }
+ }
+ case 2:
+ for _, inNode := range inNodes {
+ if len(chainHops) > 0 {
+ for _, firstNode := range chainHops[0] {
+ description := fmt.Sprintf("入口(%s)->第1跳(%s)", inNode.NodeName, firstNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, firstNode, description, map[string]interface{}{
+ "fromChainType": 1,
+ "toChainType": 2,
+ "toInx": firstNode.Inx,
+ })
+ }
+ } else {
+ for _, outNode := range outNodes {
+ description := fmt.Sprintf("入口(%s)->出口(%s)", inNode.NodeName, outNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, outNode, description, map[string]interface{}{
+ "fromChainType": 1,
+ "toChainType": 3,
+ })
+ }
+ }
+ }
+
+ for i, hop := range chainHops {
+ for _, currentNode := range hop {
+ if i+1 < len(chainHops) {
+ for _, nextNode := range chainHops[i+1] {
+ description := fmt.Sprintf("第%d跳(%s)->第%d跳(%s)", i+1, currentNode.NodeName, i+2, nextNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, nextNode, description, map[string]interface{}{
+ "fromChainType": 2,
+ "fromInx": currentNode.Inx,
+ "toChainType": 2,
+ "toInx": nextNode.Inx,
+ })
+ }
+ } else {
+ for _, outNode := range outNodes {
+ description := fmt.Sprintf("第%d跳(%s)->出口(%s)", i+1, currentNode.NodeName, outNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, outNode, description, map[string]interface{}{
+ "fromChainType": 2,
+ "fromInx": currentNode.Inx,
+ "toChainType": 3,
+ })
+ }
+ }
+ }
+ }
+
+ for _, outNode := range outNodes {
+ for _, target := range targets {
+ description := fmt.Sprintf("出口(%s)->目标(%s)", outNode.NodeName, target.Address)
+ h.appendPathDiagnosis(&results, nodeCache, outNode.NodeID, target.IP, target.Port, description, map[string]interface{}{
+ "fromChainType": 3,
+ })
+ }
+ }
+ default:
+ for _, inNode := range inNodes {
+ for _, target := range targets {
+ description := fmt.Sprintf("入口(%s)->目标(%s)", inNode.NodeName, target.Address)
+ h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, target.IP, target.Port, description, map[string]interface{}{
+ "fromChainType": 1,
+ })
+ }
+ }
+ }
+
+ payload := map[string]interface{}{
+ "forwardName": forward.Name,
+ "timestamp": time.Now().UnixMilli(),
+ "results": results,
+ }
+ return payload, nil
+}
+
+func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{}, error) {
+ tunnel, err := h.getTunnelRecord(tunnelID)
+ if err != nil {
+ return nil, err
+ }
+
+ var tunnelName string
+ if err := h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName); err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, errors.New("隧道不存在")
+ }
+ return nil, err
+ }
+
+ chainRows, err := h.listChainNodesForTunnel(tunnelID)
+ if err != nil {
+ return nil, err
+ }
+ if len(chainRows) == 0 {
+ return nil, errors.New("隧道配置不完整")
+ }
+
+ inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
+ results := make([]map[string]interface{}, 0, len(chainRows)*2)
+ nodeCache := map[int64]*nodeRecord{}
+
+ switch tunnel.Type {
+ case 1:
+ for _, inNode := range inNodes {
+ description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
+ h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, "www.google.com", 443, description, map[string]interface{}{
+ "fromChainType": 1,
+ })
+ }
+ case 2:
+ for _, inNode := range inNodes {
+ if len(chainHops) > 0 {
+ for _, firstNode := range chainHops[0] {
+ description := fmt.Sprintf("入口(%s)->第1跳(%s)", inNode.NodeName, firstNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, firstNode, description, map[string]interface{}{
+ "fromChainType": 1,
+ "toChainType": 2,
+ "toInx": firstNode.Inx,
+ })
+ }
+ } else {
+ for _, outNode := range outNodes {
+ description := fmt.Sprintf("入口(%s)->出口(%s)", inNode.NodeName, outNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, inNode.NodeID, outNode, description, map[string]interface{}{
+ "fromChainType": 1,
+ "toChainType": 3,
+ })
+ }
+ }
+ }
+
+ for i, hop := range chainHops {
+ for _, currentNode := range hop {
+ if i+1 < len(chainHops) {
+ for _, nextNode := range chainHops[i+1] {
+ description := fmt.Sprintf("第%d跳(%s)->第%d跳(%s)", i+1, currentNode.NodeName, i+2, nextNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, nextNode, description, map[string]interface{}{
+ "fromChainType": 2,
+ "fromInx": currentNode.Inx,
+ "toChainType": 2,
+ "toInx": nextNode.Inx,
+ })
+ }
+ } else {
+ for _, outNode := range outNodes {
+ description := fmt.Sprintf("第%d跳(%s)->出口(%s)", i+1, currentNode.NodeName, outNode.NodeName)
+ h.appendChainHopDiagnosis(&results, nodeCache, currentNode.NodeID, outNode, description, map[string]interface{}{
+ "fromChainType": 2,
+ "fromInx": currentNode.Inx,
+ "toChainType": 3,
+ })
+ }
+ }
+ }
+ }
+
+ for _, outNode := range outNodes {
+ description := fmt.Sprintf("出口(%s)->外网", outNode.NodeName)
+ h.appendPathDiagnosis(&results, nodeCache, outNode.NodeID, "www.google.com", 443, description, map[string]interface{}{
+ "fromChainType": 3,
+ })
+ }
+ default:
+ for _, inNode := range inNodes {
+ description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
+ h.appendPathDiagnosis(&results, nodeCache, inNode.NodeID, "www.google.com", 443, description, map[string]interface{}{
+ "fromChainType": 1,
+ })
+ }
+ }
+
+ payload := map[string]interface{}{
+ "tunnelName": tunnelName,
+ "tunnelType": map[bool]string{true: "端口转发", false: "隧道转发"}[tunnel.Type == 1],
+ "timestamp": time.Now().UnixMilli(),
+ "results": results,
+ }
+ return payload, nil
+}
+
+func splitChainNodeGroups(rows []chainNodeRecord) ([]chainNodeRecord, [][]chainNodeRecord, []chainNodeRecord) {
+ inNodes := make([]chainNodeRecord, 0)
+ outNodes := make([]chainNodeRecord, 0)
+ chainByInx := map[int64][]chainNodeRecord{}
+ hopOrder := make([]int64, 0)
+
+ for _, row := range rows {
+ switch row.ChainType {
+ case 1:
+ inNodes = append(inNodes, row)
+ case 2:
+ if _, ok := chainByInx[row.Inx]; !ok {
+ hopOrder = append(hopOrder, row.Inx)
+ }
+ chainByInx[row.Inx] = append(chainByInx[row.Inx], row)
+ case 3:
+ outNodes = append(outNodes, row)
+ }
+ }
+
+ sort.Slice(hopOrder, func(i, j int) bool { return hopOrder[i] < hopOrder[j] })
+ chainHops := make([][]chainNodeRecord, 0, len(hopOrder))
+ for _, inx := range hopOrder {
+ chainHops = append(chainHops, chainByInx[inx])
+ }
+
+ return inNodes, chainHops, outNodes
+}
+
+func resolveDiagnosisTargets(remoteAddr string) ([]diagnosisTarget, error) {
+ rawTargets := splitRemoteTargets(remoteAddr)
+ if len(rawTargets) == 0 {
+ return nil, errors.New("目标地址不能为空")
+ }
+
+ targets := make([]diagnosisTarget, 0, len(rawTargets))
+ for _, raw := range rawTargets {
+ ip, port, err := parseTargetAddress(raw)
+ if err != nil {
+ continue
+ }
+ targets = append(targets, diagnosisTarget{Address: raw, IP: ip, Port: port})
+ }
+ if len(targets) == 0 {
+ return nil, errors.New("目标地址格式错误")
+ }
+ return targets, nil
+}
+
+func (h *Handler) cachedNode(nodeCache map[int64]*nodeRecord, nodeID int64) (*nodeRecord, error) {
+ if node, ok := nodeCache[nodeID]; ok {
+ return node, nil
+ }
+ node, err := h.getNodeRecord(nodeID)
+ if err != nil {
+ return nil, err
+ }
+ nodeCache[nodeID] = node
+ return node, nil
+}
+
+func newDiagnosisResultItem(fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}) map[string]interface{} {
+ item := map[string]interface{}{
+ "nodeName": fmt.Sprintf("node_%d", fromNodeID),
+ "nodeId": strconv.FormatInt(fromNodeID, 10),
+ "targetIp": targetIP,
+ "targetPort": targetPort,
+ "description": description,
+ "averageTime": 0,
+ "packetLoss": 100,
+ }
+ for k, v := range metadata {
+ item[k] = v
+ }
+ return item
+}
+
+func (h *Handler) appendFailedDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}, message string) {
+ item := newDiagnosisResultItem(fromNodeID, targetIP, targetPort, description, metadata)
+ if node, err := h.cachedNode(nodeCache, fromNodeID); err == nil {
+ item["nodeName"] = node.Name
+ }
+ if strings.TrimSpace(message) == "" {
+ message = "TCP连接失败"
+ }
+ item["success"] = false
+ item["message"] = message
+ *results = append(*results, item)
+}
+
+func (h *Handler) appendPathDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}) {
+ item := newDiagnosisResultItem(fromNodeID, targetIP, targetPort, description, metadata)
+
+ fromNode, err := h.cachedNode(nodeCache, fromNodeID)
+ if err != nil {
+ item["success"] = false
+ item["message"] = err.Error()
+ *results = append(*results, item)
+ return
+ }
+ item["nodeName"] = fromNode.Name
+
+ pingData, pingErr := h.tcpPingViaNode(fromNodeID, targetIP, targetPort)
+ if pingErr != nil {
+ item["success"] = false
+ item["message"] = pingErr.Error()
+ *results = append(*results, item)
+ return
+ }
+
+ success := asBool(pingData["success"], false)
+ item["success"] = success
+ item["averageTime"] = asFloat(pingData["averageTime"], 0)
+ item["packetLoss"] = asFloat(pingData["packetLoss"], 100)
+
+ message := strings.TrimSpace(asString(pingData["message"]))
+ if success {
+ if message == "" {
+ message = "TCP连接成功"
+ }
+ } else {
+ if message == "" {
+ message = strings.TrimSpace(asString(pingData["errorMessage"]))
+ }
+ if message == "" {
+ message = "TCP连接失败"
+ }
+ }
+ item["message"] = message
+ *results = append(*results, item)
+}
+
+func (h *Handler) appendChainHopDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, toNode chainNodeRecord, description string, metadata map[string]interface{}) {
+ targetNode, err := h.cachedNode(nodeCache, toNode.NodeID)
+ if err != nil {
+ h.appendFailedDiagnosis(results, nodeCache, fromNodeID, "", 0, description, metadata, err.Error())
+ return
+ }
+ targetIP, targetPort, err := resolveChainProbeTarget(targetNode, toNode.Port)
+ if err != nil {
+ h.appendFailedDiagnosis(results, nodeCache, fromNodeID, strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]"), toNode.Port, description, metadata, err.Error())
+ return
+ }
+ h.appendPathDiagnosis(results, nodeCache, fromNodeID, targetIP, targetPort, description, metadata)
+}
+
+func resolveChainProbeTarget(targetNode *nodeRecord, preferredPort int) (string, int, error) {
+ if targetNode == nil {
+ return "", 0, errors.New("目标节点不存在")
+ }
+ host := strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]")
+ if host == "" {
+ return "", 0, errors.New("目标节点地址为空")
+ }
+ port := preferredPort
+ if port <= 0 {
+ port = firstPortFromRange(targetNode.PortRange)
+ }
+ if port <= 0 {
+ port = 443
+ }
+ return host, port, nil
+}
+
+func firstPortFromRange(portRange string) int {
+ portRange = strings.TrimSpace(portRange)
+ if portRange == "" {
+ return 0
+ }
+ first := strings.Split(portRange, ",")[0]
+ first = strings.TrimSpace(first)
+ if strings.Contains(first, "-") {
+ parts := strings.SplitN(first, "-", 2)
+ if len(parts) != 2 {
+ return 0
+ }
+ p, err := strconv.Atoi(strings.TrimSpace(parts[0]))
+ if err != nil || p <= 0 {
+ return 0
+ }
+ return p
+ }
+ p, err := strconv.Atoi(first)
+ if err != nil || p <= 0 {
+ return 0
+ }
+ return p
+}
+
+func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
+ rows, err := h.repo.DB().Query(`
+ SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name
+ FROM chain_tunnel ct
+ LEFT JOIN node n ON n.id = ct.node_id
+ WHERE ct.tunnel_id = ?
+ ORDER BY ct.chain_type ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
+ `, tunnelID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ result := make([]chainNodeRecord, 0)
+ for rows.Next() {
+ var item chainNodeRecord
+ var name sql.NullString
+ if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &item.Port, &name); err != nil {
+ return nil, err
+ }
+ if strings.TrimSpace(name.String) == "" {
+ item.NodeName = fmt.Sprintf("node_%d", item.NodeID)
+ } else {
+ item.NodeName = name.String
+ }
+ result = append(result, item)
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return result, nil
+}
+
+func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int) (map[string]interface{}, error) {
+ res, err := h.sendNodeCommand(nodeID, "TcpPing", map[string]interface{}{
+ "ip": ip,
+ "port": port,
+ "count": 4,
+ "timeout": 5000,
+ }, false, false)
+ if err != nil {
+ return nil, err
+ }
+ if res.Data == nil {
+ return nil, errors.New("节点未返回诊断数据")
+ }
+ return res.Data, nil
+}
+
+func splitRemoteTargets(remoteAddr string) []string {
+ parts := strings.Split(remoteAddr, ",")
+ out := make([]string, 0, len(parts))
+ for _, part := range parts {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ continue
+ }
+ out = append(out, processServerAddress(part))
+ }
+ return out
+}
+
+func parseTargetAddress(addr string) (string, int, error) {
+ addr = strings.TrimSpace(addr)
+ if addr == "" {
+ return "", 0, errors.New("empty address")
+ }
+ host, portStr, err := net.SplitHostPort(addr)
+ if err != nil {
+ idx := strings.LastIndex(addr, ":")
+ if idx <= 0 || idx >= len(addr)-1 {
+ return "", 0, err
+ }
+ host = strings.TrimSpace(addr[:idx])
+ portStr = strings.TrimSpace(addr[idx+1:])
+ }
+ port, err := strconv.Atoi(strings.TrimSpace(portStr))
+ if err != nil || port <= 0 || port > 65535 {
+ return "", 0, errors.New("invalid port")
+ }
+ host = strings.Trim(strings.TrimSpace(host), "[]")
+ if host == "" {
+ return "", 0, errors.New("invalid host")
+ }
+ return host, port, nil
+}
+
+func buildForwardServiceBase(forwardID, userID, userTunnelID int64) string {
+ return fmt.Sprintf("%d_%d_%d", forwardID, userID, userTunnelID)
+}
+
+func buildForwardControlServiceNames(base, commandType string) []string {
+ names := []string{base + "_tcp", base + "_udp"}
+ if strings.EqualFold(strings.TrimSpace(commandType), "DeleteService") {
+ return append([]string{base}, names...)
+ }
+ return names
+}
+
+func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiter *int) []map[string]interface{} {
+ protocols := []string{"tcp", "udp"}
+ services := make([]map[string]interface{}, 0, 2)
+ targets := splitRemoteTargets(forward.RemoteAddr)
+ strategy := strings.TrimSpace(forward.Strategy)
+ if strategy == "" {
+ strategy = "fifo"
+ }
+
+ for _, protocol := range protocols {
+ listenerAddr := node.TCPListenAddr
+ if protocol == "udp" {
+ listenerAddr = node.UDPListenAddr
+ }
+ service := map[string]interface{}{
+ "name": fmt.Sprintf("%s_%s", baseName, protocol),
+ "addr": fmt.Sprintf("%s:%d", listenerAddr, port),
+ "handler": map[string]interface{}{
+ "type": protocol,
+ },
+ "listener": map[string]interface{}{
+ "type": protocol,
+ },
+ "forwarder": map[string]interface{}{
+ "nodes": buildForwarderNodes(targets),
+ "selector": map[string]interface{}{
+ "strategy": strategy,
+ "maxFails": 1,
+ "failTimeout": "600s",
+ },
+ },
+ }
+ if protocol == "udp" {
+ service["listener"].(map[string]interface{})["metadata"] = map[string]interface{}{"keepAlive": true}
+ }
+ if tunnel != nil && tunnel.Type == 2 {
+ service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
+ }
+ if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
+ service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
+ }
+ if limiter != nil && *limiter > 0 {
+ service["limiter"] = strconv.Itoa(*limiter)
+ }
+ services = append(services, service)
+ }
+
+ return services
+}
+
+func buildForwarderNodes(targets []string) []map[string]interface{} {
+ nodes := make([]map[string]interface{}, 0, len(targets))
+ for i, addr := range targets {
+ nodes = append(nodes, map[string]interface{}{
+ "name": fmt.Sprintf("node_%d", i+1),
+ "addr": addr,
+ })
+ }
+ return nodes
+}
+
+func processServerAddress(serverAddr string) string {
+ serverAddr = strings.TrimSpace(serverAddr)
+ if serverAddr == "" {
+ return serverAddr
+ }
+ if strings.HasPrefix(serverAddr, "[") {
+ return serverAddr
+ }
+ idx := strings.LastIndex(serverAddr, ":")
+ if idx < 0 {
+ if looksLikeIPv6(serverAddr) {
+ return "[" + serverAddr + "]"
+ }
+ return serverAddr
+ }
+ host := strings.TrimSpace(serverAddr[:idx])
+ port := strings.TrimSpace(serverAddr[idx+1:])
+ if host == "" || port == "" {
+ return serverAddr
+ }
+ if looksLikeIPv6(host) {
+ return "[" + host + "]:" + port
+ }
+ return serverAddr
+}
+
+func looksLikeIPv6(address string) bool {
+ return strings.Count(address, ":") >= 2
+}
+
+func asBool(v interface{}, def bool) bool {
+ s := strings.TrimSpace(strings.ToLower(asString(v)))
+ if s == "" {
+ return def
+ }
+ switch s {
+ case "1", "t", "true", "yes", "y":
+ return true
+ case "0", "f", "false", "no", "n":
+ return false
+ default:
+ return def
+ }
+}
diff --git a/go-backend/internal/http/handler/control_plane_test.go b/go-backend/internal/http/handler/control_plane_test.go
new file mode 100644
index 0000000..f34cb92
--- /dev/null
+++ b/go-backend/internal/http/handler/control_plane_test.go
@@ -0,0 +1,27 @@
+package handler
+
+import (
+ "reflect"
+ "testing"
+)
+
+func TestBuildForwardControlServiceNamesPauseResume(t *testing.T) {
+ base := "12_34_56"
+ want := []string{base + "_tcp", base + "_udp"}
+
+ for _, command := range []string{"PauseService", "ResumeService"} {
+ got := buildForwardControlServiceNames(base, command)
+ if !reflect.DeepEqual(got, want) {
+ t.Fatalf("command %s expected %v, got %v", command, want, got)
+ }
+ }
+}
+
+func TestBuildForwardControlServiceNamesDelete(t *testing.T) {
+ base := "12_34_56"
+ want := []string{base, base + "_tcp", base + "_udp"}
+ got := buildForwardControlServiceNames(base, " DeleteService ")
+ if !reflect.DeepEqual(got, want) {
+ t.Fatalf("expected %v, got %v", want, got)
+ }
+}
diff --git a/go-backend/internal/http/handler/flow_policy.go b/go-backend/internal/http/handler/flow_policy.go
new file mode 100644
index 0000000..5efdb18
--- /dev/null
+++ b/go-backend/internal/http/handler/flow_policy.go
@@ -0,0 +1,338 @@
+package handler
+
+import (
+ "database/sql"
+ "encoding/json"
+ "strconv"
+ "strings"
+ "time"
+)
+
+const bytesPerGB int64 = 1024 * 1024 * 1024
+
+type userTunnelPolicy struct {
+ ID int64
+ UserID int64
+ TunnelID int64
+ Flow int64
+ InFlow int64
+ OutFlow int64
+ ExpTime int64
+ Status int
+}
+
+type gostConfigSnapshot struct {
+ Services []namedConfigItem `json:"services"`
+ Chains []namedConfigItem `json:"chains"`
+ Limiters []namedConfigItem `json:"limiters"`
+}
+
+type namedConfigItem struct {
+ Name string `json:"name"`
+}
+
+func (h *Handler) processFlowItem(item flowItem) {
+ serviceName := strings.TrimSpace(item.N)
+ if serviceName == "" || serviceName == "web_api" {
+ return
+ }
+
+ forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
+ if !ok {
+ return
+ }
+
+ inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
+ _ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
+
+ if userTunnelID > 0 {
+ h.enforceFlowPolicies(userID, userTunnelID)
+ }
+}
+
+func parseFlowServiceIDs(serviceName string) (int64, int64, int64, bool) {
+ parts := strings.Split(serviceName, "_")
+ if len(parts) < 3 {
+ return 0, 0, 0, false
+ }
+
+ 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 || forwardID <= 0 || userID <= 0 {
+ return 0, 0, 0, false
+ }
+
+ return forwardID, userID, userTunnelID, true
+}
+
+func (h *Handler) scaleFlowByTunnel(forwardID int64, inFlow int64, outFlow int64) (int64, int64) {
+ forward, err := h.getForwardRecord(forwardID)
+ if err != nil || forward == nil {
+ return inFlow, outFlow
+ }
+
+ tunnel, err := h.getTunnelRecord(forward.TunnelID)
+ if err != nil || tunnel == nil {
+ return inFlow, outFlow
+ }
+
+ scaledIn := int64(float64(inFlow)*tunnel.TrafficRatio) * tunnel.Flow
+ scaledOut := int64(float64(outFlow)*tunnel.TrafficRatio) * tunnel.Flow
+ return scaledIn, scaledOut
+}
+
+func (h *Handler) enforceFlowPolicies(userID int64, userTunnelID int64) {
+ now := time.Now().UnixMilli()
+
+ if h.shouldPauseUser(userID, now) {
+ h.pauseUserForwards(userID, now)
+ }
+
+ policy, err := h.getUserTunnelPolicy(userTunnelID)
+ if err != nil || policy == nil {
+ return
+ }
+
+ if shouldPauseUserTunnel(policy, now) {
+ h.pauseUserTunnelForwards(policy.UserID, policy.TunnelID, now)
+ }
+}
+
+func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
+ user, err := h.repo.GetUserByID(userID)
+ if err != nil || user == nil {
+ return false
+ }
+
+ flowLimit := user.Flow * bytesPerGB
+ current := user.InFlow + user.OutFlow
+ if flowLimit < current {
+ return true
+ }
+ if user.ExpTime > 0 && user.ExpTime <= now {
+ return true
+ }
+ return user.Status != 1
+}
+
+func shouldPauseUserTunnel(policy *userTunnelPolicy, now int64) bool {
+ if policy == nil {
+ return false
+ }
+
+ flowLimit := policy.Flow * bytesPerGB
+ current := policy.InFlow + policy.OutFlow
+ if current >= flowLimit {
+ return true
+ }
+ if policy.ExpTime > 0 && policy.ExpTime <= now {
+ return true
+ }
+ return policy.Status != 1
+}
+
+func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, error) {
+ if userTunnelID <= 0 {
+ return nil, nil
+ }
+
+ row := h.repo.DB().QueryRow(`
+ SELECT id, user_id, tunnel_id, flow, in_flow, out_flow, exp_time, status
+ FROM user_tunnel
+ WHERE id = ?
+ LIMIT 1
+ `, userTunnelID)
+
+ var policy userTunnelPolicy
+ if err := row.Scan(&policy.ID, &policy.UserID, &policy.TunnelID, &policy.Flow, &policy.InFlow, &policy.OutFlow, &policy.ExpTime, &policy.Status); err != nil {
+ if err == sql.ErrNoRows {
+ return nil, nil
+ }
+ return nil, err
+ }
+ return &policy, nil
+}
+
+func (h *Handler) pauseUserForwards(userID int64, now int64) {
+ forwards, err := h.listActiveForwardsByUser(userID)
+ if err != nil {
+ return
+ }
+ h.pauseForwardRecords(forwards, now)
+}
+
+func (h *Handler) pauseUserTunnelForwards(userID int64, tunnelID int64, now int64) {
+ forwards, err := h.listActiveForwardsByUserTunnel(userID, tunnelID)
+ if err != nil {
+ return
+ }
+ h.pauseForwardRecords(forwards, now)
+}
+
+func (h *Handler) pauseForwardRecords(forwards []forwardRecord, now int64) {
+ for i := range forwards {
+ forward := forwards[i]
+ _ = h.controlForwardServices(&forward, "PauseService", false)
+ _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, now, forward.ID)
+ }
+}
+
+func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error) {
+ rows, err := h.repo.DB().Query(`
+ SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
+ FROM forward
+ WHERE user_id = ? AND status = 1
+ ORDER BY id ASC
+ `, userID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ return scanForwardRecords(rows)
+}
+
+func (h *Handler) listActiveForwardsByUserTunnel(userID int64, tunnelID int64) ([]forwardRecord, error) {
+ rows, err := h.repo.DB().Query(`
+ SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
+ FROM forward
+ WHERE user_id = ? AND tunnel_id = ? AND status = 1
+ ORDER BY id ASC
+ `, userID, tunnelID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ return scanForwardRecords(rows)
+}
+
+func scanForwardRecords(rows *sql.Rows) ([]forwardRecord, error) {
+ out := make([]forwardRecord, 0)
+ for rows.Next() {
+ var record forwardRecord
+ if err := rows.Scan(&record.ID, &record.UserID, &record.UserName, &record.Name, &record.TunnelID, &record.RemoteAddr, &record.Strategy, &record.Status); err != nil {
+ return nil, err
+ }
+ if strings.TrimSpace(record.Strategy) == "" {
+ record.Strategy = "fifo"
+ }
+ out = append(out, record)
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return out, nil
+}
+
+func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) {
+ if h == nil || h.repo == nil || h.repo.DB() == nil || nodeID <= 0 {
+ return
+ }
+ if strings.TrimSpace(rawConfig) == "" {
+ return
+ }
+
+ var snapshot gostConfigSnapshot
+ if err := json.Unmarshal([]byte(rawConfig), &snapshot); err != nil {
+ return
+ }
+
+ h.cleanOrphanedServices(nodeID, snapshot.Services)
+ h.cleanOrphanedChains(nodeID, snapshot.Chains)
+ h.cleanOrphanedLimiters(nodeID, snapshot.Limiters)
+}
+
+func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) {
+ for _, item := range services {
+ name := strings.TrimSpace(item.Name)
+ if name == "" || name == "web_api" {
+ continue
+ }
+
+ parts := strings.Split(name, "_")
+ if len(parts) >= 3 {
+ forwardID, err := strconv.ParseInt(parts[0], 10, 64)
+ if err == nil && forwardID > 0 && !h.forwardExists(forwardID) {
+ _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name, parts[0] + "_" + parts[1] + "_" + parts[2], parts[0] + "_" + parts[1] + "_" + parts[2] + "_tcp", parts[0] + "_" + parts[1] + "_" + parts[2] + "_udp"}}, false, true)
+ continue
+ }
+ }
+ suffix := parts[len(parts)-1]
+
+ switch suffix {
+ case "tls":
+ tunnelID, err := strconv.ParseInt(parts[0], 10, 64)
+ if err != nil || tunnelID <= 0 || h.tunnelExists(tunnelID) {
+ continue
+ }
+ _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name}}, false, true)
+ case "tcp":
+ if len(parts) < 4 {
+ continue
+ }
+ forwardID, err := strconv.ParseInt(parts[0], 10, 64)
+ if err != nil || forwardID <= 0 || h.forwardExists(forwardID) {
+ continue
+ }
+ base := strings.TrimSuffix(name, "_tcp")
+ _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{base + "_tcp", base + "_udp"}}, false, true)
+ }
+ }
+}
+
+func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem) {
+ for _, item := range chains {
+ name := strings.TrimSpace(item.Name)
+ if name == "" {
+ continue
+ }
+
+ idx := strings.LastIndex(name, "_")
+ if idx <= 0 || idx >= len(name)-1 {
+ continue
+ }
+ tunnelID, err := strconv.ParseInt(name[idx+1:], 10, 64)
+ if err != nil || tunnelID <= 0 || h.tunnelExists(tunnelID) {
+ continue
+ }
+ _, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": name}, false, true)
+ }
+}
+
+func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem) {
+ for _, item := range limiters {
+ name := strings.TrimSpace(item.Name)
+ if name == "" || h.speedLimiterExists(name) {
+ continue
+ }
+ _, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", map[string]interface{}{"limiter": name}, false, true)
+ }
+}
+
+func (h *Handler) tunnelExists(tunnelID int64) bool {
+ var count int
+ err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE id = ?`, tunnelID).Scan(&count)
+ return err == nil && count > 0
+}
+
+func (h *Handler) forwardExists(forwardID int64) bool {
+ var count int
+ err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM forward WHERE id = ?`, forwardID).Scan(&count)
+ return err == nil && count > 0
+}
+
+func (h *Handler) speedLimiterExists(name string) bool {
+ if name == "" {
+ return false
+ }
+ id, err := strconv.ParseInt(name, 10, 64)
+ if err != nil || id <= 0 {
+ return false
+ }
+
+ var count int
+ err = h.repo.DB().QueryRow(`SELECT COUNT(1) FROM speed_limit WHERE id = ?`, id).Scan(&count)
+ return err == nil && count > 0
+}
diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go
new file mode 100644
index 0000000..570e961
--- /dev/null
+++ b/go-backend/internal/http/handler/handler.go
@@ -0,0 +1,963 @@
+package handler
+
+import (
+ "context"
+ "database/sql"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "sort"
+ "strconv"
+ "strings"
+ "sync"
+ "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
+
+ captchaMu sync.Mutex
+ captchaTokens map[string]int64
+
+ jobsMu sync.Mutex
+ jobsCancel context.CancelFunc
+ jobsStarted bool
+ jobsWG sync.WaitGroup
+}
+
+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),
+ captchaTokens: make(map[string]int64),
+ }
+}
+
+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/user/list", h.userList)
+ mux.HandleFunc("/api/v1/user/create", h.userCreate)
+ mux.HandleFunc("/api/v1/user/update", h.userUpdate)
+ mux.HandleFunc("/api/v1/user/delete", h.userDelete)
+ mux.HandleFunc("/api/v1/user/reset", h.userResetFlow)
+ 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/captcha/generate", h.captchaGenerate)
+ mux.HandleFunc("/api/v1/captcha/verify", h.captchaVerify)
+ mux.HandleFunc("/api/v1/user/package", h.userPackage)
+ mux.HandleFunc("/api/v1/user/updatePassword", h.updatePassword)
+ mux.HandleFunc("/api/v1/node/list", h.nodeList)
+ mux.HandleFunc("/api/v1/node/create", h.nodeCreate)
+ mux.HandleFunc("/api/v1/node/update", h.nodeUpdate)
+ mux.HandleFunc("/api/v1/node/delete", h.nodeDelete)
+ mux.HandleFunc("/api/v1/node/install", h.nodeInstall)
+ mux.HandleFunc("/api/v1/node/update-order", h.nodeUpdateOrder)
+ mux.HandleFunc("/api/v1/node/batch-delete", h.nodeBatchDelete)
+ mux.HandleFunc("/api/v1/node/check-status", h.nodeCheckStatus)
+ mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
+ mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
+ mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
+ mux.HandleFunc("/api/v1/tunnel/update", h.tunnelUpdate)
+ mux.HandleFunc("/api/v1/tunnel/delete", h.tunnelDelete)
+ mux.HandleFunc("/api/v1/tunnel/diagnose", h.tunnelDiagnose)
+ mux.HandleFunc("/api/v1/tunnel/update-order", h.tunnelUpdateOrder)
+ mux.HandleFunc("/api/v1/tunnel/batch-delete", h.tunnelBatchDelete)
+ mux.HandleFunc("/api/v1/tunnel/batch-redeploy", h.tunnelBatchRedeploy)
+ mux.HandleFunc("/api/v1/tunnel/user/assign", h.userTunnelAssign)
+ mux.HandleFunc("/api/v1/tunnel/user/batch-assign", h.userTunnelBatchAssign)
+ mux.HandleFunc("/api/v1/tunnel/user/remove", h.userTunnelRemove)
+ mux.HandleFunc("/api/v1/tunnel/user/update", h.userTunnelUpdate)
+ mux.HandleFunc("/api/v1/forward/list", h.forwardList)
+ mux.HandleFunc("/api/v1/forward/create", h.forwardCreate)
+ mux.HandleFunc("/api/v1/forward/update", h.forwardUpdate)
+ mux.HandleFunc("/api/v1/forward/delete", h.forwardDelete)
+ mux.HandleFunc("/api/v1/forward/force-delete", h.forwardForceDelete)
+ mux.HandleFunc("/api/v1/forward/pause", h.forwardPause)
+ mux.HandleFunc("/api/v1/forward/resume", h.forwardResume)
+ mux.HandleFunc("/api/v1/forward/diagnose", h.forwardDiagnose)
+ mux.HandleFunc("/api/v1/forward/update-order", h.forwardUpdateOrder)
+ mux.HandleFunc("/api/v1/forward/batch-delete", h.forwardBatchDelete)
+ mux.HandleFunc("/api/v1/forward/batch-pause", h.forwardBatchPause)
+ mux.HandleFunc("/api/v1/forward/batch-resume", h.forwardBatchResume)
+ mux.HandleFunc("/api/v1/forward/batch-redeploy", h.forwardBatchRedeploy)
+ mux.HandleFunc("/api/v1/forward/batch-change-tunnel", h.forwardBatchChangeTunnel)
+ mux.HandleFunc("/api/v1/speed-limit/list", h.speedLimitList)
+ mux.HandleFunc("/api/v1/speed-limit/create", h.speedLimitCreate)
+ mux.HandleFunc("/api/v1/speed-limit/update", h.speedLimitUpdate)
+ mux.HandleFunc("/api/v1/speed-limit/delete", h.speedLimitDelete)
+ mux.HandleFunc("/api/v1/speed-limit/tunnels", h.tunnelList)
+ mux.HandleFunc("/api/v1/tunnel/user/tunnel", h.userTunnelVisibleList)
+ mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList)
+ mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList)
+ mux.HandleFunc("/api/v1/group/tunnel/create", h.groupTunnelCreate)
+ mux.HandleFunc("/api/v1/group/tunnel/update", h.groupTunnelUpdate)
+ mux.HandleFunc("/api/v1/group/tunnel/delete", h.groupTunnelDelete)
+ mux.HandleFunc("/api/v1/group/tunnel/assign", h.groupTunnelAssign)
+ mux.HandleFunc("/api/v1/group/user/list", h.userGroupList)
+ mux.HandleFunc("/api/v1/group/user/create", h.groupUserCreate)
+ mux.HandleFunc("/api/v1/group/user/update", h.groupUserUpdate)
+ mux.HandleFunc("/api/v1/group/user/delete", h.groupUserDelete)
+ mux.HandleFunc("/api/v1/group/user/assign", h.groupUserAssign)
+ mux.HandleFunc("/api/v1/group/permission/list", h.groupPermissionList)
+ mux.HandleFunc("/api/v1/group/permission/assign", h.groupPermissionAssign)
+ mux.HandleFunc("/api/v1/group/permission/remove", h.groupPermissionRemove)
+ mux.HandleFunc("/api/v1/open_api/sub_store", h.openAPISubStore)
+
+ mux.HandleFunc("/flow/test", h.flowTest)
+ mux.HandleFunc("/flow/config", h.flowConfig)
+ mux.HandleFunc("/flow/upload", h.flowUpload)
+ mux.HandleFunc("/error", h.errorPage)
+}
+
+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
+ }
+ if captchaEnabled && !h.consumeCaptchaToken(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) userList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ users, err := h.repo.ListUsers()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(users))
+}
+
+func (h *Handler) nodeList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ items, err := h.repo.ListNodes()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) tunnelList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ items, err := h.repo.ListTunnels()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) forwardList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ userID, roleID, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+
+ items, err := h.repo.ListForwards()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if roleID != 0 {
+ filtered := make([]map[string]interface{}, 0, len(items))
+ for _, item := range items {
+ if asInt64(item["userId"], 0) == userID {
+ filtered = append(filtered, item)
+ }
+ }
+ items = filtered
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) speedLimitList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ items, err := h.repo.ListSpeedLimits()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodGet {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ if h == nil || h.repo == nil || h.repo.DB() == nil {
+ response.WriteJSON(w, response.Err(-2, "database unavailable"))
+ return
+ }
+
+ username := strings.TrimSpace(r.URL.Query().Get("user"))
+ password := strings.TrimSpace(r.URL.Query().Get("pwd"))
+ tunnel := strings.TrimSpace(r.URL.Query().Get("tunnel"))
+ if tunnel == "" {
+ tunnel = "-1"
+ }
+
+ if username == "" {
+ response.WriteJSON(w, response.ErrDefault("用户不能为空"))
+ return
+ }
+ if password == "" {
+ response.WriteJSON(w, response.ErrDefault("密码不能为空"))
+ return
+ }
+
+ user, err := h.repo.GetUserByUsername(username)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if user == nil || user.Pwd != security.MD5(password) {
+ response.WriteJSON(w, response.ErrDefault("鉴权失败"))
+ return
+ }
+
+ const giga = int64(1024 * 1024 * 1024)
+ headerValue := ""
+
+ if tunnel == "-1" {
+ headerValue = buildSubscriptionHeader(user.OutFlow, user.InFlow, user.Flow*giga, user.ExpTime/1000)
+ } else {
+ tunnelID, parseErr := strconv.ParseInt(tunnel, 10, 64)
+ if parseErr != nil || tunnelID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+ return
+ }
+
+ var userID int64
+ var inFlow int64
+ var outFlow int64
+ var flow int64
+ var expTime int64
+ err = h.repo.DB().QueryRow(`SELECT user_id, in_flow, out_flow, flow, exp_time FROM user_tunnel WHERE id = ? LIMIT 1`, tunnelID).
+ Scan(&userID, &inFlow, &outFlow, &flow, &expTime)
+ if err != nil {
+ if err == sql.ErrNoRows {
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if userID != user.ID {
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+ return
+ }
+
+ headerValue = buildSubscriptionHeader(outFlow, inFlow, flow*giga, expTime/1000)
+ }
+
+ w.Header().Set("subscription-userinfo", headerValue)
+ w.Header().Set("Content-Type", "text/plain; charset=utf-8")
+ _, _ = w.Write([]byte(headerValue))
+}
+
+func (h *Handler) errorPage(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "text/html; charset=UTF-8")
+ w.WriteHeader(http.StatusNotFound)
+ _, _ = w.Write([]byte("
错误 404"))
+}
+
+func buildSubscriptionHeader(upload, download, total, expire int64) string {
+ return fmt.Sprintf("upload=%d; download=%d; total=%d; expire=%d", download, upload, total, expire)
+}
+
+func (h *Handler) userTunnelVisibleList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ userID, roleID, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+
+ items := make([]map[string]interface{}, 0)
+ if roleID == 0 {
+ items, err = h.repo.ListEnabledTunnelSummaries()
+ } else {
+ items, err = h.repo.ListUserAccessibleTunnels(userID)
+ }
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) userTunnelList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ var req struct {
+ UserID int64 `json:"userId"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ if req.UserID <= 0 {
+ response.WriteJSON(w, response.OK([]interface{}{}))
+ return
+ }
+
+ tunnels, err := h.repo.GetUserPackageTunnels(req.UserID)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+
+ out := 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,
+ "status": 1,
+ "flow": t.Flow,
+ "num": t.Num,
+ "expTime": t.ExpTime,
+ "flowResetTime": t.FlowResetTime,
+ "inFlow": t.InFlow,
+ "outFlow": t.OutFlow,
+ "tunnelFlow": t.TunnelFlow,
+ "speedId": nil,
+ "speedLimitName": nil,
+ }
+ if t.SpeedID.Valid {
+ item["speedId"] = t.SpeedID.Int64
+ }
+ if t.SpeedLimit.Valid {
+ item["speedLimitName"] = t.SpeedLimit.String
+ }
+ out = append(out, item)
+ }
+ response.WriteJSON(w, response.OK(out))
+}
+
+func (h *Handler) tunnelGroupList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ items, err := h.repo.ListTunnelGroups()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) userGroupList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ items, err := h.repo.ListUserGroups()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) groupPermissionList(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+
+ items, err := h.repo.ListGroupPermissions()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+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")
+ node, err := h.repo.GetNodeBySecret(secret)
+ if err != nil || node == nil {
+ w.Header().Set("Content-Type", "text/plain; charset=utf-8")
+ _, _ = w.Write([]byte("ok"))
+ return
+ }
+
+ rawData, err := readAndDecryptFlowBody(r.Body, secret)
+ if err == nil && strings.TrimSpace(rawData) != "" {
+ h.cleanNodeConfigs(node.ID, rawData)
+ }
+ 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 {
+ h.processFlowItem(item)
+ }
+ }
+ }
+
+ 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 userIDFromRequest(r *http.Request) (int64, error) {
+ claims, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims)
+ if !ok {
+ return 0, strconv.ErrSyntax
+ }
+ return parseUserID(claims.Sub)
+}
+
+func userRoleFromRequest(r *http.Request) (int64, int, error) {
+ claims, ok := r.Context().Value(middleware.ClaimsContextKey).(auth.Claims)
+ if !ok {
+ return 0, 0, strconv.ErrSyntax
+ }
+ userID, err := parseUserID(claims.Sub)
+ if err != nil {
+ return 0, 0, err
+ }
+ return userID, claims.RoleID, nil
+}
+
+func nullableNullInt64(v sql.NullInt64) interface{} {
+ if v.Valid {
+ return v.Int64
+ }
+ return nil
+}
+
+func readAndDecryptFlowBody(body io.ReadCloser, secret string) (string, error) {
+ defer body.Close()
+ raw, err := io.ReadAll(body)
+ if err != nil {
+ return "", err
+ }
+ text := strings.TrimSpace(string(raw))
+ if text == "" {
+ return "", nil
+ }
+
+ var wrap struct {
+ Encrypted bool `json:"encrypted"`
+ Data string `json:"data"`
+ Timestamp int64 `json:"timestamp"`
+ }
+ if err := json.Unmarshal(raw, &wrap); err != nil || !wrap.Encrypted || strings.TrimSpace(wrap.Data) == "" {
+ return text, nil
+ }
+
+ crypto, err := security.NewAESCrypto(secret)
+ if err != nil {
+ return text, nil
+ }
+ plain, err := crypto.Decrypt(wrap.Data)
+ if err != nil {
+ return text, nil
+ }
+ return string(plain), nil
+}
diff --git a/go-backend/internal/http/handler/jobs.go b/go-backend/internal/http/handler/jobs.go
new file mode 100644
index 0000000..f665918
--- /dev/null
+++ b/go-backend/internal/http/handler/jobs.go
@@ -0,0 +1,270 @@
+package handler
+
+import (
+ "context"
+ "database/sql"
+ "time"
+)
+
+func (h *Handler) StartBackgroundJobs() {
+ if h == nil || h.repo == nil || h.repo.DB() == nil {
+ return
+ }
+
+ h.jobsMu.Lock()
+ if h.jobsStarted {
+ h.jobsMu.Unlock()
+ return
+ }
+ ctx, cancel := context.WithCancel(context.Background())
+ h.jobsCancel = cancel
+ h.jobsStarted = true
+ h.jobsWG.Add(2)
+ h.jobsMu.Unlock()
+
+ go h.runHourlyStatsLoop(ctx)
+ go h.runDailyMaintenanceLoop(ctx)
+}
+
+func (h *Handler) StopBackgroundJobs() {
+ if h == nil {
+ return
+ }
+
+ h.jobsMu.Lock()
+ if !h.jobsStarted {
+ h.jobsMu.Unlock()
+ return
+ }
+ cancel := h.jobsCancel
+ h.jobsCancel = nil
+ h.jobsStarted = false
+ h.jobsMu.Unlock()
+
+ if cancel != nil {
+ cancel()
+ }
+ h.jobsWG.Wait()
+}
+
+func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
+ defer h.jobsWG.Done()
+
+ for {
+ wait := durationUntilNextHour(time.Now())
+ timer := time.NewTimer(wait)
+ select {
+ case <-ctx.Done():
+ if !timer.Stop() {
+ <-timer.C
+ }
+ return
+ case <-timer.C:
+ h.runStatisticsFlowJob(time.Now())
+ }
+ }
+}
+
+func (h *Handler) runDailyMaintenanceLoop(ctx context.Context) {
+ defer h.jobsWG.Done()
+
+ for {
+ wait := durationUntilNextDailyMaintenance(time.Now())
+ timer := time.NewTimer(wait)
+ select {
+ case <-ctx.Done():
+ if !timer.Stop() {
+ <-timer.C
+ }
+ return
+ case <-timer.C:
+ h.runResetAndExpiryJob(time.Now())
+ }
+ }
+}
+
+func durationUntilNextHour(now time.Time) time.Duration {
+ next := now.Truncate(time.Hour).Add(time.Hour)
+ return next.Sub(now)
+}
+
+func durationUntilNextDailyMaintenance(now time.Time) time.Duration {
+ next := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 5, 0, now.Location())
+ if !next.After(now) {
+ next = next.Add(24 * time.Hour)
+ }
+ return next.Sub(now)
+}
+
+func (h *Handler) runStatisticsFlowJob(now time.Time) {
+ if h == nil || h.repo == nil || h.repo.DB() == nil {
+ return
+ }
+
+ db := h.repo.DB()
+ nowMs := now.UnixMilli()
+ cutoffMs := nowMs - int64((48*time.Hour)/time.Millisecond)
+ _, _ = db.Exec(`DELETE FROM statistics_flow WHERE created_time < ?`, cutoffMs)
+
+ hourMark := now.Truncate(time.Hour)
+ hourText := hourMark.Format("15:04")
+ createdTime := hourMark.UnixMilli()
+
+ rows, err := db.Query(`SELECT id, in_flow, out_flow FROM user ORDER BY id ASC`)
+ if err != nil {
+ return
+ }
+ type userFlowSnapshot struct {
+ userID int64
+ inFlow int64
+ outFlow int64
+ }
+ users := make([]userFlowSnapshot, 0)
+
+ for rows.Next() {
+ var userID int64
+ var inFlow int64
+ var outFlow int64
+ if err := rows.Scan(&userID, &inFlow, &outFlow); err != nil {
+ continue
+ }
+ users = append(users, userFlowSnapshot{userID: userID, inFlow: inFlow, outFlow: outFlow})
+ }
+ _ = rows.Close()
+
+ for _, user := range users {
+ currentTotal := user.inFlow + user.outFlow
+ increment := currentTotal
+
+ var lastTotal sql.NullInt64
+ err := db.QueryRow(`SELECT total_flow FROM statistics_flow WHERE user_id = ? ORDER BY id DESC LIMIT 1`, user.userID).Scan(&lastTotal)
+ if err == nil && lastTotal.Valid {
+ increment = currentTotal - lastTotal.Int64
+ if increment < 0 {
+ increment = currentTotal
+ }
+ }
+
+ _, _ = db.Exec(`
+ INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time)
+ VALUES(?, ?, ?, ?, ?)
+ `, user.userID, increment, currentTotal, hourText, createdTime)
+ }
+}
+
+func (h *Handler) runResetAndExpiryJob(now time.Time) {
+ if h == nil || h.repo == nil || h.repo.DB() == nil {
+ return
+ }
+
+ h.resetMonthlyFlow(now)
+ h.disableExpiredUsers(now.UnixMilli())
+ h.disableExpiredUserTunnels(now.UnixMilli())
+}
+
+func (h *Handler) resetMonthlyFlow(now time.Time) {
+ db := h.repo.DB()
+ currentDay := now.Day()
+ lastDay := time.Date(now.Year(), now.Month()+1, 0, 0, 0, 0, 0, now.Location()).Day()
+
+ if currentDay == lastDay {
+ _, _ = db.Exec(`
+ UPDATE user
+ SET in_flow = 0, out_flow = 0
+ WHERE flow_reset_time != 0
+ AND (flow_reset_time = ? OR flow_reset_time > ?)
+ `, currentDay, lastDay)
+ _, _ = db.Exec(`
+ UPDATE user_tunnel
+ SET in_flow = 0, out_flow = 0
+ WHERE flow_reset_time != 0
+ AND (flow_reset_time = ? OR flow_reset_time > ?)
+ `, currentDay, lastDay)
+ return
+ }
+
+ _, _ = db.Exec(`
+ UPDATE user
+ SET in_flow = 0, out_flow = 0
+ WHERE flow_reset_time != 0
+ AND flow_reset_time = ?
+ `, currentDay)
+ _, _ = db.Exec(`
+ UPDATE user_tunnel
+ SET in_flow = 0, out_flow = 0
+ WHERE flow_reset_time != 0
+ AND flow_reset_time = ?
+ `, currentDay)
+}
+
+func (h *Handler) disableExpiredUsers(nowMs int64) {
+ db := h.repo.DB()
+ rows, err := db.Query(`
+ SELECT id
+ FROM user
+ WHERE role_id != 0
+ AND status = 1
+ AND exp_time IS NOT NULL
+ AND exp_time < ?
+ `, nowMs)
+ if err != nil {
+ return
+ }
+ userIDs := make([]int64, 0)
+
+ for rows.Next() {
+ var userID int64
+ if err := rows.Scan(&userID); err != nil {
+ continue
+ }
+ userIDs = append(userIDs, userID)
+ }
+ _ = rows.Close()
+
+ for _, userID := range userIDs {
+ forwards, err := h.listActiveForwardsByUser(userID)
+ if err == nil {
+ h.pauseForwardRecords(forwards, nowMs)
+ }
+ _, _ = db.Exec(`UPDATE user SET status = 0 WHERE id = ?`, userID)
+ }
+}
+
+func (h *Handler) disableExpiredUserTunnels(nowMs int64) {
+ db := h.repo.DB()
+ rows, err := db.Query(`
+ SELECT id, user_id, tunnel_id
+ FROM user_tunnel
+ WHERE status = 1
+ AND exp_time IS NOT NULL
+ AND exp_time < ?
+ `, nowMs)
+ if err != nil {
+ return
+ }
+ type expiredUserTunnel struct {
+ userTunnelID int64
+ userID int64
+ tunnelID int64
+ }
+ items := make([]expiredUserTunnel, 0)
+
+ for rows.Next() {
+ var userTunnelID int64
+ var userID int64
+ var tunnelID int64
+ if err := rows.Scan(&userTunnelID, &userID, &tunnelID); err != nil {
+ continue
+ }
+ items = append(items, expiredUserTunnel{userTunnelID: userTunnelID, userID: userID, tunnelID: tunnelID})
+ }
+ _ = rows.Close()
+
+ for _, item := range items {
+ forwards, err := h.listActiveForwardsByUserTunnel(item.userID, item.tunnelID)
+ if err == nil {
+ h.pauseForwardRecords(forwards, nowMs)
+ }
+ _, _ = db.Exec(`UPDATE user_tunnel SET status = 0 WHERE id = ?`, item.userTunnelID)
+ }
+}
diff --git a/go-backend/internal/http/handler/jobs_test.go b/go-backend/internal/http/handler/jobs_test.go
new file mode 100644
index 0000000..d349f8c
--- /dev/null
+++ b/go-backend/internal/http/handler/jobs_test.go
@@ -0,0 +1,128 @@
+package handler
+
+import (
+ "path/filepath"
+ "testing"
+ "time"
+
+ "go-backend/internal/store/sqlite"
+)
+
+func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) {
+ dbPath := filepath.Join(t.TempDir(), "jobs-stats.db")
+ repo, err := sqlite.Open(dbPath)
+ if err != nil {
+ t.Fatalf("open sqlite: %v", err)
+ }
+ t.Cleanup(func() { _ = repo.Close() })
+
+ h := New(repo, "secret")
+ now := time.Date(2026, 2, 7, 12, 0, 0, 0, time.UTC)
+ nowMs := now.UnixMilli()
+
+ if _, err := repo.DB().Exec(`UPDATE user SET in_flow = 100, out_flow = 200 WHERE id = 1`); err != nil {
+ t.Fatalf("seed user flow: %v", err)
+ }
+
+ if _, err := repo.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 250, 250, '11:00', ?)`, now.Add(-time.Hour).UnixMilli()); err != nil {
+ t.Fatalf("seed recent statistics row: %v", err)
+ }
+ if _, err := repo.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 10, 10, '00:00', ?)`, now.Add(-49*time.Hour).UnixMilli()); err != nil {
+ t.Fatalf("seed stale statistics row: %v", err)
+ }
+
+ h.runStatisticsFlowJob(now)
+
+ var staleCount int
+ if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM statistics_flow WHERE created_time < ?`, nowMs-int64((48*time.Hour)/time.Millisecond)).Scan(&staleCount); err != nil {
+ t.Fatalf("query stale statistics rows: %v", err)
+ }
+ if staleCount != 0 {
+ t.Fatalf("expected stale statistics rows to be pruned, got %d", staleCount)
+ }
+
+ var flow int64
+ var total int64
+ var hour string
+ if err := repo.DB().QueryRow(`SELECT flow, total_flow, time FROM statistics_flow WHERE user_id = 1 ORDER BY id DESC LIMIT 1`).Scan(&flow, &total, &hour); err != nil {
+ t.Fatalf("query latest statistics row: %v", err)
+ }
+ if flow != 50 {
+ t.Fatalf("expected increment flow 50, got %d", flow)
+ }
+ if total != 300 {
+ t.Fatalf("expected total flow 300, got %d", total)
+ }
+ if hour != "12:00" {
+ t.Fatalf("expected hour mark 12:00, got %s", hour)
+ }
+}
+
+func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
+ dbPath := filepath.Join(t.TempDir(), "jobs-reset.db")
+ repo, err := sqlite.Open(dbPath)
+ if err != nil {
+ t.Fatalf("open sqlite: %v", err)
+ }
+ t.Cleanup(func() { _ = repo.Close() })
+
+ h := New(repo, "secret")
+ now := time.Date(2026, 3, 15, 0, 0, 5, 0, time.UTC)
+ nowMs := now.UnixMilli()
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
+ VALUES(2, 'expired_user', 'x', 1, ?, 100, 1000, 2000, 15, 1, ?, ?, 1)
+ `, nowMs-1000, nowMs, nowMs); err != nil {
+ t.Fatalf("insert expired user: %v", err)
+ }
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
+ VALUES(1, 't1', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
+ `, nowMs, nowMs); err != nil {
+ t.Fatalf("insert tunnel: %v", err)
+ }
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
+ VALUES(10, 2, 1, NULL, 1, 1, 300, 400, 15, ?, 1)
+ `, nowMs-1000); err != nil {
+ t.Fatalf("insert expired user_tunnel: %v", err)
+ }
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
+ VALUES(20, 2, 'expired_user', 'f1', 1, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
+ `, nowMs, nowMs); err != nil {
+ t.Fatalf("insert forward: %v", err)
+ }
+
+ h.runResetAndExpiryJob(now)
+
+ var userIn, userOut int64
+ var userStatus int
+ if err := repo.DB().QueryRow(`SELECT in_flow, out_flow, status FROM user WHERE id = 2`).Scan(&userIn, &userOut, &userStatus); err != nil {
+ t.Fatalf("query user after maintenance: %v", err)
+ }
+ if userIn != 0 || userOut != 0 || userStatus != 0 {
+ t.Fatalf("expected user reset+disabled, got in=%d out=%d status=%d", userIn, userOut, userStatus)
+ }
+
+ var utIn, utOut int64
+ var utStatus int
+ if err := repo.DB().QueryRow(`SELECT in_flow, out_flow, status FROM user_tunnel WHERE id = 10`).Scan(&utIn, &utOut, &utStatus); err != nil {
+ t.Fatalf("query user_tunnel after maintenance: %v", err)
+ }
+ if utIn != 0 || utOut != 0 || utStatus != 0 {
+ t.Fatalf("expected user_tunnel reset+disabled, got in=%d out=%d status=%d", utIn, utOut, utStatus)
+ }
+
+ var forwardStatus int
+ if err := repo.DB().QueryRow(`SELECT status FROM forward WHERE id = 20`).Scan(&forwardStatus); err != nil {
+ t.Fatalf("query forward after maintenance: %v", err)
+ }
+ if forwardStatus != 0 {
+ t.Fatalf("expected forward status=0 after expiry handling, got %d", forwardStatus)
+ }
+}
diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go
new file mode 100644
index 0000000..3569756
--- /dev/null
+++ b/go-backend/internal/http/handler/mutations.go
@@ -0,0 +1,2642 @@
+package handler
+
+import (
+ "crypto/rand"
+ "database/sql"
+ "encoding/hex"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "net"
+ "net/http"
+ "sort"
+ "strconv"
+ "strings"
+ "time"
+
+ "go-backend/internal/http/response"
+ "go-backend/internal/security"
+)
+
+func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+
+ username := asString(req["user"])
+ pwd := asString(req["pwd"])
+ if username == "" || pwd == "" {
+ response.WriteJSON(w, response.ErrDefault("用户名或密码不能为空"))
+ return
+ }
+
+ db := h.repo.DB()
+ if db == nil {
+ response.WriteJSON(w, response.Err(-2, "database unavailable"))
+ return
+ }
+
+ var cnt int
+ if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&cnt); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if cnt > 0 {
+ response.WriteJSON(w, response.ErrDefault("用户名已存在"))
+ return
+ }
+
+ status := asInt(req["status"], 1)
+ flow := asInt64(req["flow"], 100)
+ num := asInt(req["num"], 10)
+ expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
+ flowResetTime := asInt64(req["flowResetTime"], 1)
+ roleID := 1
+ now := time.Now().UnixMilli()
+
+ _, err := db.Exec(`
+ INSERT INTO user(user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
+ VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?, ?, ?)
+ `, username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, now, now, status)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ if id <= 0 {
+ response.WriteJSON(w, response.ErrDefault("用户ID不能为空"))
+ return
+ }
+ username := asString(req["user"])
+ if username == "" {
+ response.WriteJSON(w, response.ErrDefault("用户名不能为空"))
+ return
+ }
+
+ db := h.repo.DB()
+ if db == nil {
+ response.WriteJSON(w, response.Err(-2, "database unavailable"))
+ return
+ }
+
+ var roleID int
+ if err := db.QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil {
+ if err == sql.ErrNoRows {
+ response.WriteJSON(w, response.ErrDefault("用户不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if roleID == 0 {
+ response.WriteJSON(w, response.ErrDefault("请不要作死"))
+ return
+ }
+
+ var cnt int
+ if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, id).Scan(&cnt); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if cnt > 0 {
+ response.WriteJSON(w, response.ErrDefault("用户名已存在"))
+ return
+ }
+
+ flow := asInt64(req["flow"], 100)
+ num := asInt(req["num"], 10)
+ expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
+ flowResetTime := asInt64(req["flowResetTime"], 1)
+ status := asInt(req["status"], 1)
+ now := time.Now().UnixMilli()
+
+ pwd := asString(req["pwd"])
+ if strings.TrimSpace(pwd) == "" {
+ _, err := db.Exec(`
+ UPDATE user
+ SET user = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ?
+ WHERE id = ?
+ `, username, flow, num, expTime, flowResetTime, status, now, id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ } else {
+ _, err := db.Exec(`
+ UPDATE user
+ SET user = ?, pwd = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ?
+ WHERE id = ?
+ `, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, now, id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ }
+
+ _, _ = db.Exec(`UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ? WHERE user_id = ?`, flow, num, expTime, flowResetTime, id)
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+
+ var roleID int
+ if err := h.repo.DB().QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil {
+ if err == sql.ErrNoRows {
+ response.WriteJSON(w, response.ErrDefault("用户不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if roleID == 0 {
+ response.WriteJSON(w, response.ErrDefault("请不要作死"))
+ return
+ }
+
+ db := h.repo.DB()
+ tx, err := db.Begin()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ defer func() { _ = tx.Rollback() }()
+
+ if _, err = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE user_id = ?)`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if _, err = tx.Exec(`DELETE FROM forward WHERE user_id = ?`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if _, err = tx.Exec(`DELETE FROM group_permission_grant WHERE user_tunnel_id IN (SELECT id FROM user_tunnel WHERE user_id = ?)`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if _, err = tx.Exec(`DELETE FROM user_tunnel WHERE user_id = ?`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if _, err = tx.Exec(`DELETE FROM user_group_user WHERE user_id = ?`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if _, err = tx.Exec(`DELETE FROM statistics_flow WHERE user_id = ?`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if _, err = tx.Exec(`DELETE FROM user WHERE id = ?`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+
+ if err = tx.Commit(); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) userResetFlow(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ typeVal := asInt(req["type"], 0)
+ if id <= 0 || (typeVal != 1 && typeVal != 2) {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+
+ db := h.repo.DB()
+ if typeVal == 1 {
+ _, _ = db.Exec(`UPDATE user SET in_flow = 0, out_flow = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id)
+ _, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE user_id = ?`, id)
+ } else {
+ _, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE id = ?`, id)
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) captchaGenerate(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ w.WriteHeader(http.StatusBadRequest)
+ _, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
+ return
+ }
+ token := randomToken(16)
+ payload := map[string]interface{}{
+ "id": token,
+ "data": map[string]interface{}{
+ "id": token,
+ },
+ "success": true,
+ }
+ w.Header().Set("Content-Type", "application/json; charset=utf-8")
+ _ = json.NewEncoder(w).Encode(payload)
+}
+
+func (h *Handler) captchaVerify(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ w.WriteHeader(http.StatusBadRequest)
+ _, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ w.WriteHeader(http.StatusBadRequest)
+ _, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
+ return
+ }
+ id := asString(req["captchaId"])
+ if id == "" {
+ id = asString(req["id"])
+ }
+ trackData := asString(req["data"])
+ if trackData == "" {
+ trackData = asString(req["trackData"])
+ }
+ if id == "" || trackData == "" {
+ w.WriteHeader(http.StatusBadRequest)
+ _, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
+ return
+ }
+ h.storeCaptchaToken(id)
+ payload := map[string]interface{}{
+ "success": true,
+ "data": map[string]interface{}{"validToken": id},
+ }
+ w.Header().Set("Content-Type", "application/json; charset=utf-8")
+ _ = json.NewEncoder(w).Encode(payload)
+}
+
+func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ name := asString(req["name"])
+ serverIP := asString(req["serverIp"])
+ if name == "" || serverIP == "" {
+ response.WriteJSON(w, response.ErrDefault("节点名称和地址不能为空"))
+ return
+ }
+
+ db := h.repo.DB()
+ now := time.Now().UnixMilli()
+ inx := nextIndex(db, "node")
+ _, err := db.Exec(`
+ INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `,
+ name,
+ randomToken(16),
+ serverIP,
+ nullableText(asString(req["serverIpV4"])),
+ nullableText(asString(req["serverIpV6"])),
+ defaultString(asString(req["port"]), "1000-65535"),
+ nullableText(asString(req["interfaceName"])),
+ nullableText(""),
+ asInt(req["http"], 0),
+ asInt(req["tls"], 0),
+ asInt(req["socks"], 0),
+ now,
+ now,
+ 0,
+ defaultString(asString(req["tcpListenAddr"]), "[::]"),
+ defaultString(asString(req["udpListenAddr"]), "[::]"),
+ inx,
+ )
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ if id <= 0 {
+ response.WriteJSON(w, response.ErrDefault("节点ID不能为空"))
+ return
+ }
+
+ var currentStatus int
+ var currentHTTP int
+ var currentTLS int
+ var currentSocks int
+ if err := h.repo.DB().QueryRow(`SELECT status, http, tls, socks FROM node WHERE id = ?`, id).Scan(¤tStatus, ¤tHTTP, ¤tTLS, ¤tSocks); err != nil {
+ if err == sql.ErrNoRows {
+ response.WriteJSON(w, response.ErrDefault("节点不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+
+ newHTTP := asInt(req["http"], currentHTTP)
+ newTLS := asInt(req["tls"], currentTLS)
+ newSocks := asInt(req["socks"], currentSocks)
+ if currentStatus == 1 && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
+ if err := h.applyNodeProtocolChange(id, newHTTP, newTLS, newSocks); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ }
+
+ now := time.Now().UnixMilli()
+ _, err := h.repo.DB().Exec(`
+ UPDATE node
+ SET name = ?, server_ip = ?, server_ip_v4 = ?, server_ip_v6 = ?, port = ?, interface_name = ?, http = ?, tls = ?, socks = ?, tcp_listen_addr = ?, udp_listen_addr = ?, updated_time = ?
+ WHERE id = ?
+ `,
+ asString(req["name"]),
+ asString(req["serverIp"]),
+ nullableText(asString(req["serverIpV4"])),
+ nullableText(asString(req["serverIpV6"])),
+ defaultString(asString(req["port"]), "1000-65535"),
+ nullableText(asString(req["interfaceName"])),
+ newHTTP,
+ newTLS,
+ newSocks,
+ defaultString(asString(req["tcpListenAddr"]), "[::]"),
+ defaultString(asString(req["udpListenAddr"]), "[::]"),
+ now,
+ id,
+ )
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) nodeDelete(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ if err := h.deleteNodeByID(id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) nodeInstall(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ db := h.repo.DB()
+ var secret string
+ if err := db.QueryRow(`SELECT secret FROM node WHERE id = ?`, id).Scan(&secret); err != nil {
+ response.WriteJSON(w, response.ErrDefault("节点不存在"))
+ return
+ }
+ var panelAddr string
+ if err := db.QueryRow(`SELECT value FROM vite_config WHERE name = 'ip' LIMIT 1`).Scan(&panelAddr); err != nil {
+ if err == sql.ErrNoRows {
+ response.WriteJSON(w, response.ErrDefault("请先前往网站配置中设置ip"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ cmd := fmt.Sprintf("curl -L https://github.com/Sagit-chu/flux-panel/releases/latest/download/install.sh -o ./install.sh && chmod +x ./install.sh && ./install.sh -a %s -s %s", processServerAddress(panelAddr), secret)
+ response.WriteJSON(w, response.OK(cmd))
+}
+
+func (h *Handler) nodeUpdateOrder(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req struct {
+ Nodes []struct {
+ ID int64 `json:"id"`
+ Inx int `json:"inx"`
+ } `json:"nodes"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ for _, n := range req.Nodes {
+ _, _ = h.repo.DB().Exec(`UPDATE node SET inx = ?, updated_time = ? WHERE id = ?`, n.Inx, time.Now().UnixMilli(), n.ID)
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) nodeBatchDelete(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ ids := idsFromBody(r, w)
+ if ids == nil {
+ return
+ }
+ for _, id := range ids {
+ _ = h.deleteNodeByID(id)
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) nodeCheckStatus(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ items, err := h.repo.ListNodes()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(items))
+}
+
+func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ name := asString(req["name"])
+ if name == "" {
+ response.WriteJSON(w, response.ErrDefault("隧道名称不能为空"))
+ return
+ }
+ var tunnelNameDup int
+ if err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, name).Scan(&tunnelNameDup); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if tunnelNameDup > 0 {
+ response.WriteJSON(w, response.ErrDefault("隧道名称重复"))
+ return
+ }
+
+ typeVal := asInt(req["type"], 1)
+ flow := asInt64(req["flow"], 1)
+ status := asInt(req["status"], 1)
+ trafficRatio := asFloat(req["trafficRatio"], 1.0)
+ inIP := asString(req["inIp"])
+ now := time.Now().UnixMilli()
+ inx := nextIndex(h.repo.DB(), "tunnel")
+
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ defer func() { _ = tx.Rollback() }()
+
+ runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal)
+ if err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ if strings.TrimSpace(inIP) == "" {
+ inIP = buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes)
+ }
+
+ res, err := tx.Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
+ name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ tunnelID, _ := res.LastInsertId()
+ runtimeState.TunnelID = tunnelID
+ applyTunnelPortsToRequest(req, runtimeState)
+ if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := tx.Commit(); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if typeVal == 2 {
+ createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
+ if applyErr != nil {
+ h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID)
+ _ = h.deleteTunnelByID(tunnelID)
+ response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
+ return
+ }
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) tunnelGet(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ items, err := h.repo.ListTunnels()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ for _, it := range items {
+ if asInt64(it["id"], 0) == id {
+ response.WriteJSON(w, response.OK(it))
+ return
+ }
+ }
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+}
+
+func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ if id <= 0 {
+ response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
+ return
+ }
+ now := time.Now().UnixMilli()
+ _, err := h.repo.DB().Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, updated_time=? WHERE id=?`,
+ asString(req["name"]), asInt(req["type"], 1), asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(asString(req["inIp"])), now, id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ defer func() { _ = tx.Rollback() }()
+ if _, err := tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := replaceTunnelChainsTx(tx, id, req); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := tx.Commit(); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) tunnelDelete(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ if err := h.deleteTunnelByID(id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) tunnelDiagnose(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ id := asInt64FromBodyKey(r, w, "tunnelId")
+ if id <= 0 {
+ return
+ }
+ result, err := h.diagnoseTunnelRuntime(id)
+ if err != nil {
+ if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不完整") {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(result))
+}
+
+func (h *Handler) tunnelUpdateOrder(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return
+ }
+ var req struct {
+ Tunnels []struct {
+ ID int64 `json:"id"`
+ Inx int `json:"inx"`
+ } `json:"tunnels"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ for _, t := range req.Tunnels {
+ _, _ = h.repo.DB().Exec(`UPDATE tunnel SET inx = ?, updated_time = ? WHERE id = ?`, t.Inx, time.Now().UnixMilli(), t.ID)
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) {
+ ids := idsFromBody(r, w)
+ if ids == nil {
+ return
+ }
+ success := 0
+ fail := 0
+ for _, id := range ids {
+ if err := h.deleteTunnelByID(id); err != nil {
+ fail++
+ } else {
+ success++
+ }
+ }
+ response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
+}
+
+func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
+ ids := idsFromBody(r, w)
+ if ids == nil {
+ return
+ }
+ success := 0
+ fail := 0
+ for _, tunnelID := range ids {
+ forwards, err := h.listForwardsByTunnel(tunnelID)
+ if err != nil {
+ fail++
+ continue
+ }
+ if len(forwards) == 0 {
+ success++
+ continue
+ }
+ ok := true
+ for i := range forwards {
+ if err := h.syncForwardServices(&forwards[i], "UpdateService", true); err != nil {
+ ok = false
+ break
+ }
+ }
+ if ok {
+ success++
+ } else {
+ fail++
+ }
+ }
+ response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
+}
+
+func (h *Handler) userTunnelAssign(w http.ResponseWriter, r *http.Request) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ if err := h.upsertUserTunnel(req); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) userTunnelBatchAssign(w http.ResponseWriter, r *http.Request) {
+ var req struct {
+ UserID int64 `json:"userId"`
+ Tunnels []struct {
+ TunnelID int64 `json:"tunnelId"`
+ SpeedID *int64 `json:"speedId"`
+ } `json:"tunnels"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil || req.UserID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ for _, t := range req.Tunnels {
+ m := map[string]interface{}{"userId": req.UserID, "tunnelId": t.TunnelID}
+ if t.SpeedID != nil {
+ m["speedId"] = *t.SpeedID
+ }
+ if err := h.upsertUserTunnel(m); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) userTunnelRemove(w http.ResponseWriter, r *http.Request) {
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ _, err := h.repo.DB().Exec(`DELETE FROM user_tunnel WHERE id = ?`, id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ if id <= 0 {
+ response.WriteJSON(w, response.ErrDefault("权限ID不能为空"))
+ return
+ }
+ _, err := h.repo.DB().Exec(`
+ UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, speed_id = ?, status = ? WHERE id = ?
+ `,
+ asInt64(req["flow"], 0),
+ asInt(req["num"], 0),
+ asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
+ asInt64(req["flowResetTime"], 1),
+ nullableInt(asAnyToInt64Ptr(req["speedId"])),
+ asInt(req["status"], 1),
+ id,
+ )
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ userID, roleID, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+ tunnelID := asInt64(req["tunnelId"], 0)
+ if tunnelID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
+ return
+ }
+ if err := h.ensureTunnelPermission(userID, roleID, tunnelID); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ tunnel, err := h.getTunnelRecord(tunnelID)
+ if err != nil {
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+ return
+ }
+ if tunnel.Status != 1 {
+ response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法创建转发"))
+ return
+ }
+ name := asString(req["name"])
+ remoteAddr := asString(req["remoteAddr"])
+ if name == "" || remoteAddr == "" {
+ response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
+ return
+ }
+ port := asInt(req["inPort"], 0)
+ if port <= 0 {
+ port = h.pickTunnelPort(tunnelID)
+ }
+ if port <= 0 {
+ port = 10000
+ }
+ now := time.Now().UnixMilli()
+ inx := nextIndex(h.repo.DB(), "forward")
+ var userName string
+ _ = h.repo.DB().QueryRow(`SELECT user FROM user WHERE id = ?`, userID).Scan(&userName)
+ if userName == "" {
+ userName = "user"
+ }
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ defer func() { _ = tx.Rollback() }()
+ res, err := tx.Exec(`
+ INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
+ VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
+ `, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ forwardID, _ := res.LastInsertId()
+ entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
+ for _, nodeID := range entryNodes {
+ _, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port)
+ }
+ if err := tx.Commit(); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ createdForward, err := h.getForwardRecord(forwardID)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := h.syncForwardServices(createdForward, "AddService", false); err != nil {
+ _ = h.deleteForwardByID(forwardID)
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ if id <= 0 {
+ response.WriteJSON(w, response.ErrDefault("转发ID不能为空"))
+ return
+ }
+ forward, actorUserID, actorRole, err := h.resolveForwardAccess(r, id)
+ if err != nil {
+ if errors.Is(err, errForwardNotFound) {
+ response.WriteJSON(w, response.ErrDefault("转发不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+
+ tunnelID := asInt64(req["tunnelId"], forward.TunnelID)
+ if tunnelID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
+ return
+ }
+ if err := h.ensureTunnelPermission(actorUserID, actorRole, tunnelID); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ tunnel, err := h.getTunnelRecord(tunnelID)
+ if err != nil {
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+ return
+ }
+ if tunnel.Status != 1 {
+ response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法更新转发"))
+ return
+ }
+
+ name := strings.TrimSpace(asString(req["name"]))
+ if name == "" {
+ name = forward.Name
+ }
+ remoteAddr := strings.TrimSpace(asString(req["remoteAddr"]))
+ if remoteAddr == "" {
+ remoteAddr = forward.RemoteAddr
+ }
+ strategy := strings.TrimSpace(asString(req["strategy"]))
+ if strategy == "" {
+ strategy = forward.Strategy
+ }
+
+ port := asInt(req["inPort"], 0)
+ if port <= 0 {
+ var minPort sql.NullInt64
+ _ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&minPort)
+ if minPort.Valid {
+ port = int(minPort.Int64)
+ }
+ if port <= 0 {
+ port = h.pickTunnelPort(tunnelID)
+ }
+ }
+ now := time.Now().UnixMilli()
+ _, err = h.repo.DB().Exec(`
+ UPDATE forward SET name = ?, tunnel_id = ?, remote_addr = ?, strategy = ?, updated_time = ? WHERE id = ?
+ `, name, tunnelID, remoteAddr, strategy, now, id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ _ = h.replaceForwardPorts(id, tunnelID, port)
+ updatedForward, err := h.getForwardRecord(id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ forward, _, _, err := h.resolveForwardAccess(r, id)
+ if err != nil {
+ if errors.Is(err, errForwardNotFound) {
+ response.WriteJSON(w, response.ErrDefault("转发不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ if err := h.deleteForwardByID(id); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
+ h.forwardDelete(w, r)
+}
+
+func (h *Handler) forwardPause(w http.ResponseWriter, r *http.Request) {
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ forward, _, _, err := h.resolveForwardAccess(r, id)
+ if err != nil {
+ if errors.Is(err, errForwardNotFound) {
+ response.WriteJSON(w, response.ErrDefault("转发不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id)
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) forwardResume(w http.ResponseWriter, r *http.Request) {
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ forward, _, _, err := h.resolveForwardAccess(r, id)
+ if err != nil {
+ if errors.Is(err, errForwardNotFound) {
+ response.WriteJSON(w, response.ErrDefault("转发不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id)
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) forwardDiagnose(w http.ResponseWriter, r *http.Request) {
+ id := asInt64FromBodyKey(r, w, "forwardId")
+ if id <= 0 {
+ return
+ }
+ forward, _, _, err := h.resolveForwardAccess(r, id)
+ if err != nil {
+ if errors.Is(err, errForwardNotFound) {
+ response.WriteJSON(w, response.ErrDefault("转发不存在"))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ payload, err := h.diagnoseForwardRuntime(forward)
+ if err != nil {
+ if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不能为空") || strings.Contains(err.Error(), "错误") {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OK(payload))
+}
+
+func (h *Handler) forwardUpdateOrder(w http.ResponseWriter, r *http.Request) {
+ var req struct {
+ Forwards []struct {
+ ID int64 `json:"id"`
+ Inx int `json:"inx"`
+ } `json:"forwards"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ for _, f := range req.Forwards {
+ _, _ = h.repo.DB().Exec(`UPDATE forward SET inx = ?, updated_time = ? WHERE id = ?`, f.Inx, time.Now().UnixMilli(), f.ID)
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
+ ids := idsFromBody(r, w)
+ if ids == nil {
+ return
+ }
+ actorUserID, actorRole, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+ s := 0
+ f := 0
+ for _, id := range ids {
+ forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
+ if accessErr != nil {
+ f++
+ continue
+ }
+ if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
+ f++
+ continue
+ }
+ if err := h.deleteForwardByID(id); err != nil {
+ f++
+ } else {
+ s++
+ }
+ }
+ response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
+}
+
+func (h *Handler) forwardBatchPause(w http.ResponseWriter, r *http.Request) {
+ ids := idsFromBody(r, w)
+ if ids == nil {
+ return
+ }
+ actorUserID, actorRole, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+ s := 0
+ f := 0
+ for _, id := range ids {
+ forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
+ if accessErr != nil {
+ f++
+ continue
+ }
+ if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
+ f++
+ continue
+ }
+ if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil {
+ f++
+ } else {
+ s++
+ }
+ }
+ response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
+}
+
+func (h *Handler) forwardBatchResume(w http.ResponseWriter, r *http.Request) {
+ ids := idsFromBody(r, w)
+ if ids == nil {
+ return
+ }
+ actorUserID, actorRole, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+ s := 0
+ f := 0
+ for _, id := range ids {
+ forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
+ if accessErr != nil {
+ f++
+ continue
+ }
+ if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
+ f++
+ continue
+ }
+ if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil {
+ f++
+ } else {
+ s++
+ }
+ }
+ response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
+}
+
+func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
+ ids := idsFromBody(r, w)
+ if ids == nil {
+ return
+ }
+ actorUserID, actorRole, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+ s := 0
+ f := 0
+ for _, id := range ids {
+ forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
+ if accessErr != nil {
+ f++
+ continue
+ }
+ if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
+ f++
+ } else {
+ s++
+ }
+ }
+ response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
+}
+
+func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Request) {
+ var req struct {
+ ForwardIDs []int64 `json:"forwardIds"`
+ TargetTunnelID int64 `json:"targetTunnelId"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil || req.TargetTunnelID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ actorUserID, actorRole, err := userRoleFromRequest(r)
+ if err != nil {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+ if err := h.ensureTunnelPermission(actorUserID, actorRole, req.TargetTunnelID); err != nil {
+ response.WriteJSON(w, response.ErrDefault(err.Error()))
+ return
+ }
+ targetTunnel, err := h.getTunnelRecord(req.TargetTunnelID)
+ if err != nil {
+ response.WriteJSON(w, response.ErrDefault("目标隧道不存在"))
+ return
+ }
+ if targetTunnel.Status != 1 {
+ response.WriteJSON(w, response.ErrDefault("目标隧道已禁用"))
+ return
+ }
+ success := 0
+ fail := 0
+ for _, id := range req.ForwardIDs {
+ if id <= 0 {
+ continue
+ }
+ forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
+ if accessErr != nil {
+ fail++
+ continue
+ }
+ if forward.TunnelID == req.TargetTunnelID {
+ fail++
+ continue
+ }
+ var port sql.NullInt64
+ _ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&port)
+ _, err := h.repo.DB().Exec(`UPDATE forward SET tunnel_id = ?, updated_time = ? WHERE id = ?`, req.TargetTunnelID, time.Now().UnixMilli(), id)
+ if err != nil {
+ fail++
+ continue
+ }
+ p := 0
+ if port.Valid {
+ p = int(port.Int64)
+ }
+ if p <= 0 {
+ p = h.pickTunnelPort(req.TargetTunnelID)
+ }
+ _ = h.replaceForwardPorts(id, req.TargetTunnelID, p)
+ updatedForward, fetchErr := h.getForwardRecord(id)
+ if fetchErr != nil {
+ fail++
+ continue
+ }
+ if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil {
+ fail++
+ continue
+ }
+ success++
+ }
+ response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
+}
+
+func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ tunnelID := asInt64(req["tunnelId"], 0)
+ if tunnelID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
+ return
+ }
+ name := asString(req["name"])
+ if name == "" {
+ response.WriteJSON(w, response.ErrDefault("名称不能为空"))
+ return
+ }
+ var tunnelName string
+ _ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName)
+ if tunnelName == "" {
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+ return
+ }
+ now := time.Now().UnixMilli()
+ _, err := h.repo.DB().Exec(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`,
+ name, asInt(req["speed"], 100), tunnelID, tunnelName, now, now, asInt(req["status"], 1))
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ tunnelID := asInt64(req["tunnelId"], 0)
+ if id <= 0 || tunnelID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ var tunnelName string
+ _ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName)
+ if tunnelName == "" {
+ response.WriteJSON(w, response.ErrDefault("隧道不存在"))
+ return
+ }
+ _, err := h.repo.DB().Exec(`UPDATE speed_limit SET name=?, speed=?, tunnel_id=?, tunnel_name=?, status=?, updated_time=? WHERE id=?`,
+ asString(req["name"]), asInt(req["speed"], 100), tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli(), id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ _, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) groupTunnelCreate(w http.ResponseWriter, r *http.Request) {
+ h.groupCreate(w, r, "tunnel_group")
+}
+
+func (h *Handler) groupTunnelUpdate(w http.ResponseWriter, r *http.Request) {
+ h.groupUpdate(w, r, "tunnel_group")
+}
+
+func (h *Handler) groupTunnelDelete(w http.ResponseWriter, r *http.Request) {
+ h.groupDelete(w, r, "tunnel_group")
+}
+
+func (h *Handler) groupUserCreate(w http.ResponseWriter, r *http.Request) {
+ h.groupCreate(w, r, "user_group")
+}
+
+func (h *Handler) groupUserUpdate(w http.ResponseWriter, r *http.Request) {
+ h.groupUpdate(w, r, "user_group")
+}
+
+func (h *Handler) groupUserDelete(w http.ResponseWriter, r *http.Request) {
+ h.groupDelete(w, r, "user_group")
+}
+
+func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) {
+ var req struct {
+ GroupID int64 `json:"groupId"`
+ TunnelIDs []int64 `json:"tunnelIds"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil || req.GroupID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ defer func() { _ = tx.Rollback() }()
+ _, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID)
+ for _, tid := range req.TunnelIDs {
+ _, _ = tx.Exec(`INSERT OR IGNORE INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, req.GroupID, tid, time.Now().UnixMilli())
+ }
+ if err := tx.Commit(); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ _ = h.syncPermissionsByTunnelGroup(req.GroupID)
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) {
+ var req struct {
+ GroupID int64 `json:"groupId"`
+ UserIDs []int64 `json:"userIds"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil || req.GroupID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ defer func() { _ = tx.Rollback() }()
+ _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID)
+ for _, uid := range req.UserIDs {
+ _, _ = tx.Exec(`INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)`, req.GroupID, uid, time.Now().UnixMilli())
+ }
+ if err := tx.Commit(); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ _ = h.syncPermissionsByUserGroup(req.GroupID)
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request) {
+ var req struct {
+ UserGroupID int64 `json:"userGroupId"`
+ TunnelGroupID int64 `json:"tunnelGroupId"`
+ }
+ if err := decodeJSON(r.Body, &req); err != nil || req.UserGroupID <= 0 || req.TunnelGroupID <= 0 {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ _, err := h.repo.DB().Exec(`INSERT OR IGNORE INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ _ = h.applyGroupPermission(req.UserGroupID, req.TunnelGroupID)
+ response.WriteJSON(w, response.OK("权限分配成功"))
+}
+
+func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) {
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ var ug, tg int64
+ _ = h.repo.DB().QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg)
+ _, _ = h.repo.DB().Exec(`DELETE FROM group_permission WHERE id = ?`, id)
+ _, _ = h.repo.DB().Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, ug, tg)
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) groupCreate(w http.ResponseWriter, r *http.Request, table string) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ name := asString(req["name"])
+ if name == "" {
+ response.WriteJSON(w, response.ErrDefault("分组名称不能为空"))
+ return
+ }
+ now := time.Now().UnixMilli()
+ _, err := h.repo.DB().Exec(`INSERT INTO `+table+`(name, created_time, updated_time, status) VALUES(?, ?, ?, ?)`, name, now, now, asInt(req["status"], 1))
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) groupUpdate(w http.ResponseWriter, r *http.Request, table string) {
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return
+ }
+ id := asInt64(req["id"], 0)
+ if id <= 0 {
+ response.WriteJSON(w, response.ErrDefault("分组ID不能为空"))
+ return
+ }
+ _, err := h.repo.DB().Exec(`UPDATE `+table+` SET name = ?, status = ?, updated_time = ? WHERE id = ?`, asString(req["name"]), asInt(req["status"], 1), time.Now().UnixMilli(), id)
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) groupDelete(w http.ResponseWriter, r *http.Request, table string) {
+ id := idFromBody(r, w)
+ if id <= 0 {
+ return
+ }
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ defer func() { _ = tx.Rollback() }()
+ if table == "tunnel_group" {
+ _, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM group_permission WHERE tunnel_group_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE tunnel_group_id = ?`, id)
+ } else {
+ _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM group_permission WHERE user_group_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ?`, id)
+ }
+ _, _ = tx.Exec(`DELETE FROM `+table+` WHERE id = ?`, id)
+ if err := tx.Commit(); err != nil {
+ response.WriteJSON(w, response.Err(-2, err.Error()))
+ return
+ }
+ response.WriteJSON(w, response.OKEmpty())
+}
+
+func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error {
+ db := h.repo.DB()
+ userIDs, _ := queryInt64List(db, `SELECT user_id FROM user_group_user WHERE user_group_id = ?`, userGroupID)
+ tunnelIDs, _ := queryInt64List(db, `SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tunnelGroupID)
+ for _, uid := range userIDs {
+ for _, tid := range tunnelIDs {
+ utID, created, err := ensureUserTunnelGrant(db, uid, tid)
+ if err != nil {
+ continue
+ }
+ createdByGroup := 0
+ if created {
+ createdByGroup = 1
+ }
+ _, _ = db.Exec(`INSERT OR IGNORE INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?)`,
+ userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli())
+ }
+ }
+ return nil
+}
+
+func (h *Handler) syncPermissionsByUserGroup(userGroupID int64) error {
+ db := h.repo.DB()
+ pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE user_group_id = ?`, userGroupID)
+ if err != nil {
+ return err
+ }
+ for _, p := range pairs {
+ _ = h.applyGroupPermission(p[0], p[1])
+ }
+ return nil
+}
+
+func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error {
+ db := h.repo.DB()
+ pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE tunnel_group_id = ?`, tunnelGroupID)
+ if err != nil {
+ return err
+ }
+ for _, p := range pairs {
+ _ = h.applyGroupPermission(p[0], p[1])
+ }
+ return nil
+}
+
+func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, error) {
+ var id int64
+ err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&id)
+ if err == nil {
+ return id, false, nil
+ }
+ if err != sql.ErrNoRows {
+ return 0, false, err
+ }
+ var flow int64
+ var num int
+ var expTime int64
+ var flowReset int64
+ if err := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset); err != nil {
+ return 0, false, err
+ }
+ res, err := db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1)`,
+ userID, tunnelID, num, flow, flowReset, expTime)
+ if err != nil {
+ return 0, false, err
+ }
+ id, _ = res.LastInsertId()
+ return id, true, nil
+}
+
+func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error) {
+ rows, err := db.Query(q, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ out := make([]int64, 0)
+ for rows.Next() {
+ var v int64
+ if err := rows.Scan(&v); err != nil {
+ return nil, err
+ }
+ out = append(out, v)
+ }
+ return out, rows.Err()
+}
+
+func queryPairs(db *sql.DB, q string, args ...interface{}) ([][2]int64, error) {
+ rows, err := db.Query(q, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ out := make([][2]int64, 0)
+ for rows.Next() {
+ var a, b int64
+ if err := rows.Scan(&a, &b); err != nil {
+ return nil, err
+ }
+ out = append(out, [2]int64{a, b})
+ }
+ return out, rows.Err()
+}
+
+type tunnelRuntimeNode struct {
+ NodeID int64
+ Protocol string
+ Strategy string
+ Inx int
+ ChainType int
+ Port int
+}
+
+type tunnelCreateState struct {
+ TunnelID int64
+ Type int
+ InNodes []tunnelRuntimeNode
+ ChainHops [][]tunnelRuntimeNode
+ OutNodes []tunnelRuntimeNode
+ Nodes map[int64]*nodeRecord
+ NodeIDList []int64
+}
+
+func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{}, tunnelType int) (*tunnelCreateState, error) {
+ state := &tunnelCreateState{
+ Type: tunnelType,
+ InNodes: make([]tunnelRuntimeNode, 0),
+ ChainHops: make([][]tunnelRuntimeNode, 0),
+ OutNodes: make([]tunnelRuntimeNode, 0),
+ Nodes: make(map[int64]*nodeRecord),
+ }
+ nodeIDs := make([]int64, 0)
+
+ for _, item := range asMapSlice(req["inNodeId"]) {
+ nodeID := asInt64(item["nodeId"], 0)
+ if nodeID <= 0 {
+ continue
+ }
+ nodeIDs = append(nodeIDs, nodeID)
+ state.InNodes = append(state.InNodes, tunnelRuntimeNode{
+ NodeID: nodeID,
+ Protocol: defaultString(asString(item["protocol"]), "tls"),
+ Strategy: defaultString(asString(item["strategy"]), "round"),
+ ChainType: 1,
+ })
+ }
+ if len(state.InNodes) == 0 {
+ return nil, errors.New("入口不能为空")
+ }
+
+ if tunnelType == 2 {
+ outNodesRaw := asMapSlice(req["outNodeId"])
+ if len(outNodesRaw) == 0 {
+ return nil, errors.New("出口不能为空")
+ }
+
+ allocated := map[int64]int{}
+ for _, item := range outNodesRaw {
+ nodeID := asInt64(item["nodeId"], 0)
+ if nodeID <= 0 {
+ continue
+ }
+ nodeIDs = append(nodeIDs, nodeID)
+ port := asInt(item["port"], 0)
+ if port <= 0 {
+ var err error
+ port, err = pickNodePortTx(tx, nodeID, allocated)
+ if err != nil {
+ return nil, err
+ }
+ }
+ state.OutNodes = append(state.OutNodes, tunnelRuntimeNode{
+ NodeID: nodeID,
+ Protocol: defaultString(asString(item["protocol"]), "tls"),
+ Strategy: defaultString(asString(item["strategy"]), "round"),
+ ChainType: 3,
+ Port: port,
+ })
+ }
+ if len(state.OutNodes) == 0 {
+ return nil, errors.New("出口不能为空")
+ }
+
+ for hopIdx, hopRaw := range asAnySlice(req["chainNodes"]) {
+ hop := make([]tunnelRuntimeNode, 0)
+ for _, item := range asMapSlice(hopRaw) {
+ nodeID := asInt64(item["nodeId"], 0)
+ if nodeID <= 0 {
+ continue
+ }
+ nodeIDs = append(nodeIDs, nodeID)
+ port := asInt(item["port"], 0)
+ if port <= 0 {
+ var err error
+ port, err = pickNodePortTx(tx, nodeID, allocated)
+ if err != nil {
+ return nil, err
+ }
+ }
+ hop = append(hop, tunnelRuntimeNode{
+ NodeID: nodeID,
+ Protocol: defaultString(asString(item["protocol"]), "tls"),
+ Strategy: defaultString(asString(item["strategy"]), "round"),
+ Inx: hopIdx + 1,
+ ChainType: 2,
+ Port: port,
+ })
+ }
+ if len(hop) > 0 {
+ state.ChainHops = append(state.ChainHops, hop)
+ }
+ }
+ }
+
+ seen := make(map[int64]struct{}, len(nodeIDs))
+ for _, nodeID := range nodeIDs {
+ if _, ok := seen[nodeID]; ok {
+ return nil, errors.New("节点重复")
+ }
+ seen[nodeID] = struct{}{}
+ state.NodeIDList = append(state.NodeIDList, nodeID)
+ node, err := h.getNodeRecord(nodeID)
+ if err != nil {
+ if strings.Contains(err.Error(), "不存在") {
+ return nil, errors.New("节点不存在")
+ }
+ return nil, err
+ }
+ if node.Status != 1 {
+ return nil, errors.New("部分节点不在线")
+ }
+ state.Nodes[nodeID] = node
+ }
+
+ return state, nil
+}
+
+func buildTunnelInIP(inNodes []tunnelRuntimeNode, nodes map[int64]*nodeRecord) string {
+ set := make(map[string]struct{})
+ ordered := make([]string, 0)
+ for _, inNode := range inNodes {
+ node := nodes[inNode.NodeID]
+ if node == nil {
+ continue
+ }
+ if v := strings.TrimSpace(node.ServerIPv4); v != "" {
+ if _, ok := set[v]; !ok {
+ set[v] = struct{}{}
+ ordered = append(ordered, v)
+ }
+ }
+ if v := strings.TrimSpace(node.ServerIPv6); v != "" {
+ if _, ok := set[v]; !ok {
+ set[v] = struct{}{}
+ ordered = append(ordered, v)
+ }
+ }
+ if strings.TrimSpace(node.ServerIPv4) == "" && strings.TrimSpace(node.ServerIPv6) == "" {
+ if v := strings.TrimSpace(node.ServerIP); v != "" {
+ if _, ok := set[v]; !ok {
+ set[v] = struct{}{}
+ ordered = append(ordered, v)
+ }
+ }
+ }
+ }
+ return strings.Join(ordered, ",")
+}
+
+func applyTunnelPortsToRequest(req map[string]interface{}, state *tunnelCreateState) {
+ if req == nil || state == nil {
+ return
+ }
+ outPorts := make(map[int64]int)
+ for _, n := range state.OutNodes {
+ outPorts[n.NodeID] = n.Port
+ }
+ for _, item := range asMapSlice(req["outNodeId"]) {
+ nodeID := asInt64(item["nodeId"], 0)
+ if port, ok := outPorts[nodeID]; ok && port > 0 {
+ item["port"] = port
+ }
+ }
+
+ chainPorts := make(map[int64]int)
+ for _, hop := range state.ChainHops {
+ for _, n := range hop {
+ chainPorts[n.NodeID] = n.Port
+ }
+ }
+ for _, hopRaw := range asAnySlice(req["chainNodes"]) {
+ for _, item := range asMapSlice(hopRaw) {
+ nodeID := asInt64(item["nodeId"], 0)
+ if port, ok := chainPorts[nodeID]; ok && port > 0 {
+ item["port"] = port
+ }
+ }
+ }
+}
+
+func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) {
+ if h == nil || state == nil {
+ return nil, nil, errors.New("invalid tunnel runtime state")
+ }
+ createdChains := make([]int64, 0)
+ createdServices := make([]int64, 0)
+ if state.Type != 2 {
+ return createdChains, createdServices, nil
+ }
+
+ for _, inNode := range state.InNodes {
+ targets := state.OutNodes
+ if len(state.ChainHops) > 0 {
+ targets = state.ChainHops[0]
+ }
+ chainData, err := buildTunnelChainConfig(state.TunnelID, inNode.NodeID, targets, state.Nodes)
+ if err != nil {
+ return createdChains, createdServices, err
+ }
+ if _, err := h.sendNodeCommand(inNode.NodeID, "AddChains", chainData, true, false); err != nil {
+ return createdChains, createdServices, fmt.Errorf("入口节点 %s 下发转发链失败: %w", nodeDisplayName(state.Nodes[inNode.NodeID]), err)
+ }
+ createdChains = append(createdChains, inNode.NodeID)
+ }
+
+ for i, hop := range state.ChainHops {
+ nextTargets := state.OutNodes
+ if i+1 < len(state.ChainHops) {
+ nextTargets = state.ChainHops[i+1]
+ }
+ for _, chainNode := range hop {
+ chainData, err := buildTunnelChainConfig(state.TunnelID, chainNode.NodeID, nextTargets, state.Nodes)
+ if err != nil {
+ return createdChains, createdServices, err
+ }
+ if _, err := h.sendNodeCommand(chainNode.NodeID, "AddChains", chainData, true, false); err != nil {
+ return createdChains, createdServices, fmt.Errorf("转发链节点 %s 下发转发链失败: %w", nodeDisplayName(state.Nodes[chainNode.NodeID]), err)
+ }
+ createdChains = append(createdChains, chainNode.NodeID)
+
+ serviceData := buildTunnelChainServiceConfig(state.TunnelID, chainNode, state.Nodes[chainNode.NodeID])
+ if _, err := h.sendNodeCommand(chainNode.NodeID, "AddService", serviceData, true, false); err != nil {
+ return createdChains, createdServices, fmt.Errorf("转发链节点 %s 下发服务失败: %w", nodeDisplayName(state.Nodes[chainNode.NodeID]), err)
+ }
+ createdServices = append(createdServices, chainNode.NodeID)
+ }
+ }
+
+ for _, outNode := range state.OutNodes {
+ serviceData := buildTunnelChainServiceConfig(state.TunnelID, outNode, state.Nodes[outNode.NodeID])
+ if _, err := h.sendNodeCommand(outNode.NodeID, "AddService", serviceData, true, false); err != nil {
+ return createdChains, createdServices, fmt.Errorf("出口节点 %s 下发服务失败: %w", nodeDisplayName(state.Nodes[outNode.NodeID]), err)
+ }
+ createdServices = append(createdServices, outNode.NodeID)
+ }
+
+ return createdChains, createdServices, nil
+}
+
+func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tunnelID int64) {
+ if h == nil || tunnelID <= 0 {
+ return
+ }
+ seenServices := make(map[int64]struct{})
+ serviceName := fmt.Sprintf("%d_tls", tunnelID)
+ for i := len(serviceNodeIDs) - 1; i >= 0; i-- {
+ nodeID := serviceNodeIDs[i]
+ if _, ok := seenServices[nodeID]; ok {
+ continue
+ }
+ seenServices[nodeID] = struct{}{}
+ _, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{serviceName}}, false, true)
+ }
+
+ seenChains := make(map[int64]struct{})
+ chainName := fmt.Sprintf("chains_%d", tunnelID)
+ for i := len(chainNodeIDs) - 1; i >= 0; i-- {
+ nodeID := chainNodeIDs[i]
+ if _, ok := seenChains[nodeID]; ok {
+ continue
+ }
+ seenChains[nodeID] = struct{}{}
+ _, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
+ }
+}
+
+func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRuntimeNode, nodes map[int64]*nodeRecord) (map[string]interface{}, error) {
+ fromNode := nodes[fromNodeID]
+ if fromNode == nil {
+ return nil, errors.New("节点不存在")
+ }
+ if len(targets) == 0 {
+ return nil, errors.New("转发链目标不能为空")
+ }
+ nodeItems := make([]map[string]interface{}, 0, len(targets))
+ for idx, target := range targets {
+ targetNode := nodes[target.NodeID]
+ if targetNode == nil {
+ return nil, errors.New("节点不存在")
+ }
+ host, err := selectTunnelDialHost(fromNode, targetNode)
+ if err != nil {
+ return nil, err
+ }
+ port := target.Port
+ if port <= 0 {
+ return nil, errors.New("节点端口不能为空")
+ }
+ nodeItems = append(nodeItems, map[string]interface{}{
+ "name": fmt.Sprintf("node_%d", idx+1),
+ "addr": processServerAddress(fmt.Sprintf("%s:%d", host, port)),
+ "connector": map[string]interface{}{
+ "type": "relay",
+ },
+ "dialer": map[string]interface{}{
+ "type": defaultString(target.Protocol, "tls"),
+ },
+ })
+ }
+
+ strategy := defaultString(strings.TrimSpace(targets[0].Strategy), "round")
+ hop := map[string]interface{}{
+ "name": fmt.Sprintf("hop_%d", tunnelID),
+ "selector": map[string]interface{}{
+ "strategy": strategy,
+ "maxFails": 1,
+ "failTimeout": int64(600000000000),
+ },
+ "nodes": nodeItems,
+ }
+ if strings.TrimSpace(fromNode.InterfaceName) != "" {
+ hop["interface"] = fromNode.InterfaceName
+ }
+
+ return map[string]interface{}{
+ "name": fmt.Sprintf("chains_%d", tunnelID),
+ "hops": []map[string]interface{}{hop},
+ }, nil
+}
+
+func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode, node *nodeRecord) []map[string]interface{} {
+ if node == nil {
+ return nil
+ }
+ service := map[string]interface{}{
+ "name": fmt.Sprintf("%d_tls", tunnelID),
+ "addr": fmt.Sprintf("%s:%d", node.TCPListenAddr, chainNode.Port),
+ "handler": map[string]interface{}{
+ "type": "relay",
+ },
+ "listener": map[string]interface{}{
+ "type": defaultString(chainNode.Protocol, "tls"),
+ },
+ }
+ if chainNode.ChainType == 2 {
+ service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", tunnelID)
+ }
+ if chainNode.ChainType == 3 && strings.TrimSpace(node.InterfaceName) != "" {
+ service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
+ }
+ return []map[string]interface{}{service}
+}
+
+func selectTunnelDialHost(fromNode, toNode *nodeRecord) (string, error) {
+ if fromNode == nil || toNode == nil {
+ return "", errors.New("节点不存在")
+ }
+ fromV4 := nodeSupportsV4(fromNode)
+ fromV6 := nodeSupportsV6(fromNode)
+ toV4 := nodeSupportsV4(toNode)
+ toV6 := nodeSupportsV6(toNode)
+
+ if fromV4 && toV4 {
+ host := pickNodeAddressV4(toNode)
+ if host != "" {
+ return host, nil
+ }
+ }
+ if fromV6 && toV6 {
+ host := pickNodeAddressV6(toNode)
+ if host != "" {
+ return host, nil
+ }
+ }
+ return "", fmt.Errorf("节点链路不兼容:%s(v4=%t,v6=%t) -> %s(v4=%t,v6=%t)", nodeDisplayName(fromNode), fromV4, fromV6, nodeDisplayName(toNode), toV4, toV6)
+}
+
+func nodeDisplayName(node *nodeRecord) string {
+ if node == nil {
+ return "node"
+ }
+ if strings.TrimSpace(node.Name) != "" {
+ return strings.TrimSpace(node.Name)
+ }
+ return fmt.Sprintf("node_%d", node.ID)
+}
+
+func nodeSupportsV4(node *nodeRecord) bool {
+ if node == nil {
+ return false
+ }
+ if strings.TrimSpace(node.ServerIPv4) != "" {
+ return true
+ }
+ if strings.TrimSpace(node.ServerIPv6) != "" {
+ return false
+ }
+ legacy := strings.Trim(strings.TrimSpace(node.ServerIP), "[]")
+ if legacy == "" {
+ return false
+ }
+ if ip := net.ParseIP(legacy); ip != nil {
+ return ip.To4() != nil
+ }
+ return true
+}
+
+func nodeSupportsV6(node *nodeRecord) bool {
+ if node == nil {
+ return false
+ }
+ if strings.TrimSpace(node.ServerIPv6) != "" {
+ return true
+ }
+ if strings.TrimSpace(node.ServerIPv4) != "" {
+ return false
+ }
+ legacy := strings.Trim(strings.TrimSpace(node.ServerIP), "[]")
+ if legacy == "" {
+ return false
+ }
+ if ip := net.ParseIP(legacy); ip != nil {
+ return ip.To4() == nil
+ }
+ return true
+}
+
+func pickNodeAddressV4(node *nodeRecord) string {
+ if node == nil {
+ return ""
+ }
+ if v := strings.TrimSpace(node.ServerIPv4); v != "" {
+ return v
+ }
+ return strings.TrimSpace(node.ServerIP)
+}
+
+func pickNodeAddressV6(node *nodeRecord) string {
+ if node == nil {
+ return ""
+ }
+ if v := strings.TrimSpace(node.ServerIPv6); v != "" {
+ return v
+ }
+ return strings.TrimSpace(node.ServerIP)
+}
+
+func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int) (int, error) {
+ if tx == nil {
+ return 0, errors.New("database unavailable")
+ }
+ if nodeID <= 0 {
+ return 0, errors.New("节点不存在")
+ }
+ if port, ok := allocated[nodeID]; ok && port > 0 {
+ return port, nil
+ }
+
+ var portRange string
+ if err := tx.QueryRow(`SELECT port FROM node WHERE id = ? LIMIT 1`, nodeID).Scan(&portRange); err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return 0, errors.New("节点不存在")
+ }
+ return 0, err
+ }
+ candidates := parsePortRangeSpec(portRange)
+ if len(candidates) == 0 {
+ return 0, errors.New("节点端口已满,无可用端口")
+ }
+
+ used := map[int]struct{}{}
+ chainRows, err := tx.Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL`, nodeID)
+ if err != nil {
+ return 0, err
+ }
+ for chainRows.Next() {
+ var p sql.NullInt64
+ if scanErr := chainRows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 {
+ used[int(p.Int64)] = struct{}{}
+ }
+ }
+ _ = chainRows.Close()
+
+ forwardRows, err := tx.Query(`SELECT port FROM forward_port WHERE node_id = ?`, nodeID)
+ if err != nil {
+ return 0, err
+ }
+ for forwardRows.Next() {
+ var p sql.NullInt64
+ if scanErr := forwardRows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 {
+ used[int(p.Int64)] = struct{}{}
+ }
+ }
+ _ = forwardRows.Close()
+
+ for _, candidate := range candidates {
+ if candidate <= 0 {
+ continue
+ }
+ if _, ok := used[candidate]; ok {
+ continue
+ }
+ allocated[nodeID] = candidate
+ return candidate, nil
+ }
+ return 0, errors.New("节点端口已满,无可用端口")
+}
+
+func parsePortRangeSpec(input string) []int {
+ input = strings.TrimSpace(input)
+ if input == "" {
+ return nil
+ }
+ set := make(map[int]struct{})
+ parts := strings.Split(input, ",")
+ for _, part := range parts {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ continue
+ }
+ if strings.Contains(part, "-") {
+ r := strings.SplitN(part, "-", 2)
+ if len(r) != 2 {
+ continue
+ }
+ start, err1 := strconv.Atoi(strings.TrimSpace(r[0]))
+ end, err2 := strconv.Atoi(strings.TrimSpace(r[1]))
+ if err1 != nil || err2 != nil || start <= 0 || end <= 0 {
+ continue
+ }
+ if end < start {
+ start, end = end, start
+ }
+ for p := start; p <= end; p++ {
+ set[p] = struct{}{}
+ }
+ continue
+ }
+ p, err := strconv.Atoi(part)
+ if err != nil || p <= 0 {
+ continue
+ }
+ set[p] = struct{}{}
+ }
+ out := make([]int, 0, len(set))
+ for p := range set {
+ out = append(out, p)
+ }
+ sort.Ints(out)
+ return out
+}
+
+func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error {
+ allocated := map[int64]int{}
+ inNodes := asMapSlice(req["inNodeId"])
+ for _, n := range inNodes {
+ nodeID := asInt64(n["nodeId"], 0)
+ if nodeID <= 0 {
+ continue
+ }
+ _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, NULL, NULL, 0, ?)`,
+ tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls"))
+ if err != nil {
+ return err
+ }
+ }
+ for _, n := range asMapSlice(req["outNodeId"]) {
+ nodeID := asInt64(n["nodeId"], 0)
+ if nodeID <= 0 {
+ continue
+ }
+ port := asInt(n["port"], 0)
+ if port <= 0 {
+ var pickErr error
+ port, pickErr = pickNodePortTx(tx, nodeID, allocated)
+ if pickErr != nil {
+ return pickErr
+ }
+ }
+ _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, ?, NULL, 0, ?)`,
+ tunnelID, nodeID, port, defaultString(asString(n["protocol"]), "tls"))
+ if err != nil {
+ return err
+ }
+ }
+ chainNodes := asAnySlice(req["chainNodes"])
+ for i, grp := range chainNodes {
+ for _, n := range asMapSlice(grp) {
+ nodeID := asInt64(n["nodeId"], 0)
+ if nodeID <= 0 {
+ continue
+ }
+ port := asInt(n["port"], 0)
+ if port <= 0 {
+ var pickErr error
+ port, pickErr = pickNodePortTx(tx, nodeID, allocated)
+ if pickErr != nil {
+ return pickErr
+ }
+ }
+ _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, ?, ?, ?, ?)`,
+ tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls"))
+ if err != nil {
+ return err
+ }
+ }
+ }
+ return nil
+}
+
+func (h *Handler) deleteNodeByID(id int64) error {
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ return err
+ }
+ defer func() { _ = tx.Rollback() }()
+ _, _ = tx.Exec(`DELETE FROM forward_port WHERE node_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE node_id = ?`, id)
+ _, err = tx.Exec(`DELETE FROM node WHERE id = ?`, id)
+ if err != nil {
+ return err
+ }
+ return tx.Commit()
+}
+
+func (h *Handler) deleteTunnelByID(id int64) error {
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ return err
+ }
+ defer func() { _ = tx.Rollback() }()
+ _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE tunnel_id = ?)`, id)
+ _, _ = tx.Exec(`DELETE FROM forward WHERE tunnel_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM user_tunnel WHERE tunnel_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM speed_limit WHERE tunnel_id = ?`, id)
+ _, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id)
+ _, err = tx.Exec(`DELETE FROM tunnel WHERE id = ?`, id)
+ if err != nil {
+ return err
+ }
+ return tx.Commit()
+}
+
+func (h *Handler) deleteForwardByID(id int64) error {
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ return err
+ }
+ defer func() { _ = tx.Rollback() }()
+ _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, id)
+ _, err = tx.Exec(`DELETE FROM forward WHERE id = ?`, id)
+ if err != nil {
+ return err
+ }
+ return tx.Commit()
+}
+
+func (h *Handler) batchForwardDelete(ids []int64) (int, int) {
+ s := 0
+ f := 0
+ for _, id := range ids {
+ if err := h.deleteForwardByID(id); err != nil {
+ f++
+ } else {
+ s++
+ }
+ }
+ return s, f
+}
+
+func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) {
+ s := 0
+ f := 0
+ for _, id := range ids {
+ if _, err := h.repo.DB().Exec(`UPDATE forward SET status = ?, updated_time = ? WHERE id = ?`, status, time.Now().UnixMilli(), id); err != nil {
+ f++
+ } else {
+ s++
+ }
+ }
+ return s, f
+}
+
+func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
+ rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 ORDER BY inx ASC, id ASC`, tunnelID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ out := make([]int64, 0)
+ for rows.Next() {
+ var id int64
+ if err := rows.Scan(&id); err == nil {
+ out = append(out, id)
+ }
+ }
+ return out, rows.Err()
+}
+
+func (h *Handler) pickTunnelPort(tunnelID int64) int {
+ entry, _ := h.tunnelEntryNodeIDs(tunnelID)
+ if len(entry) == 0 {
+ return 10000
+ }
+ var portRange string
+ _ = h.repo.DB().QueryRow(`SELECT port FROM node WHERE id = ?`, entry[0]).Scan(&portRange)
+ if portRange == "" {
+ return 10000
+ }
+ first := strings.Split(portRange, ",")[0]
+ first = strings.TrimSpace(first)
+ if strings.Contains(first, "-") {
+ parts := strings.SplitN(first, "-", 2)
+ p, _ := strconv.Atoi(strings.TrimSpace(parts[0]))
+ if p > 0 {
+ return p
+ }
+ }
+ if p, err := strconv.Atoi(first); err == nil && p > 0 {
+ return p
+ }
+ return 10000
+}
+
+func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error {
+ tx, err := h.repo.DB().Begin()
+ if err != nil {
+ return err
+ }
+ defer func() { _ = tx.Rollback() }()
+ _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID)
+ entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
+ for _, nodeID := range entryNodes {
+ _, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port)
+ }
+ return tx.Commit()
+}
+
+func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
+ userID := asInt64(req["userId"], 0)
+ tunnelID := asInt64(req["tunnelId"], 0)
+ if userID <= 0 || tunnelID <= 0 {
+ return fmt.Errorf("userId or tunnelId missing")
+ }
+ db := h.repo.DB()
+ var existingID int64
+ err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&existingID)
+ flow := asInt64(req["flow"], -1)
+ num := asInt(req["num"], -1)
+ expTime := asInt64(req["expTime"], -1)
+ flowReset := asInt64(req["flowResetTime"], -1)
+ status := asInt(req["status"], 1)
+ speedID := asAnyToInt64Ptr(req["speedId"])
+ if err == sql.ErrNoRows {
+ if flow < 0 || num < 0 || expTime < 0 || flowReset < 0 {
+ _ = db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset)
+ }
+ _, err = db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?)`,
+ userID, tunnelID, nullableInt(speedID), num, flow, flowReset, expTime, status)
+ return err
+ }
+ if err != nil {
+ return err
+ }
+ if flow < 0 {
+ flow = 0
+ }
+ if num < 0 {
+ num = 0
+ }
+ if expTime < 0 {
+ expTime = time.Now().Add(365 * 24 * time.Hour).UnixMilli()
+ }
+ if flowReset < 0 {
+ flowReset = 1
+ }
+ _, err = db.Exec(`UPDATE user_tunnel SET speed_id = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ? WHERE id = ?`,
+ nullableInt(speedID), flow, num, expTime, flowReset, status, existingID)
+ return err
+}
+
+func asAnySlice(v interface{}) []interface{} {
+ if v == nil {
+ return nil
+ }
+ if arr, ok := v.([]interface{}); ok {
+ return arr
+ }
+ return nil
+}
+
+func asMapSlice(v interface{}) []map[string]interface{} {
+ arr := asAnySlice(v)
+ if arr == nil {
+ return nil
+ }
+ out := make([]map[string]interface{}, 0, len(arr))
+ for _, it := range arr {
+ if m, ok := it.(map[string]interface{}); ok {
+ out = append(out, m)
+ }
+ }
+ return out
+}
+
+func asString(v interface{}) string {
+ switch t := v.(type) {
+ case nil:
+ return ""
+ case string:
+ return strings.TrimSpace(t)
+ case float64:
+ if t == float64(int64(t)) {
+ return strconv.FormatInt(int64(t), 10)
+ }
+ return strconv.FormatFloat(t, 'f', -1, 64)
+ case int, int32, int64:
+ return fmt.Sprintf("%v", t)
+ default:
+ b, _ := json.Marshal(t)
+ return strings.Trim(string(b), "\"")
+ }
+}
+
+func asInt(v interface{}, def int) int {
+ s := asString(v)
+ if s == "" {
+ return def
+ }
+ i, err := strconv.Atoi(s)
+ if err != nil {
+ return def
+ }
+ return i
+}
+
+func asInt64(v interface{}, def int64) int64 {
+ s := asString(v)
+ if s == "" {
+ return def
+ }
+ i, err := strconv.ParseInt(s, 10, 64)
+ if err != nil {
+ return def
+ }
+ return i
+}
+
+func asFloat(v interface{}, def float64) float64 {
+ s := asString(v)
+ if s == "" {
+ return def
+ }
+ f, err := strconv.ParseFloat(s, 64)
+ if err != nil {
+ return def
+ }
+ return f
+}
+
+func asAnyToInt64Ptr(v interface{}) *int64 {
+ s := asString(v)
+ if s == "" || strings.EqualFold(s, "null") {
+ return nil
+ }
+ i, err := strconv.ParseInt(s, 10, 64)
+ if err != nil {
+ return nil
+ }
+ return &i
+}
+
+func idFromBody(r *http.Request, w http.ResponseWriter) int64 {
+ return asInt64FromBodyKey(r, w, "id")
+}
+
+func asInt64FromBodyKey(r *http.Request, w http.ResponseWriter, key string) int64 {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return 0
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return 0
+ }
+ id := asInt64(req[key], 0)
+ if id <= 0 {
+ response.WriteJSON(w, response.ErrDefault("参数错误"))
+ return 0
+ }
+ return id
+}
+
+func idsFromBody(r *http.Request, w http.ResponseWriter) []int64 {
+ if r.Method != http.MethodPost {
+ response.WriteJSON(w, response.ErrDefault("请求失败"))
+ return nil
+ }
+ var req map[string]interface{}
+ if err := decodeJSON(r.Body, &req); err != nil {
+ response.WriteJSON(w, response.ErrDefault("请求参数错误"))
+ return nil
+ }
+ arr := asAnySlice(req["ids"])
+ if len(arr) == 0 {
+ response.WriteJSON(w, response.ErrDefault("ids不能为空"))
+ return nil
+ }
+ ids := make([]int64, 0, len(arr))
+ for _, x := range arr {
+ id := asInt64(x, 0)
+ if id > 0 {
+ ids = append(ids, id)
+ }
+ }
+ sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
+ return ids
+}
+
+func nullableText(s string) interface{} {
+ if strings.TrimSpace(s) == "" {
+ return nil
+ }
+ return s
+}
+
+func nullableInt(v *int64) interface{} {
+ if v == nil {
+ return nil
+ }
+ return *v
+}
+
+func defaultString(v, def string) string {
+ if strings.TrimSpace(v) == "" {
+ return def
+ }
+ return v
+}
+
+func randomToken(n int) string {
+ buf := make([]byte, n)
+ if _, err := rand.Read(buf); err != nil {
+ return strconv.FormatInt(time.Now().UnixNano(), 16)
+ }
+ return hex.EncodeToString(buf)
+}
+
+func nextIndex(db *sql.DB, table string) int {
+ if db == nil {
+ return 0
+ }
+ row := db.QueryRow(`SELECT COALESCE(MAX(inx), -1) + 1 FROM ` + table)
+ var n int
+ if err := row.Scan(&n); err != nil {
+ return 0
+ }
+ if n < 0 {
+ return 0
+ }
+ return n
+}
diff --git a/go-backend/internal/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go
new file mode 100644
index 0000000..6493afa
--- /dev/null
+++ b/go-backend/internal/http/middleware/auth.go
@@ -0,0 +1,117 @@
+package middleware
+
+import (
+ "context"
+ "net/http"
+ "strings"
+
+ "go-backend/internal/auth"
+ "go-backend/internal/http/response"
+)
+
+type contextKey string
+
+const ClaimsContextKey contextKey = "claims"
+
+type AuthOptions struct {
+ JWTSecret string
+}
+
+func JWT(opts AuthOptions) func(http.Handler) http.Handler {
+ return func(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if shouldSkip(r.URL.Path) {
+ next.ServeHTTP(w, r)
+ return
+ }
+
+ if !strings.HasPrefix(r.URL.Path, "/api/") {
+ next.ServeHTTP(w, r)
+ return
+ }
+
+ token := strings.TrimSpace(r.Header.Get("Authorization"))
+ if token == "" {
+ response.WriteJSON(w, response.Err(401, "未登录或token已过期"))
+ return
+ }
+
+ claims, ok := auth.ValidateToken(token, opts.JWTSecret)
+ if !ok {
+ response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
+ return
+ }
+
+ if requiresAdmin(r.URL.Path) && claims.RoleID != 0 {
+ response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作"))
+ return
+ }
+
+ ctx := context.WithValue(r.Context(), ClaimsContextKey, claims)
+ next.ServeHTTP(w, r.WithContext(ctx))
+ })
+ }
+}
+
+func RequireAdmin(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ raw := r.Context().Value(ClaimsContextKey)
+ claims, ok := raw.(auth.Claims)
+ if !ok {
+ response.WriteJSON(w, response.Err(401, "无法获取用户权限信息"))
+ return
+ }
+ if claims.RoleID != 0 {
+ response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作"))
+ return
+ }
+ next.ServeHTTP(w, r)
+ })
+}
+
+func shouldSkip(path string) bool {
+ switch {
+ case strings.HasPrefix(path, "/flow/"):
+ return true
+ case strings.HasPrefix(path, "/api/v1/open_api/"):
+ return true
+ case strings.HasPrefix(path, "/api/v1/captcha/"):
+ return true
+ case path == "/api/v1/config/get":
+ return true
+ case path == "/api/v1/user/login":
+ return true
+ default:
+ return false
+ }
+}
+
+func requiresAdmin(path string) bool {
+ if strings.HasPrefix(path, "/api/v1/group/") {
+ return true
+ }
+
+ if strings.HasPrefix(path, "/api/v1/node/") {
+ return true
+ }
+
+ if strings.HasPrefix(path, "/api/v1/speed-limit/") {
+ return true
+ }
+
+ if strings.HasPrefix(path, "/api/v1/tunnel/") {
+ if strings.HasPrefix(path, "/api/v1/tunnel/user/tunnel") {
+ return false
+ }
+ return true
+ }
+
+ switch path {
+ case "/api/v1/user/create", "/api/v1/user/list", "/api/v1/user/update", "/api/v1/user/delete", "/api/v1/user/reset":
+ return true
+ case "/api/v1/config/update", "/api/v1/config/update-single":
+ return true
+ default:
+ return false
+ }
+}
diff --git a/go-backend/internal/http/middleware/cors.go b/go-backend/internal/http/middleware/cors.go
new file mode 100644
index 0000000..b33100c
--- /dev/null
+++ b/go-backend/internal/http/middleware/cors.go
@@ -0,0 +1,17 @@
+package middleware
+
+import "net/http"
+
+func CORS(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Access-Control-Allow-Origin", "*")
+ w.Header().Set("Access-Control-Allow-Headers", "*")
+ w.Header().Set("Access-Control-Allow-Methods", "GET, POST, DELETE, PUT, OPTIONS")
+ w.Header().Set("Access-Control-Expose-Headers", "Authorization")
+ if r.Method == http.MethodOptions {
+ w.WriteHeader(http.StatusNoContent)
+ return
+ }
+ next.ServeHTTP(w, r)
+ })
+}
diff --git a/go-backend/internal/http/middleware/recover.go b/go-backend/internal/http/middleware/recover.go
new file mode 100644
index 0000000..0728ac5
--- /dev/null
+++ b/go-backend/internal/http/middleware/recover.go
@@ -0,0 +1,19 @@
+package middleware
+
+import (
+ "fmt"
+ "net/http"
+
+ "go-backend/internal/http/response"
+)
+
+func Recover(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ defer func() {
+ if rec := recover(); rec != nil {
+ response.WriteJSON(w, response.Err(-2, fmt.Sprint(rec)))
+ }
+ }()
+ next.ServeHTTP(w, r)
+ })
+}
diff --git a/go-backend/internal/http/middleware/request_log.go b/go-backend/internal/http/middleware/request_log.go
new file mode 100644
index 0000000..506f325
--- /dev/null
+++ b/go-backend/internal/http/middleware/request_log.go
@@ -0,0 +1,57 @@
+package middleware
+
+import (
+ "bufio"
+ "io"
+ "log"
+ "net"
+ "net/http"
+ "time"
+)
+
+type statusWriter struct {
+ http.ResponseWriter
+ status int
+}
+
+func (w *statusWriter) WriteHeader(code int) {
+ w.status = code
+ w.ResponseWriter.WriteHeader(code)
+}
+
+func (w *statusWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
+ hj, ok := w.ResponseWriter.(http.Hijacker)
+ if !ok {
+ return nil, nil, http.ErrNotSupported
+ }
+ return hj.Hijack()
+}
+
+func (w *statusWriter) Flush() {
+ if f, ok := w.ResponseWriter.(http.Flusher); ok {
+ f.Flush()
+ }
+}
+
+func (w *statusWriter) ReadFrom(r io.Reader) (int64, error) {
+ if rf, ok := w.ResponseWriter.(io.ReaderFrom); ok {
+ return rf.ReadFrom(r)
+ }
+ return io.Copy(w.ResponseWriter, r)
+}
+
+func (w *statusWriter) Push(target string, opts *http.PushOptions) error {
+ if p, ok := w.ResponseWriter.(http.Pusher); ok {
+ return p.Push(target, opts)
+ }
+ return http.ErrNotSupported
+}
+
+func RequestLog(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ sw := &statusWriter{ResponseWriter: w, status: http.StatusOK}
+ start := time.Now()
+ next.ServeHTTP(sw, r)
+ log.Printf("%s %s -> %d (%s)", r.Method, r.URL.Path, sw.status, time.Since(start).String())
+ })
+}
diff --git a/go-backend/internal/http/response/r.go b/go-backend/internal/http/response/r.go
new file mode 100644
index 0000000..be21806
--- /dev/null
+++ b/go-backend/internal/http/response/r.go
@@ -0,0 +1,48 @@
+package response
+
+import (
+ "encoding/json"
+ "net/http"
+ "time"
+)
+
+type R struct {
+ Code int `json:"code"`
+ Msg string `json:"msg"`
+ TS int64 `json:"ts"`
+ Data interface{} `json:"data,omitempty"`
+}
+
+func OK(data interface{}) R {
+ return R{
+ Code: 0,
+ Msg: "操作成功",
+ TS: time.Now().UnixMilli(),
+ Data: data,
+ }
+}
+
+func OKEmpty() R {
+ return R{
+ Code: 0,
+ Msg: "操作成功",
+ TS: time.Now().UnixMilli(),
+ }
+}
+
+func Err(code int, msg string) R {
+ return R{
+ Code: code,
+ Msg: msg,
+ TS: time.Now().UnixMilli(),
+ }
+}
+
+func ErrDefault(msg string) R {
+ return Err(-1, msg)
+}
+
+func WriteJSON(w http.ResponseWriter, payload R) {
+ w.Header().Set("Content-Type", "application/json; charset=utf-8")
+ _ = json.NewEncoder(w).Encode(payload)
+}
diff --git a/go-backend/internal/http/router.go b/go-backend/internal/http/router.go
new file mode 100644
index 0000000..a18ed47
--- /dev/null
+++ b/go-backend/internal/http/router.go
@@ -0,0 +1,20 @@
+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.RequestLog(wrapped)
+ wrapped = middleware.CORS(wrapped)
+ return wrapped
+}
diff --git a/go-backend/internal/security/aes.go b/go-backend/internal/security/aes.go
new file mode 100644
index 0000000..56f1ec4
--- /dev/null
+++ b/go-backend/internal/security/aes.go
@@ -0,0 +1,65 @@
+package security
+
+import (
+ "crypto/aes"
+ "crypto/cipher"
+ "crypto/rand"
+ "crypto/sha256"
+ "encoding/base64"
+ "fmt"
+)
+
+type AESCrypto struct {
+ key []byte
+}
+
+func NewAESCrypto(secret string) (*AESCrypto, error) {
+ if secret == "" {
+ return nil, fmt.Errorf("secret is empty")
+ }
+ hash := sha256.Sum256([]byte(secret))
+ return &AESCrypto{key: hash[:]}, nil
+}
+
+func (a *AESCrypto) Encrypt(plain []byte) (string, error) {
+ if len(plain) == 0 {
+ return "", fmt.Errorf("empty plaintext")
+ }
+ block, err := aes.NewCipher(a.key)
+ if err != nil {
+ return "", err
+ }
+ gcm, err := cipher.NewGCM(block)
+ if err != nil {
+ return "", err
+ }
+ nonce := make([]byte, gcm.NonceSize())
+ if _, err := rand.Read(nonce); err != nil {
+ return "", err
+ }
+ sealed := gcm.Seal(nil, nonce, plain, nil)
+ data := append(nonce, sealed...)
+ return base64.StdEncoding.EncodeToString(data), nil
+}
+
+func (a *AESCrypto) Decrypt(cipherText string) ([]byte, error) {
+ raw, err := base64.StdEncoding.DecodeString(cipherText)
+ if err != nil {
+ return nil, err
+ }
+ block, err := aes.NewCipher(a.key)
+ if err != nil {
+ return nil, err
+ }
+ gcm, err := cipher.NewGCM(block)
+ if err != nil {
+ return nil, err
+ }
+ nonceSize := gcm.NonceSize()
+ if len(raw) < nonceSize {
+ return nil, fmt.Errorf("ciphertext too short")
+ }
+ nonce := raw[:nonceSize]
+ data := raw[nonceSize:]
+ return gcm.Open(nil, nonce, data, nil)
+}
diff --git a/go-backend/internal/security/md5.go b/go-backend/internal/security/md5.go
new file mode 100644
index 0000000..17263cc
--- /dev/null
+++ b/go-backend/internal/security/md5.go
@@ -0,0 +1,11 @@
+package security
+
+import (
+ "crypto/md5"
+ "fmt"
+)
+
+func MD5(input string) string {
+ hash := md5.Sum([]byte(input))
+ return fmt.Sprintf("%x", hash)
+}
diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go
new file mode 100644
index 0000000..2c81562
--- /dev/null
+++ b/go-backend/internal/store/sqlite/repository.go
@@ -0,0 +1,1191 @@
+package sqlite
+
+import (
+ "database/sql"
+ _ "embed"
+ "errors"
+ "fmt"
+ "log"
+ "os"
+ "path/filepath"
+ "sort"
+ "strings"
+ "time"
+
+ _ "modernc.org/sqlite"
+)
+
+//go:embed sql/schema.sql
+var embeddedSchema string
+
+//go:embed sql/data.sql
+var embeddedSeedData string
+
+type Repository struct {
+ db *sql.DB
+}
+
+func (r *Repository) DB() *sql.DB {
+ if r == nil {
+ return nil
+ }
+ return r.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) {
+ if err := ensureParentDir(path); err != nil {
+ return nil, err
+ }
+
+ db, err := sql.Open("sqlite", path)
+ if err != nil {
+ return nil, err
+ }
+
+ if err := db.Ping(); err != nil {
+ _ = db.Close()
+ return nil, err
+ }
+
+ if err := bootstrapSchema(db); 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
+ FROM forward f
+ LEFT JOIN tunnel t ON t.id = f.tunnel_id
+ WHERE f.user_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,
+ ); err != nil {
+ return nil, err
+ }
+
+ inIP, inPort, err := resolveForwardIngress(r.db, item.ID, item.TunnelID)
+ if err != nil {
+ return nil, err
+ }
+ item.InIP = inIP
+ item.InPort = inPort
+
+ 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 (r *Repository) ListNodes() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT id, inx, name, server_ip, server_ip_v4, server_ip_v6, port, tcp_listen_addr, udp_listen_addr, version, http, tls, socks, status
+ FROM node
+ ORDER BY inx ASC, id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ items := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id, inx int64
+ var name, serverIP, port string
+ var serverIPV4, serverIPV6, tcpListen, udpListen, version sql.NullString
+ var httpVal, tlsVal, socksVal, status int
+
+ if err := rows.Scan(&id, &inx, &name, &serverIP, &serverIPV4, &serverIPV6, &port, &tcpListen, &udpListen, &version, &httpVal, &tlsVal, &socksVal, &status); err != nil {
+ return nil, err
+ }
+
+ items = append(items, map[string]interface{}{
+ "id": id,
+ "inx": inx,
+ "name": name,
+ "ip": serverIP,
+ "serverIp": serverIP,
+ "serverIpV4": nullableString(serverIPV4),
+ "serverIpV6": nullableString(serverIPV6),
+ "port": port,
+ "tcpListenAddr": nullableString(tcpListen),
+ "udpListenAddr": nullableString(udpListen),
+ "version": nullableString(version),
+ "http": httpVal,
+ "tls": tlsVal,
+ "socks": socksVal,
+ "status": status,
+ })
+ }
+
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
+func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT id, user, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status
+ FROM user
+ WHERE role_id != 0
+ ORDER BY id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ items := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id int64
+ var user string
+ var roleID int
+ var expTime, flow, inFlow, outFlow, flowResetTime, createdTime int64
+ var num, status int
+ var updatedTime sql.NullInt64
+
+ if err := rows.Scan(&id, &user, &roleID, &expTime, &flow, &inFlow, &outFlow, &flowResetTime, &num, &createdTime, &updatedTime, &status); err != nil {
+ return nil, err
+ }
+
+ items = append(items, map[string]interface{}{
+ "id": id,
+ "user": user,
+ "name": user,
+ "roleId": roleID,
+ "status": status,
+ "flow": flow,
+ "num": num,
+ "expTime": expTime,
+ "flowResetTime": flowResetTime,
+ "createdTime": createdTime,
+ "updatedTime": nullableInt64(updatedTime),
+ "inFlow": inFlow,
+ "outFlow": outFlow,
+ })
+ }
+
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
+func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT id, name, speed, tunnel_id, tunnel_name, status, created_time, updated_time
+ FROM speed_limit
+ ORDER BY id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ items := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id, tunnelID, createdTime int64
+ var name, tunnelName string
+ var speed, status int
+ var updatedTime sql.NullInt64
+ if err := rows.Scan(&id, &name, &speed, &tunnelID, &tunnelName, &status, &createdTime, &updatedTime); err != nil {
+ return nil, err
+ }
+ items = append(items, map[string]interface{}{
+ "id": id,
+ "name": name,
+ "speed": speed,
+ "tunnelId": tunnelID,
+ "tunnelName": tunnelName,
+ "status": status,
+ "createdTime": createdTime,
+ "updatedTime": nullableInt64(updatedTime),
+ })
+ }
+
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
+func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, t.name, f.remote_addr, f.strategy,
+ f.in_flow, f.out_flow, f.created_time, f.status, f.inx
+ FROM forward f
+ LEFT JOIN tunnel t ON t.id = f.tunnel_id
+ ORDER BY f.inx ASC, f.id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ items := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id, userID, tunnelID, inFlow, outFlow, createdTime, inx int64
+ var userName, name, tunnelName, remoteAddr, strategy string
+ var status int
+
+ if err := rows.Scan(&id, &userID, &userName, &name, &tunnelID, &tunnelName, &remoteAddr, &strategy, &inFlow, &outFlow, &createdTime, &status, &inx); err != nil {
+ return nil, err
+ }
+
+ inIP, inPort, err := resolveForwardIngress(r.db, id, tunnelID)
+ if err != nil {
+ return nil, err
+ }
+
+ items = append(items, map[string]interface{}{
+ "id": id,
+ "userId": userID,
+ "userName": userName,
+ "name": name,
+ "tunnelId": tunnelID,
+ "tunnelName": tunnelName,
+ "inIp": nullableForwardIngress(inIP),
+ "inPort": nullableInt64(inPort),
+ "remoteAddr": remoteAddr,
+ "strategy": strategy,
+ "inFlow": inFlow,
+ "outFlow": outFlow,
+ "createdTime": createdTime,
+ "status": status,
+ "inx": inx,
+ })
+ }
+
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
+func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT DISTINCT t.id, t.name
+ FROM user_tunnel ut
+ JOIN tunnel t ON t.id = ut.tunnel_id
+ WHERE ut.user_id = ? AND t.status = 1
+ ORDER BY t.inx ASC, t.id ASC
+ `, userID)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ items := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id int64
+ var name string
+ if err := rows.Scan(&id, &name); err != nil {
+ return nil, err
+ }
+ items = append(items, map[string]interface{}{"id": id, "name": name})
+ }
+
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
+func (r *Repository) ListEnabledTunnelSummaries() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT id, name
+ FROM tunnel
+ WHERE status = 1
+ ORDER BY inx ASC, id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ items := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id int64
+ var name string
+ if err := rows.Scan(&id, &name); err != nil {
+ return nil, err
+ }
+ items = append(items, map[string]interface{}{"id": id, "name": name})
+ }
+
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
+func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT id, inx, name, type, flow, traffic_ratio, status, created_time, in_ip
+ FROM tunnel
+ ORDER BY inx ASC, id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ tunnelMap := make(map[int64]map[string]interface{})
+ orderedIDs := make([]int64, 0)
+
+ for rows.Next() {
+ var id, inx, flow, createdTime int64
+ var name string
+ var typ, status int
+ var trafficRatio float64
+ var inIP sql.NullString
+ if err := rows.Scan(&id, &inx, &name, &typ, &flow, &trafficRatio, &status, &createdTime, &inIP); err != nil {
+ return nil, err
+ }
+
+ tunnelMap[id] = map[string]interface{}{
+ "id": id,
+ "inx": inx,
+ "name": name,
+ "type": typ,
+ "flow": flow,
+ "trafficRatio": trafficRatio,
+ "status": status,
+ "createdTime": createdTime,
+ "inIp": nullableString(inIP),
+ "inNodeId": make([]map[string]interface{}, 0),
+ "outNodeId": make([]map[string]interface{}, 0),
+ "chainNodes": make([][]map[string]interface{}, 0),
+ }
+ orderedIDs = append(orderedIDs, id)
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+
+ nodeIPMap := map[int64]string{}
+ nRows, err := r.db.Query(`SELECT id, server_ip FROM node`)
+ if err == nil {
+ for nRows.Next() {
+ var id int64
+ var ip string
+ if scanErr := nRows.Scan(&id, &ip); scanErr == nil {
+ nodeIPMap[id] = ip
+ }
+ }
+ _ = nRows.Close()
+ }
+
+ chainRows, err := r.db.Query(`
+ SELECT tunnel_id, chain_type, node_id, protocol, strategy, COALESCE(inx, 0)
+ FROM chain_tunnel
+ ORDER BY tunnel_id ASC, chain_type ASC, inx ASC, id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer chainRows.Close()
+
+ chainBucket := map[int64]map[int][]map[string]interface{}{}
+ inNodeIPs := map[int64][]string{}
+
+ for chainRows.Next() {
+ var tunnelID, nodeID, inx int64
+ var chainType int
+ var protocol, strategy sql.NullString
+ if err := chainRows.Scan(&tunnelID, &chainType, &nodeID, &protocol, &strategy, &inx); err != nil {
+ return nil, err
+ }
+
+ t, ok := tunnelMap[tunnelID]
+ if !ok {
+ continue
+ }
+
+ nodeObj := map[string]interface{}{
+ "nodeId": nodeID,
+ "chainType": chainType,
+ "inx": inx,
+ }
+ if protocol.Valid {
+ nodeObj["protocol"] = protocol.String
+ }
+ if strategy.Valid {
+ nodeObj["strategy"] = strategy.String
+ }
+
+ switch chainType {
+ case 1:
+ t["inNodeId"] = append(t["inNodeId"].([]map[string]interface{}), nodeObj)
+ if ip, ok := nodeIPMap[nodeID]; ok && ip != "" {
+ inNodeIPs[tunnelID] = append(inNodeIPs[tunnelID], ip)
+ }
+ case 2:
+ if _, ok := chainBucket[tunnelID]; !ok {
+ chainBucket[tunnelID] = map[int][]map[string]interface{}{}
+ }
+ chainBucket[tunnelID][int(inx)] = append(chainBucket[tunnelID][int(inx)], nodeObj)
+ case 3:
+ t["outNodeId"] = append(t["outNodeId"].([]map[string]interface{}), nodeObj)
+ }
+ }
+ if err := chainRows.Err(); err != nil {
+ return nil, err
+ }
+
+ for tunnelID, groups := range chainBucket {
+ t := tunnelMap[tunnelID]
+ if t == nil {
+ continue
+ }
+ keys := make([]int, 0, len(groups))
+ for k := range groups {
+ keys = append(keys, k)
+ }
+ sort.Ints(keys)
+ ordered := make([][]map[string]interface{}, 0, len(keys))
+ for _, k := range keys {
+ ordered = append(ordered, groups[k])
+ }
+ t["chainNodes"] = ordered
+
+ if s, ok := t["inIp"].(string); !ok || strings.TrimSpace(s) == "" {
+ if ips := inNodeIPs[tunnelID]; len(ips) > 0 {
+ t["inIp"] = strings.Join(ips, ",")
+ }
+ }
+ }
+
+ result := make([]map[string]interface{}, 0, len(orderedIDs))
+ for _, id := range orderedIDs {
+ if t, ok := tunnelMap[id]; ok {
+ result = append(result, t)
+ }
+ }
+ return result, nil
+}
+
+func (r *Repository) ListTunnelGroups() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`SELECT id, name, status, created_time FROM tunnel_group ORDER BY id ASC`)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ result := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id, createdTime int64
+ var name string
+ var status int
+ if err := rows.Scan(&id, &name, &status, &createdTime); err != nil {
+ return nil, err
+ }
+
+ ids, names, err := r.listTunnelGroupMembers(id)
+ if err != nil {
+ return nil, err
+ }
+
+ result = append(result, map[string]interface{}{
+ "id": id,
+ "name": name,
+ "status": status,
+ "tunnelIds": ids,
+ "tunnelNames": names,
+ "createdTime": createdTime,
+ })
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return result, nil
+}
+
+func (r *Repository) ListUserGroups() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`SELECT id, name, status, created_time FROM user_group ORDER BY id ASC`)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ result := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id, createdTime int64
+ var name string
+ var status int
+ if err := rows.Scan(&id, &name, &status, &createdTime); err != nil {
+ return nil, err
+ }
+
+ ids, names, err := r.listUserGroupMembers(id)
+ if err != nil {
+ return nil, err
+ }
+
+ result = append(result, map[string]interface{}{
+ "id": id,
+ "name": name,
+ "status": status,
+ "userIds": ids,
+ "userNames": names,
+ "createdTime": createdTime,
+ })
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return result, nil
+}
+
+func (r *Repository) ListGroupPermissions() ([]map[string]interface{}, error) {
+ if r == nil || r.db == nil {
+ return nil, errors.New("repository not initialized")
+ }
+
+ rows, err := r.db.Query(`
+ SELECT gp.id, gp.user_group_id, ug.name, gp.tunnel_group_id, tg.name, gp.created_time
+ FROM group_permission gp
+ LEFT JOIN user_group ug ON ug.id = gp.user_group_id
+ LEFT JOIN tunnel_group tg ON tg.id = gp.tunnel_group_id
+ ORDER BY gp.id ASC
+ `)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ result := make([]map[string]interface{}, 0)
+ for rows.Next() {
+ var id, userGroupID, tunnelGroupID, createdTime int64
+ var userGroupName, tunnelGroupName sql.NullString
+ if err := rows.Scan(&id, &userGroupID, &userGroupName, &tunnelGroupID, &tunnelGroupName, &createdTime); err != nil {
+ return nil, err
+ }
+
+ result = append(result, map[string]interface{}{
+ "id": id,
+ "userGroupId": userGroupID,
+ "userGroupName": nullableString(userGroupName),
+ "tunnelGroupId": tunnelGroupID,
+ "tunnelGroupName": nullableString(tunnelGroupName),
+ "createdTime": createdTime,
+ })
+ }
+ if err := rows.Err(); err != nil {
+ return nil, err
+ }
+ return result, nil
+}
+
+func (r *Repository) listTunnelGroupMembers(groupID int64) ([]int64, []string, error) {
+ rows, err := r.db.Query(`
+ SELECT t.id, t.name
+ FROM tunnel_group_tunnel tgt
+ JOIN tunnel t ON t.id = tgt.tunnel_id
+ WHERE tgt.tunnel_group_id = ?
+ ORDER BY t.id ASC
+ `, groupID)
+ if err != nil {
+ return nil, nil, err
+ }
+ defer rows.Close()
+
+ ids := make([]int64, 0)
+ names := make([]string, 0)
+ for rows.Next() {
+ var id int64
+ var name string
+ if err := rows.Scan(&id, &name); err != nil {
+ return nil, nil, err
+ }
+ ids = append(ids, id)
+ names = append(names, name)
+ }
+ if err := rows.Err(); err != nil {
+ return nil, nil, err
+ }
+ return ids, names, nil
+}
+
+func (r *Repository) listUserGroupMembers(groupID int64) ([]int64, []string, error) {
+ rows, err := r.db.Query(`
+ SELECT u.id, u.user
+ FROM user_group_user ugu
+ JOIN user u ON u.id = ugu.user_id
+ WHERE ugu.user_group_id = ?
+ ORDER BY u.id ASC
+ `, groupID)
+ if err != nil {
+ return nil, nil, err
+ }
+ defer rows.Close()
+
+ ids := make([]int64, 0)
+ names := make([]string, 0)
+ for rows.Next() {
+ var id int64
+ var name string
+ if err := rows.Scan(&id, &name); err != nil {
+ return nil, nil, err
+ }
+ ids = append(ids, id)
+ names = append(names, name)
+ }
+ if err := rows.Err(); err != nil {
+ return nil, nil, err
+ }
+ return ids, names, nil
+}
+
+func nullableString(v sql.NullString) interface{} {
+ if v.Valid {
+ return v.String
+ }
+ return nil
+}
+
+func nullableForwardIngress(v string) interface{} {
+ v = strings.TrimSpace(v)
+ if v == "" {
+ return nil
+ }
+ return v
+}
+
+func resolveForwardIngress(db *sql.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) {
+ var tunnelInIP sql.NullString
+ if err := db.QueryRow(`SELECT in_ip FROM tunnel WHERE id = ? LIMIT 1`, tunnelID).Scan(&tunnelInIP); err != nil {
+ if !errors.Is(err, sql.ErrNoRows) {
+ return "", sql.NullInt64{}, err
+ }
+ }
+
+ rows, err := db.Query(`
+ SELECT fp.port, n.server_ip
+ FROM forward_port fp
+ LEFT JOIN node n ON n.id = fp.node_id
+ WHERE fp.forward_id = ?
+ ORDER BY fp.id ASC
+ `, forwardID)
+ if err != nil {
+ return "", sql.NullInt64{}, err
+ }
+ defer rows.Close()
+
+ ports := make([]int64, 0)
+ nodePairs := make([]string, 0)
+ seenPorts := make(map[int64]struct{})
+ seenPairs := make(map[string]struct{})
+
+ for rows.Next() {
+ var port sql.NullInt64
+ var nodeIP sql.NullString
+ if err := rows.Scan(&port, &nodeIP); err != nil {
+ return "", sql.NullInt64{}, err
+ }
+ if !port.Valid {
+ continue
+ }
+ if _, ok := seenPorts[port.Int64]; !ok {
+ seenPorts[port.Int64] = struct{}{}
+ ports = append(ports, port.Int64)
+ }
+ if nodeIP.Valid && strings.TrimSpace(nodeIP.String) != "" {
+ pair := fmt.Sprintf("%s:%d", strings.TrimSpace(nodeIP.String), port.Int64)
+ if _, ok := seenPairs[pair]; !ok {
+ seenPairs[pair] = struct{}{}
+ nodePairs = append(nodePairs, pair)
+ }
+ }
+ }
+ if err := rows.Err(); err != nil {
+ return "", sql.NullInt64{}, err
+ }
+
+ if len(ports) == 0 {
+ return "", sql.NullInt64{}, nil
+ }
+
+ inPort := sql.NullInt64{Int64: ports[0], Valid: true}
+
+ entries := make([]string, 0)
+ if tunnelInIP.Valid && strings.TrimSpace(tunnelInIP.String) != "" {
+ tunnelIPs := strings.Split(tunnelInIP.String, ",")
+ seen := make(map[string]struct{})
+ for _, ip := range tunnelIPs {
+ ip = strings.TrimSpace(ip)
+ if ip == "" {
+ continue
+ }
+ if _, ok := seen[ip]; ok {
+ continue
+ }
+ seen[ip] = struct{}{}
+ for _, port := range ports {
+ entries = append(entries, fmt.Sprintf("%s:%d", ip, port))
+ }
+ }
+ } else {
+ entries = append(entries, nodePairs...)
+ }
+
+ return strings.Join(entries, ","), inPort, nil
+}
+
+func nullableInt64(v sql.NullInt64) interface{} {
+ if v.Valid {
+ return v.Int64
+ }
+ return nil
+}
+
+func unixMilliNow() int64 {
+ return time.Now().UnixMilli()
+}
+
+func ensureParentDir(dbPath string) error {
+ if dbPath == "" {
+ return fmt.Errorf("empty db path")
+ }
+ dir := filepath.Dir(dbPath)
+ if dir == "" || dir == "." {
+ return nil
+ }
+ return osMkdirAll(dir)
+}
+
+func bootstrapSchema(db *sql.DB) error {
+ if db == nil {
+ return errors.New("nil db")
+ }
+
+ var exists int
+ err := db.QueryRow(`SELECT COUNT(1) FROM sqlite_master WHERE type='table' AND name='user'`).Scan(&exists)
+ if err != nil {
+ return fmt.Errorf("check schema: %w", err)
+ }
+ if exists > 0 {
+ return nil
+ }
+
+ log.Printf("sqlite schema not found, bootstrapping embedded schema")
+ if _, err := db.Exec(embeddedSchema); err != nil {
+ return fmt.Errorf("apply schema.sql: %w", err)
+ }
+ if _, err := db.Exec(embeddedSeedData); err != nil {
+ return fmt.Errorf("apply data.sql: %w", err)
+ }
+ return nil
+}
+
+var osMkdirAll = func(path string) error {
+ return os.MkdirAll(path, 0o755)
+}
diff --git a/go-backend/internal/store/sqlite/sql/data.sql b/go-backend/internal/store/sqlite/sql/data.sql
new file mode 100644
index 0000000..ed78932
--- /dev/null
+++ b/go-backend/internal/store/sqlite/sql/data.sql
@@ -0,0 +1,5 @@
+INSERT OR IGNORE INTO user (id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
+VALUES (1, 'admin_user', '3c85cdebade1c51cf64ca9f3c09d182d', 0, 2727251700000, 99999, 0, 0, 1, 99999, 1748914865000, 1754011744252, 1);
+
+INSERT OR IGNORE INTO vite_config (id, name, value, time)
+VALUES (1, 'app_name', 'flux', 1755147963000);
diff --git a/go-backend/internal/store/sqlite/sql/schema.sql b/go-backend/internal/store/sqlite/sql/schema.sql
new file mode 100644
index 0000000..9330f04
--- /dev/null
+++ b/go-backend/internal/store/sqlite/sql/schema.sql
@@ -0,0 +1,182 @@
+-- SQLite Auto-generated schema
+-- This will be executed automatically on startup if tables don't exist
+
+CREATE TABLE IF NOT EXISTS forward (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ user_id INTEGER NOT NULL,
+ user_name VARCHAR(100) NOT NULL,
+ name VARCHAR(100) NOT NULL,
+ tunnel_id INTEGER NOT NULL,
+ remote_addr TEXT NOT NULL,
+ strategy VARCHAR(100) NOT NULL DEFAULT 'fifo',
+ in_flow INTEGER NOT NULL DEFAULT 0,
+ out_flow INTEGER NOT NULL DEFAULT 0,
+ created_time INTEGER NOT NULL,
+ updated_time INTEGER NOT NULL,
+ status INTEGER NOT NULL,
+ inx INTEGER NOT NULL DEFAULT 0
+);
+
+CREATE TABLE IF NOT EXISTS forward_port (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ forward_id INTEGER NOT NULL,
+ node_id INTEGER NOT NULL,
+ port INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS node (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ name VARCHAR(100) NOT NULL,
+ secret VARCHAR(100) NOT NULL,
+ server_ip VARCHAR(100) NOT NULL,
+ server_ip_v4 VARCHAR(100),
+ server_ip_v6 VARCHAR(100),
+ port TEXT NOT NULL,
+ interface_name VARCHAR(200),
+ version VARCHAR(100),
+ http INTEGER NOT NULL DEFAULT 0,
+ tls INTEGER NOT NULL DEFAULT 0,
+ socks INTEGER NOT NULL DEFAULT 0,
+ created_time INTEGER NOT NULL,
+ updated_time INTEGER,
+ status INTEGER NOT NULL,
+ tcp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
+ udp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
+ inx INTEGER NOT NULL DEFAULT 0
+);
+
+CREATE TABLE IF NOT EXISTS speed_limit (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ name VARCHAR(100) NOT NULL,
+ speed INTEGER NOT NULL,
+ tunnel_id INTEGER NOT NULL,
+ tunnel_name VARCHAR(100) NOT NULL,
+ created_time INTEGER NOT NULL,
+ updated_time INTEGER,
+ status INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS statistics_flow (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ user_id INTEGER NOT NULL,
+ flow INTEGER NOT NULL,
+ total_flow INTEGER NOT NULL,
+ time VARCHAR(100) NOT NULL,
+ created_time INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS tunnel (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ name VARCHAR(100) NOT NULL,
+ traffic_ratio REAL NOT NULL DEFAULT 1.0,
+ type INTEGER NOT NULL,
+ protocol VARCHAR(10) NOT NULL DEFAULT 'tls',
+ flow INTEGER NOT NULL,
+ created_time INTEGER NOT NULL,
+ updated_time INTEGER NOT NULL,
+ status INTEGER NOT NULL,
+ in_ip TEXT,
+ inx INTEGER NOT NULL DEFAULT 0
+);
+
+CREATE TABLE IF NOT EXISTS chain_tunnel (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ tunnel_id INTEGER NOT NULL ,
+ chain_type VARCHAR(10) NOT NULL,
+ node_id INTEGER NOT NULL ,
+ port INTEGER,
+ strategy VARCHAR(10),
+ inx INTEGER,
+ protocol VARCHAR(10)
+);
+
+
+CREATE TABLE IF NOT EXISTS user (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ user VARCHAR(100) NOT NULL,
+ pwd VARCHAR(100) NOT NULL,
+ role_id INTEGER NOT NULL,
+ exp_time INTEGER NOT NULL,
+ flow INTEGER NOT NULL,
+ in_flow INTEGER NOT NULL DEFAULT 0,
+ out_flow INTEGER NOT NULL DEFAULT 0,
+ flow_reset_time INTEGER NOT NULL,
+ num INTEGER NOT NULL,
+ created_time INTEGER NOT NULL,
+ updated_time INTEGER,
+ status INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS user_tunnel (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ user_id INTEGER NOT NULL,
+ tunnel_id INTEGER NOT NULL,
+ speed_id INTEGER,
+ num INTEGER NOT NULL,
+ flow INTEGER NOT NULL,
+ in_flow INTEGER NOT NULL DEFAULT 0,
+ out_flow INTEGER NOT NULL DEFAULT 0,
+ flow_reset_time INTEGER NOT NULL,
+ exp_time INTEGER NOT NULL,
+ status INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS tunnel_group (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ name VARCHAR(100) NOT NULL,
+ created_time INTEGER NOT NULL,
+ updated_time INTEGER NOT NULL,
+ status INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS user_group (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ name VARCHAR(100) NOT NULL,
+ created_time INTEGER NOT NULL,
+ updated_time INTEGER NOT NULL,
+ status INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS tunnel_group_tunnel (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ tunnel_group_id INTEGER NOT NULL,
+ tunnel_id INTEGER NOT NULL,
+ created_time INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS user_group_user (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ user_group_id INTEGER NOT NULL,
+ user_id INTEGER NOT NULL,
+ created_time INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS group_permission (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ user_group_id INTEGER NOT NULL,
+ tunnel_group_id INTEGER NOT NULL,
+ created_time INTEGER NOT NULL
+);
+
+CREATE TABLE IF NOT EXISTS group_permission_grant (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ user_group_id INTEGER NOT NULL,
+ tunnel_group_id INTEGER NOT NULL,
+ user_tunnel_id INTEGER NOT NULL,
+ created_by_group INTEGER NOT NULL DEFAULT 0,
+ created_time INTEGER NOT NULL
+);
+
+CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_name ON tunnel_group(name);
+CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_name ON user_group(name);
+CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_tunnel_unique ON tunnel_group_tunnel(tunnel_group_id, tunnel_id);
+CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_user_unique ON user_group_user(user_group_id, user_id);
+CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_unique ON group_permission(user_group_id, tunnel_group_id);
+CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_grant_unique ON group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id);
+
+CREATE TABLE IF NOT EXISTS vite_config (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ name VARCHAR(200) NOT NULL UNIQUE,
+ value VARCHAR(200) NOT NULL,
+ time INTEGER NOT NULL
+);
diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go
new file mode 100644
index 0000000..f36ec97
--- /dev/null
+++ b/go-backend/internal/ws/server.go
@@ -0,0 +1,430 @@
+package ws
+
+import (
+ "encoding/json"
+ "errors"
+ "fmt"
+ "log"
+ "net/http"
+ "strconv"
+ "strings"
+ "sync"
+ "time"
+
+ "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 commandResponse struct {
+ Type string `json:"type"`
+ Success bool `json:"success"`
+ Message string `json:"message"`
+ Data json.RawMessage `json:"data,omitempty"`
+ RequestID string `json:"requestId,omitempty"`
+}
+
+type pendingRequest struct {
+ nodeID int64
+ ch chan CommandResult
+}
+
+type CommandResult struct {
+ Type string `json:"type"`
+ Success bool `json:"success"`
+ Message string `json:"message"`
+ Data map[string]interface{} `json:"data,omitempty"`
+}
+
+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
+ pending map[string]pendingRequest
+}
+
+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),
+ pending: make(map[string]pendingRequest),
+ }
+}
+
+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.failPendingForNode(nodeID, "节点连接已断开")
+ _ = 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.tryResolvePending(nodeID, msg)
+ s.broadcastInfo(nodeID, msg)
+ }
+}
+
+func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, timeout time.Duration) (CommandResult, error) {
+ if s == nil {
+ return CommandResult{}, errors.New("server not initialized")
+ }
+ if strings.TrimSpace(cmdType) == "" {
+ return CommandResult{}, errors.New("command type is empty")
+ }
+ if timeout <= 0 {
+ timeout = 10 * time.Second
+ }
+
+ s.mu.RLock()
+ ns, ok := s.nodes[nodeID]
+ s.mu.RUnlock()
+ if !ok || ns == nil || ns.conn == nil || ns.conn.conn == nil {
+ return CommandResult{}, errors.New("节点不在线")
+ }
+
+ requestID := fmt.Sprintf("%d_%d", nodeID, time.Now().UnixNano())
+ ch := make(chan CommandResult, 1)
+
+ s.mu.Lock()
+ s.pending[requestID] = pendingRequest{nodeID: nodeID, ch: ch}
+ s.mu.Unlock()
+
+ cleanup := func() {
+ s.mu.Lock()
+ if p, exists := s.pending[requestID]; exists {
+ delete(s.pending, requestID)
+ close(p.ch)
+ }
+ s.mu.Unlock()
+ }
+
+ cmdPayload := map[string]interface{}{
+ "type": cmdType,
+ "data": data,
+ "requestId": requestID,
+ }
+ rawCmd, err := json.Marshal(cmdPayload)
+ if err != nil {
+ cleanup()
+ return CommandResult{}, err
+ }
+
+ messageData := rawCmd
+ if strings.TrimSpace(ns.secret) != "" {
+ crypto, err := security.NewAESCrypto(ns.secret)
+ if err != nil {
+ cleanup()
+ return CommandResult{}, err
+ }
+ encrypted, err := crypto.Encrypt(rawCmd)
+ if err != nil {
+ cleanup()
+ return CommandResult{}, err
+ }
+ wrapper := map[string]interface{}{
+ "encrypted": true,
+ "data": encrypted,
+ "timestamp": time.Now().UnixMilli(),
+ }
+ messageData, err = json.Marshal(wrapper)
+ if err != nil {
+ cleanup()
+ return CommandResult{}, err
+ }
+ }
+
+ ns.conn.mu.Lock()
+ err = ns.conn.conn.WriteMessage(websocket.TextMessage, messageData)
+ ns.conn.mu.Unlock()
+ if err != nil {
+ cleanup()
+ return CommandResult{}, err
+ }
+
+ select {
+ case result, ok := <-ch:
+ if !ok {
+ return CommandResult{}, errors.New("命令通道已关闭")
+ }
+ if !result.Success {
+ if strings.TrimSpace(result.Message) == "" {
+ result.Message = "命令执行失败"
+ }
+ return result, errors.New(result.Message)
+ }
+ return result, nil
+ case <-time.After(timeout):
+ cleanup()
+ return CommandResult{}, errors.New("等待节点响应超时")
+ }
+}
+
+func (s *Server) tryResolvePending(nodeID int64, message string) {
+ if s == nil || strings.TrimSpace(message) == "" {
+ return
+ }
+
+ var resp commandResponse
+ if err := json.Unmarshal([]byte(message), &resp); err != nil {
+ return
+ }
+ if strings.TrimSpace(resp.RequestID) == "" {
+ return
+ }
+
+ s.mu.Lock()
+ p, ok := s.pending[resp.RequestID]
+ if ok {
+ delete(s.pending, resp.RequestID)
+ }
+ s.mu.Unlock()
+ if !ok {
+ return
+ }
+ if p.nodeID != nodeID {
+ select {
+ case p.ch <- CommandResult{Type: resp.Type, Success: false, Message: "节点响应与请求不匹配"}:
+ default:
+ }
+ close(p.ch)
+ return
+ }
+
+ result := CommandResult{
+ Type: resp.Type,
+ Success: resp.Success,
+ Message: resp.Message,
+ }
+ if len(resp.Data) > 0 {
+ var data map[string]interface{}
+ if err := json.Unmarshal(resp.Data, &data); err == nil {
+ result.Data = data
+ }
+ }
+
+ select {
+ case p.ch <- result:
+ default:
+ }
+ close(p.ch)
+}
+
+func (s *Server) failPendingForNode(nodeID int64, message string) {
+ if s == nil {
+ return
+ }
+
+ type pair struct {
+ id string
+ pr pendingRequest
+ }
+ items := make([]pair, 0)
+
+ s.mu.Lock()
+ for id, pr := range s.pending {
+ if pr.nodeID != nodeID {
+ continue
+ }
+ items = append(items, pair{id: id, pr: pr})
+ delete(s.pending, id)
+ }
+ s.mu.Unlock()
+
+ for _, item := range items {
+ select {
+ case item.pr.ch <- CommandResult{Success: false, Message: message}:
+ default:
+ }
+ close(item.pr.ch)
+ }
+}
+
+func (s *Server) broadcastStatus(nodeID int64, status int) {
+ payload := map[string]interface{}{
+ "id": strconv.FormatInt(nodeID, 10),
+ "type": "status",
+ "data": status,
+ }
+ raw, _ := json.Marshal(payload)
+ s.broadcastToAdmins(string(raw))
+}
+
+func (s *Server) broadcastInfo(nodeID int64, data string) {
+ payload := broadcastMessage{ID: nodeID, Type: "info", Data: data}
+ raw, _ := json.Marshal(payload)
+ s.broadcastToAdmins(string(raw))
+}
+
+func (s *Server) broadcastToAdmins(message string) {
+ s.mu.RLock()
+ admins := make([]*connWrap, 0, len(s.admins))
+ for c := range s.admins {
+ admins = append(admins, c)
+ }
+ s.mu.RUnlock()
+
+ for _, c := range admins {
+ c.mu.Lock()
+ err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
+ c.mu.Unlock()
+ if err != nil {
+ log.Printf("websocket broadcast failed: %v", err)
+ }
+ }
+}
+
+func decryptIfNeeded(payload []byte, secret string) string {
+ text := string(payload)
+ var wrap encryptedMessage
+ if err := json.Unmarshal(payload, &wrap); err != nil || !wrap.Encrypted || strings.TrimSpace(wrap.Data) == "" {
+ return text
+ }
+
+ crypto, err := security.NewAESCrypto(secret)
+ if err != nil {
+ return text
+ }
+ plain, err := crypto.Decrypt(wrap.Data)
+ if err != nil {
+ return text
+ }
+ return string(plain)
+}
+
+func parseIntDefault(v string, fallback int) int {
+ x, err := strconv.Atoi(v)
+ if err != nil {
+ return fallback
+ }
+ return x
+}
diff --git a/go-backend/tests/contract/auth_contract_test.go b/go-backend/tests/contract/auth_contract_test.go
new file mode 100644
index 0000000..43d01c8
--- /dev/null
+++ b/go-backend/tests/contract/auth_contract_test.go
@@ -0,0 +1,90 @@
+package contract_test
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "go-backend/internal/auth"
+ "go-backend/internal/http/middleware"
+ "go-backend/internal/http/response"
+)
+
+func TestJWTMiddlewareContracts(t *testing.T) {
+ secret := "unit-test-secret"
+
+ next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ response.WriteJSON(w, response.OK("pass"))
+ })
+
+ wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret})(next)
+
+ t.Run("login path is excluded", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", nil)
+ res := httptest.NewRecorder()
+ wrapped.ServeHTTP(res, req)
+ assertCode(t, res, 0)
+ })
+
+ t.Run("missing token returns 401 contract message", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
+ res := httptest.NewRecorder()
+ wrapped.ServeHTTP(res, req)
+ assertCodeMsg(t, res, 401, "未登录或token已过期")
+ })
+
+ t.Run("invalid token returns 401 contract message", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
+ req.Header.Set("Authorization", "invalid.token.value")
+ res := httptest.NewRecorder()
+ wrapped.ServeHTTP(res, req)
+ assertCodeMsg(t, res, 401, "无效的token或token已过期")
+ })
+
+ t.Run("valid token reaches next", func(t *testing.T) {
+ token, err := auth.GenerateToken(1, "admin_user", 0, secret)
+ if err != nil {
+ t.Fatalf("generate token: %v", err)
+ }
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
+ req.Header.Set("Authorization", token)
+ res := httptest.NewRecorder()
+ wrapped.ServeHTTP(res, req)
+ assertCode(t, res, 0)
+ })
+
+ t.Run("non-admin blocked on admin path", func(t *testing.T) {
+ token, err := auth.GenerateToken(2, "normal_user", 1, secret)
+ if err != nil {
+ t.Fatalf("generate token: %v", err)
+ }
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", nil)
+ req.Header.Set("Authorization", token)
+ res := httptest.NewRecorder()
+ wrapped.ServeHTTP(res, req)
+ assertCodeMsg(t, res, 403, "权限不足,仅管理员可操作")
+ })
+}
+
+func assertCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
+ t.Helper()
+ var out response.R
+ if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != expected {
+ t.Fatalf("expected code %d, got %d", expected, out.Code)
+ }
+}
+
+func assertCodeMsg(t *testing.T, rec *httptest.ResponseRecorder, expectedCode int, expectedMsg string) {
+ t.Helper()
+ var out response.R
+ if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != expectedCode || out.Msg != expectedMsg {
+ t.Fatalf("expected (%d,%q), got (%d,%q)", expectedCode, expectedMsg, out.Code, out.Msg)
+ }
+}
diff --git a/go-backend/tests/contract/diagnosis_contract_test.go b/go-backend/tests/contract/diagnosis_contract_test.go
new file mode 100644
index 0000000..cc3c76a
--- /dev/null
+++ b/go-backend/tests/contract/diagnosis_contract_test.go
@@ -0,0 +1,239 @@
+package contract
+
+import (
+ "bytes"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "path/filepath"
+ "strconv"
+ "strings"
+ "testing"
+ "time"
+
+ "go-backend/internal/auth"
+ httpserver "go-backend/internal/http"
+ "go-backend/internal/http/handler"
+ "go-backend/internal/http/response"
+ "go-backend/internal/store/sqlite"
+)
+
+func TestDiagnosisChainCoverageContracts(t *testing.T) {
+ secret := "contract-jwt-secret"
+ router, repo := setupDiagnosisContractRouter(t, secret)
+ now := time.Now().UnixMilli()
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
+ VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
+ `, now, now); err != nil {
+ t.Fatalf("insert user: %v", err)
+ }
+
+ tunnelRes, err := repo.DB().Exec(`
+ INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, "diagnose-chain-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0)
+ if err != nil {
+ t.Fatalf("insert tunnel: %v", err)
+ }
+ tunnelID, err := tunnelRes.LastInsertId()
+ if err != nil {
+ t.Fatalf("get tunnel id: %v", err)
+ }
+
+ insertNode := func(name, ip string) int64 {
+ res, err := repo.DB().Exec(`
+ INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0)
+ if err != nil {
+ t.Fatalf("insert node %s: %v", name, err)
+ }
+ id, err := res.LastInsertId()
+ if err != nil {
+ t.Fatalf("get node id %s: %v", name, err)
+ }
+ return id
+ }
+
+ entryNodeID := insertNode("entry-node", "10.0.1.10")
+ chainNodeID := insertNode("chain-node", "10.0.1.20")
+ exitNodeID := insertNode("exit-node", "10.0.1.30")
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
+ VALUES(?, 1, ?, 30001, 'round', 1, 'tls')
+ `, tunnelID, entryNodeID); err != nil {
+ t.Fatalf("insert entry chain: %v", err)
+ }
+ if _, err := repo.DB().Exec(`
+ INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
+ VALUES(?, 2, ?, 30002, 'round', 1, 'tls')
+ `, tunnelID, chainNodeID); err != nil {
+ t.Fatalf("insert middle chain: %v", err)
+ }
+ if _, err := repo.DB().Exec(`
+ INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
+ VALUES(?, 3, ?, 30003, 'round', 1, 'tls')
+ `, tunnelID, exitNodeID); err != nil {
+ t.Fatalf("insert exit chain: %v", err)
+ }
+
+ forwardRes, err := repo.DB().Exec(`
+ INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
+ VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
+ `, 2, "normal_user", "chain-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 0)
+ if err != nil {
+ t.Fatalf("insert forward: %v", err)
+ }
+ forwardID, err := forwardRes.LastInsertId()
+ if err != nil {
+ t.Fatalf("get forward id: %v", err)
+ }
+
+ userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
+ if err != nil {
+ t.Fatalf("generate user token: %v", err)
+ }
+ adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
+ if err != nil {
+ t.Fatalf("generate admin token: %v", err)
+ }
+
+ t.Run("forward diagnose includes entry chain exit paths", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+strconv.FormatInt(forwardID, 10)+`}`))
+ req.Header.Set("Authorization", userToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+
+ payload, ok := out.Data.(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected object payload, got %T", out.Data)
+ }
+ results, ok := payload["results"].([]interface{})
+ if !ok || len(results) == 0 {
+ t.Fatalf("expected non-empty results, got %v", payload["results"])
+ }
+
+ hasEntryToChain := false
+ hasChainToExit := false
+ hasExitToTarget := false
+ for _, raw := range results {
+ item, ok := raw.(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected result object, got %T", raw)
+ }
+ if strings.TrimSpace(valueAsString(item["message"])) == "" {
+ t.Fatalf("expected non-empty message field")
+ }
+ from := valueAsInt(item["fromChainType"])
+ to := valueAsInt(item["toChainType"])
+ if from == 1 && to == 2 {
+ hasEntryToChain = true
+ }
+ if from == 2 && to == 3 {
+ hasChainToExit = true
+ }
+ if from == 3 {
+ hasExitToTarget = true
+ }
+ }
+
+ if !hasEntryToChain || !hasChainToExit || !hasExitToTarget {
+ t.Fatalf("expected entry->chain, chain->exit, exit->target coverage; got entry=%v chain=%v exit=%v", hasEntryToChain, hasChainToExit, hasExitToTarget)
+ }
+ })
+
+ t.Run("tunnel diagnose includes entry chain exit groups", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose", bytes.NewBufferString(`{"tunnelId":`+strconv.FormatInt(tunnelID, 10)+`}`))
+ req.Header.Set("Authorization", adminToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+
+ payload, ok := out.Data.(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected object payload, got %T", out.Data)
+ }
+ results, ok := payload["results"].([]interface{})
+ if !ok || len(results) == 0 {
+ t.Fatalf("expected non-empty results, got %v", payload["results"])
+ }
+
+ hasEntry := false
+ hasChain := false
+ hasExit := false
+ for _, raw := range results {
+ item, ok := raw.(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected result object, got %T", raw)
+ }
+ if strings.TrimSpace(valueAsString(item["message"])) == "" {
+ t.Fatalf("expected non-empty message field")
+ }
+ switch valueAsInt(item["fromChainType"]) {
+ case 1:
+ hasEntry = true
+ case 2:
+ hasChain = true
+ case 3:
+ hasExit = true
+ }
+ }
+
+ if !hasEntry || !hasChain || !hasExit {
+ t.Fatalf("expected entry/chain/exit groups, got entry=%v chain=%v exit=%v", hasEntry, hasChain, hasExit)
+ }
+ })
+}
+
+func valueAsInt(v interface{}) int {
+ switch n := v.(type) {
+ case float64:
+ return int(n)
+ case int:
+ return n
+ case int64:
+ return int(n)
+ default:
+ return 0
+ }
+}
+
+func valueAsString(v interface{}) string {
+ s, _ := v.(string)
+ return s
+}
+
+func setupDiagnosisContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) {
+ t.Helper()
+ dbPath := filepath.Join(t.TempDir(), "diagnosis-contract.db")
+ repo, err := sqlite.Open(dbPath)
+ if err != nil {
+ t.Fatalf("open sqlite: %v", err)
+ }
+ t.Cleanup(func() {
+ _ = repo.Close()
+ })
+
+ h := handler.New(repo, jwtSecret)
+ return httpserver.NewRouter(h, jwtSecret), repo
+}
diff --git a/go-backend/tests/contract/flow_contract_test.go b/go-backend/tests/contract/flow_contract_test.go
new file mode 100644
index 0000000..f45f247
--- /dev/null
+++ b/go-backend/tests/contract/flow_contract_test.go
@@ -0,0 +1,44 @@
+package contract_test
+
+import (
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "go-backend/internal/http/handler"
+)
+
+func TestFlowEndpointsStringResponses(t *testing.T) {
+ h := handler.New(nil, "secret")
+ mux := http.NewServeMux()
+ h.Register(mux)
+
+ tests := []struct {
+ name string
+ method string
+ path string
+ expected string
+ }{
+ {name: "flow test", method: http.MethodGet, path: "/flow/test", expected: "test"},
+ {name: "flow config", method: http.MethodPost, path: "/flow/config?secret=abc", expected: "ok"},
+ {name: "flow upload", method: http.MethodPost, path: "/flow/upload?secret=abc", expected: "ok"},
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ req := httptest.NewRequest(tc.method, tc.path, nil)
+ res := httptest.NewRecorder()
+ mux.ServeHTTP(res, req)
+
+ body, err := io.ReadAll(res.Body)
+ if err != nil {
+ t.Fatalf("read body: %v", err)
+ }
+
+ if string(body) != tc.expected {
+ t.Fatalf("expected %q, got %q", tc.expected, string(body))
+ }
+ })
+ }
+}
diff --git a/go-backend/tests/contract/forward_contract_test.go b/go-backend/tests/contract/forward_contract_test.go
new file mode 100644
index 0000000..7868eb9
--- /dev/null
+++ b/go-backend/tests/contract/forward_contract_test.go
@@ -0,0 +1,202 @@
+package contract_test
+
+import (
+ "bytes"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "strconv"
+ "testing"
+ "time"
+
+ "go-backend/internal/auth"
+ "go-backend/internal/http/response"
+)
+
+func TestForwardOwnershipAndScopeContracts(t *testing.T) {
+ secret := "contract-jwt-secret"
+ router, repo := setupContractRouter(t, secret)
+ now := time.Now().UnixMilli()
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
+ VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
+ `, now, now); err != nil {
+ t.Fatalf("insert user: %v", err)
+ }
+
+ res, err := repo.DB().Exec(`
+ INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, "contract-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0)
+ if err != nil {
+ t.Fatalf("insert tunnel: %v", err)
+ }
+ tunnelID, err := res.LastInsertId()
+ if err != nil {
+ t.Fatalf("get tunnel id: %v", err)
+ }
+
+ nodeRes, err := repo.DB().Exec(`
+ INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, "entry-node", "entry-secret", "10.0.0.10", "10.0.0.10", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0)
+ if err != nil {
+ t.Fatalf("insert node: %v", err)
+ }
+ entryNodeID, err := nodeRes.LastInsertId()
+ if err != nil {
+ t.Fatalf("get node id: %v", err)
+ }
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
+ VALUES(?, 1, ?, 20001, 'round', 1, 'tls')
+ `, tunnelID, entryNodeID); err != nil {
+ t.Fatalf("insert chain_tunnel: %v", err)
+ }
+
+ resAdmin, err := repo.DB().Exec(`
+ INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
+ VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
+ `, 1, "admin_user", "admin-forward", tunnelID, "1.1.1.1:443", "fifo", now, now, 0)
+ if err != nil {
+ t.Fatalf("insert admin forward: %v", err)
+ }
+ adminForwardID, err := resAdmin.LastInsertId()
+ if err != nil {
+ t.Fatalf("get admin forward id: %v", err)
+ }
+
+ resUser, err := repo.DB().Exec(`
+ INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
+ VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
+ `, 2, "normal_user", "user-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 1)
+ if err != nil {
+ t.Fatalf("insert user forward: %v", err)
+ }
+ userForwardID, err := resUser.LastInsertId()
+ if err != nil {
+ t.Fatalf("get user forward id: %v", err)
+ }
+
+ userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
+ if err != nil {
+ t.Fatalf("generate user token: %v", err)
+ }
+ adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
+ if err != nil {
+ t.Fatalf("generate admin token: %v", err)
+ }
+
+ t.Run("non-owner cannot delete another user's forward", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/delete", bytes.NewBufferString(`{"id":`+jsonNumber(adminForwardID)+`}`))
+ req.Header.Set("Authorization", userToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ assertCodeMsg(t, res, -1, "转发不存在")
+ })
+
+ t.Run("non-admin forward list is scoped to owner", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/list", bytes.NewBufferString(`{}`))
+ req.Header.Set("Authorization", userToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+ arr, ok := out.Data.([]interface{})
+ if !ok {
+ t.Fatalf("expected array data, got %T", out.Data)
+ }
+ if len(arr) != 1 {
+ t.Fatalf("expected 1 forward, got %d", len(arr))
+ }
+ item, ok := arr[0].(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected object item, got %T", arr[0])
+ }
+ if got := int64(item["id"].(float64)); got != userForwardID {
+ t.Fatalf("expected forward id %d, got %d", userForwardID, got)
+ }
+ })
+
+ t.Run("forward diagnose returns structured payload", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+jsonNumber(userForwardID)+`}`))
+ req.Header.Set("Authorization", userToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+
+ payload, ok := out.Data.(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected object payload, got %T", out.Data)
+ }
+ results, ok := payload["results"].([]interface{})
+ if !ok || len(results) == 0 {
+ t.Fatalf("expected non-empty results, got %v", payload["results"])
+ }
+ first, ok := results[0].(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected result object, got %T", results[0])
+ }
+ if _, ok := first["message"]; !ok {
+ t.Fatalf("expected message field in diagnosis result")
+ }
+ if got := int(first["fromChainType"].(float64)); got != 1 {
+ t.Fatalf("expected fromChainType=1, got %d", got)
+ }
+ })
+
+ t.Run("tunnel diagnose returns structured payload", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose", bytes.NewBufferString(`{"tunnelId":`+jsonNumber(tunnelID)+`}`))
+ req.Header.Set("Authorization", adminToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+
+ payload, ok := out.Data.(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected object payload, got %T", out.Data)
+ }
+ results, ok := payload["results"].([]interface{})
+ if !ok || len(results) == 0 {
+ t.Fatalf("expected non-empty results, got %v", payload["results"])
+ }
+ first, ok := results[0].(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected result object, got %T", results[0])
+ }
+ if _, ok := first["message"]; !ok {
+ t.Fatalf("expected message field in tunnel diagnosis result")
+ }
+ })
+}
+
+func jsonNumber(v int64) string {
+ return strconv.FormatInt(v, 10)
+}
diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go
new file mode 100644
index 0000000..76ec67e
--- /dev/null
+++ b/go-backend/tests/contract/migration_contract_test.go
@@ -0,0 +1,215 @@
+package contract_test
+
+import (
+ "bytes"
+ "encoding/json"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "path/filepath"
+ "strconv"
+ "strings"
+ "testing"
+ "time"
+
+ "go-backend/internal/auth"
+ httpserver "go-backend/internal/http"
+ "go-backend/internal/http/handler"
+ "go-backend/internal/http/response"
+ "go-backend/internal/store/sqlite"
+)
+
+func TestCaptchaVerifyLoginContract(t *testing.T) {
+ secret := "contract-jwt-secret"
+ router, repo := setupContractRouter(t, secret)
+
+ _, err := repo.DB().Exec(`
+ INSERT INTO vite_config(name, value, time)
+ VALUES(?, ?, ?)
+ ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
+ `, "captcha_enabled", "true", time.Now().UnixMilli())
+ if err != nil {
+ t.Fatalf("enable captcha: %v", err)
+ }
+
+ t.Run("login denied without verified captcha token", func(t *testing.T) {
+ body := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":""}`)
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
+ req.Header.Set("Content-Type", "application/json")
+ resp := httptest.NewRecorder()
+
+ router.ServeHTTP(resp, req)
+
+ assertCodeMsg(t, resp, -1, "验证码校验失败")
+ })
+
+ t.Run("captcha token is one-time and consumed by login", func(t *testing.T) {
+ verifyReq := httptest.NewRequest(http.MethodPost, "/api/v1/captcha/verify", bytes.NewBufferString(`{"id":"captcha-token-1","data":"ok"}`))
+ verifyReq.Header.Set("Content-Type", "application/json")
+ verifyResp := httptest.NewRecorder()
+
+ router.ServeHTTP(verifyResp, verifyReq)
+
+ var verifyOut struct {
+ Success bool `json:"success"`
+ Data struct {
+ ValidToken string `json:"validToken"`
+ } `json:"data"`
+ }
+ if err := json.NewDecoder(verifyResp.Body).Decode(&verifyOut); err != nil {
+ t.Fatalf("decode captcha verify response: %v", err)
+ }
+ if !verifyOut.Success || verifyOut.Data.ValidToken != "captcha-token-1" {
+ t.Fatalf("unexpected captcha verify payload: success=%v token=%q", verifyOut.Success, verifyOut.Data.ValidToken)
+ }
+
+ loginBody := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":"captcha-token-1"}`)
+ loginReq := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", loginBody)
+ loginReq.Header.Set("Content-Type", "application/json")
+ loginResp := httptest.NewRecorder()
+ router.ServeHTTP(loginResp, loginReq)
+ assertCode(t, loginResp, 0)
+
+ replayBody := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":"captcha-token-1"}`)
+ replayReq := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", replayBody)
+ replayReq.Header.Set("Content-Type", "application/json")
+ replayResp := httptest.NewRecorder()
+ router.ServeHTTP(replayResp, replayReq)
+ assertCodeMsg(t, replayResp, -1, "验证码校验失败")
+ })
+}
+
+func TestOpenAPISubStoreContracts(t *testing.T) {
+ router, repo := setupContractRouter(t, "contract-jwt-secret")
+
+ const tunnelFlowGB = int64(500)
+ const tunnelInFlow = int64(123)
+ const tunnelOutFlow = int64(456)
+ const tunnelExpTimeMs = int64(2727251700000)
+
+ now := time.Now().UnixMilli()
+ res, err := repo.DB().Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
+ "contract-tunnel", 1.0, 1, "tls", 1, now, now, 1, nil, 0)
+ if err != nil {
+ t.Fatalf("insert tunnel: %v", err)
+ }
+ tunnelID, err := res.LastInsertId()
+ if err != nil {
+ t.Fatalf("last insert id: %v", err)
+ }
+ if _, err := repo.DB().Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, ?, ?, ?, ?, ?)`,
+ 1, tunnelID, 99999, tunnelFlowGB, tunnelInFlow, tunnelOutFlow, 1, tunnelExpTimeMs, 1); err != nil {
+ t.Fatalf("insert user_tunnel: %v", err)
+ }
+
+ t.Run("default user subscription payload", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user", nil)
+ resp := httptest.NewRecorder()
+
+ router.ServeHTTP(resp, req)
+
+ body, err := io.ReadAll(resp.Body)
+ if err != nil {
+ t.Fatalf("read body: %v", err)
+ }
+
+ expected := "upload=0; download=0; total=107373108658176; expire=2727251700"
+ if string(body) != expected {
+ t.Fatalf("expected body %q, got %q", expected, string(body))
+ }
+ if got := resp.Header().Get("subscription-userinfo"); got != expected {
+ t.Fatalf("expected subscription-userinfo %q, got %q", expected, got)
+ }
+ if !strings.Contains(resp.Header().Get("Content-Type"), "text/plain") {
+ t.Fatalf("expected text/plain content type, got %q", resp.Header().Get("Content-Type"))
+ }
+ })
+
+ t.Run("tunnel scoped subscription payload", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user&tunnel="+strconv.FormatInt(tunnelID, 10), nil)
+ resp := httptest.NewRecorder()
+
+ router.ServeHTTP(resp, req)
+
+ body, err := io.ReadAll(resp.Body)
+ if err != nil {
+ t.Fatalf("read body: %v", err)
+ }
+
+ expected := "upload=123; download=456; total=536870912000; expire=2727251700"
+ if string(body) != expected {
+ t.Fatalf("expected body %q, got %q", expected, string(body))
+ }
+ if got := resp.Header().Get("subscription-userinfo"); got != expected {
+ t.Fatalf("expected subscription-userinfo %q, got %q", expected, got)
+ }
+ })
+
+ t.Run("invalid credentials returns contract error", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=wrong", nil)
+ resp := httptest.NewRecorder()
+
+ router.ServeHTTP(resp, req)
+
+ assertCodeMsg(t, resp, -1, "鉴权失败")
+ })
+
+ t.Run("missing tunnel returns contract error", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=admin_user&pwd=admin_user&tunnel=999999", nil)
+ resp := httptest.NewRecorder()
+
+ router.ServeHTTP(resp, req)
+
+ assertCodeMsg(t, resp, -1, "隧道不存在")
+ })
+}
+
+func TestSpeedLimitTunnelsRouteAlias(t *testing.T) {
+ secret := "contract-jwt-secret"
+ router, _ := setupContractRouter(t, secret)
+
+ t.Run("missing token blocked", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil)
+ resp := httptest.NewRecorder()
+
+ router.ServeHTTP(resp, req)
+
+ assertCodeMsg(t, resp, 401, "未登录或token已过期")
+ })
+
+ t.Run("admin token receives success envelope", 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/speed-limit/tunnels", nil)
+ req.Header.Set("Authorization", token)
+ resp := httptest.NewRecorder()
+
+ router.ServeHTTP(resp, req)
+
+ var out response.R
+ if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+ })
+}
+
+func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) {
+ t.Helper()
+ dbPath := filepath.Join(t.TempDir(), "contract.db")
+ repo, err := sqlite.Open(dbPath)
+ if err != nil {
+ t.Fatalf("open sqlite: %v", err)
+ }
+ t.Cleanup(func() {
+ _ = repo.Close()
+ })
+
+ h := handler.New(repo, jwtSecret)
+ return httpserver.NewRouter(h, jwtSecret), repo
+}
diff --git a/go-backend/tests/contract/tunnel_create_contract_test.go b/go-backend/tests/contract/tunnel_create_contract_test.go
new file mode 100644
index 0000000..407a744
--- /dev/null
+++ b/go-backend/tests/contract/tunnel_create_contract_test.go
@@ -0,0 +1,151 @@
+package contract_test
+
+import (
+ "bytes"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "strconv"
+ "strings"
+ "testing"
+ "time"
+
+ "go-backend/internal/auth"
+ "go-backend/internal/http/response"
+)
+
+func TestTunnelCreateRuntimeRollbackContract(t *testing.T) {
+ secret := "contract-jwt-secret"
+ router, repo := setupContractRouter(t, secret)
+ now := time.Now().UnixMilli()
+
+ adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
+ if err != nil {
+ t.Fatalf("generate admin token: %v", err)
+ }
+
+ insertNode := func(name, ip, portRange string) int64 {
+ res, err := repo.DB().Exec(`
+ INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0)
+ if err != nil {
+ t.Fatalf("insert node %s: %v", name, err)
+ }
+ id, err := res.LastInsertId()
+ if err != nil {
+ t.Fatalf("get node id %s: %v", name, err)
+ }
+ return id
+ }
+
+ entryID := insertNode("create-entry", "10.20.0.1", "30000-30010")
+ chainID := insertNode("create-chain", "10.20.0.2", "31000-31010")
+ exitID := insertNode("create-exit", "10.20.0.3", "32000-32010")
+
+ payload := `{"name":"runtime-rollback-tunnel","type":2,"flow":99999,"status":1,"inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[[{"nodeId":` + jsonInt(chainID) + `,"protocol":"tls","strategy":"round"}]],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}`
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", bytes.NewBufferString(payload))
+ req.Header.Set("Authorization", adminToken)
+ req.Header.Set("Content-Type", "application/json")
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code == 0 {
+ t.Fatalf("expected create failure when nodes are offline")
+ }
+ if !strings.Contains(out.Msg, "节点") {
+ t.Fatalf("expected node-related error, got %q", out.Msg)
+ }
+
+ var tunnelCount int
+ if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, "runtime-rollback-tunnel").Scan(&tunnelCount); err != nil {
+ t.Fatalf("count tunnel: %v", err)
+ }
+ if tunnelCount != 0 {
+ t.Fatalf("expected tunnel rollback, found %d records", tunnelCount)
+ }
+
+ var chainCount int
+ if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM chain_tunnel`).Scan(&chainCount); err != nil {
+ t.Fatalf("count chain_tunnel: %v", err)
+ }
+ if chainCount != 0 {
+ t.Fatalf("expected chain_tunnel rollback, found %d records", chainCount)
+ }
+}
+
+func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) {
+ secret := "contract-jwt-secret"
+ router, repo := setupContractRouter(t, secret)
+ now := time.Now().UnixMilli()
+
+ adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
+ if err != nil {
+ t.Fatalf("generate admin token: %v", err)
+ }
+
+ insertNode := func(name, ip, portRange string) int64 {
+ res, err := repo.DB().Exec(`
+ INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0)
+ if err != nil {
+ t.Fatalf("insert node %s: %v", name, err)
+ }
+ id, err := res.LastInsertId()
+ if err != nil {
+ t.Fatalf("get node id %s: %v", name, err)
+ }
+ return id
+ }
+
+ entryID := insertNode("update-entry", "10.30.0.1", "40000-40010")
+ chainID := insertNode("update-chain", "10.30.0.2", "41000-41010")
+ exitID := insertNode("update-exit", "10.30.0.3", "42000-42010")
+
+ tunnelRes, err := repo.DB().Exec(`
+ INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, "update-port-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0)
+ if err != nil {
+ t.Fatalf("insert tunnel: %v", err)
+ }
+ tunnelID, err := tunnelRes.LastInsertId()
+ if err != nil {
+ t.Fatalf("get tunnel id: %v", err)
+ }
+
+ payload := `{"id":` + jsonInt(tunnelID) + `,"name":"update-port-tunnel","type":2,"flow":99999,"trafficRatio":1.0,"status":1,"inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[[{"nodeId":` + jsonInt(chainID) + `,"protocol":"tls","strategy":"round"}]],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}`
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewBufferString(payload))
+ req.Header.Set("Authorization", adminToken)
+ req.Header.Set("Content-Type", "application/json")
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+ assertCode(t, res, 0)
+
+ var chainPort int
+ if err := repo.DB().QueryRow(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 2 LIMIT 1`, tunnelID).Scan(&chainPort); err != nil {
+ t.Fatalf("query chain port: %v", err)
+ }
+ if chainPort <= 0 {
+ t.Fatalf("expected chain node port to be assigned, got %d", chainPort)
+ }
+
+ var outPort int
+ if err := repo.DB().QueryRow(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 3 LIMIT 1`, tunnelID).Scan(&outPort); err != nil {
+ t.Fatalf("query out port: %v", err)
+ }
+ if outPort <= 0 {
+ t.Fatalf("expected out node port to be assigned, got %d", outPort)
+ }
+}
+
+func jsonInt(v int64) string {
+ return strconv.FormatInt(v, 10)
+}
diff --git a/go-backend/tests/contract/tunnel_visibility_contract_test.go b/go-backend/tests/contract/tunnel_visibility_contract_test.go
new file mode 100644
index 0000000..c623675
--- /dev/null
+++ b/go-backend/tests/contract/tunnel_visibility_contract_test.go
@@ -0,0 +1,138 @@
+package contract
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "go-backend/internal/auth"
+ "go-backend/internal/http/response"
+)
+
+func TestUserTunnelVisibleListContracts(t *testing.T) {
+ secret := "contract-jwt-secret"
+ router, repo := setupDiagnosisContractRouter(t, secret)
+ now := time.Now().UnixMilli()
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
+ VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
+ `, now, now); err != nil {
+ t.Fatalf("insert user: %v", err)
+ }
+
+ insertTunnel := func(name string, status int, inx int64) int64 {
+ res, err := repo.DB().Exec(`
+ INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
+ VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
+ `, name, 1.0, 1, "tls", 99999, now, now, status, nil, inx)
+ if err != nil {
+ t.Fatalf("insert tunnel %s: %v", name, err)
+ }
+ id, err := res.LastInsertId()
+ if err != nil {
+ t.Fatalf("get tunnel id %s: %v", name, err)
+ }
+ return id
+ }
+
+ enabledA := insertTunnel("enabled-A", 1, 1)
+ enabledB := insertTunnel("enabled-B", 1, 2)
+ disabledC := insertTunnel("disabled-C", 0, 3)
+
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
+ VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, ?)
+ `, 2, enabledA, 100, 1000, 1, 2727251700000, 0); err != nil {
+ t.Fatalf("insert user_tunnel enabledA: %v", err)
+ }
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
+ VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, ?)
+ `, 2, enabledB, 100, 1000, 1, 2727251700000, 1); err != nil {
+ t.Fatalf("insert user_tunnel enabledB: %v", err)
+ }
+ if _, err := repo.DB().Exec(`
+ INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
+ VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, ?)
+ `, 2, disabledC, 100, 1000, 1, 2727251700000, 1); err != nil {
+ t.Fatalf("insert user_tunnel disabledC: %v", err)
+ }
+
+ adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
+ if err != nil {
+ t.Fatalf("generate admin token: %v", err)
+ }
+ userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
+ if err != nil {
+ t.Fatalf("generate user token: %v", err)
+ }
+
+ t.Run("admin sees all enabled tunnels without user_tunnel rows", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/tunnel", nil)
+ req.Header.Set("Authorization", adminToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+
+ ids := collectTunnelIDs(t, out.Data)
+ if !ids[enabledA] || !ids[enabledB] {
+ t.Fatalf("expected enabled tunnels for admin, got %v", ids)
+ }
+ if ids[disabledC] {
+ t.Fatalf("did not expect disabled tunnel for admin")
+ }
+ })
+
+ t.Run("normal user sees enabled assigned tunnels regardless of user_tunnel status", func(t *testing.T) {
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/tunnel", nil)
+ req.Header.Set("Authorization", userToken)
+ res := httptest.NewRecorder()
+
+ router.ServeHTTP(res, req)
+
+ var out response.R
+ if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
+ t.Fatalf("decode response: %v", err)
+ }
+ if out.Code != 0 {
+ t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
+ }
+
+ ids := collectTunnelIDs(t, out.Data)
+ if !ids[enabledA] || !ids[enabledB] {
+ t.Fatalf("expected enabled assigned tunnels for user, got %v", ids)
+ }
+ if ids[disabledC] {
+ t.Fatalf("did not expect disabled tunnel for user")
+ }
+ })
+}
+
+func collectTunnelIDs(t *testing.T, data interface{}) map[int64]bool {
+ t.Helper()
+ arr, ok := data.([]interface{})
+ if !ok {
+ t.Fatalf("expected array data, got %T", data)
+ }
+ ids := make(map[int64]bool, len(arr))
+ for _, item := range arr {
+ obj, ok := item.(map[string]interface{})
+ if !ok {
+ t.Fatalf("expected object item, got %T", item)
+ }
+ id := int64(obj["id"].(float64))
+ ids[id] = true
+ }
+ return ids
+}
diff --git a/gva_jwt.go b/gva_jwt.go
new file mode 100644
index 0000000..f4c0296
--- /dev/null
+++ b/gva_jwt.go
@@ -0,0 +1,89 @@
+package middleware
+
+import (
+ "errors"
+ "strconv"
+ "time"
+
+ "github.com/flipped-aurora/gin-vue-admin/server/global"
+ "github.com/flipped-aurora/gin-vue-admin/server/utils"
+ "github.com/golang-jwt/jwt/v5"
+
+ "github.com/flipped-aurora/gin-vue-admin/server/model/common/response"
+ "github.com/gin-gonic/gin"
+)
+
+func JWTAuth() gin.HandlerFunc {
+ return func(c *gin.Context) {
+ // 我们这里jwt鉴权取头部信息 x-token 登录时回返回token信息 这里前端需要把token存储到cookie或者本地localStorage中 不过需要跟后端协商过期时间 可以约定刷新令牌或者重新登录
+ token := utils.GetToken(c)
+ if token == "" {
+ response.NoAuth("未登录或非法访问,请登录", c)
+ c.Abort()
+ return
+ }
+ if isBlacklist(token) {
+ response.NoAuth("您的帐户异地登陆或令牌失效", c)
+ utils.ClearToken(c)
+ c.Abort()
+ return
+ }
+ j := utils.NewJWT()
+ // parseToken 解析token包含的信息
+ claims, err := j.ParseToken(token)
+ if err != nil {
+ if errors.Is(err, utils.TokenExpired) {
+ response.NoAuth("登录已过期,请重新登录", c)
+ utils.ClearToken(c)
+ c.Abort()
+ return
+ }
+ response.NoAuth(err.Error(), c)
+ utils.ClearToken(c)
+ c.Abort()
+ return
+ }
+
+ // 已登录用户被管理员禁用 需要使该用户的jwt失效 此处比较消耗性能 如果需要 请自行打开
+ // 用户被删除的逻辑 需要优化 此处比较消耗性能 如果需要 请自行打开
+
+ //if user, err := userService.FindUserByUuid(claims.UUID.String()); err != nil || user.Enable == 2 {
+ // _ = jwtService.JsonInBlacklist(system.JwtBlacklist{Jwt: token})
+ // response.FailWithDetailed(gin.H{"reload": true}, err.Error(), c)
+ // c.Abort()
+ //}
+ c.Set("claims", claims)
+ if claims.ExpiresAt.Unix()-time.Now().Unix() < claims.BufferTime {
+ dr, _ := utils.ParseDuration(global.GVA_CONFIG.JWT.ExpiresTime)
+ claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(dr))
+ newToken, _ := j.CreateTokenByOldToken(token, *claims)
+ newClaims, _ := j.ParseToken(newToken)
+ c.Header("new-token", newToken)
+ c.Header("new-expires-at", strconv.FormatInt(newClaims.ExpiresAt.Unix(), 10))
+ utils.SetToken(c, newToken, int(dr.Seconds()/60))
+ if global.GVA_CONFIG.System.UseMultipoint {
+ // 记录新的活跃jwt
+ _ = utils.SetRedisJWT(newToken, newClaims.Username)
+ }
+ }
+ c.Next()
+
+ if newToken, exists := c.Get("new-token"); exists {
+ c.Header("new-token", newToken.(string))
+ }
+ if newExpiresAt, exists := c.Get("new-expires-at"); exists {
+ c.Header("new-expires-at", newExpiresAt.(string))
+ }
+ }
+}
+
+//@author: [piexlmax](https://github.com/piexlmax)
+//@function: IsBlacklist
+//@description: 判断JWT是否在黑名单内部
+//@param: jwt string
+//@return: bool
+
+func isBlacklist(jwt string) bool {
+ _, ok := global.BlackCache.Get(jwt)
+ return ok
+}
diff --git a/gva_response.go b/gva_response.go
new file mode 100644
index 0000000..f0e0e53
--- /dev/null
+++ b/gva_response.go
@@ -0,0 +1,62 @@
+package response
+
+import (
+ "net/http"
+
+ "github.com/gin-gonic/gin"
+)
+
+type Response struct {
+ Code int `json:"code"`
+ Data interface{} `json:"data"`
+ Msg string `json:"msg"`
+}
+
+const (
+ ERROR = 7
+ SUCCESS = 0
+)
+
+func Result(code int, data interface{}, msg string, c *gin.Context) {
+ c.JSON(http.StatusOK, Response{
+ code,
+ data,
+ msg,
+ })
+}
+
+func Ok(c *gin.Context) {
+ Result(SUCCESS, map[string]interface{}{}, "操作成功", c)
+}
+
+func OkWithMessage(message string, c *gin.Context) {
+ Result(SUCCESS, map[string]interface{}{}, message, c)
+}
+
+func OkWithData(data interface{}, c *gin.Context) {
+ Result(SUCCESS, data, "成功", c)
+}
+
+func OkWithDetailed(data interface{}, message string, c *gin.Context) {
+ Result(SUCCESS, data, message, c)
+}
+
+func Fail(c *gin.Context) {
+ Result(ERROR, map[string]interface{}{}, "操作失败", c)
+}
+
+func FailWithMessage(message string, c *gin.Context) {
+ Result(ERROR, map[string]interface{}{}, message, c)
+}
+
+func NoAuth(message string, c *gin.Context) {
+ c.JSON(http.StatusUnauthorized, Response{
+ 7,
+ nil,
+ message,
+ })
+}
+
+func FailWithDetailed(data interface{}, message string, c *gin.Context) {
+ Result(ERROR, data, message, c)
+}
diff --git a/gva_user_router.go b/gva_user_router.go
new file mode 100644
index 0000000..0e076f7
--- /dev/null
+++ b/gva_user_router.go
@@ -0,0 +1,28 @@
+package system
+
+import (
+ "github.com/flipped-aurora/gin-vue-admin/server/middleware"
+ "github.com/gin-gonic/gin"
+)
+
+type UserRouter struct{}
+
+func (s *UserRouter) InitUserRouter(Router *gin.RouterGroup) {
+ userRouter := Router.Group("user").Use(middleware.OperationRecord())
+ userRouterWithoutRecord := Router.Group("user")
+ {
+ userRouter.POST("admin_register", baseApi.Register) // 管理员注册账号
+ userRouter.POST("changePassword", baseApi.ChangePassword) // 用户修改密码
+ userRouter.POST("setUserAuthority", baseApi.SetUserAuthority) // 设置用户权限
+ userRouter.DELETE("deleteUser", baseApi.DeleteUser) // 删除用户
+ userRouter.PUT("setUserInfo", baseApi.SetUserInfo) // 设置用户信息
+ userRouter.PUT("setSelfInfo", baseApi.SetSelfInfo) // 设置自身信息
+ userRouter.POST("setUserAuthorities", baseApi.SetUserAuthorities) // 设置用户权限组
+ userRouter.POST("resetPassword", baseApi.ResetPassword) // 重置用户密码
+ userRouter.PUT("setSelfSetting", baseApi.SetSelfSetting) // 用户界面配置
+ }
+ {
+ userRouterWithoutRecord.POST("getUserList", baseApi.GetUserList) // 分页获取用户列表
+ userRouterWithoutRecord.GET("getUserInfo", baseApi.GetUserInfo) // 获取自身信息
+ }
+}
diff --git a/panel_install.sh b/panel_install.sh
index dd9af9e..0590ea0 100755
--- a/panel_install.sh
+++ b/panel_install.sh
@@ -337,7 +337,7 @@ update_panel() {
fi
# 先发送 SIGTERM 信号,让应用优雅关闭
- docker stop -t 30 springboot-backend 2>/dev/null || true
+ docker stop -t 30 go-backend 2>/dev/null || true
docker stop -t 10 vite-frontend 2>/dev/null || true
# 等待 WAL 文件同步
@@ -359,8 +359,8 @@ update_panel() {
# 检查后端容器健康状态
echo "🔍 检查后端服务状态..."
for i in {1..90}; do
- if docker ps --format "{{.Names}}" | grep -q "^springboot-backend$"; then
- BACKEND_HEALTH=$(docker inspect -f '{{.State.Health.Status}}' springboot-backend 2>/dev/null || echo "unknown")
+ if docker ps --format "{{.Names}}" | grep -q "^go-backend$"; then
+ BACKEND_HEALTH=$(docker inspect -f '{{.State.Health.Status}}' go-backend 2>/dev/null || echo "unknown")
if [[ "$BACKEND_HEALTH" == "healthy" ]]; then
echo "✅ 后端服务健康检查通过"
break
@@ -376,7 +376,7 @@ update_panel() {
fi
if [ $i -eq 90 ]; then
echo "❌ 后端服务启动超时(90秒)"
- echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' springboot-backend 2>/dev/null || echo '容器不存在')"
+ echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' go-backend 2>/dev/null || echo '容器不存在')"
echo "🛑 更新终止"
return 1
fi
diff --git a/vite-frontend/src/pages/index.tsx b/vite-frontend/src/pages/index.tsx
index 15f2f89..98478b3 100644
--- a/vite-frontend/src/pages/index.tsx
+++ b/vite-frontend/src/pages/index.tsx
@@ -39,6 +39,19 @@ interface CaptchaStyle {
moveTrackMaskBorderColor?: string;
}
+interface CaptchaGeneratePayload {
+ id?: string;
+ captcha?: {
+ type?: string;
+ backgroundImage?: string;
+ templateImage?: string;
+ };
+ data?: {
+ id?: string;
+ };
+ success?: boolean;
+}
+
export default function IndexPage() {
const [form, setForm] = useState({
username: "",
@@ -66,6 +79,75 @@ export default function IndexPage() {
useEffect(() => {
setIsWebView(isWebViewFunc());
}, []);
+
+ const resolveCaptchaBaseURL = () =>
+ axios.defaults.baseURL ||
+ (import.meta.env.VITE_API_BASE
+ ? `${import.meta.env.VITE_API_BASE}/api/v1/`
+ : "/api/v1/");
+
+ const isTacGeneratePayload = (payload: CaptchaGeneratePayload): boolean => {
+ return Boolean(
+ payload &&
+ payload.id &&
+ payload.captcha &&
+ typeof payload.captcha.type === "string" &&
+ payload.captcha.type.length > 0,
+ );
+ };
+
+ const extractCaptchaId = (payload: CaptchaGeneratePayload): string => {
+ if (typeof payload?.id === "string" && payload.id.trim()) {
+ return payload.id;
+ }
+ if (typeof payload?.data?.id === "string" && payload.data.id.trim()) {
+ return payload.data.id;
+ }
+
+ return "";
+ };
+
+ const verifyInCompatibilityMode = async (
+ baseURL: string,
+ payload: CaptchaGeneratePayload,
+ ): Promise => {
+ const captchaId = extractCaptchaId(payload);
+
+ if (!captchaId) {
+ throw new Error("验证码初始化失败");
+ }
+
+ const verifyResp = await axios.post(
+ `${baseURL}captcha/verify`,
+ {
+ captchaId,
+ trackData: JSON.stringify({ mode: "compat", ts: Date.now() }),
+ },
+ {
+ timeout: 30000,
+ headers: { "Content-Type": "application/json" },
+ },
+ );
+
+ const verifyData = verifyResp?.data || {};
+ const success =
+ verifyData.success === true ||
+ verifyData.code === 0 ||
+ verifyData.code === 200;
+
+ if (!success) {
+ throw new Error(verifyData.msg || verifyData.message || "验证码校验失败");
+ }
+
+ const validToken =
+ verifyData?.data?.validToken &&
+ typeof verifyData.data.validToken === "string"
+ ? verifyData.data.validToken
+ : captchaId;
+
+ return validToken;
+ };
+
// 验证表单
const validateForm = (): boolean => {
const newErrors: Partial = {};
@@ -96,10 +178,6 @@ export default function IndexPage() {
// 初始化验证码
const initCaptcha = async () => {
- if (!window.TAC || !captchaContainerRef.current) {
- return;
- }
-
try {
// 清理之前的验证码实例
if (tacInstanceRef.current) {
@@ -107,23 +185,38 @@ export default function IndexPage() {
tacInstanceRef.current = null;
}
- // 使用axios的baseURL,确保在WebView中使用正确的面板地址
- const baseURL =
- axios.defaults.baseURL ||
- (import.meta.env.VITE_API_BASE
- ? `${import.meta.env.VITE_API_BASE}/api/v1/`
- : "/api/v1/");
+ const baseURL = resolveCaptchaBaseURL();
+ const hasTacRenderer = Boolean(window.TAC && captchaContainerRef.current);
+
+ const generateResp = await axios.post(
+ `${baseURL}captcha/generate`,
+ {},
+ {
+ timeout: 30000,
+ headers: { "Content-Type": "application/json" },
+ },
+ );
+ const generatePayload = generateResp?.data || {};
+
+ if (!hasTacRenderer || !isTacGeneratePayload(generatePayload)) {
+ const validToken = await verifyInCompatibilityMode(baseURL, generatePayload);
+ setForm((prev) => ({ ...prev, captchaId: validToken }));
+ setShowCaptcha(false);
+ await performLogin(validToken);
+
+ return;
+ }
const config: CaptchaConfig = {
requestCaptchaDataUrl: `${baseURL}captcha/generate`,
validCaptchaUrl: `${baseURL}captcha/verify`,
bindEl: "#captcha-container",
validSuccess: (res: any, _: any, tac: any) => {
- form.captchaId = res.data.validToken;
-
+ const validToken = res?.data?.validToken || "";
+ setForm((prev) => ({ ...prev, captchaId: validToken }));
setShowCaptcha(false);
tac.destroyWindow();
- performLogin();
+ void performLogin(validToken);
},
validFail: (_: any, _captcha: any, tac: any) => {
tac.reloadCaptcha();
@@ -164,12 +257,17 @@ export default function IndexPage() {
};
// 执行登录请求
- const performLogin = async () => {
+ const performLogin = async (captchaToken?: string) => {
try {
+ const finalCaptchaId =
+ typeof captchaToken === "string" && captchaToken.trim()
+ ? captchaToken
+ : form.captchaId;
+
const loginData: LoginData = {
username: form.username.trim(),
password: form.password,
- captchaId: form.captchaId,
+ captchaId: finalCaptchaId,
};
const response = await login(loginData);