mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 23:56:36 +08:00
Compare commits
12 Commits
2.0.16
...
2.1.0-beta
| Author | SHA1 | Date | |
|---|---|---|---|
| ae6c65e9a9 | |||
| 8a85ab2844 | |||
| b9b4312768 | |||
| a2b819dbcc | |||
| 6f0412de0f | |||
| e2ae241f8c | |||
| 1a8f424d53 | |||
| 1977ba9eab | |||
| b2407c3442 | |||
| f9bc165c16 | |||
| b6333d81a2 | |||
| 0f37017760 |
@@ -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
|
||||
|
||||
@@ -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}/flux-panel-backend:latest \
|
||||
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-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: .*flux-panel-backend:[^[:space:]]*|image: ${{ env.REGISTRY }}/${OWNER}/flux-panel-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: .*flux-panel-backend:[^[:space:]]*|image: ${{ env.REGISTRY }}/${OWNER}/flux-panel-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}/flux-panel-backend:${VERSION}
|
||||
|
||||
# Frontend
|
||||
docker pull ${{ env.REGISTRY }}/${OWNER}/vite-frontend:${VERSION}
|
||||
|
||||
+4
-1
@@ -257,4 +257,7 @@ gitee/
|
||||
doraemon.jks
|
||||
device.id
|
||||
commit.sh
|
||||
sql/
|
||||
sql/
|
||||
!go-backend/internal/store/sqlite/sql/
|
||||
!go-backend/internal/store/sqlite/sql/schema.sql
|
||||
!go-backend/internal/store/sqlite/sql/data.sql
|
||||
|
||||
+42
@@ -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
|
||||
}
|
||||
+424
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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/flux-panel-backend:${FLUX_VERSION:-latest}
|
||||
container_name: flux-panel-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}
|
||||
|
||||
@@ -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/flux-panel-backend:${FLUX_VERSION:-latest}
|
||||
container_name: flux-panel-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}
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
FROM golang:1.23-bookworm AS builder
|
||||
WORKDIR /src
|
||||
|
||||
COPY go.mod ./
|
||||
RUN go mod download
|
||||
|
||||
COPY . .
|
||||
RUN CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -o /out/paneld ./cmd/paneld
|
||||
|
||||
FROM debian:bookworm-slim
|
||||
WORKDIR /app
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates wget && rm -rf /var/lib/apt/lists/*
|
||||
COPY --from=builder /out/paneld /app/paneld
|
||||
|
||||
ENV SERVER_ADDR=:6365
|
||||
EXPOSE 6365
|
||||
ENTRYPOINT ["/app/paneld"]
|
||||
@@ -0,0 +1,12 @@
|
||||
GO ?= go
|
||||
|
||||
.PHONY: test build run
|
||||
|
||||
test:
|
||||
$(GO) test ./...
|
||||
|
||||
build:
|
||||
$(GO) build ./cmd/paneld
|
||||
|
||||
run:
|
||||
SERVER_ADDR=:6365 $(GO) run ./cmd/paneld
|
||||
@@ -0,0 +1,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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
module go-backend
|
||||
|
||||
go 1.23.0
|
||||
|
||||
toolchain go1.24.4
|
||||
|
||||
require (
|
||||
github.com/gorilla/websocket v1.5.3
|
||||
modernc.org/sqlite v1.37.1
|
||||
)
|
||||
|
||||
require (
|
||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||
github.com/google/uuid v1.6.0 // indirect
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/ncruces/go-strftime v0.1.9 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect
|
||||
golang.org/x/sys v0.33.0 // indirect
|
||||
modernc.org/libc v1.65.7 // indirect
|
||||
modernc.org/mathutil v1.7.1 // indirect
|
||||
modernc.org/memory v1.11.0 // indirect
|
||||
)
|
||||
@@ -0,0 +1,49 @@
|
||||
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
|
||||
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
|
||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||
github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4=
|
||||
github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM=
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8=
|
||||
golang.org/x/mod v0.24.0 h1:ZfthKaKaT4NrhGVZHO1/WDTwGES4De8KtWO0SIbNJMU=
|
||||
golang.org/x/mod v0.24.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww=
|
||||
golang.org/x/sync v0.14.0 h1:woo0S4Yywslg6hp4eUFjTVOyKt0RookbpAHG4c1HmhQ=
|
||||
golang.org/x/sync v0.14.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
|
||||
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
|
||||
golang.org/x/tools v0.33.0 h1:4qz2S3zmRxbGIhDIAgjxvFutSvH5EfnsYrRBj0UI0bc=
|
||||
golang.org/x/tools v0.33.0/go.mod h1:CIJMaWEY88juyUfo7UbgPqbC8rU2OqfAV1h2Qp0oMYI=
|
||||
modernc.org/cc/v4 v4.26.1 h1:+X5NtzVBn0KgsBCBe+xkDC7twLb/jNVj9FPgiwSQO3s=
|
||||
modernc.org/cc/v4 v4.26.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
||||
modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU=
|
||||
modernc.org/ccgo/v4 v4.28.0/go.mod h1:JygV3+9AV6SmPhDasu4JgquwU81XAKLd3OKTUDNOiKE=
|
||||
modernc.org/fileutil v1.3.1 h1:8vq5fe7jdtEvoCf3Zf9Nm0Q05sH6kGx0Op2CPx1wTC8=
|
||||
modernc.org/fileutil v1.3.1/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
|
||||
modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
|
||||
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
||||
modernc.org/libc v1.65.7 h1:Ia9Z4yzZtWNtUIuiPuQ7Qf7kxYrxP1/jeHZzG8bFu00=
|
||||
modernc.org/libc v1.65.7/go.mod h1:011EQibzzio/VX3ygj1qGFt5kMjP0lHb0qCW5/D/pQU=
|
||||
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
||||
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
||||
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
||||
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
||||
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
|
||||
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
||||
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
|
||||
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
||||
modernc.org/sqlite v1.37.1 h1:EgHJK/FPoqC+q2YBXg7fUmES37pCHFc97sI7zSayBEs=
|
||||
modernc.org/sqlite v1.37.1/go.mod h1:XwdRtsE1MpiBcL54+MbKcaDvcuej+IYSMfLN6gSKV8g=
|
||||
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
|
||||
modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
|
||||
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
|
||||
modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM=
|
||||
@@ -0,0 +1,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
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
algorithm = "HmacSHA256"
|
||||
expireTime = 90 * 24 * time.Hour
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
Sub string `json:"sub"`
|
||||
Iat int64 `json:"iat"`
|
||||
Exp int64 `json:"exp"`
|
||||
User string `json:"user"`
|
||||
Name string `json:"name"`
|
||||
RoleID int `json:"role_id"`
|
||||
}
|
||||
|
||||
type tokenHeader struct {
|
||||
Alg string `json:"alg"`
|
||||
Typ string `json:"typ"`
|
||||
}
|
||||
|
||||
func GenerateToken(userID int64, username string, roleID int, secret string) (string, error) {
|
||||
now := time.Now()
|
||||
header := tokenHeader{Alg: algorithm, Typ: "JWT"}
|
||||
claims := Claims{
|
||||
Sub: strconv.FormatInt(userID, 10),
|
||||
Iat: now.Unix(),
|
||||
Exp: now.Add(expireTime).Unix(),
|
||||
User: username,
|
||||
Name: username,
|
||||
RoleID: roleID,
|
||||
}
|
||||
|
||||
headerPart, err := encodeJSON(header)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
payloadPart, err := encodeJSON(claims)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
sig := sign(headerPart+"."+payloadPart, secret)
|
||||
|
||||
return headerPart + "." + payloadPart + "." + sig, nil
|
||||
}
|
||||
|
||||
func ValidateToken(token, secret string) (Claims, bool) {
|
||||
claims, err := ParseClaims(token, secret)
|
||||
if err != nil {
|
||||
return Claims{}, false
|
||||
}
|
||||
return claims, true
|
||||
}
|
||||
|
||||
func ParseClaims(token, secret string) (Claims, error) {
|
||||
parts := splitToken(token)
|
||||
if len(parts) != 3 {
|
||||
return Claims{}, errors.New("invalid token")
|
||||
}
|
||||
|
||||
signedContent := parts[0] + "." + parts[1]
|
||||
expected := sign(signedContent, secret)
|
||||
if !hmac.Equal([]byte(expected), []byte(parts[2])) {
|
||||
return Claims{}, errors.New("invalid signature")
|
||||
}
|
||||
|
||||
payloadBytes, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil {
|
||||
return Claims{}, err
|
||||
}
|
||||
|
||||
var claims Claims
|
||||
if err := json.Unmarshal(payloadBytes, &claims); err != nil {
|
||||
return Claims{}, err
|
||||
}
|
||||
|
||||
if claims.Exp <= time.Now().Unix() {
|
||||
return Claims{}, errors.New("token expired")
|
||||
}
|
||||
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
func splitToken(token string) []string {
|
||||
parts := make([]string, 0, 3)
|
||||
current := ""
|
||||
for i := 0; i < len(token); i++ {
|
||||
if token[i] == '.' {
|
||||
parts = append(parts, current)
|
||||
current = ""
|
||||
continue
|
||||
}
|
||||
current += string(token[i])
|
||||
}
|
||||
parts = append(parts, current)
|
||||
return parts
|
||||
}
|
||||
|
||||
func encodeJSON(v interface{}) (string, error) {
|
||||
raw, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(raw), nil
|
||||
}
|
||||
|
||||
func sign(content, secret string) string {
|
||||
h := hmac.New(sha256.New, []byte(secret))
|
||||
h.Write([]byte(content))
|
||||
return base64.RawURLEncoding.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package config
|
||||
|
||||
import "os"
|
||||
|
||||
type Config struct {
|
||||
Addr string
|
||||
DBPath string
|
||||
JWTSecret string
|
||||
LogDir string
|
||||
}
|
||||
|
||||
func FromEnv() Config {
|
||||
cfg := Config{
|
||||
Addr: getEnv("SERVER_ADDR", ":6365"),
|
||||
DBPath: getEnv("DB_PATH", "/app/data/gost.db"),
|
||||
JWTSecret: getEnv("JWT_SECRET", ""),
|
||||
LogDir: getEnv("LOG_DIR", "/app/logs"),
|
||||
}
|
||||
|
||||
return cfg
|
||||
}
|
||||
|
||||
func getEnv(key, fallback string) string {
|
||||
if v := os.Getenv(key); v != "" {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
@@ -0,0 +1,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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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("<!DOCTYPE html><html lang='zh-CN'><head><meta charset='UTF-8'><meta name='viewport' content='width=device-width, initial-scale=1.0'><title>错误 404</title></head><body><div style='min-height:100vh;display:flex;align-items:center;justify-content:center;flex-direction:column;font-family:-apple-system,BlinkMacSystemFont,Segoe UI,Arial,sans-serif;'><div style='font-size:6rem;color:#333;font-weight:300;'>404</div><div style='font-size:1.2rem;color:#666;'>你推开了后端的大门,却发现里面只有寂寞。</div></div></body></html>"))
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,117 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
type contextKey string
|
||||
|
||||
const ClaimsContextKey contextKey = "claims"
|
||||
|
||||
type AuthOptions struct {
|
||||
JWTSecret string
|
||||
}
|
||||
|
||||
func JWT(opts AuthOptions) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if shouldSkip(r.URL.Path) {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
if !strings.HasPrefix(r.URL.Path, "/api/") {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
token := strings.TrimSpace(r.Header.Get("Authorization"))
|
||||
if token == "" {
|
||||
response.WriteJSON(w, response.Err(401, "未登录或token已过期"))
|
||||
return
|
||||
}
|
||||
|
||||
claims, ok := auth.ValidateToken(token, opts.JWTSecret)
|
||||
if !ok {
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
|
||||
if requiresAdmin(r.URL.Path) && claims.RoleID != 0 {
|
||||
response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作"))
|
||||
return
|
||||
}
|
||||
|
||||
ctx := context.WithValue(r.Context(), ClaimsContextKey, claims)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func RequireAdmin(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
raw := r.Context().Value(ClaimsContextKey)
|
||||
claims, ok := raw.(auth.Claims)
|
||||
if !ok {
|
||||
response.WriteJSON(w, response.Err(401, "无法获取用户权限信息"))
|
||||
return
|
||||
}
|
||||
if claims.RoleID != 0 {
|
||||
response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作"))
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func shouldSkip(path string) bool {
|
||||
switch {
|
||||
case strings.HasPrefix(path, "/flow/"):
|
||||
return true
|
||||
case strings.HasPrefix(path, "/api/v1/open_api/"):
|
||||
return true
|
||||
case strings.HasPrefix(path, "/api/v1/captcha/"):
|
||||
return true
|
||||
case path == "/api/v1/config/get":
|
||||
return true
|
||||
case path == "/api/v1/user/login":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func requiresAdmin(path string) bool {
|
||||
if strings.HasPrefix(path, "/api/v1/group/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/node/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/speed-limit/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/tunnel/") {
|
||||
if strings.HasPrefix(path, "/api/v1/tunnel/user/tunnel") {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
switch path {
|
||||
case "/api/v1/user/create", "/api/v1/user/list", "/api/v1/user/update", "/api/v1/user/delete", "/api/v1/user/reset":
|
||||
return true
|
||||
case "/api/v1/config/update", "/api/v1/config/update-single":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package middleware
|
||||
|
||||
import "net/http"
|
||||
|
||||
func CORS(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
w.Header().Set("Access-Control-Allow-Headers", "*")
|
||||
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, DELETE, PUT, OPTIONS")
|
||||
w.Header().Set("Access-Control-Expose-Headers", "Authorization")
|
||||
if r.Method == http.MethodOptions {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func Recover(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprint(rec)))
|
||||
}
|
||||
}()
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,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())
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package response
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
type R struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
TS int64 `json:"ts"`
|
||||
Data interface{} `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
func OK(data interface{}) R {
|
||||
return R{
|
||||
Code: 0,
|
||||
Msg: "操作成功",
|
||||
TS: time.Now().UnixMilli(),
|
||||
Data: data,
|
||||
}
|
||||
}
|
||||
|
||||
func OKEmpty() R {
|
||||
return R{
|
||||
Code: 0,
|
||||
Msg: "操作成功",
|
||||
TS: time.Now().UnixMilli(),
|
||||
}
|
||||
}
|
||||
|
||||
func Err(code int, msg string) R {
|
||||
return R{
|
||||
Code: code,
|
||||
Msg: msg,
|
||||
TS: time.Now().UnixMilli(),
|
||||
}
|
||||
}
|
||||
|
||||
func ErrDefault(msg string) R {
|
||||
return Err(-1, msg)
|
||||
}
|
||||
|
||||
func WriteJSON(w http.ResponseWriter, payload R) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
_ = json.NewEncoder(w).Encode(payload)
|
||||
}
|
||||
@@ -0,0 +1,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
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
type AESCrypto struct {
|
||||
key []byte
|
||||
}
|
||||
|
||||
func NewAESCrypto(secret string) (*AESCrypto, error) {
|
||||
if secret == "" {
|
||||
return nil, fmt.Errorf("secret is empty")
|
||||
}
|
||||
hash := sha256.Sum256([]byte(secret))
|
||||
return &AESCrypto{key: hash[:]}, nil
|
||||
}
|
||||
|
||||
func (a *AESCrypto) Encrypt(plain []byte) (string, error) {
|
||||
if len(plain) == 0 {
|
||||
return "", fmt.Errorf("empty plaintext")
|
||||
}
|
||||
block, err := aes.NewCipher(a.key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
nonce := make([]byte, gcm.NonceSize())
|
||||
if _, err := rand.Read(nonce); err != nil {
|
||||
return "", err
|
||||
}
|
||||
sealed := gcm.Seal(nil, nonce, plain, nil)
|
||||
data := append(nonce, sealed...)
|
||||
return base64.StdEncoding.EncodeToString(data), nil
|
||||
}
|
||||
|
||||
func (a *AESCrypto) Decrypt(cipherText string) ([]byte, error) {
|
||||
raw, err := base64.StdEncoding.DecodeString(cipherText)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
block, err := aes.NewCipher(a.key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nonceSize := gcm.NonceSize()
|
||||
if len(raw) < nonceSize {
|
||||
return nil, fmt.Errorf("ciphertext too short")
|
||||
}
|
||||
nonce := raw[:nonceSize]
|
||||
data := raw[nonceSize:]
|
||||
return gcm.Open(nil, nonce, data, nil)
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"crypto/md5"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
func MD5(input string) string {
|
||||
hash := md5.Sum([]byte(input))
|
||||
return fmt.Sprintf("%x", hash)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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);
|
||||
@@ -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
|
||||
);
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
|
||||
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
})
|
||||
|
||||
wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret})(next)
|
||||
|
||||
t.Run("login path is excluded", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", nil)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("missing token returns 401 contract message", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCodeMsg(t, res, 401, "未登录或token已过期")
|
||||
})
|
||||
|
||||
t.Run("invalid token returns 401 contract message", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", "invalid.token.value")
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCodeMsg(t, res, 401, "无效的token或token已过期")
|
||||
})
|
||||
|
||||
t.Run("valid token reaches next", func(t *testing.T) {
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("non-admin blocked on admin path", func(t *testing.T) {
|
||||
token, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCodeMsg(t, res, 403, "权限不足,仅管理员可操作")
|
||||
})
|
||||
}
|
||||
|
||||
func assertCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expected {
|
||||
t.Fatalf("expected code %d, got %d", expected, out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCodeMsg(t *testing.T, rec *httptest.ResponseRecorder, expectedCode int, expectedMsg string) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expectedCode || out.Msg != expectedMsg {
|
||||
t.Fatalf("expected (%d,%q), got (%d,%q)", expectedCode, expectedMsg, out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+89
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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) // 获取自身信息
|
||||
}
|
||||
}
|
||||
+4
-4
@@ -337,7 +337,7 @@ update_panel() {
|
||||
fi
|
||||
|
||||
# 先发送 SIGTERM 信号,让应用优雅关闭
|
||||
docker stop -t 30 springboot-backend 2>/dev/null || true
|
||||
docker stop -t 30 flux-panel-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 "^flux-panel-backend$"; then
|
||||
BACKEND_HEALTH=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-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}}' flux-panel-backend 2>/dev/null || echo '容器不存在')"
|
||||
echo "🛑 更新终止"
|
||||
return 1
|
||||
fi
|
||||
|
||||
@@ -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<LoginForm>({
|
||||
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<string> => {
|
||||
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<LoginForm> = {};
|
||||
@@ -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<CaptchaGeneratePayload>(
|
||||
`${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);
|
||||
|
||||
Reference in New Issue
Block a user