mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-05 01:26:37 +08:00
Merge branch 'origin/main' into opencode/gentle-comet
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
FROM golang:1.23-bookworm AS builder
|
||||
FROM golang:1.24-bookworm AS builder
|
||||
WORKDIR /src
|
||||
|
||||
COPY go.mod ./
|
||||
COPY go.mod go.sum ./
|
||||
RUN go mod download
|
||||
|
||||
COPY . .
|
||||
|
||||
+8
-1
@@ -1,22 +1,29 @@
|
||||
module go-backend
|
||||
|
||||
go 1.23.0
|
||||
go 1.24.0
|
||||
|
||||
toolchain go1.24.4
|
||||
|
||||
require (
|
||||
github.com/gorilla/websocket v1.5.3
|
||||
github.com/jackc/pgx/v5 v5.7.3
|
||||
modernc.org/sqlite v1.37.1
|
||||
)
|
||||
|
||||
require (
|
||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||
github.com/google/uuid v1.6.0 // indirect
|
||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/ncruces/go-strftime v0.1.9 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
golang.org/x/crypto v0.31.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect
|
||||
golang.org/x/sync v0.17.0 // indirect
|
||||
golang.org/x/sys v0.33.0 // indirect
|
||||
golang.org/x/text v0.29.0 // indirect
|
||||
modernc.org/libc v1.65.7 // indirect
|
||||
modernc.org/mathutil v1.7.1 // indirect
|
||||
modernc.org/memory v1.11.0 // indirect
|
||||
|
||||
+32
-6
@@ -1,3 +1,6 @@
|
||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
|
||||
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
||||
@@ -6,23 +9,46 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
|
||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||
github.com/jackc/pgx/v5 v5.7.3 h1:PO1wNKj/bTAwxSJnO1Z4Ai8j4magtqg2SLNjEDzcXQo=
|
||||
github.com/jackc/pgx/v5 v5.7.3/go.mod h1:ncY89UGWxg82EykZUwSpUKEfccBGGYq1xjrOpsbsfGQ=
|
||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||
github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4=
|
||||
github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||
github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk=
|
||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
||||
golang.org/x/crypto v0.31.0 h1:ihbySMvVjLAeSH1IbfcRTkD/iNscyz8rGzjF/E5hV6U=
|
||||
golang.org/x/crypto v0.31.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM=
|
||||
golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8=
|
||||
golang.org/x/mod v0.24.0 h1:ZfthKaKaT4NrhGVZHO1/WDTwGES4De8KtWO0SIbNJMU=
|
||||
golang.org/x/mod v0.24.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww=
|
||||
golang.org/x/sync v0.14.0 h1:woo0S4Yywslg6hp4eUFjTVOyKt0RookbpAHG4c1HmhQ=
|
||||
golang.org/x/sync v0.14.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
|
||||
golang.org/x/mod v0.27.0 h1:kb+q2PyFnEADO2IEF935ehFUXlWiNjJWtRNgBLSfbxQ=
|
||||
golang.org/x/mod v0.27.0/go.mod h1:rWI627Fq0DEoudcK+MBkNkCe0EetEaDSwJJkCcjpazc=
|
||||
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
|
||||
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
|
||||
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
|
||||
golang.org/x/tools v0.33.0 h1:4qz2S3zmRxbGIhDIAgjxvFutSvH5EfnsYrRBj0UI0bc=
|
||||
golang.org/x/tools v0.33.0/go.mod h1:CIJMaWEY88juyUfo7UbgPqbC8rU2OqfAV1h2Qp0oMYI=
|
||||
golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
|
||||
golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
|
||||
golang.org/x/tools v0.36.0 h1:kWS0uv/zsvHEle1LbV5LE8QujrxB3wfQyxHfhOk0Qkg=
|
||||
golang.org/x/tools v0.36.0/go.mod h1:WBDiHKJK8YgLHlcQPYQzNCkUxUypCaa5ZegCVutKm+s=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
modernc.org/cc/v4 v4.26.1 h1:+X5NtzVBn0KgsBCBe+xkDC7twLb/jNVj9FPgiwSQO3s=
|
||||
modernc.org/cc/v4 v4.26.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
||||
modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU=
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/config"
|
||||
@@ -20,9 +21,24 @@ type App struct {
|
||||
}
|
||||
|
||||
func New(cfg config.Config) (*App, error) {
|
||||
repo, err := sqlite.Open(cfg.DBPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open sqlite: %w", err)
|
||||
var (
|
||||
repo *sqlite.Repository
|
||||
err error
|
||||
)
|
||||
|
||||
switch strings.ToLower(strings.TrimSpace(cfg.DBType)) {
|
||||
case "", "sqlite":
|
||||
repo, err = sqlite.Open(cfg.DBPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open sqlite: %w", err)
|
||||
}
|
||||
case "postgres", "postgresql":
|
||||
repo, err = sqlite.OpenPostgres(cfg.DatabaseURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open postgres: %w", err)
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported DB_TYPE %q", cfg.DBType)
|
||||
}
|
||||
|
||||
h := handler.New(repo, cfg.JWTSecret)
|
||||
|
||||
@@ -3,16 +3,22 @@ package config
|
||||
import "os"
|
||||
|
||||
type Config struct {
|
||||
Addr string
|
||||
DBPath string
|
||||
JWTSecret string
|
||||
Addr string
|
||||
DBType string
|
||||
DBPath string
|
||||
DatabaseURL string
|
||||
JWTSecret string
|
||||
LogDir string
|
||||
}
|
||||
|
||||
func FromEnv() Config {
|
||||
cfg := Config{
|
||||
Addr: getEnv("SERVER_ADDR", ":6365"),
|
||||
DBPath: getEnv("DB_PATH", "/app/data/gost.db"),
|
||||
JWTSecret: getEnv("JWT_SECRET", ""),
|
||||
Addr: getEnv("SERVER_ADDR", ":6365"),
|
||||
DBType: getEnv("DB_TYPE", "sqlite"),
|
||||
DBPath: getEnv("DB_PATH", "/app/data/gost.db"),
|
||||
DatabaseURL: getEnv("DATABASE_URL", ""),
|
||||
JWTSecret: getEnv("JWT_SECRET", ""),
|
||||
LogDir: getEnv("LOG_DIR", "/app/logs"),
|
||||
}
|
||||
|
||||
return cfg
|
||||
|
||||
@@ -85,6 +85,14 @@ func NewFederationClient() *FederationClient {
|
||||
}
|
||||
}
|
||||
|
||||
func NewFederationClientWithTimeout(timeout time.Duration) *FederationClient {
|
||||
return &FederationClient{
|
||||
client: &http.Client{
|
||||
Timeout: timeout,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *FederationClient) Connect(url, token, localDomain string) (*RemoteNodeInfo, error) {
|
||||
url = strings.TrimSuffix(url, "/")
|
||||
req, err := http.NewRequest("POST", url+"/api/v1/federation/connect", nil)
|
||||
|
||||
@@ -889,11 +889,11 @@ func firstPortFromRange(portRange string) int {
|
||||
|
||||
func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
|
||||
rows, err := h.repo.DB().Query(`
|
||||
SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
|
||||
SELECT CAST(ct.chain_type AS INTEGER), COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
|
||||
FROM chain_tunnel ct
|
||||
LEFT JOIN node n ON n.id = ct.node_id
|
||||
WHERE ct.tunnel_id = ?
|
||||
ORDER BY ct.chain_type ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
|
||||
ORDER BY CAST(ct.chain_type AS INTEGER) ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
|
||||
`, tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/client"
|
||||
@@ -40,6 +41,17 @@ type resetPeerShareFlowRequest struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
|
||||
type updatePeerShareRequest struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
MaxBandwidth int64 `json:"maxBandwidth"`
|
||||
ExpiryTime int64 `json:"expiryTime"`
|
||||
PortRangeStart int `json:"portRangeStart"`
|
||||
PortRangeEnd int `json:"portRangeEnd"`
|
||||
AllowedDomains string `json:"allowedDomains"`
|
||||
AllowedIPs string `json:"allowedIps"`
|
||||
}
|
||||
|
||||
type nodeImportRequest struct {
|
||||
RemoteURL string `json:"remoteUrl"`
|
||||
Token string `json:"token"`
|
||||
@@ -325,6 +337,80 @@ func (h *Handler) federationShareResetFlow(w http.ResponseWriter, r *http.Reques
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) federationShareUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid method"))
|
||||
return
|
||||
}
|
||||
|
||||
var req updatePeerShareRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid JSON"))
|
||||
return
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Share ID is required"))
|
||||
return
|
||||
}
|
||||
|
||||
share, err := h.repo.GetPeerShare(req.ID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if share == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("Share not found"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("Name is required"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.MaxBandwidth < 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Max bandwidth cannot be negative"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.ExpiryTime < 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Expiry time cannot be negative"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.PortRangeStart < 0 || req.PortRangeStart > 65535 || req.PortRangeEnd < 0 || req.PortRangeEnd > 65535 {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid port range"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.PortRangeStart > req.PortRangeEnd {
|
||||
response.WriteJSON(w, response.ErrDefault("Port range start cannot be greater than end"))
|
||||
return
|
||||
}
|
||||
|
||||
allowedIPs, err := normalizePeerShareAllowedIPs(req.AllowedIPs)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
share.Name = req.Name
|
||||
share.MaxBandwidth = req.MaxBandwidth
|
||||
share.ExpiryTime = req.ExpiryTime
|
||||
share.PortRangeStart = req.PortRangeStart
|
||||
share.PortRangeEnd = req.PortRangeEnd
|
||||
share.AllowedDomains = req.AllowedDomains
|
||||
share.AllowedIPs = allowedIPs
|
||||
share.UpdatedTime = time.Now().UnixMilli()
|
||||
|
||||
if err := h.repo.UpdatePeerShare(share); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid method"))
|
||||
@@ -700,21 +786,20 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
|
||||
defer tx.Rollback()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
var tunnelID int64
|
||||
err = tx.QueryRow(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?) RETURNING id`,
|
||||
tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?)`,
|
||||
fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort),
|
||||
tunnelType,
|
||||
req.Protocol,
|
||||
now,
|
||||
now,
|
||||
"",
|
||||
).Scan(&tunnelID)
|
||||
)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, 1, ?, ?, 'fifo', 0, ?)`,
|
||||
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, '1', ?, ?, 'fifo', 0, ?)`,
|
||||
tunnelID,
|
||||
share.NodeID,
|
||||
req.RemotePort,
|
||||
@@ -1316,6 +1401,74 @@ func isPeerIPAllowed(clientIP net.IP, whitelist string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *Handler) syncRemoteNodeStatuses(items []map[string]interface{}) {
|
||||
type remoteEntry struct {
|
||||
index int
|
||||
remoteURL string
|
||||
remoteToken string
|
||||
}
|
||||
|
||||
var remotes []remoteEntry
|
||||
for i, item := range items {
|
||||
isRemote, _ := item["isRemote"].(int)
|
||||
if isRemote != 1 {
|
||||
continue
|
||||
}
|
||||
url, _ := item["remoteUrl"].(string)
|
||||
token, _ := item["remoteToken"].(string)
|
||||
url = strings.TrimSpace(url)
|
||||
token = strings.TrimSpace(token)
|
||||
if url == "" || token == "" {
|
||||
continue
|
||||
}
|
||||
remotes = append(remotes, remoteEntry{index: i, remoteURL: url, remoteToken: token})
|
||||
}
|
||||
if len(remotes) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
localDomain := h.federationLocalDomain()
|
||||
fc := client.NewFederationClientWithTimeout(5 * time.Second)
|
||||
|
||||
type syncResult struct {
|
||||
index int
|
||||
status int
|
||||
syncError string
|
||||
}
|
||||
|
||||
results := make([]syncResult, len(remotes))
|
||||
var wg sync.WaitGroup
|
||||
for i, entry := range remotes {
|
||||
wg.Add(1)
|
||||
go func(idx int, e remoteEntry) {
|
||||
defer wg.Done()
|
||||
info, err := fc.Connect(e.remoteURL, e.remoteToken, localDomain)
|
||||
if err != nil {
|
||||
errMsg := err.Error()
|
||||
if strings.Contains(errMsg, "401") || strings.Contains(errMsg, "Invalid token") || strings.Contains(errMsg, "Unauthorized") {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: "provider_share_deleted"}
|
||||
} else if strings.Contains(errMsg, "403") || strings.Contains(errMsg, "Share is disabled") {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: "provider_share_disabled"}
|
||||
} else if strings.Contains(errMsg, "Share expired") {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: "provider_share_expired"}
|
||||
} else {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: errMsg}
|
||||
}
|
||||
} else {
|
||||
results[idx] = syncResult{index: e.index, status: info.Status, syncError: ""}
|
||||
}
|
||||
}(i, entry)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for _, r := range results {
|
||||
items[r.index]["status"] = r.status
|
||||
if r.syncError != "" {
|
||||
items[r.index]["syncError"] = r.syncError
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) cleanupPeerShareRuntimes(shareID int64) {
|
||||
if h == nil || h.repo == nil || shareID <= 0 {
|
||||
return
|
||||
|
||||
@@ -105,6 +105,10 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/node/update-order", h.nodeUpdateOrder)
|
||||
mux.HandleFunc("/api/v1/node/batch-delete", h.nodeBatchDelete)
|
||||
mux.HandleFunc("/api/v1/node/check-status", h.nodeCheckStatus)
|
||||
mux.HandleFunc("/api/v1/node/upgrade", h.nodeUpgrade)
|
||||
mux.HandleFunc("/api/v1/node/batch-upgrade", h.nodeBatchUpgrade)
|
||||
mux.HandleFunc("/api/v1/node/rollback", h.nodeRollback)
|
||||
mux.HandleFunc("/api/v1/node/releases", h.listReleases)
|
||||
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
|
||||
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
|
||||
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||
@@ -155,6 +159,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/open_api/sub_store", h.openAPISubStore)
|
||||
mux.HandleFunc("/api/v1/federation/share/list", h.federationShareList)
|
||||
mux.HandleFunc("/api/v1/federation/share/create", h.federationShareCreate)
|
||||
mux.HandleFunc("/api/v1/federation/share/update", h.federationShareUpdate)
|
||||
mux.HandleFunc("/api/v1/federation/share/delete", h.federationShareDelete)
|
||||
mux.HandleFunc("/api/v1/federation/share/reset-flow", h.federationShareResetFlow)
|
||||
mux.HandleFunc("/api/v1/federation/share/remote-usage/list", h.federationRemoteUsageList)
|
||||
@@ -320,6 +325,9 @@ func (h *Handler) nodeList(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
h.syncRemoteNodeStatuses(items)
|
||||
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"go-backend/internal/http/client"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store"
|
||||
"go-backend/internal/store/sqlite"
|
||||
)
|
||||
|
||||
@@ -413,7 +414,7 @@ func (h *Handler) nodeInstall(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
cmd := fmt.Sprintf("curl -L https://github.com/Sagit-chu/flux-panel/releases/latest/download/install.sh -o ./install.sh && chmod +x ./install.sh && ./install.sh -a %s -s %s", processServerAddress(panelAddr), secret)
|
||||
cmd := fmt.Sprintf("curl -L https://gcode.hostcentral.cc/https://github.com/Sagit-chu/flvx/releases/latest/download/install.sh -o ./install.sh && chmod +x ./install.sh && ./install.sh -a %s -s %s", processServerAddress(panelAddr), secret)
|
||||
response.WriteJSON(w, response.OK(cmd))
|
||||
}
|
||||
|
||||
@@ -559,9 +560,8 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
|
||||
var tunnelID int64
|
||||
err = tx.QueryRow(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) RETURNING id`,
|
||||
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx).Scan(&tunnelID)
|
||||
tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1122,11 +1122,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
var forwardID int64
|
||||
err = tx.QueryRow(`
|
||||
forwardID, err := tx.ExecReturningID(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?) RETURNING id
|
||||
`, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx).Scan(&forwardID)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||
`, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1590,9 +1589,8 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
speed := asInt(req["speed"], 100)
|
||||
var id int64
|
||||
err := h.repo.DB().QueryRow(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?) RETURNING id`,
|
||||
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1)).Scan(&id)
|
||||
id, err := h.repo.DB().ExecReturningID(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`,
|
||||
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1))
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1690,7 +1688,7 @@ func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
_, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID)
|
||||
for _, tid := range req.TunnelIDs {
|
||||
_, _ = tx.Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?) ON CONFLICT(tunnel_group_id, tunnel_id) DO NOTHING`, req.GroupID, tid, time.Now().UnixMilli())
|
||||
_, _ = tx.Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.GroupID, tid, time.Now().UnixMilli())
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -1717,7 +1715,7 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
_, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID)
|
||||
for _, uid := range req.UserIDs {
|
||||
_, _ = tx.Exec(`INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT(user_group_id, user_id) DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli())
|
||||
_, _ = tx.Exec(`INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli())
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -1736,7 +1734,7 @@ func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request)
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
_, err := h.repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?) ON CONFLICT(user_group_id, tunnel_group_id) DO NOTHING`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
|
||||
_, err := h.repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1838,7 +1836,7 @@ func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error {
|
||||
if created {
|
||||
createdByGroup = 1
|
||||
}
|
||||
_, _ = db.Exec(`INSERT INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?) ON CONFLICT(user_group_id, tunnel_group_id, user_tunnel_id) DO NOTHING`,
|
||||
_, _ = db.Exec(`INSERT INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?) ON CONFLICT DO NOTHING`,
|
||||
userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli())
|
||||
}
|
||||
}
|
||||
@@ -1869,7 +1867,7 @@ func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, error) {
|
||||
func ensureUserTunnelGrant(db *store.DB, userID, tunnelID int64) (int64, bool, error) {
|
||||
var id int64
|
||||
err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&id)
|
||||
if err == nil {
|
||||
@@ -1885,15 +1883,15 @@ func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, err
|
||||
if err := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset); err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
err = db.QueryRow(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1) RETURNING id`,
|
||||
userID, tunnelID, num, flow, flowReset, expTime).Scan(&id)
|
||||
id, err = db.ExecReturningID(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1)`,
|
||||
userID, tunnelID, num, flow, flowReset, expTime)
|
||||
if err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
return id, true, nil
|
||||
}
|
||||
|
||||
func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error) {
|
||||
func queryInt64List(db *store.DB, q string, args ...interface{}) ([]int64, error) {
|
||||
rows, err := db.Query(q, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1910,7 +1908,7 @@ func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error)
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func queryPairs(db *sql.DB, q string, args ...interface{}) ([][2]int64, error) {
|
||||
func queryPairs(db *store.DB, q string, args ...interface{}) ([][2]int64, error) {
|
||||
rows, err := db.Query(q, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1946,7 +1944,7 @@ type tunnelCreateState struct {
|
||||
NodeIDList []int64
|
||||
}
|
||||
|
||||
func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
|
||||
func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
|
||||
state := &tunnelCreateState{
|
||||
Type: tunnelType,
|
||||
InNodes: make([]tunnelRuntimeNode, 0),
|
||||
@@ -2405,7 +2403,7 @@ func (h *Handler) cleanupFederationRuntime(tunnelID int64) {
|
||||
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID)
|
||||
}
|
||||
|
||||
func replaceFederationTunnelBindingsTx(tx *sql.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error {
|
||||
func replaceFederationTunnelBindingsTx(tx *store.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
@@ -2715,7 +2713,7 @@ func pickNodeAddressV6(node *nodeRecord) string {
|
||||
return strings.TrimSpace(node.ServerIP)
|
||||
}
|
||||
|
||||
func isRemoteNodeTx(tx *sql.Tx, nodeID int64) (bool, error) {
|
||||
func isRemoteNodeTx(tx *store.Tx, nodeID int64) (bool, error) {
|
||||
if tx == nil {
|
||||
return false, errors.New("database unavailable")
|
||||
}
|
||||
@@ -2732,7 +2730,7 @@ func isRemoteNodeTx(tx *sql.Tx, nodeID int64) (bool, error) {
|
||||
return isRemote == 1, nil
|
||||
}
|
||||
|
||||
func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
|
||||
func pickNodePortTx(tx *store.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
|
||||
if tx == nil {
|
||||
return 0, errors.New("database unavailable")
|
||||
}
|
||||
@@ -2843,7 +2841,7 @@ func parsePortRangeSpec(input string) []int {
|
||||
return out
|
||||
}
|
||||
|
||||
func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error {
|
||||
func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interface{}) error {
|
||||
allocated := map[int64]int{}
|
||||
inNodes := asMapSlice(req["inNodeId"])
|
||||
for _, n := range inNodes {
|
||||
@@ -2851,7 +2849,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
|
||||
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, ?)`,
|
||||
_, 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"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2870,7 +2868,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
|
||||
return pickErr
|
||||
}
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, ?, ?, 0, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '3', ?, ?, ?, 0, ?)`,
|
||||
tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2891,7 +2889,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
|
||||
return pickErr
|
||||
}
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, ?, ?, ?, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '2', ?, ?, ?, ?, ?)`,
|
||||
tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2977,7 +2975,7 @@ func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) {
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
|
||||
rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 ORDER BY inx ASC, id ASC`, tunnelID)
|
||||
rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = '1' ORDER BY inx ASC, id ASC`, tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -3464,7 +3462,7 @@ func randomToken(n int) string {
|
||||
return hex.EncodeToString(buf)
|
||||
}
|
||||
|
||||
func nextIndex(db *sql.DB, table string) int {
|
||||
func nextIndex(db *store.DB, table string) int {
|
||||
if db == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
@@ -0,0 +1,291 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
const (
|
||||
githubRepo = "Sagit-chu/flvx"
|
||||
githubProxy = "https://gcode.hostcentral.cc"
|
||||
githubAPIBase = "https://api.github.com"
|
||||
githubHTMLBase = "https://github.com"
|
||||
upgradeTimeout = 5 * time.Minute
|
||||
batchWorkers = 5
|
||||
)
|
||||
|
||||
func (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("节点ID无效"))
|
||||
return
|
||||
}
|
||||
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = resolveLatestRelease()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新版本失败: %v", err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
downloadURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
checksumURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}.sha256",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
|
||||
result, err := h.wsServer.SendCommand(req.ID, "UpgradeAgent", map[string]interface{}{
|
||||
"downloadUrl": downloadURL,
|
||||
"checksumUrl": checksumURL,
|
||||
}, upgradeTimeout)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("升级失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"version": version,
|
||||
"message": result.Message,
|
||||
}))
|
||||
}
|
||||
|
||||
func resolveLatestRelease() (string, error) {
|
||||
client := &http.Client{
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
resp, err := client.Get(githubProxy + "/" + githubHTMLBase + "/" + githubRepo + "/releases/latest")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求GitHub失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusFound && resp.StatusCode != http.StatusMovedPermanently {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
location := resp.Header.Get("Location")
|
||||
if location == "" {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
parts := strings.Split(location, "/")
|
||||
tag := parts[len(parts)-1]
|
||||
if tag == "" || tag == "latest" {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
return tag, nil
|
||||
}
|
||||
|
||||
func resolveLatestReleaseAPI() (string, error) {
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
resp, err := client.Get(githubAPIBase + "/repos/" + githubRepo + "/releases/latest")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求GitHub API失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
return "", fmt.Errorf("GitHub API返回 %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&release); err != nil {
|
||||
return "", fmt.Errorf("解析GitHub API响应失败: %v", err)
|
||||
}
|
||||
if strings.TrimSpace(release.TagName) == "" {
|
||||
return "", fmt.Errorf("无法从GitHub获取最新版本号")
|
||||
}
|
||||
|
||||
return release.TagName, nil
|
||||
}
|
||||
|
||||
func (h *Handler) nodeBatchUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
IDs []int64 `json:"ids"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if len(req.IDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("ids不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = resolveLatestRelease()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新版本失败: %v", err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
downloadURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
checksumURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}.sha256",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
|
||||
type upgradeResult struct {
|
||||
ID int64 `json:"id"`
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
results := make([]upgradeResult, len(req.IDs))
|
||||
sem := make(chan struct{}, batchWorkers)
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for i, id := range req.IDs {
|
||||
wg.Add(1)
|
||||
go func(index int, nodeID int64) {
|
||||
defer wg.Done()
|
||||
sem <- struct{}{}
|
||||
defer func() { <-sem }()
|
||||
|
||||
result, err := h.wsServer.SendCommand(nodeID, "UpgradeAgent", map[string]interface{}{
|
||||
"downloadUrl": downloadURL,
|
||||
"checksumUrl": checksumURL,
|
||||
}, upgradeTimeout)
|
||||
if err != nil {
|
||||
results[index] = upgradeResult{ID: nodeID, Success: false, Message: err.Error()}
|
||||
return
|
||||
}
|
||||
results[index] = upgradeResult{ID: nodeID, Success: true, Message: result.Message}
|
||||
}(i, id)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"version": version,
|
||||
"results": results,
|
||||
}))
|
||||
}
|
||||
|
||||
func (h *Handler) listReleases(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 15 * time.Second}
|
||||
resp, err := client.Get(githubAPIBase + "/repos/" + githubRepo + "/releases?per_page=20")
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: GitHub API返回 %d: %s", resp.StatusCode, string(body))))
|
||||
return
|
||||
}
|
||||
|
||||
var releases []struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("解析版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
type releaseItem struct {
|
||||
Version string `json:"version"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"publishedAt"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
}
|
||||
|
||||
items := make([]releaseItem, 0, len(releases))
|
||||
for _, r := range releases {
|
||||
if r.Draft {
|
||||
continue
|
||||
}
|
||||
items = append(items, releaseItem{
|
||||
Version: r.TagName,
|
||||
Name: r.Name,
|
||||
PublishedAt: r.PublishedAt,
|
||||
Prerelease: r.Prerelease,
|
||||
})
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
func (h *Handler) nodeRollback(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("节点ID无效"))
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.wsServer.SendCommand(req.ID, "RollbackAgent", map[string]interface{}{}, 30*time.Second)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("回退失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"message": result.Message,
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,466 @@
|
||||
// Package store provides a thin dialect-aware wrapper around database/sql,
|
||||
// enabling transparent use of both SQLite and PostgreSQL.
|
||||
package store
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Dialect identifies the underlying database engine.
|
||||
type Dialect int
|
||||
|
||||
const (
|
||||
DialectSQLite Dialect = iota
|
||||
DialectPostgres
|
||||
)
|
||||
|
||||
// String returns a human-readable dialect name.
|
||||
func (d Dialect) String() string {
|
||||
switch d {
|
||||
case DialectSQLite:
|
||||
return "sqlite"
|
||||
case DialectPostgres:
|
||||
return "postgres"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
// DB wraps *sql.DB with dialect awareness.
|
||||
type DB struct {
|
||||
raw *sql.DB
|
||||
dialect Dialect
|
||||
}
|
||||
|
||||
// Wrap creates a new dialect-aware DB from an existing *sql.DB.
|
||||
func Wrap(raw *sql.DB, dialect Dialect) *DB {
|
||||
return &DB{raw: raw, dialect: dialect}
|
||||
}
|
||||
|
||||
// Dialect returns the database dialect.
|
||||
func (db *DB) Dialect() Dialect {
|
||||
if db == nil {
|
||||
return DialectSQLite
|
||||
}
|
||||
return db.dialect
|
||||
}
|
||||
|
||||
// RawDB returns the underlying *sql.DB.
|
||||
func (db *DB) RawDB() *sql.DB {
|
||||
if db == nil {
|
||||
return nil
|
||||
}
|
||||
return db.raw
|
||||
}
|
||||
|
||||
// Close closes the underlying connection.
|
||||
func (db *DB) Close() error {
|
||||
if db == nil || db.raw == nil {
|
||||
return nil
|
||||
}
|
||||
return db.raw.Close()
|
||||
}
|
||||
|
||||
// Ping verifies the connection is alive.
|
||||
func (db *DB) Ping() error {
|
||||
return db.raw.Ping()
|
||||
}
|
||||
|
||||
// Exec executes a query with transparent placeholder and syntax rewriting.
|
||||
func (db *DB) Exec(query string, args ...any) (sql.Result, error) {
|
||||
return db.raw.Exec(db.rewrite(query), args...)
|
||||
}
|
||||
|
||||
// Query executes a query that returns rows, with transparent rewriting.
|
||||
func (db *DB) Query(query string, args ...any) (*sql.Rows, error) {
|
||||
return db.raw.Query(db.rewrite(query), args...)
|
||||
}
|
||||
|
||||
// QueryRow executes a query that returns at most one row, with transparent rewriting.
|
||||
func (db *DB) QueryRow(query string, args ...any) *sql.Row {
|
||||
return db.raw.QueryRow(db.rewrite(query), args...)
|
||||
}
|
||||
|
||||
// Begin starts a transaction, returning a dialect-aware Tx.
|
||||
func (db *DB) Begin() (*Tx, error) {
|
||||
tx, err := db.raw.Begin()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Tx{raw: tx, dialect: db.dialect}, nil
|
||||
}
|
||||
|
||||
// ExecReturningID executes an INSERT and returns the auto-generated id.
|
||||
// - SQLite: uses LastInsertId()
|
||||
// - PostgreSQL: appends RETURNING id and uses QueryRow().Scan()
|
||||
func (db *DB) ExecReturningID(query string, args ...any) (int64, error) {
|
||||
q := db.rewrite(query)
|
||||
if db.dialect == DialectPostgres {
|
||||
q = ensureReturningID(q)
|
||||
var id int64
|
||||
if err := db.raw.QueryRow(q, args...).Scan(&id); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
res, err := db.raw.Exec(q, args...)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
// Tx wraps *sql.Tx with dialect awareness.
|
||||
type Tx struct {
|
||||
raw *sql.Tx
|
||||
dialect Dialect
|
||||
}
|
||||
|
||||
// Exec executes a query inside the transaction with transparent rewriting.
|
||||
func (tx *Tx) Exec(query string, args ...any) (sql.Result, error) {
|
||||
return tx.raw.Exec(rewriteQuery(tx.dialect, query), args...)
|
||||
}
|
||||
|
||||
// Query executes a query that returns rows inside the transaction.
|
||||
func (tx *Tx) Query(query string, args ...any) (*sql.Rows, error) {
|
||||
return tx.raw.Query(rewriteQuery(tx.dialect, query), args...)
|
||||
}
|
||||
|
||||
// QueryRow executes a query that returns at most one row inside the transaction.
|
||||
func (tx *Tx) QueryRow(query string, args ...any) *sql.Row {
|
||||
return tx.raw.QueryRow(rewriteQuery(tx.dialect, query), args...)
|
||||
}
|
||||
|
||||
// Commit commits the transaction.
|
||||
func (tx *Tx) Commit() error { return tx.raw.Commit() }
|
||||
|
||||
// Rollback aborts the transaction.
|
||||
func (tx *Tx) Rollback() error { return tx.raw.Rollback() }
|
||||
|
||||
// ExecReturningID executes an INSERT inside the transaction and returns the id.
|
||||
func (tx *Tx) ExecReturningID(query string, args ...any) (int64, error) {
|
||||
q := rewriteQuery(tx.dialect, query)
|
||||
if tx.dialect == DialectPostgres {
|
||||
q = ensureReturningID(q)
|
||||
var id int64
|
||||
if err := tx.raw.QueryRow(q, args...).Scan(&id); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
res, err := tx.raw.Exec(q, args...)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
func (db *DB) rewrite(query string) string {
|
||||
return rewriteQuery(db.dialect, query)
|
||||
}
|
||||
|
||||
func rewriteQuery(dialect Dialect, query string) string {
|
||||
if dialect != DialectPostgres {
|
||||
return query
|
||||
}
|
||||
query = rewriteUserIdentifier(query)
|
||||
query = rewriteInsertOrIgnore(query)
|
||||
query = rewritePlaceholders(query)
|
||||
return query
|
||||
}
|
||||
|
||||
func rewriteUserIdentifier(query string) string {
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(query) + 16)
|
||||
i := 0
|
||||
for i < len(query) {
|
||||
if end, ok := skipSQLProtectedSegment(query, i); ok {
|
||||
buf.WriteString(query[i:end])
|
||||
i = end
|
||||
continue
|
||||
}
|
||||
|
||||
ch := query[i]
|
||||
if isIdentifierChar(ch) {
|
||||
j := i + 1
|
||||
for j < len(query) && isIdentifierChar(query[j]) {
|
||||
j++
|
||||
}
|
||||
tok := query[i:j]
|
||||
if strings.EqualFold(tok, "user") {
|
||||
buf.WriteString(`"user"`)
|
||||
} else {
|
||||
buf.WriteString(tok)
|
||||
}
|
||||
i = j
|
||||
continue
|
||||
}
|
||||
|
||||
buf.WriteByte(ch)
|
||||
i++
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func isIdentifierChar(ch byte) bool {
|
||||
if ch >= 'a' && ch <= 'z' {
|
||||
return true
|
||||
}
|
||||
if ch >= 'A' && ch <= 'Z' {
|
||||
return true
|
||||
}
|
||||
if ch >= '0' && ch <= '9' {
|
||||
return true
|
||||
}
|
||||
return ch == '_'
|
||||
}
|
||||
|
||||
func rewriteInsertOrIgnore(query string) string {
|
||||
start, end, ok := findKeywordSequenceOutside(query, []string{"INSERT", "OR", "IGNORE", "INTO"}, 0)
|
||||
if !ok {
|
||||
return query
|
||||
}
|
||||
|
||||
rewritten := query[:start] + "INSERT INTO" + query[end:]
|
||||
rewritten = strings.TrimRight(rewritten, "; \t\n")
|
||||
|
||||
insertIntoEnd := start + len("INSERT INTO")
|
||||
if _, _, hasOnConflict := findKeywordSequenceOutside(rewritten, []string{"ON", "CONFLICT"}, insertIntoEnd); hasOnConflict {
|
||||
return rewritten
|
||||
}
|
||||
|
||||
if retStart, _, hasReturning := findKeywordSequenceOutside(rewritten, []string{"RETURNING"}, insertIntoEnd); hasReturning {
|
||||
prefix := strings.TrimRight(rewritten[:retStart], " \t\n")
|
||||
suffix := strings.TrimLeft(rewritten[retStart:], " \t\n")
|
||||
return prefix + " ON CONFLICT DO NOTHING " + suffix
|
||||
}
|
||||
|
||||
return rewritten + " ON CONFLICT DO NOTHING"
|
||||
}
|
||||
|
||||
func rewritePlaceholders(query string) string {
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(query) + 16)
|
||||
n := 1
|
||||
for i := 0; i < len(query); i++ {
|
||||
if end, ok := skipSQLProtectedSegment(query, i); ok {
|
||||
buf.WriteString(query[i:end])
|
||||
i = end - 1
|
||||
continue
|
||||
}
|
||||
|
||||
ch := query[i]
|
||||
if ch == '?' {
|
||||
buf.WriteByte('$')
|
||||
buf.WriteString(strconv.Itoa(n))
|
||||
n++
|
||||
continue
|
||||
}
|
||||
buf.WriteByte(ch)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func ensureReturningID(query string) string {
|
||||
trimmed := strings.TrimRight(query, "; \t\n")
|
||||
if _, _, ok := findKeywordSequenceOutside(trimmed, []string{"RETURNING"}, 0); ok {
|
||||
return trimmed
|
||||
}
|
||||
return trimmed + " RETURNING id"
|
||||
}
|
||||
|
||||
func findKeywordSequenceOutside(query string, keywords []string, from int) (int, int, bool) {
|
||||
if len(keywords) == 0 {
|
||||
return 0, 0, false
|
||||
}
|
||||
if from < 0 {
|
||||
from = 0
|
||||
}
|
||||
if from >= len(query) {
|
||||
return 0, 0, false
|
||||
}
|
||||
|
||||
matched := 0
|
||||
seqStart := -1
|
||||
|
||||
for i := from; i < len(query); {
|
||||
if end, ok := skipSQLProtectedSegment(query, i); ok {
|
||||
i = end
|
||||
continue
|
||||
}
|
||||
|
||||
ch := query[i]
|
||||
if isIdentifierChar(ch) {
|
||||
j := i + 1
|
||||
for j < len(query) && isIdentifierChar(query[j]) {
|
||||
j++
|
||||
}
|
||||
tok := query[i:j]
|
||||
|
||||
if strings.EqualFold(tok, keywords[matched]) {
|
||||
if matched == 0 {
|
||||
seqStart = i
|
||||
}
|
||||
matched++
|
||||
if matched == len(keywords) {
|
||||
return seqStart, j, true
|
||||
}
|
||||
} else if strings.EqualFold(tok, keywords[0]) {
|
||||
seqStart = i
|
||||
matched = 1
|
||||
} else {
|
||||
matched = 0
|
||||
seqStart = -1
|
||||
}
|
||||
|
||||
i = j
|
||||
continue
|
||||
}
|
||||
|
||||
if !isSQLSpace(ch) {
|
||||
matched = 0
|
||||
seqStart = -1
|
||||
}
|
||||
i++
|
||||
}
|
||||
|
||||
return 0, 0, false
|
||||
}
|
||||
|
||||
func skipSQLProtectedSegment(query string, i int) (int, bool) {
|
||||
if i < 0 || i >= len(query) {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
switch query[i] {
|
||||
case '\'':
|
||||
return skipSingleQuotedLiteral(query, i), true
|
||||
case '"':
|
||||
return skipDoubleQuotedIdentifier(query, i), true
|
||||
case '-':
|
||||
if i+1 < len(query) && query[i+1] == '-' {
|
||||
return skipLineComment(query, i), true
|
||||
}
|
||||
case '/':
|
||||
if i+1 < len(query) && query[i+1] == '*' {
|
||||
return skipBlockComment(query, i), true
|
||||
}
|
||||
case '$':
|
||||
if end, ok := skipDollarQuotedLiteral(query, i); ok {
|
||||
return end, true
|
||||
}
|
||||
}
|
||||
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func skipSingleQuotedLiteral(query string, i int) int {
|
||||
for j := i + 1; j < len(query); j++ {
|
||||
if query[j] != '\'' {
|
||||
continue
|
||||
}
|
||||
if j+1 < len(query) && query[j+1] == '\'' {
|
||||
j++
|
||||
continue
|
||||
}
|
||||
return j + 1
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipDoubleQuotedIdentifier(query string, i int) int {
|
||||
for j := i + 1; j < len(query); j++ {
|
||||
if query[j] != '"' {
|
||||
continue
|
||||
}
|
||||
if j+1 < len(query) && query[j+1] == '"' {
|
||||
j++
|
||||
continue
|
||||
}
|
||||
return j + 1
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipLineComment(query string, i int) int {
|
||||
for j := i + 2; j < len(query); j++ {
|
||||
if query[j] == '\n' {
|
||||
return j
|
||||
}
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipBlockComment(query string, i int) int {
|
||||
depth := 1
|
||||
for j := i + 2; j < len(query)-1; j++ {
|
||||
if query[j] == '/' && query[j+1] == '*' {
|
||||
depth++
|
||||
j++
|
||||
continue
|
||||
}
|
||||
if query[j] == '*' && query[j+1] == '/' {
|
||||
depth--
|
||||
j++
|
||||
if depth == 0 {
|
||||
return j + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
return len(query)
|
||||
}
|
||||
|
||||
func skipDollarQuotedLiteral(query string, i int) (int, bool) {
|
||||
if i < 0 || i >= len(query) || query[i] != '$' {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
if i+1 >= len(query) {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
var endTag int
|
||||
if query[i+1] == '$' {
|
||||
endTag = i + 1
|
||||
} else {
|
||||
if !isDollarTagStart(query[i+1]) {
|
||||
return 0, false
|
||||
}
|
||||
j := i + 2
|
||||
for j < len(query) && isDollarTagChar(query[j]) {
|
||||
j++
|
||||
}
|
||||
if j >= len(query) || query[j] != '$' {
|
||||
return 0, false
|
||||
}
|
||||
endTag = j
|
||||
}
|
||||
|
||||
tag := query[i : endTag+1]
|
||||
if closeIdx := strings.Index(query[endTag+1:], tag); closeIdx >= 0 {
|
||||
return endTag + 1 + closeIdx + len(tag), true
|
||||
}
|
||||
return len(query), true
|
||||
}
|
||||
|
||||
func isDollarTagStart(ch byte) bool {
|
||||
return ch == '_' || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z')
|
||||
}
|
||||
|
||||
func isDollarTagChar(ch byte) bool {
|
||||
if isDollarTagStart(ch) {
|
||||
return true
|
||||
}
|
||||
return ch >= '0' && ch <= '9'
|
||||
}
|
||||
|
||||
func isSQLSpace(ch byte) bool {
|
||||
switch ch {
|
||||
case ' ', '\t', '\n', '\r', '\f':
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package store
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestRewritePlaceholdersSkipsProtectedSegments(t *testing.T) {
|
||||
q := `SELECT ?, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ?
|
||||
FROM t /* block ? */ WHERE id = ?`
|
||||
got := rewritePlaceholders(q)
|
||||
want := `SELECT $1, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ?
|
||||
FROM t /* block ? */ WHERE id = $2`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreBasic(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreBeforeReturning(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO x(a) VALUES(?) RETURNING id`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `INSERT INTO x(a) VALUES(?) ON CONFLICT DO NOTHING RETURNING id`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreNotDuplicatingOnConflict(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `INSERT INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureReturningID(t *testing.T) {
|
||||
if got := ensureReturningID(`INSERT INTO x(a) VALUES($1)`); got != `INSERT INTO x(a) VALUES($1) RETURNING id` {
|
||||
t.Fatalf("missing RETURNING append: %s", got)
|
||||
}
|
||||
if got := ensureReturningID(`INSERT INTO x(a) VALUES($1) RETURNING other_id`); got != `INSERT INTO x(a) VALUES($1) RETURNING other_id` {
|
||||
t.Fatalf("RETURNING should not be duplicated: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteUserIdentifierSafety(t *testing.T) {
|
||||
q := `SELECT user, user_id, 'user', "user", note FROM user -- user
|
||||
WHERE owner='user'`
|
||||
got := rewriteUserIdentifier(q)
|
||||
want := `SELECT "user", user_id, 'user', "user", note FROM "user" -- user
|
||||
WHERE owner='user'`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteQueryPostgresPipeline(t *testing.T) {
|
||||
q := `INSERT OR IGNORE INTO user(name, note) VALUES(?, '?')`
|
||||
got := rewriteQuery(DialectPostgres, q)
|
||||
want := `INSERT INTO "user"(name, note) VALUES($1, '?') ON CONFLICT DO NOTHING`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreSkipsStringLiteral(t *testing.T) {
|
||||
q := `SELECT 'INSERT OR IGNORE INTO t(a) VALUES(?)' AS q`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
if got != q {
|
||||
t.Fatalf("string literal should stay unchanged\nwant: %s\ngot: %s", q, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteInsertOrIgnoreSkipsCommentedKeyword(t *testing.T) {
|
||||
q := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?)
|
||||
INSERT OR IGNORE INTO real_t(a) VALUES(?)`
|
||||
got := rewriteInsertOrIgnore(q)
|
||||
want := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?)
|
||||
INSERT INTO real_t(a) VALUES(?) ON CONFLICT DO NOTHING`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewritePlaceholdersSkipsNestedBlockComment(t *testing.T) {
|
||||
q := `SELECT ? /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = ?`
|
||||
got := rewritePlaceholders(q)
|
||||
want := `SELECT $1 /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = $2`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewritePlaceholdersSkipsUnterminatedBlockComment(t *testing.T) {
|
||||
q := `SELECT ? /* unterminated ? comment`
|
||||
got := rewritePlaceholders(q)
|
||||
want := `SELECT $1 /* unterminated ? comment`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewriteUserIdentifierSkipsDollarQuotedAndComment(t *testing.T) {
|
||||
q := `SELECT user, $$user ?$$ AS body, col FROM user /* user */ -- user`
|
||||
got := rewriteUserIdentifier(q)
|
||||
want := `SELECT "user", $$user ?$$ AS body, col FROM "user" /* user */ -- user`
|
||||
if got != want {
|
||||
t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
package postgres
|
||||
|
||||
import _ "embed"
|
||||
|
||||
//go:embed sql/schema.sql
|
||||
var EmbeddedSchema string
|
||||
|
||||
//go:embed sql/data.sql
|
||||
var EmbeddedSeedData string
|
||||
@@ -0,0 +1,18 @@
|
||||
INSERT INTO "user" (id, "user", pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES (1, 'admin_user', '3c85cdebade1c51cf64ca9f3c09d182d', 0, 2727251700000, 99999, 0, 0, 1, 99999, 1748914865000, 1754011744252, 1)
|
||||
ON CONFLICT DO NOTHING;
|
||||
|
||||
INSERT INTO vite_config (id, name, value, time)
|
||||
VALUES (1, 'app_name', 'flux', 1755147963000)
|
||||
ON CONFLICT DO NOTHING;
|
||||
|
||||
DO $$
|
||||
BEGIN
|
||||
IF to_regclass('public.user_id_seq') IS NOT NULL THEN
|
||||
PERFORM setval('user_id_seq', (SELECT COALESCE(MAX(id), 0) FROM "user"));
|
||||
END IF;
|
||||
IF to_regclass('public.vite_config_id_seq') IS NOT NULL THEN
|
||||
PERFORM setval('vite_config_id_seq', (SELECT COALESCE(MAX(id), 0) FROM vite_config));
|
||||
END IF;
|
||||
END
|
||||
$$;
|
||||
@@ -0,0 +1,241 @@
|
||||
CREATE TABLE IF NOT EXISTS forward (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
user_name VARCHAR(100) NOT NULL,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
tunnel_id INTEGER NOT NULL,
|
||||
remote_addr TEXT NOT NULL,
|
||||
strategy VARCHAR(100) NOT NULL DEFAULT 'fifo',
|
||||
in_flow BIGINT NOT NULL DEFAULT 0,
|
||||
out_flow BIGINT NOT NULL DEFAULT 0,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT NOT NULL,
|
||||
status INTEGER NOT NULL,
|
||||
inx INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS forward_port (
|
||||
id SERIAL PRIMARY KEY,
|
||||
forward_id INTEGER NOT NULL,
|
||||
node_id INTEGER NOT NULL,
|
||||
port INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS node (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
secret VARCHAR(100) NOT NULL,
|
||||
server_ip VARCHAR(100) NOT NULL,
|
||||
server_ip_v4 VARCHAR(100),
|
||||
server_ip_v6 VARCHAR(100),
|
||||
port TEXT NOT NULL,
|
||||
interface_name VARCHAR(200),
|
||||
version VARCHAR(100),
|
||||
http INTEGER NOT NULL DEFAULT 0,
|
||||
tls INTEGER NOT NULL DEFAULT 0,
|
||||
socks INTEGER NOT NULL DEFAULT 0,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT,
|
||||
status INTEGER NOT NULL,
|
||||
tcp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
|
||||
udp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
|
||||
inx INTEGER NOT NULL DEFAULT 0,
|
||||
is_remote INTEGER DEFAULT 0,
|
||||
remote_url TEXT,
|
||||
remote_token TEXT,
|
||||
remote_config TEXT
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS speed_limit (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
speed INTEGER NOT NULL,
|
||||
tunnel_id INTEGER NOT NULL,
|
||||
tunnel_name VARCHAR(100) NOT NULL,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT,
|
||||
status INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS statistics_flow (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
flow BIGINT NOT NULL,
|
||||
total_flow BIGINT NOT NULL,
|
||||
time VARCHAR(100) NOT NULL,
|
||||
created_time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tunnel (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
traffic_ratio DOUBLE PRECISION NOT NULL DEFAULT 1.0,
|
||||
type INTEGER NOT NULL,
|
||||
protocol VARCHAR(10) NOT NULL DEFAULT 'tls',
|
||||
flow BIGINT NOT NULL,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT NOT NULL,
|
||||
status INTEGER NOT NULL,
|
||||
in_ip TEXT,
|
||||
inx INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS chain_tunnel (
|
||||
id SERIAL PRIMARY KEY,
|
||||
tunnel_id INTEGER NOT NULL,
|
||||
chain_type VARCHAR(10) NOT NULL,
|
||||
node_id INTEGER NOT NULL,
|
||||
port INTEGER,
|
||||
strategy VARCHAR(10),
|
||||
inx INTEGER,
|
||||
protocol VARCHAR(10)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS "user" (
|
||||
id SERIAL PRIMARY KEY,
|
||||
"user" VARCHAR(100) NOT NULL,
|
||||
pwd VARCHAR(100) NOT NULL,
|
||||
role_id INTEGER NOT NULL,
|
||||
exp_time BIGINT NOT NULL,
|
||||
flow BIGINT NOT NULL,
|
||||
in_flow BIGINT NOT NULL DEFAULT 0,
|
||||
out_flow BIGINT NOT NULL DEFAULT 0,
|
||||
flow_reset_time BIGINT NOT NULL,
|
||||
num INTEGER NOT NULL,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT,
|
||||
status INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_tunnel (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
tunnel_id INTEGER NOT NULL,
|
||||
speed_id INTEGER,
|
||||
num INTEGER NOT NULL,
|
||||
flow BIGINT NOT NULL,
|
||||
in_flow BIGINT NOT NULL DEFAULT 0,
|
||||
out_flow BIGINT NOT NULL DEFAULT 0,
|
||||
flow_reset_time BIGINT NOT NULL,
|
||||
exp_time BIGINT NOT NULL,
|
||||
status INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tunnel_group (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT NOT NULL,
|
||||
status INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_group (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT NOT NULL,
|
||||
status INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tunnel_group_tunnel (
|
||||
id SERIAL PRIMARY KEY,
|
||||
tunnel_group_id INTEGER NOT NULL,
|
||||
tunnel_id INTEGER NOT NULL,
|
||||
created_time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_group_user (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_group_id INTEGER NOT NULL,
|
||||
user_id INTEGER NOT NULL,
|
||||
created_time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS group_permission (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_group_id INTEGER NOT NULL,
|
||||
tunnel_group_id INTEGER NOT NULL,
|
||||
created_time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS group_permission_grant (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_group_id INTEGER NOT NULL,
|
||||
tunnel_group_id INTEGER NOT NULL,
|
||||
user_tunnel_id INTEGER NOT NULL,
|
||||
created_by_group INTEGER NOT NULL DEFAULT 0,
|
||||
created_time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_name ON tunnel_group(name);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_name ON user_group(name);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_tunnel_unique ON tunnel_group_tunnel(tunnel_group_id, tunnel_id);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_user_unique ON user_group_user(user_group_id, user_id);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_unique ON group_permission(user_group_id, tunnel_group_id);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_grant_unique ON group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_tunnel_unique ON user_tunnel(user_id, tunnel_id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS vite_config (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name VARCHAR(200) NOT NULL UNIQUE,
|
||||
value VARCHAR(200) NOT NULL,
|
||||
time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS peer_share (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
node_id INTEGER NOT NULL,
|
||||
token TEXT NOT NULL UNIQUE,
|
||||
max_bandwidth INTEGER DEFAULT 0,
|
||||
expiry_time BIGINT DEFAULT 0,
|
||||
port_range_start INTEGER DEFAULT 0,
|
||||
port_range_end INTEGER DEFAULT 0,
|
||||
current_flow BIGINT DEFAULT 0,
|
||||
is_active INTEGER DEFAULT 1,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT NOT NULL,
|
||||
allowed_domains TEXT DEFAULT '',
|
||||
allowed_ips TEXT DEFAULT ''
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS peer_share_runtime (
|
||||
id SERIAL PRIMARY KEY,
|
||||
share_id INTEGER NOT NULL,
|
||||
node_id INTEGER NOT NULL,
|
||||
reservation_id TEXT NOT NULL UNIQUE,
|
||||
resource_key TEXT NOT NULL UNIQUE,
|
||||
binding_id TEXT NOT NULL DEFAULT '',
|
||||
role TEXT NOT NULL DEFAULT '',
|
||||
chain_name TEXT NOT NULL DEFAULT '',
|
||||
service_name TEXT NOT NULL DEFAULT '',
|
||||
protocol TEXT NOT NULL DEFAULT 'tls',
|
||||
strategy TEXT NOT NULL DEFAULT 'round',
|
||||
port INTEGER NOT NULL DEFAULT 0,
|
||||
target TEXT NOT NULL DEFAULT '',
|
||||
applied INTEGER NOT NULL DEFAULT 0,
|
||||
status INTEGER NOT NULL DEFAULT 1,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_peer_share_runtime_share_node_status ON peer_share_runtime(share_id, node_id, status);
|
||||
CREATE INDEX IF NOT EXISTS idx_peer_share_runtime_binding_id ON peer_share_runtime(binding_id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS federation_tunnel_binding (
|
||||
id SERIAL PRIMARY KEY,
|
||||
tunnel_id INTEGER NOT NULL,
|
||||
node_id INTEGER NOT NULL,
|
||||
chain_type INTEGER NOT NULL,
|
||||
hop_inx INTEGER NOT NULL DEFAULT 0,
|
||||
remote_url TEXT NOT NULL,
|
||||
resource_key TEXT NOT NULL UNIQUE,
|
||||
remote_binding_id TEXT NOT NULL,
|
||||
allocated_port INTEGER NOT NULL,
|
||||
status INTEGER NOT NULL DEFAULT 1,
|
||||
created_time BIGINT NOT NULL,
|
||||
updated_time BIGINT NOT NULL
|
||||
);
|
||||
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_federation_tunnel_binding_unique ON federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx);
|
||||
CREATE INDEX IF NOT EXISTS idx_federation_tunnel_binding_tunnel ON federation_tunnel_binding(tunnel_id, status);
|
||||
@@ -12,6 +12,9 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
"go-backend/internal/store"
|
||||
pgstore "go-backend/internal/store/postgres"
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
@@ -22,10 +25,10 @@ var embeddedSchema string
|
||||
var embeddedSeedData string
|
||||
|
||||
type Repository struct {
|
||||
db *sql.DB
|
||||
db *store.DB
|
||||
}
|
||||
|
||||
func (r *Repository) DB() *sql.DB {
|
||||
func (r *Repository) DB() *store.DB {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -165,17 +168,53 @@ func Open(path string) (*Repository, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
db, err := sql.Open("sqlite", path)
|
||||
// Use _pragma DSN parameters so every connection from the pool gets
|
||||
// the same settings (busy_timeout and synchronous are per-connection).
|
||||
dsn := "file:" + path +
|
||||
"?_pragma=busy_timeout(5000)" +
|
||||
"&_pragma=journal_mode(WAL)" +
|
||||
"&_pragma=synchronous(NORMAL)"
|
||||
raw, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db := store.Wrap(raw, store.DialectSQLite)
|
||||
|
||||
if err := db.Ping(); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := bootstrapSchema(db); err != nil {
|
||||
if err := bootstrapSchema(db, embeddedSchema, embeddedSeedData); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := migrateSchema(db); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Repository{db: db}, nil
|
||||
}
|
||||
|
||||
func OpenPostgres(dsn string) (*Repository, error) {
|
||||
if strings.TrimSpace(dsn) == "" {
|
||||
return nil, fmt.Errorf("empty postgres dsn")
|
||||
}
|
||||
|
||||
raw, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db := store.Wrap(raw, store.DialectPostgres)
|
||||
|
||||
if err := db.Ping(); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := bootstrapSchema(db, pgstore.EmbeddedSchema, pgstore.EmbeddedSeedData); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
@@ -860,7 +899,7 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
chainRows, err := r.db.Query(`
|
||||
SELECT tunnel_id, CAST(chain_type AS INTEGER), node_id, protocol, strategy, COALESCE(inx, 0)
|
||||
FROM chain_tunnel
|
||||
ORDER BY tunnel_id ASC, chain_type ASC, inx ASC, id ASC
|
||||
ORDER BY tunnel_id ASC, CAST(chain_type AS INTEGER) ASC, inx ASC, id ASC
|
||||
`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1141,7 +1180,7 @@ func nullableForwardIngress(v string) interface{} {
|
||||
return v
|
||||
}
|
||||
|
||||
func resolveForwardIngress(db *sql.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) {
|
||||
func resolveForwardIngress(db *store.DB, forwardID int64, tunnelID int64) (string, sql.NullInt64, error) {
|
||||
var tunnelInIP sql.NullString
|
||||
if err := db.QueryRow(`SELECT in_ip FROM tunnel WHERE id = ? LIMIT 1`, tunnelID).Scan(&tunnelInIP); err != nil {
|
||||
if !errors.Is(err, sql.ErrNoRows) {
|
||||
@@ -1243,16 +1282,16 @@ func ensureParentDir(dbPath string) error {
|
||||
return osMkdirAll(dir)
|
||||
}
|
||||
|
||||
func bootstrapSchema(db *sql.DB) error {
|
||||
func bootstrapSchema(db *store.DB, schemaSQL, seedSQL string) error {
|
||||
if db == nil {
|
||||
return errors.New("nil db")
|
||||
}
|
||||
|
||||
if _, err := db.Exec(embeddedSchema); err != nil {
|
||||
if _, err := db.Exec(schemaSQL); err != nil {
|
||||
return fmt.Errorf("apply schema.sql: %w", err)
|
||||
}
|
||||
|
||||
if _, err := db.Exec(embeddedSeedData); err != nil {
|
||||
if _, err := db.Exec(seedSQL); err != nil {
|
||||
return fmt.Errorf("apply data.sql: %w", err)
|
||||
}
|
||||
return nil
|
||||
@@ -1260,7 +1299,7 @@ func bootstrapSchema(db *sql.DB) error {
|
||||
|
||||
const currentSchemaVersion = 1
|
||||
|
||||
func getSchemaVersion(db *sql.DB) int {
|
||||
func getSchemaVersion(db *store.DB) int {
|
||||
_, _ = db.Exec(`CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL DEFAULT 0)`)
|
||||
var v int
|
||||
if err := db.QueryRow(`SELECT version FROM schema_version LIMIT 1`).Scan(&v); err != nil {
|
||||
@@ -1270,11 +1309,11 @@ func getSchemaVersion(db *sql.DB) int {
|
||||
return v
|
||||
}
|
||||
|
||||
func setSchemaVersion(db *sql.DB, v int) {
|
||||
func setSchemaVersion(db *store.DB, v int) {
|
||||
_, _ = db.Exec(`UPDATE schema_version SET version = ?`, v)
|
||||
}
|
||||
|
||||
func migrateSchema(db *sql.DB) error {
|
||||
func migrateSchema(db *store.DB) error {
|
||||
if db == nil {
|
||||
return errors.New("nil db")
|
||||
}
|
||||
@@ -1290,9 +1329,10 @@ func migrateSchema(db *sql.DB) error {
|
||||
if err == nil || errors.Is(err, sql.ErrNoRows) {
|
||||
return
|
||||
}
|
||||
// Column likely missing (SQLite: "no such column", PG: "does not exist", etc.)
|
||||
if _, alterErr := db.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, col, typ)); alterErr != nil {
|
||||
log.Printf("failed to add column %s to %s: %v", col, table, alterErr)
|
||||
if isMissingColumnError(db.Dialect(), err) {
|
||||
if _, alterErr := db.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, col, typ)); alterErr != nil {
|
||||
log.Printf("failed to add column %s to %s: %v", col, table, alterErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1327,10 +1367,176 @@ func migrateSchema(db *sql.DB) error {
|
||||
}
|
||||
}
|
||||
|
||||
if db.Dialect() == store.DialectPostgres {
|
||||
if err := ensurePostgresIDDefaults(db); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
setSchemaVersion(db, currentSchemaVersion)
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensurePostgresIDDefaults(db *store.DB) error {
|
||||
rows, err := db.Query(`
|
||||
SELECT c.table_schema, c.table_name
|
||||
FROM information_schema.table_constraints tc
|
||||
JOIN information_schema.key_column_usage kcu
|
||||
ON tc.constraint_name = kcu.constraint_name
|
||||
AND tc.table_schema = kcu.table_schema
|
||||
JOIN information_schema.columns c
|
||||
ON c.table_schema = kcu.table_schema
|
||||
AND c.table_name = kcu.table_name
|
||||
AND c.column_name = kcu.column_name
|
||||
WHERE tc.constraint_type = 'PRIMARY KEY'
|
||||
AND kcu.column_name = 'id'
|
||||
AND c.data_type IN ('integer', 'bigint')
|
||||
AND c.is_identity = 'NO'
|
||||
AND c.table_schema = current_schema()
|
||||
ORDER BY c.table_name ASC
|
||||
`)
|
||||
if err != nil {
|
||||
return fmt.Errorf("discover postgres id columns: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
var schemaName string
|
||||
var tableName string
|
||||
if err := rows.Scan(&schemaName, &tableName); err != nil {
|
||||
return fmt.Errorf("scan postgres id table row: %w", err)
|
||||
}
|
||||
if err := ensurePostgresTableIDDefault(db, schemaName, tableName); err != nil {
|
||||
return fmt.Errorf("repair %s.%s id default: %w", schemaName, tableName, err)
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return fmt.Errorf("iterate postgres id tables: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensurePostgresTableIDDefault(db *store.DB, schemaName, tableName string) error {
|
||||
var defaultExpr sql.NullString
|
||||
if err := db.QueryRow(`
|
||||
SELECT column_default
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = ?
|
||||
AND table_name = ?
|
||||
AND column_name = 'id'
|
||||
LIMIT 1
|
||||
`, schemaName, tableName).Scan(&defaultExpr); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
hasNextvalDefault := defaultExpr.Valid && strings.Contains(strings.ToLower(defaultExpr.String), "nextval(")
|
||||
|
||||
var serialSeq sql.NullString
|
||||
if err := db.QueryRow(`
|
||||
SELECT pg_get_serial_sequence(quote_ident(?) || '.' || quote_ident(?), 'id')
|
||||
`, schemaName, tableName).Scan(&serialSeq); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
seqRef := strings.TrimSpace(serialSeq.String)
|
||||
if seqRef == "" && hasNextvalDefault {
|
||||
seqRef = extractNextvalRegclass(defaultExpr.String)
|
||||
}
|
||||
|
||||
if !hasNextvalDefault || seqRef == "" {
|
||||
seqName := tableName + "_id_seq"
|
||||
if _, err := db.Exec(fmt.Sprintf("CREATE SEQUENCE IF NOT EXISTS %s.%s", quoteSQLIdentifier(schemaName), quoteSQLIdentifier(seqName))); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
seqRef = schemaName + "." + seqName
|
||||
if _, err := db.Exec(fmt.Sprintf(
|
||||
"ALTER TABLE %s.%s ALTER COLUMN id SET DEFAULT nextval(%s::regclass)",
|
||||
quoteSQLIdentifier(schemaName),
|
||||
quoteSQLIdentifier(tableName),
|
||||
quoteSQLLiteral(seqRef),
|
||||
)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := db.Exec(fmt.Sprintf(
|
||||
"ALTER SEQUENCE %s.%s OWNED BY %s.%s.id",
|
||||
quoteSQLIdentifier(schemaName),
|
||||
quoteSQLIdentifier(seqName),
|
||||
quoteSQLIdentifier(schemaName),
|
||||
quoteSQLIdentifier(tableName),
|
||||
)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return syncPostgresTableIDSequence(db, schemaName, tableName, seqRef)
|
||||
}
|
||||
|
||||
func syncPostgresTableIDSequence(db *store.DB, schemaName, tableName, seqRef string) error {
|
||||
var maxID int64
|
||||
if err := db.QueryRow(fmt.Sprintf(
|
||||
"SELECT COALESCE(MAX(id), 0) FROM %s.%s",
|
||||
quoteSQLIdentifier(schemaName),
|
||||
quoteSQLIdentifier(tableName),
|
||||
)).Scan(&maxID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
setVal := maxID
|
||||
isCalled := true
|
||||
if maxID <= 0 {
|
||||
setVal = 1
|
||||
isCalled = false
|
||||
}
|
||||
|
||||
if _, err := db.Exec(`SELECT setval(?::regclass, ?, ?)`, seqRef, setVal, isCalled); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func extractNextvalRegclass(defaultExpr string) string {
|
||||
nextvalIdx := strings.Index(strings.ToLower(defaultExpr), "nextval(")
|
||||
if nextvalIdx < 0 {
|
||||
return ""
|
||||
}
|
||||
expr := defaultExpr[nextvalIdx:]
|
||||
firstQuote := strings.Index(expr, "'")
|
||||
if firstQuote < 0 {
|
||||
return ""
|
||||
}
|
||||
expr = expr[firstQuote+1:]
|
||||
secondQuote := strings.Index(expr, "'")
|
||||
if secondQuote < 0 {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(expr[:secondQuote])
|
||||
}
|
||||
|
||||
func quoteSQLIdentifier(ident string) string {
|
||||
return `"` + strings.ReplaceAll(ident, `"`, `""`) + `"`
|
||||
}
|
||||
|
||||
func quoteSQLLiteral(value string) string {
|
||||
return "'" + strings.ReplaceAll(value, "'", "''") + "'"
|
||||
}
|
||||
|
||||
func isMissingColumnError(dialect store.Dialect, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
if dialect == store.DialectPostgres {
|
||||
return strings.Contains(msg, "column") && strings.Contains(msg, "does not exist")
|
||||
}
|
||||
return strings.Contains(msg, "no such column")
|
||||
}
|
||||
|
||||
func (r *Repository) CreatePeerShare(share *PeerShare) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
|
||||
@@ -54,6 +54,12 @@ type pendingRequest struct {
|
||||
ch chan CommandResult
|
||||
}
|
||||
|
||||
const (
|
||||
wsPingPeriod = 15 * time.Second
|
||||
wsPongWait = 45 * time.Second
|
||||
wsWriteWait = 5 * time.Second
|
||||
)
|
||||
|
||||
type CommandResult struct {
|
||||
Type string `json:"type"`
|
||||
Success bool `json:"success"`
|
||||
@@ -120,12 +126,19 @@ func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
cw := &connWrap{conn: conn}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
go startKeepalive(cw, done)
|
||||
|
||||
s.mu.Lock()
|
||||
s.admins[cw] = struct{}{}
|
||||
s.mu.Unlock()
|
||||
|
||||
defer func() {
|
||||
close(done)
|
||||
s.mu.Lock()
|
||||
delete(s.admins, cw)
|
||||
s.mu.Unlock()
|
||||
@@ -145,6 +158,12 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
return
|
||||
}
|
||||
cw := &connWrap{conn: conn}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||
})
|
||||
done := make(chan struct{})
|
||||
go startKeepalive(cw, done)
|
||||
|
||||
version := r.URL.Query().Get("version")
|
||||
httpVal := parseIntDefault(r.URL.Query().Get("http"), 0)
|
||||
@@ -165,6 +184,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
s.broadcastStatus(nodeID, 1)
|
||||
|
||||
defer func() {
|
||||
close(done)
|
||||
needOfflineBroadcast := false
|
||||
s.mu.Lock()
|
||||
current, ok := s.nodes[nodeID]
|
||||
@@ -190,7 +210,15 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
|
||||
msg := decryptIfNeeded(payload, secret)
|
||||
s.tryResolvePending(nodeID, msg)
|
||||
s.broadcastInfo(nodeID, msg)
|
||||
|
||||
var parsed struct {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
if json.Unmarshal([]byte(msg), &parsed) == nil && parsed.Type == "UpgradeProgress" {
|
||||
s.broadcastTyped(nodeID, "upgrade_progress", msg)
|
||||
} else {
|
||||
s.broadcastInfo(nodeID, msg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -264,7 +292,9 @@ func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, tim
|
||||
}
|
||||
|
||||
ns.conn.mu.Lock()
|
||||
_ = ns.conn.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err = ns.conn.conn.WriteMessage(websocket.TextMessage, messageData)
|
||||
_ = ns.conn.conn.SetWriteDeadline(time.Time{})
|
||||
ns.conn.mu.Unlock()
|
||||
if err != nil {
|
||||
cleanup()
|
||||
@@ -385,6 +415,12 @@ func (s *Server) broadcastInfo(nodeID int64, data string) {
|
||||
s.broadcastToAdmins(string(raw))
|
||||
}
|
||||
|
||||
func (s *Server) broadcastTyped(nodeID int64, msgType string, data string) {
|
||||
payload := broadcastMessage{ID: nodeID, Type: msgType, Data: data}
|
||||
raw, _ := json.Marshal(payload)
|
||||
s.broadcastToAdmins(string(raw))
|
||||
}
|
||||
|
||||
func (s *Server) broadcastToAdmins(message string) {
|
||||
s.mu.RLock()
|
||||
admins := make([]*connWrap, 0, len(s.admins))
|
||||
@@ -395,7 +431,9 @@ func (s *Server) broadcastToAdmins(message string) {
|
||||
|
||||
for _, c := range admins {
|
||||
c.mu.Lock()
|
||||
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
|
||||
_ = c.conn.SetWriteDeadline(time.Time{})
|
||||
c.mu.Unlock()
|
||||
if err != nil {
|
||||
log.Printf("websocket broadcast failed: %v", err)
|
||||
@@ -428,3 +466,28 @@ func parseIntDefault(v string, fallback int) int {
|
||||
}
|
||||
return x
|
||||
}
|
||||
|
||||
func startKeepalive(cw *connWrap, done <-chan struct{}) {
|
||||
if cw == nil || cw.conn == nil {
|
||||
return
|
||||
}
|
||||
ticker := time.NewTicker(wsPingPeriod)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
cw.mu.Lock()
|
||||
_ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||
err := cw.conn.WriteMessage(websocket.PingMessage, nil)
|
||||
_ = cw.conn.SetWriteDeadline(time.Time{})
|
||||
cw.mu.Unlock()
|
||||
if err != nil {
|
||||
_ = cw.conn.Close()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
httpserver "go-backend/internal/http"
|
||||
"go-backend/internal/http/handler"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store"
|
||||
"go-backend/internal/store/sqlite"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
@@ -305,7 +306,7 @@ func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func readTableColumns(t *testing.T, db *sql.DB, table string) map[string]bool {
|
||||
func readTableColumns(t *testing.T, db *store.DB, table string) map[string]bool {
|
||||
t.Helper()
|
||||
|
||||
rows, err := db.Query("PRAGMA table_info(" + table + ")")
|
||||
|
||||
Reference in New Issue
Block a user