feat(postgres): add postgres backend support and migration docs

This commit is contained in:
sagit
2026-02-12 06:38:25 +00:00
parent 69f62188cf
commit cedcaebd1f
16 changed files with 819 additions and 53 deletions
+3
View File
@@ -262,3 +262,6 @@ sql/
!go-backend/internal/store/sqlite/sql/
!go-backend/internal/store/sqlite/sql/schema.sql
!go-backend/internal/store/sqlite/sql/data.sql
!go-backend/internal/store/postgres/sql/
!go-backend/internal/store/postgres/sql/schema.sql
!go-backend/internal/store/postgres/sql/data.sql
+64
View File
@@ -66,6 +66,70 @@ curl -L https://github.com/Sagit-chu/flux-panel/releases/download/2.1.0/panel_in
curl -L https://github.com/Sagit-chu/flux-panel/releases/download/2.1.0/install.sh -o install.sh && chmod +x install.sh && ./install.sh
```
#### PostgreSQL 部署(Docker Compose)
`docker-compose-v4.yml` / `docker-compose-v6.yml` 已包含 PostgreSQL 服务。默认仍使用 SQLite,切换到 PostgreSQL 只需要配置环境变量。
1) 在 `docker-compose` 同目录创建或修改 `.env`:
```bash
JWT_SECRET=replace_with_your_secret
BACKEND_PORT=6365
FRONTEND_PORT=6366
DB_TYPE=postgres
DATABASE_URL=postgres://flux_panel:replace_with_strong_password@postgres:5432/flux_panel?sslmode=disable
POSTGRES_DB=flux_panel
POSTGRES_USER=flux_panel
POSTGRES_PASSWORD=replace_with_strong_password
```
2) 启动(IPv4/IPv6 二选一):
```bash
docker compose -f docker-compose-v4.yml up -d
```
```bash
docker compose -f docker-compose-v6.yml up -d
```
3) 如果你想继续使用 SQLite,保留 `DB_TYPE=sqlite`(或不设置 `DB_TYPE`)即可。
#### 从 SQLite 迁移到 PostgreSQL
以下示例基于 Docker Volume `sqlite_data`(项目默认配置)与 `pgloader`:
1) 停止服务并备份 SQLite 数据:
```bash
docker compose -f docker-compose-v4.yml down
docker run --rm -v sqlite_data:/data -v "$(pwd)":/backup alpine sh -c "cp /data/gost.db /backup/gost.db.bak"
```
2) 仅启动 PostgreSQL:
```bash
docker compose -f docker-compose-v4.yml up -d postgres
```
3) 使用 `pgloader` 迁移:
```bash
docker run --rm --network gost-network -v sqlite_data:/sqlite dimitri/pgloader:latest pgloader /sqlite/gost.db postgresql://flux_panel:replace_with_strong_password@postgres:5432/flux_panel
```
4) 切换后端到 PostgreSQL 并启动:
```bash
export DB_TYPE=postgres
export DATABASE_URL="postgres://flux_panel:replace_with_strong_password@postgres:5432/flux_panel?sslmode=disable"
docker compose -f docker-compose-v4.yml up -d
```
5) 迁移完成后,登录面板检查用户、隧道、转发、节点数据是否正确。
#### 默认管理员账号
- **账号**: admin_user
+29
View File
@@ -8,7 +8,9 @@ services:
options:
max-size: "20m"
environment:
DB_TYPE: ${DB_TYPE:-sqlite}
DB_PATH: /app/data/gost.db
DATABASE_URL: ${DATABASE_URL:-}
JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
SERVER_ADDR: :6365
@@ -29,6 +31,30 @@ services:
retries: 5
start_period: 30s
postgres:
image: postgres:16-alpine
container_name: flux-panel-postgres
restart: unless-stopped
logging:
driver: json-file
options:
max-size: "20m"
environment:
POSTGRES_DB: ${POSTGRES_DB:-flux_panel}
POSTGRES_USER: ${POSTGRES_USER:-flux_panel}
POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:-flux_panel_change_me}
TZ: Asia/Shanghai
volumes:
- postgres_data:/var/lib/postgresql/data
networks:
- gost-network
healthcheck:
test: ["CMD-SHELL", "pg_isready -U ${POSTGRES_USER:-flux_panel} -d ${POSTGRES_DB:-flux_panel}"]
interval: 10s
timeout: 5s
retries: 10
start_period: 20s
frontend:
image: ghcr.io/sagit-chu/vite-frontend:${FLUX_VERSION:-latest}
container_name: vite-frontend
@@ -50,6 +76,9 @@ volumes:
sqlite_data:
name: sqlite_data
driver: local
postgres_data:
name: postgres_data
driver: local
backend_logs:
name: backend_logs
driver: local
+29
View File
@@ -8,7 +8,9 @@ services:
options:
max-size: "20m"
environment:
DB_TYPE: ${DB_TYPE:-sqlite}
DB_PATH: /app/data/gost.db
DATABASE_URL: ${DATABASE_URL:-}
JWT_SECRET: ${JWT_SECRET}
LOG_DIR: /app/logs
SERVER_ADDR: :6365
@@ -29,6 +31,30 @@ services:
retries: 5
start_period: 30s
postgres:
image: postgres:16-alpine
container_name: flux-panel-postgres
restart: unless-stopped
logging:
driver: json-file
options:
max-size: "20m"
environment:
POSTGRES_DB: ${POSTGRES_DB:-flux_panel}
POSTGRES_USER: ${POSTGRES_USER:-flux_panel}
POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:-flux_panel_change_me}
TZ: Asia/Shanghai
volumes:
- postgres_data:/var/lib/postgresql/data
networks:
- gost-network
healthcheck:
test: ["CMD-SHELL", "pg_isready -U ${POSTGRES_USER:-flux_panel} -d ${POSTGRES_DB:-flux_panel}"]
interval: 10s
timeout: 5s
retries: 10
start_period: 20s
frontend:
image: ghcr.io/sagit-chu/vite-frontend:${FLUX_VERSION:-latest}
container_name: vite-frontend
@@ -50,6 +76,9 @@ volumes:
sqlite_data:
name: sqlite_data
driver: local
postgres_data:
name: postgres_data
driver: local
backend_logs:
name: backend_logs
driver: local
+8 -1
View File
@@ -1,22 +1,29 @@
module go-backend
go 1.23.0
go 1.24.0
toolchain go1.24.4
require (
github.com/gorilla/websocket v1.5.3
github.com/jackc/pgx/v5 v5.7.3
modernc.org/sqlite v1.37.1
)
require (
github.com/dustin/go-humanize v1.0.1 // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/puddle/v2 v2.2.2 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
github.com/ncruces/go-strftime v0.1.9 // indirect
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
golang.org/x/crypto v0.31.0 // indirect
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect
golang.org/x/sync v0.17.0 // indirect
golang.org/x/sys v0.33.0 // indirect
golang.org/x/text v0.29.0 // indirect
modernc.org/libc v1.65.7 // indirect
modernc.org/mathutil v1.7.1 // indirect
modernc.org/memory v1.11.0 // indirect
+32 -6
View File
@@ -1,3 +1,6 @@
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
@@ -6,23 +9,46 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.7.3 h1:PO1wNKj/bTAwxSJnO1Z4Ai8j4magtqg2SLNjEDzcXQo=
github.com/jackc/pgx/v5 v5.7.3/go.mod h1:ncY89UGWxg82EykZUwSpUKEfccBGGYq1xjrOpsbsfGQ=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4=
github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk=
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
golang.org/x/crypto v0.31.0 h1:ihbySMvVjLAeSH1IbfcRTkD/iNscyz8rGzjF/E5hV6U=
golang.org/x/crypto v0.31.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM=
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8=
golang.org/x/mod v0.24.0 h1:ZfthKaKaT4NrhGVZHO1/WDTwGES4De8KtWO0SIbNJMU=
golang.org/x/mod v0.24.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww=
golang.org/x/sync v0.14.0 h1:woo0S4Yywslg6hp4eUFjTVOyKt0RookbpAHG4c1HmhQ=
golang.org/x/sync v0.14.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
golang.org/x/mod v0.27.0 h1:kb+q2PyFnEADO2IEF935ehFUXlWiNjJWtRNgBLSfbxQ=
golang.org/x/mod v0.27.0/go.mod h1:rWI627Fq0DEoudcK+MBkNkCe0EetEaDSwJJkCcjpazc=
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/tools v0.33.0 h1:4qz2S3zmRxbGIhDIAgjxvFutSvH5EfnsYrRBj0UI0bc=
golang.org/x/tools v0.33.0/go.mod h1:CIJMaWEY88juyUfo7UbgPqbC8rU2OqfAV1h2Qp0oMYI=
golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
golang.org/x/tools v0.36.0 h1:kWS0uv/zsvHEle1LbV5LE8QujrxB3wfQyxHfhOk0Qkg=
golang.org/x/tools v0.36.0/go.mod h1:WBDiHKJK8YgLHlcQPYQzNCkUxUypCaa5ZegCVutKm+s=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
modernc.org/cc/v4 v4.26.1 h1:+X5NtzVBn0KgsBCBe+xkDC7twLb/jNVj9FPgiwSQO3s=
modernc.org/cc/v4 v4.26.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU=
+19 -3
View File
@@ -4,6 +4,7 @@ import (
"context"
"fmt"
"net/http"
"strings"
"time"
"go-backend/internal/config"
@@ -20,9 +21,24 @@ type App struct {
}
func New(cfg config.Config) (*App, error) {
repo, err := sqlite.Open(cfg.DBPath)
if err != nil {
return nil, fmt.Errorf("open sqlite: %w", err)
var (
repo *sqlite.Repository
err error
)
switch strings.ToLower(strings.TrimSpace(cfg.DBType)) {
case "", "sqlite":
repo, err = sqlite.Open(cfg.DBPath)
if err != nil {
return nil, fmt.Errorf("open sqlite: %w", err)
}
case "postgres", "postgresql":
repo, err = sqlite.OpenPostgres(cfg.DatabaseURL)
if err != nil {
return nil, fmt.Errorf("open postgres: %w", err)
}
default:
return nil, fmt.Errorf("unsupported DB_TYPE %q", cfg.DBType)
}
h := handler.New(repo, cfg.JWTSecret)
+12 -8
View File
@@ -3,18 +3,22 @@ package config
import "os"
type Config struct {
Addr string
DBPath string
JWTSecret string
LogDir string
Addr string
DBType string
DBPath string
DatabaseURL string
JWTSecret string
LogDir string
}
func FromEnv() Config {
cfg := Config{
Addr: getEnv("SERVER_ADDR", ":6365"),
DBPath: getEnv("DB_PATH", "/app/data/gost.db"),
JWTSecret: getEnv("JWT_SECRET", ""),
LogDir: getEnv("LOG_DIR", "/app/logs"),
Addr: getEnv("SERVER_ADDR", ":6365"),
DBType: getEnv("DB_TYPE", "sqlite"),
DBPath: getEnv("DB_PATH", "/app/data/gost.db"),
DatabaseURL: getEnv("DATABASE_URL", ""),
JWTSecret: getEnv("JWT_SECRET", ""),
LogDir: getEnv("LOG_DIR", "/app/logs"),
}
return cfg
@@ -700,7 +700,7 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
defer tx.Rollback()
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, ?)`,
tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?)`,
fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort),
tunnelType,
req.Protocol,
@@ -713,8 +713,6 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
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, ?)`,
tunnelID,
share.NodeID,
+18 -21
View File
@@ -18,6 +18,7 @@ import (
"go-backend/internal/http/client"
"go-backend/internal/http/response"
"go-backend/internal/security"
"go-backend/internal/store"
"go-backend/internal/store/sqlite"
)
@@ -559,13 +560,12 @@ 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
tunnelID, _ := res.LastInsertId()
runtimeState.TunnelID = tunnelID
var federationBindings []sqlite.FederationTunnelBinding
var federationReleaseRefs []federationRuntimeReleaseRef
@@ -1119,7 +1119,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
return
}
defer func() { _ = tx.Rollback() }()
res, err := tx.Exec(`
forwardID, err := tx.ExecReturningID(`
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, ?)
`, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx)
@@ -1127,7 +1127,6 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
forwardID, _ := res.LastInsertId()
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
for _, nodeID := range entryNodes {
_, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port)
@@ -1587,13 +1586,12 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
}
now := time.Now().UnixMilli()
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(?, ?, ?, ?, ?, ?, ?)`,
id, err := h.repo.DB().ExecReturningID(`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())
}
@@ -1687,7 +1685,7 @@ func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) {
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID)
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 DO NOTHING`, req.GroupID, tid, time.Now().UnixMilli())
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
@@ -1714,7 +1712,7 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) {
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID)
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 DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli())
}
if err := tx.Commit(); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
@@ -1733,7 +1731,7 @@ func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
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 DO NOTHING`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -1835,7 +1833,7 @@ func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error {
if created {
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 DO NOTHING`,
userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli())
}
}
@@ -1866,7 +1864,7 @@ func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error {
return nil
}
func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, error) {
func ensureUserTunnelGrant(db *store.DB, userID, tunnelID int64) (int64, bool, error) {
var id int64
err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&id)
if err == nil {
@@ -1882,16 +1880,15 @@ 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 {
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)`,
id, err = db.ExecReturningID(`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)`,
userID, tunnelID, num, flow, flowReset, expTime)
if err != nil {
return 0, false, err
}
id, _ = res.LastInsertId()
return id, true, nil
}
func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error) {
func queryInt64List(db *store.DB, q string, args ...interface{}) ([]int64, error) {
rows, err := db.Query(q, args...)
if err != nil {
return nil, err
@@ -1908,7 +1905,7 @@ func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error)
return out, rows.Err()
}
func queryPairs(db *sql.DB, q string, args ...interface{}) ([][2]int64, error) {
func queryPairs(db *store.DB, q string, args ...interface{}) ([][2]int64, error) {
rows, err := db.Query(q, args...)
if err != nil {
return nil, err
@@ -1944,7 +1941,7 @@ type tunnelCreateState struct {
NodeIDList []int64
}
func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
state := &tunnelCreateState{
Type: tunnelType,
InNodes: make([]tunnelRuntimeNode, 0),
@@ -2403,7 +2400,7 @@ func (h *Handler) cleanupFederationRuntime(tunnelID int64) {
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID)
}
func replaceFederationTunnelBindingsTx(tx *sql.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error {
func replaceFederationTunnelBindingsTx(tx *store.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error {
if tx == nil {
return errors.New("database unavailable")
}
@@ -2713,7 +2710,7 @@ func pickNodeAddressV6(node *nodeRecord) string {
return strings.TrimSpace(node.ServerIP)
}
func isRemoteNodeTx(tx *sql.Tx, nodeID int64) (bool, error) {
func isRemoteNodeTx(tx *store.Tx, nodeID int64) (bool, error) {
if tx == nil {
return false, errors.New("database unavailable")
}
@@ -2730,7 +2727,7 @@ func isRemoteNodeTx(tx *sql.Tx, nodeID int64) (bool, error) {
return isRemote == 1, nil
}
func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
func pickNodePortTx(tx *store.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
if tx == nil {
return 0, errors.New("database unavailable")
}
@@ -2841,7 +2838,7 @@ func parsePortRangeSpec(input string) []int {
return out
}
func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error {
func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interface{}) error {
allocated := map[int64]int{}
inNodes := asMapSlice(req["inNodeId"])
for _, n := range inNodes {
@@ -3462,7 +3459,7 @@ func randomToken(n int) string {
return hex.EncodeToString(buf)
}
func nextIndex(db *sql.DB, table string) int {
func nextIndex(db *store.DB, table string) int {
if db == nil {
return 0
}
+280
View File
@@ -0,0 +1,280 @@
// Package store provides a thin dialect-aware wrapper around database/sql,
// enabling transparent use of both SQLite and PostgreSQL.
package store
import (
"database/sql"
"fmt"
"strings"
)
// Dialect identifies the underlying database engine.
type Dialect int
const (
DialectSQLite Dialect = iota
DialectPostgres
)
// String returns a human-readable dialect name.
func (d Dialect) String() string {
switch d {
case DialectSQLite:
return "sqlite"
case DialectPostgres:
return "postgres"
default:
return "unknown"
}
}
// DB wraps *sql.DB with dialect awareness.
type DB struct {
raw *sql.DB
dialect Dialect
}
// Wrap creates a new dialect-aware DB from an existing *sql.DB.
func Wrap(raw *sql.DB, dialect Dialect) *DB {
return &DB{raw: raw, dialect: dialect}
}
// Dialect returns the database dialect.
func (db *DB) Dialect() Dialect {
if db == nil {
return DialectSQLite
}
return db.dialect
}
// RawDB returns the underlying *sql.DB.
func (db *DB) RawDB() *sql.DB {
if db == nil {
return nil
}
return db.raw
}
// Close closes the underlying connection.
func (db *DB) Close() error {
if db == nil || db.raw == nil {
return nil
}
return db.raw.Close()
}
// Ping verifies the connection is alive.
func (db *DB) Ping() error {
return db.raw.Ping()
}
// Exec executes a query with transparent placeholder and syntax rewriting.
func (db *DB) Exec(query string, args ...any) (sql.Result, error) {
return db.raw.Exec(db.rewrite(query), args...)
}
// Query executes a query that returns rows, with transparent rewriting.
func (db *DB) Query(query string, args ...any) (*sql.Rows, error) {
return db.raw.Query(db.rewrite(query), args...)
}
// QueryRow executes a query that returns at most one row, with transparent rewriting.
func (db *DB) QueryRow(query string, args ...any) *sql.Row {
return db.raw.QueryRow(db.rewrite(query), args...)
}
// Begin starts a transaction, returning a dialect-aware Tx.
func (db *DB) Begin() (*Tx, error) {
tx, err := db.raw.Begin()
if err != nil {
return nil, err
}
return &Tx{raw: tx, dialect: db.dialect}, nil
}
// ExecReturningID executes an INSERT and returns the auto-generated id.
// - SQLite: uses LastInsertId()
// - PostgreSQL: appends RETURNING id and uses QueryRow().Scan()
func (db *DB) ExecReturningID(query string, args ...any) (int64, error) {
q := db.rewrite(query)
if db.dialect == DialectPostgres {
q = strings.TrimRight(q, "; \t\n") + " RETURNING id"
var id int64
if err := db.raw.QueryRow(q, args...).Scan(&id); err != nil {
return 0, err
}
return id, nil
}
res, err := db.raw.Exec(q, args...)
if err != nil {
return 0, err
}
return res.LastInsertId()
}
// Tx wraps *sql.Tx with dialect awareness.
type Tx struct {
raw *sql.Tx
dialect Dialect
}
// Exec executes a query inside the transaction with transparent rewriting.
func (tx *Tx) Exec(query string, args ...any) (sql.Result, error) {
return tx.raw.Exec(rewriteQuery(tx.dialect, query), args...)
}
// Query executes a query that returns rows inside the transaction.
func (tx *Tx) Query(query string, args ...any) (*sql.Rows, error) {
return tx.raw.Query(rewriteQuery(tx.dialect, query), args...)
}
// QueryRow executes a query that returns at most one row inside the transaction.
func (tx *Tx) QueryRow(query string, args ...any) *sql.Row {
return tx.raw.QueryRow(rewriteQuery(tx.dialect, query), args...)
}
// Commit commits the transaction.
func (tx *Tx) Commit() error { return tx.raw.Commit() }
// Rollback aborts the transaction.
func (tx *Tx) Rollback() error { return tx.raw.Rollback() }
// ExecReturningID executes an INSERT inside the transaction and returns the id.
func (tx *Tx) ExecReturningID(query string, args ...any) (int64, error) {
q := rewriteQuery(tx.dialect, query)
if tx.dialect == DialectPostgres {
q = strings.TrimRight(q, "; \t\n") + " RETURNING id"
var id int64
if err := tx.raw.QueryRow(q, args...).Scan(&id); err != nil {
return 0, err
}
return id, nil
}
res, err := tx.raw.Exec(q, args...)
if err != nil {
return 0, err
}
return res.LastInsertId()
}
func (db *DB) rewrite(query string) string {
return rewriteQuery(db.dialect, query)
}
func rewriteQuery(dialect Dialect, query string) string {
if dialect != DialectPostgres {
return query
}
query = rewriteUserIdentifier(query)
query = rewriteInsertOrIgnore(query)
query = rewritePlaceholders(query)
return query
}
func rewriteUserIdentifier(query string) string {
var buf strings.Builder
buf.Grow(len(query) + 16)
inSingle := false
inDouble := false
i := 0
for i < len(query) {
ch := query[i]
if ch == '\'' && !inDouble {
if inSingle && i+1 < len(query) && query[i+1] == '\'' {
buf.WriteByte(ch)
buf.WriteByte(query[i+1])
i += 2
continue
}
inSingle = !inSingle
buf.WriteByte(ch)
i++
continue
}
if ch == '"' && !inSingle {
inDouble = !inDouble
buf.WriteByte(ch)
i++
continue
}
if inSingle || inDouble {
buf.WriteByte(ch)
i++
continue
}
if isIdentifierChar(ch) {
j := i + 1
for j < len(query) && isIdentifierChar(query[j]) {
j++
}
tok := query[i:j]
if strings.EqualFold(tok, "user") {
buf.WriteString(`"user"`)
} else {
buf.WriteString(tok)
}
i = j
continue
}
buf.WriteByte(ch)
i++
}
return buf.String()
}
func isIdentifierChar(ch byte) bool {
if ch >= 'a' && ch <= 'z' {
return true
}
if ch >= 'A' && ch <= 'Z' {
return true
}
if ch >= '0' && ch <= '9' {
return true
}
return ch == '_'
}
func rewriteInsertOrIgnore(query string) string {
upper := strings.ToUpper(query)
idx := strings.Index(upper, "INSERT OR IGNORE INTO")
if idx < 0 {
return query
}
prefix := query[:idx]
suffix := query[idx+len("INSERT OR IGNORE INTO"):]
result := prefix + "INSERT INTO" + suffix
trimmed := strings.TrimRight(result, "; \t\n")
return trimmed + " ON CONFLICT DO NOTHING"
}
func rewritePlaceholders(query string) string {
var buf strings.Builder
buf.Grow(len(query) + 16)
n := 1
inString := false
for i := 0; i < len(query); i++ {
ch := query[i]
if ch == '\'' {
if inString && i+1 < len(query) && query[i+1] == '\'' {
buf.WriteByte(ch)
buf.WriteByte(query[i+1])
i++
continue
}
inString = !inString
buf.WriteByte(ch)
continue
}
if ch == '?' && !inString {
buf.WriteString(fmt.Sprintf("$%d", n))
n++
continue
}
buf.WriteByte(ch)
}
return buf.String()
}
@@ -0,0 +1,9 @@
package postgres
import _ "embed"
//go:embed sql/schema.sql
var EmbeddedSchema string
//go:embed sql/data.sql
var EmbeddedSeedData string
@@ -0,0 +1,18 @@
INSERT INTO "user" (id, "user", pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES (1, 'admin_user', '3c85cdebade1c51cf64ca9f3c09d182d', 0, 2727251700000, 99999, 0, 0, 1, 99999, 1748914865000, 1754011744252, 1)
ON CONFLICT DO NOTHING;
INSERT INTO vite_config (id, name, value, time)
VALUES (1, 'app_name', 'flux', 1755147963000)
ON CONFLICT DO NOTHING;
DO $$
BEGIN
IF to_regclass('public.user_id_seq') IS NOT NULL THEN
PERFORM setval('user_id_seq', (SELECT COALESCE(MAX(id), 0) FROM "user"));
END IF;
IF to_regclass('public.vite_config_id_seq') IS NOT NULL THEN
PERFORM setval('vite_config_id_seq', (SELECT COALESCE(MAX(id), 0) FROM vite_config));
END IF;
END
$$;
@@ -0,0 +1,241 @@
CREATE TABLE IF NOT EXISTS forward (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
user_name VARCHAR(100) NOT NULL,
name VARCHAR(100) NOT NULL,
tunnel_id INTEGER NOT NULL,
remote_addr TEXT NOT NULL,
strategy VARCHAR(100) NOT NULL DEFAULT 'fifo',
in_flow BIGINT NOT NULL DEFAULT 0,
out_flow BIGINT NOT NULL DEFAULT 0,
created_time BIGINT NOT NULL,
updated_time BIGINT NOT NULL,
status INTEGER NOT NULL,
inx INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE IF NOT EXISTS forward_port (
id SERIAL PRIMARY KEY,
forward_id INTEGER NOT NULL,
node_id INTEGER NOT NULL,
port INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS node (
id SERIAL PRIMARY KEY,
name VARCHAR(100) NOT NULL,
secret VARCHAR(100) NOT NULL,
server_ip VARCHAR(100) NOT NULL,
server_ip_v4 VARCHAR(100),
server_ip_v6 VARCHAR(100),
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 BIGINT NOT NULL,
updated_time BIGINT,
status INTEGER NOT NULL,
tcp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
udp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
inx INTEGER NOT NULL DEFAULT 0,
is_remote INTEGER DEFAULT 0,
remote_url TEXT,
remote_token TEXT,
remote_config TEXT
);
CREATE TABLE IF NOT EXISTS speed_limit (
id SERIAL PRIMARY KEY,
name VARCHAR(100) NOT NULL,
speed INTEGER NOT NULL,
tunnel_id INTEGER NOT NULL,
tunnel_name VARCHAR(100) NOT NULL,
created_time BIGINT NOT NULL,
updated_time BIGINT,
status INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS statistics_flow (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
flow BIGINT NOT NULL,
total_flow BIGINT NOT NULL,
time VARCHAR(100) NOT NULL,
created_time BIGINT NOT NULL
);
CREATE TABLE IF NOT EXISTS tunnel (
id SERIAL PRIMARY KEY,
name VARCHAR(100) NOT NULL,
traffic_ratio DOUBLE PRECISION NOT NULL DEFAULT 1.0,
type INTEGER NOT NULL,
protocol VARCHAR(10) NOT NULL DEFAULT 'tls',
flow BIGINT NOT NULL,
created_time BIGINT NOT NULL,
updated_time BIGINT NOT NULL,
status INTEGER NOT NULL,
in_ip TEXT,
inx INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE IF NOT EXISTS chain_tunnel (
id SERIAL PRIMARY KEY,
tunnel_id INTEGER NOT NULL,
chain_type VARCHAR(10) NOT NULL,
node_id INTEGER NOT NULL,
port INTEGER,
strategy VARCHAR(10),
inx INTEGER,
protocol VARCHAR(10)
);
CREATE TABLE IF NOT EXISTS "user" (
id SERIAL PRIMARY KEY,
"user" VARCHAR(100) NOT NULL,
pwd VARCHAR(100) NOT NULL,
role_id INTEGER NOT NULL,
exp_time BIGINT NOT NULL,
flow BIGINT NOT NULL,
in_flow BIGINT NOT NULL DEFAULT 0,
out_flow BIGINT NOT NULL DEFAULT 0,
flow_reset_time BIGINT NOT NULL,
num INTEGER NOT NULL,
created_time BIGINT NOT NULL,
updated_time BIGINT,
status INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS user_tunnel (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
tunnel_id INTEGER NOT NULL,
speed_id INTEGER,
num INTEGER NOT NULL,
flow BIGINT NOT NULL,
in_flow BIGINT NOT NULL DEFAULT 0,
out_flow BIGINT NOT NULL DEFAULT 0,
flow_reset_time BIGINT NOT NULL,
exp_time BIGINT NOT NULL,
status INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS tunnel_group (
id SERIAL PRIMARY KEY,
name VARCHAR(100) NOT NULL,
created_time BIGINT NOT NULL,
updated_time BIGINT NOT NULL,
status INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS user_group (
id SERIAL PRIMARY KEY,
name VARCHAR(100) NOT NULL,
created_time BIGINT NOT NULL,
updated_time BIGINT NOT NULL,
status INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS tunnel_group_tunnel (
id SERIAL PRIMARY KEY,
tunnel_group_id INTEGER NOT NULL,
tunnel_id INTEGER NOT NULL,
created_time BIGINT NOT NULL
);
CREATE TABLE IF NOT EXISTS user_group_user (
id SERIAL PRIMARY KEY,
user_group_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
created_time BIGINT NOT NULL
);
CREATE TABLE IF NOT EXISTS group_permission (
id SERIAL PRIMARY KEY,
user_group_id INTEGER NOT NULL,
tunnel_group_id INTEGER NOT NULL,
created_time BIGINT NOT NULL
);
CREATE TABLE IF NOT EXISTS group_permission_grant (
id SERIAL PRIMARY KEY,
user_group_id INTEGER NOT NULL,
tunnel_group_id INTEGER NOT NULL,
user_tunnel_id INTEGER NOT NULL,
created_by_group INTEGER NOT NULL DEFAULT 0,
created_time BIGINT NOT NULL
);
CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_name ON tunnel_group(name);
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_name ON user_group(name);
CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_tunnel_unique ON tunnel_group_tunnel(tunnel_group_id, tunnel_id);
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_user_unique ON user_group_user(user_group_id, user_id);
CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_unique ON group_permission(user_group_id, tunnel_group_id);
CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_grant_unique ON group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id);
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_tunnel_unique ON user_tunnel(user_id, tunnel_id);
CREATE TABLE IF NOT EXISTS vite_config (
id SERIAL PRIMARY KEY,
name VARCHAR(200) NOT NULL UNIQUE,
value VARCHAR(200) NOT NULL,
time BIGINT NOT NULL
);
CREATE TABLE IF NOT EXISTS peer_share (
id SERIAL PRIMARY KEY,
name TEXT NOT NULL,
node_id INTEGER NOT NULL,
token TEXT NOT NULL UNIQUE,
max_bandwidth INTEGER DEFAULT 0,
expiry_time BIGINT DEFAULT 0,
port_range_start INTEGER DEFAULT 0,
port_range_end INTEGER DEFAULT 0,
current_flow BIGINT DEFAULT 0,
is_active INTEGER DEFAULT 1,
created_time BIGINT NOT NULL,
updated_time BIGINT NOT NULL,
allowed_domains TEXT DEFAULT '',
allowed_ips TEXT DEFAULT ''
);
CREATE TABLE IF NOT EXISTS peer_share_runtime (
id SERIAL PRIMARY KEY,
share_id INTEGER NOT NULL,
node_id INTEGER NOT NULL,
reservation_id TEXT NOT NULL UNIQUE,
resource_key TEXT NOT NULL UNIQUE,
binding_id TEXT NOT NULL DEFAULT '',
role TEXT NOT NULL DEFAULT '',
chain_name TEXT NOT NULL DEFAULT '',
service_name TEXT NOT NULL DEFAULT '',
protocol TEXT NOT NULL DEFAULT 'tls',
strategy TEXT NOT NULL DEFAULT 'round',
port INTEGER NOT NULL DEFAULT 0,
target TEXT NOT NULL DEFAULT '',
applied INTEGER NOT NULL DEFAULT 0,
status INTEGER NOT NULL DEFAULT 1,
created_time BIGINT NOT NULL,
updated_time BIGINT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_peer_share_runtime_share_node_status ON peer_share_runtime(share_id, node_id, status);
CREATE INDEX IF NOT EXISTS idx_peer_share_runtime_binding_id ON peer_share_runtime(binding_id);
CREATE TABLE IF NOT EXISTS federation_tunnel_binding (
id SERIAL PRIMARY KEY,
tunnel_id INTEGER NOT NULL,
node_id INTEGER NOT NULL,
chain_type INTEGER NOT NULL,
hop_inx INTEGER NOT NULL DEFAULT 0,
remote_url TEXT NOT NULL,
resource_key TEXT NOT NULL UNIQUE,
remote_binding_id TEXT NOT NULL,
allocated_port INTEGER NOT NULL,
status INTEGER NOT NULL DEFAULT 1,
created_time BIGINT NOT NULL,
updated_time BIGINT NOT NULL
);
CREATE UNIQUE INDEX IF NOT EXISTS idx_federation_tunnel_binding_unique ON federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx);
CREATE INDEX IF NOT EXISTS idx_federation_tunnel_binding_tunnel ON federation_tunnel_binding(tunnel_id, status);
+54 -10
View File
@@ -12,6 +12,9 @@ import (
"strings"
"time"
_ "github.com/jackc/pgx/v5/stdlib"
"go-backend/internal/store"
pgstore "go-backend/internal/store/postgres"
_ "modernc.org/sqlite"
)
@@ -22,10 +25,10 @@ var embeddedSchema string
var embeddedSeedData string
type Repository struct {
db *sql.DB
db *store.DB
}
func (r *Repository) DB() *sql.DB {
func (r *Repository) DB() *store.DB {
if r == nil {
return nil
}
@@ -165,17 +168,47 @@ func Open(path string) (*Repository, error) {
return nil, err
}
db, err := sql.Open("sqlite", path)
raw, err := sql.Open("sqlite", path)
if err != nil {
return nil, err
}
db := store.Wrap(raw, store.DialectSQLite)
if err := db.Ping(); err != nil {
_ = db.Close()
return nil, err
}
if err := bootstrapSchema(db); err != nil {
if err := bootstrapSchema(db, embeddedSchema, embeddedSeedData); err != nil {
_ = db.Close()
return nil, err
}
if err := migrateSchema(db); err != nil {
_ = db.Close()
return nil, err
}
return &Repository{db: db}, nil
}
func OpenPostgres(dsn string) (*Repository, error) {
if strings.TrimSpace(dsn) == "" {
return nil, fmt.Errorf("empty postgres dsn")
}
raw, err := sql.Open("pgx", dsn)
if err != nil {
return nil, err
}
db := store.Wrap(raw, store.DialectPostgres)
if err := db.Ping(); err != nil {
_ = db.Close()
return nil, err
}
if err := bootstrapSchema(db, pgstore.EmbeddedSchema, pgstore.EmbeddedSeedData); err != nil {
_ = db.Close()
return nil, err
}
@@ -1141,7 +1174,7 @@ func nullableForwardIngress(v string) interface{} {
return v
}
func resolveForwardIngress(db *sql.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) {
func resolveForwardIngress(db *store.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) {
var tunnelInIP sql.NullString
if err := db.QueryRow(`SELECT in_ip FROM tunnel WHERE id = ? LIMIT 1`, tunnelID).Scan(&tunnelInIP); err != nil {
if !errors.Is(err, sql.ErrNoRows) {
@@ -1243,22 +1276,22 @@ func ensureParentDir(dbPath string) error {
return osMkdirAll(dir)
}
func bootstrapSchema(db *sql.DB) error {
func bootstrapSchema(db *store.DB, schemaSQL, seedSQL string) error {
if db == nil {
return errors.New("nil db")
}
if _, err := db.Exec(embeddedSchema); err != nil {
if _, err := db.Exec(schemaSQL); err != nil {
return fmt.Errorf("apply schema.sql: %w", err)
}
if _, err := db.Exec(embeddedSeedData); err != nil {
if _, err := db.Exec(seedSQL); err != nil {
return fmt.Errorf("apply data.sql: %w", err)
}
return nil
}
func migrateSchema(db *sql.DB) error {
func migrateSchema(db *store.DB) error {
if db == nil {
return errors.New("nil db")
}
@@ -1269,7 +1302,7 @@ func migrateSchema(db *sql.DB) error {
if err == nil || errors.Is(err, sql.ErrNoRows) {
return
}
if strings.Contains(err.Error(), "no such column") {
if isMissingColumnError(db.Dialect(), err) {
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)
}
@@ -1309,6 +1342,17 @@ func migrateSchema(db *sql.DB) error {
return nil
}
func isMissingColumnError(dialect store.Dialect, err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
if dialect == store.DialectPostgres {
return strings.Contains(msg, "column") && strings.Contains(msg, "does not exist")
}
return strings.Contains(msg, "no such column")
}
func (r *Repository) CreatePeerShare(share *PeerShare) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
@@ -17,6 +17,7 @@ import (
httpserver "go-backend/internal/http"
"go-backend/internal/http/handler"
"go-backend/internal/http/response"
"go-backend/internal/store"
"go-backend/internal/store/sqlite"
_ "modernc.org/sqlite"
@@ -305,7 +306,7 @@ func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
}
}
func readTableColumns(t *testing.T, db *sql.DB, table string) map[string]bool {
func readTableColumns(t *testing.T, db *store.DB, table string) map[string]bool {
t.Helper()
rows, err := db.Query("PRAGMA table_info(" + table + ")")