Compare commits

...

15 Commits

Author SHA1 Message Date
sagit 065b23d9c3 Merge pull request #59 from Sagit-chu/opencode/calm-orchid
refactor(backend): reimplement speed limit logic
2026-02-09 16:15:59 +08:00
sagit 7919dfde59 Merge branch 'main' into opencode/calm-orchid 2026-02-09 16:13:40 +08:00
sagit 565d732967 refactor(backend): reimplement speed limit logic
1. Refactor speed limit CRUD to sync with agents immediately via WebSocket (AddLimiters/DeleteLimiters).
2. Update unit conversion to match GOST v3 requirements (Mbps -> MB/s).
3. Update service config generation to reference Limiter IDs instead of hardcoded values.
2026-02-09 08:12:10 +00:00
sagit 20dc151aec Merge pull request #58 from Sagit-chu/opencode/calm-orchid
fix(backend): fix tunnel batch redeploy logic for type 2 tunnels and speed limit
2026-02-09 14:28:01 +08:00
sagit 630ed969d3 Merge branch 'main' into opencode/calm-orchid 2026-02-09 14:23:20 +08:00
sagit e94aa01213 fix(gost): append 'B' suffix to speed limit values for correct unit parsing 2026-02-09 06:22:47 +00:00
sagit 67d8f7a381 fix(backend): correct speed limit unit conversion from Mbps to Bytes/s 2026-02-09 06:13:56 +00:00
sagit 0c7b7deaf5 fix(backend): fix tunnel batch redeploy logic for type 2 tunnels 2026-02-09 05:17:58 +00:00
sagit a4def9c5f3 Merge pull request #57 from Sagit-chu/opencode/calm-orchid
fix: prevent nil pointer dereference in listener config parsing
2026-02-09 12:42:16 +08:00
sagit 6582348da2 Merge branch 'main' into opencode/calm-orchid 2026-02-09 12:40:57 +08:00
sagit 3a14b22ebc fix: prevent nil pointer dereference in listener config parsing 2026-02-09 04:39:27 +00:00
sagit d7b44916bf Merge pull request #56 from Sagit-chu/opencode/calm-orchid
fix(limiter): fix traffic limiter ScopeClient behavior to allow per-u…
2026-02-09 11:39:23 +08:00
sagit f8a0bda3fd Merge branch 'main' into opencode/calm-orchid 2026-02-09 11:37:49 +08:00
sagit 634562e56d fix(config): support raw number string for limiter configuration 2026-02-09 03:17:02 +00:00
sagit d06e02998b fix(limiter): fix traffic limiter ScopeClient behavior to allow per-user limits 2026-02-09 03:12:33 +00:00
4 changed files with 202 additions and 24 deletions
@@ -59,6 +59,8 @@ type chainNodeRecord struct {
NodeID int64
Port int
NodeName string
Protocol string
Strategy string
}
type diagnosisTarget struct {
@@ -232,9 +234,9 @@ func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
return &n, nil
}
func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int, error) {
func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int64, error) {
row := h.repo.DB().QueryRow(`
SELECT ut.id, sl.speed
SELECT ut.id, sl.id
FROM user_tunnel ut
LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
WHERE ut.user_id = ? AND ut.tunnel_id = ?
@@ -242,18 +244,18 @@ func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *i
LIMIT 1
`, userID, tunnelID)
var userTunnelID int64
var speed sql.NullInt64
err := row.Scan(&userTunnelID, &speed)
var limiterID sql.NullInt64
err := row.Scan(&userTunnelID, &limiterID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return 0, nil, nil
}
return 0, nil, err
}
if !speed.Valid || speed.Int64 <= 0 {
if !limiterID.Valid || limiterID.Int64 <= 0 {
return userTunnelID, nil, nil
}
v := int(speed.Int64)
v := limiterID.Int64
return userTunnelID, &v, nil
}
@@ -326,7 +328,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
return errors.New("转发入口端口不存在")
}
userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
userTunnelID, limiterID, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil {
return err
}
@@ -337,7 +339,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
if err != nil {
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)
if err != nil && allowFallbackAdd && method == "UpdateService" {
_, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
@@ -862,7 +864,7 @@ func firstPortFromRange(portRange string) int {
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
SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
FROM chain_tunnel ct
LEFT JOIN node n ON n.id = ct.node_id
WHERE ct.tunnel_id = ?
@@ -877,7 +879,9 @@ func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, er
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 {
var protocol sql.NullString
var strategy sql.NullString
if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &item.Port, &name, &protocol, &strategy); err != nil {
return nil, err
}
if strings.TrimSpace(name.String) == "" {
@@ -885,6 +889,8 @@ func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, er
} else {
item.NodeName = name.String
}
item.Protocol = defaultString(protocol.String, "tls")
item.Strategy = defaultString(strategy.String, "round")
result = append(result, item)
}
if err := rows.Err(); err != nil {
@@ -997,7 +1003,7 @@ func isNotFoundError(err error) bool {
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"}
services := make([]map[string]interface{}, 0, 2)
targets := splitRemoteTargets(forward.RemoteAddr)
@@ -1038,8 +1044,8 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
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)
if limiterID != nil && *limiterID > 0 {
service["limiter"] = strconv.FormatInt(*limiterID, 10)
}
services = append(services, service)
}
@@ -1102,3 +1108,39 @@ func asBool(v interface{}, def bool) bool {
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
}
+110 -3
View File
@@ -731,6 +731,82 @@ func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.OK(map[string]interface{}{"successCount": success, "failCount": fail}))
}
func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, error) {
tunnel, err := h.getTunnelRecord(tunnelID)
if err != nil {
return nil, err
}
chainRows, err := h.listChainNodesForTunnel(tunnelID)
if err != nil {
return nil, err
}
state := &tunnelCreateState{
TunnelID: tunnelID,
Type: tunnel.Type,
InNodes: make([]tunnelRuntimeNode, 0),
ChainHops: make([][]tunnelRuntimeNode, 0),
OutNodes: make([]tunnelRuntimeNode, 0),
Nodes: make(map[int64]*nodeRecord),
NodeIDList: make([]int64, 0),
}
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
for _, r := range inNodes {
state.InNodes = append(state.InNodes, tunnelRuntimeNode{
NodeID: r.NodeID,
Protocol: r.Protocol,
Strategy: r.Strategy,
ChainType: 1,
})
state.NodeIDList = append(state.NodeIDList, r.NodeID)
}
for _, r := range outNodes {
state.OutNodes = append(state.OutNodes, tunnelRuntimeNode{
NodeID: r.NodeID,
Protocol: r.Protocol,
Strategy: r.Strategy,
ChainType: 3,
Port: r.Port,
})
state.NodeIDList = append(state.NodeIDList, r.NodeID)
}
for _, hop := range chainHops {
stateHop := make([]tunnelRuntimeNode, 0)
for _, r := range hop {
stateHop = append(stateHop, tunnelRuntimeNode{
NodeID: r.NodeID,
Protocol: r.Protocol,
Strategy: r.Strategy,
ChainType: 2,
Inx: int(r.Inx),
Port: r.Port,
})
state.NodeIDList = append(state.NodeIDList, r.NodeID)
}
state.ChainHops = append(state.ChainHops, stateHop)
}
seen := make(map[int64]struct{})
for _, id := range state.NodeIDList {
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
node, err := h.getNodeRecord(id)
if err != nil {
return nil, err
}
state.Nodes[id] = node
}
return state, nil
}
func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
ids := idsFromBody(r, w)
if ids == nil {
@@ -739,6 +815,26 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
success := 0
fail := 0
for _, tunnelID := range ids {
tunnel, err := h.getTunnelRecord(tunnelID)
if err != nil {
fail++
continue
}
if tunnel.Type == 2 {
h.cleanupTunnelRuntime(tunnelID)
state, err := h.reconstructTunnelState(tunnelID)
if err != nil {
fail++
continue
}
_, _, applyErr := h.applyTunnelRuntime(state)
if applyErr != nil {
fail++
continue
}
}
forwards, err := h.listForwardsByTunnel(tunnelID)
if err != nil {
fail++
@@ -1373,12 +1469,15 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
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))
speed := asInt(req["speed"], 100)
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 {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
id, _ := res.LastInsertId()
_ = h.sendLimiterConfig(id, speed, tunnelID)
response.WriteJSON(w, response.OKEmpty())
}
@@ -1400,12 +1499,14 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
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=?`,
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 {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.sendLimiterConfig(id, speed, tunnelID)
response.WriteJSON(w, response.OKEmpty())
}
@@ -1414,11 +1515,17 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
if id <= 0 {
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)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if tunnelID > 0 {
_ = h.sendDeleteLimiterConfig(id, tunnelID)
}
response.WriteJSON(w, response.OKEmpty())
}
+31 -8
View File
@@ -3,6 +3,7 @@ package service
import (
"fmt"
"runtime"
"strconv"
"strings"
"time"
@@ -30,6 +31,7 @@ import (
logger_parser "github.com/go-gost/x/config/parsing/logger"
selector_parser "github.com/go-gost/x/config/parsing/selector"
tls_util "github.com/go-gost/x/internal/util/tls"
xtraffic "github.com/go-gost/x/limiter/traffic"
cache_limiter "github.com/go-gost/x/limiter/traffic/cache"
"github.com/go-gost/x/metadata"
mdutil "github.com/go-gost/x/metadata/util"
@@ -181,6 +183,32 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
)
}
var trafficLimiter listener.Option
if cfg.Limiter != "" {
lim := registry.TrafficLimiterRegistry().Get(cfg.Limiter)
if lim == nil {
// Try to parse as simple number (bandwidth in bytes/sec)
if val, err := strconv.Atoi(cfg.Limiter); err == nil && val > 0 {
lim = xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %dB %dB", xtraffic.ServiceLimitKey, val, val)),
)
}
if lim == nil {
lim = xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %s %s", xtraffic.ServiceLimitKey, cfg.Limiter, cfg.Limiter)),
)
}
}
trafficLimiter = listener.TrafficLimiterOption(
cache_limiter.NewCachedTrafficLimiter(
lim,
cache_limiter.RefreshIntervalOption(limiterRefreshInterval),
cache_limiter.CleanupIntervalOption(limiterCleanupInterval),
cache_limiter.ScopeOption(limiterScope),
),
)
}
listenOpts := []listener.Option{
listener.AddrOption(cfg.Addr),
listener.RouterOption(xchain.NewRouter(routerOpts...)),
@@ -188,14 +216,6 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
listener.AuthOption(auth_parser.Info(cfg.Listener.Auth)),
listener.TLSConfigOption(tlsConfig),
listener.AdmissionOption(xadmission.AdmissionGroup(admissions...)),
listener.TrafficLimiterOption(
cache_limiter.NewCachedTrafficLimiter(
registry.TrafficLimiterRegistry().Get(cfg.Limiter),
cache_limiter.RefreshIntervalOption(limiterRefreshInterval),
cache_limiter.CleanupIntervalOption(limiterCleanupInterval),
cache_limiter.ScopeOption(limiterScope),
),
),
listener.ConnLimiterOption(registry.ConnLimiterRegistry().Get(cfg.CLimiter)),
listener.ServiceOption(cfg.Name),
listener.ProxyProtocolOption(ppv),
@@ -203,6 +223,9 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
listener.NetnsOption(netnsIn),
listener.LoggerOption(listenerLogger),
}
if trafficLimiter != nil {
listenOpts = append(listenOpts, trafficLimiter)
}
if netnsIn != "" {
runtime.LockOSThread()
+6
View File
@@ -136,6 +136,9 @@ func (l *trafficLimiter) In(ctx context.Context, key string, opts ...limiter.Opt
return nil
case limiter.ScopeClient:
if lim, ok := l.inLimits.Get(key); ok && lim != nil {
return lim.(traffic.Limiter)
}
return nil
case limiter.ScopeConn:
@@ -215,6 +218,9 @@ func (l *trafficLimiter) Out(ctx context.Context, key string, opts ...limiter.Op
return nil
case limiter.ScopeClient:
if lim, ok := l.outLimits.Get(key); ok && lim != nil {
return lim.(traffic.Limiter)
}
return nil
case limiter.ScopeConn: