From cedcaebd1feb56888c8c4a87553cca05ab682ac0 Mon Sep 17 00:00:00 2001 From: sagit Date: Thu, 12 Feb 2026 06:38:25 +0000 Subject: [PATCH] feat(postgres): add postgres backend support and migration docs --- .gitignore | 3 + README.md | 64 ++++ docker-compose-v4.yml | 29 ++ docker-compose-v6.yml | 29 ++ go-backend/go.mod | 9 +- go-backend/go.sum | 38 ++- go-backend/internal/app/app.go | 22 +- go-backend/internal/config/config.go | 20 +- .../internal/http/handler/federation.go | 4 +- go-backend/internal/http/handler/mutations.go | 39 ++- go-backend/internal/store/db.go | 280 ++++++++++++++++++ go-backend/internal/store/postgres/embed.go | 9 + .../internal/store/postgres/sql/data.sql | 18 ++ .../internal/store/postgres/sql/schema.sql | 241 +++++++++++++++ .../internal/store/sqlite/repository.go | 64 +++- .../tests/contract/migration_contract_test.go | 3 +- 16 files changed, 819 insertions(+), 53 deletions(-) create mode 100644 go-backend/internal/store/db.go create mode 100644 go-backend/internal/store/postgres/embed.go create mode 100644 go-backend/internal/store/postgres/sql/data.sql create mode 100644 go-backend/internal/store/postgres/sql/schema.sql diff --git a/.gitignore b/.gitignore index a422608..c7b3db5 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/README.md b/README.md index 66ae164..5133b80 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/docker-compose-v4.yml b/docker-compose-v4.yml index 1169509..0da0631 100644 --- a/docker-compose-v4.yml +++ b/docker-compose-v4.yml @@ -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 diff --git a/docker-compose-v6.yml b/docker-compose-v6.yml index e0832eb..e8080ec 100644 --- a/docker-compose-v6.yml +++ b/docker-compose-v6.yml @@ -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 diff --git a/go-backend/go.mod b/go-backend/go.mod index 2c26ac9..4295d62 100644 --- a/go-backend/go.mod +++ b/go-backend/go.mod @@ -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 diff --git a/go-backend/go.sum b/go-backend/go.sum index fa6b48b..207fe9e 100644 --- a/go-backend/go.sum +++ b/go-backend/go.sum @@ -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= diff --git a/go-backend/internal/app/app.go b/go-backend/internal/app/app.go index 425b05a..527b866 100644 --- a/go-backend/internal/app/app.go +++ b/go-backend/internal/app/app.go @@ -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) diff --git a/go-backend/internal/config/config.go b/go-backend/internal/config/config.go index 043730c..f94c721 100644 --- a/go-backend/internal/config/config.go +++ b/go-backend/internal/config/config.go @@ -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 diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index e3bea0d..af616a1 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -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, diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 9784d04..bafe95c 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -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 } diff --git a/go-backend/internal/store/db.go b/go-backend/internal/store/db.go new file mode 100644 index 0000000..50107c3 --- /dev/null +++ b/go-backend/internal/store/db.go @@ -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() +} diff --git a/go-backend/internal/store/postgres/embed.go b/go-backend/internal/store/postgres/embed.go new file mode 100644 index 0000000..749852d --- /dev/null +++ b/go-backend/internal/store/postgres/embed.go @@ -0,0 +1,9 @@ +package postgres + +import _ "embed" + +//go:embed sql/schema.sql +var EmbeddedSchema string + +//go:embed sql/data.sql +var EmbeddedSeedData string diff --git a/go-backend/internal/store/postgres/sql/data.sql b/go-backend/internal/store/postgres/sql/data.sql new file mode 100644 index 0000000..ee3f9dd --- /dev/null +++ b/go-backend/internal/store/postgres/sql/data.sql @@ -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 +$$; diff --git a/go-backend/internal/store/postgres/sql/schema.sql b/go-backend/internal/store/postgres/sql/schema.sql new file mode 100644 index 0000000..fcc6cc1 --- /dev/null +++ b/go-backend/internal/store/postgres/sql/schema.sql @@ -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); diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index 7f6bc00..ec68fd2 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -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") diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go index 5795474..89780cc 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -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 + ")")