From b2407c3442bbee3283d7a01ae35b751a94e8d9ec Mon Sep 17 00:00:00 2001 From: sagit Date: Sat, 7 Feb 2026 11:03:12 +0000 Subject: [PATCH] feat: finalize Go backend migration and deployment cutover --- .github/workflows/ci-build.yml | 20 +- .github/workflows/docker-build.yml | 47 +- 3x_base.go | 42 ++ 3x_inbound.go | 424 ++++++++++++++++++ 3x_xui.go | 54 +++ docker-compose-v4.yml | 37 +- docker-compose-v6.yml | 37 +- go-backend/internal/app/app.go | 9 +- .../internal/http/handler/captcha_state.go | 67 +++ .../internal/http/handler/control_plane.go | 20 +- .../internal/http/handler/flow_policy.go | 338 ++++++++++++++ go-backend/internal/http/handler/handler.go | 41 +- go-backend/internal/http/handler/jobs.go | 270 +++++++++++ go-backend/internal/http/handler/jobs_test.go | 128 ++++++ go-backend/internal/http/handler/mutations.go | 48 +- .../internal/store/sqlite/repository.go | 123 ++++- .../tests/contract/migration_contract_test.go | 61 +++ gva_jwt.go | 89 ++++ gva_response.go | 62 +++ gva_user_router.go | 28 ++ panel_install.sh | 8 +- 21 files changed, 1813 insertions(+), 140 deletions(-) create mode 100644 3x_base.go create mode 100644 3x_inbound.go create mode 100644 3x_xui.go create mode 100644 go-backend/internal/http/handler/captcha_state.go create mode 100644 go-backend/internal/http/handler/flow_policy.go create mode 100644 go-backend/internal/http/handler/jobs.go create mode 100644 go-backend/internal/http/handler/jobs_test.go create mode 100644 gva_jwt.go create mode 100644 gva_response.go create mode 100644 gva_user_router.go 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/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 f234503..8bf8af2 100644 --- a/docker-compose-v4.yml +++ b/docker-compose-v4.yml @@ -1,37 +1,6 @@ services: backend: - image: ghcr.io/sagit-chu/springboot-backend:2.0.7-beta - container_name: springboot-backend - restart: unless-stopped - logging: - driver: json-file - options: - max-size: "20m" - environment: - DB_PATH: /app/data/gost.db - JWT_SECRET: ${JWT_SECRET} - LOG_DIR: /app/logs - JAVA_OPTS: "-Xms256m -Xmx512m -Dfile.encoding=UTF-8 -Duser.timezone=Asia/Shanghai" - ports: - - "${BACKEND_PORT}:6365" - volumes: - - backend_logs:/app/logs - - sqlite_data:/app/data - networks: - - gost-network - stop_grace_period: 30s - stop_signal: SIGTERM - healthcheck: - test: ["CMD", "sh", "-c", "wget --no-verbose --tries=1 --spider http://localhost:6365/flow/test || exit 1"] - interval: 30s - timeout: 10s - retries: 5 - start_period: 60s - - backend-go: - profiles: ["go-backend"] - build: - context: ./go-backend + image: ghcr.io/sagit-chu/go-backend:2.0.7-beta container_name: go-backend restart: unless-stopped logging: @@ -44,12 +13,14 @@ services: LOG_DIR: /app/logs SERVER_ADDR: :6365 ports: - - "${GO_BACKEND_PORT:-6366}:6365" + - "${BACKEND_PORT}:6365" volumes: - backend_logs:/app/logs - sqlite_data:/app/data networks: - gost-network + stop_grace_period: 30s + stop_signal: SIGTERM healthcheck: test: ["CMD", "sh", "-c", "wget --no-verbose --tries=1 --spider http://localhost:6365/flow/test || exit 1"] interval: 30s diff --git a/docker-compose-v6.yml b/docker-compose-v6.yml index fc8795d..6b8b2a7 100644 --- a/docker-compose-v6.yml +++ b/docker-compose-v6.yml @@ -1,37 +1,6 @@ services: backend: - image: ghcr.io/sagit-chu/springboot-backend:2.0.7-beta - container_name: springboot-backend - restart: unless-stopped - logging: - driver: json-file - options: - max-size: "20m" - environment: - DB_PATH: /app/data/gost.db - JWT_SECRET: ${JWT_SECRET} - LOG_DIR: /app/logs - JAVA_OPTS: "-Xms256m -Xmx512m -Dfile.encoding=UTF-8 -Duser.timezone=Asia/Shanghai" - ports: - - "${BACKEND_PORT}:6365" - volumes: - - backend_logs:/app/logs - - sqlite_data:/app/data - networks: - - gost-network - stop_grace_period: 30s - stop_signal: SIGTERM - healthcheck: - test: ["CMD", "sh", "-c", "wget --no-verbose --tries=1 --spider http://localhost:6365/flow/test || exit 1"] - interval: 30s - timeout: 10s - retries: 5 - start_period: 60s - - backend-go: - profiles: ["go-backend"] - build: - context: ./go-backend + image: ghcr.io/sagit-chu/go-backend:2.0.7-beta container_name: go-backend restart: unless-stopped logging: @@ -44,12 +13,14 @@ services: LOG_DIR: /app/logs SERVER_ADDR: :6365 ports: - - "${GO_BACKEND_PORT:-6366}:6365" + - "${BACKEND_PORT}:6365" volumes: - backend_logs:/app/logs - sqlite_data:/app/data networks: - gost-network + stop_grace_period: 30s + stop_signal: SIGTERM healthcheck: test: ["CMD", "sh", "-c", "wget --no-verbose --tries=1 --spider http://localhost:6365/flow/test || exit 1"] interval: 30s diff --git a/go-backend/internal/app/app.go b/go-backend/internal/app/app.go index 42be18e..425b05a 100644 --- a/go-backend/internal/app/app.go +++ b/go-backend/internal/app/app.go @@ -16,6 +16,7 @@ type App struct { cfg config.Config server *http.Server repo *sqlite.Repository + h *handler.Handler } func New(cfg config.Config) (*App, error) { @@ -36,14 +37,20 @@ func New(cfg config.Config) (*App, error) { IdleTimeout: 60 * time.Second, } - return &App{cfg: cfg, server: s, repo: repo}, nil + 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 { 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 index 5830c14..7bf1f6a 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -27,9 +27,11 @@ type forwardRecord struct { } type tunnelRecord struct { - ID int64 - Type int - Status int + ID int64 + Type int + Status int + Flow int64 + TrafficRatio float64 } type forwardPortRecord struct { @@ -111,15 +113,21 @@ func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) { } func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) { - row := h.repo.DB().QueryRow(`SELECT id, type, status FROM tunnel WHERE id = ? LIMIT 1`, tunnelID) + 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) + 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 } @@ -287,7 +295,7 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str } base := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID) payload := map[string]interface{}{ - "services": []string{base + "_tcp", base + "_udp"}, + "services": []string{base, base + "_tcp", base + "_udp"}, } seen := map[int64]struct{}{} for _, fp := range ports { 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 index 391dd0c..b4cc0a0 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -1,6 +1,7 @@ package handler import ( + "context" "database/sql" "encoding/json" "fmt" @@ -9,6 +10,7 @@ import ( "sort" "strconv" "strings" + "sync" "time" "go-backend/internal/auth" @@ -23,6 +25,14 @@ 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 { @@ -54,7 +64,12 @@ type flowItem struct { } func New(repo *sqlite.Repository, jwtSecret string) *Handler { - return &Handler{repo: repo, jwtSecret: jwtSecret, wsServer: ws.NewServer(repo, jwtSecret)} + return &Handler{ + repo: repo, + jwtSecret: jwtSecret, + wsServer: ws.NewServer(repo, jwtSecret), + captchaTokens: make(map[string]int64), + } } func (h *Handler) WebSocketHandler() http.Handler { @@ -170,6 +185,10 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) { 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 { @@ -558,13 +577,17 @@ func (h *Handler) flowTest(w http.ResponseWriter, _ *http.Request) { func (h *Handler) flowConfig(w http.ResponseWriter, r *http.Request) { secret := r.URL.Query().Get("secret") - if ok, _ := h.repo.NodeExistsBySecret(secret); !ok { + 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 } - _, _ = readAndDecryptFlowBody(r.Body, secret) + 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")) } @@ -582,17 +605,7 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) { var items []flowItem if json.Unmarshal([]byte(raw), &items) == nil { for _, item := range items { - parts := strings.Split(item.N, "_") - if len(parts) < 3 || item.N == "web_api" { - continue - } - forwardID, err1 := strconv.ParseInt(parts[0], 10, 64) - userID, err2 := strconv.ParseInt(parts[1], 10, 64) - userTunnelID, err3 := strconv.ParseInt(parts[2], 10, 64) - if err1 != nil || err2 != nil || err3 != nil { - continue - } - _ = h.repo.AddFlow(forwardID, userID, userTunnelID, item.D, item.U) + h.processFlowItem(item) } } } 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 index 92dcfa1..84f85da 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -56,7 +56,7 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) { num := asInt(req["num"], 10) expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()) flowResetTime := asInt64(req["flowResetTime"], 1) - roleID := asInt(req["roleId"], asInt(req["role_id"], 1)) + roleID := 1 now := time.Now().UnixMilli() _, err := db.Exec(` @@ -97,6 +97,20 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) { 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())) @@ -151,6 +165,20 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) { 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 { @@ -167,6 +195,10 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) { 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 @@ -175,6 +207,10 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) { 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 @@ -248,6 +284,16 @@ func (h *Handler) captchaVerify(w http.ResponseWriter, r *http.Request) { 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}, diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index 8fd5806..55baf77 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -300,14 +300,10 @@ func (r *Repository) GetUserPackageForwards(userID int64) ([]UserForwardDetail, } rows, err := r.db.Query(` - SELECT f.id, f.name, f.tunnel_id, t.name, f.remote_addr, f.in_flow, f.out_flow, f.status, f.created_time, - GROUP_CONCAT(n.server_ip || ':' || fp.port), MIN(fp.port) + 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 - LEFT JOIN forward_port fp ON fp.forward_id = f.id - LEFT JOIN node n ON n.id = fp.node_id WHERE f.user_id = ? - GROUP BY f.id ORDER BY f.id ASC `, userID) if err != nil { @@ -320,10 +316,18 @@ func (r *Repository) GetUserPackageForwards(userID int64) ([]UserForwardDetail, var item UserForwardDetail if err := rows.Scan( &item.ID, &item.Name, &item.TunnelID, &item.TunnelName, &item.RemoteAddr, - &item.InFlow, &item.OutFlow, &item.Status, &item.CreatedAt, &item.InIP, &item.InPort, + &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) } @@ -503,6 +507,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) { 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 { @@ -595,14 +600,9 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) { 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, - GROUP_CONCAT(CASE WHEN n.server_ip IS NOT NULL AND fp.port IS NOT NULL THEN n.server_ip || ':' || fp.port END), - MIN(fp.port) + 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 - LEFT JOIN forward_port fp ON fp.forward_id = f.id - LEFT JOIN node n ON n.id = fp.node_id - GROUP BY f.id ORDER BY f.inx ASC, f.id ASC `) if err != nil { @@ -615,10 +615,13 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) { var id, userID, tunnelID, inFlow, outFlow, createdTime, inx int64 var userName, name, tunnelName, remoteAddr, strategy string var status int - var inIP sql.NullString - var inPort sql.NullInt64 - if err := rows.Scan(&id, &userID, &userName, &name, &tunnelID, &tunnelName, &remoteAddr, &strategy, &inFlow, &outFlow, &createdTime, &status, &inx, &inIP, &inPort); err != nil { + 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 } @@ -629,7 +632,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) { "name": name, "tunnelId": tunnelID, "tunnelName": tunnelName, - "inIp": nullableString(inIP), + "inIp": nullableForwardIngress(inIP), "inPort": nullableInt64(inPort), "remoteAddr": remoteAddr, "strategy": strategy, @@ -1017,6 +1020,94 @@ func nullableString(v sql.NullString) interface{} { 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 diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go index becb59e..76ec67e 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -1,6 +1,7 @@ package contract_test import ( + "bytes" "encoding/json" "io" "net/http" @@ -18,6 +19,66 @@ import ( "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") 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 21f8ade..da082bb 100755 --- a/panel_install.sh +++ b/panel_install.sh @@ -301,7 +301,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 文件同步 @@ -323,8 +323,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 @@ -340,7 +340,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