mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
Compare commits
39 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 45d7970177 | |||
| 1b4500202a | |||
| 9a9e83dda0 | |||
| e5e22baf43 | |||
| 961c06655a | |||
| 8dc31383e0 | |||
| 184ac3c3e5 | |||
| 77dbd719ed | |||
| 04ce125416 | |||
| fd5cfc2a40 | |||
| 4e4193e0b0 | |||
| 3f80278dd4 | |||
| 6abe3e7713 | |||
| 47c05c3d02 | |||
| d05c8a2ea4 | |||
| e00e41bb64 | |||
| 5271efec1e | |||
| 2d39cb3005 | |||
| 3e52c8eace | |||
| 7808d57a79 | |||
| 46bc4ca6e4 | |||
| 28e66ab172 | |||
| f19bccec4c | |||
| e37d6cf666 | |||
| 177c2bc35f | |||
| 76c0978763 | |||
| fd1168d855 | |||
| 92c9590c1a | |||
| 2afb1d275a | |||
| 880cd4cac5 | |||
| a69a0f040b | |||
| cf6294a77d | |||
| 524ee4cd95 | |||
| c049ceaacf | |||
| 3424221176 | |||
| 5a9715eb26 | |||
| 1b79213aed | |||
| c0d71125f4 | |||
| f01c0481cd |
@@ -48,6 +48,43 @@ jobs:
|
||||
- name: Build
|
||||
run: go build -v ./...
|
||||
|
||||
backend-postgres-contract:
|
||||
name: Go Backend PostgreSQL Contract
|
||||
runs-on: ubuntu-latest
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:17
|
||||
env:
|
||||
POSTGRES_USER: flux_test
|
||||
POSTGRES_PASSWORD: flux_test_pass
|
||||
POSTGRES_DB: flux_test
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd "pg_isready -U flux_test -d flux_test"
|
||||
--health-interval 10s
|
||||
--health-timeout 5s
|
||||
--health-retries 10
|
||||
defaults:
|
||||
run:
|
||||
working-directory: go-backend
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup Go
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: '1.23'
|
||||
cache-dependency-path: go-backend/go.sum
|
||||
|
||||
- name: Download dependencies
|
||||
run: go mod download
|
||||
|
||||
- name: Run PostgreSQL contract test
|
||||
env:
|
||||
FLVX_POSTGRES_TEST_DSN: 'postgres://flux_test:flux_test_pass@127.0.0.1:5432/flux_test?sslmode=disable'
|
||||
run: go test ./tests/contract -run TestPostgresNodeCreateRepairsMissingIDDefaultContract -count=1
|
||||
|
||||
agent:
|
||||
name: Build Agent
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
# PROJECT KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Mon Feb 02 2026
|
||||
**Commit:** 7ca01ab
|
||||
**Branch:** beta
|
||||
**Generated:** Sun Feb 15 2026
|
||||
**Commit:** e5e22ba
|
||||
**Branch:** main
|
||||
|
||||
## OVERVIEW
|
||||
FLVX (formerly Flux Panel) is a traffic forwarding management system built on a forked GOST v3 stack. It ships as a Go-based admin API (SQLite) + Vite/React UI + Go forwarding agent, with optional mobile WebView wrappers.
|
||||
@@ -43,11 +43,17 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a
|
||||
|
||||
|
||||
## CONVENTIONS
|
||||
- `Authorization` header carries the raw JWT token (no `Bearer` prefix) between `vite-frontend/` and `go-backend/`.
|
||||
- `go-gost/` uses `replace github.com/go-gost/x => ./x` and `go-gost/x/` is also its own Go module.
|
||||
- **Auth**: `Authorization` header carries the raw JWT token (no `Bearer` prefix) between `vite-frontend/` and `go-backend/`.
|
||||
- **Module Fork**: `go-gost/` uses `replace github.com/go-gost/x => ./x` and `go-gost/x/` is also its own Go module.
|
||||
- **Encryption**: Agent-to-panel communication uses AES encryption with node `secret` as PSK.
|
||||
- **API Envelope**: All REST responses follow `{code, msg, data, ts}` structure (code 0 = success).
|
||||
|
||||
## ANTI-PATTERNS (THIS PROJECT)
|
||||
- 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 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 ADD** frontend tests - project has no test infrastructure (Vitest/Jest not configured).
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
@@ -65,6 +71,19 @@ docker compose -f docker-compose-v6.yml up -d
|
||||
(cd go-gost && go run .)
|
||||
```
|
||||
|
||||
## UNIQUE STYLES
|
||||
- **Flat Monorepo**: Language-prefixed dirs (`go-backend`, `go-gost`, `vite-frontend`) instead of `apps/`/`libs/`.
|
||||
- **Asymmetric Go Layout**: `go-backend` follows `cmd/<app>/main.go` while `go-gost` uses `root/main.go`.
|
||||
- **Frontend Hybrid Mode**: `App.tsx` detects "H5 mode" (mobile WebView) vs desktop, dictating layout strategy.
|
||||
|
||||
## NOTES
|
||||
- LSP servers are not installed in this environment (gopls/jdtls/typescript-language-server); rely on grep-based navigation.
|
||||
- `vite-frontend/vite.config.ts` sets `minify: false` and disables treeshake; expect larger bundles.
|
||||
- `vite-frontend` uses `rolldown-vite` (experimental Rust bundler) instead of standard Vite.
|
||||
- Install scripts (`install.sh`, `panel_install.sh`) self-delete after execution - common pattern in one-liner installs.
|
||||
- CI uses UPX compression (`--best --lzma`) on Go binaries before release.
|
||||
- CI dynamically injects `PINNED_VERSION` into install scripts and docker-compose files during releases.
|
||||
- `panel_install.sh` auto-detects IPv6 and modifies `/etc/docker/daemon.json` to enable IPv6 bridge.
|
||||
- Download proxy `https://gcode.hostcentral.cc/` used for GitHub downloads in China/restricted environments.
|
||||
- Backend has contract tests in `go-backend/tests/contract/` - frontend has no test infrastructure (Vitest/Jest not configured).
|
||||
- `analysis/3x-ui/` contains a separate git repo for reference/comparison - not part of FLVX core.
|
||||
|
||||
+10
-3
@@ -1,8 +1,8 @@
|
||||
# GO BACKEND KNOWLEDGE BASE
|
||||
|
||||
## OVERVIEW
|
||||
Go-based Admin API for FLVX (formerly Flux Panel). Replaces the legacy Spring Boot backend.
|
||||
**Stack:** Go 1.23, net/http (std lib), SQLite (modernc.org/sqlite).
|
||||
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).
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
@@ -34,14 +34,21 @@ go-backend/
|
||||
|
||||
## 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.
|
||||
- **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.
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- **DO NOT USE** ORM - uses raw SQL throughout.
|
||||
- **DO NOT CHANGE** handler signatures without updating `router.go`.
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
cd go-backend
|
||||
go run ./cmd/paneld
|
||||
go run ./cmd/paneld # Default: SERVER_ADDR=:6365
|
||||
go test ./...
|
||||
make build
|
||||
```
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# BACKEND HTTP HANDLER KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
|
||||
## OVERVIEW
|
||||
HTTP request handlers for FLVX Admin API. Core business logic layer.
|
||||
**Stack:** Go 1.23, net/http, raw SQL (no ORM).
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
handler/
|
||||
├── handler.go # Main Handler struct, login/captcha, job scheduling
|
||||
├── control_plane.go # Node control plane API (add/delete/list)
|
||||
├── federation.go # Federation/cluster sync API
|
||||
├── flow_policy.go # Traffic policy API
|
||||
├── jobs.go # Background job management (sync, cleanup)
|
||||
├── mutations.go # CRUD for users, tunnels, forwards (largest: 100k+ LOC)
|
||||
└── upgrade.go # System upgrade API
|
||||
```
|
||||
|
||||
## WHERE TO LOOK
|
||||
| Task | Location | Notes |
|
||||
|------|----------|-------|
|
||||
| **User/Tunnel CRUD** | `mutations.go` | Largest file; all create/update/delete ops |
|
||||
| **Login/Captcha** | `handler.go` | Login flow, captcha verification |
|
||||
| **Federation Sync** | `federation.go` | Panel-to-panel sync |
|
||||
| **Traffic Policies** | `flow_policy.go` | Flow limiting, quota management |
|
||||
| **Background Jobs** | `jobs.go` | Scheduled sync/cleanup tasks |
|
||||
|
||||
## CONVENTIONS
|
||||
- Inherits from parent: raw SQL, no ORM, JWT in Authorization header.
|
||||
- Large files expected (`mutations.go` 3716 LOC - central mutation hub).
|
||||
- Uses `sqlite.Repository` for DB access via `repo.XXX()` 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 change handler signatures without updating router.go.
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
cd go-backend
|
||||
go test ./internal/http/handler/...
|
||||
```
|
||||
@@ -114,7 +114,7 @@ func (h *Handler) ensureTunnelPermission(userID int64, roleID int, tunnelID int6
|
||||
|
||||
func (h *Handler) getForwardRecord(forwardID int64) (*forwardRecord, error) {
|
||||
row := h.repo.DB().QueryRow(`
|
||||
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, strategy, status
|
||||
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
|
||||
@@ -152,7 +152,7 @@ func (h *Handler) getTunnelRecord(tunnelID int64) (*tunnelRecord, error) {
|
||||
|
||||
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, strategy, status
|
||||
SELECT id, user_id, user_name, name, tunnel_id, remote_addr, COALESCE(strategy, 'fifo'), status
|
||||
FROM forward
|
||||
WHERE tunnel_id = ?
|
||||
ORDER BY id ASC
|
||||
@@ -200,6 +200,26 @@ func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
|
||||
return isTLSTunnelProtocol(protocol.String), 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
|
||||
@@ -346,6 +366,10 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
|
||||
return err
|
||||
}
|
||||
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
|
||||
tunnelTLSProtocol, err := h.isTunnelSelectedTLSProtocol(forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, fp := range ports {
|
||||
if limiterID != nil && speed != nil {
|
||||
@@ -356,7 +380,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiterID)
|
||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, limiterID, tunnelTLSProtocol)
|
||||
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
|
||||
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
||||
_, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
|
||||
@@ -1095,7 +1119,7 @@ func isNotFoundError(err error) bool {
|
||||
return strings.Contains(msg, "not found") || strings.Contains(msg, "不存在")
|
||||
}
|
||||
|
||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiterID *int64) []map[string]interface{} {
|
||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, limiterID *int64, tunnelTLSProtocol bool) []map[string]interface{} {
|
||||
protocols := []string{"tcp", "udp"}
|
||||
services := make([]map[string]interface{}, 0, 2)
|
||||
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||
@@ -1128,7 +1152,11 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
},
|
||||
}
|
||||
if protocol == "udp" {
|
||||
service["listener"].(map[string]interface{})["metadata"] = map[string]interface{}{"keepAlive": true}
|
||||
listenerMetadata := map[string]interface{}{"keepAlive": true}
|
||||
if tunnelTLSProtocol {
|
||||
listenerMetadata["ttl"] = "10s"
|
||||
}
|
||||
service["listener"].(map[string]interface{})["metadata"] = listenerMetadata
|
||||
}
|
||||
if tunnel != nil && tunnel.Type == 2 {
|
||||
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
|
||||
|
||||
@@ -0,0 +1,377 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// nodeSupportsV4 / nodeSupportsV6
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestNodeSupportsV4_Nil(t *testing.T) {
|
||||
if nodeSupportsV4(nil) {
|
||||
t.Fatal("nil node must not support v4")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV6_Nil(t *testing.T) {
|
||||
if nodeSupportsV6(nil) {
|
||||
t.Fatal("nil node must not support v6")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV4_ExplicitV4(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv4: "10.0.0.1"}
|
||||
if !nodeSupportsV4(n) {
|
||||
t.Fatal("explicit server_ip_v4 must support v4")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV6_ExplicitV6(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv6: "2001:db8::1"}
|
||||
if !nodeSupportsV6(n) {
|
||||
t.Fatal("explicit server_ip_v6 must support v6")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV4_OnlyV6Set(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv6: "2001:db8::1"}
|
||||
if nodeSupportsV4(n) {
|
||||
t.Fatal("node with only v6 should not support v4")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV6_OnlyV4Set(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv4: "10.0.0.1"}
|
||||
if nodeSupportsV6(n) {
|
||||
t.Fatal("node with only v4 should not support v6")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV4_DualStack(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv4: "10.0.0.1", ServerIPv6: "2001:db8::1"}
|
||||
if !nodeSupportsV4(n) {
|
||||
t.Fatal("dual-stack node must support v4")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV6_DualStack(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv4: "10.0.0.1", ServerIPv6: "2001:db8::1"}
|
||||
if !nodeSupportsV6(n) {
|
||||
t.Fatal("dual-stack node must support v6")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV4_LegacyV4Only(t *testing.T) {
|
||||
n := &nodeRecord{ServerIP: "192.168.1.1"}
|
||||
if !nodeSupportsV4(n) {
|
||||
t.Fatal("legacy v4 ip in server_ip must support v4")
|
||||
}
|
||||
if nodeSupportsV6(n) {
|
||||
t.Fatal("legacy v4 ip in server_ip must not support v6")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV6_LegacyV6Only(t *testing.T) {
|
||||
n := &nodeRecord{ServerIP: "2001:db8::1"}
|
||||
if !nodeSupportsV6(n) {
|
||||
t.Fatal("legacy v6 ip in server_ip must support v6")
|
||||
}
|
||||
if nodeSupportsV4(n) {
|
||||
t.Fatal("legacy v6 ip in server_ip must not support v4")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV4_EmptyNode(t *testing.T) {
|
||||
n := &nodeRecord{}
|
||||
if nodeSupportsV4(n) {
|
||||
t.Fatal("empty node must not support v4")
|
||||
}
|
||||
if nodeSupportsV6(n) {
|
||||
t.Fatal("empty node must not support v6")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV4_LegacyBracketed(t *testing.T) {
|
||||
n := &nodeRecord{ServerIP: "[::1]"}
|
||||
if nodeSupportsV4(n) {
|
||||
t.Fatal("bracketed ipv6 must not support v4")
|
||||
}
|
||||
if !nodeSupportsV6(n) {
|
||||
t.Fatal("bracketed ipv6 must support v6")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// pickNodeAddressV4 / pickNodeAddressV6
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestPickNodeAddressV4_Nil(t *testing.T) {
|
||||
if pickNodeAddressV4(nil) != "" {
|
||||
t.Fatal("nil node must return empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickNodeAddressV6_Nil(t *testing.T) {
|
||||
if pickNodeAddressV6(nil) != "" {
|
||||
t.Fatal("nil node must return empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickNodeAddressV4_PreferExplicit(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv4: "10.0.0.1", ServerIP: "192.168.0.1"}
|
||||
got := pickNodeAddressV4(n)
|
||||
if got != "10.0.0.1" {
|
||||
t.Fatalf("expected explicit v4 10.0.0.1, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickNodeAddressV4_FallbackLegacy(t *testing.T) {
|
||||
n := &nodeRecord{ServerIP: "192.168.0.1"}
|
||||
got := pickNodeAddressV4(n)
|
||||
if got != "192.168.0.1" {
|
||||
t.Fatalf("expected legacy 192.168.0.1, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickNodeAddressV6_PreferExplicit(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv6: "2001:db8::1", ServerIP: "::1"}
|
||||
got := pickNodeAddressV6(n)
|
||||
if got != "2001:db8::1" {
|
||||
t.Fatalf("expected explicit v6 2001:db8::1, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickNodeAddressV6_FallbackLegacy(t *testing.T) {
|
||||
n := &nodeRecord{ServerIP: "::1"}
|
||||
got := pickNodeAddressV6(n)
|
||||
if got != "::1" {
|
||||
t.Fatalf("expected legacy ::1, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// selectTunnelDialHost — core IP preference selection logic
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func dualStackNode(name, v4, v6 string) *nodeRecord {
|
||||
return &nodeRecord{
|
||||
Name: name,
|
||||
ServerIPv4: v4,
|
||||
ServerIPv6: v6,
|
||||
}
|
||||
}
|
||||
|
||||
func v4OnlyNode(name, v4 string) *nodeRecord {
|
||||
return &nodeRecord{
|
||||
Name: name,
|
||||
ServerIPv4: v4,
|
||||
}
|
||||
}
|
||||
|
||||
func v6OnlyNode(name, v6 string) *nodeRecord {
|
||||
return &nodeRecord{
|
||||
Name: name,
|
||||
ServerIPv6: v6,
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_NilNodes(t *testing.T) {
|
||||
_, err := selectTunnelDialHost(nil, nil, "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil nodes")
|
||||
}
|
||||
_, err = selectTunnelDialHost(dualStackNode("a", "1.1.1.1", "::1"), nil, "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil toNode")
|
||||
}
|
||||
_, err = selectTunnelDialHost(nil, dualStackNode("b", "1.1.1.1", "::1"), "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil fromNode")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_DualStack_DefaultPreference(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
host, err := selectTunnelDialHost(from, to, "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
// Default prefers v4 when both available
|
||||
if host != "10.0.0.2" {
|
||||
t.Fatalf("default preference should pick v4, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_DualStack_PreferV4(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "10.0.0.2" {
|
||||
t.Fatalf("v4 preference should pick v4 address, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_DualStack_PreferV6(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "2001:db8::2" {
|
||||
t.Fatalf("v6 preference should pick v6 address, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_V4Only_PreferV6Fallback(t *testing.T) {
|
||||
from := v4OnlyNode("from", "10.0.0.1")
|
||||
to := v4OnlyNode("to", "10.0.0.2")
|
||||
|
||||
// User prefers v6, but both nodes are v4-only — should fallback to v4
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "10.0.0.2" {
|
||||
t.Fatalf("v6 preference on v4-only nodes should fallback to v4, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_V6Only_PreferV4Fallback(t *testing.T) {
|
||||
from := v6OnlyNode("from", "2001:db8::1")
|
||||
to := v6OnlyNode("to", "2001:db8::2")
|
||||
|
||||
// User prefers v4, but both nodes are v6-only — should fallback to v6
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "2001:db8::2" {
|
||||
t.Fatalf("v4 preference on v6-only nodes should fallback to v6, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_Incompatible(t *testing.T) {
|
||||
from := v4OnlyNode("from", "10.0.0.1")
|
||||
to := v6OnlyNode("to", "2001:db8::2")
|
||||
|
||||
_, err := selectTunnelDialHost(from, to, "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for incompatible nodes (v4-only -> v6-only)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_Incompatible_Reverse(t *testing.T) {
|
||||
from := v6OnlyNode("from", "2001:db8::1")
|
||||
to := v4OnlyNode("to", "10.0.0.2")
|
||||
|
||||
_, err := selectTunnelDialHost(from, to, "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for incompatible nodes (v6-only -> v4-only)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_WhitespacePreference(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
// Whitespace should be trimmed, treated as "v6"
|
||||
host, err := selectTunnelDialHost(from, to, " v6 ")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "2001:db8::2" {
|
||||
t.Fatalf("trimmed v6 preference should pick v6 address, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_MixedStack_FromDualToV4(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := v4OnlyNode("to", "10.0.0.2")
|
||||
|
||||
// v6 preferred, but target only has v4 — should succeed with v4
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "10.0.0.2" {
|
||||
t.Fatalf("should fallback to v4 when target is v4-only, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_MixedStack_FromDualToV6(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := v6OnlyNode("to", "2001:db8::2")
|
||||
|
||||
// v4 preferred, but target only has v6 — should succeed with v6
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "2001:db8::2" {
|
||||
t.Fatalf("should fallback to v6 when target is v6-only, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_MixedStack_FromV4ToDual(t *testing.T) {
|
||||
from := v4OnlyNode("from", "10.0.0.1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
// v6 preferred, but from only has v4 — should use v4 (from can only reach v4 of target)
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "10.0.0.2" {
|
||||
t.Fatalf("should use v4 when from is v4-only, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_MixedStack_FromV6ToDual(t *testing.T) {
|
||||
from := v6OnlyNode("from", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
// v4 preferred, but from only has v6 — should use v6
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "2001:db8::2" {
|
||||
t.Fatalf("should use v6 when from is v6-only, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// nodeDisplayName
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestNodeDisplayName_Nil(t *testing.T) {
|
||||
got := nodeDisplayName(nil)
|
||||
if got != "node" {
|
||||
t.Fatalf("nil node display name should be 'node', got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeDisplayName_Named(t *testing.T) {
|
||||
n := &nodeRecord{ID: 42, Name: "hk-node"}
|
||||
got := nodeDisplayName(n)
|
||||
if got != "hk-node" {
|
||||
t.Fatalf("expected 'hk-node', got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeDisplayName_Unnamed(t *testing.T) {
|
||||
n := &nodeRecord{ID: 42}
|
||||
got := nodeDisplayName(n)
|
||||
if got != "node_42" {
|
||||
t.Fatalf("expected 'node_42', got %q", got)
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"net"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -298,13 +299,20 @@ func (h *Handler) federationShareDelete(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
share, _ := h.repo.GetPeerShare(req.ID)
|
||||
|
||||
h.cleanupPeerShareRuntimes(req.ID)
|
||||
h.cleanupFederationTunnels(req.ID)
|
||||
|
||||
if err := h.repo.DeletePeerShare(req.ID); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if share != nil && h.wsServer != nil {
|
||||
h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5)
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -543,6 +551,28 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
func remoteNodePortRange(node *nodeRecord) (int, int) {
|
||||
if node == nil || node.IsRemote != 1 || node.RemoteConfig == "" {
|
||||
return 0, 0
|
||||
}
|
||||
_, _, _, _, portRangeStart, portRangeEnd := parseRemoteShareUsageConfig(node.RemoteConfig)
|
||||
return portRangeStart, portRangeEnd
|
||||
}
|
||||
|
||||
func validateRemoteNodePort(node *nodeRecord, port int) error {
|
||||
if node == nil || node.IsRemote != 1 || port <= 0 {
|
||||
return nil
|
||||
}
|
||||
start, end := remoteNodePortRange(node)
|
||||
if start <= 0 || end <= 0 {
|
||||
return nil
|
||||
}
|
||||
if port < start || port > end {
|
||||
return fmt.Errorf("远程节点端口 %d 超出允许范围 %d-%d", port, start, end)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseRemoteShareUsageConfig(raw string) (int64, int64, int64, int64, int, int) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
@@ -779,9 +809,6 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
|
||||
}
|
||||
|
||||
tunnelType := 1
|
||||
if strings.ToLower(req.Protocol) == "udp" {
|
||||
tunnelType = 2
|
||||
}
|
||||
|
||||
tx, err := h.repo.DB().Begin()
|
||||
if err != nil {
|
||||
@@ -982,6 +1009,13 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
return
|
||||
}
|
||||
|
||||
if share.PortRangeStart > 0 && share.PortRangeEnd > 0 && runtime.Port > 0 {
|
||||
if runtime.Port < share.PortRangeStart || runtime.Port > share.PortRangeEnd {
|
||||
response.WriteJSON(w, response.Err(403, fmt.Sprintf("port %d out of allowed range %d-%d", runtime.Port, share.PortRangeStart, share.PortRangeEnd)))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(share.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
@@ -1232,6 +1266,13 @@ func (h *Handler) federationRuntimeCommand(w http.ResponseWriter, r *http.Reques
|
||||
return
|
||||
}
|
||||
|
||||
if isFederationServiceCommand(cmd) {
|
||||
if err := validateFederationCommandPorts(share, req.Data); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
res, err := h.sendNodeCommand(share.NodeID, cmd, req.Data, false, false)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
@@ -1249,6 +1290,55 @@ func isFederationRuntimeCommandAllowed(commandType string) bool {
|
||||
}
|
||||
}
|
||||
|
||||
func isFederationServiceCommand(commandType string) bool {
|
||||
switch strings.ToLower(strings.TrimSpace(commandType)) {
|
||||
case "addservice", "updateservice":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func validateFederationCommandPorts(share *sqlite.PeerShare, data interface{}) error {
|
||||
if share == nil || (share.PortRangeStart <= 0 && share.PortRangeEnd <= 0) {
|
||||
return nil
|
||||
}
|
||||
dataMap, ok := data.(map[string]interface{})
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
if services, ok := dataMap["services"]; ok {
|
||||
serviceList, ok := services.([]interface{})
|
||||
if !ok {
|
||||
return fmt.Errorf("invalid services format")
|
||||
}
|
||||
for _, svc := range serviceList {
|
||||
svcMap, ok := svc.(map[string]interface{})
|
||||
if !ok {
|
||||
return fmt.Errorf("invalid service entry format")
|
||||
}
|
||||
addr, ok := svcMap["addr"].(string)
|
||||
if !ok || addr == "" {
|
||||
continue
|
||||
}
|
||||
_, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid service address: %s", addr)
|
||||
}
|
||||
port, err := strconv.Atoi(portStr)
|
||||
if err != nil || port <= 0 {
|
||||
return fmt.Errorf("invalid port in service address: %s", addr)
|
||||
}
|
||||
if port < share.PortRangeStart || port > share.PortRangeEnd {
|
||||
return fmt.Errorf("port %d out of allowed range %d-%d", port, share.PortRangeStart, share.PortRangeEnd)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) (int, error) {
|
||||
if share == nil {
|
||||
return 0, fmt.Errorf("share not found")
|
||||
@@ -1406,7 +1496,7 @@ func parseIPLiteral(raw string) net.IP {
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(value); ip != nil {
|
||||
return ip
|
||||
return normalizeIPAddress(ip)
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(value)
|
||||
@@ -1418,7 +1508,17 @@ func parseIPLiteral(raw string) net.IP {
|
||||
if host == "" {
|
||||
return nil
|
||||
}
|
||||
return net.ParseIP(host)
|
||||
return normalizeIPAddress(net.ParseIP(host))
|
||||
}
|
||||
|
||||
func normalizeIPAddress(ip net.IP) net.IP {
|
||||
if ip == nil {
|
||||
return nil
|
||||
}
|
||||
if v4 := ip.To4(); v4 != nil {
|
||||
return v4
|
||||
}
|
||||
return ip.To16()
|
||||
}
|
||||
|
||||
func isTrustedProxyIP(ip net.IP) bool {
|
||||
@@ -1549,3 +1649,27 @@ func (h *Handler) cleanupPeerShareRuntimes(shareID int64) {
|
||||
_ = h.repo.MarkPeerShareRuntimeReleased(runtime.ID, now)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) cleanupFederationTunnels(shareID int64) {
|
||||
if h == nil || h.repo == nil || shareID <= 0 {
|
||||
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 {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -56,14 +56,35 @@ func TestPickPeerSharePortUsesRuntimeReservations(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyTunnelRuntimeSkipsRemoteNodes(t *testing.T) {
|
||||
h := &Handler{}
|
||||
func TestApplyTunnelRuntimeSkipsRemoteChainAndOutNodes(t *testing.T) {
|
||||
repo, err := sqlite.Open(filepath.Join(t.TempDir(), "rt-skip.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer repo.Close()
|
||||
|
||||
h := &Handler{repo: repo}
|
||||
now := time.Now().UnixMilli()
|
||||
for _, n := range []struct {
|
||||
id int64
|
||||
name string
|
||||
ip string
|
||||
}{
|
||||
{12, "remote-chain", "10.99.0.2"},
|
||||
{13, "remote-out", "10.99.0.3"},
|
||||
} {
|
||||
if _, err := repo.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 {
|
||||
t.Fatalf("insert node %s: %v", n.name, err)
|
||||
}
|
||||
}
|
||||
|
||||
state := &tunnelCreateState{
|
||||
TunnelID: 1,
|
||||
Type: 2,
|
||||
InNodes: []tunnelRuntimeNode{
|
||||
{NodeID: 11, ChainType: 1, Protocol: "tls"},
|
||||
},
|
||||
InNodes: []tunnelRuntimeNode{},
|
||||
ChainHops: [][]tunnelRuntimeNode{
|
||||
{
|
||||
{NodeID: 12, ChainType: 2, Inx: 1, Port: 41000, Protocol: "tls", Strategy: "round"},
|
||||
@@ -73,9 +94,8 @@ func TestApplyTunnelRuntimeSkipsRemoteNodes(t *testing.T) {
|
||||
{NodeID: 13, ChainType: 3, Port: 42000, Protocol: "tls", Strategy: "round"},
|
||||
},
|
||||
Nodes: map[int64]*nodeRecord{
|
||||
11: {ID: 11, Name: "remote-in", IsRemote: 1},
|
||||
12: {ID: 12, Name: "remote-chain", IsRemote: 1},
|
||||
13: {ID: 13, Name: "remote-out", IsRemote: 1},
|
||||
12: {ID: 12, Name: "remote-chain", IsRemote: 1, ServerIPv4: "10.99.0.2"},
|
||||
13: {ID: 13, Name: "remote-out", IsRemote: 1, ServerIPv4: "10.99.0.3"},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -84,10 +104,10 @@ func TestApplyTunnelRuntimeSkipsRemoteNodes(t *testing.T) {
|
||||
t.Fatalf("apply runtime: %v", err)
|
||||
}
|
||||
if len(chains) != 0 {
|
||||
t.Fatalf("expected no local chains created, got %d", len(chains))
|
||||
t.Fatalf("expected no local chains for remote-only nodes, got %d", len(chains))
|
||||
}
|
||||
if len(services) != 0 {
|
||||
t.Fatalf("expected no local services created, got %d", len(services))
|
||||
t.Fatalf("expected no local services for remote-only nodes, got %d", len(services))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -567,6 +567,13 @@ func TestAuthPeerAllowedIPs(t *testing.T) {
|
||||
xff: "198.51.100.20, 172.20.0.3",
|
||||
wantAllowed: true,
|
||||
},
|
||||
{
|
||||
name: "ipv4-mapped proxy xff allowed",
|
||||
allowedIPs: "198.51.100.20",
|
||||
remoteAddr: "[::ffff:172.20.0.3]:34567",
|
||||
xff: "198.51.100.20, 172.20.0.3",
|
||||
wantAllowed: true,
|
||||
},
|
||||
{
|
||||
name: "non whitelisted ip denied",
|
||||
allowedIPs: "203.0.113.10",
|
||||
|
||||
@@ -251,7 +251,7 @@ func (h *Handler) pauseForwardRecords(forwards []forwardRecord, now int64) {
|
||||
|
||||
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, strategy, status
|
||||
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
|
||||
@@ -266,7 +266,7 @@ func (h *Handler) listActiveForwardsByUser(userID int64) ([]forwardRecord, error
|
||||
|
||||
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, strategy, status
|
||||
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
|
||||
|
||||
@@ -93,6 +93,12 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig)
|
||||
mux.HandleFunc("/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/backup/import", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/backup/restore", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/api/v1/backup/import", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/api/v1/backup/restore", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/captcha/check", h.checkCaptcha)
|
||||
mux.HandleFunc("/api/v1/captcha/verify", h.captchaVerify)
|
||||
mux.HandleFunc("/api/v1/user/package", h.userPackage)
|
||||
@@ -171,9 +177,8 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/federation/runtime/diagnose", h.authPeer(h.federationRuntimeDiagnose))
|
||||
mux.HandleFunc("/api/v1/federation/runtime/command", h.authPeer(h.federationRuntimeCommand))
|
||||
mux.HandleFunc("/api/v1/federation/node/import", h.nodeImport)
|
||||
|
||||
mux.HandleFunc("/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/backup/import", h.backupImport)
|
||||
mux.HandleFunc("/api/v1/announcement/get", h.getAnnouncement)
|
||||
mux.HandleFunc("/api/v1/announcement/update", h.updateAnnouncement)
|
||||
|
||||
mux.HandleFunc("/flow/test", h.flowTest)
|
||||
mux.HandleFunc("/flow/config", h.flowConfig)
|
||||
@@ -1225,3 +1230,53 @@ func (h *Handler) backupImport(w http.ResponseWriter, r *http.Request) {
|
||||
result.AutoBackup = autoBackup
|
||||
response.WriteJSON(w, response.OK(result))
|
||||
}
|
||||
|
||||
func (h *Handler) getAnnouncement(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
ann, err := h.repo.GetAnnouncement()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-1, fmt.Sprintf("获取公告失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
if ann == nil {
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"content": "",
|
||||
"enabled": 0,
|
||||
}))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"content": ann.Content,
|
||||
"enabled": ann.Enabled,
|
||||
}))
|
||||
}
|
||||
|
||||
func (h *Handler) updateAnnouncement(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Content string `json:"content"`
|
||||
Enabled int `json:"enabled"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.Err(500, "请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.UpsertAnnouncement(req.Content, req.Enabled, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-1, fmt.Sprintf("更新公告失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -497,6 +497,7 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
status := asInt(req["status"], 1)
|
||||
trafficRatio := asFloat(req["trafficRatio"], 1.0)
|
||||
inIP := asString(req["inIp"])
|
||||
ipPreference := asString(req["ipPreference"])
|
||||
now := time.Now().UnixMilli()
|
||||
inx := nextIndex(h.repo.DB(), "tunnel")
|
||||
|
||||
@@ -512,6 +513,7 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
runtimeState.IPPreference = ipPreference
|
||||
if strings.TrimSpace(inIP) == "" {
|
||||
inIP = buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes)
|
||||
}
|
||||
@@ -545,6 +547,11 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if targetPort > 0 && targetAddr != "" {
|
||||
inNodeRec := runtimeState.Nodes[firstNodeID]
|
||||
if err := validateRemoteNodePort(inNodeRec, targetPort); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
domainCfg, _ := h.repo.GetConfigByName("panel_domain")
|
||||
localDomain := ""
|
||||
if domainCfg != nil {
|
||||
@@ -560,8 +567,8 @@ 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) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx)
|
||||
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 {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -569,12 +576,10 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
runtimeState.TunnelID = tunnelID
|
||||
var federationBindings []sqlite.FederationTunnelBinding
|
||||
var federationReleaseRefs []federationRuntimeReleaseRef
|
||||
if typeVal == 2 {
|
||||
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
applyTunnelPortsToRequest(req, runtimeState)
|
||||
if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil {
|
||||
@@ -674,6 +679,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
typeVal := asInt(req["type"], 1)
|
||||
ipPreference := asString(req["ipPreference"])
|
||||
|
||||
tx, err := h.repo.DB().Begin()
|
||||
if err != nil {
|
||||
@@ -688,22 +694,21 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
runtimeState.TunnelID = id
|
||||
runtimeState.IPPreference = ipPreference
|
||||
|
||||
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes)
|
||||
|
||||
var federationBindings []sqlite.FederationTunnelBinding
|
||||
var federationReleaseRefs []federationRuntimeReleaseRef
|
||||
if typeVal == 2 {
|
||||
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState)
|
||||
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=?, updated_time=? WHERE id=?`,
|
||||
asString(req["name"]), typeVal, asInt64(req["flow"], 1), asFloat(req["trafficRatio"], 1.0), asInt(req["status"], 1), nullableText(inIp), now, id)
|
||||
_, 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 {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -837,14 +842,18 @@ 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)
|
||||
|
||||
state := &tunnelCreateState{
|
||||
TunnelID: tunnelID,
|
||||
Type: tunnel.Type,
|
||||
InNodes: make([]tunnelRuntimeNode, 0),
|
||||
ChainHops: make([][]tunnelRuntimeNode, 0),
|
||||
OutNodes: make([]tunnelRuntimeNode, 0),
|
||||
Nodes: make(map[int64]*nodeRecord),
|
||||
NodeIDList: make([]int64, 0),
|
||||
TunnelID: tunnelID,
|
||||
Type: tunnel.Type,
|
||||
IPPreference: ipPreference,
|
||||
InNodes: make([]tunnelRuntimeNode, 0),
|
||||
ChainHops: make([][]tunnelRuntimeNode, 0),
|
||||
OutNodes: make([]tunnelRuntimeNode, 0),
|
||||
Nodes: make(map[int64]*nodeRecord),
|
||||
NodeIDList: make([]int64, 0),
|
||||
}
|
||||
|
||||
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
|
||||
@@ -1109,6 +1118,17 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if port <= 0 {
|
||||
port = 10000
|
||||
}
|
||||
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
|
||||
for _, nodeID := range entryNodes {
|
||||
node, nodeErr := h.getNodeRecord(nodeID)
|
||||
if nodeErr != nil {
|
||||
continue
|
||||
}
|
||||
if err := validateRemoteNodePort(node, port); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
inx := nextIndex(h.repo.DB(), "forward")
|
||||
var userName string
|
||||
@@ -1130,7 +1150,6 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
|
||||
for _, nodeID := range entryNodes {
|
||||
_, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port)
|
||||
}
|
||||
@@ -1220,6 +1239,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
port = h.pickTunnelPort(tunnelID)
|
||||
}
|
||||
}
|
||||
fwdEntryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
|
||||
for _, nodeID := range fwdEntryNodes {
|
||||
node, nodeErr := h.getNodeRecord(nodeID)
|
||||
if nodeErr != nil {
|
||||
continue
|
||||
}
|
||||
if err := validateRemoteNodePort(node, port); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
_, err = h.repo.DB().Exec(`
|
||||
UPDATE forward SET name = ?, tunnel_id = ?, remote_addr = ?, strategy = ?, updated_time = ? WHERE id = ?
|
||||
@@ -1544,6 +1574,22 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques
|
||||
if p <= 0 {
|
||||
p = h.pickTunnelPort(req.TargetTunnelID)
|
||||
}
|
||||
bctEntryNodes, _ := h.tunnelEntryNodeIDs(req.TargetTunnelID)
|
||||
portRangeOk := true
|
||||
for _, nid := range bctEntryNodes {
|
||||
nd, ndErr := h.getNodeRecord(nid)
|
||||
if ndErr != nil {
|
||||
continue
|
||||
}
|
||||
if validateRemoteNodePort(nd, p) != nil {
|
||||
portRangeOk = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if !portRangeOk {
|
||||
fail++
|
||||
continue
|
||||
}
|
||||
if err := h.replaceForwardPorts(id, req.TargetTunnelID, p); err != nil {
|
||||
h.rollbackForwardMutation(forward, oldPorts)
|
||||
fail++
|
||||
@@ -2107,13 +2153,14 @@ type tunnelRuntimeNode struct {
|
||||
}
|
||||
|
||||
type tunnelCreateState struct {
|
||||
TunnelID int64
|
||||
Type int
|
||||
InNodes []tunnelRuntimeNode
|
||||
ChainHops [][]tunnelRuntimeNode
|
||||
OutNodes []tunnelRuntimeNode
|
||||
Nodes map[int64]*nodeRecord
|
||||
NodeIDList []int64
|
||||
TunnelID int64
|
||||
Type int
|
||||
IPPreference string // "" = auto, "v4" = prefer IPv4, "v6" = prefer IPv6
|
||||
InNodes []tunnelRuntimeNode
|
||||
ChainHops [][]tunnelRuntimeNode
|
||||
OutNodes []tunnelRuntimeNode
|
||||
Nodes map[int64]*nodeRecord
|
||||
NodeIDList []int64
|
||||
}
|
||||
|
||||
func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
|
||||
@@ -2239,6 +2286,19 @@ func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interfac
|
||||
state.Nodes[nodeID] = node
|
||||
}
|
||||
|
||||
for _, outNode := range state.OutNodes {
|
||||
if err := validateRemoteNodePort(state.Nodes[outNode.NodeID], outNode.Port); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
for _, hop := range state.ChainHops {
|
||||
for _, chainNode := range hop {
|
||||
if err := validateRemoteNodePort(state.Nodes[chainNode.NodeID], chainNode.Port); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
@@ -2340,7 +2400,7 @@ func (h *Handler) federationLocalDomain() string {
|
||||
func (h *Handler) applyFederationRuntime(state *tunnelCreateState) ([]sqlite.FederationTunnelBinding, []federationRuntimeReleaseRef, error) {
|
||||
bindings := make([]sqlite.FederationTunnelBinding, 0)
|
||||
releaseRefs := make([]federationRuntimeReleaseRef, 0)
|
||||
if h == nil || state == nil || state.Type != 2 {
|
||||
if h == nil || state == nil {
|
||||
return bindings, releaseRefs, nil
|
||||
}
|
||||
fc := client.NewFederationClient()
|
||||
@@ -2462,7 +2522,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState) ([]sqlite.Fed
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, errors.New("节点不存在")
|
||||
}
|
||||
host, hostErr := selectTunnelDialHost(node, targetNode)
|
||||
host, hostErr := selectTunnelDialHost(node, targetNode, state.IPPreference)
|
||||
if hostErr != nil {
|
||||
h.releaseFederationRuntimeRefs(releaseRefs)
|
||||
return nil, nil, hostErr
|
||||
@@ -2613,18 +2673,19 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
}
|
||||
|
||||
for _, inNode := range state.InNodes {
|
||||
if node := state.Nodes[inNode.NodeID]; node != nil && node.IsRemote == 1 {
|
||||
continue
|
||||
}
|
||||
node := state.Nodes[inNode.NodeID]
|
||||
targets := state.OutNodes
|
||||
if len(state.ChainHops) > 0 {
|
||||
targets = state.ChainHops[0]
|
||||
}
|
||||
chainData, err := buildTunnelChainConfig(state.TunnelID, inNode.NodeID, targets, state.Nodes)
|
||||
chainData, err := buildTunnelChainConfig(state.TunnelID, inNode.NodeID, targets, state.Nodes, state.IPPreference)
|
||||
if err != nil {
|
||||
return createdChains, createdServices, err
|
||||
}
|
||||
if _, err := h.sendNodeCommand(inNode.NodeID, "AddChains", chainData, true, false); err != nil {
|
||||
if node != nil && node.IsRemote == 1 && shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
return createdChains, createdServices, fmt.Errorf("入口节点 %s 下发转发链失败: %w", nodeDisplayName(state.Nodes[inNode.NodeID]), err)
|
||||
}
|
||||
createdChains = append(createdChains, inNode.NodeID)
|
||||
@@ -2639,7 +2700,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
if node := state.Nodes[chainNode.NodeID]; node != nil && node.IsRemote == 1 {
|
||||
continue
|
||||
}
|
||||
chainData, err := buildTunnelChainConfig(state.TunnelID, chainNode.NodeID, nextTargets, state.Nodes)
|
||||
chainData, err := buildTunnelChainConfig(state.TunnelID, chainNode.NodeID, nextTargets, state.Nodes, state.IPPreference)
|
||||
if err != nil {
|
||||
return createdChains, createdServices, err
|
||||
}
|
||||
@@ -2714,7 +2775,7 @@ func shouldDeferTunnelRuntimeApplyError(err error) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRuntimeNode, nodes map[int64]*nodeRecord) (map[string]interface{}, error) {
|
||||
func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRuntimeNode, nodes map[int64]*nodeRecord, ipPreference string) (map[string]interface{}, error) {
|
||||
fromNode := nodes[fromNodeID]
|
||||
if fromNode == nil {
|
||||
return nil, errors.New("节点不存在")
|
||||
@@ -2728,7 +2789,7 @@ func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRu
|
||||
if targetNode == nil {
|
||||
return nil, errors.New("节点不存在")
|
||||
}
|
||||
host, err := selectTunnelDialHost(fromNode, targetNode)
|
||||
host, err := selectTunnelDialHost(fromNode, targetNode, ipPreference)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -2801,7 +2862,7 @@ func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode,
|
||||
return []map[string]interface{}{service}
|
||||
}
|
||||
|
||||
func selectTunnelDialHost(fromNode, toNode *nodeRecord) (string, error) {
|
||||
func selectTunnelDialHost(fromNode, toNode *nodeRecord, ipPreference string) (string, error) {
|
||||
if fromNode == nil || toNode == nil {
|
||||
return "", errors.New("节点不存在")
|
||||
}
|
||||
@@ -2810,16 +2871,39 @@ func selectTunnelDialHost(fromNode, toNode *nodeRecord) (string, error) {
|
||||
toV4 := nodeSupportsV4(toNode)
|
||||
toV6 := nodeSupportsV6(toNode)
|
||||
|
||||
if fromV4 && toV4 {
|
||||
host := pickNodeAddressV4(toNode)
|
||||
if host != "" {
|
||||
return host, nil
|
||||
switch strings.TrimSpace(ipPreference) {
|
||||
case "v6":
|
||||
if fromV6 && toV6 {
|
||||
if host := pickNodeAddressV6(toNode); host != "" {
|
||||
return host, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
if fromV6 && toV6 {
|
||||
host := pickNodeAddressV6(toNode)
|
||||
if host != "" {
|
||||
return host, nil
|
||||
if fromV4 && toV4 {
|
||||
if host := pickNodeAddressV4(toNode); host != "" {
|
||||
return host, nil
|
||||
}
|
||||
}
|
||||
case "v4":
|
||||
if fromV4 && toV4 {
|
||||
if host := pickNodeAddressV4(toNode); host != "" {
|
||||
return host, nil
|
||||
}
|
||||
}
|
||||
if fromV6 && toV6 {
|
||||
if host := pickNodeAddressV6(toNode); host != "" {
|
||||
return host, nil
|
||||
}
|
||||
}
|
||||
default:
|
||||
if fromV4 && toV4 {
|
||||
if host := pickNodeAddressV4(toNode); host != "" {
|
||||
return host, nil
|
||||
}
|
||||
}
|
||||
if fromV6 && toV6 {
|
||||
if host := pickNodeAddressV6(toNode); host != "" {
|
||||
return host, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("节点链路不兼容:%s(v4=%t,v6=%t) -> %s(v4=%t,v6=%t)", nodeDisplayName(fromNode), fromV4, fromV6, nodeDisplayName(toNode), toV4, toV6)
|
||||
@@ -3035,8 +3119,8 @@ 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, NULL, 0, ?)`,
|
||||
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls"))
|
||||
_, 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 {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -117,6 +117,14 @@ func requiresAdmin(path string) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/backup/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/api/v1/backup/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/tunnel/") {
|
||||
if strings.HasPrefix(path, "/api/v1/tunnel/user/tunnel") {
|
||||
return false
|
||||
@@ -129,6 +137,8 @@ func requiresAdmin(path string) bool {
|
||||
return true
|
||||
case "/api/v1/config/update", "/api/v1/config/update-single":
|
||||
return true
|
||||
case "/api/v1/announcement/update":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -239,3 +239,11 @@ CREATE TABLE IF NOT EXISTS federation_tunnel_binding (
|
||||
|
||||
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
|
||||
);
|
||||
|
||||
@@ -66,6 +66,14 @@ type ViteConfig struct {
|
||||
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
|
||||
@@ -319,6 +327,46 @@ func (r *Repository) UpsertConfig(name, value string, now int64) error {
|
||||
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")
|
||||
@@ -727,7 +775,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
}
|
||||
|
||||
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, f.strategy,
|
||||
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
|
||||
@@ -784,7 +832,7 @@ func (r *Repository) ListUserAccessibleTunnels(userID int64) ([]map[string]inter
|
||||
}
|
||||
|
||||
rows, err := r.db.Query(`
|
||||
SELECT DISTINCT t.id, t.name
|
||||
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
|
||||
@@ -849,7 +897,7 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
}
|
||||
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, inx, name, type, flow, traffic_ratio, status, created_time, in_ip
|
||||
SELECT id, inx, name, type, flow, traffic_ratio, status, created_time, in_ip, COALESCE(ip_preference, '')
|
||||
FROM tunnel
|
||||
ORDER BY inx ASC, id ASC
|
||||
`)
|
||||
@@ -867,7 +915,8 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
var typ, status int
|
||||
var trafficRatio float64
|
||||
var inIP sql.NullString
|
||||
if err := rows.Scan(&id, &inx, &name, &typ, &flow, &trafficRatio, &status, &createdTime, &inIP); err != nil {
|
||||
var ipPreference string
|
||||
if err := rows.Scan(&id, &inx, &name, &typ, &flow, &trafficRatio, &status, &createdTime, &inIP, &ipPreference); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -881,6 +930,7 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
"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),
|
||||
@@ -1305,7 +1355,9 @@ func bootstrapSchema(db *store.DB, schemaSQL, seedSQL string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
const currentSchemaVersion = 1
|
||||
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)`)
|
||||
@@ -1327,6 +1379,11 @@ func migrateSchema(db *store.DB) error {
|
||||
}
|
||||
|
||||
ver := getSchemaVersion(db)
|
||||
if db.Dialect() == store.DialectPostgres {
|
||||
if err := ensurePostgresIDDefaultsFn(db); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if ver >= currentSchemaVersion {
|
||||
return nil
|
||||
}
|
||||
@@ -1359,7 +1416,8 @@ func migrateSchema(db *store.DB) error {
|
||||
"remote_config": "TEXT",
|
||||
},
|
||||
"tunnel": {
|
||||
"inx": "INTEGER NOT NULL DEFAULT 0",
|
||||
"inx": "INTEGER NOT NULL DEFAULT 0",
|
||||
"ip_preference": "VARCHAR(10) NOT NULL DEFAULT ''",
|
||||
},
|
||||
"forward": {
|
||||
"inx": "INTEGER NOT NULL DEFAULT 0",
|
||||
@@ -1375,11 +1433,27 @@ func migrateSchema(db *store.DB) error {
|
||||
}
|
||||
}
|
||||
|
||||
if db.Dialect() == store.DialectPostgres {
|
||||
if err := ensurePostgresIDDefaults(db); err != nil {
|
||||
return err
|
||||
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
|
||||
}
|
||||
@@ -1545,6 +1619,17 @@ func isMissingColumnError(dialect store.Dialect, err error) bool {
|
||||
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")
|
||||
@@ -1963,6 +2048,7 @@ type TunnelBackup struct {
|
||||
Status int `json:"status"`
|
||||
InIP string `json:"inIp,omitempty"`
|
||||
Inx int `json:"inx"`
|
||||
IPPreference string `json:"ipPreference,omitempty"`
|
||||
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
|
||||
}
|
||||
|
||||
@@ -1978,19 +2064,25 @@ type ChainTunnelBackup struct {
|
||||
}
|
||||
|
||||
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"`
|
||||
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 {
|
||||
@@ -2287,7 +2379,7 @@ func (r *Repository) exportNodes() ([]NodeBackup, error) {
|
||||
|
||||
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
|
||||
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 {
|
||||
@@ -2298,13 +2390,25 @@ func (r *Repository) exportTunnels() ([]TunnelBackup, error) {
|
||||
var tunnels []TunnelBackup
|
||||
for rows.Next() {
|
||||
var t TunnelBackup
|
||||
var protocol sql.NullString
|
||||
var updatedTime sql.NullInt64
|
||||
var inIP sql.NullString
|
||||
if err := rows.Scan(&t.ID, &t.Name, &t.TrafficRatio, &t.Type, &t.Protocol, &t.Flow, &t.CreatedTime, &t.UpdatedTime, &t.Status, &inIP, &t.Inx); err != nil {
|
||||
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 {
|
||||
@@ -2330,12 +2434,23 @@ func (r *Repository) exportChainTunnels(tunnelID int64) ([]ChainTunnelBackup, er
|
||||
for rows.Next() {
|
||||
var ct ChainTunnelBackup
|
||||
var port sql.NullInt64
|
||||
if err := rows.Scan(&ct.ID, &ct.TunnelID, &ct.ChainType, &ct.NodeID, &port, &ct.Strategy, &ct.Inx, &ct.Protocol); err != nil {
|
||||
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()
|
||||
@@ -2354,14 +2469,62 @@ func (r *Repository) exportForwards() ([]ForwardBackup, error) {
|
||||
var forwards []ForwardBackup
|
||||
for rows.Next() {
|
||||
var f ForwardBackup
|
||||
if err := rows.Scan(&f.ID, &f.UserID, &f.UserName, &f.Name, &f.TunnelID, &f.RemoteAddr, &f.Strategy, &f.InFlow, &f.OutFlow, &f.CreatedTime, &f.UpdatedTime, &f.Status, &f.Inx); err != nil {
|
||||
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
|
||||
@@ -2484,7 +2647,7 @@ func (r *Repository) exportUserGroups() ([]UserGroupBackup, error) {
|
||||
|
||||
func (r *Repository) exportPermissions() ([]PermissionBackup, error) {
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, user_group_id, tunnel_group_id, created_time, created_by_group
|
||||
SELECT id, user_group_id, tunnel_group_id, created_time
|
||||
FROM group_permission ORDER BY id ASC
|
||||
`)
|
||||
if err != nil {
|
||||
@@ -2495,9 +2658,10 @@ func (r *Repository) exportPermissions() ([]PermissionBackup, error) {
|
||||
var permissions []PermissionBackup
|
||||
for rows.Next() {
|
||||
var p PermissionBackup
|
||||
if err := rows.Scan(&p.ID, &p.UserGroupID, &p.TunnelGroupID, &p.CreatedTime, &p.CreatedByGroup); err != nil {
|
||||
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 {
|
||||
@@ -2714,8 +2878,8 @@ func (r *Repository) importTunnels(db Execer, tunnels []TunnelBackup, now int64)
|
||||
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)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
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,
|
||||
@@ -2725,8 +2889,9 @@ func (r *Repository) importTunnels(db Execer, tunnels []TunnelBackup, now int64)
|
||||
updated_time = excluded.updated_time,
|
||||
status = excluded.status,
|
||||
in_ip = excluded.in_ip,
|
||||
inx = excluded.inx
|
||||
`, t.ID, t.Name, t.TrafficRatio, t.Type, t.Protocol, t.Flow, t.CreatedTime, now, t.Status, t.InIP, t.Inx)
|
||||
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
|
||||
}
|
||||
@@ -2775,6 +2940,18 @@ func (r *Repository) importForwards(db Execer, forwards []ForwardBackup, now int
|
||||
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
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
package sqlite
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
func TestMigrateSchemaRunsPostgresIDRepairEvenAtCurrentVersion(t *testing.T) {
|
||||
raw, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = raw.Close()
|
||||
})
|
||||
|
||||
db := store.Wrap(raw, store.DialectPostgres)
|
||||
if _, err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`); err != nil {
|
||||
t.Fatalf("create schema_version: %v", err)
|
||||
}
|
||||
if _, err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, currentSchemaVersion); err != nil {
|
||||
t.Fatalf("seed schema_version: %v", err)
|
||||
}
|
||||
|
||||
called := 0
|
||||
original := ensurePostgresIDDefaultsFn
|
||||
ensurePostgresIDDefaultsFn = func(db *store.DB) error {
|
||||
called++
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePostgresIDDefaultsFn = original
|
||||
})
|
||||
|
||||
if err := migrateSchema(db); err != nil {
|
||||
t.Fatalf("migrateSchema: %v", err)
|
||||
}
|
||||
if called != 1 {
|
||||
t.Fatalf("expected postgres id repair to run once, got %d", called)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaReturnsPostgresIDRepairError(t *testing.T) {
|
||||
raw, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = raw.Close()
|
||||
})
|
||||
|
||||
db := store.Wrap(raw, store.DialectPostgres)
|
||||
if _, err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`); err != nil {
|
||||
t.Fatalf("create schema_version: %v", err)
|
||||
}
|
||||
if _, err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, currentSchemaVersion); err != nil {
|
||||
t.Fatalf("seed schema_version: %v", err)
|
||||
}
|
||||
|
||||
wantErr := errors.New("repair failed")
|
||||
original := ensurePostgresIDDefaultsFn
|
||||
ensurePostgresIDDefaultsFn = func(db *store.DB) error {
|
||||
return wantErr
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePostgresIDDefaultsFn = original
|
||||
})
|
||||
|
||||
err = migrateSchema(db)
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("expected error %v, got %v", wantErr, err)
|
||||
}
|
||||
}
|
||||
@@ -80,7 +80,8 @@ CREATE TABLE IF NOT EXISTS tunnel (
|
||||
updated_time INTEGER NOT NULL,
|
||||
status INTEGER NOT NULL,
|
||||
in_ip TEXT,
|
||||
inx INTEGER NOT NULL DEFAULT 0
|
||||
inx INTEGER NOT NULL DEFAULT 0,
|
||||
ip_preference VARCHAR(10) NOT NULL DEFAULT ''
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS chain_tunnel (
|
||||
@@ -243,3 +244,11 @@ CREATE TABLE IF NOT EXISTS federation_tunnel_binding (
|
||||
|
||||
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
|
||||
);
|
||||
|
||||
@@ -316,6 +316,136 @@ func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationDualPanelRemoteEntryRuntimeContract(t *testing.T) {
|
||||
providerSecret := "provider-contract-jwt"
|
||||
providerRouter, providerRepo := setupContractRouter(t, providerSecret)
|
||||
providerServer := httptest.NewServer(providerRouter)
|
||||
defer providerServer.Close()
|
||||
|
||||
consumerSecret := "consumer-contract-jwt"
|
||||
consumerRouter, consumerRepo := setupContractRouter(t, consumerSecret)
|
||||
|
||||
consumerAdminToken, err := auth.GenerateToken(1, "consumer-admin", 0, consumerSecret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate consumer admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
providerEntryNodeID := insertContractNode(t, providerRepo, "provider-entry-rt", "198.51.100.21", "43020-43030", "provider-entry-rt-secret", 1)
|
||||
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{
|
||||
Name: "entry-share-rt",
|
||||
NodeID: providerEntryNodeID,
|
||||
Token: "share-entry-rt-token",
|
||||
PortRangeStart: 43020,
|
||||
PortRangeEnd: 43030,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
})
|
||||
insertPeerShare(t, providerRepo, &sqlite.PeerShare{
|
||||
Name: "middle-share-rt",
|
||||
NodeID: providerMiddleNodeID,
|
||||
Token: "share-middle-rt-token",
|
||||
PortRangeStart: 44020,
|
||||
PortRangeEnd: 44030,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
})
|
||||
insertPeerShare(t, providerRepo, &sqlite.PeerShare{
|
||||
Name: "exit-share-rt",
|
||||
NodeID: providerExitNodeID,
|
||||
Token: "share-exit-rt-token",
|
||||
PortRangeStart: 45020,
|
||||
PortRangeEnd: 45030,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
})
|
||||
|
||||
importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-entry-rt-token")
|
||||
importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-middle-rt-token")
|
||||
importRemoteNodeForContract(t, consumerRouter, consumerAdminToken, providerServer.URL, "share-exit-rt-token")
|
||||
|
||||
entryRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-entry-rt-token")
|
||||
middleRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-middle-rt-token")
|
||||
exitRemoteNodeID := queryRemoteNodeIDByToken(t, consumerRepo, "share-exit-rt-token")
|
||||
|
||||
var commandMu sync.Mutex
|
||||
entryCommands := make([]string, 0, 8)
|
||||
stopEntry := startMockNodeSessionWithHook(t, providerServer.URL, "provider-entry-rt-secret", func(cmdType string) {
|
||||
commandMu.Lock()
|
||||
entryCommands = append(entryCommands, cmdType)
|
||||
commandMu.Unlock()
|
||||
})
|
||||
defer stopEntry()
|
||||
stopMiddle := startMockNodeSession(t, providerServer.URL, "provider-middle-rt-secret")
|
||||
defer stopMiddle()
|
||||
stopExit := startMockNodeSession(t, providerServer.URL, "provider-exit-rt-secret")
|
||||
defer stopExit()
|
||||
|
||||
createTunnel := func(name string) int64 {
|
||||
payload := map[string]interface{}{
|
||||
"name": name,
|
||||
"type": 2,
|
||||
"flow": 99999,
|
||||
"status": 1,
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": entryRemoteNodeID, "protocol": "tls", "strategy": "round"},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{
|
||||
{{"nodeId": middleRemoteNodeID, "protocol": "tls", "strategy": "round"}},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": exitRemoteNodeID, "protocol": "tls", "strategy": "round"},
|
||||
},
|
||||
}
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", bytes.NewReader(body))
|
||||
req.Header.Set("Authorization", consumerAdminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
consumerRouter.ServeHTTP(res, req)
|
||||
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 {
|
||||
t.Fatalf("query tunnel id (%s): %v", name, err)
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
t.Fatalf("invalid tunnel id for %s", name)
|
||||
}
|
||||
return tunnelID
|
||||
}
|
||||
|
||||
createTunnel("dual-panel-remote-entry-online")
|
||||
|
||||
commandMu.Lock()
|
||||
seenAddChains := false
|
||||
seenCommands := append([]string(nil), entryCommands...)
|
||||
for _, cmdType := range entryCommands {
|
||||
if strings.EqualFold(strings.TrimSpace(cmdType), "AddChains") {
|
||||
seenAddChains = true
|
||||
break
|
||||
}
|
||||
}
|
||||
commandMu.Unlock()
|
||||
if !seenAddChains {
|
||||
t.Fatalf("expected entry remote node to receive AddChains, commands=%v", seenCommands)
|
||||
}
|
||||
|
||||
stopEntry()
|
||||
waitNodeStatus(t, providerRepo, providerEntryNodeID, 0)
|
||||
|
||||
createTunnel("dual-panel-remote-entry-offline")
|
||||
}
|
||||
|
||||
func insertContractNode(t *testing.T, repo *sqlite.Repository, name, ip, portRange, secret string, status int) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
@@ -409,6 +539,10 @@ func assertCount(t *testing.T, repo *sqlite.Repository, query string, arg interf
|
||||
}
|
||||
|
||||
func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func() {
|
||||
return startMockNodeSessionWithHook(t, baseURL, nodeSecret, nil)
|
||||
}
|
||||
|
||||
func startMockNodeSessionWithHook(t *testing.T, baseURL string, nodeSecret string, onCommand func(cmdType string)) func() {
|
||||
t.Helper()
|
||||
u, err := url.Parse(baseURL)
|
||||
if err != nil {
|
||||
@@ -468,6 +602,9 @@ func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func(
|
||||
if strings.TrimSpace(cmd.RequestID) == "" {
|
||||
continue
|
||||
}
|
||||
if onCommand != nil {
|
||||
onCommand(strings.TrimSpace(cmd.Type))
|
||||
}
|
||||
|
||||
respType := fmt.Sprintf("%sResponse", cmd.Type)
|
||||
respPayload := map[string]interface{}{
|
||||
@@ -492,9 +629,27 @@ func startMockNodeSession(t *testing.T, baseURL string, nodeSecret string) func(
|
||||
}
|
||||
}()
|
||||
|
||||
var stopOnce sync.Once
|
||||
return func() {
|
||||
_ = conn.Close()
|
||||
wg.Wait()
|
||||
stopOnce.Do(func() {
|
||||
_ = conn.Close()
|
||||
wg.Wait()
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func waitNodeStatus(t *testing.T, repo *sqlite.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 {
|
||||
return
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("node %d status did not reach %d before timeout", nodeID, expectedStatus)
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -533,3 +688,112 @@ func valueAsBool(v interface{}) bool {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) {
|
||||
providerSecret := "provider-portrange-jwt"
|
||||
providerRouter, providerRepo := setupContractRouter(t, providerSecret)
|
||||
providerServer := httptest.NewServer(providerRouter)
|
||||
defer providerServer.Close()
|
||||
|
||||
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{
|
||||
Name: "portrange-share",
|
||||
NodeID: providerNodeID,
|
||||
Token: "share-portrange-token",
|
||||
PortRangeStart: 44000,
|
||||
PortRangeEnd: 44010,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
})
|
||||
|
||||
stopNode := startMockNodeSession(t, providerServer.URL, "provider-portrange-secret")
|
||||
defer stopNode()
|
||||
|
||||
sendCommand := func(token string, cmdType string, data interface{}) *httptest.ResponseRecorder {
|
||||
payload := map[string]interface{}{
|
||||
"commandType": cmdType,
|
||||
"data": data,
|
||||
}
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal command payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/runtime/command", bytes.NewReader(body))
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
providerRouter.ServeHTTP(res, req)
|
||||
return res
|
||||
}
|
||||
|
||||
// Test: AddService with port OUTSIDE allowed range should be rejected
|
||||
outOfRangeData := map[string]interface{}{
|
||||
"services": []map[string]interface{}{
|
||||
{
|
||||
"name": "test_service_tcp",
|
||||
"addr": "[::]:55555",
|
||||
"handler": map[string]interface{}{
|
||||
"type": "tcp",
|
||||
},
|
||||
"listener": map[string]interface{}{
|
||||
"type": "tcp",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
res := sendCommand("share-portrange-token", "AddService", outOfRangeData)
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 403 {
|
||||
t.Fatalf("expected code 403 for out-of-range port, got %d (msg: %s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
// Test: UpdateService with port OUTSIDE allowed range should be rejected
|
||||
res = sendCommand("share-portrange-token", "UpdateService", outOfRangeData)
|
||||
out = response.R{}
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 403 {
|
||||
t.Fatalf("expected code 403 for out-of-range UpdateService, got %d (msg: %s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
// Test: AddService with port INSIDE allowed range should succeed
|
||||
inRangeData := map[string]interface{}{
|
||||
"services": []map[string]interface{}{
|
||||
{
|
||||
"name": "test_service_ok_tcp",
|
||||
"addr": "[::]:44005",
|
||||
"handler": map[string]interface{}{
|
||||
"type": "tcp",
|
||||
},
|
||||
"listener": map[string]interface{}{
|
||||
"type": "tcp",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
res = sendCommand("share-portrange-token", "AddService", inRangeData)
|
||||
out = response.R{}
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0 for in-range port, got %d (msg: %s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
// Test: Non-service commands should pass through without port validation
|
||||
res = sendCommand("share-portrange-token", "reload", nil)
|
||||
out = response.R{}
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0 for reload command, got %d (msg: %s)", out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -203,6 +203,399 @@ func TestSpeedLimitTunnelsRouteAlias(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate user token: %v", err)
|
||||
}
|
||||
|
||||
key := "backup_contract_key"
|
||||
if _, err := repo.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 {
|
||||
t.Fatalf("seed config for backup contract: %v", err)
|
||||
}
|
||||
|
||||
t.Run("non-admin is blocked on backup export", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/export", nil)
|
||||
req.Header.Set("Authorization", userToken)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCodeMsg(t, resp, 403, "权限不足,仅管理员可操作")
|
||||
})
|
||||
|
||||
t.Run("standard and duplicate export routes both work", func(t *testing.T) {
|
||||
payloadA := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
if len(payloadA.Configs) == 0 {
|
||||
t.Fatalf("expected exported configs, got none")
|
||||
}
|
||||
if _, ok := payloadA.Configs[key]; !ok {
|
||||
t.Fatalf("expected %q in exported configs", key)
|
||||
}
|
||||
|
||||
payloadB := exportBackupPayload(t, router, "/api/v1/api/v1/backup/export", adminToken)
|
||||
if len(payloadB.Configs) == 0 {
|
||||
t.Fatalf("expected exported configs from duplicate-prefix route, got none")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("backup import applies exported data", func(t *testing.T) {
|
||||
payload := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
payload.Configs[key] = "v2"
|
||||
raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: payload})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal import payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/import", bytes.NewReader(raw))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
var out response.R
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode import response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected import code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
cfg, err := repo.GetConfigByName(key)
|
||||
if err != nil {
|
||||
t.Fatalf("query imported config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "v2" {
|
||||
t.Fatalf("expected imported config value v2, got %+v", cfg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("backup restore alias applies exported data", func(t *testing.T) {
|
||||
payload := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
payload.Configs[key] = "v3"
|
||||
raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: payload})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal restore payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/restore", bytes.NewReader(raw))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
var out response.R
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode restore response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected restore code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
cfg, err := repo.GetConfigByName(key)
|
||||
if err != nil {
|
||||
t.Fatalf("query restored config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "v3" {
|
||||
t.Fatalf("expected restored config value v3, got %+v", cfg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("backup export and import preserve forward ports", func(t *testing.T) {
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
tunnelRes, err := repo.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 {
|
||||
t.Fatalf("seed tunnel for forward backup: %v", err)
|
||||
}
|
||||
tunnelID, err := tunnelRes.LastInsertId()
|
||||
if err != nil {
|
||||
t.Fatalf("read tunnel id for forward backup: %v", err)
|
||||
}
|
||||
|
||||
forwardRes, 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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88)
|
||||
if err != nil {
|
||||
t.Fatalf("seed forward for backup: %v", err)
|
||||
}
|
||||
forwardID, err := forwardRes.LastInsertId()
|
||||
if err != nil {
|
||||
t.Fatalf("read forward id for backup: %v", err)
|
||||
}
|
||||
|
||||
expected := map[int64]int{
|
||||
2001: 21001,
|
||||
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 {
|
||||
t.Fatalf("seed forward_port %d:%d: %v", nodeID, port, err)
|
||||
}
|
||||
}
|
||||
|
||||
exportReq := httptest.NewRequest(http.MethodPost, "/api/v1/backup/export", bytes.NewBufferString(`{"types":["forwards"]}`))
|
||||
exportReq.Header.Set("Authorization", adminToken)
|
||||
exportReq.Header.Set("Content-Type", "application/json")
|
||||
exportResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(exportResp, exportReq)
|
||||
|
||||
if exportResp.Code != http.StatusOK {
|
||||
t.Fatalf("expected export status 200, got %d", exportResp.Code)
|
||||
}
|
||||
|
||||
exportBody, err := io.ReadAll(exportResp.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("read forwards backup body: %v", err)
|
||||
}
|
||||
|
||||
var payload map[string]interface{}
|
||||
if err := json.Unmarshal(exportBody, &payload); err != nil {
|
||||
t.Fatalf("decode forwards backup payload: %v", err)
|
||||
}
|
||||
version, _ := payload["version"].(string)
|
||||
if strings.TrimSpace(version) == "" {
|
||||
t.Fatalf("expected backup payload version, body=%s", string(exportBody))
|
||||
}
|
||||
|
||||
forwardsRaw, ok := payload["forwards"].([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected forwards array in payload, body=%s", string(exportBody))
|
||||
}
|
||||
|
||||
foundForward := false
|
||||
foundPorts := map[int64]int{}
|
||||
for _, item := range forwardsRaw {
|
||||
forwardMap, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
idValue, ok := forwardMap["id"].(float64)
|
||||
if !ok || int64(idValue) != forwardID {
|
||||
continue
|
||||
}
|
||||
foundForward = true
|
||||
|
||||
portsRaw, ok := forwardMap["forwardPorts"].([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected forwardPorts for forward %d in payload", forwardID)
|
||||
}
|
||||
for _, p := range portsRaw {
|
||||
portMap, ok := p.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
nodeID, nodeOK := portMap["nodeId"].(float64)
|
||||
port, portOK := portMap["port"].(float64)
|
||||
if nodeOK && portOK {
|
||||
foundPorts[int64(nodeID)] = int(port)
|
||||
}
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
if !foundForward {
|
||||
t.Fatalf("expected forward %d in exported forwards payload", forwardID)
|
||||
}
|
||||
if len(foundPorts) != len(expected) {
|
||||
t.Fatalf("expected %d exported forward ports, got %d", len(expected), len(foundPorts))
|
||||
}
|
||||
for nodeID, port := range expected {
|
||||
if got, ok := foundPorts[nodeID]; !ok || got != port {
|
||||
t.Fatalf("expected exported forward port node=%d port=%d, got %v", nodeID, port, foundPorts)
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := repo.DB().Exec(`DELETE FROM forward_port WHERE forward_id = ?`, forwardID); 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 {
|
||||
t.Fatalf("seed wrong forward_port before import: %v", err)
|
||||
}
|
||||
|
||||
payload["types"] = []string{"forwards"}
|
||||
importBody, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal forwards import payload: %v", err)
|
||||
}
|
||||
|
||||
importReq := httptest.NewRequest(http.MethodPost, "/api/v1/backup/import", bytes.NewReader(importBody))
|
||||
importReq.Header.Set("Authorization", adminToken)
|
||||
importReq.Header.Set("Content-Type", "application/json")
|
||||
importResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(importResp, importReq)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(importResp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode forwards import response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
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)
|
||||
if err != nil {
|
||||
t.Fatalf("query forward ports after import: %v", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
after := make(map[int64]int)
|
||||
for rows.Next() {
|
||||
var nodeID int64
|
||||
var port int
|
||||
if err := rows.Scan(&nodeID, &port); err != nil {
|
||||
t.Fatalf("scan forward_port row: %v", err)
|
||||
}
|
||||
after[nodeID] = port
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatalf("iterate forward_port rows: %v", err)
|
||||
}
|
||||
|
||||
if len(after) != len(expected) {
|
||||
t.Fatalf("expected %d forward ports after import, got %d (%v)", len(expected), len(after), after)
|
||||
}
|
||||
for nodeID, port := range expected {
|
||||
if got, ok := after[nodeID]; !ok || got != port {
|
||||
t.Fatalf("expected forward_port node=%d port=%d after import, got %v", nodeID, port, after)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) {
|
||||
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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "legacy-null-chain", 1.0, 1, "tls", 1000, now, now, 1, nil, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("seed tunnel for nullable chain export: %v", err)
|
||||
}
|
||||
tunnelID, err := res.LastInsertId()
|
||||
if err != nil {
|
||||
t.Fatalf("read tunnel id for nullable chain export: %v", err)
|
||||
}
|
||||
|
||||
if _, err := repo.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 {
|
||||
t.Fatalf("seed nullable chain_tunnel row: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/export", bytes.NewBufferString(`{"types":["tunnels"]}`))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
if resp.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d", resp.Code)
|
||||
}
|
||||
|
||||
var payload struct {
|
||||
Version string `json:"version"`
|
||||
Tunnels []struct {
|
||||
ID int64 `json:"id"`
|
||||
ChainTunnels []struct {
|
||||
Inx int `json:"inx"`
|
||||
Strategy string `json:"strategy"`
|
||||
Protocol string `json:"protocol"`
|
||||
} `json:"chainTunnels"`
|
||||
} `json:"tunnels"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode tunnels backup payload: %v", err)
|
||||
}
|
||||
if strings.TrimSpace(payload.Version) == "" {
|
||||
t.Fatalf("expected backup payload version, got empty")
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, tunnel := range payload.Tunnels {
|
||||
if tunnel.ID != tunnelID {
|
||||
continue
|
||||
}
|
||||
if len(tunnel.ChainTunnels) != 1 {
|
||||
t.Fatalf("expected one chain tunnel for seeded tunnel %d, got %d", tunnelID, len(tunnel.ChainTunnels))
|
||||
}
|
||||
if tunnel.ChainTunnels[0].Inx != 0 {
|
||||
t.Fatalf("expected nullable chain inx to export as 0, got %d", tunnel.ChainTunnels[0].Inx)
|
||||
}
|
||||
if tunnel.ChainTunnels[0].Strategy != "" {
|
||||
t.Fatalf("expected nullable chain strategy to export as empty string, got %q", tunnel.ChainTunnels[0].Strategy)
|
||||
}
|
||||
if tunnel.ChainTunnels[0].Protocol != "" {
|
||||
t.Fatalf("expected nullable chain protocol to export as empty string, got %q", tunnel.ChainTunnels[0].Protocol)
|
||||
}
|
||||
found = true
|
||||
break
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("expected seeded tunnel %d in backup export", tunnelID)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
type backupExportPayload struct {
|
||||
Version string `json:"version"`
|
||||
ExportedAt int64 `json:"exportedAt"`
|
||||
Configs map[string]string `json:"configs"`
|
||||
}
|
||||
|
||||
type backupImportPayload struct {
|
||||
Types []string `json:"types"`
|
||||
backupExportPayload
|
||||
}
|
||||
|
||||
func exportBackupPayload(t *testing.T, router http.Handler, path, token string) backupExportPayload {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, path, bytes.NewBufferString(`{"types":["configs"]}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
if resp.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200 on %s, got %d", path, resp.Code)
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("read backup payload from %s: %v", path, err)
|
||||
}
|
||||
|
||||
var payload backupExportPayload
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
t.Fatalf("decode backup payload from %s: %v", path, err)
|
||||
}
|
||||
if strings.TrimSpace(payload.Version) == "" {
|
||||
var out response.R
|
||||
if err := json.Unmarshal(body, &out); err == nil {
|
||||
t.Fatalf("expected backup payload on %s, got envelope code=%d msg=%q", path, out.Code, out.Msg)
|
||||
}
|
||||
t.Fatalf("expected non-empty backup payload version on %s, body=%s", path, string(body))
|
||||
}
|
||||
if payload.Configs == nil {
|
||||
t.Fatalf("expected configs map in backup payload on %s", path)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.Repository) {
|
||||
t.Helper()
|
||||
dbPath := filepath.Join(t.TempDir(), "contract.db")
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
httpserver "go-backend/internal/http"
|
||||
"go-backend/internal/http/handler"
|
||||
"go-backend/internal/store/sqlite"
|
||||
)
|
||||
|
||||
func TestPostgresNodeCreateRepairsMissingIDDefaultContract(t *testing.T) {
|
||||
baseDSN := strings.TrimSpace(os.Getenv("FLVX_POSTGRES_TEST_DSN"))
|
||||
if baseDSN == "" {
|
||||
t.Skip("set FLVX_POSTGRES_TEST_DSN to run postgres contract tests")
|
||||
}
|
||||
|
||||
schemaName := "contract_node_id_" + strconv.FormatInt(time.Now().UnixNano(), 36)
|
||||
adminDB, err := sql.Open("pgx", baseDSN)
|
||||
if err != nil {
|
||||
t.Fatalf("open postgres admin connection: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = adminDB.Exec(`DROP SCHEMA IF EXISTS "` + schemaName + `" CASCADE`)
|
||||
_ = adminDB.Close()
|
||||
})
|
||||
|
||||
if _, err := adminDB.Exec(`CREATE SCHEMA "` + schemaName + `"`); err != nil {
|
||||
t.Fatalf("create schema %s: %v", schemaName, err)
|
||||
}
|
||||
|
||||
testDSN, err := withSearchPath(baseDSN, schemaName)
|
||||
if err != nil {
|
||||
t.Fatalf("build schema dsn: %v", err)
|
||||
}
|
||||
|
||||
repo, err := sqlite.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()
|
||||
t.Fatalf("drop node.id default to simulate drift: %v", err)
|
||||
}
|
||||
if err := repo.Close(); err != nil {
|
||||
t.Fatalf("close repository before reopen: %v", err)
|
||||
}
|
||||
|
||||
repo, err = sqlite.OpenPostgres(testDSN)
|
||||
if err != nil {
|
||||
t.Fatalf("reopen postgres repository: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = repo.Close()
|
||||
})
|
||||
|
||||
var columnDefault sql.NullString
|
||||
if err := repo.DB().QueryRow(`
|
||||
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 {
|
||||
t.Fatalf("query node.id default: %v", err)
|
||||
}
|
||||
if !columnDefault.Valid || !strings.Contains(strings.ToLower(columnDefault.String), "nextval(") {
|
||||
t.Fatalf("expected node.id default to be nextval(...), got %q", columnDefault.String)
|
||||
}
|
||||
|
||||
jwtSecret := "postgres-contract-secret"
|
||||
router := httpserver.NewRouter(handler.New(repo, jwtSecret), jwtSecret)
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, jwtSecret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
body := strings.NewReader(`{"name":"pg-repair-node","serverIp":"10.77.0.10"}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/node/create", body)
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
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 {
|
||||
t.Fatalf("query created node: %v", err)
|
||||
}
|
||||
if nodeID <= 0 {
|
||||
t.Fatalf("expected positive node id, got %d", nodeID)
|
||||
}
|
||||
}
|
||||
|
||||
func withSearchPath(dsn, schema string) (string, error) {
|
||||
u, err := url.Parse(dsn)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("search_path", schema)
|
||||
u.RawQuery = q.Encode()
|
||||
return u.String(), nil
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -144,6 +145,17 @@ func TestTunnelUpdateAssignsChainPortsContract(t *testing.T) {
|
||||
if outPort <= 0 {
|
||||
t.Fatalf("expected out node port to be assigned, got %d", outPort)
|
||||
}
|
||||
|
||||
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 {
|
||||
t.Fatalf("query entry strategy: %v", err)
|
||||
}
|
||||
if !entryStrategy.Valid || strings.TrimSpace(entryStrategy.String) == "" {
|
||||
t.Fatalf("expected entry strategy to be non-null and non-empty")
|
||||
}
|
||||
if entryStrategy.String != "round" {
|
||||
t.Fatalf("expected entry strategy round, got %q", entryStrategy.String)
|
||||
}
|
||||
}
|
||||
|
||||
func jsonInt(v int64) string {
|
||||
|
||||
@@ -0,0 +1,301 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestTunnelCreateWithIPPreferenceContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
insertDualStackNode := func(name, v4, v6, portRange string) int64 {
|
||||
res, 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 {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
id, err := res.LastInsertId()
|
||||
if err != nil {
|
||||
t.Fatalf("get node id %s: %v", name, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
entryID := insertDualStackNode("ip-pref-entry", "10.50.0.1", "2001:db8::1", "50000-50010")
|
||||
exitID := insertDualStackNode("ip-pref-exit", "10.50.0.2", "2001:db8::2", "51000-51010")
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
preference string
|
||||
}{
|
||||
{"v4-preference", "v4"},
|
||||
{"v6-preference", "v6"},
|
||||
{"empty-preference", ""},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
payload := `{"name":"tunnel-` + tc.name + `","type":2,"flow":99999,"status":1,"ipPreference":"` + tc.preference + `","inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
|
||||
var stored string
|
||||
err := repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE name = ?`, "tunnel-"+tc.name).Scan(&stored)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
t.Skipf("tunnel not created (nodes offline), skipping DB verification")
|
||||
}
|
||||
t.Fatalf("query ip_preference: %v", err)
|
||||
}
|
||||
if stored != tc.preference {
|
||||
t.Fatalf("expected ip_preference=%q in DB, got %q", tc.preference, stored)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateIPPreferenceContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
insertDualStackNode := func(name, v4, v6, portRange string) int64 {
|
||||
res, 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 {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
id, err := res.LastInsertId()
|
||||
if err != nil {
|
||||
t.Fatalf("get node id %s: %v", name, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
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(`
|
||||
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 {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID, err := tunnelRes.LastInsertId()
|
||||
if err != nil {
|
||||
t.Fatalf("get tunnel id: %v", err)
|
||||
}
|
||||
|
||||
payload := `{"id":` + jsonInt(tunnelID) + `,"name":"update-ip-pref-tunnel","type":2,"flow":99999,"trafficRatio":1.0,"status":1,"ipPreference":"v6","inNodeId":[{"nodeId":` + jsonInt(entryID) + `,"protocol":"tls"}],"chainNodes":[],"outNodeId":[{"nodeId":` + jsonInt(exitID) + `,"protocol":"tls"}]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
|
||||
var stored string
|
||||
if err := repo.DB().QueryRow(`SELECT COALESCE(ip_preference, '') FROM tunnel WHERE id = ?`, tunnelID).Scan(&stored); err != nil {
|
||||
t.Fatalf("query ip_preference: %v", err)
|
||||
}
|
||||
if stored != "v6" {
|
||||
t.Fatalf("expected ip_preference='v6' after update, got %q", stored)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelListReturnsIPPreferenceContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
_, 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")
|
||||
if err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d (msg=%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
tunnels, ok := out.Data.([]interface{})
|
||||
if !ok || len(tunnels) == 0 {
|
||||
t.Fatalf("expected non-empty tunnel list, got %v", out.Data)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, raw := range tunnels {
|
||||
tm, ok := raw.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if tm["name"] == "list-ip-pref-tunnel" {
|
||||
found = true
|
||||
pref, _ := tm["ipPreference"].(string)
|
||||
if pref != "v6" {
|
||||
t.Fatalf("expected ipPreference='v6' in list response, got %q", pref)
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("tunnel 'list-ip-pref-tunnel' not found in list response")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPPreferenceColumnDefaultContract(t *testing.T) {
|
||||
_, repo := setupContractRouter(t, "contract-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
_, 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 {
|
||||
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 {
|
||||
t.Fatalf("query ip_preference: %v", err)
|
||||
}
|
||||
if stored != "" {
|
||||
t.Fatalf("expected default ip_preference='', got %q", stored)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
if err != nil {
|
||||
t.Fatalf("check column existence: %v", err)
|
||||
}
|
||||
if colCount != 1 {
|
||||
t.Fatalf("expected ip_preference column to exist in tunnel table, found %d", colCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPPreferenceCoalesceNullSafety(t *testing.T) {
|
||||
_, repo := setupContractRouter(t, "contract-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
_, 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 {
|
||||
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 {
|
||||
t.Fatalf("query ip_preference: %v", err)
|
||||
}
|
||||
if stored != "" {
|
||||
t.Fatalf("COALESCE should convert NULL to empty string, got %q", stored)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDualStackNodeIPFieldsStoredContract(t *testing.T) {
|
||||
_, repo := setupContractRouter(t, "contract-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
_, 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 {
|
||||
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 {
|
||||
t.Fatalf("query node IPs: %v", err)
|
||||
}
|
||||
if !v4.Valid || v4.String != "10.70.0.1" {
|
||||
t.Fatalf("expected server_ip_v4='10.70.0.1', got %v", v4)
|
||||
}
|
||||
if !v6.Valid || v6.String != "2001:db8:2::1" {
|
||||
t.Fatalf("expected server_ip_v6='2001:db8:2::1', got %v", v6)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPPreferenceValidValuesContract(t *testing.T) {
|
||||
_, repo := setupContractRouter(t, "contract-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
for _, pref := range []string{"", "v4", "v6"} {
|
||||
name := "valid-pref-" + pref
|
||||
if pref == "" {
|
||||
name = "valid-pref-empty"
|
||||
}
|
||||
_, 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 {
|
||||
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 {
|
||||
t.Fatalf("query ip_preference for %s: %v", name, err)
|
||||
}
|
||||
if stored != pref {
|
||||
t.Fatalf("expected ip_preference=%q, got %q for %s", pref, stored, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
+6
-1
@@ -1,6 +1,6 @@
|
||||
# GO-GOST SERVICE KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Mon Feb 02 2026
|
||||
**Generated:** Sun Feb 15 2026
|
||||
|
||||
## OVERVIEW
|
||||
Forwarding agent built on GOST v3 with a local fork of `github.com/go-gost/x` under `x/`.
|
||||
@@ -27,6 +27,11 @@ go-gost/
|
||||
## CONVENTIONS
|
||||
- Two configs exist: panel integration uses `config.json`; forwarding services use GOST config (defaults to `gost.{json,yaml}` via viper search paths).
|
||||
- `go-gost/x/` is the primary extension surface; avoid editing vendored deps.
|
||||
- Agent communicates with panel via WebSocket (real-time commands) + HTTP (batch traffic reports).
|
||||
- All panel communication uses AES encryption with node `secret` as PSK.
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- **DO NOT EDIT** generated protobuf in `x/internal/util/grpc/proto/`.
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
|
||||
+3
-1
@@ -1,7 +1,7 @@
|
||||
# GO-GOST/X KNOWLEDGE BASE
|
||||
|
||||
## OVERVIEW
|
||||
Local fork of `github.com/go-gost/x` used by `go-gost/` via `replace github.com/go-gost/x => ./x`. Most protocol/runtime behavior changes happen here.
|
||||
Local fork of `github.com/go-gost/x` used by `go-gost/` via `replace github.com/go-gost/x => ./x`. Most protocol/runtime behavior changes happen here. 30+ top-level packages - framework-style layout.
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
@@ -31,6 +31,8 @@ go-gost/x/
|
||||
## CONVENTIONS
|
||||
- `go-gost/x/` is a standalone Go module (`go-gost/x/go.mod`); run go tooling from this dir when debugging module resolution.
|
||||
- Generated gRPC/proto code lives under `go-gost/x/internal/util/grpc/proto/`.
|
||||
- Handlers/listeners/dialers follow consistent pattern: `{type}.go` + `metadata.go` per protocol.
|
||||
- OS-specific code uses `name_[os].go` suffix (e.g., `tun_linux.go`, `tun_darwin.go`).
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- Do not edit generated files in `go-gost/x/internal/util/grpc/proto/` (`*.pb.go`, `*_grpc.pb.go`).
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
# GOST CONNECTOR KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Feb 13 2026
|
||||
|
||||
## OVERVIEW
|
||||
Connection initiators (clients) for various protocols in GOST forwarding.
|
||||
**Stack:** Go, GOST core.
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
connector/
|
||||
├── direct/ # Direct connection
|
||||
├── forward/ # Forward proxy
|
||||
├── http/ # HTTP connector
|
||||
├── http2/ # HTTP/2 connector
|
||||
├── relay/ # Relay protocol
|
||||
├── router/ # Router connector
|
||||
├── serial/ # Serial port
|
||||
├── sni/ # SNI routing
|
||||
├── socks/ # SOCKS4/5
|
||||
├── ss/ # Shadowsocks
|
||||
├── sshd/ # SSH daemon
|
||||
├── tcp/ # TCP connector
|
||||
├── tunnel/ # Tunnel mode
|
||||
└── unix/ # Unix socket
|
||||
```
|
||||
|
||||
## CONVENTIONS
|
||||
- Inherits from parent `go-gost/x/` conventions.
|
||||
- Each subdir implements `Connector` interface from GOST core.
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- DO NOT EDIT generated protobuf in `go-gost/x/internal/util/grpc/proto/`.
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
cd go-gost
|
||||
go test ./x/connector/...
|
||||
```
|
||||
@@ -0,0 +1,38 @@
|
||||
# GOST SOCKET KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
|
||||
## OVERVIEW
|
||||
WebSocket reporter and socket utilities for panel integration.
|
||||
**Stack:** Go, GOST core, gorilla/websocket.
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
socket/
|
||||
├── websocket_reporter.go # Agent-to-panel telemetry (1504 LOC)
|
||||
├── service.go # Socket service orchestration (534 LOC)
|
||||
├── socket.go # Core socket interface
|
||||
├── udp.go # UDP socket handling
|
||||
├── packet.go # Packet framing
|
||||
└── packetconn.go # Packet connection wrapper
|
||||
```
|
||||
|
||||
## WHERE TO LOOK
|
||||
| Task | Location | Notes |
|
||||
|------|----------|-------|
|
||||
| **Panel Reporting** | `websocket_reporter.go` | Real-time system info (CPU, mem, uptime) every 2s |
|
||||
| **Command Handling** | `websocket_reporter.go` | Processes `AddService`, `UpgradeAgent`, etc. |
|
||||
|
||||
## CONVENTIONS
|
||||
- Inherits from parent `go-gost/x/` conventions.
|
||||
- Low-level network primitives.
|
||||
- All panel communication is AES-encrypted using node `secret`.
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- DO NOT EDIT generated protobuf.
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
cd go-gost
|
||||
go test ./x/socket/...
|
||||
```
|
||||
+14
-4
@@ -1,10 +1,10 @@
|
||||
# VITE FRONTEND KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Mon Feb 02 2026
|
||||
**Generated:** Sun Feb 15 2026
|
||||
|
||||
## OVERVIEW
|
||||
Web management console for FLVX (formerly Flux Panel).
|
||||
**Stack:** React 18, Vite 5, TypeScript, TailwindCSS 4, HeroUI.
|
||||
Web management console for FLVX.
|
||||
**Stack:** React 18, Vite 5 (rolldown-vite), TypeScript, TailwindCSS 4, HeroUI.
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
@@ -37,9 +37,19 @@ vite-frontend/
|
||||
|
||||
## CONVENTIONS
|
||||
- **Auth**: JWT stored as `localStorage.token`. Sent in `Authorization` header (no "Bearer" prefix).
|
||||
- **API**: Default base URL is `/api/v1/`.
|
||||
- **API**: Default base URL is `/api/v1/`. Responses follow `{code, msg, data, ts}` structure.
|
||||
- **WebView**: In WebView mode, base URL is derived from selected panel address. If unset, API returns `code: -1`.
|
||||
- **Routing**: URL query param `h5=true` forces mobile layout.
|
||||
- **Build**: `minify: false`, `treeshake: false` - unoptimized production bundles for debugging.
|
||||
- **ESLint**: `react-hooks/exhaustive-deps` disabled, unused vars starting with `_` ignored.
|
||||
- **Large Pages**: `forward.tsx` (3263 LOC), `tunnel.tsx` (2552 LOC), `node.tsx` (2194 LOC).
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- **DO NOT ADD** tests - no test infrastructure (Vitest/Jest not configured).
|
||||
|
||||
## NOTES
|
||||
- Uses `rolldown-vite` (experimental Rust bundler) instead of standard Vite.
|
||||
- ESLint Flat Config format with custom import ordering rules.
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
|
||||
@@ -140,6 +140,10 @@ export const updateConfigs = (configMap: Record<string, string>) =>
|
||||
export const updateConfig = (name: string, value: string) =>
|
||||
Network.post("/config/update-single", { name, value });
|
||||
|
||||
export const exportBackupData = () => Network.post("/backup/export");
|
||||
export const importBackupData = (data: any) => Network.post("/backup/import", data);
|
||||
export const restoreBackupData = (data: any) => Network.post("/backup/restore", data);
|
||||
|
||||
// 验证码相关接口
|
||||
export const checkCaptcha = () => Network.post("/captcha/check");
|
||||
export const generateCaptcha = () => Network.post(`/captcha/generate`);
|
||||
@@ -281,3 +285,11 @@ export const exportBackup = async (types: string[] = []) => {
|
||||
|
||||
export const importBackup = (data: { types: string[]; [key: string]: any }) =>
|
||||
Network.post("/backup/import", data);
|
||||
|
||||
export interface AnnouncementData {
|
||||
content: string;
|
||||
enabled: number;
|
||||
}
|
||||
|
||||
export const getAnnouncement = () => Network.get<AnnouncementData>("/announcement/get");
|
||||
export const updateAnnouncement = (data: AnnouncementData) => Network.post("/announcement/update", data);
|
||||
|
||||
@@ -3,6 +3,7 @@ import { useNavigate } from "react-router-dom";
|
||||
import { Button } from "@heroui/button";
|
||||
import { Card, CardBody, CardHeader } from "@heroui/card";
|
||||
import { Input } from "@heroui/input";
|
||||
import { Textarea } from "@heroui/input";
|
||||
import { Spinner } from "@heroui/spinner";
|
||||
import { Divider } from "@heroui/divider";
|
||||
import { Switch } from "@heroui/switch";
|
||||
@@ -10,7 +11,7 @@ import { Select, SelectItem } from "@heroui/select";
|
||||
import { Checkbox, CheckboxGroup } from "@heroui/checkbox";
|
||||
import toast from "react-hot-toast";
|
||||
|
||||
import { updateConfigs, exportBackup, importBackup } from "@/api";
|
||||
import { updateConfigs, exportBackup, importBackup, getAnnouncement, updateAnnouncement, type AnnouncementData } from "@/api";
|
||||
import { SettingsIcon } from "@/components/icons";
|
||||
import { isAdmin } from "@/utils/auth";
|
||||
import {
|
||||
@@ -144,6 +145,13 @@ export default function ConfigPage() {
|
||||
const [importFileName, setImportFileName] = useState("");
|
||||
const fileInputRef = useRef<HTMLInputElement>(null);
|
||||
|
||||
const [announcement, setAnnouncement] = useState<AnnouncementData>({
|
||||
content: "",
|
||||
enabled: 0,
|
||||
});
|
||||
const [announcementLoading, setAnnouncementLoading] = useState(true);
|
||||
const [announcementSaving, setAnnouncementSaving] = useState(false);
|
||||
|
||||
// 权限检查
|
||||
useEffect(() => {
|
||||
if (!isAdmin()) {
|
||||
@@ -188,21 +196,51 @@ export default function ConfigPage() {
|
||||
};
|
||||
|
||||
useEffect(() => {
|
||||
// 延迟加载,避免阻塞初始渲染
|
||||
const timer = setTimeout(() => {
|
||||
loadConfigs(initialConfigs);
|
||||
loadAnnouncement();
|
||||
}, 100);
|
||||
|
||||
return () => clearTimeout(timer);
|
||||
}, []); // 只在组件挂载时执行一次
|
||||
}, []);
|
||||
|
||||
const loadAnnouncement = async () => {
|
||||
setAnnouncementLoading(true);
|
||||
try {
|
||||
const res = await getAnnouncement();
|
||||
|
||||
if (res.code === 0 && res.data) {
|
||||
setAnnouncement(res.data);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Failed to load announcement:", error);
|
||||
} finally {
|
||||
setAnnouncementLoading(false);
|
||||
}
|
||||
};
|
||||
|
||||
const saveAnnouncement = async () => {
|
||||
setAnnouncementSaving(true);
|
||||
try {
|
||||
const res = await updateAnnouncement(announcement);
|
||||
|
||||
if (res.code === 0) {
|
||||
toast.success("公告保存成功");
|
||||
} else {
|
||||
toast.error(res.msg || "保存失败");
|
||||
}
|
||||
} catch {
|
||||
toast.error("保存公告失败,请重试");
|
||||
} finally {
|
||||
setAnnouncementSaving(false);
|
||||
}
|
||||
};
|
||||
|
||||
// 处理配置项变更
|
||||
const handleConfigChange = (key: string, value: string) => {
|
||||
const newConfigs = { ...configs, [key]: value };
|
||||
|
||||
setConfigs(newConfigs);
|
||||
|
||||
// 检查是否有变更
|
||||
const hasChangesNow =
|
||||
Object.keys(newConfigs).some(
|
||||
(k) => newConfigs[k] !== originalConfigs[k],
|
||||
@@ -479,7 +517,6 @@ export default function ConfigPage() {
|
||||
</CardBody>
|
||||
</Card>
|
||||
|
||||
{/* 操作提示 */}
|
||||
{hasChanges && (
|
||||
<Card className="mt-4 bg-warning-50 dark:bg-warning-900/20 border-warning-200 dark:border-warning-800">
|
||||
<CardBody className="py-3">
|
||||
@@ -493,6 +530,69 @@ export default function ConfigPage() {
|
||||
</Card>
|
||||
)}
|
||||
|
||||
<Card className="mt-6 shadow-md">
|
||||
<CardHeader className="pb-4">
|
||||
<div className="flex justify-between items-center w-full">
|
||||
<div>
|
||||
<h2 className="text-xl font-semibold">公告管理</h2>
|
||||
<p className="text-sm text-gray-600 dark:text-gray-400">
|
||||
设置首页显示的公告内容
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
</CardHeader>
|
||||
|
||||
<Divider />
|
||||
|
||||
<CardBody className="space-y-4 pt-6">
|
||||
{announcementLoading ? (
|
||||
<div className="flex justify-center py-8">
|
||||
<Spinner size="lg" />
|
||||
</div>
|
||||
) : (
|
||||
<>
|
||||
<div className="space-y-2">
|
||||
<Switch
|
||||
isSelected={announcement.enabled === 1}
|
||||
onValueChange={(checked) =>
|
||||
setAnnouncement({ ...announcement, enabled: checked ? 1 : 0 })
|
||||
}
|
||||
>
|
||||
<span className="text-sm text-gray-700 dark:text-gray-300">
|
||||
{announcement.enabled === 1 ? "已启用" : "已禁用"}
|
||||
</span>
|
||||
</Switch>
|
||||
<p className="text-xs text-gray-500 dark:text-gray-400">
|
||||
启用后,公告将在首页顶部显示
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<Textarea
|
||||
label="公告内容"
|
||||
placeholder="请输入公告内容"
|
||||
value={announcement.content}
|
||||
variant="bordered"
|
||||
minRows={4}
|
||||
onChange={(e) =>
|
||||
setAnnouncement({ ...announcement, content: e.target.value })
|
||||
}
|
||||
/>
|
||||
|
||||
<div className="flex justify-end">
|
||||
<Button
|
||||
color="primary"
|
||||
isLoading={announcementSaving}
|
||||
startContent={<SaveIcon className="w-4 h-4" />}
|
||||
onClick={saveAnnouncement}
|
||||
>
|
||||
保存公告
|
||||
</Button>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
</CardBody>
|
||||
</Card>
|
||||
|
||||
{/* 备份与恢复 */}
|
||||
<Card className="mt-6 shadow-md">
|
||||
<CardHeader className="pb-4">
|
||||
|
||||
@@ -13,7 +13,7 @@ import {
|
||||
ResponsiveContainer,
|
||||
} from "recharts";
|
||||
|
||||
import { getUserPackageInfo } from "@/api";
|
||||
import { getUserPackageInfo, getAnnouncement, type AnnouncementData } from "@/api";
|
||||
|
||||
interface UserInfo {
|
||||
flow: number;
|
||||
@@ -71,6 +71,7 @@ export default function DashboardPage() {
|
||||
const [forwardList, setForwardList] = useState<Forward[]>([]);
|
||||
const [statisticsFlows, setStatisticsFlows] = useState<StatisticsFlow[]>([]);
|
||||
const [isAdmin, setIsAdmin] = useState(false);
|
||||
const [announcement, setAnnouncement] = useState<AnnouncementData | null>(null);
|
||||
|
||||
const [addressModalOpen, setAddressModalOpen] = useState(false);
|
||||
const [addressModalTitle, setAddressModalTitle] = useState("");
|
||||
@@ -170,22 +171,33 @@ export default function DashboardPage() {
|
||||
};
|
||||
|
||||
useEffect(() => {
|
||||
// 重置状态并加载数据,防止页面切换时显示旧数据
|
||||
setLoading(true);
|
||||
setUserInfo({} as UserInfo);
|
||||
setUserTunnels([]);
|
||||
setForwardList([]);
|
||||
setStatisticsFlows([]);
|
||||
|
||||
// 检查用户是否是管理员
|
||||
const adminStatus = localStorage.getItem("admin");
|
||||
|
||||
setIsAdmin(adminStatus === "true");
|
||||
|
||||
loadPackageData();
|
||||
loadAnnouncement();
|
||||
localStorage.setItem("e", "/dashboard");
|
||||
}, []);
|
||||
|
||||
const loadAnnouncement = async () => {
|
||||
try {
|
||||
const res = await getAnnouncement();
|
||||
|
||||
if (res.code === 0 && res.data && res.data.enabled === 1) {
|
||||
setAnnouncement(res.data);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Failed to load announcement:", error);
|
||||
}
|
||||
};
|
||||
|
||||
const loadPackageData = async () => {
|
||||
setLoading(true);
|
||||
try {
|
||||
@@ -703,7 +715,35 @@ export default function DashboardPage() {
|
||||
|
||||
return (
|
||||
<div className="px-3 lg:px-6 py-2 lg:py-4">
|
||||
{/* 响应式统计卡片 */}
|
||||
{announcement && announcement.content && (
|
||||
<Card className="mb-4 lg:mb-6 border border-blue-200 dark:border-blue-500/30 bg-gradient-to-r from-blue-50 to-purple-50 dark:from-blue-500/10 dark:to-purple-500/10">
|
||||
<CardBody className="p-4">
|
||||
<div className="flex items-start gap-3">
|
||||
<div className="p-2 bg-blue-100 dark:bg-blue-500/20 rounded-lg flex-shrink-0">
|
||||
<svg
|
||||
className="w-5 h-5 text-blue-600 dark:text-blue-400"
|
||||
fill="currentColor"
|
||||
viewBox="0 0 20 20"
|
||||
>
|
||||
<path
|
||||
clipRule="evenodd"
|
||||
d="M18 10a8 8 0 11-16 0 8 8 0 0116 0zm-7-4a1 1 0 11-2 0 1 1 0 012 0zM9 9a1 1 0 000 2v3a1 1 0 001 1h1a1 1 0 100-2v-3a1 1 0 00-1-1H9z"
|
||||
fillRule="evenodd"
|
||||
/>
|
||||
</svg>
|
||||
</div>
|
||||
<div className="flex-1 min-w-0">
|
||||
<h3 className="text-sm lg:text-base font-semibold text-blue-900 dark:text-blue-100 mb-1">
|
||||
公告
|
||||
</h3>
|
||||
<p className="text-xs lg:text-sm text-blue-800 dark:text-blue-200 whitespace-pre-wrap break-words">
|
||||
{announcement.content}
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
</CardBody>
|
||||
</Card>
|
||||
)}
|
||||
<div className="grid grid-cols-2 lg:grid-cols-4 gap-3 lg:gap-4 mb-6 lg:mb-8">
|
||||
<Card className="border border-gray-200 dark:border-default-200 shadow-md hover:shadow-lg transition-shadow">
|
||||
<CardBody className="p-3 lg:p-4">
|
||||
|
||||
@@ -67,6 +67,7 @@ interface Tunnel {
|
||||
protocol?: string;
|
||||
flow: number; // 1: 单向, 2: 双向
|
||||
trafficRatio: number;
|
||||
ipPreference?: string;
|
||||
status: number;
|
||||
createdTime: string;
|
||||
}
|
||||
@@ -87,6 +88,7 @@ interface TunnelForm {
|
||||
flow: number;
|
||||
trafficRatio: number;
|
||||
inIp: string; // 入口IP
|
||||
ipPreference: string;
|
||||
status: number;
|
||||
}
|
||||
|
||||
@@ -141,6 +143,7 @@ export default function TunnelPage() {
|
||||
flow: 1,
|
||||
trafficRatio: 1.0,
|
||||
inIp: "",
|
||||
ipPreference: "",
|
||||
status: 1,
|
||||
});
|
||||
|
||||
@@ -301,6 +304,7 @@ export default function TunnelPage() {
|
||||
flow: 1,
|
||||
trafficRatio: 1.0,
|
||||
inIp: "",
|
||||
ipPreference: "",
|
||||
status: 1,
|
||||
});
|
||||
setErrors({});
|
||||
@@ -313,21 +317,22 @@ export default function TunnelPage() {
|
||||
|
||||
// 直接使用列表数据,getAllTunnels 已经包含完整的节点信息
|
||||
setForm({
|
||||
id: tunnel.id,
|
||||
name: tunnel.name,
|
||||
type: tunnel.type,
|
||||
inNodeId: tunnel.inNodeId || [],
|
||||
outNodeId: tunnel.outNodeId || [],
|
||||
chainNodes: tunnel.chainNodes || [],
|
||||
flow: tunnel.flow,
|
||||
trafficRatio: tunnel.trafficRatio,
|
||||
inIp: tunnel.inIp
|
||||
? tunnel.inIp
|
||||
.split(",")
|
||||
.map((ip) => ip.trim())
|
||||
.join("\n")
|
||||
: "",
|
||||
status: tunnel.status,
|
||||
id: tunnel.id,
|
||||
name: tunnel.name,
|
||||
type: tunnel.type,
|
||||
inNodeId: tunnel.inNodeId || [],
|
||||
outNodeId: tunnel.outNodeId || [],
|
||||
chainNodes: tunnel.chainNodes || [],
|
||||
flow: tunnel.flow,
|
||||
trafficRatio: tunnel.trafficRatio,
|
||||
inIp: tunnel.inIp
|
||||
? tunnel.inIp
|
||||
.split(",")
|
||||
.map((ip: string) => ip.trim())
|
||||
.join("\n")
|
||||
: "",
|
||||
ipPreference: tunnel.ipPreference || "",
|
||||
status: tunnel.status,
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
@@ -1045,7 +1050,7 @@ export default function TunnelPage() {
|
||||
</div>
|
||||
|
||||
{/* 流量配置 */}
|
||||
<div className="grid grid-cols-2 gap-2">
|
||||
<div className={`grid gap-2 ${tunnel.type === 2 && tunnel.ipPreference ? "grid-cols-3" : "grid-cols-2"}`}>
|
||||
<div className="text-center p-1.5 bg-default-50 dark:bg-default-100/30 rounded">
|
||||
<div className="text-xs text-default-500">
|
||||
流量计算
|
||||
@@ -1062,6 +1067,16 @@ export default function TunnelPage() {
|
||||
{tunnel.trafficRatio}x
|
||||
</div>
|
||||
</div>
|
||||
{tunnel.type === 2 && tunnel.ipPreference && (
|
||||
<div className="text-center p-1.5 bg-default-50 dark:bg-default-100/30 rounded">
|
||||
<div className="text-xs text-default-500">
|
||||
连接偏好
|
||||
</div>
|
||||
<div className="text-sm font-semibold text-foreground mt-0.5">
|
||||
{tunnel.ipPreference === "v4" ? "IPv4" : "IPv6"}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -1293,6 +1308,27 @@ export default function TunnelPage() {
|
||||
}
|
||||
/>
|
||||
|
||||
{form.type === 2 && (
|
||||
<Select
|
||||
description="当节点同时拥有IPv4和IPv6地址时,选择隧道连接使用的地址类型"
|
||||
label="隧道连接地址偏好"
|
||||
selectedKeys={[form.ipPreference || ""]}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as string;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
ipPreference: selectedKey || "",
|
||||
}));
|
||||
}}
|
||||
>
|
||||
<SelectItem key="">自动选择</SelectItem>
|
||||
<SelectItem key="v4">优先IPv4</SelectItem>
|
||||
<SelectItem key="v6">优先IPv6</SelectItem>
|
||||
</Select>
|
||||
)}
|
||||
|
||||
<Divider />
|
||||
<h3 className="text-lg font-semibold">入口配置</h3>
|
||||
|
||||
|
||||
@@ -46,6 +46,28 @@ html, body {
|
||||
--safe-area-bottom: env(safe-area-inset-bottom, 0px);
|
||||
}
|
||||
|
||||
[data-slot="input-wrapper"] {
|
||||
box-shadow: none;
|
||||
}
|
||||
|
||||
[data-slot="input-wrapper"]:focus-within:not([data-invalid="true"]),
|
||||
[data-slot="input-wrapper"][data-focus="true"]:not([data-invalid="true"]),
|
||||
[data-slot="input-wrapper"][data-focused="true"]:not([data-invalid="true"]),
|
||||
button[data-slot="trigger"][data-focus="true"]:not([data-invalid="true"]),
|
||||
button[data-slot="trigger"][data-open="true"]:not([data-invalid="true"]) {
|
||||
border-color: var(--heroui-default-200, #e5e7eb) !important;
|
||||
outline: none !important;
|
||||
outline-offset: 0 !important;
|
||||
box-shadow: none !important;
|
||||
}
|
||||
|
||||
[data-slot="input-wrapper"] input:focus,
|
||||
[data-slot="input-wrapper"] textarea:focus {
|
||||
outline: none !important;
|
||||
box-shadow: none !important;
|
||||
border-color: transparent !important;
|
||||
}
|
||||
|
||||
.safe-top {
|
||||
padding-top: var(--safe-area-top);
|
||||
}
|
||||
@@ -85,4 +107,4 @@ html, body {
|
||||
}
|
||||
}
|
||||
|
||||
@config "../../tailwind.config.js"
|
||||
@config "../../tailwind.config.js"
|
||||
|
||||
Reference in New Issue
Block a user