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
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);