Merge branch 'main' into opencode/silent-wizard

This commit is contained in:
sagit
2026-02-09 14:35:25 +08:00
committed by GitHub
4 changed files with 145 additions and 11 deletions
@@ -59,6 +59,8 @@ type chainNodeRecord struct {
NodeID int64
Port int
NodeName string
Protocol string
Strategy string
}
type diagnosisTarget struct {
@@ -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 {
@@ -1039,7 +1045,10 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
}
if limiter != nil && *limiter > 0 {
service["limiter"] = strconv.Itoa(*limiter)
// Convert Mbps to Bytes/s
// 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)
}
@@ -775,6 +775,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 {
@@ -783,6 +859,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++
+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: