Files
flvx/go-backend/internal/http/handler/mutations.go
T

2037 lines
58 KiB
Go

package handler
import (
"crypto/rand"
"database/sql"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"net/http"
"sort"
"strconv"
"strings"
"time"
"go-backend/internal/http/response"
"go-backend/internal/security"
)
func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
username := asString(req["user"])
pwd := asString(req["pwd"])
if username == "" || pwd == "" {
response.WriteJSON(w, response.ErrDefault("用户名或密码不能为空"))
return
}
db := h.repo.DB()
if db == nil {
response.WriteJSON(w, response.Err(-2, "database unavailable"))
return
}
var cnt int
if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&cnt); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if cnt > 0 {
response.WriteJSON(w, response.ErrDefault("用户名已存在"))
return
}
status := asInt(req["status"], 1)
flow := asInt64(req["flow"], 100)
num := asInt(req["num"], 10)
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
flowResetTime := asInt64(req["flowResetTime"], 1)
roleID := 1
now := time.Now().UnixMilli()
_, err := db.Exec(`
INSERT INTO user(user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?, ?, ?)
`, username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, now, now, status)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("用户ID不能为空"))
return
}
username := asString(req["user"])
if username == "" {
response.WriteJSON(w, response.ErrDefault("用户名不能为空"))
return
}
db := h.repo.DB()
if db == nil {
response.WriteJSON(w, response.Err(-2, "database unavailable"))
return
}
var roleID int
if err := db.QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil {
if err == sql.ErrNoRows {
response.WriteJSON(w, response.ErrDefault("用户不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if roleID == 0 {
response.WriteJSON(w, response.ErrDefault("请不要作死"))
return
}
var cnt int
if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, id).Scan(&cnt); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if cnt > 0 {
response.WriteJSON(w, response.ErrDefault("用户名已存在"))
return
}
flow := asInt64(req["flow"], 100)
num := asInt(req["num"], 10)
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
flowResetTime := asInt64(req["flowResetTime"], 1)
status := asInt(req["status"], 1)
now := time.Now().UnixMilli()
pwd := asString(req["pwd"])
if strings.TrimSpace(pwd) == "" {
_, err := db.Exec(`
UPDATE user
SET user = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ?
WHERE id = ?
`, username, flow, num, expTime, flowResetTime, status, now, id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
} else {
_, err := db.Exec(`
UPDATE user
SET user = ?, pwd = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ?
WHERE id = ?
`, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, now, id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
}
_, _ = db.Exec(`UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ? WHERE user_id = ?`, flow, num, expTime, flowResetTime, id)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
id := idFromBody(r, w)
if id <= 0 {
return
}
var roleID int
if err := h.repo.DB().QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil {
if err == sql.ErrNoRows {
response.WriteJSON(w, response.ErrDefault("用户不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if roleID == 0 {
response.WriteJSON(w, response.ErrDefault("请不要作死"))
return
}
db := h.repo.DB()
tx, err := db.Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer func() { _ = tx.Rollback() }()
if _, err = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE user_id = ?)`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if _, err = tx.Exec(`DELETE FROM forward WHERE user_id = ?`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if _, err = tx.Exec(`DELETE FROM group_permission_grant WHERE user_tunnel_id IN (SELECT id FROM user_tunnel WHERE user_id = ?)`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if _, err = tx.Exec(`DELETE FROM user_tunnel WHERE user_id = ?`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if _, err = tx.Exec(`DELETE FROM user_group_user WHERE user_id = ?`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if _, err = tx.Exec(`DELETE FROM statistics_flow WHERE user_id = ?`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if _, err = tx.Exec(`DELETE FROM user WHERE id = ?`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err = tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) userResetFlow(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
typeVal := asInt(req["type"], 0)
if id <= 0 || (typeVal != 1 && typeVal != 2) {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
db := h.repo.DB()
if typeVal == 1 {
_, _ = db.Exec(`UPDATE user SET in_flow = 0, out_flow = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id)
_, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE user_id = ?`, id)
} else {
_, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE id = ?`, id)
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) captchaGenerate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
return
}
token := randomToken(16)
payload := map[string]interface{}{
"id": token,
"data": map[string]interface{}{
"id": token,
},
"success": true,
}
w.Header().Set("Content-Type", "application/json; charset=utf-8")
_ = json.NewEncoder(w).Encode(payload)
}
func (h *Handler) captchaVerify(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
return
}
id := asString(req["captchaId"])
if id == "" {
id = asString(req["id"])
}
trackData := asString(req["data"])
if trackData == "" {
trackData = asString(req["trackData"])
}
if id == "" || trackData == "" {
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"success":false,"message":"bad request"}`))
return
}
h.storeCaptchaToken(id)
payload := map[string]interface{}{
"success": true,
"data": map[string]interface{}{"validToken": id},
}
w.Header().Set("Content-Type", "application/json; charset=utf-8")
_ = json.NewEncoder(w).Encode(payload)
}
func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
name := asString(req["name"])
serverIP := asString(req["serverIp"])
if name == "" || serverIP == "" {
response.WriteJSON(w, response.ErrDefault("节点名称和地址不能为空"))
return
}
db := h.repo.DB()
now := time.Now().UnixMilli()
inx := nextIndex(db, "node")
_, err := db.Exec(`
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`,
name,
randomToken(16),
serverIP,
nullableText(asString(req["serverIpV4"])),
nullableText(asString(req["serverIpV6"])),
defaultString(asString(req["port"]), "1000-65535"),
nullableText(asString(req["interfaceName"])),
nullableText(""),
asInt(req["http"], 0),
asInt(req["tls"], 0),
asInt(req["socks"], 0),
now,
now,
0,
defaultString(asString(req["tcpListenAddr"]), "[::]"),
defaultString(asString(req["udpListenAddr"]), "[::]"),
inx,
)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("节点ID不能为空"))
return
}
var currentStatus int
var currentHTTP int
var currentTLS int
var currentSocks int
if err := h.repo.DB().QueryRow(`SELECT status, http, tls, socks FROM node WHERE id = ?`, id).Scan(&currentStatus, &currentHTTP, &currentTLS, &currentSocks); err != nil {
if err == sql.ErrNoRows {
response.WriteJSON(w, response.ErrDefault("节点不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
newHTTP := asInt(req["http"], currentHTTP)
newTLS := asInt(req["tls"], currentTLS)
newSocks := asInt(req["socks"], currentSocks)
if currentStatus == 1 && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
if err := h.applyNodeProtocolChange(id, newHTTP, newTLS, newSocks); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
now := time.Now().UnixMilli()
_, err := h.repo.DB().Exec(`
UPDATE node
SET name = ?, server_ip = ?, server_ip_v4 = ?, server_ip_v6 = ?, port = ?, interface_name = ?, http = ?, tls = ?, socks = ?, tcp_listen_addr = ?, udp_listen_addr = ?, updated_time = ?
WHERE id = ?
`,
asString(req["name"]),
asString(req["serverIp"]),
nullableText(asString(req["serverIpV4"])),
nullableText(asString(req["serverIpV6"])),
defaultString(asString(req["port"]), "1000-65535"),
nullableText(asString(req["interfaceName"])),
newHTTP,
newTLS,
newSocks,
defaultString(asString(req["tcpListenAddr"]), "[::]"),
defaultString(asString(req["udpListenAddr"]), "[::]"),
now,
id,
)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) nodeDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
id := idFromBody(r, w)
if id <= 0 {
return
}
if err := h.deleteNodeByID(id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) nodeInstall(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
id := idFromBody(r, w)
if id <= 0 {
return
}
db := h.repo.DB()
var secret string
if err := db.QueryRow(`SELECT secret FROM node WHERE id = ?`, id).Scan(&secret); err != nil {
response.WriteJSON(w, response.ErrDefault("节点不存在"))
return
}
var panelAddr string
if err := db.QueryRow(`SELECT value FROM vite_config WHERE name = 'ip' LIMIT 1`).Scan(&panelAddr); err != nil {
if err == sql.ErrNoRows {
response.WriteJSON(w, response.ErrDefault("请先前往网站配置中设置ip"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
cmd := fmt.Sprintf("curl -L https://github.com/Sagit-chu/flux-panel/releases/latest/download/install.sh -o ./install.sh && chmod +x ./install.sh && ./install.sh -a %s -s %s", processServerAddress(panelAddr), secret)
response.WriteJSON(w, response.OK(cmd))
}
func (h *Handler) nodeUpdateOrder(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req struct {
Nodes []struct {
ID int64 `json:"id"`
Inx int `json:"inx"`
} `json:"nodes"`
}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
for _, n := range req.Nodes {
_, _ = h.repo.DB().Exec(`UPDATE node SET inx = ?, updated_time = ? WHERE id = ?`, n.Inx, time.Now().UnixMilli(), n.ID)
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) nodeBatchDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
ids := idsFromBody(r, w)
if ids == nil {
return
}
for _, id := range ids {
_ = h.deleteNodeByID(id)
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) nodeCheckStatus(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
items, err := h.repo.ListNodes()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(items))
}
func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
name := asString(req["name"])
if name == "" {
response.WriteJSON(w, response.ErrDefault("隧道名称不能为空"))
return
}
typeVal := asInt(req["type"], 1)
flow := asInt64(req["flow"], 1)
status := asInt(req["status"], 1)
trafficRatio := asFloat(req["trafficRatio"], 1.0)
inIP := asString(req["inIp"])
now := time.Now().UnixMilli()
inx := nextIndex(h.repo.DB(), "tunnel")
tx, err := h.repo.DB().Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer func() { _ = tx.Rollback() }()
res, err := tx.Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
tunnelID, _ := res.LastInsertId()
if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) tunnelGet(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
id := idFromBody(r, w)
if id <= 0 {
return
}
items, err := h.repo.ListTunnels()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
for _, it := range items {
if asInt64(it["id"], 0) == id {
response.WriteJSON(w, response.OK(it))
return
}
}
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
}
func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
return
}
now := time.Now().UnixMilli()
_, err := h.repo.DB().Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, updated_time=? WHERE id=?`,
asString(req["name"]), asInt(req["type"], 1), asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(asString(req["inIp"])), now, id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
tx, err := h.repo.DB().Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer func() { _ = tx.Rollback() }()
if _, err := tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := replaceTunnelChainsTx(tx, id, req); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) tunnelDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
id := idFromBody(r, w)
if id <= 0 {
return
}
if err := h.deleteTunnelByID(id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) tunnelDiagnose(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
id := asInt64FromBodyKey(r, w, "tunnelId")
if id <= 0 {
return
}
result, err := h.diagnoseTunnelRuntime(id)
if err != nil {
if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不完整") {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(result))
}
func (h *Handler) tunnelUpdateOrder(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req struct {
Tunnels []struct {
ID int64 `json:"id"`
Inx int `json:"inx"`
} `json:"tunnels"`
}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
for _, t := range req.Tunnels {
_, _ = h.repo.DB().Exec(`UPDATE tunnel SET inx = ?, updated_time = ? WHERE id = ?`, t.Inx, time.Now().UnixMilli(), t.ID)
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) {
ids := idsFromBody(r, w)
if ids == nil {
return
}
success := 0
fail := 0
for _, id := range ids {
if err := h.deleteTunnelByID(id); err != nil {
fail++
} else {
success++
}
}
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
}
func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
ids := idsFromBody(r, w)
if ids == nil {
return
}
success := 0
fail := 0
for _, tunnelID := range ids {
forwards, err := h.listForwardsByTunnel(tunnelID)
if err != nil {
fail++
continue
}
if len(forwards) == 0 {
success++
continue
}
ok := true
for i := range forwards {
if err := h.syncForwardServices(&forwards[i], "UpdateService", true); err != nil {
ok = false
break
}
}
if ok {
success++
} else {
fail++
}
}
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
}
func (h *Handler) userTunnelAssign(w http.ResponseWriter, r *http.Request) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
if err := h.upsertUserTunnel(req); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) userTunnelBatchAssign(w http.ResponseWriter, r *http.Request) {
var req struct {
UserID int64 `json:"userId"`
Tunnels []struct {
TunnelID int64 `json:"tunnelId"`
SpeedID *int64 `json:"speedId"`
} `json:"tunnels"`
}
if err := decodeJSON(r.Body, &req); err != nil || req.UserID <= 0 {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
for _, t := range req.Tunnels {
m := map[string]interface{}{"userId": req.UserID, "tunnelId": t.TunnelID}
if t.SpeedID != nil {
m["speedId"] = *t.SpeedID
}
if err := h.upsertUserTunnel(m); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) userTunnelRemove(w http.ResponseWriter, r *http.Request) {
id := idFromBody(r, w)
if id <= 0 {
return
}
_, err := h.repo.DB().Exec(`DELETE FROM user_tunnel WHERE id = ?`, id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("权限ID不能为空"))
return
}
_, err := h.repo.DB().Exec(`
UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, speed_id = ?, status = ? WHERE id = ?
`,
asInt64(req["flow"], 0),
asInt(req["num"], 0),
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
asInt64(req["flowResetTime"], 1),
nullableInt(asAnyToInt64Ptr(req["speedId"])),
asInt(req["status"], 1),
id,
)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
userID, roleID, err := userRoleFromRequest(r)
if err != nil {
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
return
}
tunnelID := asInt64(req["tunnelId"], 0)
if tunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
return
}
if err := h.ensureTunnelPermission(userID, roleID, tunnelID); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
tunnel, err := h.getTunnelRecord(tunnelID)
if err != nil {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
if tunnel.Status != 1 {
response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法创建转发"))
return
}
name := asString(req["name"])
remoteAddr := asString(req["remoteAddr"])
if name == "" || remoteAddr == "" {
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
return
}
port := asInt(req["inPort"], 0)
if port <= 0 {
port = h.pickTunnelPort(tunnelID)
}
if port <= 0 {
port = 10000
}
now := time.Now().UnixMilli()
inx := nextIndex(h.repo.DB(), "forward")
var userName string
_ = h.repo.DB().QueryRow(`SELECT user FROM user WHERE id = ?`, userID).Scan(&userName)
if userName == "" {
userName = "user"
}
tx, err := h.repo.DB().Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer func() { _ = tx.Rollback() }()
res, err := tx.Exec(`
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
`, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
forwardID, _ := res.LastInsertId()
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
for _, nodeID := range entryNodes {
_, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port)
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
createdForward, err := h.getForwardRecord(forwardID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.syncForwardServices(createdForward, "AddService", false); err != nil {
_ = h.deleteForwardByID(forwardID)
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("转发ID不能为空"))
return
}
forward, actorUserID, actorRole, err := h.resolveForwardAccess(r, id)
if err != nil {
if errors.Is(err, errForwardNotFound) {
response.WriteJSON(w, response.ErrDefault("转发不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
tunnelID := asInt64(req["tunnelId"], forward.TunnelID)
if tunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
return
}
if err := h.ensureTunnelPermission(actorUserID, actorRole, tunnelID); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
tunnel, err := h.getTunnelRecord(tunnelID)
if err != nil {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
if tunnel.Status != 1 {
response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法更新转发"))
return
}
name := strings.TrimSpace(asString(req["name"]))
if name == "" {
name = forward.Name
}
remoteAddr := strings.TrimSpace(asString(req["remoteAddr"]))
if remoteAddr == "" {
remoteAddr = forward.RemoteAddr
}
strategy := strings.TrimSpace(asString(req["strategy"]))
if strategy == "" {
strategy = forward.Strategy
}
port := asInt(req["inPort"], 0)
if port <= 0 {
var minPort sql.NullInt64
_ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&minPort)
if minPort.Valid {
port = int(minPort.Int64)
}
if port <= 0 {
port = h.pickTunnelPort(tunnelID)
}
}
now := time.Now().UnixMilli()
_, err = h.repo.DB().Exec(`
UPDATE forward SET name = ?, tunnel_id = ?, remote_addr = ?, strategy = ?, updated_time = ? WHERE id = ?
`, name, tunnelID, remoteAddr, strategy, now, id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.replaceForwardPorts(id, tunnelID, port)
updatedForward, err := h.getForwardRecord(id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
id := idFromBody(r, w)
if id <= 0 {
return
}
forward, _, _, err := h.resolveForwardAccess(r, id)
if err != nil {
if errors.Is(err, errForwardNotFound) {
response.WriteJSON(w, response.ErrDefault("转发不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
if err := h.deleteForwardByID(id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
h.forwardDelete(w, r)
}
func (h *Handler) forwardPause(w http.ResponseWriter, r *http.Request) {
id := idFromBody(r, w)
if id <= 0 {
return
}
forward, _, _, err := h.resolveForwardAccess(r, id)
if err != nil {
if errors.Is(err, errForwardNotFound) {
response.WriteJSON(w, response.ErrDefault("转发不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
_, _ = h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) forwardResume(w http.ResponseWriter, r *http.Request) {
id := idFromBody(r, w)
if id <= 0 {
return
}
forward, _, _, err := h.resolveForwardAccess(r, id)
if err != nil {
if errors.Is(err, errForwardNotFound) {
response.WriteJSON(w, response.ErrDefault("转发不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
_, _ = h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) forwardDiagnose(w http.ResponseWriter, r *http.Request) {
id := asInt64FromBodyKey(r, w, "forwardId")
if id <= 0 {
return
}
forward, _, _, err := h.resolveForwardAccess(r, id)
if err != nil {
if errors.Is(err, errForwardNotFound) {
response.WriteJSON(w, response.ErrDefault("转发不存在"))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
payload, err := h.diagnoseForwardRuntime(forward)
if err != nil {
if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不能为空") || strings.Contains(err.Error(), "错误") {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(payload))
}
func (h *Handler) forwardUpdateOrder(w http.ResponseWriter, r *http.Request) {
var req struct {
Forwards []struct {
ID int64 `json:"id"`
Inx int `json:"inx"`
} `json:"forwards"`
}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
for _, f := range req.Forwards {
_, _ = h.repo.DB().Exec(`UPDATE forward SET inx = ?, updated_time = ? WHERE id = ?`, f.Inx, time.Now().UnixMilli(), f.ID)
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
ids := idsFromBody(r, w)
if ids == nil {
return
}
actorUserID, actorRole, err := userRoleFromRequest(r)
if err != nil {
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
return
}
s := 0
f := 0
for _, id := range ids {
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
if accessErr != nil {
f++
continue
}
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
f++
continue
}
if err := h.deleteForwardByID(id); err != nil {
f++
} else {
s++
}
}
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
}
func (h *Handler) forwardBatchPause(w http.ResponseWriter, r *http.Request) {
ids := idsFromBody(r, w)
if ids == nil {
return
}
actorUserID, actorRole, err := userRoleFromRequest(r)
if err != nil {
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
return
}
s := 0
f := 0
for _, id := range ids {
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
if accessErr != nil {
f++
continue
}
if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
f++
continue
}
if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil {
f++
} else {
s++
}
}
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
}
func (h *Handler) forwardBatchResume(w http.ResponseWriter, r *http.Request) {
ids := idsFromBody(r, w)
if ids == nil {
return
}
actorUserID, actorRole, err := userRoleFromRequest(r)
if err != nil {
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
return
}
s := 0
f := 0
for _, id := range ids {
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
if accessErr != nil {
f++
continue
}
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
f++
continue
}
if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil {
f++
} else {
s++
}
}
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
}
func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
ids := idsFromBody(r, w)
if ids == nil {
return
}
actorUserID, actorRole, err := userRoleFromRequest(r)
if err != nil {
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
return
}
s := 0
f := 0
for _, id := range ids {
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
if accessErr != nil {
f++
continue
}
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
f++
} else {
s++
}
}
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": s, "failCount": f}))
}
func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Request) {
var req struct {
ForwardIDs []int64 `json:"forwardIds"`
TargetTunnelID int64 `json:"targetTunnelId"`
}
if err := decodeJSON(r.Body, &req); err != nil || req.TargetTunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
actorUserID, actorRole, err := userRoleFromRequest(r)
if err != nil {
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
return
}
if err := h.ensureTunnelPermission(actorUserID, actorRole, req.TargetTunnelID); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
targetTunnel, err := h.getTunnelRecord(req.TargetTunnelID)
if err != nil {
response.WriteJSON(w, response.ErrDefault("目标隧道不存在"))
return
}
if targetTunnel.Status != 1 {
response.WriteJSON(w, response.ErrDefault("目标隧道已禁用"))
return
}
success := 0
fail := 0
for _, id := range req.ForwardIDs {
if id <= 0 {
continue
}
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
if accessErr != nil {
fail++
continue
}
if forward.TunnelID == req.TargetTunnelID {
fail++
continue
}
var port sql.NullInt64
_ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&port)
_, err := h.repo.DB().Exec(`UPDATE forward SET tunnel_id = ?, updated_time = ? WHERE id = ?`, req.TargetTunnelID, time.Now().UnixMilli(), id)
if err != nil {
fail++
continue
}
p := 0
if port.Valid {
p = int(port.Int64)
}
if p <= 0 {
p = h.pickTunnelPort(req.TargetTunnelID)
}
_ = h.replaceForwardPorts(id, req.TargetTunnelID, p)
updatedForward, fetchErr := h.getForwardRecord(id)
if fetchErr != nil {
fail++
continue
}
if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil {
fail++
continue
}
success++
}
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
}
func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
tunnelID := asInt64(req["tunnelId"], 0)
if tunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
return
}
name := asString(req["name"])
if name == "" {
response.WriteJSON(w, response.ErrDefault("名称不能为空"))
return
}
var tunnelName string
_ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName)
if tunnelName == "" {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
now := time.Now().UnixMilli()
_, err := h.repo.DB().Exec(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`,
name, asInt(req["speed"], 100), tunnelID, tunnelName, now, now, asInt(req["status"], 1))
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
tunnelID := asInt64(req["tunnelId"], 0)
if id <= 0 || tunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
var tunnelName string
_ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName)
if tunnelName == "" {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
_, err := h.repo.DB().Exec(`UPDATE speed_limit SET name=?, speed=?, tunnel_id=?, tunnel_name=?, status=?, updated_time=? WHERE id=?`,
asString(req["name"]), asInt(req["speed"], 100), tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli(), id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
id := idFromBody(r, w)
if id <= 0 {
return
}
_, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) groupTunnelCreate(w http.ResponseWriter, r *http.Request) {
h.groupCreate(w, r, "tunnel_group")
}
func (h *Handler) groupTunnelUpdate(w http.ResponseWriter, r *http.Request) {
h.groupUpdate(w, r, "tunnel_group")
}
func (h *Handler) groupTunnelDelete(w http.ResponseWriter, r *http.Request) {
h.groupDelete(w, r, "tunnel_group")
}
func (h *Handler) groupUserCreate(w http.ResponseWriter, r *http.Request) {
h.groupCreate(w, r, "user_group")
}
func (h *Handler) groupUserUpdate(w http.ResponseWriter, r *http.Request) {
h.groupUpdate(w, r, "user_group")
}
func (h *Handler) groupUserDelete(w http.ResponseWriter, r *http.Request) {
h.groupDelete(w, r, "user_group")
}
func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) {
var req struct {
GroupID int64 `json:"groupId"`
TunnelIDs []int64 `json:"tunnelIds"`
}
if err := decodeJSON(r.Body, &req); err != nil || req.GroupID <= 0 {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
tx, err := h.repo.DB().Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID)
for _, tid := range req.TunnelIDs {
_, _ = tx.Exec(`INSERT OR IGNORE INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, req.GroupID, tid, time.Now().UnixMilli())
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.syncPermissionsByTunnelGroup(req.GroupID)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) {
var req struct {
GroupID int64 `json:"groupId"`
UserIDs []int64 `json:"userIds"`
}
if err := decodeJSON(r.Body, &req); err != nil || req.GroupID <= 0 {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
tx, err := h.repo.DB().Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID)
for _, uid := range req.UserIDs {
_, _ = tx.Exec(`INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)`, req.GroupID, uid, time.Now().UnixMilli())
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.syncPermissionsByUserGroup(req.GroupID)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request) {
var req struct {
UserGroupID int64 `json:"userGroupId"`
TunnelGroupID int64 `json:"tunnelGroupId"`
}
if err := decodeJSON(r.Body, &req); err != nil || req.UserGroupID <= 0 || req.TunnelGroupID <= 0 {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
_, err := h.repo.DB().Exec(`INSERT OR IGNORE INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.applyGroupPermission(req.UserGroupID, req.TunnelGroupID)
response.WriteJSON(w, response.OK("权限分配成功"))
}
func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) {
id := idFromBody(r, w)
if id <= 0 {
return
}
var ug, tg int64
_ = h.repo.DB().QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg)
_, _ = h.repo.DB().Exec(`DELETE FROM group_permission WHERE id = ?`, id)
_, _ = h.repo.DB().Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, ug, tg)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) groupCreate(w http.ResponseWriter, r *http.Request, table string) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
name := asString(req["name"])
if name == "" {
response.WriteJSON(w, response.ErrDefault("分组名称不能为空"))
return
}
now := time.Now().UnixMilli()
_, err := h.repo.DB().Exec(`INSERT INTO `+table+`(name, created_time, updated_time, status) VALUES(?, ?, ?, ?)`, name, now, now, asInt(req["status"], 1))
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) groupUpdate(w http.ResponseWriter, r *http.Request, table string) {
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("分组ID不能为空"))
return
}
_, err := h.repo.DB().Exec(`UPDATE `+table+` SET name = ?, status = ?, updated_time = ? WHERE id = ?`, asString(req["name"]), asInt(req["status"], 1), time.Now().UnixMilli(), id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) groupDelete(w http.ResponseWriter, r *http.Request, table string) {
id := idFromBody(r, w)
if id <= 0 {
return
}
tx, err := h.repo.DB().Begin()
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
defer func() { _ = tx.Rollback() }()
if table == "tunnel_group" {
_, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM group_permission WHERE tunnel_group_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE tunnel_group_id = ?`, id)
} else {
_, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM group_permission WHERE user_group_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ?`, id)
}
_, _ = tx.Exec(`DELETE FROM `+table+` WHERE id = ?`, id)
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error {
db := h.repo.DB()
userIDs, _ := queryInt64List(db, `SELECT user_id FROM user_group_user WHERE user_group_id = ?`, userGroupID)
tunnelIDs, _ := queryInt64List(db, `SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tunnelGroupID)
for _, uid := range userIDs {
for _, tid := range tunnelIDs {
utID, created, err := ensureUserTunnelGrant(db, uid, tid)
if err != nil {
continue
}
createdByGroup := 0
if created {
createdByGroup = 1
}
_, _ = db.Exec(`INSERT OR IGNORE INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?)`,
userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli())
}
}
return nil
}
func (h *Handler) syncPermissionsByUserGroup(userGroupID int64) error {
db := h.repo.DB()
pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE user_group_id = ?`, userGroupID)
if err != nil {
return err
}
for _, p := range pairs {
_ = h.applyGroupPermission(p[0], p[1])
}
return nil
}
func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error {
db := h.repo.DB()
pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE tunnel_group_id = ?`, tunnelGroupID)
if err != nil {
return err
}
for _, p := range pairs {
_ = h.applyGroupPermission(p[0], p[1])
}
return nil
}
func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, error) {
var id int64
err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&id)
if err == nil {
return id, false, nil
}
if err != sql.ErrNoRows {
return 0, false, err
}
var flow int64
var num int
var expTime int64
var flowReset int64
if err := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset); err != nil {
return 0, false, err
}
res, err := db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1)`,
userID, tunnelID, num, flow, flowReset, expTime)
if err != nil {
return 0, false, err
}
id, _ = res.LastInsertId()
return id, true, nil
}
func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error) {
rows, err := db.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := make([]int64, 0)
for rows.Next() {
var v int64
if err := rows.Scan(&v); err != nil {
return nil, err
}
out = append(out, v)
}
return out, rows.Err()
}
func queryPairs(db *sql.DB, q string, args ...interface{}) ([][2]int64, error) {
rows, err := db.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := make([][2]int64, 0)
for rows.Next() {
var a, b int64
if err := rows.Scan(&a, &b); err != nil {
return nil, err
}
out = append(out, [2]int64{a, b})
}
return out, rows.Err()
}
func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error {
inNodes := asMapSlice(req["inNodeId"])
for _, n := range inNodes {
nodeID := asInt64(n["nodeId"], 0)
if nodeID <= 0 {
continue
}
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, NULL, NULL, 0, ?)`,
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls"))
if err != nil {
return err
}
}
for _, n := range asMapSlice(req["outNodeId"]) {
nodeID := asInt64(n["nodeId"], 0)
if nodeID <= 0 {
continue
}
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, NULL, NULL, 0, ?)`,
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls"))
if err != nil {
return err
}
}
chainNodes := asAnySlice(req["chainNodes"])
for i, grp := range chainNodes {
for _, n := range asMapSlice(grp) {
nodeID := asInt64(n["nodeId"], 0)
if nodeID <= 0 {
continue
}
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, NULL, ?, ?, ?)`,
tunnelID, nodeID, defaultString(asString(n["strategy"]), "round"), i, defaultString(asString(n["protocol"]), "tls"))
if err != nil {
return err
}
}
}
return nil
}
func (h *Handler) deleteNodeByID(id int64) error {
tx, err := h.repo.DB().Begin()
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM forward_port WHERE node_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE node_id = ?`, id)
_, err = tx.Exec(`DELETE FROM node WHERE id = ?`, id)
if err != nil {
return err
}
return tx.Commit()
}
func (h *Handler) deleteTunnelByID(id int64) error {
tx, err := h.repo.DB().Begin()
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE tunnel_id = ?)`, id)
_, _ = tx.Exec(`DELETE FROM forward WHERE tunnel_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM user_tunnel WHERE tunnel_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM speed_limit WHERE tunnel_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id)
_, err = tx.Exec(`DELETE FROM tunnel WHERE id = ?`, id)
if err != nil {
return err
}
return tx.Commit()
}
func (h *Handler) deleteForwardByID(id int64) error {
tx, err := h.repo.DB().Begin()
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, id)
_, err = tx.Exec(`DELETE FROM forward WHERE id = ?`, id)
if err != nil {
return err
}
return tx.Commit()
}
func (h *Handler) batchForwardDelete(ids []int64) (int, int) {
s := 0
f := 0
for _, id := range ids {
if err := h.deleteForwardByID(id); err != nil {
f++
} else {
s++
}
}
return s, f
}
func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) {
s := 0
f := 0
for _, id := range ids {
if _, err := h.repo.DB().Exec(`UPDATE forward SET status = ?, updated_time = ? WHERE id = ?`, status, time.Now().UnixMilli(), id); err != nil {
f++
} else {
s++
}
}
return s, f
}
func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 ORDER BY inx ASC, id ASC`, tunnelID)
if err != nil {
return nil, err
}
defer rows.Close()
out := make([]int64, 0)
for rows.Next() {
var id int64
if err := rows.Scan(&id); err == nil {
out = append(out, id)
}
}
return out, rows.Err()
}
func (h *Handler) pickTunnelPort(tunnelID int64) int {
entry, _ := h.tunnelEntryNodeIDs(tunnelID)
if len(entry) == 0 {
return 10000
}
var portRange string
_ = h.repo.DB().QueryRow(`SELECT port FROM node WHERE id = ?`, entry[0]).Scan(&portRange)
if portRange == "" {
return 10000
}
first := strings.Split(portRange, ",")[0]
first = strings.TrimSpace(first)
if strings.Contains(first, "-") {
parts := strings.SplitN(first, "-", 2)
p, _ := strconv.Atoi(strings.TrimSpace(parts[0]))
if p > 0 {
return p
}
}
if p, err := strconv.Atoi(first); err == nil && p > 0 {
return p
}
return 10000
}
func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error {
tx, err := h.repo.DB().Begin()
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID)
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
for _, nodeID := range entryNodes {
_, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port)
}
return tx.Commit()
}
func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
userID := asInt64(req["userId"], 0)
tunnelID := asInt64(req["tunnelId"], 0)
if userID <= 0 || tunnelID <= 0 {
return fmt.Errorf("userId or tunnelId missing")
}
db := h.repo.DB()
var existingID int64
err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&existingID)
flow := asInt64(req["flow"], -1)
num := asInt(req["num"], -1)
expTime := asInt64(req["expTime"], -1)
flowReset := asInt64(req["flowResetTime"], -1)
status := asInt(req["status"], 1)
speedID := asAnyToInt64Ptr(req["speedId"])
if err == sql.ErrNoRows {
if flow < 0 || num < 0 || expTime < 0 || flowReset < 0 {
_ = db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset)
}
_, err = db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?)`,
userID, tunnelID, nullableInt(speedID), num, flow, flowReset, expTime, status)
return err
}
if err != nil {
return err
}
if flow < 0 {
flow = 0
}
if num < 0 {
num = 0
}
if expTime < 0 {
expTime = time.Now().Add(365 * 24 * time.Hour).UnixMilli()
}
if flowReset < 0 {
flowReset = 1
}
_, err = db.Exec(`UPDATE user_tunnel SET speed_id = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ? WHERE id = ?`,
nullableInt(speedID), flow, num, expTime, flowReset, status, existingID)
return err
}
func asAnySlice(v interface{}) []interface{} {
if v == nil {
return nil
}
if arr, ok := v.([]interface{}); ok {
return arr
}
return nil
}
func asMapSlice(v interface{}) []map[string]interface{} {
arr := asAnySlice(v)
if arr == nil {
return nil
}
out := make([]map[string]interface{}, 0, len(arr))
for _, it := range arr {
if m, ok := it.(map[string]interface{}); ok {
out = append(out, m)
}
}
return out
}
func asString(v interface{}) string {
switch t := v.(type) {
case nil:
return ""
case string:
return strings.TrimSpace(t)
case float64:
if t == float64(int64(t)) {
return strconv.FormatInt(int64(t), 10)
}
return strconv.FormatFloat(t, 'f', -1, 64)
case int, int32, int64:
return fmt.Sprintf("%v", t)
default:
b, _ := json.Marshal(t)
return strings.Trim(string(b), "\"")
}
}
func asInt(v interface{}, def int) int {
s := asString(v)
if s == "" {
return def
}
i, err := strconv.Atoi(s)
if err != nil {
return def
}
return i
}
func asInt64(v interface{}, def int64) int64 {
s := asString(v)
if s == "" {
return def
}
i, err := strconv.ParseInt(s, 10, 64)
if err != nil {
return def
}
return i
}
func asFloat(v interface{}, def float64) float64 {
s := asString(v)
if s == "" {
return def
}
f, err := strconv.ParseFloat(s, 64)
if err != nil {
return def
}
return f
}
func asAnyToInt64Ptr(v interface{}) *int64 {
s := asString(v)
if s == "" || strings.EqualFold(s, "null") {
return nil
}
i, err := strconv.ParseInt(s, 10, 64)
if err != nil {
return nil
}
return &i
}
func idFromBody(r *http.Request, w http.ResponseWriter) int64 {
return asInt64FromBodyKey(r, w, "id")
}
func asInt64FromBodyKey(r *http.Request, w http.ResponseWriter, key string) int64 {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return 0
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return 0
}
id := asInt64(req[key], 0)
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("参数错误"))
return 0
}
return id
}
func idsFromBody(r *http.Request, w http.ResponseWriter) []int64 {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return nil
}
var req map[string]interface{}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return nil
}
arr := asAnySlice(req["ids"])
if len(arr) == 0 {
response.WriteJSON(w, response.ErrDefault("ids不能为空"))
return nil
}
ids := make([]int64, 0, len(arr))
for _, x := range arr {
id := asInt64(x, 0)
if id > 0 {
ids = append(ids, id)
}
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids
}
func nullableText(s string) interface{} {
if strings.TrimSpace(s) == "" {
return nil
}
return s
}
func nullableInt(v *int64) interface{} {
if v == nil {
return nil
}
return *v
}
func defaultString(v, def string) string {
if strings.TrimSpace(v) == "" {
return def
}
return v
}
func randomToken(n int) string {
buf := make([]byte, n)
if _, err := rand.Read(buf); err != nil {
return strconv.FormatInt(time.Now().UnixNano(), 16)
}
return hex.EncodeToString(buf)
}
func nextIndex(db *sql.DB, table string) int {
if db == nil {
return 0
}
row := db.QueryRow(`SELECT COALESCE(MAX(inx), -1) + 1 FROM ` + table)
var n int
if err := row.Scan(&n); err != nil {
return 0
}
if n < 0 {
return 0
}
return n
}