From 66be07750f163a60fa88968df97788e694813539 Mon Sep 17 00:00:00 2001 From: Antigravity Date: Tue, 17 Feb 2026 04:47:11 +0000 Subject: [PATCH] refactor(backend): migrate to modular repository pattern with separated concerns - Extract database layer into model and repo packages - Split repository into focused modules (control, federation, flow, groups, mutations) - Remove monolithic db.go and sqlite/repository.go - Update handlers to use new repository structure - Migrate contract tests to new patterns - Add migration plan documentation --- .sisyphus/ralph-loop.local.md | 9 + AGENTS.md | 4 +- go-backend/AGENTS.md | 29 +- go-backend/MIGRATION_PLAN.md | 536 +++ go-backend/go.mod | 6 + go-backend/go.sum | 12 + go-backend/internal/app/app.go | 16 +- go-backend/internal/http/handler/AGENTS.md | 9 +- .../internal/http/handler/control_plane.go | 316 +- .../internal/http/handler/federation.go | 190 +- .../http/handler/federation_runtime_test.go | 76 +- .../http/handler/federation_share_test.go | 135 +- .../internal/http/handler/flow_policy.go | 89 +- .../handler/flow_policy_federation_test.go | 20 +- go-backend/internal/http/handler/handler.go | 30 +- go-backend/internal/http/handler/jobs.go | 122 +- go-backend/internal/http/handler/jobs_test.go | 46 +- go-backend/internal/http/handler/mutations.go | 1056 ++---- go-backend/internal/store/db.go | 466 --- go-backend/internal/store/db_test.go | 116 - go-backend/internal/store/model/model.go | 591 ++++ go-backend/internal/store/postgres/embed.go | 9 - .../internal/store/postgres/sql/data.sql | 18 - .../internal/store/postgres/sql/schema.sql | 250 -- go-backend/internal/store/repo/repository.go | 2545 ++++++++++++++ .../internal/store/repo/repository_control.go | 308 ++ .../store/repo/repository_federation.go | 269 ++ .../internal/store/repo/repository_flow.go | 171 + .../internal/store/repo/repository_groups.go | 69 + .../repository_migrate_test.go | 41 +- .../store/repo/repository_mutations.go | 1408 ++++++++ .../internal/store/sqlite/repository.go | 3117 ----------------- go-backend/internal/store/sqlite/sql/data.sql | 5 - .../internal/store/sqlite/sql/schema.sql | 254 -- go-backend/internal/ws/server.go | 6 +- .../tests/contract/diagnosis_contract_test.go | 98 +- .../federation_dual_panel_contract_test.go | 63 +- .../tests/contract/forward_contract_test.go | 185 +- .../group_permission_contract_test.go | 78 +- .../tests/contract/migration_contract_test.go | 100 +- .../postgres_node_id_repair_contract_test.go | 22 +- .../contract/tunnel_create_contract_test.go | 37 +- .../tunnel_ip_preference_contract_test.go | 65 +- .../tunnel_visibility_contract_test.go | 25 +- 44 files changed, 6829 insertions(+), 6188 deletions(-) create mode 100644 .sisyphus/ralph-loop.local.md create mode 100644 go-backend/MIGRATION_PLAN.md delete mode 100644 go-backend/internal/store/db.go delete mode 100644 go-backend/internal/store/db_test.go create mode 100644 go-backend/internal/store/model/model.go delete mode 100644 go-backend/internal/store/postgres/embed.go delete mode 100644 go-backend/internal/store/postgres/sql/data.sql delete mode 100644 go-backend/internal/store/postgres/sql/schema.sql create mode 100644 go-backend/internal/store/repo/repository.go create mode 100644 go-backend/internal/store/repo/repository_control.go create mode 100644 go-backend/internal/store/repo/repository_federation.go create mode 100644 go-backend/internal/store/repo/repository_flow.go create mode 100644 go-backend/internal/store/repo/repository_groups.go rename go-backend/internal/store/{sqlite => repo}/repository_migrate_test.go (53%) create mode 100644 go-backend/internal/store/repo/repository_mutations.go delete mode 100644 go-backend/internal/store/sqlite/repository.go delete mode 100644 go-backend/internal/store/sqlite/sql/data.sql delete mode 100644 go-backend/internal/store/sqlite/sql/schema.sql diff --git a/.sisyphus/ralph-loop.local.md b/.sisyphus/ralph-loop.local.md new file mode 100644 index 0000000..af00a2a --- /dev/null +++ b/.sisyphus/ralph-loop.local.md @@ -0,0 +1,9 @@ +--- +active: true +iteration: 1 +max_iterations: 100 +completion_promise: "DONE" +started_at: "2026-02-15T16:38:58.212Z" +session_id: "ses_39dd49703ffeveg711aA1D1YAk" +--- +现状后端数据库兼容sqlite和postgresql,每一次新增功能需要维护两套数据库sql,需求是使用一个数据库驱动能同时兼容两个数据库,请仔细分析,列出计划,全量迁移,并且写好所有的测试,确保重构后所有的功能都能正常运行,由于工程量大,请写一个计划列表的markdown记录,每次完成一个就记录一下进度 diff --git a/AGENTS.md b/AGENTS.md index 206e2ef..955a62e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -12,7 +12,7 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a ./ ├── go-gost/ # Go forwarding agent (forked gost + local x/) │ └── x/ # Local fork of github.com/go-gost/x (replace => ./x) -├── go-backend/ # Go Admin API (SQLite, net/http) +├── go-backend/ # Go Admin API (GORM + SQLite/PostgreSQL, net/http) ├── vite-frontend/ # React/Vite dashboard (HeroUI + Tailwind) ├── docker-compose-v4.yml # Panel deploy (IPv4-only bridge) ├── docker-compose-v6.yml # Panel deploy (IPv6-enabled bridge) @@ -52,7 +52,7 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a - **DO NOT EDIT** generated protobuf output: `go-gost/x/internal/util/grpc/proto/*.pb.go`, `go-gost/x/internal/util/grpc/proto/*_grpc.pb.go`. - **DO NOT ADD** `Bearer` prefix to Authorization header - expects raw JWT token. - **DO NOT MODIFY** `install.sh` or `panel_install.sh` locally - CI overwrites these on release. -- **DO NOT USE** ORM in backend - uses raw SQL with `database/sql`. +- **DO NOT** let backend handlers call `repo.DB()` directly — add a Repository method instead. - **DO NOT ADD** frontend tests - project has no test infrastructure (Vitest/Jest not configured). ## COMMANDS diff --git a/go-backend/AGENTS.md b/go-backend/AGENTS.md index 7efa6dc..0a5201c 100644 --- a/go-backend/AGENTS.md +++ b/go-backend/AGENTS.md @@ -2,7 +2,7 @@ ## OVERVIEW Go-based Admin API for FLVX. Replaced legacy Spring Boot backend. -**Stack:** Go 1.23, net/http (std lib), SQLite/PostgreSQL (modernc.org/sqlite - CGO-free). +**Stack:** Go 1.23, net/http (std lib), GORM + SQLite/PostgreSQL (glebarez/sqlite - CGO-free). ## STRUCTURE ``` @@ -14,9 +14,14 @@ go-backend/ │ │ ├── handler/ # API Handlers (User, Tunnel, Node, etc.) │ │ ├── middleware/ # JWT, CORS, Logging, Recover │ │ └── response/ # JSON response helpers -│ ├── store/sqlite/ # Data Access Layer (Repository pattern) -│ │ ├── repository.go # SQL queries & Struct definitions -│ │ └── sql/ # Embedded schema.sql & data.sql +│ ├── store/ +│ │ ├── model/model.go # GORM model structs (single source of truth) +│ │ └── repo/ # Data Access Layer (Repository pattern, GORM) +│ │ ├── repository.go # Core queries, Open/OpenPostgres, AutoMigrate +│ │ ├── repository_mutations.go # Mutation helpers (user/node/tunnel/forward CRUD) +│ │ ├── repository_federation.go# Federation-specific queries +│ │ ├── repository_flow.go # Flow/forward status queries +│ │ └── repository_control.go # Control plane queries │ └── auth/ # Auth logic ├── tests/ # Integration/Contract tests ├── Dockerfile # Multi-stage build (alpine) @@ -27,23 +32,27 @@ go-backend/ | Task | Location | Notes | |------|----------|-------| | **API Routes** | `go-backend/internal/http/router.go` | Registers handlers to `http.ServeMux` | -| **DB Schema** | `go-backend/internal/store/sqlite/sql/schema.sql` | Embedded in binary | -| **SQL Queries** | `go-backend/internal/store/sqlite/repository.go` | Raw SQL, no ORM | +| **DB Models** | `go-backend/internal/store/model/model.go` | GORM structs with `TableName()` methods | +| **Repository** | `go-backend/internal/store/repo/` | GORM-based queries, all DB ops encapsulated | | **Auth Middleware** | `go-backend/internal/http/middleware/jwt.go` | Extracts `Authorization` header | | **WebSocket** | `go-backend/internal/ws/` | Real-time updates (traffic, status) | ## CONVENTIONS -- **No ORM**: Uses raw SQL with `database/sql` and `modernc.org/sqlite`. -- **CGO-Free SQLite**: `modernc.org/sqlite` instead of `mattn/go-sqlite3` - builds without CGO. +- **GORM ORM**: Uses GORM with `glebarez/sqlite` (CGO-free) and `gorm.io/driver/postgres`. +- **AutoMigrate**: Schema created at startup via `autoMigrateAll()` — no hand-written DDL. +- **TableName()**: All models define explicit `TableName()` returning singular snake_case names. +- **Repository Pattern**: Handlers never access `*gorm.DB` directly — all queries go through `repo.Repository` methods. - **Standard Lib**: Uses `net/http` for routing (Go 1.22+ patterns). - **Auth**: Expects raw JWT in `Authorization` header (no `Bearer` prefix). - **API Envelope**: All responses use `response.R{code, msg, data, ts}` structure. - **Config**: Loaded from environment variables (see `cmd/paneld/main.go`). -- **SQL Idempotency**: Prefer `ON CONFLICT DO NOTHING` for inserts in migrations/sync. +- **SQLite Constraints**: `MaxOpenConns(1)`, WAL mode, busy_timeout=5000. ## ANTI-PATTERNS -- **DO NOT USE** ORM - uses raw SQL throughout. +- **DO NOT** let handlers call `repo.DB()` directly — add a Repository method instead. - **DO NOT CHANGE** handler signatures without updating `router.go`. +- **DO NOT** use `type:jsonb` or `type:serial` in GORM tags (SQLite incompatible). +- **DO NOT** omit `TableName()` on new models — GORM pluralizes by default. ## COMMANDS ```bash diff --git a/go-backend/MIGRATION_PLAN.md b/go-backend/MIGRATION_PLAN.md new file mode 100644 index 0000000..5a9fe10 --- /dev/null +++ b/go-backend/MIGRATION_PLAN.md @@ -0,0 +1,536 @@ +# 数据库 GORM ORM 迁移计划 + +**创建时间:** 2026-02-15 +**更新时间:** 2026-02-17 (实施:完成 P1 + P2 + P3 + P5(Repo 查询层 + schema 收尾) + 测试/构建收尾) +**分支:** main (commit e5e22ba) +**状态:** 基本完成(保留 4 处 PG 序列修复 DDL `Exec`) + +--- + +## 一、现状分析 + +### 1.1 迁移前架构 (已归档) + +项目原使用 `database/sql` + 手写 raw SQL,通过 `internal/store/db.go` 中的运行时 SQL 重写层实现 SQLite/PostgreSQL 双数据库兼容。 + +| 组件 | 行数 | 角色 | 当前状态 | +|------|------|------|----------| +| `store/db.go` | ~520 | SQL 方言重写层 | **已删除** | +| `store/sqlite/repository.go` | ~3118 | Repository 查询方法 | **已重写为 store/repo/** | +| `handler/mutations.go` | ~3748 | Handler 内直接写 raw SQL | **已迁移到 repo(生产 SQL=0)** | +| `handler/handler.go` | ~1283 | 部分方法用 `repo.DB()` | **大部分已迁移** | +| `handler/federation.go` | ~若干 | Federation 相关 SQL | **已迁移到 repo** | +| `handler/control_plane.go` | ~若干 | 控制面相关 SQL | **已迁移到 repo** | +| `handler/flow_policy.go` | ~若干 | 流量策略相关 SQL | **已迁移到 repo** | +| `handler/jobs.go` | ~若干 | 后台任务相关 SQL | **已迁移到 repo** | +| `store/postgres/` | 目录 | PostgreSQL 专用 schema/data | **已删除** | + +### 1.2 痛点 (迁移目标) + +1. ~~**双 Schema 维护**~~:已通过 AutoMigrate 解决 +2. ~~**SQL 重写层复杂**~~:db.go 已删除 +3. ~~**handler 直接写 SQL**~~:`mutations.go` 生产路径 `tx.Exec`/`tx.Raw` 已清零(测试代码除外) +4. ~~**无类型安全**~~:repo 业务查询已 GORM 化;剩余 4 处为 PG 序列修复 DDL `Exec`(设计保留) +5. ~~**模型定义分散**~~:已集中到 model/model.go + +--- + +## 二、方案:引入 GORM ORM(全面重写) + +### 2.1 方案变更说明 + +原计划为 **方案 D(扩展现有 DDL 重写层)**,现变更为 **方案 A(GORM 全面重写)**。 + +### 2.2 选择 GORM 的理由 + +1. Go 生态最成熟的 ORM,社区庞大,文档完善 +2. 原生支持 SQLite + PostgreSQL 双数据库,自动处理方言差异 +3. AutoMigrate 消除双 schema 维护,自动处理 AUTOINCREMENT ↔ SERIAL 等 +4. 类型安全的模型定义,编译期检查字段映射 +5. 内置事务管理(closure pattern 自动 rollback/commit) +6. 自动处理 `"user"` 保留字引号 + +### 2.3 GORM 驱动选择 + +| 数据库 | 驱动 | 包 | 备注 | +|--------|------|-----|------| +| SQLite | modernc.org/sqlite (CGO-free) | `github.com/glebarez/sqlite` | 纯 Go,无需 CGO | +| PostgreSQL | pgx/v5 | `gorm.io/driver/postgres` | 默认使用 pgx | + +> **注意**:标准 `gorm.io/driver/sqlite` 依赖 CGO,必须使用 `glebarez/sqlite` 包装器。 + +### 2.4 核心设计原则 + +1. **Model 集中定义**:所有 GORM Model 在 `internal/store/model/` 包中 +2. **Repository 模式保留**:Repository struct 持有 `*gorm.DB`,对外方法签名尽量不变 +3. **Handler 不直接操作 DB**:所有数据库操作必须封装在 Repository 方法中 +4. **AutoMigrate 替代 schema.sql**:启动时自动迁移,不再维护手写 DDL +5. **保留 PG 序列修复**:pgloader 迁移场景仍需 `ensurePostgresIDDefaults()` +6. **Package 重命名**:`store/sqlite` → `store/repo` + +--- + +## 三、Model 设计 + +### 3.1 GORM 类型映射 + +| Go 类型 | GORM 行为 | PostgreSQL | SQLite | +|---------|-----------|------------|--------| +| `int64` + `primaryKey` | 自增主键 | `bigserial` | `INTEGER PRIMARY KEY AUTOINCREMENT` | +| `int64` | 64位整数 | `bigint` | `integer` (SQLite 自动 64位) | +| `int` | 整数 | `integer` | `integer` | +| `float64` | 浮点 | `double precision` | `real` | +| `string` + `size:100` | 变长字符 | `varchar(100)` | `varchar(100)` | +| `string` (无 size) | 文本 | `text` | `text` | +| `sql.NullInt64` | 可空整数 | `bigint NULL` | `integer NULL` | +| `sql.NullString` | 可空文本 | `text NULL` | `text NULL` | + +### 3.2 表清单(21 张表) + +| 表名 | Model | 特殊处理 | +|------|-------|----------| +| `user` | `User` | `TableName()` 返回 `"user"` (PG 保留字) | +| `forward` | `Forward` | | +| `forward_port` | `ForwardPort` | | +| `node` | `Node` | | +| `speed_limit` | `SpeedLimit` | | +| `statistics_flow` | `StatisticsFlow` | | +| `tunnel` | `Tunnel` | | +| `chain_tunnel` | `ChainTunnel` | | +| `user_tunnel` | `UserTunnel` | 复合唯一索引 (user_id, tunnel_id) | +| `tunnel_group` | `TunnelGroup` | | +| `user_group` | `UserGroup` | | +| `tunnel_group_tunnel` | `TunnelGroupTunnel` | 复合唯一索引 | +| `user_group_user` | `UserGroupUser` | 复合唯一索引 | +| `group_permission` | `GroupPermission` | 复合唯一索引 | +| `group_permission_grant` | `GroupPermissionGrant` | 复合唯一索引 | +| `vite_config` | `ViteConfig` | name 唯一 | +| `peer_share` | `PeerShare` | token 唯一 | +| `peer_share_runtime` | `PeerShareRuntime` | reservation_id, resource_key 唯一 | +| `federation_tunnel_binding` | `FederationTunnelBinding` | 复合唯一索引 + resource_key 唯一 | +| `announcement` | `Announcement` | | +| `schema_version` | `SchemaVersion` | | + +--- + +## 四、详细实施步骤 + +### 阶段 1:基础设施 — 添加依赖 + 定义 Model ✅ 已完成 + +| 步骤 | 任务 | 文件 | 状态 | +|------|------|------|------| +| 1.1 | `go get gorm.io/gorm gorm.io/driver/postgres github.com/glebarez/sqlite` | `go.mod` | ✅ | +| 1.2 | 创建 `internal/store/model/model.go`,定义全部 21 个表 Model | 新文件 | ✅ | +| 1.3 | 为 `user` 表添加 `TableName()` 处理 PG 保留字 | model.go | ✅ | +| 1.4 | 为复合唯一索引的表添加 GORM 索引 tag | model.go | ✅ | +| 1.5 | 将 Backup 相关 struct 也迁移到 model/ | model.go | ✅ | +| 1.6 | 验证 `go build ./...` 编译通过 | - | ✅ | + +### 阶段 2:GORM DB 初始化 ✅ 已完成 + +| 步骤 | 任务 | 文件 | 状态 | +|------|------|------|------| +| 2.1 | 修改 Repository struct,`*store.DB` → `*gorm.DB` | repository.go | ✅ | +| 2.2 | 重写 `Open()` — 用 `glebarez/sqlite` 打开 SQLite | repository.go | ✅ | +| 2.3 | 重写 `OpenPostgres()` — 用 `gorm.io/driver/postgres` 打开 PG | repository.go | ✅ | +| 2.4 | 用 `db.AutoMigrate()` 替代 `bootstrapSchema()` | repository.go | ✅ | +| 2.5 | 实现种子数据逻辑(FirstOrCreate 替代 data.sql) | repository.go | ✅ | +| 2.6 | 保留并适配 `ensurePostgresIDDefaults()`(用 `db.Exec()`) | repository.go | ✅ | +| 2.7 | 保留并适配 `migrateSchema()` 增量迁移 | repository.go | ✅ | +| 2.8 | `DB()` 方法返回 `*gorm.DB` | repository.go | ✅ | +| 2.9 | SQLite 连接池设置 `MaxOpenConns(1)` 防锁 | repository.go | ✅ | + +### 阶段 3:重写 repository 查询方法 ⚠️ ~97% 完成 + +将所有 raw SQL 查询替换为 GORM 链式调用。 + +> **2026-02-16 审计**:基础 CRUD 查询已 GORM 化,但 mutation、JOIN 查询、import/export 仍大量使用 raw SQL。 +> **2026-02-17 更新**:已完成 `repository_mutations.go`、Import、以及 `repository_federation/control/flow` 查询层 GORM 化;`repository.go` 中 Raw 已清零,当前仅保留 4 处 PG 序列修复 DDL `Exec`。 + +| 步骤 | 任务 | 方法数 | 状态 | +|------|------|--------|------| +| 3.1 | 用户查询:GetUserByUsername, GetUserByID, UsernameExists* 等 | ~5 | ✅ | +| 3.2 | 配置查询:GetConfigByName, ListConfigs, UpsertConfig | ~3 | ✅ | +| 3.3 | 公告查询:GetAnnouncement, UpsertAnnouncement | ~2 | ✅ | +| 3.4 | 节点查询:GetNodeBy*, ListNodes, UpdateNode* | ~6 | ✅ | +| 3.5 | 隧道查询:ListTunnels, ListTunnelGroups 等 (含 chain_tunnel 关联) | ~5 | ✅ | +| 3.6 | 转发查询:ListForwards, resolveForwardIngress | ~3 | ✅ | +| 3.7 | 用户隧道:GetUserPackageTunnels, GetUserPackageForwards | ~3 | ✅ | +| 3.8 | 统计/限速:GetStatisticsFlows, ListSpeedLimits, AddFlow | ~4 | ✅ | +| 3.9 | 分组查询:ListUserGroups, ListGroupPermissions 等 | ~4 | ✅ | +| 3.10 | PeerShare 全部方法 (CRUD + Runtime) | ~15 | ✅ | +| 3.11 | FederationTunnelBinding 全部方法 | ~4 | ✅ (Upsert 用 clause.OnConflict) | +| 3.12 | Export 全部方法 | ~10 | ✅ | +| 3.13 | Import 全部方法 | ~10 | ✅ 已全部改为 GORM `Clauses(clause.OnConflict)`(见 §9.6) | +| **3.14** | **repository_mutations.go 全部方法 (~40 个)** | **~40** | **✅ 已全量改为 GORM 链式调用(见 §9.3)** | +| **3.15** | **repository_federation.go 查询方法** | **~8** | **✅ 已全部改为 GORM 链式调用** | +| **3.16** | **repository_control.go 复杂查询** | **~5** | **✅ 已全部改为 GORM 链式调用** | +| **3.17** | **repository_flow.go 查询方法** | **~5** | **✅ 已全部改为 GORM 链式调用** | +| **3.18** | **Jobs 查询方法 (repository.go 尾部)** | **~8** | **✅ 已 GORM 化** | + +### 阶段 4:消除 handler 中直接 SQL — 提取为 Repository 方法 ✅ 已完成 + +> **2026-02-16 审计**:handler 中的 SQL 已大部分提取到 repo 层,但这些 repo 方法本身仍使用 raw SQL(见阶段 3)。 +> **2026-02-17 更新**:`mutations.go` 直接 `tx.Exec`/`tx.Raw` 已从 27 处降至 0 处(生产代码),详见 §9.4。 + +mutations.go 和其他 handler 文件中大量直接操作 `h.repo.DB()` 执行 raw SQL,需要: +1. 将 SQL 逻辑提取为 Repository 方法 +2. Handler 只调用 Repository 方法 + +| 步骤 | 任务 | 文件 | 状态 | +|------|------|------|------| +| 4.1 | 用户 CRUD:userCreate, userUpdate, userDelete, userResetFlow | mutations.go | ✅ 已提取到 repo 方法 | +| 4.2 | 节点 CRUD:nodeCreate, nodeUpdate, nodeDelete, nodeBatch* | mutations.go | ✅ 已提取到 repo 方法 | +| 4.3 | 隧道 CRUD:tunnelCreate, tunnelUpdate, tunnelDelete, tunnelBatch* | mutations.go | ✅ tunnelCreate/Update 的 SQL 已下沉 repo | +| 4.4 | 转发 CRUD:forwardCreate, forwardUpdate, forwardDelete, forwardBatch* | mutations.go | ✅ 已提取到 repo (CreateForwardTx 等) | +| 4.5 | 限速 CRUD:speedLimitCreate, speedLimitUpdate, speedLimitDelete | mutations.go | ✅ 已提取到 repo 方法 | +| 4.6 | 分组 CRUD:所有 group* 方法 | mutations.go | ✅ 成员同步/权限管理 SQL 已下沉 repo | +| 4.7 | 用户隧道:userTunnelAssign, userTunnelRemove, userTunnelUpdate | mutations.go | ✅ 已提取到 repo 方法 | +| 4.8 | handler.go 中的直接 SQL (openAPISubStore 等) | handler.go | ✅ 已迁移(含 nil 检查清理) | +| 4.9 | federation.go 中的 raw SQL | federation.go | ✅ 已提取到 repo_federation.go | +| 4.10 | control_plane.go 中的 raw SQL | control_plane.go | ✅ 已提取到 repo_control.go | +| 4.11 | flow_policy.go 中的 raw SQL | flow_policy.go | ✅ 已提取到 repo_flow.go | +| 4.12 | jobs.go 中的 raw SQL | jobs.go | ✅ 已提取到 repo 方法(含 nil 检查清理) | + +### 阶段 5:清理旧代码 ✅ 已完成 + +| 步骤 | 任务 | 文件 | 状态 | +|------|------|------|------| +| 5.1 | 删除 `internal/store/postgres/` 整个目录 | 目录删除 | ✅ | +| 5.2 | 删除 `internal/store/sqlite/sql/` 目录 | 目录删除 | ✅ | +| 5.3 | 删除 `internal/store/db.go` SQL 重写层 | 文件删除 | ✅ | +| 5.4 | 删除 `internal/store/db_test.go` | 文件删除 | ✅ | +| 5.5 | 清理 repository.go 中不再需要的 embed 指令 | 清理 | ✅ | + +### 阶段 6:Package 重命名 ✅ 已完成 + +| 步骤 | 任务 | 文件 | 状态 | +|------|------|------|------| +| 6.1 | `internal/store/sqlite/` → `internal/store/repo/` | 目录重命名 | ✅ | +| 6.2 | 更新所有 import 路径:`store/sqlite` → `store/repo` (13处) | 全局替换 | ✅ | + +### 阶段 7:测试 + 验证 ⚠️ 部分完成 + +| 步骤 | 任务 | 状态 | +|------|------|------| +| 7.1 | 更新所有现有测试适配 GORM | ✅ 测试已适配 (使用 repo.DB() 做数据准备) | +| 7.2 | `go test ./...` 全部通过 | ✅ 已通过(含 `internal/http/handler`、`tests/contract`) | +| 7.3 | `make build` 构建成功 | ✅ 已通过 | + +### 阶段 8:文档更新 ✅ 已完成 + +| 步骤 | 任务 | 文件 | 状态 | +|------|------|------|------| +| 8.1 | 更新 `go-backend/AGENTS.md` — 移除 "DO NOT USE ORM",记录 GORM 规范 | AGENTS.md | ✅ | +| 8.2 | 更新根 `AGENTS.md` | AGENTS.md | ✅ | +| 8.3 | 更新 `handler/AGENTS.md` | AGENTS.md | ✅ | + +--- + +## 五、GORM 使用规范 + +### 5.1 查询模式 + +```go +// 单条查询 - 未找到返回 nil, nil (保持现有语义) +var user model.User +err := r.db.Where("id = ?", id).First(&user).Error +if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil +} + +// 列表查询 +var users []model.User +err := r.db.Where("role_id != ?", 0).Order("id ASC").Find(&users).Error + +// 创建 +err := r.db.Create(&user).Error + +// 更新 (部分字段) +err := r.db.Model(&model.User{}).Where("id = ?", id).Updates(map[string]interface{}{ + "user": username, "flow": flow, "updated_time": now, +}).Error + +// 事务 (closure pattern - 自动 rollback/commit) +err := r.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("user_id = ?", id).Delete(&model.Forward{}).Error; err != nil { + return err + } + return tx.Where("id = ?", id).Delete(&model.User{}).Error +}) + +// 原生 SQL (仅用于复杂查询和 PG 特有操作) +r.db.Exec("SELECT setval(?::regclass, ?, ?)", seqRef, maxID, true) +``` + +### 5.2 关键注意事项 + +1. **user 保留字**:通过 `TableName()` 返回 `"user"`,GORM 自动处理引号 +2. **SQLite MaxOpenConns**:必须设为 1 防止 "database locked" +3. **SQLite WAL 模式**:DSN 中配置 `_pragma=journal_mode(WAL)` +4. **不要用 `type:jsonb`**:SQLite 不支持,用 `serializer:json` +5. **不要用 `type:serial`**:让 GORM 从 `primaryKey` 自动推断 +6. **AutoMigrate 在 SQLite 中使用 copy-swap-drop**:大表慎用 + +--- + +## 六、影响范围 + +### 需要修改的文件 + +| 文件 | 修改类型 | 描述 | 当前状态 | +|------|----------|------|----------| +| `go.mod` / `go.sum` | 修改 | 添加 GORM + 驱动依赖 | ✅ | +| `internal/store/model/model.go` | **新增** | 全部 21 个 GORM Model | ✅ | +| `internal/store/repo/repository.go` | **重写** | 全部查询 GORM 化 | ⚠️ 业务查询已 GORM;仅剩 PG 序列修复 DDL `Exec` 4 处 | +| `internal/store/repo/repository_mutations.go` | **重写** | Mutation helpers | ✅ 全量 GORM(Raw=0) | +| `internal/store/repo/repository_federation.go` | **重写** | Federation 查询 | ✅ 已 GORM 化(Raw=0) | +| `internal/store/repo/repository_control.go` | **重写** | 控制面查询 | ✅ 已 GORM 化(Raw=0) | +| `internal/store/repo/repository_flow.go` | **重写** | 流量/转发查询 | ✅ 已 GORM 化(Raw=0) | +| `internal/http/handler/mutations.go` | **重写** | 全部 CRUD 提取到 repo | ✅ 生产代码 `tx.Exec/tx.Raw` = 0 | +| `internal/http/handler/handler.go` | 修改 | 更新 import、移除直接 SQL | ✅ (仅剩 nil check) | +| `internal/http/handler/federation.go` | 修改 | GORM 替代 raw SQL | ✅ | +| `internal/http/handler/control_plane.go` | 修改 | GORM 替代 raw SQL | ✅ | +| `internal/http/handler/flow_policy.go` | 修改 | GORM 替代 raw SQL | ✅ | +| `internal/http/handler/jobs.go` | 修改 | GORM 替代 raw SQL | ✅ (仅剩 nil check) | +| `internal/ws/server.go` | 修改 | 更新 import | ✅ | +| `internal/app/app.go` | 修改 | 更新 import | ✅ | +| `internal/store/postgres/` | **删除** | 不再需要 | ✅ | +| `internal/store/db.go` | **删除** | GORM 自动处理方言 | ✅ | +| `internal/store/db_test.go` | **删除** | 旧重写层测试 | ✅ | +| `internal/store/sqlite/sql/` | **删除** | AutoMigrate 替代 | ✅ | +| `tests/contract/*.go` | 修改 | 适配 GORM | ✅ | +| `AGENTS.md` (3处) | 更新 | 反映新架构 | ✅ | + +### 不需要修改的文件 + +- `internal/http/router.go` — 路由不变 +- `internal/config/config.go` — 配置不变 +- `internal/auth/` — 认证不变 +- `internal/security/` — 加密不变 +- `internal/http/middleware/` — 中间件不变 +- `internal/http/response/` — 响应格式不变 +- `Dockerfile`, `Makefile` — 构建不变 + +--- + +## 七、风险与缓解 + +| 风险 | 可能性 | 影响 | 缓解措施 | +|------|--------|------|----------| +| GORM AutoMigrate SQLite/PG 行为差异 | 中 | 高 | 先写 Model 验证双数据库 AutoMigrate | +| handler 中散落 raw SQL 遗漏 | 中 | 高 | 全局搜索 `.Exec(`, `.Query(`, `.QueryRow(` | +| 事务语义变化 | 低 | 中 | 逐方法对比旧代码事务边界 | +| 大量代码变更导致回归 | 高 | 高 | 分阶段提交,每阶段 `go test` | +| GORM 性能开销 | 低 | 低 | 此场景下可忽略 | +| SQLite "database locked" | 中 | 高 | `MaxOpenConns(1)` + WAL 模式 | + +--- + +## 八、迁移顺序原则 + +1. **先 Model 后查询**:确保 AutoMigrate 双数据库通过 +2. **先 Repository 后 Handler**:Handler 依赖 Repository +3. **先核心后边缘**:User → Node → Tunnel → Forward → 分组 → Federation +4. **每步编译**:每完成一组方法确保 `go build ./...` 通过 +5. **最后清理**:全部重写完成后再删除旧代码和重命名 package + +--- + +--- + +## 九、2026-02-16 审计发现 + 2026-02-17 进展记录 + +### 9.1 总体完成度 + +| 指标 | 数值 | +|------|------| +| 阶段完成数 | 7/8 完成 (1, 2, 4, 5, 6, 7, 8),1/8 部分完成 (3) | +| GORM 链式调用 | ~226 处 | +| Raw SQL 调用 (`.Exec`/`.Raw`+`.Scan`) | 4 处(生产代码) | +| GORM 占比 | ~98% | +| Handler 内 `tx.Exec`/`tx.Raw` | 0 处(生产代码) | +| `last_insert_rowid()` 生产代码 | 0 处(已消灭) | + +### 9.2 ✅ P0:`last_insert_rowid()`(生产代码)已清零 + +`last_insert_rowid()` 已从生产路径移除,创建主键统一改为 `Create(&model)` 自动回填 ID, +确保 SQLite / PostgreSQL 双数据库行为一致。 + +> 备注:测试代码中的历史 SQL 兼容性用例可在后续测试清理阶段单独处理。 + +### 9.3 ✅ P1:`repository_mutations.go` 已全量 GORM 化 + +本次已完成 `repository_mutations.go` 的集中清理: + +1. User / Node / Tunnel / Forward / UserTunnel / SpeedLimit / Group / Permission 全部 mutation 方法改为 GORM 链式调用。 +2. 事务内级联删除统一为 `tx.Where(...).Delete(&Model{})` 模式。 +3. `ON CONFLICT DO NOTHING` 统一替换为 `Clauses(clause.OnConflict{DoNothing: true})`。 +4. 保留原有调用语义(含 `sql.ErrNoRows` 行为兼容)并完成 `go build ./...` 验证。 + +> 当前 `repository_mutations.go` 中生产代码 `.Raw(`/`.Exec(` 调用已降为 0。 + +### 9.4 ✅ P2:Handler `mutations.go` 直接 SQL 已清零 + +2026-02-17 本轮静态扫描结果:`mutations.go` **0 处** `tx.Exec`/`tx.Raw`(生产代码)。 + +本轮完成下沉到 repo 的逻辑: + +- `tunnelUpdate` 中 `UPDATE tunnel` + `DELETE chain_tunnel` +- `isRemoteNodeTx` 查询 +- `pickNodePortTx` 的 node/chain_tunnel/forward_port 端口占用查询 +- `replaceTunnelChainsTx` 的 chain_tunnel 写入 +- 分组成员同步(`tunnel_group_tunnel` / `user_group_user`) +- 权限删除与 grant 回收(`group_permission` / `group_permission_grant` / `user_tunnel`) +- federation 绑定替换(`federation_tunnel_binding`) + +### 9.5 ✅ P3(部分):已移除 `QueryInt64List` / `QueryPairs` SQL 透传 + +- `repository_mutations.go` 中两个 SQL 透传入口已删除。 +- Handler 已切换为语义化 repo 方法: + - `ListUserIDsByUserGroup` + - `ListTunnelIDsByTunnelGroup` + - `ListGroupPermissionPairsByUserGroup` + - `ListGroupPermissionPairsByTunnelGroup` + +### 9.6 ✅ P3:Import 函数已全部 GORM 化 + +`repository.go` 中 Import 相关函数已完成迁移: + +- `importUsers` +- `importNodes` +- `importTunnels`(含 `chain_tunnel` 子项 upsert) +- `importForwards`(含 `forward_port` 覆盖写入) +- `importUserTunnels` +- `importSpeedLimits` +- `importTunnelGroups` +- `importUserGroups` +- `importPermissions` +- `importConfigs`(原本已是 GORM) + +迁移后统一采用 `Clauses(clause.OnConflict{Columns: id/name, DoUpdates: ...}).Create(&model)` 模式, +保留原 `ON CONFLICT ... DO UPDATE` 语义;Import 区段 `tx.Exec`/`tx.Raw` 已清零。 + +### 9.7 ✅ P4:`h.repo.DB() == nil` 检查已清理 + +`internal/http/handler/` 下已无 `h.repo.DB()` 直接访问;handler 仅通过语义化 repo 方法进行数据访问。 + +### 9.8 ✅ P5:Repository 层 Raw 已收敛(仅保留 PG 序列修复 DDL) + +当前生产代码中 `.Raw()` 已清零;仅剩 `repository.go` 的 4 处 `Exec()`,全部位于 PG 序列修复 DDL: + +- `CREATE SEQUENCE IF NOT EXISTS ...` +- `ALTER TABLE ... ALTER COLUMN id SET DEFAULT nextval(...)` +- `ALTER SEQUENCE ... OWNED BY ...` +- `SELECT setval(...::regclass, ?, ?)` + +以上 4 处属于数据库管理 DDL/序列同步语义,当前保留,不再继续向 GORM 链式调用替换。 + +`repository_federation.go` / `repository_control.go` / `repository_flow.go` 已完成 GORM 化(Raw=0)。 + +--- + +## 十、后续工作优先级 + +| 优先级 | 任务 | 影响范围 | 工作量 | +|--------|------|----------|--------| +| **P0** | ✅ 已完成:生产代码中 `last_insert_rowid()` 清零(测试用例待单独清理) | 6 处生产(已完成) | 完成 | +| **P1** | ✅ 已完成:`repository_mutations.go` ~40 方法改为 GORM 链式调用 | 659 行(已完成) | 完成 | +| **P2** | ✅ 已完成:`mutations.go` handler 直接 SQL 全部提取为 repo 方法 | mutations.go | 完成 | +| **P3** | ✅ 已完成:移除 `QueryInt64List`/`QueryPairs` 透传,切换语义化 repo 方法 | 2 个方法 + 调用方(已完成) | 完成 | +| **P3** | ✅ 已完成:Import 函数 Raw SQL 改为 GORM `Clauses(clause.OnConflict{}).Create()` | 9 个函数(已完成) | 完成 | +| **P4** | ✅ 已完成:`h.repo.DB() == nil` 检查清理完毕 | 4 处(已完成) | 完成 | +| **P5** | ✅ 已完成:repo 查询层 Raw 清零,`repository.go` 保留 4 处 PG 序列修复 DDL `Exec`(设计保留) | repository.go | 完成 | +| **P5** | ✅ 已完成:更新 MIGRATION_PLAN.md 状态标记与收尾记录 | 本文件 | 完成 | + +### 10.5 本轮执行记录(2026-02-17,P5 schema 收尾) + +1. 完成 `repository.go` schema 迁移段去 Raw: + - `normalizeStrategy` 改为 `Model(...).Where(...).Update(...)` + - `ensurePostgresIDDefaults`/`ensurePostgresTableIDDefault` 的 information_schema 查询改为 GORM `Table+Joins+Where+Scan` + - `syncPostgresTableIDSequence` 的 `MAX(id)` 查询改为 GORM `Table+Select+Scan` +2. 复扫结果: + - `repository.go` `.Raw()` = 0 + - repo 生产路径剩余 `.Exec()` = 4(全部为 PG 序列修复 DDL) +3. 验证结果: + - `go build ./...` ✅ + - `go test ./internal/store/repo/...` ✅ + +### 10.6 本轮执行记录(2026-02-17,测试/构建收尾) + +1. 修复事务内 SQLite 连接阻塞(`MaxOpenConns(1)` 场景): + - 新增 `GetNodeRecordTx` 并在 `prepareTunnelCreateState` 使用事务句柄读取节点。 + - 新增 `GetNodeRemoteFieldsTx` 并在 `tunnelCreate` 事务内改用事务句柄读取远端字段。 + - `applyFederationRuntime` 改为显式接收 `localDomain`,避免事务内再次走 `repo.GetConfigByName`。 +2. 修复 legacy SQLite schema 迁移契约: + - 新增 `prepareSQLiteLegacyColumns` 预补齐 `node/tunnel` 关键列。 + - SQLite 模式下对已存在 `node/tunnel` 表跳过对应 `AutoMigrate` 重建流程,避免 `node__temp.name` 约束失败。 +3. 验证结果: + - `go test ./internal/http/handler/...` ✅ + - `go test ./tests/contract/...` ✅ + - `go test ./...` ✅ + - `go build ./...` ✅ + - `make build` ✅ + +### 10.1 本轮执行记录(2026-02-17,P5 查询层) + +1. 完成 `repository_federation.go` 全量 GORM 化: + - `ListRemoteNodes` / `UpdateNodeRemoteConfig` + - `ListActiveBindingsForNode` / `GetNodeBasicInfo` + - `ListUsedPortsOnNode` / `ListTunnelIDsByNamePrefix` / `NextIndex` +2. 完成 `repository_control.go` 全量 GORM 化: + - `ListForwardsByTunnel` / `ListForwardPorts` / `GetTunnelOutProtocol` + - `ResolveUserTunnelAndLimiter` / `ListChainNodesForTunnel` +3. 完成 `repository_flow.go` 全量 GORM 化: + - `ListActiveForwardsByUser` / `ListActiveForwardsByUserTunnel` + - `GetForwardRecord` / `GetTunnelRecord` +4. 复扫结果: + - `repository_federation.go` Raw/Exec = 0 + - `repository_control.go` Raw/Exec = 0 + - `repository_flow.go` Raw/Exec = 0 + - repo 生产路径剩余 Raw/Exec = 9(全部在 `repository.go`) +5. 验证结果: + - `go build ./...` ✅ + - `go test ./internal/store/repo/...` ✅ + +### 10.2 本轮执行记录(2026-02-17) + +1. 完成 P3 Import 9 个函数的 GORM 化(`repository.go`),并保持 `ON CONFLICT` 语义一致。 +2. 复扫确认:`repository.go` Import 区段 `tx.Exec`/`tx.Raw` 已清零。 +3. 验证结果: + - `go build ./...` ✅(使用显式 `GOMODCACHE/GOPATH/GOCACHE/HOME` 环境) + - `go test ./internal/store/repo/...` ✅ + +### 10.3 本轮执行记录(2026-02-17,P2 部分) + +1. 将 tunnel 更新/chain 重建路径 SQL 下沉到 `repository_mutations.go`: + - 新增 `UpdateTunnelTx` + - 新增 `DeleteChainTunnelsByTunnelTx` + - 新增 `CreateChainTunnelTx` +2. 将 handler 内部 SQL helper 迁移到 repo: + - 新增 `IsRemoteNodeTx` + - 新增 `PickNodePortTx` + - `replaceTunnelChainsTx` 改为 handler 方法并改用 repo 调用,不再直接 SQL +3. 复扫结果:`mutations.go` 直接 SQL 从 27 处降至 17 处。 +4. 验证结果: + - `go build ./...` ✅ + - `go test ./internal/store/repo/...` ✅ + +### 10.4 本轮执行记录(2026-02-17,P2 收尾) + +1. 新增并落地事务语义化 repo 方法: + - `ReplaceTunnelGroupMembersTx` / `ReplaceUserGroupMembersTx` + - `ListUserIDsByUserGroupTx` + - `GetGroupPermissionPairByIDTx` / `DeleteGroupPermissionByIDTx` + - `RevokeGroupGrantsForRemovedUsersTx` / `RevokeGroupPermissionPairTx` + - `ReplaceFederationTunnelBindingsTx` +2. 删除 handler 内 SQL helper(`queryInt64ListTx` / `revokeGroupGrantsForRemovedUsersTx` / `revokeGroupPermissionPairTx` / `replaceFederationTunnelBindingsTx`)。 +3. 复扫确认:`mutations.go` 生产路径 `tx.Exec`/`tx.Raw` = 0。 +4. 验证结果: + - `go build ./...` ✅ + - `go test ./internal/store/repo/...` ✅ + +--- + +*本文档将随迁移进展实时更新状态标记。* +*最后审计时间:2026-02-17,审计工具:代码静态分析 (grep/AST) + go build/go test 验证* diff --git a/go-backend/go.mod b/go-backend/go.mod index 4295d62..1dbadb3 100644 --- a/go-backend/go.mod +++ b/go-backend/go.mod @@ -12,10 +12,14 @@ require ( require ( github.com/dustin/go-humanize v1.0.1 // indirect + github.com/glebarez/go-sqlite v1.21.2 // indirect + github.com/glebarez/sqlite v1.11.0 // 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/jinzhu/inflection v1.0.0 // indirect + github.com/jinzhu/now v1.1.5 // 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 @@ -24,6 +28,8 @@ require ( golang.org/x/sync v0.17.0 // indirect golang.org/x/sys v0.33.0 // indirect golang.org/x/text v0.29.0 // indirect + gorm.io/driver/postgres v1.6.0 // indirect + gorm.io/gorm v1.31.1 // 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 207fe9e..3077428 100644 --- a/go-backend/go.sum +++ b/go-backend/go.sum @@ -3,6 +3,10 @@ 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/glebarez/go-sqlite v1.21.2 h1:3a6LFC4sKahUunAmynQKLZceZCOzUthkRkEAl9gAXWo= +github.com/glebarez/go-sqlite v1.21.2/go.mod h1:sfxdZyhQjTM2Wry3gVYWaW072Ri1WMdWJi0k6+3382k= +github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GMw= +github.com/glebarez/sqlite v1.11.0/go.mod h1:h8/o8j5wiAsqSPoWELDUdJXhjAhsVliSn7bWZjOhrgQ= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= @@ -17,6 +21,10 @@ 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/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E= +github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= +github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= +github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= 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= @@ -49,6 +57,10 @@ gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8 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= +gorm.io/driver/postgres v1.6.0 h1:2dxzU8xJ+ivvqTRph34QX+WrRaJlmfyPqXmoGVjMBa4= +gorm.io/driver/postgres v1.6.0/go.mod h1:vUw0mrGgrTK+uPHEhAdV4sfFELrByKVGnaVRkXDhtWo= +gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg= +gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs= 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 527b866..5c013a8 100644 --- a/go-backend/internal/app/app.go +++ b/go-backend/internal/app/app.go @@ -10,30 +10,30 @@ import ( "go-backend/internal/config" httpserver "go-backend/internal/http" "go-backend/internal/http/handler" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) type App struct { cfg config.Config server *http.Server - repo *sqlite.Repository + repo *repo.Repository h *handler.Handler } func New(cfg config.Config) (*App, error) { var ( - repo *sqlite.Repository - err error + r *repo.Repository + err error ) switch strings.ToLower(strings.TrimSpace(cfg.DBType)) { case "", "sqlite": - repo, err = sqlite.Open(cfg.DBPath) + r, err = repo.Open(cfg.DBPath) if err != nil { return nil, fmt.Errorf("open sqlite: %w", err) } case "postgres", "postgresql": - repo, err = sqlite.OpenPostgres(cfg.DatabaseURL) + r, err = repo.OpenPostgres(cfg.DatabaseURL) if err != nil { return nil, fmt.Errorf("open postgres: %w", err) } @@ -41,7 +41,7 @@ func New(cfg config.Config) (*App, error) { return nil, fmt.Errorf("unsupported DB_TYPE %q", cfg.DBType) } - h := handler.New(repo, cfg.JWTSecret) + h := handler.New(r, cfg.JWTSecret) router := httpserver.NewRouter(h, cfg.JWTSecret) s := &http.Server{ @@ -53,7 +53,7 @@ func New(cfg config.Config) (*App, error) { IdleTimeout: 60 * time.Second, } - return &App{cfg: cfg, server: s, repo: repo, h: h}, nil + return &App{cfg: cfg, server: s, repo: r, h: h}, nil } func (a *App) Run() error { diff --git a/go-backend/internal/http/handler/AGENTS.md b/go-backend/internal/http/handler/AGENTS.md index d447ea5..5d14ed9 100644 --- a/go-backend/internal/http/handler/AGENTS.md +++ b/go-backend/internal/http/handler/AGENTS.md @@ -4,7 +4,7 @@ ## OVERVIEW HTTP request handlers for FLVX Admin API. Core business logic layer. -**Stack:** Go 1.23, net/http, raw SQL (no ORM). +**Stack:** Go 1.23, net/http, GORM via Repository pattern. ## STRUCTURE ``` @@ -28,13 +28,14 @@ handler/ | **Background Jobs** | `jobs.go` | Scheduled sync/cleanup tasks | ## CONVENTIONS -- Inherits from parent: raw SQL, no ORM, JWT in Authorization header. +- Inherits from parent: GORM via Repository pattern, JWT in Authorization header. - Large files expected (`mutations.go` 3716 LOC - central mutation hub). -- Uses `sqlite.Repository` for DB access via `repo.XXX()` methods. +- Uses `repo.Repository` for DB access via `h.repo.XXX()` methods. +- Handlers never call `repo.DB()` directly — all queries go through Repository methods. - Domain-driven file split: one file per functional area (federation, jobs, etc.). ## ANTI-PATTERNS -- Do NOT add ORM here - uses raw SQL throughout. +- Do NOT let handlers call `repo.DB()` directly — add a Repository method instead. - Do NOT change handler signatures without updating router.go. ## COMMANDS diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 89f82e0..0c0e3f6 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -1,7 +1,6 @@ package handler import ( - "database/sql" "errors" "fmt" "net" @@ -12,61 +11,18 @@ import ( "time" "go-backend/internal/http/client" + "go-backend/internal/store/model" "go-backend/internal/ws" ) var errForwardNotFound = errors.New("forward not found") -type forwardRecord struct { - ID int64 - UserID int64 - UserName string - Name string - TunnelID int64 - RemoteAddr string - Strategy string - Status int -} +type forwardRecord = model.ForwardRecord +type tunnelRecord = model.TunnelRecord +type forwardPortRecord = model.ForwardPortRecord +type nodeRecord = model.NodeRecord -type tunnelRecord struct { - ID int64 - Type int - Status int - Flow int64 - TrafficRatio float64 -} - -type forwardPortRecord struct { - NodeID int64 - Port int -} - -type nodeRecord struct { - ID int64 - Name string - ServerIP string - ServerIPv4 string - ServerIPv6 string - Status int - PortRange string - TCPListenAddr string - UDPListenAddr string - InterfaceName string - IsRemote int - RemoteURL string - RemoteToken string - RemoteConfig string -} - -type chainNodeRecord struct { - ChainType int - Inx int64 - NodeID int64 - Port int - NodeName string - Protocol string - Strategy string -} +type chainNodeRecord = model.ChainNodeRecord type diagnosisTarget struct { Address string @@ -101,247 +57,82 @@ func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int6 if roleID == 0 { return nil } - var count int - err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? AND status = 1`, userID, tunnelID).Scan(&count) + ok, err := h.repo.UserTunnelExistsByUserAndTunnel(userID, tunnelID) if err != nil { return err } - if count <= 0 { + if !ok { return errors.New("你没有该隧道的权限") } return nil } func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) { - row := h.repo.DB().QueryRow(` - SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status - FROM forward WHERE id = ? LIMIT 1 - `, forwardID) - var fr forwardRecord - err := row.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status) + fr, err := h.repo.GetForwardRecord(forwardID) if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, errForwardNotFound - } return nil, err } - if strings.TrimSpace(fr.Strategy) == "" { - fr.Strategy = "fifo" + if fr == nil { + return nil, errForwardNotFound } - return &fr, nil + return fr, nil } func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) { - row := h.repo.DB().QueryRow(`SELECT id, type, status, flow, traffic_ratio FROM tunnel WHERE id = ? LIMIT 1`, tunnelID) - var tr tunnelRecord - err := row.Scan(&tr.ID, &tr.Type, &tr.Status, &tr.Flow, &tr.TrafficRatio) + tr, err := h.repo.GetTunnelRecord(tunnelID) if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, errors.New("隧道不存在") - } return nil, err } - if tr.Flow <= 0 { - tr.Flow = 1 + if tr == nil { + return nil, errors.New("隧道不存在") } - if tr.TrafficRatio <= 0 { - tr.TrafficRatio = 1 - } - return &tr, nil + return tr, nil } func (h *Handler) listForwardsByTunnel(tunnelID int64) ([]forwardRecord, error) { - rows, err := h.repo.DB().Query(` - SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status - FROM forward - WHERE tunnel_id = ? - ORDER BY id ASC - `, tunnelID) - if err != nil { - return nil, err - } - defer rows.Close() - - result := make([]forwardRecord, 0) - for rows.Next() { - var fr forwardRecord - if err := rows.Scan(&fr.ID, &fr.UserID, &fr.UserName, &fr.Name, &fr.TunnelID, &fr.RemoteAddr, &fr.Strategy, &fr.Status); err != nil { - return nil, err - } - if strings.TrimSpace(fr.Strategy) == "" { - fr.Strategy = "fifo" - } - result = append(result, fr) - } - if err := rows.Err(); err != nil { - return nil, err - } - return result, nil + return h.repo.ListForwardsByTunnel(tunnelID) } func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error) { - rows, err := h.repo.DB().Query(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID) - if err != nil { - return nil, err - } - defer rows.Close() - - result := make([]forwardPortRecord, 0) - for rows.Next() { - var item forwardPortRecord - if err := rows.Scan(&item.NodeID, &item.Port); err != nil { - return nil, err - } - result = append(result, item) - } - if err := rows.Err(); err != nil { - return nil, err - } - return result, nil + return h.repo.ListForwardPorts(forwardID) } func (h *Handler) isTunnelSelectedTLSProtocol(tunnelID int64) (bool, error) { - row := h.repo.DB().QueryRow(` - SELECT protocol - FROM chain_tunnel - WHERE tunnel_id = ? AND chain_type = '3' - ORDER BY id ASC - LIMIT 1 - `, tunnelID) - - var protocol sql.NullString - if err := row.Scan(&protocol); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return false, nil - } + protocol, err := h.repo.GetTunnelOutProtocol(tunnelID) + if err != nil { return false, err } - - return isTLSTunnelProtocol(protocol.String), nil + return isTLSTunnelProtocol(protocol), nil } func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) { - row := h.repo.DB().QueryRow(` - SELECT id, name, server_ip, server_ip_v4, server_ip_v6, status, port, tcp_listen_addr, udp_listen_addr, interface_name, is_remote, remote_url, remote_token, remote_config - FROM node - WHERE id = ? - LIMIT 1 - `, nodeID) - var n nodeRecord - var serverIPv4 sql.NullString - var serverIPv6 sql.NullString - var portRange sql.NullString - var tcpListen sql.NullString - var udpListen sql.NullString - var iface sql.NullString - var remoteURL sql.NullString - var remoteToken sql.NullString - var remoteConfig sql.NullString - err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Status, &portRange, &tcpListen, &udpListen, &iface, &n.IsRemote, &remoteURL, &remoteToken, &remoteConfig) + n, err := h.repo.GetNodeRecord(nodeID) if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, errors.New("节点不存在") - } return nil, err } - n.ServerIPv4 = strings.TrimSpace(serverIPv4.String) - n.ServerIPv6 = strings.TrimSpace(serverIPv6.String) - n.PortRange = strings.TrimSpace(portRange.String) - n.TCPListenAddr = strings.TrimSpace(tcpListen.String) - n.UDPListenAddr = strings.TrimSpace(udpListen.String) - n.InterfaceName = strings.TrimSpace(iface.String) - n.RemoteURL = strings.TrimSpace(remoteURL.String) - n.RemoteToken = strings.TrimSpace(remoteToken.String) - n.RemoteConfig = strings.TrimSpace(remoteConfig.String) - if n.TCPListenAddr == "" { - n.TCPListenAddr = "[::]" + if n == nil { + return nil, errors.New("节点不存在") } - if n.UDPListenAddr == "" { - n.UDPListenAddr = "[::]" - } - if strings.TrimSpace(n.Name) == "" { - n.Name = fmt.Sprintf("node_%d", n.ID) - } - return &n, nil + return n, nil } func (h *Handler) resolveUserTunnelAndLimiter(userID, tunnelID int64) (int64, *int64, *int, error) { - row := h.repo.DB().QueryRow(` - SELECT ut.id, sl.id, sl.speed - FROM user_tunnel ut - LEFT JOIN speed_limit sl ON sl.id = ut.speed_id - WHERE ut.user_id = ? AND ut.tunnel_id = ? - ORDER BY ut.id ASC - LIMIT 1 - `, userID, tunnelID) - var userTunnelID int64 - var limiterID sql.NullInt64 - var speed sql.NullInt64 - err := row.Scan(&userTunnelID, &limiterID, &speed) + info, err := h.repo.ResolveUserTunnelAndLimiter(userID, tunnelID) if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return 0, nil, nil, nil - } return 0, nil, nil, err } - if !limiterID.Valid || limiterID.Int64 <= 0 { - return userTunnelID, nil, nil, nil + if info == nil { + return 0, nil, nil, nil } - v := limiterID.Int64 - s := int(speed.Int64) - return userTunnelID, &v, &s, nil + return info.UserTunnelID, info.LimiterID, info.Speed, nil } func (h *Handler) listUserTunnelIDs(userID, tunnelID int64) ([]int64, error) { - rows, err := h.repo.DB().Query(` - SELECT id - FROM user_tunnel - WHERE user_id = ? AND tunnel_id = ? - ORDER BY id ASC - `, userID, tunnelID) - if err != nil { - return nil, err - } - defer rows.Close() - - out := make([]int64, 0) - for rows.Next() { - var id int64 - if err := rows.Scan(&id); err != nil { - return nil, err - } - out = append(out, id) - } - if err := rows.Err(); err != nil { - return nil, err - } - return out, nil + return h.repo.ListUserTunnelIDs(userID, tunnelID) } func (h *Handler) listUserTunnelIDsByUser(userID int64) ([]int64, error) { - rows, err := h.repo.DB().Query(` - SELECT id - FROM user_tunnel - WHERE user_id = ? - ORDER BY id ASC - `, userID) - if err != nil { - return nil, err - } - defer rows.Close() - - out := make([]int64, 0) - for rows.Next() { - var id int64 - if err := rows.Scan(&id); err != nil { - return nil, err - } - out = append(out, id) - } - if err := rows.Err(); err != nil { - return nil, err - } - return out, nil + return h.repo.ListUserTunnelIDsByUser(userID) } func (h *Handler) syncForwardServices(forward *forwardRecord, method string, allowFallbackAdd bool) error { @@ -663,13 +454,13 @@ func (h *Handler) diagnoseTunnelRuntime(tunnelID int64) (map[string]interface{}, return nil, err } - var tunnelName string - if err := h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, errors.New("隧道不存在") - } + tunnelName, err := h.repo.GetTunnelName(tunnelID) + if err != nil { return nil, err } + if tunnelName == "" { + return nil, errors.New("隧道不存在") + } chainRows, err := h.listChainNodesForTunnel(tunnelID) if err != nil { @@ -960,40 +751,7 @@ func firstPortFromRange(portRange string) int { } func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) { - rows, err := h.repo.DB().Query(` - SELECT CAST(ct.chain_type AS INTEGER), COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy - FROM chain_tunnel ct - LEFT JOIN node n ON n.id = ct.node_id - WHERE ct.tunnel_id = ? - ORDER BY CAST(ct.chain_type AS INTEGER) ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC - `, tunnelID) - if err != nil { - return nil, err - } - defer rows.Close() - - result := make([]chainNodeRecord, 0) - for rows.Next() { - var item chainNodeRecord - var name sql.NullString - var protocol sql.NullString - var strategy sql.NullString - if err := rows.Scan(&item.ChainType, &item.Inx, &item.NodeID, &item.Port, &name, &protocol, &strategy); err != nil { - return nil, err - } - if strings.TrimSpace(name.String) == "" { - item.NodeName = fmt.Sprintf("node_%d", item.NodeID) - } else { - item.NodeName = name.String - } - item.Protocol = defaultString(protocol.String, "tls") - item.Strategy = defaultString(strategy.String, "round") - result = append(result, item) - } - if err := rows.Err(); err != nil { - return nil, err - } - return result, nil + return h.repo.ListChainNodesForTunnel(tunnelID) } func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int) (map[string]interface{}, error) { diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 9cfdc7f..4f03c99 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -1,7 +1,6 @@ package handler import ( - "database/sql" "encoding/json" "fmt" "net" @@ -14,7 +13,7 @@ import ( "go-backend/internal/http/client" "go-backend/internal/http/response" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) type federationTunnelRequest struct { @@ -108,7 +107,7 @@ type peerShareUsedPort struct { } type peerShareListItem struct { - sqlite.PeerShare + repo.PeerShare UsedPorts []int `json:"usedPorts"` UsedPortDetails []peerShareUsedPort `json:"usedPortDetails"` ActiveRuntimeNum int `json:"activeRuntimeNum"` @@ -264,7 +263,7 @@ func (h *Handler) federationShareCreate(w http.ResponseWriter, r *http.Request) now := time.Now().UnixMilli() token := randomToken(32) - share := &sqlite.PeerShare{ + share := &repo.PeerShare{ Name: req.Name, NodeID: req.NodeID, Token: token, @@ -430,40 +429,25 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque return } - rows, err := h.repo.DB().Query(` - SELECT id, name, remote_url, remote_token, remote_config - FROM node - WHERE is_remote = 1 - ORDER BY id DESC - `) + remoteNodes, err := h.repo.ListRemoteNodes() if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - defer rows.Close() fc := client.NewFederationClient() localDomain := h.federationLocalDomain() items := make([]remoteUsageNodeItem, 0) - for rows.Next() { - var ( - nodeID int64 - nodeName string - remoteURL sql.NullString - remoteToken sql.NullString - remoteConfig sql.NullString - ) - if err := rows.Scan(&nodeID, &nodeName, &remoteURL, &remoteToken, &remoteConfig); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } + for _, node := range remoteNodes { + nodeID := node.ID + nodeName := node.Name - shareID, maxBandwidth, currentFlow, expiryTime, portRangeStart, portRangeEnd := parseRemoteShareUsageConfig(remoteConfig.String) + shareID, maxBandwidth, currentFlow, expiryTime, portRangeStart, portRangeEnd := parseRemoteShareUsageConfig(node.RemoteConfig.String) var syncError string - url := strings.TrimSpace(remoteURL.String) - token := strings.TrimSpace(remoteToken.String) + url := strings.TrimSpace(node.RemoteURL.String) + token := strings.TrimSpace(node.RemoteToken.String) if url != "" && token != "" { info, connectErr := fc.Connect(url, token, localDomain) if connectErr != nil { @@ -484,42 +468,34 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque "portRangeStart": info.PortRangeStart, "portRangeEnd": info.PortRangeEnd, }) - _, _ = h.repo.DB().Exec(`UPDATE node SET remote_config = ? WHERE id = ?`, string(configData), nodeID) + _ = h.repo.UpdateNodeRemoteConfig(nodeID, string(configData)) } } - bindingRows, err := h.repo.DB().Query(` - SELECT fb.id, fb.tunnel_id, COALESCE(t.name, ''), fb.chain_type, fb.hop_inx, fb.allocated_port, fb.resource_key, fb.remote_binding_id, fb.updated_time - FROM federation_tunnel_binding fb - LEFT JOIN tunnel t ON t.id = fb.tunnel_id - WHERE fb.node_id = ? AND fb.status = 1 - ORDER BY fb.allocated_port ASC, fb.id ASC - `, nodeID) + bindingRows, err := h.repo.ListActiveBindingsForNode(nodeID) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } usedSet := make(map[int]struct{}) - bindings := make([]remoteUsageBindingItem, 0) - for bindingRows.Next() { - var item remoteUsageBindingItem - if err := bindingRows.Scan(&item.BindingID, &item.TunnelID, &item.TunnelName, &item.ChainType, &item.HopInx, &item.AllocatedPort, &item.ResourceKey, &item.RemoteBindingID, &item.UpdatedTime); err != nil { - _ = bindingRows.Close() - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - bindings = append(bindings, item) - if item.AllocatedPort > 0 { - usedSet[item.AllocatedPort] = struct{}{} + bindings := make([]remoteUsageBindingItem, 0, len(bindingRows)) + for _, b := range bindingRows { + bindings = append(bindings, remoteUsageBindingItem{ + BindingID: b.ID, + TunnelID: b.TunnelID, + TunnelName: b.TunnelName, + ChainType: b.ChainType, + HopInx: b.HopInx, + AllocatedPort: b.AllocatedPort, + ResourceKey: b.ResourceKey, + RemoteBindingID: b.RemoteBindingID, + UpdatedTime: b.UpdatedTime, + }) + if b.AllocatedPort > 0 { + usedSet[b.AllocatedPort] = struct{}{} } } - if err := bindingRows.Err(); err != nil { - _ = bindingRows.Close() - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - _ = bindingRows.Close() usedPorts := make([]int, 0, len(usedSet)) for port := range usedSet { @@ -543,10 +519,6 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque SyncError: syncError, }) } - if err := rows.Err(); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } response.WriteJSON(w, response.OK(items)) } @@ -639,31 +611,21 @@ func (h *Handler) nodeImport(w http.ResponseWriter, r *http.Request) { portRange = fmt.Sprintf("%d-%d", info.PortRangeStart, info.PortRangeEnd) } - db := h.repo.DB() - inx := nextIndex(db, "node") + inx := h.repo.NextIndex("node") now := time.Now().UnixMilli() - _, err = db.Exec(` - INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, 0, 0, 0, ?, ?, ?, ?, ?, ?, 1, ?, ?, ?) - `, + if err = h.repo.CreateRemoteNode( fmt.Sprintf("%s (Remote)", info.NodeName), - randomToken(16), // Dummy secret + randomToken(16), info.ServerIP, - "", "", // v4/v6 unknown, use server_ip portRange, - "", - "", - now, now, + now, info.Status, - "[::]", "[::]", inx, req.RemoteURL, req.Token, string(configBytes), - ) - - if err != nil { + ); err != nil { response.WriteJSON(w, response.Err(-2, "Database error: "+err.Error())) return } @@ -755,11 +717,7 @@ func (h *Handler) federationConnect(w http.ResponseWriter, r *http.Request) { return } - var nodeName string - var serverIP string - var status int - - err = h.repo.DB().QueryRow("SELECT name, server_ip, status FROM node WHERE id = ?", share.NodeID).Scan(&nodeName, &serverIP, &status) + nodeInfo, err := h.repo.GetNodeBasicInfo(share.NodeID) if err != nil { response.WriteJSON(w, response.Err(-2, "Node not found")) return @@ -769,9 +727,9 @@ func (h *Handler) federationConnect(w http.ResponseWriter, r *http.Request) { "shareId": share.ID, "shareName": share.Name, "nodeId": share.NodeID, - "nodeName": nodeName, - "serverIp": serverIP, - "status": status, + "nodeName": nodeInfo.Name, + "serverIp": nodeInfo.ServerIP, + "status": nodeInfo.Status, "maxBandwidth": share.MaxBandwidth, "currentFlow": share.CurrentFlow, "expiryTime": share.ExpiryTime, @@ -808,45 +766,20 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request) return } - tunnelType := 1 - - tx, err := h.repo.DB().Begin() - if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - defer tx.Rollback() - now := time.Now().UnixMilli() - tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?)`, + tunnelID, err := h.repo.CreateFederationTunnel( fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort), - tunnelType, + 1, req.Protocol, now, - now, - "", - ) - if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - - _, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, '1', ?, ?, 'fifo', 0, ?)`, - tunnelID, share.NodeID, req.RemotePort, - req.Protocol, ) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := tx.Commit(); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5) response.WriteJSON(w, response.OK(map[string]interface{}{ @@ -927,7 +860,7 @@ func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Re return } - runtime := &sqlite.PeerShareRuntime{ + runtime := &repo.PeerShareRuntime{ ShareID: share.ID, NodeID: share.NodeID, ReservationID: randomToken(24), @@ -981,7 +914,7 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ return } - var runtime *sqlite.PeerShareRuntime + var runtime *repo.PeerShareRuntime if strings.TrimSpace(req.ReservationID) != "" { runtime, err = h.repo.GetPeerShareRuntimeByReservationID(share.ID, strings.TrimSpace(req.ReservationID)) } else { @@ -1152,7 +1085,7 @@ func (h *Handler) federationRuntimeReleaseRole(w http.ResponseWriter, r *http.Re return } - var runtime *sqlite.PeerShareRuntime + var runtime *repo.PeerShareRuntime if strings.TrimSpace(req.BindingID) != "" { runtime, err = h.repo.GetPeerShareRuntimeByBindingID(share.ID, strings.TrimSpace(req.BindingID)) } else if strings.TrimSpace(req.ReservationID) != "" { @@ -1299,7 +1232,7 @@ func isFederationServiceCommand(commandType string) bool { } } -func validateFederationCommandPorts(share *sqlite.PeerShare, data interface{}) error { +func validateFederationCommandPorts(share *repo.PeerShare, data interface{}) error { if share == nil || (share.PortRangeStart <= 0 && share.PortRangeEnd <= 0) { return nil } @@ -1339,7 +1272,7 @@ func validateFederationCommandPorts(share *sqlite.PeerShare, data interface{}) e return nil } -func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) (int, error) { +func (h *Handler) pickPeerSharePort(share *repo.PeerShare, requestedPort int) (int, error) { if share == nil { return 0, fmt.Errorf("share not found") } @@ -1349,29 +1282,13 @@ func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) used := make(map[int]struct{}) - rows, err := h.repo.DB().Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL AND port > 0`, share.NodeID) + nodePorts, err := h.repo.ListUsedPortsOnNode(share.NodeID) if err != nil { return 0, err } - for rows.Next() { - var p sql.NullInt64 - if scanErr := rows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 { - used[int(p.Int64)] = struct{}{} - } + for _, p := range nodePorts { + used[p] = struct{}{} } - _ = rows.Close() - - rows, err = h.repo.DB().Query(`SELECT port FROM forward_port WHERE node_id = ? AND port > 0`, share.NodeID) - if err != nil { - return 0, err - } - for rows.Next() { - var p sql.NullInt64 - if scanErr := rows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 { - used[int(p.Int64)] = struct{}{} - } - } - _ = rows.Close() ports, err := h.repo.ListActivePeerShareRuntimePorts(share.ID, share.NodeID) if err != nil { @@ -1412,7 +1329,7 @@ func extractBearerToken(r *http.Request) string { return "" } -func isPeerShareFlowExceeded(share *sqlite.PeerShare) bool { +func isPeerShareFlowExceeded(share *repo.PeerShare) bool { if share == nil { return false } @@ -1655,19 +1572,10 @@ func (h *Handler) cleanupFederationTunnels(shareID int64) { return } namePrefix := fmt.Sprintf("Share-%d-Port-", shareID) - rows, err := h.repo.DB().Query(`SELECT id FROM tunnel WHERE name LIKE ?`, namePrefix+"%") - if err != nil { + tunnelIDs, err := h.repo.ListTunnelIDsByNamePrefix(namePrefix) + if err != nil || len(tunnelIDs) == 0 { return } - defer rows.Close() - - var tunnelIDs []int64 - for rows.Next() { - var id int64 - if err := rows.Scan(&id); err == nil { - tunnelIDs = append(tunnelIDs, id) - } - } for _, tid := range tunnelIDs { _ = h.deleteTunnelByID(tid) diff --git a/go-backend/internal/http/handler/federation_runtime_test.go b/go-backend/internal/http/handler/federation_runtime_test.go index 980f380..1047eae 100644 --- a/go-backend/internal/http/handler/federation_runtime_test.go +++ b/go-backend/internal/http/handler/federation_runtime_test.go @@ -10,33 +10,33 @@ import ( "time" "go-backend/internal/http/response" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) func TestPickPeerSharePortUsesRuntimeReservations(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open repo: %v", err) } - defer repo.Close() + defer r.Close() - h := &Handler{repo: repo} + h := &Handler{repo: r} now := time.Now().UnixMilli() - if _, err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, ?, ?, ?, ?, ?, ?)`, 1, 2, 1, 3000, "round", 1, "tls"); err != nil { + if err := r.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, ?, ?, ?, ?, ?, ?)`, 1, 2, 1, 3000, "round", 1, "tls").Error; err != nil { t.Fatalf("insert chain_tunnel: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, 1, 3001); err != nil { + if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, 1, 3001).Error; err != nil { t.Fatalf("insert forward_port: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, 77, 1, "res-1", "rk-1", "b-1", "exit", "", "fed_svc_1", "tls", "round", 3002, "", 1, 1, now, now); err != nil { + `, 77, 1, "res-1", "rk-1", "b-1", "exit", "", "fed_svc_1", "tls", "round", 3002, "", 1, 1, now, now).Error; err != nil { t.Fatalf("insert peer_share_runtime: %v", err) } - share := &sqlite.PeerShare{ + share := &repo.PeerShare{ ID: 77, NodeID: 1, PortRangeStart: 3000, @@ -57,13 +57,13 @@ func TestPickPeerSharePortUsesRuntimeReservations(t *testing.T) { } func TestApplyTunnelRuntimeSkipsRemoteChainAndOutNodes(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "rt-skip.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "rt-skip.db")) if err != nil { t.Fatalf("open repo: %v", err) } - defer repo.Close() + defer r.Close() - h := &Handler{repo: repo} + h := &Handler{repo: r} now := time.Now().UnixMilli() for _, n := range []struct { id int64 @@ -73,10 +73,10 @@ func TestApplyTunnelRuntimeSkipsRemoteChainAndOutNodes(t *testing.T) { {12, "remote-chain", "10.99.0.2"}, {13, "remote-out", "10.99.0.3"}, } { - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, n.id, n.name, n.name+"-secret", n.ip, n.ip, "", "40000-40010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://remote-peer", "remote-token"); err != nil { + `, n.id, n.name, n.name+"-secret", n.ip, n.ip, "", "40000-40010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://remote-peer", "remote-token").Error; err != nil { t.Fatalf("insert node %s: %v", n.name, err) } } @@ -112,25 +112,24 @@ func TestApplyTunnelRuntimeSkipsRemoteChainAndOutNodes(t *testing.T) { } func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open repo: %v", err) } - defer repo.Close() + defer r.Close() - h := &Handler{repo: repo} + h := &Handler{repo: r} now := time.Now().UnixMilli() insertNode := func(name string, status int, portRange string, isRemote int) int64 { - res, execErr := repo.DB().Exec(` + if execErr := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":1}`) - if execErr != nil { + `, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":1}`).Error; execErr != nil { t.Fatalf("insert node %s: %v", name, execErr) } - id, idErr := res.LastInsertId() - if idErr != nil { + var id int64 + if idErr := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); idErr != nil { t.Fatalf("node id %s: %v", name, idErr) } return id @@ -139,12 +138,12 @@ func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T) entryID := insertNode("entry", 1, "31000-31010", 0) remoteOutID := insertNode("remote-out", 1, "30000", 1) - if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, remoteOutID, 30000); err != nil { + if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, 1, remoteOutID, 30000).Error; err != nil { t.Fatalf("insert forward_port: %v", err) } - tx, err := repo.DB().Begin() - if err != nil { + tx := r.DB().Begin() + if tx.Error != nil { t.Fatalf("begin tx: %v", err) } defer tx.Rollback() @@ -173,25 +172,24 @@ func TestPrepareTunnelCreateStateRemoteAutoPortDefersToFederation(t *testing.T) } func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open repo: %v", err) } - defer repo.Close() + defer r.Close() - h := &Handler{repo: repo} + h := &Handler{repo: r} now := time.Now().UnixMilli() insertNode := func(name string, status int, portRange string, isRemote int) int64 { - res, execErr := repo.DB().Exec(` + if execErr := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":2}`) - if execErr != nil { + `, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0, isRemote, "http://peer", "peer-token", `{"shareId":2}`).Error; execErr != nil { t.Fatalf("insert node %s: %v", name, execErr) } - id, idErr := res.LastInsertId() - if idErr != nil { + var id int64 + if idErr := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); idErr != nil { t.Fatalf("node id %s: %v", name, idErr) } return id @@ -201,8 +199,8 @@ func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) { remoteMiddleID := insertNode("middle-remote", 0, "33000-33010", 1) outID := insertNode("out-local", 1, "34000-34010", 0) - tx, err := repo.DB().Begin() - if err != nil { + tx := r.DB().Begin() + if tx.Error != nil { t.Fatalf("begin tx: %v", err) } defer tx.Rollback() @@ -238,16 +236,16 @@ func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) { } func TestFederationRuntimeReservePortRejectsWhenShareFlowExceeded(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open repo: %v", err) } - defer repo.Close() + defer r.Close() - h := &Handler{repo: repo} + h := &Handler{repo: r} now := time.Now().UnixMilli() - if err := repo.CreatePeerShare(&sqlite.PeerShare{ + if err := r.CreatePeerShare(&repo.PeerShare{ Name: "limited-share", NodeID: 1, Token: "limited-token", diff --git a/go-backend/internal/http/handler/federation_share_test.go b/go-backend/internal/http/handler/federation_share_test.go index d5e8c95..27df2f5 100644 --- a/go-backend/internal/http/handler/federation_share_test.go +++ b/go-backend/internal/http/handler/federation_share_test.go @@ -12,28 +12,27 @@ import ( "time" "go-backend/internal/http/response" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) func TestFederationShareCreateRejectsRemoteNode(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() - insertRes, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "remote-share-node", "remote-share-secret", "10.10.10.1", "10.10.10.1", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":1}`) - if err != nil { + `, "remote-share-node", "remote-share-secret", "10.10.10.1", "10.10.10.1", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":1}`).Error; err != nil { t.Fatalf("insert remote node: %v", err) } - remoteNodeID, err := insertRes.LastInsertId() - if err != nil { + var remoteNodeID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&remoteNodeID); err != nil { t.Fatalf("get remote node id: %v", err) } @@ -71,7 +70,7 @@ func TestFederationShareCreateRejectsRemoteNode(t *testing.T) { } var shareCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, remoteNodeID).Scan(&shareCount); err != nil { + if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, remoteNodeID).Row().Scan(&shareCount); err != nil { t.Fatalf("query peer_share count: %v", err) } if shareCount != 0 { @@ -80,24 +79,23 @@ func TestFederationShareCreateRejectsRemoteNode(t *testing.T) { } func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() - insertRes, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "local-share-node", "local-share-secret", "10.20.30.40", "10.20.30.40", "", "21000-21010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "") - if err != nil { + `, "local-share-node", "local-share-secret", "10.20.30.40", "10.20.30.40", "", "21000-21010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "").Error; err != nil { t.Fatalf("insert local node: %v", err) } - localNodeID, err := insertRes.LastInsertId() - if err != nil { + var localNodeID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&localNodeID); err != nil { t.Fatalf("get local node id: %v", err) } @@ -136,7 +134,7 @@ func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) { } var shareCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, localNodeID).Scan(&shareCount); err != nil { + if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE node_id = ?`, localNodeID).Row().Scan(&shareCount); err != nil { t.Fatalf("query peer_share count: %v", err) } if shareCount != 0 { @@ -145,16 +143,16 @@ func TestFederationShareCreateRejectsInvalidAllowedIPs(t *testing.T) { } func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() - if err := repo.CreatePeerShare(&sqlite.PeerShare{ + if err := r.CreatePeerShare(&repo.PeerShare{ Name: "provider-share", NodeID: 9, Token: "share-list-token", @@ -169,12 +167,12 @@ func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) { t.Fatalf("create peer share: %v", err) } - share, err := repo.GetPeerShareByToken("share-list-token") + share, err := r.GetPeerShareByToken("share-list-token") if err != nil || share == nil { t.Fatalf("load peer share: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?), (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?), @@ -183,7 +181,7 @@ func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) { share.ID, share.NodeID, "r-1", "rk-1", "b-1", "middle", "fed_chain_1", "fed_svc_1", "tls", "round", 22001, "", 1, 1, now, now, share.ID, share.NodeID, "r-2", "rk-2", "b-2", "exit", "", "fed_svc_2", "tls", "round", 22002, "", 1, 1, now, now, share.ID, share.NodeID, "r-3", "rk-3", "", "", "", "", "tls", "round", 22003, "", 0, 0, now, now, - ); err != nil { + ).Error; err != nil { t.Fatalf("insert peer_share_runtime rows: %v", err) } @@ -238,16 +236,16 @@ func TestFederationShareListIncludesRemoteUsedPorts(t *testing.T) { } func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() - if err := repo.CreatePeerShare(&sqlite.PeerShare{ + if err := r.CreatePeerShare(&repo.PeerShare{ Name: "delete-cleanup-share", NodeID: 99, Token: "delete-cleanup-token", @@ -261,24 +259,24 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) { t.Fatalf("create peer share: %v", err) } - share, err := repo.GetPeerShareByToken("delete-cleanup-token") + share, err := r.GetPeerShareByToken("delete-cleanup-token") if err != nil || share == nil { t.Fatalf("load peer share: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?), (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) `, share.ID, 99, "dc-r1", "dc-rk1", "dc-b1", "exit", "", "fed_svc_dc1", "tls", "round", 40001, "", 1, 1, now, now, share.ID, 99, "dc-r2", "dc-rk2", "dc-b2", "middle", "fed_chain_dc2", "fed_svc_dc2", "tls", "round", 40002, "", 1, 1, now, now, - ); err != nil { + ).Error; err != nil { t.Fatalf("insert peer_share_runtime rows: %v", err) } var runtimeCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1`, share.ID).Scan(&runtimeCount); err != nil { + if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1`, share.ID).Row().Scan(&runtimeCount); err != nil { t.Fatalf("count active runtimes before: %v", err) } if runtimeCount != 2 { @@ -307,7 +305,7 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) { } var shareCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share WHERE id = ?`, share.ID).Scan(&shareCount); err != nil { + if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share WHERE id = ?`, share.ID).Row().Scan(&shareCount); err != nil { t.Fatalf("count peer_share after: %v", err) } if shareCount != 0 { @@ -315,7 +313,7 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) { } var runtimeCountAfter int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, share.ID).Scan(&runtimeCountAfter); err != nil { + if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, share.ID).Row().Scan(&runtimeCountAfter); err != nil { t.Fatalf("count peer_share_runtime after: %v", err) } if runtimeCountAfter != 0 { @@ -324,19 +322,19 @@ func TestFederationShareDeleteCleansUpRuntimes(t *testing.T) { } func TestFederationRemoteUsageListSyncErrorFallback(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "sync-error-node", "sync-error-secret", "10.50.60.70", "10.50.60.70", "", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://unreachable.invalid:9999", "bad-token", `{"shareId":42,"maxBandwidth":5368709120,"currentFlow":999999,"portRangeStart":32000,"portRangeEnd":32010}`); err != nil { + `, "sync-error-node", "sync-error-secret", "10.50.60.70", "10.50.60.70", "", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://unreachable.invalid:9999", "bad-token", `{"shareId":42,"maxBandwidth":5368709120,"currentFlow":999999,"portRangeStart":32000,"portRangeEnd":32010}`).Error; err != nil { t.Fatalf("insert remote node: %v", err) } @@ -380,15 +378,15 @@ func TestFederationRemoteUsageListSyncErrorFallback(t *testing.T) { } func TestFederationShareResetFlow(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() - if err := repo.CreatePeerShare(&sqlite.PeerShare{ + if err := r.CreatePeerShare(&repo.PeerShare{ Name: "reset-flow-share", NodeID: 11, Token: "reset-flow-token", @@ -402,7 +400,7 @@ func TestFederationShareResetFlow(t *testing.T) { }); err != nil { t.Fatalf("create peer share: %v", err) } - share, err := repo.GetPeerShareByToken("reset-flow-token") + share, err := r.GetPeerShareByToken("reset-flow-token") if err != nil || share == nil { t.Fatalf("load peer share: %v", err) } @@ -428,7 +426,7 @@ func TestFederationShareResetFlow(t *testing.T) { t.Fatalf("expected response code 0, got %d (%s)", payload.Code, payload.Msg) } - updated, err := repo.GetPeerShare(share.ID) + updated, err := r.GetPeerShare(share.ID) if err != nil || updated == nil { t.Fatalf("reload peer share: %v", err) } @@ -438,47 +436,50 @@ func TestFederationShareResetFlow(t *testing.T) { } func TestFederationRemoteUsageList(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() - resNode, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "remote-consumer-node", "remote-consumer-secret", "10.30.40.50", "10.30.40.50", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":88,"maxBandwidth":2147483648,"currentFlow":1073741824,"portRangeStart":31000,"portRangeEnd":31010}`) - if err != nil { + `, "remote-consumer-node", "remote-consumer-secret", "10.30.40.50", "10.30.40.50", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "http://peer.example", "peer-token", `{"shareId":88,"maxBandwidth":2147483648,"currentFlow":1073741824,"portRangeStart":31000,"portRangeEnd":31010}`).Error; err != nil { t.Fatalf("insert remote node: %v", err) } - nodeID, err := resNode.LastInsertId() - if err != nil { + var nodeID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&nodeID); err != nil { t.Fatalf("remote node id: %v", err) } - resTunnelA, err := repo.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-a", 2, "tls", 1, now, now, 1, "", 0) - if err != nil { + if err := r.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-a", 2, "tls", 1, now, now, 1, "", 0).Error; err != nil { t.Fatalf("insert tunnel a: %v", err) } - tunnelAID, _ := resTunnelA.LastInsertId() + var tunnelAID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelAID); err != nil { + t.Fatal(err) + } - resTunnelB, err := repo.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-b", 2, "tls", 1, now, now, 1, "", 0) - if err != nil { + if err := r.DB().Exec(`INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)`, "consumer-tunnel-b", 2, "tls", 1, now, now, 1, "", 0).Error; err != nil { t.Fatalf("insert tunnel b: %v", err) } - tunnelBID, _ := resTunnelB.LastInsertId() + var tunnelBID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelBID); err != nil { + t.Fatal(err) + } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?), (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) `, tunnelAID, nodeID, 2, 1, "http://peer.example", "rk-a", "rb-a", 31001, 1, now, now, tunnelBID, nodeID, 3, 0, "http://peer.example", "rk-b", "rb-b", 31002, 1, now, now, - ); err != nil { + ).Error; err != nil { t.Fatalf("insert federation bindings: %v", err) } @@ -532,13 +533,13 @@ func TestFederationRemoteUsageList(t *testing.T) { } func TestAuthPeerAllowedIPs(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "test-jwt-secret") + h := New(r, "test-jwt-secret") now := time.Now().UnixMilli() tests := []struct { @@ -585,7 +586,7 @@ func TestAuthPeerAllowedIPs(t *testing.T) { for idx, tt := range tests { t.Run(tt.name, func(t *testing.T) { token := fmt.Sprintf("share-token-%d", idx) - if err := repo.CreatePeerShare(&sqlite.PeerShare{ + if err := r.CreatePeerShare(&repo.PeerShare{ Name: "share-" + tt.name, NodeID: 1, Token: token, diff --git a/go-backend/internal/http/handler/flow_policy.go b/go-backend/internal/http/handler/flow_policy.go index e8ca759..3f8abee 100644 --- a/go-backend/internal/http/handler/flow_policy.go +++ b/go-backend/internal/http/handler/flow_policy.go @@ -1,7 +1,6 @@ package handler import ( - "database/sql" "encoding/json" "strconv" "strings" @@ -207,22 +206,18 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er if userTunnelID <= 0 { return nil, nil } - - row := h.repo.DB().QueryRow(` - SELECT id, user_id, tunnel_id, flow, in_flow, out_flow, exp_time, status - FROM user_tunnel - WHERE id = ? - LIMIT 1 - `, userTunnelID) - - var policy userTunnelPolicy - if err := row.Scan(&policy.ID, &policy.UserID, &policy.TunnelID, &policy.Flow, &policy.InFlow, &policy.OutFlow, &policy.ExpTime, &policy.Status); err != nil { - if err == sql.ErrNoRows { - return nil, nil - } + ut, err := h.repo.GetUserTunnelByID(userTunnelID) + if err != nil { return nil, err } - return &policy, nil + if ut == nil { + return nil, nil + } + return &userTunnelPolicy{ + ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID, + Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow, + ExpTime: ut.ExpTime, Status: ut.Status, + }, nil } func (h *Handler) pauseUserForwards(userID int64, now int64) { @@ -245,60 +240,20 @@ func (h *Handler) pauseForwardRecords(forwards []forwardRecord, now int64) { for i := range forwards { forward := forwards[i] _ = h.controlForwardServices(&forward, "PauseService", false) - _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, now, forward.ID) + _ = h.repo.UpdateForwardStatus(forward.ID, 0, now) } } func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error) { - rows, err := h.repo.DB().Query(` - SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status - FROM forward - WHERE user_id = ? AND status = 1 - ORDER BY id ASC - `, userID) - if err != nil { - return nil, err - } - defer rows.Close() - - return scanForwardRecords(rows) + return h.repo.ListActiveForwardsByUser(userID) } func (h *Handler) listActiveForwardsByUserTunnel(userID int64, tunnelID int64) ([]forwardRecord, error) { - rows, err := h.repo.DB().Query(` - SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status - FROM forward - WHERE user_id = ? AND tunnel_id = ? AND status = 1 - ORDER BY id ASC - `, userID, tunnelID) - if err != nil { - return nil, err - } - defer rows.Close() - - return scanForwardRecords(rows) -} - -func scanForwardRecords(rows *sql.Rows) ([]forwardRecord, error) { - out := make([]forwardRecord, 0) - for rows.Next() { - var record forwardRecord - if err := rows.Scan(&record.ID, &record.UserID, &record.UserName, &record.Name, &record.TunnelID, &record.RemoteAddr, &record.Strategy, &record.Status); err != nil { - return nil, err - } - if strings.TrimSpace(record.Strategy) == "" { - record.Strategy = "fifo" - } - out = append(out, record) - } - if err := rows.Err(); err != nil { - return nil, err - } - return out, nil + return h.repo.ListActiveForwardsByUserTunnel(userID, tunnelID) } func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) { - if h == nil || h.repo == nil || h.repo.DB() == nil || nodeID <= 0 { + if h == nil || h.repo == nil || nodeID <= 0 { return } if strings.TrimSpace(rawConfig) == "" { @@ -383,15 +338,13 @@ func (h *Handler) cleanOrphanedLimiters(nodeID int64, limiters []namedConfigItem } func (h *Handler) tunnelExists(tunnelID int64) bool { - var count int - err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE id = ?`, tunnelID).Scan(&count) - return err == nil && count > 0 + ok, _ := h.repo.TunnelExists(tunnelID) + return ok } func (h *Handler) forwardExists(forwardID int64) bool { - var count int - err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM forward WHERE id = ?`, forwardID).Scan(&count) - return err == nil && count > 0 + ok, _ := h.repo.ForwardExists(forwardID) + return ok } func (h *Handler) speedLimiterExists(name string) bool { @@ -402,8 +355,6 @@ func (h *Handler) speedLimiterExists(name string) bool { if err != nil || id <= 0 { return false } - - var count int - err = h.repo.DB().QueryRow(`SELECT COUNT(1) FROM speed_limit WHERE id = ?`, id).Scan(&count) - return err == nil && count > 0 + ok, _ := h.repo.SpeedLimitExists(id) + return ok } diff --git a/go-backend/internal/http/handler/flow_policy_federation_test.go b/go-backend/internal/http/handler/flow_policy_federation_test.go index d3814f5..b9ca83a 100644 --- a/go-backend/internal/http/handler/flow_policy_federation_test.go +++ b/go-backend/internal/http/handler/flow_policy_federation_test.go @@ -5,18 +5,18 @@ import ( "testing" "time" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) { - repo, err := sqlite.Open(filepath.Join(t.TempDir(), "panel.db")) + r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db")) if err != nil { t.Fatalf("open repo: %v", err) } - defer repo.Close() + defer r.Close() now := time.Now().UnixMilli() - if err := repo.CreatePeerShare(&sqlite.PeerShare{ + if err := r.CreatePeerShare(&repo.PeerShare{ Name: "flow-share", NodeID: 1, Token: "flow-share-token", @@ -30,22 +30,22 @@ func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) { }); err != nil { t.Fatalf("create peer share: %v", err) } - share, err := repo.GetPeerShareByToken("flow-share-token") + share, err := r.GetPeerShareByToken("flow-share-token") if err != nil || share == nil { t.Fatalf("load peer share: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO peer_share_runtime(id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, 17, share.ID, share.NodeID, "res-17", "rk-17", "17", "exit", "", "fed_svc_17", "tls", "round", 32001, "", 1, 1, now, now); err != nil { + `, 17, share.ID, share.NodeID, "res-17", "rk-17", "17", "exit", "", "fed_svc_17", "tls", "round", 32001, "", 1, 1, now, now).Error; err != nil { t.Fatalf("insert peer_share_runtime: %v", err) } - h := &Handler{repo: repo} + h := &Handler{repo: r} h.processFlowItem(flowItem{N: "fed_svc_17", U: 1200, D: 900}) - updatedShare, err := repo.GetPeerShare(share.ID) + updatedShare, err := r.GetPeerShare(share.ID) if err != nil || updatedShare == nil { t.Fatalf("reload share: %v", err) } @@ -53,7 +53,7 @@ func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) { t.Fatalf("expected current_flow=3100, got %d", updatedShare.CurrentFlow) } - runtime, err := repo.GetPeerShareRuntimeByID(17) + runtime, err := r.GetPeerShareRuntimeByID(17) if err != nil || runtime == nil { t.Fatalf("reload runtime: %v", err) } diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 37f314a..bec3b3b 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -18,12 +18,12 @@ import ( "go-backend/internal/http/middleware" "go-backend/internal/http/response" "go-backend/internal/security" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" "go-backend/internal/ws" ) type Handler struct { - repo *sqlite.Repository + repo *repo.Repository jwtSecret string wsServer *ws.Server @@ -69,7 +69,7 @@ type flowItem struct { D int64 `json:"d"` } -func New(repo *sqlite.Repository, jwtSecret string) *Handler { +func New(repo *repo.Repository, jwtSecret string) *Handler { return &Handler{ repo: repo, jwtSecret: jwtSecret, @@ -426,7 +426,7 @@ func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("请求失败")) return } - if h == nil || h.repo == nil || h.repo.DB() == nil { + if h == nil || h.repo == nil { response.WriteJSON(w, response.Err(-2, "database unavailable")) return } @@ -469,27 +469,21 @@ func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) { return } - var userID int64 - var inFlow int64 - var outFlow int64 - var flow int64 - var expTime int64 - err = h.repo.DB().QueryRow(`SELECT user_id, in_flow, out_flow, flow, exp_time FROM user_tunnel WHERE id = ? LIMIT 1`, tunnelID). - Scan(&userID, &inFlow, &outFlow, &flow, &expTime) + ut, err := h.repo.GetUserTunnelByID(tunnelID) if err != nil { - if err == sql.ErrNoRows { - response.WriteJSON(w, response.ErrDefault("隧道不存在")) - return - } response.WriteJSON(w, response.Err(-2, err.Error())) return } - if userID != user.ID { + if ut == nil { + response.WriteJSON(w, response.ErrDefault("隧道不存在")) + return + } + if ut.UserID != user.ID { response.WriteJSON(w, response.ErrDefault("隧道不存在")) return } - headerValue = buildSubscriptionHeader(outFlow, inFlow, flow*giga, expTime/1000) + headerValue = buildSubscriptionHeader(ut.OutFlow, ut.InFlow, ut.Flow*giga, ut.ExpTime/1000) } w.Header().Set("subscription-userinfo", headerValue) @@ -1190,7 +1184,7 @@ func (h *Handler) backupExport(w http.ResponseWriter, r *http.Request) { type backupImportRequest struct { Types []string `json:"types"` - sqlite.BackupData + repo.BackupData } func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) { diff --git a/go-backend/internal/http/handler/jobs.go b/go-backend/internal/http/handler/jobs.go index f665918..d7c697e 100644 --- a/go-backend/internal/http/handler/jobs.go +++ b/go-backend/internal/http/handler/jobs.go @@ -2,12 +2,11 @@ package handler import ( "context" - "database/sql" "time" ) func (h *Handler) StartBackgroundJobs() { - if h == nil || h.repo == nil || h.repo.DB() == nil { + if h == nil || h.repo == nil { return } @@ -97,47 +96,28 @@ func durationUntilNextDailyMaintenance(now time.Time) time.Duration { } func (h *Handler) runStatisticsFlowJob(now time.Time) { - if h == nil || h.repo == nil || h.repo.DB() == nil { + if h == nil || h.repo == nil { return } - db := h.repo.DB() nowMs := now.UnixMilli() cutoffMs := nowMs - int64((48*time.Hour)/time.Millisecond) - _, _ = db.Exec(`DELETE FROM statistics_flow WHERE created_time < ?`, cutoffMs) + _ = h.repo.PurgeOldStatisticsFlows(cutoffMs) hourMark := now.Truncate(time.Hour) hourText := hourMark.Format("15:04") createdTime := hourMark.UnixMilli() - rows, err := db.Query(`SELECT id, in_flow, out_flow FROM user ORDER BY id ASC`) + users, err := h.repo.ListAllUserFlowSnapshots() if err != nil { return } - type userFlowSnapshot struct { - userID int64 - inFlow int64 - outFlow int64 - } - users := make([]userFlowSnapshot, 0) - - for rows.Next() { - var userID int64 - var inFlow int64 - var outFlow int64 - if err := rows.Scan(&userID, &inFlow, &outFlow); err != nil { - continue - } - users = append(users, userFlowSnapshot{userID: userID, inFlow: inFlow, outFlow: outFlow}) - } - _ = rows.Close() for _, user := range users { - currentTotal := user.inFlow + user.outFlow + currentTotal := user.InFlow + user.OutFlow increment := currentTotal - var lastTotal sql.NullInt64 - err := db.QueryRow(`SELECT total_flow FROM statistics_flow WHERE user_id = ? ORDER BY id DESC LIMIT 1`, user.userID).Scan(&lastTotal) + lastTotal, err := h.repo.GetLastStatisticsFlowTotal(user.UserID) if err == nil && lastTotal.Valid { increment = currentTotal - lastTotal.Int64 if increment < 0 { @@ -145,15 +125,12 @@ func (h *Handler) runStatisticsFlowJob(now time.Time) { } } - _, _ = db.Exec(` - INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) - VALUES(?, ?, ?, ?, ?) - `, user.userID, increment, currentTotal, hourText, createdTime) + _ = h.repo.CreateStatisticsFlow(user.UserID, increment, currentTotal, hourText, createdTime) } } func (h *Handler) runResetAndExpiryJob(now time.Time) { - if h == nil || h.repo == nil || h.repo.DB() == nil { + if h == nil || h.repo == nil { return } @@ -163,108 +140,39 @@ func (h *Handler) runResetAndExpiryJob(now time.Time) { } func (h *Handler) resetMonthlyFlow(now time.Time) { - db := h.repo.DB() currentDay := now.Day() lastDay := time.Date(now.Year(), now.Month()+1, 0, 0, 0, 0, 0, now.Location()).Day() - if currentDay == lastDay { - _, _ = db.Exec(` - UPDATE user - SET in_flow = 0, out_flow = 0 - WHERE flow_reset_time != 0 - AND (flow_reset_time = ? OR flow_reset_time > ?) - `, currentDay, lastDay) - _, _ = db.Exec(` - UPDATE user_tunnel - SET in_flow = 0, out_flow = 0 - WHERE flow_reset_time != 0 - AND (flow_reset_time = ? OR flow_reset_time > ?) - `, currentDay, lastDay) - return - } - - _, _ = db.Exec(` - UPDATE user - SET in_flow = 0, out_flow = 0 - WHERE flow_reset_time != 0 - AND flow_reset_time = ? - `, currentDay) - _, _ = db.Exec(` - UPDATE user_tunnel - SET in_flow = 0, out_flow = 0 - WHERE flow_reset_time != 0 - AND flow_reset_time = ? - `, currentDay) + _ = h.repo.ResetUserMonthlyFlow(currentDay, lastDay) + _ = h.repo.ResetUserTunnelMonthlyFlow(currentDay, lastDay) } func (h *Handler) disableExpiredUsers(nowMs int64) { - db := h.repo.DB() - rows, err := db.Query(` - SELECT id - FROM user - WHERE role_id != 0 - AND status = 1 - AND exp_time IS NOT NULL - AND exp_time < ? - `, nowMs) + userIDs, err := h.repo.ListExpiredActiveUserIDs(nowMs) if err != nil { return } - userIDs := make([]int64, 0) - - for rows.Next() { - var userID int64 - if err := rows.Scan(&userID); err != nil { - continue - } - userIDs = append(userIDs, userID) - } - _ = rows.Close() for _, userID := range userIDs { forwards, err := h.listActiveForwardsByUser(userID) if err == nil { h.pauseForwardRecords(forwards, nowMs) } - _, _ = db.Exec(`UPDATE user SET status = 0 WHERE id = ?`, userID) + _ = h.repo.DisableUser(userID) } } func (h *Handler) disableExpiredUserTunnels(nowMs int64) { - db := h.repo.DB() - rows, err := db.Query(` - SELECT id, user_id, tunnel_id - FROM user_tunnel - WHERE status = 1 - AND exp_time IS NOT NULL - AND exp_time < ? - `, nowMs) + items, err := h.repo.ListExpiredActiveUserTunnels(nowMs) if err != nil { return } - type expiredUserTunnel struct { - userTunnelID int64 - userID int64 - tunnelID int64 - } - items := make([]expiredUserTunnel, 0) - - for rows.Next() { - var userTunnelID int64 - var userID int64 - var tunnelID int64 - if err := rows.Scan(&userTunnelID, &userID, &tunnelID); err != nil { - continue - } - items = append(items, expiredUserTunnel{userTunnelID: userTunnelID, userID: userID, tunnelID: tunnelID}) - } - _ = rows.Close() for _, item := range items { - forwards, err := h.listActiveForwardsByUserTunnel(item.userID, item.tunnelID) + forwards, err := h.listActiveForwardsByUserTunnel(item.UserID, item.TunnelID) if err == nil { h.pauseForwardRecords(forwards, nowMs) } - _, _ = db.Exec(`UPDATE user_tunnel SET status = 0 WHERE id = ?`, item.userTunnelID) + _ = h.repo.DisableUserTunnel(item.ID) } } diff --git a/go-backend/internal/http/handler/jobs_test.go b/go-backend/internal/http/handler/jobs_test.go index d349f8c..4f68205 100644 --- a/go-backend/internal/http/handler/jobs_test.go +++ b/go-backend/internal/http/handler/jobs_test.go @@ -5,36 +5,36 @@ import ( "testing" "time" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) { dbPath := filepath.Join(t.TempDir(), "jobs-stats.db") - repo, err := sqlite.Open(dbPath) + r, err := repo.Open(dbPath) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "secret") + h := New(r, "secret") now := time.Date(2026, 2, 7, 12, 0, 0, 0, time.UTC) nowMs := now.UnixMilli() - if _, err := repo.DB().Exec(`UPDATE user SET in_flow = 100, out_flow = 200 WHERE id = 1`); err != nil { + if err := r.DB().Exec(`UPDATE user SET in_flow = 100, out_flow = 200 WHERE id = 1`).Error; err != nil { t.Fatalf("seed user flow: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 250, 250, '11:00', ?)`, now.Add(-time.Hour).UnixMilli()); err != nil { + if err := r.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 250, 250, '11:00', ?)`, now.Add(-time.Hour).UnixMilli()).Error; err != nil { t.Fatalf("seed recent statistics row: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 10, 10, '00:00', ?)`, now.Add(-49*time.Hour).UnixMilli()); err != nil { + if err := r.DB().Exec(`INSERT INTO statistics_flow(user_id, flow, total_flow, time, created_time) VALUES(1, 10, 10, '00:00', ?)`, now.Add(-49*time.Hour).UnixMilli()).Error; err != nil { t.Fatalf("seed stale statistics row: %v", err) } h.runStatisticsFlowJob(now) var staleCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM statistics_flow WHERE created_time < ?`, nowMs-int64((48*time.Hour)/time.Millisecond)).Scan(&staleCount); err != nil { + if err := r.DB().Raw(`SELECT COUNT(1) FROM statistics_flow WHERE created_time < ?`, nowMs-int64((48*time.Hour)/time.Millisecond)).Row().Scan(&staleCount); err != nil { t.Fatalf("query stale statistics rows: %v", err) } if staleCount != 0 { @@ -44,7 +44,7 @@ func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) { var flow int64 var total int64 var hour string - if err := repo.DB().QueryRow(`SELECT flow, total_flow, time FROM statistics_flow WHERE user_id = 1 ORDER BY id DESC LIMIT 1`).Scan(&flow, &total, &hour); err != nil { + if err := r.DB().Raw(`SELECT flow, total_flow, time FROM statistics_flow WHERE user_id = 1 ORDER BY id DESC LIMIT 1`).Row().Scan(&flow, &total, &hour); err != nil { t.Fatalf("query latest statistics row: %v", err) } if flow != 50 { @@ -60,41 +60,41 @@ func TestRunStatisticsFlowJobTracksIncrementAndPrunes(t *testing.T) { func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) { dbPath := filepath.Join(t.TempDir(), "jobs-reset.db") - repo, err := sqlite.Open(dbPath) + r, err := repo.Open(dbPath) if err != nil { t.Fatalf("open sqlite: %v", err) } - t.Cleanup(func() { _ = repo.Close() }) + t.Cleanup(func() { _ = r.Close() }) - h := New(repo, "secret") + h := New(r, "secret") now := time.Date(2026, 3, 15, 0, 0, 5, 0, time.UTC) nowMs := now.UnixMilli() - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` 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(2, 'expired_user', 'x', 1, ?, 100, 1000, 2000, 15, 1, ?, ?, 1) - `, nowMs-1000, nowMs, nowMs); err != nil { + `, nowMs-1000, nowMs, nowMs).Error; err != nil { t.Fatalf("insert expired user: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 't1', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0) - `, nowMs, nowMs); err != nil { + `, nowMs, nowMs).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 1, 1, 300, 400, 15, ?, 1) - `, nowMs-1000); err != nil { + `, nowMs-1000).Error; err != nil { t.Fatalf("insert expired user_tunnel: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(20, 2, 'expired_user', 'f1', 1, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0) - `, nowMs, nowMs); err != nil { + `, nowMs, nowMs).Error; err != nil { t.Fatalf("insert forward: %v", err) } @@ -102,7 +102,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) { var userIn, userOut int64 var userStatus int - if err := repo.DB().QueryRow(`SELECT in_flow, out_flow, status FROM user WHERE id = 2`).Scan(&userIn, &userOut, &userStatus); err != nil { + if err := r.DB().Raw(`SELECT in_flow, out_flow, status FROM user WHERE id = 2`).Row().Scan(&userIn, &userOut, &userStatus); err != nil { t.Fatalf("query user after maintenance: %v", err) } if userIn != 0 || userOut != 0 || userStatus != 0 { @@ -111,7 +111,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) { var utIn, utOut int64 var utStatus int - if err := repo.DB().QueryRow(`SELECT in_flow, out_flow, status FROM user_tunnel WHERE id = 10`).Scan(&utIn, &utOut, &utStatus); err != nil { + if err := r.DB().Raw(`SELECT in_flow, out_flow, status FROM user_tunnel WHERE id = 10`).Row().Scan(&utIn, &utOut, &utStatus); err != nil { t.Fatalf("query user_tunnel after maintenance: %v", err) } if utIn != 0 || utOut != 0 || utStatus != 0 { @@ -119,7 +119,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) { } var forwardStatus int - if err := repo.DB().QueryRow(`SELECT status FROM forward WHERE id = 20`).Scan(&forwardStatus); err != nil { + if err := r.DB().Raw(`SELECT status FROM forward WHERE id = 20`).Row().Scan(&forwardStatus); err != nil { t.Fatalf("query forward after maintenance: %v", err) } if forwardStatus != 0 { diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index a61d2d4..cd25093 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -18,8 +18,10 @@ import ( "go-backend/internal/http/client" "go-backend/internal/http/response" "go-backend/internal/security" - "go-backend/internal/store" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/model" + "go-backend/internal/store/repo" + + "gorm.io/gorm" ) func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) { @@ -40,18 +42,12 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) { return } - db := h.repo.DB() - if db == nil { - response.WriteJSON(w, response.Err(-2, "database unavailable")) - return - } - - var cnt int - if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&cnt); err != nil { + exists, err := h.repo.UserExists(username) + if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if cnt > 0 { + if exists { response.WriteJSON(w, response.ErrDefault("用户名已存在")) return } @@ -64,11 +60,7 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) { roleID := 1 now := time.Now().UnixMilli() - _, err := db.Exec(` - INSERT INTO user(user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) - VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?, ?, ?) - `, username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, now, now, status) - if err != nil { + if err := h.repo.CreateUser(username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, status, now); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -96,14 +88,8 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) { return } - db := h.repo.DB() - if db == nil { - response.WriteJSON(w, response.Err(-2, "database unavailable")) - return - } - - var roleID int - if err := db.QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil { + roleID, err := h.repo.GetUserRoleID(id) + if err != nil { if err == sql.ErrNoRows { response.WriteJSON(w, response.ErrDefault("用户不存在")) return @@ -116,12 +102,12 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) { return } - var cnt int - if err := db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, id).Scan(&cnt); err != nil { + dup, err := h.repo.UserExistsExcluding(username, id) + if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if cnt > 0 { + if dup { response.WriteJSON(w, response.ErrDefault("用户名已存在")) return } @@ -135,28 +121,18 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) { pwd := asString(req["pwd"]) if strings.TrimSpace(pwd) == "" { - _, err := db.Exec(` - UPDATE user - SET user = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ? - WHERE id = ? - `, username, flow, num, expTime, flowResetTime, status, now, id) - if err != nil { + if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, now); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } } else { - _, err := db.Exec(` - UPDATE user - SET user = ?, pwd = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ?, updated_time = ? - WHERE id = ? - `, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, now, id) - if err != nil { + if err := h.repo.UpdateUserWithPassword(id, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, now); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } } - _, _ = db.Exec(`UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ? WHERE user_id = ?`, flow, num, expTime, flowResetTime, id) + h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime) response.WriteJSON(w, response.OKEmpty()) } @@ -170,8 +146,8 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) { return } - var roleID int - if err := h.repo.DB().QueryRow(`SELECT role_id FROM user WHERE id = ?`, id).Scan(&roleID); err != nil { + roleID, err := h.repo.GetUserRoleID(id) + if err != nil { if err == sql.ErrNoRows { response.WriteJSON(w, response.ErrDefault("用户不存在")) return @@ -184,44 +160,7 @@ func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) { return } - db := h.repo.DB() - tx, err := db.Begin() - if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - defer func() { _ = tx.Rollback() }() - - if _, err = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE user_id = ?)`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - if _, err = tx.Exec(`DELETE FROM forward WHERE user_id = ?`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - if _, err = tx.Exec(`DELETE FROM group_permission_grant WHERE user_tunnel_id IN (SELECT id FROM user_tunnel WHERE user_id = ?)`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - if _, err = tx.Exec(`DELETE FROM user_tunnel WHERE user_id = ?`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - if _, err = tx.Exec(`DELETE FROM user_group_user WHERE user_id = ?`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - if _, err = tx.Exec(`DELETE FROM statistics_flow WHERE user_id = ?`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - if _, err = tx.Exec(`DELETE FROM user WHERE id = ?`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - - if err = tx.Commit(); err != nil { + if err := h.repo.DeleteUserCascade(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -245,12 +184,10 @@ func (h *Handler) userResetFlow(w http.ResponseWriter, r *http.Request) { return } - db := h.repo.DB() if typeVal == 1 { - _, _ = db.Exec(`UPDATE user SET in_flow = 0, out_flow = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id) - _, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE user_id = ?`, id) + h.repo.ResetUserFlowByUser(id, time.Now().UnixMilli()) } else { - _, _ = db.Exec(`UPDATE user_tunnel SET in_flow = 0, out_flow = 0 WHERE id = ?`, id) + h.repo.ResetUserFlowByUserTunnel(id) } response.WriteJSON(w, response.OKEmpty()) } @@ -272,13 +209,9 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) { return } - db := h.repo.DB() now := time.Now().UnixMilli() - inx := nextIndex(db, "node") - _, err := db.Exec(` - INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, + inx := h.repo.NextIndex("node") + if err := h.repo.CreateNode( name, randomToken(16), serverIP, @@ -291,7 +224,6 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) { asInt(req["tls"], 0), asInt(req["socks"], 0), now, - now, 0, defaultString(asString(req["tcpListenAddr"]), "[::]"), defaultString(asString(req["udpListenAddr"]), "[::]"), @@ -300,8 +232,7 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) { nullableText(asString(req["remoteUrl"])), nullableText(asString(req["remoteToken"])), nullableText(asString(req["remoteConfig"])), - ) - if err != nil { + ); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -324,11 +255,8 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) { return } - var currentStatus int - var currentHTTP int - var currentTLS int - var currentSocks int - if err := h.repo.DB().QueryRow(`SELECT status, http, tls, socks FROM node WHERE id = ?`, id).Scan(¤tStatus, ¤tHTTP, ¤tTLS, ¤tSocks); err != nil { + currentStatus, currentHTTP, currentTLS, currentSocks, err := h.repo.GetNodeStatusFields(id) + if err != nil { if err == sql.ErrNoRows { response.WriteJSON(w, response.ErrDefault("节点不存在")) return @@ -348,11 +276,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) { } now := time.Now().UnixMilli() - _, err := h.repo.DB().Exec(` - UPDATE node - SET name = ?, server_ip = ?, server_ip_v4 = ?, server_ip_v6 = ?, port = ?, interface_name = ?, http = ?, tls = ?, socks = ?, tcp_listen_addr = ?, udp_listen_addr = ?, updated_time = ? - WHERE id = ? - `, + if err := h.repo.UpdateNode(id, asString(req["name"]), asString(req["serverIp"]), nullableText(asString(req["serverIpV4"])), @@ -365,9 +289,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) { defaultString(asString(req["tcpListenAddr"]), "[::]"), defaultString(asString(req["udpListenAddr"]), "[::]"), now, - id, - ) - if err != nil { + ); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -399,14 +321,13 @@ func (h *Handler) nodeInstall(w http.ResponseWriter, r *http.Request) { if id <= 0 { return } - db := h.repo.DB() - var secret string - if err := db.QueryRow(`SELECT secret FROM node WHERE id = ?`, id).Scan(&secret); err != nil { + secret, err := h.repo.GetNodeSecret(id) + if err != nil { response.WriteJSON(w, response.ErrDefault("节点不存在")) return } - var panelAddr string - if err := db.QueryRow(`SELECT value FROM vite_config WHERE name = 'ip' LIMIT 1`).Scan(&panelAddr); err != nil { + panelAddr, err := h.repo.GetViteConfigValue("ip") + if err != nil { if err == sql.ErrNoRows { response.WriteJSON(w, response.ErrDefault("请先前往网站配置中设置ip")) return @@ -434,7 +355,7 @@ func (h *Handler) nodeUpdateOrder(w http.ResponseWriter, r *http.Request) { return } for _, n := range req.Nodes { - _, _ = h.repo.DB().Exec(`UPDATE node SET inx = ?, updated_time = ? WHERE id = ?`, n.Inx, time.Now().UnixMilli(), n.ID) + h.repo.UpdateNodeOrder(n.ID, n.Inx, time.Now().UnixMilli()) } response.WriteJSON(w, response.OKEmpty()) } @@ -482,12 +403,12 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("隧道名称不能为空")) return } - var tunnelNameDup int - if err := h.repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, name).Scan(&tunnelNameDup); err != nil { + nameDup, err := h.repo.TunnelNameExists(name) + if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if tunnelNameDup > 0 { + if nameDup { response.WriteJSON(w, response.ErrDefault("隧道名称重复")) return } @@ -499,14 +420,15 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { inIP := asString(req["inIp"]) ipPreference := asString(req["ipPreference"]) now := time.Now().UnixMilli() - inx := nextIndex(h.repo.DB(), "tunnel") + inx := h.repo.NextIndex("tunnel") + localDomain := h.federationLocalDomain() - tx, err := h.repo.DB().Begin() - if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) + tx := h.repo.BeginTx() + if tx.Error != nil { + response.WriteJSON(w, response.Err(-2, tx.Error.Error())) return } - defer func() { _ = tx.Rollback() }() + defer func() { tx.Rollback() }() runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal, 0) if err != nil { @@ -520,9 +442,8 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { if len(runtimeState.InNodes) > 0 { firstNodeID := runtimeState.InNodes[0].NodeID - var isRemote int - var rUrl, rToken sql.NullString - if err := h.repo.DB().QueryRow("SELECT is_remote, remote_url, remote_token FROM node WHERE id = ?", firstNodeID).Scan(&isRemote, &rUrl, &rToken); err == nil && isRemote == 1 { + isRemote, rUrl, rToken, _ := h.repo.GetNodeRemoteFieldsTx(tx, firstNodeID) + if isRemote == 1 { fc := client.NewFederationClient() targetProto := "tcp" @@ -567,32 +488,48 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { } } - tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, ip_preference) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx, ipPreference) - if err != nil { + var tunnelInIP sql.NullString + if trimmed := strings.TrimSpace(inIP); trimmed != "" { + tunnelInIP = sql.NullString{String: trimmed, Valid: true} + } + tunnel := model.Tunnel{ + Name: name, + TrafficRatio: trafficRatio, + Type: typeVal, + Protocol: "tls", + Flow: flow, + CreatedTime: now, + UpdatedTime: now, + Status: status, + InIP: tunnelInIP, + Inx: inx, + IPPreference: ipPreference, + } + if err := tx.Create(&tunnel).Error; err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } + tunnelID := tunnel.ID runtimeState.TunnelID = tunnelID - var federationBindings []sqlite.FederationTunnelBinding + var federationBindings []repo.FederationTunnelBinding var federationReleaseRefs []federationRuntimeReleaseRef - federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState) + federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain) if err != nil { response.WriteJSON(w, response.ErrDefault(err.Error())) return } applyTunnelPortsToRequest(req, runtimeState) - if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil { + if err := h.replaceTunnelChainsTx(tx, tunnelID, req); err != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := replaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); err != nil { + if err := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); err != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := tx.Commit(); err != nil { + if err := tx.Commit().Error; err != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -680,13 +617,14 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { now := time.Now().UnixMilli() typeVal := asInt(req["type"], 1) ipPreference := asString(req["ipPreference"]) + localDomain := h.federationLocalDomain() - tx, err := h.repo.DB().Begin() - if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) + tx := h.repo.BeginTx() + if tx.Error != nil { + response.WriteJSON(w, response.Err(-2, tx.Error.Error())) return } - defer func() { _ = tx.Rollback() }() + defer func() { tx.Rollback() }() runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal, id) if err != nil { @@ -698,37 +636,46 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) { inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes) - var federationBindings []sqlite.FederationTunnelBinding + var federationBindings []repo.FederationTunnelBinding var federationReleaseRefs []federationRuntimeReleaseRef - federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState) + federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState, localDomain) if err != nil { response.WriteJSON(w, response.ErrDefault(err.Error())) return } applyTunnelPortsToRequest(req, runtimeState) - _, err = tx.Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, ip_preference=?, updated_time=? WHERE id=?`, - asString(req["name"]), typeVal, asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(inIp), ipPreference, now, id) - if err != nil { + if err := h.repo.UpdateTunnelTx( + tx, + id, + asString(req["name"]), + typeVal, + asInt64(req["flow"], 1), + asFloat(req["trafficRatio"], 1.0), + asInt(req["status"], 1), + inIp, + ipPreference, + now, + ); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if _, err := tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id); err != nil { + if err := h.repo.DeleteChainTunnelsByTunnelTx(tx, id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := replaceTunnelChainsTx(tx, id, req); err != nil { + if err := h.replaceTunnelChainsTx(tx, id, req); err != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := replaceFederationTunnelBindingsTx(tx, id, federationBindings); err != nil { + if err := h.repo.ReplaceFederationTunnelBindingsTx(tx, id, federationBindings); err != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := tx.Commit(); err != nil { + if err := tx.Commit().Error; err != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -807,7 +754,7 @@ func (h *Handler) tunnelUpdateOrder(w http.ResponseWriter, r *http.Request) { return } for _, t := range req.Tunnels { - _, _ = h.repo.DB().Exec(`UPDATE tunnel SET inx = ?, updated_time = ? WHERE id = ?`, t.Inx, time.Now().UnixMilli(), t.ID) + h.repo.UpdateTunnelOrder(t.ID, t.Inx, time.Now().UnixMilli()) } response.WriteJSON(w, response.OKEmpty()) } @@ -842,8 +789,7 @@ func (h *Handler) reconstructTunnelState(tunnelID int64) (*tunnelCreateState, er return nil, err } - var ipPreference string - _ = h.repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE id = ?`, tunnelID).Scan(&ipPreference) + ipPreference := h.repo.GetTunnelIPPreference(tunnelID) state := &tunnelCreateState{ TunnelID: tunnelID, @@ -933,24 +879,24 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) { fail++ continue } - federationBindings, federationReleaseRefs, fedErr := h.applyFederationRuntime(state) + federationBindings, federationReleaseRefs, fedErr := h.applyFederationRuntime(state, h.federationLocalDomain()) if fedErr != nil { fail++ continue } - tx, txErr := h.repo.DB().Begin() - if txErr != nil { + tx := h.repo.BeginTx() + if tx.Error != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) fail++ continue } - if replaceErr := replaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); replaceErr != nil { - _ = tx.Rollback() + if replaceErr := h.repo.ReplaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); replaceErr != nil { + tx.Rollback() h.releaseFederationRuntimeRefs(federationReleaseRefs) fail++ continue } - if commitErr := tx.Commit(); commitErr != nil { + if commitErr := tx.Commit().Error; commitErr != nil { h.releaseFederationRuntimeRefs(federationReleaseRefs) fail++ continue @@ -1032,8 +978,7 @@ func (h *Handler) userTunnelRemove(w http.ResponseWriter, r *http.Request) { if id <= 0 { return } - _, err := h.repo.DB().Exec(`DELETE FROM user_tunnel WHERE id = ?`, id) - if err != nil { + if err := h.repo.DeleteUserTunnel(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1051,25 +996,20 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("权限ID不能为空")) return } - _, err := h.repo.DB().Exec(` - UPDATE user_tunnel SET flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, speed_id = ?, status = ? WHERE id = ? - `, + if err := h.repo.UpdateUserTunnel(id, asInt64(req["flow"], 0), asInt(req["num"], 0), asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()), asInt64(req["flowResetTime"], 1), nullableInt(asAnyToInt64Ptr(req["speedId"])), asInt(req["status"], 1), - id, - ) - if err != nil { + ); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - // Fetch details to sync forwards - var userID, tunnelID int64 - if err := h.repo.DB().QueryRow("SELECT user_id, tunnel_id FROM user_tunnel WHERE id = ?", id).Scan(&userID, &tunnelID); err == nil { + userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id) + if utErr == nil { h.syncUserTunnelForwards(userID, tunnelID) } @@ -1130,33 +1070,16 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) { } } now := time.Now().UnixMilli() - inx := nextIndex(h.repo.DB(), "forward") - var userName string - _ = h.repo.DB().QueryRow(`SELECT user FROM user WHERE id = ?`, userID).Scan(&userName) + inx := h.repo.NextIndex("forward") + userName := h.repo.GetUsernameByID(userID) if userName == "" { userName = "user" } - tx, err := h.repo.DB().Begin() + forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - defer func() { _ = tx.Rollback() }() - 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) - if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - for _, nodeID := range entryNodes { - _, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port) - } - if err := tx.Commit(); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } createdForward, err := h.getForwardRecord(forwardID) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) @@ -1230,8 +1153,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { port := asInt(req["inPort"], 0) if port <= 0 { - var minPort sql.NullInt64 - _ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&minPort) + minPort := h.repo.GetMinForwardPort(id) if minPort.Valid { port = int(minPort.Int64) } @@ -1251,10 +1173,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { } } now := time.Now().UnixMilli() - _, err = h.repo.DB().Exec(` - UPDATE forward SET name = ?, tunnel_id = ?, remote_addr = ?, strategy = ?, updated_time = ? WHERE id = ? - `, name, tunnelID, remoteAddr, strategy, now, id) - if err != nil { + if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1324,7 +1243,7 @@ func (h *Handler) forwardPause(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault(err.Error())) return } - _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id) + _ = h.repo.UpdateForwardStatus(id, 0, time.Now().UnixMilli()) response.WriteJSON(w, response.OKEmpty()) } @@ -1346,7 +1265,7 @@ func (h *Handler) forwardResume(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault(err.Error())) return } - _, _ = h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id) + _ = h.repo.UpdateForwardStatus(id, 1, time.Now().UnixMilli()) response.WriteJSON(w, response.OKEmpty()) } @@ -1388,7 +1307,7 @@ func (h *Handler) forwardUpdateOrder(w http.ResponseWriter, r *http.Request) { return } for _, f := range req.Forwards { - _, _ = h.repo.DB().Exec(`UPDATE forward SET inx = ?, updated_time = ? WHERE id = ?`, f.Inx, time.Now().UnixMilli(), f.ID) + h.repo.UpdateForwardOrder(f.ID, f.Inx, time.Now().UnixMilli()) } response.WriteJSON(w, response.OKEmpty()) } @@ -1446,7 +1365,7 @@ func (h *Handler) forwardBatchPause(w http.ResponseWriter, r *http.Request) { f++ continue } - if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 0, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil { + if err := h.repo.UpdateForwardStatus(id, 0, time.Now().UnixMilli()); err != nil { f++ } else { s++ @@ -1477,7 +1396,7 @@ func (h *Handler) forwardBatchResume(w http.ResponseWriter, r *http.Request) { f++ continue } - if _, err := h.repo.DB().Exec(`UPDATE forward SET status = 1, updated_time = ? WHERE id = ?`, time.Now().UnixMilli(), id); err != nil { + if err := h.repo.UpdateForwardStatus(id, 1, time.Now().UnixMilli()); err != nil { f++ } else { s++ @@ -1560,10 +1479,8 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques fail++ continue } - var port sql.NullInt64 - _ = h.repo.DB().QueryRow(`SELECT MIN(port) FROM forward_port WHERE forward_id = ?`, id).Scan(&port) - _, err := h.repo.DB().Exec(`UPDATE forward SET tunnel_id = ?, updated_time = ? WHERE id = ?`, req.TargetTunnelID, time.Now().UnixMilli(), id) - if err != nil { + port := h.repo.GetMinForwardPort(id) + if err := h.repo.UpdateForwardTunnel(id, req.TargetTunnelID, time.Now().UnixMilli()); err != nil { fail++ continue } @@ -1627,16 +1544,14 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("名称不能为空")) return } - var tunnelName string - _ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName) + tunnelName := h.repo.GetTunnelNameByID(tunnelID) if tunnelName == "" { response.WriteJSON(w, response.ErrDefault("隧道不存在")) return } now := time.Now().UnixMilli() speed := asInt(req["speed"], 100) - 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)) + id, err := h.repo.CreateSpeedLimit(name, speed, tunnelID, tunnelName, now, asInt(req["status"], 1)) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -1657,16 +1572,13 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("请求参数错误")) return } - var tunnelName string - _ = h.repo.DB().QueryRow(`SELECT name FROM tunnel WHERE id = ?`, tunnelID).Scan(&tunnelName) + tunnelName := h.repo.GetTunnelNameByID(tunnelID) if tunnelName == "" { response.WriteJSON(w, response.ErrDefault("隧道不存在")) return } speed := asInt(req["speed"], 100) - _, err := h.repo.DB().Exec(`UPDATE speed_limit SET name=?, speed=?, tunnel_id=?, tunnel_name=?, status=?, updated_time=? WHERE id=?`, - asString(req["name"]), speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli(), id) - if err != nil { + if err := h.repo.UpdateSpeedLimit(id, asString(req["name"]), speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1679,11 +1591,9 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) { if id <= 0 { return } - var tunnelID int64 - _ = h.repo.DB().QueryRow(`SELECT tunnel_id FROM speed_limit WHERE id = ?`, id).Scan(&tunnelID) + tunnelID := h.repo.GetSpeedLimitTunnelID(id) - _, err := h.repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, id) - if err != nil { + if err := h.repo.DeleteSpeedLimit(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1726,17 +1636,17 @@ func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("请求参数错误")) return } - tx, err := h.repo.DB().Begin() - if err != nil { + tx := h.repo.BeginTx() + if tx.Error != nil { + response.WriteJSON(w, response.Err(-2, tx.Error.Error())) + return + } + defer func() { tx.Rollback() }() + if err := h.repo.ReplaceTunnelGroupMembersTx(tx, req.GroupID, req.TunnelIDs, time.Now().UnixMilli()); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - 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 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 { + if err := tx.Commit().Error; err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1753,26 +1663,26 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.ErrDefault("请求参数错误")) return } - tx, err := h.repo.DB().Begin() + tx := h.repo.BeginTx() + if tx.Error != nil { + response.WriteJSON(w, response.Err(-2, tx.Error.Error())) + return + } + defer func() { tx.Rollback() }() + previousUserIDs, err := h.repo.ListUserIDsByUserGroupTx(tx, req.GroupID) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - defer func() { _ = tx.Rollback() }() - previousUserIDs, err := queryInt64ListTx(tx, `SELECT user_id FROM user_group_user WHERE user_group_id = ?`, req.GroupID) - if err != nil { + if err := h.repo.ReplaceUserGroupMembersTx(tx, req.GroupID, req.UserIDs, time.Now().UnixMilli()); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID) - for _, uid := range req.UserIDs { - _, _ = 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 := revokeGroupGrantsForRemovedUsersTx(tx, req.GroupID, previousUserIDs, req.UserIDs); err != nil { + if err := h.repo.RevokeGroupGrantsForRemovedUsersTx(tx, req.GroupID, previousUserIDs, req.UserIDs); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := tx.Commit(); err != nil { + if err := tx.Commit().Error; err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1789,8 +1699,7 @@ func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request) response.WriteJSON(w, response.ErrDefault("请求参数错误")) return } - _, 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 { + if err := h.repo.InsertGroupPermission(req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli()); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1803,32 +1712,31 @@ func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) if id <= 0 { return } - tx, err := h.repo.DB().Begin() + tx := h.repo.BeginTx() + if tx.Error != nil { + response.WriteJSON(w, response.Err(-2, tx.Error.Error())) + return + } + defer func() { tx.Rollback() }() + + ug, tg, exists, err := h.repo.GetGroupPermissionPairByIDTx(tx, id) if err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - defer func() { _ = tx.Rollback() }() - var ug, tg int64 - err = tx.QueryRow(`SELECT user_group_id, tunnel_group_id FROM group_permission WHERE id = ?`, id).Scan(&ug, &tg) - if err != nil && err != sql.ErrNoRows { + if err := h.repo.DeleteGroupPermissionByIDTx(tx, id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } - - if _, err := tx.Exec(`DELETE FROM group_permission WHERE id = ?`, id); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - if err == nil { - if err := revokeGroupPermissionPairTx(tx, ug, tg); err != nil { + if exists { + if err := h.repo.RevokeGroupPermissionPairTx(tx, ug, tg); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } } - if err := tx.Commit(); err != nil { + if err := tx.Commit().Error; err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1847,8 +1755,7 @@ func (h *Handler) groupCreate(w http.ResponseWriter, r *http.Request, table stri return } now := time.Now().UnixMilli() - _, err := h.repo.DB().Exec(`INSERT INTO `+table+`(name, created_time, updated_time, status) VALUES(?, ?, ?, ?)`, name, now, now, asInt(req["status"], 1)) - if err != nil { + if err := h.repo.GroupCreate(table, name, asInt(req["status"], 1), now); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1866,8 +1773,7 @@ func (h *Handler) groupUpdate(w http.ResponseWriter, r *http.Request, table stri response.WriteJSON(w, response.ErrDefault("分组ID不能为空")) return } - _, err := h.repo.DB().Exec(`UPDATE `+table+` SET name = ?, status = ?, updated_time = ? WHERE id = ?`, asString(req["name"]), asInt(req["status"], 1), time.Now().UnixMilli(), id) - if err != nil { + if err := h.repo.GroupUpdate(table, id, asString(req["name"]), asInt(req["status"], 1), time.Now().UnixMilli()); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1879,23 +1785,7 @@ func (h *Handler) groupDelete(w http.ResponseWriter, r *http.Request, table stri if id <= 0 { return } - tx, err := h.repo.DB().Begin() - if err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) - return - } - defer func() { _ = tx.Rollback() }() - if table == "tunnel_group" { - _, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM group_permission WHERE tunnel_group_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE tunnel_group_id = ?`, id) - } else { - _, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM group_permission WHERE user_group_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ?`, id) - } - _, _ = tx.Exec(`DELETE FROM `+table+` WHERE id = ?`, id) - if err := tx.Commit(); err != nil { + if err := h.repo.GroupDeleteCascade(table, id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } @@ -1903,12 +1793,11 @@ func (h *Handler) groupDelete(w http.ResponseWriter, r *http.Request, table stri } func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error { - db := h.repo.DB() - userIDs, _ := queryInt64List(db, `SELECT user_id FROM user_group_user WHERE user_group_id = ?`, userGroupID) - tunnelIDs, _ := queryInt64List(db, `SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tunnelGroupID) + userIDs, _ := h.repo.ListUserIDsByUserGroup(userGroupID) + tunnelIDs, _ := h.repo.ListTunnelIDsByTunnelGroup(tunnelGroupID) for _, uid := range userIDs { for _, tid := range tunnelIDs { - utID, created, err := ensureUserTunnelGrant(db, uid, tid) + utID, created, err := h.repo.EnsureUserTunnelGrant(uid, tid) if err != nil { continue } @@ -1916,16 +1805,14 @@ func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error { if created { createdByGroup = 1 } - _, _ = 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()) + h.repo.InsertGroupPermissionGrant(userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli()) } } return nil } func (h *Handler) syncPermissionsByUserGroup(userGroupID int64) error { - db := h.repo.DB() - pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE user_group_id = ?`, userGroupID) + pairs, err := h.repo.ListGroupPermissionPairsByUserGroup(userGroupID) if err != nil { return err } @@ -1936,8 +1823,7 @@ func (h *Handler) syncPermissionsByUserGroup(userGroupID int64) error { } func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error { - db := h.repo.DB() - pairs, err := queryPairs(db, `SELECT user_group_id, tunnel_group_id FROM group_permission WHERE tunnel_group_id = ?`, tunnelGroupID) + pairs, err := h.repo.ListGroupPermissionPairsByTunnelGroup(tunnelGroupID) if err != nil { return err } @@ -1947,202 +1833,6 @@ func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error { return nil } -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 { - return id, false, nil - } - if err != sql.ErrNoRows { - return 0, false, err - } - var flow int64 - var num int - var expTime int64 - var flowReset int64 - 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 - } - 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 - } - return id, true, nil -} - -func queryInt64List(db *store.DB, q string, args ...interface{}) ([]int64, error) { - rows, err := db.Query(q, args...) - if err != nil { - return nil, err - } - defer rows.Close() - out := make([]int64, 0) - for rows.Next() { - var v int64 - if err := rows.Scan(&v); err != nil { - return nil, err - } - out = append(out, v) - } - return out, rows.Err() -} - -func queryInt64ListTx(tx *store.Tx, q string, args ...interface{}) ([]int64, error) { - rows, err := tx.Query(q, args...) - if err != nil { - return nil, err - } - defer rows.Close() - out := make([]int64, 0) - for rows.Next() { - var v int64 - if err := rows.Scan(&v); err != nil { - return nil, err - } - out = append(out, v) - } - return out, rows.Err() -} - -func revokeGroupGrantsForRemovedUsersTx(tx *store.Tx, userGroupID int64, previousUserIDs, currentUserIDs []int64) error { - currentSet := make(map[int64]struct{}, len(currentUserIDs)) - for _, uid := range currentUserIDs { - if uid > 0 { - currentSet[uid] = struct{}{} - } - } - - removedUserIDs := make([]int64, 0) - for _, uid := range previousUserIDs { - if uid <= 0 { - continue - } - if _, ok := currentSet[uid]; !ok { - removedUserIDs = append(removedUserIDs, uid) - } - } - if len(removedUserIDs) == 0 { - return nil - } - - for _, userID := range removedUserIDs { - rows, err := tx.Query(` - SELECT g.user_tunnel_id, g.created_by_group - FROM group_permission_grant g - JOIN user_tunnel ut ON ut.id = g.user_tunnel_id - WHERE g.user_group_id = ? AND ut.user_id = ? - `, userGroupID, userID) - if err != nil { - return err - } - - groupCreatedTunnelIDs := make(map[int64]struct{}) - for rows.Next() { - var userTunnelID int64 - var createdByGroup int - if err := rows.Scan(&userTunnelID, &createdByGroup); err != nil { - rows.Close() - return err - } - if createdByGroup == 1 && userTunnelID > 0 { - groupCreatedTunnelIDs[userTunnelID] = struct{}{} - } - } - if err := rows.Err(); err != nil { - rows.Close() - return err - } - rows.Close() - - if _, err := tx.Exec(` - DELETE FROM group_permission_grant - WHERE user_group_id = ? - AND user_tunnel_id IN (SELECT id FROM user_tunnel WHERE user_id = ?) - `, userGroupID, userID); err != nil { - return err - } - - for userTunnelID := range groupCreatedTunnelIDs { - var remaining int - if err := tx.QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&remaining); err != nil { - return err - } - if remaining == 0 { - if _, err := tx.Exec(`DELETE FROM user_tunnel WHERE id = ?`, userTunnelID); err != nil { - return err - } - } - } - } - - return nil -} - -func revokeGroupPermissionPairTx(tx *store.Tx, userGroupID, tunnelGroupID int64) error { - rows, err := tx.Query(` - SELECT user_tunnel_id, created_by_group - FROM group_permission_grant - WHERE user_group_id = ? AND tunnel_group_id = ? - `, userGroupID, tunnelGroupID) - if err != nil { - return err - } - - groupCreatedTunnelIDs := make(map[int64]struct{}) - for rows.Next() { - var userTunnelID int64 - var createdByGroup int - if err := rows.Scan(&userTunnelID, &createdByGroup); err != nil { - rows.Close() - return err - } - if createdByGroup == 1 && userTunnelID > 0 { - groupCreatedTunnelIDs[userTunnelID] = struct{}{} - } - } - if err := rows.Err(); err != nil { - rows.Close() - return err - } - rows.Close() - - if _, err := tx.Exec(`DELETE FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID); err != nil { - return err - } - - for userTunnelID := range groupCreatedTunnelIDs { - var remaining int - if err := tx.QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&remaining); err != nil { - return err - } - if remaining == 0 { - if _, err := tx.Exec(`DELETE FROM user_tunnel WHERE id = ?`, userTunnelID); err != nil { - return err - } - } - } - - return nil -} - -func queryPairs(db *store.DB, q string, args ...interface{}) ([][2]int64, error) { - rows, err := db.Query(q, args...) - if err != nil { - return nil, err - } - defer rows.Close() - out := make([][2]int64, 0) - for rows.Next() { - var a, b int64 - if err := rows.Scan(&a, &b); err != nil { - return nil, err - } - out = append(out, [2]int64{a, b}) - } - return out, rows.Err() -} - type tunnelRuntimeNode struct { NodeID int64 Protocol string @@ -2163,7 +1853,7 @@ type tunnelCreateState struct { NodeIDList []int64 } -func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) { +func (h *Handler) prepareTunnelCreateState(tx *gorm.DB, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) { state := &tunnelCreateState{ Type: tunnelType, InNodes: make([]tunnelRuntimeNode, 0), @@ -2205,13 +1895,13 @@ func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interfac nodeIDs = append(nodeIDs, nodeID) port := asInt(item["port"], 0) if port <= 0 { - isRemote, remoteErr := isRemoteNodeTx(tx, nodeID) + isRemote, remoteErr := h.repo.IsRemoteNodeTx(tx, nodeID) if remoteErr != nil { return nil, remoteErr } if !isRemote { var err error - port, err = pickNodePortTx(tx, nodeID, allocated, excludeTunnelID) + port, err = h.repo.PickNodePortTx(tx, nodeID, allocated, excludeTunnelID) if err != nil { return nil, err } @@ -2239,13 +1929,13 @@ func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interfac nodeIDs = append(nodeIDs, nodeID) port := asInt(item["port"], 0) if port <= 0 { - isRemote, remoteErr := isRemoteNodeTx(tx, nodeID) + isRemote, remoteErr := h.repo.IsRemoteNodeTx(tx, nodeID) if remoteErr != nil { return nil, remoteErr } if !isRemote { var err error - port, err = pickNodePortTx(tx, nodeID, allocated, excludeTunnelID) + port, err = h.repo.PickNodePortTx(tx, nodeID, allocated, excludeTunnelID) if err != nil { return nil, err } @@ -2273,13 +1963,16 @@ func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interfac } seen[nodeID] = struct{}{} state.NodeIDList = append(state.NodeIDList, nodeID) - node, err := h.getNodeRecord(nodeID) + node, err := h.repo.GetNodeRecordTx(tx, nodeID) if err != nil { if strings.Contains(err.Error(), "不存在") { return nil, errors.New("节点不存在") } return nil, err } + if node == nil { + return nil, errors.New("节点不存在") + } if node.IsRemote != 1 && node.Status != 1 { return nil, errors.New("部分节点不在线") } @@ -2397,14 +2090,13 @@ func (h *Handler) federationLocalDomain() string { return strings.TrimSpace(cfg.Value) } -func (h *Handler) applyFederationRuntime(state *tunnelCreateState) ([]sqlite.FederationTunnelBinding, []federationRuntimeReleaseRef, error) { - bindings := make([]sqlite.FederationTunnelBinding, 0) +func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain string) ([]repo.FederationTunnelBinding, []federationRuntimeReleaseRef, error) { + bindings := make([]repo.FederationTunnelBinding, 0) releaseRefs := make([]federationRuntimeReleaseRef, 0) if h == nil || state == nil { return bindings, releaseRefs, nil } fc := client.NewFederationClient() - localDomain := h.federationLocalDomain() now := time.Now().UnixMilli() for outIdx := range state.OutNodes { @@ -2456,7 +2148,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState) ([]sqlite.Fed outNode = state.OutNodes[outIdx] } - bindings = append(bindings, sqlite.FederationTunnelBinding{ + bindings = append(bindings, repo.FederationTunnelBinding{ TunnelID: state.TunnelID, NodeID: outNode.NodeID, ChainType: 3, @@ -2556,7 +2248,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState) ([]sqlite.Fed chainNode = state.ChainHops[hopIdx][nodeIdx] } - bindings = append(bindings, sqlite.FederationTunnelBinding{ + bindings = append(bindings, repo.FederationTunnelBinding{ TunnelID: state.TunnelID, NodeID: chainNode.NodeID, ChainType: 2, @@ -2635,33 +2327,6 @@ func (h *Handler) cleanupFederationRuntime(tunnelID int64) { _ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID) } -func replaceFederationTunnelBindingsTx(tx *store.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error { - if tx == nil { - return errors.New("database unavailable") - } - if _, err := tx.Exec(`DELETE FROM federation_tunnel_binding WHERE tunnel_id = ?`, tunnelID); err != nil { - return err - } - for _, b := range bindings { - created := b.CreatedTime - if created <= 0 { - created = time.Now().UnixMilli() - } - updated := b.UpdatedTime - if updated <= 0 { - updated = created - } - _, err := tx.Exec(` - INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, tunnelID, b.NodeID, b.ChainType, b.HopInx, b.RemoteURL, b.ResourceKey, b.RemoteBindingID, b.AllocatedPort, b.Status, created, updated) - if err != nil { - return err - } - } - return nil -} - func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) { if h == nil || state == nil { return nil, nil, errors.New("invalid tunnel runtime state") @@ -2983,135 +2648,7 @@ func pickNodeAddressV6(node *nodeRecord) string { return strings.TrimSpace(node.ServerIP) } -func isRemoteNodeTx(tx *store.Tx, nodeID int64) (bool, error) { - if tx == nil { - return false, errors.New("database unavailable") - } - if nodeID <= 0 { - return false, errors.New("节点不存在") - } - var isRemote int - if err := tx.QueryRow(`SELECT is_remote FROM node WHERE id = ? LIMIT 1`, nodeID).Scan(&isRemote); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return false, errors.New("节点不存在") - } - return false, err - } - return isRemote == 1, nil -} - -func pickNodePortTx(tx *store.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) { - if tx == nil { - return 0, errors.New("database unavailable") - } - if nodeID <= 0 { - return 0, errors.New("节点不存在") - } - if port, ok := allocated[nodeID]; ok && port > 0 { - return port, nil - } - - var portRange string - if err := tx.QueryRow(`SELECT port FROM node WHERE id = ? LIMIT 1`, nodeID).Scan(&portRange); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return 0, errors.New("节点不存在") - } - return 0, err - } - candidates := parsePortRangeSpec(portRange) - if len(candidates) == 0 { - return 0, errors.New("节点端口已满,无可用端口") - } - - used := map[int]struct{}{} - var chainRows *sql.Rows - var err error - if excludeTunnelID > 0 { - chainRows, err = tx.Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL AND tunnel_id != ?`, nodeID, excludeTunnelID) - } else { - chainRows, err = tx.Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL`, nodeID) - } - if err != nil { - return 0, err - } - for chainRows.Next() { - var p sql.NullInt64 - if scanErr := chainRows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 { - used[int(p.Int64)] = struct{}{} - } - } - _ = chainRows.Close() - - forwardRows, err := tx.Query(`SELECT port FROM forward_port WHERE node_id = ?`, nodeID) - if err != nil { - return 0, err - } - for forwardRows.Next() { - var p sql.NullInt64 - if scanErr := forwardRows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 { - used[int(p.Int64)] = struct{}{} - } - } - _ = forwardRows.Close() - - for _, candidate := range candidates { - if candidate <= 0 { - continue - } - if _, ok := used[candidate]; ok { - continue - } - allocated[nodeID] = candidate - return candidate, nil - } - return 0, errors.New("节点端口已满,无可用端口") -} - -func parsePortRangeSpec(input string) []int { - input = strings.TrimSpace(input) - if input == "" { - return nil - } - set := make(map[int]struct{}) - parts := strings.Split(input, ",") - for _, part := range parts { - part = strings.TrimSpace(part) - if part == "" { - continue - } - if strings.Contains(part, "-") { - r := strings.SplitN(part, "-", 2) - if len(r) != 2 { - continue - } - start, err1 := strconv.Atoi(strings.TrimSpace(r[0])) - end, err2 := strconv.Atoi(strings.TrimSpace(r[1])) - if err1 != nil || err2 != nil || start <= 0 || end <= 0 { - continue - } - if end < start { - start, end = end, start - } - for p := start; p <= end; p++ { - set[p] = struct{}{} - } - continue - } - p, err := strconv.Atoi(part) - if err != nil || p <= 0 { - continue - } - set[p] = struct{}{} - } - out := make([]int, 0, len(set)) - for p := range set { - out = append(out, p) - } - sort.Ints(out) - return out -} - -func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interface{}) error { +func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[string]interface{}) error { allocated := map[int64]int{} inNodes := asMapSlice(req["inNodeId"]) for _, n := range inNodes { @@ -3119,9 +2656,16 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac if nodeID <= 0 { continue } - _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '1', ?, NULL, ?, 0, ?)`, - tunnelID, nodeID, defaultString(asString(n["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls")) - if err != nil { + if err := h.repo.CreateChainTunnelTx( + tx, + tunnelID, + "1", + nodeID, + sql.NullInt64{}, + defaultString(asString(n["strategy"]), "round"), + 0, + defaultString(asString(n["protocol"]), "tls"), + ); err != nil { return err } } @@ -3133,14 +2677,21 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac port := asInt(n["port"], 0) if port <= 0 { var pickErr error - port, pickErr = pickNodePortTx(tx, nodeID, allocated, 0) + port, pickErr = h.repo.PickNodePortTx(tx, nodeID, allocated, 0) if pickErr != nil { return pickErr } } - _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '3', ?, ?, ?, 0, ?)`, - tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls")) - if err != nil { + if err := h.repo.CreateChainTunnelTx( + tx, + tunnelID, + "3", + nodeID, + sql.NullInt64{Int64: int64(port), Valid: true}, + defaultString(asString(n["strategy"]), "round"), + 0, + defaultString(asString(n["protocol"]), "tls"), + ); err != nil { return err } } @@ -3154,14 +2705,21 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac port := asInt(n["port"], 0) if port <= 0 { var pickErr error - port, pickErr = pickNodePortTx(tx, nodeID, allocated, 0) + port, pickErr = h.repo.PickNodePortTx(tx, nodeID, allocated, 0) if pickErr != nil { return pickErr } } - _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '2', ?, ?, ?, ?, ?)`, - tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls")) - if err != nil { + if err := h.repo.CreateChainTunnelTx( + tx, + tunnelID, + "2", + nodeID, + sql.NullInt64{Int64: int64(port), Valid: true}, + defaultString(asString(n["strategy"]), "round"), + i+1, + defaultString(asString(n["protocol"]), "tls"), + ); err != nil { return err } } @@ -3170,52 +2728,15 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac } func (h *Handler) deleteNodeByID(id int64) error { - tx, err := h.repo.DB().Begin() - if err != nil { - return err - } - defer func() { _ = tx.Rollback() }() - _, _ = tx.Exec(`DELETE FROM forward_port WHERE node_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE node_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM federation_tunnel_binding WHERE node_id = ?`, id) - _, err = tx.Exec(`DELETE FROM node WHERE id = ?`, id) - if err != nil { - return err - } - return tx.Commit() + return h.repo.DeleteNodeCascade(id) } func (h *Handler) deleteTunnelByID(id int64) error { - tx, err := h.repo.DB().Begin() - if err != nil { - return err - } - defer func() { _ = tx.Rollback() }() - _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id IN (SELECT id FROM forward WHERE tunnel_id = ?)`, id) - _, _ = tx.Exec(`DELETE FROM forward WHERE tunnel_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM user_tunnel WHERE tunnel_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM speed_limit WHERE tunnel_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id) - _, _ = tx.Exec(`DELETE FROM federation_tunnel_binding WHERE tunnel_id = ?`, id) - _, err = tx.Exec(`DELETE FROM tunnel WHERE id = ?`, id) - if err != nil { - return err - } - return tx.Commit() + return h.repo.DeleteTunnelCascade(id) } func (h *Handler) deleteForwardByID(id int64) error { - tx, err := h.repo.DB().Begin() - if err != nil { - return err - } - defer func() { _ = tx.Rollback() }() - _, _ = tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, id) - _, err = tx.Exec(`DELETE FROM forward WHERE id = ?`, id) - if err != nil { - return err - } - return tx.Commit() + return h.repo.DeleteForwardCascade(id) } func (h *Handler) batchForwardDelete(ids []int64) (int, int) { @@ -3232,32 +2753,11 @@ func (h *Handler) batchForwardDelete(ids []int64) (int, int) { } func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) { - s := 0 - f := 0 - for _, id := range ids { - if _, err := h.repo.DB().Exec(`UPDATE forward SET status = ?, updated_time = ? WHERE id = ?`, status, time.Now().UnixMilli(), id); err != nil { - f++ - } else { - s++ - } - } - return s, f + return h.repo.BatchUpdateForwardStatus(ids, status) } func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) { - rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = '1' ORDER BY inx ASC, id ASC`, tunnelID) - if err != nil { - return nil, err - } - defer rows.Close() - out := make([]int64, 0) - for rows.Next() { - var id int64 - if err := rows.Scan(&id); err == nil { - out = append(out, id) - } - } - return out, rows.Err() + return h.repo.TunnelEntryNodeIDs(tunnelID) } func (h *Handler) pickTunnelPort(tunnelID int64) int { @@ -3270,8 +2770,8 @@ func (h *Handler) pickTunnelPort(tunnelID int64) int { firstNode := true for _, nodeID := range entryNodes { - var portRange string - if err := h.repo.DB().QueryRow("SELECT port FROM node WHERE id = ?", nodeID).Scan(&portRange); err != nil { + portRange, err := h.repo.GetNodePortRange(nodeID) + if err != nil { continue } if portRange == "" { @@ -3326,30 +2826,7 @@ func (h *Handler) pickTunnelPort(tunnelID int64) int { } func (h *Handler) getUsedPorts(nodeID int64) (map[int]bool, error) { - used := make(map[int]bool) - rows, err := h.repo.DB().Query("SELECT port FROM forward_port WHERE node_id = ?", nodeID) - if err != nil { - return nil, err - } - defer rows.Close() - for rows.Next() { - var p int - if err := rows.Scan(&p); err == nil { - used[p] = true - } - } - rows2, err := h.repo.DB().Query("SELECT port FROM chain_tunnel WHERE node_id = ? AND port > 0", nodeID) - if err != nil { - return nil, err - } - defer rows2.Close() - for rows2.Next() { - var p int - if err := rows2.Scan(&p); err == nil { - used[p] = true - } - } - return used, nil + return h.repo.GetUsedPortsOnNodeAsMap(nodeID) } func parsePorts(portRange string) ([]int, error) { @@ -3385,55 +2862,47 @@ func parsePorts(portRange string) ([]int, error) { } func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error { - tx, err := h.repo.DB().Begin() - if err != nil { - return err - } - defer func() { _ = tx.Rollback() }() - if _, err := tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID); err != nil { - return err - } entryNodes, err := h.tunnelEntryNodeIDs(tunnelID) if err != nil { return err } - for _, nodeID := range entryNodes { - if _, err := tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port); err != nil { - return err - } + entries := make([]struct { + NodeID int64 + Port int + }, len(entryNodes)) + for i, nid := range entryNodes { + entries[i] = struct { + NodeID int64 + Port int + }{NodeID: nid, Port: port} } - return tx.Commit() + return h.repo.ReplaceForwardPorts(forwardID, entries) } func (h *Handler) replaceForwardPortsWithRecords(forwardID int64, ports []forwardPortRecord) error { - tx, err := h.repo.DB().Begin() - if err != nil { - return err + entries := make([]struct { + NodeID int64 + Port int + }, len(ports)) + for i, fp := range ports { + entries[i] = struct { + NodeID int64 + Port int + }{NodeID: fp.NodeID, Port: fp.Port} } - defer func() { _ = tx.Rollback() }() - - if _, err := tx.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID); err != nil { - return err - } - for _, fp := range ports { - if _, err := tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, fp.NodeID, fp.Port); err != nil { - return err - } - } - - return tx.Commit() + return h.repo.ReplaceForwardPorts(forwardID, entries) } func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []forwardPortRecord) { - if h == nil || oldForward == nil || h.repo == nil || h.repo.DB() == nil { + if h == nil || oldForward == nil || h.repo == nil { return } - _, _ = h.repo.DB().Exec(` - UPDATE forward - SET user_id = ?, user_name = ?, name = ?, tunnel_id = ?, remote_addr = ?, strategy = ?, status = ?, updated_time = ? - WHERE id = ? - `, oldForward.UserID, oldForward.UserName, oldForward.Name, oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status, time.Now().UnixMilli(), oldForward.ID) + h.repo.RollbackForwardFields( + oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name, + oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status, + time.Now().UnixMilli(), + ) if err := h.replaceForwardPortsWithRecords(oldForward.ID, oldPorts); err != nil { return @@ -3448,18 +2917,9 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { if userID <= 0 || tunnelID <= 0 { return fmt.Errorf("userId or tunnelId missing") } - db := h.repo.DB() - var existingID int64 - var currentFlow, currentNum, currentExpTime, currentFlowReset int64 - var currentSpeedID sql.NullInt64 - var currentStatus int - err := db.QueryRow(` - SELECT id, flow, num, exp_time, flow_reset_time, speed_id, status - FROM user_tunnel - WHERE user_id = ? AND tunnel_id = ? - LIMIT 1 - `, userID, tunnelID).Scan(&existingID, ¤tFlow, ¤tNum, ¤tExpTime, ¤tFlowReset, ¤tSpeedID, ¤tStatus) + existingID, currentFlow, currentNum, currentExpTime, currentFlowReset, currentSpeedID, currentStatus, err := + h.repo.GetExistingUserTunnel(userID, tunnelID) speedID := asAnyToInt64Ptr(req["speedId"]) reqFlow := asInt64(req["flow"], -1) @@ -3470,13 +2930,13 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { if err == sql.ErrNoRows { if reqFlow < 0 || reqNum < 0 || reqExpTime < 0 || reqFlowReset < 0 { - var uFlow, uNum, uExp, uReset int64 - if uErr := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&uFlow, &uNum, &uExp, &uReset); uErr == nil { + uFlow, uNum, uExp, uReset, uErr := h.repo.GetUserDefaultsForTunnel(userID) + if uErr == nil { if reqFlow < 0 { reqFlow = uFlow } if reqNum < 0 { - reqNum = int(uNum) + reqNum = uNum } if reqExpTime < 0 { reqExpTime = uExp @@ -3502,9 +2962,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { reqStatus = 1 } - _, 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(?, ?, ?, ?, ?, 0, 0, ?, ?, ?)`, - userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus) - return err + return h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus) } if err != nil { return err @@ -3542,8 +3000,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error { newSpeedID = sql.NullInt64{Valid: false} } - _, err = db.Exec(`UPDATE user_tunnel SET speed_id = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ? WHERE id = ?`, - newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus, existingID) + err = h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus) if err == nil { h.syncUserTunnelForwards(userID, tunnelID) @@ -3731,18 +3188,3 @@ func randomToken(n int) string { } return hex.EncodeToString(buf) } - -func nextIndex(db *store.DB, table string) int { - if db == nil { - return 0 - } - row := db.QueryRow(`SELECT COALESCE(MAX(inx), -1) + 1 FROM ` + table) - var n int - if err := row.Scan(&n); err != nil { - return 0 - } - if n < 0 { - return 0 - } - return n -} diff --git a/go-backend/internal/store/db.go b/go-backend/internal/store/db.go deleted file mode 100644 index a597eef..0000000 --- a/go-backend/internal/store/db.go +++ /dev/null @@ -1,466 +0,0 @@ -// Package store provides a thin dialect-aware wrapper around database/sql, -// enabling transparent use of both SQLite and PostgreSQL. -package store - -import ( - "database/sql" - "strconv" - "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 = ensureReturningID(q) - 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 = ensureReturningID(q) - 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) - i := 0 - for i < len(query) { - if end, ok := skipSQLProtectedSegment(query, i); ok { - buf.WriteString(query[i:end]) - i = end - continue - } - - ch := query[i] - 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 { - start, end, ok := findKeywordSequenceOutside(query, []string{"INSERT", "OR", "IGNORE", "INTO"}, 0) - if !ok { - return query - } - - rewritten := query[:start] + "INSERT INTO" + query[end:] - rewritten = strings.TrimRight(rewritten, "; \t\n") - - insertIntoEnd := start + len("INSERT INTO") - if _, _, hasOnConflict := findKeywordSequenceOutside(rewritten, []string{"ON", "CONFLICT"}, insertIntoEnd); hasOnConflict { - return rewritten - } - - if retStart, _, hasReturning := findKeywordSequenceOutside(rewritten, []string{"RETURNING"}, insertIntoEnd); hasReturning { - prefix := strings.TrimRight(rewritten[:retStart], " \t\n") - suffix := strings.TrimLeft(rewritten[retStart:], " \t\n") - return prefix + " ON CONFLICT DO NOTHING " + suffix - } - - return rewritten + " ON CONFLICT DO NOTHING" -} - -func rewritePlaceholders(query string) string { - var buf strings.Builder - buf.Grow(len(query) + 16) - n := 1 - for i := 0; i < len(query); i++ { - if end, ok := skipSQLProtectedSegment(query, i); ok { - buf.WriteString(query[i:end]) - i = end - 1 - continue - } - - ch := query[i] - if ch == '?' { - buf.WriteByte('$') - buf.WriteString(strconv.Itoa(n)) - n++ - continue - } - buf.WriteByte(ch) - } - return buf.String() -} - -func ensureReturningID(query string) string { - trimmed := strings.TrimRight(query, "; \t\n") - if _, _, ok := findKeywordSequenceOutside(trimmed, []string{"RETURNING"}, 0); ok { - return trimmed - } - return trimmed + " RETURNING id" -} - -func findKeywordSequenceOutside(query string, keywords []string, from int) (int, int, bool) { - if len(keywords) == 0 { - return 0, 0, false - } - if from < 0 { - from = 0 - } - if from >= len(query) { - return 0, 0, false - } - - matched := 0 - seqStart := -1 - - for i := from; i < len(query); { - if end, ok := skipSQLProtectedSegment(query, i); ok { - i = end - continue - } - - ch := query[i] - if isIdentifierChar(ch) { - j := i + 1 - for j < len(query) && isIdentifierChar(query[j]) { - j++ - } - tok := query[i:j] - - if strings.EqualFold(tok, keywords[matched]) { - if matched == 0 { - seqStart = i - } - matched++ - if matched == len(keywords) { - return seqStart, j, true - } - } else if strings.EqualFold(tok, keywords[0]) { - seqStart = i - matched = 1 - } else { - matched = 0 - seqStart = -1 - } - - i = j - continue - } - - if !isSQLSpace(ch) { - matched = 0 - seqStart = -1 - } - i++ - } - - return 0, 0, false -} - -func skipSQLProtectedSegment(query string, i int) (int, bool) { - if i < 0 || i >= len(query) { - return 0, false - } - - switch query[i] { - case '\'': - return skipSingleQuotedLiteral(query, i), true - case '"': - return skipDoubleQuotedIdentifier(query, i), true - case '-': - if i+1 < len(query) && query[i+1] == '-' { - return skipLineComment(query, i), true - } - case '/': - if i+1 < len(query) && query[i+1] == '*' { - return skipBlockComment(query, i), true - } - case '$': - if end, ok := skipDollarQuotedLiteral(query, i); ok { - return end, true - } - } - - return 0, false -} - -func skipSingleQuotedLiteral(query string, i int) int { - for j := i + 1; j < len(query); j++ { - if query[j] != '\'' { - continue - } - if j+1 < len(query) && query[j+1] == '\'' { - j++ - continue - } - return j + 1 - } - return len(query) -} - -func skipDoubleQuotedIdentifier(query string, i int) int { - for j := i + 1; j < len(query); j++ { - if query[j] != '"' { - continue - } - if j+1 < len(query) && query[j+1] == '"' { - j++ - continue - } - return j + 1 - } - return len(query) -} - -func skipLineComment(query string, i int) int { - for j := i + 2; j < len(query); j++ { - if query[j] == '\n' { - return j - } - } - return len(query) -} - -func skipBlockComment(query string, i int) int { - depth := 1 - for j := i + 2; j < len(query)-1; j++ { - if query[j] == '/' && query[j+1] == '*' { - depth++ - j++ - continue - } - if query[j] == '*' && query[j+1] == '/' { - depth-- - j++ - if depth == 0 { - return j + 1 - } - } - } - return len(query) -} - -func skipDollarQuotedLiteral(query string, i int) (int, bool) { - if i < 0 || i >= len(query) || query[i] != '$' { - return 0, false - } - - if i+1 >= len(query) { - return 0, false - } - - var endTag int - if query[i+1] == '$' { - endTag = i + 1 - } else { - if !isDollarTagStart(query[i+1]) { - return 0, false - } - j := i + 2 - for j < len(query) && isDollarTagChar(query[j]) { - j++ - } - if j >= len(query) || query[j] != '$' { - return 0, false - } - endTag = j - } - - tag := query[i : endTag+1] - if closeIdx := strings.Index(query[endTag+1:], tag); closeIdx >= 0 { - return endTag + 1 + closeIdx + len(tag), true - } - return len(query), true -} - -func isDollarTagStart(ch byte) bool { - return ch == '_' || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') -} - -func isDollarTagChar(ch byte) bool { - if isDollarTagStart(ch) { - return true - } - return ch >= '0' && ch <= '9' -} - -func isSQLSpace(ch byte) bool { - switch ch { - case ' ', '\t', '\n', '\r', '\f': - return true - default: - return false - } -} diff --git a/go-backend/internal/store/db_test.go b/go-backend/internal/store/db_test.go deleted file mode 100644 index 4a3f580..0000000 --- a/go-backend/internal/store/db_test.go +++ /dev/null @@ -1,116 +0,0 @@ -package store - -import "testing" - -func TestRewritePlaceholdersSkipsProtectedSegments(t *testing.T) { - q := `SELECT ?, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ? -FROM t /* block ? */ WHERE id = ?` - got := rewritePlaceholders(q) - want := `SELECT $1, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ? -FROM t /* block ? */ WHERE id = $2` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewriteInsertOrIgnoreBasic(t *testing.T) { - q := `INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)` - got := rewriteInsertOrIgnore(q) - want := `INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewriteInsertOrIgnoreBeforeReturning(t *testing.T) { - q := `INSERT OR IGNORE INTO x(a) VALUES(?) RETURNING id` - got := rewriteInsertOrIgnore(q) - want := `INSERT INTO x(a) VALUES(?) ON CONFLICT DO NOTHING RETURNING id` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewriteInsertOrIgnoreNotDuplicatingOnConflict(t *testing.T) { - q := `INSERT OR IGNORE INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a` - got := rewriteInsertOrIgnore(q) - want := `INSERT INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestEnsureReturningID(t *testing.T) { - if got := ensureReturningID(`INSERT INTO x(a) VALUES($1)`); got != `INSERT INTO x(a) VALUES($1) RETURNING id` { - t.Fatalf("missing RETURNING append: %s", got) - } - if got := ensureReturningID(`INSERT INTO x(a) VALUES($1) RETURNING other_id`); got != `INSERT INTO x(a) VALUES($1) RETURNING other_id` { - t.Fatalf("RETURNING should not be duplicated: %s", got) - } -} - -func TestRewriteUserIdentifierSafety(t *testing.T) { - q := `SELECT user, user_id, 'user', "user", note FROM user -- user -WHERE owner='user'` - got := rewriteUserIdentifier(q) - want := `SELECT "user", user_id, 'user', "user", note FROM "user" -- user -WHERE owner='user'` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewriteQueryPostgresPipeline(t *testing.T) { - q := `INSERT OR IGNORE INTO user(name, note) VALUES(?, '?')` - got := rewriteQuery(DialectPostgres, q) - want := `INSERT INTO "user"(name, note) VALUES($1, '?') ON CONFLICT DO NOTHING` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewriteInsertOrIgnoreSkipsStringLiteral(t *testing.T) { - q := `SELECT 'INSERT OR IGNORE INTO t(a) VALUES(?)' AS q` - got := rewriteInsertOrIgnore(q) - if got != q { - t.Fatalf("string literal should stay unchanged\nwant: %s\ngot: %s", q, got) - } -} - -func TestRewriteInsertOrIgnoreSkipsCommentedKeyword(t *testing.T) { - q := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?) -INSERT OR IGNORE INTO real_t(a) VALUES(?)` - got := rewriteInsertOrIgnore(q) - want := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?) -INSERT INTO real_t(a) VALUES(?) ON CONFLICT DO NOTHING` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewritePlaceholdersSkipsNestedBlockComment(t *testing.T) { - q := `SELECT ? /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = ?` - got := rewritePlaceholders(q) - want := `SELECT $1 /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = $2` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewritePlaceholdersSkipsUnterminatedBlockComment(t *testing.T) { - q := `SELECT ? /* unterminated ? comment` - got := rewritePlaceholders(q) - want := `SELECT $1 /* unterminated ? comment` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} - -func TestRewriteUserIdentifierSkipsDollarQuotedAndComment(t *testing.T) { - q := `SELECT user, $$user ?$$ AS body, col FROM user /* user */ -- user` - got := rewriteUserIdentifier(q) - want := `SELECT "user", $$user ?$$ AS body, col FROM "user" /* user */ -- user` - if got != want { - t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) - } -} diff --git a/go-backend/internal/store/model/model.go b/go-backend/internal/store/model/model.go new file mode 100644 index 0000000..dd2bc10 --- /dev/null +++ b/go-backend/internal/store/model/model.go @@ -0,0 +1,591 @@ +// Package model defines GORM model structs for all database tables, +// providing a single source of truth for the schema that works +// transparently with both SQLite and PostgreSQL. +package model + +import "database/sql" + +// ─── Core Business Tables ──────────────────────────────────────────── + +// User maps to the "user" table. PostgreSQL treats "user" as a reserved +// word, so TableName() is required for correct quoting. +type User struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + User string `gorm:"column:user;type:varchar(100);not null"` + Pwd string `gorm:"type:varchar(100);not null"` + RoleID int `gorm:"column:role_id;not null"` + ExpTime int64 `gorm:"column:exp_time;not null"` + Flow int64 `gorm:"not null"` + InFlow int64 `gorm:"column:in_flow;not null;default:0"` + OutFlow int64 `gorm:"column:out_flow;not null;default:0"` + FlowResetTime int64 `gorm:"column:flow_reset_time;not null"` + Num int `gorm:"not null"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime sql.NullInt64 `gorm:"column:updated_time"` + Status int `gorm:"not null"` +} + +func (User) TableName() string { return "user" } + +// Forward maps to the "forward" table. +type Forward struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + UserID int64 `gorm:"column:user_id;not null"` + UserName string `gorm:"column:user_name;type:varchar(100);not null"` + Name string `gorm:"type:varchar(100);not null"` + TunnelID int64 `gorm:"column:tunnel_id;not null"` + RemoteAddr string `gorm:"column:remote_addr;type:text;not null"` + Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"` + InFlow int64 `gorm:"column:in_flow;not null;default:0"` + OutFlow int64 `gorm:"column:out_flow;not null;default:0"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` + Status int `gorm:"not null"` + Inx int `gorm:"not null;default:0"` +} + +func (Forward) TableName() string { return "forward" } + +type ForwardPort struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + ForwardID int64 `gorm:"column:forward_id;not null"` + NodeID int64 `gorm:"column:node_id;not null"` + Port int `gorm:"not null"` +} + +func (ForwardPort) TableName() string { return "forward_port" } + +type Node struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + Name string `gorm:"type:varchar(100);not null"` + Secret string `gorm:"type:varchar(100);not null"` + ServerIP string `gorm:"column:server_ip;type:varchar(100);not null"` + ServerIPV4 sql.NullString `gorm:"column:server_ip_v4;type:varchar(100)"` + ServerIPV6 sql.NullString `gorm:"column:server_ip_v6;type:varchar(100)"` + Port string `gorm:"type:text;not null"` + InterfaceName sql.NullString `gorm:"column:interface_name;type:varchar(200)"` + Version sql.NullString `gorm:"type:varchar(100)"` + HTTP int `gorm:"column:http;not null;default:0"` + TLS int `gorm:"column:tls;not null;default:0"` + Socks int `gorm:"not null;default:0"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime sql.NullInt64 `gorm:"column:updated_time"` + Status int `gorm:"not null"` + TCPListenAddr string `gorm:"column:tcp_listen_addr;type:varchar(100);not null;default:'[::]'"` + UDPListenAddr string `gorm:"column:udp_listen_addr;type:varchar(100);not null;default:'[::]'"` + Inx int `gorm:"not null;default:0"` + IsRemote int `gorm:"column:is_remote;default:0"` + RemoteURL sql.NullString `gorm:"column:remote_url;type:text"` + RemoteToken sql.NullString `gorm:"column:remote_token;type:text"` + RemoteConfig sql.NullString `gorm:"column:remote_config;type:text"` +} + +func (Node) TableName() string { return "node" } + +type SpeedLimit struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + Name string `gorm:"type:varchar(100);not null"` + Speed int `gorm:"not null"` + TunnelID int64 `gorm:"column:tunnel_id;not null"` + TunnelName string `gorm:"column:tunnel_name;type:varchar(100);not null"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime sql.NullInt64 `gorm:"column:updated_time"` + Status int `gorm:"not null"` +} + +func (SpeedLimit) TableName() string { return "speed_limit" } + +type StatisticsFlow struct { + ID int64 `gorm:"primaryKey;autoIncrement" json:"id"` + UserID int64 `gorm:"column:user_id;not null" json:"userId"` + Flow int64 `gorm:"not null" json:"flow"` + TotalFlow int64 `gorm:"column:total_flow;not null" json:"totalFlow"` + Time string `gorm:"type:varchar(100);not null" json:"time"` + CreatedTime int64 `gorm:"column:created_time;not null" json:"-"` +} + +func (StatisticsFlow) TableName() string { return "statistics_flow" } + +type Tunnel struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + Name string `gorm:"type:varchar(100);not null"` + TrafficRatio float64 `gorm:"column:traffic_ratio;not null;default:1.0"` + Type int `gorm:"not null"` + Protocol string `gorm:"type:varchar(10);not null;default:'tls'"` + Flow int64 `gorm:"not null"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` + Status int `gorm:"not null"` + InIP sql.NullString `gorm:"column:in_ip;type:text"` + Inx int `gorm:"not null;default:0"` + IPPreference string `gorm:"column:ip_preference;type:varchar(10);not null;default:''"` +} + +func (Tunnel) TableName() string { return "tunnel" } + +type ChainTunnel struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + TunnelID int64 `gorm:"column:tunnel_id;not null"` + ChainType string `gorm:"column:chain_type;type:varchar(10);not null"` + NodeID int64 `gorm:"column:node_id;not null"` + Port sql.NullInt64 `gorm:"column:port"` + Strategy sql.NullString `gorm:"type:varchar(10)"` + Inx sql.NullInt64 `gorm:"column:inx"` + Protocol sql.NullString `gorm:"type:varchar(10)"` +} + +func (ChainTunnel) TableName() string { return "chain_tunnel" } + +type UserTunnel struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + UserID int64 `gorm:"column:user_id;not null;uniqueIndex:idx_user_tunnel_unique"` + TunnelID int64 `gorm:"column:tunnel_id;not null;uniqueIndex:idx_user_tunnel_unique"` + SpeedID sql.NullInt64 `gorm:"column:speed_id"` + Num int `gorm:"not null"` + Flow int64 `gorm:"not null"` + InFlow int64 `gorm:"column:in_flow;not null;default:0"` + OutFlow int64 `gorm:"column:out_flow;not null;default:0"` + FlowResetTime int64 `gorm:"column:flow_reset_time;not null"` + ExpTime int64 `gorm:"column:exp_time;not null"` + Status int `gorm:"not null"` +} + +func (UserTunnel) TableName() string { return "user_tunnel" } + +type TunnelGroup struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + Name string `gorm:"type:varchar(100);not null;uniqueIndex:idx_tunnel_group_name"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` + Status int `gorm:"not null"` +} + +func (TunnelGroup) TableName() string { return "tunnel_group" } + +type UserGroup struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + Name string `gorm:"type:varchar(100);not null;uniqueIndex:idx_user_group_name"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` + Status int `gorm:"not null"` +} + +func (UserGroup) TableName() string { return "user_group" } + +type TunnelGroupTunnel struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + TunnelGroupID int64 `gorm:"column:tunnel_group_id;not null;uniqueIndex:idx_tunnel_group_tunnel_unique"` + TunnelID int64 `gorm:"column:tunnel_id;not null;uniqueIndex:idx_tunnel_group_tunnel_unique"` + CreatedTime int64 `gorm:"column:created_time;not null"` +} + +func (TunnelGroupTunnel) TableName() string { return "tunnel_group_tunnel" } + +type UserGroupUser struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + UserGroupID int64 `gorm:"column:user_group_id;not null;uniqueIndex:idx_user_group_user_unique"` + UserID int64 `gorm:"column:user_id;not null;uniqueIndex:idx_user_group_user_unique"` + CreatedTime int64 `gorm:"column:created_time;not null"` +} + +func (UserGroupUser) TableName() string { return "user_group_user" } + +type GroupPermission struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + UserGroupID int64 `gorm:"column:user_group_id;not null;uniqueIndex:idx_group_permission_unique"` + TunnelGroupID int64 `gorm:"column:tunnel_group_id;not null;uniqueIndex:idx_group_permission_unique"` + CreatedTime int64 `gorm:"column:created_time;not null"` +} + +func (GroupPermission) TableName() string { return "group_permission" } + +type GroupPermissionGrant struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + UserGroupID int64 `gorm:"column:user_group_id;not null;uniqueIndex:idx_group_permission_grant_unique"` + TunnelGroupID int64 `gorm:"column:tunnel_group_id;not null;uniqueIndex:idx_group_permission_grant_unique"` + UserTunnelID int64 `gorm:"column:user_tunnel_id;not null;uniqueIndex:idx_group_permission_grant_unique"` + CreatedByGroup int `gorm:"column:created_by_group;not null;default:0"` + CreatedTime int64 `gorm:"column:created_time;not null"` +} + +func (GroupPermissionGrant) TableName() string { return "group_permission_grant" } + +type ViteConfig struct { + ID int64 `gorm:"primaryKey;autoIncrement" json:"id"` + Name string `gorm:"type:varchar(200);not null;uniqueIndex" json:"name"` + Value string `gorm:"type:varchar(200);not null" json:"value"` + Time int64 `gorm:"not null" json:"time"` +} + +func (ViteConfig) TableName() string { return "vite_config" } + +type Announcement struct { + ID int64 `gorm:"primaryKey;autoIncrement" json:"id"` + Content string `gorm:"type:text;not null" json:"content"` + Enabled int `gorm:"not null;default:1" json:"enabled"` + CreatedTime int64 `gorm:"column:created_time;not null" json:"created_time"` + UpdatedTime sql.NullInt64 `gorm:"column:updated_time" json:"updated_time,omitempty"` +} + +func (Announcement) TableName() string { return "announcement" } + +type SchemaVersion struct { + Version int `gorm:"not null;default:0"` +} + +func (SchemaVersion) TableName() string { return "schema_version" } + +type PeerShare struct { + ID int64 `gorm:"primaryKey;autoIncrement" json:"id"` + Name string `gorm:"type:text;not null" json:"name"` + NodeID int64 `gorm:"column:node_id;not null" json:"nodeId"` + Token string `gorm:"type:text;not null;uniqueIndex" json:"token"` + MaxBandwidth int64 `gorm:"column:max_bandwidth;default:0" json:"maxBandwidth"` + ExpiryTime int64 `gorm:"column:expiry_time;default:0" json:"expiryTime"` + PortRangeStart int `gorm:"column:port_range_start;default:0" json:"portRangeStart"` + PortRangeEnd int `gorm:"column:port_range_end;default:0" json:"portRangeEnd"` + CurrentFlow int64 `gorm:"column:current_flow;default:0" json:"currentFlow"` + IsActive int `gorm:"column:is_active;default:1" json:"isActive"` + CreatedTime int64 `gorm:"column:created_time;not null" json:"createdTime"` + UpdatedTime int64 `gorm:"column:updated_time;not null" json:"updatedTime"` + AllowedDomains string `gorm:"column:allowed_domains;type:text;default:''" json:"allowedDomains"` + AllowedIPs string `gorm:"column:allowed_ips;type:text;default:''" json:"allowedIps"` +} + +func (PeerShare) TableName() string { return "peer_share" } + +type PeerShareRuntime struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + ShareID int64 `gorm:"column:share_id;not null;index:idx_peer_share_runtime_share_node_status"` + NodeID int64 `gorm:"column:node_id;not null;index:idx_peer_share_runtime_share_node_status"` + ReservationID string `gorm:"column:reservation_id;type:text;not null;uniqueIndex"` + ResourceKey string `gorm:"column:resource_key;type:text;not null;uniqueIndex"` + BindingID string `gorm:"column:binding_id;type:text;not null;default:'';index:idx_peer_share_runtime_binding_id"` + Role string `gorm:"type:text;not null;default:''"` + ChainName string `gorm:"column:chain_name;type:text;not null;default:''"` + ServiceName string `gorm:"column:service_name;type:text;not null;default:''"` + Protocol string `gorm:"type:text;not null;default:'tls'"` + Strategy string `gorm:"type:text;not null;default:'round'"` + Port int `gorm:"not null;default:0"` + Target string `gorm:"type:text;not null;default:''"` + Applied int `gorm:"not null;default:0"` + Status int `gorm:"not null;default:1;index:idx_peer_share_runtime_share_node_status"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` +} + +func (PeerShareRuntime) TableName() string { return "peer_share_runtime" } + +type FederationTunnelBinding struct { + ID int64 `gorm:"primaryKey;autoIncrement"` + TunnelID int64 `gorm:"column:tunnel_id;not null;uniqueIndex:idx_federation_tunnel_binding_unique;index:idx_federation_tunnel_binding_tunnel"` + NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_federation_tunnel_binding_unique"` + ChainType int `gorm:"column:chain_type;not null;uniqueIndex:idx_federation_tunnel_binding_unique"` + HopInx int `gorm:"column:hop_inx;not null;default:0;uniqueIndex:idx_federation_tunnel_binding_unique"` + RemoteURL string `gorm:"column:remote_url;type:text;not null"` + ResourceKey string `gorm:"column:resource_key;type:text;not null;uniqueIndex"` + RemoteBindingID string `gorm:"column:remote_binding_id;type:text;not null"` + AllocatedPort int `gorm:"column:allocated_port;not null"` + Status int `gorm:"not null;default:1;index:idx_federation_tunnel_binding_tunnel"` + CreatedTime int64 `gorm:"column:created_time;not null"` + UpdatedTime int64 `gorm:"column:updated_time;not null"` +} + +func (FederationTunnelBinding) TableName() string { return "federation_tunnel_binding" } + +// ─── Backup / Import-Export Structs ────────────────────────────────── +// These are not GORM models; they define the JSON wire format for the +// backup/restore API and MUST keep their existing json tags unchanged. + +// BackupData represents the full backup structure. +type BackupData struct { + Version string `json:"version"` + ExportedAt int64 `json:"exportedAt"` + Users []UserBackup `json:"users,omitempty"` + Nodes []NodeBackup `json:"nodes,omitempty"` + Tunnels []TunnelBackup `json:"tunnels,omitempty"` + Forwards []ForwardBackup `json:"forwards,omitempty"` + UserTunnels []UserTunnelBackup `json:"userTunnels,omitempty"` + SpeedLimits []SpeedLimitBackup `json:"speedLimits,omitempty"` + TunnelGroups []TunnelGroupBackup `json:"tunnelGroups,omitempty"` + UserGroups []UserGroupBackup `json:"userGroups,omitempty"` + Permissions []PermissionBackup `json:"permissions,omitempty"` + Configs map[string]string `json:"configs,omitempty"` +} + +type UserBackup struct { + ID int64 `json:"id"` + User string `json:"user"` + Pwd string `json:"pwd"` + RoleID int `json:"roleId"` + ExpTime int64 `json:"expTime"` + Flow int64 `json:"flow"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + FlowResetTime int64 `json:"flowResetTime"` + Num int `json:"num"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` +} + +type NodeBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + Secret string `json:"secret"` + ServerIP string `json:"serverIp"` + ServerIPv4 string `json:"serverIpV4,omitempty"` + ServerIPv6 string `json:"serverIpV6,omitempty"` + Port string `json:"port"` + InterfaceName string `json:"interfaceName,omitempty"` + Version string `json:"version,omitempty"` + HTTP int `json:"http"` + TLS int `json:"tls"` + Socks int `json:"socks"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` + TCPListenAddr string `json:"tcpListenAddr"` + UDPListenAddr string `json:"udpListenAddr"` + Inx int `json:"inx"` + IsRemote int `json:"isRemote"` + RemoteURL string `json:"remoteUrl,omitempty"` + RemoteToken string `json:"remoteToken,omitempty"` + RemoteConfig string `json:"remoteConfig,omitempty"` +} + +type TunnelBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + TrafficRatio float64 `json:"trafficRatio"` + Type int `json:"type"` + Protocol string `json:"protocol"` + Flow int64 `json:"flow"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + InIP string `json:"inIp,omitempty"` + Inx int `json:"inx"` + IPPreference string `json:"ipPreference,omitempty"` + ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"` +} + +type ChainTunnelBackup struct { + ID int64 `json:"id"` + TunnelID int64 `json:"tunnelId"` + ChainType string `json:"chainType"` + NodeID int64 `json:"nodeId"` + Port int `json:"port,omitempty"` + Strategy string `json:"strategy,omitempty"` + Inx int `json:"inx,omitempty"` + Protocol string `json:"protocol,omitempty"` +} + +type ForwardBackup struct { + ID int64 `json:"id"` + UserID int64 `json:"userId"` + UserName string `json:"userName"` + Name string `json:"name"` + TunnelID int64 `json:"tunnelId"` + RemoteAddr string `json:"remoteAddr"` + Strategy string `json:"strategy"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Inx int `json:"inx"` + ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"` +} + +type ForwardPortBackup struct { + NodeID int64 `json:"nodeId"` + Port int `json:"port"` +} + +type UserTunnelBackup struct { + ID int64 `json:"id"` + UserID int64 `json:"userId"` + TunnelID int64 `json:"tunnelId"` + SpeedID int64 `json:"speedId,omitempty"` + Num int `json:"num"` + Flow int64 `json:"flow"` + InFlow int64 `json:"inFlow"` + OutFlow int64 `json:"outFlow"` + FlowResetTime int64 `json:"flowResetTime"` + ExpTime int64 `json:"expTime"` + Status int `json:"status"` +} + +type SpeedLimitBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + Speed int64 `json:"speed"` + TunnelID int64 `json:"tunnelId"` + TunnelName string `json:"tunnelName"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime,omitempty"` + Status int `json:"status"` +} + +type TunnelGroupBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Tunnels []int64 `json:"tunnels,omitempty"` +} + +type UserGroupBackup struct { + ID int64 `json:"id"` + Name string `json:"name"` + CreatedTime int64 `json:"createdTime"` + UpdatedTime int64 `json:"updatedTime"` + Status int `json:"status"` + Users []int64 `json:"users,omitempty"` +} + +type PermissionBackup struct { + ID int64 `json:"id"` + UserGroupID int64 `json:"userGroupId"` + TunnelGroupID int64 `json:"tunnelGroupId"` + CreatedTime int64 `json:"createdTime"` + CreatedByGroup int `json:"createdByGroup"` + Grants []PermissionGrantBackup `json:"grants,omitempty"` +} + +type PermissionGrantBackup struct { + ID int64 `json:"id"` + UserGroupID int64 `json:"userGroupId"` + TunnelGroupID int64 `json:"tunnelGroupId"` + UserTunnelID int64 `json:"userTunnelId"` + CreatedTime int64 `json:"createdTime"` + CreatedByGroup int `json:"createdByGroup"` +} + +// ImportResult contains the result of an import operation. +type ImportResult struct { + UsersImported int `json:"usersImported"` + NodesImported int `json:"nodesImported"` + TunnelsImported int `json:"tunnelsImported"` + ForwardsImported int `json:"forwardsImported"` + UserTunnelsImported int `json:"userTunnelsImported"` + SpeedLimitsImported int `json:"speedLimitsImported"` + TunnelGroupsImported int `json:"tunnelGroupsImported"` + UserGroupsImported int `json:"userGroupsImported"` + PermissionsImported int `json:"permissionsImported"` + ConfigsImported int `json:"configsImported"` + AutoBackup *BackupData `json:"autoBackup,omitempty"` +} + +// ─── View Structs (used by Repository, not GORM models) ───────────── +// These are used for JOIN query results that don't map 1:1 to a table. + +// ForwardRecord is a minimal forward view used by control plane and flow policy. +type ForwardRecord struct { + ID int64 + UserID int64 + UserName string + Name string + TunnelID int64 + RemoteAddr string + Strategy string + Status int +} + +// TunnelRecord is a minimal tunnel view used by control plane. +type TunnelRecord struct { + ID int64 + Type int + Status int + Flow int64 + TrafficRatio float64 +} + +// ForwardPortRecord is a forward port mapping used by control plane. +type ForwardPortRecord struct { + NodeID int64 + Port int +} + +// NodeRecord is a node view used by control plane. +type NodeRecord struct { + ID int64 + Name string + ServerIP string + ServerIPv4 string + ServerIPv6 string + Status int + PortRange string + TCPListenAddr string + UDPListenAddr string + InterfaceName string + IsRemote int + RemoteURL string + RemoteToken string + RemoteConfig string +} + +type ChainNodeRecord struct { + ChainType int + Inx int64 + NodeID int64 + Port int + NodeName string + Protocol string + Strategy string +} + +type UserTunnelLimiterInfo struct { + UserTunnelID int64 + LimiterID *int64 + Speed *int +} + +// UserFlowSnapshot holds a user's current flow counters (used by stats job). +type UserFlowSnapshot struct { + UserID int64 + InFlow int64 + OutFlow int64 +} + +// ExpiredUserTunnel holds minimal info for an expired user_tunnel row. +type ExpiredUserTunnel struct { + ID int64 + UserID int64 + TunnelID int64 +} + +// UserTunnelDetail is a joined view of user_tunnel + tunnel + speed_limit. +type UserTunnelDetail struct { + ID int64 + UserID int64 + TunnelID int64 + TunnelName string + TunnelFlow int + Flow int64 + InFlow int64 + OutFlow int64 + Num int + FlowResetTime int64 + ExpTime int64 + SpeedID sql.NullInt64 + SpeedLimit sql.NullString + Speed sql.NullInt64 +} + +// UserForwardDetail is a joined view of forward + tunnel. +type UserForwardDetail struct { + ID int64 + Name string + TunnelID int64 + TunnelName string + InIP string + InPort sql.NullInt64 + RemoteAddr string + InFlow int64 + OutFlow int64 + Status int + CreatedAt int64 +} diff --git a/go-backend/internal/store/postgres/embed.go b/go-backend/internal/store/postgres/embed.go deleted file mode 100644 index 749852d..0000000 --- a/go-backend/internal/store/postgres/embed.go +++ /dev/null @@ -1,9 +0,0 @@ -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 deleted file mode 100644 index ee3f9dd..0000000 --- a/go-backend/internal/store/postgres/sql/data.sql +++ /dev/null @@ -1,18 +0,0 @@ -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 deleted file mode 100644 index c7deb24..0000000 --- a/go-backend/internal/store/postgres/sql/schema.sql +++ /dev/null @@ -1,250 +0,0 @@ -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, - ip_preference VARCHAR(10) NOT NULL DEFAULT '' -); - -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); - -CREATE TABLE IF NOT EXISTS announcement ( - id SERIAL PRIMARY KEY, - content TEXT NOT NULL, - enabled INTEGER NOT NULL DEFAULT 1, - created_time BIGINT NOT NULL, - updated_time BIGINT -); diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go new file mode 100644 index 0000000..2c6672c --- /dev/null +++ b/go-backend/internal/store/repo/repository.go @@ -0,0 +1,2545 @@ +package repo + +import ( + "database/sql" + "errors" + "fmt" + "log" + "os" + "path/filepath" + "sort" + "strings" + "time" + + gsqlite "github.com/glebarez/sqlite" + "gorm.io/driver/postgres" + "gorm.io/gorm" + "gorm.io/gorm/clause" + "gorm.io/gorm/logger" + + "go-backend/internal/store/model" +) + +// ─── Type aliases for backward compatibility ───────────────────────── +// Handlers still reference repo.User, repo.BackupData, etc. + +type User = model.User +type ViteConfig = model.ViteConfig +type Announcement = model.Announcement +type UserTunnelDetail = model.UserTunnelDetail +type UserForwardDetail = model.UserForwardDetail +type StatisticsFlow = model.StatisticsFlow +type Node = model.Node +type PeerShare = model.PeerShare +type PeerShareRuntime = model.PeerShareRuntime +type FederationTunnelBinding = model.FederationTunnelBinding +type BackupData = model.BackupData +type UserBackup = model.UserBackup +type NodeBackup = model.NodeBackup +type TunnelBackup = model.TunnelBackup +type ChainTunnelBackup = model.ChainTunnelBackup +type ForwardBackup = model.ForwardBackup +type ForwardPortBackup = model.ForwardPortBackup +type UserTunnelBackup = model.UserTunnelBackup +type SpeedLimitBackup = model.SpeedLimitBackup +type TunnelGroupBackup = model.TunnelGroupBackup +type UserGroupBackup = model.UserGroupBackup +type PermissionBackup = model.PermissionBackup +type PermissionGrantBackup = model.PermissionGrantBackup +type ImportResult = model.ImportResult + +// ─── Repository ────────────────────────────────────────────────────── + +type Repository struct { + db *gorm.DB +} + +func (r *Repository) DB() *gorm.DB { + if r == nil { + return nil + } + return r.db +} + +// ─── Open / Close ──────────────────────────────────────────────────── + +func Open(path string) (*Repository, error) { + if err := ensureParentDir(path); err != nil { + return nil, err + } + + dsn := "file:" + path + + "?_pragma=busy_timeout(5000)" + + "&_pragma=journal_mode(WAL)" + + "&_pragma=synchronous(NORMAL)" + + db, err := gorm.Open(gsqlite.Open(dsn), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + return nil, err + } + + sqlDB, err := db.DB() + if err != nil { + return nil, err + } + sqlDB.SetMaxOpenConns(1) + + if err := prepareSQLiteLegacyColumns(db); err != nil { + _ = sqlDB.Close() + return nil, fmt.Errorf("prepare sqlite legacy schema: %w", err) + } + + if err := autoMigrateAll(db); err != nil { + _ = sqlDB.Close() + return nil, fmt.Errorf("auto migrate: %w", err) + } + + seedData(db) + + if err := migrateSchema(db); err != nil { + _ = sqlDB.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") + } + + db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + return nil, err + } + + sqlDB, err := db.DB() + if err != nil { + return nil, err + } + if err := sqlDB.Ping(); err != nil { + _ = sqlDB.Close() + return nil, err + } + + if err := autoMigrateAll(db); err != nil { + _ = sqlDB.Close() + return nil, fmt.Errorf("auto migrate: %w", err) + } + + seedData(db) + + if err := migrateSchema(db); err != nil { + _ = sqlDB.Close() + return nil, err + } + + return &Repository{db: db}, nil +} + +func (r *Repository) Close() error { + if r == nil || r.db == nil { + return nil + } + sqlDB, err := r.db.DB() + if err != nil { + return err + } + return sqlDB.Close() +} + +func autoMigrateAll(db *gorm.DB) error { + models := []interface{}{ + &model.User{}, + &model.Forward{}, + &model.ForwardPort{}, + &model.Node{}, + &model.SpeedLimit{}, + &model.StatisticsFlow{}, + &model.Tunnel{}, + &model.ChainTunnel{}, + &model.UserTunnel{}, + &model.TunnelGroup{}, + &model.UserGroup{}, + &model.TunnelGroupTunnel{}, + &model.UserGroupUser{}, + &model.GroupPermission{}, + &model.GroupPermissionGrant{}, + &model.ViteConfig{}, + &model.PeerShare{}, + &model.PeerShareRuntime{}, + &model.FederationTunnelBinding{}, + &model.Announcement{}, + &model.SchemaVersion{}, + } + + if db.Dialector.Name() != "sqlite" { + return db.AutoMigrate(models...) + } + + m := db.Migrator() + hasNode := m.HasTable(&model.Node{}) + hasTunnel := m.HasTable(&model.Tunnel{}) + + for _, item := range models { + if hasNode { + if _, ok := item.(*model.Node); ok { + continue + } + } + if hasTunnel { + if _, ok := item.(*model.Tunnel); ok { + continue + } + } + if err := db.AutoMigrate(item); err != nil { + return err + } + } + + return nil +} + +func prepareSQLiteLegacyColumns(db *gorm.DB) error { + if db == nil || db.Dialector.Name() != "sqlite" { + return nil + } + m := db.Migrator() + + if m.HasTable(&model.Node{}) { + for _, field := range []string{"ServerIPV4", "ServerIPV6", "Inx", "IsRemote", "RemoteURL", "RemoteToken", "RemoteConfig"} { + if m.HasColumn(&model.Node{}, field) { + continue + } + if err := m.AddColumn(&model.Node{}, field); err != nil { + return fmt.Errorf("add node.%s: %w", field, err) + } + } + } + + if m.HasTable(&model.Tunnel{}) { + for _, field := range []string{"Inx", "IPPreference"} { + if m.HasColumn(&model.Tunnel{}, field) { + continue + } + if err := m.AddColumn(&model.Tunnel{}, field); err != nil { + return fmt.Errorf("add tunnel.%s: %w", field, err) + } + } + } + + return nil +} + +func seedData(db *gorm.DB) { + adminUser := model.User{ + ID: 1, User: "admin_user", Pwd: "3c85cdebade1c51cf64ca9f3c09d182d", + RoleID: 0, ExpTime: 2727251700000, Flow: 99999, InFlow: 0, OutFlow: 0, + FlowResetTime: 1, Num: 99999, CreatedTime: 1748914865000, + UpdatedTime: sql.NullInt64{Int64: 1754011744252, Valid: true}, Status: 1, + } + db.Where("id = ?", 1).FirstOrCreate(&adminUser) + + appNameConfig := model.ViteConfig{ID: 1, Name: "app_name", Value: "flux", Time: 1755147963000} + db.Where("id = ?", 1).FirstOrCreate(&appNameConfig) +} + +// ─── User Queries ──────────────────────────────────────────────────── + +func (r *Repository) GetUserByUsername(username string) (*model.User, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var user model.User + err := r.db.Where(`"user" = ?`, username).First(&user).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &user, nil +} + +func (r *Repository) GetUserByID(id int64) (*model.User, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var user model.User + err := r.db.Where("id = ?", id).First(&user).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &user, nil +} + +func (r *Repository) UsernameExists(username string) (bool, error) { + var count int64 + err := r.db.Model(&model.User{}).Where(`"user" = ?`, username).Count(&count).Error + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) UsernameExistsExceptID(username string, exceptID int64) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.User{}).Where(`"user" = ? AND id != ?`, username, exceptID).Count(&count).Error + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) UpdateUserNameAndPassword(userID int64, username, passwordMD5 string, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.User{}).Where("id = ?", userID).Updates(map[string]interface{}{ + "user": username, + "pwd": passwordMD5, + "updated_time": now, + }).Error +} + +// ─── Config Queries ────────────────────────────────────────────────── + +func (r *Repository) GetConfigByName(name string) (*model.ViteConfig, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var cfg model.ViteConfig + err := r.db.Where("name = ?", name).First(&cfg).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &cfg, nil +} + +func (r *Repository) ListConfigs() (map[string]string, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var configs []model.ViteConfig + if err := r.db.Find(&configs).Error; err != nil { + return nil, err + } + result := make(map[string]string) + for _, c := range configs { + result[c.Name] = c.Value + } + return result, nil +} + +func (r *Repository) UpsertConfig(name, value string, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "name"}}, + DoUpdates: clause.AssignmentColumns([]string{"value", "time"}), + }).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error +} + +// ─── Announcement Queries ──────────────────────────────────────────── + +func (r *Repository) GetAnnouncement() (*model.Announcement, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ann model.Announcement + err := r.db.Order("id DESC").First(&ann).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &ann, nil +} + +func (r *Repository) UpsertAnnouncement(content string, enabled int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + var count int64 + if err := r.db.Model(&model.Announcement{}).Count(&count).Error; err != nil { + return err + } + if count == 0 { + return r.db.Create(&model.Announcement{ + Content: content, Enabled: enabled, + CreatedTime: now, UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + }).Error + } + return r.db.Model(&model.Announcement{}).Where("1=1").Updates(map[string]interface{}{ + "content": content, "enabled": enabled, "updated_time": now, + }).Error +} + +// ─── User Package Queries ──────────────────────────────────────────── + +func (r *Repository) GetUserPackageTunnels(userID int64) ([]model.UserTunnelDetail, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var items []model.UserTunnelDetail + err := r.db.Model(&model.UserTunnel{}). + Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed"). + Joins("LEFT JOIN tunnel ON tunnel.id = user_tunnel.tunnel_id"). + Joins("LEFT JOIN speed_limit ON speed_limit.id = user_tunnel.speed_id"). + Where("user_tunnel.user_id = ?", userID). + Order("user_tunnel.id ASC"). + Find(&items).Error + if err != nil { + return nil, err + } + if items == nil { + items = make([]model.UserTunnelDetail, 0) + } + return items, nil +} + +func (r *Repository) GetUserPackageForwards(userID int64) ([]model.UserForwardDetail, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + type fwdRow struct { + ID int64 + Name string + TunnelID int64 + TunnelName string + RemoteAddr string + InFlow int64 + OutFlow int64 + Status int + CreatedAt int64 + } + + var rows []fwdRow + err := r.db.Model(&model.Forward{}). + Select("forward.id, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, forward.in_flow, forward.out_flow, forward.status, forward.created_time AS created_at"). + Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id"). + Where("forward.user_id = ?", userID). + Order("forward.id ASC"). + Find(&rows).Error + if err != nil { + return nil, err + } + + items := make([]model.UserForwardDetail, 0, len(rows)) + for _, row := range rows { + inIP, inPort, err := resolveForwardIngress(r.db, row.ID, row.TunnelID) + if err != nil { + return nil, err + } + items = append(items, model.UserForwardDetail{ + ID: row.ID, Name: row.Name, TunnelID: row.TunnelID, + TunnelName: row.TunnelName, InIP: inIP, InPort: inPort, + RemoteAddr: row.RemoteAddr, InFlow: row.InFlow, OutFlow: row.OutFlow, + Status: row.Status, CreatedAt: row.CreatedAt, + }) + } + return items, nil +} + +// ─── Statistics Queries ────────────────────────────────────────────── + +func (r *Repository) GetStatisticsFlows(userID int64, limit int) ([]model.StatisticsFlow, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var items []model.StatisticsFlow + err := r.db.Where("user_id = ?", userID).Order("id DESC").Limit(limit).Find(&items).Error + if err != nil { + return nil, err + } + if items == nil { + items = make([]model.StatisticsFlow, 0) + } + return items, nil +} + +// ─── Node Queries ──────────────────────────────────────────────────── + +func (r *Repository) NodeExistsBySecret(secret string) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.Node{}).Where("secret = ?", secret).Count(&count).Error + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) GetNodeBySecret(secret string) (*model.Node, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var n model.Node + err := r.db.Where("secret = ?", secret).First(&n).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &n, nil +} + +func (r *Repository) GetNodeByID(id int64) (*model.Node, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var n model.Node + err := r.db.Where("id = ?", id).First(&n).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &n, nil +} + +func (r *Repository) UpdateNodeOnline(nodeID int64, status int, version string, httpVal, tlsVal, socksVal int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.Node{}).Where("id = ?", nodeID).Updates(map[string]interface{}{ + "status": status, "version": version, "http": httpVal, "tls": tlsVal, + "socks": socksVal, "updated_time": unixMilliNow(), + }).Error +} + +func (r *Repository) UpdateNodeStatus(nodeID int64, status int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.Node{}).Where("id = ?", nodeID).Updates(map[string]interface{}{ + "status": status, "updated_time": unixMilliNow(), + }).Error +} + +// ─── Flow ──────────────────────────────────────────────────────────── + +func (r *Repository) AddFlow(forwardID, userID int64, userTunnelID int64, inFlow, outFlow int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID). + UpdateColumns(map[string]interface{}{ + "in_flow": gorm.Expr("in_flow + ?", inFlow), + "out_flow": gorm.Expr("out_flow + ?", outFlow), + }).Error; err != nil { + return err + } + if err := tx.Model(&model.User{}).Where("id = ?", userID). + UpdateColumns(map[string]interface{}{ + "in_flow": gorm.Expr("in_flow + ?", inFlow), + "out_flow": gorm.Expr("out_flow + ?", outFlow), + }).Error; err != nil { + return err + } + if userTunnelID > 0 { + if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID). + UpdateColumns(map[string]interface{}{ + "in_flow": gorm.Expr("in_flow + ?", inFlow), + "out_flow": gorm.Expr("out_flow + ?", outFlow), + }).Error; err != nil { + return err + } + } + return nil + }) +} + +// ─── List Methods (return map[string]interface{}) ──────────────────── + +func (r *Repository) ListNodes() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var nodes []model.Node + if err := r.db.Order("inx ASC, id ASC").Find(&nodes).Error; err != nil { + return nil, err + } + items := make([]map[string]interface{}, 0, len(nodes)) + for _, n := range nodes { + items = append(items, map[string]interface{}{ + "id": n.ID, "inx": n.Inx, "name": n.Name, + "ip": n.ServerIP, "serverIp": n.ServerIP, + "serverIpV4": nullableString(n.ServerIPV4), + "serverIpV6": nullableString(n.ServerIPV6), + "port": n.Port, + "tcpListenAddr": n.TCPListenAddr, + "udpListenAddr": n.UDPListenAddr, + "version": nullableString(n.Version), + "http": n.HTTP, "tls": n.TLS, "socks": n.Socks, + "status": n.Status, "isRemote": n.IsRemote, + "remoteUrl": nullableString(n.RemoteURL), + "remoteToken": nullableString(n.RemoteToken), + "remoteConfig": nullableString(n.RemoteConfig), + }) + } + return items, nil +} + +func (r *Repository) ListUsers() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var users []model.User + if err := r.db.Where("role_id != ?", 0).Order("id ASC").Find(&users).Error; err != nil { + return nil, err + } + items := make([]map[string]interface{}, 0, len(users)) + for _, u := range users { + items = append(items, map[string]interface{}{ + "id": u.ID, "user": u.User, "name": u.User, + "roleId": u.RoleID, "status": u.Status, + "flow": u.Flow, "num": u.Num, "expTime": u.ExpTime, + "flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime, + "updatedTime": nullableInt64(u.UpdatedTime), + "inFlow": u.InFlow, "outFlow": u.OutFlow, + }) + } + return items, nil +} + +func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var limits []model.SpeedLimit + if err := r.db.Order("id ASC").Find(&limits).Error; err != nil { + return nil, err + } + items := make([]map[string]interface{}, 0, len(limits)) + for _, sl := range limits { + items = append(items, map[string]interface{}{ + "id": sl.ID, "name": sl.Name, "speed": sl.Speed, + "tunnelId": sl.TunnelID, "tunnelName": sl.TunnelName, + "status": sl.Status, "createdTime": sl.CreatedTime, + "updatedTime": nullableInt64(sl.UpdatedTime), + }) + } + return items, nil +} + +func (r *Repository) ListForwards() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + type fwdRow struct { + ID int64 + UserID int64 + UserName string + Name string + TunnelID int64 + TunnelName string + RemoteAddr string + Strategy string + InFlow int64 + OutFlow int64 + CreatedTime int64 + Status int + Inx int + } + + var rows []fwdRow + err := r.db.Model(&model.Forward{}). + Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx"). + Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id"). + Order("forward.inx ASC, forward.id ASC"). + Find(&rows).Error + if err != nil { + return nil, err + } + + items := make([]map[string]interface{}, 0, len(rows)) + for _, row := range rows { + inIP, inPort, err := resolveForwardIngress(r.db, row.ID, row.TunnelID) + if err != nil { + return nil, err + } + items = append(items, map[string]interface{}{ + "id": row.ID, "userId": row.UserID, "userName": row.UserName, + "name": row.Name, "tunnelId": row.TunnelID, "tunnelName": row.TunnelName, + "inIp": nullableForwardIngress(inIP), "inPort": nullableInt64(inPort), + "remoteAddr": row.RemoteAddr, "strategy": row.Strategy, + "inFlow": row.InFlow, "outFlow": row.OutFlow, + "createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx), + }) + } + return items, nil +} + +func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + type row struct { + ID int64 + Name string + } + var rows []row + err := r.db.Model(&model.UserTunnel{}). + Select("tunnel.id, tunnel.name"). + Joins("JOIN tunnel ON tunnel.id = user_tunnel.tunnel_id"). + Where("user_tunnel.user_id = ? AND tunnel.status = 1", userID). + Order("tunnel.inx ASC, tunnel.id ASC"). + Find(&rows).Error + if err != nil { + return nil, err + } + items := make([]map[string]interface{}, 0, len(rows)) + for _, r := range rows { + items = append(items, map[string]interface{}{"id": r.ID, "name": r.Name}) + } + return items, nil +} + +func (r *Repository) ListEnabledTunnelSummaries() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + type row struct { + ID int64 + Name string + } + var rows []row + err := r.db.Model(&model.Tunnel{}).Select("id, name").Where("status = 1").Order("inx ASC, id ASC").Find(&rows).Error + if err != nil { + return nil, err + } + items := make([]map[string]interface{}, 0, len(rows)) + for _, r := range rows { + items = append(items, map[string]interface{}{"id": r.ID, "name": r.Name}) + } + return items, nil +} + +func (r *Repository) ListTunnels() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + var tunnels []model.Tunnel + if err := r.db.Order("inx ASC, id ASC").Find(&tunnels).Error; err != nil { + return nil, err + } + + tunnelMap := make(map[int64]map[string]interface{}) + orderedIDs := make([]int64, 0, len(tunnels)) + + for _, t := range tunnels { + tunnelMap[t.ID] = map[string]interface{}{ + "id": t.ID, "inx": t.Inx, "name": t.Name, + "type": t.Type, "flow": t.Flow, "trafficRatio": t.TrafficRatio, + "status": t.Status, "createdTime": t.CreatedTime, + "inIp": nullableString(t.InIP), + "ipPreference": t.IPPreference, + "inNodeId": make([]map[string]interface{}, 0), + "outNodeId": make([]map[string]interface{}, 0), + "chainNodes": make([][]map[string]interface{}, 0), + } + orderedIDs = append(orderedIDs, t.ID) + } + + // Build node IP map + nodeIPMap := map[int64]string{} + var nodeList []model.Node + if err := r.db.Select("id, server_ip").Find(&nodeList).Error; err == nil { + for _, n := range nodeList { + nodeIPMap[n.ID] = n.ServerIP + } + } + + // Load chain tunnels + var chains []model.ChainTunnel + if err := r.db.Order("tunnel_id ASC, chain_type ASC, inx ASC, id ASC").Find(&chains).Error; err != nil { + return nil, err + } + + chainBucket := map[int64]map[int][]map[string]interface{}{} + inNodeIPs := map[int64][]string{} + + for _, c := range chains { + t, ok := tunnelMap[c.TunnelID] + if !ok { + continue + } + + chainTypeInt := 0 + fmt.Sscanf(c.ChainType, "%d", &chainTypeInt) + + inx := int64(0) + if c.Inx.Valid { + inx = c.Inx.Int64 + } + + nodeObj := map[string]interface{}{ + "nodeId": c.NodeID, + "chainType": chainTypeInt, + "inx": inx, + } + if c.Protocol.Valid { + nodeObj["protocol"] = c.Protocol.String + } + if c.Strategy.Valid { + nodeObj["strategy"] = c.Strategy.String + } + + switch chainTypeInt { + case 1: + t["inNodeId"] = append(t["inNodeId"].([]map[string]interface{}), nodeObj) + if ip, ok := nodeIPMap[c.NodeID]; ok && ip != "" { + inNodeIPs[c.TunnelID] = append(inNodeIPs[c.TunnelID], ip) + } + case 2: + if _, ok := chainBucket[c.TunnelID]; !ok { + chainBucket[c.TunnelID] = map[int][]map[string]interface{}{} + } + chainBucket[c.TunnelID][int(inx)] = append(chainBucket[c.TunnelID][int(inx)], nodeObj) + case 3: + t["outNodeId"] = append(t["outNodeId"].([]map[string]interface{}), nodeObj) + } + } + + for tunnelID, groups := range chainBucket { + t := tunnelMap[tunnelID] + if t == nil { + continue + } + keys := make([]int, 0, len(groups)) + for k := range groups { + keys = append(keys, k) + } + sort.Ints(keys) + ordered := make([][]map[string]interface{}, 0, len(keys)) + for _, k := range keys { + ordered = append(ordered, groups[k]) + } + t["chainNodes"] = ordered + + if s, ok := t["inIp"].(string); !ok || strings.TrimSpace(s) == "" { + if ips := inNodeIPs[tunnelID]; len(ips) > 0 { + t["inIp"] = strings.Join(ips, ",") + } + } + } + + result := make([]map[string]interface{}, 0, len(orderedIDs)) + for _, id := range orderedIDs { + if t, ok := tunnelMap[id]; ok { + result = append(result, t) + } + } + return result, nil +} + +// ─── Group Queries ─────────────────────────────────────────────────── + +func (r *Repository) ListTunnelGroups() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var groups []model.TunnelGroup + if err := r.db.Order("id ASC").Find(&groups).Error; err != nil { + return nil, err + } + result := make([]map[string]interface{}, 0, len(groups)) + for _, g := range groups { + ids, names, err := r.listTunnelGroupMembers(g.ID) + if err != nil { + return nil, err + } + result = append(result, map[string]interface{}{ + "id": g.ID, "name": g.Name, "status": g.Status, + "tunnelIds": ids, "tunnelNames": names, + "createdTime": g.CreatedTime, + }) + } + return result, nil +} + +func (r *Repository) ListUserGroups() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var groups []model.UserGroup + if err := r.db.Order("id ASC").Find(&groups).Error; err != nil { + return nil, err + } + result := make([]map[string]interface{}, 0, len(groups)) + for _, g := range groups { + ids, names, err := r.listUserGroupMembers(g.ID) + if err != nil { + return nil, err + } + result = append(result, map[string]interface{}{ + "id": g.ID, "name": g.Name, "status": g.Status, + "userIds": ids, "userNames": names, + "createdTime": g.CreatedTime, + }) + } + return result, nil +} + +func (r *Repository) ListGroupPermissions() ([]map[string]interface{}, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + + type permRow struct { + ID int64 + UserGroupID int64 + UserGroupName sql.NullString + TunnelGroupID int64 + TunnelGroupName sql.NullString + CreatedTime int64 + } + var rows []permRow + err := r.db.Model(&model.GroupPermission{}). + Select("group_permission.id, group_permission.user_group_id, user_group.name AS user_group_name, group_permission.tunnel_group_id, tunnel_group.name AS tunnel_group_name, group_permission.created_time"). + Joins("LEFT JOIN user_group ON user_group.id = group_permission.user_group_id"). + Joins("LEFT JOIN tunnel_group ON tunnel_group.id = group_permission.tunnel_group_id"). + Order("group_permission.id ASC"). + Find(&rows).Error + if err != nil { + return nil, err + } + result := make([]map[string]interface{}, 0, len(rows)) + for _, r := range rows { + result = append(result, map[string]interface{}{ + "id": r.ID, "userGroupId": r.UserGroupID, + "userGroupName": nullableString(r.UserGroupName), + "tunnelGroupId": r.TunnelGroupID, + "tunnelGroupName": nullableString(r.TunnelGroupName), + "createdTime": r.CreatedTime, + }) + } + return result, nil +} + +func (r *Repository) listTunnelGroupMembers(groupID int64) ([]int64, []string, error) { + type row struct { + ID int64 + Name string + } + var rows []row + err := r.db.Model(&model.TunnelGroupTunnel{}). + Select("tunnel.id, tunnel.name"). + Joins("JOIN tunnel ON tunnel.id = tunnel_group_tunnel.tunnel_id"). + Where("tunnel_group_tunnel.tunnel_group_id = ?", groupID). + Order("tunnel.id ASC"). + Find(&rows).Error + if err != nil { + return nil, nil, err + } + ids := make([]int64, 0, len(rows)) + names := make([]string, 0, len(rows)) + for _, r := range rows { + ids = append(ids, r.ID) + names = append(names, r.Name) + } + return ids, names, nil +} + +func (r *Repository) listUserGroupMembers(groupID int64) ([]int64, []string, error) { + type row struct { + ID int64 + Name string + } + var rows []row + err := r.db.Model(&model.UserGroupUser{}). + Select(`"user".id, "user"."user" AS name`). + Joins(`JOIN "user" ON "user".id = user_group_user.user_id`). + Where("user_group_user.user_group_id = ?", groupID). + Order(`"user".id ASC`). + Find(&rows).Error + if err != nil { + return nil, nil, err + } + ids := make([]int64, 0, len(rows)) + names := make([]string, 0, len(rows)) + for _, r := range rows { + ids = append(ids, r.ID) + names = append(names, r.Name) + } + return ids, names, nil +} + +// ─── PeerShare CRUD ────────────────────────────────────────────────── + +func (r *Repository) CreatePeerShare(share *model.PeerShare) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Create(share).Error +} + +func (r *Repository) UpdatePeerShare(share *model.PeerShare) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.PeerShare{}).Where("id = ?", share.ID).Updates(map[string]interface{}{ + "name": share.Name, "max_bandwidth": share.MaxBandwidth, + "expiry_time": share.ExpiryTime, "port_range_start": share.PortRangeStart, + "port_range_end": share.PortRangeEnd, "is_active": share.IsActive, + "updated_time": share.UpdatedTime, "allowed_domains": share.AllowedDomains, + "allowed_ips": share.AllowedIPs, + }).Error +} + +func (r *Repository) DeletePeerShare(id int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + tx.Where("share_id = ?", id).Delete(&model.PeerShareRuntime{}) + return tx.Where("id = ?", id).Delete(&model.PeerShare{}).Error + }) +} + +func (r *Repository) GetPeerShare(id int64) (*model.PeerShare, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var s model.PeerShare + err := r.db.Where("id = ?", id).First(&s).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &s, nil +} + +func (r *Repository) GetPeerShareByToken(token string) (*model.PeerShare, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var s model.PeerShare + err := r.db.Where("token = ?", token).First(&s).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &s, nil +} + +func (r *Repository) ListPeerShares() ([]model.PeerShare, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var shares []model.PeerShare + err := r.db.Order("id DESC").Find(&shares).Error + return shares, err +} + +func (r *Repository) AddPeerShareCurrentFlow(shareID int64, delta int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if shareID <= 0 || delta <= 0 { + return nil + } + return r.db.Model(&model.PeerShare{}).Where("id = ?", shareID). + UpdateColumns(map[string]interface{}{ + "current_flow": gorm.Expr("current_flow + ?", delta), + "updated_time": unixMilliNow(), + }).Error +} + +func (r *Repository) ResetPeerShareCurrentFlow(shareID int64, updatedTime int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if shareID <= 0 { + return nil + } + if updatedTime <= 0 { + updatedTime = unixMilliNow() + } + return r.db.Model(&model.PeerShare{}).Where("id = ?", shareID).Updates(map[string]interface{}{ + "current_flow": 0, "updated_time": updatedTime, + }).Error +} + +// ─── PeerShareRuntime CRUD ─────────────────────────────────────────── + +func (r *Repository) CreatePeerShareRuntime(item *model.PeerShareRuntime) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if item == nil { + return errors.New("runtime item is nil") + } + return r.db.Create(item).Error +} + +func (r *Repository) UpdatePeerShareRuntime(item *model.PeerShareRuntime) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if item == nil { + return errors.New("runtime item is nil") + } + return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", item.ID).Updates(map[string]interface{}{ + "binding_id": item.BindingID, "role": item.Role, + "chain_name": item.ChainName, "service_name": item.ServiceName, + "protocol": item.Protocol, "strategy": item.Strategy, + "port": item.Port, "target": item.Target, + "applied": item.Applied, "status": item.Status, + "updated_time": item.UpdatedTime, + }).Error +} + +func (r *Repository) MarkPeerShareRuntimeReleased(id int64, updatedTime int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.PeerShareRuntime{}).Where("id = ?", id).Updates(map[string]interface{}{ + "status": 0, "updated_time": updatedTime, + }).Error +} + +func (r *Repository) GetPeerShareRuntimeByResourceKey(shareID int64, resourceKey string) (*model.PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var item model.PeerShareRuntime + err := r.db.Where("share_id = ? AND resource_key = ?", shareID, resourceKey).First(&item).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &item, nil +} + +func (r *Repository) GetPeerShareRuntimeByReservationID(shareID int64, reservationID string) (*model.PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var item model.PeerShareRuntime + err := r.db.Where("share_id = ? AND reservation_id = ?", shareID, reservationID).First(&item).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &item, nil +} + +func (r *Repository) GetPeerShareRuntimeByBindingID(shareID int64, bindingID string) (*model.PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var item model.PeerShareRuntime + err := r.db.Where("share_id = ? AND binding_id = ?", shareID, bindingID).First(&item).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &item, nil +} + +func (r *Repository) GetPeerShareRuntimeByID(id int64) (*model.PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var item model.PeerShareRuntime + err := r.db.Where("id = ?", id).First(&item).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &item, nil +} + +func (r *Repository) ListActivePeerShareRuntimesByShareID(shareID int64) ([]model.PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var out []model.PeerShareRuntime + err := r.db.Where("share_id = ? AND status = 1", shareID).Order("port ASC, id ASC").Find(&out).Error + if err != nil { + return nil, err + } + if out == nil { + out = make([]model.PeerShareRuntime, 0) + } + return out, nil +} + +func (r *Repository) ListActivePeerShareRuntimePorts(shareID int64, nodeID int64) ([]int, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ports []int + err := r.db.Model(&model.PeerShareRuntime{}). + Where("share_id = ? AND node_id = ? AND status = 1 AND port > 0", shareID, nodeID). + Pluck("port", &ports).Error + if err != nil { + return nil, err + } + if ports == nil { + ports = make([]int, 0) + } + return ports, nil +} + +// ─── FederationTunnelBinding ───────────────────────────────────────── + +func (r *Repository) UpsertFederationTunnelBinding(item *model.FederationTunnelBinding) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if item == nil { + return errors.New("binding item is nil") + } + return r.db.Clauses(clause.OnConflict{ + Columns: []clause.Column{ + {Name: "tunnel_id"}, {Name: "node_id"}, {Name: "chain_type"}, {Name: "hop_inx"}, + }, + DoUpdates: clause.AssignmentColumns([]string{ + "remote_url", "resource_key", "remote_binding_id", + "allocated_port", "status", "updated_time", + }), + }).Create(item).Error +} + +func (r *Repository) ListActiveFederationTunnelBindingsByTunnel(tunnelID int64) ([]model.FederationTunnelBinding, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var out []model.FederationTunnelBinding + err := r.db.Where("tunnel_id = ? AND status = 1", tunnelID). + Order("chain_type ASC, hop_inx ASC, id ASC").Find(&out).Error + if err != nil { + return nil, err + } + if out == nil { + out = make([]model.FederationTunnelBinding, 0) + } + return out, nil +} + +func (r *Repository) DeleteFederationTunnelBindingsByTunnel(tunnelID int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Where("tunnel_id = ?", tunnelID).Delete(&model.FederationTunnelBinding{}).Error +} + +// ─── Export Methods ────────────────────────────────────────────────── + +func (r *Repository) ExportAll() (*model.BackupData, error) { + backup := &model.BackupData{Version: "1.0", ExportedAt: unixMilliNow()} + + users, err := r.exportUsers() + if err != nil { + return nil, fmt.Errorf("export users failed: %w", err) + } + backup.Users = users + + nodes, err := r.exportNodes() + if err != nil { + return nil, fmt.Errorf("export nodes failed: %w", err) + } + backup.Nodes = nodes + + tunnels, err := r.exportTunnels() + if err != nil { + return nil, fmt.Errorf("export tunnels failed: %w", err) + } + backup.Tunnels = tunnels + + forwards, err := r.exportForwards() + if err != nil { + return nil, fmt.Errorf("export forwards failed: %w", err) + } + backup.Forwards = forwards + + userTunnels, err := r.exportUserTunnels() + if err != nil { + return nil, fmt.Errorf("export user tunnels failed: %w", err) + } + backup.UserTunnels = userTunnels + + speedLimits, err := r.exportSpeedLimits() + if err != nil { + return nil, fmt.Errorf("export speed limits failed: %w", err) + } + backup.SpeedLimits = speedLimits + + tunnelGroups, err := r.exportTunnelGroups() + if err != nil { + return nil, fmt.Errorf("export tunnel groups failed: %w", err) + } + backup.TunnelGroups = tunnelGroups + + userGroups, err := r.exportUserGroups() + if err != nil { + return nil, fmt.Errorf("export user groups failed: %w", err) + } + backup.UserGroups = userGroups + + permissions, err := r.exportPermissions() + if err != nil { + return nil, fmt.Errorf("export permissions failed: %w", err) + } + backup.Permissions = permissions + + configs, err := r.ListConfigs() + if err != nil { + return nil, fmt.Errorf("export configs failed: %w", err) + } + backup.Configs = configs + + return backup, nil +} + +func (r *Repository) ExportPartial(types []string) (*model.BackupData, error) { + backup := &model.BackupData{Version: "1.0", ExportedAt: unixMilliNow()} + typeSet := make(map[string]bool) + for _, t := range types { + typeSet[t] = true + } + + if typeSet["users"] { + v, err := r.exportUsers() + if err != nil { + return nil, fmt.Errorf("export users failed: %w", err) + } + backup.Users = v + } + if typeSet["nodes"] { + v, err := r.exportNodes() + if err != nil { + return nil, fmt.Errorf("export nodes failed: %w", err) + } + backup.Nodes = v + } + if typeSet["tunnels"] { + v, err := r.exportTunnels() + if err != nil { + return nil, fmt.Errorf("export tunnels failed: %w", err) + } + backup.Tunnels = v + } + if typeSet["forwards"] { + v, err := r.exportForwards() + if err != nil { + return nil, fmt.Errorf("export forwards failed: %w", err) + } + backup.Forwards = v + } + if typeSet["userTunnels"] { + v, err := r.exportUserTunnels() + if err != nil { + return nil, fmt.Errorf("export user tunnels failed: %w", err) + } + backup.UserTunnels = v + } + if typeSet["speedLimits"] { + v, err := r.exportSpeedLimits() + if err != nil { + return nil, fmt.Errorf("export speed limits failed: %w", err) + } + backup.SpeedLimits = v + } + if typeSet["tunnelGroups"] { + v, err := r.exportTunnelGroups() + if err != nil { + return nil, fmt.Errorf("export tunnel groups failed: %w", err) + } + backup.TunnelGroups = v + } + if typeSet["userGroups"] { + v, err := r.exportUserGroups() + if err != nil { + return nil, fmt.Errorf("export user groups failed: %w", err) + } + backup.UserGroups = v + } + if typeSet["permissions"] { + v, err := r.exportPermissions() + if err != nil { + return nil, fmt.Errorf("export permissions failed: %w", err) + } + backup.Permissions = v + } + if typeSet["configs"] { + v, err := r.ListConfigs() + if err != nil { + return nil, fmt.Errorf("export configs failed: %w", err) + } + backup.Configs = v + } + return backup, nil +} + +func (r *Repository) exportUsers() ([]model.UserBackup, error) { + var users []model.User + if err := r.db.Order("id ASC").Find(&users).Error; err != nil { + return nil, err + } + out := make([]model.UserBackup, 0, len(users)) + for _, u := range users { + b := model.UserBackup{ + ID: u.ID, User: u.User, Pwd: u.Pwd, RoleID: u.RoleID, + ExpTime: u.ExpTime, Flow: u.Flow, InFlow: u.InFlow, OutFlow: u.OutFlow, + FlowResetTime: u.FlowResetTime, Num: u.Num, + CreatedTime: u.CreatedTime, Status: u.Status, + } + if u.UpdatedTime.Valid { + b.UpdatedTime = u.UpdatedTime.Int64 + } + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportNodes() ([]model.NodeBackup, error) { + var nodes []model.Node + if err := r.db.Order("inx ASC, id ASC").Find(&nodes).Error; err != nil { + return nil, err + } + out := make([]model.NodeBackup, 0, len(nodes)) + for _, n := range nodes { + b := model.NodeBackup{ + ID: n.ID, Name: n.Name, Secret: n.Secret, ServerIP: n.ServerIP, + Port: n.Port, HTTP: n.HTTP, TLS: n.TLS, Socks: n.Socks, + CreatedTime: n.CreatedTime, Status: n.Status, + TCPListenAddr: n.TCPListenAddr, UDPListenAddr: n.UDPListenAddr, + Inx: n.Inx, IsRemote: n.IsRemote, + } + if n.UpdatedTime.Valid { + b.UpdatedTime = n.UpdatedTime.Int64 + } + if n.ServerIPV4.Valid { + b.ServerIPv4 = n.ServerIPV4.String + } + if n.ServerIPV6.Valid { + b.ServerIPv6 = n.ServerIPV6.String + } + if n.InterfaceName.Valid { + b.InterfaceName = n.InterfaceName.String + } + if n.Version.Valid { + b.Version = n.Version.String + } + if n.RemoteURL.Valid { + b.RemoteURL = n.RemoteURL.String + } + if n.RemoteToken.Valid { + b.RemoteToken = n.RemoteToken.String + } + if n.RemoteConfig.Valid { + b.RemoteConfig = n.RemoteConfig.String + } + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportTunnels() ([]model.TunnelBackup, error) { + var tunnels []model.Tunnel + if err := r.db.Order("inx ASC, id ASC").Find(&tunnels).Error; err != nil { + return nil, err + } + out := make([]model.TunnelBackup, 0, len(tunnels)) + for _, t := range tunnels { + b := model.TunnelBackup{ + ID: t.ID, Name: t.Name, TrafficRatio: t.TrafficRatio, + Type: t.Type, Protocol: t.Protocol, Flow: t.Flow, + CreatedTime: t.CreatedTime, UpdatedTime: t.UpdatedTime, + Status: t.Status, Inx: t.Inx, IPPreference: t.IPPreference, + } + if t.InIP.Valid { + b.InIP = t.InIP.String + } + chains, err := r.exportChainTunnels(t.ID) + if err != nil { + return nil, err + } + b.ChainTunnels = chains + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportChainTunnels(tunnelID int64) ([]model.ChainTunnelBackup, error) { + var chains []model.ChainTunnel + if err := r.db.Where("tunnel_id = ?", tunnelID).Order("inx ASC, id ASC").Find(&chains).Error; err != nil { + return nil, err + } + out := make([]model.ChainTunnelBackup, 0, len(chains)) + for _, c := range chains { + b := model.ChainTunnelBackup{ + ID: c.ID, TunnelID: c.TunnelID, ChainType: c.ChainType, NodeID: c.NodeID, + } + if c.Port.Valid { + b.Port = int(c.Port.Int64) + } + if c.Strategy.Valid { + b.Strategy = c.Strategy.String + } + if c.Inx.Valid { + b.Inx = int(c.Inx.Int64) + } + if c.Protocol.Valid { + b.Protocol = c.Protocol.String + } + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportForwards() ([]model.ForwardBackup, error) { + var forwards []model.Forward + if err := r.db.Order("id ASC").Find(&forwards).Error; err != nil { + return nil, err + } + out := make([]model.ForwardBackup, 0, len(forwards)) + for _, f := range forwards { + b := model.ForwardBackup{ + ID: f.ID, UserID: f.UserID, UserName: f.UserName, Name: f.Name, + TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy, + InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime, + UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx, + } + ports, err := r.exportForwardPorts(f.ID) + if err != nil { + return nil, err + } + portsCopy := append([]model.ForwardPortBackup(nil), ports...) + b.ForwardPorts = &portsCopy + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportForwardPorts(forwardID int64) ([]model.ForwardPortBackup, error) { + var fps []model.ForwardPort + if err := r.db.Where("forward_id = ?", forwardID).Order("id ASC").Find(&fps).Error; err != nil { + return nil, err + } + out := make([]model.ForwardPortBackup, 0, len(fps)) + for _, fp := range fps { + out = append(out, model.ForwardPortBackup{NodeID: fp.NodeID, Port: fp.Port}) + } + return out, nil +} + +func (r *Repository) exportUserTunnels() ([]model.UserTunnelBackup, error) { + var uts []model.UserTunnel + if err := r.db.Order("id ASC").Find(&uts).Error; err != nil { + return nil, err + } + out := make([]model.UserTunnelBackup, 0, len(uts)) + for _, ut := range uts { + b := model.UserTunnelBackup{ + ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID, + Num: ut.Num, Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow, + FlowResetTime: ut.FlowResetTime, ExpTime: ut.ExpTime, Status: ut.Status, + } + if ut.SpeedID.Valid { + b.SpeedID = ut.SpeedID.Int64 + } + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportSpeedLimits() ([]model.SpeedLimitBackup, error) { + var sls []model.SpeedLimit + if err := r.db.Order("id ASC").Find(&sls).Error; err != nil { + return nil, err + } + out := make([]model.SpeedLimitBackup, 0, len(sls)) + for _, sl := range sls { + b := model.SpeedLimitBackup{ + ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed), + TunnelID: sl.TunnelID, TunnelName: sl.TunnelName, + CreatedTime: sl.CreatedTime, Status: sl.Status, + } + if sl.UpdatedTime.Valid { + b.UpdatedTime = sl.UpdatedTime.Int64 + } + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportTunnelGroups() ([]model.TunnelGroupBackup, error) { + var groups []model.TunnelGroup + if err := r.db.Order("id ASC").Find(&groups).Error; err != nil { + return nil, err + } + out := make([]model.TunnelGroupBackup, 0, len(groups)) + for _, tg := range groups { + b := model.TunnelGroupBackup{ + ID: tg.ID, Name: tg.Name, CreatedTime: tg.CreatedTime, + UpdatedTime: tg.UpdatedTime, Status: tg.Status, + } + var tunnelIDs []int64 + r.db.Model(&model.TunnelGroupTunnel{}).Where("tunnel_group_id = ?", tg.ID).Pluck("tunnel_id", &tunnelIDs) + b.Tunnels = tunnelIDs + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportUserGroups() ([]model.UserGroupBackup, error) { + var groups []model.UserGroup + if err := r.db.Order("id ASC").Find(&groups).Error; err != nil { + return nil, err + } + out := make([]model.UserGroupBackup, 0, len(groups)) + for _, ug := range groups { + b := model.UserGroupBackup{ + ID: ug.ID, Name: ug.Name, CreatedTime: ug.CreatedTime, + UpdatedTime: ug.UpdatedTime, Status: ug.Status, + } + var userIDs []int64 + r.db.Model(&model.UserGroupUser{}).Where("user_group_id = ?", ug.ID).Pluck("user_id", &userIDs) + b.Users = userIDs + out = append(out, b) + } + return out, nil +} + +func (r *Repository) exportPermissions() ([]model.PermissionBackup, error) { + var perms []model.GroupPermission + if err := r.db.Order("id ASC").Find(&perms).Error; err != nil { + return nil, err + } + out := make([]model.PermissionBackup, 0, len(perms)) + for _, p := range perms { + b := model.PermissionBackup{ + ID: p.ID, UserGroupID: p.UserGroupID, TunnelGroupID: p.TunnelGroupID, + CreatedTime: p.CreatedTime, + } + var grants []model.GroupPermissionGrant + r.db.Where("user_group_id = ? AND tunnel_group_id = ?", p.UserGroupID, p.TunnelGroupID).Find(&grants) + for _, g := range grants { + b.Grants = append(b.Grants, model.PermissionGrantBackup{ + ID: g.ID, UserGroupID: g.UserGroupID, TunnelGroupID: g.TunnelGroupID, + UserTunnelID: g.UserTunnelID, CreatedTime: g.CreatedTime, + CreatedByGroup: g.CreatedByGroup, + }) + } + out = append(out, b) + } + return out, nil +} + +// ─── Import Methods ────────────────────────────────────────────────── + +func (r *Repository) Import(backup *model.BackupData, types []string) (*model.ImportResult, error) { + result := &model.ImportResult{} + typeSet := make(map[string]bool) + for _, t := range types { + typeSet[t] = true + } + + err := r.db.Transaction(func(tx *gorm.DB) error { + now := unixMilliNow() + + if typeSet["users"] && len(backup.Users) > 0 { + count, err := importUsers(tx, backup.Users, now) + if err != nil { + return fmt.Errorf("import users failed: %w", err) + } + result.UsersImported = count + } + if typeSet["nodes"] && len(backup.Nodes) > 0 { + count, err := importNodes(tx, backup.Nodes, now) + if err != nil { + return fmt.Errorf("import nodes failed: %w", err) + } + result.NodesImported = count + } + if typeSet["tunnels"] && len(backup.Tunnels) > 0 { + count, err := importTunnels(tx, backup.Tunnels, now) + if err != nil { + return fmt.Errorf("import tunnels failed: %w", err) + } + result.TunnelsImported = count + } + if typeSet["forwards"] && len(backup.Forwards) > 0 { + count, err := importForwards(tx, backup.Forwards, now) + if err != nil { + return fmt.Errorf("import forwards failed: %w", err) + } + result.ForwardsImported = count + } + if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 { + count, err := importUserTunnels(tx, backup.UserTunnels, now) + if err != nil { + return fmt.Errorf("import user tunnels failed: %w", err) + } + result.UserTunnelsImported = count + } + if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 { + count, err := importSpeedLimits(tx, backup.SpeedLimits, now) + if err != nil { + return fmt.Errorf("import speed limits failed: %w", err) + } + result.SpeedLimitsImported = count + } + if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 { + count, err := importTunnelGroups(tx, backup.TunnelGroups, now) + if err != nil { + return fmt.Errorf("import tunnel groups failed: %w", err) + } + result.TunnelGroupsImported = count + } + if typeSet["userGroups"] && len(backup.UserGroups) > 0 { + count, err := importUserGroups(tx, backup.UserGroups, now) + if err != nil { + return fmt.Errorf("import user groups failed: %w", err) + } + result.UserGroupsImported = count + } + if typeSet["permissions"] && len(backup.Permissions) > 0 { + count, err := importPermissions(tx, backup.Permissions, now) + if err != nil { + return fmt.Errorf("import permissions failed: %w", err) + } + result.PermissionsImported = count + } + if typeSet["configs"] && len(backup.Configs) > 0 { + count, err := importConfigs(tx, backup.Configs, now) + if err != nil { + return fmt.Errorf("import configs failed: %w", err) + } + result.ConfigsImported = count + } + return nil + }) + if err != nil { + return nil, err + } + return result, nil +} + +func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error) { + count := 0 + for _, u := range users { + item := model.User{ + ID: u.ID, + User: u.User, + Pwd: u.Pwd, + RoleID: u.RoleID, + ExpTime: u.ExpTime, + Flow: u.Flow, + InFlow: u.InFlow, + OutFlow: u.OutFlow, + FlowResetTime: u.FlowResetTime, + Num: u.Num, + CreatedTime: u.CreatedTime, + UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + Status: u.Status, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "user", "pwd", "role_id", "exp_time", "flow", "in_flow", "out_flow", + "flow_reset_time", "num", "updated_time", "status", + }), + }).Create(&item).Error + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func importNodes(tx *gorm.DB, nodes []model.NodeBackup, now int64) (int, error) { + count := 0 + for _, n := range nodes { + item := model.Node{ + ID: n.ID, + Name: n.Name, + Secret: n.Secret, + ServerIP: n.ServerIP, + ServerIPV4: sql.NullString{String: n.ServerIPv4, Valid: true}, + ServerIPV6: sql.NullString{String: n.ServerIPv6, Valid: true}, + Port: n.Port, + InterfaceName: sql.NullString{String: n.InterfaceName, Valid: true}, + Version: sql.NullString{String: n.Version, Valid: true}, + HTTP: n.HTTP, + TLS: n.TLS, + Socks: n.Socks, + CreatedTime: n.CreatedTime, + UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + Status: n.Status, + TCPListenAddr: n.TCPListenAddr, + UDPListenAddr: n.UDPListenAddr, + Inx: n.Inx, + IsRemote: n.IsRemote, + RemoteURL: sql.NullString{String: n.RemoteURL, Valid: true}, + RemoteToken: sql.NullString{String: n.RemoteToken, Valid: true}, + RemoteConfig: sql.NullString{String: n.RemoteConfig, Valid: true}, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "name", "secret", "server_ip", "server_ip_v4", "server_ip_v6", "port", "interface_name", "version", + "http", "tls", "socks", "updated_time", "status", "tcp_listen_addr", "udp_listen_addr", + "inx", "is_remote", "remote_url", "remote_token", "remote_config", + }), + }).Create(&item).Error + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, error) { + count := 0 + for _, t := range tunnels { + item := model.Tunnel{ + ID: t.ID, + Name: t.Name, + TrafficRatio: t.TrafficRatio, + Type: t.Type, + Protocol: t.Protocol, + Flow: t.Flow, + CreatedTime: t.CreatedTime, + UpdatedTime: now, + Status: t.Status, + InIP: sql.NullString{String: t.InIP, Valid: true}, + Inx: t.Inx, + IPPreference: t.IPPreference, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "name", "traffic_ratio", "type", "protocol", "flow", "updated_time", "status", "in_ip", "inx", "ip_preference", + }), + }).Create(&item).Error + if err != nil { + return count, err + } + for _, ct := range t.ChainTunnels { + chainItem := model.ChainTunnel{ + ID: ct.ID, + TunnelID: ct.TunnelID, + ChainType: ct.ChainType, + NodeID: ct.NodeID, + Port: sql.NullInt64{Int64: int64(ct.Port), Valid: true}, + Strategy: sql.NullString{String: ct.Strategy, Valid: true}, + Inx: sql.NullInt64{Int64: int64(ct.Inx), Valid: true}, + Protocol: sql.NullString{String: ct.Protocol, Valid: true}, + } + err = tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "chain_type", "node_id", "port", "strategy", "inx", "protocol", + }), + }).Create(&chainItem).Error + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int, error) { + count := 0 + for _, f := range forwards { + item := model.Forward{ + ID: f.ID, + UserID: f.UserID, + UserName: f.UserName, + Name: f.Name, + TunnelID: f.TunnelID, + RemoteAddr: f.RemoteAddr, + Strategy: f.Strategy, + InFlow: f.InFlow, + OutFlow: f.OutFlow, + CreatedTime: f.CreatedTime, + UpdatedTime: now, + Status: f.Status, + Inx: f.Inx, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy", + "in_flow", "out_flow", "updated_time", "status", "inx", + }), + }).Create(&item).Error + if err != nil { + return count, err + } + if f.ForwardPorts != nil { + if err := tx.Where("forward_id = ?", f.ID).Delete(&model.ForwardPort{}).Error; err != nil { + return count, err + } + for _, fp := range *f.ForwardPorts { + if err := tx.Create(&model.ForwardPort{ForwardID: f.ID, NodeID: fp.NodeID, Port: fp.Port}).Error; err != nil { + return count, err + } + } + } + count++ + } + return count, nil +} + +func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int64) (int, error) { + count := 0 + for _, ut := range userTunnels { + item := model.UserTunnel{ + ID: ut.ID, + UserID: ut.UserID, + TunnelID: ut.TunnelID, + SpeedID: sql.NullInt64{Int64: ut.SpeedID, Valid: ut.SpeedID > 0}, + Num: ut.Num, + Flow: ut.Flow, + InFlow: ut.InFlow, + OutFlow: ut.OutFlow, + FlowResetTime: ut.FlowResetTime, + ExpTime: ut.ExpTime, + Status: ut.Status, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "user_id", "tunnel_id", "speed_id", "num", "flow", "in_flow", "out_flow", + "flow_reset_time", "exp_time", "status", + }), + }).Create(&item).Error + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func importSpeedLimits(tx *gorm.DB, speedLimits []model.SpeedLimitBackup, now int64) (int, error) { + count := 0 + for _, sl := range speedLimits { + item := model.SpeedLimit{ + ID: sl.ID, + Name: sl.Name, + Speed: int(sl.Speed), + TunnelID: sl.TunnelID, + TunnelName: sl.TunnelName, + CreatedTime: sl.CreatedTime, + UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + Status: sl.Status, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "name", "speed", "tunnel_id", "tunnel_name", "updated_time", "status", + }), + }).Create(&item).Error + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +func importTunnelGroups(tx *gorm.DB, tunnelGroups []model.TunnelGroupBackup, now int64) (int, error) { + count := 0 + for _, tg := range tunnelGroups { + item := model.TunnelGroup{ + ID: tg.ID, + Name: tg.Name, + CreatedTime: tg.CreatedTime, + UpdatedTime: now, + Status: tg.Status, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{"name", "updated_time", "status"}), + }).Create(&item).Error + if err != nil { + return count, err + } + if err := tx.Where("tunnel_group_id = ?", tg.ID).Delete(&model.TunnelGroupTunnel{}).Error; err != nil { + return count, err + } + for _, tunnelID := range tg.Tunnels { + if err := tx.Create(&model.TunnelGroupTunnel{TunnelGroupID: tg.ID, TunnelID: tunnelID, CreatedTime: now}).Error; err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func importUserGroups(tx *gorm.DB, userGroups []model.UserGroupBackup, now int64) (int, error) { + count := 0 + for _, ug := range userGroups { + item := model.UserGroup{ + ID: ug.ID, + Name: ug.Name, + CreatedTime: ug.CreatedTime, + UpdatedTime: now, + Status: ug.Status, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{"name", "updated_time", "status"}), + }).Create(&item).Error + if err != nil { + return count, err + } + if err := tx.Where("user_group_id = ?", ug.ID).Delete(&model.UserGroupUser{}).Error; err != nil { + return count, err + } + for _, userID := range ug.Users { + if err := tx.Create(&model.UserGroupUser{UserGroupID: ug.ID, UserID: userID, CreatedTime: now}).Error; err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func importPermissions(tx *gorm.DB, permissions []model.PermissionBackup, _ int64) (int, error) { + count := 0 + for _, p := range permissions { + item := model.GroupPermission{ + ID: p.ID, + UserGroupID: p.UserGroupID, + TunnelGroupID: p.TunnelGroupID, + CreatedTime: p.CreatedTime, + } + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{"user_group_id", "tunnel_group_id"}), + }).Create(&item).Error + if err != nil { + return count, err + } + for _, g := range p.Grants { + grantItem := model.GroupPermissionGrant{ + ID: g.ID, + UserGroupID: g.UserGroupID, + TunnelGroupID: g.TunnelGroupID, + UserTunnelID: g.UserTunnelID, + CreatedTime: g.CreatedTime, + CreatedByGroup: g.CreatedByGroup, + } + err = tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{"user_tunnel_id", "created_by_group"}), + }).Create(&grantItem).Error + if err != nil { + return count, err + } + } + count++ + } + return count, nil +} + +func importConfigs(tx *gorm.DB, configs map[string]string, now int64) (int, error) { + count := 0 + for name, value := range configs { + err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "name"}}, + DoUpdates: clause.AssignmentColumns([]string{"value", "time"}), + }).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error + if err != nil { + return count, err + } + count++ + } + return count, nil +} + +// ─── Jobs Queries (background stats / expiry) ─────────────────────── + +func (r *Repository) PurgeOldStatisticsFlows(cutoffMs int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Where("created_time < ?", cutoffMs).Delete(&model.StatisticsFlow{}).Error +} + +func (r *Repository) ListAllUserFlowSnapshots() ([]model.UserFlowSnapshot, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var users []model.User + err := r.db.Order("id ASC").Find(&users).Error + if err != nil { + return nil, err + } + out := make([]model.UserFlowSnapshot, len(users)) + for i, u := range users { + out[i] = model.UserFlowSnapshot{UserID: u.ID, InFlow: u.InFlow, OutFlow: u.OutFlow} + } + return out, nil +} + +func (r *Repository) GetLastStatisticsFlowTotal(userID int64) (sql.NullInt64, error) { + if r == nil || r.db == nil { + return sql.NullInt64{}, errors.New("repository not initialized") + } + var sf model.StatisticsFlow + err := r.db.Where("user_id = ?", userID).Order("id DESC").First(&sf).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return sql.NullInt64{}, nil + } + if err != nil { + return sql.NullInt64{}, err + } + return sql.NullInt64{Int64: sf.TotalFlow, Valid: true}, nil +} + +func (r *Repository) CreateStatisticsFlow(userID, flow, totalFlow int64, timeText string, createdTime int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Create(&model.StatisticsFlow{ + UserID: userID, Flow: flow, TotalFlow: totalFlow, + Time: timeText, CreatedTime: createdTime, + }).Error +} + +func (r *Repository) ResetUserMonthlyFlow(day int, lastDay int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + updates := map[string]interface{}{"in_flow": 0, "out_flow": 0} + if day == lastDay { + return r.db.Model(&model.User{}). + Where("flow_reset_time != 0 AND (flow_reset_time = ? OR flow_reset_time > ?)", day, lastDay). + Updates(updates).Error + } + return r.db.Model(&model.User{}). + Where("flow_reset_time != 0 AND flow_reset_time = ?", day). + Updates(updates).Error +} + +func (r *Repository) ResetUserTunnelMonthlyFlow(day int, lastDay int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + updates := map[string]interface{}{"in_flow": 0, "out_flow": 0} + if day == lastDay { + return r.db.Model(&model.UserTunnel{}). + Where("flow_reset_time != 0 AND (flow_reset_time = ? OR flow_reset_time > ?)", day, lastDay). + Updates(updates).Error + } + return r.db.Model(&model.UserTunnel{}). + Where("flow_reset_time != 0 AND flow_reset_time = ?", day). + Updates(updates).Error +} + +func (r *Repository) ListExpiredActiveUserIDs(nowMs int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.User{}). + Where("role_id != 0 AND status = 1 AND exp_time IS NOT NULL AND exp_time < ?", nowMs). + Pluck("id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + +func (r *Repository) DisableUser(userID int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.User{}).Where("id = ?", userID).Update("status", 0).Error +} + +func (r *Repository) ListExpiredActiveUserTunnels(nowMs int64) ([]model.ExpiredUserTunnel, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var uts []model.UserTunnel + err := r.db.Where("status = 1 AND exp_time IS NOT NULL AND exp_time < ?", nowMs).Find(&uts).Error + if err != nil { + return nil, err + } + out := make([]model.ExpiredUserTunnel, len(uts)) + for i, ut := range uts { + out[i] = model.ExpiredUserTunnel{ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID} + } + return out, nil +} + +func (r *Repository) DisableUserTunnel(id int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.UserTunnel{}).Where("id = ?", id).Update("status", 0).Error +} + +func (r *Repository) GetUserTunnelByID(id int64) (*model.UserTunnel, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ut model.UserTunnel + err := r.db.Where("id = ?", id).First(&ut).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &ut, nil +} + +// ─── Migration ─────────────────────────────────────────────────────── + +const currentSchemaVersion = 2 + +var ensurePostgresIDDefaultsFn = ensurePostgresIDDefaults + +func getSchemaVersion(db *gorm.DB) int { + var v model.SchemaVersion + if err := db.First(&v).Error; err != nil { + db.Create(&model.SchemaVersion{Version: 0}) + return 0 + } + return v.Version +} + +func setSchemaVersion(db *gorm.DB, ver int) { + db.Model(&model.SchemaVersion{}).Where("1=1").Update("version", ver) +} + +func migrateSchema(db *gorm.DB) error { + if db == nil { + return errors.New("nil db") + } + + if err := ensurePostgresIDDefaultsFn(db); err != nil { + return err + } + + ver := getSchemaVersion(db) + if ver >= currentSchemaVersion { + return nil + } + + // Normalize strategy columns + normalizeStrategy := func(modelRef interface{}, table, defaultValue string) error { + result := db.Model(modelRef).Where("strategy IS NULL").Update("strategy", defaultValue) + if result.Error != nil { + msg := strings.ToLower(result.Error.Error()) + if strings.Contains(msg, "no such table") || (strings.Contains(msg, "relation") && strings.Contains(msg, "does not exist")) { + return nil + } + return fmt.Errorf("normalize %s.strategy: %w", table, result.Error) + } + return nil + } + + if err := normalizeStrategy(&model.Forward{}, "forward", "fifo"); err != nil { + return err + } + if err := normalizeStrategy(&model.ChainTunnel{}, "chain_tunnel", "round"); err != nil { + return err + } + if err := normalizeStrategy(&model.PeerShareRuntime{}, "peer_share_runtime", "round"); err != nil { + return err + } + + setSchemaVersion(db, currentSchemaVersion) + return nil +} + +func ensurePostgresIDDefaults(db *gorm.DB) error { + if db.Dialector.Name() != "postgres" { + return nil + } + type idRow struct { + TableSchema string + TableName string + } + var rows []idRow + err := db.Table("information_schema.table_constraints AS tc"). + Select("c.table_schema, c.table_name"). + Joins("JOIN information_schema.key_column_usage AS kcu ON tc.constraint_name = kcu.constraint_name AND tc.table_schema = kcu.table_schema"). + Joins("JOIN information_schema.columns AS c ON c.table_schema = kcu.table_schema AND c.table_name = kcu.table_name AND c.column_name = kcu.column_name"). + Where("tc.constraint_type = ?", "PRIMARY KEY"). + Where("kcu.column_name = ?", "id"). + Where("c.data_type IN ?", []string{"integer", "bigint"}). + Where("c.is_identity = ?", "NO"). + Where("c.table_schema = current_schema()"). + Order("c.table_name ASC"). + Scan(&rows).Error + if err != nil { + return fmt.Errorf("discover postgres id columns: %w", err) + } + + for _, r := range rows { + if err := ensurePostgresTableIDDefault(db, r.TableSchema, r.TableName); err != nil { + return fmt.Errorf("repair %s.%s id default: %w", r.TableSchema, r.TableName, err) + } + } + return nil +} + +func ensurePostgresTableIDDefault(db *gorm.DB, schemaName, tableName string) error { + type defaultRow struct { + ColumnDefault sql.NullString `gorm:"column:column_default"` + } + var row defaultRow + err := db.Table("information_schema.columns"). + Select("column_default"). + Where("table_schema = ? AND table_name = ? AND column_name = 'id'", schemaName, tableName). + Limit(1). + Scan(&row).Error + if err != nil { + return err + } + defaultExpr := row.ColumnDefault + + hasNextvalDefault := defaultExpr.Valid && strings.Contains(strings.ToLower(defaultExpr.String), "nextval(") + seqRef := "" + if hasNextvalDefault { + seqRef = extractNextvalRegclass(defaultExpr.String) + } + + if !hasNextvalDefault || seqRef == "" { + seqName := tableName + "_id_seq" + if err := db.Exec(fmt.Sprintf("CREATE SEQUENCE IF NOT EXISTS %s.%s", quoteSQLIdentifier(schemaName), quoteSQLIdentifier(seqName))).Error; err != nil { + return err + } + seqRef = schemaName + "." + seqName + if err := db.Exec(fmt.Sprintf( + "ALTER TABLE %s.%s ALTER COLUMN id SET DEFAULT nextval(%s::regclass)", + quoteSQLIdentifier(schemaName), quoteSQLIdentifier(tableName), quoteSQLLiteral(seqRef), + )).Error; err != nil { + return err + } + if err := db.Exec(fmt.Sprintf( + "ALTER SEQUENCE %s.%s OWNED BY %s.%s.id", + quoteSQLIdentifier(schemaName), quoteSQLIdentifier(seqName), + quoteSQLIdentifier(schemaName), quoteSQLIdentifier(tableName), + )).Error; err != nil { + return err + } + } + + return syncPostgresTableIDSequence(db, schemaName, tableName, seqRef) +} + +func syncPostgresTableIDSequence(db *gorm.DB, schemaName, tableName, seqRef string) error { + type maxRow struct { + MaxID int64 `gorm:"column:max_id"` + } + var row maxRow + qualifiedTable := fmt.Sprintf("%s.%s", quoteSQLIdentifier(schemaName), quoteSQLIdentifier(tableName)) + err := db.Table(qualifiedTable). + Select("COALESCE(MAX(id), 0) AS max_id"). + Scan(&row).Error + if err != nil { + return err + } + maxID := row.MaxID + + setVal := maxID + isCalled := true + if maxID <= 0 { + setVal = 1 + isCalled = false + } + + return db.Exec(`SELECT setval(?::regclass, ?, ?)`, seqRef, setVal, isCalled).Error +} + +func extractNextvalRegclass(defaultExpr string) string { + nextvalIdx := strings.Index(strings.ToLower(defaultExpr), "nextval(") + if nextvalIdx < 0 { + return "" + } + expr := defaultExpr[nextvalIdx:] + firstQuote := strings.Index(expr, "'") + if firstQuote < 0 { + return "" + } + expr = expr[firstQuote+1:] + secondQuote := strings.Index(expr, "'") + if secondQuote < 0 { + return "" + } + return strings.TrimSpace(expr[:secondQuote]) +} + +func quoteSQLIdentifier(ident string) string { + return `"` + strings.ReplaceAll(ident, `"`, `""`) + `"` +} + +func quoteSQLLiteral(value string) string { + return "'" + strings.ReplaceAll(value, "'", "''") + "'" +} + +// ─── Helper Functions ──────────────────────────────────────────────── + +func resolveForwardIngress(db *gorm.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) { + var tunnelInIP sql.NullString + db.Model(&model.Tunnel{}).Select("in_ip").Where("id = ?", tunnelID).Limit(1).Scan(&tunnelInIP) + + type fpRow struct { + Port sql.NullInt64 + ServerIP sql.NullString + } + var fpRows []fpRow + err := db.Model(&model.ForwardPort{}). + Select("forward_port.port, node.server_ip"). + Joins("LEFT JOIN node ON node.id = forward_port.node_id"). + Where("forward_port.forward_id = ?", forwardID). + Order("forward_port.id ASC"). + Find(&fpRows).Error + if err != nil { + return "", sql.NullInt64{}, err + } + + ports := make([]int64, 0) + nodePairs := make([]string, 0) + seenPorts := make(map[int64]struct{}) + seenPairs := make(map[string]struct{}) + + for _, row := range fpRows { + if !row.Port.Valid { + continue + } + if _, ok := seenPorts[row.Port.Int64]; !ok { + seenPorts[row.Port.Int64] = struct{}{} + ports = append(ports, row.Port.Int64) + } + if row.ServerIP.Valid && strings.TrimSpace(row.ServerIP.String) != "" { + pair := fmt.Sprintf("%s:%d", strings.TrimSpace(row.ServerIP.String), row.Port.Int64) + if _, ok := seenPairs[pair]; !ok { + seenPairs[pair] = struct{}{} + nodePairs = append(nodePairs, pair) + } + } + } + + if len(ports) == 0 { + return "", sql.NullInt64{}, nil + } + + inPort := sql.NullInt64{Int64: ports[0], Valid: true} + + entries := make([]string, 0) + if tunnelInIP.Valid && strings.TrimSpace(tunnelInIP.String) != "" { + tunnelIPs := strings.Split(tunnelInIP.String, ",") + seen := make(map[string]struct{}) + for _, ip := range tunnelIPs { + ip = strings.TrimSpace(ip) + if ip == "" { + continue + } + if _, ok := seen[ip]; ok { + continue + } + seen[ip] = struct{}{} + for _, port := range ports { + entries = append(entries, fmt.Sprintf("%s:%d", ip, port)) + } + } + } else { + entries = append(entries, nodePairs...) + } + + return strings.Join(entries, ","), inPort, nil +} + +func nullableString(v sql.NullString) interface{} { + if v.Valid { + return v.String + } + return nil +} + +func nullableForwardIngress(v string) interface{} { + v = strings.TrimSpace(v) + if v == "" { + return nil + } + return v +} + +func nullableInt64(v sql.NullInt64) interface{} { + if v.Valid { + return v.Int64 + } + return nil +} + +func unixMilliNow() int64 { + return time.Now().UnixMilli() +} + +func ensureParentDir(dbPath string) error { + if dbPath == "" { + return fmt.Errorf("empty db path") + } + dir := filepath.Dir(dbPath) + if dir == "" || dir == "." { + return nil + } + return osMkdirAll(dir) +} + +var osMkdirAll = func(path string) error { + return os.MkdirAll(path, 0o755) +} + +// Suppress unused import warning for log +var _ = log.Printf diff --git a/go-backend/internal/store/repo/repository_control.go b/go-backend/internal/store/repo/repository_control.go new file mode 100644 index 0000000..5333180 --- /dev/null +++ b/go-backend/internal/store/repo/repository_control.go @@ -0,0 +1,308 @@ +package repo + +import ( + "database/sql" + "errors" + "fmt" + "strconv" + "strings" + + "gorm.io/gorm" + + "go-backend/internal/store/model" +) + +func (r *Repository) UserTunnelExistsByUserAndTunnel(userID, tunnelID int64) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.UserTunnel{}). + Where("user_id = ? AND tunnel_id = ? AND status = 1", userID, tunnelID). + Count(&count).Error + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) ListForwardsByTunnel(tunnelID int64) ([]model.ForwardRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var forwards []model.Forward + err := r.db.Where("tunnel_id = ?", tunnelID).Order("id ASC").Find(&forwards).Error + if err != nil { + return nil, err + } + rows := make([]model.ForwardRecord, 0, len(forwards)) + for _, f := range forwards { + rows = append(rows, model.ForwardRecord{ + ID: f.ID, + UserID: f.UserID, + UserName: f.UserName, + Name: f.Name, + TunnelID: f.TunnelID, + RemoteAddr: f.RemoteAddr, + Strategy: f.Strategy, + Status: f.Status, + }) + } + for i := range rows { + if strings.TrimSpace(rows[i].Strategy) == "" { + rows[i].Strategy = "fifo" + } + } + return rows, nil +} + +func (r *Repository) ListForwardPorts(forwardID int64) ([]model.ForwardPortRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ports []model.ForwardPort + err := r.db.Where("forward_id = ?", forwardID).Order("id ASC").Find(&ports).Error + if err != nil { + return nil, err + } + rows := make([]model.ForwardPortRecord, 0, len(ports)) + for _, p := range ports { + rows = append(rows, model.ForwardPortRecord{NodeID: p.NodeID, Port: p.Port}) + } + return rows, nil +} + +func (r *Repository) GetTunnelOutProtocol(tunnelID int64) (string, error) { + if r == nil || r.db == nil { + return "", errors.New("repository not initialized") + } + var ct model.ChainTunnel + err := r.db.Select("protocol"). + Where("tunnel_id = ? AND chain_type = ?", tunnelID, "3"). + Order("id ASC"). + Take(&ct).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return "", nil + } + return "", err + } + if ct.Protocol.Valid { + return ct.Protocol.String, nil + } + return "", nil +} + +func (r *Repository) GetNodeRecord(nodeID int64) (*model.NodeRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var n model.Node + err := r.db.Where("id = ?", nodeID).First(&n).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return nodeRecordFromModel(&n), nil +} + +func (r *Repository) GetNodeRecordTx(tx *gorm.DB, nodeID int64) (*model.NodeRecord, error) { + if tx == nil { + return nil, errors.New("database unavailable") + } + var n model.Node + err := tx.Where("id = ?", nodeID).First(&n).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return nodeRecordFromModel(&n), nil +} + +func nodeRecordFromModel(n *model.Node) *model.NodeRecord { + if n == nil { + return nil + } + rec := &model.NodeRecord{ + ID: n.ID, + Name: n.Name, + ServerIP: n.ServerIP, + Status: n.Status, + PortRange: n.Port, + TCPListenAddr: n.TCPListenAddr, UDPListenAddr: n.UDPListenAddr, + IsRemote: n.IsRemote, + } + if n.ServerIPV4.Valid { + rec.ServerIPv4 = strings.TrimSpace(n.ServerIPV4.String) + } + if n.ServerIPV6.Valid { + rec.ServerIPv6 = strings.TrimSpace(n.ServerIPV6.String) + } + if n.InterfaceName.Valid { + rec.InterfaceName = strings.TrimSpace(n.InterfaceName.String) + } + if n.RemoteURL.Valid { + rec.RemoteURL = strings.TrimSpace(n.RemoteURL.String) + } + if n.RemoteToken.Valid { + rec.RemoteToken = strings.TrimSpace(n.RemoteToken.String) + } + if n.RemoteConfig.Valid { + rec.RemoteConfig = strings.TrimSpace(n.RemoteConfig.String) + } + if rec.TCPListenAddr == "" { + rec.TCPListenAddr = "[::]" + } + if rec.UDPListenAddr == "" { + rec.UDPListenAddr = "[::]" + } + if strings.TrimSpace(rec.Name) == "" { + rec.Name = fmt.Sprintf("node_%d", rec.ID) + } + return rec +} + +func (r *Repository) ResolveUserTunnelAndLimiter(userID, tunnelID int64) (*model.UserTunnelLimiterInfo, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + type row struct { + UserTunnelID int64 `gorm:"column:user_tunnel_id"` + LimiterID sql.NullInt64 `gorm:"column:limiter_id"` + Speed sql.NullInt64 `gorm:"column:speed"` + } + var rec row + err := r.db.Model(&model.UserTunnel{}). + Select("user_tunnel.id AS user_tunnel_id, speed_limit.id AS limiter_id, speed_limit.speed AS speed"). + Joins("LEFT JOIN speed_limit ON speed_limit.id = user_tunnel.speed_id"). + Where("user_tunnel.user_id = ? AND user_tunnel.tunnel_id = ?", userID, tunnelID). + Order("user_tunnel.id ASC"). + Limit(1). + Take(&rec).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return &model.UserTunnelLimiterInfo{}, nil + } + return nil, err + } + info := &model.UserTunnelLimiterInfo{UserTunnelID: rec.UserTunnelID} + if rec.LimiterID.Valid && rec.LimiterID.Int64 > 0 { + v := rec.LimiterID.Int64 + info.LimiterID = &v + s := int(rec.Speed.Int64) + info.Speed = &s + } + return info, nil +} + +func (r *Repository) ListUserTunnelIDs(userID, tunnelID int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.UserTunnel{}). + Where("user_id = ? AND tunnel_id = ?", userID, tunnelID). + Order("id ASC").Pluck("id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + +func (r *Repository) ListUserTunnelIDsByUser(userID int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.UserTunnel{}). + Where("user_id = ?", userID). + Order("id ASC").Pluck("id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + +func (r *Repository) GetTunnelName(tunnelID int64) (string, error) { + if r == nil || r.db == nil { + return "", errors.New("repository not initialized") + } + var name string + err := r.db.Model(&model.Tunnel{}).Where("id = ?", tunnelID).Pluck("name", &name).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return "", nil + } + return "", err + } + return name, nil +} + +func (r *Repository) ListChainNodesForTunnel(tunnelID int64) ([]model.ChainNodeRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + type row struct { + ChainType string + Inx sql.NullInt64 + NodeID int64 + Port sql.NullInt64 + Name sql.NullString + Protocol sql.NullString + Strategy sql.NullString + } + var rows []row + err := r.db.Model(&model.ChainTunnel{}). + Select("chain_tunnel.chain_type, chain_tunnel.inx, chain_tunnel.node_id, chain_tunnel.port, node.name, chain_tunnel.protocol, chain_tunnel.strategy"). + Joins("LEFT JOIN node ON node.id = chain_tunnel.node_id"). + Where("chain_tunnel.tunnel_id = ?", tunnelID). + Order("chain_tunnel.chain_type ASC, chain_tunnel.inx ASC, chain_tunnel.id ASC"). + Find(&rows).Error + if err != nil { + return nil, err + } + result := make([]model.ChainNodeRecord, 0, len(rows)) + for _, row := range rows { + chainType := 0 + if v := strings.TrimSpace(row.ChainType); v != "" { + if parsed, parseErr := strconv.Atoi(v); parseErr == nil { + chainType = parsed + } + } + inx := int64(0) + if row.Inx.Valid { + inx = row.Inx.Int64 + } + port := 0 + if row.Port.Valid { + port = int(row.Port.Int64) + } + item := model.ChainNodeRecord{ + ChainType: chainType, + Inx: inx, + NodeID: row.NodeID, + Port: port, + } + if strings.TrimSpace(row.Name.String) == "" { + item.NodeName = fmt.Sprintf("node_%d", row.NodeID) + } else { + item.NodeName = row.Name.String + } + if strings.TrimSpace(row.Protocol.String) == "" { + item.Protocol = "tls" + } else { + item.Protocol = row.Protocol.String + } + if strings.TrimSpace(row.Strategy.String) == "" { + item.Strategy = "round" + } else { + item.Strategy = row.Strategy.String + } + result = append(result, item) + } + return result, nil +} diff --git a/go-backend/internal/store/repo/repository_federation.go b/go-backend/internal/store/repo/repository_federation.go new file mode 100644 index 0000000..c3f34ab --- /dev/null +++ b/go-backend/internal/store/repo/repository_federation.go @@ -0,0 +1,269 @@ +package repo + +import ( + "database/sql" + "errors" + + "go-backend/internal/store/model" + + "gorm.io/gorm" +) + +// RemoteNodeRow holds the columns fetched for a remote node listing. +type RemoteNodeRow struct { + ID int64 + Name string + RemoteURL sql.NullString + RemoteToken sql.NullString + RemoteConfig sql.NullString +} + +// NodeBasicInfo holds name, server_ip, and status for a node. +type NodeBasicInfo struct { + Name string + ServerIP string + Status int +} + +// FederationBindingRow holds the columns for an active federation tunnel binding. +type FederationBindingRow struct { + ID int64 + TunnelID int64 + TunnelName string + ChainType int + HopInx int + AllocatedPort int + ResourceKey string + RemoteBindingID string + UpdatedTime int64 +} + +// ListRemoteNodes returns all nodes with is_remote=1, ordered by id desc. +func (r *Repository) ListRemoteNodes() ([]RemoteNodeRow, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var result []RemoteNodeRow + err := r.db.Model(&model.Node{}). + Select("id, name, remote_url, remote_token, remote_config"). + Where("is_remote = 1"). + Order("id DESC"). + Find(&result).Error + if err != nil { + return nil, err + } + if result == nil { + result = make([]RemoteNodeRow, 0) + } + return result, nil +} + +// UpdateNodeRemoteConfig sets the remote_config JSON for a given node. +func (r *Repository) UpdateNodeRemoteConfig(nodeID int64, configJSON string) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.Node{}).Where("id = ?", nodeID).Update("remote_config", configJSON).Error +} + +// ListActiveBindingsForNode returns active federation tunnel bindings for a node. +func (r *Repository) ListActiveBindingsForNode(nodeID int64) ([]FederationBindingRow, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var result []FederationBindingRow + err := r.db.Model(&model.FederationTunnelBinding{}). + Select("federation_tunnel_binding.id, federation_tunnel_binding.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, federation_tunnel_binding.chain_type, federation_tunnel_binding.hop_inx, federation_tunnel_binding.allocated_port, federation_tunnel_binding.resource_key, federation_tunnel_binding.remote_binding_id, federation_tunnel_binding.updated_time"). + Joins("LEFT JOIN tunnel ON tunnel.id = federation_tunnel_binding.tunnel_id"). + Where("federation_tunnel_binding.node_id = ? AND federation_tunnel_binding.status = 1", nodeID). + Order("federation_tunnel_binding.allocated_port ASC, federation_tunnel_binding.id ASC"). + Find(&result).Error + if err != nil { + return nil, err + } + if result == nil { + result = make([]FederationBindingRow, 0) + } + return result, nil +} + +// GetNodeBasicInfo returns the name, server_ip, and status for a given node. +func (r *Repository) GetNodeBasicInfo(nodeID int64) (*NodeBasicInfo, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var n model.Node + err := r.db.Select("name", "server_ip", "status").Where("id = ?", nodeID).First(&n).Error + if err != nil { + return nil, err + } + return &NodeBasicInfo{Name: n.Name, ServerIP: n.ServerIP, Status: n.Status}, nil +} + +// CreateFederationTunnel creates a tunnel and chain_tunnel entry in a transaction, +// returning the new tunnel ID. +func (r *Repository) CreateFederationTunnel(name string, tunnelType int, protocol string, now int64, nodeID int64, remotePort int) (int64, error) { + if r == nil || r.db == nil { + return 0, errors.New("repository not initialized") + } + tunnel := model.Tunnel{ + Name: name, + Type: tunnelType, + Protocol: protocol, + Flow: 0, + CreatedTime: now, + UpdatedTime: now, + Status: 1, + InIP: sql.NullString{String: "", Valid: false}, + } + err := r.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Create(&tunnel).Error; err != nil { + return err + } + ct := model.ChainTunnel{ + TunnelID: tunnel.ID, + ChainType: "1", + NodeID: nodeID, + Port: sql.NullInt64{Int64: int64(remotePort), Valid: true}, + Strategy: sql.NullString{String: "fifo", Valid: true}, + Inx: sql.NullInt64{Int64: 0, Valid: true}, + Protocol: sql.NullString{String: protocol, Valid: true}, + } + if err := tx.Create(&ct).Error; err != nil { + return err + } + return nil + }) + if err != nil { + return 0, err + } + return tunnel.ID, nil +} + +// ListUsedPortsOnNode returns all ports in use on a given node from chain_tunnel and forward_port tables. +func (r *Repository) ListUsedPortsOnNode(nodeID int64) ([]int, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + used := make(map[int]struct{}) + + var chainPorts []int + err := r.db.Model(&model.ChainTunnel{}). + Where("node_id = ? AND port > 0", nodeID). + Pluck("port", &chainPorts).Error + if err != nil { + return nil, err + } + for _, p := range chainPorts { + if p > 0 { + used[p] = struct{}{} + } + } + + var forwardPorts []int + err = r.db.Model(&model.ForwardPort{}). + Where("node_id = ? AND port > 0", nodeID). + Pluck("port", &forwardPorts).Error + if err != nil { + return nil, err + } + for _, p := range forwardPorts { + if p > 0 { + used[p] = struct{}{} + } + } + + result := make([]int, 0, len(used)) + for p := range used { + result = append(result, p) + } + return result, nil +} + +// ListTunnelIDsByNamePrefix returns all tunnel IDs whose name starts with the given prefix. +func (r *Repository) ListTunnelIDsByNamePrefix(prefix string) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.Tunnel{}). + Where("name LIKE ?", prefix+"%"). + Order("id ASC"). + Pluck("id", &ids).Error + if err != nil { + return nil, err + } + if ids == nil { + ids = make([]int64, 0) + } + return ids, nil +} + +// NextIndex returns COALESCE(MAX(inx), -1) + 1 for the given table. +func (r *Repository) NextIndex(table string) int { + if r == nil || r.db == nil { + return 0 + } + var modelRef interface{} + switch table { + case "node": + modelRef = &model.Node{} + case "tunnel": + modelRef = &model.Tunnel{} + case "forward": + modelRef = &model.Forward{} + default: + return 0 + } + + type inxRow struct { + Inx int + } + var row inxRow + err := r.db.Model(modelRef). + Select("inx"). + Order("inx DESC"). + Limit(1). + Take(&row).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0 + } + if err != nil { + return 0 + } + if row.Inx < 0 { + return 0 + } + return row.Inx + 1 +} + +// CreateRemoteNode inserts a new remote node. +func (r *Repository) CreateRemoteNode(name, secret, serverIP, portRange string, now int64, status int, inx int, remoteURL, remoteToken, remoteConfigJSON string) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + node := model.Node{ + Name: name, + Secret: secret, + ServerIP: serverIP, + ServerIPV4: sql.NullString{}, + ServerIPV6: sql.NullString{}, + Port: portRange, + InterfaceName: sql.NullString{}, + Version: sql.NullString{}, + HTTP: 0, + TLS: 0, + Socks: 0, + CreatedTime: now, + UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + Status: status, + TCPListenAddr: "[::]", + UDPListenAddr: "[::]", + Inx: inx, + IsRemote: 1, + RemoteURL: sql.NullString{String: remoteURL, Valid: remoteURL != ""}, + RemoteToken: sql.NullString{String: remoteToken, Valid: remoteToken != ""}, + RemoteConfig: sql.NullString{String: remoteConfigJSON, Valid: remoteConfigJSON != ""}, + } + return r.db.Create(&node).Error +} diff --git a/go-backend/internal/store/repo/repository_flow.go b/go-backend/internal/store/repo/repository_flow.go new file mode 100644 index 0000000..cacc55c --- /dev/null +++ b/go-backend/internal/store/repo/repository_flow.go @@ -0,0 +1,171 @@ +package repo + +import ( + "errors" + "strings" + + "gorm.io/gorm" + + "go-backend/internal/store/model" +) + +func (r *Repository) UpdateForwardStatus(forwardID int64, status int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.Forward{}).Where("id = ?", forwardID).Updates(map[string]interface{}{ + "status": status, "updated_time": now, + }).Error +} + +func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var forwards []model.Forward + err := r.db.Where("user_id = ? AND status = 1", userID).Order("id ASC").Find(&forwards).Error + if err != nil { + return nil, err + } + rows := make([]model.ForwardRecord, 0, len(forwards)) + for _, f := range forwards { + rows = append(rows, model.ForwardRecord{ + ID: f.ID, + UserID: f.UserID, + UserName: f.UserName, + Name: f.Name, + TunnelID: f.TunnelID, + RemoteAddr: f.RemoteAddr, + Strategy: f.Strategy, + Status: f.Status, + }) + } + for i := range rows { + if strings.TrimSpace(rows[i].Strategy) == "" { + rows[i].Strategy = "fifo" + } + } + return rows, nil +} + +func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]model.ForwardRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var forwards []model.Forward + err := r.db.Where("user_id = ? AND tunnel_id = ? AND status = 1", userID, tunnelID).Order("id ASC").Find(&forwards).Error + if err != nil { + return nil, err + } + rows := make([]model.ForwardRecord, 0, len(forwards)) + for _, f := range forwards { + rows = append(rows, model.ForwardRecord{ + ID: f.ID, + UserID: f.UserID, + UserName: f.UserName, + Name: f.Name, + TunnelID: f.TunnelID, + RemoteAddr: f.RemoteAddr, + Strategy: f.Strategy, + Status: f.Status, + }) + } + for i := range rows { + if strings.TrimSpace(rows[i].Strategy) == "" { + rows[i].Strategy = "fifo" + } + } + return rows, nil +} + +func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var f model.Forward + err := r.db.Where("id = ?", forwardID).First(&f).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, err + } + fr := model.ForwardRecord{ + ID: f.ID, + UserID: f.UserID, + UserName: f.UserName, + Name: f.Name, + TunnelID: f.TunnelID, + RemoteAddr: f.RemoteAddr, + Strategy: f.Strategy, + Status: f.Status, + } + if strings.TrimSpace(fr.Strategy) == "" { + fr.Strategy = "fifo" + } + return &fr, nil +} + +func (r *Repository) GetTunnelRecord(tunnelID int64) (*model.TunnelRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var t model.Tunnel + err := r.db.Where("id = ?", tunnelID).First(&t).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, err + } + tr := model.TunnelRecord{ + ID: t.ID, + Type: t.Type, + Status: t.Status, + Flow: t.Flow, + TrafficRatio: t.TrafficRatio, + } + if tr.Flow <= 0 { + tr.Flow = 1 + } + if tr.TrafficRatio <= 0 { + tr.TrafficRatio = 1 + } + return &tr, nil +} + +func (r *Repository) TunnelExists(tunnelID int64) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.Tunnel{}).Where("id = ?", tunnelID).Count(&count).Error + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) ForwardExists(forwardID int64) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.Forward{}).Where("id = ?", forwardID).Count(&count).Error + if err != nil { + return false, err + } + return count > 0, nil +} + +func (r *Repository) SpeedLimitExists(id int64) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var count int64 + err := r.db.Model(&model.SpeedLimit{}).Where("id = ?", id).Count(&count).Error + if err != nil { + return false, err + } + return count > 0, nil +} diff --git a/go-backend/internal/store/repo/repository_groups.go b/go-backend/internal/store/repo/repository_groups.go new file mode 100644 index 0000000..24683b3 --- /dev/null +++ b/go-backend/internal/store/repo/repository_groups.go @@ -0,0 +1,69 @@ +package repo + +import ( + "errors" + + "go-backend/internal/store/model" +) + +// ─── Semantic Group Queries (replacing QueryInt64List/QueryPairs passthrough) ─ + +// ListUserIDsByUserGroup returns all user IDs belonging to a user group. +func (r *Repository) ListUserIDsByUserGroup(userGroupID int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.UserGroupUser{}). + Where("user_group_id = ?", userGroupID). + Pluck("user_id", &ids).Error + return ids, err +} + +// ListTunnelIDsByTunnelGroup returns all tunnel IDs belonging to a tunnel group. +func (r *Repository) ListTunnelIDsByTunnelGroup(tunnelGroupID int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.TunnelGroupTunnel{}). + Where("tunnel_group_id = ?", tunnelGroupID). + Pluck("tunnel_id", &ids).Error + return ids, err +} + +// ListGroupPermissionPairsByUserGroup returns [userGroupID, tunnelGroupID] pairs +// for all group permissions associated with a user group. +func (r *Repository) ListGroupPermissionPairsByUserGroup(userGroupID int64) ([][2]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var perms []model.GroupPermission + err := r.db.Where("user_group_id = ?", userGroupID).Find(&perms).Error + if err != nil { + return nil, err + } + result := make([][2]int64, len(perms)) + for i, p := range perms { + result[i] = [2]int64{p.UserGroupID, p.TunnelGroupID} + } + return result, err +} + +// ListGroupPermissionPairsByTunnelGroup returns [userGroupID, tunnelGroupID] pairs +// for all group permissions associated with a tunnel group. +func (r *Repository) ListGroupPermissionPairsByTunnelGroup(tunnelGroupID int64) ([][2]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var perms []model.GroupPermission + err := r.db.Where("tunnel_group_id = ?", tunnelGroupID).Find(&perms).Error + if err != nil { + return nil, err + } + result := make([][2]int64, len(perms)) + for i, p := range perms { + result[i] = [2]int64{p.UserGroupID, p.TunnelGroupID} + } + return result, err +} diff --git a/go-backend/internal/store/sqlite/repository_migrate_test.go b/go-backend/internal/store/repo/repository_migrate_test.go similarity index 53% rename from go-backend/internal/store/sqlite/repository_migrate_test.go rename to go-backend/internal/store/repo/repository_migrate_test.go index 5d31913..4cbd612 100644 --- a/go-backend/internal/store/sqlite/repository_migrate_test.go +++ b/go-backend/internal/store/repo/repository_migrate_test.go @@ -1,35 +1,38 @@ -package sqlite +package repo import ( - "database/sql" "errors" "testing" - "go-backend/internal/store" - - _ "modernc.org/sqlite" + gsqlite "github.com/glebarez/sqlite" + "gorm.io/gorm" + "gorm.io/gorm/logger" ) func TestMigrateSchemaRunsPostgresIDRepairEvenAtCurrentVersion(t *testing.T) { - raw, err := sql.Open("sqlite", ":memory:") + db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) if err != nil { t.Fatalf("open sqlite: %v", err) } t.Cleanup(func() { - _ = raw.Close() + sqlDB, _ := db.DB() + if sqlDB != nil { + _ = sqlDB.Close() + } }) - db := store.Wrap(raw, store.DialectPostgres) - if _, err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`); err != nil { + if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil { t.Fatalf("create schema_version: %v", err) } - if _, err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, currentSchemaVersion); err != nil { + if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, currentSchemaVersion).Error; err != nil { t.Fatalf("seed schema_version: %v", err) } called := 0 original := ensurePostgresIDDefaultsFn - ensurePostgresIDDefaultsFn = func(db *store.DB) error { + ensurePostgresIDDefaultsFn = func(db *gorm.DB) error { called++ return nil } @@ -46,25 +49,29 @@ func TestMigrateSchemaRunsPostgresIDRepairEvenAtCurrentVersion(t *testing.T) { } func TestMigrateSchemaReturnsPostgresIDRepairError(t *testing.T) { - raw, err := sql.Open("sqlite", ":memory:") + db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) if err != nil { t.Fatalf("open sqlite: %v", err) } t.Cleanup(func() { - _ = raw.Close() + sqlDB, _ := db.DB() + if sqlDB != nil { + _ = sqlDB.Close() + } }) - db := store.Wrap(raw, store.DialectPostgres) - if _, err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`); err != nil { + if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil { t.Fatalf("create schema_version: %v", err) } - if _, err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, currentSchemaVersion); err != nil { + if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, currentSchemaVersion).Error; err != nil { t.Fatalf("seed schema_version: %v", err) } wantErr := errors.New("repair failed") original := ensurePostgresIDDefaultsFn - ensurePostgresIDDefaultsFn = func(db *store.DB) error { + ensurePostgresIDDefaultsFn = func(db *gorm.DB) error { return wantErr } t.Cleanup(func() { diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go new file mode 100644 index 0000000..21e2817 --- /dev/null +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -0,0 +1,1408 @@ +package repo + +import ( + "database/sql" + "errors" + "sort" + "strconv" + "strings" + "time" + + "go-backend/internal/store/model" + + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +func (r *Repository) UserExists(username string) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var cnt int64 + err := r.db.Model(&model.User{}).Where(`"user" = ?`, username).Count(&cnt).Error + return cnt > 0, err +} + +func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var cnt int64 + err := r.db.Model(&model.User{}). + Where(`"user" = ? AND id != ?`, username, excludeID). + Count(&cnt).Error + return cnt > 0, err +} + +func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + user := model.User{ + User: username, + Pwd: pwdHash, + RoleID: roleID, + ExpTime: expTime, + Flow: flow, + InFlow: 0, + OutFlow: 0, + FlowResetTime: flowResetTime, + Num: num, + CreatedTime: now, + UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + Status: status, + } + return r.db.Create(&user).Error +} + +func (r *Repository) GetUserRoleID(userID int64) (int, error) { + if r == nil || r.db == nil { + return 0, errors.New("repository not initialized") + } + var user model.User + err := r.db.Select("role_id").Where("id = ?", userID).First(&user).Error + if err != nil { + return 0, normalizeNotFoundErr(err) + } + return user.RoleID, nil +} + +func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.User{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "user": username, + "pwd": pwdHash, + "flow": flow, + "num": num, + "exp_time": expTime, + "flow_reset_time": flowResetTime, + "status": status, + "updated_time": sql.NullInt64{Int64: now, Valid: true}, + }).Error +} + +func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.User{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "user": username, + "flow": flow, + "num": num, + "exp_time": expTime, + "flow_reset_time": flowResetTime, + "status": status, + "updated_time": sql.NullInt64{Int64: now, Valid: true}, + }).Error +} + +func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64) { + if r == nil || r.db == nil { + return + } + _ = r.db.Model(&model.UserTunnel{}). + Where("user_id = ?", userID). + Updates(map[string]interface{}{ + "flow": flow, + "num": num, + "exp_time": expTime, + "flow_reset_time": flowResetTime, + }).Error +} + +func (r *Repository) DeleteUserCascade(userID int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + forwardIDs := tx.Model(&model.Forward{}).Select("id").Where("user_id = ?", userID) + if err := tx.Where("forward_id IN (?)", forwardIDs).Delete(&model.ForwardPort{}).Error; err != nil { + return err + } + if err := tx.Where("user_id = ?", userID).Delete(&model.Forward{}).Error; err != nil { + return err + } + userTunnelIDs := tx.Model(&model.UserTunnel{}).Select("id").Where("user_id = ?", userID) + if err := tx.Where("user_tunnel_id IN (?)", userTunnelIDs).Delete(&model.GroupPermissionGrant{}).Error; err != nil { + return err + } + if err := tx.Where("user_id = ?", userID).Delete(&model.UserTunnel{}).Error; err != nil { + return err + } + if err := tx.Where("user_id = ?", userID).Delete(&model.UserGroupUser{}).Error; err != nil { + return err + } + if err := tx.Where("user_id = ?", userID).Delete(&model.StatisticsFlow{}).Error; err != nil { + return err + } + return tx.Where("id = ?", userID).Delete(&model.User{}).Error + }) +} + +func (r *Repository) ResetUserFlowByUser(userID int64, now int64) { + if r == nil || r.db == nil { + return + } + _ = r.db.Model(&model.User{}). + Where("id = ?", userID). + Updates(map[string]interface{}{ + "in_flow": 0, + "out_flow": 0, + "updated_time": sql.NullInt64{Int64: now, Valid: true}, + }).Error + _ = r.db.Model(&model.UserTunnel{}). + Where("user_id = ?", userID). + Updates(map[string]interface{}{"in_flow": 0, "out_flow": 0}).Error +} + +func (r *Repository) ResetUserFlowByUserTunnel(userTunnelID int64) { + if r == nil || r.db == nil { + return + } + _ = r.db.Model(&model.UserTunnel{}). + Where("id = ?", userTunnelID). + Updates(map[string]interface{}{"in_flow": 0, "out_flow": 0}).Error +} + +func (r *Repository) GetUsernameByID(userID int64) string { + if r == nil || r.db == nil { + return "" + } + var user model.User + if err := r.db.Select("user").Where("id = ?", userID).First(&user).Error; err != nil { + return "" + } + return user.User +} + +func (r *Repository) GetUserDefaultsForTunnel(userID int64) (flow int64, num int, expTime int64, flowReset int64, err error) { + if r == nil || r.db == nil { + return 0, 0, 0, 0, errors.New("repository not initialized") + } + var user model.User + err = r.db.Select("flow", "num", "exp_time", "flow_reset_time").Where("id = ?", userID).First(&user).Error + if err != nil { + return 0, 0, 0, 0, normalizeNotFoundErr(err) + } + return user.Flow, user.Num, user.ExpTime, user.FlowResetTime, nil +} + +func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig interface{}) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + node := model.Node{ + Name: name, + Secret: secret, + ServerIP: serverIP, + ServerIPV4: nullStringFromInterface(serverIPV4), + ServerIPV6: nullStringFromInterface(serverIPV6), + Port: stringFromInterface(port), + InterfaceName: nullStringFromInterface(interfaceName), + Version: nullStringFromInterface(version), + HTTP: httpFlag, + TLS: tlsFlag, + Socks: socksFlag, + CreatedTime: now, + UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + Status: status, + TCPListenAddr: tcpAddr, + UDPListenAddr: udpAddr, + Inx: inx, + IsRemote: isRemote, + RemoteURL: nullStringFromInterface(remoteURL), + RemoteToken: nullStringFromInterface(remoteToken), + RemoteConfig: nullStringFromInterface(remoteConfig), + } + return r.db.Create(&node).Error +} + +func (r *Repository) GetNodeStatusFields(nodeID int64) (status, httpFlag, tlsFlag, socksFlag int, err error) { + if r == nil || r.db == nil { + return 0, 0, 0, 0, errors.New("repository not initialized") + } + var node model.Node + err = r.db.Select("status", "http", "tls", "socks").Where("id = ?", nodeID).First(&node).Error + if err != nil { + return 0, 0, 0, 0, normalizeNotFoundErr(err) + } + return node.Status, node.HTTP, node.TLS, node.Socks, nil +} + +func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.Node{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "name": name, + "server_ip": serverIP, + "server_ip_v4": nullStringFromInterface(serverIPV4), + "server_ip_v6": nullStringFromInterface(serverIPV6), + "port": stringFromInterface(port), + "interface_name": nullStringFromInterface(interfaceName), + "http": httpFlag, + "tls": tlsFlag, + "socks": socksFlag, + "tcp_listen_addr": tcpAddr, + "udp_listen_addr": udpAddr, + "updated_time": sql.NullInt64{Int64: now, Valid: true}, + }).Error +} + +func (r *Repository) GetNodeSecret(nodeID int64) (string, error) { + if r == nil || r.db == nil { + return "", errors.New("repository not initialized") + } + var node model.Node + err := r.db.Select("secret").Where("id = ?", nodeID).First(&node).Error + if err != nil { + return "", normalizeNotFoundErr(err) + } + return node.Secret, nil +} + +func (r *Repository) GetViteConfigValue(name string) (string, error) { + if r == nil || r.db == nil { + return "", errors.New("repository not initialized") + } + var cfg model.ViteConfig + err := r.db.Select("value").Where("name = ?", name).First(&cfg).Error + if err != nil { + return "", normalizeNotFoundErr(err) + } + return cfg.Value, nil +} + +func (r *Repository) UpdateNodeOrder(nodeID int64, inx int, now int64) { + if r == nil || r.db == nil { + return + } + _ = r.db.Model(&model.Node{}). + Where("id = ?", nodeID). + Updates(map[string]interface{}{ + "inx": inx, + "updated_time": sql.NullInt64{Int64: now, Valid: true}, + }).Error +} + +func (r *Repository) DeleteNodeCascade(nodeID int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("node_id = ?", nodeID).Delete(&model.ForwardPort{}).Error; err != nil { + return err + } + if err := tx.Where("node_id = ?", nodeID).Delete(&model.ChainTunnel{}).Error; err != nil { + return err + } + if err := tx.Where("node_id = ?", nodeID).Delete(&model.FederationTunnelBinding{}).Error; err != nil { + return err + } + return tx.Where("id = ?", nodeID).Delete(&model.Node{}).Error + }) +} + +func (r *Repository) GetNodeRemoteFields(nodeID int64) (isRemote int, remoteURL, remoteToken sql.NullString, err error) { + if r == nil || r.db == nil { + return 0, sql.NullString{}, sql.NullString{}, errors.New("repository not initialized") + } + return r.GetNodeRemoteFieldsTx(r.db, nodeID) +} + +func (r *Repository) GetNodeRemoteFieldsTx(tx *gorm.DB, nodeID int64) (isRemote int, remoteURL, remoteToken sql.NullString, err error) { + if tx == nil { + return 0, sql.NullString{}, sql.NullString{}, errors.New("database unavailable") + } + var node model.Node + err = tx.Select("is_remote", "remote_url", "remote_token").Where("id = ?", nodeID).First(&node).Error + if err != nil { + return 0, sql.NullString{}, sql.NullString{}, normalizeNotFoundErr(err) + } + return node.IsRemote, node.RemoteURL, node.RemoteToken, nil +} + +func (r *Repository) GetNodePortRange(nodeID int64) (string, error) { + if r == nil || r.db == nil { + return "", errors.New("repository not initialized") + } + var node model.Node + err := r.db.Select("port").Where("id = ?", nodeID).First(&node).Error + if err != nil { + return "", normalizeNotFoundErr(err) + } + return node.Port, nil +} + +func (r *Repository) TunnelNameExists(name string) (bool, error) { + if r == nil || r.db == nil { + return false, errors.New("repository not initialized") + } + var cnt int64 + err := r.db.Model(&model.Tunnel{}).Where("name = ?", name).Count(&cnt).Error + return cnt > 0, err +} + +func (r *Repository) BeginTx() *gorm.DB { + if r == nil || r.db == nil { + return nil + } + return r.db.Begin() +} + +func (r *Repository) UpdateTunnelOrder(tunnelID int64, inx int, now int64) { + if r == nil || r.db == nil { + return + } + _ = r.db.Model(&model.Tunnel{}). + Where("id = ?", tunnelID). + Updates(map[string]interface{}{"inx": inx, "updated_time": now}).Error +} + +func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, typeVal int, flow int64, trafficRatio float64, status int, inIP, ipPreference string, now int64) error { + if tx == nil { + return errors.New("database unavailable") + } + return tx.Model(&model.Tunnel{}). + Where("id = ?", tunnelID). + Updates(map[string]interface{}{ + "name": name, + "type": typeVal, + "flow": flow, + "traffic_ratio": trafficRatio, + "status": status, + "in_ip": nullStringFromInterface(inIP), + "ip_preference": ipPreference, + "updated_time": now, + }).Error +} + +func (r *Repository) DeleteChainTunnelsByTunnelTx(tx *gorm.DB, tunnelID int64) error { + if tx == nil { + return errors.New("database unavailable") + } + return tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error +} + +func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType string, nodeID int64, port sql.NullInt64, strategy string, inx int, protocol string) error { + if tx == nil { + return errors.New("database unavailable") + } + ct := model.ChainTunnel{ + TunnelID: tunnelID, + ChainType: chainType, + NodeID: nodeID, + Port: port, + Strategy: nullStringFromInterface(strategy), + Inx: nullInt64FromInterface(inx), + Protocol: nullStringFromInterface(protocol), + } + return tx.Create(&ct).Error +} + +func (r *Repository) IsRemoteNodeTx(tx *gorm.DB, nodeID int64) (bool, error) { + if tx == nil { + return false, errors.New("database unavailable") + } + if nodeID <= 0 { + return false, errors.New("节点不存在") + } + var node model.Node + err := tx.Select("is_remote").Where("id = ?", nodeID).First(&node).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, errors.New("节点不存在") + } + if err != nil { + return false, err + } + return node.IsRemote == 1, nil +} + +func (r *Repository) PickNodePortTx(tx *gorm.DB, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) { + if tx == nil { + return 0, errors.New("database unavailable") + } + if nodeID <= 0 { + return 0, errors.New("节点不存在") + } + if port, ok := allocated[nodeID]; ok && port > 0 { + return port, nil + } + + var node model.Node + err := tx.Select("port").Where("id = ?", nodeID).First(&node).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0, errors.New("节点不存在") + } + if err != nil { + return 0, err + } + + candidates := parsePortRangeSpec(node.Port) + if len(candidates) == 0 { + return 0, errors.New("节点端口已满,无可用端口") + } + + used := make(map[int]struct{}) + + chainQuery := tx.Model(&model.ChainTunnel{}).Where("node_id = ? AND port > 0", nodeID) + if excludeTunnelID > 0 { + chainQuery = chainQuery.Where("tunnel_id != ?", excludeTunnelID) + } + var chainPorts []int + if err := chainQuery.Pluck("port", &chainPorts).Error; err != nil { + return 0, err + } + for _, p := range chainPorts { + if p > 0 { + used[p] = struct{}{} + } + } + + var forwardPorts []int + if err := tx.Model(&model.ForwardPort{}). + Where("node_id = ? AND port > 0", nodeID). + Pluck("port", &forwardPorts).Error; err != nil { + return 0, err + } + for _, p := range forwardPorts { + if p > 0 { + used[p] = struct{}{} + } + } + + for _, candidate := range candidates { + if candidate <= 0 { + continue + } + if _, ok := used[candidate]; ok { + continue + } + allocated[nodeID] = candidate + return candidate, nil + } + + return 0, errors.New("节点端口已满,无可用端口") +} + +func (r *Repository) GetTunnelIPPreference(tunnelID int64) string { + if r == nil || r.db == nil { + return "" + } + var tunnel model.Tunnel + if err := r.db.Select("ip_preference").Where("id = ?", tunnelID).First(&tunnel).Error; err != nil { + return "" + } + return tunnel.IPPreference +} + +func (r *Repository) DeleteTunnelCascade(tunnelID int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + forwardIDs := tx.Model(&model.Forward{}).Select("id").Where("tunnel_id = ?", tunnelID) + if err := tx.Where("forward_id IN (?)", forwardIDs).Delete(&model.ForwardPort{}).Error; err != nil { + return err + } + if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.Forward{}).Error; err != nil { + return err + } + if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.UserTunnel{}).Error; err != nil { + return err + } + if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.SpeedLimit{}).Error; err != nil { + return err + } + if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error; err != nil { + return err + } + if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.FederationTunnelBinding{}).Error; err != nil { + return err + } + return tx.Where("id = ?", tunnelID).Delete(&model.Tunnel{}).Error + }) +} + +func (r *Repository) GetTunnelNameByID(tunnelID int64) string { + if r == nil || r.db == nil { + return "" + } + var tunnel model.Tunnel + if err := r.db.Select("name").Where("id = ?", tunnelID).First(&tunnel).Error; err != nil { + return "" + } + return tunnel.Name +} + +func (r *Repository) TunnelEntryNodeIDs(tunnelID int64) ([]int64, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var ids []int64 + err := r.db.Model(&model.ChainTunnel{}). + Where("tunnel_id = ? AND chain_type = ?", tunnelID, "1"). + Order("inx ASC, id ASC"). + Pluck("node_id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + +func (r *Repository) DeleteUserTunnel(id int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Where("id = ?", id).Delete(&model.UserTunnel{}).Error +} + +func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.UserTunnel{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "flow": flow, + "num": num, + "exp_time": expTime, + "flow_reset_time": flowResetTime, + "speed_id": nullInt64FromInterface(speedID), + "status": status, + }).Error +} + +func (r *Repository) GetUserTunnelUserAndTunnel(id int64) (userID, tunnelID int64, err error) { + if r == nil || r.db == nil { + return 0, 0, errors.New("repository not initialized") + } + var ut model.UserTunnel + err = r.db.Select("user_id", "tunnel_id").Where("id = ?", id).First(&ut).Error + if err != nil { + return 0, 0, normalizeNotFoundErr(err) + } + return ut.UserID, ut.TunnelID, nil +} + +func (r *Repository) GetExistingUserTunnel(userID, tunnelID int64) (id int64, flow, num, expTime, flowReset int64, speedID sql.NullInt64, status int, err error) { + if r == nil || r.db == nil { + return 0, 0, 0, 0, 0, sql.NullInt64{}, 0, errors.New("repository not initialized") + } + var ut model.UserTunnel + err = r.db.Select("id", "flow", "num", "exp_time", "flow_reset_time", "speed_id", "status"). + Where("user_id = ? AND tunnel_id = ?", userID, tunnelID). + First(&ut).Error + if err != nil { + return 0, 0, 0, 0, 0, sql.NullInt64{}, 0, normalizeNotFoundErr(err) + } + return ut.ID, ut.Flow, int64(ut.Num), ut.ExpTime, ut.FlowResetTime, ut.SpeedID, ut.Status, nil +} + +func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + ut := model.UserTunnel{ + UserID: userID, + TunnelID: tunnelID, + SpeedID: nullInt64FromInterface(speedID), + Num: num, + Flow: flow, + InFlow: 0, + OutFlow: 0, + FlowResetTime: flowResetTime, + ExpTime: expTime, + Status: status, + } + return r.db.Create(&ut).Error +} + +func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.UserTunnel{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "speed_id": nullInt64FromInterface(speedID), + "flow": flow, + "num": num, + "exp_time": expTime, + "flow_reset_time": flowResetTime, + "status": status, + }).Error +} + +func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 { + if r == nil || r.db == nil { + return sql.NullInt64{} + } + var p sql.NullInt64 + _ = r.db.Model(&model.ForwardPort{}). + Select("MIN(port)"). + Where("forward_id = ?", forwardID). + Scan(&p).Error + return p +} + +func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.Forward{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "name": name, + "tunnel_id": tunnelID, + "remote_addr": remoteAddr, + "strategy": strategy, + "updated_time": now, + }).Error +} + +func (r *Repository) UpdateForwardOrder(forwardID int64, inx int, now int64) { + if r == nil || r.db == nil { + return + } + _ = r.db.Model(&model.Forward{}). + Where("id = ?", forwardID). + Updates(map[string]interface{}{"inx": inx, "updated_time": now}).Error +} + +func (r *Repository) UpdateForwardTunnel(forwardID, tunnelID int64, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.Forward{}). + Where("id = ?", forwardID). + Updates(map[string]interface{}{"tunnel_id": tunnelID, "updated_time": now}).Error +} + +func (r *Repository) DeleteForwardCascade(forwardID int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("forward_id = ?", forwardID).Delete(&model.ForwardPort{}).Error; err != nil { + return err + } + return tx.Where("id = ?", forwardID).Delete(&model.Forward{}).Error + }) +} + +func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct { + NodeID int64 + Port int +}) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("forward_id = ?", forwardID).Delete(&model.ForwardPort{}).Error; err != nil { + return err + } + if len(entries) == 0 { + return nil + } + rows := make([]model.ForwardPort, 0, len(entries)) + for _, e := range entries { + rows = append(rows, model.ForwardPort{ForwardID: forwardID, NodeID: e.NodeID, Port: e.Port}) + } + return tx.Create(&rows).Error + }) +} + +func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, now int64) { + if r == nil || r.db == nil { + return + } + _ = r.db.Model(&model.Forward{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "user_id": userID, + "user_name": userName, + "name": name, + "tunnel_id": tunnelID, + "remote_addr": remoteAddr, + "strategy": strategy, + "status": status, + "updated_time": now, + }).Error +} + +func (r *Repository) GetUsedPortsOnNodeAsMap(nodeID int64) (map[int]bool, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + used := make(map[int]bool) + var forwardPorts []int + if err := r.db.Model(&model.ForwardPort{}).Where("node_id = ?", nodeID).Pluck("port", &forwardPorts).Error; err != nil { + return nil, err + } + for _, p := range forwardPorts { + used[p] = true + } + var chainPorts []int + if err := r.db.Model(&model.ChainTunnel{}).Where("node_id = ? AND port > 0", nodeID).Pluck("port", &chainPorts).Error; err != nil { + return nil, err + } + for _, p := range chainPorts { + used[p] = true + } + return used, nil +} + +func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID int64, tunnelName string, now int64, status int) (int64, error) { + if r == nil || r.db == nil { + return 0, errors.New("repository not initialized") + } + sl := model.SpeedLimit{ + Name: name, + Speed: speed, + TunnelID: tunnelID, + TunnelName: tunnelName, + CreatedTime: now, + UpdatedTime: sql.NullInt64{Int64: now, Valid: true}, + Status: status, + } + if err := r.db.Create(&sl).Error; err != nil { + return 0, err + } + return sl.ID, nil +} + +func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, tunnelID int64, tunnelName string, status int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Model(&model.SpeedLimit{}). + Where("id = ?", id). + Updates(map[string]interface{}{ + "name": name, + "speed": speed, + "tunnel_id": tunnelID, + "tunnel_name": tunnelName, + "status": status, + "updated_time": sql.NullInt64{ + Int64: now, + Valid: true, + }, + }).Error +} + +func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) int64 { + if r == nil || r.db == nil { + return 0 + } + var sl model.SpeedLimit + if err := r.db.Select("tunnel_id").Where("id = ?", speedLimitID).First(&sl).Error; err != nil { + return 0 + } + return sl.TunnelID +} + +func (r *Repository) DeleteSpeedLimit(id int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Where("id = ?", id).Delete(&model.SpeedLimit{}).Error +} + +func (r *Repository) GroupCreate(table, name string, status int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + switch table { + case "tunnel_group": + return r.db.Create(&model.TunnelGroup{ + Name: name, + CreatedTime: now, + UpdatedTime: now, + Status: status, + }).Error + case "user_group": + return r.db.Create(&model.UserGroup{ + Name: name, + CreatedTime: now, + UpdatedTime: now, + Status: status, + }).Error + default: + return errors.New("invalid group table") + } +} + +func (r *Repository) GroupUpdate(table string, id int64, name string, status int, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + updates := map[string]interface{}{"name": name, "status": status, "updated_time": now} + switch table { + case "tunnel_group": + return r.db.Model(&model.TunnelGroup{}).Where("id = ?", id).Updates(updates).Error + case "user_group": + return r.db.Model(&model.UserGroup{}).Where("id = ?", id).Updates(updates).Error + default: + return errors.New("invalid group table") + } +} + +func (r *Repository) GroupDeleteCascade(table string, id int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + return r.db.Transaction(func(tx *gorm.DB) error { + switch table { + case "tunnel_group": + if err := tx.Where("tunnel_group_id = ?", id).Delete(&model.TunnelGroupTunnel{}).Error; err != nil { + return err + } + if err := tx.Where("tunnel_group_id = ?", id).Delete(&model.GroupPermission{}).Error; err != nil { + return err + } + if err := tx.Where("tunnel_group_id = ?", id).Delete(&model.GroupPermissionGrant{}).Error; err != nil { + return err + } + return tx.Where("id = ?", id).Delete(&model.TunnelGroup{}).Error + case "user_group": + if err := tx.Where("user_group_id = ?", id).Delete(&model.UserGroupUser{}).Error; err != nil { + return err + } + if err := tx.Where("user_group_id = ?", id).Delete(&model.GroupPermission{}).Error; err != nil { + return err + } + if err := tx.Where("user_group_id = ?", id).Delete(&model.GroupPermissionGrant{}).Error; err != nil { + return err + } + return tx.Where("id = ?", id).Delete(&model.UserGroup{}).Error + default: + return errors.New("invalid group table") + } + }) +} + +func (r *Repository) ListUserIDsByUserGroupTx(tx *gorm.DB, userGroupID int64) ([]int64, error) { + if tx == nil { + return nil, errors.New("database unavailable") + } + var ids []int64 + err := tx.Model(&model.UserGroupUser{}). + Where("user_group_id = ?", userGroupID). + Pluck("user_id", &ids).Error + if err != nil { + return nil, err + } + if ids == nil { + ids = make([]int64, 0) + } + return ids, nil +} + +func (r *Repository) ReplaceTunnelGroupMembersTx(tx *gorm.DB, groupID int64, tunnelIDs []int64, now int64) error { + if tx == nil { + return errors.New("database unavailable") + } + if err := tx.Where("tunnel_group_id = ?", groupID).Delete(&model.TunnelGroupTunnel{}).Error; err != nil { + return err + } + if len(tunnelIDs) == 0 { + return nil + } + rows := make([]model.TunnelGroupTunnel, 0, len(tunnelIDs)) + for _, tunnelID := range tunnelIDs { + if tunnelID <= 0 { + continue + } + rows = append(rows, model.TunnelGroupTunnel{TunnelGroupID: groupID, TunnelID: tunnelID, CreatedTime: now}) + } + if len(rows) == 0 { + return nil + } + return tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error +} + +func (r *Repository) ReplaceUserGroupMembersTx(tx *gorm.DB, groupID int64, userIDs []int64, now int64) error { + if tx == nil { + return errors.New("database unavailable") + } + if err := tx.Where("user_group_id = ?", groupID).Delete(&model.UserGroupUser{}).Error; err != nil { + return err + } + if len(userIDs) == 0 { + return nil + } + rows := make([]model.UserGroupUser, 0, len(userIDs)) + for _, userID := range userIDs { + if userID <= 0 { + continue + } + rows = append(rows, model.UserGroupUser{UserGroupID: groupID, UserID: userID, CreatedTime: now}) + } + if len(rows) == 0 { + return nil + } + return tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error +} + +func (r *Repository) GetGroupPermissionPairByIDTx(tx *gorm.DB, id int64) (userGroupID int64, tunnelGroupID int64, exists bool, err error) { + if tx == nil { + return 0, 0, false, errors.New("database unavailable") + } + var gp model.GroupPermission + err = tx.Select("user_group_id", "tunnel_group_id").Where("id = ?", id).First(&gp).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0, 0, false, nil + } + if err != nil { + return 0, 0, false, err + } + return gp.UserGroupID, gp.TunnelGroupID, true, nil +} + +func (r *Repository) DeleteGroupPermissionByIDTx(tx *gorm.DB, id int64) error { + if tx == nil { + return errors.New("database unavailable") + } + return tx.Where("id = ?", id).Delete(&model.GroupPermission{}).Error +} + +func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID int64, previousUserIDs, currentUserIDs []int64) error { + if tx == nil { + return errors.New("database unavailable") + } + currentSet := make(map[int64]struct{}, len(currentUserIDs)) + for _, uid := range currentUserIDs { + if uid > 0 { + currentSet[uid] = struct{}{} + } + } + + removedUserIDs := make([]int64, 0) + for _, uid := range previousUserIDs { + if uid <= 0 { + continue + } + if _, ok := currentSet[uid]; !ok { + removedUserIDs = append(removedUserIDs, uid) + } + } + if len(removedUserIDs) == 0 { + return nil + } + + type grantRow struct { + UserTunnelID int64 + CreatedByGroup int + } + + for _, userID := range removedUserIDs { + var rows []grantRow + if err := tx.Model(&model.GroupPermissionGrant{}). + Select("group_permission_grant.user_tunnel_id, group_permission_grant.created_by_group"). + Joins("JOIN user_tunnel ON user_tunnel.id = group_permission_grant.user_tunnel_id"). + Where("group_permission_grant.user_group_id = ? AND user_tunnel.user_id = ?", userGroupID, userID). + Find(&rows).Error; err != nil { + return err + } + + groupCreatedTunnelIDs := make(map[int64]struct{}) + for _, row := range rows { + if row.CreatedByGroup == 1 && row.UserTunnelID > 0 { + groupCreatedTunnelIDs[row.UserTunnelID] = struct{}{} + } + } + + userTunnelIDs := tx.Model(&model.UserTunnel{}).Select("id").Where("user_id = ?", userID) + if err := tx.Where("user_group_id = ? AND user_tunnel_id IN (?)", userGroupID, userTunnelIDs). + Delete(&model.GroupPermissionGrant{}).Error; err != nil { + return err + } + + for userTunnelID := range groupCreatedTunnelIDs { + var remaining int64 + if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil { + return err + } + if remaining == 0 { + if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil { + return err + } + } + } + } + + return nil +} + +func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunnelGroupID int64) error { + if tx == nil { + return errors.New("database unavailable") + } + + type grantRow struct { + UserTunnelID int64 + CreatedByGroup int + } + + var rows []grantRow + if err := tx.Model(&model.GroupPermissionGrant{}). + Select("user_tunnel_id, created_by_group"). + Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID). + Find(&rows).Error; err != nil { + return err + } + + groupCreatedTunnelIDs := make(map[int64]struct{}) + for _, row := range rows { + if row.CreatedByGroup == 1 && row.UserTunnelID > 0 { + groupCreatedTunnelIDs[row.UserTunnelID] = struct{}{} + } + } + + if err := tx.Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID). + Delete(&model.GroupPermissionGrant{}).Error; err != nil { + return err + } + + for userTunnelID := range groupCreatedTunnelIDs { + var remaining int64 + if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil { + return err + } + if remaining == 0 { + if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil { + return err + } + } + } + + return nil +} + +func (r *Repository) ReplaceFederationTunnelBindingsTx(tx *gorm.DB, tunnelID int64, bindings []FederationTunnelBinding) error { + if tx == nil { + return errors.New("database unavailable") + } + if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.FederationTunnelBinding{}).Error; err != nil { + return err + } + if len(bindings) == 0 { + return nil + } + + rows := make([]model.FederationTunnelBinding, 0, len(bindings)) + now := time.Now().UnixMilli() + for _, b := range bindings { + created := b.CreatedTime + if created <= 0 { + created = now + } + updated := b.UpdatedTime + if updated <= 0 { + updated = created + } + rows = append(rows, model.FederationTunnelBinding{ + TunnelID: tunnelID, + NodeID: b.NodeID, + ChainType: b.ChainType, + HopInx: b.HopInx, + RemoteURL: b.RemoteURL, + ResourceKey: b.ResourceKey, + RemoteBindingID: b.RemoteBindingID, + AllocatedPort: b.AllocatedPort, + Status: b.Status, + CreatedTime: created, + UpdatedTime: updated, + }) + } + + return tx.Create(&rows).Error +} + +func (r *Repository) InsertGroupPermission(userGroupID, tunnelGroupID int64, now int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + gp := model.GroupPermission{UserGroupID: userGroupID, TunnelGroupID: tunnelGroupID, CreatedTime: now} + return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&gp).Error +} + +func (r *Repository) InsertGroupPermissionGrant(userGroupID, tunnelGroupID, userTunnelID int64, createdByGroup int, now int64) { + if r == nil || r.db == nil { + return + } + g := model.GroupPermissionGrant{ + UserGroupID: userGroupID, + TunnelGroupID: tunnelGroupID, + UserTunnelID: userTunnelID, + CreatedByGroup: createdByGroup, + CreatedTime: now, + } + _ = r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&g).Error +} + +func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool, error) { + if r == nil || r.db == nil { + return 0, false, errors.New("repository not initialized") + } + var existing model.UserTunnel + err := r.db.Select("id").Where("user_id = ? AND tunnel_id = ?", userID, tunnelID).First(&existing).Error + if err == nil { + return existing.ID, false, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return 0, false, err + } + var user model.User + if err := r.db.Select("flow, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil { + return 0, false, err + } + flow := user.Flow + num := user.Num + expTime := user.ExpTime + flowReset := user.FlowResetTime + ut := model.UserTunnel{ + UserID: userID, + TunnelID: tunnelID, + Num: num, + Flow: flow, + InFlow: 0, + OutFlow: 0, + FlowResetTime: flowReset, + ExpTime: expTime, + Status: 1, + } + if err := r.db.Create(&ut).Error; err != nil { + return 0, false, err + } + return ut.ID, true, nil +} + +func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int) (int64, error) { + if r == nil || r.db == nil { + return 0, errors.New("repository not initialized") + } + var forwardID int64 + err := r.db.Transaction(func(tx *gorm.DB) error { + fwd := model.Forward{ + UserID: userID, + UserName: userName, + Name: name, + TunnelID: tunnelID, + RemoteAddr: remoteAddr, + Strategy: strategy, + InFlow: 0, + OutFlow: 0, + CreatedTime: now, + UpdatedTime: now, + Status: 1, + Inx: inx, + } + if err := tx.Create(&fwd).Error; err != nil { + return err + } + forwardID = fwd.ID + for _, nodeID := range entryNodeIDs { + fp := model.ForwardPort{ + ForwardID: forwardID, + NodeID: nodeID, + Port: port, + } + if err := tx.Create(&fp).Error; err != nil { + return err + } + } + return nil + }) + return forwardID, err +} + +func (r *Repository) BatchUpdateForwardStatus(ids []int64, status int) (int, int) { + if r == nil || r.db == nil { + return 0, len(ids) + } + s := 0 + f := 0 + now := time.Now().UnixMilli() + for _, id := range ids { + if err := r.db.Model(&model.Forward{}).Where("id = ?", id).Updates(map[string]interface{}{"status": status, "updated_time": now}).Error; err != nil { + f++ + } else { + s++ + } + } + return s, f +} + +func (r *Repository) CreateTunnelTx(tx *gorm.DB, name string, trafficRatio float64, typeVal int, flow int64, now int64, status int, inIP interface{}, inx int, ipPreference string) (int64, error) { + inIPVal := nullStringFromInterface(inIP) + tunnel := model.Tunnel{ + Name: name, + TrafficRatio: trafficRatio, + Type: typeVal, + Protocol: "tls", + Flow: flow, + CreatedTime: now, + UpdatedTime: now, + Status: status, + InIP: inIPVal, + Inx: inx, + IPPreference: ipPreference, + } + if err := tx.Create(&tunnel).Error; err != nil { + return 0, err + } + return tunnel.ID, nil +} + +func normalizeNotFoundErr(err error) error { + if errors.Is(err, gorm.ErrRecordNotFound) { + return sql.ErrNoRows + } + return err +} + +func nullStringFromInterface(v interface{}) sql.NullString { + switch t := v.(type) { + case nil: + return sql.NullString{} + case sql.NullString: + return t + case *sql.NullString: + if t == nil { + return sql.NullString{} + } + return *t + case string: + if t == "" { + return sql.NullString{} + } + return sql.NullString{String: t, Valid: true} + case *string: + if t == nil || *t == "" { + return sql.NullString{} + } + return sql.NullString{String: *t, Valid: true} + default: + return sql.NullString{} + } +} + +func nullInt64FromInterface(v interface{}) sql.NullInt64 { + switch t := v.(type) { + case nil: + return sql.NullInt64{} + case sql.NullInt64: + return t + case *sql.NullInt64: + if t == nil { + return sql.NullInt64{} + } + return *t + case int64: + return sql.NullInt64{Int64: t, Valid: true} + case *int64: + if t == nil { + return sql.NullInt64{} + } + return sql.NullInt64{Int64: *t, Valid: true} + case int: + return sql.NullInt64{Int64: int64(t), Valid: true} + case int32: + return sql.NullInt64{Int64: int64(t), Valid: true} + case int16: + return sql.NullInt64{Int64: int64(t), Valid: true} + case int8: + return sql.NullInt64{Int64: int64(t), Valid: true} + case uint64: + return sql.NullInt64{Int64: int64(t), Valid: true} + case uint: + return sql.NullInt64{Int64: int64(t), Valid: true} + case uint32: + return sql.NullInt64{Int64: int64(t), Valid: true} + case uint16: + return sql.NullInt64{Int64: int64(t), Valid: true} + case uint8: + return sql.NullInt64{Int64: int64(t), Valid: true} + case float64: + return sql.NullInt64{Int64: int64(t), Valid: true} + default: + return sql.NullInt64{} + } +} + +func stringFromInterface(v interface{}) string { + switch t := v.(type) { + case nil: + return "" + case string: + return t + case sql.NullString: + if t.Valid { + return t.String + } + return "" + case *sql.NullString: + if t != nil && t.Valid { + return t.String + } + return "" + case []byte: + return string(t) + default: + return "" + } +} + +func parsePortRangeSpec(input string) []int { + input = strings.TrimSpace(input) + if input == "" { + return nil + } + set := make(map[int]struct{}) + parts := strings.Split(input, ",") + for _, part := range parts { + part = strings.TrimSpace(part) + if part == "" { + continue + } + if strings.Contains(part, "-") { + r := strings.SplitN(part, "-", 2) + if len(r) != 2 { + continue + } + start, err1 := strconv.Atoi(strings.TrimSpace(r[0])) + end, err2 := strconv.Atoi(strings.TrimSpace(r[1])) + if err1 != nil || err2 != nil || start <= 0 || end <= 0 { + continue + } + if end < start { + start, end = end, start + } + for p := start; p <= end; p++ { + set[p] = struct{}{} + } + continue + } + p, err := strconv.Atoi(part) + if err != nil || p <= 0 { + continue + } + set[p] = struct{}{} + } + out := make([]int, 0, len(set)) + for p := range set { + out = append(out, p) + } + sort.Ints(out) + return out +} diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go deleted file mode 100644 index fa9d7d6..0000000 --- a/go-backend/internal/store/sqlite/repository.go +++ /dev/null @@ -1,3117 +0,0 @@ -package sqlite - -import ( - "database/sql" - _ "embed" - "errors" - "fmt" - "log" - "os" - "path/filepath" - "sort" - "strings" - "time" - - _ "github.com/jackc/pgx/v5/stdlib" - "go-backend/internal/store" - pgstore "go-backend/internal/store/postgres" - _ "modernc.org/sqlite" -) - -//go:embed sql/schema.sql -var embeddedSchema string - -//go:embed sql/data.sql -var embeddedSeedData string - -// Execer is an interface that both *store.DB and *store.Tx satisfy. -// Used to allow import functions to work with both regular DB and transactions. -type Execer interface { - Exec(query string, args ...any) (sql.Result, error) - Query(query string, args ...any) (*sql.Rows, error) - QueryRow(query string, args ...any) *sql.Row -} - -type Repository struct { - db *store.DB -} - -func (r *Repository) DB() *store.DB { - if r == nil { - return nil - } - return r.db -} - -type User struct { - ID int64 - User string - Pwd string - RoleID int - ExpTime int64 - Flow int64 - InFlow int64 - OutFlow int64 - FlowResetTime int64 - Num int - CreatedTime int64 - UpdatedTime sql.NullInt64 - Status int -} - -type ViteConfig struct { - ID int64 `json:"id"` - Name string `json:"name"` - Value string `json:"value"` - Time int64 `json:"time"` -} - -type Announcement struct { - ID int64 `json:"id"` - Content string `json:"content"` - Enabled int `json:"enabled"` - CreatedTime int64 `json:"created_time"` - UpdatedTime sql.NullInt64 `json:"updated_time,omitempty"` -} - -type UserTunnelDetail struct { - ID int64 - UserID int64 - TunnelID int64 - TunnelName string - TunnelFlow int - Flow int64 - InFlow int64 - OutFlow int64 - Num int - FlowResetTime int64 - ExpTime int64 - SpeedID sql.NullInt64 - SpeedLimit sql.NullString - Speed sql.NullInt64 -} - -type UserForwardDetail struct { - ID int64 - Name string - TunnelID int64 - TunnelName string - InIP string - InPort sql.NullInt64 - RemoteAddr string - InFlow int64 - OutFlow int64 - Status int - CreatedAt int64 -} - -type StatisticsFlow struct { - ID int64 `json:"id"` - UserID int64 `json:"userId"` - Flow int64 `json:"flow"` - TotalFlow int64 `json:"totalFlow"` - Time string `json:"time"` -} - -type Node struct { - ID int64 - Secret string - Version sql.NullString - HTTP int - TLS int - Socks int - Status int - IsRemote int - RemoteURL sql.NullString - RemoteToken sql.NullString - RemoteConfig sql.NullString -} - -type PeerShare struct { - ID int64 `json:"id"` - Name string `json:"name"` - NodeID int64 `json:"nodeId"` - Token string `json:"token"` - MaxBandwidth int64 `json:"maxBandwidth"` - ExpiryTime int64 `json:"expiryTime"` - PortRangeStart int `json:"portRangeStart"` - PortRangeEnd int `json:"portRangeEnd"` - CurrentFlow int64 `json:"currentFlow"` - IsActive int `json:"isActive"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime"` - AllowedDomains string `json:"allowedDomains"` - AllowedIPs string `json:"allowedIps"` -} - -type PeerShareRuntime struct { - ID int64 - ShareID int64 - NodeID int64 - ReservationID string - ResourceKey string - BindingID string - Role string - ChainName string - ServiceName string - Protocol string - Strategy string - Port int - Target string - Applied int - Status int - CreatedTime int64 - UpdatedTime int64 -} - -type FederationTunnelBinding struct { - ID int64 - TunnelID int64 - NodeID int64 - ChainType int - HopInx int - RemoteURL string - ResourceKey string - RemoteBindingID string - AllocatedPort int - Status int - CreatedTime int64 - UpdatedTime int64 -} - -func Open(path string) (*Repository, error) { - if err := ensureParentDir(path); err != nil { - return nil, err - } - - // Use _pragma DSN parameters so every connection from the pool gets - // the same settings (busy_timeout and synchronous are per-connection). - dsn := "file:" + path + - "?_pragma=busy_timeout(5000)" + - "&_pragma=journal_mode(WAL)" + - "&_pragma=synchronous(NORMAL)" - raw, err := sql.Open("sqlite", dsn) - 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, 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 - } - - if err := migrateSchema(db); err != nil { - _ = db.Close() - return nil, err - } - - return &Repository{db: db}, nil -} - -func (r *Repository) Close() error { - if r == nil || r.db == nil { - return nil - } - return r.db.Close() -} - -func (r *Repository) GetUserByUsername(username string) (*User, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - row := r.db.QueryRow(` - SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status - FROM user WHERE user = ? LIMIT 1 - `, username) - user := &User{} - if err := row.Scan( - &user.ID, &user.User, &user.Pwd, &user.RoleID, &user.ExpTime, - &user.Flow, &user.InFlow, &user.OutFlow, &user.FlowResetTime, - &user.Num, &user.CreatedTime, &user.UpdatedTime, &user.Status, - ); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return user, nil -} - -func (r *Repository) GetConfigByName(name string) (*ViteConfig, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - row := r.db.QueryRow(`SELECT id, name, value, time FROM vite_config WHERE name = ? LIMIT 1`, name) - cfg := &ViteConfig{} - if err := row.Scan(&cfg.ID, &cfg.Name, &cfg.Value, &cfg.Time); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return cfg, nil -} - -func (r *Repository) ListConfigs() (map[string]string, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(`SELECT name, value FROM vite_config`) - if err != nil { - return nil, err - } - defer rows.Close() - - result := make(map[string]string) - for rows.Next() { - var name, value string - if err := rows.Scan(&name, &value); err != nil { - return nil, err - } - result[name] = value - } - if err := rows.Err(); err != nil { - return nil, err - } - return result, nil -} - -func (r *Repository) UpsertConfig(name, value string, now int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - - _, err := r.db.Exec(` - INSERT INTO vite_config(name, value, time) - VALUES(?, ?, ?) - ON CONFLICT(name) DO UPDATE SET value=excluded.value, time=excluded.time - `, name, value, now) - return err -} - -func (r *Repository) GetAnnouncement() (*Announcement, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - row := r.db.QueryRow(`SELECT id, content, enabled, created_time, updated_time FROM announcement ORDER BY id DESC LIMIT 1`) - ann := &Announcement{} - if err := row.Scan(&ann.ID, &ann.Content, &ann.Enabled, &ann.CreatedTime, &ann.UpdatedTime); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return ann, nil -} - -func (r *Repository) UpsertAnnouncement(content string, enabled int, now int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - - var count int - err := r.db.QueryRow(`SELECT COUNT(*) FROM announcement`).Scan(&count) - if err != nil { - return err - } - - if count == 0 { - _, err = r.db.Exec(` - INSERT INTO announcement(content, enabled, created_time, updated_time) - VALUES(?, ?, ?, ?) - `, content, enabled, now, now) - } else { - _, err = r.db.Exec(` - UPDATE announcement SET content = ?, enabled = ?, updated_time = ? - `, content, enabled, now) - } - return err -} - -func (r *Repository) GetUserByID(id int64) (*User, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - row := r.db.QueryRow(` - SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status - FROM user WHERE id = ? LIMIT 1 - `, id) - user := &User{} - if err := row.Scan( - &user.ID, &user.User, &user.Pwd, &user.RoleID, &user.ExpTime, - &user.Flow, &user.InFlow, &user.OutFlow, &user.FlowResetTime, - &user.Num, &user.CreatedTime, &user.UpdatedTime, &user.Status, - ); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return user, nil -} - -func (r *Repository) UsernameExistsExceptID(username string, exceptID int64) (bool, error) { - if r == nil || r.db == nil { - return false, errors.New("repository not initialized") - } - - row := r.db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ? AND id != ?`, username, exceptID) - var count int - if err := row.Scan(&count); err != nil { - return false, err - } - return count > 0, nil -} - -func (r *Repository) UpdateUserNameAndPassword(userID int64, username, passwordMD5 string, now int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - _, err := r.db.Exec(`UPDATE user SET user = ?, pwd = ?, updated_time = ? WHERE id = ?`, username, passwordMD5, now, userID) - return err -} - -func (r *Repository) GetUserPackageTunnels(userID int64) ([]UserTunnelDetail, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT ut.id, ut.user_id, ut.tunnel_id, t.name, t.flow, ut.flow, ut.in_flow, ut.out_flow, - ut.num, ut.flow_reset_time, ut.exp_time, ut.speed_id, sl.name, sl.speed - FROM user_tunnel ut - LEFT JOIN tunnel t ON t.id = ut.tunnel_id - LEFT JOIN speed_limit sl ON sl.id = ut.speed_id - WHERE ut.user_id = ? - ORDER BY ut.id ASC - `, userID) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]UserTunnelDetail, 0) - for rows.Next() { - var item UserTunnelDetail - if err := rows.Scan( - &item.ID, &item.UserID, &item.TunnelID, &item.TunnelName, &item.TunnelFlow, - &item.Flow, &item.InFlow, &item.OutFlow, &item.Num, &item.FlowResetTime, - &item.ExpTime, &item.SpeedID, &item.SpeedLimit, &item.Speed, - ); err != nil { - return nil, err - } - items = append(items, item) - } - - if err := rows.Err(); err != nil { - return nil, err - } - - return items, nil -} - -func (r *Repository) GetUserPackageForwards(userID int64) ([]UserForwardDetail, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - 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 - LEFT JOIN tunnel t ON t.id = f.tunnel_id - WHERE f.user_id = ? - ORDER BY f.id ASC - `, userID) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]UserForwardDetail, 0) - for rows.Next() { - var item UserForwardDetail - if err := rows.Scan( - &item.ID, &item.Name, &item.TunnelID, &item.TunnelName, &item.RemoteAddr, - &item.InFlow, &item.OutFlow, &item.Status, &item.CreatedAt, - ); err != nil { - return nil, err - } - - inIP, inPort, err := resolveForwardIngress(r.db, item.ID, item.TunnelID) - if err != nil { - return nil, err - } - item.InIP = inIP - item.InPort = inPort - - items = append(items, item) - } - - if err := rows.Err(); err != nil { - return nil, err - } - - return items, nil -} - -func (r *Repository) GetStatisticsFlows(userID int64, limit int) ([]StatisticsFlow, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT id, user_id, flow, total_flow, time - FROM statistics_flow - WHERE user_id = ? - ORDER BY id DESC - LIMIT ? - `, userID, limit) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]StatisticsFlow, 0) - for rows.Next() { - var item StatisticsFlow - if err := rows.Scan(&item.ID, &item.UserID, &item.Flow, &item.TotalFlow, &item.Time); err != nil { - return nil, err - } - items = append(items, item) - } - - if err := rows.Err(); err != nil { - return nil, err - } - - return items, nil -} - -func (r *Repository) NodeExistsBySecret(secret string) (bool, error) { - if r == nil || r.db == nil { - return false, errors.New("repository not initialized") - } - - row := r.db.QueryRow(`SELECT COUNT(1) FROM node WHERE secret = ?`, secret) - var count int - if err := row.Scan(&count); err != nil { - return false, err - } - return count > 0, nil -} - -func (r *Repository) GetNodeBySecret(secret string) (*Node, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - row := r.db.QueryRow(`SELECT id, secret, version, http, tls, socks, status, is_remote, remote_url, remote_token, remote_config FROM node WHERE secret = ? LIMIT 1`, secret) - var n Node - if err := row.Scan(&n.ID, &n.Secret, &n.Version, &n.HTTP, &n.TLS, &n.Socks, &n.Status, &n.IsRemote, &n.RemoteURL, &n.RemoteToken, &n.RemoteConfig); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &n, nil -} - -func (r *Repository) GetNodeByID(id int64) (*Node, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - row := r.db.QueryRow(`SELECT id, secret, version, http, tls, socks, status, is_remote, remote_url, remote_token, remote_config FROM node WHERE id = ? LIMIT 1`, id) - var n Node - if err := row.Scan(&n.ID, &n.Secret, &n.Version, &n.HTTP, &n.TLS, &n.Socks, &n.Status, &n.IsRemote, &n.RemoteURL, &n.RemoteToken, &n.RemoteConfig); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &n, nil -} - -func (r *Repository) UpdateNodeOnline(nodeID int64, status int, version string, httpVal, tlsVal, socksVal int) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - _, err := r.db.Exec(`UPDATE node SET status = ?, version = ?, http = ?, tls = ?, socks = ?, updated_time = ? WHERE id = ?`, - status, version, httpVal, tlsVal, socksVal, unixMilliNow(), nodeID) - return err -} - -func (r *Repository) UpdateNodeStatus(nodeID int64, status int) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - _, err := r.db.Exec(`UPDATE node SET status = ?, updated_time = ? WHERE id = ?`, status, unixMilliNow(), nodeID) - return err -} - -func (r *Repository) AddFlow(forwardID, userID int64, userTunnelID int64, inFlow, outFlow int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - - tx, err := r.db.Begin() - if err != nil { - return err - } - defer func() { - if err != nil { - _ = tx.Rollback() - } - }() - - if _, err = tx.Exec(`UPDATE forward SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, forwardID); err != nil { - return err - } - if _, err = tx.Exec(`UPDATE user SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, userID); err != nil { - return err - } - if userTunnelID > 0 { - if _, err = tx.Exec(`UPDATE user_tunnel SET in_flow = in_flow + ?, out_flow = out_flow + ? WHERE id = ?`, inFlow, outFlow, userTunnelID); err != nil { - return err - } - } - - err = tx.Commit() - return err -} - -func (r *Repository) ListNodes() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT id, inx, name, server_ip, server_ip_v4, server_ip_v6, port, tcp_listen_addr, udp_listen_addr, version, http, tls, socks, status, is_remote, remote_url, remote_token, remote_config - FROM node - ORDER BY inx ASC, id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]map[string]interface{}, 0) - for rows.Next() { - var id, inx int64 - var name, serverIP, port string - var serverIPV4, serverIPV6, tcpListen, udpListen, version, remoteURL, remoteToken, remoteConfig sql.NullString - var httpVal, tlsVal, socksVal, status, isRemote int - - if err := rows.Scan(&id, &inx, &name, &serverIP, &serverIPV4, &serverIPV6, &port, &tcpListen, &udpListen, &version, &httpVal, &tlsVal, &socksVal, &status, &isRemote, &remoteURL, &remoteToken, &remoteConfig); err != nil { - return nil, err - } - - items = append(items, map[string]interface{}{ - "id": id, - "inx": inx, - "name": name, - "ip": serverIP, - "serverIp": serverIP, - "serverIpV4": nullableString(serverIPV4), - "serverIpV6": nullableString(serverIPV6), - "port": port, - "tcpListenAddr": nullableString(tcpListen), - "udpListenAddr": nullableString(udpListen), - "version": nullableString(version), - "http": httpVal, - "tls": tlsVal, - "socks": socksVal, - "status": status, - "isRemote": isRemote, - "remoteUrl": nullableString(remoteURL), - "remoteToken": nullableString(remoteToken), - "remoteConfig": nullableString(remoteConfig), - }) - } - - if err := rows.Err(); err != nil { - return nil, err - } - return items, nil -} - -func (r *Repository) ListUsers() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT id, user, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status - FROM user - WHERE role_id != 0 - ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]map[string]interface{}, 0) - for rows.Next() { - var id int64 - var user string - var roleID int - var expTime, flow, inFlow, outFlow, flowResetTime, createdTime int64 - var num, status int - var updatedTime sql.NullInt64 - - if err := rows.Scan(&id, &user, &roleID, &expTime, &flow, &inFlow, &outFlow, &flowResetTime, &num, &createdTime, &updatedTime, &status); err != nil { - return nil, err - } - - items = append(items, map[string]interface{}{ - "id": id, - "user": user, - "name": user, - "roleId": roleID, - "status": status, - "flow": flow, - "num": num, - "expTime": expTime, - "flowResetTime": flowResetTime, - "createdTime": createdTime, - "updatedTime": nullableInt64(updatedTime), - "inFlow": inFlow, - "outFlow": outFlow, - }) - } - - if err := rows.Err(); err != nil { - return nil, err - } - return items, nil -} - -func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT id, name, speed, tunnel_id, tunnel_name, status, created_time, updated_time - FROM speed_limit - ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]map[string]interface{}, 0) - for rows.Next() { - var id, tunnelID, createdTime int64 - var name, tunnelName string - var speed, status int - var updatedTime sql.NullInt64 - if err := rows.Scan(&id, &name, &speed, &tunnelID, &tunnelName, &status, &createdTime, &updatedTime); err != nil { - return nil, err - } - items = append(items, map[string]interface{}{ - "id": id, - "name": name, - "speed": speed, - "tunnelId": tunnelID, - "tunnelName": tunnelName, - "status": status, - "createdTime": createdTime, - "updatedTime": nullableInt64(updatedTime), - }) - } - - if err := rows.Err(); err != nil { - return nil, err - } - return items, nil -} - -func (r *Repository) ListForwards() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT f.id, f.user_id, f.user_name, f.name, f.tunnel_id, COALESCE(t.name, ''), f.remote_addr, COALESCE(f.strategy, 'fifo'), - f.in_flow, f.out_flow, f.created_time, f.status, f.inx - FROM forward f - LEFT JOIN tunnel t ON t.id = f.tunnel_id - ORDER BY f.inx ASC, f.id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]map[string]interface{}, 0) - for rows.Next() { - var id, userID, tunnelID, inFlow, outFlow, createdTime, inx int64 - var userName, name, tunnelName, remoteAddr, strategy string - var status int - - if err := rows.Scan(&id, &userID, &userName, &name, &tunnelID, &tunnelName, &remoteAddr, &strategy, &inFlow, &outFlow, &createdTime, &status, &inx); err != nil { - return nil, err - } - - inIP, inPort, err := resolveForwardIngress(r.db, id, tunnelID) - if err != nil { - return nil, err - } - - items = append(items, map[string]interface{}{ - "id": id, - "userId": userID, - "userName": userName, - "name": name, - "tunnelId": tunnelID, - "tunnelName": tunnelName, - "inIp": nullableForwardIngress(inIP), - "inPort": nullableInt64(inPort), - "remoteAddr": remoteAddr, - "strategy": strategy, - "inFlow": inFlow, - "outFlow": outFlow, - "createdTime": createdTime, - "status": status, - "inx": inx, - }) - } - - if err := rows.Err(); err != nil { - return nil, err - } - return items, nil -} - -func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT t.id, t.name - FROM user_tunnel ut - JOIN tunnel t ON t.id = ut.tunnel_id - WHERE ut.user_id = ? AND t.status = 1 - ORDER BY t.inx ASC, t.id ASC - `, userID) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]map[string]interface{}, 0) - for rows.Next() { - var id int64 - var name string - if err := rows.Scan(&id, &name); err != nil { - return nil, err - } - items = append(items, map[string]interface{}{"id": id, "name": name}) - } - - if err := rows.Err(); err != nil { - return nil, err - } - return items, nil -} - -func (r *Repository) ListEnabledTunnelSummaries() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT id, name - FROM tunnel - WHERE status = 1 - ORDER BY inx ASC, id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - items := make([]map[string]interface{}, 0) - for rows.Next() { - var id int64 - var name string - if err := rows.Scan(&id, &name); err != nil { - return nil, err - } - items = append(items, map[string]interface{}{"id": id, "name": name}) - } - - if err := rows.Err(); err != nil { - return nil, err - } - return items, nil -} - -func (r *Repository) ListTunnels() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT id, inx, name, type, flow, traffic_ratio, status, created_time, in_ip, COALESCE(ip_preference, '') - FROM tunnel - ORDER BY inx ASC, id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - tunnelMap := make(map[int64]map[string]interface{}) - orderedIDs := make([]int64, 0) - - for rows.Next() { - var id, inx, flow, createdTime int64 - var name string - var typ, status int - var trafficRatio float64 - var inIP sql.NullString - var ipPreference string - if err := rows.Scan(&id, &inx, &name, &typ, &flow, &trafficRatio, &status, &createdTime, &inIP, &ipPreference); err != nil { - return nil, err - } - - tunnelMap[id] = map[string]interface{}{ - "id": id, - "inx": inx, - "name": name, - "type": typ, - "flow": flow, - "trafficRatio": trafficRatio, - "status": status, - "createdTime": createdTime, - "inIp": nullableString(inIP), - "ipPreference": ipPreference, - "inNodeId": make([]map[string]interface{}, 0), - "outNodeId": make([]map[string]interface{}, 0), - "chainNodes": make([][]map[string]interface{}, 0), - } - orderedIDs = append(orderedIDs, id) - } - if err := rows.Err(); err != nil { - return nil, err - } - - nodeIPMap := map[int64]string{} - nRows, err := r.db.Query(`SELECT id, server_ip FROM node`) - if err == nil { - for nRows.Next() { - var id int64 - var ip string - if scanErr := nRows.Scan(&id, &ip); scanErr == nil { - nodeIPMap[id] = ip - } - } - _ = nRows.Close() - } - - chainRows, err := r.db.Query(` - SELECT tunnel_id, CAST(chain_type AS INTEGER), node_id, protocol, strategy, COALESCE(inx, 0) - FROM chain_tunnel - ORDER BY tunnel_id ASC, CAST(chain_type AS INTEGER) ASC, inx ASC, id ASC - `) - if err != nil { - return nil, err - } - defer chainRows.Close() - - chainBucket := map[int64]map[int][]map[string]interface{}{} - inNodeIPs := map[int64][]string{} - - for chainRows.Next() { - var tunnelID, nodeID, inx int64 - var chainType int - var protocol, strategy sql.NullString - if err := chainRows.Scan(&tunnelID, &chainType, &nodeID, &protocol, &strategy, &inx); err != nil { - return nil, err - } - - t, ok := tunnelMap[tunnelID] - if !ok { - continue - } - - nodeObj := map[string]interface{}{ - "nodeId": nodeID, - "chainType": chainType, - "inx": inx, - } - if protocol.Valid { - nodeObj["protocol"] = protocol.String - } - if strategy.Valid { - nodeObj["strategy"] = strategy.String - } - - switch chainType { - case 1: - t["inNodeId"] = append(t["inNodeId"].([]map[string]interface{}), nodeObj) - if ip, ok := nodeIPMap[nodeID]; ok && ip != "" { - inNodeIPs[tunnelID] = append(inNodeIPs[tunnelID], ip) - } - case 2: - if _, ok := chainBucket[tunnelID]; !ok { - chainBucket[tunnelID] = map[int][]map[string]interface{}{} - } - chainBucket[tunnelID][int(inx)] = append(chainBucket[tunnelID][int(inx)], nodeObj) - case 3: - t["outNodeId"] = append(t["outNodeId"].([]map[string]interface{}), nodeObj) - } - } - if err := chainRows.Err(); err != nil { - return nil, err - } - - for tunnelID, groups := range chainBucket { - t := tunnelMap[tunnelID] - if t == nil { - continue - } - keys := make([]int, 0, len(groups)) - for k := range groups { - keys = append(keys, k) - } - sort.Ints(keys) - ordered := make([][]map[string]interface{}, 0, len(keys)) - for _, k := range keys { - ordered = append(ordered, groups[k]) - } - t["chainNodes"] = ordered - - if s, ok := t["inIp"].(string); !ok || strings.TrimSpace(s) == "" { - if ips := inNodeIPs[tunnelID]; len(ips) > 0 { - t["inIp"] = strings.Join(ips, ",") - } - } - } - - result := make([]map[string]interface{}, 0, len(orderedIDs)) - for _, id := range orderedIDs { - if t, ok := tunnelMap[id]; ok { - result = append(result, t) - } - } - return result, nil -} - -func (r *Repository) ListTunnelGroups() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(`SELECT id, name, status, created_time FROM tunnel_group ORDER BY id ASC`) - if err != nil { - return nil, err - } - defer rows.Close() - - result := make([]map[string]interface{}, 0) - for rows.Next() { - var id, createdTime int64 - var name string - var status int - if err := rows.Scan(&id, &name, &status, &createdTime); err != nil { - return nil, err - } - - ids, names, err := r.listTunnelGroupMembers(id) - if err != nil { - return nil, err - } - - result = append(result, map[string]interface{}{ - "id": id, - "name": name, - "status": status, - "tunnelIds": ids, - "tunnelNames": names, - "createdTime": createdTime, - }) - } - if err := rows.Err(); err != nil { - return nil, err - } - return result, nil -} - -func (r *Repository) ListUserGroups() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(`SELECT id, name, status, created_time FROM user_group ORDER BY id ASC`) - if err != nil { - return nil, err - } - defer rows.Close() - - result := make([]map[string]interface{}, 0) - for rows.Next() { - var id, createdTime int64 - var name string - var status int - if err := rows.Scan(&id, &name, &status, &createdTime); err != nil { - return nil, err - } - - ids, names, err := r.listUserGroupMembers(id) - if err != nil { - return nil, err - } - - result = append(result, map[string]interface{}{ - "id": id, - "name": name, - "status": status, - "userIds": ids, - "userNames": names, - "createdTime": createdTime, - }) - } - if err := rows.Err(); err != nil { - return nil, err - } - return result, nil -} - -func (r *Repository) ListGroupPermissions() ([]map[string]interface{}, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - - rows, err := r.db.Query(` - SELECT gp.id, gp.user_group_id, ug.name, gp.tunnel_group_id, tg.name, gp.created_time - FROM group_permission gp - LEFT JOIN user_group ug ON ug.id = gp.user_group_id - LEFT JOIN tunnel_group tg ON tg.id = gp.tunnel_group_id - ORDER BY gp.id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - result := make([]map[string]interface{}, 0) - for rows.Next() { - var id, userGroupID, tunnelGroupID, createdTime int64 - var userGroupName, tunnelGroupName sql.NullString - if err := rows.Scan(&id, &userGroupID, &userGroupName, &tunnelGroupID, &tunnelGroupName, &createdTime); err != nil { - return nil, err - } - - result = append(result, map[string]interface{}{ - "id": id, - "userGroupId": userGroupID, - "userGroupName": nullableString(userGroupName), - "tunnelGroupId": tunnelGroupID, - "tunnelGroupName": nullableString(tunnelGroupName), - "createdTime": createdTime, - }) - } - if err := rows.Err(); err != nil { - return nil, err - } - return result, nil -} - -func (r *Repository) listTunnelGroupMembers(groupID int64) ([]int64, []string, error) { - rows, err := r.db.Query(` - SELECT t.id, t.name - FROM tunnel_group_tunnel tgt - JOIN tunnel t ON t.id = tgt.tunnel_id - WHERE tgt.tunnel_group_id = ? - ORDER BY t.id ASC - `, groupID) - if err != nil { - return nil, nil, err - } - defer rows.Close() - - ids := make([]int64, 0) - names := make([]string, 0) - for rows.Next() { - var id int64 - var name string - if err := rows.Scan(&id, &name); err != nil { - return nil, nil, err - } - ids = append(ids, id) - names = append(names, name) - } - if err := rows.Err(); err != nil { - return nil, nil, err - } - return ids, names, nil -} - -func (r *Repository) listUserGroupMembers(groupID int64) ([]int64, []string, error) { - rows, err := r.db.Query(` - SELECT u.id, u.user - FROM user_group_user ugu - JOIN user u ON u.id = ugu.user_id - WHERE ugu.user_group_id = ? - ORDER BY u.id ASC - `, groupID) - if err != nil { - return nil, nil, err - } - defer rows.Close() - - ids := make([]int64, 0) - names := make([]string, 0) - for rows.Next() { - var id int64 - var name string - if err := rows.Scan(&id, &name); err != nil { - return nil, nil, err - } - ids = append(ids, id) - names = append(names, name) - } - if err := rows.Err(); err != nil { - return nil, nil, err - } - return ids, names, nil -} - -func nullableString(v sql.NullString) interface{} { - if v.Valid { - return v.String - } - return nil -} - -func nullableForwardIngress(v string) interface{} { - v = strings.TrimSpace(v) - if v == "" { - return nil - } - return v -} - -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) { - return "", sql.NullInt64{}, err - } - } - - rows, err := db.Query(` - SELECT fp.port, n.server_ip - FROM forward_port fp - LEFT JOIN node n ON n.id = fp.node_id - WHERE fp.forward_id = ? - ORDER BY fp.id ASC - `, forwardID) - if err != nil { - return "", sql.NullInt64{}, err - } - defer rows.Close() - - ports := make([]int64, 0) - nodePairs := make([]string, 0) - seenPorts := make(map[int64]struct{}) - seenPairs := make(map[string]struct{}) - - for rows.Next() { - var port sql.NullInt64 - var nodeIP sql.NullString - if err := rows.Scan(&port, &nodeIP); err != nil { - return "", sql.NullInt64{}, err - } - if !port.Valid { - continue - } - if _, ok := seenPorts[port.Int64]; !ok { - seenPorts[port.Int64] = struct{}{} - ports = append(ports, port.Int64) - } - if nodeIP.Valid && strings.TrimSpace(nodeIP.String) != "" { - pair := fmt.Sprintf("%s:%d", strings.TrimSpace(nodeIP.String), port.Int64) - if _, ok := seenPairs[pair]; !ok { - seenPairs[pair] = struct{}{} - nodePairs = append(nodePairs, pair) - } - } - } - if err := rows.Err(); err != nil { - return "", sql.NullInt64{}, err - } - - if len(ports) == 0 { - return "", sql.NullInt64{}, nil - } - - inPort := sql.NullInt64{Int64: ports[0], Valid: true} - - entries := make([]string, 0) - if tunnelInIP.Valid && strings.TrimSpace(tunnelInIP.String) != "" { - tunnelIPs := strings.Split(tunnelInIP.String, ",") - seen := make(map[string]struct{}) - for _, ip := range tunnelIPs { - ip = strings.TrimSpace(ip) - if ip == "" { - continue - } - if _, ok := seen[ip]; ok { - continue - } - seen[ip] = struct{}{} - for _, port := range ports { - entries = append(entries, fmt.Sprintf("%s:%d", ip, port)) - } - } - } else { - entries = append(entries, nodePairs...) - } - - return strings.Join(entries, ","), inPort, nil -} - -func nullableInt64(v sql.NullInt64) interface{} { - if v.Valid { - return v.Int64 - } - return nil -} - -func unixMilliNow() int64 { - return time.Now().UnixMilli() -} - -func ensureParentDir(dbPath string) error { - if dbPath == "" { - return fmt.Errorf("empty db path") - } - dir := filepath.Dir(dbPath) - if dir == "" || dir == "." { - return nil - } - return osMkdirAll(dir) -} - -func bootstrapSchema(db *store.DB, schemaSQL, seedSQL string) error { - if db == nil { - return errors.New("nil db") - } - - if _, err := db.Exec(schemaSQL); err != nil { - return fmt.Errorf("apply schema.sql: %w", err) - } - - if _, err := db.Exec(seedSQL); err != nil { - return fmt.Errorf("apply data.sql: %w", err) - } - return nil -} - -const currentSchemaVersion = 2 - -var ensurePostgresIDDefaultsFn = ensurePostgresIDDefaults - -func getSchemaVersion(db *store.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 *store.DB, v int) { - _, _ = db.Exec(`UPDATE schema_version SET version = ?`, v) -} - -func migrateSchema(db *store.DB) error { - if db == nil { - return errors.New("nil db") - } - - ver := getSchemaVersion(db) - if db.Dialect() == store.DialectPostgres { - if err := ensurePostgresIDDefaultsFn(db); err != nil { - return err - } - } - if ver >= currentSchemaVersion { - return nil - } - - ensureColumn := func(table, col, typ string) { - var dummy interface{} - err := db.QueryRow(fmt.Sprintf("SELECT %s FROM %s LIMIT 1", col, table)).Scan(&dummy) - if err == nil || errors.Is(err, sql.ErrNoRows) { - return - } - 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) - } - } - } - - columnsByTable := map[string]map[string]string{ - "peer_share": { - "allowed_domains": "TEXT DEFAULT ''", - "allowed_ips": "TEXT DEFAULT ''", - }, - "node": { - "server_ip_v4": "VARCHAR(100)", - "server_ip_v6": "VARCHAR(100)", - "inx": "INTEGER NOT NULL DEFAULT 0", - "is_remote": "INTEGER DEFAULT 0", - "remote_url": "TEXT", - "remote_token": "TEXT", - "remote_config": "TEXT", - }, - "tunnel": { - "inx": "INTEGER NOT NULL DEFAULT 0", - "ip_preference": "VARCHAR(10) NOT NULL DEFAULT ''", - }, - "forward": { - "inx": "INTEGER NOT NULL DEFAULT 0", - }, - "chain_tunnel": { - "inx": "INTEGER", - }, - } - - for table, columns := range columnsByTable { - for col, typ := range columns { - ensureColumn(table, col, typ) - } - } - - normalizeStrategy := func(table, defaultValue string) error { - _, err := db.Exec(fmt.Sprintf("UPDATE %s SET strategy = ? WHERE strategy IS NULL", table), defaultValue) - if err != nil { - if isMissingTableError(db.Dialect(), err) { - return nil - } - return fmt.Errorf("normalize %s.strategy: %w", table, err) - } - return nil - } - - if err := normalizeStrategy("forward", "fifo"); err != nil { - return err - } - if err := normalizeStrategy("chain_tunnel", "round"); err != nil { - return err - } - if err := normalizeStrategy("peer_share_runtime", "round"); err != nil { - return err - } - - setSchemaVersion(db, currentSchemaVersion) - return nil -} - -func ensurePostgresIDDefaults(db *store.DB) error { - rows, err := db.Query(` - SELECT c.table_schema, c.table_name - FROM information_schema.table_constraints tc - JOIN information_schema.key_column_usage kcu - ON tc.constraint_name = kcu.constraint_name - AND tc.table_schema = kcu.table_schema - JOIN information_schema.columns c - ON c.table_schema = kcu.table_schema - AND c.table_name = kcu.table_name - AND c.column_name = kcu.column_name - WHERE tc.constraint_type = 'PRIMARY KEY' - AND kcu.column_name = 'id' - AND c.data_type IN ('integer', 'bigint') - AND c.is_identity = 'NO' - AND c.table_schema = current_schema() - ORDER BY c.table_name ASC - `) - if err != nil { - return fmt.Errorf("discover postgres id columns: %w", err) - } - defer rows.Close() - - for rows.Next() { - var schemaName string - var tableName string - if err := rows.Scan(&schemaName, &tableName); err != nil { - return fmt.Errorf("scan postgres id table row: %w", err) - } - if err := ensurePostgresTableIDDefault(db, schemaName, tableName); err != nil { - return fmt.Errorf("repair %s.%s id default: %w", schemaName, tableName, err) - } - } - if err := rows.Err(); err != nil { - return fmt.Errorf("iterate postgres id tables: %w", err) - } - - return nil -} - -func ensurePostgresTableIDDefault(db *store.DB, schemaName, tableName string) error { - var defaultExpr sql.NullString - if err := db.QueryRow(` - SELECT column_default - FROM information_schema.columns - WHERE table_schema = ? - AND table_name = ? - AND column_name = 'id' - LIMIT 1 - `, schemaName, tableName).Scan(&defaultExpr); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil - } - return err - } - - hasNextvalDefault := defaultExpr.Valid && strings.Contains(strings.ToLower(defaultExpr.String), "nextval(") - - var serialSeq sql.NullString - if err := db.QueryRow(` - SELECT pg_get_serial_sequence(quote_ident(?) || '.' || quote_ident(?), 'id') - `, schemaName, tableName).Scan(&serialSeq); err != nil { - return err - } - - seqRef := strings.TrimSpace(serialSeq.String) - if seqRef == "" && hasNextvalDefault { - seqRef = extractNextvalRegclass(defaultExpr.String) - } - - if !hasNextvalDefault || seqRef == "" { - seqName := tableName + "_id_seq" - if _, err := db.Exec(fmt.Sprintf("CREATE SEQUENCE IF NOT EXISTS %s.%s", quoteSQLIdentifier(schemaName), quoteSQLIdentifier(seqName))); err != nil { - return err - } - - seqRef = schemaName + "." + seqName - if _, err := db.Exec(fmt.Sprintf( - "ALTER TABLE %s.%s ALTER COLUMN id SET DEFAULT nextval(%s::regclass)", - quoteSQLIdentifier(schemaName), - quoteSQLIdentifier(tableName), - quoteSQLLiteral(seqRef), - )); err != nil { - return err - } - - if _, err := db.Exec(fmt.Sprintf( - "ALTER SEQUENCE %s.%s OWNED BY %s.%s.id", - quoteSQLIdentifier(schemaName), - quoteSQLIdentifier(seqName), - quoteSQLIdentifier(schemaName), - quoteSQLIdentifier(tableName), - )); err != nil { - return err - } - } - - return syncPostgresTableIDSequence(db, schemaName, tableName, seqRef) -} - -func syncPostgresTableIDSequence(db *store.DB, schemaName, tableName, seqRef string) error { - var maxID int64 - if err := db.QueryRow(fmt.Sprintf( - "SELECT COALESCE(MAX(id), 0) FROM %s.%s", - quoteSQLIdentifier(schemaName), - quoteSQLIdentifier(tableName), - )).Scan(&maxID); err != nil { - return err - } - - setVal := maxID - isCalled := true - if maxID <= 0 { - setVal = 1 - isCalled = false - } - - if _, err := db.Exec(`SELECT setval(?::regclass, ?, ?)`, seqRef, setVal, isCalled); err != nil { - return err - } - - return nil -} - -func extractNextvalRegclass(defaultExpr string) string { - nextvalIdx := strings.Index(strings.ToLower(defaultExpr), "nextval(") - if nextvalIdx < 0 { - return "" - } - expr := defaultExpr[nextvalIdx:] - firstQuote := strings.Index(expr, "'") - if firstQuote < 0 { - return "" - } - expr = expr[firstQuote+1:] - secondQuote := strings.Index(expr, "'") - if secondQuote < 0 { - return "" - } - return strings.TrimSpace(expr[:secondQuote]) -} - -func quoteSQLIdentifier(ident string) string { - return `"` + strings.ReplaceAll(ident, `"`, `""`) + `"` -} - -func quoteSQLLiteral(value string) string { - return "'" + strings.ReplaceAll(value, "'", "''") + "'" -} - -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 isMissingTableError(dialect store.Dialect, err error) bool { - if err == nil { - return false - } - msg := strings.ToLower(err.Error()) - if dialect == store.DialectPostgres { - return strings.Contains(msg, "relation") && strings.Contains(msg, "does not exist") - } - return strings.Contains(msg, "no such table") -} - -func (r *Repository) CreatePeerShare(share *PeerShare) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - _, err := r.db.Exec(` - INSERT INTO peer_share(name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, share.Name, share.NodeID, share.Token, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.CurrentFlow, share.IsActive, share.CreatedTime, share.UpdatedTime, share.AllowedDomains, share.AllowedIPs) - return err -} - -func (r *Repository) UpdatePeerShare(share *PeerShare) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - _, err := r.db.Exec(` - UPDATE peer_share SET name=?, max_bandwidth=?, expiry_time=?, port_range_start=?, port_range_end=?, is_active=?, updated_time=?, allowed_domains=?, allowed_ips=? - WHERE id=? - `, share.Name, share.MaxBandwidth, share.ExpiryTime, share.PortRangeStart, share.PortRangeEnd, share.IsActive, share.UpdatedTime, share.AllowedDomains, share.AllowedIPs, share.ID) - return err -} - -func (r *Repository) DeletePeerShare(id int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - tx, err := r.db.Begin() - if err != nil { - return err - } - defer func() { _ = tx.Rollback() }() - _, _ = tx.Exec(`DELETE FROM peer_share_runtime WHERE share_id = ?`, id) - if _, err := tx.Exec(`DELETE FROM peer_share WHERE id=?`, id); err != nil { - return err - } - return tx.Commit() -} - -func (r *Repository) GetPeerShare(id int64) (*PeerShare, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share WHERE id = ?`, id) - var s PeerShare - if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &s, nil -} - -func (r *Repository) GetPeerShareByToken(token string) (*PeerShare, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - row := r.db.QueryRow(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share WHERE token = ?`, token) - var s PeerShare - if err := row.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &s, nil -} - -func (r *Repository) ListPeerShares() ([]PeerShare, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - rows, err := r.db.Query(`SELECT id, name, node_id, token, max_bandwidth, expiry_time, port_range_start, port_range_end, current_flow, is_active, created_time, updated_time, allowed_domains, allowed_ips FROM peer_share ORDER BY id DESC`) - if err != nil { - return nil, err - } - defer rows.Close() - - var shares []PeerShare - for rows.Next() { - var s PeerShare - if err := rows.Scan(&s.ID, &s.Name, &s.NodeID, &s.Token, &s.MaxBandwidth, &s.ExpiryTime, &s.PortRangeStart, &s.PortRangeEnd, &s.CurrentFlow, &s.IsActive, &s.CreatedTime, &s.UpdatedTime, &s.AllowedDomains, &s.AllowedIPs); err != nil { - return nil, err - } - shares = append(shares, s) - } - return shares, nil -} - -func (r *Repository) GetPeerShareRuntimeByResourceKey(shareID int64, resourceKey string) (*PeerShareRuntime, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - row := r.db.QueryRow(` - SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time - FROM peer_share_runtime - WHERE share_id = ? AND resource_key = ? - LIMIT 1 - `, shareID, resourceKey) - var item PeerShareRuntime - if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &item, nil -} - -func (r *Repository) GetPeerShareRuntimeByReservationID(shareID int64, reservationID string) (*PeerShareRuntime, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - row := r.db.QueryRow(` - SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time - FROM peer_share_runtime - WHERE share_id = ? AND reservation_id = ? - LIMIT 1 - `, shareID, reservationID) - var item PeerShareRuntime - if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &item, nil -} - -func (r *Repository) GetPeerShareRuntimeByBindingID(shareID int64, bindingID string) (*PeerShareRuntime, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - row := r.db.QueryRow(` - SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time - FROM peer_share_runtime - WHERE share_id = ? AND binding_id = ? - LIMIT 1 - `, shareID, bindingID) - var item PeerShareRuntime - if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &item, nil -} - -func (r *Repository) GetPeerShareRuntimeByID(id int64) (*PeerShareRuntime, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - row := r.db.QueryRow(` - SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time - FROM peer_share_runtime - WHERE id = ? - LIMIT 1 - `, id) - var item PeerShareRuntime - if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - return nil, err - } - return &item, nil -} - -func (r *Repository) ListActivePeerShareRuntimesByShareID(shareID int64) ([]PeerShareRuntime, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - rows, err := r.db.Query(` - SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time - FROM peer_share_runtime - WHERE share_id = ? AND status = 1 - ORDER BY port ASC, id ASC - `, shareID) - if err != nil { - return nil, err - } - defer rows.Close() - - out := make([]PeerShareRuntime, 0) - for rows.Next() { - var item PeerShareRuntime - if err := rows.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { - return nil, err - } - out = append(out, item) - } - if err := rows.Err(); err != nil { - return nil, err - } - return out, nil -} - -func (r *Repository) AddPeerShareCurrentFlow(shareID int64, delta int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - if shareID <= 0 || delta <= 0 { - return nil - } - _, err := r.db.Exec(`UPDATE peer_share SET current_flow = current_flow + ?, updated_time = ? WHERE id = ?`, delta, unixMilliNow(), shareID) - return err -} - -func (r *Repository) ResetPeerShareCurrentFlow(shareID int64, updatedTime int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - if shareID <= 0 { - return nil - } - if updatedTime <= 0 { - updatedTime = unixMilliNow() - } - _, err := r.db.Exec(`UPDATE peer_share SET current_flow = 0, updated_time = ? WHERE id = ?`, updatedTime, shareID) - return err -} - -func (r *Repository) CreatePeerShareRuntime(item *PeerShareRuntime) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - if item == nil { - return errors.New("runtime item is nil") - } - _, err := r.db.Exec(` - INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, item.ShareID, item.NodeID, item.ReservationID, item.ResourceKey, item.BindingID, item.Role, item.ChainName, item.ServiceName, item.Protocol, item.Strategy, item.Port, item.Target, item.Applied, item.Status, item.CreatedTime, item.UpdatedTime) - return err -} - -func (r *Repository) UpdatePeerShareRuntime(item *PeerShareRuntime) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - if item == nil { - return errors.New("runtime item is nil") - } - _, err := r.db.Exec(` - UPDATE peer_share_runtime - SET binding_id = ?, role = ?, chain_name = ?, service_name = ?, protocol = ?, strategy = ?, port = ?, target = ?, applied = ?, status = ?, updated_time = ? - WHERE id = ? - `, item.BindingID, item.Role, item.ChainName, item.ServiceName, item.Protocol, item.Strategy, item.Port, item.Target, item.Applied, item.Status, item.UpdatedTime, item.ID) - return err -} - -func (r *Repository) MarkPeerShareRuntimeReleased(id int64, updatedTime int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - _, err := r.db.Exec(`UPDATE peer_share_runtime SET status = 0, updated_time = ? WHERE id = ?`, updatedTime, id) - return err -} - -func (r *Repository) ListActivePeerShareRuntimePorts(shareID int64, nodeID int64) ([]int, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - rows, err := r.db.Query(`SELECT port FROM peer_share_runtime WHERE share_id = ? AND node_id = ? AND status = 1 AND port > 0`, shareID, nodeID) - if err != nil { - return nil, err - } - defer rows.Close() - out := make([]int, 0) - for rows.Next() { - var port int - if err := rows.Scan(&port); err != nil { - return nil, err - } - if port > 0 { - out = append(out, port) - } - } - if err := rows.Err(); err != nil { - return nil, err - } - return out, nil -} - -func (r *Repository) UpsertFederationTunnelBinding(item *FederationTunnelBinding) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - if item == nil { - return errors.New("binding item is nil") - } - _, err := r.db.Exec(` - INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(tunnel_id, node_id, chain_type, hop_inx) - DO UPDATE SET - remote_url = excluded.remote_url, - resource_key = excluded.resource_key, - remote_binding_id = excluded.remote_binding_id, - allocated_port = excluded.allocated_port, - status = excluded.status, - updated_time = excluded.updated_time - `, item.TunnelID, item.NodeID, item.ChainType, item.HopInx, item.RemoteURL, item.ResourceKey, item.RemoteBindingID, item.AllocatedPort, item.Status, item.CreatedTime, item.UpdatedTime) - return err -} - -func (r *Repository) ListActiveFederationTunnelBindingsByTunnel(tunnelID int64) ([]FederationTunnelBinding, error) { - if r == nil || r.db == nil { - return nil, errors.New("repository not initialized") - } - rows, err := r.db.Query(` - SELECT id, tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time - FROM federation_tunnel_binding - WHERE tunnel_id = ? AND status = 1 - ORDER BY chain_type ASC, hop_inx ASC, id ASC - `, tunnelID) - if err != nil { - return nil, err - } - defer rows.Close() - out := make([]FederationTunnelBinding, 0) - for rows.Next() { - var item FederationTunnelBinding - if err := rows.Scan(&item.ID, &item.TunnelID, &item.NodeID, &item.ChainType, &item.HopInx, &item.RemoteURL, &item.ResourceKey, &item.RemoteBindingID, &item.AllocatedPort, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { - return nil, err - } - out = append(out, item) - } - if err := rows.Err(); err != nil { - return nil, err - } - return out, nil -} - -func (r *Repository) DeleteFederationTunnelBindingsByTunnel(tunnelID int64) error { - if r == nil || r.db == nil { - return errors.New("repository not initialized") - } - _, err := r.db.Exec(`DELETE FROM federation_tunnel_binding WHERE tunnel_id = ?`, tunnelID) - return err -} - -var osMkdirAll = func(path string) error { - return os.MkdirAll(path, 0o755) -} - -// ============ Backup/Export Data Structures ============ - -// BackupData represents the full backup structure -type BackupData struct { - Version string `json:"version"` - ExportedAt int64 `json:"exportedAt"` - Users []UserBackup `json:"users,omitempty"` - Nodes []NodeBackup `json:"nodes,omitempty"` - Tunnels []TunnelBackup `json:"tunnels,omitempty"` - Forwards []ForwardBackup `json:"forwards,omitempty"` - UserTunnels []UserTunnelBackup `json:"userTunnels,omitempty"` - SpeedLimits []SpeedLimitBackup `json:"speedLimits,omitempty"` - TunnelGroups []TunnelGroupBackup `json:"tunnelGroups,omitempty"` - UserGroups []UserGroupBackup `json:"userGroups,omitempty"` - Permissions []PermissionBackup `json:"permissions,omitempty"` - Configs map[string]string `json:"configs,omitempty"` -} - -type UserBackup struct { - ID int64 `json:"id"` - User string `json:"user"` - Pwd string `json:"pwd"` - RoleID int `json:"roleId"` - ExpTime int64 `json:"expTime"` - Flow int64 `json:"flow"` - InFlow int64 `json:"inFlow"` - OutFlow int64 `json:"outFlow"` - FlowResetTime int64 `json:"flowResetTime"` - Num int `json:"num"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime,omitempty"` - Status int `json:"status"` -} - -type NodeBackup struct { - ID int64 `json:"id"` - Name string `json:"name"` - Secret string `json:"secret"` - ServerIP string `json:"serverIp"` - ServerIPv4 string `json:"serverIpV4,omitempty"` - ServerIPv6 string `json:"serverIpV6,omitempty"` - Port string `json:"port"` - InterfaceName string `json:"interfaceName,omitempty"` - Version string `json:"version,omitempty"` - HTTP int `json:"http"` - TLS int `json:"tls"` - Socks int `json:"socks"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime,omitempty"` - Status int `json:"status"` - TCPListenAddr string `json:"tcpListenAddr"` - UDPListenAddr string `json:"udpListenAddr"` - Inx int `json:"inx"` - IsRemote int `json:"isRemote"` - RemoteURL string `json:"remoteUrl,omitempty"` - RemoteToken string `json:"remoteToken,omitempty"` - RemoteConfig string `json:"remoteConfig,omitempty"` -} - -type TunnelBackup struct { - ID int64 `json:"id"` - Name string `json:"name"` - TrafficRatio float64 `json:"trafficRatio"` - Type int `json:"type"` - Protocol string `json:"protocol"` - Flow int64 `json:"flow"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime"` - Status int `json:"status"` - InIP string `json:"inIp,omitempty"` - Inx int `json:"inx"` - IPPreference string `json:"ipPreference,omitempty"` - ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"` -} - -type ChainTunnelBackup struct { - ID int64 `json:"id"` - TunnelID int64 `json:"tunnelId"` - ChainType string `json:"chainType"` - NodeID int64 `json:"nodeId"` - Port int `json:"port,omitempty"` - Strategy string `json:"strategy,omitempty"` - Inx int `json:"inx,omitempty"` - Protocol string `json:"protocol,omitempty"` -} - -type ForwardBackup struct { - ID int64 `json:"id"` - UserID int64 `json:"userId"` - UserName string `json:"userName"` - Name string `json:"name"` - TunnelID int64 `json:"tunnelId"` - RemoteAddr string `json:"remoteAddr"` - Strategy string `json:"strategy"` - InFlow int64 `json:"inFlow"` - OutFlow int64 `json:"outFlow"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime"` - Status int `json:"status"` - Inx int `json:"inx"` - ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"` -} - -type ForwardPortBackup struct { - NodeID int64 `json:"nodeId"` - Port int `json:"port"` -} - -type UserTunnelBackup struct { - ID int64 `json:"id"` - UserID int64 `json:"userId"` - TunnelID int64 `json:"tunnelId"` - SpeedID int64 `json:"speedId,omitempty"` - Num int `json:"num"` - Flow int64 `json:"flow"` - InFlow int64 `json:"inFlow"` - OutFlow int64 `json:"outFlow"` - FlowResetTime int64 `json:"flowResetTime"` - ExpTime int64 `json:"expTime"` - Status int `json:"status"` -} - -type SpeedLimitBackup struct { - ID int64 `json:"id"` - Name string `json:"name"` - Speed int64 `json:"speed"` - TunnelID int64 `json:"tunnelId"` - TunnelName string `json:"tunnelName"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime,omitempty"` - Status int `json:"status"` -} - -type TunnelGroupBackup struct { - ID int64 `json:"id"` - Name string `json:"name"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime"` - Status int `json:"status"` - Tunnels []int64 `json:"tunnels,omitempty"` -} - -type UserGroupBackup struct { - ID int64 `json:"id"` - Name string `json:"name"` - CreatedTime int64 `json:"createdTime"` - UpdatedTime int64 `json:"updatedTime"` - Status int `json:"status"` - Users []int64 `json:"users,omitempty"` -} - -type PermissionBackup struct { - ID int64 `json:"id"` - UserGroupID int64 `json:"userGroupId"` - TunnelGroupID int64 `json:"tunnelGroupId"` - CreatedTime int64 `json:"createdTime"` - CreatedByGroup int `json:"createdByGroup"` - Grants []PermissionGrantBackup `json:"grants,omitempty"` -} - -type PermissionGrantBackup struct { - ID int64 `json:"id"` - UserGroupID int64 `json:"userGroupId"` - TunnelGroupID int64 `json:"tunnelGroupId"` - UserTunnelID int64 `json:"userTunnelId"` - CreatedTime int64 `json:"createdTime"` - CreatedByGroup int `json:"createdByGroup"` -} - -// ============ Export Methods ============ - -// ExportAll exports all data as BackupData -func (r *Repository) ExportAll() (*BackupData, error) { - backup := &BackupData{ - Version: "1.0", - ExportedAt: unixMilliNow(), - } - - // Export all data types - users, err := r.exportUsers() - if err != nil { - return nil, fmt.Errorf("export users failed: %w", err) - } - backup.Users = users - - nodes, err := r.exportNodes() - if err != nil { - return nil, fmt.Errorf("export nodes failed: %w", err) - } - backup.Nodes = nodes - - tunnels, err := r.exportTunnels() - if err != nil { - return nil, fmt.Errorf("export tunnels failed: %w", err) - } - backup.Tunnels = tunnels - - forwards, err := r.exportForwards() - if err != nil { - return nil, fmt.Errorf("export forwards failed: %w", err) - } - backup.Forwards = forwards - - userTunnels, err := r.exportUserTunnels() - if err != nil { - return nil, fmt.Errorf("export user tunnels failed: %w", err) - } - backup.UserTunnels = userTunnels - - speedLimits, err := r.exportSpeedLimits() - if err != nil { - return nil, fmt.Errorf("export speed limits failed: %w", err) - } - backup.SpeedLimits = speedLimits - - tunnelGroups, err := r.exportTunnelGroups() - if err != nil { - return nil, fmt.Errorf("export tunnel groups failed: %w", err) - } - backup.TunnelGroups = tunnelGroups - - userGroups, err := r.exportUserGroups() - if err != nil { - return nil, fmt.Errorf("export user groups failed: %w", err) - } - backup.UserGroups = userGroups - - permissions, err := r.exportPermissions() - if err != nil { - return nil, fmt.Errorf("export permissions failed: %w", err) - } - backup.Permissions = permissions - - configs, err := r.ListConfigs() - if err != nil { - return nil, fmt.Errorf("export configs failed: %w", err) - } - backup.Configs = configs - - return backup, nil -} - -// ExportPartial exports selected data types -func (r *Repository) ExportPartial(types []string) (*BackupData, error) { - backup := &BackupData{ - Version: "1.0", - ExportedAt: unixMilliNow(), - } - - typeSet := make(map[string]bool) - for _, t := range types { - typeSet[t] = true - } - - if typeSet["users"] { - users, err := r.exportUsers() - if err != nil { - return nil, fmt.Errorf("export users failed: %w", err) - } - backup.Users = users - } - if typeSet["nodes"] { - nodes, err := r.exportNodes() - if err != nil { - return nil, fmt.Errorf("export nodes failed: %w", err) - } - backup.Nodes = nodes - } - if typeSet["tunnels"] { - tunnels, err := r.exportTunnels() - if err != nil { - return nil, fmt.Errorf("export tunnels failed: %w", err) - } - backup.Tunnels = tunnels - } - if typeSet["forwards"] { - forwards, err := r.exportForwards() - if err != nil { - return nil, fmt.Errorf("export forwards failed: %w", err) - } - backup.Forwards = forwards - } - if typeSet["userTunnels"] { - userTunnels, err := r.exportUserTunnels() - if err != nil { - return nil, fmt.Errorf("export user tunnels failed: %w", err) - } - backup.UserTunnels = userTunnels - } - if typeSet["speedLimits"] { - speedLimits, err := r.exportSpeedLimits() - if err != nil { - return nil, fmt.Errorf("export speed limits failed: %w", err) - } - backup.SpeedLimits = speedLimits - } - if typeSet["tunnelGroups"] { - tunnelGroups, err := r.exportTunnelGroups() - if err != nil { - return nil, fmt.Errorf("export tunnel groups failed: %w", err) - } - backup.TunnelGroups = tunnelGroups - } - if typeSet["userGroups"] { - userGroups, err := r.exportUserGroups() - if err != nil { - return nil, fmt.Errorf("export user groups failed: %w", err) - } - backup.UserGroups = userGroups - } - if typeSet["permissions"] { - permissions, err := r.exportPermissions() - if err != nil { - return nil, fmt.Errorf("export permissions failed: %w", err) - } - backup.Permissions = permissions - } - if typeSet["configs"] { - configs, err := r.ListConfigs() - if err != nil { - return nil, fmt.Errorf("export configs failed: %w", err) - } - backup.Configs = configs - } - - return backup, nil -} - -func (r *Repository) exportUsers() ([]UserBackup, error) { - rows, err := r.db.Query(` - SELECT id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status - FROM user ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var users []UserBackup - for rows.Next() { - var u UserBackup - var updatedTime sql.NullInt64 - if err := rows.Scan(&u.ID, &u.User, &u.Pwd, &u.RoleID, &u.ExpTime, &u.Flow, &u.InFlow, &u.OutFlow, &u.FlowResetTime, &u.Num, &u.CreatedTime, &updatedTime, &u.Status); err != nil { - return nil, err - } - if updatedTime.Valid { - u.UpdatedTime = updatedTime.Int64 - } - users = append(users, u) - } - return users, rows.Err() -} - -func (r *Repository) exportNodes() ([]NodeBackup, error) { - rows, err := r.db.Query(` - SELECT id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config - FROM node ORDER BY inx ASC, id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var nodes []NodeBackup - for rows.Next() { - var n NodeBackup - var updatedTime sql.NullInt64 - var serverIPv4, serverIPv6, interfaceName, version, remoteURL, remoteToken, remoteConfig sql.NullString - if err := rows.Scan(&n.ID, &n.Name, &n.Secret, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Port, &interfaceName, &version, &n.HTTP, &n.TLS, &n.Socks, &n.CreatedTime, &updatedTime, &n.Status, &n.TCPListenAddr, &n.UDPListenAddr, &n.Inx, &n.IsRemote, &remoteURL, &remoteToken, &remoteConfig); err != nil { - return nil, err - } - if updatedTime.Valid { - n.UpdatedTime = updatedTime.Int64 - } - if serverIPv4.Valid { - n.ServerIPv4 = serverIPv4.String - } - if serverIPv6.Valid { - n.ServerIPv6 = serverIPv6.String - } - if interfaceName.Valid { - n.InterfaceName = interfaceName.String - } - if version.Valid { - n.Version = version.String - } - if remoteURL.Valid { - n.RemoteURL = remoteURL.String - } - if remoteToken.Valid { - n.RemoteToken = remoteToken.String - } - if remoteConfig.Valid { - n.RemoteConfig = remoteConfig.String - } - nodes = append(nodes, n) - } - return nodes, rows.Err() -} - -func (r *Repository) exportTunnels() ([]TunnelBackup, error) { - rows, err := r.db.Query(` - SELECT id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, COALESCE(ip_preference, '') - FROM tunnel ORDER BY inx ASC, id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var tunnels []TunnelBackup - for rows.Next() { - var t TunnelBackup - var protocol sql.NullString - var updatedTime sql.NullInt64 - var inIP sql.NullString - var inx sql.NullInt64 - if err := rows.Scan(&t.ID, &t.Name, &t.TrafficRatio, &t.Type, &protocol, &t.Flow, &t.CreatedTime, &updatedTime, &t.Status, &inIP, &inx, &t.IPPreference); err != nil { - return nil, err - } - if protocol.Valid { - t.Protocol = protocol.String - } - if updatedTime.Valid { - t.UpdatedTime = updatedTime.Int64 - } - if inIP.Valid { - t.InIP = inIP.String - } - if inx.Valid { - t.Inx = int(inx.Int64) - } - // Export chain tunnels - chainTunnels, err := r.exportChainTunnels(t.ID) - if err != nil { - return nil, err - } - t.ChainTunnels = chainTunnels - tunnels = append(tunnels, t) - } - return tunnels, rows.Err() -} - -func (r *Repository) exportChainTunnels(tunnelID int64) ([]ChainTunnelBackup, error) { - rows, err := r.db.Query(` - SELECT id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol - FROM chain_tunnel WHERE tunnel_id = ? ORDER BY inx ASC, id ASC - `, tunnelID) - if err != nil { - return nil, err - } - defer rows.Close() - - var chainTunnels []ChainTunnelBackup - for rows.Next() { - var ct ChainTunnelBackup - var port sql.NullInt64 - var strategy, protocol sql.NullString - var inx sql.NullInt64 - if err := rows.Scan(&ct.ID, &ct.TunnelID, &ct.ChainType, &ct.NodeID, &port, &strategy, &inx, &protocol); err != nil { - return nil, err - } - if port.Valid { - ct.Port = int(port.Int64) - } - if strategy.Valid { - ct.Strategy = strategy.String - } - if inx.Valid { - ct.Inx = int(inx.Int64) - } - if protocol.Valid { - ct.Protocol = protocol.String - } - chainTunnels = append(chainTunnels, ct) - } - return chainTunnels, rows.Err() -} - -func (r *Repository) exportForwards() ([]ForwardBackup, error) { - rows, err := r.db.Query(` - SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx - FROM forward ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var forwards []ForwardBackup - for rows.Next() { - var f ForwardBackup - var strategy sql.NullString - var updatedTime sql.NullInt64 - var inx sql.NullInt64 - if err := rows.Scan(&f.ID, &f.UserID, &f.UserName, &f.Name, &f.TunnelID, &f.RemoteAddr, &strategy, &f.InFlow, &f.OutFlow, &f.CreatedTime, &updatedTime, &f.Status, &inx); err != nil { - return nil, err - } - if strategy.Valid { - f.Strategy = strategy.String - } - if updatedTime.Valid { - f.UpdatedTime = updatedTime.Int64 - } - if inx.Valid { - f.Inx = int(inx.Int64) - } - - forwardPorts, err := r.exportForwardPorts(f.ID) - if err != nil { - return nil, err - } - portsCopy := append([]ForwardPortBackup(nil), forwardPorts...) - f.ForwardPorts = &portsCopy - - forwards = append(forwards, f) - } - return forwards, rows.Err() -} - -func (r *Repository) exportForwardPorts(forwardID int64) ([]ForwardPortBackup, error) { - rows, err := r.db.Query(` - SELECT node_id, port - FROM forward_port - WHERE forward_id = ? - ORDER BY id ASC - `, forwardID) - if err != nil { - return nil, err - } - defer rows.Close() - - ports := make([]ForwardPortBackup, 0) - for rows.Next() { - var fp ForwardPortBackup - if err := rows.Scan(&fp.NodeID, &fp.Port); err != nil { - return nil, err - } - ports = append(ports, fp) - } - - if err := rows.Err(); err != nil { - return nil, err - } - - return ports, nil -} - -func (r *Repository) exportUserTunnels() ([]UserTunnelBackup, error) { - rows, err := r.db.Query(` - SELECT id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status - FROM user_tunnel ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var userTunnels []UserTunnelBackup - for rows.Next() { - var ut UserTunnelBackup - var speedID sql.NullInt64 - if err := rows.Scan(&ut.ID, &ut.UserID, &ut.TunnelID, &speedID, &ut.Num, &ut.Flow, &ut.InFlow, &ut.OutFlow, &ut.FlowResetTime, &ut.ExpTime, &ut.Status); err != nil { - return nil, err - } - if speedID.Valid { - ut.SpeedID = speedID.Int64 - } - userTunnels = append(userTunnels, ut) - } - return userTunnels, rows.Err() -} - -func (r *Repository) exportSpeedLimits() ([]SpeedLimitBackup, error) { - rows, err := r.db.Query(` - SELECT id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status - FROM speed_limit ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var speedLimits []SpeedLimitBackup - for rows.Next() { - var sl SpeedLimitBackup - var updatedTime sql.NullInt64 - if err := rows.Scan(&sl.ID, &sl.Name, &sl.Speed, &sl.TunnelID, &sl.TunnelName, &sl.CreatedTime, &updatedTime, &sl.Status); err != nil { - return nil, err - } - if updatedTime.Valid { - sl.UpdatedTime = updatedTime.Int64 - } - speedLimits = append(speedLimits, sl) - } - return speedLimits, rows.Err() -} - -func (r *Repository) exportTunnelGroups() ([]TunnelGroupBackup, error) { - rows, err := r.db.Query(` - SELECT id, name, created_time, updated_time, status - FROM tunnel_group ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var groups []TunnelGroupBackup - for rows.Next() { - var tg TunnelGroupBackup - if err := rows.Scan(&tg.ID, &tg.Name, &tg.CreatedTime, &tg.UpdatedTime, &tg.Status); err != nil { - return nil, err - } - // Get tunnel IDs for this group - tunnelRows, err := r.db.Query(`SELECT tunnel_id FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) - if err != nil { - return nil, err - } - for tunnelRows.Next() { - var tunnelID int64 - if err := tunnelRows.Scan(&tunnelID); err != nil { - tunnelRows.Close() - return nil, err - } - tg.Tunnels = append(tg.Tunnels, tunnelID) - } - tunnelRows.Close() - groups = append(groups, tg) - } - return groups, rows.Err() -} - -func (r *Repository) exportUserGroups() ([]UserGroupBackup, error) { - rows, err := r.db.Query(` - SELECT id, name, created_time, updated_time, status - FROM user_group ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var groups []UserGroupBackup - for rows.Next() { - var ug UserGroupBackup - if err := rows.Scan(&ug.ID, &ug.Name, &ug.CreatedTime, &ug.UpdatedTime, &ug.Status); err != nil { - return nil, err - } - // Get user IDs for this group - userRows, err := r.db.Query(`SELECT user_id FROM user_group_user WHERE user_group_id = ?`, ug.ID) - if err != nil { - return nil, err - } - for userRows.Next() { - var userID int64 - if err := userRows.Scan(&userID); err != nil { - userRows.Close() - return nil, err - } - ug.Users = append(ug.Users, userID) - } - userRows.Close() - groups = append(groups, ug) - } - return groups, rows.Err() -} - -func (r *Repository) exportPermissions() ([]PermissionBackup, error) { - rows, err := r.db.Query(` - SELECT id, user_group_id, tunnel_group_id, created_time - FROM group_permission ORDER BY id ASC - `) - if err != nil { - return nil, err - } - defer rows.Close() - - var permissions []PermissionBackup - for rows.Next() { - var p PermissionBackup - if err := rows.Scan(&p.ID, &p.UserGroupID, &p.TunnelGroupID, &p.CreatedTime); err != nil { - return nil, err - } - p.CreatedByGroup = 0 - // Get grants for this permission - grantRows, err := r.db.Query(`SELECT id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group FROM group_permission_grant WHERE user_group_id = ? AND tunnel_group_id = ?`, p.UserGroupID, p.TunnelGroupID) - if err != nil { - return nil, err - } - for grantRows.Next() { - var g PermissionGrantBackup - if err := grantRows.Scan(&g.ID, &g.UserGroupID, &g.TunnelGroupID, &g.UserTunnelID, &g.CreatedTime, &g.CreatedByGroup); err != nil { - grantRows.Close() - return nil, err - } - p.Grants = append(p.Grants, g) - } - grantRows.Close() - permissions = append(permissions, p) - } - return permissions, rows.Err() -} - -// ============ Import Methods ============ - -// ImportResult contains the result of an import operation -type ImportResult struct { - UsersImported int `json:"usersImported"` - NodesImported int `json:"nodesImported"` - TunnelsImported int `json:"tunnelsImported"` - ForwardsImported int `json:"forwardsImported"` - UserTunnelsImported int `json:"userTunnelsImported"` - SpeedLimitsImported int `json:"speedLimitsImported"` - TunnelGroupsImported int `json:"tunnelGroupsImported"` - UserGroupsImported int `json:"userGroupsImported"` - PermissionsImported int `json:"permissionsImported"` - ConfigsImported int `json:"configsImported"` - AutoBackup *BackupData `json:"autoBackup,omitempty"` -} - -// Import imports data from BackupData with transaction support -func (r *Repository) Import(backup *BackupData, types []string) (*ImportResult, error) { - result := &ImportResult{} - - typeSet := make(map[string]bool) - for _, t := range types { - typeSet[t] = true - } - - tx, err := r.db.Begin() - if err != nil { - return nil, fmt.Errorf("failed to begin transaction: %w", err) - } - defer func() { _ = tx.Rollback() }() - - now := unixMilliNow() - - if typeSet["users"] && len(backup.Users) > 0 { - count, err := r.importUsers(tx, backup.Users, now) - if err != nil { - return nil, fmt.Errorf("import users failed: %w", err) - } - result.UsersImported = count - } - - if typeSet["nodes"] && len(backup.Nodes) > 0 { - count, err := r.importNodes(tx, backup.Nodes, now) - if err != nil { - return nil, fmt.Errorf("import nodes failed: %w", err) - } - result.NodesImported = count - } - - if typeSet["tunnels"] && len(backup.Tunnels) > 0 { - count, err := r.importTunnels(tx, backup.Tunnels, now) - if err != nil { - return nil, fmt.Errorf("import tunnels failed: %w", err) - } - result.TunnelsImported = count - } - - if typeSet["forwards"] && len(backup.Forwards) > 0 { - count, err := r.importForwards(tx, backup.Forwards, now) - if err != nil { - return nil, fmt.Errorf("import forwards failed: %w", err) - } - result.ForwardsImported = count - } - - if typeSet["userTunnels"] && len(backup.UserTunnels) > 0 { - count, err := r.importUserTunnels(tx, backup.UserTunnels, now) - if err != nil { - return nil, fmt.Errorf("import user tunnels failed: %w", err) - } - result.UserTunnelsImported = count - } - - if typeSet["speedLimits"] && len(backup.SpeedLimits) > 0 { - count, err := r.importSpeedLimits(tx, backup.SpeedLimits, now) - if err != nil { - return nil, fmt.Errorf("import speed limits failed: %w", err) - } - result.SpeedLimitsImported = count - } - - if typeSet["tunnelGroups"] && len(backup.TunnelGroups) > 0 { - count, err := r.importTunnelGroups(tx, backup.TunnelGroups, now) - if err != nil { - return nil, fmt.Errorf("import tunnel groups failed: %w", err) - } - result.TunnelGroupsImported = count - } - - if typeSet["userGroups"] && len(backup.UserGroups) > 0 { - count, err := r.importUserGroups(tx, backup.UserGroups, now) - if err != nil { - return nil, fmt.Errorf("import user groups failed: %w", err) - } - result.UserGroupsImported = count - } - - if typeSet["permissions"] && len(backup.Permissions) > 0 { - count, err := r.importPermissions(tx, backup.Permissions, now) - if err != nil { - return nil, fmt.Errorf("import permissions failed: %w", err) - } - result.PermissionsImported = count - } - - if typeSet["configs"] && len(backup.Configs) > 0 { - count, err := r.importConfigs(tx, backup.Configs, now) - if err != nil { - return nil, fmt.Errorf("import configs failed: %w", err) - } - result.ConfigsImported = count - } - - if err := tx.Commit(); err != nil { - return nil, fmt.Errorf("failed to commit transaction: %w", err) - } - - return result, nil -} - -func (r *Repository) importUsers(db Execer, users []UserBackup, now int64) (int, error) { - count := 0 - for _, u := range users { - _, err := db.Exec(` - 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - user = excluded.user, - pwd = excluded.pwd, - role_id = excluded.role_id, - exp_time = excluded.exp_time, - flow = excluded.flow, - in_flow = excluded.in_flow, - out_flow = excluded.out_flow, - flow_reset_time = excluded.flow_reset_time, - num = excluded.num, - updated_time = excluded.updated_time, - status = excluded.status - `, u.ID, u.User, u.Pwd, u.RoleID, u.ExpTime, u.Flow, u.InFlow, u.OutFlow, u.FlowResetTime, u.Num, u.CreatedTime, now, u.Status) - if err != nil { - return count, err - } - count++ - } - return count, nil -} - -func (r *Repository) UsernameExists(username string) (bool, error) { - var count int - err := r.db.QueryRow(`SELECT COUNT(1) FROM user WHERE user = ?`, username).Scan(&count) - if err != nil { - return false, err - } - return count > 0, nil -} - -func (r *Repository) importNodes(db Execer, nodes []NodeBackup, now int64) (int, error) { - count := 0 - for _, n := range nodes { - _, err := db.Exec(` - INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - name = excluded.name, - secret = excluded.secret, - server_ip = excluded.server_ip, - server_ip_v4 = excluded.server_ip_v4, - server_ip_v6 = excluded.server_ip_v6, - port = excluded.port, - interface_name = excluded.interface_name, - version = excluded.version, - http = excluded.http, - tls = excluded.tls, - socks = excluded.socks, - updated_time = excluded.updated_time, - status = excluded.status, - tcp_listen_addr = excluded.tcp_listen_addr, - udp_listen_addr = excluded.udp_listen_addr, - inx = excluded.inx, - is_remote = excluded.is_remote, - remote_url = excluded.remote_url, - remote_token = excluded.remote_token, - remote_config = excluded.remote_config - `, n.ID, n.Name, n.Secret, n.ServerIP, n.ServerIPv4, n.ServerIPv6, n.Port, n.InterfaceName, n.Version, n.HTTP, n.TLS, n.Socks, n.CreatedTime, now, n.Status, n.TCPListenAddr, n.UDPListenAddr, n.Inx, n.IsRemote, n.RemoteURL, n.RemoteToken, n.RemoteConfig) - if err != nil { - return count, err - } - count++ - } - return count, nil -} - -func (r *Repository) importTunnels(db Execer, tunnels []TunnelBackup, now int64) (int, error) { - count := 0 - for _, t := range tunnels { - _, err := db.Exec(` - INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, ip_preference) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - name = excluded.name, - traffic_ratio = excluded.traffic_ratio, - type = excluded.type, - protocol = excluded.protocol, - flow = excluded.flow, - updated_time = excluded.updated_time, - status = excluded.status, - in_ip = excluded.in_ip, - inx = excluded.inx, - ip_preference = excluded.ip_preference - `, t.ID, t.Name, t.TrafficRatio, t.Type, t.Protocol, t.Flow, t.CreatedTime, now, t.Status, t.InIP, t.Inx, t.IPPreference) - if err != nil { - return count, err - } - if len(t.ChainTunnels) > 0 { - for _, ct := range t.ChainTunnels { - _, err = db.Exec(` - INSERT INTO chain_tunnel(id, tunnel_id, chain_type, node_id, port, strategy, inx, protocol) - VALUES(?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - chain_type = excluded.chain_type, - node_id = excluded.node_id, - port = excluded.port, - strategy = excluded.strategy, - inx = excluded.inx, - protocol = excluded.protocol - `, ct.ID, ct.TunnelID, ct.ChainType, ct.NodeID, ct.Port, ct.Strategy, ct.Inx, ct.Protocol) - if err != nil { - return count, err - } - } - } - count++ - } - return count, nil -} - -func (r *Repository) importForwards(db Execer, forwards []ForwardBackup, now int64) (int, error) { - count := 0 - for _, f := range forwards { - _, err := db.Exec(` - INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - user_id = excluded.user_id, - user_name = excluded.user_name, - name = excluded.name, - tunnel_id = excluded.tunnel_id, - remote_addr = excluded.remote_addr, - strategy = excluded.strategy, - in_flow = excluded.in_flow, - out_flow = excluded.out_flow, - updated_time = excluded.updated_time, - status = excluded.status, - inx = excluded.inx - `, f.ID, f.UserID, f.UserName, f.Name, f.TunnelID, f.RemoteAddr, f.Strategy, f.InFlow, f.OutFlow, f.CreatedTime, now, f.Status, f.Inx) - if err != nil { - return count, err - } - - if f.ForwardPorts != nil { - if _, err := db.Exec(`DELETE FROM forward_port WHERE forward_id = ?`, f.ID); err != nil { - return count, err - } - for _, fp := range *f.ForwardPorts { - if _, err := db.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, f.ID, fp.NodeID, fp.Port); err != nil { - return count, err - } - } - } - - count++ - } - return count, nil -} - -func (r *Repository) importUserTunnels(db Execer, userTunnels []UserTunnelBackup, now int64) (int, error) { - count := 0 - for _, ut := range userTunnels { - var speedID interface{} - if ut.SpeedID > 0 { - speedID = ut.SpeedID - } - _, err := db.Exec(` - INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) - VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - user_id = excluded.user_id, - tunnel_id = excluded.tunnel_id, - speed_id = excluded.speed_id, - num = excluded.num, - flow = excluded.flow, - in_flow = excluded.in_flow, - out_flow = excluded.out_flow, - flow_reset_time = excluded.flow_reset_time, - exp_time = excluded.exp_time, - status = excluded.status - `, ut.ID, ut.UserID, ut.TunnelID, speedID, ut.Num, ut.Flow, ut.InFlow, ut.OutFlow, ut.FlowResetTime, ut.ExpTime, ut.Status) - if err != nil { - return count, err - } - count++ - } - return count, nil -} - -func (r *Repository) importSpeedLimits(db Execer, speedLimits []SpeedLimitBackup, now int64) (int, error) { - count := 0 - for _, sl := range speedLimits { - _, err := db.Exec(` - INSERT INTO speed_limit(id, name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) - VALUES(?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - name = excluded.name, - speed = excluded.speed, - tunnel_id = excluded.tunnel_id, - tunnel_name = excluded.tunnel_name, - updated_time = excluded.updated_time, - status = excluded.status - `, sl.ID, sl.Name, sl.Speed, sl.TunnelID, sl.TunnelName, sl.CreatedTime, now, sl.Status) - if err != nil { - return count, err - } - count++ - } - return count, nil -} - -func (r *Repository) importTunnelGroups(db Execer, tunnelGroups []TunnelGroupBackup, now int64) (int, error) { - count := 0 - for _, tg := range tunnelGroups { - _, err := db.Exec(` - INSERT INTO tunnel_group(id, name, created_time, updated_time, status) - VALUES(?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - name = excluded.name, - updated_time = excluded.updated_time, - status = excluded.status - `, tg.ID, tg.Name, tg.CreatedTime, now, tg.Status) - if err != nil { - return count, err - } - _, err = db.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, tg.ID) - if err != nil { - return count, err - } - for _, tunnelID := range tg.Tunnels { - _, err = db.Exec(` - INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) - VALUES(?, ?, ?) - `, tg.ID, tunnelID, now) - if err != nil { - return count, err - } - } - count++ - } - return count, nil -} - -func (r *Repository) importUserGroups(db Execer, userGroups []UserGroupBackup, now int64) (int, error) { - count := 0 - for _, ug := range userGroups { - _, err := db.Exec(` - INSERT INTO user_group(id, name, created_time, updated_time, status) - VALUES(?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - name = excluded.name, - updated_time = excluded.updated_time, - status = excluded.status - `, ug.ID, ug.Name, ug.CreatedTime, now, ug.Status) - if err != nil { - return count, err - } - _, err = db.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, ug.ID) - if err != nil { - return count, err - } - for _, userID := range ug.Users { - _, err = db.Exec(` - INSERT INTO user_group_user(user_group_id, user_id, created_time) - VALUES(?, ?, ?) - `, ug.ID, userID, now) - if err != nil { - return count, err - } - } - count++ - } - return count, nil -} - -func (r *Repository) importPermissions(db Execer, permissions []PermissionBackup, now int64) (int, error) { - count := 0 - for _, p := range permissions { - _, err := db.Exec(` - INSERT INTO group_permission(id, user_group_id, tunnel_group_id, created_time, created_by_group) - VALUES(?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - user_group_id = excluded.user_group_id, - tunnel_group_id = excluded.tunnel_group_id, - created_by_group = excluded.created_by_group - `, p.ID, p.UserGroupID, p.TunnelGroupID, p.CreatedTime, p.CreatedByGroup) - if err != nil { - return count, err - } - for _, g := range p.Grants { - _, err = db.Exec(` - INSERT INTO group_permission_grant(id, user_group_id, tunnel_group_id, user_tunnel_id, created_time, created_by_group) - VALUES(?, ?, ?, ?, ?, ?) - ON CONFLICT(id) DO UPDATE SET - user_tunnel_id = excluded.user_tunnel_id, - created_by_group = excluded.created_by_group - `, g.ID, g.UserGroupID, g.TunnelGroupID, g.UserTunnelID, g.CreatedTime, g.CreatedByGroup) - if err != nil { - return count, err - } - } - count++ - } - return count, nil -} - -func (r *Repository) importConfigs(db Execer, configs map[string]string, now int64) (int, error) { - count := 0 - for name, value := range configs { - err := r.UpsertConfig(name, value, now) - if err != nil { - return count, err - } - count++ - } - return count, nil -} diff --git a/go-backend/internal/store/sqlite/sql/data.sql b/go-backend/internal/store/sqlite/sql/data.sql deleted file mode 100644 index ed78932..0000000 --- a/go-backend/internal/store/sqlite/sql/data.sql +++ /dev/null @@ -1,5 +0,0 @@ -INSERT OR IGNORE 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); - -INSERT OR IGNORE INTO vite_config (id, name, value, time) -VALUES (1, 'app_name', 'flux', 1755147963000); diff --git a/go-backend/internal/store/sqlite/sql/schema.sql b/go-backend/internal/store/sqlite/sql/schema.sql deleted file mode 100644 index 68677ab..0000000 --- a/go-backend/internal/store/sqlite/sql/schema.sql +++ /dev/null @@ -1,254 +0,0 @@ --- SQLite Auto-generated schema --- This will be executed automatically on startup if tables don't exist - -CREATE TABLE IF NOT EXISTS forward ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - 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 INTEGER NOT NULL DEFAULT 0, - out_flow INTEGER NOT NULL DEFAULT 0, - created_time INTEGER NOT NULL, - updated_time INTEGER NOT NULL, - status INTEGER NOT NULL, - inx INTEGER NOT NULL DEFAULT 0 -); - -CREATE TABLE IF NOT EXISTS forward_port ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - forward_id INTEGER NOT NULL, - node_id INTEGER NOT NULL, - port INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS node ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name VARCHAR(100) NOT NULL, - secret VARCHAR(100) NOT NULL, - server_ip VARCHAR(100) NOT NULL, - 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 INTEGER NOT NULL, - updated_time INTEGER, - status INTEGER NOT NULL, - tcp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]', - udp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]', - 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 INTEGER PRIMARY KEY AUTOINCREMENT, - name VARCHAR(100) NOT NULL, - speed INTEGER NOT NULL, - tunnel_id INTEGER NOT NULL, - tunnel_name VARCHAR(100) NOT NULL, - created_time INTEGER NOT NULL, - updated_time INTEGER, - status INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS statistics_flow ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - user_id INTEGER NOT NULL, - flow INTEGER NOT NULL, - total_flow INTEGER NOT NULL, - time VARCHAR(100) NOT NULL, - created_time INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS tunnel ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name VARCHAR(100) NOT NULL, - traffic_ratio REAL NOT NULL DEFAULT 1.0, - type INTEGER NOT NULL, - protocol VARCHAR(10) NOT NULL DEFAULT 'tls', - flow INTEGER NOT NULL, - created_time INTEGER NOT NULL, - updated_time INTEGER NOT NULL, - status INTEGER NOT NULL, - in_ip TEXT, - inx INTEGER NOT NULL DEFAULT 0, - ip_preference VARCHAR(10) NOT NULL DEFAULT '' -); - -CREATE TABLE IF NOT EXISTS chain_tunnel ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - 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 INTEGER PRIMARY KEY AUTOINCREMENT, - user VARCHAR(100) NOT NULL, - pwd VARCHAR(100) NOT NULL, - role_id INTEGER NOT NULL, - exp_time INTEGER NOT NULL, - flow INTEGER NOT NULL, - in_flow INTEGER NOT NULL DEFAULT 0, - out_flow INTEGER NOT NULL DEFAULT 0, - flow_reset_time INTEGER NOT NULL, - num INTEGER NOT NULL, - created_time INTEGER NOT NULL, - updated_time INTEGER, - status INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS user_tunnel ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - user_id INTEGER NOT NULL, - tunnel_id INTEGER NOT NULL, - speed_id INTEGER, - num INTEGER NOT NULL, - flow INTEGER NOT NULL, - in_flow INTEGER NOT NULL DEFAULT 0, - out_flow INTEGER NOT NULL DEFAULT 0, - flow_reset_time INTEGER NOT NULL, - exp_time INTEGER NOT NULL, - status INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS tunnel_group ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name VARCHAR(100) NOT NULL, - created_time INTEGER NOT NULL, - updated_time INTEGER NOT NULL, - status INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS user_group ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name VARCHAR(100) NOT NULL, - created_time INTEGER NOT NULL, - updated_time INTEGER NOT NULL, - status INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS tunnel_group_tunnel ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - tunnel_group_id INTEGER NOT NULL, - tunnel_id INTEGER NOT NULL, - created_time INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS user_group_user ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - user_group_id INTEGER NOT NULL, - user_id INTEGER NOT NULL, - created_time INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS group_permission ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - user_group_id INTEGER NOT NULL, - tunnel_group_id INTEGER NOT NULL, - created_time INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS group_permission_grant ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - 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 INTEGER 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 INTEGER PRIMARY KEY AUTOINCREMENT, - name VARCHAR(200) NOT NULL UNIQUE, - value VARCHAR(200) NOT NULL, - time INTEGER NOT NULL -); - -CREATE TABLE IF NOT EXISTS peer_share ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - name TEXT NOT NULL, - node_id INTEGER NOT NULL, - token TEXT NOT NULL UNIQUE, - max_bandwidth INTEGER DEFAULT 0, - expiry_time INTEGER DEFAULT 0, - port_range_start INTEGER DEFAULT 0, - port_range_end INTEGER DEFAULT 0, - current_flow INTEGER DEFAULT 0, - is_active INTEGER DEFAULT 1, - created_time INTEGER NOT NULL, - updated_time INTEGER NOT NULL, - allowed_domains TEXT DEFAULT '', - allowed_ips TEXT DEFAULT '' -); - -CREATE TABLE IF NOT EXISTS peer_share_runtime ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - 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 INTEGER NOT NULL, - updated_time INTEGER 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 INTEGER PRIMARY KEY AUTOINCREMENT, - 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 INTEGER NOT NULL, - updated_time INTEGER 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); - -CREATE TABLE IF NOT EXISTS announcement ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - content TEXT NOT NULL, - enabled INTEGER NOT NULL DEFAULT 1, - created_time INTEGER NOT NULL, - updated_time INTEGER -); diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go index 54edfa0..63e194c 100644 --- a/go-backend/internal/ws/server.go +++ b/go-backend/internal/ws/server.go @@ -15,7 +15,7 @@ import ( "go-backend/internal/auth" "go-backend/internal/security" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) type encryptedMessage struct { @@ -68,7 +68,7 @@ type CommandResult struct { } type Server struct { - repo *sqlite.Repository + repo *repo.Repository jwtSecret string upgrader websocket.Upgrader @@ -79,7 +79,7 @@ type Server struct { pending map[string]pendingRequest } -func NewServer(repo *sqlite.Repository, jwtSecret string) *Server { +func NewServer(repo *repo.Repository, jwtSecret string) *Server { return &Server{ repo: repo, jwtSecret: jwtSecret, diff --git a/go-backend/tests/contract/diagnosis_contract_test.go b/go-backend/tests/contract/diagnosis_contract_test.go index 621b55d..bedec1b 100644 --- a/go-backend/tests/contract/diagnosis_contract_test.go +++ b/go-backend/tests/contract/diagnosis_contract_test.go @@ -16,43 +16,41 @@ import ( httpserver "go-backend/internal/http" "go-backend/internal/http/handler" "go-backend/internal/http/response" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) func TestDiagnosisChainCoverageContracts(t *testing.T) { secret := "contract-jwt-secret" - router, repo := setupDiagnosisContractRouter(t, secret) + router, r := setupDiagnosisContractRouter(t, secret) now := time.Now().UnixMilli() - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` 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(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert user: %v", err) } - tunnelRes, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "diagnose-chain-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0) - if err != nil { + `, "diagnose-chain-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := tunnelRes.LastInsertId() - if err != nil { + var tunnelID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("get tunnel id: %v", err) } insertNode := func(name, ip string) int64 { - res, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get node id %s: %v", name, err) } return id @@ -62,34 +60,33 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) { chainNodeID := insertNode("chain-node", "10.0.1.20") exitNodeID := insertNode("exit-node", "10.0.1.30") - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 30001, 'round', 1, 'tls') - `, tunnelID, entryNodeID); err != nil { + `, tunnelID, entryNodeID).Error; err != nil { t.Fatalf("insert entry chain: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, 30002, 'round', 1, 'tls') - `, tunnelID, chainNodeID); err != nil { + `, tunnelID, chainNodeID).Error; err != nil { t.Fatalf("insert middle chain: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, 30003, 'round', 1, 'tls') - `, tunnelID, exitNodeID); err != nil { + `, tunnelID, exitNodeID).Error; err != nil { t.Fatalf("insert exit chain: %v", err) } - forwardRes, err := repo.DB().Exec(` + if err := r.DB().Exec(` 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, ?) - `, 2, "normal_user", "chain-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 0) - if err != nil { + `, 2, "normal_user", "chain-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 0).Error; err != nil { t.Fatalf("insert forward: %v", err) } - forwardID, err := forwardRes.LastInsertId() - if err != nil { + var forwardID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil { t.Fatalf("get forward id: %v", err) } @@ -208,7 +205,7 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) { func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) { secret := "contract-jwt-secret" - router, repo := setupDiagnosisContractRouter(t, secret) + router, r := setupDiagnosisContractRouter(t, secret) now := time.Now().UnixMilli() remoteToken := "remote-diagnose-token" @@ -256,30 +253,28 @@ func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) { defer remoteServer.Close() insertLocalNode := func(name, ip string) int64 { - res, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert local node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get local node id %s: %v", name, err) } return id } insertRemoteNode := func(name, ip string) int64 { - res, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, is_remote, remote_url, remote_token, remote_config) VALUES(?, ?, ?, ?, ?, ?, ?, ?, 0, 0, 0, ?, ?, 1, ?, ?, ?, 1, ?, ?, ?) - `, name, name+"-secret", ip, "", "", "31000-31010", "", "", now, now, "[::]", "[::]", 1, remoteServer.URL, remoteToken, `{"shareId": 123}`) - if err != nil { + `, name, name+"-secret", ip, "", "", "31000-31010", "", "", now, now, "[::]", "[::]", 1, remoteServer.URL, remoteToken, `{"shareId": 123}`).Error; err != nil { t.Fatalf("insert remote node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get remote node id %s: %v", name, err) } return id @@ -289,34 +284,33 @@ func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) { remoteChainNodeID := insertRemoteNode("middle-remote", "10.50.0.20") exitNodeID := insertLocalNode("exit-local", "10.50.0.30") - tunnelRes, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "diagnose-remote-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0) - if err != nil { + `, "diagnose-remote-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := tunnelRes.LastInsertId() - if err != nil { + var tunnelID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("get tunnel id: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 30001, 'round', 1, 'tls') - `, tunnelID, entryNodeID); err != nil { + `, tunnelID, entryNodeID).Error; err != nil { t.Fatalf("insert entry chain: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, 30002, 'round', 1, 'tls') - `, tunnelID, remoteChainNodeID); err != nil { + `, tunnelID, remoteChainNodeID).Error; err != nil { t.Fatalf("insert middle chain: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, 30003, 'round', 1, 'tls') - `, tunnelID, exitNodeID); err != nil { + `, tunnelID, exitNodeID).Error; err != nil { t.Fatalf("insert exit chain: %v", err) } @@ -409,17 +403,17 @@ func valueAsBool(v interface{}) bool { } } -func setupDiagnosisContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) { +func setupDiagnosisContractRouter(t *testing.T, jwtSecret string) (http.Handler, *repo.Repository) { t.Helper() dbPath := filepath.Join(t.TempDir(), "diagnosis-contract.db") - repo, err := sqlite.Open(dbPath) + r, err := repo.Open(dbPath) if err != nil { t.Fatalf("open sqlite: %v", err) } t.Cleanup(func() { - _ = repo.Close() + _ = r.Close() }) - h := handler.New(repo, jwtSecret) - return httpserver.NewRouter(h, jwtSecret), repo + h := handler.New(r, jwtSecret) + return httpserver.NewRouter(h, jwtSecret), r } diff --git a/go-backend/tests/contract/federation_dual_panel_contract_test.go b/go-backend/tests/contract/federation_dual_panel_contract_test.go index f5e3d49..0008104 100644 --- a/go-backend/tests/contract/federation_dual_panel_contract_test.go +++ b/go-backend/tests/contract/federation_dual_panel_contract_test.go @@ -17,7 +17,7 @@ import ( "go-backend/internal/auth" "go-backend/internal/http/response" "go-backend/internal/security" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { @@ -39,7 +39,7 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { providerMiddleNodeID := insertContractNode(t, providerRepo, "provider-middle", "198.51.100.12", "44000-44010", "provider-middle-secret", 1) providerExitNodeID := insertContractNode(t, providerRepo, "provider-exit", "198.51.100.13", "45000-45010", "provider-exit-secret", 1) - entryShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + entryShareID := insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "entry-share", NodeID: providerEntryNodeID, Token: "share-entry-token", @@ -49,7 +49,7 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { CreatedTime: now, UpdatedTime: now, }) - middleShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + middleShareID := insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "middle-share", NodeID: providerMiddleNodeID, Token: "share-middle-token", @@ -59,7 +59,7 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { CreatedTime: now, UpdatedTime: now, }) - exitShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + exitShareID := insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "exit-share", NodeID: providerExitNodeID, Token: "share-exit-token", @@ -113,7 +113,7 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) { assertCode(t, res, 0) var tunnelID int64 - if err := consumerRepo.DB().QueryRow(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name).Scan(&tunnelID); err != nil { + if err := consumerRepo.DB().Raw(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name).Row().Scan(&tunnelID); err != nil { t.Fatalf("query tunnel id (%s): %v", name, err) } if tunnelID <= 0 { @@ -191,7 +191,7 @@ func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) { providerMiddleNodeID := insertContractNode(t, providerRepo, "provider-middle-dx", "203.0.113.12", "54000-54010", "provider-middle-dx-secret", 1) providerExitNodeID := insertContractNode(t, providerRepo, "provider-exit-dx", "203.0.113.13", "55000-55010", "provider-exit-dx-secret", 1) - entryShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + entryShareID := insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "entry-share-dx", NodeID: providerEntryNodeID, Token: "share-entry-dx-token", @@ -201,7 +201,7 @@ func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) { CreatedTime: now, UpdatedTime: now, }) - middleShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + middleShareID := insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "middle-share-dx", NodeID: providerMiddleNodeID, Token: "share-middle-dx-token", @@ -211,7 +211,7 @@ func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) { CreatedTime: now, UpdatedTime: now, }) - exitShareID := insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + exitShareID := insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "exit-share-dx", NodeID: providerExitNodeID, Token: "share-exit-dx-token", @@ -262,7 +262,7 @@ func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) { assertCode(t, createRes, 0) var tunnelID int64 - if err := consumerRepo.DB().QueryRow(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, "dual-panel-diagnose-remote").Scan(&tunnelID); err != nil { + if err := consumerRepo.DB().Raw(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, "dual-panel-diagnose-remote").Row().Scan(&tunnelID); err != nil { t.Fatalf("query tunnel id: %v", err) } if tunnelID <= 0 { @@ -335,7 +335,7 @@ func TestFederationDualPanelRemoteEntryRuntimeContract(t *testing.T) { providerMiddleNodeID := insertContractNode(t, providerRepo, "provider-middle-rt", "198.51.100.22", "44020-44030", "provider-middle-rt-secret", 1) providerExitNodeID := insertContractNode(t, providerRepo, "provider-exit-rt", "198.51.100.23", "45020-45030", "provider-exit-rt-secret", 1) - insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "entry-share-rt", NodeID: providerEntryNodeID, Token: "share-entry-rt-token", @@ -345,7 +345,7 @@ func TestFederationDualPanelRemoteEntryRuntimeContract(t *testing.T) { CreatedTime: now, UpdatedTime: now, }) - insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "middle-share-rt", NodeID: providerMiddleNodeID, Token: "share-middle-rt-token", @@ -355,7 +355,7 @@ func TestFederationDualPanelRemoteEntryRuntimeContract(t *testing.T) { CreatedTime: now, UpdatedTime: now, }) - insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "exit-share-rt", NodeID: providerExitNodeID, Token: "share-exit-rt-token", @@ -415,7 +415,7 @@ func TestFederationDualPanelRemoteEntryRuntimeContract(t *testing.T) { assertCode(t, res, 0) var tunnelID int64 - if err := consumerRepo.DB().QueryRow(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name).Scan(&tunnelID); err != nil { + if err := consumerRepo.DB().Raw(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, name).Row().Scan(&tunnelID); err != nil { t.Fatalf("query tunnel id (%s): %v", name, err) } if tunnelID <= 0 { @@ -446,32 +446,31 @@ func TestFederationDualPanelRemoteEntryRuntimeContract(t *testing.T) { createTunnel("dual-panel-remote-entry-offline") } -func insertContractNode(t *testing.T, repo *sqlite.Repository, name, ip, portRange, secret string, status int) int64 { +func insertContractNode(t *testing.T, r *repo.Repository, name, ip, portRange, secret string, status int) int64 { t.Helper() now := time.Now().UnixMilli() - res, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, secret, ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0) - if err != nil { + `, name, secret, ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, status, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("node id %s: %v", name, err) } return id } -func insertPeerShare(t *testing.T, repo *sqlite.Repository, share *sqlite.PeerShare) int64 { +func insertPeerShare(t *testing.T, r *repo.Repository, share *repo.PeerShare) int64 { t.Helper() if share == nil { t.Fatalf("share is nil") } - if err := repo.CreatePeerShare(share); err != nil { + if err := r.CreatePeerShare(share); err != nil { t.Fatalf("create peer share %s: %v", share.Name, err) } - saved, err := repo.GetPeerShareByToken(share.Token) + saved, err := r.GetPeerShareByToken(share.Token) if err != nil { t.Fatalf("query peer share %s: %v", share.Name, err) } @@ -498,10 +497,10 @@ func importRemoteNodeForContract(t *testing.T, router http.Handler, adminToken, assertCode(t, res, 0) } -func queryRemoteNodeIDByToken(t *testing.T, repo *sqlite.Repository, token string) int64 { +func queryRemoteNodeIDByToken(t *testing.T, r *repo.Repository, token string) int64 { t.Helper() var id int64 - if err := repo.DB().QueryRow(`SELECT id FROM node WHERE is_remote = 1 AND remote_token = ? ORDER BY id DESC LIMIT 1`, token).Scan(&id); err != nil { + if err := r.DB().Raw(`SELECT id FROM node WHERE is_remote = 1 AND remote_token = ? ORDER BY id DESC LIMIT 1`, token).Row().Scan(&id); err != nil { t.Fatalf("query remote node by token %s: %v", token, err) } if id <= 0 { @@ -510,15 +509,15 @@ func queryRemoteNodeIDByToken(t *testing.T, repo *sqlite.Repository, token strin return id } -func assertTunnelPortInRange(t *testing.T, repo *sqlite.Repository, tunnelID int64, chainType int, nodeID int64, minPort int, maxPort int) { +func assertTunnelPortInRange(t *testing.T, r *repo.Repository, tunnelID int64, chainType int, nodeID int64, minPort int, maxPort int) { t.Helper() var port int - err := repo.DB().QueryRow(` + err := r.DB().Raw(` SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = ? AND node_id = ? LIMIT 1 - `, tunnelID, chainType, nodeID).Scan(&port) + `, tunnelID, chainType, nodeID).Row().Scan(&port) if err != nil { t.Fatalf("query tunnel=%d chainType=%d node=%d port: %v", tunnelID, chainType, nodeID, err) } @@ -527,10 +526,10 @@ func assertTunnelPortInRange(t *testing.T, repo *sqlite.Repository, tunnelID int } } -func assertCount(t *testing.T, repo *sqlite.Repository, query string, arg interface{}, expected int) { +func assertCount(t *testing.T, r *repo.Repository, query string, arg interface{}, expected int) { t.Helper() var got int - if err := repo.DB().QueryRow(query, arg).Scan(&got); err != nil { + if err := r.DB().Raw(query, arg).Row().Scan(&got); err != nil { t.Fatalf("count query failed: %v", err) } if got != expected { @@ -638,12 +637,12 @@ func startMockNodeSessionWithHook(t *testing.T, baseURL string, nodeSecret strin } } -func waitNodeStatus(t *testing.T, repo *sqlite.Repository, nodeID int64, expectedStatus int) { +func waitNodeStatus(t *testing.T, r *repo.Repository, nodeID int64, expectedStatus int) { t.Helper() deadline := time.Now().Add(2 * time.Second) for { var status int - if err := repo.DB().QueryRow(`SELECT status FROM node WHERE id = ?`, nodeID).Scan(&status); err == nil && status == expectedStatus { + if err := r.DB().Raw(`SELECT status FROM node WHERE id = ?`, nodeID).Row().Scan(&status); err == nil && status == expectedStatus { return } if time.Now().After(deadline) { @@ -698,7 +697,7 @@ func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) { now := time.Now().UnixMilli() providerNodeID := insertContractNode(t, providerRepo, "provider-portrange-node", "198.51.100.50", "44000-44010", "provider-portrange-secret", 1) - insertPeerShare(t, providerRepo, &sqlite.PeerShare{ + insertPeerShare(t, providerRepo, &repo.PeerShare{ Name: "portrange-share", NodeID: providerNodeID, Token: "share-portrange-token", diff --git a/go-backend/tests/contract/forward_contract_test.go b/go-backend/tests/contract/forward_contract_test.go index b438761..8204e0f 100644 --- a/go-backend/tests/contract/forward_contract_test.go +++ b/go-backend/tests/contract/forward_contract_test.go @@ -18,65 +18,61 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) { router, repo := setupContractRouter(t, secret) now := time.Now().UnixMilli() - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert user: %v", err) } - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "contract-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0) - if err != nil { + `, "contract-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := res.LastInsertId() - if err != nil { + var tunnelID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("get tunnel id: %v", err) } - nodeRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "entry-node", "entry-secret", "10.0.0.10", "10.0.0.10", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, "entry-node", "entry-secret", "10.0.0.10", "10.0.0.10", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert node: %v", err) } - entryNodeID, err := nodeRes.LastInsertId() - if err != nil { + var entryNodeID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&entryNodeID); err != nil { t.Fatalf("get node id: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 20001, 'round', 1, 'tls') - `, tunnelID, entryNodeID); err != nil { + `, tunnelID, entryNodeID).Error; err != nil { t.Fatalf("insert chain_tunnel: %v", err) } - resAdmin, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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, ?) - `, 1, "admin_user", "admin-forward", tunnelID, "1.1.1.1:443", "fifo", now, now, 0) - if err != nil { + `, 1, "admin_user", "admin-forward", tunnelID, "1.1.1.1:443", "fifo", now, now, 0).Error; err != nil { t.Fatalf("insert admin forward: %v", err) } - adminForwardID, err := resAdmin.LastInsertId() - if err != nil { + var adminForwardID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&adminForwardID); err != nil { t.Fatalf("get admin forward id: %v", err) } - resUser, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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, ?) - `, 2, "normal_user", "user-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 1) - if err != nil { + `, 2, "normal_user", "user-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 1).Error; err != nil { t.Fatalf("insert user forward: %v", err) } - userForwardID, err := resUser.LastInsertId() - if err != nil { + var userForwardID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&userForwardID); err != nil { t.Fatalf("get user forward id: %v", err) } @@ -207,38 +203,36 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) { t.Fatalf("generate admin token: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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(2, 'switch_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert user: %v", err) } insertTunnel := func(name string, inx int) int64 { - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, 1.0, 1, "tls", 99999, now, now, 1, nil, inx) - if err != nil { + `, name, 1.0, 1, "tls", 99999, now, now, 1, nil, inx).Error; err != nil { t.Fatalf("insert tunnel %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get tunnel id %s: %v", name, err) } return id } insertNode := func(name, ip, portRange string, inx int) int64 { - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", inx) - if err != nil { + `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", inx).Error; err != nil { t.Fatalf("insert node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get node id %s: %v", name, err) } return id @@ -249,45 +243,44 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) { nodeA := insertNode("switch-node-a", "10.10.0.1", "21000-21010", 0) nodeB := insertNode("switch-node-b", "10.10.0.2", "22000-22010", 1) - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 21001, 'round', 1, 'tls') - `, tunnelA, nodeA); err != nil { + `, tunnelA, nodeA).Error; err != nil { t.Fatalf("insert chain_tunnel tunnelA: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 22001, 'round', 1, 'tls') - `, tunnelB, nodeB); err != nil { + `, tunnelB, nodeB).Error; err != nil { t.Fatalf("insert chain_tunnel tunnelB: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1) - `, tunnelA); err != nil { + `, tunnelA).Error; err != nil { t.Fatalf("insert user_tunnel A: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(11, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1) - `, tunnelB); err != nil { + `, tunnelB).Error; err != nil { t.Fatalf("insert user_tunnel B: %v", err) } - forwardRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(2, 'switch_user', 'switch-forward', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0) - `, tunnelA, now, now) - if err != nil { + `, tunnelA, now, now).Error; err != nil { t.Fatalf("insert forward: %v", err) } - forwardID, err := forwardRes.LastInsertId() - if err != nil { + var forwardID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil { t.Fatalf("get forward id: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeA, 21001); err != nil { + if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeA, 21001).Error; err != nil { t.Fatalf("insert forward_port: %v", err) } @@ -308,7 +301,7 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) { } var tunnelAfter int64 - if err := repo.DB().QueryRow(`SELECT tunnel_id FROM forward WHERE id = ?`, forwardID).Scan(&tunnelAfter); err != nil { + if err := repo.DB().Raw(`SELECT tunnel_id FROM forward WHERE id = ?`, forwardID).Row().Scan(&tunnelAfter); err != nil { t.Fatalf("query forward tunnel_id: %v", err) } if tunnelAfter != tunnelA { @@ -317,7 +310,7 @@ func TestForwardSwitchTunnelRollbackOnSyncFailure(t *testing.T) { var nodeAfter int64 var portAfter int - if err := repo.DB().QueryRow(`SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID).Scan(&nodeAfter, &portAfter); err != nil { + if err := repo.DB().Raw(`SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID).Row().Scan(&nodeAfter, &portAfter); err != nil { t.Fatalf("query forward_port: %v", err) } if nodeAfter != nodeA || portAfter != 21001 { @@ -335,73 +328,83 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) { t.Fatalf("generate admin token: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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(2, 'batch_switch_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert user: %v", err) } - tunnelResA, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES('batch-switch-tunnel-a', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) - `, now, now) - if err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert tunnel A: %v", err) } - tunnelA, _ := tunnelResA.LastInsertId() + var tunnelA int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelA); err != nil { + t.Fatal(err) + } - tunnelResB, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES('batch-switch-tunnel-b', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 1) - `, now, now) - if err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert tunnel B: %v", err) } - tunnelB, _ := tunnelResB.LastInsertId() + var tunnelB int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelB); err != nil { + t.Fatal(err) + } - nodeResA, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES('batch-switch-node-a', 'batch-switch-node-a-secret', '10.11.0.1', '10.11.0.1', '', '23000-23010', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0) - `, now, now) - if err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert node A: %v", err) } - nodeA, _ := nodeResA.LastInsertId() + var nodeA int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&nodeA); err != nil { + t.Fatal(err) + } - nodeResB, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES('batch-switch-node-b', 'batch-switch-node-b-secret', '10.11.0.2', '10.11.0.2', '', '24000-24010', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 1) - `, now, now) - if err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert node B: %v", err) } - nodeB, _ := nodeResB.LastInsertId() + var nodeB int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&nodeB); err != nil { + t.Fatal(err) + } - if _, err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 23001, 'round', 1, 'tls')`, tunnelA, nodeA); err != nil { + if err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 23001, 'round', 1, 'tls')`, tunnelA, nodeA).Error; err != nil { t.Fatalf("insert chain_tunnel A: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 24001, 'round', 1, 'tls')`, tunnelB, nodeB); err != nil { + if err := repo.DB().Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, 24001, 'round', 1, 'tls')`, tunnelB, nodeB).Error; err != nil { t.Fatalf("insert chain_tunnel B: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(20, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnelA); err != nil { + if err := repo.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(20, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnelA).Error; err != nil { t.Fatalf("insert user_tunnel A: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(21, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnelB); err != nil { + if err := repo.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(21, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnelB).Error; err != nil { t.Fatalf("insert user_tunnel B: %v", err) } - forwardRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(2, 'batch_switch_user', 'batch-switch-forward', ?, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0) - `, tunnelA, now, now) - if err != nil { + `, tunnelA, now, now).Error; err != nil { t.Fatalf("insert forward: %v", err) } - forwardID, _ := forwardRes.LastInsertId() + var forwardID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil { + t.Fatal(err) + } - if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeA, 23001); err != nil { + if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeA, 23001).Error; err != nil { t.Fatalf("insert forward_port: %v", err) } @@ -430,7 +433,7 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) { } var tunnelAfter int64 - if err := repo.DB().QueryRow(`SELECT tunnel_id FROM forward WHERE id = ?`, forwardID).Scan(&tunnelAfter); err != nil { + if err := repo.DB().Raw(`SELECT tunnel_id FROM forward WHERE id = ?`, forwardID).Row().Scan(&tunnelAfter); err != nil { t.Fatalf("query forward tunnel_id: %v", err) } if tunnelAfter != tunnelA { @@ -439,7 +442,7 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) { var nodeAfter int64 var portAfter int - if err := repo.DB().QueryRow(`SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID).Scan(&nodeAfter, &portAfter); err != nil { + if err := repo.DB().Raw(`SELECT node_id, port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID).Row().Scan(&nodeAfter, &portAfter); err != nil { t.Fatalf("query forward_port: %v", err) } if nodeAfter != nodeA || portAfter != 23001 { @@ -457,21 +460,23 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) { t.Fatalf("generate admin token: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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(100, 'stable_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert user: %v", err) } - tunnelRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES('stable-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) - `, now, now) - if err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, _ := tunnelRes.LastInsertId() + var tunnelID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { + t.Fatal(err) + } // 1. Assign permission (creates new user_tunnel) // userTunnelBatchAssign expects structure: {userId: 123, tunnels: [{tunnelId: 456, ...}]} @@ -491,7 +496,7 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) { } var initialID int64 - if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(&initialID); err != nil { + if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Row().Scan(&initialID); err != nil { t.Fatalf("query initial user_tunnel id: %v", err) } @@ -513,7 +518,7 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) { // 3. Verify stable ID and no duplicates var count int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(&count); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Row().Scan(&count); err != nil { t.Fatalf("query count: %v", err) } if count != 1 { @@ -521,7 +526,7 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) { } var currentID int64 - if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(¤tID); err != nil { + if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Row().Scan(¤tID); err != nil { t.Fatalf("query current user_tunnel: %v", err) } diff --git a/go-backend/tests/contract/group_permission_contract_test.go b/go-backend/tests/contract/group_permission_contract_test.go index 530acdd..127648e 100644 --- a/go-backend/tests/contract/group_permission_contract_test.go +++ b/go-backend/tests/contract/group_permission_contract_test.go @@ -15,47 +15,44 @@ func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) { router, repo := setupContractRouter(t, secret) now := time.Now().UnixMilli() - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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(200, 'group_user_contract', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert test user: %v", err) } - tunnelRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES('group-contract-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) - `, now, now) - if err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := tunnelRes.LastInsertId() - if err != nil { + var tunnelID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("read tunnel id: %v", err) } - ugRes, err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-contract', ?, ?, 1)`, now, now) - if err != nil { + if err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-contract', ?, ?, 1)`, now, now).Error; err != nil { t.Fatalf("insert user_group: %v", err) } - userGroupID, err := ugRes.LastInsertId() - if err != nil { + var userGroupID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&userGroupID); err != nil { t.Fatalf("read user_group id: %v", err) } - tgRes, err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-contract', ?, ?, 1)`, now, now) - if err != nil { + if err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-contract', ?, ?, 1)`, now, now).Error; err != nil { t.Fatalf("insert tunnel_group: %v", err) } - tunnelGroupID, err := tgRes.LastInsertId() - if err != nil { + var tunnelGroupID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelGroupID); err != nil { t.Fatalf("read tunnel_group id: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, tunnelGroupID, tunnelID, now); err != nil { + if err := repo.DB().Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?)`, tunnelGroupID, tunnelID, now).Error; err != nil { t.Fatalf("insert tunnel_group_tunnel: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, userGroupID, tunnelGroupID, now); err != nil { + if err := repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?)`, userGroupID, tunnelGroupID, now).Error; err != nil { t.Fatalf("insert group_permission: %v", err) } @@ -71,12 +68,12 @@ func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) { assertCode(t, bindRes, 0) var userTunnelID int64 - if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 200 AND tunnel_id = ?`, tunnelID).Scan(&userTunnelID); err != nil { + if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 200 AND tunnel_id = ?`, tunnelID).Row().Scan(&userTunnelID); err != nil { t.Fatalf("query user_tunnel after bind: %v", err) } var grantCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil { t.Fatalf("query group_permission_grant after bind: %v", err) } if grantCount == 0 { @@ -89,7 +86,7 @@ func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) { router.ServeHTTP(unbindRes, unbindReq) assertCode(t, unbindRes, 0) - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil { t.Fatalf("query group_permission_grant after unbind: %v", err) } if grantCount != 0 { @@ -97,7 +94,7 @@ func TestGroupUserUnbindRevokesInheritedTunnelPermission(t *testing.T) { } var userTunnelCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Scan(&userTunnelCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Row().Scan(&userTunnelCount); err != nil { t.Fatalf("query user_tunnel after unbind: %v", err) } if userTunnelCount != 0 { @@ -110,40 +107,37 @@ func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) { router, repo := setupContractRouter(t, secret) now := time.Now().UnixMilli() - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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(201, 'group_user_permission_remove', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert test user: %v", err) } - tunnelRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES('group-remove-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0) - `, now, now) - if err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := tunnelRes.LastInsertId() - if err != nil { + var tunnelID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("read tunnel id: %v", err) } - ugRes, err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-remove-contract', ?, ?, 1)`, now, now) - if err != nil { + if err := repo.DB().Exec(`INSERT INTO user_group(name, created_time, updated_time, status) VALUES('ug-remove-contract', ?, ?, 1)`, now, now).Error; err != nil { t.Fatalf("insert user_group: %v", err) } - userGroupID, err := ugRes.LastInsertId() - if err != nil { + var userGroupID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&userGroupID); err != nil { t.Fatalf("read user_group id: %v", err) } - tgRes, err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-remove-contract', ?, ?, 1)`, now, now) - if err != nil { + if err := repo.DB().Exec(`INSERT INTO tunnel_group(name, created_time, updated_time, status) VALUES('tg-remove-contract', ?, ?, 1)`, now, now).Error; err != nil { t.Fatalf("insert tunnel_group: %v", err) } - tunnelGroupID, err := tgRes.LastInsertId() - if err != nil { + var tunnelGroupID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelGroupID); err != nil { t.Fatalf("read tunnel_group id: %v", err) } @@ -171,17 +165,17 @@ func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) { assertCode(t, assignPermissionRes, 0) var permissionID int64 - if err := repo.DB().QueryRow(`SELECT id FROM group_permission WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID).Scan(&permissionID); err != nil { + if err := repo.DB().Raw(`SELECT id FROM group_permission WHERE user_group_id = ? AND tunnel_group_id = ?`, userGroupID, tunnelGroupID).Row().Scan(&permissionID); err != nil { t.Fatalf("query group_permission id: %v", err) } var userTunnelID int64 - if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 201 AND tunnel_id = ?`, tunnelID).Scan(&userTunnelID); err != nil { + if err := repo.DB().Raw(`SELECT id FROM user_tunnel WHERE user_id = 201 AND tunnel_id = ?`, tunnelID).Row().Scan(&userTunnelID); err != nil { t.Fatalf("query user_tunnel after assign: %v", err) } var grantCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil { t.Fatalf("query group_permission_grant after assign: %v", err) } if grantCount == 0 { @@ -195,14 +189,14 @@ func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) { assertCode(t, removeRes, 0) var permissionCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission WHERE id = ?`, permissionID).Scan(&permissionCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission WHERE id = ?`, permissionID).Row().Scan(&permissionCount); err != nil { t.Fatalf("query group_permission after remove: %v", err) } if permissionCount != 0 { t.Fatalf("expected group_permission removed, got %d", permissionCount) } - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Scan(&grantCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM group_permission_grant WHERE user_tunnel_id = ?`, userTunnelID).Row().Scan(&grantCount); err != nil { t.Fatalf("query group_permission_grant after remove: %v", err) } if grantCount != 0 { @@ -210,7 +204,7 @@ func TestGroupPermissionRemoveRevokesInheritedTunnelPermission(t *testing.T) { } var userTunnelCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Scan(&userTunnelCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM user_tunnel WHERE id = ?`, userTunnelID).Row().Scan(&userTunnelCount); err != nil { t.Fatalf("query user_tunnel after permission remove: %v", err) } if userTunnelCount != 0 { diff --git a/go-backend/tests/contract/migration_contract_test.go b/go-backend/tests/contract/migration_contract_test.go index e58c7ae..710c118 100644 --- a/go-backend/tests/contract/migration_contract_test.go +++ b/go-backend/tests/contract/migration_contract_test.go @@ -17,22 +17,20 @@ 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" + "go-backend/internal/store/repo" - _ "modernc.org/sqlite" + "gorm.io/gorm" ) func TestCaptchaVerifyLoginContract(t *testing.T) { secret := "contract-jwt-secret" - router, repo := setupContractRouter(t, secret) + router, r := setupContractRouter(t, secret) - _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO vite_config(name, value, time) VALUES(?, ?, ?) ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time - `, "captcha_enabled", "true", time.Now().UnixMilli()) - if err != nil { + `, "captcha_enabled", "true", time.Now().UnixMilli()).Error; err != nil { t.Fatalf("enable captcha: %v", err) } @@ -84,7 +82,7 @@ func TestCaptchaVerifyLoginContract(t *testing.T) { } func TestOpenAPISubStoreContracts(t *testing.T) { - router, repo := setupContractRouter(t, "contract-jwt-secret") + router, r := setupContractRouter(t, "contract-jwt-secret") const tunnelFlowGB = int64(500) const tunnelInFlow = int64(123) @@ -92,17 +90,16 @@ func TestOpenAPISubStoreContracts(t *testing.T) { const tunnelExpTimeMs = int64(2727251700000) now := time.Now().UnixMilli() - res, err := repo.DB().Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - "contract-tunnel", 1.0, 1, "tls", 1, now, now, 1, nil, 0) - if err != nil { + if err := r.DB().Exec(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + "contract-tunnel", 1.0, 1, "tls", 1, now, now, 1, nil, 0).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := res.LastInsertId() - if err != nil { + var tunnelID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("last insert id: %v", err) } - if _, err := repo.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, ?, ?, ?, ?, ?, ?, ?)`, - 1, tunnelID, 99999, tunnelFlowGB, tunnelInFlow, tunnelOutFlow, 1, tunnelExpTimeMs, 1); err != nil { + if err := r.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, ?, ?, ?, ?, ?, ?, ?)`, + 1, tunnelID, 99999, tunnelFlowGB, tunnelInFlow, tunnelOutFlow, 1, tunnelExpTimeMs, 1).Error; err != nil { t.Fatalf("insert user_tunnel: %v", err) } @@ -205,7 +202,7 @@ func TestSpeedLimitTunnelsRouteAlias(t *testing.T) { func TestBackupExportImportRestoreContracts(t *testing.T) { secret := "contract-jwt-secret" - router, repo := setupContractRouter(t, secret) + router, r := setupContractRouter(t, secret) adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret) if err != nil { @@ -217,11 +214,11 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { } key := "backup_contract_key" - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO vite_config(name, value, time) VALUES(?, ?, ?) ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time - `, key, "v1", time.Now().UnixMilli()); err != nil { + `, key, "v1", time.Now().UnixMilli()).Error; err != nil { t.Fatalf("seed config for backup contract: %v", err) } @@ -271,7 +268,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { t.Fatalf("expected import code 0, got %d (%s)", out.Code, out.Msg) } - cfg, err := repo.GetConfigByName(key) + cfg, err := r.GetConfigByName(key) if err != nil { t.Fatalf("query imported config: %v", err) } @@ -302,7 +299,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { t.Fatalf("expected restore code 0, got %d (%s)", out.Code, out.Msg) } - cfg, err := repo.GetConfigByName(key) + cfg, err := r.GetConfigByName(key) if err != nil { t.Fatalf("query restored config: %v", err) } @@ -314,27 +311,25 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { t.Run("backup export and import preserve forward ports", func(t *testing.T) { now := time.Now().UnixMilli() - tunnelRes, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "backup-forward-tunnel", 1.0, 1, "tls", 0, now, now, 1, "", 88) - if err != nil { + `, "backup-forward-tunnel", 1.0, 1, "tls", 0, now, now, 1, "", 88).Error; err != nil { t.Fatalf("seed tunnel for forward backup: %v", err) } - tunnelID, err := tunnelRes.LastInsertId() - if err != nil { + var tunnelID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("read tunnel id for forward backup: %v", err) } - forwardRes, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88) - if err != nil { + `, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88).Error; err != nil { t.Fatalf("seed forward for backup: %v", err) } - forwardID, err := forwardRes.LastInsertId() - if err != nil { + var forwardID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&forwardID); err != nil { t.Fatalf("read forward id for backup: %v", err) } @@ -343,7 +338,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { 2002: 21002, } for nodeID, port := range expected { - if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port); err != nil { + if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port).Error; err != nil { t.Fatalf("seed forward_port %d:%d: %v", nodeID, port, err) } } @@ -420,10 +415,10 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { } } - if _, err := repo.DB().Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID); err != nil { + if err := r.DB().Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID).Error; err != nil { t.Fatalf("clear forward_port before import: %v", err) } - if _, err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, 9999, 39999); err != nil { + if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, 9999, 39999).Error; err != nil { t.Fatalf("seed wrong forward_port before import: %v", err) } @@ -447,7 +442,7 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { t.Fatalf("expected forwards import code 0, got %d (%s)", out.Code, out.Msg) } - rows, err := repo.DB().Query(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID) + rows, err := r.DB().Raw(`SELECT node_id, port FROM forward_port WHERE forward_id = ? ORDER BY id ASC`, forwardID).Rows() if err != nil { t.Fatalf("query forward ports after import: %v", err) } @@ -478,22 +473,21 @@ func TestBackupExportImportRestoreContracts(t *testing.T) { t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) { now := time.Now().UnixMilli() - res, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "legacy-null-chain", 1.0, 1, "tls", 1000, now, now, 1, nil, 1) - if err != nil { + `, "legacy-null-chain", 1.0, 1, "tls", 1000, now, now, 1, nil, 1).Error; err != nil { t.Fatalf("seed tunnel for nullable chain export: %v", err) } - tunnelID, err := res.LastInsertId() - if err != nil { + var tunnelID int64 + if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("read tunnel id for nullable chain export: %v", err) } - if _, err := repo.DB().Exec(` + if err := r.DB().Exec(` INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, ?, ?, ?, ?, ?, ?) - `, tunnelID, "1", 1, nil, nil, nil, nil); err != nil { + `, tunnelID, "1", 1, nil, nil, nil, nil).Error; err != nil { t.Fatalf("seed nullable chain_tunnel row: %v", err) } @@ -596,19 +590,19 @@ func exportBackupPayload(t *testing.T, router http.Handler, path, token string) return payload } -func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) { +func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *repo.Repository) { t.Helper() dbPath := filepath.Join(t.TempDir(), "contract.db") - repo, err := sqlite.Open(dbPath) + r, err := repo.Open(dbPath) if err != nil { t.Fatalf("open sqlite: %v", err) } t.Cleanup(func() { - _ = repo.Close() + _ = r.Close() }) - h := handler.New(repo, jwtSecret) - return httpserver.NewRouter(h, jwtSecret), repo + h := handler.New(r, jwtSecret) + return httpserver.NewRouter(h, jwtSecret), r } func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) { @@ -669,15 +663,15 @@ func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) { t.Fatalf("seed legacy node row: %v", err) } - repo, err := sqlite.Open(dbPath) + r, err := repo.Open(dbPath) if err != nil { t.Fatalf("open migrated sqlite: %v", err) } t.Cleanup(func() { - _ = repo.Close() + _ = r.Close() }) - nodes, err := repo.ListNodes() + nodes, err := r.ListNodes() if err != nil { t.Fatalf("list nodes after migration: %v", err) } @@ -685,7 +679,7 @@ func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) { t.Fatalf("expected 1 node after migration, got %d", len(nodes)) } - columns := readTableColumns(t, repo.DB(), "node") + columns := readTableColumns(t, r.DB(), "node") for _, required := range []string{"server_ip_v4", "server_ip_v6", "inx"} { if !columns[required] { @@ -693,16 +687,16 @@ func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) { } } - tunnelColumns := readTableColumns(t, repo.DB(), "tunnel") + tunnelColumns := readTableColumns(t, r.DB(), "tunnel") if !tunnelColumns["inx"] { t.Fatalf("expected tunnel column %q to exist after migration", "inx") } } -func readTableColumns(t *testing.T, db *store.DB, table string) map[string]bool { +func readTableColumns(t *testing.T, db *gorm.DB, table string) map[string]bool { t.Helper() - rows, err := db.Query("PRAGMA table_info(" + table + ")") + rows, err := db.Raw("PRAGMA table_info(" + table + ")").Rows() if err != nil { t.Fatalf("inspect %s columns: %v", table, err) } diff --git a/go-backend/tests/contract/postgres_node_id_repair_contract_test.go b/go-backend/tests/contract/postgres_node_id_repair_contract_test.go index 2659473..bdc6655 100644 --- a/go-backend/tests/contract/postgres_node_id_repair_contract_test.go +++ b/go-backend/tests/contract/postgres_node_id_repair_contract_test.go @@ -16,7 +16,7 @@ import ( "go-backend/internal/auth" httpserver "go-backend/internal/http" "go-backend/internal/http/handler" - "go-backend/internal/store/sqlite" + "go-backend/internal/store/repo" ) func TestPostgresNodeCreateRepairsMissingIDDefaultContract(t *testing.T) { @@ -44,35 +44,35 @@ func TestPostgresNodeCreateRepairsMissingIDDefaultContract(t *testing.T) { t.Fatalf("build schema dsn: %v", err) } - repo, err := sqlite.OpenPostgres(testDSN) + r, err := repo.OpenPostgres(testDSN) if err != nil { t.Fatalf("open postgres repository: %v", err) } - if _, err := repo.DB().Exec(`ALTER TABLE node ALTER COLUMN id DROP DEFAULT`); err != nil { - _ = repo.Close() + if err := r.DB().Exec(`ALTER TABLE node ALTER COLUMN id DROP DEFAULT`).Error; err != nil { + _ = r.Close() t.Fatalf("drop node.id default to simulate drift: %v", err) } - if err := repo.Close(); err != nil { + if err := r.Close(); err != nil { t.Fatalf("close repository before reopen: %v", err) } - repo, err = sqlite.OpenPostgres(testDSN) + r, err = repo.OpenPostgres(testDSN) if err != nil { t.Fatalf("reopen postgres repository: %v", err) } t.Cleanup(func() { - _ = repo.Close() + _ = r.Close() }) var columnDefault sql.NullString - if err := repo.DB().QueryRow(` + if err := r.DB().Raw(` SELECT column_default FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = 'node' AND column_name = 'id' LIMIT 1 - `).Scan(&columnDefault); err != nil { + `).Row().Scan(&columnDefault); err != nil { t.Fatalf("query node.id default: %v", err) } if !columnDefault.Valid || !strings.Contains(strings.ToLower(columnDefault.String), "nextval(") { @@ -80,7 +80,7 @@ func TestPostgresNodeCreateRepairsMissingIDDefaultContract(t *testing.T) { } jwtSecret := "postgres-contract-secret" - router := httpserver.NewRouter(handler.New(repo, jwtSecret), jwtSecret) + router := httpserver.NewRouter(handler.New(r, jwtSecret), jwtSecret) token, err := auth.GenerateToken(1, "admin_user", 0, jwtSecret) if err != nil { t.Fatalf("generate admin token: %v", err) @@ -95,7 +95,7 @@ func TestPostgresNodeCreateRepairsMissingIDDefaultContract(t *testing.T) { assertCode(t, resp, 0) var nodeID int64 - if err := repo.DB().QueryRow(`SELECT id FROM node WHERE name = ? ORDER BY id DESC LIMIT 1`, "pg-repair-node").Scan(&nodeID); err != nil { + if err := r.DB().Raw(`SELECT id FROM node WHERE name = ? ORDER BY id DESC LIMIT 1`, "pg-repair-node").Row().Scan(&nodeID); err != nil { t.Fatalf("query created node: %v", err) } if nodeID <= 0 { diff --git a/go-backend/tests/contract/tunnel_create_contract_test.go b/go-backend/tests/contract/tunnel_create_contract_test.go index c1e8d86..4d3ce62 100644 --- a/go-backend/tests/contract/tunnel_create_contract_test.go +++ b/go-backend/tests/contract/tunnel_create_contract_test.go @@ -26,15 +26,14 @@ func TestTunnelCreateRuntimeRollbackContract(t *testing.T) { } insertNode := func(name, ip, portRange string) int64 { - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get node id %s: %v", name, err) } return id @@ -64,7 +63,7 @@ func TestTunnelCreateRuntimeRollbackContract(t *testing.T) { } var tunnelCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, "runtime-rollback-tunnel").Scan(&tunnelCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM tunnel WHERE name = ?`, "runtime-rollback-tunnel").Row().Scan(&tunnelCount); err != nil { t.Fatalf("count tunnel: %v", err) } if tunnelCount != 0 { @@ -72,7 +71,7 @@ func TestTunnelCreateRuntimeRollbackContract(t *testing.T) { } var chainCount int - if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM chain_tunnel`).Scan(&chainCount); err != nil { + if err := repo.DB().Raw(`SELECT COUNT(1) FROM chain_tunnel`).Row().Scan(&chainCount); err != nil { t.Fatalf("count chain_tunnel: %v", err) } if chainCount != 0 { @@ -91,15 +90,14 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) { } insertNode := func(name, ip, portRange string) int64 { - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get node id %s: %v", name, err) } return id @@ -109,15 +107,14 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) { chainID := insertNode("update-chain", "10.30.0.2", "41000-41010") exitID := insertNode("update-exit", "10.30.0.3", "42000-42010") - tunnelRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "update-port-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0) - if err != nil { + `, "update-port-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := tunnelRes.LastInsertId() - if err != nil { + var tunnelID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("get tunnel id: %v", err) } @@ -131,7 +128,7 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) { assertCode(t, res, 0) var chainPort int - if err := repo.DB().QueryRow(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 2 LIMIT 1`, tunnelID).Scan(&chainPort); err != nil { + if err := repo.DB().Raw(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 2 LIMIT 1`, tunnelID).Row().Scan(&chainPort); err != nil { t.Fatalf("query chain port: %v", err) } if chainPort <= 0 { @@ -139,7 +136,7 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) { } var outPort int - if err := repo.DB().QueryRow(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 3 LIMIT 1`, tunnelID).Scan(&outPort); err != nil { + if err := repo.DB().Raw(`SELECT port FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 3 LIMIT 1`, tunnelID).Row().Scan(&outPort); err != nil { t.Fatalf("query out port: %v", err) } if outPort <= 0 { @@ -147,7 +144,7 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) { } var entryStrategy sql.NullString - if err := repo.DB().QueryRow(`SELECT strategy FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 LIMIT 1`, tunnelID).Scan(&entryStrategy); err != nil { + if err := repo.DB().Raw(`SELECT strategy FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 LIMIT 1`, tunnelID).Row().Scan(&entryStrategy); err != nil { t.Fatalf("query entry strategy: %v", err) } if !entryStrategy.Valid || strings.TrimSpace(entryStrategy.String) == "" { diff --git a/go-backend/tests/contract/tunnel_ip_preference_contract_test.go b/go-backend/tests/contract/tunnel_ip_preference_contract_test.go index a98cf8d..6871958 100644 --- a/go-backend/tests/contract/tunnel_ip_preference_contract_test.go +++ b/go-backend/tests/contract/tunnel_ip_preference_contract_test.go @@ -24,15 +24,14 @@ func TestTunnelCreateWithIPPreferenceContract(t *testing.T) { } insertDualStackNode := func(name, v4, v6, portRange string) int64 { - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", v4, v4, v6, portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, name, name+"-secret", v4, v4, v6, portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get node id %s: %v", name, err) } return id @@ -63,7 +62,7 @@ func TestTunnelCreateWithIPPreferenceContract(t *testing.T) { } var stored string - err := repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "tunnel-"+tc.name).Scan(&stored) + err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "tunnel-"+tc.name).Row().Scan(&stored) if err != nil { if err == sql.ErrNoRows { t.Skipf("tunnel not created (nodes offline), skipping DB verification") @@ -88,15 +87,14 @@ func TestTunnelUpdateIPPreferenceContract(t *testing.T) { } insertDualStackNode := func(name, v4, v6, portRange string) int64 { - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, name+"-secret", v4, v4, v6, portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, name, name+"-secret", v4, v4, v6, portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert node %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get node id %s: %v", name, err) } return id @@ -105,15 +103,14 @@ func TestTunnelUpdateIPPreferenceContract(t *testing.T) { entryID := insertDualStackNode("upd-entry", "10.60.0.1", "2001:db8:1::1", "60000-60010") exitID := insertDualStackNode("upd-exit", "10.60.0.2", "2001:db8:1::2", "61000-61010") - tunnelRes, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, ip_preference) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "update-ip-pref-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0, "") - if err != nil { + `, "update-ip-pref-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0, "").Error; err != nil { t.Fatalf("insert tunnel: %v", err) } - tunnelID, err := tunnelRes.LastInsertId() - if err != nil { + var tunnelID int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&tunnelID); err != nil { t.Fatalf("get tunnel id: %v", err) } @@ -130,7 +127,7 @@ func TestTunnelUpdateIPPreferenceContract(t *testing.T) { } var stored string - if err := repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE id = ?`, tunnelID).Scan(&stored); err != nil { + if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE id = ?`, tunnelID).Row().Scan(&stored); err != nil { t.Fatalf("query ip_preference: %v", err) } if stored != "v6" { @@ -148,10 +145,10 @@ func TestTunnelListReturnsIPPreferenceContract(t *testing.T) { t.Fatalf("generate admin token: %v", err) } - _, err = repo.DB().Exec(` + err = repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, ip_preference) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "list-ip-pref-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0, "v6") + `, "list-ip-pref-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0, "v6").Error if err != nil { t.Fatalf("insert tunnel: %v", err) } @@ -198,16 +195,15 @@ func TestIPPreferenceColumnDefaultContract(t *testing.T) { _, repo := setupContractRouter(t, "contract-jwt-secret") now := time.Now().UnixMilli() - _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "no-pref-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0) - if err != nil { + `, "no-pref-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { t.Fatalf("insert tunnel without ip_preference: %v", err) } var stored string - if err := repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "no-pref-tunnel").Scan(&stored); err != nil { + if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "no-pref-tunnel").Row().Scan(&stored); err != nil { t.Fatalf("query ip_preference: %v", err) } if stored != "" { @@ -219,7 +215,7 @@ func TestIPPreferenceColumnMigrationContract(t *testing.T) { _, repo := setupContractRouter(t, "contract-jwt-secret") var colCount int - err := repo.DB().QueryRow(`SELECT COUNT(*) FROM pragma_table_info('tunnel') WHERE name = 'ip_preference'`).Scan(&colCount) + err := repo.DB().Raw(`SELECT COUNT(*) FROM pragma_table_info('tunnel') WHERE name = 'ip_preference'`).Row().Scan(&colCount) if err != nil { t.Fatalf("check column existence: %v", err) } @@ -232,16 +228,15 @@ func TestIPPreferenceCoalesceNullSafety(t *testing.T) { _, repo := setupContractRouter(t, "contract-jwt-secret") now := time.Now().UnixMilli() - _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, ip_preference) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL) - `, "null-pref-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0) - if err != nil { + `, "null-pref-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil { t.Skipf("DB does not allow NULL ip_preference (NOT NULL constraint): %v", err) } var stored string - if err := repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "null-pref-tunnel").Scan(&stored); err != nil { + if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "null-pref-tunnel").Row().Scan(&stored); err != nil { t.Fatalf("query ip_preference: %v", err) } if stored != "" { @@ -253,16 +248,15 @@ func TestDualStackNodeIPFieldsStoredContract(t *testing.T) { _, repo := setupContractRouter(t, "contract-jwt-secret") now := time.Now().UnixMilli() - _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, "ds-verify-node", "ds-secret", "10.70.0.1", "10.70.0.1", "2001:db8:2::1", "70000-70010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0) - if err != nil { + `, "ds-verify-node", "ds-secret", "10.70.0.1", "10.70.0.1", "2001:db8:2::1", "70000-70010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil { t.Fatalf("insert dual-stack node: %v", err) } var v4, v6 sql.NullString - if err := repo.DB().QueryRow(`SELECT server_ip_v4, server_ip_v6 FROM node WHERE name = ?`, "ds-verify-node").Scan(&v4, &v6); err != nil { + if err := repo.DB().Raw(`SELECT server_ip_v4, server_ip_v6 FROM node WHERE name = ?`, "ds-verify-node").Row().Scan(&v4, &v6); err != nil { t.Fatalf("query node IPs: %v", err) } if !v4.Valid || v4.String != "10.70.0.1" { @@ -282,16 +276,15 @@ func TestIPPreferenceValidValuesContract(t *testing.T) { if pref == "" { name = "valid-pref-empty" } - _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, ip_preference) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, 1.0, 2, "tls", 99999, now, now, 1, nil, 0, pref) - if err != nil { + `, name, 1.0, 2, "tls", 99999, now, now, 1, nil, 0, pref).Error; err != nil { t.Fatalf("insert tunnel with ip_preference=%q: %v", pref, err) } var stored string - if err := repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, name).Scan(&stored); err != nil { + if err := repo.DB().Raw(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, name).Row().Scan(&stored); err != nil { t.Fatalf("query ip_preference for %s: %v", name, err) } if stored != pref { diff --git a/go-backend/tests/contract/tunnel_visibility_contract_test.go b/go-backend/tests/contract/tunnel_visibility_contract_test.go index c623675..8efda3d 100644 --- a/go-backend/tests/contract/tunnel_visibility_contract_test.go +++ b/go-backend/tests/contract/tunnel_visibility_contract_test.go @@ -16,23 +16,22 @@ func TestUserTunnelVisibleListContracts(t *testing.T) { router, repo := setupDiagnosisContractRouter(t, secret) now := time.Now().UnixMilli() - if _, err := repo.DB().Exec(` + if err := repo.DB().Exec(` 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(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1) - `, now, now); err != nil { + `, now, now).Error; err != nil { t.Fatalf("insert user: %v", err) } insertTunnel := func(name string, status int, inx int64) int64 { - res, err := repo.DB().Exec(` + if err := repo.DB().Exec(` INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - `, name, 1.0, 1, "tls", 99999, now, now, status, nil, inx) - if err != nil { + `, name, 1.0, 1, "tls", 99999, now, now, status, nil, inx).Error; err != nil { t.Fatalf("insert tunnel %s: %v", name, err) } - id, err := res.LastInsertId() - if err != nil { + var id int64 + if err := repo.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil { t.Fatalf("get tunnel id %s: %v", name, err) } return id @@ -42,22 +41,22 @@ func TestUserTunnelVisibleListContracts(t *testing.T) { enabledB := insertTunnel("enabled-B", 1, 2) disabledC := insertTunnel("disabled-C", 0, 3) - if _, err := repo.DB().Exec(` + if err := repo.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, ?, ?, ?) - `, 2, enabledA, 100, 1000, 1, 2727251700000, 0); err != nil { + `, 2, enabledA, 100, 1000, 1, 2727251700000, 0).Error; err != nil { t.Fatalf("insert user_tunnel enabledA: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.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, ?, ?, ?) - `, 2, enabledB, 100, 1000, 1, 2727251700000, 1); err != nil { + `, 2, enabledB, 100, 1000, 1, 2727251700000, 1).Error; err != nil { t.Fatalf("insert user_tunnel enabledB: %v", err) } - if _, err := repo.DB().Exec(` + if err := repo.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, ?, ?, ?) - `, 2, disabledC, 100, 1000, 1, 2727251700000, 1); err != nil { + `, 2, disabledC, 100, 1000, 1, 2727251700000, 1).Error; err != nil { t.Fatalf("insert user_tunnel disabledC: %v", err) }