Compare commits

...

11 Commits

Author SHA1 Message Date
sagit abb8591b11 fix(backend): backfill legacy node dual-stack columns in sqlite migration 2026-02-10 04:06:00 +00:00
sagit 2acee481f0 fix(backend): backfill legacy inx columns during sqlite migration 2026-02-10 03:14:08 +00:00
sagit 76f443f900 Merge pull request #62 from Sagit-chu/update-tz-mirror
feat: add Shanghai timezone to docker-compose and update github mirror
2026-02-09 18:52:59 +08:00
sagit 47b1663938 feat: add Shanghai timezone to docker-compose and update github mirror 2026-02-09 10:47:13 +00:00
sagit 406f5bb380 Merge pull request #61 from Sagit-chu/opencode/calm-orchid
fix: limit speed
2026-02-09 17:13:40 +08:00
sagit 85ea6c17a4 Merge branch 'main' into opencode/calm-orchid 2026-02-09 17:12:38 +08:00
sagit 3420dc5460 fix(backend): sync limiter on association instead of connection
Reverted the full sync on connection hook. Instead, ensureLimiterOnNode is called within syncForwardServices to push limiter configuration immediately before pushing the service configuration that references it.
2026-02-09 09:06:59 +00:00
sagit 3d7a0b697d feat(backend): sync limiters on agent connect
Implemented full sync of speed limit configurations when an Agent connects via WebSocket. This ensures that even fresh or restarted agents receive the necessary limiter configurations.
2026-02-09 08:49:47 +00:00
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
8 changed files with 255 additions and 23 deletions
+1
View File
@@ -12,6 +12,7 @@ services:
JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
SERVER_ADDR: :6365
TZ: Asia/Shanghai
ports:
- "${BACKEND_PORT}:6365"
volumes:
+1
View File
@@ -12,6 +12,7 @@ services:
JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
SERVER_ADDR: :6365
TZ: Asia/Shanghai
ports:
- "${BACKEND_PORT}:6365"
volumes:
@@ -234,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, *int, error) {
row := h.repo.DB().QueryRow(`
SELECT ut.id, sl.speed
SELECT ut.id, sl.id, sl.speed
FROM user_tunnel ut
LEFT JOIN speed_limit sl ON sl.id = ut.speed_id
WHERE ut.user_id = ? AND ut.tunnel_id = ?
@@ -244,19 +244,21 @@ func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *i
LIMIT 1
`, userID, tunnelID)
var userTunnelID int64
var limiterID sql.NullInt64
var speed sql.NullInt64
err := row.Scan(&userTunnelID, &speed)
err := row.Scan(&userTunnelID, &limiterID, &speed)
if err != nil {
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 {
return userTunnelID, nil, nil
if !limiterID.Valid || limiterID.Int64 <= 0 {
return userTunnelID, nil, nil, nil
}
v := int(speed.Int64)
return userTunnelID, &v, nil
v := limiterID.Int64
s := int(speed.Int64)
return userTunnelID, &v, &s, nil
}
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("转发入口端口不存在")
}
userTunnelID, limiter, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
userTunnelID, limiterID, speed, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil {
return err
}
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
for _, fp := range ports {
if limiterID != nil && speed != nil {
h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed)
}
node, err := h.getNodeRecord(fp.NodeID)
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)
@@ -362,7 +368,7 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str
if len(ports) == 0 {
return nil
}
userTunnelID, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil {
return err
}
@@ -1003,7 +1009,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)
@@ -1044,11 +1050,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 {
// 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)
if limiterID != nil && *limiterID > 0 {
service["limiter"] = strconv.FormatInt(*limiterID, 10)
}
services = append(services, service)
}
@@ -1111,3 +1114,49 @@ 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
}
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
@@ -1469,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())
}
@@ -1496,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())
}
@@ -1510,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())
}
@@ -123,6 +123,11 @@ func Open(path string) (*Repository, error) {
return nil, err
}
if err := migrateSchema(db); err != nil {
_ = db.Close()
return nil, err
}
return &Repository{db: db}, nil
}
@@ -1176,6 +1181,54 @@ func bootstrapSchema(db *sql.DB) error {
return nil
}
func migrateSchema(db *sql.DB) error {
if db == nil {
return errors.New("nil db")
}
ensureColumn := func(table, col, typ string) error {
var dummy interface{}
err := db.QueryRow(fmt.Sprintf("SELECT %s FROM %s LIMIT 1", col, table)).Scan(&dummy)
if err == nil || errors.Is(err, sql.ErrNoRows) {
return nil
}
if !strings.Contains(err.Error(), "no such column") {
return nil
}
if _, err := db.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, col, typ)); err != nil {
return fmt.Errorf("add %s.%s: %w", table, col, err)
}
return nil
}
columnsByTable := map[string]map[string]string{
"node": {
"server_ip_v4": "VARCHAR(100)",
"server_ip_v6": "VARCHAR(100)",
"inx": "INTEGER NOT NULL DEFAULT 0",
},
"tunnel": {
"inx": "INTEGER NOT NULL DEFAULT 0",
},
"forward": {
"inx": "INTEGER NOT NULL DEFAULT 0",
},
"chain_tunnel": {
"inx": "INTEGER",
},
}
for table, cols := range columnsByTable {
for col, typ := range cols {
if err := ensureColumn(table, col, typ); err != nil {
return err
}
}
}
return nil
}
var osMkdirAll = func(path string) error {
return os.MkdirAll(path, 0o755)
}
@@ -2,6 +2,7 @@ package contract_test
import (
"bytes"
"database/sql"
"encoding/json"
"io"
"net/http"
@@ -17,6 +18,8 @@ import (
"go-backend/internal/http/handler"
"go-backend/internal/http/response"
"go-backend/internal/store/sqlite"
_ "modernc.org/sqlite"
)
func TestCaptchaVerifyLoginContract(t *testing.T) {
@@ -213,3 +216,117 @@ func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.
h := handler.New(repo, jwtSecret)
return httpserver.NewRouter(h, jwtSecret), repo
}
func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "legacy-2.0.7-beta.db")
legacyDB, err := sql.Open("sqlite", dbPath)
if err != nil {
t.Fatalf("open legacy sqlite: %v", err)
}
t.Cleanup(func() {
_ = legacyDB.Close()
})
if _, err := legacyDB.Exec(`
CREATE TABLE IF NOT EXISTS node (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name VARCHAR(100) NOT NULL,
secret VARCHAR(100) NOT NULL,
server_ip VARCHAR(100) NOT NULL,
port TEXT NOT NULL,
interface_name VARCHAR(200),
version VARCHAR(100),
http INTEGER NOT NULL DEFAULT 0,
tls INTEGER NOT NULL DEFAULT 0,
socks INTEGER NOT NULL DEFAULT 0,
created_time INTEGER NOT NULL,
updated_time INTEGER,
status INTEGER NOT NULL,
tcp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
udp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]'
)
`); err != nil {
t.Fatalf("create legacy node table: %v", err)
}
if _, err := legacyDB.Exec(`
CREATE TABLE IF NOT EXISTS tunnel (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name VARCHAR(100) NOT NULL,
traffic_ratio REAL NOT NULL DEFAULT 1.0,
type INTEGER NOT NULL,
protocol VARCHAR(10) NOT NULL DEFAULT 'tls',
flow INTEGER NOT NULL,
created_time INTEGER NOT NULL,
updated_time INTEGER NOT NULL,
status INTEGER NOT NULL,
in_ip TEXT
)
`); err != nil {
t.Fatalf("create legacy tunnel table: %v", err)
}
now := time.Now().UnixMilli()
if _, err := legacyDB.Exec(`
INSERT INTO node(name, secret, server_ip, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "legacy-node", "legacy-secret", "10.10.0.1", "10000-10010", "eth0", "v-old", 1, 1, 1, now, now, 1, "[::]", "[::]"); err != nil {
t.Fatalf("seed legacy node row: %v", err)
}
repo, err := sqlite.Open(dbPath)
if err != nil {
t.Fatalf("open migrated sqlite: %v", err)
}
t.Cleanup(func() {
_ = repo.Close()
})
nodes, err := repo.ListNodes()
if err != nil {
t.Fatalf("list nodes after migration: %v", err)
}
if len(nodes) != 1 {
t.Fatalf("expected 1 node after migration, got %d", len(nodes))
}
columns := readTableColumns(t, repo.DB(), "node")
for _, required := range []string{"server_ip_v4", "server_ip_v6", "inx"} {
if !columns[required] {
t.Fatalf("expected node column %q to exist after migration", required)
}
}
tunnelColumns := readTableColumns(t, repo.DB(), "tunnel")
if !tunnelColumns["inx"] {
t.Fatalf("expected tunnel column %q to exist after migration", "inx")
}
}
func readTableColumns(t *testing.T, db *sql.DB, table string) map[string]bool {
t.Helper()
rows, err := db.Query("PRAGMA table_info(" + table + ")")
if err != nil {
t.Fatalf("inspect %s columns: %v", table, err)
}
defer rows.Close()
columns := map[string]bool{}
for rows.Next() {
var cid, notNull, pk int
var name, typ string
var defaultValue sql.NullString
if err := rows.Scan(&cid, &name, &typ, &notNull, &defaultValue, &pk); err != nil {
t.Fatalf("scan %s pragma row: %v", table, err)
}
columns[name] = true
}
if err := rows.Err(); err != nil {
t.Fatalf("iterate %s pragma rows: %v", table, err)
}
return columns
}
+1 -1
View File
@@ -28,7 +28,7 @@ COUNTRY=$(curl -s https://ipinfo.io/country)
maybe_proxy_url() {
local url="$1"
if [ "$COUNTRY" = "CN" ]; then
echo "https://ghfast.top/${url}"
echo "https://gcode.hostcentral.cc/${url}"
else
echo "$url"
fi
+1 -1
View File
@@ -15,7 +15,7 @@ COUNTRY=$(curl -s https://ipinfo.io/country)
maybe_proxy_url() {
local url="$1"
if [ "$COUNTRY" = "CN" ]; then
echo "https://ghfast.top/${url}"
echo "https://gcode.hostcentral.cc/${url}"
else
echo "$url"
fi