Merge branch 'main' into opencode/silent-wizard

This commit is contained in:
sagit
2026-02-09 19:42:48 +08:00
committed by GitHub
4 changed files with 83 additions and 21 deletions
+1
View File
@@ -12,6 +12,7 @@ services:
JWT_SECRET: ${JWT_SECRET} JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs LOG_DIR: /app/logs
SERVER_ADDR: :6365 SERVER_ADDR: :6365
TZ: Asia/Shanghai
ports: ports:
- "${BACKEND_PORT}:6365" - "${BACKEND_PORT}:6365"
volumes: volumes:
+1
View File
@@ -12,6 +12,7 @@ services:
JWT_SECRET: ${JWT_SECRET} JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs LOG_DIR: /app/logs
SERVER_ADDR: :6365 SERVER_ADDR: :6365
TZ: Asia/Shanghai
ports: ports:
- "${BACKEND_PORT}:6365" - "${BACKEND_PORT}:6365"
volumes: volumes:
@@ -234,9 +234,9 @@ func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
return &n, nil return &n, nil
} }
func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int, error) { func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int64, *int, error) {
row := h.repo.DB().QueryRow(` row := h.repo.DB().QueryRow(`
SELECT ut.id, sl.speed SELECT ut.id, sl.id, sl.speed
FROM user_tunnel ut FROM user_tunnel ut
LEFT JOIN speed_limit sl ON sl.id = ut.speed_id LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
WHERE ut.user_id = ? AND ut.tunnel_id = ? WHERE ut.user_id = ? AND ut.tunnel_id = ?
@@ -244,19 +244,21 @@ func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *i
LIMIT 1 LIMIT 1
`, userID, tunnelID) `, userID, tunnelID)
var userTunnelID int64 var userTunnelID int64
var limiterID sql.NullInt64
var speed sql.NullInt64 var speed sql.NullInt64
err := row.Scan(&userTunnelID, &speed) err := row.Scan(&userTunnelID, &limiterID, &speed)
if err != nil { if err != nil {
if errors.Is(err, sql.ErrNoRows) { if errors.Is(err, sql.ErrNoRows) {
return 0, nil, nil return 0, nil, nil, nil
} }
return 0, nil, err return 0, nil, nil, err
} }
if !speed.Valid || speed.Int64 <= 0 { if !limiterID.Valid || limiterID.Int64 <= 0 {
return userTunnelID, nil, nil return userTunnelID, nil, nil, nil
} }
v := int(speed.Int64) v := limiterID.Int64
return userTunnelID, &v, nil s := int(speed.Int64)
return userTunnelID, &v, &s, nil
} }
func (h *Handler) listUserTunnelIDs(userID, tunnelID int64) ([]int64, error) { func (h *Handler) listUserTunnelIDs(userID, tunnelID int64) ([]int64, error) {
@@ -328,18 +330,22 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
return errors.New("转发入口端口不存在") return errors.New("转发入口端口不存在")
} }
userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) userTunnelID, limiterID, speed, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil { if err != nil {
return err return err
} }
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID) serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
for _, fp := range ports { for _, fp := range ports {
if limiterID != nil && speed != nil {
h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed)
}
node, err := h.getNodeRecord(fp.NodeID) node, err := h.getNodeRecord(fp.NodeID)
if err != nil { if err != nil {
return err return err
} }
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiter) services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiterID)
_, err = h.sendNodeCommand(node.ID, method, services, true, false) _, err = h.sendNodeCommand(node.ID, method, services, true, false)
if err != nil && allowFallbackAdd && method == "UpdateService" { if err != nil && allowFallbackAdd && method == "UpdateService" {
_, err = h.sendNodeCommand(node.ID, "AddService", services, true, false) _, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
@@ -362,7 +368,7 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str
if len(ports) == 0 { if len(ports) == 0 {
return nil return nil
} }
userTunnelID, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID) userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil { if err != nil {
return err return err
} }
@@ -1003,7 +1009,7 @@ func isNotFoundError(err error) bool {
return strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") return strings.Contains(msg, "not found") || strings.Contains(msg, "不存在")
} }
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiter *int) []map[string]interface{} { func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiterID *int64) []map[string]interface{} {
protocols := []string{"tcp", "udp"} protocols := []string{"tcp", "udp"}
services := make([]map[string]interface{}, 0, 2) services := make([]map[string]interface{}, 0, 2)
targets := splitRemoteTargets(forward.RemoteAddr) targets := splitRemoteTargets(forward.RemoteAddr)
@@ -1044,11 +1050,8 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" { if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName} service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
} }
if limiter != nil && *limiter > 0 { if limiterID != nil && *limiterID > 0 {
// Convert Mbps to Bytes/s service["limiter"] = strconv.FormatInt(*limiterID, 10)
// 1 Mbps = 1,000,000 bits/s = 125,000 Bytes/s
// We use decimal Mbps standard as is common in networking
service["limiter"] = strconv.Itoa(*limiter * 125000)
} }
services = append(services, service) services = append(services, service)
} }
@@ -1111,3 +1114,49 @@ func asBool(v interface{}, def bool) bool {
return def return def
} }
} }
func (h *Handler) sendLimiterConfig(limiterID int64, speedMbps int, tunnelID int64) error {
rate := float64(speedMbps) / 8.0
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
payload := map[string]interface{}{
"name": strconv.FormatInt(limiterID, 10),
"limits": []string{limitStr},
}
nodes, err := h.tunnelEntryNodeIDs(tunnelID)
if err != nil {
return err
}
for _, nodeID := range nodes {
_, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false)
}
return nil
}
func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error {
payload := map[string]interface{}{
"limiter": strconv.FormatInt(limiterID, 10),
}
nodes, err := h.tunnelEntryNodeIDs(tunnelID)
if err != nil {
return err
}
for _, nodeID := range nodes {
_, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", payload, false, true)
}
return nil
}
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) {
rate := float64(speed) / 8.0
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
payload := map[string]interface{}{
"name": strconv.FormatInt(limiterID, 10),
"limits": []string{limitStr},
}
_, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false)
}
+14 -3
View File
@@ -1518,12 +1518,15 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
return return
} }
now := time.Now().UnixMilli() 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(?, ?, ?, ?, ?, ?, ?)`, speed := asInt(req["speed"], 100)
name, asInt(req["speed"], 100), tunnelID, tunnelName, now, now, asInt(req["status"], 1)) res, err := h.repo.DB().Exec(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`,
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1))
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
id, _ := res.LastInsertId()
_ = h.sendLimiterConfig(id, speed, tunnelID)
response.WriteJSON(w, response.OKEmpty()) response.WriteJSON(w, response.OKEmpty())
} }
@@ -1545,12 +1548,14 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("隧道不存在")) response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return return
} }
speed := asInt(req["speed"], 100)
_, err := h.repo.DB().Exec(`UPDATE speed_limit SET name=?, speed=?, tunnel_id=?, tunnel_name=?, status=?, updated_time=? WHERE id=?`, _, 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) asString(req["name"]), speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli(), id)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
_ = h.sendLimiterConfig(id, speed, tunnelID)
response.WriteJSON(w, response.OKEmpty()) response.WriteJSON(w, response.OKEmpty())
} }
@@ -1559,11 +1564,17 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
if id <= 0 { if id <= 0 {
return return
} }
var tunnelID int64
_ = h.repo.DB().QueryRow(`SELECT tunnel_id FROM speed_limit WHERE id = ?`, id).Scan(&tunnelID)
_, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id) _, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
if tunnelID > 0 {
_ = h.sendDeleteLimiterConfig(id, tunnelID)
}
response.WriteJSON(w, response.OKEmpty()) response.WriteJSON(w, response.OKEmpty())
} }