fix(tunnel): auto-update entry node IP on every tunnel update

This commit is contained in:
sagit
2026-02-13 01:24:48 +00:00
parent 69f62188cf
commit 9ed875b7ef
6 changed files with 56 additions and 41 deletions
+2 -5
View File
@@ -7,16 +7,15 @@ services:
driver: json-file driver: json-file
options: options:
max-size: "20m" max-size: "20m"
max-file: "3"
environment: environment:
DB_PATH: /app/data/gost.db DB_PATH: /app/data/gost.db
JWT_SECRET: ${JWT_SECRET} JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
SERVER_ADDR: :6365 SERVER_ADDR: :6365
TZ: Asia/Shanghai TZ: Asia/Shanghai
ports: ports:
- "${BACKEND_PORT}:6365" - "${BACKEND_PORT}:6365"
volumes: volumes:
- backend_logs:/app/logs
- sqlite_data:/app/data - sqlite_data:/app/data
networks: networks:
- gost-network - gost-network
@@ -37,6 +36,7 @@ services:
driver: json-file driver: json-file
options: options:
max-size: "20m" max-size: "20m"
max-file: "3"
ports: ports:
- "${FRONTEND_PORT}:80" - "${FRONTEND_PORT}:80"
depends_on: depends_on:
@@ -50,9 +50,6 @@ volumes:
sqlite_data: sqlite_data:
name: sqlite_data name: sqlite_data
driver: local driver: local
backend_logs:
name: backend_logs
driver: local
networks: networks:
+2 -5
View File
@@ -7,16 +7,15 @@ services:
driver: json-file driver: json-file
options: options:
max-size: "20m" max-size: "20m"
max-file: "3"
environment: environment:
DB_PATH: /app/data/gost.db DB_PATH: /app/data/gost.db
JWT_SECRET: ${JWT_SECRET} JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
SERVER_ADDR: :6365 SERVER_ADDR: :6365
TZ: Asia/Shanghai TZ: Asia/Shanghai
ports: ports:
- "${BACKEND_PORT}:6365" - "${BACKEND_PORT}:6365"
volumes: volumes:
- backend_logs:/app/logs
- sqlite_data:/app/data - sqlite_data:/app/data
networks: networks:
- gost-network - gost-network
@@ -37,6 +36,7 @@ services:
driver: json-file driver: json-file
options: options:
max-size: "20m" max-size: "20m"
max-file: "3"
ports: ports:
- "${FRONTEND_PORT}:80" - "${FRONTEND_PORT}:80"
depends_on: depends_on:
@@ -50,9 +50,6 @@ volumes:
sqlite_data: sqlite_data:
name: sqlite_data name: sqlite_data
driver: local driver: local
backend_logs:
name: backend_logs
driver: local
networks: networks:
-2
View File
@@ -6,7 +6,6 @@ type Config struct {
Addr string Addr string
DBPath string DBPath string
JWTSecret string JWTSecret string
LogDir string
} }
func FromEnv() Config { func FromEnv() Config {
@@ -14,7 +13,6 @@ func FromEnv() Config {
Addr: getEnv("SERVER_ADDR", ":6365"), Addr: getEnv("SERVER_ADDR", ":6365"),
DBPath: getEnv("DB_PATH", "/app/data/gost.db"), DBPath: getEnv("DB_PATH", "/app/data/gost.db"),
JWTSecret: getEnv("JWT_SECRET", ""), JWTSecret: getEnv("JWT_SECRET", ""),
LogDir: getEnv("LOG_DIR", "/app/logs"),
} }
return cfg return cfg
@@ -700,21 +700,20 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
defer tx.Rollback() defer tx.Rollback()
now := time.Now().UnixMilli() now := time.Now().UnixMilli()
res, err := tx.Exec(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?)`, var tunnelID int64
err = tx.QueryRow(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?) RETURNING id`,
fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort), fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort),
tunnelType, tunnelType,
req.Protocol, req.Protocol,
now, now,
now, now,
"", "",
) ).Scan(&tunnelID)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
tunnelID, _ := res.LastInsertId()
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, 1, ?, ?, 'fifo', 0, ?)`, _, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, 1, ?, ?, 'fifo', 0, ?)`,
tunnelID, tunnelID,
share.NodeID, share.NodeID,
+20 -18
View File
@@ -559,13 +559,13 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
} }
} }
res, err := tx.Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, var tunnelID int64
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx) err = tx.QueryRow(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) RETURNING id`,
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx).Scan(&tunnelID)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
tunnelID, _ := res.LastInsertId()
runtimeState.TunnelID = tunnelID runtimeState.TunnelID = tunnelID
var federationBindings []sqlite.FederationTunnelBinding var federationBindings []sqlite.FederationTunnelBinding
var federationReleaseRefs []federationRuntimeReleaseRef var federationReleaseRefs []federationRuntimeReleaseRef
@@ -688,6 +688,9 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
return return
} }
runtimeState.TunnelID = id runtimeState.TunnelID = id
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes)
var federationBindings []sqlite.FederationTunnelBinding var federationBindings []sqlite.FederationTunnelBinding
var federationReleaseRefs []federationRuntimeReleaseRef var federationReleaseRefs []federationRuntimeReleaseRef
if typeVal == 2 { if typeVal == 2 {
@@ -700,7 +703,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
applyTunnelPortsToRequest(req, runtimeState) applyTunnelPortsToRequest(req, runtimeState)
_, err = tx.Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, updated_time=? WHERE id=?`, _, err = tx.Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, updated_time=? WHERE id=?`,
asString(req["name"]), typeVal, asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(asString(req["inIp"])), now, id) asString(req["name"]), typeVal, asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(inIp), now, 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
@@ -1119,15 +1122,15 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
return return
} }
defer func() { _ = tx.Rollback() }() defer func() { _ = tx.Rollback() }()
res, err := tx.Exec(` var forwardID int64
err = tx.QueryRow(`
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?) VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?) RETURNING id
`, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx) `, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx).Scan(&forwardID)
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
forwardID, _ := res.LastInsertId()
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
for _, nodeID := range entryNodes { for _, nodeID := range entryNodes {
_, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port) _, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port)
@@ -1587,13 +1590,13 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
} }
now := time.Now().UnixMilli() now := time.Now().UnixMilli()
speed := asInt(req["speed"], 100) 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(?, ?, ?, ?, ?, ?, ?)`, var id int64
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1)) err := h.repo.DB().QueryRow(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?) RETURNING id`,
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1)).Scan(&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
} }
id, _ := res.LastInsertId()
_ = h.sendLimiterConfig(id, speed, tunnelID) _ = h.sendLimiterConfig(id, speed, tunnelID)
response.WriteJSON(w, response.OKEmpty()) response.WriteJSON(w, response.OKEmpty())
} }
@@ -1687,7 +1690,7 @@ func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) {
defer func() { _ = tx.Rollback() }() defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID) _, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID)
for _, tid := range req.TunnelIDs { for _, tid := range req.TunnelIDs {
_, _ = tx.Exec(`INSERT OR IGNORE INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, req.GroupID, tid, time.Now().UnixMilli()) _, _ = tx.Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?) ON CONFLICT(tunnel_group_id, tunnel_id) DO NOTHING`, req.GroupID, tid, time.Now().UnixMilli())
} }
if err := tx.Commit(); err != nil { if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
@@ -1714,7 +1717,7 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) {
defer func() { _ = tx.Rollback() }() defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID) _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID)
for _, uid := range req.UserIDs { for _, uid := range req.UserIDs {
_, _ = tx.Exec(`INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)`, req.GroupID, uid, time.Now().UnixMilli()) _, _ = tx.Exec(`INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT(user_group_id, user_id) DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli())
} }
if err := tx.Commit(); err != nil { if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
@@ -1733,7 +1736,7 @@ func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.ErrDefault("请求参数错误")) response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return return
} }
_, err := h.repo.DB().Exec(`INSERT OR IGNORE INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli()) _, err := h.repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?) ON CONFLICT(user_group_id, tunnel_group_id) DO NOTHING`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
if err != nil { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
@@ -1835,7 +1838,7 @@ func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error {
if created { if created {
createdByGroup = 1 createdByGroup = 1
} }
_, _ = db.Exec(`INSERT OR IGNORE INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?)`, _, _ = db.Exec(`INSERT INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?) ON CONFLICT(user_group_id, tunnel_group_id, user_tunnel_id) DO NOTHING`,
userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli()) userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli())
} }
} }
@@ -1882,12 +1885,11 @@ func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, err
if err := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset); err != nil { if err := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset); err != nil {
return 0, false, err return 0, false, err
} }
res, err := db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1)`, err = db.QueryRow(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1) RETURNING id`,
userID, tunnelID, num, flow, flowReset, expTime) userID, tunnelID, num, flow, flowReset, expTime).Scan(&id)
if err != nil { if err != nil {
return 0, false, err return 0, false, err
} }
id, _ = res.LastInsertId()
return id, true, nil return id, true, nil
} }
+29 -7
View File
@@ -361,7 +361,7 @@ func (r *Repository) GetUserPackageForwards(userID int64) ([]UserForwardDetail,
} }
rows, err := r.db.Query(` rows, err := r.db.Query(`
SELECT f.id, f.name, f.tunnel_id, t.name, f.remote_addr, f.in_flow, f.out_flow, f.status, f.created_time SELECT f.id, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, f.in_flow, f.out_flow, f.status, f.created_time
FROM forward f FROM forward f
LEFT JOIN tunnel t ON t.id = f.tunnel_id LEFT JOIN tunnel t ON t.id = f.tunnel_id
WHERE f.user_id = ? WHERE f.user_id = ?
@@ -680,7 +680,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
} }
rows, err := r.db.Query(` rows, err := r.db.Query(`
SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, t.name, f.remote_addr, f.strategy, SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, f.strategy,
f.in_flow, f.out_flow, f.created_time, f.status, f.inx f.in_flow, f.out_flow, f.created_time, f.status, f.inx
FROM forward f FROM forward f
LEFT JOIN tunnel t ON t.id = f.tunnel_id LEFT JOIN tunnel t ON t.id = f.tunnel_id
@@ -858,7 +858,7 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
} }
chainRows, err := r.db.Query(` chainRows, err := r.db.Query(`
SELECT tunnel_id, chain_type, node_id, protocol, strategy, COALESCE(inx, 0) SELECT tunnel_id, CAST(chain_type AS INTEGER), node_id, protocol, strategy, COALESCE(inx, 0)
FROM chain_tunnel FROM chain_tunnel
ORDER BY tunnel_id ASC, chain_type ASC, inx ASC, id ASC ORDER BY tunnel_id ASC, chain_type ASC, inx ASC, id ASC
`) `)
@@ -1258,21 +1258,41 @@ func bootstrapSchema(db *sql.DB) error {
return nil return nil
} }
const currentSchemaVersion = 1
func getSchemaVersion(db *sql.DB) int {
_, _ = db.Exec(`CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL DEFAULT 0)`)
var v int
if err := db.QueryRow(`SELECT version FROM schema_version LIMIT 1`).Scan(&v); err != nil {
_, _ = db.Exec(`INSERT INTO schema_version(version) VALUES(0)`)
return 0
}
return v
}
func setSchemaVersion(db *sql.DB, v int) {
_, _ = db.Exec(`UPDATE schema_version SET version = ?`, v)
}
func migrateSchema(db *sql.DB) error { func migrateSchema(db *sql.DB) error {
if db == nil { if db == nil {
return errors.New("nil db") return errors.New("nil db")
} }
ver := getSchemaVersion(db)
if ver >= currentSchemaVersion {
return nil
}
ensureColumn := func(table, col, typ string) { ensureColumn := func(table, col, typ string) {
var dummy interface{} var dummy interface{}
err := db.QueryRow(fmt.Sprintf("SELECT %s FROM %s LIMIT 1", col, table)).Scan(&dummy) err := db.QueryRow(fmt.Sprintf("SELECT %s FROM %s LIMIT 1", col, table)).Scan(&dummy)
if err == nil || errors.Is(err, sql.ErrNoRows) { if err == nil || errors.Is(err, sql.ErrNoRows) {
return return
} }
if strings.Contains(err.Error(), "no such column") { // Column likely missing (SQLite: "no such column", PG: "does not exist", etc.)
if _, alterErr := db.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, col, typ)); alterErr != nil { if _, alterErr := db.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, col, typ)); alterErr != nil {
log.Printf("failed to add column %s to %s: %v", col, table, alterErr) log.Printf("failed to add column %s to %s: %v", col, table, alterErr)
}
} }
} }
@@ -1306,6 +1326,8 @@ func migrateSchema(db *sql.DB) error {
ensureColumn(table, col, typ) ensureColumn(table, col, typ)
} }
} }
setSchemaVersion(db, currentSchemaVersion)
return nil return nil
} }