Files

703 lines
20 KiB
Go

package handler
import (
"encoding/json"
"errors"
"log"
"math"
"strconv"
"strings"
"time"
)
const bytesPerGB int64 = 1024 * 1024 * 1024
const bytesPerMiB int64 = 1024 * 1024
func flowLimitBytes(flowGB, flowMiB int64) int64 {
if flowMiB > 0 {
if flowMiB > math.MaxInt64/bytesPerMiB {
return math.MaxInt64
}
return flowMiB * bytesPerMiB
}
if flowGB > math.MaxInt64/bytesPerGB {
return math.MaxInt64
}
return flowGB * bytesPerGB
}
type userTunnelPolicy struct {
ID int64
UserID int64
TunnelID int64
Flow int64
FlowMiB int64
InFlow int64
OutFlow int64
ExpTime int64
Status int
Num int
}
type gostConfigSnapshot struct {
Services []namedConfigItem `json:"services"`
Chains []namedConfigItem `json:"chains"`
Limiters []namedConfigItem `json:"limiters"`
}
type namedConfigItem struct {
Name string `json:"name"`
Limiter string `json:"limiter,omitempty"`
Handler *struct {
Chain string `json:"chain"`
} `json:"handler,omitempty"`
}
func (h *Handler) processFlowItem(nodeID int64, item flowItem) {
if h == nil || h.repo == nil || nodeID <= 0 {
return
}
metas, err := h.repo.GetFlowUploadForwardMetas(collectFlowUploadForwardIDs([]flowItem{item}))
if err != nil {
metas = nil
}
h.applyFlowUploadBatch(nodeID, h.buildNodeFlowUploadBatch(nodeID, []flowItem{item}, metas), time.Now())
}
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 parsePeerShareRuntimeServiceID(serviceName string) (int64, bool) {
const prefix = "fed_svc_"
if !strings.HasPrefix(serviceName, prefix) {
return 0, false
}
raw := strings.TrimPrefix(serviceName, prefix)
if raw == "" {
return 0, false
}
parts := strings.SplitN(raw, "_", 2)
runtimeID, err := strconv.ParseInt(parts[0], 10, 64)
if err != nil || runtimeID <= 0 {
return 0, false
}
return runtimeID, true
}
func parsePeerShareInfoFromFederationTunnelName(tunnelName string) (int64, int, bool) {
tunnelName = strings.TrimSpace(tunnelName)
if !strings.HasPrefix(tunnelName, "Share-") {
return 0, 0, false
}
raw := strings.TrimPrefix(tunnelName, "Share-")
idx := strings.Index(raw, "-Port-")
if idx <= 0 {
return 0, 0, false
}
shareID, err := strconv.ParseInt(raw[:idx], 10, 64)
if err != nil || shareID <= 0 {
return 0, 0, false
}
portValue := strings.TrimSpace(raw[idx+len("-Port-"):])
port, err := strconv.Atoi(portValue)
if err != nil || port <= 0 {
return 0, 0, false
}
return shareID, port, true
}
func parsePeerShareIDFromFederationTunnelName(tunnelName string) (int64, bool) {
tunnelName = strings.TrimSpace(tunnelName)
if !strings.HasPrefix(tunnelName, "Share-") {
return 0, false
}
raw := strings.TrimPrefix(tunnelName, "Share-")
idx := strings.Index(raw, "-Port-")
if idx <= 0 {
return 0, false
}
shareID, err := strconv.ParseInt(raw[:idx], 10, 64)
if err != nil || shareID <= 0 {
return 0, false
}
return shareID, true
}
func (h *Handler) processPeerShareFlow(nodeID, runtimeID int64, item flowItem) {
if h == nil || h.repo == nil || nodeID <= 0 || runtimeID <= 0 {
return
}
runtime, err := h.repo.GetPeerShareRuntimeByID(runtimeID)
if err != nil || runtime == nil || runtime.NodeID != nodeID || runtime.Status != 1 {
return
}
h.addPeerShareFlow(nodeID, runtime.ShareID, item.D+item.U)
}
func (h *Handler) processPeerShareFlowFromForward(forwardID int64, nodeID int64, serviceName string, item flowItem) {
if h == nil || h.repo == nil || forwardID <= 0 || nodeID <= 0 {
return
}
// Prefer the reporting node's explicit shared ownership over a coincidentally
// equal local forward ID. Never fall back to a service on another node.
runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID)
if err != nil {
return
}
for _, runtime := range runtimes {
if normalizeForwardRuntimeServiceName(runtime.ServiceName) == normalizeForwardRuntimeServiceName(serviceName) {
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
return
}
}
forward, err := h.getForwardRecord(forwardID)
if err != nil || forward == nil {
return
}
_, userID, _, ok := parseFlowServiceIDs(serviceName)
if !ok || userID != forward.UserID {
return
}
tunnelName, err := h.repo.GetTunnelName(forward.TunnelID)
if err != nil {
return
}
shareID, ok := parsePeerShareIDFromFederationTunnelName(tunnelName)
if ok {
h.addPeerShareFlow(nodeID, shareID, item.D+item.U)
}
}
func normalizeForwardRuntimeServiceName(serviceName string) string {
name := strings.TrimSpace(serviceName)
if strings.HasSuffix(name, "_tcp") {
return strings.TrimSuffix(name, "_tcp")
}
if strings.HasSuffix(name, "_udp") {
return strings.TrimSuffix(name, "_udp")
}
return name
}
func (h *Handler) processPeerShareFlowByServiceName(nodeID int64, serviceName string, item flowItem) {
if h == nil || h.repo == nil || nodeID <= 0 || strings.TrimSpace(serviceName) == "" {
return
}
runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID)
if err != nil {
return
}
var shareID int64
for _, runtime := range runtimes {
if normalizeForwardRuntimeServiceName(runtime.ServiceName) != normalizeForwardRuntimeServiceName(serviceName) {
continue
}
if shareID != 0 {
log.Printf("ambiguous peer share runtime service=%s node_id=%d", serviceName, nodeID)
return
}
shareID = runtime.ShareID
}
if shareID > 0 {
h.addPeerShareFlow(nodeID, shareID, item.D+item.U)
}
}
func (h *Handler) enforcePeerShareFlowLimit(shareID int64) {
if h == nil || h.repo == nil || shareID <= 0 {
return
}
if err := h.cleanupPeerShareRuntimes(shareID); err != nil {
log.Printf("peer share quota cleanup pending share_id=%d err=%v", shareID, err)
}
}
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) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, now int64) error {
if h == nil || h.repo == nil {
return errors.New("invalid flow policy context")
}
if userID <= 0 || tunnelID <= 0 {
return nil
}
user, err := h.repo.GetUserByID(userID)
if err != nil {
return err
}
if user == nil {
return errors.New("用户不存在")
}
if user.Status != 1 {
return errors.New("账号已禁用")
}
if user.ExpTime > 0 && user.ExpTime <= now {
return errors.New("账号已过期")
}
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
current := user.InFlow + user.OutFlow
if flowLimit < current {
return errors.New("流量已超额,禁止开启转发")
}
if err := h.ensureUserForwardAllowedByQuota(userID, now); err != nil {
return err
}
if user.Num > 0 {
currentForwardCount, err := h.repo.CountActiveForwardsByUser(userID)
if err != nil {
return err
}
if currentForwardCount >= int64(user.Num) {
return errors.New("转发数量已达上限")
}
}
userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(userID, tunnelID)
if err != nil {
return err
}
if userTunnelID <= 0 {
return nil
}
policy, err := h.getUserTunnelPolicy(userTunnelID)
if err != nil {
return err
}
if policy == nil {
return nil
}
if policy.Status != 1 {
return errors.New("该隧道已禁用")
}
if policy.ExpTime > 0 && policy.ExpTime <= now {
return errors.New("该隧道已过期")
}
utFlowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
utCurrent := policy.InFlow + policy.OutFlow
if utCurrent >= utFlowLimit {
return errors.New("该隧道流量已超额,禁止开启转发")
}
if policy.Num > 0 {
currentTunnelForwardCount, err := h.repo.CountActiveForwardsByUserTunnel(userID, tunnelID)
if err != nil {
return err
}
if currentTunnelForwardCount >= int64(policy.Num) {
return errors.New("该隧道转发数量已达上限")
}
}
return nil
}
func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
user, err := h.repo.GetUserByID(userID)
if err != nil || user == nil {
return false
}
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
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 := flowLimitBytes(policy.Flow, policy.FlowMiB)
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
}
ut, err := h.repo.GetUserTunnelByID(userTunnelID)
if err != nil {
return nil, err
}
if ut == nil {
return nil, nil
}
return &userTunnelPolicy{
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
ExpTime: ut.ExpTime, Status: ut.Status, Num: ut.Num,
}, 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.UpdateForwardStatus(forward.ID, 0, now)
}
}
func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error) {
return h.repo.ListActiveForwardsByUser(userID)
}
func (h *Handler) listActiveForwardsByUserTunnel(userID int64, tunnelID int64) ([]forwardRecord, error) {
return h.repo.ListActiveForwardsByUserTunnel(userID, tunnelID)
}
func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) {
if h == nil || h.repo == nil || nodeID <= 0 {
return
}
if strings.TrimSpace(rawConfig) == "" {
return
}
var snapshot gostConfigSnapshot
if err := json.Unmarshal([]byte(rawConfig), &snapshot); err != nil {
return
}
protection, err := h.loadForwardServiceProtection(nodeID)
if err != nil {
return
}
h.cleanOrphanedServicesWithProtection(nodeID, snapshot.Services, protection)
// Dependencies are sent before services. A pending shared reservation may
// therefore have chains/limiters that are not referenced in this snapshot yet.
if protection.unbound {
return
}
chainsInUse := make(map[string]struct{})
limitersInUse := make(map[string]struct{})
for _, service := range snapshot.Services {
if service.Handler != nil {
if chain := strings.TrimSpace(service.Handler.Chain); chain != "" {
chainsInUse[chain] = struct{}{}
}
}
for _, limiter := range strings.Split(service.Limiter, ",") {
if limiter = strings.TrimSpace(limiter); limiter != "" {
limitersInUse[limiter] = struct{}{}
}
}
}
// Keep dependencies referenced by the reported services, even when their
// IDs belong to a different panel. Orphan dependencies can be collected on
// the next report after their services have actually disappeared.
h.cleanOrphanedChains(nodeID, snapshot.Chains, chainsInUse)
h.cleanOrphanedLimiters(nodeID, snapshot.Limiters, limitersInUse)
}
type forwardServiceProtection struct {
sharedNames map[string]struct{}
unbound bool
}
// Shared forward IDs belong to another panel and need not exist in our forward
// table. Use the same node-scoped ownership check for config and flow reports.
func (h *Handler) loadForwardServiceProtection(nodeID int64) (forwardServiceProtection, error) {
protection := forwardServiceProtection{sharedNames: make(map[string]struct{})}
// Read names and pending bindings in one snapshot, so a concurrent bind
// cannot fall between two queries and disappear from both protections.
runtimes, err := h.repo.ListActiveForwardPeerShareRuntimesByNode(nodeID)
if err != nil {
return protection, err
}
minUpdatedTime := time.Now().Add(-10 * time.Minute).UnixMilli()
for _, runtime := range runtimes {
serviceName := normalizeForwardRuntimeServiceName(runtime.ServiceName)
if serviceName == "" {
if runtime.Applied == 0 && runtime.UpdatedTime >= minUpdatedTime {
protection.unbound = true
}
continue
}
protection.sharedNames[serviceName] = struct{}{}
}
resources, err := h.repo.ListPeerShareResourcesByNode(nodeID)
if err != nil {
return protection, err
}
for _, resource := range resources {
if base := normalizeForwardRuntimeServiceName(resource.LegacyServiceBase); base != "" {
protection.sharedNames[base] = struct{}{}
}
}
return protection, nil
}
func (p forwardServiceProtection) preserves(serviceName string) bool {
_, shared := p.sharedNames[normalizeForwardRuntimeServiceName(serviceName)]
return shared || p.unbound
}
func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) {
if h == nil || h.repo == nil || nodeID <= 0 {
return
}
protection, err := h.loadForwardServiceProtection(nodeID)
if err != nil {
// A failed ownership lookup must never authorize deletion.
return
}
h.cleanOrphanedServicesWithProtection(nodeID, services, protection)
}
func (h *Handler) cleanOrphanedServicesWithProtection(nodeID int64, services []namedConfigItem, protection forwardServiceProtection) {
for _, item := range services {
name := strings.TrimSpace(item.Name)
if name == "" || name == "web_api" {
continue
}
if strings.HasPrefix(name, "fed_svc_") || strings.HasPrefix(name, "peer-share-") {
continue
}
if _, ok := protection.sharedNames[normalizeForwardRuntimeServiceName(name)]; ok {
continue
}
parts := strings.Split(name, "_")
if len(parts) == 2 && parts[0] == "tunnel" {
tunnelID, err := strconv.ParseInt(parts[1], 10, 64)
if err == nil && tunnelID > 0 && !h.tunnelExists(tunnelID) {
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name}}, false, true)
}
continue
}
if _, _, _, ok := parseFlowServiceIDs(name); ok {
h.deleteOrphanedForwardService(nodeID, name, protection)
continue
}
suffix := parts[len(parts)-1]
switch suffix {
case "tls", "kcp", "wss", "mtls", "mwss", "mtcp":
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 {
tunnelID, err := strconv.ParseInt(parts[0], 10, 64)
if err == nil && tunnelID > 0 && !h.tunnelExists(tunnelID) {
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name}}, false, true)
}
continue
}
}
}
}
func (h *Handler) cleanOrphanedChains(nodeID int64, chains []namedConfigItem, inUse map[string]struct{}) {
for _, item := range chains {
name := strings.TrimSpace(item.Name)
if name == "" || strings.HasPrefix(name, "fed_chain_") || strings.HasPrefix(name, "peer-share-") {
continue
}
if _, ok := inUse[name]; ok {
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 {
continue
}
exists, err := h.repo.TunnelExists(tunnelID)
if err != nil || exists {
continue
}
_, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": name}, false, true)
}
}
func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem, inUse map[string]struct{}) {
for _, item := range limiters {
name := strings.TrimSpace(item.Name)
if name == "" || strings.HasPrefix(name, "peer-share-") {
continue
}
if _, ok := inUse[name]; ok {
continue
}
exists, err := h.lookupSpeedLimiter(name)
if err != nil || exists {
continue
}
_, _ = h.sendNodeCommand(nodeID, "DeleteLimiters", map[string]interface{}{"limiter": name}, false, true)
}
}
func (h *Handler) tunnelExists(tunnelID int64) bool {
ok, _ := h.repo.TunnelExists(tunnelID)
return ok
}
func (h *Handler) forwardExists(forwardID int64) bool {
ok, _ := h.repo.ForwardExists(forwardID)
return ok
}
func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName string) {
h.sendDeleteOrphanedForwardServices(nodeID, []string{serviceName})
}
func (h *Handler) sendDeleteOrphanedForwardServices(nodeID int64, serviceNames []string) {
if h == nil || h.repo == nil || nodeID <= 0 || len(serviceNames) == 0 {
return
}
protection, err := h.loadForwardServiceProtection(nodeID)
if err != nil {
return
}
seen := make(map[string]struct{}, len(serviceNames))
for _, serviceName := range serviceNames {
serviceName = normalizeForwardRuntimeServiceName(serviceName)
if _, ok := seen[serviceName]; ok {
continue
}
seen[serviceName] = struct{}{}
h.deleteOrphanedForwardService(nodeID, serviceName, protection)
}
}
func (h *Handler) deleteOrphanedForwardService(nodeID int64, serviceName string, protection forwardServiceProtection) {
forwardID, _, _, ok := parseFlowServiceIDs(serviceName)
if !ok || protection.preserves(serviceName) {
return
}
parts := strings.Split(serviceName, "_")
base := parts[0] + "_" + parts[1] + "_" + parts[2]
// Parsing accepts legacy suffixes, while deletion targets the entire base
// family. Verify the actual targets cannot include a protected share.
if protection.preserves(base) {
return
}
// Batch metadata can be missing after a read failure, or stale by the time
// cleanup runs. Confirm absence before issuing a destructive command.
exists, err := h.repo.ForwardExists(forwardID)
if err != nil || exists {
return
}
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{
"services": buildForwardServiceDeleteNames([]string{base}),
}, false, true)
}
func (h *Handler) speedLimiterExists(name string) bool {
exists, _ := h.lookupSpeedLimiter(name)
return exists
}
func (h *Handler) lookupSpeedLimiter(name string) (bool, error) {
name = strings.TrimSpace(name)
if name == "" {
return false, nil
}
const forwardRulePrefix = "rule_traffic_limit_"
if strings.HasPrefix(name, forwardRulePrefix) {
forwardID, err := strconv.ParseInt(strings.TrimPrefix(name, forwardRulePrefix), 10, 64)
if err != nil || forwardID <= 0 {
return false, nil
}
forward, err := h.getForwardRecord(forwardID)
if errors.Is(err, errForwardNotFound) {
return false, nil
}
return forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0, err
}
id, err := strconv.ParseInt(name, 10, 64)
if err != nil || id <= 0 {
return false, nil
}
return h.repo.SpeedLimitExists(id)
}