diff --git a/Wavelet/go.mod b/Wavelet/go.mod index 8a79def9..b7441453 100644 --- a/Wavelet/go.mod +++ b/Wavelet/go.mod @@ -1,44 +1,52 @@ module github.com/Rain-kl/Wavelet -go 1.25.5 +go 1.25.7 require ( - github.com/ClickHouse/clickhouse-go/v2 v2.37.2 + github.com/ClickHouse/clickhouse-go/v2 v2.45.0 github.com/alicebob/miniredis/v2 v2.38.0 + github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.1 github.com/aws/aws-sdk-go-v2 v1.41.5 github.com/aws/aws-sdk-go-v2/config v1.32.14 github.com/aws/aws-sdk-go-v2/credentials v1.19.14 github.com/aws/aws-sdk-go-v2/service/s3 v1.99.0 github.com/bwmarrin/snowflake v0.3.0 github.com/coreos/go-oidc/v3 v3.17.0 + github.com/deepteams/webp v1.2.3 github.com/gin-contrib/sessions v1.0.4 github.com/gin-gonic/gin v1.11.0 github.com/glebarez/sqlite v1.11.0 - github.com/go-jose/go-jose/v4 v4.1.3 + github.com/go-jose/go-jose/v4 v4.1.4 github.com/google/uuid v1.6.0 github.com/gorilla/websocket v1.5.3 github.com/hibiken/asynq v0.25.1 - github.com/pressly/goose/v3 v3.15.1 + github.com/maypok86/otter/v2 v2.3.0 + github.com/peterbourgon/diskv/v3 v3.0.1 + github.com/pressly/goose/v3 v3.27.1 + github.com/rain-kl/openflare v0.0.0 github.com/redis/go-redis/extra/redisotel/v9 v9.16.0 github.com/redis/go-redis/v9 v9.16.0 + github.com/robfig/cron/v3 v3.0.1 github.com/shopspring/decimal v1.4.0 github.com/spf13/cobra v1.10.1 github.com/spf13/viper v1.21.0 github.com/stretchr/testify v1.11.1 + github.com/studio-b12/gowebdav v0.12.0 github.com/swaggo/files v1.0.1 github.com/swaggo/gin-swagger v1.6.1 github.com/swaggo/swag v1.16.6 github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2 go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin v0.61.0 - go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0 - go.opentelemetry.io/otel v1.36.0 + go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.68.0 + go.opentelemetry.io/otel v1.43.0 go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.36.0 - go.opentelemetry.io/otel/sdk v1.36.0 - go.opentelemetry.io/otel/trace v1.36.0 - go.uber.org/zap v1.27.0 + go.opentelemetry.io/otel/sdk v1.43.0 + go.opentelemetry.io/otel/trace v1.43.0 + go.uber.org/zap v1.27.1 golang.org/x/crypto v0.51.0 + golang.org/x/image v0.42.0 golang.org/x/mod v0.36.0 - golang.org/x/oauth2 v0.32.0 + golang.org/x/oauth2 v0.34.0 golang.org/x/sync v0.21.0 gopkg.in/natefinch/lumberjack.v2 v2.2.1 gorm.io/driver/postgres v1.6.0 @@ -48,12 +56,13 @@ require ( gorm.io/plugin/opentelemetry v0.1.14 ) +replace github.com/rain-kl/openflare => ../ + require ( - filippo.io/edwards25519 v1.1.0 // indirect - github.com/ClickHouse/ch-go v0.66.1 // indirect + filippo.io/edwards25519 v1.2.0 // indirect + github.com/ClickHouse/ch-go v0.71.0 // indirect github.com/KyleBanks/depth v1.2.1 // indirect - github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.1 // indirect - github.com/andybalholm/brotli v1.2.0 // indirect + github.com/andybalholm/brotli v1.2.1 // indirect github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.8 // indirect github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.21 // indirect github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.21 // indirect @@ -77,12 +86,12 @@ require ( github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/cloudwego/base64x v0.1.6 // indirect github.com/davecgh/go-spew v1.1.1 // indirect - github.com/deepteams/webp v1.2.3 // indirect + github.com/dgraph-io/ristretto/v2 v2.2.0 // indirect github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/felixge/httpsnoop v1.0.4 // indirect github.com/fsnotify/fsnotify v1.9.0 // indirect - github.com/gabriel-vasile/mimetype v1.4.11 // indirect + github.com/gabriel-vasile/mimetype v1.4.13 // indirect github.com/gin-contrib/sse v1.1.0 // indirect github.com/glebarez/go-sqlite v1.21.2 // indirect github.com/go-faster/city v1.0.1 // indirect @@ -106,74 +115,73 @@ require ( github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/goccy/go-json v0.10.5 // indirect github.com/goccy/go-yaml v1.18.0 // indirect - github.com/gomodule/redigo v1.9.3 // indirect + github.com/gomodule/redigo v2.0.0+incompatible // indirect github.com/google/btree v1.0.0 // indirect github.com/gorilla/context v1.1.2 // indirect github.com/gorilla/securecookie v1.1.2 // indirect github.com/gorilla/sessions v1.4.0 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 // indirect - github.com/hashicorp/go-version v1.7.0 // indirect + github.com/hashicorp/go-version v1.8.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect - github.com/jackc/pgx/v5 v5.7.6 // indirect + github.com/jackc/pgx/v5 v5.9.2 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect - github.com/json-iterator/go v1.1.12 // indirect - github.com/klauspost/compress v1.18.1 // indirect + github.com/json-iterator/go v1.1.13-0.20220915233716-71ac16282d12 // indirect + github.com/klauspost/compress v1.18.5 // indirect github.com/klauspost/cpuid/v2 v2.3.0 // indirect github.com/leodido/go-urn v1.4.0 // indirect - github.com/mattn/go-isatty v0.0.20 // indirect + github.com/mattn/go-isatty v0.0.21 // indirect github.com/mattn/go-sqlite3 v1.14.22 // indirect - github.com/maypok86/otter/v2 v2.3.0 // indirect + github.com/mfridman/interpolate v0.0.2 // indirect github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect - github.com/modern-go/reflect2 v1.0.2 // indirect - github.com/paulmach/orb v0.12.0 // indirect + github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect + github.com/oschwald/maxminddb-golang v1.13.1 // indirect + github.com/paulmach/orb v0.13.0 // indirect github.com/pelletier/go-toml/v2 v2.2.4 // indirect - github.com/peterbourgon/diskv/v3 v3.0.1 // indirect - github.com/pierrec/lz4/v4 v4.1.22 // indirect + github.com/pierrec/lz4/v4 v4.1.26 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/quic-go/qpack v0.5.1 // indirect github.com/quic-go/quic-go v0.55.0 // indirect github.com/redis/go-redis/extra/rediscmd/v9 v9.16.0 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect - github.com/robfig/cron/v3 v3.0.1 // indirect github.com/sagikazarmark/locafero v0.12.0 // indirect github.com/segmentio/asm v1.2.1 // indirect + github.com/sethvargo/go-retry v0.3.0 // indirect github.com/spf13/afero v1.15.0 // indirect github.com/spf13/cast v1.10.0 // indirect github.com/spf13/pflag v1.0.10 // indirect - github.com/studio-b12/gowebdav v0.12.0 // indirect github.com/subosito/gotenv v1.6.0 // indirect github.com/twitchyliquid64/golang-asm v0.15.1 // indirect github.com/ugorji/go/codec v1.3.1 // indirect github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2 // indirect github.com/yuin/gopher-lua v1.1.1 // indirect - go.opentelemetry.io/auto/sdk v1.1.0 // indirect + go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.36.0 // indirect go.opentelemetry.io/otel/log v0.12.2 // indirect - go.opentelemetry.io/otel/metric v1.36.0 // indirect + go.opentelemetry.io/otel/metric v1.43.0 // indirect go.opentelemetry.io/proto/otlp v1.6.0 // indirect go.uber.org/mock v0.6.0 // indirect go.uber.org/multierr v1.11.0 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/arch v0.22.0 // indirect - golang.org/x/image v0.42.0 // indirect golang.org/x/net v0.54.0 // indirect golang.org/x/sys v0.44.0 // indirect golang.org/x/text v0.38.0 // indirect golang.org/x/time v0.14.0 // indirect golang.org/x/tools v0.45.0 // indirect - google.golang.org/genproto/googleapis/api v0.0.0-20251103181224-f26f9409b101 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20251103181224-f26f9409b101 // indirect - google.golang.org/grpc v1.72.1 // indirect - google.golang.org/protobuf v1.36.10 // indirect + google.golang.org/genproto/googleapis/api v0.0.0-20260120221211-b8f7ae30c516 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260420184626-e10c466a9529 // indirect + google.golang.org/grpc v1.80.0 // indirect + google.golang.org/protobuf v1.36.11 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect gorm.io/driver/clickhouse v0.7.0 // indirect gorm.io/driver/mysql v1.6.0 // indirect - modernc.org/libc v1.24.1 // indirect - modernc.org/mathutil v1.6.0 // indirect - modernc.org/memory v1.7.2 // indirect - modernc.org/sqlite v1.26.0 // indirect + modernc.org/libc v1.72.1 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect + modernc.org/sqlite v1.49.1 // indirect ) diff --git a/Wavelet/go.sum b/Wavelet/go.sum index 900a42c4..78f1e86b 100644 --- a/Wavelet/go.sum +++ b/Wavelet/go.sum @@ -1,17 +1,17 @@ -filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= -filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= -github.com/ClickHouse/ch-go v0.66.1 h1:LQHFslfVYZsISOY0dnOYOXGkOUvpv376CCm8g7W74A4= -github.com/ClickHouse/ch-go v0.66.1/go.mod h1:NEYcg3aOFv2EmTJfo4m2WF7sHB/YFbLUuIWv9iq76xY= -github.com/ClickHouse/clickhouse-go/v2 v2.37.2 h1:wRLNKoynvHQEN4znnVHNLaYnrqVc9sGJmGYg+GGCfto= -github.com/ClickHouse/clickhouse-go/v2 v2.37.2/go.mod h1:pH2zrBGp5Y438DMwAxXMm1neSXPPjSI7tD4MURVULw8= +filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= +filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= +github.com/ClickHouse/ch-go v0.71.0 h1:bUdZ/EZj/LcVHsMqaRUP2holqygrPWQKeMjc6nZoyRM= +github.com/ClickHouse/ch-go v0.71.0/go.mod h1:NwbNc+7jaqfY58dmdDUbG4Jl22vThgx1cYjBw0vtgXw= +github.com/ClickHouse/clickhouse-go/v2 v2.45.0 h1:iHt15nA4iYhfde5bDQAcLAat9BAh7B5ksPRNRa4UI7s= +github.com/ClickHouse/clickhouse-go/v2 v2.45.0/go.mod h1:giJfUVlMkcfUEPVfRpt51zZaGEx9i17gCos8gBl392c= github.com/KyleBanks/depth v1.2.1 h1:5h8fQADFrWtarTdtDudMmGsC7GPbOAu6RVB3ffsVFHc= github.com/KyleBanks/depth v1.2.1/go.mod h1:jzSb9d0L43HxTQfT+oSA1EEp2q+ne2uh6XgeJcm8brE= github.com/alicebob/miniredis/v2 v2.38.0 h1:nZAzCR+Lj+Vxk4ZXzm2NuKq2O33RXj1XxJ2e2uP9jiw= github.com/alicebob/miniredis/v2 v2.38.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.1 h1:vtiFd0hhPAbyYJjztl0wYUq/PqEGkIlDmVuTIy6zw8Y= github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.1/go.mod h1:FTzydeQVmR24FI0D6XWUOMKckjXehM/jgMn1xC+DA9M= -github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwToPjQ= -github.com/andybalholm/brotli v1.2.0/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= +github.com/andybalholm/brotli v1.2.1 h1:R+f5xP285VArJDRgowrfb9DqL18yVK0gKAW/F+eTWro= +github.com/andybalholm/brotli v1.2.1/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= github.com/aws/aws-sdk-go-v2 v1.41.5 h1:dj5kopbwUsVUVFgO4Fi5BIT3t4WyqIDjGKCangnV/yY= github.com/aws/aws-sdk-go-v2 v1.41.5/go.mod h1:mwsPRE8ceUUpiTgF7QmQIJ7lgsKUPQOUl3o72QBrE1o= github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.8 h1:eBMB84YGghSocM7PsjmmPffTa+1FBUeNvGvFou6V/4o= @@ -78,6 +78,10 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/deepteams/webp v1.2.3 h1:TedmP3+U8/xBrC0dTkbd3+wbmnM4EIaMowQJpLz6UqU= github.com/deepteams/webp v1.2.3/go.mod h1:J8Ap+HAixxpKKRN9IpEeSKlfvhsef1v43jKTO7m3f4c= +github.com/dgraph-io/ristretto/v2 v2.2.0 h1:bkY3XzJcXoMuELV8F+vS8kzNgicwQFAaGINAEJdWGOM= +github.com/dgraph-io/ristretto/v2 v2.2.0/go.mod h1:RZrm63UmcBAaYWC1DotLYBmTvgkrs0+XhBd7Npn7/zI= +github.com/dgryski/go-farm v0.0.0-20240924180020-3414d57e47da h1:aIftn67I1fkbMa512G+w+Pxci9hJPB8oMnkcP3iZF38= +github.com/dgryski/go-farm v0.0.0-20240924180020-3414d57e47da/go.mod h1:SqUrOPUnsFjfmXRMNPybcSiG0BgUW2AuFH8PAnS2iTw= github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78= github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= @@ -88,8 +92,8 @@ github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHk github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0= github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= -github.com/gabriel-vasile/mimetype v1.4.11 h1:AQvxbp830wPhHTqc1u7nzoLT+ZFxGY7emj5DR5DYFik= -github.com/gabriel-vasile/mimetype v1.4.11/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s= +github.com/gabriel-vasile/mimetype v1.4.13 h1:46nXokslUBsAJE/wMsp5gtO500a4F3Nkz9Ufpk2AcUM= +github.com/gabriel-vasile/mimetype v1.4.13/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s= github.com/gin-contrib/gzip v0.0.6 h1:NjcunTcGAj5CO1gn4N8jHOSIeRFHIbn51z6K+xaN4d4= github.com/gin-contrib/gzip v0.0.6/go.mod h1:QOJlmV2xmayAjkNS2Y8NQsMneuRShOU/kjovCXNuzzk= github.com/gin-contrib/sessions v1.0.4 h1:ha6CNdpYiTOK/hTp05miJLbpTSNfOnFg5Jm2kbcqy8U= @@ -106,8 +110,8 @@ github.com/go-faster/city v1.0.1 h1:4WAxSZ3V2Ws4QRDrscLEDcibJY8uf41H6AhXDrNDcGw= github.com/go-faster/city v1.0.1/go.mod h1:jKcUJId49qdW3L1qKHH/3wPeUstCVpVSXTM6vO3VcTw= github.com/go-faster/errors v0.7.1 h1:MkJTnDoEdi9pDabt1dpWf7AA8/BaSYZqibYyhZ20AYg= github.com/go-faster/errors v0.7.1/go.mod h1:5ySTjWFiphBs07IKuiL69nxdfd5+fzh1u7FPGZP2quo= -github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs= -github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= +github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA= +github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A= github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= @@ -152,24 +156,19 @@ github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4= github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= github.com/goccy/go-yaml v1.18.0 h1:8W7wMFS12Pcas7KU+VVkaiCng+kG8QiFeFwzFb+rwuw= github.com/goccy/go-yaml v1.18.0/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= -github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= -github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk= github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= -github.com/golang/snappy v0.0.1/go.mod h1:/XxbfmMg8lxefKM7IXC3fBNl/7bRcc72aCRzEWrmP2Q= -github.com/gomodule/redigo v1.9.3 h1:dNPSXeXv6HCq2jdyWfjgmhBdqnR6PRO3m/G05nvpPC8= -github.com/gomodule/redigo v1.9.3/go.mod h1:KsU3hiK/Ay8U42qpaJk+kuNa3C+spxapWpM+ywhcgtw= +github.com/gomodule/redigo v2.0.0+incompatible h1:K/R+8tc58AaqLkqG2Ol3Qk+DR/TlNuhuh457pBFPtt0= +github.com/gomodule/redigo v2.0.0+incompatible/go.mod h1:B4C85qUVwatsJoIUNIfCRsp7qO0iAmpGFZ4EELWSbC4= github.com/google/btree v1.0.0 h1:0udJVsspx3VBr5FwtLhQQtuAsVc79tTq0ocGIPAU6qo= github.com/google/btree v1.0.0/go.mod h1:lNA+9X1NB3Zf8V7Ke586lFgjr2dZNuvo3lPJSGZ5JPQ= -github.com/google/go-cmp v0.5.2/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= -github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/gofuzz v1.2.0 h1:xRy4A+RhZaiKjJ1bPfwQ8sedCA+YS2YcCHW6ec7JMi0= github.com/google/gofuzz v1.2.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= -github.com/google/pprof v0.0.0-20221118152302-e6195bd50e26 h1:Xim43kblpZXfIBQsbuBVKCudVG457BR2GZFIz3uw3hQ= -github.com/google/pprof v0.0.0-20221118152302-e6195bd50e26/go.mod h1:dDKJzRmX4S37WGHujM7tX//fmj1uioxKzKxz3lo4HJo= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gorilla/context v1.1.2 h1:WRkNAv2uoa03QNIc1A6u4O7DAGMUVoopZhkiXWA2V1o= @@ -182,8 +181,10 @@ github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aN github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 h1:5ZPtiqj0JL5oKWmcsq4VMaAW5ukBEgSGXEN89zeH1Jo= github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3/go.mod h1:ndYquD05frm2vACXE1nsccT4oJzjhw2arTS2cpUD1PI= -github.com/hashicorp/go-version v1.7.0 h1:5tqGy27NaOTB8yJKUZELlFAS/LTKJkrmONwQKeRZfjY= -github.com/hashicorp/go-version v1.7.0/go.mod h1:fltr4n8CU8Ke44wwGCBoEymUuxUHl09ZGVZPK5anwXA= +github.com/hashicorp/go-version v1.8.0 h1:KAkNb1HAiZd1ukkxDFGmokVZe1Xy9HG6NUp+bPle2i4= +github.com/hashicorp/go-version v1.8.0/go.mod h1:fltr4n8CU8Ke44wwGCBoEymUuxUHl09ZGVZPK5anwXA= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/hibiken/asynq v0.25.1 h1:phj028N0nm15n8O2ims+IvJ2gz4k2auvermngh9JhTw= github.com/hibiken/asynq v0.25.1/go.mod h1:pazWNOLBu0FEynQRBvHA26qdIKRSmfdIfUm4HdsLmXg= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= @@ -192,60 +193,56 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI 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.6 h1:rWQc5FwZSPX58r1OQmkuaNicxdmExaEz5A2DO2hUuTk= -github.com/jackc/pgx/v5 v5.7.6/go.mod h1:aruU7o91Tc2q2cFp5h4uP3f6ztExVpyVv88Xl/8Vl8M= +github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw= +github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E= github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= -github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= -github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= -github.com/kballard/go-shellquote v0.0.0-20180428030007-95032a82bc51 h1:Z9n2FFNUXsshfwJMBgNA0RU6/i7WVaAegv3PtuIHPMs= -github.com/kballard/go-shellquote v0.0.0-20180428030007-95032a82bc51/go.mod h1:CzGEWj7cYgsdH8dAjBGEr58BoE7ScuLd+fwFZ44+/x8= -github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= -github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= -github.com/klauspost/compress v1.13.6/go.mod h1:/3/Vjq9QcHkK5uEr5lBEmyoZ1iFhe47etQ6QUkpK6sk= -github.com/klauspost/compress v1.18.1 h1:bcSGx7UbpBqMChDtsF28Lw6v/G94LPrrbMbdC3JH2co= -github.com/klauspost/compress v1.18.1/go.mod h1:ZQFFVG+MdnR0P+l6wpXgIL4NTtwiKIdBnrBd8Nrxr+0= +github.com/json-iterator/go v1.1.13-0.20220915233716-71ac16282d12 h1:9Nu54bhS/H/Kgo2/7xNSUuC5G28VR8ljfrLKU2G4IjU= +github.com/json-iterator/go v1.1.13-0.20220915233716-71ac16282d12/go.mod h1:TBzl5BIHNXfS9+C35ZyJaklL7mLDbgUkcgXzSLa8Tk0= +github.com/klauspost/compress v1.18.5 h1:/h1gH5Ce+VWNLSWqPzOVn6XBO+vJbCNGvjoaGBFW2IE= +github.com/klauspost/compress v1.18.5/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= -github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= -github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= -github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ= github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI= -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/mattn/go-isatty v0.0.21 h1:xYae+lCNBP7QuW4PUnNG61ffM4hVIfm+zUzDuSzYLGs= +github.com/mattn/go-isatty v0.0.21/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4= github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU= github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= github.com/maypok86/otter/v2 v2.3.0 h1:8H8AVVFUSzJwIegKwv1uF5aGitTY+AIrtktg7OcLs8w= github.com/maypok86/otter/v2 v2.3.0/go.mod h1:XgIdlpmL6jYz882/CAx1E4C1ukfgDKSaw4mWq59+7l8= +github.com/mfridman/interpolate v0.0.2 h1:pnuTK7MQIxxFz1Gr+rjSIx9u7qVjf5VOoM/u6BbAxPY= +github.com/mfridman/interpolate v0.0.2/go.mod h1:p+7uk6oE07mpE/Ik1b8EckO0O4ZXiGAfshKBWLUM9Xg= github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= -github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M= github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= -github.com/montanaflynn/stats v0.0.0-20171201202039-1bf9dbcd8cbe/go.mod h1:wL8QJuTMNUDYhXwkmfOly8iTdp5TEcJFWZD2D7SIkUc= -github.com/paulmach/orb v0.12.0 h1:z+zOwjmG3MyEEqzv92UN49Lg1JFYx0L9GpGKNVDKk1s= -github.com/paulmach/orb v0.12.0/go.mod h1:5mULz1xQfs3bmQm63QEJA6lNGujuRafwA5S/EnuLaLU= -github.com/paulmach/protoscan v0.2.1/go.mod h1:SpcSwydNLrxUGSDvXvO0P7g7AuhJ7lcKfDlhJCDw2gY= +github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee h1:W5t00kpgFdJifH4BDsTlE89Zl93FEloxaWZfGcifgq8= +github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/oschwald/maxminddb-golang v1.13.1 h1:G3wwjdN9JmIK2o/ermkHM+98oX5fS+k5MbwsmL4MRQE= +github.com/oschwald/maxminddb-golang v1.13.1/go.mod h1:K4pgV9N/GcK694KSTmVSDTODk4IsCNThNdTmnaBZ/F8= +github.com/paulmach/orb v0.13.0 h1:r7n7mQGGF+cj/CbcivEj9J3HGK+XR+yXnvzRdq9saIw= +github.com/paulmach/orb v0.13.0/go.mod h1:6scRWINywA2Jf05dcjOfLfxrUIMECvTSG2MVbRLxu/k= github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= github.com/peterbourgon/diskv/v3 v3.0.1 h1:x06SQA46+PKIUftmEujdwSEpIx8kR+M9eLYsUxeYveU= github.com/peterbourgon/diskv/v3 v3.0.1/go.mod h1:kJ5Ny7vLdARGU3WUuy6uzO6T0nb/2gWcT1JiBvRmb5o= -github.com/pierrec/lz4/v4 v4.1.22 h1:cKFw6uJDK+/gfw5BcDL0JL5aBsAFdsIT18eRtLj7VIU= -github.com/pierrec/lz4/v4 v4.1.22/go.mod h1:gZWDp/Ze/IJXGXf23ltt2EXimqmTUXEy0GFuRQyBid4= -github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= +github.com/pierrec/lz4/v4 v4.1.26 h1:GrpZw1gZttORinvzBdXPUXATeqlJjqUG/D87TKMnhjY= +github.com/pierrec/lz4/v4 v4.1.26/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= 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/pressly/goose/v3 v3.15.1 h1:dKaJ1SdLvS/+HtS8PzFT0KBEtICC1jewLXM+b3emlv8= -github.com/pressly/goose/v3 v3.15.1/go.mod h1:0E3Yg/+EwYzO6Rz2P98MlClFgIcoujbVRs575yi3iIM= +github.com/pressly/goose/v3 v3.27.1 h1:6uEvcprBybDmW4hcz3gYujhARhye+GoWKhEWyzD5sh4= +github.com/pressly/goose/v3 v3.27.1/go.mod h1:maruOxsPnIG2yHHyo8UqKWXYKFcH7Q76csUV7+7KYoM= github.com/quic-go/qpack v0.5.1 h1:giqksBPnT/HDtZ6VhtFKgoLOWmlyo9Ei6u9PqzIMbhI= github.com/quic-go/qpack v0.5.1/go.mod h1:+PC4XFrEskIVkcLzpEkbLqq1uCoxPhQuvK5rH1ZgaEg= github.com/quic-go/quic-go v0.55.0 h1:zccPQIqYCXDt5NmcEabyYvOnomjs8Tlwl7tISjJh9Mk= @@ -260,13 +257,15 @@ github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94 github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= -github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII= -github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o= +github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= +github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4= github.com/sagikazarmark/locafero v0.12.0/go.mod h1:sZh36u/YSZ918v0Io+U9ogLYQJ9tLLBmM4eneO6WwsI= github.com/segmentio/asm v1.2.1 h1:DTNbBqs57ioxAD4PrArqftgypG4/qNpXoJx8TVXxPR0= github.com/segmentio/asm v1.2.1/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs= +github.com/sethvargo/go-retry v0.3.0 h1:EEt31A35QhrcRZtrYFDTBg91cqZVnFL2navjDrah2SE= +github.com/sethvargo/go-retry v0.3.0/go.mod h1:mNX17F0C/HguQMyMyJxcnU471gOZGxCLyYaFyAZraas= github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k= github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+DMd9qYNcwME= github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I= @@ -285,7 +284,6 @@ github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSS github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= -github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= @@ -303,7 +301,6 @@ github.com/swaggo/gin-swagger v1.6.1 h1:Ri06G4gc9N4t4k8hekMigJ9zKTFSlqj/9paAQCQs github.com/swaggo/gin-swagger v1.6.1/go.mod h1:LQ+hJStHakCWRiK/YNYtJOu4mR2FP+pxLnILT/qNiTw= github.com/swaggo/swag v1.16.6 h1:qBNcx53ZaX+M5dxVyTrgQ0PJ/ACK+NzhwcbieTt+9yI= github.com/swaggo/swag v1.16.6/go.mod h1:ngP2etMK5a0P3QBizic5MEwpRmluJZPHjXcMoj4Xesg= -github.com/tidwall/pretty v1.0.0/go.mod h1:XNkn88O1ChpSDQmQeStsy+sBenx6DDtFZJxhVysOjyk= github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI= github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY= @@ -312,26 +309,19 @@ github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2 h1:3/aHKUq7qaFMWxyQV0W github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2/go.mod h1:Zit4b8AQXaXvA68+nzmbyDzqiyFRISyw1JiD5JqUBjw= github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2 h1:cj/Z6FKTTYBnstI0Lni9PA+k2foounKIPUmj1LBwNiQ= github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2/go.mod h1:LDaXk90gKEC2nC7JH3Lpnhfu+2V7o/TsqomJJmqA39o= -github.com/xdg-go/pbkdf2 v1.0.0/go.mod h1:jrpuAogTd400dnrH08LKmI/xc1MbPOebTwRqcT5RDeI= -github.com/xdg-go/scram v1.1.1/go.mod h1:RaEWvsqvNKKvBPvcKeFjrG2cJqOkHTiyTpzz23ni57g= -github.com/xdg-go/stringprep v1.0.3/go.mod h1:W3f5j4i+9rC0kuIEJL0ky1VpHXQU3ocBgklLGvcBnW8= github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU= github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= -github.com/youmark/pkcs8 v0.0.0-20181117223130-1be2e3e5546d/go.mod h1:rHwXgn7JulP+udvsHwJoVG1YGAP6VLg4y9I5dyZdqmA= -github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= -github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M= github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw= -go.mongodb.org/mongo-driver v1.11.4/go.mod h1:PTSz5yu21bkT/wXpkS7WR5f0ddqw5quethTUn9WM+2g= -go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA= -go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin v0.61.0 h1:VkrF0D14uQrCmPqBkYlwWnhgcwzXvIRAjX8eXO7vy6M= go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin v0.61.0/go.mod h1:p/mVr/Hs7gQnguNPXUyuiMRNtisyc9y/Oo7Kqr/6wbU= -go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0 h1:F7Jx+6hwnZ41NSFTO5q4LYDtJRXBf2PD0rNBkeB/lus= -go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0/go.mod h1:UHB22Z8QsdRDrnAtX4PntOl36ajSxcdUMt1sF7Y6E7Q= -go.opentelemetry.io/otel v1.36.0 h1:UumtzIklRBY6cI/lllNZlALOF5nNIzJVb16APdvgTXg= -go.opentelemetry.io/otel v1.36.0/go.mod h1:/TcFMXYjyRNh8khOAO9ybYkqaDBb/70aVwkNML4pP8E= +go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.68.0 h1:CqXxU8VOmDefoh0+ztfGaymYbhdB/tT3zs79QaZTNGY= +go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.68.0/go.mod h1:BuhAPThV8PBHBvg8ZzZ/Ok3idOdhWIodywz2xEcRbJo= +go.opentelemetry.io/otel v1.43.0 h1:mYIM03dnh5zfN7HautFE4ieIig9amkNANT+xcVxAj9I= +go.opentelemetry.io/otel v1.43.0/go.mod h1:JuG+u74mvjvcm8vj8pI5XiHy1zDeoCS2LB1spIq7Ay0= go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.36.0 h1:dNzwXjZKpMpE2JhmO+9HsPl42NIXFIFSUSSs0fiqra0= go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.36.0/go.mod h1:90PoxvaEB5n6AOdZvi+yWJQoE95U8Dhhw2bSyRqnTD0= go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.36.0 h1:JgtbA0xkWHnTmYk7YusopJFX6uleBmAuZ8n05NEh8nQ= @@ -340,14 +330,14 @@ go.opentelemetry.io/otel/exporters/stdout/stdouttrace v1.36.0 h1:G8Xec/SgZQricwW go.opentelemetry.io/otel/exporters/stdout/stdouttrace v1.36.0/go.mod h1:PD57idA/AiFD5aqoxGxCvT/ILJPeHy3MjqU/NS7KogY= go.opentelemetry.io/otel/log v0.12.2 h1:yob9JVHn2ZY24byZeaXpTVoPS6l+UrrxmxmPKohXTwc= go.opentelemetry.io/otel/log v0.12.2/go.mod h1:ShIItIxSYxufUMt+1H5a2wbckGli3/iCfuEbVZi/98E= -go.opentelemetry.io/otel/metric v1.36.0 h1:MoWPKVhQvJ+eeXWHFBOPoBOi20jh6Iq2CcCREuTYufE= -go.opentelemetry.io/otel/metric v1.36.0/go.mod h1:zC7Ks+yeyJt4xig9DEw9kuUFe5C3zLbVjV2PzT6qzbs= -go.opentelemetry.io/otel/sdk v1.36.0 h1:b6SYIuLRs88ztox4EyrvRti80uXIFy+Sqzoh9kFULbs= -go.opentelemetry.io/otel/sdk v1.36.0/go.mod h1:+lC+mTgD+MUWfjJubi2vvXWcVxyr9rmlshZni72pXeY= -go.opentelemetry.io/otel/sdk/metric v1.36.0 h1:r0ntwwGosWGaa0CrSt8cuNuTcccMXERFwHX4dThiPis= -go.opentelemetry.io/otel/sdk/metric v1.36.0/go.mod h1:qTNOhFDfKRwX0yXOqJYegL5WRaW376QbB7P4Pb0qva4= -go.opentelemetry.io/otel/trace v1.36.0 h1:ahxWNuqZjpdiFAyrIoQ4GIiAIhxAunQR6MUoKrsNd4w= -go.opentelemetry.io/otel/trace v1.36.0/go.mod h1:gQ+OnDZzrybY4k4seLzPAWNwVBBVlF2szhehOBB/tGA= +go.opentelemetry.io/otel/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM= +go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY= +go.opentelemetry.io/otel/sdk v1.43.0 h1:pi5mE86i5rTeLXqoF/hhiBtUNcrAGHLKQdhg4h4V9Dg= +go.opentelemetry.io/otel/sdk v1.43.0/go.mod h1:P+IkVU3iWukmiit/Yf9AWvpyRDlUeBaRg6Y+C58QHzg= +go.opentelemetry.io/otel/sdk/metric v1.43.0 h1:S88dyqXjJkuBNLeMcVPRFXpRw2fuwdvfCGLEo89fDkw= +go.opentelemetry.io/otel/sdk/metric v1.43.0/go.mod h1:C/RJtwSEJ5hzTiUz5pXF1kILHStzb9zFlIEe85bhj6A= +go.opentelemetry.io/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09nk+3A= +go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0= go.opentelemetry.io/proto/otlp v1.6.0 h1:jQjP+AQyTf+Fe7OKj/MfkDrmK4MNVtw2NpXsf9fefDI= go.opentelemetry.io/proto/otlp v1.6.0/go.mod h1:cicgGehlFuNdgZkcALOCh3VE6K/u2tAjzlRhDwmVpZc= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= @@ -356,65 +346,39 @@ go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y= go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU= go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0= go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y= -go.uber.org/zap v1.27.0 h1:aJMhYGrd5QSmlpLMr2MftRKl7t8J8PTZPA732ud/XR8= -go.uber.org/zap v1.27.0/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E= +go.uber.org/zap v1.27.1 h1:08RqriUEv8+ArZRYSTXy1LeBScaMpVSTBhCeaZYfMYc= +go.uber.org/zap v1.27.1/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E= go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= golang.org/x/arch v0.22.0 h1:c/Zle32i5ttqRXjdLyyHZESLD/bB90DCU1g9l/0YBDI= golang.org/x/arch v0.22.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= -golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= -golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= -golang.org/x/crypto v0.0.0-20220622213112-05595931fe9d/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4= -golang.org/x/crypto v0.47.0 h1:V6e3FRj+n4dbpw86FJ8Fv7XVOql7TEwpHapKoMJ/GO8= -golang.org/x/crypto v0.47.0/go.mod h1:ff3Y9VzzKbwSSEzWqJsJVBnWmRwRSHt/6Op5n9bQc4A= golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= golang.org/x/image v0.42.0 h1:1gSs6ehNWXLbkHBIPcWztk3D/6aIA/8hauiAYtlodVY= golang.org/x/image v0.42.0/go.mod h1:rrpelvGFt+kLPAjPM4HeWPgrl0FtafueU//e5N0qk/Q= -golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= -golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= -golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c= -golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU= golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= -golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= -golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= -golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= -golang.org/x/net v0.0.0-20211112202133-69e39bad7dc2/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= -golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o= -golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8= golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w= golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ= -golang.org/x/oauth2 v0.32.0 h1:jsCblLleRMDrxMN29H3z/k1KliIvpLgCkE6R8FXXNgY= -golang.org/x/oauth2 v0.32.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA= +golang.org/x/oauth2 v0.34.0 h1:hqK/t4AKgbqWkdkcAeI8XLmbK+4m4G5YeQRrmiotGlw= +golang.org/x/oauth2 v0.34.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= -golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= -golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ= -golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= @@ -422,40 +386,29 @@ golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuX golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= -golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= -golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= -golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= golang.org/x/text v0.38.0 h1:sXmwo9DwP3OK9EZ7PqAdaooSGozfl/3a6/xJcbzPRhE= golang.org/x/text v0.38.0/go.mod h1:YXZt3QhHUKYT53r2lLKFIVi6Ao1jdzrTR/KQ09qyxF4= golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI= golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= -golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE= -golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA= golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= -golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc= -golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg= golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -google.golang.org/genproto/googleapis/api v0.0.0-20251103181224-f26f9409b101 h1:vk5TfqZHNn0obhPIYeS+cxIFKFQgser/M2jnI+9c6MM= -google.golang.org/genproto/googleapis/api v0.0.0-20251103181224-f26f9409b101/go.mod h1:E17fc4PDhkr22dE3RgnH2hEubUaky6ZwW4VhANxyspg= -google.golang.org/genproto/googleapis/rpc v0.0.0-20251103181224-f26f9409b101 h1:tRPGkdGHuewF4UisLzzHHr1spKw92qLM98nIzxbC0wY= -google.golang.org/genproto/googleapis/rpc v0.0.0-20251103181224-f26f9409b101/go.mod h1:7i2o+ce6H/6BluujYR+kqX3GKH+dChPTQU19wjRPiGk= -google.golang.org/grpc v1.72.1 h1:HR03wO6eyZ7lknl75XlxABNVLLFc2PAb6mHlYh756mA= -google.golang.org/grpc v1.72.1/go.mod h1:wH5Aktxcg25y1I3w7H69nHfXdOG3UiadoBtjh3izSDM= -google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw= -google.golang.org/protobuf v1.27.1/go.mod h1:9q0QmTI4eRPtz6boOQmLYwt+qCgq0jsYwAQnmE0givc= -google.golang.org/protobuf v1.36.10 h1:AYd7cD/uASjIL6Q9LiTjz8JLcrh/88q5UObnmY3aOOE= -google.golang.org/protobuf v1.36.10/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +google.golang.org/genproto/googleapis/api v0.0.0-20260120221211-b8f7ae30c516 h1:vmC/ws+pLzWjj/gzApyoZuSVrDtF1aod4u/+bbj8hgM= +google.golang.org/genproto/googleapis/api v0.0.0-20260120221211-b8f7ae30c516/go.mod h1:p3MLuOwURrGBRoEyFHBT3GjUwaCQVKeNqqWxlcISGdw= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260420184626-e10c466a9529 h1:XF8+t6QQiS0o9ArVan/HW8Q7cycNPGsJf6GA2nXxYAg= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260420184626-e10c466a9529/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/grpc v1.80.0 h1:Xr6m2WmWZLETvUNvIUmeD5OAagMw3FiKmMlTdViWsHM= +google.golang.org/grpc v1.80.0/go.mod h1:ho/dLnxwi3EDJA4Zghp7k2Ec1+c2jqup0bFkw07bwF4= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc= @@ -477,23 +430,31 @@ gorm.io/plugin/dbresolver v1.6.2 h1:F4b85TenghUeITqe3+epPSUtHH7RIk3fXr5l83DF8Pc= gorm.io/plugin/dbresolver v1.6.2/go.mod h1:tctw63jdrOezFR9HmrKnPkmig3m5Edem9fdxk9bQSzM= gorm.io/plugin/opentelemetry v0.1.14 h1:xivP39t/0JgcceDl+BLwVAJHihjFEUj0ZocMSBwZ7ZY= gorm.io/plugin/opentelemetry v0.1.14/go.mod h1:ZAp4v5vU1CCcK9Oo8/va5rl6NStrzpSU+a70evd+W/g= -lukechampine.com/uint128 v1.3.0 h1:cDdUVfRwDUDovz610ABgFD17nXD4/uDgVHl2sC3+sbo= -lukechampine.com/uint128 v1.3.0/go.mod h1:c4eWIwlEGaxC/+H1VguhU4PHXNWDCDMUlWdIWl2j1gk= -modernc.org/cc/v3 v3.41.0 h1:QoR1Sn3YWlmA1T4vLaKZfawdVtSiGx8H+cEojbC7v1Q= -modernc.org/cc/v3 v3.41.0/go.mod h1:Ni4zjJYJ04CDOhG7dn640WGfwBzfE0ecX8TyMB0Fv0Y= -modernc.org/ccgo/v3 v3.16.15 h1:KbDR3ZAVU+wiLyMESPtbtE/Add4elztFyfsWoNTgxS0= -modernc.org/ccgo/v3 v3.16.15/go.mod h1:yT7B+/E2m43tmMOT51GMoM98/MtHIcQQSleGnddkUNI= -modernc.org/libc v1.24.1 h1:uvJSeCKL/AgzBo2yYIPPTy82v21KgGnizcGYfBHaNuM= -modernc.org/libc v1.24.1/go.mod h1:FmfO1RLrU3MHJfyi9eYYmZBfi/R+tqZ6+hQ3yQQUkak= -modernc.org/mathutil v1.6.0 h1:fRe9+AmYlaej+64JsEEhoWuAYBkOtQiMEU7n/XgfYi4= -modernc.org/mathutil v1.6.0/go.mod h1:Ui5Q9q1TR2gFm0AQRqQUaBWFLAhQpCwNcuhBOSedWPo= -modernc.org/memory v1.7.2 h1:Klh90S215mmH8c9gO98QxQFsY+W451E8AnzjoE2ee1E= -modernc.org/memory v1.7.2/go.mod h1:NO4NVCQy0N7ln+T9ngWqOQfi7ley4vpwvARR+Hjw95E= -modernc.org/opt v0.1.3 h1:3XOZf2yznlhC+ibLltsDGzABUGVx8J6pnFMS3E4dcq4= -modernc.org/opt v0.1.3/go.mod h1:WdSiB5evDcignE70guQKxYUl14mgWtbClRi5wmkkTX0= -modernc.org/sqlite v1.26.0 h1:SocQdLRSYlA8W99V8YH0NES75thx19d9sB/aFc4R8Lw= -modernc.org/sqlite v1.26.0/go.mod h1:FL3pVXie73rg3Rii6V/u5BoHlSoyeZeIgKZEgHARyCU= -modernc.org/strutil v1.2.0 h1:agBi9dp1I+eOnxXeiZawM8F4LawKv4NzGWSaLfyeNZA= -modernc.org/strutil v1.2.0/go.mod h1:/mdcBmfOibveCTBxUl5B5l6W+TTH1FXPLHZE6bTosX0= +modernc.org/cc/v4 v4.28.1 h1:XpLbkYVQ24E8tX5u8+yWGvaxerxkR/S4zqxI8ZoSBuc= +modernc.org/cc/v4 v4.28.1/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI= +modernc.org/ccgo/v4 v4.33.0 h1:dspBCm75jsj8Y/ufwAMVfe375L2iYdMyQ2QG/v3hL54= +modernc.org/ccgo/v4 v4.33.0/go.mod h1:+RhXBoRYzRwaH21mV/aj6XvQRDtfjcZfAlPMsQo8CR0= +modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM= +modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU= +modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI= +modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito= +modernc.org/gc/v3 v3.1.2 h1:ZtDCnhonXSZexk/AYsegNRV1lJGgaNZJuKjJSWKyEqo= +modernc.org/gc/v3 v3.1.2/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= +modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= +modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= +modernc.org/libc v1.72.1 h1:db1xwJ6u1kE3KHTFTTbe2GCrczHPKzlURP0aDC4NGD0= +modernc.org/libc v1.72.1/go.mod h1:HRMiC/PhPGLIPM7GzAFCbI+oSgE3dhZ8FWftmRrHVlY= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg= +modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= +modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= +modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= +modernc.org/sqlite v1.49.1 h1:dYGHTKcX1sJ+EQDnUzvz4TJ5GbuvhNJa8Fg6ElGx73U= +modernc.org/sqlite v1.49.1/go.mod h1:m0w8xhwYUVY3H6pSDwc3gkJ/irZT/0YEXwBlhaxQEew= +modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= +modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM= diff --git a/Wavelet/internal/apps/admin/updater/export.go b/Wavelet/internal/apps/admin/updater/export.go new file mode 100644 index 00000000..ee7999d6 --- /dev/null +++ b/Wavelet/internal/apps/admin/updater/export.go @@ -0,0 +1,40 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package updater + +import "context" + +// GetStatus returns the current build and the newest compatible upstream release. +func GetStatus(ctx context.Context) (Status, error) { + status, _, err := defaultManager.status(ctx) + return status, err +} + +// PrepareUpgrade downloads and stages the upgrade binary for the current platform. +func PrepareUpgrade(ctx context.Context) (executable string, stagedBinary string, status Status, err error) { + status, _, err = defaultManager.status(ctx) + if err != nil { + return "", "", Status{}, err + } + + executable, stagedBinary, err = defaultManager.prepareUpgrade(ctx) + return executable, stagedBinary, status, err +} + +// ApplyPreparedUpgrade replaces the running binary and restarts the process. +func ApplyPreparedUpgrade(executable, stagedBinary string) error { + return replaceAndRestart(executable, stagedBinary) +} + +// FinishUpgrade clears the in-progress upgrade flag after a failed restart. +func FinishUpgrade() { + defaultManager.finishUpgrade() +} + +// IsUpgrading reports whether an upgrade task is currently running. +func IsUpgrading() bool { + defaultManager.mu.Lock() + defer defaultManager.mu.Unlock() + return defaultManager.upgrading +} diff --git a/Wavelet/internal/apps/openflare/agent/auth_cache.go b/Wavelet/internal/apps/openflare/agent/auth_cache.go new file mode 100644 index 00000000..758a060a --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/auth_cache.go @@ -0,0 +1,143 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "errors" + "strings" + "sync" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +const ( + agentTokenPositiveCacheTTL = 2 * time.Minute + agentTokenNegativeCacheTTL = 10 * time.Minute +) + +type cachedAgentNode struct { + node *model.OpenFlareNode + expiresAt time.Time +} + +type accessTokenAuthCache struct { + mu sync.RWMutex + positive map[string]cachedAgentNode + negative map[string]time.Time + now func() time.Time + loadNodeByToken func(context.Context, string) (*model.OpenFlareNode, error) +} + +var tokenCache = newAccessTokenAuthCache() + +func newAccessTokenAuthCache() *accessTokenAuthCache { + return &accessTokenAuthCache{ + positive: make(map[string]cachedAgentNode), + negative: make(map[string]time.Time), + now: time.Now, + loadNodeByToken: func(ctx context.Context, token string) (*model.OpenFlareNode, error) { + return model.GetOpenFlareNodeByAccessToken(ctx, token) + }, + } +} + +func (c *accessTokenAuthCache) authenticate(ctx context.Context, token string) (*model.OpenFlareNode, error) { + now := c.now() + if node, ok := c.getNode(token, now); ok { + return node, nil + } + if c.isMissing(token, now) { + return nil, gorm.ErrRecordNotFound + } + + node, err := c.loadNodeByToken(ctx, token) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL)) + } + return nil, err + } + + c.storeNode(token, node) + return cloneNode(node), nil +} + +func (c *accessTokenAuthCache) getNode(token string, now time.Time) (*model.OpenFlareNode, bool) { + c.mu.RLock() + entry, ok := c.positive[token] + c.mu.RUnlock() + if !ok { + return nil, false + } + if now.After(entry.expiresAt) { + c.mu.Lock() + delete(c.positive, token) + c.mu.Unlock() + return nil, false + } + return cloneNode(entry.node), true +} + +func (c *accessTokenAuthCache) isMissing(token string, now time.Time) bool { + c.mu.RLock() + expiresAt, ok := c.negative[token] + c.mu.RUnlock() + if !ok { + return false + } + if now.After(expiresAt) { + c.mu.Lock() + delete(c.negative, token) + c.mu.Unlock() + return false + } + return true +} + +func (c *accessTokenAuthCache) storeNode(token string, node *model.OpenFlareNode) { + if token == "" || node == nil { + return + } + c.mu.Lock() + defer c.mu.Unlock() + delete(c.negative, token) + c.positive[token] = cachedAgentNode{ + node: cloneNode(node), + expiresAt: c.now().Add(agentTokenPositiveCacheTTL), + } +} + +func (c *accessTokenAuthCache) storeMissing(token string, expiresAt time.Time) { + if token == "" { + return + } + c.mu.Lock() + defer c.mu.Unlock() + delete(c.positive, token) + c.negative[token] = expiresAt +} + +func (c *accessTokenAuthCache) reset() { + c.mu.Lock() + defer c.mu.Unlock() + c.positive = make(map[string]cachedAgentNode) + c.negative = make(map[string]time.Time) +} + +// ResetAuthCacheForTest clears the in-memory access token cache for integration tests. +func ResetAuthCacheForTest() { + tokenCache.reset() +} + +// AuthenticateAccessToken validates X-Agent-Token against of_nodes.access_token. +func AuthenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) { + token = strings.TrimSpace(token) + if token == "" { + return nil, errors.New(errMissingAgentToken) + } + return tokenCache.authenticate(ctx, token) +} diff --git a/Wavelet/internal/apps/openflare/agent/config.go b/Wavelet/internal/apps/openflare/agent/config.go new file mode 100644 index 00000000..e8470356 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/config.go @@ -0,0 +1,98 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "encoding/json" + "errors" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +type configVersionRecord struct { + ID uint `gorm:"primaryKey"` + Version string `gorm:"column:version"` + SnapshotJSON string `gorm:"column:snapshot_json"` + SupportFilesJSON string `gorm:"column:support_files_json"` + Checksum string `gorm:"column:checksum"` + IsActive bool `gorm:"column:is_active"` + CreatedAt time.Time `gorm:"column:created_at"` +} + +func (configVersionRecord) TableName() string { + return "of_config_versions" +} + +func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { + version, err := loadActiveConfigVersion(ctx) + if err != nil { + return nil, err + } + return &ActiveConfigMeta{ + Version: version.Version, + Checksum: version.Checksum, + }, nil +} + +func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) { + version, err := loadActiveConfigVersion(ctx) + if err != nil { + return nil, err + } + + var supportFiles []SupportFile + if strings.TrimSpace(version.SupportFilesJSON) != "" { + if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil { + return nil, err + } + } + + return &ConfigResponse{ + Version: version.Version, + Checksum: version.Checksum, + SourceConfigJSON: version.SnapshotJSON, + SupportFiles: sourceSupportFiles(supportFiles), + CreatedAt: version.CreatedAt, + }, nil +} + +func loadActiveConfigVersion(ctx context.Context) (*configVersionRecord, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New("database not initialized") + } + version := &configVersionRecord{} + err := conn.Where("is_active = ?", true).Order("id desc").First(version).Error + if err != nil { + return nil, err + } + return version, nil +} + +func sourceSupportFiles(files []SupportFile) []SupportFile { + if len(files) == 0 { + return nil + } + result := make([]SupportFile, 0, len(files)) + for _, file := range files { + if isRuntimeGeneratedSupportFile(file.Path) { + continue + } + result = append(result, file) + } + return result +} + +func isRuntimeGeneratedSupportFile(path string) bool { + path = strings.TrimSpace(path) + return strings.HasPrefix(path, "runtime/") +} + +func isActiveConfigNotFound(err error) bool { + return errors.Is(err, gorm.ErrRecordNotFound) +} diff --git a/Wavelet/internal/apps/openflare/agent/errs.go b/Wavelet/internal/apps/openflare/agent/errs.go new file mode 100644 index 00000000..efb2a961 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/errs.go @@ -0,0 +1,21 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +const ( + errMissingAgentToken = "缺少 Agent Token" + errInvalidAgentToken = "无权进行此操作,Agent Token 无效" + errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效" + errNodeMissingFromContext = "Node object missing from context" + errNoActiveConfig = "当前没有激活版本" + errNodeNotFound = "节点不存在" + errNodeIDRequired = "node_id 不能为空" + errVersionRequired = "version 不能为空" + errInvalidApplyResult = "result 仅支持 success、warning 或 failed" + errIPRequired = "ip 不能为空" + errIPInvalid = "ip 格式无效" + errAgentVersionRequired = "version 不能为空" + errNodeIDConflict = "节点标识生成冲突,请重试" + errPagesPackageNotFound = "Pages 部署包尚未实现" +) diff --git a/Wavelet/internal/apps/openflare/agent/helpers.go b/Wavelet/internal/apps/openflare/agent/helpers.go new file mode 100644 index 00000000..54dd0138 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/helpers.go @@ -0,0 +1,257 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "crypto/rand" + "encoding/hex" + "net" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + openrestyStatusHealthy = "healthy" + openrestyStatusUnhealthy = "unhealthy" + openrestyStatusUnknown = "unknown" + releaseChannelStable = "stable" +) + +func newRandomToken() (string, error) { + buf := make([]byte, 16) + if _, err := rand.Read(buf); err != nil { + return "", err + } + return hex.EncodeToString(buf), nil +} + +func newServerNodeID() (string, error) { + token, err := newRandomToken() + if err != nil { + return "", err + } + return "node-" + token, nil +} + +func normalizeOpenrestyStatus(status string) string { + switch strings.ToLower(strings.TrimSpace(status)) { + case openrestyStatusHealthy: + return openrestyStatusHealthy + case openrestyStatusUnhealthy: + return openrestyStatusUnhealthy + default: + return openrestyStatusUnknown + } +} + +func normalizeNodePayload(payload NodePayload) NodePayload { + payload.Name = strings.TrimSpace(payload.Name) + payload.IP = strings.TrimSpace(payload.IP) + payload.Version = strings.TrimSpace(payload.Version) + payload.ExtVersion = strings.TrimSpace(payload.ExtVersion) + payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion) + payload.LastError = truncateForDatabase(payload.LastError, 16000) + payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus) + payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000) + return payload +} + +func validateNodePayload(payload NodePayload) error { + if payload.IP == "" { + return errPayload(errIPRequired) + } + if net.ParseIP(payload.IP) == nil { + return errPayload(errIPInvalid) + } + if payload.Version == "" { + return errPayload(errAgentVersionRequired) + } + return nil +} + +type payloadError string + +func (e payloadError) Error() string { return string(e) } + +func errPayload(message string) error { return payloadError(message) } + +func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveName bool) { + if !preserveName || strings.TrimSpace(node.Name) == "" { + if strings.TrimSpace(payload.Name) != "" { + node.Name = strings.TrimSpace(payload.Name) + } + } + if !node.IPManualOverride { + node.IP = strings.TrimSpace(payload.IP) + } + node.Version = strings.TrimSpace(payload.Version) + node.ExtVersion = strings.TrimSpace(payload.ExtVersion) + node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus) + node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000) + node.Status = nodeStatusOnline + node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion) + now := time.Now() + node.LastSeenAt = &now + node.LastError = truncateForDatabase(payload.LastError, 16000) +} + +func truncateForDatabase(value string, max int) string { + if max <= 0 { + return "" + } + runes := []rune(strings.TrimSpace(value)) + if len(runes) <= max { + return string(runes) + } + return string(runes[:max]) +} + +func resolveReportedNodeIP(reportedIP string, remoteAddr string) string { + reported := normalizeIP(reportedIP) + remote := normalizeRemoteAddr(remoteAddr) + if reported == "" { + return remote + } + if isPublicNodeIP(reported) { + return reported + } + if isPublicNodeIP(remote) { + return remote + } + return reported +} + +func normalizeIP(raw string) string { + raw = strings.TrimSpace(raw) + if raw == "" { + return "" + } + host := raw + if strings.Contains(raw, ":") { + if h, _, err := net.SplitHostPort(raw); err == nil { + host = h + } + } + host = strings.TrimPrefix(host, "[") + host = strings.TrimSuffix(host, "]") + if ip := net.ParseIP(host); ip != nil { + return ip.String() + } + return "" +} + +func normalizeRemoteAddr(remoteAddr string) string { + remoteAddr = strings.TrimSpace(remoteAddr) + if remoteAddr == "" { + return "" + } + host, _, err := net.SplitHostPort(remoteAddr) + if err != nil { + return normalizeIP(remoteAddr) + } + return normalizeIP(host) +} + +func isPublicNodeIP(raw string) bool { + ip := net.ParseIP(strings.TrimSpace(raw)) + if ip == nil { + return false + } + if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() { + return false + } + return true +} + +func buildAgentSettings(node *model.OpenFlareNode, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *Settings { + autoUpdate := false + if node != nil { + autoUpdate = node.AutoUpdateEnabled + } + if strings.TrimSpace(updateChannel) == "" { + updateChannel = releaseChannelStable + } + return &Settings{ + HeartbeatInterval: model.AgentHeartbeatInterval, + WebsocketUpgradeEnabled: model.AgentWebsocketUpgradeEnabled, + AutoUpdate: autoUpdate, + UpdateRepo: model.AgentUpdateRepo, + UpdateNow: updateNow, + UpdateChannel: updateChannel, + UpdateTag: strings.TrimSpace(updateTag), + RestartOpenrestyNow: restartOpenrestyNow, + } +} + +func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenFlareNode) map[string]any { + if previous == nil || current == nil { + return map[string]any{} + } + changes := make(map[string]any) + appendIfChanged := func(key string, before any, after any) { + if before != after { + changes[key] = after + } + } + appendIfChanged("name", previous.Name, current.Name) + appendIfChanged("ip", previous.IP, current.IP) + appendIfChanged("version", previous.Version, current.Version) + appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion) + appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus) + appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage) + appendIfChanged("status", previous.Status, current.Status) + appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion) + appendIfChanged("last_error", previous.LastError, current.LastError) + appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested) + appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel) + appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag) + appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested) + if !lastSeenAtEqual(previous.LastSeenAt, current.LastSeenAt) { + changes["last_seen_at"] = current.LastSeenAt + } + return changes +} + +func lastSeenAtEqual(before *time.Time, after *time.Time) bool { + if before == nil || after == nil { + return before == after + } + return before.Equal(*after) +} + +func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload { + payload.NodeID = strings.TrimSpace(payload.NodeID) + payload.Version = strings.TrimSpace(payload.Version) + payload.Result = strings.ToLower(strings.TrimSpace(payload.Result)) + payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), 16000) + payload.Checksum = strings.TrimSpace(payload.Checksum) + payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum) + payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum) + return payload +} + +func isUniqueConstraintError(err error) bool { + if err == nil { + return false + } + return strings.Contains(strings.ToLower(err.Error()), "unique") +} + +func refreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) { + if node == nil { + return + } + tokenCache.storeNode(node.AccessToken, cloneNode(node)) +} + +func cloneNode(node *model.OpenFlareNode) *model.OpenFlareNode { + if node == nil { + return nil + } + cloned := *node + return &cloned +} diff --git a/Wavelet/internal/apps/openflare/agent/logics.go b/Wavelet/internal/apps/openflare/agent/logics.go new file mode 100644 index 00000000..843d08df --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/logics.go @@ -0,0 +1,209 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/node" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// RegisterWithAccessToken registers an agent on a reserved node token. +func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*RegistrationResponse, error) { + payload = normalizeNodePayload(payload) + if authNode == nil { + return nil, errors.New(errNodeNotFound) + } + if err := validateNodePayload(payload); err != nil { + return nil, err + } + applyNodeRuntime(authNode, payload, true) + if err := model.SaveOpenFlareNode(ctx, authNode); err != nil { + return nil, err + } + refreshAccessTokenCache(ctx, authNode) + return &RegistrationResponse{ + NodeID: authNode.NodeID, + AccessToken: authNode.AccessToken, + Name: authNode.Name, + }, nil +} + +// RegisterWithDiscovery registers a new node using the global discovery token. +func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*RegistrationResponse, error) { + payload = normalizeNodePayload(payload) + if err := validateNodePayload(payload); err != nil { + return nil, err + } + + nodeID, err := newServerNodeID() + if err != nil { + return nil, err + } + accessToken, err := newRandomToken() + if err != nil { + return nil, err + } + + nodeName := payload.Name + if nodeName == "" { + nodeName = nodeID + } + + record := &model.OpenFlareNode{ + NodeID: nodeID, + Name: nodeName, + AccessToken: accessToken, + Status: nodeStatusOnline, + NodeType: "edge_node", + CapabilitiesJSON: "[]", + UpdateChannel: releaseChannelStable, + } + applyNodeRuntime(record, payload, false) + + if err = model.CreateOpenFlareNode(ctx, record); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errNodeIDConflict) + } + return nil, err + } + refreshAccessTokenCache(ctx, record) + return &RegistrationResponse{ + NodeID: record.NodeID, + AccessToken: record.AccessToken, + Name: record.Name, + }, nil +} + +// HeartbeatNode updates runtime state and returns agent settings. +func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*HeartbeatResponse, error) { + if authNode == nil { + return nil, errors.New(errNodeNotFound) + } + payload.NodeID = authNode.NodeID + payload = normalizeNodePayload(payload) + if err := validateNodePayload(payload); err != nil { + return nil, err + } + + previous := *authNode + updateNow := authNode.UpdateRequested + restartOpenrestyNow := authNode.RestartOpenrestyRequested + updateChannel := strings.TrimSpace(authNode.UpdateChannel) + updateTag := strings.TrimSpace(authNode.UpdateTag) + + applyNodeRuntime(authNode, payload, true) + authNode.UpdateRequested = false + authNode.UpdateChannel = releaseChannelStable + authNode.UpdateTag = "" + authNode.RestartOpenrestyRequested = false + + changes := collectHeartbeatChanges(&previous, authNode) + if len(changes) > 0 { + fields := make([]string, 0, len(changes)) + for field := range changes { + fields = append(fields, field) + } + if err := model.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil { + return nil, err + } + } + + refreshAccessTokenCache(ctx, authNode) + + activeConfig, err := getActiveConfigMeta(ctx) + if err != nil && !isActiveConfigNotFound(err) { + return nil, err + } + + return &HeartbeatResponse{ + Node: authNode, + AgentSettings: buildAgentSettings(authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow), + ActiveConfig: activeConfig, + WAFIPGroups: nil, + }, nil +} + +// GetActiveConfig returns the active configuration for an agent. +func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) { + config, err := getActiveConfigForAgent(ctx) + if err != nil { + if isActiveConfigNotFound(err) { + return nil, errors.New(errNoActiveConfig) + } + return nil, err + } + return config, nil +} + +// SyncWAFIPGroups is a stub until full WAF agent sync is migrated. +func SyncWAFIPGroups(_ context.Context, _ WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) { + return &WAFIPGroupSyncResult{Groups: []WAFIPGroup{}}, nil +} + +// ReportApplyLog records an agent apply result. +func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) { + now := time.Now() + payload = normalizeApplyLogPayload(payload) + if payload.NodeID == "" { + return nil, errors.New(errNodeIDRequired) + } + if payload.Version == "" { + return nil, errors.New(errVersionRequired) + } + if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFailed { + return nil, errors.New(errInvalidApplyResult) + } + + log := &model.OpenFlareApplyLog{ + NodeID: payload.NodeID, + Version: payload.Version, + Result: payload.Result, + Message: payload.Message, + Checksum: payload.Checksum, + MainConfigChecksum: payload.MainConfigChecksum, + RouteConfigChecksum: payload.RouteConfigChecksum, + SupportFileCount: payload.SupportFileCount, + CreatedAt: now, + } + + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New("database not initialized") + } + + err := conn.Transaction(func(tx *gorm.DB) error { + record := &model.OpenFlareNode{} + if err := tx.Where("node_id = ?", payload.NodeID).First(record).Error; err != nil { + return err + } + record.Status = nodeStatusOnline + record.LastSeenAt = &now + if payload.Result == applyResultOK { + record.CurrentVersion = payload.Version + record.LastError = "" + } else { + record.LastError = payload.Message + } + if err := tx.Create(log).Error; err != nil { + return err + } + return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).Error + }) + if err != nil { + return nil, err + } + return log, nil +} + +// ValidateDiscoveryToken delegates to the node package discovery token helper. +func ValidateDiscoveryToken(ctx context.Context, token string) error { + return node.ValidateDiscoveryToken(ctx, token) +} diff --git a/Wavelet/internal/apps/openflare/agent/middleware.go b/Wavelet/internal/apps/openflare/agent/middleware.go new file mode 100644 index 00000000..ed5f6da8 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/middleware.go @@ -0,0 +1,61 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +const ( + agentTokenHeader = "X-Agent-Token" + agentNodeContextKey = "agent_node" +) + +// AgentAuth validates X-Agent-Token against of_nodes.access_token. +func AgentAuth() gin.HandlerFunc { + return func(c *gin.Context) { + token := strings.TrimSpace(c.GetHeader(agentTokenHeader)) + node, err := AuthenticateAccessToken(c.Request.Context(), token) + if err != nil { + compat.Unauthorized(c, errInvalidAgentToken) + c.Abort() + return + } + c.Set(agentNodeContextKey, node) + c.Next() + } +} + +// AgentRegisterAuth accepts either a node access token or the global discovery token. +func AgentRegisterAuth() gin.HandlerFunc { + return func(c *gin.Context) { + token := strings.TrimSpace(c.GetHeader(agentTokenHeader)) + if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil { + c.Set(agentNodeContextKey, node) + c.Next() + return + } + if err := ValidateDiscoveryToken(c.Request.Context(), token); err != nil { + compat.Unauthorized(c, errInvalidDiscoveryToken) + c.Abort() + return + } + c.Set("discovery_enabled", true) + c.Next() + } +} + +// AgentNodeFromContext returns the authenticated agent node. +func AgentNodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) { + value, ok := c.Get(agentNodeContextKey) + if !ok { + return nil, false + } + node, ok := value.(*model.OpenFlareNode) + return node, ok +} diff --git a/Wavelet/internal/apps/openflare/agent/middleware_test.go b/Wavelet/internal/apps/openflare/agent/middleware_test.go new file mode 100644 index 00000000..d4f3a573 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/middleware_test.go @@ -0,0 +1,209 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupAgentAuthTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.OpenFlareNode{}, + &model.OpenFlareOption{}, + )) + + db.SetDB(sqliteDB) + option.ResetInitializationForTest() + tokenCache.reset() + + return func() { + db.SetDB(nil) + option.ResetInitializationForTest() + tokenCache.reset() + } +} + +func TestAuthenticateAccessToken(t *testing.T) { + cleanup := setupAgentAuthTestDB(t) + defer cleanup() + + ctx := context.Background() + now := time.Now() + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + NodeID: "node-auth-1", + Name: "edge", + AccessToken: "valid-agent-token", + Status: nodeStatusOnline, + LastSeenAt: &now, + NodeType: "edge_node", + }).Error) + + t.Run("valid token", func(t *testing.T) { + node, err := AuthenticateAccessToken(ctx, "valid-agent-token") + require.NoError(t, err) + assert.Equal(t, "node-auth-1", node.NodeID) + }) + + t.Run("cached token", func(t *testing.T) { + originalLoader := tokenCache.loadNodeByToken + t.Cleanup(func() { + tokenCache.loadNodeByToken = originalLoader + }) + tokenCache.loadNodeByToken = func(context.Context, string) (*model.OpenFlareNode, error) { + t.Fatal("db should not be queried for cached token") + return nil, nil + } + node, err := AuthenticateAccessToken(ctx, "valid-agent-token") + require.NoError(t, err) + assert.Equal(t, "node-auth-1", node.NodeID) + }) + + t.Run("missing token", func(t *testing.T) { + _, err := AuthenticateAccessToken(ctx, "") + require.Error(t, err) + assert.Contains(t, err.Error(), errMissingAgentToken) + }) + + t.Run("invalid token", func(t *testing.T) { + _, err := AuthenticateAccessToken(ctx, "invalid-token") + require.Error(t, err) + }) +} + +func TestAgentAuthMiddleware(t *testing.T) { + cleanup := setupAgentAuthTestDB(t) + defer cleanup() + + ctx := context.Background() + now := time.Now() + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + NodeID: "node-mw-1", + Name: "edge", + AccessToken: "middleware-token", + Status: nodeStatusOnline, + LastSeenAt: &now, + NodeType: "edge_node", + }).Error) + + gin.SetMode(gin.TestMode) + router := gin.New() + router.GET("/protected", AgentAuth(), func(c *gin.Context) { + node, ok := AgentNodeFromContext(c) + if !ok { + c.Status(http.StatusInternalServerError) + return + } + compat.OK(c, gin.H{"node_id": node.NodeID}) + }) + + t.Run("authorized request", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/protected", nil) + req.Header.Set(agentTokenHeader, "middleware-token") + resp := httptest.NewRecorder() + router.ServeHTTP(resp, req) + + assert.Equal(t, http.StatusOK, resp.Code) + var envelope compat.Envelope + require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &envelope)) + assert.True(t, envelope.Success) + }) + + t.Run("unauthorized request", func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/protected", nil) + req.Header.Set(agentTokenHeader, "bad-token") + resp := httptest.NewRecorder() + router.ServeHTTP(resp, req) + + assert.Equal(t, http.StatusUnauthorized, resp.Code) + }) +} + +func TestAgentRegisterAuthMiddleware(t *testing.T) { + cleanup := setupAgentAuthTestDB(t) + defer cleanup() + + ctx := context.Background() + now := time.Now() + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + NodeID: "node-register-1", + Name: "edge", + AccessToken: "existing-node-token", + Status: nodeStatusOnline, + LastSeenAt: &now, + NodeType: "edge_node", + }).Error) + require.NoError(t, model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", "discovery-token")) + + gin.SetMode(gin.TestMode) + router := gin.New() + router.POST("/register", AgentRegisterAuth(), func(c *gin.Context) { + if node, ok := AgentNodeFromContext(c); ok { + compat.OK(c, gin.H{"mode": "node", "node_id": node.NodeID}) + return + } + if _, ok := c.Get("discovery_enabled"); ok { + compat.OK(c, gin.H{"mode": "discovery"}) + return + } + c.Status(http.StatusInternalServerError) + }) + + t.Run("existing node token", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/register", nil) + req.Header.Set(agentTokenHeader, "existing-node-token") + resp := httptest.NewRecorder() + router.ServeHTTP(resp, req) + + assert.Equal(t, http.StatusOK, resp.Code) + var envelope compat.Envelope + require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &envelope)) + data, ok := envelope.Data.(map[string]any) + require.True(t, ok) + assert.Equal(t, "node", data["mode"]) + }) + + t.Run("discovery token", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/register", nil) + req.Header.Set(agentTokenHeader, "discovery-token") + resp := httptest.NewRecorder() + router.ServeHTTP(resp, req) + + assert.Equal(t, http.StatusOK, resp.Code) + var envelope compat.Envelope + require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &envelope)) + data, ok := envelope.Data.(map[string]any) + require.True(t, ok) + assert.Equal(t, "discovery", data["mode"]) + }) + + t.Run("invalid token", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/register", nil) + req.Header.Set(agentTokenHeader, "invalid-token") + resp := httptest.NewRecorder() + router.ServeHTTP(resp, req) + + assert.Equal(t, http.StatusUnauthorized, resp.Code) + }) +} diff --git a/Wavelet/internal/apps/openflare/agent/routers.go b/Wavelet/internal/apps/openflare/agent/routers.go new file mode 100644 index 00000000..5e2836d4 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/routers.go @@ -0,0 +1,161 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "net/http" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" + "github.com/gin-gonic/gin" +) + +// RegisterRoutes mounts agent API routes under /agent. +func RegisterRoutes(apiGroup *gin.RouterGroup) { + agentRoute := apiGroup.Group("/agent") + { + discoveryRoute := agentRoute.Group("/") + discoveryRoute.Use(AgentRegisterAuth()) + { + discoveryRoute.POST("/nodes/register", RegisterHandler) + } + + authorizedRoute := agentRoute.Group("/") + authorizedRoute.Use(AgentAuth()) + { + authorizedRoute.GET("/ws", AgentWebSocketHandler) + authorizedRoute.POST("/nodes/heartbeat", HeartbeatHandler) + authorizedRoute.GET("/config-versions/active", GetActiveConfigHandler) + authorizedRoute.GET("/pages/deployments/:deployment_id/package", DownloadPagesPackageHandler) + authorizedRoute.POST("/waf/ip-groups/sync", SyncWAFIPGroupsHandler) + authorizedRoute.POST("/apply-logs", ReportApplyLogHandler) + } + } +} + +// RegisterHandler registers or discovers an agent node. +func RegisterHandler(c *gin.Context) { + var payload NodePayload + if !compat.BindJSON(c, &payload) { + return + } + payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr) + + var ( + result *RegistrationResponse + err error + ) + if authNode, ok := AgentNodeFromContext(c); ok { + result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload) + } else { + result, err = RegisterWithDiscovery(c.Request.Context(), payload) + } + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +// HeartbeatHandler records agent heartbeat state. +func HeartbeatHandler(c *gin.Context) { + var payload NodePayload + if !compat.BindJSON(c, &payload) { + return + } + payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr) + + authNode, ok := AgentNodeFromContext(c) + if !ok { + compat.Unauthorized(c, errInvalidAgentToken) + return + } + + response, err := HeartbeatNode(c.Request.Context(), authNode, payload) + if err != nil { + compat.Fail(c, err.Error()) + return + } + okWithExtras(c, response.Node, gin.H{ + "agent_settings": response.AgentSettings, + "active_config": response.ActiveConfig, + "waf_ip_groups": response.WAFIPGroups, + }) +} + +// GetActiveConfigHandler returns the active configuration version. +func GetActiveConfigHandler(c *gin.Context) { + if _, ok := AgentNodeFromContext(c); !ok { + compat.Unauthorized(c, errNodeMissingFromContext) + return + } + config, err := GetActiveConfig(c.Request.Context()) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, config) +} + +// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent (stub). +func SyncWAFIPGroupsHandler(c *gin.Context) { + var input WAFIPGroupSyncInput + if !compat.BindJSON(c, &input) { + return + } + result, err := SyncWAFIPGroups(c.Request.Context(), input) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +// ReportApplyLogHandler records an agent apply log entry. +func ReportApplyLogHandler(c *gin.Context) { + var payload ApplyLogPayload + if !compat.BindJSON(c, &payload) { + return + } + if authNode, ok := AgentNodeFromContext(c); ok { + payload.NodeID = authNode.NodeID + } + log, err := ReportApplyLog(c.Request.Context(), payload) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, log) +} + +// DownloadPagesPackageHandler is a stub until Pages agent packaging is migrated. +func DownloadPagesPackageHandler(c *gin.Context) { + c.JSON(http.StatusNotFound, compat.Envelope{ + Success: false, + Message: errPagesPackageNotFound, + Data: nil, + }) +} + +// AgentWebSocketHandler upgrades an authenticated agent websocket connection. +func AgentWebSocketHandler(c *gin.Context) { + authNode, ok := AgentNodeFromContext(c) + if !ok { + compat.Unauthorized(c, errInvalidAgentToken) + return + } + websocket.ServeAgent(c, authNode.NodeID) +} + +func okWithExtras(c *gin.Context, data any, extras gin.H) { + payload := gin.H{ + "success": true, + "message": "", + "data": data, + } + for key, value := range extras { + payload[key] = value + } + c.JSON(http.StatusOK, payload) +} diff --git a/Wavelet/internal/apps/openflare/agent/types.go b/Wavelet/internal/apps/openflare/agent/types.go new file mode 100644 index 00000000..e5a52f50 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/types.go @@ -0,0 +1,112 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + nodeStatusOnline = "online" + applyResultOK = "success" + applyResultWarn = "warning" + applyResultFailed = "failed" +) + +// NodePayload is the agent register/heartbeat payload. +type NodePayload struct { + NodeID string `json:"node_id"` + Name string `json:"name"` + IP string `json:"ip"` + Version string `json:"version"` + ExtVersion string `json:"ext_version"` + CurrentVersion string `json:"current_version"` + LastError string `json:"last_error"` + OpenrestyStatus string `json:"openresty_status"` + OpenrestyMessage string `json:"openresty_message"` + WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"` +} + +// ApplyLogPayload is the agent apply log report payload. +type ApplyLogPayload struct { + NodeID string `json:"node_id"` + Version string `json:"version"` + Result string `json:"result"` + Message string `json:"message"` + Checksum string `json:"checksum"` + MainConfigChecksum string `json:"main_config_checksum"` + RouteConfigChecksum string `json:"route_config_checksum"` + SupportFileCount int `json:"support_file_count"` +} + +// RegistrationResponse is returned after agent registration. +type RegistrationResponse struct { + NodeID string `json:"node_id"` + AccessToken string `json:"access_token"` + Name string `json:"name"` +} + +// Settings carries remote agent control flags. +type Settings struct { + HeartbeatInterval int `json:"heartbeat_interval"` + WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"` + AutoUpdate bool `json:"auto_update"` + UpdateRepo string `json:"update_repo"` + UpdateNow bool `json:"update_now"` + UpdateChannel string `json:"update_channel"` + UpdateTag string `json:"update_tag"` + RestartOpenrestyNow bool `json:"restart_openresty_now"` +} + +// ActiveConfigMeta summarizes the active configuration version. +type ActiveConfigMeta struct { + Version string `json:"version"` + Checksum string `json:"checksum"` +} + +// SupportFile is a configuration support artifact shipped to agents. +type SupportFile struct { + Path string `json:"path"` + Content string `json:"content"` +} + +// ConfigResponse is the full active config payload for agents. +type ConfigResponse struct { + Version string `json:"version"` + Checksum string `json:"checksum"` + SourceConfigJSON string `json:"source_config_json"` + SupportFiles []SupportFile `json:"support_files"` + CreatedAt time.Time `json:"created_at"` +} + +// WAFIPGroup is a WAF IP group snapshot for agents. +type WAFIPGroup struct { + ID uint `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + Enabled bool `json:"enabled"` + IPList []string `json:"ip_list"` + Checksum string `json:"checksum"` +} + +// WAFIPGroupSyncInput requests changed WAF IP groups. +type WAFIPGroupSyncInput struct { + IDs []uint `json:"ids"` + Checksums map[string]string `json:"checksums"` +} + +// WAFIPGroupSyncResult returns synced WAF IP groups. +type WAFIPGroupSyncResult struct { + Groups []WAFIPGroup `json:"groups"` +} + +// HeartbeatResponse is the heartbeat handler result. +type HeartbeatResponse struct { + Node *model.OpenFlareNode `json:"node"` + AgentSettings *Settings `json:"agent_settings"` + ActiveConfig *ActiveConfigMeta `json:"active_config"` + WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"` +} diff --git a/Wavelet/internal/apps/openflare/apply_log/errs.go b/Wavelet/internal/apps/openflare/apply_log/errs.go new file mode 100644 index 00000000..914b3635 --- /dev/null +++ b/Wavelet/internal/apps/openflare/apply_log/errs.go @@ -0,0 +1,8 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package apply_log + +const ( + errRetentionDaysOutOfRange = "retention_days 必须在 1 到 3650 之间" +) diff --git a/Wavelet/internal/apps/openflare/apply_log/logics.go b/Wavelet/internal/apps/openflare/apply_log/logics.go new file mode 100644 index 00000000..f61dbce7 --- /dev/null +++ b/Wavelet/internal/apps/openflare/apply_log/logics.go @@ -0,0 +1,128 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package apply_log + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + defaultApplyLogPageSize = 20 + maxApplyLogPageSize = 200 + maxApplyLogRetentionDays = 3650 +) + +// ListQuery filters apply logs for paginated listing. +type ListQuery struct { + NodeID string `json:"node_id"` + PageNo int `json:"pageNo"` + PageSize int `json:"pageSize"` +} + +// ListResult is the paginated apply log list response. +type ListResult struct { + Rows []*model.OpenFlareApplyLog `json:"rows"` + Current int `json:"current"` + Total int `json:"total"` + TotalPage int `json:"totalPage"` +} + +// CleanupInput controls apply log cleanup behavior. +type CleanupInput struct { + DeleteAll bool `json:"delete_all"` + RetentionDays int `json:"retention_days"` +} + +// CleanupResult reports apply log cleanup outcome. +type CleanupResult struct { + DeleteAll bool `json:"delete_all"` + RetentionDays int `json:"retention_days"` + DeletedCount int64 `json:"deleted_count"` + Cutoff *time.Time `json:"cutoff,omitempty"` +} + +// ListPage returns paginated apply logs with optional node_id filter. +func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) { + pageNo := normalizePageNo(input.PageNo) + pageSize := normalizePageSize(input.PageSize) + nodeID := strings.TrimSpace(input.NodeID) + + rows, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ + NodeID: nodeID, + PageNo: pageNo, + PageSize: pageSize, + }) + if err != nil { + return nil, err + } + + total, err := model.CountOpenFlareApplyLogs(ctx, nodeID) + if err != nil { + return nil, err + } + + totalPage := 0 + if total > 0 { + totalPage = int((total + int64(pageSize) - 1) / int64(pageSize)) + } + + return &ListResult{ + Rows: rows, + Current: pageNo, + Total: int(total), + TotalPage: totalPage, + }, nil +} + +// Cleanup removes old apply logs or deletes all records. +func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) { + if input.DeleteAll { + deleted, err := model.DeleteAllOpenFlareApplyLogs(ctx) + if err != nil { + return nil, err + } + return &CleanupResult{ + DeleteAll: true, + DeletedCount: deleted, + }, nil + } + + if input.RetentionDays <= 0 || input.RetentionDays > maxApplyLogRetentionDays { + return nil, errors.New(errRetentionDaysOutOfRange) + } + + cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour) + deleted, err := model.DeleteOpenFlareApplyLogsBefore(ctx, cutoff) + if err != nil { + return nil, err + } + + return &CleanupResult{ + RetentionDays: input.RetentionDays, + DeletedCount: deleted, + Cutoff: &cutoff, + }, nil +} + +func normalizePageNo(pageNo int) int { + if pageNo <= 0 { + return 1 + } + return pageNo +} + +func normalizePageSize(pageSize int) int { + if pageSize <= 0 { + return defaultApplyLogPageSize + } + if pageSize > maxApplyLogPageSize { + return maxApplyLogPageSize + } + return pageSize +} diff --git a/Wavelet/internal/apps/openflare/apply_log/logics_test.go b/Wavelet/internal/apps/openflare/apply_log/logics_test.go new file mode 100644 index 00000000..6f14220a --- /dev/null +++ b/Wavelet/internal/apps/openflare/apply_log/logics_test.go @@ -0,0 +1,105 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package apply_log + +import ( + "context" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupApplyLogTestDB(t *testing.T) func() { + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + + err = sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{}) + require.NoError(t, err) + + db.SetDB(sqliteDB) + + return func() { + db.SetDB(nil) + } +} + +func TestListPageAndCleanup(t *testing.T) { + cleanup := setupApplyLogTestDB(t) + defer cleanup() + + ctx := context.Background() + now := time.Now().UTC() + + logs := []model.OpenFlareApplyLog{ + {NodeID: "node-logs", Version: "v1", Result: "success", Message: "1", CreatedAt: now.Add(-10 * 24 * time.Hour)}, + {NodeID: "node-logs", Version: "v2", Result: "success", Message: "2", CreatedAt: now.Add(-5 * 24 * time.Hour)}, + {NodeID: "node-logs", Version: "v3", Result: "success", Message: "3", CreatedAt: now}, + } + for i := range logs { + require.NoError(t, db.DB(ctx).Create(&logs[i]).Error) + } + + pageResult, err := ListPage(ctx, ListQuery{ + NodeID: "node-logs", + PageNo: 1, + PageSize: 2, + }) + require.NoError(t, err) + assert.Equal(t, 3, pageResult.Total) + assert.Len(t, pageResult.Rows, 2) + assert.Equal(t, 2, pageResult.TotalPage) + assert.Equal(t, 1, pageResult.Current) + + cleanupResult, err := Cleanup(ctx, CleanupInput{ + DeleteAll: false, + RetentionDays: 7, + }) + require.NoError(t, err) + assert.Equal(t, int64(1), cleanupResult.DeletedCount) + assert.NotNil(t, cleanupResult.Cutoff) + + remaining, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ + NodeID: "node-logs", + PageNo: 1, + PageSize: 10, + }) + require.NoError(t, err) + assert.Len(t, remaining, 2) + + cleanupAll, err := Cleanup(ctx, CleanupInput{DeleteAll: true}) + require.NoError(t, err) + assert.Equal(t, int64(2), cleanupAll.DeletedCount) + assert.True(t, cleanupAll.DeleteAll) + + finalLogs, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ + NodeID: "node-logs", + PageNo: 1, + PageSize: 10, + }) + require.NoError(t, err) + assert.Empty(t, finalLogs) +} + +func TestCleanupInvalidRetentionDays(t *testing.T) { + cleanup := setupApplyLogTestDB(t) + defer cleanup() + + ctx := context.Background() + + _, err := Cleanup(ctx, CleanupInput{RetentionDays: 0}) + require.Error(t, err) + assert.Equal(t, errRetentionDaysOutOfRange, err.Error()) + + _, err = Cleanup(ctx, CleanupInput{RetentionDays: 4000}) + require.Error(t, err) + assert.Equal(t, errRetentionDaysOutOfRange, err.Error()) +} diff --git a/Wavelet/internal/apps/openflare/apply_log/routers.go b/Wavelet/internal/apps/openflare/apply_log/routers.go new file mode 100644 index 00000000..bca22f90 --- /dev/null +++ b/Wavelet/internal/apps/openflare/apply_log/routers.go @@ -0,0 +1,49 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package apply_log + +import ( + "strconv" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +// GetApplyLogs lists apply logs with pagination and optional node_id filter. +func GetApplyLogs(c *gin.Context) { + result, err := ListPage(c.Request.Context(), ListQuery{ + NodeID: c.Query("node_id"), + PageNo: readIntQuery(c, "pageNo", "page_no"), + PageSize: readIntQuery(c, "pageSize", "page_size"), + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +// CleanupApplyLogs removes old apply logs or deletes all records. +func CleanupApplyLogs(c *gin.Context) { + var input CleanupInput + if !compat.BindJSON(c, &input) { + return + } + + result, err := Cleanup(c.Request.Context(), input) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +func readIntQuery(c *gin.Context, primary, secondary string) int { + value := c.Query(primary) + if value == "" { + value = c.Query(secondary) + } + parsed, _ := strconv.Atoi(value) + return parsed +} diff --git a/Wavelet/internal/apps/openflare/auth/errs.go b/Wavelet/internal/apps/openflare/auth/errs.go new file mode 100644 index 00000000..1a8ace6f --- /dev/null +++ b/Wavelet/internal/apps/openflare/auth/errs.go @@ -0,0 +1,33 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +const ( + errInvalidParams = "无效的参数" + errPasswordLoginDisabled = "管理员关闭了密码登录" + errUsernameOrPasswordWrong = "用户名或密码错误" + errBannedAccount = "用户已被封禁" + errSaveSessionFailed = "无法保存会话信息,请重试" + errRegistrationDisabled = "管理员关闭了注册" + errPasswordTooShort = "密码长度不能少于 8 位" + errEmailRequired = "邮箱地址不能为空" + errEmailAlreadyRegistered = "邮箱地址已被占用" + errEmailNotRegistered = "该邮箱地址未注册" + errEmailCodeInvalid = "验证码错误或已过期" + errResetLinkInvalid = "重置链接非法或已过期" + errUserNotFound = "用户不存在" + errGenerateTokenFailed = "生成 Token 失败" + errInsufficientPermission = "无权进行此操作,权限不足" + errCannotDisableRoot = "无法禁用超级管理员用户" + errCannotDeleteRoot = "无法删除超级管理员用户" + errCannotPromoteAdmin = "普通管理员用户无法提升其他用户为管理员" + errAlreadyAdmin = "该用户已经是管理员" + errAlreadyCommonUser = "该用户已经是普通用户" + errAuthSourceDisabled = "认证源未启用" + errInvalidAuthSourceID = "认证源 ID 无效" + errPendingOAuthExpired = "待绑定第三方账号已失效,请重新登录" + errPendingOAuthInvalid = "待绑定第三方账号无效,请重新登录" + errCapTokenMissing = "缺少人机验证凭证" + errCapTokenInvalid = "人机验证凭证无效或已过期" +) diff --git a/Wavelet/internal/apps/openflare/auth/legacy_user.go b/Wavelet/internal/apps/openflare/auth/legacy_user.go new file mode 100644 index 00000000..cb092fce --- /dev/null +++ b/Wavelet/internal/apps/openflare/auth/legacy_user.go @@ -0,0 +1,82 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package auth provides OpenFlare legacy auth business logic. +package auth + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// LegacyUser mirrors the old OpenFlare frontend user shape. +type LegacyUser struct { + ID int `json:"id"` + Username string `json:"username"` + DisplayName string `json:"display_name"` + Role int `json:"role"` + Status int `json:"status"` + Token string `json:"token,omitempty"` + Email string `json:"email,omitempty"` +} + +const ( + legacyUserStatusEnabled = 1 + legacyUserStatusDisabled = 2 +) + +// RoleFromUser maps a Wavelet user to the legacy role value. +func RoleFromUser(user *model.User) int { + if user == nil { + return 0 + } + if user.IsAdmin { + return compat.RoleRootUser + } + return compat.RoleCommonUser +} + +// StatusFromUser maps is_active to legacy status. +func StatusFromUser(user *model.User) int { + if user == nil || !user.IsActive { + return legacyUserStatusDisabled + } + return legacyUserStatusEnabled +} + +// ToLegacyUser converts a Wavelet user to the legacy response shape. +func ToLegacyUser(user *model.User, token string) LegacyUser { + if user == nil { + return LegacyUser{} + } + return LegacyUser{ + ID: int(user.ID), + Username: user.Username, + DisplayName: displayName(user), + Role: RoleFromUser(user), + Status: StatusFromUser(user), + Token: token, + Email: user.Email, + } +} + +// ToLegacyUsers converts a slice of users. +func ToLegacyUsers(users []model.User) []LegacyUser { + result := make([]LegacyUser, 0, len(users)) + for i := range users { + result = append(result, ToLegacyUser(&users[i], "")) + } + return result +} + +func displayName(user *model.User) string { + if user.Nickname != "" { + return user.Nickname + } + return user.Username +} + +// IsAdminRole reports whether a legacy role has admin privileges. +func IsAdminRole(role int) bool { + return role >= compat.RoleAdminUser +} diff --git a/Wavelet/internal/apps/openflare/auth/logics.go b/Wavelet/internal/apps/openflare/auth/logics.go new file mode 100644 index 00000000..2c49df3e --- /dev/null +++ b/Wavelet/internal/apps/openflare/auth/logics.go @@ -0,0 +1,801 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "context" + "crypto/rand" + "encoding/json" + "errors" + "fmt" + "math/big" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/buildinfo" + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/db/idgen" + "github.com/Rain-kl/Wavelet/internal/listener" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/task" + "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +const ( + legacyTokenName = "openflare-legacy" + minPasswordLength = 8 + legacyItemsPerPage = 10 + passwordResetKeyFmt = "of_password_reset:%s" + passwordResetExpiry = 15 * time.Minute + verificationCodeRange = 900000 + verificationOffset = 100000 +) + +// LoginInput holds legacy login credentials. +type LoginInput struct { + Username string + Password string + Code string +} + +// RegisterInput holds legacy registration fields. +type RegisterInput struct { + Username string + Password string + Nickname string + DisplayName string + Email string + Code string +} + +// ManageUserInput holds legacy user management actions. +type ManageUserInput struct { + Username string + Action string +} + +// UpdateUserInput holds legacy admin user update fields. +type UpdateUserInput struct { + ID int + Username string + Password string + DisplayName string + Role int + Email string +} + +// UpdateSelfInput holds legacy self-update fields. +type UpdateSelfInput struct { + Username string + Password string + DisplayName string + Email string +} + +// CreateUserInput holds legacy admin create-user fields. +type CreateUserInput struct { + Username string + Password string + DisplayName string + Role int + Email string +} + +// PasswordResetInput holds password reset confirmation. +type PasswordResetInput struct { + Email string + Token string +} + +// LinkExistingInput binds a pending OAuth account to an existing user. +type LinkExistingInput struct { + Username string + Password string +} + +// Login authenticates a user and returns the legacy user shape with an access token. +func Login(ctx context.Context, c *gin.Context, input LoginInput) (LegacyUser, error) { + if !isPasswordLoginEnabled(ctx) { + return LegacyUser{}, errors.New(errPasswordLoginDisabled) + } + + input.Username = strings.TrimSpace(input.Username) + if input.Username == "" || input.Password == "" { + return LegacyUser{}, errors.New(errInvalidParams) + } + + var user model.User + if err := db.DB(ctx).Where("username = ? OR email = ?", input.Username, input.Username).First(&user).Error; err != nil { + logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", input.Username, c.ClientIP()) + return LegacyUser{}, errors.New(errUsernameOrPasswordWrong) + } + if !user.IsActive { + return LegacyUser{}, errors.New(errBannedAccount) + } + if !user.CheckPassword(input.Password) { + logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) + return LegacyUser{}, errors.New(errUsernameOrPasswordWrong) + } + + if isEmailLoginVerificationEnabled(ctx) { + if err := verifyLoginEmailCode(ctx, user.Email, input.Code); err != nil { + return LegacyUser{}, err + } + } + + user.LastLoginAt = time.Now() + if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil { + return LegacyUser{}, err + } + if err := setLoginSession(ctx, c, &user); err != nil { + return LegacyUser{}, errors.New(errSaveSessionFailed) + } + + token, err := issueLegacyAccessToken(ctx, &user) + if err != nil { + return LegacyUser{}, err + } + + logger.InfoF(ctx, "[LoginAudit] successful legacy login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) + listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP()) + + return ToLegacyUser(&user, token), nil +} + +// Register creates a new user and logs them in. +func Register(ctx context.Context, c *gin.Context, input RegisterInput) (LegacyUser, error) { + if !isRegistrationEnabled(ctx) || !isPasswordRegisterEnabled(ctx) { + return LegacyUser{}, errors.New(errRegistrationDisabled) + } + + input.Username = strings.TrimSpace(input.Username) + input.Password = strings.TrimSpace(input.Password) + input.Nickname = strings.TrimSpace(input.Nickname) + input.DisplayName = strings.TrimSpace(input.DisplayName) + input.Email = strings.TrimSpace(input.Email) + input.Code = strings.TrimSpace(input.Code) + + if input.Username == "" || input.Password == "" { + return LegacyUser{}, errors.New(errInvalidParams) + } + if len(input.Password) < minPasswordLength { + return LegacyUser{}, errors.New(errPasswordTooShort) + } + if input.Email == "" { + return LegacyUser{}, errors.New(errEmailRequired) + } + if isEmailRegisterVerificationEnabled(ctx) { + if input.Code == "" || !verifyEmailCode(ctx, input.Email, "register", input.Code) { + return LegacyUser{}, errors.New(errEmailCodeInvalid) + } + } + + user := model.User{ + ID: idgen.NextUint64ID(), + Username: input.Username, + Nickname: input.Nickname, + Email: input.Email, + IsActive: true, + IsAdmin: false, + LastLoginAt: time.Now(), + } + if user.Nickname == "" { + user.Nickname = input.DisplayName + } + if user.Nickname == "" { + user.Nickname = input.Username + } + if err := user.SetEncryptedPassword(input.Password); err != nil { + return LegacyUser{}, err + } + if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil { + return LegacyUser{}, err + } + if err := setLoginSession(ctx, c, &user); err != nil { + return LegacyUser{}, errors.New(errSaveSessionFailed) + } + + token, err := issueLegacyAccessToken(ctx, &user) + if err != nil { + return LegacyUser{}, err + } + return ToLegacyUser(&user, token), nil +} + +// Logout clears session and revokes the legacy access token when provided. +func Logout(ctx context.Context, c *gin.Context) error { + token := strings.TrimSpace(c.GetHeader(compat.OpenFlareTokenHeader())) + if token == "" { + token = strings.TrimSpace(c.GetHeader("X-Access-Token")) + } + if token != "" { + tokenHash := model.HashToken(token) + _ = db.DB(ctx).Where("token_hash = ? AND name = ?", tokenHash, legacyTokenName).Delete(&model.AccessToken{}).Error + } + + session := sessions.Default(c) + session.Options(oauth.GetSessionOptions(-1)) + session.Clear() + return session.Save() +} + +// GetSelf returns the current user's legacy profile. +func GetSelf(ctx context.Context, userID uint64) (LegacyUser, error) { + user, err := repository.GetUserByID(ctx, userID) + if err != nil { + return LegacyUser{}, errors.New(errUserNotFound) + } + return ToLegacyUser(&user, ""), nil +} + +// GenerateUserToken issues a fresh legacy access token for the user. +func GenerateUserToken(ctx context.Context, userID uint64) (string, error) { + user, err := repository.GetUserByID(ctx, userID) + if err != nil { + return "", errors.New(errUserNotFound) + } + return issueLegacyAccessToken(ctx, &user) +} + +// UpdateSelf updates the logged-in user's profile. +func UpdateSelf(ctx context.Context, userID uint64, input UpdateSelfInput) error { + user, err := repository.GetUserByID(ctx, userID) + if err != nil { + return errors.New(errUserNotFound) + } + + if input.DisplayName != "" { + user.Nickname = strings.TrimSpace(input.DisplayName) + } + if input.Username != "" { + user.Username = strings.TrimSpace(input.Username) + } + if input.Email != "" { + user.Email = strings.TrimSpace(input.Email) + } + if input.Password != "" { + if len(input.Password) < minPasswordLength { + return errors.New(errPasswordTooShort) + } + if err := user.SetEncryptedPassword(input.Password); err != nil { + return err + } + } + return db.DB(ctx).Save(&user).Error +} + +// DeleteSelf removes the logged-in user. +func DeleteSelf(ctx context.Context, userID uint64) error { + return repository.DeleteUserWithRelations(ctx, userID) +} + +// ListUsers returns a paginated legacy user list. +func ListUsers(ctx context.Context, page int) ([]LegacyUser, error) { + if page < 0 { + page = 0 + } + _, users, err := repository.ListAdminUsers(ctx, repository.AdminUserListFilter{ + Page: page + 1, + PageSize: legacyItemsPerPage, + }) + if err != nil { + return nil, err + } + return ToLegacyUsers(users), nil +} + +// SearchUsers searches users by keyword. +func SearchUsers(ctx context.Context, keyword string) ([]LegacyUser, error) { + keyword = strings.TrimSpace(keyword) + var users []model.User + query := db.DB(ctx).Model(&model.User{}). + Select("id, username, nickname, email, is_active, is_admin") + if keyword != "" { + like := keyword + "%" + query = query.Where( + "CAST(id AS TEXT) = ? OR username LIKE ? OR email LIKE ? OR nickname LIKE ?", + keyword, like, like, like, + ) + } + if err := query.Order("id DESC").Find(&users).Error; err != nil { + return nil, err + } + return ToLegacyUsers(users), nil +} + +// GetUserByID returns a legacy user if the caller has sufficient role. +func GetUserByID(ctx context.Context, callerRole int, id uint64) (LegacyUser, error) { + user, err := repository.GetAdminUserDetail(ctx, id) + if err != nil { + return LegacyUser{}, errors.New(errUserNotFound) + } + targetRole := RoleFromUser(&user) + if callerRole <= targetRole { + return LegacyUser{}, errors.New(errInsufficientPermission) + } + return ToLegacyUser(&user, ""), nil +} + +// CreateUser creates a user from legacy admin input. +func CreateUser(ctx context.Context, callerRole int, input CreateUserInput) error { + input.Username = strings.TrimSpace(input.Username) + input.Password = strings.TrimSpace(input.Password) + input.DisplayName = strings.TrimSpace(input.DisplayName) + input.Email = strings.TrimSpace(input.Email) + + if input.Username == "" || input.Password == "" { + return errors.New(errInvalidParams) + } + if input.Role >= callerRole { + return errors.New(errInsufficientPermission) + } + if input.Email == "" { + input.Email = input.Username + "@openflare.local" + } + + newUser := model.User{ + ID: idgen.NextUint64ID(), + Username: input.Username, + Nickname: input.DisplayName, + Email: input.Email, + IsActive: true, + IsAdmin: input.Role >= compat.RoleAdminUser, + } + if newUser.Nickname == "" { + newUser.Nickname = input.Username + } + if err := newUser.SetEncryptedPassword(input.Password); err != nil { + return err + } + return repository.CreateUser(ctx, &newUser) +} + +// UpdateUser updates another user with legacy role checks. +func UpdateUser(ctx context.Context, callerRole int, input UpdateUserInput) error { + if input.ID == 0 { + return errors.New(errInvalidParams) + } + + origin, err := repository.GetAdminUserDetail(ctx, uint64(input.ID)) + if err != nil { + return errors.New(errUserNotFound) + } + originRole := RoleFromUser(&origin) + if callerRole <= originRole { + return errors.New(errInsufficientPermission) + } + if input.Role > 0 && callerRole <= input.Role { + return errors.New(errInsufficientPermission) + } + + if trimmed := strings.TrimSpace(input.Username); trimmed != "" { + origin.Username = trimmed + } + if input.DisplayName != "" { + origin.Nickname = strings.TrimSpace(input.DisplayName) + } + if input.Email != "" { + origin.Email = strings.TrimSpace(input.Email) + } + if input.Role > 0 { + origin.IsAdmin = input.Role >= compat.RoleAdminUser + } + if input.Password != "" { + if len(input.Password) < minPasswordLength { + return errors.New(errPasswordTooShort) + } + if err := origin.SetEncryptedPassword(input.Password); err != nil { + return err + } + } + return db.DB(ctx).Save(&origin).Error +} + +// DeleteUserByID deletes a user when the caller has sufficient role. +func DeleteUserByID(ctx context.Context, callerRole int, id uint64) error { + origin, err := repository.GetAdminUserDetail(ctx, id) + if err != nil { + return errors.New(errUserNotFound) + } + if callerRole <= RoleFromUser(&origin) { + return errors.New(errInsufficientPermission) + } + if RoleFromUser(&origin) >= compat.RoleRootUser { + return errors.New(errCannotDeleteRoot) + } + return repository.DeleteUserWithRelations(ctx, id) +} + +// ManageUser performs enable/disable/delete/promote/demote actions. +func ManageUser(ctx context.Context, callerRole int, input ManageUserInput) (LegacyUser, error) { + input.Username = strings.TrimSpace(input.Username) + input.Action = strings.TrimSpace(input.Action) + if input.Username == "" || input.Action == "" { + return LegacyUser{}, errors.New(errInvalidParams) + } + + user, err := repository.GetUserByUsername(ctx, input.Username) + if err != nil { + return LegacyUser{}, errors.New(errUserNotFound) + } + targetRole := RoleFromUser(&user) + if callerRole <= targetRole && callerRole != compat.RoleRootUser { + return LegacyUser{}, errors.New(errInsufficientPermission) + } + + switch input.Action { + case "disable": + if targetRole >= compat.RoleRootUser { + return LegacyUser{}, errors.New(errCannotDisableRoot) + } + user.IsActive = false + case "enable": + user.IsActive = true + case "delete": + if targetRole >= compat.RoleRootUser { + return LegacyUser{}, errors.New(errCannotDeleteRoot) + } + if err := repository.DeleteUserWithRelations(ctx, user.ID); err != nil { + return LegacyUser{}, err + } + return LegacyUser{Role: compat.RoleCommonUser, Status: legacyUserStatusDisabled}, nil + case "promote": + if callerRole != compat.RoleRootUser { + return LegacyUser{}, errors.New(errCannotPromoteAdmin) + } + if user.IsAdmin { + return LegacyUser{}, errors.New(errAlreadyAdmin) + } + user.IsAdmin = true + case "demote": + if targetRole >= compat.RoleRootUser { + return LegacyUser{}, errors.New(errCannotDisableRoot) + } + if !user.IsAdmin { + return LegacyUser{}, errors.New(errAlreadyCommonUser) + } + user.IsAdmin = false + default: + return LegacyUser{}, errors.New(errInvalidParams) + } + + if err := db.DB(ctx).Save(&user).Error; err != nil { + return LegacyUser{}, err + } + return LegacyUser{ + Role: RoleFromUser(&user), + Status: StatusFromUser(&user), + }, nil +} + +// SendRegisterVerificationEmail sends a registration verification code. +func SendRegisterVerificationEmail(ctx context.Context, email string) error { + email = strings.TrimSpace(email) + if email == "" || !strings.Contains(email, "@") { + return errors.New(errInvalidParams) + } + var count int64 + if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil { + return err + } + if count > 0 { + return errors.New(errEmailAlreadyRegistered) + } + return sendEmailVerificationCode(ctx, email, "register", "register_email") +} + +// SendPasswordResetEmail stores a reset token and emails the user. +func SendPasswordResetEmail(ctx context.Context, email string) error { + email = strings.TrimSpace(email) + if email == "" || !strings.Contains(email, "@") { + return errors.New(errInvalidParams) + } + var user model.User + if err := db.DB(ctx).Where("email = ?", email).First(&user).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return errors.New(errEmailNotRegistered) + } + return err + } + + token, err := generateResetToken() + if err != nil { + return err + } + key := fmt.Sprintf(passwordResetKeyFmt, email) + if err := db.SetJSON(ctx, key, token, passwordResetExpiry); err != nil { + return err + } + + serverAddr, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress) + base := strings.TrimRight(serverAddr.Value, "/") + link := fmt.Sprintf("%s/user/reset?email=%s&token=%s", base, email, token) + body := fmt.Sprintf("
您好,你正在进行密码重置。
点击此处进行密码重置。
", link) + return dispatchEmail(ctx, email, "密码重置", body) +} + +// ResetPassword validates a reset token and returns a new random password. +func ResetPassword(ctx context.Context, input PasswordResetInput) (string, error) { + input.Email = strings.TrimSpace(input.Email) + input.Token = strings.TrimSpace(input.Token) + if input.Email == "" || input.Token == "" { + return "", errors.New(errInvalidParams) + } + + key := fmt.Sprintf(passwordResetKeyFmt, input.Email) + var stored string + if err := db.GetJSON(ctx, key, &stored); err != nil || stored != input.Token { + return "", errors.New(errResetLinkInvalid) + } + + password, err := generateResetToken() + if err != nil { + return "", err + } + if len(password) < minPasswordLength { + password = password + "Aa1!" + } + + var user model.User + if err := db.DB(ctx).Where("email = ?", input.Email).First(&user).Error; err != nil { + return "", errors.New(errEmailNotRegistered) + } + if err := user.SetEncryptedPassword(password); err != nil { + return "", err + } + if err := db.DB(ctx).Model(&user).Update("password", user.Password).Error; err != nil { + return "", err + } + _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() + return password, nil +} + +// BuildPublicStatus assembles the legacy /api/status payload. +func BuildPublicStatus(ctx context.Context) (map[string]any, error) { + authSources, err := publicAuthSources(ctx, "/api") + if err != nil { + authSources = []map[string]any{} + } + + siteName, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName) + serverAddr, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress) + + return map[string]any{ + "version": buildVersion(), + "start_time": appStartUnix(), + "email_verification": isEmailRegisterVerificationEnabled(ctx), + "github_oauth": false, + "github_client_id": "", + "system_name": siteName.Value, + "home_page_link": "", + "footer_html": "", + "wechat_qrcode": "", + "wechat_login": false, + "server_address": serverAddr.Value, + "password_register_enabled": isPasswordRegisterEnabled(ctx), + "cap_login_enabled": capLoginEnabled(ctx), + "auth_sources": authSources, + }, nil +} + +// GetNotice returns the legacy notice option value. +func GetNotice(ctx context.Context) string { + return getOptionValue(ctx, "notice") +} + +// GetAbout returns the legacy about option value. +func GetAbout(ctx context.Context) string { + return getOptionValue(ctx, "about") +} + +func issueLegacyAccessToken(ctx context.Context, user *model.User) (string, error) { + tokenStr, err := model.GenerateTokenString() + if err != nil { + return "", errors.New(errGenerateTokenFailed) + } + record := model.AccessToken{ + UserID: user.ID, + Name: legacyTokenName, + TokenHash: model.HashToken(tokenStr), + MaskedToken: model.MaskTokenString(tokenStr), + IsAdmin: user.IsAdmin, + } + if err := db.DB(ctx).Create(&record).Error; err != nil { + return "", err + } + return tokenStr, nil +} + +func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error { + session := sessions.Default(c) + session.Set(oauth.UserIDKey, user.ID) + session.Set(oauth.UserNameKey, user.Username) + session.Set(oauth.PasswordHashKey, user.Password) + + maxAge := config.Config.App.SessionAge + isSessionCookie := false + ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours) + if err == nil { + switch { + case ttlHours == -1: + maxAge = 10 * 365 * 24 * 3600 + case ttlHours > 0: + maxAge = ttlHours * 3600 + case ttlHours == 0: + isSessionCookie = true + } + } + session.Options(oauth.GetSessionOptions(maxAge)) + if err := session.Save(); err != nil { + return err + } + if isSessionCookie { + oauth.StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName) + } + return nil +} + +func isPasswordLoginEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled) + return err != nil || enabled +} + +func isPasswordRegisterEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled) + return err != nil || enabled +} + +func isRegistrationEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled) + return err != nil || enabled +} + +func isEmailLoginVerificationEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled) + return err == nil && enabled +} + +func isEmailRegisterVerificationEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailRegisterVerificationEnabled) + return err == nil && enabled +} + +func capLoginEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyCapLoginEnabled) + return err == nil && enabled +} + +func verifyLoginEmailCode(ctx context.Context, email, code string) error { + if code == "" { + return errors.New("need_email_code:" + email) + } + if !verifyEmailCode(ctx, email, "login", code) { + return errors.New(errEmailCodeInvalid) + } + return nil +} + +func verifyEmailCode(ctx context.Context, email, scene, code string) bool { + key := fmt.Sprintf("email_code:%s:%s", scene, email) + var stored string + if err := db.GetJSON(ctx, key, &stored); err != nil { + return false + } + if stored != code { + return false + } + _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() + return true +} + +func sendEmailVerificationCode(ctx context.Context, email, scene, templateName string) error { + scHost, errHost := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost) + scPort, errPort := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort) + scUser, errUser := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername) + scPass, errPass := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword) + if errHost != nil || errPort != nil || errUser != nil || errPass != nil || + scHost.Value == "" || scPort.Value == "" || scUser.Value == "" || scPass.Value == "" { + return errors.New("系统 SMTP 邮件服务配置不完整") + } + + code, err := generateVerificationCode() + if err != nil { + return err + } + codeKey := fmt.Sprintf("email_code:%s:%s", scene, email) + if err := db.SetJSON(ctx, codeKey, code, 5*time.Minute); err != nil { + return err + } + + tmpl, err := repository.GetTemplateByKey(ctx, templateName) + if err != nil { + body := fmt.Sprintf("您的验证码为: %s
", code) + return dispatchEmail(ctx, email, "邮箱验证", body) + } + subject, body, err := tmpl.Render(map[string]any{"Code": code}) + if err != nil { + return err + } + return dispatchEmail(ctx, email, subject, body) +} + +type sendEmailPayload struct { + To string `json:"to"` + Subject string `json:"subject"` + Body string `json:"body"` +} + +func dispatchEmail(ctx context.Context, to, subject, body string) error { + payload := sendEmailPayload{To: to, Subject: subject, Body: body} + payloadBytes, err := json.Marshal(payload) + if err != nil { + return err + } + _, err = task.DispatchTask(ctx, "mail:send", payloadBytes, "system") + return err +} + +func generateVerificationCode() (string, error) { + n, err := rand.Int(rand.Reader, big.NewInt(verificationCodeRange)) + if err != nil { + return "", err + } + return fmt.Sprintf("%06d", n.Int64()+verificationOffset), nil +} + +func generateResetToken() (string, error) { + n, err := rand.Int(rand.Reader, big.NewInt(1<<62)) + if err != nil { + return "", err + } + return fmt.Sprintf("%x", n.Int64()), nil +} + +func publicAuthSources(ctx context.Context, baseAPIPath string) ([]map[string]any, error) { + sources, err := model.GetActiveAuthSources(ctx) + if err != nil { + return nil, err + } + result := make([]map[string]any, 0, len(sources)) + base := strings.TrimRight(baseAPIPath, "/") + for _, source := range sources { + result = append(result, map[string]any{ + "id": source.ID, + "name": source.Name, + "type": source.Type, + "display_name": source.DisplayName, + "authorize_url": fmt.Sprintf("%s/oauth/%s/authorize", base, source.Name), + "icon_url": source.IconURL, + }) + } + return result, nil +} + +func getOptionValue(ctx context.Context, key string) string { + sc, err := repository.GetSystemConfigByKey(ctx, key) + if err != nil { + return "" + } + return sc.Value +} + +var appStart = time.Now() + +func appStartUnix() int64 { + return appStart.Unix() +} + +func buildVersion() string { + if buildinfo.Version != "" { + return buildinfo.Version + } + return "dev" +} diff --git a/Wavelet/internal/apps/openflare/auth/oauth.go b/Wavelet/internal/apps/openflare/auth/oauth.go new file mode 100644 index 00000000..700d1aeb --- /dev/null +++ b/Wavelet/internal/apps/openflare/auth/oauth.go @@ -0,0 +1,509 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/url" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/listener" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" + "github.com/google/uuid" + "gorm.io/gorm" +) + +const pendingExternalAccountSessionKey = "pending_external_account" + +// OAuthCallbackResult is the legacy OAuth callback payload. +type OAuthCallbackResult struct { + Status string `json:"status"` + User *LegacyUser `json:"user,omitempty"` +} + +// PendingExternalAccount stores OAuth bind-pending state in session. +type PendingExternalAccount struct { + AuthSourceID uint64 `json:"auth_source_id"` + ExternalID string `json:"external_id"` + ExternalUsername string `json:"external_username"` + DisplayName string `json:"display_name"` + Email string `json:"email"` +} + +// OAuthAuthorize builds an authorize URL for a legacy auth source route param. +func OAuthAuthorize(ctx context.Context, c *gin.Context, sourceKey string) (string, error) { + source, err := resolveAuthSourceByRoute(ctx, sourceKey) + if err != nil { + return "", err + } + if !source.IsActive { + return "", errors.New(errAuthSourceDisabled) + } + if err := source.Validate(); err != nil { + return "", err + } + + state := uuid.NewString() + session := sessions.Default(c) + session.Set(oauthStateSessionKey(source.ID), state) + if err := session.Save(); err != nil { + return "", errors.New(errSaveSessionFailed) + } + + redirectURL := legacyOAuthCallbackURL(c, source) + return buildLegacyAuthorizeURL(ctx, source, redirectURL, state) +} + +// OAuthCallback handles GET /oauth/:source/callback for the legacy frontend. +func OAuthCallback(ctx context.Context, c *gin.Context, sourceKey string) (OAuthCallbackResult, error) { + source, err := resolveAuthSourceByRoute(ctx, sourceKey) + if err != nil { + return OAuthCallbackResult{}, err + } + if !source.IsActive { + return OAuthCallbackResult{}, errors.New(errAuthSourceDisabled) + } + + session := sessions.Default(c) + expectedState, _ := session.Get(oauthStateSessionKey(source.ID)).(string) + state := c.Query("state") + if expectedState == "" || state == "" || state != expectedState { + return OAuthCallbackResult{}, errors.New("授权状态无效,请重新登录") + } + session.Delete(oauthStateSessionKey(source.ID)) + if err := session.Save(); err != nil { + return OAuthCallbackResult{}, errors.New(errSaveSessionFailed) + } + if oauthError := c.Query("error"); oauthError != "" { + description := c.Query("error_description") + if description == "" { + description = oauthError + } + return OAuthCallbackResult{}, errors.New(description) + } + + redirectURL := legacyOAuthCallbackURL(c, source) + userInfo, err := exchangeLegacyOAuthProfile(ctx, source, c.Query("code"), state, redirectURL) + if err != nil { + return OAuthCallbackResult{}, err + } + + var currentUserID *uint64 + if current := currentUserFromLegacyToken(ctx, c); current != nil { + currentUserID = ¤t.ID + } + + result, pending, err := completeLegacyOAuthLogin(ctx, source, userInfo, currentUserID) + if err != nil { + return OAuthCallbackResult{}, err + } + if pending != nil { + raw, marshalErr := json.Marshal(pending) + if marshalErr != nil { + return OAuthCallbackResult{}, marshalErr + } + session.Set(pendingExternalAccountSessionKey, string(raw)) + if err := session.Save(); err != nil { + return OAuthCallbackResult{}, errors.New(errSaveSessionFailed) + } + return result, nil + } + if result.User != nil { + var dbUser model.User + if err := db.DB(ctx).Where("id = ?", result.User.ID).First(&dbUser).Error; err != nil { + return OAuthCallbackResult{}, err + } + if err := setLoginSession(ctx, c, &dbUser); err != nil { + return OAuthCallbackResult{}, errors.New(errSaveSessionFailed) + } + token, tokenErr := issueLegacyAccessToken(ctx, &dbUser) + if tokenErr != nil { + return OAuthCallbackResult{}, tokenErr + } + legacy := ToLegacyUser(&dbUser, token) + result.User = &legacy + listener.EmitAdminLoggedIn(ctx, &dbUser, c.ClientIP()) + } + return result, nil +} + +// LinkExistingOAuthAccount binds a pending external account to an existing user. +func LinkExistingOAuthAccount(ctx context.Context, c *gin.Context, input LinkExistingInput) (OAuthCallbackResult, error) { + session := sessions.Default(c) + raw, _ := session.Get(pendingExternalAccountSessionKey).(string) + if raw == "" { + return OAuthCallbackResult{}, errors.New(errPendingOAuthExpired) + } + var pending PendingExternalAccount + if err := json.Unmarshal([]byte(raw), &pending); err != nil { + return OAuthCallbackResult{}, errors.New(errPendingOAuthInvalid) + } + + user, err := linkPendingExternalAccount(ctx, &pending, input) + if err != nil { + return OAuthCallbackResult{}, err + } + + session.Delete(pendingExternalAccountSessionKey) + if err := session.Save(); err != nil { + return OAuthCallbackResult{}, errors.New(errSaveSessionFailed) + } + if err := setLoginSession(ctx, c, user); err != nil { + return OAuthCallbackResult{}, errors.New(errSaveSessionFailed) + } + token, err := issueLegacyAccessToken(ctx, user) + if err != nil { + return OAuthCallbackResult{}, err + } + legacy := ToLegacyUser(user, token) + return OAuthCallbackResult{Status: "linked", User: &legacy}, nil +} + +func resolveAuthSourceByRoute(ctx context.Context, raw string) (*model.AuthSource, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return nil, errors.New("认证源不能为空") + } + if parsed, err := parseUint64(raw); err == nil && parsed > 0 { + return model.GetAuthSourceByID(ctx, parsed) + } + return model.GetAuthSourceByName(ctx, raw) +} + +func oauthStateSessionKey(sourceID uint64) string { + return fmt.Sprintf("oauth_state_%d", sourceID) +} + +func legacyOAuthCallbackURL(c *gin.Context, source *model.AuthSource) string { + ctx := c.Request.Context() + base := "" + if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil { + base = strings.TrimRight(sc.Value, "/") + } + if base == "" { + scheme := "http" + if c.Request.TLS != nil || c.GetHeader("X-Forwarded-Proto") == "https" { + scheme = "https" + } + host := c.Request.Host + if forwardedHost := c.GetHeader("X-Forwarded-Host"); forwardedHost != "" { + host = forwardedHost + } + base = scheme + "://" + host + } + sourceName := source.Name + if sourceName == "" { + sourceName = fmt.Sprintf("%d", source.ID) + } + callback, _ := url.JoinPath(base, "oauth", sourceName) + return callback +} + +func buildLegacyAuthorizeURL(ctx context.Context, source *model.AuthSource, redirectURL, state string) (string, error) { + payloadValue, err := encodeLegacyOAuthState(source.Name, state) + if err != nil { + return "", err + } + stateKey := fmt.Sprintf("of_oauth_state:%s", state) + if err := db.Redis.Set(ctx, db.PrefixedKey(stateKey), payloadValue, 10*time.Minute).Err(); err != nil { + return "", err + } + return oauthBuildAuthorizeURL(ctx, source, redirectURL, state) +} + +func encodeLegacyOAuthState(sourceName, state string) (string, error) { + payload := map[string]string{ + "source_name": sourceName, + "state": state, + } + raw, err := json.Marshal(payload) + if err != nil { + return "", err + } + return string(raw), nil +} + +func legacyFrontendLoginRedirectURL(ctx context.Context, source *model.AuthSource) (string, error) { + sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress) + if err != nil || strings.TrimSpace(sc.Value) == "" { + return "", errors.New("server_address 未配置") + } + base := strings.TrimRight(sc.Value, "/") + name := source.Name + if name == "" { + name = fmt.Sprintf("%d", source.ID) + } + return base + "/oauth/" + url.PathEscape(name), nil +} + +func exchangeLegacyOAuthProfile(ctx context.Context, source *model.AuthSource, code, state, redirectURL string) (*model.OAuthUserInfo, error) { + if strings.TrimSpace(code) == "" { + return nil, errors.New("授权 code 不能为空") + } + // Validate state from Redis cache written during authorize. + stateKey := fmt.Sprintf("of_oauth_state:%s", state) + payloadRaw, err := db.Redis.Get(ctx, db.PrefixedKey(stateKey)).Result() + if err != nil { + return nil, errors.New("授权状态无效,请重新登录") + } + _ = db.Redis.Del(ctx, db.PrefixedKey(stateKey)).Err() + + var payload map[string]string + if err := json.Unmarshal([]byte(payloadRaw), &payload); err != nil { + return nil, err + } + if payload["source_name"] != source.Name { + return nil, errors.New("授权状态无效,请重新登录") + } + + userInfo, err := buildOAuthUserInfo(ctx, source, code, state, redirectURL) + if err != nil { + return nil, err + } + if err := normalizeOAuthUserInfo(userInfo); err != nil { + return nil, err + } + if userInfo.Sub == "" { + userInfo.Sub = userInfo.Username + } + return userInfo, nil +} + +func completeLegacyOAuthLogin(ctx context.Context, source *model.AuthSource, profile *model.OAuthUserInfo, currentUserID *uint64) (OAuthCallbackResult, *PendingExternalAccount, error) { + if source == nil || profile == nil || strings.TrimSpace(profile.Sub) == "" { + return OAuthCallbackResult{}, nil, errors.New("第三方账号资料不完整") + } + + account, err := model.FindExternalAccount(ctx, source.ID, profile.Sub) + if err == nil { + var user model.User + if err := db.DB(ctx).Where("id = ?", account.UserID).First(&user).Error; err != nil { + return OAuthCallbackResult{}, nil, err + } + if !user.IsActive { + return OAuthCallbackResult{}, nil, errors.New(errBannedAccount) + } + legacy := ToLegacyUser(&user, "") + return OAuthCallbackResult{Status: "logged_in", User: &legacy}, nil, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return OAuthCallbackResult{}, nil, err + } + + if currentUserID != nil && *currentUserID > 0 { + var user model.User + if err := db.DB(ctx).Where("id = ?", *currentUserID).First(&user).Error; err != nil { + return OAuthCallbackResult{}, nil, err + } + if !user.IsActive { + return OAuthCallbackResult{}, nil, errors.New(errBannedAccount) + } + if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ + AuthSourceID: source.ID, + UserID: user.ID, + ExternalID: profile.Sub, + ExternalUsername: profile.Username, + Email: profile.Email, + }); err != nil { + return OAuthCallbackResult{}, nil, err + } + legacy := ToLegacyUser(&user, "") + return OAuthCallbackResult{Status: "linked", User: &legacy}, nil, nil + } + + registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled) + if regErr != nil { + registrationEnabled = true + } + if !registrationEnabled { + pending := &PendingExternalAccount{ + AuthSourceID: source.ID, + ExternalID: profile.Sub, + ExternalUsername: profile.Username, + DisplayName: profile.Name, + Email: profile.Email, + } + return OAuthCallbackResult{Status: "link_required"}, pending, nil + } + + user, err := createUserFromOAuthProfile(ctx, source, profile) + if err != nil { + return OAuthCallbackResult{}, nil, err + } + legacy := ToLegacyUser(&user, "") + return OAuthCallbackResult{Status: "logged_in", User: &legacy}, nil, nil +} + +func createUserFromOAuthProfile(ctx context.Context, source *model.AuthSource, profile *model.OAuthUserInfo) (model.User, error) { + username, err := uniqueLegacyUsername(ctx, profile.Username) + if err != nil { + return model.User{}, err + } + profile.Username = username + + var user model.User + if err := user.CreateUser(ctx, db.DB(ctx), profile); err != nil { + return model.User{}, err + } + if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ + AuthSourceID: source.ID, + UserID: user.ID, + ExternalID: profile.Sub, + ExternalUsername: profile.Username, + Email: profile.Email, + }); err != nil { + return model.User{}, err + } + logger.InfoF(ctx, "[LoginAudit] successful legacy OAuth registration via source: %s, user: %s, ID: %d", source.Name, user.Username, user.ID) + return user, nil +} + +func linkPendingExternalAccount(ctx context.Context, pending *PendingExternalAccount, input LinkExistingInput) (*model.User, error) { + if pending == nil || pending.AuthSourceID == 0 || pending.ExternalID == "" { + return nil, errors.New(errPendingOAuthExpired) + } + input.Username = strings.TrimSpace(input.Username) + if input.Username == "" || input.Password == "" { + return nil, errors.New(errInvalidParams) + } + + var user model.User + if err := db.DB(ctx).Where("username = ? OR email = ?", input.Username, input.Username).First(&user).Error; err != nil { + return nil, errors.New(errUsernameOrPasswordWrong) + } + if !user.IsActive { + return nil, errors.New(errBannedAccount) + } + if !user.CheckPassword(input.Password) { + return nil, errors.New(errUsernameOrPasswordWrong) + } + + if existing, err := model.FindExternalAccount(ctx, pending.AuthSourceID, pending.ExternalID); err == nil { + if existing.UserID != user.ID { + return nil, errors.New("该第三方账号已绑定其他用户") + } + return &user, nil + } else if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + + if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ + AuthSourceID: pending.AuthSourceID, + UserID: user.ID, + ExternalID: pending.ExternalID, + ExternalUsername: pending.ExternalUsername, + Email: pending.Email, + }); err != nil { + return nil, err + } + return &user, nil +} + +func currentUserFromLegacyToken(ctx context.Context, c *gin.Context) *model.User { + token := strings.TrimSpace(c.GetHeader(compat.OpenFlareTokenHeader())) + if token == "" { + return nil + } + tokenHash := model.HashToken(token) + var record model.AccessToken + if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&record).Error; err != nil { + return nil + } + var user model.User + if err := db.DB(ctx).Where("id = ? AND is_active = ?", record.UserID, true).First(&user).Error; err != nil { + return nil + } + return &user +} + +func uniqueLegacyUsername(ctx context.Context, base string) (string, error) { + base = strings.TrimSpace(base) + if base == "" { + base = "user" + } + candidate := base + for i := 0; i <= 1000; i++ { + if i > 0 { + candidate = fmt.Sprintf("%s-%d", base, i) + } + count, err := repository.CountUsersByUsername(ctx, candidate) + if err != nil { + return "", err + } + if count == 0 { + return candidate, nil + } + } + return "", errors.New("无法生成唯一用户名") +} + +func parseUint64(raw string) (uint64, error) { + var id uint64 + _, err := fmt.Sscanf(raw, "%d", &id) + return id, err +} + +func isOIDCLoginEnabled(ctx context.Context) bool { + enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled) + return err != nil || enabled +} + +// The following functions mirror oauth package internals for legacy GET callback support. +// They intentionally duplicate minimal logic to avoid modifying the core oauth module. + +func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code, nonce, redirectURL string) (*model.OAuthUserInfo, error) { + authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) + if err != nil { + return nil, err + } + token, err := authConfig.Exchange(ctx, code) + if err != nil { + return nil, err + } + userInfo := &model.OAuthUserInfo{Active: true} + if verifier != nil { + if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil { + return nil, verifyErr + } + } + if userInfo.Username == "" && userInfo.PreferredUsername != "" { + userInfo.Username = userInfo.PreferredUsername + } + if userInfo.Username == "" && userInfo.Email != "" { + userInfo.Username = strings.Split(userInfo.Email, "@")[0] + } + if userInfo.Username == "" && userInfo.Sub != "" { + userInfo.Username = userInfo.Sub + } + if userInfo.Name == "" { + userInfo.Name = userInfo.Username + } + return userInfo, nil +} + +func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error { + userInfo.Username = strings.TrimSpace(userInfo.Username) + userInfo.Email = strings.TrimSpace(userInfo.Email) + userInfo.Name = strings.TrimSpace(userInfo.Name) + if userInfo.Username == "" { + return errors.New("无法从认证源获取用户名") + } + if userInfo.Name == "" { + userInfo.Name = userInfo.Username + } + if !userInfo.Active { + userInfo.Active = true + } + return nil +} diff --git a/Wavelet/internal/apps/openflare/auth/oauth_compat.go b/Wavelet/internal/apps/openflare/auth/oauth_compat.go new file mode 100644 index 00000000..8048a5b0 --- /dev/null +++ b/Wavelet/internal/apps/openflare/auth/oauth_compat.go @@ -0,0 +1,134 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "context" + "errors" + "fmt" + "net/http" + "strings" + "sync" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" + "golang.org/x/sync/singleflight" +) + +type legacyOIDCProviderCacheType struct { + mu sync.RWMutex + entries map[string]*oidc.Provider + sfGroup singleflight.Group +} + +var legacyOIDCProviderCache = &legacyOIDCProviderCacheType{ + entries: make(map[string]*oidc.Provider), +} + +func (c *legacyOIDCProviderCacheType) get(ctx context.Context, issuer string) (*oidc.Provider, error) { + c.mu.RLock() + if p, ok := c.entries[issuer]; ok { + c.mu.RUnlock() + return p, nil + } + c.mu.RUnlock() + + bg := context.Background() + if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil { + bg = oidc.ClientContext(bg, client) + } + + v, err, _ := c.sfGroup.Do(issuer, func() (any, error) { + c.mu.RLock() + if p, ok := c.entries[issuer]; ok { + c.mu.RUnlock() + return p, nil + } + c.mu.RUnlock() + + p, err := oidc.NewProvider(bg, issuer) + if err != nil { + return nil, err + } + c.mu.Lock() + c.entries[issuer] = p + c.mu.Unlock() + return p, nil + }) + if err != nil { + return nil, err + } + return v.(*oidc.Provider), nil //nolint:forcetypeassert // singleflight value type is fixed +} + +// oauthBuildAuthorizeURL builds an OAuth authorize URL using the same rules as apps/oauth. +func oauthBuildAuthorizeURL(ctx context.Context, source *model.AuthSource, redirectURL, state string) (string, error) { + authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) + if err != nil { + return "", err + } + if verifier != nil { + return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil + } + return authConfig.AuthCodeURL(state), nil +} + +func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) { + if source == nil { + return nil, nil, errors.New("认证源不能为空") + } + if source.OpenIDDiscoveryURL == "" { + return nil, nil, errors.New("认证源未配置 OpenID Discovery URL") + } + + issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/") + issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration") + issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server") + + provider, err := legacyOIDCProviderCache.get(ctx, issuer) + if err != nil { + return nil, nil, err + } + verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID}) + scopes := strings.Fields(source.Scopes) + if len(scopes) == 0 { + scopes = []string{oidc.ScopeOpenID, "profile", "email"} + } + if !containsScope(scopes, oidc.ScopeOpenID) { + scopes = append([]string{oidc.ScopeOpenID}, scopes...) + } + + return &oauth2.Config{ + ClientID: source.ClientID, + ClientSecret: source.ClientSecret, + RedirectURL: redirectURL, + Scopes: scopes, + Endpoint: provider.Endpoint(), + }, verifier, nil +} + +func containsScope(scopes []string, scope string) bool { + for _, item := range scopes { + if item == scope { + return true + } + } + return false +} + +func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error { + rawIDToken, ok := token.Extra("id_token").(string) + if !ok { + return nil + } + idToken, verifyErr := verifier.Verify(ctx, rawIDToken) + if verifyErr != nil { + return fmt.Errorf("ID Token 验证失败: %w", verifyErr) + } + if nonce != "" && idToken.Nonce != nonce { + return errors.New("nonce 不匹配") + } + return idToken.Claims(userInfo) +} diff --git a/Wavelet/internal/apps/openflare/compat/auth.go b/Wavelet/internal/apps/openflare/compat/auth.go new file mode 100644 index 00000000..d8091775 --- /dev/null +++ b/Wavelet/internal/apps/openflare/compat/auth.go @@ -0,0 +1,78 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package compat + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +const openFlareTokenHeader = "OpenFlare-Token" + +// Role constants mirror legacy OpenFlare role values. +const ( + RoleCommonUser = 1 + RoleAdminUser = 10 + RoleRootUser = 100 +) + +// RequireRole ensures the caller is authenticated with at least minRole. +// Phase 1: supports Wavelet session/access-token via oauth.GetUserFromRequest. +func RequireRole(minRole int) gin.HandlerFunc { + return func(c *gin.Context) { + user, err := oauth.GetUserFromRequest(c) + if err != nil || user == nil { + Unauthorized(c, "无权进行此操作,未登录或 token 无效") + c.Abort() + return + } + role := resolveRole(user) + if role < minRole { + Fail(c, "无权进行此操作,权限不足") + c.Abort() + return + } + c.Set("of_user_id", user.ID) + c.Set("of_role", role) + c.Set("of_is_admin", user.IsAdmin) + c.Next() + } +} + +// UserAuth requires role >= CommonUser. +func UserAuth() gin.HandlerFunc { return RequireRole(RoleCommonUser) } + +// AdminAuth requires role >= AdminUser. +func AdminAuth() gin.HandlerFunc { return RequireRole(RoleAdminUser) } + +// RootAuth requires role >= RootUser. +func RootAuth() gin.HandlerFunc { return RequireRole(RoleRootUser) } + +func resolveRole(user *model.User) int { + if user == nil { + return 0 + } + if user.IsAdmin { + return RoleRootUser + } + return RoleCommonUser +} + +// OpenFlareTokenHeader returns the legacy auth header name. +func OpenFlareTokenHeader() string { + return openFlareTokenHeader +} + +// BridgeOpenFlareToken maps OpenFlare-Token to X-Access-Token for legacy clients. +func BridgeOpenFlareToken() gin.HandlerFunc { + return func(c *gin.Context) { + if c.GetHeader("X-Access-Token") == "" { + if token := c.GetHeader(openFlareTokenHeader); token != "" { + c.Request.Header.Set("X-Access-Token", token) + } + } + c.Next() + } +} diff --git a/Wavelet/internal/apps/openflare/compat/bind.go b/Wavelet/internal/apps/openflare/compat/bind.go new file mode 100644 index 00000000..c285c6d0 --- /dev/null +++ b/Wavelet/internal/apps/openflare/compat/bind.go @@ -0,0 +1,34 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package compat + +import ( + "strconv" + + "github.com/gin-gonic/gin" +) + +// IDParam parses :id from the URL path. +func IDParam(c *gin.Context) (uint, bool) { + raw := c.Param("id") + if raw == "" { + Fail(c, "无效的 ID") + return 0, false + } + id64, err := strconv.ParseUint(raw, 10, 64) + if err != nil || id64 == 0 { + Fail(c, "无效的 ID") + return 0, false + } + return uint(id64), true +} + +// BindJSON binds JSON body; returns false after writing a failure response. +func BindJSON(c *gin.Context, dst any) bool { + if err := c.ShouldBindJSON(dst); err != nil { + Fail(c, "参数错误") + return false + } + return true +} diff --git a/Wavelet/internal/apps/openflare/compat/response.go b/Wavelet/internal/apps/openflare/compat/response.go new file mode 100644 index 00000000..ac2387c2 --- /dev/null +++ b/Wavelet/internal/apps/openflare/compat/response.go @@ -0,0 +1,43 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package compat provides OpenFlare legacy API compatibility helpers. +package compat + +import ( + "net/http" + + "github.com/gin-gonic/gin" +) + +// Envelope is the legacy OpenFlare frontend response format. +type Envelope struct { + Success bool `json:"success"` + Message string `json:"message"` + Data any `json:"data"` +} + +// OK sends a successful legacy response. +func OK(c *gin.Context, data any) { + c.JSON(http.StatusOK, Envelope{Success: true, Message: "", Data: data}) +} + +// OKMessage sends a successful legacy response with a message. +func OKMessage(c *gin.Context, message string) { + c.JSON(http.StatusOK, Envelope{Success: true, Message: message, Data: nil}) +} + +// Fail sends a failed legacy response with HTTP 200 (OpenFlare convention). +func Fail(c *gin.Context, message string) { + c.JSON(http.StatusOK, Envelope{Success: false, Message: message, Data: nil}) +} + +// Unauthorized sends a 401 legacy response. +func Unauthorized(c *gin.Context, message string) { + c.JSON(http.StatusUnauthorized, Envelope{Success: false, Message: message, Data: nil}) +} + +// Forbidden sends a 403 legacy response. +func Forbidden(c *gin.Context, message string) { + c.JSON(http.StatusForbidden, Envelope{Success: false, Message: message, Data: nil}) +} diff --git a/Wavelet/internal/apps/openflare/config_version/errs.go b/Wavelet/internal/apps/openflare/config_version/errs.go new file mode 100644 index 00000000..2ccd9bb6 --- /dev/null +++ b/Wavelet/internal/apps/openflare/config_version/errs.go @@ -0,0 +1,12 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package config_version + +const ( + errNoActiveVersion = "当前没有激活版本" + errNoEnabledRoutes = "没有可发布的启用规则" + errNoChangesToPublish = "当前规则没有变更,不能重复发布" + errVersionConflict = "版本号生成冲突,请重试" + errInvalidSnapshotFormat = "历史版本快照格式不合法" +) diff --git a/Wavelet/internal/apps/openflare/config_version/helpers.go b/Wavelet/internal/apps/openflare/config_version/helpers.go new file mode 100644 index 00000000..41bc37f1 --- /dev/null +++ b/Wavelet/internal/apps/openflare/config_version/helpers.go @@ -0,0 +1,352 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package config_version + +import ( + "context" + "encoding/json" + "fmt" + "net" + "strconv" + "strings" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +type customHeaderInput struct { + Key string `json:"key"` + Value string `json:"value"` +} + +func isUniqueConstraintError(err error) bool { + if err == nil { + return false + } + return strings.Contains(strings.ToLower(err.Error()), "unique") +} + +func normalizeProxyRouteSiteName(route *model.ProxyRoute, raw, primaryDomain string) string { + siteName := strings.TrimSpace(raw) + if siteName != "" { + return siteName + } + if route != nil && strings.TrimSpace(route.SiteName) != "" { + return strings.TrimSpace(route.SiteName) + } + return primaryDomain +} + +func normalizeProxyRouteDomains(rawDomains []string) ([]string, error) { + normalized := make([]string, 0, len(rawDomains)) + seen := make(map[string]struct{}, len(rawDomains)) + for _, rawDomain := range rawDomains { + domain := strings.ToLower(strings.TrimSpace(rawDomain)) + if domain == "" { + continue + } + if strings.Contains(domain, "://") || strings.Contains(domain, "/") { + return nil, fmt.Errorf("domain %q is invalid", rawDomain) + } + if _, ok := seen[domain]; ok { + continue + } + seen[domain] = struct{}{} + normalized = append(normalized, domain) + } + if len(normalized) == 0 { + return nil, fmt.Errorf("domain is required") + } + return normalized, nil +} + +func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return normalizeProxyRouteDomains([]string{fallbackDomain}) + } + var domains []string + if err := json.Unmarshal([]byte(text), &domains); err != nil { + return nil, fmt.Errorf("domains payload is invalid") + } + return normalizeProxyRouteDomains(domains) +} + +func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return normalizeUpstreams(fallbackOriginURL, nil) + } + var upstreams []string + if err := json.Unmarshal([]byte(text), &upstreams); err != nil { + return nil, fmt.Errorf("upstreams payload is invalid") + } + return normalizeUpstreams(fallbackOriginURL, upstreams) +} + +func normalizeUpstreams(originURL string, upstreams []string) ([]string, error) { + candidates := upstreams + if len(candidates) == 0 { + candidates = []string{originURL} + } + normalized := make([]string, 0, len(candidates)) + seen := make(map[string]struct{}, len(candidates)) + for _, item := range candidates { + value := strings.TrimSpace(item) + if value == "" { + continue + } + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + normalized = append(normalized, value) + } + if len(normalized) == 0 { + return nil, fmt.Errorf("upstream is required") + } + return normalized, nil +} + +func decodeStoredCustomHeaders(raw string) ([]customHeaderInput, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []customHeaderInput{}, nil + } + var headers []customHeaderInput + if err := json.Unmarshal([]byte(text), &headers); err != nil { + return nil, fmt.Errorf("custom_headers payload is invalid") + } + return headers, nil +} + +func decodeStoredCacheRules(raw string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []string{}, nil + } + var rules []string + if err := json.Unmarshal([]byte(text), &rules); err != nil { + return nil, fmt.Errorf("cache_rules payload is invalid") + } + normalized := make([]string, 0, len(rules)) + for _, rule := range rules { + item := strings.TrimSpace(rule) + if item == "" { + continue + } + normalized = append(normalized, item) + } + return normalized, nil +} + +func decodeStoredCertIDs(raw string, fallbackCertID *uint) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + if fallbackCertID == nil || *fallbackCertID == 0 { + return []uint{}, nil + } + return []uint{*fallbackCertID}, nil + } + var certIDs []uint + if err := json.Unmarshal([]byte(text), &certIDs); err != nil { + return nil, fmt.Errorf("cert_ids payload is invalid") + } + normalized := make([]uint, 0, len(certIDs)) + seen := make(map[uint]struct{}, len(certIDs)) + for _, certID := range certIDs { + if certID == 0 { + continue + } + if _, ok := seen[certID]; ok { + continue + } + seen[certID] = struct{}{} + normalized = append(normalized, certID) + } + if len(normalized) == 0 && fallbackCertID != nil && *fallbackCertID != 0 { + return []uint{*fallbackCertID}, nil + } + return normalized, nil +} + +func resolveDomainCertIDs(domains []string, certIDs []uint, rawDomainCertIDs string) ([]uint, error) { + text := strings.TrimSpace(rawDomainCertIDs) + if text != "" { + var domainCertIDs []uint + if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil { + return nil, fmt.Errorf("domain_cert_ids payload is invalid") + } + if len(domains) > 0 && len(domainCertIDs) != len(domains) { + return nil, fmt.Errorf("domain_cert_ids length is invalid") + } + return domainCertIDs, nil + } + if len(certIDs) == 0 { + return []uint{}, nil + } + if len(certIDs) == 1 { + result := make([]uint, len(domains)) + for index := range result { + result[index] = certIDs[0] + } + return result, nil + } + if len(certIDs) == len(domains) { + result := make([]uint, len(certIDs)) + copy(result, certIDs) + return result, nil + } + return []uint{}, nil +} + +func mustDecodeCertIDs(route *model.ProxyRoute) []uint { + if route == nil { + return []uint{} + } + certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) + if err != nil { + return []uint{} + } + return certIDs +} + +func mustDecodeDomainCertIDs(route *model.ProxyRoute, domains []string) []uint { + if route == nil { + return []uint{} + } + certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) + if err != nil { + return []uint{} + } + domainCertIDs, err := resolveDomainCertIDs(domains, certIDs, route.DomainCertIDs) + if err != nil { + return []uint{} + } + return domainCertIDs +} + +func normalizeUpstreamType(raw string) string { + value := strings.ToLower(strings.TrimSpace(raw)) + switch value { + case "tunnel", "pages": + return value + default: + return "direct" + } +} + +func normalizeTunnelTargetProtocol(raw string) string { + value := strings.ToLower(strings.TrimSpace(raw)) + switch value { + case "http", "https", "tcp": + return value + default: + return "http" + } +} + +func normalizePEM(content string) string { + return strings.TrimSpace(content) + "\n" +} + +func certificateCertFileName(id uint) string { + return fmt.Sprintf("%d.crt", id) +} + +func certificateKeyFileName(id uint) string { + return fmt.Sprintf("%d.key", id) +} + +func dedupeSupportFiles(files []SupportFile) []SupportFile { + if len(files) == 0 { + return nil + } + unique := make(map[string]SupportFile, len(files)) + for _, file := range files { + unique[file.Path] = file + } + result := make([]SupportFile, 0, len(unique)) + for _, file := range unique { + result = append(result, file) + } + return result +} + +func uintPtrEqual(left *uint, right *uint) bool { + if left == nil || right == nil { + return left == nil && right == nil + } + return *left == *right +} + +func uintSliceEqual(left []uint, right []uint) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] != right[index] { + return false + } + } + return true +} + +func relayAgentAddress(node *model.OpenFlareNode) string { + if node == nil { + return "" + } + port := node.RelayVhostHTTPPort + if port <= 0 { + port = 8080 + } + addr := strings.TrimSpace(node.RelayAgentAccessAddr) + if addr == "" { + addr = strings.TrimSpace(node.RelayClientAccessAddr) + } + if addr == "" { + addr = strings.TrimSpace(node.IP) + } + if addr == "" { + return fmt.Sprintf("127.0.0.1:%d", port) + } + if _, _, err := net.SplitHostPort(addr); err == nil { + return addr + } + if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 { + return net.JoinHostPort(addr, strconv.Itoa(port)) + } + return fmt.Sprintf("%s:%d", addr, port) +} + +func resolveTunnelOpenRestyUpstreamURL(ctx context.Context) string { + nodes, err := model.ListOpenFlareNodes(ctx) + if err == nil { + for index := range nodes { + node := &nodes[index] + if node.NodeType != "tunnel_relay" { + continue + } + addr := relayAgentAddress(node) + if addr != "" { + return "http://" + addr + } + } + } + return "http://127.0.0.1:8080" +} + +func listWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*model.OpenFlareWAFIPGroup, error) { + if len(ids) == 0 { + return []*model.OpenFlareWAFIPGroup{}, nil + } + groups := make([]*model.OpenFlareWAFIPGroup, 0, len(ids)) + for _, id := range ids { + group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + if err != nil { + return nil, err + } + groups = append(groups, group) + } + return groups, nil +} diff --git a/Wavelet/internal/apps/openflare/config_version/logics.go b/Wavelet/internal/apps/openflare/config_version/logics.go new file mode 100644 index 00000000..0102e664 --- /dev/null +++ b/Wavelet/internal/apps/openflare/config_version/logics.go @@ -0,0 +1,559 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package config_version + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sort" + "strconv" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// ConfigPreviewResult is the preview response payload. +type ConfigPreviewResult struct { + SnapshotJSON string `json:"snapshot_json"` + MainConfig string `json:"main_config"` + RouteConfig string `json:"route_config"` + RenderedConfig string `json:"rendered_config"` + SupportFiles []SupportFile `json:"support_files"` + Checksum string `json:"checksum"` + RouteCount int `json:"route_count"` + WebsiteCount int `json:"website_count"` +} + +// ConfigDiffResult is the diff response payload. +type ConfigDiffResult struct { + ActiveVersion string `json:"active_version,omitempty"` + AddedSites []string `json:"added_sites"` + RemovedSites []string `json:"removed_sites"` + ModifiedSites []string `json:"modified_sites"` + AddedDomains []string `json:"added_domains"` + RemovedDomains []string `json:"removed_domains"` + ModifiedDomains []string `json:"modified_domains"` + MainConfigChanged bool `json:"main_config_changed"` + WAFConfigChanged bool `json:"waf_config_changed"` + ChangedOptionKeys []string `json:"changed_option_keys"` + ChangedOptionDetails []ConfigOptionDiffItem `json:"changed_option_details"` + CurrentWebsiteCount int `json:"current_website_count"` + ActiveWebsiteCount int `json:"active_website_count"` +} + +// ConfigOptionDiffItem describes a changed OpenResty option. +type ConfigOptionDiffItem struct { + Key string `json:"key"` + PreviousValue string `json:"previous_value"` + CurrentValue string `json:"current_value"` +} + +// CleanupInput is the cleanup request payload. +type CleanupInput struct { + KeepCount int `json:"keep_count"` +} + +// CleanupResult is the cleanup response payload. +type CleanupResult struct { + DeletedCount int64 `json:"deleted_count"` + Message string `json:"message"` +} + +// ListConfigVersions returns all config version summaries. +func ListConfigVersions(ctx context.Context) ([]*model.ConfigVersionSummary, error) { + return model.ListConfigVersionSummaries(ctx) +} + +// GetConfigVersionDetail returns a config version by id. +func GetConfigVersionDetail(ctx context.Context, id uint) (*model.ConfigVersion, error) { + return model.GetConfigVersionByID(ctx, id) +} + +// GetActiveConfigVersion returns the active config version. +func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) { + return model.GetActiveConfigVersion(ctx) +} + +// PreviewConfigVersion renders the current draft configuration. +func PreviewConfigVersion(ctx context.Context) (*ConfigPreviewResult, error) { + bundle, err := buildCurrentConfigBundle(ctx, false) + if err != nil { + return nil, err + } + return &ConfigPreviewResult{ + SnapshotJSON: bundle.SnapshotJSON, + MainConfig: bundle.MainConfig, + RouteConfig: bundle.RouteConfig, + RenderedConfig: bundle.RouteConfig, + SupportFiles: bundle.SupportFiles, + Checksum: bundle.Checksum, + RouteCount: len(bundle.Routes), + WebsiteCount: len(bundle.SnapshotRoutes), + }, nil +} + +// DiffConfigVersion compares the current draft against the active version. +func DiffConfigVersion(ctx context.Context) (*ConfigDiffResult, error) { + bundle, err := buildCurrentConfigBundle(ctx, false) + if err != nil { + return nil, err + } + result := &ConfigDiffResult{ + AddedSites: []string{}, + RemovedSites: []string{}, + ModifiedSites: []string{}, + AddedDomains: []string{}, + RemovedDomains: []string{}, + ModifiedDomains: []string{}, + ChangedOptionKeys: []string{}, + ChangedOptionDetails: []ConfigOptionDiffItem{}, + CurrentWebsiteCount: len(bundle.SnapshotRoutes), + } + activeVersion, err := model.GetActiveConfigVersion(ctx) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + for _, route := range bundle.SnapshotRoutes { + result.AddedSites = append(result.AddedSites, route.SiteName) + result.AddedDomains = append(result.AddedDomains, route.Domains...) + } + result.MainConfigChanged = true + result.ChangedOptionKeys = openRestyOptionKeys() + result.ChangedOptionDetails = buildInitialOpenRestyOptionDiffs(bundle.OpenRestyConfig) + sort.Strings(result.AddedSites) + sort.Strings(result.AddedDomains) + sort.Strings(result.ChangedOptionKeys) + return result, nil + } + return nil, err + } + result.ActiveVersion = activeVersion.Version + activeSnapshot, err := parseSnapshotDocument(activeVersion.SnapshotJSON) + if err != nil { + return nil, err + } + result.ActiveWebsiteCount = len(activeSnapshot.Routes) + currentSiteMap := flattenSnapshotRoutesBySite(bundle.SnapshotRoutes) + activeSiteMap := flattenSnapshotRoutesBySite(activeSnapshot.Routes) + for siteName, currentRoute := range currentSiteMap { + activeRoute, ok := activeSiteMap[siteName] + if !ok { + result.AddedSites = append(result.AddedSites, siteName) + continue + } + if !snapshotRouteConfigEqual(activeRoute, currentRoute) { + result.ModifiedSites = append(result.ModifiedSites, siteName) + } + } + for siteName := range activeSiteMap { + if _, ok := currentSiteMap[siteName]; !ok { + result.RemovedSites = append(result.RemovedSites, siteName) + } + } + currentMap := flattenSnapshotRoutesByDomain(bundle.SnapshotRoutes) + activeMap := flattenSnapshotRoutesByDomain(activeSnapshot.Routes) + for domain, currentRoute := range currentMap { + activeRoute, ok := activeMap[domain] + if !ok { + result.AddedDomains = append(result.AddedDomains, domain) + continue + } + if !snapshotRouteConfigEqual(activeRoute, currentRoute) { + result.ModifiedDomains = append(result.ModifiedDomains, domain) + } + } + for domain := range activeMap { + if _, ok := currentMap[domain]; !ok { + result.RemovedDomains = append(result.RemovedDomains, domain) + } + } + result.MainConfigChanged = activeVersion.MainConfig != bundle.MainConfig + result.WAFConfigChanged = !snapshotWAFConfigEqual(activeSnapshot.WAF, bundle.WAFSnapshot) + result.ChangedOptionDetails = diffOpenRestyOptionDetails(activeSnapshot.OpenRestyConfig, bundle.OpenRestyConfig) + result.ChangedOptionKeys = extractOptionDiffKeys(result.ChangedOptionDetails) + sort.Strings(result.AddedSites) + sort.Strings(result.RemovedSites) + sort.Strings(result.ModifiedSites) + sort.Strings(result.AddedDomains) + sort.Strings(result.RemovedDomains) + sort.Strings(result.ModifiedDomains) + sort.Strings(result.ChangedOptionKeys) + return result, nil +} + +// PublishConfigVersion publishes the current draft as a new active version. +func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*model.ConfigVersion, error) { + bundle, err := buildCurrentConfigBundle(ctx, true) + if err != nil { + return nil, err + } + if len(bundle.Routes) == 0 { + return nil, errors.New(errNoEnabledRoutes) + } + activeVersion, err := model.GetActiveConfigVersion(ctx) + if !force && err == nil && activeVersion.Checksum == bundle.Checksum { + return nil, errors.New(errNoChangesToPublish) + } + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + supportFilesJSON, err := json.Marshal(bundle.SupportFiles) + if err != nil { + return nil, err + } + version, err := nextVersionNumber(ctx, time.Now()) + if err != nil { + return nil, err + } + record := &model.ConfigVersion{ + Version: version, + SnapshotJSON: bundle.SnapshotJSON, + MainConfig: bundle.MainConfig, + RenderedConfig: bundle.RouteConfig, + SupportFilesJSON: string(supportFilesJSON), + Checksum: bundle.Checksum, + IsActive: true, + CreatedBy: createdBy, + } + if err = model.PublishConfigVersionTx(ctx, record); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errVersionConflict) + } + return nil, err + } + return record, nil +} + +// ActivateConfigVersion activates an existing config version. +func ActivateConfigVersion(ctx context.Context, id uint) (*model.ConfigVersion, error) { + version, err := model.GetConfigVersionByID(ctx, id) + if err != nil { + return nil, err + } + if err = model.ActivateConfigVersionTx(ctx, id); err != nil { + return nil, err + } + version.IsActive = true + return version, nil +} + +// CleanupConfigVersions removes old inactive config versions. +func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, error) { + if keepCount < 3 { + keepCount = 3 + } + versions, err := model.ListConfigVersionSummaries(ctx) + if err != nil { + return nil, err + } + if len(versions) <= keepCount { + return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil + } + var deleteIDs []uint + for index, version := range versions { + if index < keepCount { + continue + } + if version.IsActive { + continue + } + deleteIDs = append(deleteIDs, version.ID) + } + if len(deleteIDs) == 0 { + return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil + } + deletedCount, err := model.DeleteConfigVersionsByIDs(ctx, deleteIDs) + if err != nil { + return nil, err + } + return &CleanupResult{DeletedCount: deletedCount, Message: "清理成功"}, nil +} + +func nextVersionNumber(ctx context.Context, now time.Time) (string, error) { + prefix := now.Format("20060102") + latest, err := model.GetLatestConfigVersionByPrefix(ctx, prefix) + if errors.Is(err, gorm.ErrRecordNotFound) { + return fmt.Sprintf("%s-%03d", prefix, 1), nil + } + if err != nil { + return "", err + } + suffix := strings.TrimPrefix(latest, prefix+"-") + sequence, err := strconv.Atoi(suffix) + if err != nil { + return "", fmt.Errorf("invalid config version sequence %q: %w", latest, err) + } + return fmt.Sprintf("%s-%03d", prefix, sequence+1), nil +} + +func parseSnapshotDocument(snapshotJSON string) (*snapshotDocument, error) { + text := strings.TrimSpace(snapshotJSON) + if text == "" { + return &snapshotDocument{Routes: []snapshotRoute{}}, nil + } + if strings.HasPrefix(text, "[") { + var routes []snapshotRoute + if err := json.Unmarshal([]byte(text), &routes); err != nil { + return nil, errors.New(errInvalidSnapshotFormat) + } + return &snapshotDocument{Routes: normalizeSnapshotRoutes(routes)}, nil + } + var snapshot snapshotDocument + if err := json.Unmarshal([]byte(text), &snapshot); err != nil { + return nil, errors.New(errInvalidSnapshotFormat) + } + snapshot.Routes = normalizeSnapshotRoutes(snapshot.Routes) + return &snapshot, nil +} + +func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute { + if len(routes) == 0 { + return []snapshotRoute{} + } + for index := range routes { + normalizedDomains, err := decodeStoredDomains("", routes[index].Domain) + if len(routes[index].Domains) > 0 { + normalizedDomains, err = normalizeProxyRouteDomains(routes[index].Domains) + } + if err == nil && len(normalizedDomains) > 0 { + routes[index].Domains = normalizedDomains + routes[index].Domain = normalizedDomains[0] + routes[index].SiteName = normalizeProxyRouteSiteName(nil, routes[index].SiteName, normalizedDomains[0]) + } + normalizedCertIDs, primaryCertID, certErr := normalizeSnapshotCertificateIDs(routes[index].CertID, routes[index].CertIDs) + if certErr == nil { + routes[index].CertID = primaryCertID + routes[index].CertIDs = normalizedCertIDs + } + normalizedDomainCertIDs, domainCertErr := resolveDomainCertIDs(routes[index].Domains, routes[index].CertIDs, "") + if domainCertErr == nil && len(routes[index].DomainCertIDs) == 0 { + routes[index].DomainCertIDs = normalizedDomainCertIDs + } + normalizedUpstreams, upstreamErr := normalizeUpstreams(routes[index].OriginURL, routes[index].Upstreams) + if upstreamErr == nil { + routes[index].OriginURL = normalizedUpstreams[0] + routes[index].Upstreams = normalizedUpstreams + } + if !routes[index].BasicAuthEnabled { + routes[index].BasicAuthUsername = "" + routes[index].BasicAuthPassword = "" + } + routes[index].UpstreamType = normalizeUpstreamType(routes[index].UpstreamType) + } + return routes +} + +func flattenSnapshotRoutesBySite(routes []snapshotRoute) map[string]snapshotRoute { + siteMap := make(map[string]snapshotRoute) + for _, route := range normalizeSnapshotRoutes(routes) { + siteMap[route.SiteName] = route + } + return siteMap +} + +func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRoute { + domainMap := make(map[string]snapshotRoute) + for _, route := range normalizeSnapshotRoutes(routes) { + for _, domain := range route.Domains { + item := route + item.Domain = domain + domainMap[domain] = item + } + } + return domainMap +} + +func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool { + if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL || + left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || + left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP || + left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || + left.BasicAuthEnabled != right.BasicAuthEnabled || left.BasicAuthUsername != right.BasicAuthUsername || + left.BasicAuthPassword != right.BasicAuthPassword || left.UpstreamType != right.UpstreamType || + !uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) || left.TunnelTargetAddr != right.TunnelTargetAddr || + left.TunnelTargetProto != right.TunnelTargetProto || !uintPtrEqual(left.PagesProjectID, right.PagesProjectID) || + !uintSliceEqual(left.CertIDs, right.CertIDs) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) { + return false + } + if len(left.Domains) != len(right.Domains) { + return false + } + for index := range left.Domains { + if left.Domains[index] != right.Domains[index] { + return false + } + } + if len(left.Upstreams) != len(right.Upstreams) { + return false + } + for index := range left.Upstreams { + if left.Upstreams[index] != right.Upstreams[index] { + return false + } + } + if len(left.CacheRules) != len(right.CacheRules) { + return false + } + for index := range left.CacheRules { + if left.CacheRules[index] != right.CacheRules[index] { + return false + } + } + if len(left.CustomHeaders) != len(right.CustomHeaders) { + return false + } + for index := range left.CustomHeaders { + if left.CustomHeaders[index] != right.CustomHeaders[index] { + return false + } + } + return true +} + +func snapshotWAFConfigEqual(left snapshotWAFDocument, right snapshotWAFDocument) bool { + leftJSON, err := json.Marshal(left) + if err != nil { + return false + } + rightJSON, err := json.Marshal(right) + if err != nil { + return false + } + return string(leftJSON) == string(rightJSON) +} + +func normalizeSnapshotCertificateIDs(primaryCertID *uint, certIDs []uint) ([]uint, *uint, error) { + candidates := make([]uint, 0, len(certIDs)+1) + if primaryCertID != nil && *primaryCertID != 0 { + candidates = append(candidates, *primaryCertID) + } + candidates = append(candidates, certIDs...) + normalized := make([]uint, 0, len(candidates)) + seen := make(map[uint]struct{}, len(candidates)) + for _, certID := range candidates { + if certID == 0 { + continue + } + if _, ok := seen[certID]; ok { + continue + } + seen[certID] = struct{}{} + normalized = append(normalized, certID) + } + var normalizedPrimary *uint + if len(normalized) > 0 { + normalizedPrimary = &normalized[0] + } + return normalized, normalizedPrimary, nil +} + +func buildInitialOpenRestyOptionDiffs(current openRestyConfigSnapshot) []ConfigOptionDiffItem { + details := diffOpenRestyOptionDetails(openRestyConfigSnapshot{}, current) + for index := range details { + details[index].PreviousValue = "" + } + return details +} + +func diffOpenRestyOptionDetails(left openRestyConfigSnapshot, right openRestyConfigSnapshot) []ConfigOptionDiffItem { + changes := make([]ConfigOptionDiffItem, 0) + appendIfChanged := func(key string, previous string, current string) { + if previous == current { + return + } + changes = append(changes, ConfigOptionDiffItem{ + Key: key, + PreviousValue: previous, + CurrentValue: current, + }) + } + appendIfChanged("OpenRestyDefaultServerReturnStatus", fmt.Sprintf("%d", left.DefaultServerReturnStatus), fmt.Sprintf("%d", right.DefaultServerReturnStatus)) + appendIfChanged("OpenRestyWorkerProcesses", left.WorkerProcesses, right.WorkerProcesses) + appendIfChanged("OpenRestyWorkerConnections", fmt.Sprintf("%d", left.WorkerConnections), fmt.Sprintf("%d", right.WorkerConnections)) + appendIfChanged("OpenRestyWorkerRlimitNofile", fmt.Sprintf("%d", left.WorkerRlimitNofile), fmt.Sprintf("%d", right.WorkerRlimitNofile)) + appendIfChanged("OpenRestyEventsUse", left.EventsUse, right.EventsUse) + appendIfChanged("OpenRestyEventsMultiAcceptEnabled", fmt.Sprintf("%t", left.EventsMultiAcceptEnabled), fmt.Sprintf("%t", right.EventsMultiAcceptEnabled)) + appendIfChanged("OpenRestyKeepaliveTimeout", fmt.Sprintf("%d", left.KeepaliveTimeout), fmt.Sprintf("%d", right.KeepaliveTimeout)) + appendIfChanged("OpenRestyKeepaliveRequests", fmt.Sprintf("%d", left.KeepaliveRequests), fmt.Sprintf("%d", right.KeepaliveRequests)) + appendIfChanged("OpenRestyClientHeaderTimeout", fmt.Sprintf("%d", left.ClientHeaderTimeout), fmt.Sprintf("%d", right.ClientHeaderTimeout)) + appendIfChanged("OpenRestyClientBodyTimeout", fmt.Sprintf("%d", left.ClientBodyTimeout), fmt.Sprintf("%d", right.ClientBodyTimeout)) + appendIfChanged("OpenRestyClientMaxBodySize", left.ClientMaxBodySize, right.ClientMaxBodySize) + appendIfChanged("OpenRestyLargeClientHeaderBuffers", left.LargeClientHeaderBuffers, right.LargeClientHeaderBuffers) + appendIfChanged("OpenRestySendTimeout", fmt.Sprintf("%d", left.SendTimeout), fmt.Sprintf("%d", right.SendTimeout)) + appendIfChanged("OpenRestyProxyConnectTimeout", fmt.Sprintf("%d", left.ProxyConnectTimeout), fmt.Sprintf("%d", right.ProxyConnectTimeout)) + appendIfChanged("OpenRestyProxySendTimeout", fmt.Sprintf("%d", left.ProxySendTimeout), fmt.Sprintf("%d", right.ProxySendTimeout)) + appendIfChanged("OpenRestyProxyReadTimeout", fmt.Sprintf("%d", left.ProxyReadTimeout), fmt.Sprintf("%d", right.ProxyReadTimeout)) + appendIfChanged("OpenRestyWebsocketEnabled", fmt.Sprintf("%t", left.WebsocketEnabled), fmt.Sprintf("%t", right.WebsocketEnabled)) + appendIfChanged("OpenRestyHTTP3Enabled", fmt.Sprintf("%t", left.HTTP3Enabled), fmt.Sprintf("%t", right.HTTP3Enabled)) + appendIfChanged("OpenRestyProxyRequestBufferingEnabled", fmt.Sprintf("%t", left.ProxyRequestBuffering), fmt.Sprintf("%t", right.ProxyRequestBuffering)) + appendIfChanged("OpenRestyProxyBufferingEnabled", fmt.Sprintf("%t", left.ProxyBufferingEnabled), fmt.Sprintf("%t", right.ProxyBufferingEnabled)) + appendIfChanged("OpenRestyProxyBuffers", left.ProxyBuffers, right.ProxyBuffers) + appendIfChanged("OpenRestyProxyBufferSize", left.ProxyBufferSize, right.ProxyBufferSize) + appendIfChanged("OpenRestyProxyBusyBuffersSize", left.ProxyBusyBuffersSize, right.ProxyBusyBuffersSize) + appendIfChanged("OpenRestyGzipEnabled", fmt.Sprintf("%t", left.GzipEnabled), fmt.Sprintf("%t", right.GzipEnabled)) + appendIfChanged("OpenRestyGzipMinLength", fmt.Sprintf("%d", left.GzipMinLength), fmt.Sprintf("%d", right.GzipMinLength)) + appendIfChanged("OpenRestyGzipCompLevel", fmt.Sprintf("%d", left.GzipCompLevel), fmt.Sprintf("%d", right.GzipCompLevel)) + appendIfChanged("OpenRestyResolvers", left.Resolvers, right.Resolvers) + appendIfChanged("OpenRestyCacheEnabled", fmt.Sprintf("%t", left.CacheEnabled), fmt.Sprintf("%t", right.CacheEnabled)) + appendIfChanged("OpenRestyCachePath", left.CachePath, right.CachePath) + appendIfChanged("OpenRestyCacheLevels", left.CacheLevels, right.CacheLevels) + appendIfChanged("OpenRestyCacheInactive", left.CacheInactive, right.CacheInactive) + appendIfChanged("OpenRestyCacheMaxSize", left.CacheMaxSize, right.CacheMaxSize) + appendIfChanged("OpenRestyCacheKeyTemplate", left.CacheKeyTemplate, right.CacheKeyTemplate) + appendIfChanged("OpenRestyCacheLockEnabled", fmt.Sprintf("%t", left.CacheLockEnabled), fmt.Sprintf("%t", right.CacheLockEnabled)) + appendIfChanged("OpenRestyCacheLockTimeout", left.CacheLockTimeout, right.CacheLockTimeout) + appendIfChanged("OpenRestyCacheUseStale", left.CacheUseStale, right.CacheUseStale) + return changes +} + +func extractOptionDiffKeys(details []ConfigOptionDiffItem) []string { + keys := make([]string, 0, len(details)) + for _, item := range details { + keys = append(keys, item.Key) + } + return keys +} + +func openRestyOptionKeys() []string { + return []string{ + "OpenRestyDefaultServerReturnStatus", + "OpenRestyWorkerProcesses", + "OpenRestyWorkerConnections", + "OpenRestyWorkerRlimitNofile", + "OpenRestyEventsUse", + "OpenRestyEventsMultiAcceptEnabled", + "OpenRestyKeepaliveTimeout", + "OpenRestyKeepaliveRequests", + "OpenRestyClientHeaderTimeout", + "OpenRestyClientBodyTimeout", + "OpenRestyClientMaxBodySize", + "OpenRestyLargeClientHeaderBuffers", + "OpenRestySendTimeout", + "OpenRestyProxyConnectTimeout", + "OpenRestyProxySendTimeout", + "OpenRestyProxyReadTimeout", + "OpenRestyWebsocketEnabled", + "OpenRestyHTTP3Enabled", + "OpenRestyProxyRequestBufferingEnabled", + "OpenRestyProxyBufferingEnabled", + "OpenRestyProxyBuffers", + "OpenRestyProxyBufferSize", + "OpenRestyProxyBusyBuffersSize", + "OpenRestyGzipEnabled", + "OpenRestyGzipMinLength", + "OpenRestyGzipCompLevel", + "OpenRestyCacheEnabled", + "OpenRestyCachePath", + "OpenRestyCacheLevels", + "OpenRestyCacheInactive", + "OpenRestyCacheMaxSize", + "OpenRestyCacheKeyTemplate", + "OpenRestyCacheLockEnabled", + "OpenRestyCacheLockTimeout", + "OpenRestyCacheUseStale", + } +} diff --git a/Wavelet/internal/apps/openflare/config_version/logics_test.go b/Wavelet/internal/apps/openflare/config_version/logics_test.go new file mode 100644 index 00000000..6ab93049 --- /dev/null +++ b/Wavelet/internal/apps/openflare/config_version/logics_test.go @@ -0,0 +1,83 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package config_version + +import ( + "context" + "encoding/json" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupConfigVersionTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.ProxyRoute{}, + &model.ConfigVersion{}, + &model.OpenFlareWAFRuleGroup{}, + &model.OpenFlareWAFRuleGroupBinding{}, + &model.OpenFlareWAFIPGroup{}, + )) + + db.SetDB(sqliteDB) + return func() { + db.SetDB(nil) + } +} + +func TestPublishConfigVersionCreatesVersion(t *testing.T) { + cleanup := setupConfigVersionTestDB(t) + defer cleanup() + ctx := context.Background() + + route := &model.ProxyRoute{ + SiteName: "publish-site", + Domain: "publish.example.com", + Domains: `["publish.example.com"]`, + OriginURL: "http://origin.publish.example.com:8080", + Upstreams: `["http://origin.publish.example.com:8080"]`, + Enabled: true, + } + require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + + version, err := PublishConfigVersion(ctx, "tester", false) + require.NoError(t, err) + require.NotNil(t, version) + assert.NotZero(t, version.ID) + assert.True(t, version.IsActive) + assert.Equal(t, "tester", version.CreatedBy) + assert.NotEmpty(t, version.Version) + assert.NotEmpty(t, version.Checksum) + assert.NotEmpty(t, version.SnapshotJSON) + assert.NotEmpty(t, version.RenderedConfig) + + var snapshot snapshotDocument + require.NoError(t, json.Unmarshal([]byte(version.SnapshotJSON), &snapshot)) + require.Len(t, snapshot.Routes, 1) + assert.Equal(t, "publish-site", snapshot.Routes[0].SiteName) + assert.Equal(t, "publish.example.com", snapshot.Routes[0].Domain) + + active, err := GetActiveConfigVersion(ctx) + require.NoError(t, err) + assert.Equal(t, version.ID, active.ID) + + _, err = PublishConfigVersion(ctx, "tester", false) + require.Error(t, err) + assert.Contains(t, err.Error(), errNoChangesToPublish) + + forced, err := PublishConfigVersion(ctx, "tester", true) + require.NoError(t, err) + assert.NotEqual(t, version.ID, forced.ID) +} diff --git a/Wavelet/internal/apps/openflare/config_version/renderer.go b/Wavelet/internal/apps/openflare/config_version/renderer.go new file mode 100644 index 00000000..a5a70e2a --- /dev/null +++ b/Wavelet/internal/apps/openflare/config_version/renderer.go @@ -0,0 +1,53 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package config_version + +import ( + openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty" +) + +// SupportFile is a rendered configuration support artifact. +type SupportFile struct { + Path string `json:"path"` + Content string `json:"content"` +} + +func renderSnapshotConfig(sourceJSON string, certificateFiles []SupportFile) (*openrestyrender.Result, error) { + return openrestyrender.RenderJSON(sourceJSON, toOpenRestySupportFiles(certificateFiles)) +} + +func toOpenRestySupportFiles(files []SupportFile) []openrestyrender.SupportFile { + if len(files) == 0 { + return nil + } + result := make([]openrestyrender.SupportFile, 0, len(files)) + for _, file := range files { + result = append(result, openrestyrender.SupportFile{ + Path: file.Path, + Content: file.Content, + }) + } + return result +} + +func fromOpenRestySupportFiles(files []openrestyrender.SupportFile) []SupportFile { + if len(files) == 0 { + return nil + } + result := make([]SupportFile, 0, len(files)) + for _, file := range files { + result = append(result, SupportFile{ + Path: file.Path, + Content: file.Content, + }) + } + return result +} + +func renderPlaceholderConfig(snapshotJSON string) (mainConfig, routeConfig, checksum string) { + mainConfig = `{"placeholder":"main_config"}` + routeConfig = snapshotJSON + checksum = openrestyrender.ChecksumBundle(mainConfig, routeConfig, nil) + return mainConfig, routeConfig, checksum +} diff --git a/Wavelet/internal/apps/openflare/config_version/routers.go b/Wavelet/internal/apps/openflare/config_version/routers.go new file mode 100644 index 00000000..b9684b64 --- /dev/null +++ b/Wavelet/internal/apps/openflare/config_version/routers.go @@ -0,0 +1,115 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package config_version + +import ( + "errors" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +func handleLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, "记录不存在") + return true + } + compat.Fail(c, err.Error()) + return true +} + +// ListConfigVersionsHandler lists config versions. +func ListConfigVersionsHandler(c *gin.Context) { + versions, err := ListConfigVersions(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, versions) +} + +// GetConfigVersionHandler returns a config version by id. +func GetConfigVersionHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + version, err := GetConfigVersionDetail(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, version) +} + +// GetActiveConfigVersionHandler returns the active config version. +func GetActiveConfigVersionHandler(c *gin.Context) { + version, err := GetActiveConfigVersion(c.Request.Context()) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, errNoActiveVersion) + return + } + handleLogicError(c, err) + return + } + compat.OK(c, version) +} + +// PreviewConfigVersionHandler previews the current draft configuration. +func PreviewConfigVersionHandler(c *gin.Context) { + preview, err := PreviewConfigVersion(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, preview) +} + +// DiffConfigVersionHandler diffs the current draft against the active version. +func DiffConfigVersionHandler(c *gin.Context) { + diff, err := DiffConfigVersion(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, diff) +} + +// PublishConfigVersionHandler publishes a new config version. +func PublishConfigVersionHandler(c *gin.Context) { + username := c.GetString("username") + force := c.Query("force") == "true" + version, err := PublishConfigVersion(c.Request.Context(), username, force) + if handleLogicError(c, err) { + return + } + compat.OK(c, version) +} + +// ActivateConfigVersionHandler activates an existing config version. +func ActivateConfigVersionHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + version, err := ActivateConfigVersion(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, version) +} + +// CleanupConfigVersionsHandler removes old inactive config versions. +func CleanupConfigVersionsHandler(c *gin.Context) { + var input CleanupInput + if !compat.BindJSON(c, &input) { + return + } + result, err := CleanupConfigVersions(c.Request.Context(), input.KeepCount) + if handleLogicError(c, err) { + return + } + compat.OK(c, result) +} diff --git a/Wavelet/internal/apps/openflare/config_version/snapshot.go b/Wavelet/internal/apps/openflare/config_version/snapshot.go new file mode 100644 index 00000000..8909237b --- /dev/null +++ b/Wavelet/internal/apps/openflare/config_version/snapshot.go @@ -0,0 +1,510 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package config_version + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sort" + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" + "github.com/Rain-kl/Wavelet/internal/model" + openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty" +) + +type snapshotRoute struct { + ID uint `json:"id,omitempty"` + SiteName string `json:"site_name,omitempty"` + Domain string `json:"domain"` + Domains []string `json:"domains,omitempty"` + OriginURL string `json:"origin_url"` + OriginHost string `json:"origin_host,omitempty"` + Upstreams []string `json:"upstreams,omitempty"` + Enabled bool `json:"enabled"` + EnableHTTPS bool `json:"enable_https"` + CertID *uint `json:"cert_id,omitempty"` + CertIDs []uint `json:"cert_ids,omitempty"` + DomainCertIDs []uint `json:"domain_cert_ids,omitempty"` + RedirectHTTP bool `json:"redirect_http"` + LimitConnPerServer int `json:"limit_conn_per_server,omitempty"` + LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"` + LimitRate string `json:"limit_rate,omitempty"` + CacheEnabled bool `json:"cache_enabled"` + CachePolicy string `json:"cache_policy,omitempty"` + CacheRules []string `json:"cache_rules,omitempty"` + CustomHeaders []customHeaderInput `json:"custom_headers,omitempty"` + BasicAuthEnabled bool `json:"basic_auth_enabled,omitempty"` + BasicAuthUsername string `json:"basic_auth_username,omitempty"` + BasicAuthPassword string `json:"basic_auth_password,omitempty"` + Remark string `json:"remark,omitempty"` + UpstreamType string `json:"upstream_type,omitempty"` + TunnelNodeID *uint `json:"tunnel_node_id,omitempty"` + TunnelTargetAddr string `json:"tunnel_target_addr,omitempty"` + TunnelTargetProto string `json:"tunnel_target_protocol,omitempty"` + PagesProjectID *uint `json:"pages_project_id,omitempty"` +} + +type snapshotWAFRuleGroup struct { + ID uint `json:"id"` + Name string `json:"name"` + Enabled bool `json:"enabled"` + IsGlobal bool `json:"is_global"` + BlockStatusCode int `json:"block_status_code"` + BlockResponseBody string `json:"block_response_body,omitempty"` + IPWhitelist []string `json:"ip_whitelist,omitempty"` + IPBlacklist []string `json:"ip_blacklist,omitempty"` + IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"` + IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"` + CountryWhitelist []string `json:"country_whitelist,omitempty"` + CountryBlacklist []string `json:"country_blacklist,omitempty"` + RegionWhitelist []string `json:"region_whitelist,omitempty"` + RegionBlacklist []string `json:"region_blacklist,omitempty"` + PoWEnabled bool `json:"pow_enabled,omitempty"` + PoWConfig *openrestyrender.PoWConfig `json:"pow_config,omitempty"` +} + +type snapshotWAFIPGroup struct { + ID uint `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + Enabled bool `json:"enabled"` + IPList []string `json:"ip_list,omitempty"` +} + +type snapshotWAFBinding struct { + RouteID uint `json:"route_id"` + SiteName string `json:"site_name"` + RuleGroupIDs []uint `json:"rule_group_ids"` +} + +type snapshotWAFDocument struct { + RuleGroups []snapshotWAFRuleGroup `json:"rule_groups"` + IPGroups []snapshotWAFIPGroup `json:"ip_groups,omitempty"` + Bindings []snapshotWAFBinding `json:"bindings"` +} + +type openRestyConfigSnapshot struct { + DefaultServerReturnStatus int `json:"default_server_return_status"` + WorkerProcesses string `json:"worker_processes"` + WorkerConnections int `json:"worker_connections"` + WorkerRlimitNofile int `json:"worker_rlimit_nofile"` + EventsUse string `json:"events_use,omitempty"` + EventsMultiAcceptEnabled bool `json:"events_multi_accept_enabled"` + KeepaliveTimeout int `json:"keepalive_timeout"` + KeepaliveRequests int `json:"keepalive_requests"` + ClientHeaderTimeout int `json:"client_header_timeout"` + ClientBodyTimeout int `json:"client_body_timeout"` + ClientMaxBodySize string `json:"client_max_body_size"` + LargeClientHeaderBuffers string `json:"large_client_header_buffers"` + SendTimeout int `json:"send_timeout"` + ProxyConnectTimeout int `json:"proxy_connect_timeout"` + ProxySendTimeout int `json:"proxy_send_timeout"` + ProxyReadTimeout int `json:"proxy_read_timeout"` + WebsocketEnabled bool `json:"websocket_enabled"` + HTTP3Enabled bool `json:"http3_enabled"` + ProxyRequestBuffering bool `json:"proxy_request_buffering"` + ProxyBufferingEnabled bool `json:"proxy_buffering_enabled"` + ProxyBuffers string `json:"proxy_buffers"` + ProxyBufferSize string `json:"proxy_buffer_size"` + ProxyBusyBuffersSize string `json:"proxy_busy_buffers_size"` + GzipEnabled bool `json:"gzip_enabled"` + GzipMinLength int `json:"gzip_min_length"` + GzipCompLevel int `json:"gzip_comp_level"` + Resolvers string `json:"resolvers,omitempty"` + CacheEnabled bool `json:"cache_enabled"` + CachePath string `json:"cache_path,omitempty"` + CacheLevels string `json:"cache_levels"` + CacheInactive string `json:"cache_inactive"` + CacheMaxSize string `json:"cache_max_size"` + CacheKeyTemplate string `json:"cache_key_template"` + CacheLockEnabled bool `json:"cache_lock_enabled"` + CacheLockTimeout string `json:"cache_lock_timeout"` + CacheUseStale string `json:"cache_use_stale"` + MainConfigTemplate string `json:"main_config_template,omitempty"` +} + +type snapshotDocument struct { + Routes []snapshotRoute `json:"routes"` + OpenRestyConfig openRestyConfigSnapshot `json:"openresty_config"` + WAF snapshotWAFDocument `json:"waf"` +} + +type configBundle struct { + Routes []*model.ProxyRoute + SnapshotRoutes []snapshotRoute + WAFSnapshot snapshotWAFDocument + OpenRestyConfig openRestyConfigSnapshot + SnapshotJSON string + MainConfig string + RouteConfig string + SupportFiles []SupportFile + Checksum string + ChangedOptionKeys []string +} + +func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configBundle, error) { + routes, err := model.ListEnabledProxyRoutes(ctx) + if err != nil { + return nil, err + } + if requireRoutes && len(routes) == 0 { + return nil, errors.New(errNoEnabledRoutes) + } + snapshotRoutes, err := buildSnapshotRoutes(ctx, routes) + if err != nil { + return nil, err + } + wafSnapshot, err := buildSnapshotWAFDocument(ctx, routes) + if err != nil { + return nil, err + } + openRestyConfig := buildOpenRestyConfigSnapshot() + snapshotDoc := snapshotDocument{ + Routes: snapshotRoutes, + OpenRestyConfig: openRestyConfig, + WAF: wafSnapshot, + } + snapshotJSON, err := json.Marshal(snapshotDoc) + if err != nil { + return nil, err + } + certificateFiles, err := buildCertificateSupportFiles(ctx, snapshotRoutes) + if err != nil { + return nil, err + } + + mainConfig := "" + routeConfig := "" + checksum := "" + supportFiles := []SupportFile(nil) + + rendered, renderErr := renderSnapshotConfig(string(snapshotJSON), certificateFiles) + if renderErr == nil { + mainConfig = rendered.MainConfig + routeConfig = rendered.RouteConfig + checksum = rendered.Checksum + supportFiles = fromOpenRestySupportFiles(rendered.SupportFiles) + } else { + mainConfig, routeConfig, checksum = renderPlaceholderConfig(string(snapshotJSON)) + } + + return &configBundle{ + Routes: routes, + SnapshotRoutes: snapshotRoutes, + WAFSnapshot: wafSnapshot, + OpenRestyConfig: openRestyConfig, + SnapshotJSON: string(snapshotJSON), + MainConfig: mainConfig, + RouteConfig: routeConfig, + SupportFiles: supportFiles, + Checksum: checksum, + ChangedOptionKeys: openRestyOptionKeys(), + }, nil +} + +func buildSnapshotRoutes(ctx context.Context, routes []*model.ProxyRoute) ([]snapshotRoute, error) { + items := make([]snapshotRoute, 0, len(routes)) + for _, route := range routes { + domains, err := decodeStoredDomains(route.Domains, route.Domain) + if err != nil { + return nil, fmt.Errorf("route %s domains are invalid", route.Domain) + } + customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders) + if err != nil { + return nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain) + } + upstreamType := normalizeUpstreamType(route.UpstreamType) + originURL := route.OriginURL + upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL) + if err != nil { + return nil, fmt.Errorf("路由 %s 上游配置无效", route.Domain) + } + var tunnelNodeID *uint + var tunnelTargetAddr string + var tunnelTargetProtocol string + var pagesProjectID *uint + if upstreamType == "tunnel" { + originURL = resolveTunnelOpenRestyUpstreamURL(ctx) + upstreams = []string{originURL} + tunnelNodeID = route.TunnelNodeID + tunnelTargetAddr = strings.TrimSpace(route.TunnelTargetAddr) + tunnelTargetProtocol = normalizeTunnelTargetProtocol(route.TunnelTargetProtocol) + } else if upstreamType == "pages" { + return nil, fmt.Errorf("路由 %s Pages 配置无效: pages module is not available", route.Domain) + } + cacheRules, err := decodeStoredCacheRules(route.CacheRules) + if err != nil { + return nil, fmt.Errorf("路由 %s 缓存规则无效", route.Domain) + } + items = append(items, snapshotRoute{ + ID: route.ID, + SiteName: normalizeProxyRouteSiteName(route, route.SiteName, domains[0]), + Domain: domains[0], + Domains: domains, + OriginURL: originURL, + OriginHost: route.OriginHost, + Upstreams: upstreams, + Enabled: route.Enabled, + EnableHTTPS: route.EnableHTTPS, + CertID: route.CertID, + CertIDs: mustDecodeCertIDs(route), + DomainCertIDs: mustDecodeDomainCertIDs(route, domains), + RedirectHTTP: route.RedirectHTTP, + LimitConnPerServer: route.LimitConnPerServer, + LimitConnPerIP: route.LimitConnPerIP, + LimitRate: route.LimitRate, + CacheEnabled: route.CacheEnabled, + CachePolicy: route.CachePolicy, + CacheRules: cacheRules, + CustomHeaders: customHeaders, + BasicAuthEnabled: route.BasicAuthEnabled, + BasicAuthUsername: route.BasicAuthUsername, + BasicAuthPassword: route.BasicAuthPassword, + Remark: route.Remark, + UpstreamType: upstreamType, + TunnelNodeID: tunnelNodeID, + TunnelTargetAddr: tunnelTargetAddr, + TunnelTargetProto: tunnelTargetProtocol, + PagesProjectID: pagesProjectID, + }) + } + return items, nil +} + +func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (snapshotWAFDocument, error) { + if err := waf.EnsureDefaultRuleGroup(ctx); err != nil { + return snapshotWAFDocument{}, err + } + views, err := waf.ListRuleGroups(ctx) + if err != nil { + return snapshotWAFDocument{}, err + } + ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views)) + for _, view := range views { + if !view.Enabled { + continue + } + ruleGroups = append(ruleGroups, snapshotWAFRuleGroup{ + ID: view.ID, + Name: view.Name, + Enabled: view.Enabled, + IsGlobal: view.IsGlobal, + BlockStatusCode: view.BlockStatusCode, + BlockResponseBody: view.BlockResponseBody, + IPWhitelist: view.IPWhitelist, + IPBlacklist: view.IPBlacklist, + IPWhitelistGroups: view.IPWhitelistGroups, + IPBlacklistGroups: view.IPBlacklistGroups, + CountryWhitelist: view.CountryWhitelist, + CountryBlacklist: view.CountryBlacklist, + RegionWhitelist: view.RegionWhitelist, + RegionBlacklist: view.RegionBlacklist, + PoWEnabled: view.PoWEnabled, + PoWConfig: convertPoWConfig(view.PoWConfig), + }) + } + ipGroups, err := buildSnapshotWAFIPGroups(ctx, ruleGroups) + if err != nil { + return snapshotWAFDocument{}, err + } + enabledRouteIDs := make(map[uint]string, len(routes)) + for _, route := range routes { + if route == nil { + continue + } + siteName := strings.TrimSpace(route.SiteName) + if siteName == "" { + siteName = route.Domain + } + enabledRouteIDs[route.ID] = siteName + } + rawBindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx) + if err != nil { + return snapshotWAFDocument{}, err + } + groupIDsByRoute := make(map[uint][]uint, len(rawBindings)) + for _, binding := range rawBindings { + if _, ok := enabledRouteIDs[binding.ProxyRouteID]; !ok { + continue + } + groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID) + } + bindings := make([]snapshotWAFBinding, 0, len(groupIDsByRoute)) + for routeID, groupIDs := range groupIDsByRoute { + sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] }) + bindings = append(bindings, snapshotWAFBinding{ + RouteID: routeID, + SiteName: enabledRouteIDs[routeID], + RuleGroupIDs: groupIDs, + }) + } + sort.Slice(bindings, func(i, j int) bool { + if bindings[i].SiteName == bindings[j].SiteName { + return bindings[i].RouteID < bindings[j].RouteID + } + return bindings[i].SiteName < bindings[j].SiteName + }) + return snapshotWAFDocument{RuleGroups: ruleGroups, IPGroups: ipGroups, Bindings: bindings}, nil +} + +func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleGroup) ([]snapshotWAFIPGroup, error) { + idSet := make(map[uint]struct{}) + for _, group := range ruleGroups { + for _, id := range group.IPWhitelistGroups { + idSet[id] = struct{}{} + } + for _, id := range group.IPBlacklistGroups { + idSet[id] = struct{}{} + } + } + if len(idSet) == 0 { + return []snapshotWAFIPGroup{}, nil + } + ids := make([]uint, 0, len(idSet)) + for id := range idSet { + ids = append(ids, id) + } + sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) + groups, err := listWAFIPGroupsByIDs(ctx, ids) + if err != nil { + return nil, err + } + groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups)) + for _, group := range groups { + groupByID[group.ID] = group + } + snapshots := make([]snapshotWAFIPGroup, 0, len(ids)) + for _, id := range ids { + group := groupByID[id] + if group == nil { + return nil, fmt.Errorf("IP 组 %d 不存在", id) + } + ipList, decodeErr := decodeIPList(group.IPList) + if decodeErr != nil { + return nil, decodeErr + } + snapshots = append(snapshots, snapshotWAFIPGroup{ + ID: group.ID, + Name: group.Name, + Type: group.Type, + Enabled: group.Enabled, + IPList: ipList, + }) + } + return snapshots, nil +} + +func decodeIPList(raw string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []string{}, nil + } + var items []string + if err := json.Unmarshal([]byte(text), &items); err != nil { + return nil, fmt.Errorf("ip_list payload is invalid") + } + return items, nil +} + +func convertPoWConfig(config *waf.PoWConfig) *openrestyrender.PoWConfig { + if config == nil { + return nil + } + return &openrestyrender.PoWConfig{ + Difficulty: config.Difficulty, + Algorithm: config.Algorithm, + SessionTTL: config.SessionTTL, + ChallengeTTL: config.ChallengeTTL, + Whitelist: openrestyrender.PoWListConfig{ + IPs: config.Whitelist.IPs, + IPCidrs: config.Whitelist.IPCidrs, + Paths: config.Whitelist.Paths, + PathRegexes: config.Whitelist.PathRegexes, + UserAgents: config.Whitelist.UserAgents, + }, + Blacklist: openrestyrender.PoWListConfig{ + IPs: config.Blacklist.IPs, + IPCidrs: config.Blacklist.IPCidrs, + Paths: config.Blacklist.Paths, + PathRegexes: config.Blacklist.PathRegexes, + UserAgents: config.Blacklist.UserAgents, + }, + } +} + +func buildOpenRestyConfigSnapshot() openRestyConfigSnapshot { + return openRestyConfigSnapshot{ + DefaultServerReturnStatus: model.OpenRestyDefaultServerReturnStatus, + WorkerProcesses: model.OpenRestyWorkerProcesses, + WorkerConnections: model.OpenRestyWorkerConnections, + WorkerRlimitNofile: model.OpenRestyWorkerRlimitNofile, + EventsUse: model.OpenRestyEventsUse, + EventsMultiAcceptEnabled: model.OpenRestyEventsMultiAcceptEnabled, + KeepaliveTimeout: model.OpenRestyKeepaliveTimeout, + KeepaliveRequests: model.OpenRestyKeepaliveRequests, + ClientHeaderTimeout: model.OpenRestyClientHeaderTimeout, + ClientBodyTimeout: model.OpenRestyClientBodyTimeout, + ClientMaxBodySize: model.OpenRestyClientMaxBodySize, + LargeClientHeaderBuffers: model.OpenRestyLargeClientHeaderBuffers, + SendTimeout: model.OpenRestySendTimeout, + ProxyConnectTimeout: model.OpenRestyProxyConnectTimeout, + ProxySendTimeout: model.OpenRestyProxySendTimeout, + ProxyReadTimeout: model.OpenRestyProxyReadTimeout, + WebsocketEnabled: model.OpenRestyWebsocketEnabled, + HTTP3Enabled: model.OpenRestyHTTP3Enabled, + ProxyRequestBuffering: model.OpenRestyProxyRequestBufferingEnabled, + ProxyBufferingEnabled: model.OpenRestyProxyBufferingEnabled, + ProxyBuffers: model.OpenRestyProxyBuffers, + ProxyBufferSize: model.OpenRestyProxyBufferSize, + ProxyBusyBuffersSize: model.OpenRestyProxyBusyBuffersSize, + GzipEnabled: model.OpenRestyGzipEnabled, + GzipMinLength: model.OpenRestyGzipMinLength, + GzipCompLevel: model.OpenRestyGzipCompLevel, + Resolvers: model.OpenRestyResolvers, + CacheEnabled: model.OpenRestyCacheEnabled, + CachePath: model.OpenRestyCachePath, + CacheLevels: model.OpenRestyCacheLevels, + CacheInactive: model.OpenRestyCacheInactive, + CacheMaxSize: model.OpenRestyCacheMaxSize, + CacheKeyTemplate: model.OpenRestyCacheKeyTemplate, + CacheLockEnabled: model.OpenRestyCacheLockEnabled, + CacheLockTimeout: model.OpenRestyCacheLockTimeout, + CacheUseStale: model.OpenRestyCacheUseStale, + MainConfigTemplate: model.OpenRestyMainConfigTemplate, + } +} + +func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) ([]SupportFile, error) { + certIDSet := make(map[uint]struct{}) + for _, route := range routes { + for _, certID := range route.CertIDs { + if certID != 0 { + certIDSet[certID] = struct{}{} + } + } + } + if len(certIDSet) == 0 { + return nil, nil + } + certIDs := make([]uint, 0, len(certIDSet)) + for certID := range certIDSet { + certIDs = append(certIDs, certID) + } + sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] }) + files := make([]SupportFile, 0, len(certIDs)*2) + for _, certID := range certIDs { + certificate, err := model.GetTLSCertificateByID(ctx, certID) + if err != nil { + return nil, err + } + files = append(files, + SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)}, + SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)}, + ) + } + return dedupeSupportFiles(files), nil +} diff --git a/Wavelet/internal/apps/openflare/dashboard/helpers.go b/Wavelet/internal/apps/openflare/dashboard/helpers.go new file mode 100644 index 00000000..befa38a6 --- /dev/null +++ b/Wavelet/internal/apps/openflare/dashboard/helpers.go @@ -0,0 +1,36 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dashboard + +import ( + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + nodeStatusOnline = "online" + nodeStatusOffline = "offline" + nodeStatusPending = "pending" +) + +func computeNodeStatus(node *model.OpenFlareNode) string { + if node == nil { + return nodeStatusOffline + } + if node.LastSeenAt == nil || node.LastSeenAt.IsZero() { + return nodeStatusPending + } + if time.Since(*node.LastSeenAt) > model.NodeOfflineThreshold { + return nodeStatusOffline + } + return nodeStatusOnline +} + +func nodeViewLastSeenAt(node *model.OpenFlareNode) any { + if node == nil || node.LastSeenAt == nil { + return time.Time{} + } + return *node.LastSeenAt +} diff --git a/Wavelet/internal/apps/openflare/dashboard/logics.go b/Wavelet/internal/apps/openflare/dashboard/logics.go new file mode 100644 index 00000000..9ee857b3 --- /dev/null +++ b/Wavelet/internal/apps/openflare/dashboard/logics.go @@ -0,0 +1,366 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dashboard + +import ( + "context" + "sort" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/observability" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// Summary is the dashboard node summary section. +type Summary struct { + TotalNodes int `json:"total_nodes"` + OnlineNodes int `json:"online_nodes"` + OfflineNodes int `json:"offline_nodes"` + PendingNodes int `json:"pending_nodes"` + UnhealthyNodes int `json:"unhealthy_nodes"` +} + +// Traffic is the dashboard traffic section. +type Traffic struct { + RequestCount int64 `json:"request_count"` + UniqueVisitors int64 `json:"unique_visitors"` + ErrorCount int64 `json:"error_count"` + EstimatedQPS float64 `json:"estimated_qps"` + ReportedNodes int `json:"reported_nodes"` +} + +// Capacity is the dashboard capacity section. +type Capacity struct { + AverageCPUUsagePercent float64 `json:"average_cpu_usage_percent"` + AverageMemoryUsagePercent float64 `json:"average_memory_usage_percent"` + HighCPUNodes int `json:"high_cpu_nodes"` + HighMemoryNodes int `json:"high_memory_nodes"` + HighStorageNodes int `json:"high_storage_nodes"` +} + +// NodeHealth is a dashboard node health row. +type NodeHealth struct { + ID uint `json:"id"` + NodeID string `json:"node_id"` + Name string `json:"name"` + GeoName string `json:"geo_name"` + GeoLatitude *float64 `json:"geo_latitude"` + GeoLongitude *float64 `json:"geo_longitude"` + Status string `json:"status"` + OpenrestyStatus string `json:"openresty_status"` + CurrentVersion string `json:"current_version"` + LastSeenAt any `json:"last_seen_at"` + ActiveEventCount int `json:"active_event_count"` + CPUUsagePercent float64 `json:"cpu_usage_percent"` + MemoryUsagePercent float64 `json:"memory_usage_percent"` + StorageUsagePercent float64 `json:"storage_usage_percent"` + RequestCount int64 `json:"request_count"` + ErrorCount int64 `json:"error_count"` + UniqueVisitorCount int64 `json:"unique_visitor_count"` +} + +// OverviewView is the expanded dashboard overview payload. +type OverviewView struct { + GeneratedAt time.Time `json:"generated_at"` + Summary Summary `json:"summary"` + Traffic Traffic `json:"traffic"` + Capacity Capacity `json:"capacity"` + Distributions observability.TrafficDistributions `json:"distributions"` + Trends observability.NodeTrends `json:"trends"` + Nodes []NodeHealth `json:"nodes"` +} + +// OverviewPayload is the compact legacy dashboard overview response. +type OverviewPayload struct { + GeneratedAt any `json:"generated_at"` + Summary Summary `json:"summary"` + Traffic Traffic `json:"traffic"` + Capacity Capacity `json:"capacity"` + Distributions distributionsPayload `json:"distributions"` + Trends trendsPayload `json:"trends"` + Nodes [][]any `json:"nodes"` +} + +type distributionsPayload struct { + StatusCodes [][]any `json:"status_codes"` + TopDomains [][]any `json:"top_domains"` + SourceCountries [][]any `json:"source_countries"` +} + +type trendsPayload struct { + Traffic24h [][]any `json:"traffic_24h"` + Capacity24h [][]any `json:"capacity_24h"` + Network24h [][]any `json:"network_24h"` + DiskIO24h [][]any `json:"disk_io_24h"` +} + +// GetOverview aggregates dashboard overview data from nodes and observability tables. +func GetOverview(ctx context.Context) (*OverviewPayload, error) { + view, err := buildOverviewView(ctx) + if err != nil { + return nil, err + } + return compressOverview(view), nil +} + +func buildOverviewView(ctx context.Context) (*OverviewView, error) { + now := time.Now() + since := now.Add(-24 * time.Hour) + + nodes, err := model.ListOpenFlareNodes(ctx) + if err != nil { + return nil, err + } + snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", since, 0) + if err != nil { + return nil, err + } + reports, err := model.ListOpenFlareRequestReportsSince(ctx, "", since, 0) + if err != nil { + return nil, err + } + accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, 8) + if err != nil { + return nil, err + } + activeEvents, err := model.ListOpenFlareActiveHealthEvents(ctx) + if err != nil { + return nil, err + } + openrestySnapshots, err := model.ListOpenFlareNodeObservationOpenresty(ctx, "", since, 0) + if err != nil { + return nil, err + } + + view := &OverviewView{ + GeneratedAt: now, + Nodes: make([]NodeHealth, 0, len(nodes)), + Distributions: observability.BuildTrafficDistributions(reports, accessLogRegions, 8), + Trends: observability.NodeTrends{ + Traffic24h: observability.BuildTrafficTrendPoints(now, reports), + Capacity24h: observability.BuildCapacityTrendPoints(now, snapshots), + Network24h: observability.BuildNetworkTrendPoints(now, snapshots, openrestySnapshots), + DiskIO24h: observability.BuildDiskIOTrendPoints(now, snapshots), + }, + } + + var cpuNodeCount int + var memoryNodeCount int + latestSnapshots := observability.LatestMetricSnapshotsByNode(snapshots) + latestTrafficReports := observability.LatestTrafficReportsByNode(reports) + activeEventsByNode := observability.ActiveHealthEventsByNode(activeEvents) + + for _, node := range nodes { + computedStatus := computeNodeStatus(&node) + switch computedStatus { + case nodeStatusOnline: + view.Summary.OnlineNodes++ + case nodeStatusOffline: + view.Summary.OfflineNodes++ + case nodeStatusPending: + view.Summary.PendingNodes++ + } + if node.OpenrestyStatus == "unhealthy" { + view.Summary.UnhealthyNodes++ + } + + latestSnapshot := latestSnapshots[node.NodeID] + latestTraffic := latestTrafficReports[node.NodeID] + nodeActiveEvents := activeEventsByNode[node.NodeID] + + nodeHealth := NodeHealth{ + ID: node.ID, + NodeID: node.NodeID, + Name: node.Name, + GeoName: node.GeoName, + GeoLatitude: node.GeoLatitude, + GeoLongitude: node.GeoLongitude, + Status: computedStatus, + OpenrestyStatus: node.OpenrestyStatus, + CurrentVersion: node.CurrentVersion, + LastSeenAt: nodeViewLastSeenAt(&node), + ActiveEventCount: len(nodeActiveEvents), + } + + if latestSnapshot != nil { + nodeHealth.CPUUsagePercent = latestSnapshot.CPUUsagePercent + nodeHealth.MemoryUsagePercent = observability.Percentage(latestSnapshot.MemoryUsedBytes, latestSnapshot.MemoryTotalBytes) + nodeHealth.StorageUsagePercent = observability.Percentage(latestSnapshot.StorageUsedBytes, latestSnapshot.StorageTotalBytes) + if latestSnapshot.CPUUsagePercent > 0 { + view.Capacity.AverageCPUUsagePercent += latestSnapshot.CPUUsagePercent + cpuNodeCount++ + } + if nodeHealth.MemoryUsagePercent > 0 { + view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent + memoryNodeCount++ + } + if latestSnapshot.CPUUsagePercent >= 80 { + view.Capacity.HighCPUNodes++ + } + if nodeHealth.MemoryUsagePercent >= 85 { + view.Capacity.HighMemoryNodes++ + } + if nodeHealth.StorageUsagePercent >= 85 { + view.Capacity.HighStorageNodes++ + } + } + + if latestTraffic != nil { + nodeHealth.RequestCount = latestTraffic.RequestCount + nodeHealth.ErrorCount = latestTraffic.ErrorCount + nodeHealth.UniqueVisitorCount = latestTraffic.UniqueVisitorCount + view.Traffic.RequestCount += latestTraffic.RequestCount + view.Traffic.UniqueVisitors += latestTraffic.UniqueVisitorCount + view.Traffic.ErrorCount += latestTraffic.ErrorCount + if duration := latestTraffic.WindowEndedAt.Sub(latestTraffic.WindowStartedAt).Seconds(); duration > 0 { + view.Traffic.EstimatedQPS += float64(latestTraffic.RequestCount) / duration + } + view.Traffic.ReportedNodes++ + } + + view.Nodes = append(view.Nodes, nodeHealth) + } + + view.Summary.TotalNodes = len(nodes) + if cpuNodeCount > 0 { + view.Capacity.AverageCPUUsagePercent /= float64(cpuNodeCount) + } + if memoryNodeCount > 0 { + view.Capacity.AverageMemoryUsagePercent /= float64(memoryNodeCount) + } + + sort.Slice(view.Nodes, func(i int, j int) bool { + if view.Nodes[i].ActiveEventCount == view.Nodes[j].ActiveEventCount { + return view.Nodes[i].CPUUsagePercent > view.Nodes[j].CPUUsagePercent + } + return view.Nodes[i].ActiveEventCount > view.Nodes[j].ActiveEventCount + }) + + return view, nil +} + +func compressOverview(view *OverviewView) *OverviewPayload { + if view == nil { + return &OverviewPayload{ + Distributions: distributionsPayload{ + StatusCodes: [][]any{}, + TopDomains: [][]any{}, + SourceCountries: [][]any{}, + }, + Trends: trendsPayload{ + Traffic24h: [][]any{}, + Capacity24h: [][]any{}, + Network24h: [][]any{}, + DiskIO24h: [][]any{}, + }, + Nodes: [][]any{}, + } + } + return &OverviewPayload{ + GeneratedAt: view.GeneratedAt, + Summary: view.Summary, + Traffic: view.Traffic, + Capacity: view.Capacity, + Distributions: distributionsPayload{ + StatusCodes: compressDistributionItems(view.Distributions.StatusCodes), + TopDomains: compressDistributionItems(view.Distributions.TopDomains), + SourceCountries: compressDistributionItems(view.Distributions.SourceCountries), + }, + Trends: trendsPayload{ + Traffic24h: compressTrafficTrendPoints(view.Trends.Traffic24h), + Capacity24h: compressCapacityTrendPoints(view.Trends.Capacity24h), + Network24h: compressNetworkTrendPoints(view.Trends.Network24h), + DiskIO24h: compressDiskIOTrendPoints(view.Trends.DiskIO24h), + }, + Nodes: compressDashboardNodes(view.Nodes), + } +} + +func compressDistributionItems(items []observability.DistributionItem) [][]any { + rows := make([][]any, 0, len(items)) + for _, item := range items { + rows = append(rows, []any{item.Key, item.Value}) + } + return rows +} + +func compressTrafficTrendPoints(points []observability.TrafficTrendPoint) [][]any { + rows := make([][]any, 0, len(points)) + for _, point := range points { + rows = append(rows, []any{ + point.BucketStartedAt, + point.RequestCount, + point.ErrorCount, + point.UniqueVisitorCount, + }) + } + return rows +} + +func compressCapacityTrendPoints(points []observability.CapacityTrendPoint) [][]any { + rows := make([][]any, 0, len(points)) + for _, point := range points { + rows = append(rows, []any{ + point.BucketStartedAt, + point.AverageCPUUsagePercent, + point.AverageMemoryUsagePercent, + point.ReportedNodes, + }) + } + return rows +} + +func compressNetworkTrendPoints(points []observability.NetworkTrendPoint) [][]any { + rows := make([][]any, 0, len(points)) + for _, point := range points { + rows = append(rows, []any{ + point.BucketStartedAt, + point.NetworkRxBytes, + point.NetworkTxBytes, + point.OpenrestyRxBytes, + point.OpenrestyTxBytes, + point.ReportedNodes, + }) + } + return rows +} + +func compressDiskIOTrendPoints(points []observability.DiskIOTrendPoint) [][]any { + rows := make([][]any, 0, len(points)) + for _, point := range points { + rows = append(rows, []any{ + point.BucketStartedAt, + point.DiskReadBytes, + point.DiskWriteBytes, + point.ReportedNodes, + }) + } + return rows +} + +func compressDashboardNodes(nodes []NodeHealth) [][]any { + rows := make([][]any, 0, len(nodes)) + for _, node := range nodes { + rows = append(rows, []any{ + node.ID, + node.NodeID, + node.Name, + node.GeoName, + node.GeoLatitude, + node.GeoLongitude, + node.Status, + node.OpenrestyStatus, + node.CurrentVersion, + node.LastSeenAt, + node.ActiveEventCount, + node.CPUUsagePercent, + node.MemoryUsagePercent, + node.StorageUsagePercent, + node.RequestCount, + node.ErrorCount, + node.UniqueVisitorCount, + }) + } + return rows +} diff --git a/Wavelet/internal/apps/openflare/dashboard/logics_test.go b/Wavelet/internal/apps/openflare/dashboard/logics_test.go new file mode 100644 index 00000000..9c8a1ed1 --- /dev/null +++ b/Wavelet/internal/apps/openflare/dashboard/logics_test.go @@ -0,0 +1,125 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dashboard + +import ( + "context" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupDashboardTestDB(t *testing.T) func() { + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) + + db.SetDB(sqliteDB) + return func() { + db.SetDB(nil) + } +} + +func TestGetOverviewStructure(t *testing.T) { + cleanup := setupDashboardTestDB(t) + defer cleanup() + + ctx := context.Background() + now := time.Now().UTC() + lastSeen := now.Add(-time.Minute) + + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + NodeID: "node-dashboard-1", + Name: "Edge 1", + IP: "10.0.0.1", + Status: "online", + OpenrestyStatus: "healthy", + CurrentVersion: "v1.0.0", + LastSeenAt: &lastSeen, + }).Error) + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + NodeID: "node-dashboard-2", + Name: "Edge 2", + IP: "10.0.0.2", + Status: "pending", + OpenrestyStatus: "unknown", + }).Error) + + overview, err := GetOverview(ctx) + require.NoError(t, err) + require.NotNil(t, overview) + + assert.False(t, overview.GeneratedAt.(time.Time).IsZero()) + assert.Equal(t, 2, overview.Summary.TotalNodes) + assert.Equal(t, 1, overview.Summary.OnlineNodes) + assert.Equal(t, 1, overview.Summary.PendingNodes) + assert.Equal(t, 0, overview.Summary.OfflineNodes) + assert.Equal(t, 0, overview.Summary.UnhealthyNodes) + + assert.Equal(t, int64(0), overview.Traffic.RequestCount) + assert.Equal(t, int64(0), overview.Traffic.UniqueVisitors) + assert.Equal(t, int64(0), overview.Traffic.ErrorCount) + assert.Equal(t, float64(0), overview.Traffic.EstimatedQPS) + assert.Equal(t, 0, overview.Traffic.ReportedNodes) + + assert.Equal(t, float64(0), overview.Capacity.AverageCPUUsagePercent) + assert.Equal(t, float64(0), overview.Capacity.AverageMemoryUsagePercent) + assert.Equal(t, 0, overview.Capacity.HighCPUNodes) + assert.Equal(t, 0, overview.Capacity.HighMemoryNodes) + assert.Equal(t, 0, overview.Capacity.HighStorageNodes) + + require.NotNil(t, overview.Distributions.StatusCodes) + require.NotNil(t, overview.Distributions.TopDomains) + require.NotNil(t, overview.Distributions.SourceCountries) + assert.Empty(t, overview.Distributions.StatusCodes) + assert.Empty(t, overview.Distributions.TopDomains) + assert.Empty(t, overview.Distributions.SourceCountries) + + require.Len(t, overview.Trends.Traffic24h, 24) + require.Len(t, overview.Trends.Capacity24h, 24) + require.Len(t, overview.Trends.Network24h, 24) + require.Len(t, overview.Trends.DiskIO24h, 24) + for _, row := range overview.Trends.Traffic24h { + require.Len(t, row, 4) + } + for _, row := range overview.Trends.Capacity24h { + require.Len(t, row, 4) + } + for _, row := range overview.Trends.Network24h { + require.Len(t, row, 6) + } + for _, row := range overview.Trends.DiskIO24h { + require.Len(t, row, 4) + } + + require.Len(t, overview.Nodes, 2) + for _, row := range overview.Nodes { + require.Len(t, row, 17) + } + + nodeByID := make(map[string][]any, len(overview.Nodes)) + for _, row := range overview.Nodes { + nodeByID[row[1].(string)] = row + } + + onlineNode := nodeByID["node-dashboard-1"] + require.NotNil(t, onlineNode) + assert.Equal(t, "Edge 1", onlineNode[2]) + assert.Equal(t, "online", onlineNode[6]) + assert.Equal(t, "healthy", onlineNode[7]) + + pendingNode := nodeByID["node-dashboard-2"] + require.NotNil(t, pendingNode) + assert.Equal(t, "Edge 2", pendingNode[2]) + assert.Equal(t, "pending", pendingNode[6]) + assert.Equal(t, "unknown", pendingNode[7]) +} diff --git a/Wavelet/internal/apps/openflare/dashboard/routers.go b/Wavelet/internal/apps/openflare/dashboard/routers.go new file mode 100644 index 00000000..30f7cf65 --- /dev/null +++ b/Wavelet/internal/apps/openflare/dashboard/routers.go @@ -0,0 +1,27 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dashboard + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +// RegisterRoutes mounts legacy OpenFlare dashboard routes. +func RegisterRoutes(apiGroup *gin.RouterGroup) { + dashboardRoute := apiGroup.Group("/dashboard") + dashboardRoute.Use(compat.AdminAuth()) + { + dashboardRoute.GET("/overview", getOverviewHandler) + } +} + +func getOverviewHandler(c *gin.Context) { + overview, err := GetOverview(c.Request.Context()) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, overview) +} diff --git a/Wavelet/internal/apps/openflare/flared/errs.go b/Wavelet/internal/apps/openflare/flared/errs.go new file mode 100644 index 00000000..84a54a43 --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/errs.go @@ -0,0 +1,9 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +const ( + errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效" + errTunnelNodeTypeMismatch = "此节点不是 TunnelClient 类型" +) diff --git a/Wavelet/internal/apps/openflare/flared/helpers.go b/Wavelet/internal/apps/openflare/flared/helpers.go new file mode 100644 index 00000000..17ce7a3a --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/helpers.go @@ -0,0 +1,176 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net" + "strconv" + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/relay" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +type configVersionRow struct { + Version string `gorm:"column:version"` + Checksum string `gorm:"column:checksum"` +} + +func (configVersionRow) TableName() string { + return "of_config_versions" +} + +func normalizeReleaseChannel(channel string) string { + if strings.ToLower(strings.TrimSpace(channel)) == "preview" { + return "preview" + } + return "stable" +} + +func normalizeFlaredHeartbeatPayload(payload HeartbeatPayload) HeartbeatPayload { + payload.ClientVersion = strings.TrimSpace(payload.ClientVersion) + payload.FrpVersion = strings.TrimSpace(payload.FrpVersion) + payload.IP = strings.TrimSpace(payload.IP) + payload.TunnelStatus = strings.ToLower(strings.TrimSpace(payload.TunnelStatus)) + payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion) + payload.CurrentChecksum = strings.TrimSpace(payload.CurrentChecksum) + + cleaned := make([]ConnectedRelay, 0, len(payload.ConnectedRelays)) + for _, item := range payload.ConnectedRelays { + item.RelayNodeID = strings.TrimSpace(item.RelayNodeID) + item.Status = strings.ToLower(strings.TrimSpace(item.Status)) + if item.RelayNodeID == "" { + continue + } + if item.Status == "" { + item.Status = "unknown" + } + cleaned = append(cleaned, item) + } + payload.ConnectedRelays = cleaned + return payload +} + +func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New("database not initialized") + } + if !conn.Migrator().HasTable(&configVersionRow{}) { + return nil, gorm.ErrRecordNotFound + } + + var version configVersionRow + err := conn.Where("is_active = ?", true).Order("id desc").First(&version).Error + if err != nil { + return nil, err + } + return &ActiveConfigMeta{ + Version: version.Version, + Checksum: version.Checksum, + }, nil +} + +func listTunnelRelayNodes(ctx context.Context) ([]model.OpenFlareNode, error) { + nodes, err := model.ListOpenFlareNodes(ctx) + if err != nil { + return nil, err + } + relays := make([]model.OpenFlareNode, 0) + for _, node := range nodes { + if node.NodeType == "tunnel_relay" { + relays = append(relays, node) + } + } + return relays, nil +} + +func relayClientAddress(node *model.OpenFlareNode) string { + if node == nil { + return "" + } + port := node.RelayBindPort + if port <= 0 { + port = 7000 + } + addr := strings.TrimSpace(node.RelayClientAccessAddr) + if addr == "" { + addr = strings.TrimSpace(node.IP) + } + if addr == "" { + return fmt.Sprintf("127.0.0.1:%d", port) + } + if _, _, err := net.SplitHostPort(addr); err == nil { + return addr + } + if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 { + return net.JoinHostPort(addr, strconv.Itoa(port)) + } + return fmt.Sprintf("%s:%d", addr, port) +} + +func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + domain := strings.ToLower(strings.TrimSpace(fallbackDomain)) + if domain == "" { + return nil, errors.New("domain is required") + } + return []string{domain}, nil + } + var domains []string + if err := json.Unmarshal([]byte(text), &domains); err != nil { + return nil, errors.New("domains payload is invalid") + } + normalized := make([]string, 0, len(domains)) + for _, item := range domains { + domain := strings.ToLower(strings.TrimSpace(item)) + if domain == "" { + continue + } + normalized = append(normalized, domain) + } + if len(normalized) == 0 { + return nil, errors.New("domain is required") + } + return normalized, nil +} + +func parseTunnelTargetAddr(addr string) (string, int) { + addr = strings.TrimSpace(addr) + if addr == "" { + return "127.0.0.1", 80 + } + host, portStr, err := net.SplitHostPort(addr) + if err != nil { + lastColon := strings.LastIndex(addr, ":") + if lastColon < 0 { + return addr, 80 + } + host = addr[:lastColon] + portStr = addr[lastColon+1:] + } + port := 80 + if _, scanErr := fmt.Sscanf(portStr, "%d", &port); scanErr != nil { + port = 80 + } + if host == "" { + host = "127.0.0.1" + } + return host, port +} + +func sanitizeProxyName(domain string) string { + return strings.ReplaceAll(strings.ReplaceAll(domain, ".", "-"), "*", "wildcard") +} + +func buildTunnelSettings(node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *relay.Settings { + return relay.BuildSettings(node, updateNow, updateChannel, updateTag) +} diff --git a/Wavelet/internal/apps/openflare/flared/logics.go b/Wavelet/internal/apps/openflare/flared/logics.go new file mode 100644 index 00000000..1f4f9272 --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/logics.go @@ -0,0 +1,300 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/relay" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +const ( + nodeStatusOnline = "online" + applyResultOK = "success" + applyResultWarn = "warning" + applyResultFail = "failed" +) + +// HeartbeatPayload is sent by OpenFlared on each heartbeat. +type HeartbeatPayload struct { + ClientVersion string `json:"client_version"` + FrpVersion string `json:"frp_version"` + IP string `json:"ip"` + TunnelStatus string `json:"tunnel_status"` + ConnectedRelays []ConnectedRelay `json:"connected_relays"` + CurrentVersion string `json:"current_version"` + CurrentChecksum string `json:"current_checksum"` +} + +// ConnectedRelay describes relay connection status from the client. +type ConnectedRelay struct { + RelayNodeID string `json:"relay_node_id"` + Status string `json:"status"` + ProxyCount int `json:"proxy_count"` +} + +// ActiveConfigMeta summarizes the active config version. +type ActiveConfigMeta struct { + Version string `json:"version"` + Checksum string `json:"checksum"` +} + +// HeartbeatResponse is returned to the OpenFlared client. +type HeartbeatResponse struct { + ActiveConfig *ActiveConfigMeta `json:"active_config"` + TunnelSettings *relay.Settings `json:"tunnel_settings"` +} + +// TunnelConfigResponse is the full tunnel routing config sent to the client. +type TunnelConfigResponse struct { + Version string `json:"version"` + Checksum string `json:"checksum"` + Relays []RelayInfo `json:"relays"` + Proxies []ProxyEntry `json:"proxies"` +} + +// RelayInfo describes a relay the client should connect to. +type RelayInfo struct { + RelayNodeID string `json:"relay_node_id"` + Address string `json:"address"` + AuthToken string `json:"auth_token"` + ProxyURL string `json:"proxy_url"` +} + +// ProxyEntry describes one frpc proxy definition. +type ProxyEntry struct { + Name string `json:"name"` + Type string `json:"type"` + LocalAddr string `json:"local_addr"` + LocalPort int `json:"local_port"` + CustomDomains []string `json:"custom_domains"` +} + +// ApplyLogPayload is the apply result reported by OpenFlared. +type ApplyLogPayload struct { + NodeID string `json:"node_id"` + Version string `json:"version"` + Result string `json:"result"` + Message string `json:"message"` + Checksum string `json:"checksum"` + MainConfigChecksum string `json:"main_config_checksum"` + RouteConfigChecksum string `json:"route_config_checksum"` + SupportFileCount int `json:"support_file_count"` +} + +// Heartbeat processes an OpenFlared heartbeat and returns runtime settings. +func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) { + if node == nil { + return nil, fmt.Errorf("tunnel client node is nil") + } + if node.NodeType != "tunnel_client" { + return nil, fmt.Errorf("node %s is not a tunnel_client", node.NodeID) + } + + payload = normalizeFlaredHeartbeatPayload(payload) + previous := *node + updateNow := node.UpdateRequested + updateChannel := normalizeReleaseChannel(node.UpdateChannel) + updateTag := strings.TrimSpace(node.UpdateTag) + + now := time.Now().UTC() + changes := map[string]any{ + "version": payload.ClientVersion, + "ext_version": payload.FrpVersion, + "current_version": payload.CurrentVersion, + "last_seen_at": now, + "status": nodeStatusOnline, + "update_requested": false, + "update_channel": "stable", + "update_tag": "", + } + if !previous.UpdateRequested { + delete(changes, "update_requested") + } + if previous.UpdateChannel == "stable" { + delete(changes, "update_channel") + } + if previous.UpdateTag == "" { + delete(changes, "update_tag") + } + if !node.IPManualOverride && payload.IP != "" && previous.IP != payload.IP { + changes["ip"] = payload.IP + node.IP = payload.IP + } + + node.Version = payload.ClientVersion + node.ExtVersion = payload.FrpVersion + node.CurrentVersion = payload.CurrentVersion + node.UpdateRequested = false + node.UpdateChannel = "stable" + node.UpdateTag = "" + lastSeen := now + node.LastSeenAt = &lastSeen + node.Status = nodeStatusOnline + + if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil { + return nil, fmt.Errorf("update flared heartbeat: %w", err) + } + + activeConfig, err := getActiveConfigMeta(ctx) + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + + return &HeartbeatResponse{ + ActiveConfig: activeConfig, + TunnelSettings: buildTunnelSettings(node, updateNow, updateChannel, updateTag), + }, nil +} + +// GetTunnelConfig builds the full tunnel routing config for an OpenFlared client. +func GetTunnelConfig(ctx context.Context, node *model.OpenFlareNode) (*TunnelConfigResponse, error) { + if node == nil { + return nil, fmt.Errorf("node is nil") + } + + activeVersion, err := getActiveConfigVersion(ctx) + if err != nil { + return nil, fmt.Errorf("no active config version: %w", err) + } + + routes, err := model.ListProxyRoutes(ctx) + if err != nil { + return nil, fmt.Errorf("get proxy routes: %w", err) + } + + relayNodes, err := listTunnelRelayNodes(ctx) + if err != nil { + return nil, fmt.Errorf("get relay nodes: %w", err) + } + + relays := make([]RelayInfo, 0, len(relayNodes)) + for i := range relayNodes { + relayNode := relayNodes[i] + if relayNode.RelayStatus == "healthy" || relayNode.Status == nodeStatusOnline { + relays = append(relays, RelayInfo{ + RelayNodeID: relayNode.NodeID, + Address: relayClientAddress(&relayNode), + AuthToken: relayNode.RelayAuthToken, + ProxyURL: strings.TrimSpace(relayNode.RelayClientProxyURL), + }) + } + } + + proxies := make([]ProxyEntry, 0) + for _, route := range routes { + if route == nil || route.UpstreamType != "tunnel" || route.TunnelNodeID == nil || *route.TunnelNodeID != node.ID { + continue + } + if !route.Enabled { + continue + } + domains, decodeErr := decodeStoredDomains(route.Domains, route.Domain) + if decodeErr != nil { + continue + } + localAddr, localPort := parseTunnelTargetAddr(route.TunnelTargetAddr) + proxies = append(proxies, ProxyEntry{ + Name: fmt.Sprintf("%s-%s", node.NodeID, sanitizeProxyName(domains[0])), + Type: "http", + LocalAddr: localAddr, + LocalPort: localPort, + CustomDomains: domains, + }) + } + + return &TunnelConfigResponse{ + Version: activeVersion.Version, + Checksum: activeVersion.Checksum, + Relays: relays, + Proxies: proxies, + }, nil +} + +// ReportApplyLog records an apply result from OpenFlared. +func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) { + now := time.Now().UTC() + payload = normalizeApplyLogPayload(payload) + if payload.NodeID == "" { + return nil, errors.New("node_id 不能为空") + } + if payload.Version == "" { + return nil, errors.New("version 不能为空") + } + if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFail { + return nil, errors.New("result 仅支持 success、warning 或 failed") + } + + log := &model.OpenFlareApplyLog{ + NodeID: payload.NodeID, + Version: payload.Version, + Result: payload.Result, + Message: payload.Message, + Checksum: payload.Checksum, + MainConfigChecksum: payload.MainConfigChecksum, + RouteConfigChecksum: payload.RouteConfigChecksum, + SupportFileCount: payload.SupportFileCount, + CreatedAt: now, + } + + err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + var node model.OpenFlareNode + if err := tx.Where("node_id = ?", payload.NodeID).First(&node).Error; err != nil { + return err + } + node.Status = nodeStatusOnline + lastSeen := now + node.LastSeenAt = &lastSeen + if payload.Result == applyResultOK { + node.CurrentVersion = payload.Version + node.LastError = "" + } else { + node.LastError = payload.Message + } + if err := tx.Create(log).Error; err != nil { + return err + } + return tx.Model(&node).Select("status", "last_seen_at", "current_version", "last_error").Updates(&node).Error + }) + if err != nil { + return nil, err + } + return log, nil +} + +func getActiveConfigVersion(ctx context.Context) (*configVersionRow, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New("database not initialized") + } + if !conn.Migrator().HasTable(&configVersionRow{}) { + return nil, gorm.ErrRecordNotFound + } + var version configVersionRow + if err := conn.Where("is_active = ?", true).Order("id desc").First(&version).Error; err != nil { + return nil, err + } + return &version, nil +} + +func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload { + payload.NodeID = strings.TrimSpace(payload.NodeID) + payload.Version = strings.TrimSpace(payload.Version) + payload.Result = strings.ToLower(strings.TrimSpace(payload.Result)) + payload.Message = strings.TrimSpace(payload.Message) + payload.Checksum = strings.TrimSpace(payload.Checksum) + payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum) + payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum) + if len(payload.Message) > 16000 { + payload.Message = payload.Message[:16000] + } + return payload +} diff --git a/Wavelet/internal/apps/openflare/flared/middleware.go b/Wavelet/internal/apps/openflare/flared/middleware.go new file mode 100644 index 00000000..28b0ac8a --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/middleware.go @@ -0,0 +1,51 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +import ( + "context" + "errors" + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +const ctxFlaredNodeKey = "flared_node" + +// TunnelAuth authenticates flared requests using X-Tunnel-Token and verifies tunnel_client type. +func TunnelAuth() gin.HandlerFunc { + return func(c *gin.Context) { + token := strings.TrimSpace(c.GetHeader("X-Tunnel-Token")) + node, err := authenticateAccessToken(c.Request.Context(), token) + if err != nil { + compat.Unauthorized(c, errTunnelTokenInvalid) + c.Abort() + return + } + if node.NodeType != "tunnel_client" { + compat.Forbidden(c, errTunnelNodeTypeMismatch) + c.Abort() + return + } + c.Set(ctxFlaredNodeKey, node) + c.Next() + } +} + +func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) { + if token == "" { + return nil, errors.New("missing tunnel token") + } + node, err := model.GetOpenFlareNodeByAccessToken(ctx, token) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, errors.New("invalid tunnel token") + } + return nil, err + } + return node, nil +} diff --git a/Wavelet/internal/apps/openflare/flared/middleware_test.go b/Wavelet/internal/apps/openflare/flared/middleware_test.go new file mode 100644 index 00000000..cffd3148 --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/middleware_test.go @@ -0,0 +1,106 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupFlaredMiddlewareTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) + db.SetDB(sqliteDB) + + return func() { + db.SetDB(nil) + } +} + +func seedFlaredNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareNode { + t.Helper() + ctx := context.Background() + node := &model.OpenFlareNode{ + NodeID: "flared-test-node", + Name: "flared-test", + Status: "pending", + NodeType: nodeType, + AccessToken: accessToken, + } + require.NoError(t, model.CreateOpenFlareNode(ctx, node)) + return node +} + +func TestTunnelAuthMissingToken(t *testing.T) { + cleanup := setupFlaredMiddlewareTestDB(t) + defer cleanup() + + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) { + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/flared/test", nil) + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusUnauthorized, rec.Code) +} + +func TestTunnelAuthRejectsWrongNodeType(t *testing.T) { + cleanup := setupFlaredMiddlewareTestDB(t) + defer cleanup() + seedFlaredNode(t, "edge_node", "edge-token-flared") + + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) { + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/flared/test", nil) + req.Header.Set("X-Tunnel-Token", "edge-token-flared") + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusForbidden, rec.Code) +} + +func TestTunnelAuthAcceptsTunnelClient(t *testing.T) { + cleanup := setupFlaredMiddlewareTestDB(t) + defer cleanup() + node := seedFlaredNode(t, "tunnel_client", "tunnel-token-valid") + + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) { + authNode, ok := c.Get(ctxFlaredNodeKey) + require.True(t, ok) + assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID) + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/flared/test", nil) + req.Header.Set("X-Tunnel-Token", "tunnel-token-valid") + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusOK, rec.Code) +} diff --git a/Wavelet/internal/apps/openflare/flared/routers.go b/Wavelet/internal/apps/openflare/flared/routers.go new file mode 100644 index 00000000..a3816837 --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/routers.go @@ -0,0 +1,79 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +// PostHeartbeat handles POST /flared/heartbeat. +func PostHeartbeat(c *gin.Context) { + var payload HeartbeatPayload + if !compat.BindJSON(c, &payload) { + return + } + + authNode, ok := c.Get(ctxFlaredNodeKey) + if !ok { + compat.Unauthorized(c, errTunnelTokenInvalid) + return + } + node := authNode.(*model.OpenFlareNode) + + result, err := Heartbeat(c.Request.Context(), node, payload) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +// GetActiveConfig handles GET /flared/config/active. +func GetActiveConfig(c *gin.Context) { + authNode, ok := c.Get(ctxFlaredNodeKey) + if !ok { + compat.Unauthorized(c, errTunnelTokenInvalid) + return + } + node := authNode.(*model.OpenFlareNode) + + config, err := GetTunnelConfig(c.Request.Context(), node) + if err != nil { + compat.Fail(c, "无法生成隧道配置: "+err.Error()) + return + } + compat.OK(c, config) +} + +// PostApplyLog handles POST /flared/apply-log. +func PostApplyLog(c *gin.Context) { + var payload ApplyLogPayload + if !compat.BindJSON(c, &payload) { + return + } + if authNode, ok := c.Get(ctxFlaredNodeKey); ok { + payload.NodeID = authNode.(*model.OpenFlareNode).NodeID + } + + log, err := ReportApplyLog(c.Request.Context(), payload) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, log) +} + +// GetWebSocket handles GET /flared/ws. +func GetWebSocket(c *gin.Context) { + authNode, ok := c.Get(ctxFlaredNodeKey) + if !ok { + compat.Unauthorized(c, errTunnelTokenInvalid) + return + } + node := authNode.(*model.OpenFlareNode) + ofws.ServeFlared(c, node.NodeID) +} diff --git a/Wavelet/internal/apps/openflare/geoip/lookup.go b/Wavelet/internal/apps/openflare/geoip/lookup.go new file mode 100644 index 00000000..ab253f8c --- /dev/null +++ b/Wavelet/internal/apps/openflare/geoip/lookup.go @@ -0,0 +1,76 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package geoip provides OpenFlare-compatible GeoIP lookup helpers. +package geoip + +import ( + "errors" + "net" + "strings" + + pkggeoip "github.com/rain-kl/openflare/pkg/geoip" +) + +// LookupView is the legacy OpenFlare GeoIP lookup response shape. +type LookupView struct { + Provider string `json:"provider"` + IP string `json:"ip"` + ISOCode string `json:"iso_code"` + Name string `json:"name"` + Latitude *float64 `json:"latitude,omitempty"` + Longitude *float64 `json:"longitude,omitempty"` +} + +const ( + errProviderInvalid = "归属方式仅支持 disabled、mmdb、ip-api、geojs、ipinfo" + errIPEmpty = "IP 不能为空" + errIPInvalid = "IP 格式无效" + errLookupEmpty = "未获取到 IP 归属结果" +) + +// IsValidProvider reports whether provider is a supported GeoIP backend. +func IsValidProvider(provider string) bool { + return pkggeoip.IsValidProvider(provider) +} + +// Lookup resolves geographic information for rawIP using the given provider. +func Lookup(provider, rawIP string) (*LookupView, error) { + trimmedProvider := strings.TrimSpace(provider) + if !pkggeoip.IsValidProvider(trimmedProvider) { + return nil, errors.New(errProviderInvalid) + } + + trimmedIP := strings.TrimSpace(rawIP) + if trimmedIP == "" { + return nil, errors.New(errIPEmpty) + } + parsedIP := net.ParseIP(trimmedIP) + if parsedIP == nil { + return nil, errors.New(errIPInvalid) + } + + if trimmedProvider == pkggeoip.ProviderDisabled { + return &LookupView{ + Provider: trimmedProvider, + IP: parsedIP.String(), + }, nil + } + + info, err := pkggeoip.LookupGeoInfoWithProvider(trimmedProvider, parsedIP) + if err != nil { + return nil, err + } + if info == nil { + return nil, errors.New(errLookupEmpty) + } + + return &LookupView{ + Provider: trimmedProvider, + IP: parsedIP.String(), + ISOCode: info.ISOCode, + Name: info.Name, + Latitude: info.Latitude, + Longitude: info.Longitude, + }, nil +} diff --git a/Wavelet/internal/apps/openflare/geoip/lookup_test.go b/Wavelet/internal/apps/openflare/geoip/lookup_test.go new file mode 100644 index 00000000..02467be7 --- /dev/null +++ b/Wavelet/internal/apps/openflare/geoip/lookup_test.go @@ -0,0 +1,73 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package geoip + +import ( + "net" + "testing" + + pkggeoip "github.com/rain-kl/openflare/pkg/geoip" +) + +type fakeLookupProvider struct{} + +func (f *fakeLookupProvider) Name() string { return "fake-lookup" } + +func (f *fakeLookupProvider) GetGeoInfo(_ net.IP) (*pkggeoip.GeoInfo, error) { + lat := 37.7749 + lon := -122.4194 + return &pkggeoip.GeoInfo{ + ISOCode: "US", + Name: "United States", + Latitude: &lat, + Longitude: &lon, + }, nil +} + +func (f *fakeLookupProvider) UpdateDatabase() error { return nil } + +func (f *fakeLookupProvider) Close() error { return nil } + +func TestLookupWithProvider(t *testing.T) { + previousFactory := pkggeoip.ProviderFactoryForTest() + pkggeoip.SetProviderFactoryForTest(func(provider string) (pkggeoip.GeoIPService, error) { + return &fakeLookupProvider{}, nil + }) + t.Cleanup(func() { + pkggeoip.SetProviderFactoryForTest(previousFactory) + }) + + view, err := Lookup("ipinfo", "8.8.8.8") + if err != nil { + t.Fatalf("Lookup() error = %v", err) + } + if view.Provider != "ipinfo" || view.IP != "8.8.8.8" { + t.Fatalf("unexpected lookup view: %+v", view) + } + if view.ISOCode != "US" || view.Name != "United States" { + t.Fatalf("unexpected geo fields: %+v", view) + } + if view.Latitude == nil || view.Longitude == nil { + t.Fatalf("expected coordinates, got %+v", view) + } +} + +func TestLookupRejectsInvalidInput(t *testing.T) { + if _, err := Lookup("invalid", "8.8.8.8"); err == nil { + t.Fatal("expected invalid provider to fail") + } + if _, err := Lookup("ipinfo", "not-an-ip"); err == nil { + t.Fatal("expected invalid IP to fail") + } +} + +func TestLookupDisabledProvider(t *testing.T) { + view, err := Lookup("disabled", "8.8.8.8") + if err != nil { + t.Fatalf("Lookup() error = %v", err) + } + if view.Provider != "disabled" || view.IP != "8.8.8.8" { + t.Fatalf("unexpected disabled view: %+v", view) + } +} diff --git a/Wavelet/internal/apps/openflare/integration/agent_protocol_test.go b/Wavelet/internal/apps/openflare/integration/agent_protocol_test.go new file mode 100644 index 00000000..d5a4b783 --- /dev/null +++ b/Wavelet/internal/apps/openflare/integration/agent_protocol_test.go @@ -0,0 +1,232 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package integration + +import ( + "context" + "encoding/json" + "net/http" + "testing" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" + oflegacy "github.com/Rain-kl/Wavelet/internal/apps/openflare/legacy" + ofnode "github.com/Rain-kl/Wavelet/internal/apps/openflare/node" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type configVersionRecord struct { + ID uint `gorm:"primaryKey"` + Version string `gorm:"column:version"` + SnapshotJSON string `gorm:"column:snapshot_json"` + SupportFilesJSON string `gorm:"column:support_files_json"` + Checksum string `gorm:"column:checksum"` + IsActive bool `gorm:"column:is_active"` +} + +func (configVersionRecord) TableName() string { + return "of_config_versions" +} + +func setupProtocolTestEnv(t *testing.T) (*gin.Engine, func()) { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.OpenFlareNode{}, + &model.OpenFlareOption{}, + &model.OpenFlareApplyLog{}, + &configVersionRecord{}, + )) + + db.SetDB(sqliteDB) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + + gin.SetMode(gin.TestMode) + engine := gin.New() + apiGroup := engine.Group("/api") + oflegacy.RegisterRoutes(apiGroup) + + cleanup := func() { + db.SetDB(nil) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + } + return engine, cleanup +} + +func TestAgentRelayFlaredProtocol(t *testing.T) { + engine, cleanup := setupProtocolTestEnv(t) + defer cleanup() + + ctx := context.Background() + + t.Run("create edge node and heartbeat with X-Agent-Token", func(t *testing.T) { + edge, err := ofnode.CreateNode(ctx, ofnode.Input{ + Name: "edge-1", + IP: "10.0.0.1", + }) + require.NoError(t, err) + require.NotEmpty(t, edge.AccessToken) + assert.Equal(t, "edge_node", edge.NodeType) + + rec := performJSONRequest(t, engine, http.MethodPost, "/api/agent/nodes/heartbeat", map[string]any{ + "name": "edge-1", + "ip": "203.0.113.10", + "version": "0.1.0", + }, map[string]string{ + "X-Agent-Token": edge.AccessToken, + }) + assert.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + assert.True(t, envelope.Success) + + var heartbeatBody struct { + Success bool `json:"success"` + Data any `json:"data"` + AgentSettings any `json:"agent_settings"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &heartbeatBody)) + assert.True(t, heartbeatBody.Success) + assert.NotNil(t, heartbeatBody.AgentSettings) + + stored, err := model.GetOpenFlareNodeByNodeID(ctx, edge.NodeID) + require.NoError(t, err) + assert.Equal(t, "online", stored.Status) + assert.Equal(t, "0.1.0", stored.Version) + }) + + t.Run("create tunnel_relay node and relay heartbeat", func(t *testing.T) { + relayNode, err := ofnode.CreateNode(ctx, ofnode.Input{ + Name: "relay-1", + NodeType: "tunnel_relay", + }) + require.NoError(t, err) + require.NotEmpty(t, relayNode.AccessToken) + + rec := performJSONRequest(t, engine, http.MethodPost, "/api/relay/heartbeat", map[string]any{ + "version": "v0.1.0", + "frp_version": "0.61.0", + "relay_status": "healthy", + "name": "relay-1", + "ip": "203.0.113.20", + }, map[string]string{ + "X-Agent-Token": relayNode.AccessToken, + }) + assert.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + assert.True(t, envelope.Success) + + var heartbeatData struct { + RelayConfig map[string]any `json:"relay_config"` + RelaySettings map[string]any `json:"relay_settings"` + } + unmarshalEnvelopeData(t, envelope.Data, &heartbeatData) + assert.NotNil(t, heartbeatData.RelayConfig) + assert.NotNil(t, heartbeatData.RelaySettings) + + stored, err := model.GetOpenFlareNodeByNodeID(ctx, relayNode.NodeID) + require.NoError(t, err) + assert.Equal(t, "online", stored.Status) + assert.Equal(t, "healthy", stored.RelayStatus) + }) + + t.Run("create tunnel_client node and flared heartbeat with X-Tunnel-Token", func(t *testing.T) { + clientNode, err := ofnode.CreateNode(ctx, ofnode.Input{ + Name: "client-1", + NodeType: "tunnel_client", + }) + require.NoError(t, err) + require.NotEmpty(t, clientNode.AccessToken) + + rec := performJSONRequest(t, engine, http.MethodPost, "/api/flared/heartbeat", map[string]any{ + "client_version": "v0.2.0", + "frp_version": "0.61.0", + "tunnel_status": "running", + }, map[string]string{ + "X-Tunnel-Token": clientNode.AccessToken, + }) + assert.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + assert.True(t, envelope.Success) + + stored, err := model.GetOpenFlareNodeByNodeID(ctx, clientNode.NodeID) + require.NoError(t, err) + assert.Equal(t, "online", stored.Status) + assert.Equal(t, "v0.2.0", stored.Version) + }) + + t.Run("agent register with discovery token from options", func(t *testing.T) { + bootstrap, err := ofnode.GetBootstrapToken(ctx) + require.NoError(t, err) + require.NotEmpty(t, bootstrap.DiscoveryToken) + + rec := performJSONRequest(t, engine, http.MethodPost, "/api/agent/nodes/register", map[string]any{ + "name": "discovered-edge", + "ip": "203.0.113.30", + "version": "0.2.0", + }, map[string]string{ + "X-Agent-Token": bootstrap.DiscoveryToken, + }) + assert.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + assert.True(t, envelope.Success) + + var registration agent.RegistrationResponse + unmarshalEnvelopeData(t, envelope.Data, ®istration) + assert.NotEmpty(t, registration.NodeID) + assert.NotEmpty(t, registration.AccessToken) + assert.Equal(t, "discovered-edge", registration.Name) + + stored, err := model.GetOpenFlareNodeByNodeID(ctx, registration.NodeID) + require.NoError(t, err) + assert.Equal(t, "online", stored.Status) + assert.Equal(t, registration.AccessToken, stored.AccessToken) + }) + + t.Run("POST agent apply-logs", func(t *testing.T) { + edge, err := ofnode.CreateNode(ctx, ofnode.Input{ + Name: "edge-apply", + IP: "10.0.0.2", + }) + require.NoError(t, err) + + rec := performJSONRequest(t, engine, http.MethodPost, "/api/agent/apply-logs", map[string]any{ + "version": "20260618-001", + "result": "success", + "message": "apply ok", + }, map[string]string{ + "X-Agent-Token": edge.AccessToken, + }) + assert.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + assert.True(t, envelope.Success) + + var applyLog model.OpenFlareApplyLog + unmarshalEnvelopeData(t, envelope.Data, &applyLog) + assert.Equal(t, edge.NodeID, applyLog.NodeID) + assert.Equal(t, "success", applyLog.Result) + assert.Equal(t, "20260618-001", applyLog.Version) + + stored, err := model.GetOpenFlareNodeByNodeID(ctx, edge.NodeID) + require.NoError(t, err) + assert.Equal(t, "online", stored.Status) + assert.Equal(t, "20260618-001", stored.CurrentVersion) + }) +} diff --git a/Wavelet/internal/apps/openflare/integration/auth_option_test.go b/Wavelet/internal/apps/openflare/integration/auth_option_test.go new file mode 100644 index 00000000..119b0cec --- /dev/null +++ b/Wavelet/internal/apps/openflare/integration/auth_option_test.go @@ -0,0 +1,241 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package integration + +import ( + "context" + "net/http" + "testing" + + "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + oflegacy "github.com/Rain-kl/Wavelet/internal/apps/openflare/legacy" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/db/idgen" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/gin-contrib/sessions" + "github.com/gin-contrib/sessions/cookie" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type statusPayload struct { + SystemName string `json:"system_name"` +} + +type legacyUserPayload struct { + Username string `json:"username"` + Token string `json:"token"` +} + +func setupAuthOptionIntegration(t *testing.T) (*gorm.DB, *gin.Engine) { + t.Helper() + + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + t.Cleanup(cleanup) + + require.NoError(t, dbConn.AutoMigrate(&model.OpenFlareOption{})) + option.ResetInitializationForTest() + t.Cleanup(option.ResetInitializationForTest) + + oldCookieName := config.Config.App.SessionCookieName + oldSecret := config.Config.App.SessionSecret + oldDomain := config.Config.App.SessionDomain + oldSecure := config.Config.App.SessionSecure + oldHTTPOnly := config.Config.App.SessionHTTPOnly + t.Cleanup(func() { + config.Config.App.SessionCookieName = oldCookieName + config.Config.App.SessionSecret = oldSecret + config.Config.App.SessionDomain = oldDomain + config.Config.App.SessionSecure = oldSecure + config.Config.App.SessionHTTPOnly = oldHTTPOnly + }) + + config.Config.App.SessionCookieName = "test_openflare_session" + config.Config.App.SessionSecret = "test_openflare_session_secret" + config.Config.App.SessionDomain = "" + config.Config.App.SessionSecure = false + config.Config.App.SessionHTTPOnly = true + + store := cookie.NewStore([]byte(config.Config.App.SessionSecret)) + store.Options(oauth.GetSessionOptions(3600)) + r := testhelper.NewTestGinEngine(sessions.Sessions(config.Config.App.SessionCookieName, store)) + + api := r.Group("/api") + oflegacy.RegisterRoutes(api) + + return dbConn, r +} + +func seedUser(t *testing.T, dbConn *gorm.DB, username, password string, isAdmin bool) *model.User { + t.Helper() + + user := &model.User{ + ID: idgen.NextUint64ID(), + Username: username, + Nickname: username, + Email: username + "@openflare.test", + IsActive: true, + IsAdmin: isAdmin, + } + require.NoError(t, user.SetEncryptedPassword(password)) + require.NoError(t, dbConn.Create(user).Error) + return user +} + +func TestGETStatusReturnsSuccessEnvelope(t *testing.T) { + _, r := setupAuthOptionIntegration(t) + + w := performJSONRequest(t, r, http.MethodGet, "/api/status", nil, nil) + + assert.Equal(t, http.StatusOK, w.Code) + env := decodeEnvelope(t, w) + assert.True(t, env.Success, "message=%s", env.Message) + + var status statusPayload + unmarshalEnvelopeData(t, env.Data, &status) + assert.NotEmpty(t, status.SystemName) +} + +func TestPOSTUserLoginWithSeededUser(t *testing.T) { + dbConn, r := setupAuthOptionIntegration(t) + seedUser(t, dbConn, "testuser", "password123", false) + + w := performJSONRequest(t, r, http.MethodPost, "/api/user/login", map[string]string{ + "username": "testuser", + "password": "password123", + }, nil) + + assert.Equal(t, http.StatusOK, w.Code) + env := decodeEnvelope(t, w) + assert.True(t, env.Success, "message=%s", env.Message) + + var user legacyUserPayload + unmarshalEnvelopeData(t, env.Data, &user) + assert.Equal(t, "testuser", user.Username) + assert.NotEmpty(t, user.Token) +} + +func TestGETUserSelfWithToken(t *testing.T) { + dbConn, r := setupAuthOptionIntegration(t) + seedUser(t, dbConn, "selfuser", "password123", false) + + loginResp := performJSONRequest(t, r, http.MethodPost, "/api/user/login", map[string]string{ + "username": "selfuser", + "password": "password123", + }, nil) + loginEnv := decodeEnvelope(t, loginResp) + require.True(t, loginEnv.Success, "login failed: %s", loginEnv.Message) + + var loginUser legacyUserPayload + unmarshalEnvelopeData(t, loginEnv.Data, &loginUser) + require.NotEmpty(t, loginUser.Token) + + w := performJSONRequest(t, r, http.MethodGet, "/api/user/self", nil, map[string]string{ + compat.OpenFlareTokenHeader(): loginUser.Token, + }) + + assert.Equal(t, http.StatusOK, w.Code) + env := decodeEnvelope(t, w) + assert.True(t, env.Success, "message=%s", env.Message) + + var self legacyUserPayload + unmarshalEnvelopeData(t, env.Data, &self) + assert.Equal(t, "selfuser", self.Username) +} + +func TestGETOptionRequiresRootAuth(t *testing.T) { + dbConn, r := setupAuthOptionIntegration(t) + seedUser(t, dbConn, "commonuser", "password123", false) + seedUser(t, dbConn, "rootuser", "password123", true) + + commonToken := loginAndGetToken(t, r, "commonuser", "password123") + rootToken := loginAndGetToken(t, r, "rootuser", "password123") + + t.Run("unauthenticated", func(t *testing.T) { + w := performJSONRequest(t, r, http.MethodGet, "/api/option/", nil, nil) + assert.Equal(t, http.StatusUnauthorized, w.Code) + env := decodeEnvelope(t, w) + assert.False(t, env.Success) + }) + + t.Run("common user forbidden", func(t *testing.T) { + w := performJSONRequest(t, r, http.MethodGet, "/api/option/", nil, map[string]string{ + compat.OpenFlareTokenHeader(): commonToken, + }) + assert.Equal(t, http.StatusOK, w.Code) + env := decodeEnvelope(t, w) + assert.False(t, env.Success) + assert.Contains(t, env.Message, "权限不足") + }) + + t.Run("root user allowed", func(t *testing.T) { + w := performJSONRequest(t, r, http.MethodGet, "/api/option/", nil, map[string]string{ + compat.OpenFlareTokenHeader(): rootToken, + }) + assert.Equal(t, http.StatusOK, w.Code) + env := decodeEnvelope(t, w) + assert.True(t, env.Success, "message=%s", env.Message) + }) +} + +func TestOptionHotReloadAfterUpdate(t *testing.T) { + dbConn, r := setupAuthOptionIntegration(t) + seedUser(t, dbConn, "admin", "password123", true) + rootToken := loginAndGetToken(t, r, "admin", "password123") + + statusBefore := getStatusSystemName(t, r, nil) + assert.NotEmpty(t, statusBefore) + + updateResp := performJSONRequest(t, r, http.MethodPost, "/api/option/update", map[string]string{ + "key": "SystemName", + "value": "HotReloadIntegration", + }, map[string]string{ + compat.OpenFlareTokenHeader(): rootToken, + }) + assert.Equal(t, http.StatusOK, updateResp.Code) + updateEnv := decodeEnvelope(t, updateResp) + assert.True(t, updateEnv.Success, "message=%s", updateEnv.Message) + + statusAfter := getStatusSystemName(t, r, nil) + assert.Equal(t, "HotReloadIntegration", statusAfter) + assert.Equal(t, "HotReloadIntegration", model.SystemName) + + ctx := context.Background() + require.NoError(t, option.EnsureInitialized(ctx)) + assert.Equal(t, "HotReloadIntegration", model.OptionValue("SystemName")) +} + +func loginAndGetToken(t *testing.T, r http.Handler, username, password string) string { + t.Helper() + + w := performJSONRequest(t, r, http.MethodPost, "/api/user/login", map[string]string{ + "username": username, + "password": password, + }, nil) + env := decodeEnvelope(t, w) + require.True(t, env.Success, "login failed: %s", env.Message) + + var user legacyUserPayload + unmarshalEnvelopeData(t, env.Data, &user) + require.NotEmpty(t, user.Token) + return user.Token +} + +func getStatusSystemName(t *testing.T, r http.Handler, headers map[string]string) string { + t.Helper() + + w := performJSONRequest(t, r, http.MethodGet, "/api/status", nil, headers) + require.Equal(t, http.StatusOK, w.Code) + env := decodeEnvelope(t, w) + require.True(t, env.Success, "message=%s", env.Message) + + var status statusPayload + unmarshalEnvelopeData(t, env.Data, &status) + return status.SystemName +} diff --git a/Wavelet/internal/apps/openflare/integration/core_chain_test.go b/Wavelet/internal/apps/openflare/integration/core_chain_test.go new file mode 100644 index 00000000..d37c5f85 --- /dev/null +++ b/Wavelet/internal/apps/openflare/integration/core_chain_test.go @@ -0,0 +1,291 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package integration + +import ( + "net/http" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" + oflegacy "github.com/Rain-kl/Wavelet/internal/apps/openflare/legacy" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +const ( + adminUserID = uint64(1001) + adminUsername = "openflare-admin" +) + +type adminSeed struct { + User model.User + Token string + TokenHash string +} + +func setupCoreChainTest(t *testing.T) (*gin.Engine, adminSeed, func()) { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.User{}, + &model.AccessToken{}, + &model.Origin{}, + &model.ProxyRoute{}, + &model.ConfigVersion{}, + &model.OpenFlareWAFRuleGroup{}, + &model.OpenFlareWAFRuleGroupBinding{}, + &model.OpenFlareWAFIPGroup{}, + &model.OpenFlareNode{}, + &model.OpenFlareOption{}, + &model.OpenFlareApplyLog{}, + )) + + db.SetDB(sqliteDB) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + + seed, err := seedAdminWithAccessToken(sqliteDB) + require.NoError(t, err) + + gin.SetMode(gin.TestMode) + engine := gin.New() + apiGroup := engine.Group("/api") + oflegacy.RegisterRoutes(apiGroup) + + cleanup := func() { + db.SetDB(nil) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + } + + return engine, seed, cleanup +} + +func seedAdminWithAccessToken(conn *gorm.DB) (adminSeed, error) { + now := time.Now().UTC() + admin := model.User{ + ID: adminUserID, + Username: adminUsername, + Nickname: "OpenFlare Admin", + IsActive: true, + IsAdmin: true, + LastLoginAt: now, + } + if err := conn.Create(&admin).Error; err != nil { + return adminSeed{}, err + } + + token, err := model.GenerateTokenString() + if err != nil { + return adminSeed{}, err + } + tokenHash := model.HashToken(token) + tokenRecord := model.AccessToken{ + UserID: adminUserID, + Name: "integration-admin-token", + TokenHash: tokenHash, + MaskedToken: model.MaskTokenString(token), + IsAdmin: true, + } + if err := conn.Create(&tokenRecord).Error; err != nil { + return adminSeed{}, err + } + + return adminSeed{ + User: admin, + Token: token, + TokenHash: tokenHash, + }, nil +} + +func TestCoreChainMigrationFlow(t *testing.T) { + engine, seed, cleanup := setupCoreChainTest(t) + defer cleanup() + + var ( + originID uint + proxyRouteID uint + configVersion string + configChecksum string + nodeID uint + nodePublicID string + agentToken string + ) + + t.Run("create origin", func(t *testing.T) { + rec := performJSONRequest(t, engine, http.MethodPost, "/api/origins/", map[string]any{ + "name": "Primary Origin", + "address": "origin.core-chain.internal", + "remark": "integration upstream", + }, map[string]string{ + "X-Access-Token": seed.Token, + }) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + originID = uint(data["id"].(float64)) + assert.NotZero(t, originID) + assert.Equal(t, "Primary Origin", data["name"]) + assert.Equal(t, "origin.core-chain.internal", data["address"]) + }) + + t.Run("create proxy route linked to origin", func(t *testing.T) { + rec := performJSONRequest(t, engine, http.MethodPost, "/api/proxy-routes/", map[string]any{ + "site_name": "core-chain-site", + "domain": "core-chain.example.com", + "origin_id": originID, + "origin_scheme": "http", + "origin_port": "8080", + "enabled": true, + }, map[string]string{ + "X-Access-Token": seed.Token, + }) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + proxyRouteID = uint(data["id"].(float64)) + assert.NotZero(t, proxyRouteID) + assert.Equal(t, "core-chain-site", data["site_name"]) + assert.Equal(t, "core-chain.example.com", data["domain"]) + assert.Equal(t, float64(originID), data["origin_id"]) + assert.Equal(t, "http://origin.core-chain.internal:8080", data["origin_url"]) + }) + + t.Run("publish config version", func(t *testing.T) { + rec := performJSONRequest(t, engine, http.MethodPost, "/api/config-versions/publish", nil, map[string]string{ + "X-Access-Token": seed.Token, + }) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + configVersion, _ = data["version"].(string) + configChecksum, _ = data["checksum"].(string) + assert.NotEmpty(t, configVersion) + assert.NotEmpty(t, configChecksum) + assert.Equal(t, true, data["is_active"]) + + activeRec := performJSONRequest(t, engine, http.MethodGet, "/api/config-versions/active", nil, map[string]string{ + "X-Access-Token": seed.Token, + }) + require.Equal(t, http.StatusOK, activeRec.Code) + activeEnvelope := decodeEnvelope(t, activeRec) + require.True(t, activeEnvelope.Success, activeEnvelope.Message) + + activeData := unmarshalEnvelopeMap(t, activeEnvelope.Data) + assert.Equal(t, configVersion, activeData["version"]) + assert.Equal(t, configChecksum, activeData["checksum"]) + }) + + t.Run("create node", func(t *testing.T) { + rec := performJSONRequest(t, engine, http.MethodPost, "/api/nodes/", map[string]any{ + "name": "edge-core-chain", + "ip": "10.10.0.1", + "auto_update_enabled": true, + }, map[string]string{ + "X-Access-Token": seed.Token, + }) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + nodeID = uint(data["id"].(float64)) + nodePublicID, _ = data["node_id"].(string) + agentToken, _ = data["access_token"].(string) + assert.NotZero(t, nodeID) + assert.NotEmpty(t, nodePublicID) + assert.Len(t, agentToken, 32) + }) + + t.Run("create apply log for node", func(t *testing.T) { + rec := performJSONRequest(t, engine, http.MethodPost, "/api/agent/apply-logs", map[string]any{ + "version": configVersion, + "result": "success", + "message": "config applied", + "checksum": configChecksum, + "main_config_checksum": "main-checksum", + "route_config_checksum": "route-checksum", + "support_file_count": 2, + }, map[string]string{ + "X-Agent-Token": agentToken, + }) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + assert.Equal(t, nodePublicID, data["node_id"]) + assert.Equal(t, configVersion, data["version"]) + assert.Equal(t, "success", data["result"]) + assert.Equal(t, configChecksum, data["checksum"]) + }) + + t.Run("verify apply log listing and node metadata", func(t *testing.T) { + listRec := performJSONRequest( + t, + engine, + http.MethodGet, + "/api/apply-logs/?node_id="+nodePublicID+"&pageNo=1&pageSize=10", + nil, + map[string]string{ + "X-Access-Token": seed.Token, + }, + ) + require.Equal(t, http.StatusOK, listRec.Code) + + listEnvelope := decodeEnvelope(t, listRec) + require.True(t, listEnvelope.Success, listEnvelope.Message) + + listData := unmarshalEnvelopeMap(t, listEnvelope.Data) + assert.Equal(t, float64(1), listData["total"]) + + rows, ok := listData["rows"].([]any) + require.True(t, ok) + require.Len(t, rows, 1) + row, ok := rows[0].(map[string]any) + require.True(t, ok) + assert.Equal(t, nodePublicID, row["node_id"]) + assert.Equal(t, configVersion, row["version"]) + assert.Equal(t, "success", row["result"]) + + nodeRec := performJSONRequest(t, engine, http.MethodGet, "/api/nodes/", nil, map[string]string{ + "X-Access-Token": seed.Token, + }) + require.Equal(t, http.StatusOK, nodeRec.Code) + nodeEnvelope := decodeEnvelope(t, nodeRec) + require.True(t, nodeEnvelope.Success, nodeEnvelope.Message) + + nodes := unmarshalEnvelopeSlice(t, nodeEnvelope.Data) + require.Len(t, nodes, 1) + nodeView, ok := nodes[0].(map[string]any) + require.True(t, ok) + assert.Equal(t, float64(nodeID), nodeView["id"]) + assert.Equal(t, nodePublicID, nodeView["node_id"]) + assert.Equal(t, "success", nodeView["latest_apply_result"]) + assert.Equal(t, configChecksum, nodeView["latest_apply_checksum"]) + assert.Equal(t, float64(2), nodeView["latest_support_file_count"]) + }) +} diff --git a/Wavelet/internal/apps/openflare/integration/helpers_test.go b/Wavelet/internal/apps/openflare/integration/helpers_test.go new file mode 100644 index 00000000..775c9fce --- /dev/null +++ b/Wavelet/internal/apps/openflare/integration/helpers_test.go @@ -0,0 +1,93 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package integration + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/stretchr/testify/require" +) + +func decodeEnvelope(t *testing.T, rec *httptest.ResponseRecorder) compat.Envelope { + t.Helper() + + var envelope compat.Envelope + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &envelope)) + return envelope +} + +func unmarshalEnvelopeData(t *testing.T, data any, target any) { + t.Helper() + + payload, err := json.Marshal(data) + require.NoError(t, err) + require.NoError(t, json.Unmarshal(payload, target)) +} + +func unmarshalEnvelopeMap(t *testing.T, data any) map[string]any { + t.Helper() + + var result map[string]any + unmarshalEnvelopeData(t, data, &result) + return result +} + +func unmarshalEnvelopeSlice(t *testing.T, data any) []any { + t.Helper() + + var result []any + unmarshalEnvelopeData(t, data, &result) + return result +} + +func performJSONRequest( + t *testing.T, + engine http.Handler, + method, path string, + body any, + headers map[string]string, +) *httptest.ResponseRecorder { + t.Helper() + + var payload []byte + if body != nil { + var err error + payload, err = json.Marshal(body) + require.NoError(t, err) + } + + req := httptest.NewRequest(method, path, bytes.NewReader(payload)) + if body != nil { + req.Header.Set("Content-Type", "application/json") + } + for key, value := range headers { + req.Header.Set(key, value) + } + + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + return rec +} + +func performLegacyRequest( + t *testing.T, + engine http.Handler, + method, path string, + body any, + headers map[string]string, +) *httptest.ResponseRecorder { + t.Helper() + return performJSONRequest(t, engine, method, path, body, headers) +} + +func adminAuthHeaders(token string) map[string]string { + return map[string]string{ + "X-Access-Token": token, + } +} diff --git a/Wavelet/internal/apps/openflare/integration/security_test.go b/Wavelet/internal/apps/openflare/integration/security_test.go new file mode 100644 index 00000000..069fd584 --- /dev/null +++ b/Wavelet/internal/apps/openflare/integration/security_test.go @@ -0,0 +1,386 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package integration + +import ( + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "fmt" + "math/big" + "net/http" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/legacy" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupSecurityTest(t *testing.T) (*gin.Engine, adminSeed, func()) { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.User{}, + &model.AccessToken{}, + &model.Origin{}, + &model.ProxyRoute{}, + &model.OpenFlareWAFRuleGroup{}, + &model.OpenFlareWAFRuleGroupBinding{}, + &model.OpenFlareWAFIPGroup{}, + &model.TLSCertificate{}, + &model.ManagedDomain{}, + &model.DNSAccount{}, + &model.AcmeAccount{}, + )) + + db.SetDB(sqliteDB) + option.ResetInitializationForTest() + + seed, err := seedAdminWithAccessToken(sqliteDB) + require.NoError(t, err) + + oldSecret := config.Config.App.SessionSecret + config.Config.App.SessionSecret = "test_session_secret_for_security_integration" + + engine := testhelper.NewTestGinEngine() + apiGroup := engine.Group("/api") + legacy.RegisterRoutes(apiGroup) + + cleanup := func() { + config.Config.App.SessionSecret = oldSecret + db.SetDB(nil) + option.ResetInitializationForTest() + } + + return engine, seed, cleanup +} + +func generateSelfSignedCertificatePair(t *testing.T, dnsNames []string) (string, string) { + t.Helper() + + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + template := &x509.Certificate{ + SerialNumber: big.NewInt(time.Now().UnixNano()), + Subject: pkix.Name{ + CommonName: dnsNames[0], + }, + DNSNames: dnsNames, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + } + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey) + require.NoError(t, err) + + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)}) + return string(certPEM), string(keyPEM) +} + +func TestSecurityWAFTLSMigrationFlow(t *testing.T) { + engine, seed, cleanup := setupSecurityTest(t) + defer cleanup() + + var ( + ruleGroupID uint + ipGroupID uint + proxyRouteID uint + certID uint + domainID uint + dnsAccountID uint + ) + + t.Run("WAF rule group create", func(t *testing.T) { + rec := performLegacyRequest(t, engine, http.MethodPost, "/api/waf/rule-groups", map[string]any{ + "name": "edge-security", + "enabled": true, + "block_status_code": 403, + "ip_whitelist": []string{"192.0.2.1"}, + "ip_blacklist": []string{"203.0.113.10"}, + "country_blacklist": []string{"CN"}, + "remark": "integration rule group", + }, adminAuthHeaders(seed.Token)) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + ruleGroupID = uint(data["id"].(float64)) + assert.NotZero(t, ruleGroupID) + assert.Equal(t, "edge-security", data["name"]) + assert.Equal(t, false, data["is_global"]) + assert.Equal(t, float64(403), data["block_status_code"]) + }) + + t.Run("WAF rule group list includes global and custom groups", func(t *testing.T) { + rec := performLegacyRequest(t, engine, http.MethodGet, "/api/waf/rule-groups", nil, adminAuthHeaders(seed.Token)) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + groups := unmarshalEnvelopeSlice(t, envelope.Data) + require.GreaterOrEqual(t, len(groups), 2) + + foundCustom := false + foundGlobal := false + for _, item := range groups { + group, ok := item.(map[string]any) + require.True(t, ok) + if group["is_global"] == true { + foundGlobal = true + } + if uint(group["id"].(float64)) == ruleGroupID { + foundCustom = true + assert.Equal(t, "edge-security", group["name"]) + } + } + assert.True(t, foundGlobal) + assert.True(t, foundCustom) + }) + + t.Run("WAF rule group get detail", func(t *testing.T) { + rec := performLegacyRequest( + t, + engine, + http.MethodGet, + fmt.Sprintf("/api/waf/rule-groups/%d", ruleGroupID), + nil, + adminAuthHeaders(seed.Token), + ) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + assert.Equal(t, float64(ruleGroupID), data["id"]) + assert.Equal(t, "edge-security", data["name"]) + }) + + t.Run("WAF rule group update", func(t *testing.T) { + rec := performLegacyRequest( + t, + engine, + http.MethodPost, + fmt.Sprintf("/api/waf/rule-groups/%d/update", ruleGroupID), + map[string]any{ + "name": "edge-security-updated", + "enabled": true, + "block_status_code": 451, + "remark": "updated by integration test", + }, + adminAuthHeaders(seed.Token), + ) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + assert.Equal(t, "edge-security-updated", data["name"]) + assert.Equal(t, float64(451), data["block_status_code"]) + }) + + t.Run("WAF IP group create", func(t *testing.T) { + rec := performLegacyRequest(t, engine, http.MethodPost, "/api/waf/ip-groups", map[string]any{ + "name": "blocked-ips", + "type": "manual", + "enabled": true, + "ip_list": []string{"203.0.113.0/24", "198.51.100.10"}, + "remark": "manual deny list", + }, adminAuthHeaders(seed.Token)) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + ipGroupID = uint(data["id"].(float64)) + assert.NotZero(t, ipGroupID) + assert.Equal(t, "blocked-ips", data["name"]) + assert.Equal(t, "manual", data["type"]) + }) + + t.Run("create proxy route for WAF binding", func(t *testing.T) { + rec := performLegacyRequest(t, engine, http.MethodPost, "/api/proxy-routes/", map[string]any{ + "site_name": "security-site", + "domain": "security.example.com", + "origin_url": "http://origin.security.internal:8080", + "enabled": true, + }, adminAuthHeaders(seed.Token)) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + proxyRouteID = uint(data["id"].(float64)) + assert.NotZero(t, proxyRouteID) + assert.Equal(t, "security.example.com", data["domain"]) + }) + + t.Run("bind WAF rule group to proxy route", func(t *testing.T) { + rec := performLegacyRequest( + t, + engine, + http.MethodPost, + fmt.Sprintf("/api/waf/sites/%d/rule-groups", proxyRouteID), + map[string]any{ + "ids": []uint{ruleGroupID}, + }, + adminAuthHeaders(seed.Token), + ) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + assert.Equal(t, float64(proxyRouteID), data["route_id"]) + + appliedIDs, ok := data["applied_ids"].([]any) + require.True(t, ok) + require.Len(t, appliedIDs, 1) + assert.Equal(t, float64(ruleGroupID), appliedIDs[0]) + }) + + t.Run("verify site rule groups binding", func(t *testing.T) { + rec := performLegacyRequest( + t, + engine, + http.MethodGet, + fmt.Sprintf("/api/waf/sites/%d/rule-groups", proxyRouteID), + nil, + adminAuthHeaders(seed.Token), + ) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + assert.NotNil(t, data["global_rule_group"]) + + appliedGroups, ok := data["applied_rule_groups"].([]any) + require.True(t, ok) + require.Len(t, appliedGroups, 1) + group, ok := appliedGroups[0].(map[string]any) + require.True(t, ok) + assert.Equal(t, float64(ruleGroupID), group["id"]) + }) + + t.Run("create TLS certificate with PEM", func(t *testing.T) { + certPEM, keyPEM := generateSelfSignedCertificatePair(t, []string{"security.example.com"}) + + rec := performLegacyRequest(t, engine, http.MethodPost, "/api/tls-certificates/", map[string]any{ + "name": "security-cert", + "cert_pem": certPEM, + "key_pem": keyPEM, + "remark": "self-signed integration cert", + }, adminAuthHeaders(seed.Token)) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + certID = uint(data["id"].(float64)) + assert.NotZero(t, certID) + assert.Equal(t, "security-cert", data["name"]) + assert.Equal(t, "upload", data["provider"]) + }) + + t.Run("create managed domain", func(t *testing.T) { + rec := performLegacyRequest(t, engine, http.MethodPost, "/api/managed-domains/", map[string]any{ + "domain": "security.example.com", + "cert_id": certID, + "enabled": true, + "remark": "primary security domain", + }, adminAuthHeaders(seed.Token)) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + domainID = uint(data["id"].(float64)) + assert.NotZero(t, domainID) + assert.Equal(t, "security.example.com", data["domain"]) + assert.Equal(t, float64(certID), data["cert_id"]) + assert.Equal(t, true, data["enabled"]) + }) + + t.Run("create DNS account", func(t *testing.T) { + rec := performLegacyRequest(t, engine, http.MethodPost, "/api/dns-accounts/", map[string]any{ + "name": "cloudflare-dns", + "type": "cloudflare", + "authorization": "test-api-token-value", + }, adminAuthHeaders(seed.Token)) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + data := unmarshalEnvelopeMap(t, envelope.Data) + dnsAccountID = uint(data["id"].(float64)) + assert.NotZero(t, dnsAccountID) + assert.Equal(t, "cloudflare-dns", data["name"]) + assert.Equal(t, "cloudflare", data["type"]) + // API 响应会脱敏 authorization,不应回显明文凭证。 + if auth, ok := data["authorization"]; ok { + assert.NotEqual(t, "test-api-token-value", auth) + } + }) + + t.Run("WAF rule group delete", func(t *testing.T) { + rec := performLegacyRequest( + t, + engine, + http.MethodPost, + fmt.Sprintf("/api/waf/rule-groups/%d/delete", ruleGroupID), + nil, + adminAuthHeaders(seed.Token), + ) + require.Equal(t, http.StatusOK, rec.Code) + + envelope := decodeEnvelope(t, rec) + require.True(t, envelope.Success, envelope.Message) + + detailRec := performLegacyRequest( + t, + engine, + http.MethodGet, + fmt.Sprintf("/api/waf/rule-groups/%d", ruleGroupID), + nil, + adminAuthHeaders(seed.Token), + ) + require.Equal(t, http.StatusOK, detailRec.Code) + detailEnvelope := decodeEnvelope(t, detailRec) + assert.False(t, detailEnvelope.Success) + }) + + _ = ipGroupID + _ = domainID + _ = dnsAccountID +} diff --git a/Wavelet/internal/apps/openflare/legacy/auth_admin.go b/Wavelet/internal/apps/openflare/legacy/auth_admin.go new file mode 100644 index 00000000..c7c3fa9d --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/auth_admin.go @@ -0,0 +1,251 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "fmt" + "strconv" + "strings" + + ofauth "github.com/Rain-kl/Wavelet/internal/apps/openflare/auth" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +type legacyUserPayload struct { + ID int `json:"id"` + Username string `json:"username"` + Password string `json:"password"` + DisplayName string `json:"display_name"` + Role int `json:"role"` + Email string `json:"email"` +} + +type manageUserRequest struct { + Username string `json:"username"` + Action string `json:"action"` +} + +type authSourcePayload struct { + Name string `json:"name"` + Type string `json:"type"` + DisplayName string `json:"display_name"` + IsActive bool `json:"is_active"` + ClientID string `json:"client_id"` + ClientSecret string `json:"client_secret"` + OpenIDDiscoveryURL string `json:"openid_discovery_url"` + Scopes string `json:"scopes"` + IconURL string `json:"icon_url"` +} + +type authSourceTogglePayload struct { + IsActive bool `json:"is_active"` +} + +func GetAllUsers(c *gin.Context) { + page, _ := strconv.Atoi(c.Query("p")) + users, err := ofauth.ListUsers(c.Request.Context(), page) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, users) +} + +func SearchUsers(c *gin.Context) { + users, err := ofauth.SearchUsers(c.Request.Context(), c.Query("keyword")) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, users) +} + +func GetUser(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + user, err := ofauth.GetUserByID(c.Request.Context(), callerRole(c), uint64(id)) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, user) +} + +func CreateUser(c *gin.Context) { + var req legacyUserPayload + if !compat.BindJSON(c, &req) { + return + } + if err := ofauth.CreateUser(c.Request.Context(), callerRole(c), ofauth.CreateUserInput{ + Username: req.Username, + Password: req.Password, + DisplayName: req.DisplayName, + Role: req.Role, + Email: req.Email, + }); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func UpdateUser(c *gin.Context) { + var req legacyUserPayload + if !compat.BindJSON(c, &req) { + return + } + if err := ofauth.UpdateUser(c.Request.Context(), callerRole(c), ofauth.UpdateUserInput{ + ID: req.ID, + Username: req.Username, + Password: req.Password, + DisplayName: req.DisplayName, + Role: req.Role, + Email: req.Email, + }); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func DeleteUser(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := ofauth.DeleteUserByID(c.Request.Context(), callerRole(c), uint64(id)); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func ManageUser(c *gin.Context) { + var req manageUserRequest + if !compat.BindJSON(c, &req) { + return + } + user, err := ofauth.ManageUser(c.Request.Context(), callerRole(c), ofauth.ManageUserInput{ + Username: req.Username, + Action: req.Action, + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, user) +} + +func ListAuthSources(c *gin.Context) { + sources, err := model.GetAuthSources(c.Request.Context()) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, sources) +} + +func CreateAuthSource(c *gin.Context) { + var payload authSourcePayload + if !compat.BindJSON(c, &payload) { + return + } + source := payload.toModel() + if err := model.CreateAuthSource(c.Request.Context(), &source); err != nil { + compat.Fail(c, err.Error()) + return + } + source.Sanitize() + compat.OK(c, source) +} + +func UpdateAuthSource(c *gin.Context) { + id, err := parseAuthSourceID(c) + if err != nil { + compat.Fail(c, err.Error()) + return + } + var payload authSourcePayload + if !compat.BindJSON(c, &payload) { + return + } + source := payload.toModel() + source.ID = id + keepSecret := strings.TrimSpace(source.ClientSecret) == "" + if err := model.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil { + compat.Fail(c, err.Error()) + return + } + updated, err := model.GetAuthSourceByID(c.Request.Context(), id) + if err != nil { + compat.Fail(c, err.Error()) + return + } + updated.Sanitize() + compat.OK(c, updated) +} + +func DeleteAuthSource(c *gin.Context) { + id, err := parseAuthSourceID(c) + if err != nil { + compat.Fail(c, err.Error()) + return + } + if err := model.DeleteAuthSource(c.Request.Context(), id); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func ToggleAuthSource(c *gin.Context) { + id, err := parseAuthSourceID(c) + if err != nil { + compat.Fail(c, err.Error()) + return + } + var payload authSourceTogglePayload + if !compat.BindJSON(c, &payload) { + return + } + if err := model.ToggleAuthSource(c.Request.Context(), id, payload.IsActive); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func (payload authSourcePayload) toModel() model.AuthSource { + return model.AuthSource{ + Name: payload.Name, + Type: payload.Type, + DisplayName: payload.DisplayName, + IsActive: payload.IsActive, + ClientID: payload.ClientID, + ClientSecret: payload.ClientSecret, + OpenIDDiscoveryURL: payload.OpenIDDiscoveryURL, + Scopes: payload.Scopes, + IconURL: payload.IconURL, + } +} + +func parseAuthSourceID(c *gin.Context) (uint64, error) { + raw := strings.TrimSpace(c.Param("id")) + if raw == "" { + return 0, fmt.Errorf("认证源 ID 无效") + } + source, err := model.GetAuthSourceByName(c.Request.Context(), raw) + if err == nil { + return source.ID, nil + } + parsed, err := strconv.ParseUint(raw, 10, 64) + if err != nil || parsed == 0 { + return 0, fmt.Errorf("认证源 ID 无效") + } + return parsed, nil +} diff --git a/Wavelet/internal/apps/openflare/legacy/auth_oauth.go b/Wavelet/internal/apps/openflare/legacy/auth_oauth.go new file mode 100644 index 00000000..7151e9c7 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/auth_oauth.go @@ -0,0 +1,87 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "strconv" + "strings" + + ofauth "github.com/Rain-kl/Wavelet/internal/apps/openflare/auth" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +type linkExistingRequest struct { + Username string `json:"username"` + Password string `json:"password"` +} + +// OAuthAuthorize starts OAuth authorization for a legacy auth source. +func OAuthAuthorize(c *gin.Context) { + url, err := ofauth.OAuthAuthorize(c.Request.Context(), c, c.Param("source")) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, gin.H{"authorize_url": url}) +} + +// OAuthCallback handles the legacy GET OAuth callback. +func OAuthCallback(c *gin.Context) { + result, err := ofauth.OAuthCallback(c.Request.Context(), c, c.Param("source")) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +// LinkExistingOAuthAccount binds a pending OAuth account to an existing user. +func LinkExistingOAuthAccount(c *gin.Context) { + var req linkExistingRequest + if !compat.BindJSON(c, &req) { + return + } + result, err := ofauth.LinkExistingOAuthAccount(c.Request.Context(), c, ofauth.LinkExistingInput{ + Username: req.Username, + Password: req.Password, + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +// ListExternalAccounts returns external account bindings for the current user. +func ListExternalAccounts(c *gin.Context) { + userID := callerUserID(c) + accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, accounts) +} + +// DeleteExternalAccount removes an external account binding. +func DeleteExternalAccount(c *gin.Context) { + userID := callerUserID(c) + if userID == 0 { + compat.Unauthorized(c, "无权进行此操作,未登录或 token 无效") + return + } + rawID := strings.TrimSpace(c.Param("id")) + id, err := strconv.ParseUint(rawID, 10, 64) + if err != nil || id == 0 { + compat.Fail(c, "绑定记录 ID 无效") + return + } + if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} diff --git a/Wavelet/internal/apps/openflare/legacy/auth_user.go b/Wavelet/internal/apps/openflare/legacy/auth_user.go new file mode 100644 index 00000000..f478ce66 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/auth_user.go @@ -0,0 +1,150 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + ofauth "github.com/Rain-kl/Wavelet/internal/apps/openflare/auth" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +type loginRequest struct { + Username string `json:"username"` + Password string `json:"password"` + Code string `json:"code"` +} + +type registerRequest struct { + Username string `json:"username"` + Password string `json:"password"` + Nickname string `json:"nickname"` + DisplayName string `json:"display_name"` + Email string `json:"email"` + Code string `json:"code"` +} + +type updateSelfRequest struct { + Username string `json:"username"` + Password string `json:"password"` + DisplayName string `json:"display_name"` + Email string `json:"email"` +} + +type passwordResetRequest struct { + Email string `json:"email"` + Token string `json:"token"` +} + +func Login(c *gin.Context) { + var req loginRequest + if !compat.BindJSON(c, &req) { + return + } + user, err := ofauth.Login(c.Request.Context(), c, ofauth.LoginInput{ + Username: req.Username, + Password: req.Password, + Code: req.Code, + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, user) +} + +func Logout(c *gin.Context) { + if err := ofauth.Logout(c.Request.Context(), c); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func GetSelf(c *gin.Context) { + user, err := ofauth.GetSelf(c.Request.Context(), callerUserID(c)) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, user) +} + +func UpdateSelf(c *gin.Context) { + var req updateSelfRequest + if !compat.BindJSON(c, &req) { + return + } + if err := ofauth.UpdateSelf(c.Request.Context(), callerUserID(c), ofauth.UpdateSelfInput{ + Username: req.Username, + Password: req.Password, + DisplayName: req.DisplayName, + Email: req.Email, + }); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func DeleteSelf(c *gin.Context) { + if err := ofauth.DeleteSelf(c.Request.Context(), callerUserID(c)); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func GenerateToken(c *gin.Context) { + token, err := ofauth.GenerateUserToken(c.Request.Context(), callerUserID(c)) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, token) +} + +func Register(c *gin.Context) { + var req registerRequest + if !compat.BindJSON(c, &req) { + return + } + user, err := ofauth.Register(c.Request.Context(), c, ofauth.RegisterInput{ + Username: req.Username, + Password: req.Password, + Nickname: req.Nickname, + DisplayName: req.DisplayName, + Email: req.Email, + Code: req.Code, + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, user) +} + +func SendEmailVerification(c *gin.Context) { + email := c.Query("email") + if err := ofauth.SendRegisterVerificationEmail(c.Request.Context(), email); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func ResetPassword(c *gin.Context) { + var req passwordResetRequest + if !compat.BindJSON(c, &req) { + return + } + password, err := ofauth.ResetPassword(c.Request.Context(), ofauth.PasswordResetInput{ + Email: req.Email, + Token: req.Token, + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, password) +} diff --git a/Wavelet/internal/apps/openflare/legacy/cap.go b/Wavelet/internal/apps/openflare/legacy/cap.go new file mode 100644 index 00000000..7e61d2aa --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/cap.go @@ -0,0 +1,73 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "net/http" + + "github.com/Rain-kl/Wavelet/internal/apps/cap" + "github.com/gin-gonic/gin" +) + +type capRedeemRequest struct { + Token string `json:"token" binding:"required"` + Solutions []int `json:"solutions" binding:"required"` +} + +// GetCapChallenge generates a CAP challenge for the legacy frontend. +func GetCapChallenge(c *gin.Context) { + scope := c.Param("scope") + if scope == "" { + scope = c.Query("scope") + } + if scope == "" { + scope = "login" + } + + mgr := cap.GetDefaultManager() + resp, err := mgr.Generate(c.Request.Context(), scope) + if err != nil { + c.JSON(http.StatusInternalServerError, cap.RedeemResponse{ + Success: false, + Error: err.Error(), + }) + return + } + c.JSON(http.StatusOK, resp) +} + +// RedeemCapChallenge redeems a CAP challenge for the legacy frontend. +func RedeemCapChallenge(c *gin.Context) { + scope := c.Param("scope") + if scope == "" { + scope = c.Query("scope") + } + if scope == "" { + scope = "login" + } + + var req capRedeemRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, cap.RedeemResponse{ + Success: false, + Error: "无效的参数", + }) + return + } + + mgr := cap.GetDefaultManager() + resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, scope) + if err != nil { + c.JSON(http.StatusInternalServerError, cap.RedeemResponse{ + Success: false, + Error: err.Error(), + }) + return + } + if !resp.Success { + c.JSON(http.StatusBadRequest, resp) + return + } + c.JSON(http.StatusOK, resp) +} diff --git a/Wavelet/internal/apps/openflare/legacy/middleware.go b/Wavelet/internal/apps/openflare/legacy/middleware.go new file mode 100644 index 00000000..f0c1dff7 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/middleware.go @@ -0,0 +1,57 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/cap" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +// bridgeOpenFlareToken maps OpenFlare-Token to X-Access-Token for compat auth middleware. +func bridgeOpenFlareToken() gin.HandlerFunc { + return compat.BridgeOpenFlareToken() +} + +// legacyCapAuth verifies PoW CAPTCHA for legacy login using OpenFlare response format. +func legacyCapAuth(scope string) gin.HandlerFunc { + mgr := cap.GetDefaultManager() + return func(c *gin.Context) { + if !cap.ProtectionEnabled(c.Request.Context()) { + c.Next() + return + } + token := c.GetHeader("X-Cap-Token") + if token == "" { + compat.Fail(c, "缺少人机验证凭证") + c.Abort() + return + } + valid, err := mgr.VerifyToken(c.Request.Context(), token, scope) + if err != nil || !valid { + compat.Fail(c, "人机验证凭证无效或已过期") + c.Abort() + return + } + c.Next() + } +} + +func callerRole(c *gin.Context) int { + if role, ok := c.Get("of_role"); ok { + if v, ok := role.(int); ok { + return v + } + } + return 0 +} + +func callerUserID(c *gin.Context) uint64 { + if id, ok := c.Get("of_user_id"); ok { + if v, ok := id.(uint64); ok { + return v + } + } + return 0 +} diff --git a/Wavelet/internal/apps/openflare/legacy/password_reset.go b/Wavelet/internal/apps/openflare/legacy/password_reset.go new file mode 100644 index 00000000..a0d7fc7e --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/password_reset.go @@ -0,0 +1,20 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + ofauth "github.com/Rain-kl/Wavelet/internal/apps/openflare/auth" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +// SendPasswordResetEmail sends a password reset email for the legacy frontend. +func SendPasswordResetEmail(c *gin.Context) { + email := c.Query("email") + if err := ofauth.SendPasswordResetEmail(c.Request.Context(), email); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} diff --git a/Wavelet/internal/apps/openflare/legacy/register.go b/Wavelet/internal/apps/openflare/legacy/register.go new file mode 100644 index 00000000..68523ca8 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register.go @@ -0,0 +1,25 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package legacy registers OpenFlare /api/* compatibility routes for the old frontend. +package legacy + +import "github.com/gin-gonic/gin" + +// RegisterRoutes mounts all OpenFlare legacy API routes under the /api group. +func RegisterRoutes(apiGroup *gin.RouterGroup) { + registerAuthRoutes(apiGroup) + registerOptionRoutes(apiGroup) + registerOriginRoutes(apiGroup) + registerApplyLogRoutes(apiGroup) + registerProxyRouteRoutes(apiGroup) + registerNodeRoutes(apiGroup) + registerWAFRoutes(apiGroup) + registerTLSRoutes(apiGroup) + registerConfigVersionRoutes(apiGroup) + registerAgentRoutes(apiGroup) + registerPagesRoutes(apiGroup) + registerRelayFlaredRoutes(apiGroup) + registerDashboardObsRoutes(apiGroup) + registerMiscRoutes(apiGroup) +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_agent.go b/Wavelet/internal/apps/openflare/legacy/register_agent.go new file mode 100644 index 00000000..7743cb90 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_agent.go @@ -0,0 +1,13 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" + "github.com/gin-gonic/gin" +) + +func registerAgentRoutes(apiGroup *gin.RouterGroup) { + agent.RegisterRoutes(apiGroup) +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_apply_log.go b/Wavelet/internal/apps/openflare/legacy/register_apply_log.go new file mode 100644 index 00000000..19e96a6f --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_apply_log.go @@ -0,0 +1,19 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/apply_log" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +func registerApplyLogRoutes(apiGroup *gin.RouterGroup) { + applyLogRoute := apiGroup.Group("/apply-logs") + applyLogRoute.Use(compat.AdminAuth()) + { + applyLogRoute.GET("/", apply_log.GetApplyLogs) + applyLogRoute.POST("/cleanup", apply_log.CleanupApplyLogs) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_auth.go b/Wavelet/internal/apps/openflare/legacy/register_auth.go new file mode 100644 index 00000000..61dd49b1 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_auth.go @@ -0,0 +1,74 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +func registerAuthRoutes(apiGroup *gin.RouterGroup) { + // /status, /notice, /about are registered by T-OPTION (option.RegisterRoutes). + apiGroup.GET("/verification", SendEmailVerification) + apiGroup.GET("/reset_password", SendPasswordResetEmail) + apiGroup.POST("/user/reset", ResetPassword) + + oauthGroup := apiGroup.Group("/oauth") + { + oauthGroup.GET("/:source/authorize", OAuthAuthorize) + oauthGroup.GET("/:source/callback", OAuthCallback) + oauthGroup.POST("/link-existing", LinkExistingOAuthAccount) + + externalAccounts := oauthGroup.Group("/external-accounts") + externalAccounts.Use(bridgeOpenFlareToken(), compat.UserAuth()) + { + externalAccounts.GET("/", ListExternalAccounts) + externalAccounts.POST("/:id/delete", DeleteExternalAccount) + } + } + + capGroup := apiGroup.Group("/cap") + { + capGroup.POST("/:scope/challenge", GetCapChallenge) + capGroup.POST("/:scope/redeem", RedeemCapChallenge) + } + + userGroup := apiGroup.Group("/user") + { + userGroup.POST("/register", Register) + userGroup.POST("/login", legacyCapAuth("login"), Login) + userGroup.GET("/logout", Logout) + + selfGroup := userGroup.Group("/") + selfGroup.Use(bridgeOpenFlareToken(), compat.UserAuth()) + { + selfGroup.GET("/self", GetSelf) + selfGroup.POST("/self/update", UpdateSelf) + selfGroup.POST("/self/delete", DeleteSelf) + selfGroup.GET("/token", GenerateToken) + } + + adminGroup := userGroup.Group("/") + adminGroup.Use(bridgeOpenFlareToken(), compat.AdminAuth()) + { + adminGroup.GET("/", GetAllUsers) + adminGroup.GET("/search", SearchUsers) + adminGroup.GET("/:id", GetUser) + adminGroup.POST("/", CreateUser) + adminGroup.POST("/manage", ManageUser) + adminGroup.POST("/update", UpdateUser) + adminGroup.POST("/:id/delete", DeleteUser) + } + } + + authSourceGroup := apiGroup.Group("/auth-sources") + authSourceGroup.Use(bridgeOpenFlareToken(), compat.RootAuth()) + { + authSourceGroup.GET("/", ListAuthSources) + authSourceGroup.POST("/", CreateAuthSource) + authSourceGroup.POST("/:id/update", UpdateAuthSource) + authSourceGroup.POST("/:id/delete", DeleteAuthSource) + authSourceGroup.POST("/:id/toggle", ToggleAuthSource) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_config_version.go b/Wavelet/internal/apps/openflare/legacy/register_config_version.go new file mode 100644 index 00000000..2ed698c8 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_config_version.go @@ -0,0 +1,25 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/config_version" + "github.com/gin-gonic/gin" +) + +func registerConfigVersionRoutes(apiGroup *gin.RouterGroup) { + configVersionGroup := apiGroup.Group("/config-versions") + configVersionGroup.Use(compat.AdminAuth()) + { + configVersionGroup.GET("/", config_version.ListConfigVersionsHandler) + configVersionGroup.GET("/active", config_version.GetActiveConfigVersionHandler) + configVersionGroup.GET("/preview", config_version.PreviewConfigVersionHandler) + configVersionGroup.GET("/diff", config_version.DiffConfigVersionHandler) + configVersionGroup.GET("/:id", config_version.GetConfigVersionHandler) + configVersionGroup.POST("/publish", config_version.PublishConfigVersionHandler) + configVersionGroup.POST("/:id/activate", config_version.ActivateConfigVersionHandler) + configVersionGroup.POST("/cleanup", config_version.CleanupConfigVersionsHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_dashboard_obs.go b/Wavelet/internal/apps/openflare/legacy/register_dashboard_obs.go new file mode 100644 index 00000000..7a610672 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_dashboard_obs.go @@ -0,0 +1,15 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/dashboard" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/observability" + "github.com/gin-gonic/gin" +) + +func registerDashboardObsRoutes(apiGroup *gin.RouterGroup) { + dashboard.RegisterRoutes(apiGroup) + observability.RegisterRoutes(apiGroup) +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_misc.go b/Wavelet/internal/apps/openflare/legacy/register_misc.go new file mode 100644 index 00000000..1d8d5fdc --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_misc.go @@ -0,0 +1,22 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/update" + "github.com/gin-gonic/gin" +) + +func registerMiscRoutes(apiGroup *gin.RouterGroup) { + updateRoute := apiGroup.Group("/update") + updateRoute.Use(compat.RootAuth()) + { + updateRoute.GET("/latest-release", update.GetLatestReleaseHandler) + updateRoute.GET("/logs/ws", update.StreamServerUpgradeLogsHandler) + updateRoute.POST("/manual-upload", update.UploadManualServerBinaryHandler) + updateRoute.POST("/manual-upgrade", update.ConfirmManualServerUpgradeHandler) + updateRoute.POST("/upgrade", update.UpgradeServerHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_node.go b/Wavelet/internal/apps/openflare/legacy/register_node.go new file mode 100644 index 00000000..0bb8b5d4 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_node.go @@ -0,0 +1,29 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/node" + "github.com/gin-gonic/gin" +) + +func registerNodeRoutes(apiGroup *gin.RouterGroup) { + nodeRoute := apiGroup.Group("/nodes") + nodeRoute.Use(compat.AdminAuth()) + { + nodeRoute.GET("/bootstrap-token", node.GetBootstrapTokenHandler) + nodeRoute.POST("/bootstrap-token/rotate", node.RotateBootstrapTokenHandler) + nodeRoute.GET("/", node.ListNodesHandler) + nodeRoute.POST("/", node.CreateNodeHandler) + nodeRoute.GET("/:id/agent-release", node.GetAgentReleaseHandler) + nodeRoute.POST("/:id/update", node.UpdateNodeHandler) + nodeRoute.POST("/:id/delete", node.DeleteNodeHandler) + nodeRoute.POST("/:id/agent-update", node.RequestAgentUpdateHandler) + nodeRoute.POST("/:id/openresty-restart", node.RequestOpenrestyRestartHandler) + nodeRoute.POST("/:id/force-sync", node.RequestForceSyncHandler) + nodeRoute.GET("/:id/observability", node.GetObservabilityHandler) + nodeRoute.POST("/:id/observability/cleanup", node.CleanupHealthEventsHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_option.go b/Wavelet/internal/apps/openflare/legacy/register_option.go new file mode 100644 index 00000000..a972ad45 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_option.go @@ -0,0 +1,13 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/gin-gonic/gin" +) + +func registerOptionRoutes(apiGroup *gin.RouterGroup) { + option.RegisterRoutes(apiGroup) +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_origin.go b/Wavelet/internal/apps/openflare/legacy/register_origin.go new file mode 100644 index 00000000..322ac018 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_origin.go @@ -0,0 +1,22 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/origin" + "github.com/gin-gonic/gin" +) + +func registerOriginRoutes(apiGroup *gin.RouterGroup) { + originRoute := apiGroup.Group("/origins") + originRoute.Use(compat.AdminAuth()) + { + originRoute.GET("/", origin.GetOrigins) + originRoute.GET("/:id", origin.GetOrigin) + originRoute.POST("/", origin.CreateOriginHandler) + originRoute.POST("/:id/update", origin.UpdateOriginHandler) + originRoute.POST("/:id/delete", origin.DeleteOriginHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_pages.go b/Wavelet/internal/apps/openflare/legacy/register_pages.go new file mode 100644 index 00000000..1e59e504 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_pages.go @@ -0,0 +1,27 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/pages" + "github.com/gin-gonic/gin" +) + +func registerPagesRoutes(apiGroup *gin.RouterGroup) { + pagesRoute := apiGroup.Group("/pages") + pagesRoute.Use(compat.AdminAuth()) + { + pagesRoute.GET("/", pages.ListProjectsHandler) + pagesRoute.GET("/:id", pages.GetProjectHandler) + pagesRoute.POST("/", pages.CreateProjectHandler) + pagesRoute.POST("/:id/update", pages.UpdateProjectHandler) + pagesRoute.POST("/:id/delete", pages.DeleteProjectHandler) + pagesRoute.GET("/:id/deployments", pages.ListDeploymentsHandler) + pagesRoute.POST("/:id/deployments/upload", pages.UploadDeploymentHandler) + pagesRoute.POST("/:id/deployments/:deployment_id/activate", pages.ActivateDeploymentHandler) + pagesRoute.POST("/:id/deployments/:deployment_id/delete", pages.DeleteDeploymentHandler) + pagesRoute.GET("/deployments/:deployment_id/files", pages.ListDeploymentFilesHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_proxy_route.go b/Wavelet/internal/apps/openflare/legacy/register_proxy_route.go new file mode 100644 index 00000000..9683b0e9 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_proxy_route.go @@ -0,0 +1,22 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/proxy_route" + "github.com/gin-gonic/gin" +) + +func registerProxyRouteRoutes(apiGroup *gin.RouterGroup) { + proxyRouteGroup := apiGroup.Group("/proxy-routes") + proxyRouteGroup.Use(compat.AdminAuth()) + { + proxyRouteGroup.GET("/", proxy_route.GetProxyRoutes) + proxyRouteGroup.GET("/:id", proxy_route.GetProxyRouteHandler) + proxyRouteGroup.POST("/", proxy_route.CreateProxyRouteHandler) + proxyRouteGroup.POST("/:id/update", proxy_route.UpdateProxyRouteHandler) + proxyRouteGroup.POST("/:id/delete", proxy_route.DeleteProxyRouteHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_relay_flared.go b/Wavelet/internal/apps/openflare/legacy/register_relay_flared.go new file mode 100644 index 00000000..a54f847b --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_relay_flared.go @@ -0,0 +1,28 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/flared" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/relay" + "github.com/gin-gonic/gin" +) + +func registerRelayFlaredRoutes(apiGroup *gin.RouterGroup) { + relayRoute := apiGroup.Group("/relay") + relayRoute.Use(relay.RelayAuth()) + { + relayRoute.POST("/heartbeat", relay.PostHeartbeat) + relayRoute.GET("/ws", relay.GetWebSocket) + } + + flaredRoute := apiGroup.Group("/flared") + flaredRoute.Use(flared.TunnelAuth()) + { + flaredRoute.POST("/heartbeat", flared.PostHeartbeat) + flaredRoute.GET("/config/active", flared.GetActiveConfig) + flaredRoute.POST("/apply-log", flared.PostApplyLog) + flaredRoute.GET("/ws", flared.GetWebSocket) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_tls.go b/Wavelet/internal/apps/openflare/legacy/register_tls.go new file mode 100644 index 00000000..35de79eb --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_tls.go @@ -0,0 +1,53 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" + "github.com/gin-gonic/gin" +) + +func registerTLSRoutes(apiGroup *gin.RouterGroup) { + managedDomainRoute := apiGroup.Group("/managed-domains") + managedDomainRoute.Use(compat.AdminAuth()) + { + managedDomainRoute.GET("/", tls.GetManagedDomains) + managedDomainRoute.GET("/match", tls.MatchManagedDomainCertificateHandler) + managedDomainRoute.POST("/", tls.CreateManagedDomainHandler) + managedDomainRoute.POST("/:id/update", tls.UpdateManagedDomainHandler) + managedDomainRoute.POST("/:id/delete", tls.DeleteManagedDomainHandler) + } + + tlsCertificateRoute := apiGroup.Group("/tls-certificates") + tlsCertificateRoute.Use(compat.AdminAuth()) + { + tlsCertificateRoute.GET("/", tls.GetCertificates) + tlsCertificateRoute.GET("/:id", tls.GetCertificateDetail) + tlsCertificateRoute.GET("/:id/content", tls.GetCertificateContentHandler) + tlsCertificateRoute.POST("/", tls.CreateCertificateHandler) + tlsCertificateRoute.POST("/:id/update", tls.UpdateCertificateHandler) + tlsCertificateRoute.POST("/:id/update-acme", tls.UpdateACMECertificateHandler) + tlsCertificateRoute.POST("/:id/convert-acme", tls.ConvertCertificateToACMEHandler) + tlsCertificateRoute.POST("/import-file", tls.ImportCertificateFile) + tlsCertificateRoute.POST("/:id/delete", tls.DeleteCertificateHandler) + tlsCertificateRoute.POST("/apply", tls.ApplyCertificateHandler) + tlsCertificateRoute.POST("/:id/renew", tls.RenewCertificateHandler) + } + + acmeAccountRoute := apiGroup.Group("/acme-accounts") + acmeAccountRoute.Use(compat.AdminAuth()) + { + acmeAccountRoute.GET("/default", tls.GetDefaultAcmeAccountHandler) + } + + dnsAccountRoute := apiGroup.Group("/dns-accounts") + dnsAccountRoute.Use(compat.AdminAuth()) + { + dnsAccountRoute.GET("/", tls.GetDNSAccounts) + dnsAccountRoute.POST("/", tls.CreateDNSAccountHandler) + dnsAccountRoute.POST("/:id/update", tls.UpdateDNSAccountHandler) + dnsAccountRoute.POST("/:id/delete", tls.DeleteDNSAccountHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/register_waf.go b/Wavelet/internal/apps/openflare/legacy/register_waf.go new file mode 100644 index 00000000..f10f505e --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/register_waf.go @@ -0,0 +1,34 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" + "github.com/gin-gonic/gin" +) + +func registerWAFRoutes(apiGroup *gin.RouterGroup) { + wafRoute := apiGroup.Group("/waf") + wafRoute.Use(compat.AdminAuth()) + { + wafRoute.GET("/ip-groups", waf.ListIPGroupsHandler) + wafRoute.GET("/ip-groups/:id", waf.GetIPGroupHandler) + wafRoute.POST("/ip-groups", waf.CreateIPGroupHandler) + wafRoute.POST("/ip-groups/test", waf.TestIPGroupAutoConfigHandler) + wafRoute.POST("/ip-groups/:id/update", waf.UpdateIPGroupHandler) + wafRoute.POST("/ip-groups/:id/delete", waf.DeleteIPGroupHandler) + wafRoute.POST("/ip-groups/:id/sync", waf.SyncIPGroupHandler) + + wafRoute.GET("/rule-groups", waf.ListRuleGroupsHandler) + wafRoute.GET("/rule-groups/:id", waf.GetRuleGroupHandler) + wafRoute.POST("/rule-groups", waf.CreateRuleGroupHandler) + wafRoute.POST("/rule-groups/:id/update", waf.UpdateRuleGroupHandler) + wafRoute.POST("/rule-groups/:id/delete", waf.DeleteRuleGroupHandler) + wafRoute.POST("/rule-groups/:id/sites", waf.ReplaceRuleGroupSitesHandler) + + wafRoute.GET("/sites/:route_id/rule-groups", waf.GetSiteRuleGroupsHandler) + wafRoute.POST("/sites/:route_id/rule-groups", waf.ReplaceSiteRuleGroupsHandler) + } +} diff --git a/Wavelet/internal/apps/openflare/legacy/status_public.go b/Wavelet/internal/apps/openflare/legacy/status_public.go new file mode 100644 index 00000000..d4f35c51 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/status_public.go @@ -0,0 +1,30 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package legacy + +import ( + ofauth "github.com/Rain-kl/Wavelet/internal/apps/openflare/auth" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +// GetStatus returns public server status for the legacy frontend. +func GetStatus(c *gin.Context) { + data, err := ofauth.BuildPublicStatus(c.Request.Context()) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, data) +} + +// GetNotice returns the legacy notice content. +func GetNotice(c *gin.Context) { + compat.OK(c, ofauth.GetNotice(c.Request.Context())) +} + +// GetAbout returns the legacy about content. +func GetAbout(c *gin.Context) { + compat.OK(c, ofauth.GetAbout(c.Request.Context())) +} diff --git a/Wavelet/internal/apps/openflare/node/errs.go b/Wavelet/internal/apps/openflare/node/errs.go new file mode 100644 index 00000000..1715423a --- /dev/null +++ b/Wavelet/internal/apps/openflare/node/errs.go @@ -0,0 +1,20 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package node + +const ( + errNodeNameRequired = "节点名不能为空" + errNodeIPTooLong = "节点 IP 不能超过 64 个字符" + errNodeIPInvalid = "节点 IP 格式无效" + errNodeIPManualRequired = "锁定节点 IP 时必须填写节点 IP" + errNodeGeoNameTooLong = "节点位置名不能超过 128 个字符" + errNodeGeoCoordinateMismatch = "地图坐标必须同时填写纬度和经度" + errNodeGeoLatitudeInvalid = "纬度必须在 -90 到 90 之间" + errNodeGeoLongitudeInvalid = "经度必须在 -180 到 180 之间" + errNodeIDConflict = "节点标识生成冲突,请重试" + errNodeNotFound = "节点不存在" + errNodeForceSyncFailed = "节点不在线或通过 WebSocket 发送同步指令失败" + errAgentPreviewTagInvalid = "指定版本不是 preview 发布" + errAgentStableTagInvalid = "正式版更新不能选择 preview 发布" +) diff --git a/Wavelet/internal/apps/openflare/node/helpers.go b/Wavelet/internal/apps/openflare/node/helpers.go new file mode 100644 index 00000000..574bdcd0 --- /dev/null +++ b/Wavelet/internal/apps/openflare/node/helpers.go @@ -0,0 +1,456 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package node + +import ( + "context" + "crypto/rand" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "strconv" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + nodeStatusOnline = "online" + nodeStatusOffline = "offline" + nodeStatusPending = "pending" + openrestyStatusHealthy = "healthy" + openrestyStatusUnhealthy = "unhealthy" + openrestyStatusUnknown = "unknown" + githubReleasesAPIBase = "https://api.github.com/repos/%s/releases" +) + +type releaseChannel string + +const ( + releaseChannelStable releaseChannel = "stable" + releaseChannelPreview releaseChannel = "preview" +) + +var releaseHTTPClient = &http.Client{Timeout: 30 * time.Second} + +type githubReleaseResponse struct { + TagName string `json:"tag_name"` + Body string `json:"body"` + HTMLURL string `json:"html_url"` + PublishedAt string `json:"published_at"` + Prerelease bool `json:"prerelease"` + Draft bool `json:"draft"` +} + +func newRandomToken() (string, error) { + buf := make([]byte, 16) + if _, err := rand.Read(buf); err != nil { + return "", err + } + return hex.EncodeToString(buf), nil +} + +func newServerNodeID() (string, error) { + token, err := newRandomToken() + if err != nil { + return "", err + } + return "node-" + token, nil +} + +func normalizeNodeType(raw string) string { + switch strings.ToLower(strings.TrimSpace(raw)) { + case "tunnel_relay": + return "tunnel_relay" + case "tunnel_client": + return "tunnel_client" + default: + return "edge_node" + } +} + +func normalizeRelayPort(port int, defaultPort int) int { + if port <= 0 || port > 65535 { + return defaultPort + } + return port +} + +func normalizeReleaseChannel(channel string) releaseChannel { + switch strings.ToLower(strings.TrimSpace(channel)) { + case string(releaseChannelPreview): + return releaseChannelPreview + default: + return releaseChannelStable + } +} + +func (channel releaseChannel) String() string { + if channel == releaseChannelPreview { + return string(releaseChannelPreview) + } + return string(releaseChannelStable) +} + +func normalizeOpenrestyStatus(status string) string { + switch strings.ToLower(strings.TrimSpace(status)) { + case openrestyStatusHealthy: + return openrestyStatusHealthy + case openrestyStatusUnhealthy: + return openrestyStatusUnhealthy + default: + return openrestyStatusUnknown + } +} + +func cloneCoordinate(value *float64) *float64 { + if value == nil { + return nil + } + cloned := *value + return &cloned +} + +func resolveNodeIPManualOverride(input Input, existing *model.OpenFlareNode, normalizedIP string) bool { + if input.IPManualOverride != nil { + return *input.IPManualOverride + } + if existing == nil { + return strings.TrimSpace(normalizedIP) != "" + } + if existing.IPManualOverride { + return true + } + return strings.TrimSpace(normalizedIP) != "" && strings.TrimSpace(normalizedIP) != strings.TrimSpace(existing.IP) +} + +func normalizeNodeInput(input Input) (string, string, string, *float64, *float64, bool, error) { + name := strings.TrimSpace(input.Name) + ip := strings.TrimSpace(input.IP) + geoName := strings.TrimSpace(input.GeoName) + manualOverride := input.GeoManualOverride || geoName != "" || input.GeoLatitude != nil || input.GeoLongitude != nil + if len(ip) > 64 { + return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPTooLong) + } + if ip != "" && net.ParseIP(ip) == nil { + return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPInvalid) + } + if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" { + return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPManualRequired) + } + if len(geoName) > 128 { + return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoNameTooLong) + } + + geoLatitude := cloneCoordinate(input.GeoLatitude) + geoLongitude := cloneCoordinate(input.GeoLongitude) + if (geoLatitude == nil) != (geoLongitude == nil) { + return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoCoordinateMismatch) + } + if geoLatitude != nil && (*geoLatitude < -90 || *geoLatitude > 90) { + return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLatitudeInvalid) + } + if geoLongitude != nil && (*geoLongitude < -180 || *geoLongitude > 180) { + return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLongitudeInvalid) + } + + if !manualOverride { + return name, ip, "", nil, nil, false, nil + } + if geoLatitude == nil && geoLongitude == nil && geoName == "" { + return name, ip, "", nil, nil, false, nil + } + + return name, ip, geoName, geoLatitude, geoLongitude, true, nil +} + +func computeNodeStatus(node *model.OpenFlareNode) string { + if node == nil { + return nodeStatusOffline + } + if node.LastSeenAt == nil || node.LastSeenAt.IsZero() { + return nodeStatusPending + } + if time.Since(*node.LastSeenAt) > model.NodeOfflineThreshold { + return nodeStatusOffline + } + return nodeStatusOnline +} + +func nodeViewLastSeenAt(node *model.OpenFlareNode) any { + if node == nil || node.LastSeenAt == nil { + return time.Time{} + } + return *node.LastSeenAt +} + +func buildNodeView(node *model.OpenFlareNode) *View { + if node == nil { + return nil + } + status := computeNodeStatus(node) + view := &View{ + ID: node.ID, + NodeID: node.NodeID, + Name: node.Name, + IP: node.IP, + IPManualOverride: node.IPManualOverride, + GeoName: strings.TrimSpace(node.GeoName), + GeoLatitude: node.GeoLatitude, + GeoLongitude: node.GeoLongitude, + GeoManualOverride: node.GeoManualOverride, + AccessToken: node.AccessToken, + UpdateChannel: strings.TrimSpace(node.UpdateChannel), + UpdateTag: strings.TrimSpace(node.UpdateTag), + RestartOpenrestyRequested: node.RestartOpenrestyRequested, + Version: node.Version, + ExtVersion: node.ExtVersion, + OpenrestyStatus: normalizeOpenrestyStatus(node.OpenrestyStatus), + OpenrestyMessage: strings.TrimSpace(node.OpenrestyMessage), + Status: status, + CurrentVersion: node.CurrentVersion, + LastSeenAt: nodeViewLastSeenAt(node), + LastError: node.LastError, + CreatedAt: node.CreatedAt, + UpdatedAt: node.UpdatedAt, + AutoUpdateEnabled: node.AutoUpdateEnabled, + UpdateRequested: node.UpdateRequested, + NodeType: node.NodeType, + RelayBindPort: node.RelayBindPort, + RelayVhostHTTPPort: node.RelayVhostHTTPPort, + RelayAgentAccessAddr: node.RelayAgentAccessAddr, + RelayClientAccessAddr: node.RelayClientAccessAddr, + RelayClientProxyURL: node.RelayClientProxyURL, + RelayStatus: node.RelayStatus, + RelayWebServerEnabled: node.RelayWebServerEnabled, + } + if view.UpdateChannel == "" { + view.UpdateChannel = releaseChannelStable.String() + } + if view.NodeType == "" { + view.NodeType = "edge_node" + } + return view +} + +func buildNodeAgentReleaseView(node *model.OpenFlareNode, release *githubReleaseResponse, channel releaseChannel) *AgentReleaseInfo { + currentVersion := strings.TrimSpace(node.Version) + view := &AgentReleaseInfo{ + CurrentVersion: currentVersion, + Channel: channel.String(), + UpdateRequested: node.UpdateRequested, + RequestedChannel: normalizeReleaseChannel(node.UpdateChannel).String(), + RequestedTag: strings.TrimSpace(node.UpdateTag), + } + if release == nil { + return view + } + view.TagName = release.TagName + view.Body = release.Body + view.HTMLURL = release.HTMLURL + view.PublishedAt = release.PublishedAt + view.Prerelease = release.Prerelease + view.HasUpdate = isVersionNewer(currentVersion, release.TagName) + return view +} + +func isVersionNewer(current string, latest string) bool { + return compareVersions(current, latest) < 0 +} + +func compareVersions(local, remote string) int { + left := parseVersionInfo(local) + right := parseVersionInfo(remote) + if left.isDev { + if right.valid { + return -1 + } + return 0 + } + if !left.valid || !right.valid { + return 0 + } + + maxLen := len(left.numbers) + if len(right.numbers) > maxLen { + maxLen = len(right.numbers) + } + for index := 0; index < maxLen; index++ { + leftValue := 0 + rightValue := 0 + if index < len(left.numbers) { + leftValue = left.numbers[index] + } + if index < len(right.numbers) { + rightValue = right.numbers[index] + } + if leftValue < rightValue { + return -1 + } + if leftValue > rightValue { + return 1 + } + } + return 0 +} + +type versionInfo struct { + valid bool + isDev bool + numbers []int +} + +func parseVersionInfo(version string) versionInfo { + normalized := strings.TrimSpace(strings.TrimPrefix(version, "v")) + if normalized == "" || normalized == "dev" { + return versionInfo{isDev: strings.EqualFold(normalized, "dev")} + } + base := normalized + if separator := strings.IndexRune(normalized, '-'); separator >= 0 { + base = normalized[:separator] + } + segments := strings.Split(base, ".") + parts := make([]int, 0, len(segments)) + for _, segment := range segments { + segment = strings.TrimSpace(segment) + if segment == "" { + parts = append(parts, 0) + continue + } + numeric := strings.Builder{} + for _, r := range segment { + if r < '0' || r > '9' { + break + } + numeric.WriteRune(r) + } + if numeric.Len() == 0 { + parts = append(parts, 0) + continue + } + value, err := strconv.Atoi(numeric.String()) + if err != nil { + return versionInfo{} + } + parts = append(parts, value) + } + return versionInfo{valid: len(parts) > 0, numbers: parts} +} + +func fetchLatestGitHubRelease(ctx context.Context, repo string, channel releaseChannel) (*githubReleaseResponse, error) { + switch normalizeReleaseChannel(string(channel)) { + case releaseChannelPreview: + return fetchLatestPreviewGitHubRelease(ctx, repo) + default: + return fetchLatestStableGitHubRelease(ctx, repo) + } +} + +func fetchLatestStableGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) { + url := fmt.Sprintf(githubReleasesAPIBase+"/latest", strings.TrimSpace(repo)) + req, err := newGitHubReleaseRequest(ctx, url) + if err != nil { + return nil, fmt.Errorf("创建更新请求失败") + } + resp, err := releaseHTTPClient.Do(req) + if err != nil { + return nil, fmt.Errorf("获取最新版本失败: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status) + } + return decodeGitHubRelease(resp.Body) +} + +func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) { + url := fmt.Sprintf(githubReleasesAPIBase+"?per_page=20", strings.TrimSpace(repo)) + req, err := newGitHubReleaseRequest(ctx, url) + if err != nil { + return nil, fmt.Errorf("创建更新请求失败") + } + resp, err := releaseHTTPClient.Do(req) + if err != nil { + return nil, fmt.Errorf("获取 preview 版本失败: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status) + } + var releases []githubReleaseResponse + if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil { + return nil, fmt.Errorf("解析 preview 版本信息失败") + } + for _, release := range releases { + if release.Draft || !release.Prerelease { + continue + } + releaseCopy := release + return &releaseCopy, nil + } + return nil, fmt.Errorf("当前没有可用的 preview 发布") +} + +func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*githubReleaseResponse, error) { + tag = strings.TrimSpace(tag) + if tag == "" { + return nil, fmt.Errorf("缺少发布版本号") + } + url := fmt.Sprintf(githubReleasesAPIBase+"/tags/%s", strings.TrimSpace(repo), tag) + req, err := newGitHubReleaseRequest(ctx, url) + if err != nil { + return nil, fmt.Errorf("创建更新请求失败") + } + resp, err := releaseHTTPClient.Do(req) + if err != nil { + return nil, fmt.Errorf("获取指定版本失败: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode == http.StatusNotFound { + return nil, fmt.Errorf("未找到指定版本: %s", tag) + } + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status) + } + return decodeGitHubRelease(resp.Body) +} + +func newGitHubReleaseRequest(ctx context.Context, url string) (*http.Request, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/vnd.github+json") + req.Header.Set("User-Agent", "OpenFlare-Server") + return req, nil +} + +func decodeGitHubRelease(reader io.Reader) (*githubReleaseResponse, error) { + var release githubReleaseResponse + if err := json.NewDecoder(reader).Decode(&release); err != nil { + return nil, fmt.Errorf("解析版本信息失败") + } + return &release, nil +} + +func isUniqueConstraintError(err error) bool { + if err == nil { + return false + } + return strings.Contains(strings.ToLower(err.Error()), "unique") +} + +func setReleaseHTTPClientForTest(client *http.Client) *http.Client { + previous := releaseHTTPClient + if client == nil { + releaseHTTPClient = &http.Client{Timeout: 30 * time.Second} + } else { + releaseHTTPClient = client + } + return previous +} diff --git a/Wavelet/internal/apps/openflare/node/logics.go b/Wavelet/internal/apps/openflare/node/logics.go new file mode 100644 index 00000000..552da996 --- /dev/null +++ b/Wavelet/internal/apps/openflare/node/logics.go @@ -0,0 +1,393 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package node + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/observability" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// Input is the create/update node payload. +type Input struct { + Name string `json:"name"` + IP string `json:"ip"` + IPManualOverride *bool `json:"ip_manual_override"` + AutoUpdateEnabled bool `json:"auto_update_enabled"` + GeoName string `json:"geo_name"` + GeoLatitude *float64 `json:"geo_latitude"` + GeoLongitude *float64 `json:"geo_longitude"` + GeoManualOverride bool `json:"geo_manual_override"` + NodeType string `json:"node_type"` + RelayBindPort int `json:"relay_bind_port"` + RelayVhostHTTPPort int `json:"relay_vhost_http_port"` + RelayAgentAccessAddr string `json:"relay_agent_access_addr"` + RelayClientAccessAddr string `json:"relay_client_access_addr"` + RelayClientProxyURL string `json:"relay_client_proxy_url"` + RelayWebServerEnabled bool `json:"relay_web_server_enabled"` +} + +// AgentUpdateInput requests an agent self-update on a node. +type AgentUpdateInput struct { + Channel string `json:"channel"` + TagName string `json:"tag_name"` +} + +// AgentReleaseInfo describes the latest agent release for a node. +type AgentReleaseInfo struct { + TagName string `json:"tag_name"` + Body string `json:"body"` + HTMLURL string `json:"html_url"` + PublishedAt string `json:"published_at"` + CurrentVersion string `json:"current_version"` + HasUpdate bool `json:"has_update"` + Channel string `json:"channel"` + Prerelease bool `json:"prerelease"` + UpdateRequested bool `json:"update_requested"` + RequestedChannel string `json:"requested_channel"` + RequestedTag string `json:"requested_tag"` +} + +// BootstrapView exposes the global discovery token. +type BootstrapView struct { + DiscoveryToken string `json:"discovery_token"` +} + +// View is the admin-facing node representation. +type View struct { + ID uint `json:"id"` + NodeID string `json:"node_id"` + Name string `json:"name"` + IP string `json:"ip"` + IPManualOverride bool `json:"ip_manual_override"` + GeoName string `json:"geo_name"` + GeoLatitude *float64 `json:"geo_latitude"` + GeoLongitude *float64 `json:"geo_longitude"` + GeoManualOverride bool `json:"geo_manual_override"` + AccessToken string `json:"access_token"` + AutoUpdateEnabled bool `json:"auto_update_enabled"` + UpdateRequested bool `json:"update_requested"` + UpdateChannel string `json:"update_channel"` + UpdateTag string `json:"update_tag"` + RestartOpenrestyRequested bool `json:"restart_openresty_requested"` + Version string `json:"version"` + ExtVersion string `json:"ext_version"` + OpenrestyStatus string `json:"openresty_status"` + OpenrestyMessage string `json:"openresty_message"` + Status string `json:"status"` + CurrentVersion string `json:"current_version"` + LastSeenAt any `json:"last_seen_at"` + LastError string `json:"last_error"` + LatestApplyResult string `json:"latest_apply_result"` + LatestApplyMessage string `json:"latest_apply_message"` + LatestApplyChecksum string `json:"latest_apply_checksum"` + LatestMainConfigChecksum string `json:"latest_main_config_checksum"` + LatestRouteConfigChecksum string `json:"latest_route_config_checksum"` + LatestSupportFileCount int `json:"latest_support_file_count"` + LatestApplyAt *time.Time `json:"latest_apply_at"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + NodeType string `json:"node_type"` + RelayBindPort int `json:"relay_bind_port"` + RelayVhostHTTPPort int `json:"relay_vhost_http_port"` + RelayAgentAccessAddr string `json:"relay_agent_access_addr"` + RelayClientAccessAddr string `json:"relay_client_access_addr"` + RelayClientProxyURL string `json:"relay_client_proxy_url"` + RelayStatus string `json:"relay_status"` + RelayWebServerEnabled bool `json:"relay_web_server_enabled"` +} + +// ObservabilityQuery filters node observability data. +type ObservabilityQuery struct { + Hours int `json:"hours"` + Limit int `json:"limit"` +} + +// ObservabilityView is the node observability API response. +type ObservabilityView = observability.NodeView + +// HealthEventCleanupResult reports health event cleanup outcome. +type HealthEventCleanupResult = observability.HealthEventCleanupResult + +// ListNodes returns all node views with latest apply log metadata. +func ListNodes(ctx context.Context) ([]*View, error) { + nodes, err := model.ListOpenFlareNodes(ctx) + if err != nil { + return nil, err + } + nodeIDs := make([]string, 0, len(nodes)) + for _, node := range nodes { + nodeIDs = append(nodeIDs, node.NodeID) + } + latestLogs, err := model.GetLatestOpenFlareApplyLogsByNodeIDs(ctx, nodeIDs) + if err != nil { + return nil, err + } + views := make([]*View, 0, len(nodes)) + for _, node := range nodes { + view := buildNodeView(&node) + view.Status = computeNodeStatus(&node) + if log, ok := latestLogs[node.NodeID]; ok { + view.LatestApplyResult = log.Result + view.LatestApplyMessage = log.Message + view.LatestApplyChecksum = log.Checksum + view.LatestMainConfigChecksum = log.MainConfigChecksum + view.LatestRouteConfigChecksum = log.RouteConfigChecksum + view.LatestSupportFileCount = log.SupportFileCount + view.LatestApplyAt = &log.CreatedAt + } + views = append(views, view) + } + return views, nil +} + +// CreateNode creates a reserved node with generated node_id and access_token. +func CreateNode(ctx context.Context, input Input) (*View, error) { + name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input) + if err != nil { + return nil, err + } + if name == "" { + return nil, errors.New(errNodeNameRequired) + } + ipManualOverride := resolveNodeIPManualOverride(input, nil, ip) + node := &model.OpenFlareNode{ + Name: name, + IP: ip, + IPManualOverride: ipManualOverride, + GeoName: geoName, + GeoLatitude: geoLatitude, + GeoLongitude: geoLongitude, + GeoManualOverride: geoManualOverride, + Version: "", + ExtVersion: "", + Status: nodeStatusPending, + AutoUpdateEnabled: input.AutoUpdateEnabled, + NodeType: normalizeNodeType(input.NodeType), + CapabilitiesJSON: "[]", + } + node.NodeID, err = newServerNodeID() + if err != nil { + return nil, err + } + node.AccessToken, err = newRandomToken() + if err != nil { + return nil, err + } + if node.NodeType == "tunnel_relay" { + node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, 7000) + node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, 8080) + node.RelayAuthToken, err = newRandomToken() + if err != nil { + return nil, err + } + node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr) + node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr) + node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL) + node.RelayWebServerEnabled = input.RelayWebServerEnabled + } + if err = model.CreateOpenFlareNode(ctx, node); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errNodeIDConflict) + } + return nil, err + } + return buildNodeView(node), nil +} + +// UpdateNode updates an existing node. +func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) { + name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input) + if err != nil { + return nil, err + } + if name == "" { + return nil, errors.New(errNodeNameRequired) + } + node, err := model.GetOpenFlareNodeByID(ctx, id) + if err != nil { + return nil, err + } + ipManualOverride := resolveNodeIPManualOverride(input, node, ip) + node.Name = name + node.IP = ip + node.IPManualOverride = ipManualOverride + node.GeoName = geoName + node.GeoLatitude = geoLatitude + node.GeoLongitude = geoLongitude + node.GeoManualOverride = geoManualOverride + node.AutoUpdateEnabled = input.AutoUpdateEnabled + if node.NodeType == "tunnel_relay" { + node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr) + node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr) + node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL) + node.RelayWebServerEnabled = input.RelayWebServerEnabled + if input.RelayBindPort > 0 { + node.RelayBindPort = input.RelayBindPort + } + if input.RelayVhostHTTPPort > 0 { + node.RelayVhostHTTPPort = input.RelayVhostHTTPPort + } + } + if err = model.SaveOpenFlareNode(ctx, node); err != nil { + return nil, err + } + return buildNodeView(node), nil +} + +// DeleteNode removes a node by id. +func DeleteNode(ctx context.Context, id uint) error { + if _, err := model.GetOpenFlareNodeByID(ctx, id); err != nil { + return err + } + return model.DeleteOpenFlareNode(ctx, id) +} + +// GetBootstrapToken returns the global discovery token, creating one if missing. +func GetBootstrapToken(ctx context.Context) (*BootstrapView, error) { + token, err := ensureGlobalDiscoveryToken(ctx) + if err != nil { + return nil, err + } + return &BootstrapView{DiscoveryToken: token}, nil +} + +// RotateBootstrapToken rotates the global discovery token. +func RotateBootstrapToken(ctx context.Context) (*BootstrapView, error) { + token, err := newRandomToken() + if err != nil { + return nil, err + } + if err = model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", token); err != nil { + return nil, err + } + return &BootstrapView{DiscoveryToken: token}, nil +} + +// GetAgentRelease checks the latest agent release for a node. +func GetAgentRelease(ctx context.Context, id uint, channel string) (*AgentReleaseInfo, error) { + node, err := model.GetOpenFlareNodeByID(ctx, id) + if err != nil { + return nil, err + } + release, err := fetchLatestGitHubRelease(ctx, model.AgentUpdateRepo, normalizeReleaseChannel(channel)) + if err != nil { + return nil, err + } + return buildNodeAgentReleaseView(node, release, normalizeReleaseChannel(channel)), nil +} + +// RequestAgentUpdate marks a node for manual agent update. +func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (*View, error) { + node, err := model.GetOpenFlareNodeByID(ctx, id) + if err != nil { + return nil, err + } + channel := normalizeReleaseChannel(input.Channel) + tagName := strings.TrimSpace(input.TagName) + if tagName != "" { + release, releaseErr := fetchGitHubReleaseByTag(ctx, model.AgentUpdateRepo, tagName) + if releaseErr != nil { + return nil, releaseErr + } + if channel == releaseChannelPreview && !release.Prerelease { + return nil, errors.New(errAgentPreviewTagInvalid) + } + if channel == releaseChannelStable && release.Prerelease { + return nil, errors.New(errAgentStableTagInvalid) + } + } + node.UpdateRequested = true + node.UpdateChannel = channel.String() + node.UpdateTag = tagName + if err = model.UpdateOpenFlareNodeFields(ctx, node, "update_requested", "update_channel", "update_tag"); err != nil { + return nil, err + } + return buildNodeView(node), nil +} + +// RequestOpenrestyRestart marks a node for openresty restart. +func RequestOpenrestyRestart(ctx context.Context, id uint) (*View, error) { + node, err := model.GetOpenFlareNodeByID(ctx, id) + if err != nil { + return nil, err + } + node.RestartOpenrestyRequested = true + if err = model.UpdateOpenFlareNodeFields(ctx, node, "restart_openresty_requested"); err != nil { + return nil, err + } + return buildNodeView(node), nil +} + +// RequestForceSync requests a force sync via websocket (stub until T-AGENT). +func RequestForceSync(ctx context.Context, id uint) (*View, error) { + node, err := model.GetOpenFlareNodeByID(ctx, id) + if err != nil { + return nil, err + } + if !sendAgentWSForceSyncConfig(node.NodeID) { + return nil, errors.New(errNodeForceSyncFailed) + } + return buildNodeView(node), nil +} + +// GetObservability returns observability details for a node. +func GetObservability(ctx context.Context, id uint, query ObservabilityQuery) (*ObservabilityView, error) { + return observability.GetNodeObservability(ctx, id, observability.NodeQuery{ + Hours: query.Hours, + Limit: query.Limit, + }) +} + +// CleanupHealthEvents removes all health events for a node. +func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) { + return observability.CleanupHealthEvents(ctx, id) +} + +func ensureGlobalDiscoveryToken(ctx context.Context) (string, error) { + if err := option.EnsureInitialized(ctx); err != nil { + return "", err + } + model.OptionMapRWMutex.RLock() + token := strings.TrimSpace(model.AgentDiscoveryToken) + model.OptionMapRWMutex.RUnlock() + if token != "" { + return token, nil + } + token, err := newRandomToken() + if err != nil { + return "", err + } + if err = model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", token); err != nil { + return "", err + } + return token, nil +} + +func sendAgentWSForceSyncConfig(nodeID string) bool { + _ = nodeID + return false +} + +// ValidateDiscoveryToken validates the global discovery token. +func ValidateDiscoveryToken(ctx context.Context, token string) error { + token = strings.TrimSpace(token) + if token == "" { + return fmt.Errorf("缺少 Discovery Token") + } + discoveryToken, err := ensureGlobalDiscoveryToken(ctx) + if err != nil { + return err + } + if token != discoveryToken { + return fmt.Errorf("Discovery Token 无效") + } + return nil +} diff --git a/Wavelet/internal/apps/openflare/node/logics_test.go b/Wavelet/internal/apps/openflare/node/logics_test.go new file mode 100644 index 00000000..dab0dfa4 --- /dev/null +++ b/Wavelet/internal/apps/openflare/node/logics_test.go @@ -0,0 +1,291 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package node + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/option" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupNodeTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.OpenFlareNode{}, + &model.OpenFlareOption{}, + &model.OpenFlareApplyLog{}, + )) + + db.SetDB(sqliteDB) + option.ResetInitializationForTest() + + return func() { + db.SetDB(nil) + option.ResetInitializationForTest() + } +} + +func TestCreateEdgeNode(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + view, err := CreateNode(ctx, Input{ + Name: "edge-1", + IP: "10.0.0.1", + AutoUpdateEnabled: true, + }) + require.NoError(t, err) + assert.NotZero(t, view.ID) + assert.True(t, strings.HasPrefix(view.NodeID, "node-")) + assert.Len(t, view.AccessToken, 32) + assert.Equal(t, "edge_node", view.NodeType) + assert.Equal(t, nodeStatusPending, view.Status) + assert.True(t, view.AutoUpdateEnabled) +} + +func TestCreateTunnelRelayNode(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + view, err := CreateNode(ctx, Input{ + Name: "relay-1", + NodeType: "tunnel_relay", + }) + require.NoError(t, err) + assert.Equal(t, "tunnel_relay", view.NodeType) + assert.Equal(t, 7000, view.RelayBindPort) + assert.Equal(t, 8080, view.RelayVhostHTTPPort) + + stored, err := model.GetOpenFlareNodeByID(ctx, view.ID) + require.NoError(t, err) + assert.NotEmpty(t, stored.RelayAuthToken) +} + +func TestCreateTunnelClientNode(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + view, err := CreateNode(ctx, Input{ + Name: "client-1", + NodeType: "tunnel_client", + }) + require.NoError(t, err) + assert.Equal(t, "tunnel_client", view.NodeType) +} + +func TestCreateNodeRequiresName(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + _, err := CreateNode(ctx, Input{IP: "10.0.0.2"}) + require.Error(t, err) + assert.Equal(t, errNodeNameRequired, err.Error()) +} + +func TestUpdateNode(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + created, err := CreateNode(ctx, Input{Name: "edge-update"}) + require.NoError(t, err) + + updated, err := UpdateNode(ctx, created.ID, Input{ + Name: "edge-updated", + IP: "192.168.1.10", + AutoUpdateEnabled: true, + }) + require.NoError(t, err) + assert.Equal(t, "edge-updated", updated.Name) + assert.Equal(t, "192.168.1.10", updated.IP) + assert.True(t, updated.AutoUpdateEnabled) +} + +func TestDeleteNode(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + created, err := CreateNode(ctx, Input{Name: "edge-delete"}) + require.NoError(t, err) + + require.NoError(t, DeleteNode(ctx, created.ID)) + _, err = model.GetOpenFlareNodeByID(ctx, created.ID) + require.Error(t, err) + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) +} + +func TestListNodesWithApplyLogMetadata(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + created, err := CreateNode(ctx, Input{Name: "edge-list"}) + require.NoError(t, err) + + applyAt := time.Now().UTC().Truncate(time.Second) + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{ + NodeID: created.NodeID, + Version: "20260618-001", + Result: "success", + Message: "ok", + Checksum: "checksum-1", + MainConfigChecksum: "main-1", + RouteConfigChecksum: "route-1", + SupportFileCount: 3, + CreatedAt: applyAt, + }).Error) + + views, err := ListNodes(ctx) + require.NoError(t, err) + require.Len(t, views, 1) + assert.Equal(t, "success", views[0].LatestApplyResult) + assert.Equal(t, "checksum-1", views[0].LatestApplyChecksum) + assert.Equal(t, 3, views[0].LatestSupportFileCount) + require.NotNil(t, views[0].LatestApplyAt) + assert.Equal(t, applyAt, views[0].LatestApplyAt.UTC()) +} + +func TestBootstrapTokenLifecycle(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + first, err := GetBootstrapToken(ctx) + require.NoError(t, err) + assert.Len(t, first.DiscoveryToken, 32) + + second, err := GetBootstrapToken(ctx) + require.NoError(t, err) + assert.Equal(t, first.DiscoveryToken, second.DiscoveryToken) + + rotated, err := RotateBootstrapToken(ctx) + require.NoError(t, err) + assert.NotEqual(t, first.DiscoveryToken, rotated.DiscoveryToken) + assert.Equal(t, rotated.DiscoveryToken, model.OptionValue("AgentDiscoveryToken")) +} + +func TestValidateDiscoveryToken(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + bootstrap, err := GetBootstrapToken(ctx) + require.NoError(t, err) + + require.NoError(t, ValidateDiscoveryToken(ctx, bootstrap.DiscoveryToken)) + require.Error(t, ValidateDiscoveryToken(ctx, "invalid-token")) +} + +func TestRequestAgentUpdateWithPreviewTag(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + created, err := CreateNode(ctx, Input{Name: "edge-update-agent"}) + require.NoError(t, err) + + originalClient := setReleaseHTTPClientForTest(&http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + expected := "https://api.github.com/repos/" + model.AgentUpdateRepo + "/releases/tags/v0.5.0-rc.1" + require.Equal(t, expected, req.URL.String()) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"tag_name":"v0.5.0-rc.1","prerelease":true}`)), + }, nil + }), + }) + t.Cleanup(func() { + setReleaseHTTPClientForTest(originalClient) + }) + + updated, err := RequestAgentUpdate(ctx, created.ID, AgentUpdateInput{ + Channel: "preview", + TagName: "v0.5.0-rc.1", + }) + require.NoError(t, err) + assert.True(t, updated.UpdateRequested) + assert.Equal(t, "preview", updated.UpdateChannel) + assert.Equal(t, "v0.5.0-rc.1", updated.UpdateTag) +} + +func TestRequestOpenrestyRestart(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + created, err := CreateNode(ctx, Input{Name: "edge-restart"}) + require.NoError(t, err) + + updated, err := RequestOpenrestyRestart(ctx, created.ID) + require.NoError(t, err) + assert.True(t, updated.RestartOpenrestyRequested) +} + +func TestRequestForceSyncStub(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + created, err := CreateNode(ctx, Input{Name: "edge-sync"}) + require.NoError(t, err) + + _, err = RequestForceSync(ctx, created.ID) + require.Error(t, err) + assert.Equal(t, errNodeForceSyncFailed, err.Error()) +} + +func TestGetObservabilityStub(t *testing.T) { + cleanup := setupNodeTestDB(t) + defer cleanup() + ctx := context.Background() + + created, err := CreateNode(ctx, Input{Name: "edge-obs"}) + require.NoError(t, err) + + view, err := GetObservability(ctx, created.ID, ObservabilityQuery{Hours: 24, Limit: 50}) + require.NoError(t, err) + assert.Equal(t, created.NodeID, view.NodeID) + assert.Empty(t, view.MetricSnapshots) +} + +func TestComputeNodeStatus(t *testing.T) { + now := time.Now() + pending := &model.OpenFlareNode{} + assert.Equal(t, nodeStatusPending, computeNodeStatus(pending)) + + online := &model.OpenFlareNode{LastSeenAt: &now} + assert.Equal(t, nodeStatusOnline, computeNodeStatus(online)) + + offlineAt := now.Add(-model.NodeOfflineThreshold - time.Minute) + offline := &model.OpenFlareNode{LastSeenAt: &offlineAt} + assert.Equal(t, nodeStatusOffline, computeNodeStatus(offline)) +} + +type roundTripFunc func(req *http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} diff --git a/Wavelet/internal/apps/openflare/node/routers.go b/Wavelet/internal/apps/openflare/node/routers.go new file mode 100644 index 00000000..de4fe555 --- /dev/null +++ b/Wavelet/internal/apps/openflare/node/routers.go @@ -0,0 +1,192 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package node + +import ( + "encoding/json" + "errors" + "io" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +func handleLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, errNodeNotFound) + return true + } + compat.Fail(c, err.Error()) + return true +} + +// ListNodesHandler lists all nodes. +func ListNodesHandler(c *gin.Context) { + nodes, err := ListNodes(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, nodes) +} + +// CreateNodeHandler creates a node. +func CreateNodeHandler(c *gin.Context) { + var input Input + if !compat.BindJSON(c, &input) { + return + } + view, err := CreateNode(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// UpdateNodeHandler updates a node. +func UpdateNodeHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input Input + if !compat.BindJSON(c, &input) { + return + } + view, err := UpdateNode(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// DeleteNodeHandler deletes a node. +func DeleteNodeHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteNode(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OKMessage(c, "") +} + +// GetBootstrapTokenHandler returns the global discovery token. +func GetBootstrapTokenHandler(c *gin.Context) { + view, err := GetBootstrapToken(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// RotateBootstrapTokenHandler rotates the global discovery token. +func RotateBootstrapTokenHandler(c *gin.Context) { + view, err := RotateBootstrapToken(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// GetAgentReleaseHandler returns the latest agent release for a node. +func GetAgentReleaseHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + release, err := GetAgentRelease(c.Request.Context(), id, c.Query("channel")) + if handleLogicError(c, err) { + return + } + compat.OK(c, release) +} + +// RequestAgentUpdateHandler requests agent self-update on a node. +func RequestAgentUpdateHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var request AgentUpdateInput + if c.Request.ContentLength > 0 { + if err := bindOptionalJSON(c.Request.Body, &request); err != nil { + compat.Fail(c, "参数错误") + return + } + } + view, err := RequestAgentUpdate(c.Request.Context(), id, request) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// RequestOpenrestyRestartHandler requests openresty restart on a node. +func RequestOpenrestyRestartHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + view, err := RequestOpenrestyRestart(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// RequestForceSyncHandler requests force sync on a node. +func RequestForceSyncHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + view, err := RequestForceSync(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// GetObservabilityHandler returns node observability details. +func GetObservabilityHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var query ObservabilityQuery + if err := c.ShouldBindQuery(&query); err != nil { + compat.Fail(c, "参数错误") + return + } + view, err := GetObservability(c.Request.Context(), id, query) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// CleanupHealthEventsHandler cleans up node health events. +func CleanupHealthEventsHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + result, err := CleanupHealthEvents(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, result) +} + +func bindOptionalJSON(body io.Reader, target any) error { + if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) { + return err + } + return nil +} diff --git a/Wavelet/internal/apps/openflare/observability/access_log_logics.go b/Wavelet/internal/apps/openflare/observability/access_log_logics.go new file mode 100644 index 00000000..0b2d44e3 --- /dev/null +++ b/Wavelet/internal/apps/openflare/observability/access_log_logics.go @@ -0,0 +1,636 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package observability + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + defaultAccessLogPageSize = 20 + maxAccessLogPageSize = 200 + defaultAccessLogSortBy = "logged_at" + defaultAccessLogSortOrder = "desc" + defaultAccessLogFoldMinute = 3 + defaultIPTrendHours = 24 + defaultIPTrendBucketMinute = 30 + maxIPTrendHours = 168 + nodeAccessLogRetentionDays = 90 +) + +var nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour + +// AccessLogQuery filters access log list queries. +type AccessLogQuery struct { + NodeID string `json:"node_id"` + RemoteAddr string `json:"remote_addr"` + Host string `json:"host"` + Path string `json:"path"` + Page int `json:"page"` + PageSize int `json:"page_size"` + SortBy string `json:"sort_by"` + SortOrder string `json:"sort_order"` + FoldMinutes int `json:"fold_minutes"` +} + +// AccessLogView is a single access log row. +type AccessLogView struct { + ID uint `json:"id"` + NodeID string `json:"node_id"` + NodeName string `json:"node_name"` + LoggedAt time.Time `json:"logged_at"` + RemoteAddr string `json:"remote_addr"` + Region string `json:"region"` + Host string `json:"host"` + Path string `json:"path"` + StatusCode int `json:"status_code"` +} + +// AccessLogList is a paginated access log response. +type AccessLogList struct { + Items []AccessLogView `json:"items"` + Page int `json:"page"` + PageSize int `json:"page_size"` + HasMore bool `json:"has_more"` + TotalRecord int64 `json:"total_record"` + TotalIP int64 `json:"total_ip"` +} + +// FoldedAccessLogView is a folded access log bucket. +type FoldedAccessLogView struct { + BucketStartedAt time.Time `json:"bucket_started_at"` + RequestCount int64 `json:"request_count"` + UniqueIPCount int64 `json:"unique_ip_count"` + UniqueHostCount int64 `json:"unique_host_count"` + SuccessCount int64 `json:"success_count"` + ClientErrorCount int64 `json:"client_error_count"` + ServerErrorCount int64 `json:"server_error_count"` +} + +// FoldedAccessLogList is a paginated folded access log response. +type FoldedAccessLogList struct { + Items []FoldedAccessLogView `json:"items"` + Page int `json:"page"` + PageSize int `json:"page_size"` + HasMore bool `json:"has_more"` + TotalBucket int64 `json:"total_bucket"` + TotalRecord int64 `json:"total_record"` + TotalIP int64 `json:"total_ip"` + FoldMinutes int `json:"fold_minutes"` +} + +// FoldedAccessLogIPQuery filters folded IP summary queries. +type FoldedAccessLogIPQuery struct { + NodeID string `json:"node_id"` + RemoteAddr string `json:"remote_addr"` + Host string `json:"host"` + Path string `json:"path"` + BucketStartedAt string `json:"bucket_started_at"` + FoldMinutes int `json:"fold_minutes"` + Page int `json:"page"` + PageSize int `json:"page_size"` + SortBy string `json:"sort_by"` + SortOrder string `json:"sort_order"` +} + +// FoldedAccessLogIPView is a folded IP row. +type FoldedAccessLogIPView struct { + RemoteAddr string `json:"remote_addr"` + RequestCount int64 `json:"request_count"` + SuccessCount int64 `json:"success_count"` + ClientErrorCount int64 `json:"client_error_count"` + ServerErrorCount int64 `json:"server_error_count"` + LastSeenAt time.Time `json:"last_seen_at"` +} + +// FoldedAccessLogIPList is a paginated folded IP response. +type FoldedAccessLogIPList struct { + Items []FoldedAccessLogIPView `json:"items"` + Page int `json:"page"` + PageSize int `json:"page_size"` + HasMore bool `json:"has_more"` + TotalIP int64 `json:"total_ip"` + BucketStartedAt time.Time `json:"bucket_started_at"` + FoldMinutes int `json:"fold_minutes"` + SortBy string `json:"sort_by"` + SortOrder string `json:"sort_order"` +} + +// AccessLogIPSummaryQuery filters IP summary list queries. +type AccessLogIPSummaryQuery struct { + NodeID string `json:"node_id"` + RemoteAddr string `json:"remote_addr"` + Host string `json:"host"` + Page int `json:"page"` + PageSize int `json:"page_size"` + SortBy string `json:"sort_by"` + SortOrder string `json:"sort_order"` +} + +// AccessLogIPSummaryView is an IP summary row. +type AccessLogIPSummaryView struct { + RemoteAddr string `json:"remote_addr"` + TotalRequests int64 `json:"total_requests"` + RecentRequests int64 `json:"recent_requests"` + LastSeenAt time.Time `json:"last_seen_at"` +} + +// AccessLogIPSummaryList is a paginated IP summary response. +type AccessLogIPSummaryList struct { + Items []AccessLogIPSummaryView `json:"items"` + Page int `json:"page"` + PageSize int `json:"page_size"` + HasMore bool `json:"has_more"` + TotalIP int64 `json:"total_ip"` + SortBy string `json:"sort_by"` + SortOrder string `json:"sort_order"` +} + +// AccessLogIPTrendQuery filters IP trend queries. +type AccessLogIPTrendQuery struct { + NodeID string `json:"node_id"` + RemoteAddr string `json:"remote_addr"` + Host string `json:"host"` + Hours int `json:"hours"` + BucketMinutes int `json:"bucket_minutes"` +} + +// AccessLogIPTrendPoint is an IP trend bucket. +type AccessLogIPTrendPoint struct { + BucketStartedAt time.Time `json:"bucket_started_at"` + RequestCount int64 `json:"request_count"` +} + +// AccessLogIPTrendView is the IP trend response. +type AccessLogIPTrendView struct { + RemoteAddr string `json:"remote_addr"` + Hours int `json:"hours"` + BucketMinutes int `json:"bucket_minutes"` + Points []AccessLogIPTrendPoint `json:"points"` +} + +// AccessLogCleanupInput is the cleanup request payload. +type AccessLogCleanupInput struct { + RetentionDays int `json:"retention_days"` +} + +// AccessLogCleanupResult is the cleanup response payload. +type AccessLogCleanupResult struct { + RetentionDays int `json:"retention_days"` + DeletedCount int64 `json:"deleted_count"` + Cutoff time.Time `json:"cutoff"` +} + +// ListAccessLogs returns paginated access logs. +func ListAccessLogs(ctx context.Context, input AccessLogQuery) (*AccessLogList, error) { + normalized := normalizeAccessLogQuery(input) + modelQuery := buildModelAccessLogQuery(normalized) + logs, err := model.ListOpenFlareAccessLogs(ctx, modelQuery) + if err != nil { + return nil, err + } + totalRecords, totalIPs, err := model.CountOpenFlareAccessLogs(ctx, modelQuery) + if err != nil { + return nil, err + } + nodeNames, err := listNodeNameMap(ctx, logs) + if err != nil { + return nil, err + } + views := make([]AccessLogView, 0, len(logs)) + for _, item := range logs { + if item == nil { + continue + } + views = append(views, AccessLogView{ + ID: item.ID, + NodeID: item.NodeID, + NodeName: nodeNames[item.NodeID], + LoggedAt: item.LoggedAt, + RemoteAddr: item.RemoteAddr, + Region: item.Region, + Host: item.Host, + Path: item.Path, + StatusCode: item.StatusCode, + }) + } + return &AccessLogList{ + Items: views, + Page: normalized.Page, + PageSize: normalized.PageSize, + HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalRecords, + TotalRecord: totalRecords, + TotalIP: totalIPs, + }, nil +} + +// ListFoldedAccessLogs returns paginated folded access logs. +func ListFoldedAccessLogs(ctx context.Context, input AccessLogQuery) (*FoldedAccessLogList, error) { + normalized := normalizeAccessLogQuery(input) + foldMinutes, err := normalizeFoldMinutes(normalized.FoldMinutes) + if err != nil { + return nil, err + } + modelQuery := buildModelAccessLogQuery(normalized) + bucketQuery := model.OpenFlareAccessLogBucketQuery{ + NodeID: modelQuery.NodeID, + RemoteAddr: modelQuery.RemoteAddr, + Host: modelQuery.Host, + Path: modelQuery.Path, + Since: modelQuery.Since, + Page: normalized.Page, + PageSize: normalized.PageSize, + SortBy: normalizeFoldSortBy(input.SortBy), + SortOrder: normalized.SortOrder, + FoldMinutes: foldMinutes, + } + items, err := model.ListOpenFlareAccessLogBuckets(ctx, bucketQuery) + if err != nil { + return nil, err + } + totalBuckets, err := model.CountOpenFlareAccessLogBuckets(ctx, bucketQuery) + if err != nil { + return nil, err + } + totalRecords, totalIPs, err := model.CountOpenFlareAccessLogs(ctx, modelQuery) + if err != nil { + return nil, err + } + views := make([]FoldedAccessLogView, 0, len(items)) + for _, item := range items { + if item == nil { + continue + } + views = append(views, FoldedAccessLogView{ + BucketStartedAt: time.Unix(item.BucketEpoch, 0).UTC(), + RequestCount: item.RequestCount, + UniqueIPCount: item.UniqueIPCount, + UniqueHostCount: item.UniqueHostCount, + SuccessCount: item.SuccessCount, + ClientErrorCount: item.ClientErrorCount, + ServerErrorCount: item.ServerErrorCount, + }) + } + return &FoldedAccessLogList{ + Items: views, + Page: normalized.Page, + PageSize: normalized.PageSize, + HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalBuckets, + TotalBucket: totalBuckets, + TotalRecord: totalRecords, + TotalIP: totalIPs, + FoldMinutes: foldMinutes, + }, nil +} + +// ListFoldedAccessLogIPs returns paginated folded IP summaries. +func ListFoldedAccessLogIPs(ctx context.Context, input FoldedAccessLogIPQuery) (*FoldedAccessLogIPList, error) { + normalized, bucketStartedAt, err := normalizeFoldedAccessLogIPQuery(input) + if err != nil { + return nil, err + } + modelQuery := model.OpenFlareAccessLogBucketIPQuery{ + NodeID: normalized.NodeID, + RemoteAddr: normalized.RemoteAddr, + Host: normalized.Host, + Path: normalized.Path, + BucketStartedAt: bucketStartedAt, + FoldMinutes: normalized.FoldMinutes, + Page: normalized.Page, + PageSize: normalized.PageSize, + SortBy: normalized.SortBy, + SortOrder: normalized.SortOrder, + } + items, err := model.ListOpenFlareAccessLogBucketIPs(ctx, modelQuery) + if err != nil { + return nil, err + } + totalIP, err := model.CountOpenFlareAccessLogBucketIPs(ctx, modelQuery) + if err != nil { + return nil, err + } + views := make([]FoldedAccessLogIPView, 0, len(items)) + for _, item := range items { + if item == nil { + continue + } + views = append(views, FoldedAccessLogIPView{ + RemoteAddr: item.RemoteAddr, + RequestCount: item.RequestCount, + SuccessCount: item.SuccessCount, + ClientErrorCount: item.ClientErrorCount, + ServerErrorCount: item.ServerErrorCount, + LastSeenAt: time.Unix(item.LastSeenEpoch, 0).UTC(), + }) + } + return &FoldedAccessLogIPList{ + Items: views, + Page: normalized.Page, + PageSize: normalized.PageSize, + HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalIP, + TotalIP: totalIP, + BucketStartedAt: bucketStartedAt, + FoldMinutes: normalized.FoldMinutes, + SortBy: normalized.SortBy, + SortOrder: normalized.SortOrder, + }, nil +} + +// ListAccessLogIPSummaries returns paginated IP summaries. +func ListAccessLogIPSummaries(ctx context.Context, input AccessLogIPSummaryQuery) (*AccessLogIPSummaryList, error) { + normalized := normalizeAccessLogIPSummaryQuery(input) + since := time.Now().UTC().Add(-nodeAccessLogRetentionWindow) + recentSince := time.Now().UTC().Add(-3 * time.Hour) + query := model.OpenFlareAccessLogIPSummaryQuery{ + NodeID: strings.TrimSpace(normalized.NodeID), + RemoteAddr: strings.TrimSpace(normalized.RemoteAddr), + Host: strings.TrimSpace(normalized.Host), + Since: since, + Page: normalized.Page, + PageSize: normalized.PageSize, + SortBy: normalized.SortBy, + SortOrder: normalized.SortOrder, + } + items, err := model.ListOpenFlareAccessLogIPSummaries(ctx, query, recentSince) + if err != nil { + return nil, err + } + totalIP, err := model.CountOpenFlareAccessLogIPSummaries(ctx, query) + if err != nil { + return nil, err + } + views := make([]AccessLogIPSummaryView, 0, len(items)) + for _, item := range items { + if item == nil { + continue + } + views = append(views, AccessLogIPSummaryView{ + RemoteAddr: item.RemoteAddr, + TotalRequests: item.TotalRequests, + RecentRequests: item.RecentRequests, + LastSeenAt: time.Unix(item.LastSeenEpoch, 0).UTC(), + }) + } + return &AccessLogIPSummaryList{ + Items: views, + Page: normalized.Page, + PageSize: normalized.PageSize, + HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalIP, + TotalIP: totalIP, + SortBy: normalized.SortBy, + SortOrder: normalized.SortOrder, + }, nil +} + +// GetAccessLogIPTrend returns IP request trend points. +func GetAccessLogIPTrend(ctx context.Context, input AccessLogIPTrendQuery) (*AccessLogIPTrendView, error) { + normalized, err := normalizeAccessLogIPTrendQuery(input) + if err != nil { + return nil, err + } + points, err := model.ListOpenFlareAccessLogIPTrend(ctx, model.OpenFlareAccessLogIPTrendQuery{ + NodeID: strings.TrimSpace(normalized.NodeID), + RemoteAddr: strings.TrimSpace(normalized.RemoteAddr), + Host: strings.TrimSpace(normalized.Host), + Since: time.Now().UTC().Add(-time.Duration(normalized.Hours) * time.Hour), + BucketMinutes: normalized.BucketMinutes, + }) + if err != nil { + return nil, err + } + pointMap := make(map[int64]int64, len(points)) + for _, item := range points { + if item == nil { + continue + } + pointMap[item.BucketEpoch] = item.RequestCount + } + bucketDuration := time.Duration(normalized.BucketMinutes) * time.Minute + start := time.Now().UTC().Add(-time.Duration(normalized.Hours) * time.Hour).Truncate(bucketDuration) + end := time.Now().UTC().Truncate(bucketDuration) + views := make([]AccessLogIPTrendPoint, 0, int(end.Sub(start)/bucketDuration)+1) + for cursor := start; !cursor.After(end); cursor = cursor.Add(bucketDuration) { + views = append(views, AccessLogIPTrendPoint{ + BucketStartedAt: cursor, + RequestCount: pointMap[cursor.Unix()], + }) + } + return &AccessLogIPTrendView{ + RemoteAddr: normalized.RemoteAddr, + Hours: normalized.Hours, + BucketMinutes: normalized.BucketMinutes, + Points: views, + }, nil +} + +// CleanupAccessLogs removes access logs older than retention days. +func CleanupAccessLogs(ctx context.Context, input AccessLogCleanupInput) (*AccessLogCleanupResult, error) { + if input.RetentionDays <= 0 || input.RetentionDays > nodeAccessLogRetentionDays { + return nil, errors.New("retention_days 必须在 1 到 90 之间") + } + cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour) + deleted, err := model.DeleteOpenFlareAccessLogsBefore(ctx, cutoff) + if err != nil { + return nil, err + } + return &AccessLogCleanupResult{ + RetentionDays: input.RetentionDays, + DeletedCount: deleted, + Cutoff: cutoff, + }, nil +} + +func buildModelAccessLogQuery(input AccessLogQuery) model.OpenFlareAccessLogQuery { + return model.OpenFlareAccessLogQuery{ + NodeID: strings.TrimSpace(input.NodeID), + RemoteAddr: strings.TrimSpace(input.RemoteAddr), + Host: strings.TrimSpace(input.Host), + Path: strings.TrimSpace(input.Path), + Since: time.Now().UTC().Add(-nodeAccessLogRetentionWindow), + Page: input.Page, + PageSize: input.PageSize, + SortBy: input.SortBy, + SortOrder: input.SortOrder, + } +} + +func listNodeNameMap(ctx context.Context, logs []*model.OpenFlareAccessLog) (map[string]string, error) { + nodeIDs := make([]string, 0, len(logs)) + seen := make(map[string]struct{}, len(logs)) + for _, item := range logs { + if item == nil || item.NodeID == "" { + continue + } + if _, exists := seen[item.NodeID]; exists { + continue + } + seen[item.NodeID] = struct{}{} + nodeIDs = append(nodeIDs, item.NodeID) + } + if len(nodeIDs) == 0 { + return map[string]string{}, nil + } + nodes, err := model.ListOpenFlareNodesByNodeIDs(ctx, nodeIDs) + if err != nil { + return nil, err + } + result := make(map[string]string, len(nodes)) + for _, node := range nodes { + result[node.NodeID] = node.Name + } + return result, nil +} + +func normalizeAccessLogQuery(input AccessLogQuery) AccessLogQuery { + return AccessLogQuery{ + NodeID: strings.TrimSpace(input.NodeID), + RemoteAddr: strings.TrimSpace(input.RemoteAddr), + Host: strings.TrimSpace(input.Host), + Path: strings.TrimSpace(input.Path), + Page: normalizeAccessLogPage(input.Page), + PageSize: normalizeAccessLogPageSize(input.PageSize), + SortBy: normalizeAccessLogSortBy(input.SortBy), + SortOrder: normalizeAccessLogSortOrder(input.SortOrder), + FoldMinutes: input.FoldMinutes, + } +} + +func normalizeAccessLogIPSummaryQuery(input AccessLogIPSummaryQuery) AccessLogIPSummaryQuery { + return AccessLogIPSummaryQuery{ + NodeID: strings.TrimSpace(input.NodeID), + RemoteAddr: strings.TrimSpace(input.RemoteAddr), + Host: strings.TrimSpace(input.Host), + Page: normalizeAccessLogPage(input.Page), + PageSize: normalizeAccessLogPageSize(input.PageSize), + SortBy: normalizeIPSummarySortBy(input.SortBy), + SortOrder: normalizeAccessLogSortOrder(input.SortOrder), + } +} + +func normalizeFoldedAccessLogIPQuery(input FoldedAccessLogIPQuery) (FoldedAccessLogIPQuery, time.Time, error) { + foldMinutes, err := normalizeFoldMinutes(input.FoldMinutes) + if err != nil { + return FoldedAccessLogIPQuery{}, time.Time{}, err + } + bucketStartedAt, err := time.Parse(time.RFC3339, strings.TrimSpace(input.BucketStartedAt)) + if err != nil { + return FoldedAccessLogIPQuery{}, time.Time{}, errors.New("bucket_started_at 必须为 RFC3339 时间") + } + normalizedSortBy := strings.TrimSpace(input.SortBy) + switch normalizedSortBy { + case "last_seen_at", "remote_addr": + default: + normalizedSortBy = "request_count" + } + return FoldedAccessLogIPQuery{ + NodeID: strings.TrimSpace(input.NodeID), + RemoteAddr: strings.TrimSpace(input.RemoteAddr), + Host: strings.TrimSpace(input.Host), + Path: strings.TrimSpace(input.Path), + BucketStartedAt: strings.TrimSpace(input.BucketStartedAt), + FoldMinutes: foldMinutes, + Page: normalizeAccessLogPage(input.Page), + PageSize: normalizeAccessLogPageSize(input.PageSize), + SortBy: normalizedSortBy, + SortOrder: normalizeAccessLogSortOrder(input.SortOrder), + }, bucketStartedAt.UTC(), nil +} + +func normalizeAccessLogIPTrendQuery(input AccessLogIPTrendQuery) (AccessLogIPTrendQuery, error) { + remoteAddr := strings.TrimSpace(input.RemoteAddr) + if remoteAddr == "" { + return AccessLogIPTrendQuery{}, errors.New("remote_addr 不能为空") + } + hours := input.Hours + if hours <= 0 { + hours = defaultIPTrendHours + } + if hours > maxIPTrendHours { + hours = maxIPTrendHours + } + bucketMinutes := input.BucketMinutes + if bucketMinutes <= 0 { + bucketMinutes = defaultIPTrendBucketMinute + } + switch bucketMinutes { + case 5, 10, 15, 30, 60: + default: + return AccessLogIPTrendQuery{}, errors.New("bucket_minutes 仅支持 5、10、15、30、60") + } + return AccessLogIPTrendQuery{ + NodeID: strings.TrimSpace(input.NodeID), + RemoteAddr: remoteAddr, + Host: strings.TrimSpace(input.Host), + Hours: hours, + BucketMinutes: bucketMinutes, + }, nil +} + +func normalizeAccessLogPage(page int) int { + if page < 0 { + return 0 + } + return page +} + +func normalizeAccessLogPageSize(pageSize int) int { + if pageSize <= 0 { + return defaultAccessLogPageSize + } + if pageSize > maxAccessLogPageSize { + return maxAccessLogPageSize + } + return pageSize +} + +func normalizeAccessLogSortBy(sortBy string) string { + switch strings.TrimSpace(sortBy) { + case "status_code", "remote_addr", "host", "path": + return strings.TrimSpace(sortBy) + default: + return defaultAccessLogSortBy + } +} + +func normalizeAccessLogSortOrder(sortOrder string) string { + if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") { + return "asc" + } + return defaultAccessLogSortOrder +} + +func normalizeFoldSortBy(sortBy string) string { + switch strings.TrimSpace(sortBy) { + case "request_count": + return "request_count" + default: + return "bucket_started_at" + } +} + +func normalizeIPSummarySortBy(sortBy string) string { + switch strings.TrimSpace(sortBy) { + case "recent_requests", "last_seen_at", "remote_addr": + return strings.TrimSpace(sortBy) + default: + return "total_requests" + } +} + +func normalizeFoldMinutes(value int) (int, error) { + if value <= 0 { + return defaultAccessLogFoldMinute, nil + } + switch value { + case 3, 5: + return value, nil + default: + return 0, errors.New("fold_minutes 仅支持 3 或 5") + } +} diff --git a/Wavelet/internal/apps/openflare/observability/analytics.go b/Wavelet/internal/apps/openflare/observability/analytics.go new file mode 100644 index 00000000..d4e1bcbb --- /dev/null +++ b/Wavelet/internal/apps/openflare/observability/analytics.go @@ -0,0 +1,482 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package observability + +import ( + "encoding/json" + "sort" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const observabilityTrendBuckets = 24 + +const ( + healthEventStatusActive = "active" + healthEventStatusResolved = "resolved" + healthSeverityCritical = "critical" + healthSeverityWarning = "warning" +) + +// DistributionItem is a key/value distribution entry. +type DistributionItem struct { + Key string `json:"key"` + Value int64 `json:"value"` +} + +// TrafficDistributions groups traffic distribution charts. +type TrafficDistributions struct { + StatusCodes []DistributionItem `json:"status_codes"` + TopDomains []DistributionItem `json:"top_domains"` + SourceCountries []DistributionItem `json:"source_countries"` +} + +// TrafficWindowSummary summarizes a traffic reporting window. +type TrafficWindowSummary struct { + WindowStartedAt time.Time `json:"window_started_at"` + WindowEndedAt time.Time `json:"window_ended_at"` + RequestCount int64 `json:"request_count"` + UniqueVisitorCount int64 `json:"unique_visitor_count"` + ErrorCount int64 `json:"error_count"` + EstimatedQPS float64 `json:"estimated_qps"` + ErrorRatePercent float64 `json:"error_rate_percent"` +} + +// HealthSummary summarizes node health alerts and risks. +type HealthSummary struct { + ActiveAlerts int `json:"active_alerts"` + CriticalAlerts int `json:"critical_alerts"` + WarningAlerts int `json:"warning_alerts"` + InfoAlerts int `json:"info_alerts"` + ResolvedAlerts int `json:"resolved_alerts"` + HasCapacityRisk bool `json:"has_capacity_risk"` + HasTrafficRisk bool `json:"has_traffic_risk"` + HasRuntimeRisk bool `json:"has_runtime_risk"` +} + +// TrafficTrendPoint is a traffic trend bucket. +type TrafficTrendPoint struct { + BucketStartedAt time.Time `json:"bucket_started_at"` + RequestCount int64 `json:"request_count"` + ErrorCount int64 `json:"error_count"` + UniqueVisitorCount int64 `json:"unique_visitor_count"` +} + +// CapacityTrendPoint is a capacity trend bucket. +type CapacityTrendPoint struct { + BucketStartedAt time.Time `json:"bucket_started_at"` + AverageCPUUsagePercent float64 `json:"average_cpu_usage_percent"` + AverageMemoryUsagePercent float64 `json:"average_memory_usage_percent"` + ReportedNodes int `json:"reported_nodes"` +} + +// NetworkTrendPoint is a network trend bucket. +type NetworkTrendPoint struct { + BucketStartedAt time.Time `json:"bucket_started_at"` + NetworkRxBytes int64 `json:"network_rx_bytes"` + NetworkTxBytes int64 `json:"network_tx_bytes"` + OpenrestyRxBytes int64 `json:"openresty_rx_bytes"` + OpenrestyTxBytes int64 `json:"openresty_tx_bytes"` + ReportedNodes int `json:"reported_nodes"` +} + +// DiskIOTrendPoint is a disk IO trend bucket. +type DiskIOTrendPoint struct { + BucketStartedAt time.Time `json:"bucket_started_at"` + DiskReadBytes int64 `json:"disk_read_bytes"` + DiskWriteBytes int64 `json:"disk_write_bytes"` + ReportedNodes int `json:"reported_nodes"` +} + +type distributionAccumulator map[string]int64 + +type capacityTrendAccumulator struct { + cpuSum float64 + cpuCount int + memSum float64 + memCount int + nodes map[string]struct{} +} + +type snapshotTrendAccumulator struct { + nodes map[string]struct{} +} + +type diskCounterState struct { + read int64 + write int64 + seen bool +} + +func buildTrafficWindowSummary(report *model.OpenFlareRequestReport) TrafficWindowSummary { + if report == nil { + return TrafficWindowSummary{} + } + summary := TrafficWindowSummary{ + WindowStartedAt: report.WindowStartedAt, + WindowEndedAt: report.WindowEndedAt, + RequestCount: report.RequestCount, + UniqueVisitorCount: report.UniqueVisitorCount, + ErrorCount: report.ErrorCount, + } + if duration := report.WindowEndedAt.Sub(report.WindowStartedAt).Seconds(); duration > 0 { + summary.EstimatedQPS = float64(report.RequestCount) / duration + } + if report.RequestCount > 0 { + summary.ErrorRatePercent = (float64(report.ErrorCount) / float64(report.RequestCount)) * 100 + } + return summary +} + +// BuildTrafficDistributions aggregates traffic distribution charts. +func BuildTrafficDistributions( + reports []*model.OpenFlareRequestReport, + accessLogRegions []*model.OpenFlareAccessLogRegionCount, + limit int, +) TrafficDistributions { + statusCodes := make(distributionAccumulator) + topDomains := make(distributionAccumulator) + reportSourceCountries := make(distributionAccumulator) + for _, report := range reports { + mergeJSONCounts(statusCodes, report.StatusCodesJSON) + mergeJSONCounts(topDomains, report.TopDomainsJSON) + mergeJSONCounts(reportSourceCountries, report.SourceCountriesJSON) + } + sourceCountries := reportSourceCountries + if len(accessLogRegions) > 0 { + sourceCountries = make(distributionAccumulator, len(accessLogRegions)) + for _, item := range accessLogRegions { + if item == nil || strings.TrimSpace(item.Region) == "" || item.Count <= 0 { + continue + } + sourceCountries[item.Region] = item.Count + } + } + return TrafficDistributions{ + StatusCodes: toDistributionItems(statusCodes, limit), + TopDomains: toDistributionItems(topDomains, limit), + SourceCountries: toDistributionItems(sourceCountries, limit), + } +} + +func buildHealthSummary( + snapshot *model.OpenFlareMetricSnapshot, + report *model.OpenFlareRequestReport, + events []*model.OpenFlareHealthEvent, +) HealthSummary { + summary := HealthSummary{} + for _, event := range events { + if event == nil { + continue + } + if event.Status == healthEventStatusResolved { + summary.ResolvedAlerts++ + continue + } + summary.ActiveAlerts++ + switch event.Severity { + case healthSeverityCritical: + summary.CriticalAlerts++ + case healthSeverityWarning: + summary.WarningAlerts++ + default: + summary.InfoAlerts++ + } + } + if snapshot != nil { + memoryUsage := Percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes) + storageUsage := Percentage(snapshot.StorageUsedBytes, snapshot.StorageTotalBytes) + summary.HasCapacityRisk = snapshot.CPUUsagePercent >= 80 || memoryUsage >= 85 || storageUsage >= 85 + } + if report != nil && report.RequestCount >= 100 { + summary.HasTrafficRisk = (float64(report.ErrorCount) / float64(report.RequestCount)) >= 0.05 + } + summary.HasRuntimeRisk = summary.ActiveAlerts > 0 || summary.HasCapacityRisk || summary.HasTrafficRisk + return summary +} + +// BuildTrafficTrendPoints builds 24h traffic trend buckets. +func BuildTrafficTrendPoints(now time.Time, reports []*model.OpenFlareRequestReport) []TrafficTrendPoint { + start := trendWindowStart(now) + points := make([]TrafficTrendPoint, observabilityTrendBuckets) + for index := range points { + points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour) + } + for _, report := range reports { + index, ok := trendBucketIndex(report.WindowEndedAt, start) + if !ok { + continue + } + points[index].RequestCount += report.RequestCount + points[index].ErrorCount += report.ErrorCount + points[index].UniqueVisitorCount += report.UniqueVisitorCount + } + return points +} + +// BuildCapacityTrendPoints builds 24h capacity trend buckets. +func BuildCapacityTrendPoints(now time.Time, snapshots []*model.OpenFlareMetricSnapshot) []CapacityTrendPoint { + start := trendWindowStart(now) + points := make([]CapacityTrendPoint, observabilityTrendBuckets) + accumulators := make([]capacityTrendAccumulator, observabilityTrendBuckets) + for index := range points { + points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour) + accumulators[index].nodes = make(map[string]struct{}) + } + for _, snapshot := range snapshots { + index, ok := trendBucketIndex(snapshot.CapturedAt, start) + if !ok { + continue + } + if snapshot.CPUUsagePercent > 0 { + accumulators[index].cpuSum += snapshot.CPUUsagePercent + accumulators[index].cpuCount++ + } + if memoryUsage := Percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes); memoryUsage > 0 { + accumulators[index].memSum += memoryUsage + accumulators[index].memCount++ + } + if snapshot.NodeID != "" { + accumulators[index].nodes[snapshot.NodeID] = struct{}{} + } + } + for index := range points { + if accumulators[index].cpuCount > 0 { + points[index].AverageCPUUsagePercent = accumulators[index].cpuSum / float64(accumulators[index].cpuCount) + } + if accumulators[index].memCount > 0 { + points[index].AverageMemoryUsagePercent = accumulators[index].memSum / float64(accumulators[index].memCount) + } + points[index].ReportedNodes = len(accumulators[index].nodes) + } + return points +} + +// BuildNetworkTrendPoints builds 24h network trend buckets. +func BuildNetworkTrendPoints( + now time.Time, + snapshots []*model.OpenFlareMetricSnapshot, + openrestyObs []*model.OpenFlareNodeObservationOpenresty, +) []NetworkTrendPoint { + start := trendWindowStart(now) + points := make([]NetworkTrendPoint, observabilityTrendBuckets) + accumulators := make([]snapshotTrendAccumulator, observabilityTrendBuckets) + for index := range points { + points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour) + accumulators[index].nodes = make(map[string]struct{}) + } + for _, snapshot := range snapshots { + index, ok := trendBucketIndex(snapshot.CapturedAt, start) + if !ok { + continue + } + points[index].NetworkRxBytes += snapshot.NetworkRxBytes + points[index].NetworkTxBytes += snapshot.NetworkTxBytes + if snapshot.NodeID != "" { + accumulators[index].nodes[snapshot.NodeID] = struct{}{} + } + } + for _, obs := range openrestyObs { + index, ok := trendBucketIndex(obs.CapturedAt, start) + if !ok { + continue + } + points[index].OpenrestyRxBytes += obs.OpenrestyRxBytes + points[index].OpenrestyTxBytes += obs.OpenrestyTxBytes + if obs.NodeID != "" { + accumulators[index].nodes[obs.NodeID] = struct{}{} + } + } + for index := range points { + points[index].ReportedNodes = len(accumulators[index].nodes) + } + return points +} + +// BuildDiskIOTrendPoints builds 24h disk IO trend buckets. +func BuildDiskIOTrendPoints(now time.Time, snapshots []*model.OpenFlareMetricSnapshot) []DiskIOTrendPoint { + start := trendWindowStart(now) + points := make([]DiskIOTrendPoint, observabilityTrendBuckets) + accumulators := make([]snapshotTrendAccumulator, observabilityTrendBuckets) + for index := range points { + points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour) + accumulators[index].nodes = make(map[string]struct{}) + } + sort.Slice(snapshots, func(i int, j int) bool { + if snapshots[i].CapturedAt.Equal(snapshots[j].CapturedAt) { + return snapshots[i].NodeID < snapshots[j].NodeID + } + return snapshots[i].CapturedAt.Before(snapshots[j].CapturedAt) + }) + previousByNode := make(map[string]diskCounterState, len(snapshots)) + for _, snapshot := range snapshots { + nodeKey := snapshot.NodeID + if nodeKey == "" { + nodeKey = "__unknown__" + } + previous := previousByNode[nodeKey] + previousByNode[nodeKey] = diskCounterState{ + read: snapshot.DiskReadBytes, + write: snapshot.DiskWriteBytes, + seen: true, + } + if !previous.seen { + continue + } + index, ok := trendBucketIndex(snapshot.CapturedAt, start) + if !ok { + continue + } + readDelta := snapshot.DiskReadBytes - previous.read + writeDelta := snapshot.DiskWriteBytes - previous.write + if readDelta < 0 { + readDelta = 0 + } + if writeDelta < 0 { + writeDelta = 0 + } + points[index].DiskReadBytes += readDelta + points[index].DiskWriteBytes += writeDelta + if snapshot.NodeID != "" { + accumulators[index].nodes[snapshot.NodeID] = struct{}{} + } + } + for index := range points { + points[index].ReportedNodes = len(accumulators[index].nodes) + } + return points +} + +func latestMetricSnapshot(snapshots []*model.OpenFlareMetricSnapshot) *model.OpenFlareMetricSnapshot { + for _, snapshot := range snapshots { + if snapshot != nil { + return snapshot + } + } + return nil +} + +func latestTrafficReport(reports []*model.OpenFlareRequestReport) *model.OpenFlareRequestReport { + for _, report := range reports { + if report != nil { + return report + } + } + return nil +} + +// LatestMetricSnapshotsByNode returns the latest snapshot per node. +func LatestMetricSnapshotsByNode(snapshots []*model.OpenFlareMetricSnapshot) map[string]*model.OpenFlareMetricSnapshot { + result := make(map[string]*model.OpenFlareMetricSnapshot, len(snapshots)) + for _, snapshot := range snapshots { + if snapshot == nil || snapshot.NodeID == "" { + continue + } + if existing, ok := result[snapshot.NodeID]; ok && !snapshot.CapturedAt.After(existing.CapturedAt) { + continue + } + result[snapshot.NodeID] = snapshot + } + return result +} + +// LatestTrafficReportsByNode returns the latest traffic report per node. +func LatestTrafficReportsByNode(reports []*model.OpenFlareRequestReport) map[string]*model.OpenFlareRequestReport { + result := make(map[string]*model.OpenFlareRequestReport, len(reports)) + for _, report := range reports { + if report == nil || report.NodeID == "" { + continue + } + if existing, ok := result[report.NodeID]; ok && !report.WindowEndedAt.After(existing.WindowEndedAt) { + continue + } + result[report.NodeID] = report + } + return result +} + +// ActiveHealthEventsByNode groups active health events by node id. +func ActiveHealthEventsByNode(events []*model.OpenFlareHealthEvent) map[string][]*model.OpenFlareHealthEvent { + result := make(map[string][]*model.OpenFlareHealthEvent) + for _, event := range events { + if event == nil || event.NodeID == "" { + continue + } + result[event.NodeID] = append(result[event.NodeID], event) + } + return result +} + +// Percentage returns used/total as a percentage. +func Percentage(used int64, total int64) float64 { + if used <= 0 || total <= 0 { + return 0 + } + return (float64(used) / float64(total)) * 100 +} + +func mergeJSONCounts(target distributionAccumulator, raw string) { + if len(target) == 0 && strings.TrimSpace(raw) == "" { + return + } + values := parseJSONCounts(raw) + for key, value := range values { + if strings.TrimSpace(key) == "" || value <= 0 { + continue + } + target[key] += value + } +} + +func parseJSONCounts(raw string) map[string]int64 { + if strings.TrimSpace(raw) == "" { + return nil + } + values := make(map[string]int64) + if err := json.Unmarshal([]byte(raw), &values); err != nil { + return nil + } + return values +} + +func toDistributionItems(values distributionAccumulator, limit int) []DistributionItem { + if len(values) == 0 { + return []DistributionItem{} + } + items := make([]DistributionItem, 0, len(values)) + for key, value := range values { + if strings.TrimSpace(key) == "" || value <= 0 { + continue + } + items = append(items, DistributionItem{Key: key, Value: value}) + } + sort.Slice(items, func(i int, j int) bool { + if items[i].Value == items[j].Value { + return items[i].Key < items[j].Key + } + return items[i].Value > items[j].Value + }) + if limit > 0 && len(items) > limit { + items = items[:limit] + } + return items +} + +func trendWindowStart(now time.Time) time.Time { + return now.Truncate(time.Hour).Add(-(observabilityTrendBuckets - 1) * time.Hour) +} + +func trendBucketIndex(timestamp time.Time, start time.Time) (int, bool) { + if timestamp.Before(start) { + return 0, false + } + delta := timestamp.Sub(start) + index := int(delta / time.Hour) + if index < 0 || index >= observabilityTrendBuckets { + return 0, false + } + return index, true +} diff --git a/Wavelet/internal/apps/openflare/observability/node_logics.go b/Wavelet/internal/apps/openflare/observability/node_logics.go new file mode 100644 index 00000000..a7275ad4 --- /dev/null +++ b/Wavelet/internal/apps/openflare/observability/node_logics.go @@ -0,0 +1,246 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package observability + +import ( + "context" + "encoding/json" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +const ( + defaultObservabilityWindow = 24 * time.Hour + defaultObservabilityLimit = 120 + maxObservabilityLimit = 500 +) + +// NodeQuery filters node observability data. +type NodeQuery struct { + Hours int `json:"hours"` + Limit int `json:"limit"` +} + +// NodeAnalytics groups node observability analytics. +type NodeAnalytics struct { + Traffic TrafficWindowSummary `json:"traffic"` + Distributions TrafficDistributions `json:"distributions"` + Health HealthSummary `json:"health"` +} + +// NodeTrends groups node observability trend series. +type NodeTrends struct { + Traffic24h []TrafficTrendPoint `json:"traffic_24h"` + Capacity24h []CapacityTrendPoint `json:"capacity_24h"` + Network24h []NetworkTrendPoint `json:"network_24h"` + DiskIO24h []DiskIOTrendPoint `json:"disk_io_24h"` +} + +// RelayDashboardSnapshot summarizes tunnel relay status. +type RelayDashboardSnapshot struct { + TotalProxies int `json:"total_proxies"` + OnlineProxies int `json:"online_proxies"` + OfflineProxies int `json:"offline_proxies"` + Proxies []RelayProxyStat `json:"proxies"` + TotalConnections int `json:"total_connections"` + ClientCounts int `json:"client_counts"` +} + +// RelayProxyStat is a single relay proxy entry. +type RelayProxyStat struct { + Name string `json:"name"` + Type string `json:"type"` + Status string `json:"status"` + ClientVersion string `json:"client_version"` + LastStartTime string `json:"last_start_time"` + LastCloseTime string `json:"last_close_time"` + ClientAddr string `json:"client_addr"` +} + +// NodeView is the node observability API response. +type NodeView struct { + NodeID string `json:"node_id"` + Profile *model.OpenFlareNodeSystemProfile `json:"profile"` + MetricSnapshots []*model.OpenFlareMetricSnapshot `json:"metric_snapshots"` + TrafficReports []*model.OpenFlareRequestReport `json:"traffic_reports"` + HealthEvents []*model.OpenFlareHealthEvent `json:"health_events"` + Analytics NodeAnalytics `json:"analytics"` + Trends NodeTrends `json:"trends"` + RelayDashboard *RelayDashboardSnapshot `json:"relay_dashboard,omitempty"` +} + +// HealthEventCleanupResult reports health event cleanup outcome. +type HealthEventCleanupResult struct { + NodeID string `json:"node_id"` + DeletedCount int64 `json:"deleted_count"` +} + +// GetNodeObservability returns observability details for a node. +func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeView, error) { + now := time.Now() + node, err := model.GetOpenFlareNodeByID(ctx, id) + if err != nil { + return nil, err + } + + limit := normalizeObservabilityLimit(query.Limit) + since := now.Add(-normalizeObservabilityWindow(query.Hours)) + + profile, err := model.GetOpenFlareNodeSystemProfile(ctx, node.NodeID) + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + if errors.Is(err, gorm.ErrRecordNotFound) { + profile = nil + } + + snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, since, limit) + if err != nil { + return nil, err + } + reports, err := model.ListOpenFlareRequestReportsSince(ctx, node.NodeID, since, limit) + if err != nil { + return nil, err + } + accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, 8) + if err != nil { + return nil, err + } + trendSnapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-24*time.Hour), 0) + if err != nil { + return nil, err + } + trendOpenresty, err := model.ListOpenFlareNodeObservationOpenresty(ctx, node.NodeID, now.Add(-24*time.Hour), 0) + if err != nil { + return nil, err + } + trendReports, err := model.ListOpenFlareRequestReportsSince(ctx, node.NodeID, now.Add(-24*time.Hour), 0) + if err != nil { + return nil, err + } + events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, limit) + if err != nil { + return nil, err + } + + view := &NodeView{ + NodeID: node.NodeID, + Profile: profile, + MetricSnapshots: snapshots, + TrafficReports: reports, + HealthEvents: events, + Analytics: NodeAnalytics{ + Traffic: buildTrafficWindowSummary(latestTrafficReport(reports)), + Distributions: BuildTrafficDistributions(reports, accessLogRegions, 8), + Health: buildHealthSummary(latestMetricSnapshot(snapshots), latestTrafficReport(reports), events), + }, + Trends: NodeTrends{ + Traffic24h: BuildTrafficTrendPoints(now, trendReports), + Capacity24h: BuildCapacityTrendPoints(now, trendSnapshots), + Network24h: BuildNetworkTrendPoints(now, trendSnapshots, trendOpenresty), + DiskIO24h: BuildDiskIOTrendPoints(now, trendSnapshots), + }, + } + if node.NodeType == "tunnel_relay" { + frpsObs, frpsErr := model.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1) + if frpsErr != nil { + return nil, frpsErr + } + var latestFrps *model.OpenFlareNodeObservationFrps + if len(frpsObs) > 0 { + latestFrps = frpsObs[0] + } + view.RelayDashboard = buildRelayDashboardSnapshot(node, latestFrps) + } + return view, nil +} + +// CleanupHealthEvents removes all health events for a node. +func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) { + node, err := model.GetOpenFlareNodeByID(ctx, id) + if err != nil { + return nil, err + } + deletedCount, err := model.DeleteOpenFlareHealthEventsByNodeID(ctx, node.NodeID) + if err != nil { + return nil, err + } + return &HealthEventCleanupResult{ + NodeID: node.NodeID, + DeletedCount: deletedCount, + }, nil +} + +func buildRelayDashboardSnapshot(node *model.OpenFlareNode, obs *model.OpenFlareNodeObservationFrps) *RelayDashboardSnapshot { + if node == nil { + return nil + } + totalProxies := 0 + totalConnections := 0 + clientCounts := 0 + proxies := []RelayProxyStat{} + + if obs != nil { + totalProxies = obs.FrpsProxyCount + totalConnections = obs.FrpsConnections + clientCounts = obs.FrpsClientCount + if obs.FrpsProxies != "" { + var decoded []RelayProxyStat + if err := json.Unmarshal([]byte(obs.FrpsProxies), &decoded); err == nil { + proxies = decoded + } + } + } + if totalProxies < 0 { + totalProxies = 0 + } + onlineProxies := 0 + for _, proxy := range proxies { + if proxy.Status == "online" { + onlineProxies++ + } + } + if len(proxies) == 0 { + onlineProxies = totalProxies + if node.RelayStatus != "healthy" { + onlineProxies = 0 + } + } + + return &RelayDashboardSnapshot{ + TotalProxies: totalProxies, + OnlineProxies: onlineProxies, + OfflineProxies: totalProxies - onlineProxies, + Proxies: proxies, + TotalConnections: maxInt(totalConnections, 0), + ClientCounts: maxInt(clientCounts, 0), + } +} + +func maxInt(a int, b int) int { + if a > b { + return a + } + return b +} + +func normalizeObservabilityLimit(limit int) int { + if limit <= 0 { + return defaultObservabilityLimit + } + if limit > maxObservabilityLimit { + return maxObservabilityLimit + } + return limit +} + +func normalizeObservabilityWindow(hours int) time.Duration { + if hours <= 0 { + return defaultObservabilityWindow + } + return time.Duration(hours) * time.Hour +} diff --git a/Wavelet/internal/apps/openflare/observability/routers.go b/Wavelet/internal/apps/openflare/observability/routers.go new file mode 100644 index 00000000..b8a93ab9 --- /dev/null +++ b/Wavelet/internal/apps/openflare/observability/routers.go @@ -0,0 +1,128 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package observability + +import ( + "strconv" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) + +// RegisterRoutes mounts legacy OpenFlare access log routes. +func RegisterRoutes(apiGroup *gin.RouterGroup) { + accessLogRoute := apiGroup.Group("/access-logs") + accessLogRoute.Use(compat.AdminAuth()) + { + accessLogRoute.GET("/", getAccessLogsHandler) + accessLogRoute.GET("/folds", getFoldedAccessLogsHandler) + accessLogRoute.GET("/folds/ip-summary", getFoldedAccessLogIPsHandler) + accessLogRoute.GET("/ip-summary", getAccessLogIPSummariesHandler) + accessLogRoute.GET("/ip-summary/trend", getAccessLogIPTrendHandler) + accessLogRoute.POST("/cleanup", cleanupAccessLogsHandler) + } +} + +func getAccessLogsHandler(c *gin.Context) { + logs, err := ListAccessLogs(c.Request.Context(), readAccessLogQuery(c)) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, logs) +} + +func getFoldedAccessLogsHandler(c *gin.Context) { + query := readAccessLogQuery(c) + query.FoldMinutes = readQueryInt(c, "fold_minutes") + logs, err := ListFoldedAccessLogs(c.Request.Context(), query) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, logs) +} + +func getFoldedAccessLogIPsHandler(c *gin.Context) { + result, err := ListFoldedAccessLogIPs(c.Request.Context(), FoldedAccessLogIPQuery{ + NodeID: c.Query("node_id"), + RemoteAddr: c.Query("remote_addr"), + Host: c.Query("host"), + Path: c.Query("path"), + BucketStartedAt: c.Query("bucket_started_at"), + FoldMinutes: readQueryInt(c, "fold_minutes"), + Page: readQueryInt(c, "p"), + PageSize: readQueryInt(c, "page_size"), + SortBy: c.Query("sort_by"), + SortOrder: c.Query("sort_order"), + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +func getAccessLogIPSummariesHandler(c *gin.Context) { + result, err := ListAccessLogIPSummaries(c.Request.Context(), AccessLogIPSummaryQuery{ + NodeID: c.Query("node_id"), + RemoteAddr: c.Query("remote_addr"), + Host: c.Query("host"), + Page: readQueryInt(c, "p"), + PageSize: readQueryInt(c, "page_size"), + SortBy: c.Query("sort_by"), + SortOrder: c.Query("sort_order"), + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +func getAccessLogIPTrendHandler(c *gin.Context) { + result, err := GetAccessLogIPTrend(c.Request.Context(), AccessLogIPTrendQuery{ + NodeID: c.Query("node_id"), + RemoteAddr: c.Query("remote_addr"), + Host: c.Query("host"), + Hours: readQueryInt(c, "hours"), + BucketMinutes: readQueryInt(c, "bucket_minutes"), + }) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +func cleanupAccessLogsHandler(c *gin.Context) { + var input AccessLogCleanupInput + if !compat.BindJSON(c, &input) { + return + } + result, err := CleanupAccessLogs(c.Request.Context(), input) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +func readAccessLogQuery(c *gin.Context) AccessLogQuery { + return AccessLogQuery{ + NodeID: c.Query("node_id"), + RemoteAddr: c.Query("remote_addr"), + Host: c.Query("host"), + Path: c.Query("path"), + Page: readQueryInt(c, "p"), + PageSize: readQueryInt(c, "page_size"), + SortBy: c.Query("sort_by"), + SortOrder: c.Query("sort_order"), + } +} + +func readQueryInt(c *gin.Context, key string) int { + value, _ := strconv.Atoi(c.DefaultQuery(key, "0")) + return value +} diff --git a/Wavelet/internal/apps/openflare/option/errs.go b/Wavelet/internal/apps/openflare/option/errs.go new file mode 100644 index 00000000..4056a8ee --- /dev/null +++ b/Wavelet/internal/apps/openflare/option/errs.go @@ -0,0 +1,13 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package option + +const ( + errInvalidParams = "无效的参数" + errOptionInitFailed = "系统选项初始化失败" + errGeoIPProvider = "归属方式仅支持 disabled、mmdb、ip-api、geojs、ipinfo" + errGeoIPIPEmpty = "IP 不能为空" + errGeoIPIPInvalid = "IP 格式无效" + errGeoIPLookupDisabled = "GeoIP 查询已禁用" +) diff --git a/Wavelet/internal/apps/openflare/option/logics.go b/Wavelet/internal/apps/openflare/option/logics.go new file mode 100644 index 00000000..5cff9e82 --- /dev/null +++ b/Wavelet/internal/apps/openflare/option/logics.go @@ -0,0 +1,244 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package option + +import ( + "context" + "errors" + "fmt" + "strings" + "sync" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip" + "github.com/Rain-kl/Wavelet/internal/buildinfo" + "github.com/Rain-kl/Wavelet/internal/model" +) + +var ( + initOnce sync.Once + initErr error +) + +// EnsureInitialized loads OptionMap from defaults and database once per process. +func EnsureInitialized(ctx context.Context) error { + initOnce.Do(func() { + initErr = model.InitOptionMap(ctx) + }) + return initErr +} + +// ResetInitializationForTest clears lazy-init state for unit tests. +func ResetInitializationForTest() { + initOnce = sync.Once{} + initErr = nil + model.ResetOptionMapForTest() +} + +type publicAuthSourceView struct { + ID uint64 `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + DisplayName string `json:"display_name"` + AuthorizeURL string `json:"authorize_url"` + IconURL string `json:"icon_url"` +} + +type statusView struct { + Version string `json:"version"` + StartTime int64 `json:"start_time"` + EmailVerification bool `json:"email_verification"` + GitHubOAuth bool `json:"github_oauth"` + GitHubClientID string `json:"github_client_id"` + SystemName string `json:"system_name"` + HomePageLink string `json:"home_page_link"` + FooterHTML string `json:"footer_html"` + WeChatQRCode string `json:"wechat_qrcode"` + WeChatLogin bool `json:"wechat_login"` + ServerAddress string `json:"server_address"` + PasswordRegisterEnabled bool `json:"password_register_enabled"` + CapLoginEnabled bool `json:"cap_login_enabled"` + AuthSources []publicAuthSourceView `json:"auth_sources"` +} + +type geoIPLookupRequest struct { + Provider string `json:"provider"` + IP string `json:"ip"` +} + +type geoIPLookupView struct { + Provider string `json:"provider"` + IP string `json:"ip"` + ISOCode string `json:"iso_code"` + Name string `json:"name"` + Latitude *float64 `json:"latitude,omitempty"` + Longitude *float64 `json:"longitude,omitempty"` +} + +type databaseCleanupInput struct { + Target string `json:"target"` + RetentionDays *int `json:"retention_days"` +} + +type databaseCleanupResult struct { + Target string `json:"target"` + TargetLabel string `json:"target_label"` + DeletedCount int64 `json:"deleted_count"` + DeleteAll bool `json:"delete_all"` + RetentionDays *int `json:"retention_days,omitempty"` +} + +type optionBatchPayload struct { + Options []model.OpenFlareOption `json:"options"` +} + +func listOptions(ctx context.Context) ([]model.OpenFlareOption, error) { + if err := EnsureInitialized(ctx); err != nil { + return nil, err + } + + model.OptionMapRWMutex.RLock() + defer model.OptionMapRWMutex.RUnlock() + + options := make([]model.OpenFlareOption, 0, len(model.OptionMap)) + for key, value := range model.OptionMap { + if isSecretOptionKey(key) { + continue + } + options = append(options, model.OpenFlareOption{ + Key: key, + Value: value, + }) + } + return options, nil +} + +func updateOption(ctx context.Context, option model.OpenFlareOption) error { + if err := EnsureInitialized(ctx); err != nil { + return err + } + return updateOptions(ctx, []model.OpenFlareOption{option}) +} + +func updateOptionsBatch(ctx context.Context, payload optionBatchPayload) error { + if err := EnsureInitialized(ctx); err != nil { + return err + } + if len(payload.Options) == 0 { + return errors.New(errInvalidParams) + } + return updateOptions(ctx, payload.Options) +} + +func updateOptions(ctx context.Context, options []model.OpenFlareOption) error { + if err := validateOptions(options); err != nil { + return err + } + return model.UpdateOpenFlareOptions(ctx, options) +} + +func getNotice(ctx context.Context) (string, error) { + if err := EnsureInitialized(ctx); err != nil { + return "", err + } + return model.OptionValue("Notice"), nil +} + +func getAbout(ctx context.Context) (string, error) { + if err := EnsureInitialized(ctx); err != nil { + return "", err + } + return model.OptionValue("About"), nil +} + +func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) { + if err := EnsureInitialized(ctx); err != nil { + return nil, err + } + + authSources, err := publicAuthSources(ctx, baseAPIPath) + if err != nil { + authSources = []publicAuthSourceView{} + } + + return &statusView{ + Version: buildinfo.Version, + StartTime: model.StartTime, + EmailVerification: model.EmailVerificationEnabled, + GitHubOAuth: model.GitHubOAuthEnabled, + GitHubClientID: model.GitHubClientId, + SystemName: model.SystemName, + HomePageLink: model.HomePageLink, + FooterHTML: model.Footer, + WeChatQRCode: model.WeChatAccountQRCodeImageURL, + WeChatLogin: model.WeChatAuthEnabled, + ServerAddress: model.ServerAddress, + PasswordRegisterEnabled: model.PasswordRegisterEnabled, + CapLoginEnabled: model.CapLoginEnabled, + AuthSources: authSources, + }, nil +} + +func publicAuthSources(ctx context.Context, baseAPIPath string) ([]publicAuthSourceView, error) { + sources, err := model.GetActiveAuthSources(ctx) + if err != nil { + return nil, err + } + result := make([]publicAuthSourceView, 0, len(sources)) + base := strings.TrimRight(baseAPIPath, "/") + for _, source := range sources { + result = append(result, publicAuthSourceView{ + ID: source.ID, + Name: source.Name, + Type: source.Type, + DisplayName: source.DisplayName, + AuthorizeURL: fmt.Sprintf("%s/oauth/%s/authorize", base, source.Name), + IconURL: source.IconURL, + }) + } + return result, nil +} + +func lookupGeoIP(_ context.Context, provider, rawIP string) (*geoIPLookupView, error) { + view, err := geoip.Lookup(provider, rawIP) + if err != nil { + return nil, err + } + return &geoIPLookupView{ + Provider: view.Provider, + IP: view.IP, + ISOCode: view.ISOCode, + Name: view.Name, + Latitude: view.Latitude, + Longitude: view.Longitude, + }, nil +} + +func cleanupDatabaseObservability(_ context.Context, input databaseCleanupInput) (*databaseCleanupResult, error) { + target := strings.TrimSpace(input.Target) + if target == "" { + return nil, errors.New(errInvalidParams) + } + if input.RetentionDays != nil && *input.RetentionDays <= 0 { + return nil, errors.New("retention_days 必须为大于 0 的整数") + } + + return &databaseCleanupResult{ + Target: target, + TargetLabel: target, + DeletedCount: 0, + DeleteAll: input.RetentionDays == nil, + RetentionDays: input.RetentionDays, + }, nil +} + +func syncUptimeKuma(_ context.Context) error { + // Stub: full Uptime Kuma sync is implemented in T-MISC. + return nil +} + +func isSecretOptionKey(key string) bool { + return strings.Contains(key, "Token") || + strings.Contains(key, "Secret") || + strings.Contains(key, "Password") +} diff --git a/Wavelet/internal/apps/openflare/option/logics_test.go b/Wavelet/internal/apps/openflare/option/logics_test.go new file mode 100644 index 00000000..f0602ed7 --- /dev/null +++ b/Wavelet/internal/apps/openflare/option/logics_test.go @@ -0,0 +1,119 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package option + +import ( + "context" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupOptionTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareOption{})) + + db.SetDB(sqliteDB) + ResetInitializationForTest() + + return func() { + db.SetDB(nil) + ResetInitializationForTest() + } +} + +func TestListOptionsFiltersSecretKeys(t *testing.T) { + cleanup := setupOptionTestDB(t) + defer cleanup() + ctx := context.Background() + + require.NoError(t, model.UpdateOpenFlareOptions(ctx, []model.OpenFlareOption{ + {Key: "SystemName", Value: "TestFlare"}, + {Key: "SMTPToken", Value: "secret-token"}, + {Key: "GitHubClientSecret", Value: "secret-id"}, + })) + + options, err := listOptions(ctx) + require.NoError(t, err) + + keys := make(map[string]string, len(options)) + for _, option := range options { + keys[option.Key] = option.Value + } + + assert.Equal(t, "TestFlare", keys["SystemName"]) + assert.NotContains(t, keys, "SMTPToken") + assert.NotContains(t, keys, "GitHubClientSecret") +} + +func TestUpdateOptionHotReloadsOptionMap(t *testing.T) { + cleanup := setupOptionTestDB(t) + defer cleanup() + ctx := context.Background() + + err := updateOption(ctx, model.OpenFlareOption{ + Key: "SystemName", + Value: "HotReloaded", + }) + require.NoError(t, err) + + assert.Equal(t, "HotReloaded", model.OptionValue("SystemName")) + assert.Equal(t, "HotReloaded", model.SystemName) +} + +func TestGetNoticeAndAbout(t *testing.T) { + cleanup := setupOptionTestDB(t) + defer cleanup() + ctx := context.Background() + + require.NoError(t, updateOption(ctx, model.OpenFlareOption{Key: "Notice", Value: "hello"})) + require.NoError(t, updateOption(ctx, model.OpenFlareOption{Key: "About", Value: "about-us"})) + + notice, err := getNotice(ctx) + require.NoError(t, err) + assert.Equal(t, "hello", notice) + + about, err := getAbout(ctx) + require.NoError(t, err) + assert.Equal(t, "about-us", about) +} + +func TestLookupGeoIPDisabledProvider(t *testing.T) { + cleanup := setupOptionTestDB(t) + defer cleanup() + ctx := context.Background() + + view, err := lookupGeoIP(ctx, "disabled", "8.8.8.8") + require.NoError(t, err) + assert.Equal(t, "disabled", view.Provider) + assert.Equal(t, "8.8.8.8", view.IP) +} + +func TestCleanupDatabaseObservabilityStub(t *testing.T) { + cleanup := setupOptionTestDB(t) + defer cleanup() + ctx := context.Background() + + retention := 7 + result, err := cleanupDatabaseObservability(ctx, databaseCleanupInput{ + Target: "node_access_logs", + RetentionDays: &retention, + }) + require.NoError(t, err) + assert.Equal(t, "node_access_logs", result.Target) + assert.Equal(t, int64(0), result.DeletedCount) + assert.False(t, result.DeleteAll) + require.NotNil(t, result.RetentionDays) + assert.Equal(t, 7, *result.RetentionDays) +} diff --git a/Wavelet/internal/apps/openflare/option/routers.go b/Wavelet/internal/apps/openflare/option/routers.go new file mode 100644 index 00000000..89f00f11 --- /dev/null +++ b/Wavelet/internal/apps/openflare/option/routers.go @@ -0,0 +1,139 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package option + +import ( + "encoding/json" + "errors" + "io" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +// RegisterRoutes mounts legacy OpenFlare option and public status routes. +func RegisterRoutes(apiGroup *gin.RouterGroup) { + apiGroup.GET("/status", getStatusHandler) + apiGroup.GET("/notice", getNoticeHandler) + apiGroup.GET("/about", getAboutHandler) + + optionRoute := apiGroup.Group("/option") + optionRoute.Use(compat.BridgeOpenFlareToken(), compat.RootAuth()) + { + optionRoute.GET("/", listOptionsHandler) + optionRoute.POST("/update", updateOptionHandler) + optionRoute.POST("/update-batch", updateOptionsBatchHandler) + optionRoute.POST("/geoip/lookup", lookupGeoIPHandler) + optionRoute.POST("/database/cleanup", cleanupDatabaseHandler) + } + + uptimeKumaRoute := apiGroup.Group("/uptimekuma") + uptimeKumaRoute.Use(compat.BridgeOpenFlareToken(), compat.RootAuth()) + { + uptimeKumaRoute.POST("/sync", syncUptimeKumaHandler) + } +} + +func getStatusHandler(c *gin.Context) { + view, err := getStatus(c.Request.Context(), "/api") + if err != nil { + compat.Fail(c, errOptionInitFailed) + return + } + compat.OK(c, view) +} + +func getNoticeHandler(c *gin.Context) { + notice, err := getNotice(c.Request.Context()) + if err != nil { + compat.Fail(c, errOptionInitFailed) + return + } + compat.OK(c, notice) +} + +func getAboutHandler(c *gin.Context) { + about, err := getAbout(c.Request.Context()) + if err != nil { + compat.Fail(c, errOptionInitFailed) + return + } + compat.OK(c, about) +} + +func listOptionsHandler(c *gin.Context) { + options, err := listOptions(c.Request.Context()) + if err != nil { + compat.Fail(c, errOptionInitFailed) + return + } + compat.OK(c, options) +} + +func updateOptionHandler(c *gin.Context) { + var option model.OpenFlareOption + if !compat.BindJSON(c, &option) { + return + } + if err := updateOption(c.Request.Context(), option); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func updateOptionsBatchHandler(c *gin.Context) { + var payload optionBatchPayload + if !compat.BindJSON(c, &payload) { + return + } + if err := updateOptionsBatch(c.Request.Context(), payload); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +func lookupGeoIPHandler(c *gin.Context) { + var request geoIPLookupRequest + if !compat.BindJSON(c, &request) { + return + } + view, err := lookupGeoIP(c.Request.Context(), request.Provider, request.IP) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, view) +} + +func cleanupDatabaseHandler(c *gin.Context) { + var input databaseCleanupInput + if err := bindOptionalJSON(c.Request.Body, &input); err != nil { + compat.Fail(c, errInvalidParams) + return + } + result, err := cleanupDatabaseObservability(c.Request.Context(), input) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +func syncUptimeKumaHandler(c *gin.Context) { + if err := syncUptimeKuma(c.Request.Context()); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "同步成功") +} + +func bindOptionalJSON(body io.Reader, target any) error { + if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) { + return err + } + return nil +} diff --git a/Wavelet/internal/apps/openflare/option/validate.go b/Wavelet/internal/apps/openflare/option/validate.go new file mode 100644 index 00000000..42139fbb --- /dev/null +++ b/Wavelet/internal/apps/openflare/option/validate.go @@ -0,0 +1,310 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package option + +import ( + "fmt" + "regexp" + "strconv" + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const rateLimitKeyExpirationSeconds = 1200 // 20 minutes + +var ( + openRestySizePattern = regexp.MustCompile(`^\d+[kKmMgG]?$`) + openRestyProxyBuffersPattern = regexp.MustCompile(`^\d+\s+\d+[kKmMgG]?$`) + openRestyCacheLevelsPattern = regexp.MustCompile(`^\d{1,2}(?::\d{1,2}){0,2}$`) + openRestyDurationTokenPattern = regexp.MustCompile(`^\d+[smhdwSMHDW]$`) +) + +func buildOptionValidationState(options []model.OpenFlareOption) map[string]string { + model.OptionMapRWMutex.RLock() + state := make(map[string]string, len(model.OptionMap)+len(options)) + for key, value := range model.OptionMap { + state[key] = value + } + model.OptionMapRWMutex.RUnlock() + + for _, option := range options { + state[option.Key] = option.Value + } + return state +} + +func validateOptionWithState(option model.OpenFlareOption, state map[string]string) error { + switch option.Key { + case "GitHubOAuthEnabled": + if option.Value == "true" && strings.TrimSpace(state["GitHubClientId"]) == "" { + return fmt.Errorf("无法启用 GitHub OAuth,请先填入 GitHub Client ID 以及 GitHub Client Secret!") + } + case "WeChatAuthEnabled": + if option.Value == "true" && strings.TrimSpace(state["WeChatServerAddress"]) == "" { + return fmt.Errorf("无法启用微信登录,请先填入微信登录相关配置信息!") + } + } + + if err := validateRateLimitOption(option.Key, option.Value); err != nil { + return err + } + if err := validateOpenRestyOption(option.Key, option.Value); err != nil { + return err + } + if err := validateGeoIPOption(option.Key, option.Value); err != nil { + return err + } + if err := validateDatabaseCleanupOption(option.Key, option.Value); err != nil { + return err + } + if err := validateAgentOption(option.Key, option.Value); err != nil { + return err + } + return validateUptimeKumaOption(option.Key, option.Value, state) +} + +func validateRateLimitOption(key, value string) error { + switch key { + case "GlobalApiRateLimitNum", "GlobalWebRateLimitNum", "CriticalRateLimitNum": + intValue, err := strconv.Atoi(value) + if err != nil || intValue <= 0 { + return fmt.Errorf("%s 必须为大于 0 的整数", key) + } + case "GlobalApiRateLimitDuration", "GlobalWebRateLimitDuration", "CriticalRateLimitDuration": + intValue, err := strconv.Atoi(value) + if err != nil || intValue <= 0 { + return fmt.Errorf("%s 必须为大于 0 的整数秒", key) + } + if intValue > rateLimitKeyExpirationSeconds { + return fmt.Errorf("%s 不能大于 %d 秒", key, rateLimitKeyExpirationSeconds) + } + } + return nil +} + +func validatePositiveIntegerOption(key, value string) error { + intValue, err := strconv.Atoi(value) + if err != nil || intValue <= 0 { + return fmt.Errorf("%s 必须为大于 0 的整数", key) + } + return nil +} + +func validateBooleanOption(key, value string) error { + switch value { + case "true", "false": + return nil + default: + return fmt.Errorf("%s 必须为 true 或 false", key) + } +} + +func validateGeoIPOption(key, value string) error { + if key != "GeoIPProvider" { + return nil + } + if geoip.IsValidProvider(value) { + return nil + } + return fmt.Errorf("%s 仅支持 disabled、mmdb、ip-api、geojs、ipinfo", key) +} + +func validateDatabaseCleanupOption(key, value string) error { + switch key { + case "DatabaseAutoCleanupEnabled": + return validateBooleanOption(key, value) + case "DatabaseAutoCleanupRetentionDays": + intValue, err := strconv.Atoi(value) + if err != nil || intValue < 1 { + return fmt.Errorf("%s 必须为大于等于 1 的整数天", key) + } + } + return nil +} + +func validateAgentOption(key, value string) error { + if key == "AgentWebsocketUpgradeEnabled" { + return validateBooleanOption(key, strings.TrimSpace(value)) + } + return nil +} + +func validateUptimeKumaOption(key, value string, state map[string]string) error { + trimmed := strings.TrimSpace(value) + switch key { + case "UptimeKumaEnabled": + if err := validateBooleanOption(key, trimmed); err != nil { + return err + } + if trimmed == "true" { + url := strings.TrimSpace(state["UptimeKumaUrl"]) + username := strings.TrimSpace(state["UptimeKumaUsername"]) + password := strings.TrimSpace(state["UptimeKumaPassword"]) + if url == "" { + return fmt.Errorf("启用 Uptime Kuma 时地址不能为空") + } + if username == "" { + return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空") + } + if password == "" && model.UptimeKumaPassword == "" { + return fmt.Errorf("启用 Uptime Kuma 时密码不能为空") + } + } + case "UptimeKumaUsername": + if trimmed == "" && state["UptimeKumaEnabled"] == "true" { + return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空") + } + case "UptimeKumaUrl": + if trimmed != "" && !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") { + return fmt.Errorf("Uptime Kuma 地址必须以 http:// 或 https:// 开头") + } + case "UptimeKumaMonitorScope": + if trimmed != "all" && trimmed != "selected" { + return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)") + } + case "UptimeKumaSyncInterval", "UptimeKumaInterval", "UptimeKumaRetryInterval", "UptimeKumaTimeout": + return validatePositiveIntegerOption(key, trimmed) + case "UptimeKumaRetry": + intValue, err := strconv.Atoi(trimmed) + if err != nil || intValue < 0 { + return fmt.Errorf("%s 必须为大于等于 0 的整数", key) + } + } + return nil +} + +func validateOpenRestyOption(key, value string) error { + trimmed := strings.TrimSpace(value) + + switch key { + case "OpenRestyDefaultServerReturnStatus": + if err := validatePositiveIntegerOption(key, trimmed); err != nil { + return err + } + statusCode, _ := strconv.Atoi(trimmed) + if statusCode < 100 || statusCode > 999 { + return fmt.Errorf("%s 必须在 100 到 999 之间", key) + } + case "OpenRestyWorkerProcesses": + if trimmed == "auto" { + return nil + } + return validatePositiveIntegerOption(key, trimmed) + case "OpenRestyWorkerConnections", + "OpenRestyWorkerRlimitNofile", + "OpenRestyKeepaliveTimeout", + "OpenRestyKeepaliveRequests", + "OpenRestyClientHeaderTimeout", + "OpenRestyClientBodyTimeout", + "OpenRestySendTimeout", + "OpenRestyProxyConnectTimeout", + "OpenRestyProxySendTimeout", + "OpenRestyProxyReadTimeout", + "OpenRestyGzipMinLength": + return validatePositiveIntegerOption(key, trimmed) + case "OpenRestyGzipCompLevel": + if err := validatePositiveIntegerOption(key, trimmed); err != nil { + return err + } + level, _ := strconv.Atoi(trimmed) + if level > 9 { + return fmt.Errorf("%s 不能大于 9", key) + } + case "OpenRestyEventsUse": + if trimmed == "" { + return nil + } + switch trimmed { + case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport": + return nil + default: + return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key) + } + case "OpenRestyResolvers": + if trimmed == "" { + return nil + } + if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) { + return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key) + } + case "OpenRestyEventsMultiAcceptEnabled", + "OpenRestyWebsocketEnabled", + "OpenRestyHTTP3Enabled", + "OpenRestyProxyRequestBufferingEnabled", + "OpenRestyProxyBufferingEnabled", + "OpenRestyGzipEnabled", + "OpenRestyCacheEnabled", + "OpenRestyCacheLockEnabled": + return validateBooleanOption(key, trimmed) + case "OpenRestyProxyBuffers", "OpenRestyLargeClientHeaderBuffers": + if openRestyProxyBuffersPattern.MatchString(trimmed) { + return nil + } + return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key) + case "OpenRestyProxyBufferSize", "OpenRestyProxyBusyBuffersSize", "OpenRestyCacheMaxSize", "OpenRestyClientMaxBodySize": + if openRestySizePattern.MatchString(trimmed) { + return nil + } + return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key) + case "OpenRestyCachePath": + if strings.ContainsAny(trimmed, "\r\n\t") { + return fmt.Errorf("%s 不能包含换行或制表符", key) + } + case "OpenRestyCacheLevels": + if openRestyCacheLevelsPattern.MatchString(trimmed) { + return nil + } + return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key) + case "OpenRestyCacheInactive", "OpenRestyCacheLockTimeout": + if openRestyDurationTokenPattern.MatchString(trimmed) { + return nil + } + return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key) + case "OpenRestyCacheKeyTemplate": + if trimmed == "" { + return fmt.Errorf("%s 不能为空", key) + } + if strings.ContainsAny(trimmed, "\r\n") { + return fmt.Errorf("%s 不能包含换行", key) + } + case "OpenRestyCacheUseStale": + if trimmed == "" { + return fmt.Errorf("%s 不能为空", key) + } + allowedTokens := map[string]struct{}{ + "error": {}, "timeout": {}, "invalid_header": {}, "updating": {}, + "http_500": {}, "http_502": {}, "http_503": {}, "http_504": {}, + "http_403": {}, "http_404": {}, "http_429": {}, "off": {}, + } + for _, token := range strings.Fields(trimmed) { + if _, ok := allowedTokens[token]; !ok { + return fmt.Errorf("%s 包含不支持的值 %q", key, token) + } + } + case "OpenRestyMainConfigTemplate": + if strings.TrimSpace(value) == "" { + return fmt.Errorf("%s 不能为空", key) + } + } + return nil +} + +func validateOptions(options []model.OpenFlareOption) error { + if len(options) == 0 { + return fmt.Errorf(errInvalidParams) + } + + state := buildOptionValidationState(options) + for _, option := range options { + if strings.TrimSpace(option.Key) == "" { + return fmt.Errorf(errInvalidParams) + } + if err := validateOptionWithState(option, state); err != nil { + return err + } + } + return nil +} diff --git a/Wavelet/internal/apps/openflare/origin/errs.go b/Wavelet/internal/apps/openflare/origin/errs.go new file mode 100644 index 00000000..43188b49 --- /dev/null +++ b/Wavelet/internal/apps/openflare/origin/errs.go @@ -0,0 +1,13 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package origin + +const ( + errOriginAddressRequired = "源站地址不能为空" + errOriginAddressInvalid = "源站地址格式不合法" + errOriginAddressExists = "源站地址已存在" + errOriginDeleteReferenced = "该源站仍被规则引用,无法删除" + errOriginMissingPort = "源站地址缺少端口" + errOriginNotFound = "源站不存在" +) diff --git a/Wavelet/internal/apps/openflare/origin/helpers.go b/Wavelet/internal/apps/openflare/origin/helpers.go new file mode 100644 index 00000000..93fbc19d --- /dev/null +++ b/Wavelet/internal/apps/openflare/origin/helpers.go @@ -0,0 +1,87 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package origin + +import ( + "errors" + "fmt" + "net" + "net/url" + "strings" + "unicode" +) + +func normalizeOriginAddress(raw string) string { + return strings.ToLower(strings.TrimSpace(raw)) +} + +func validateOriginAddress(address string) error { + if address == "" { + return errors.New(errOriginAddressRequired) + } + if strings.Contains(address, "://") || strings.ContainsAny(address, "/?#") { + return errors.New(errOriginAddressInvalid) + } + if strings.HasPrefix(address, "[") || strings.HasSuffix(address, "]") { + return errors.New(errOriginAddressInvalid) + } + if ip := net.ParseIP(address); ip != nil { + return nil + } + if len(address) > 253 { + return errors.New(errOriginAddressInvalid) + } + labels := strings.Split(address, ".") + for _, label := range labels { + if len(label) == 0 || len(label) > 63 { + return errors.New(errOriginAddressInvalid) + } + if label[0] == '-' || label[len(label)-1] == '-' { + return errors.New(errOriginAddressInvalid) + } + for _, r := range label { + if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' { + continue + } + return errors.New(errOriginAddressInvalid) + } + } + return nil +} + +func normalizeOriginName(name string, address string) string { + normalized := strings.TrimSpace(name) + if normalized != "" { + return normalized + } + return address +} + +func formatOriginHost(address string, port string) string { + return net.JoinHostPort(address, port) +} + +func rewriteOriginURLAddress(rawURL string, newAddress string) (string, error) { + parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL)) + if err != nil { + return "", fmt.Errorf("%s: %w", errOriginAddressInvalid, err) + } + address := normalizeOriginAddress(newAddress) + if err := validateOriginAddress(address); err != nil { + return "", err + } + port := parsed.Port() + if port == "" { + return "", errors.New(errOriginMissingPort) + } + parsed.Host = formatOriginHost(address, port) + return parsed.String(), nil +} + +func isUniqueConstraintError(err error) bool { + if err == nil { + return false + } + return strings.Contains(strings.ToLower(err.Error()), "unique") +} diff --git a/Wavelet/internal/apps/openflare/origin/logics.go b/Wavelet/internal/apps/openflare/origin/logics.go new file mode 100644 index 00000000..f4ea24b6 --- /dev/null +++ b/Wavelet/internal/apps/openflare/origin/logics.go @@ -0,0 +1,229 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package origin + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sort" + "strings" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// Input 源站创建/更新请求。 +type Input struct { + Name string `json:"name"` + Address string `json:"address"` + Remark string `json:"remark"` +} + +// RouteSummary 源站详情中的代理规则摘要。 +type RouteSummary struct { + ID uint `json:"id"` + Domain string `json:"domain"` + OriginURL string `json:"origin_url"` + Enabled bool `json:"enabled"` + UpdatedAt string `json:"updated_at"` +} + +// View 源站列表项。 +type View struct { + ID uint `json:"id"` + Name string `json:"name"` + Address string `json:"address"` + Remark string `json:"remark"` + RouteCount int64 `json:"route_count"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` +} + +// DetailView 源站详情。 +type DetailView struct { + View + Routes []RouteSummary `json:"routes"` +} + +// ListOrigins 列出全部源站。 +func ListOrigins(ctx context.Context) ([]View, error) { + origins, err := model.ListOrigins(ctx) + if err != nil { + return nil, err + } + return buildOriginViews(ctx, origins) +} + +// GetOriginDetail 获取源站详情。 +func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) { + origin, err := model.GetOriginByID(ctx, id) + if err != nil { + return nil, err + } + views, err := buildOriginViews(ctx, []model.Origin{*origin}) + if err != nil { + return nil, err + } + routes, err := model.ListProxyRoutesByOriginID(ctx, id) + if err != nil { + return nil, err + } + items := make([]RouteSummary, 0, len(routes)) + for _, route := range routes { + items = append(items, RouteSummary{ + ID: route.ID, + Domain: route.Domain, + OriginURL: route.OriginURL, + Enabled: route.Enabled, + UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), + }) + } + sort.Slice(items, func(i, j int) bool { + return items[i].Domain < items[j].Domain + }) + return &DetailView{ + View: views[0], + Routes: items, + }, nil +} + +// CreateOrigin 创建源站。 +func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) { + origin, err := buildOrigin(nil, input) + if err != nil { + return nil, err + } + if err = model.CreateOriginRecord(ctx, origin); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errOriginAddressExists) + } + return nil, err + } + return origin, nil +} + +// UpdateOrigin 更新源站。 +func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, error) { + origin, err := model.GetOriginByID(ctx, id) + if err != nil { + return nil, err + } + previousAddress := origin.Address + nextOrigin, err := buildOrigin(origin, input) + if err != nil { + return nil, err + } + err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Save(nextOrigin).Error; err != nil { + if isUniqueConstraintError(err) { + return errors.New(errOriginAddressExists) + } + return err + } + if previousAddress == nextOrigin.Address { + return nil + } + return updateRoutesForOriginAddress(ctx, tx, nextOrigin.ID, nextOrigin.Address) + }) + if err != nil { + return nil, err + } + return nextOrigin, nil +} + +// DeleteOrigin 删除源站。 +func DeleteOrigin(ctx context.Context, id uint) error { + count, err := model.CountProxyRoutesByOriginID(ctx, id) + if err != nil { + return err + } + if count > 0 { + return errors.New(errOriginDeleteReferenced) + } + if _, err = model.GetOriginByID(ctx, id); err != nil { + return err + } + return model.DeleteOriginRecord(ctx, id) +} + +func buildOrigin(existing *model.Origin, input Input) (*model.Origin, error) { + address := normalizeOriginAddress(input.Address) + if err := validateOriginAddress(address); err != nil { + return nil, err + } + if existing == nil { + existing = &model.Origin{} + } + existing.Address = address + existing.Name = normalizeOriginName(input.Name, address) + existing.Remark = strings.TrimSpace(input.Remark) + return existing, nil +} + +func buildOriginViews(ctx context.Context, origins []model.Origin) ([]View, error) { + countRows, err := model.ListOriginRouteCounts(ctx) + if err != nil { + return nil, err + } + countMap := make(map[uint]int64, len(countRows)) + for _, row := range countRows { + countMap[row.OriginID] = row.RouteCount + } + views := make([]View, 0, len(origins)) + for _, origin := range origins { + views = append(views, View{ + ID: origin.ID, + Name: origin.Name, + Address: origin.Address, + Remark: origin.Remark, + RouteCount: countMap[origin.ID], + CreatedAt: origin.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), + UpdatedAt: origin.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"), + }) + } + return views, nil +} + +func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uint, address string) error { + if !model.HasProxyRoutesTable(ctx) { + return nil + } + var routes []model.OriginProxyRoute + if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; err != nil { + return fmt.Errorf("query routes for origin update failed: %w", err) + } + for _, route := range routes { + rewrittenOriginURL, err := rewriteOriginURLAddress(route.OriginURL, address) + if err != nil { + return fmt.Errorf("rewrite route %d origin failed: %w", route.ID, err) + } + upstreams := make([]string, 0) + if strings.TrimSpace(route.Upstreams) != "" { + if err := json.Unmarshal([]byte(route.Upstreams), &upstreams); err != nil { + return fmt.Errorf("decode route %d upstreams failed: %w", route.ID, err) + } + } + if len(upstreams) == 0 { + upstreams = append(upstreams, rewrittenOriginURL) + } else { + upstreams[0] = rewrittenOriginURL + } + upstreamsJSON, err := json.Marshal(upstreams) + if err != nil { + return fmt.Errorf("encode route %d upstreams failed: %w", route.ID, err) + } + if err := tx.Model(&model.OriginProxyRoute{}). + Where("id = ?", route.ID). + Updates(map[string]any{ + "origin_url": rewrittenOriginURL, + "upstreams": string(upstreamsJSON), + }).Error; err != nil { + return fmt.Errorf("update route %d origin address failed: %w", route.ID, err) + } + } + return nil +} diff --git a/Wavelet/internal/apps/openflare/origin/logics_test.go b/Wavelet/internal/apps/openflare/origin/logics_test.go new file mode 100644 index 00000000..95203e86 --- /dev/null +++ b/Wavelet/internal/apps/openflare/origin/logics_test.go @@ -0,0 +1,80 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package origin + +import ( + "context" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupOriginTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate(&model.Origin{})) + + db.SetDB(sqliteDB) + return func() { + db.SetDB(nil) + } +} + +func TestCreateOrigin(t *testing.T) { + cleanup := setupOriginTestDB(t) + defer cleanup() + ctx := context.Background() + + origin, err := CreateOrigin(ctx, Input{ + Name: "Primary Origin", + Address: "origin-a.internal", + Remark: "main upstream", + }) + require.NoError(t, err) + assert.NotZero(t, origin.ID) + assert.Equal(t, "Primary Origin", origin.Name) + assert.Equal(t, "origin-a.internal", origin.Address) + assert.Equal(t, "main upstream", origin.Remark) + + _, err = CreateOrigin(ctx, Input{ + Address: "origin-a.internal", + }) + require.Error(t, err) + assert.Equal(t, errOriginAddressExists, err.Error()) +} + +func TestListOrigins(t *testing.T) { + cleanup := setupOriginTestDB(t) + defer cleanup() + ctx := context.Background() + + first, err := CreateOrigin(ctx, Input{ + Name: "first-origin", + Address: "origin-a.internal", + }) + require.NoError(t, err) + + second, err := CreateOrigin(ctx, Input{ + Name: "second-origin", + Address: "origin-b.internal", + }) + require.NoError(t, err) + + origins, err := ListOrigins(ctx) + require.NoError(t, err) + require.Len(t, origins, 2) + assert.Equal(t, second.ID, origins[0].ID) + assert.Equal(t, first.ID, origins[1].ID) + assert.Equal(t, int64(0), origins[0].RouteCount) + assert.Equal(t, int64(0), origins[1].RouteCount) +} diff --git a/Wavelet/internal/apps/openflare/origin/routers.go b/Wavelet/internal/apps/openflare/origin/routers.go new file mode 100644 index 00000000..b064fe9d --- /dev/null +++ b/Wavelet/internal/apps/openflare/origin/routers.go @@ -0,0 +1,88 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package origin + +import ( + "errors" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +func handleLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, errOriginNotFound) + return true + } + compat.Fail(c, err.Error()) + return true +} + +// GetOrigins 列出全部源站。 +func GetOrigins(c *gin.Context) { + origins, err := ListOrigins(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, origins) +} + +// GetOrigin 获取源站详情。 +func GetOrigin(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + detail, err := GetOriginDetail(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, detail) +} + +// CreateOriginHandler 创建源站。 +func CreateOriginHandler(c *gin.Context) { + var input Input + if !compat.BindJSON(c, &input) { + return + } + origin, err := CreateOrigin(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, origin) +} + +// UpdateOriginHandler 更新源站。 +func UpdateOriginHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input Input + if !compat.BindJSON(c, &input) { + return + } + origin, err := UpdateOrigin(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, origin) +} + +// DeleteOriginHandler 删除源站。 +func DeleteOriginHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteOrigin(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OK(c, nil) +} diff --git a/Wavelet/internal/apps/openflare/pages/errs.go b/Wavelet/internal/apps/openflare/pages/errs.go new file mode 100644 index 00000000..4bbc6f3f --- /dev/null +++ b/Wavelet/internal/apps/openflare/pages/errs.go @@ -0,0 +1,23 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package pages + +const ( + errPagesProjectNotFound = "Pages 项目不存在" + errPagesSlugExists = "Pages 项目标识已存在" + errPagesNameRequired = "Pages 项目名称不能为空" + errPagesSlugInvalid = "Pages 项目标识只能包含小写字母、数字和连字符" + errPagesDeleteReferenced = "Pages 项目已被规则引用,不能删除" + errPagesDeploymentNotFound = "Pages 部署不存在" + errPagesDeploymentMismatch = "Pages 部署不属于该项目" + errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署" + errPagesPackageMissing = "缺少 Pages 部署包" + errPagesPackageNotZip = "Pages 部署包必须是 .zip 文件" + errPagesPackageInvalidZip = "Pages 部署包不是有效 zip 文件" + errPagesPackageEmpty = "Pages 部署包不能为空" + errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空" + errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头" + errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空" + errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL" +) diff --git a/Wavelet/internal/apps/openflare/pages/helpers.go b/Wavelet/internal/apps/openflare/pages/helpers.go new file mode 100644 index 00000000..0b7a3832 --- /dev/null +++ b/Wavelet/internal/apps/openflare/pages/helpers.go @@ -0,0 +1,355 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package pages + +import ( + "archive/zip" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io" + "mime/multipart" + "os" + "path" + "path/filepath" + "regexp" + "strings" + + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + pagesMaxDeploymentFiles = 1000 + pagesMaxDeploymentBytes = 100 * 1024 * 1024 + defaultPagesEntryFile = "index.html" + defaultPagesFallbackPath = "/index.html" +) + +var pagesSlugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,126}[a-z0-9]$|^[a-z0-9]$`) + +type deploymentManifest struct { + Files []model.PagesDeploymentFile + FileCount int + TotalSize int64 + EntryFile string +} + +func isUniqueConstraintError(err error) bool { + if err == nil { + return false + } + return strings.Contains(strings.ToLower(err.Error()), "unique") +} + +func normalizePagesSlug(raw string) string { + value := strings.ToLower(strings.TrimSpace(raw)) + var builder strings.Builder + lastDash := false + for _, r := range value { + valid := (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') + if valid { + builder.WriteRune(r) + lastDash = false + continue + } + if !lastDash { + builder.WriteByte('-') + lastDash = true + } + } + return strings.Trim(builder.String(), "-") +} + +func validateAndNormalizePagesRootDir(raw string) (string, error) { + value := strings.TrimSpace(raw) + if value == "" { + return "", nil + } + if len(value) > 512 { + return "", errors.New("Pages 根目录长度不能超过 512") + } + if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") { + return "", errors.New("Pages 根目录包含不支持的字符") + } + for _, r := range value { + if r <= 0x20 || r == 0x7f { + return "", errors.New("Pages 根目录不能包含空白或控制字符") + } + } + cleaned := path.Clean(filepath.ToSlash(value)) + if cleaned == "." || cleaned == "/" { + return "", nil + } + for _, segment := range strings.Split(cleaned, "/") { + if segment == "." || segment == ".." { + return "", errors.New("Pages 根目录不能包含 . 或 .. 路径段") + } + } + return strings.TrimPrefix(cleaned, "/"), nil +} + +func normalizePagesFallbackPath(raw string) (string, error) { + value := strings.TrimSpace(raw) + if value == "" { + value = defaultPagesFallbackPath + } + if len(value) > 512 { + return "", errors.New("SPA fallback 回退路径长度不能超过 512") + } + if !strings.HasPrefix(value, "/") { + return "", errors.New("SPA fallback 回退路径必须以 / 开头") + } + if value == "/" || strings.HasSuffix(value, "/") { + return "", errors.New("SPA fallback 回退路径必须指向具体文件") + } + if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") { + return "", errors.New("SPA fallback 回退路径包含不支持的字符") + } + for _, r := range value { + if r <= 0x20 || r == 0x7f { + return "", errors.New("SPA fallback 回退路径不能包含空白或控制字符") + } + } + for _, segment := range strings.Split(value, "/") { + if segment == "." || segment == ".." { + return "", errors.New("SPA fallback 回退路径不能包含 . 或 .. 路径段") + } + } + cleaned := path.Clean(value) + if cleaned == "." || !strings.HasPrefix(cleaned, "/") { + return "", errors.New("SPA fallback 回退路径不合法") + } + if cleaned == "/" || strings.HasSuffix(cleaned, "/") { + return "", errors.New("SPA fallback 回退路径必须指向具体文件") + } + return cleaned, nil +} + +func normalizeStoredPagesFallbackPath(value string) string { + normalized, err := normalizePagesFallbackPath(value) + if err != nil { + return defaultPagesFallbackPath + } + return normalized +} + +func normalizePagesEntryFile(raw string) string { + value := path.Clean(strings.TrimSpace(filepath.ToSlash(raw))) + if value == "." || value == "/" { + return defaultPagesEntryFile + } + return strings.TrimPrefix(value, "/") +} + +func persistPagesUploadTemp(fileHeader *multipart.FileHeader) (string, string, error) { + file, err := fileHeader.Open() + if err != nil { + return "", "", err + } + defer file.Close() + temp, err := os.CreateTemp("", "openflare-pages-*.zip") + if err != nil { + return "", "", err + } + defer temp.Close() + hash := sha256.New() + limited := io.LimitReader(file, pagesMaxDeploymentBytes+1) + written, err := io.Copy(io.MultiWriter(temp, hash), limited) + if err != nil { + _ = os.Remove(temp.Name()) + return "", "", err + } + if written > pagesMaxDeploymentBytes { + _ = os.Remove(temp.Name()) + return "", "", fmt.Errorf("Pages 部署包不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024) + } + return temp.Name(), hex.EncodeToString(hash.Sum(nil)), nil +} + +func findCommonRootPrefix(files []*zip.File) (string, error) { + var firstFilePath string + hasMultipleFiles := false + for _, item := range files { + normalizedPath, skip, err := normalizePagesZipPath(item.Name) + if err != nil { + return "", err + } + if skip { + continue + } + if firstFilePath == "" { + firstFilePath = normalizedPath + } else { + hasMultipleFiles = true + } + } + if firstFilePath == "" { + return "", nil + } + parts := strings.Split(firstFilePath, "/") + if len(parts) <= 1 { + return "", nil + } + commonPrefix := parts[0] + "/" + if hasMultipleFiles { + for _, item := range files { + normalizedPath, skip, err := normalizePagesZipPath(item.Name) + if err != nil { + return "", err + } + if skip { + continue + } + if !strings.HasPrefix(normalizedPath, commonPrefix) { + return "", nil + } + } + } + return commonPrefix, nil +} + +func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploymentManifest, error) { + reader, err := zip.OpenReader(zipPath) + if err != nil { + return nil, errors.New(errPagesPackageInvalidZip) + } + defer reader.Close() + + commonPrefix, err := findCommonRootPrefix(reader.File) + if err != nil { + return nil, err + } + + manifest := &deploymentManifest{ + Files: []model.PagesDeploymentFile{}, + EntryFile: entryFile, + } + targetEntryPath := entryFile + if rootDir != "" { + targetEntryPath = path.Join(rootDir, entryFile) + } + entrySeen := false + for _, item := range reader.File { + normalizedPath, skip, err := normalizePagesZipPath(item.Name) + if err != nil { + return nil, err + } + if skip { + continue + } + if commonPrefix != "" { + normalizedPath = strings.TrimPrefix(normalizedPath, commonPrefix) + } + if item.FileInfo().Mode()&os.ModeSymlink != 0 { + return nil, fmt.Errorf("Pages 部署包不支持符号链接: %s", normalizedPath) + } + if item.UncompressedSize64 > pagesMaxDeploymentBytes { + return nil, fmt.Errorf("Pages 文件过大: %s", normalizedPath) + } + manifest.FileCount++ + if manifest.FileCount > pagesMaxDeploymentFiles { + return nil, fmt.Errorf("Pages 部署文件数不能超过 %d", pagesMaxDeploymentFiles) + } + manifest.TotalSize += int64(item.UncompressedSize64) + if manifest.TotalSize > pagesMaxDeploymentBytes { + return nil, fmt.Errorf("Pages 部署展开后不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024) + } + checksum, err := checksumZipFile(item) + if err != nil { + return nil, err + } + if normalizedPath == targetEntryPath { + entrySeen = true + } + manifest.Files = append(manifest.Files, model.PagesDeploymentFile{ + Path: normalizedPath, + Size: int64(item.UncompressedSize64), + Checksum: checksum, + }) + } + if manifest.FileCount == 0 { + return nil, errors.New(errPagesPackageEmpty) + } + if !entrySeen { + return nil, fmt.Errorf("Pages 部署包缺少入口文件 %s", targetEntryPath) + } + return manifest, nil +} + +func normalizePagesZipPath(raw string) (string, bool, error) { + name := strings.TrimSpace(filepath.ToSlash(raw)) + if name == "" { + return "", true, nil + } + if strings.HasSuffix(name, "/") { + return "", true, nil + } + if strings.HasPrefix(name, "/") || path.IsAbs(name) { + return "", false, fmt.Errorf("Pages 部署包不能包含绝对路径: %s", raw) + } + cleaned := path.Clean(name) + if cleaned == "." { + return "", true, nil + } + if cleaned == ".." || strings.HasPrefix(cleaned, "../") || strings.Contains(cleaned, "/../") { + return "", false, fmt.Errorf("Pages 部署包路径不能逃逸目录: %s", raw) + } + return cleaned, false, nil +} + +func checksumZipFile(item *zip.File) (string, error) { + file, err := item.Open() + if err != nil { + return "", err + } + defer file.Close() + hash := sha256.New() + if _, err = io.Copy(hash, file); err != nil { + return "", err + } + return hex.EncodeToString(hash.Sum(nil)), nil +} + +func pagesArtifactPath(projectSlug string, checksum string) (string, error) { + root, err := pagesStorageRoot() + if err != nil { + return "", err + } + return filepath.Join(root, "artifacts", projectSlug, checksum+".zip"), nil +} + +func pagesStorageRoot() (string, error) { + cfg := config.Config.Database + if cfg.Enabled { + return filepath.Abs(filepath.Join("data", "pages")) + } + dbPath := strings.TrimSpace(cfg.SQLitePath) + if dbPath == "" || dbPath == ":memory:" { + return filepath.Abs(filepath.Join("data", "pages")) + } + dir := filepath.Dir(dbPath) + if dir == "." || dir == "" { + dir = "data" + } + return filepath.Abs(filepath.Join(dir, "pages")) +} + +func copyFile(src string, dst string) error { + input, err := os.Open(src) + if err != nil { + return err + } + defer input.Close() + output, err := os.OpenFile(dst, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o644) + if err != nil { + return err + } + defer output.Close() + if _, err = io.Copy(output, input); err != nil { + return err + } + return output.Sync() +} diff --git a/Wavelet/internal/apps/openflare/pages/logics.go b/Wavelet/internal/apps/openflare/pages/logics.go new file mode 100644 index 00000000..ba1850f3 --- /dev/null +++ b/Wavelet/internal/apps/openflare/pages/logics.go @@ -0,0 +1,485 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package pages + +import ( + "context" + "errors" + "fmt" + "mime/multipart" + "net/url" + "os" + "path/filepath" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +// Input Pages 项目创建/更新请求。 +type Input struct { + Name string `json:"name"` + Slug string `json:"slug"` + Description string `json:"description"` + Enabled bool `json:"enabled"` + SPAFallbackEnabled bool `json:"spa_fallback_enabled"` + SPAFallbackPath string `json:"spa_fallback_path"` + APIProxyEnabled bool `json:"api_proxy_enabled"` + APIProxyPath string `json:"api_proxy_path"` + APIProxyPass string `json:"api_proxy_pass"` + APIProxyRewrite string `json:"api_proxy_rewrite"` + RootDir string `json:"root_dir"` + EntryFile string `json:"entry_file"` +} + +// DeploymentView Pages 部署视图。 +type DeploymentView struct { + ID uint `json:"id"` + ProjectID uint `json:"project_id"` + DeploymentNumber int `json:"deployment_number"` + Checksum string `json:"checksum"` + Status string `json:"status"` + FileCount int `json:"file_count"` + TotalSize int64 `json:"total_size"` + CreatedBy string `json:"created_by"` + CreatedAt time.Time `json:"created_at"` + ActivatedAt *time.Time `json:"activated_at"` +} + +// DeploymentFileView Pages 部署文件视图。 +type DeploymentFileView struct { + ID uint `json:"id"` + DeploymentID uint `json:"deployment_id"` + Path string `json:"path"` + Size int64 `json:"size"` + Checksum string `json:"checksum"` + CreatedAt time.Time `json:"created_at"` +} + +// View Pages 项目视图。 +type View struct { + ID uint `json:"id"` + Name string `json:"name"` + Slug string `json:"slug"` + Description string `json:"description"` + Enabled bool `json:"enabled"` + SPAFallbackEnabled bool `json:"spa_fallback_enabled"` + SPAFallbackPath string `json:"spa_fallback_path"` + APIProxyEnabled bool `json:"api_proxy_enabled"` + APIProxyPath string `json:"api_proxy_path"` + APIProxyPass string `json:"api_proxy_pass"` + APIProxyRewrite string `json:"api_proxy_rewrite"` + RootDir string `json:"root_dir"` + EntryFile string `json:"entry_file"` + ActiveDeploymentID *uint `json:"active_deployment_id"` + ActiveDeployment *DeploymentView `json:"active_deployment,omitempty"` + DeploymentCount int64 `json:"deployment_count"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// ListProjects 列出全部 Pages 项目。 +func ListProjects(ctx context.Context) ([]View, error) { + projects, err := model.ListPagesProjects(ctx) + if err != nil { + return nil, err + } + views := make([]View, 0, len(projects)) + for _, project := range projects { + view, err := buildProjectView(ctx, &project) + if err != nil { + return nil, err + } + views = append(views, *view) + } + return views, nil +} + +// GetProject 获取 Pages 项目详情。 +func GetProject(ctx context.Context, id uint) (*View, error) { + project, err := model.GetPagesProjectByID(ctx, id) + if err != nil { + return nil, err + } + return buildProjectView(ctx, project) +} + +// CreateProject 创建 Pages 项目。 +func CreateProject(ctx context.Context, input Input) (*View, error) { + project, err := buildProject(nil, input) + if err != nil { + return nil, err + } + if err = model.CreatePagesProjectRecord(ctx, project); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errPagesSlugExists) + } + return nil, err + } + return buildProjectView(ctx, project) +} + +// UpdateProject 更新 Pages 项目。 +func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) { + project, err := model.GetPagesProjectByID(ctx, id) + if err != nil { + return nil, err + } + project, err = buildProject(project, input) + if err != nil { + return nil, err + } + if err = db.DB(ctx).Model(project).Updates(map[string]any{ + "name": project.Name, + "slug": project.Slug, + "description": project.Description, + "enabled": project.Enabled, + "spa_fallback_enabled": project.SPAFallbackEnabled, + "spa_fallback_path": project.SPAFallbackPath, + "api_proxy_enabled": project.APIProxyEnabled, + "api_proxy_path": project.APIProxyPath, + "api_proxy_pass": project.APIProxyPass, + "api_proxy_rewrite": project.APIProxyRewrite, + "root_dir": project.RootDir, + "entry_file": project.EntryFile, + }).Error; err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errPagesSlugExists) + } + return nil, err + } + return buildProjectView(ctx, project) +} + +// DeleteProject 删除 Pages 项目。 +func DeleteProject(ctx context.Context, id uint) error { + project, err := model.GetPagesProjectByID(ctx, id) + if err != nil { + return err + } + routeCount, err := model.CountProxyRoutesByPagesProjectID(ctx, project.ID) + if err != nil { + return err + } + if routeCount > 0 { + return errors.New(errPagesDeleteReferenced) + } + deployments, err := model.ListPagesDeployments(ctx, project.ID) + if err != nil { + return err + } + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where( + "deployment_id IN (?)", + tx.Model(&model.PagesDeployment{}).Select("id").Where("project_id = ?", project.ID), + ).Delete(&model.PagesDeploymentFile{}).Error; err != nil { + return err + } + if err := tx.Where("project_id = ?", project.ID).Delete(&model.PagesDeployment{}).Error; err != nil { + return err + } + if err := tx.Delete(project).Error; err != nil { + return err + } + for _, deployment := range deployments { + _ = os.Remove(deployment.ArtifactPath) + } + return nil + }) +} + +// ListProjectDeployments 列出项目的全部部署。 +func ListProjectDeployments(ctx context.Context, projectID uint) ([]DeploymentView, error) { + if _, err := model.GetPagesProjectByID(ctx, projectID); err != nil { + return nil, err + } + deployments, err := model.ListPagesDeployments(ctx, projectID) + if err != nil { + return nil, err + } + views := make([]DeploymentView, 0, len(deployments)) + for _, deployment := range deployments { + views = append(views, buildDeploymentView(&deployment)) + } + return views, nil +} + +// ListDeploymentFiles 列出部署文件清单。 +func ListDeploymentFiles(ctx context.Context, deploymentID uint) ([]DeploymentFileView, error) { + if _, err := model.GetPagesDeploymentByID(ctx, deploymentID); err != nil { + return nil, err + } + files, err := model.ListPagesDeploymentFiles(ctx, deploymentID) + if err != nil { + return nil, err + } + views := make([]DeploymentFileView, 0, len(files)) + for _, file := range files { + views = append(views, DeploymentFileView{ + ID: file.ID, + DeploymentID: file.DeploymentID, + Path: file.Path, + Size: file.Size, + Checksum: file.Checksum, + CreatedAt: file.CreatedAt, + }) + } + return views, nil +} + +// UploadDeployment 上传 Pages 部署包。 +func UploadDeployment(ctx context.Context, projectID uint, fileHeader *multipart.FileHeader, createdBy string) (*DeploymentView, error) { + project, err := model.GetPagesProjectByID(ctx, projectID) + if err != nil { + return nil, err + } + if fileHeader == nil { + return nil, errors.New(errPagesPackageMissing) + } + if !strings.EqualFold(filepath.Ext(fileHeader.Filename), ".zip") { + return nil, errors.New(errPagesPackageNotZip) + } + rootDir, err := validateAndNormalizePagesRootDir(project.RootDir) + if err != nil { + return nil, err + } + entryFile := normalizePagesEntryFile(project.EntryFile) + tempPath, checksum, err := persistPagesUploadTemp(fileHeader) + if err != nil { + return nil, err + } + defer os.Remove(tempPath) + manifest, err := inspectPagesZip(tempPath, rootDir, entryFile) + if err != nil { + return nil, err + } + artifactPath, err := pagesArtifactPath(project.Slug, checksum) + if err != nil { + return nil, err + } + if err = os.MkdirAll(filepath.Dir(artifactPath), 0o755); err != nil { + return nil, fmt.Errorf("创建 Pages 存储目录失败: %w", err) + } + if err = copyFile(tempPath, artifactPath); err != nil { + return nil, err + } + deployment := &model.PagesDeployment{} + err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + var maxNumber int + if err := tx.Model(&model.PagesDeployment{}). + Where("project_id = ?", project.ID). + Select("COALESCE(MAX(deployment_number), 0)"). + Scan(&maxNumber).Error; err != nil { + return err + } + deployment = &model.PagesDeployment{ + ProjectID: project.ID, + DeploymentNumber: maxNumber + 1, + Checksum: checksum, + Status: model.PagesDeploymentStatusUploaded, + ArtifactPath: artifactPath, + FileCount: manifest.FileCount, + TotalSize: manifest.TotalSize, + CreatedBy: strings.TrimSpace(createdBy), + } + if err := tx.Create(deployment).Error; err != nil { + return err + } + for index := range manifest.Files { + manifest.Files[index].DeploymentID = deployment.ID + } + if len(manifest.Files) > 0 { + if err := tx.Create(&manifest.Files).Error; err != nil { + return err + } + } + return nil + }) + if err != nil { + _ = os.Remove(artifactPath) + return nil, err + } + view := buildDeploymentView(deployment) + return &view, nil +} + +// ActivateDeployment 激活 Pages 部署。 +func ActivateDeployment(ctx context.Context, projectID uint, deploymentID uint) (*View, error) { + project, err := model.GetPagesProjectByID(ctx, projectID) + if err != nil { + return nil, err + } + deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) + if err != nil { + return nil, err + } + if deployment.ProjectID != project.ID { + return nil, errors.New(errPagesDeploymentMismatch) + } + now := time.Now() + if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Model(&model.PagesDeployment{}). + Where("project_id = ?", project.ID). + Update("status", model.PagesDeploymentStatusUploaded).Error; err != nil { + return err + } + if err := tx.Model(deployment).Updates(map[string]any{ + "status": model.PagesDeploymentStatusActive, + "activated_at": &now, + }).Error; err != nil { + return err + } + return tx.Model(project).Updates(map[string]any{ + "active_deployment_id": deployment.ID, + }).Error + }); err != nil { + return nil, err + } + return GetProject(ctx, project.ID) +} + +// DeleteDeployment 删除 Pages 部署。 +func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error { + project, err := model.GetPagesProjectByID(ctx, projectID) + if err != nil { + return err + } + deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) + if err != nil { + return err + } + if deployment.ProjectID != project.ID { + return errors.New(errPagesDeploymentMismatch) + } + if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID == deployment.ID { + return errors.New(errPagesDeleteActiveDeploy) + } + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("deployment_id = ?", deployment.ID).Delete(&model.PagesDeploymentFile{}).Error; err != nil { + return err + } + if err := tx.Delete(deployment).Error; err != nil { + return err + } + _ = os.Remove(deployment.ArtifactPath) + return nil + }) +} + +func buildProject(existing *model.PagesProject, input Input) (*model.PagesProject, error) { + name := strings.TrimSpace(input.Name) + if name == "" { + return nil, errors.New(errPagesNameRequired) + } + slug := normalizePagesSlug(input.Slug) + if slug == "" { + slug = normalizePagesSlug(name) + } + if !pagesSlugPattern.MatchString(slug) { + return nil, errors.New(errPagesSlugInvalid) + } + if existing == nil { + existing = &model.PagesProject{} + } + existing.Name = name + existing.Slug = slug + existing.Description = strings.TrimSpace(input.Description) + existing.Enabled = input.Enabled + existing.SPAFallbackEnabled = input.SPAFallbackEnabled + fallbackPath, err := normalizePagesFallbackPath(input.SPAFallbackPath) + if err != nil { + return nil, err + } + existing.SPAFallbackPath = fallbackPath + + existing.APIProxyEnabled = input.APIProxyEnabled + apiProxyPath := strings.TrimSpace(input.APIProxyPath) + apiProxyPass := strings.TrimSpace(input.APIProxyPass) + apiProxyRewrite := strings.TrimSpace(input.APIProxyRewrite) + + if existing.APIProxyEnabled { + if apiProxyPath == "" { + return nil, errors.New(errPagesAPIProxyPathRequired) + } + if !strings.HasPrefix(apiProxyPath, "/") { + return nil, errors.New(errPagesAPIProxyPathPrefix) + } + if apiProxyPass == "" { + return nil, errors.New(errPagesAPIProxyPassRequired) + } + parsedURL, err := url.Parse(apiProxyPass) + if err != nil || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") || parsedURL.Host == "" { + return nil, errors.New(errPagesAPIProxyPassInvalid) + } + } + existing.APIProxyPath = apiProxyPath + existing.APIProxyPass = apiProxyPass + existing.APIProxyRewrite = apiProxyRewrite + + rootDir, err := validateAndNormalizePagesRootDir(input.RootDir) + if err != nil { + return nil, err + } + existing.RootDir = rootDir + existing.EntryFile = normalizePagesEntryFile(input.EntryFile) + + return existing, nil +} + +func buildProjectView(ctx context.Context, project *model.PagesProject) (*View, error) { + if project == nil { + return nil, errors.New(errPagesProjectNotFound) + } + view := &View{ + ID: project.ID, + Name: project.Name, + Slug: project.Slug, + Description: project.Description, + Enabled: project.Enabled, + SPAFallbackEnabled: project.SPAFallbackEnabled, + SPAFallbackPath: normalizeStoredPagesFallbackPath(project.SPAFallbackPath), + APIProxyEnabled: project.APIProxyEnabled, + APIProxyPath: project.APIProxyPath, + APIProxyPass: project.APIProxyPass, + APIProxyRewrite: project.APIProxyRewrite, + RootDir: project.RootDir, + EntryFile: project.EntryFile, + ActiveDeploymentID: project.ActiveDeploymentID, + CreatedAt: project.CreatedAt, + UpdatedAt: project.UpdatedAt, + } + count, err := model.CountPagesDeploymentsByProjectID(ctx, project.ID) + if err != nil { + return nil, err + } + view.DeploymentCount = count + if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID != 0 { + deployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) + if err == nil { + active := buildDeploymentView(deployment) + view.ActiveDeployment = &active + } + } + return view, nil +} + +func buildDeploymentView(deployment *model.PagesDeployment) DeploymentView { + if deployment == nil { + return DeploymentView{} + } + return DeploymentView{ + ID: deployment.ID, + ProjectID: deployment.ProjectID, + DeploymentNumber: deployment.DeploymentNumber, + Checksum: deployment.Checksum, + Status: deployment.Status, + FileCount: deployment.FileCount, + TotalSize: deployment.TotalSize, + CreatedBy: deployment.CreatedBy, + CreatedAt: deployment.CreatedAt, + ActivatedAt: deployment.ActivatedAt, + } +} diff --git a/Wavelet/internal/apps/openflare/pages/logics_test.go b/Wavelet/internal/apps/openflare/pages/logics_test.go new file mode 100644 index 00000000..77a0413c --- /dev/null +++ b/Wavelet/internal/apps/openflare/pages/logics_test.go @@ -0,0 +1,84 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package pages + +import ( + "context" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupPagesTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.PagesProject{}, + &model.PagesDeployment{}, + &model.PagesDeploymentFile{}, + )) + + db.SetDB(sqliteDB) + return func() { + db.SetDB(nil) + } +} + +func TestCreateProject(t *testing.T) { + cleanup := setupPagesTestDB(t) + defer cleanup() + ctx := context.Background() + + project, err := CreateProject(ctx, Input{ + Name: "Marketing Site", + Slug: "marketing-site", + Description: "public site", + Enabled: true, + SPAFallbackEnabled: true, + SPAFallbackPath: "/index.html", + EntryFile: "index.html", + }) + require.NoError(t, err) + assert.NotZero(t, project.ID) + assert.Equal(t, "Marketing Site", project.Name) + assert.Equal(t, "marketing-site", project.Slug) + assert.Equal(t, "public site", project.Description) + assert.True(t, project.Enabled) + assert.True(t, project.SPAFallbackEnabled) + assert.Equal(t, "/index.html", project.SPAFallbackPath) + assert.Equal(t, "index.html", project.EntryFile) + assert.Equal(t, int64(0), project.DeploymentCount) + + _, err = CreateProject(ctx, Input{ + Name: "Duplicate Slug", + Slug: "marketing-site", + }) + require.Error(t, err) + assert.Equal(t, errPagesSlugExists, err.Error()) +} + +func TestCreateProjectRejectsUnsafeFallbackPath(t *testing.T) { + cleanup := setupPagesTestDB(t) + defer cleanup() + ctx := context.Background() + + _, err := CreateProject(ctx, Input{ + Name: "Unsafe Fallback", + Slug: "unsafe-fallback", + Enabled: true, + SPAFallbackEnabled: true, + SPAFallbackPath: "/index.html; proxy_pass http://evil", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "回退路径") +} diff --git a/Wavelet/internal/apps/openflare/pages/routers.go b/Wavelet/internal/apps/openflare/pages/routers.go new file mode 100644 index 00000000..f0e527a8 --- /dev/null +++ b/Wavelet/internal/apps/openflare/pages/routers.go @@ -0,0 +1,180 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package pages + +import ( + "errors" + "strconv" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +func handleLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, errPagesProjectNotFound) + return true + } + compat.Fail(c, err.Error()) + return true +} + +func deploymentIDParam(c *gin.Context) (uint, bool) { + raw := c.Param("deployment_id") + if raw == "" { + compat.Fail(c, "无效的 ID") + return 0, false + } + id64, err := strconv.ParseUint(raw, 10, 64) + if err != nil || id64 == 0 { + compat.Fail(c, "无效的 ID") + return 0, false + } + return uint(id64), true +} + +// ListProjectsHandler 列出全部 Pages 项目。 +func ListProjectsHandler(c *gin.Context) { + projects, err := ListProjects(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, projects) +} + +// GetProjectHandler 获取 Pages 项目详情。 +func GetProjectHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + project, err := GetProject(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, project) +} + +// CreateProjectHandler 创建 Pages 项目。 +func CreateProjectHandler(c *gin.Context) { + var input Input + if !compat.BindJSON(c, &input) { + return + } + project, err := CreateProject(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, project) +} + +// UpdateProjectHandler 更新 Pages 项目。 +func UpdateProjectHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input Input + if !compat.BindJSON(c, &input) { + return + } + project, err := UpdateProject(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, project) +} + +// DeleteProjectHandler 删除 Pages 项目。 +func DeleteProjectHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteProject(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OK(c, nil) +} + +// ListDeploymentsHandler 列出项目的全部部署。 +func ListDeploymentsHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + deployments, err := ListProjectDeployments(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, deployments) +} + +// UploadDeploymentHandler 上传 Pages 部署包。 +func UploadDeploymentHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + file, err := c.FormFile("package") + if err != nil { + compat.Fail(c, errPagesPackageMissing) + return + } + deployment, err := UploadDeployment(c.Request.Context(), id, file, "") + if handleLogicError(c, err) { + return + } + compat.OK(c, deployment) +} + +// ActivateDeploymentHandler 激活 Pages 部署。 +func ActivateDeploymentHandler(c *gin.Context) { + projectID, ok := compat.IDParam(c) + if !ok { + return + } + deploymentID, ok := deploymentIDParam(c) + if !ok { + return + } + project, err := ActivateDeployment(c.Request.Context(), projectID, deploymentID) + if handleLogicError(c, err) { + return + } + compat.OK(c, project) +} + +// DeleteDeploymentHandler 删除 Pages 部署。 +func DeleteDeploymentHandler(c *gin.Context) { + projectID, ok := compat.IDParam(c) + if !ok { + return + } + deploymentID, ok := deploymentIDParam(c) + if !ok { + return + } + if err := DeleteDeployment(c.Request.Context(), projectID, deploymentID); handleLogicError(c, err) { + return + } + compat.OK(c, nil) +} + +// ListDeploymentFilesHandler 列出部署文件清单。 +func ListDeploymentFilesHandler(c *gin.Context) { + deploymentID, ok := deploymentIDParam(c) + if !ok { + return + } + files, err := ListDeploymentFiles(c.Request.Context(), deploymentID) + if handleLogicError(c, err) { + return + } + compat.OK(c, files) +} diff --git a/Wavelet/internal/apps/openflare/proxy_route/errs.go b/Wavelet/internal/apps/openflare/proxy_route/errs.go new file mode 100644 index 00000000..e458bd87 --- /dev/null +++ b/Wavelet/internal/apps/openflare/proxy_route/errs.go @@ -0,0 +1,53 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package proxy_route + +const ( + errProxyRouteNotFound = "proxy route not found" + errProxyRouteIdentityExists = "proxy route identity already exists" + errProxyRouteSiteNameExists = "site_name already exists" + errProxyRouteDomainExists = "domain %s already exists" + errProxyRouteSiteNameEmpty = "site_name cannot be empty" + errProxyRouteDomainRequired = "at least one domain is required" + errProxyRouteDomainInvalid = "domain format is invalid" + errProxyRouteDomainMismatch = "domain must match domains[0]" + errProxyRouteOriginEmpty = "origin_url cannot be empty" + errProxyRouteOriginInvalid = "origin URL format is invalid" + errProxyRouteOriginScheme = "origin URL must start with http:// or https://" + errProxyRouteOriginHostInvalid = "origin_host format is invalid" + errProxyRouteUpstreamRequired = "at least one upstream is required" + errProxyRouteUpstreamScheme = "all upstreams must use the same scheme" + errProxyRouteUpstreamPath = "multi-upstream mode does not support origin paths" + errProxyRouteUpstreamQuery = "multi-upstream mode does not support origin query strings" + errProxyRouteOriginNotFound = "selected origin does not exist" + errProxyRouteCertNotFound = "selected certificate does not exist" + errProxyRouteCertRequired = "must select a certificate when HTTPS is enabled" + errProxyRouteCertDomainLength = "domain_cert_ids must match domains length" + errProxyRouteRedirectHTTP = "redirect_http requires enable_https" + errProxyRouteBasicAuth = "basic_auth_username and basic_auth_password cannot be empty when basic auth is enabled" + errProxyRouteLimitRate = "limit_rate must be a number or use the 512k / 1m format" + errProxyRouteCachePolicy = "cache policy is not supported" + errProxyRouteCacheSuffix = "cache suffix format is invalid" + errProxyRouteCachePath = "cache path rule format is invalid" + errProxyRouteCacheSuffixReq = "at least one suffix is required" + errProxyRouteCachePrefixReq = "at least one path prefix is required" + errProxyRouteCacheExactReq = "at least one exact path is required" + errProxyRouteHeaderKeyEmpty = "custom header key cannot be empty" + errProxyRouteHeaderKeyInvalid = "custom header key format is invalid" + errProxyRouteHeaderNewline = "custom headers cannot contain newlines" + errProxyRouteTunnelNodeReq = "tunnel_node_id is required for tunnel upstream" + errProxyRouteTunnelNodeMissing = "tunnel client node does not exist" + errProxyRouteTunnelNodeType = "tunnel_node_id must reference a tunnel_client node" + errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream" + errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https" + errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream" + errProxyRoutePagesNotFound = "Pages 项目不存在" + errProxyRoutePagesDisabled = "Pages 项目未启用" + errProxyRoutePagesNoDeploy = "Pages 项目没有激活部署" + errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https" + errProxyRouteOriginPort = "端口格式不合法" + errProxyRouteOriginPortEmpty = "端口不能为空" + errProxyRouteOriginURI = "源站路径需以 / 或 ? 开头" + errProxyRouteOriginURIProto = "源站路径不能包含协议" +) diff --git a/Wavelet/internal/apps/openflare/proxy_route/helpers.go b/Wavelet/internal/apps/openflare/proxy_route/helpers.go new file mode 100644 index 00000000..4abe3673 --- /dev/null +++ b/Wavelet/internal/apps/openflare/proxy_route/helpers.go @@ -0,0 +1,1044 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package proxy_route + +import ( + "context" + "crypto/x509" + "encoding/json" + "encoding/pem" + "errors" + "fmt" + "net" + "net/url" + "regexp" + "strconv" + "strings" + "unicode" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +var proxyHeaderKeyPattern = regexp.MustCompile(`^[A-Za-z0-9_-]+$`) +var proxyRouteLimitRatePattern = regexp.MustCompile(`^\d+[kKmM]?$`) + +const ( + proxyRouteCachePolicyURL = "url" + proxyRouteCachePolicySuffix = "suffix" + proxyRouteCachePolicyPathPrefix = "path_prefix" + proxyRouteCachePolicyPathExact = "path_exact" +) + +type tlsCertificateRow struct { + ID uint `gorm:"column:id;primaryKey"` + CertPEM string `gorm:"column:cert_pem"` +} + +func (tlsCertificateRow) TableName() string { + return "of_tls_certificates" +} + +type tunnelNodeRow struct { + ID uint `gorm:"column:id;primaryKey"` + NodeType string `gorm:"column:node_type"` +} + +func (tunnelNodeRow) TableName() string { + return "of_nodes" +} + +type pagesProjectRow struct { + ID uint `gorm:"column:id;primaryKey"` + Enabled bool `gorm:"column:enabled"` + ActiveDeploymentID *uint `gorm:"column:active_deployment_id"` +} + +func (pagesProjectRow) TableName() string { + return "of_pages_projects" +} + +func uniqueStrings(items []string) []string { + if len(items) == 0 { + return items + } + seen := make(map[string]struct{}, len(items)) + result := make([]string, 0, len(items)) + for _, item := range items { + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + result = append(result, item) + } + return result +} + +func isUniqueConstraintError(err error) bool { + if err == nil { + return false + } + return strings.Contains(strings.ToLower(err.Error()), "unique") +} + +func normalizeOriginAddress(raw string) string { + return strings.ToLower(strings.TrimSpace(raw)) +} + +func validateOriginAddress(address string) error { + if address == "" { + return errors.New(errProxyRouteOriginEmpty) + } + if strings.Contains(address, "://") || strings.ContainsAny(address, "/?#") { + return errors.New(errProxyRouteOriginInvalid) + } + if strings.HasPrefix(address, "[") || strings.HasSuffix(address, "]") { + return errors.New(errProxyRouteOriginInvalid) + } + if ip := net.ParseIP(address); ip != nil { + return nil + } + if len(address) > 253 { + return errors.New(errProxyRouteOriginInvalid) + } + labels := strings.Split(address, ".") + for _, label := range labels { + if len(label) == 0 || len(label) > 63 { + return errors.New(errProxyRouteOriginInvalid) + } + if label[0] == '-' || label[len(label)-1] == '-' { + return errors.New(errProxyRouteOriginInvalid) + } + for _, r := range label { + if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' { + continue + } + return errors.New(errProxyRouteOriginInvalid) + } + } + return nil +} + +func normalizeOriginPort(raw string) (string, error) { + port := strings.TrimSpace(raw) + if port == "" { + return "", errors.New(errProxyRouteOriginPortEmpty) + } + value, err := strconv.Atoi(port) + if err != nil || value < 1 || value > 65535 { + return "", errors.New(errProxyRouteOriginPort) + } + return strconv.Itoa(value), nil +} + +func normalizeOriginScheme(raw string) (string, error) { + scheme := strings.ToLower(strings.TrimSpace(raw)) + switch scheme { + case "http", "https": + return scheme, nil + default: + return "", errors.New(errProxyRouteOriginSchemeOnly) + } +} + +func normalizeOriginURI(raw string) (string, error) { + uri := strings.TrimSpace(raw) + if uri == "" { + return "", nil + } + if strings.Contains(uri, "://") { + return "", errors.New(errProxyRouteOriginURIProto) + } + if !strings.HasPrefix(uri, "/") && !strings.HasPrefix(uri, "?") { + return "", errors.New(errProxyRouteOriginURI) + } + return uri, nil +} + +func formatOriginHost(address string, port string) string { + return net.JoinHostPort(address, port) +} + +func buildOriginURLFromParts(scheme, address, port, uri string) (string, error) { + normalizedScheme, err := normalizeOriginScheme(scheme) + if err != nil { + return "", err + } + normalizedAddress := normalizeOriginAddress(address) + if err := validateOriginAddress(normalizedAddress); err != nil { + return "", err + } + normalizedPort, err := normalizeOriginPort(port) + if err != nil { + return "", err + } + normalizedURI, err := normalizeOriginURI(uri) + if err != nil { + return "", err + } + + parsed := &url.URL{ + Scheme: normalizedScheme, + Host: formatOriginHost(normalizedAddress, normalizedPort), + } + if normalizedURI != "" { + if strings.HasPrefix(normalizedURI, "?") { + parsed.RawQuery = strings.TrimPrefix(normalizedURI, "?") + } else { + pathQuery := strings.SplitN(normalizedURI, "?", 2) + parsed.Path = pathQuery[0] + if len(pathQuery) > 1 { + parsed.RawQuery = pathQuery[1] + } + } + } + return parsed.String(), nil +} + +func extractOriginAddress(rawURL string) (string, error) { + parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL)) + if err != nil { + return "", fmt.Errorf("%s: %w", errProxyRouteOriginInvalid, err) + } + address := normalizeOriginAddress(parsed.Hostname()) + if err := validateOriginAddress(address); err != nil { + return "", err + } + return address, nil +} + +func getOrCreateOriginByAddress(ctx context.Context, address string) (*model.Origin, error) { + normalizedAddress := normalizeOriginAddress(address) + if err := validateOriginAddress(normalizedAddress); err != nil { + return nil, err + } + existing, err := model.GetOriginByAddress(ctx, normalizedAddress) + if err == nil { + return existing, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + origin := &model.Origin{ + Name: normalizedAddress, + Address: normalizedAddress, + Remark: "", + } + if err := model.CreateOriginRecord(ctx, origin); err != nil { + if isUniqueConstraintError(err) { + return model.GetOriginByAddress(ctx, normalizedAddress) + } + return nil, err + } + return origin, nil +} + +func lookupTLSCertificateByID(ctx context.Context, id uint) (*tlsCertificateRow, error) { + if !db.DB(ctx).Migrator().HasTable(&tlsCertificateRow{}) { + return nil, gorm.ErrRecordNotFound + } + var certificate tlsCertificateRow + if err := db.DB(ctx).First(&certificate, id).Error; err != nil { + return nil, err + } + return &certificate, nil +} + +func lookupTunnelNodeByID(ctx context.Context, id uint) (*tunnelNodeRow, error) { + if !db.DB(ctx).Migrator().HasTable(&tunnelNodeRow{}) { + return nil, gorm.ErrRecordNotFound + } + var node tunnelNodeRow + if err := db.DB(ctx).First(&node, id).Error; err != nil { + return nil, err + } + return &node, nil +} + +func lookupPagesProjectByID(ctx context.Context, id uint) (*pagesProjectRow, error) { + if !db.DB(ctx).Migrator().HasTable(&pagesProjectRow{}) { + return nil, gorm.ErrRecordNotFound + } + var project pagesProjectRow + if err := db.DB(ctx).First(&project, id).Error; err != nil { + return nil, err + } + return &project, nil +} + +func parseLeafCertificate(certPEM string) (*x509.Certificate, error) { + certPEMBlock, _ := pem.Decode([]byte(certPEM)) + if certPEMBlock == nil { + return nil, errors.New(errProxyRouteCertNotFound) + } + leaf, err := x509.ParseCertificate(certPEMBlock.Bytes) + if err != nil { + return nil, err + } + return leaf, nil +} + +func validateCertificateCoverage(certificate *tlsCertificateRow, domains []string) error { + if certificate == nil { + return errors.New(errProxyRouteCertNotFound) + } + leaf, err := parseLeafCertificate(certificate.CertPEM) + if err != nil { + return err + } + for _, domain := range domains { + if err := leaf.VerifyHostname(domain); err != nil { + return fmt.Errorf("certificate does not cover domain %s", domain) + } + } + return nil +} + +func loadTLSCertificates(ctx context.Context, certIDs []uint) ([]*tlsCertificateRow, error) { + certificates := make([]*tlsCertificateRow, 0, len(certIDs)) + for _, certID := range certIDs { + certificate, err := lookupTLSCertificateByID(ctx, certID) + if err != nil { + return nil, err + } + certificates = append(certificates, certificate) + } + return certificates, nil +} + +func normalizeProxyRouteSiteNameInput(route *model.ProxyRoute, raw, primaryDomain string) string { + siteName := strings.TrimSpace(raw) + if siteName != "" { + return siteName + } + if route != nil && strings.TrimSpace(route.SiteName) != "" { + return strings.TrimSpace(route.SiteName) + } + return primaryDomain +} + +func normalizeProxyRouteDomainValue(raw string) string { + return strings.ToLower(strings.TrimSpace(raw)) +} + +func normalizeProxyRouteDomains(rawDomains []string) ([]string, error) { + normalized := make([]string, 0, len(rawDomains)) + for _, rawDomain := range rawDomains { + domain := normalizeProxyRouteDomainValue(rawDomain) + if domain == "" { + continue + } + if strings.Contains(domain, "://") || strings.Contains(domain, "/") { + return nil, errors.New(errProxyRouteDomainInvalid) + } + normalized = append(normalized, domain) + } + normalized = uniqueStrings(normalized) + if len(normalized) == 0 { + return nil, errors.New(errProxyRouteDomainRequired) + } + return normalized, nil +} + +func normalizeProxyRouteDomainsInput(route *model.ProxyRoute, rawDomain string, rawDomains []string) ([]string, error) { + if len(rawDomains) > 0 { + domains, err := normalizeProxyRouteDomains(rawDomains) + if err != nil { + return nil, err + } + domain := normalizeProxyRouteDomainValue(rawDomain) + if domain != "" && domain != domains[0] { + return nil, errors.New(errProxyRouteDomainMismatch) + } + return domains, nil + } + + if route != nil { + existingDomains, err := decodeStoredDomains(route.Domains, route.Domain) + if err == nil && len(existingDomains) > 0 { + domain := normalizeProxyRouteDomainValue(rawDomain) + if domain == "" || domain == existingDomains[0] { + return existingDomains, nil + } + } + } + + return normalizeProxyRouteDomains([]string{rawDomain}) +} + +func validateProxyRouteSiteName(siteName string) error { + if strings.TrimSpace(siteName) == "" { + return errors.New(errProxyRouteSiteNameEmpty) + } + return nil +} + +func validateProxyRouteIdentityUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string, domains []string) error { + routes, err := model.ListProxyRoutes(ctx) + if err != nil { + return err + } + + currentID := uint(0) + if route != nil { + currentID = route.ID + } + + for _, item := range routes { + if item == nil || item.ID == currentID { + continue + } + existingSiteName := normalizeProxyRouteSiteNameInput(item, item.SiteName, item.Domain) + if existingSiteName == siteName { + return errors.New(errProxyRouteSiteNameExists) + } + + existingDomains, err := decodeStoredDomains(item.Domains, item.Domain) + if err != nil { + return fmt.Errorf("existing route %d domains are invalid: %w", item.ID, err) + } + existingSet := make(map[string]struct{}, len(existingDomains)) + for _, existingDomain := range existingDomains { + existingSet[existingDomain] = struct{}{} + } + for _, domain := range domains { + if _, ok := existingSet[domain]; ok { + return fmt.Errorf(errProxyRouteDomainExists, domain) + } + } + } + + return nil +} + +func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) { + if value < 0 { + return 0, fmt.Errorf("%s must be greater than or equal to 0", field) + } + return value, nil +} + +func normalizeProxyRouteCertificateIDs(ctx context.Context, enableHTTPS bool, certID *uint, certIDs []uint) ([]uint, error) { + if !enableHTTPS { + return []uint{}, nil + } + + candidates := make([]uint, 0, len(certIDs)+1) + if certID != nil && *certID != 0 { + candidates = append(candidates, *certID) + } + candidates = append(candidates, certIDs...) + + normalized := make([]uint, 0, len(candidates)) + seen := make(map[uint]struct{}, len(candidates)) + for _, item := range candidates { + if item == 0 { + continue + } + if _, ok := seen[item]; ok { + continue + } + if _, err := lookupTLSCertificateByID(ctx, item); err != nil { + return nil, errors.New(errProxyRouteCertNotFound) + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + if len(normalized) == 0 { + return nil, errors.New(errProxyRouteCertRequired) + } + return normalized, nil +} + +func normalizeProxyRouteDomainCertificateIDs( + ctx context.Context, + domains []string, + enableHTTPS bool, + rawDomainCertIDs []uint, + certID *uint, + certIDs []uint, +) ([]uint, []uint, *uint, error) { + if !enableHTTPS { + return []uint{}, []uint{}, nil, nil + } + + if len(rawDomainCertIDs) > 0 { + if len(rawDomainCertIDs) != len(domains) { + return nil, nil, nil, errors.New(errProxyRouteCertDomainLength) + } + + normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs)) + uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs)) + seen := make(map[uint]struct{}, len(rawDomainCertIDs)) + hasAssignedCertificate := false + for index, item := range rawDomainCertIDs { + if item == 0 { + continue + } + if _, err := lookupTLSCertificateByID(ctx, item); err != nil { + return nil, nil, nil, errors.New(errProxyRouteCertNotFound) + } + normalizedDomainCertIDs[index] = item + hasAssignedCertificate = true + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + uniqueCertIDs = append(uniqueCertIDs, item) + } + if !hasAssignedCertificate { + return nil, nil, nil, errors.New(errProxyRouteCertRequired) + } + + primaryCertID := &uniqueCertIDs[0] + return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil + } + + normalizedCertIDs, err := normalizeProxyRouteCertificateIDs(ctx, enableHTTPS, certID, certIDs) + if err != nil { + return nil, nil, nil, err + } + + switch { + case len(normalizedCertIDs) == 0: + return nil, nil, nil, errors.New(errProxyRouteCertRequired) + case len(normalizedCertIDs) == 1: + domainCertIDs := make([]uint, len(domains)) + for index := range domainCertIDs { + domainCertIDs[index] = normalizedCertIDs[0] + } + primaryCertID := &normalizedCertIDs[0] + return domainCertIDs, normalizedCertIDs, primaryCertID, nil + case len(normalizedCertIDs) == len(domains): + domainCertIDs := make([]uint, len(normalizedCertIDs)) + copy(domainCertIDs, normalizedCertIDs) + primaryCertID := &normalizedCertIDs[0] + return domainCertIDs, normalizedCertIDs, primaryCertID, nil + default: + domainCertIDs, err := deriveDomainCertIDsFromCertificateSet(ctx, domains, normalizedCertIDs) + if err != nil { + return nil, nil, nil, err + } + primaryCertID := &normalizedCertIDs[0] + return domainCertIDs, normalizedCertIDs, primaryCertID, nil + } +} + +func validateProxyRouteDomainCertificateCoverage(ctx context.Context, domains []string, domainCertIDs []uint) error { + if len(domainCertIDs) == 0 { + return nil + } + + domainsByCertID := make(map[uint][]string) + for index, certID := range domainCertIDs { + if certID == 0 { + continue + } + domainsByCertID[certID] = append(domainsByCertID[certID], domains[index]) + } + + for certID, assignedDomains := range domainsByCertID { + certificate, err := lookupTLSCertificateByID(ctx, certID) + if err != nil { + return errors.New(errProxyRouteCertNotFound) + } + if err := validateCertificateCoverage(certificate, assignedDomains); err != nil { + return err + } + } + return nil +} + +func deriveDomainCertIDsFromCertificateSet(ctx context.Context, domains []string, certIDs []uint) ([]uint, error) { + certificates, err := loadTLSCertificates(ctx, certIDs) + if err != nil { + return nil, err + } + + result := make([]uint, len(domains)) + for domainIndex, domain := range domains { + if domainIndex < len(certificates) && + certificates[domainIndex] != nil && + validateCertificateCoverage(certificates[domainIndex], []string{domain}) == nil { + result[domainIndex] = certificates[domainIndex].ID + continue + } + + assigned := uint(0) + for _, certificate := range certificates { + if certificate != nil && + validateCertificateCoverage(certificate, []string{domain}) == nil { + assigned = certificate.ID + break + } + } + if assigned == 0 { + return nil, fmt.Errorf("certificate does not cover domain %s", domain) + } + result[domainIndex] = assigned + } + return result, nil +} + +func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []uint{}, nil + } + + var domainCertIDs []uint + if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil { + return nil, errors.New("domain_cert_ids payload is invalid") + } + if len(domainCertIDs) == 0 { + return []uint{}, nil + } + if domainCount > 0 && len(domainCertIDs) != domainCount { + return nil, errors.New("domain_cert_ids length does not match domains") + } + + normalized := make([]uint, len(domainCertIDs)) + copy(normalized, domainCertIDs) + return normalized, nil +} + +func resolveProxyRouteDomainCertIDs(ctx context.Context, route *model.ProxyRoute, domains []string, certIDs []uint) ([]uint, error) { + domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, len(domains)) + if err != nil { + return nil, err + } + if len(domainCertIDs) > 0 || len(certIDs) == 0 { + return domainCertIDs, nil + } + return deriveDomainCertIDsFromCertificateSet(ctx, domains, certIDs) +} + +func normalizeProxyRouteLimitRate(raw string) (string, error) { + normalized := strings.ToLower(strings.TrimSpace(raw)) + if normalized == "" || normalized == "0" { + return "", nil + } + if !proxyRouteLimitRatePattern.MatchString(normalized) { + return "", errors.New(errProxyRouteLimitRate) + } + if strings.TrimRight(normalized, "km") == "" { + return "", nil + } + return normalized, nil +} + +func hasStructuredOriginInput(input Input) bool { + return (input.OriginID != nil && *input.OriginID != 0) || + strings.TrimSpace(input.OriginScheme) != "" || + strings.TrimSpace(input.OriginAddress) != "" || + strings.TrimSpace(input.OriginPort) != "" || + strings.TrimSpace(input.OriginURI) != "" +} + +func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) { + if hasStructuredOriginInput(input) { + scheme, err := normalizeOriginScheme(input.OriginScheme) + if err != nil { + return "", nil, err + } + port, err := normalizeOriginPort(input.OriginPort) + if err != nil { + return "", nil, err + } + uri, err := normalizeOriginURI(input.OriginURI) + if err != nil { + return "", nil, err + } + if input.OriginID != nil && *input.OriginID != 0 { + origin, err := model.GetOriginByID(ctx, *input.OriginID) + if err != nil { + return "", nil, errors.New(errProxyRouteOriginNotFound) + } + originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri) + if err != nil { + return "", nil, err + } + return originURL, &origin.ID, nil + } + + address := normalizeOriginAddress(input.OriginAddress) + if err := validateOriginAddress(address); err != nil { + return "", nil, err + } + originURL, err := buildOriginURLFromParts(scheme, address, port, uri) + if err != nil { + return "", nil, err + } + origin, err := getOrCreateOriginByAddress(ctx, address) + if err != nil { + return "", nil, err + } + return originURL, &origin.ID, nil + } + + originURL := strings.TrimSpace(input.OriginURL) + if originURL == "" { + return "", nil, errors.New(errProxyRouteOriginEmpty) + } + address, err := extractOriginAddress(originURL) + if err != nil { + return "", nil, err + } + origin, findErr := model.GetOriginByAddress(ctx, address) + if findErr == nil { + return originURL, &origin.ID, nil + } + if !errors.Is(findErr, gorm.ErrRecordNotFound) { + return "", nil, findErr + } + return originURL, nil, nil +} + +func normalizeCustomHeaders(headers []CustomHeaderInput) ([]CustomHeaderInput, error) { + if len(headers) == 0 { + return []CustomHeaderInput{}, nil + } + normalized := make([]CustomHeaderInput, 0, len(headers)) + for _, header := range headers { + key := strings.TrimSpace(header.Key) + value := strings.TrimSpace(header.Value) + if key == "" && value == "" { + continue + } + if key == "" { + return nil, errors.New(errProxyRouteHeaderKeyEmpty) + } + if !proxyHeaderKeyPattern.MatchString(key) { + return nil, errors.New(errProxyRouteHeaderKeyInvalid) + } + if strings.ContainsAny(key, "\r\n") || strings.ContainsAny(value, "\r\n") { + return nil, errors.New(errProxyRouteHeaderNewline) + } + normalized = append(normalized, CustomHeaderInput{Key: key, Value: value}) + } + return normalized, nil +} + +func normalizeUpstreams(originURL string, upstreams []string) ([]string, error) { + candidates := make([]string, 0, len(upstreams)+1) + if strings.TrimSpace(originURL) != "" { + candidates = append(candidates, originURL) + } + candidates = append(candidates, upstreams...) + trimmed := make([]string, 0, len(candidates)) + for _, candidate := range candidates { + item := strings.TrimSpace(candidate) + if item == "" { + continue + } + trimmed = append(trimmed, item) + } + unique := uniqueStrings(trimmed) + normalized := make([]string, 0, len(unique)) + var scheme string + multiUpstream := len(unique) > 1 + for _, item := range unique { + if err := validateOriginURL(item); err != nil { + return nil, err + } + parsed, err := url.ParseRequestURI(item) + if err != nil { + return nil, errors.New(errProxyRouteOriginInvalid) + } + if multiUpstream && parsed.Path != "" && parsed.Path != "/" { + return nil, errors.New(errProxyRouteUpstreamPath) + } + if multiUpstream && parsed.RawQuery != "" { + return nil, errors.New(errProxyRouteUpstreamQuery) + } + if scheme == "" { + scheme = parsed.Scheme + } else if scheme != parsed.Scheme { + return nil, errors.New(errProxyRouteUpstreamScheme) + } + normalized = append(normalized, item) + } + if len(normalized) == 0 { + return nil, errors.New(errProxyRouteUpstreamRequired) + } + return normalized, nil +} + +func decodeStoredCustomHeaders(raw string) ([]CustomHeaderInput, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []CustomHeaderInput{}, nil + } + var headers []CustomHeaderInput + if err := json.Unmarshal([]byte(text), &headers); err != nil { + return nil, errors.New("custom_headers payload is invalid") + } + return normalizeCustomHeaders(headers) +} + +func normalizeCachePolicy(enabled bool, raw string) string { + if !enabled { + return "" + } + policy := strings.TrimSpace(raw) + if policy == "" { + return proxyRouteCachePolicyURL + } + return policy +} + +func normalizeCacheRules(enabled bool, rawPolicy string, rules []string) ([]string, error) { + if !enabled { + return []string{}, nil + } + policy := normalizeCachePolicy(enabled, rawPolicy) + switch policy { + case proxyRouteCachePolicyURL: + return []string{}, nil + case proxyRouteCachePolicySuffix: + return normalizeCacheSuffixRules(rules) + case proxyRouteCachePolicyPathPrefix: + return normalizeCachePathRules(rules, true) + case proxyRouteCachePolicyPathExact: + return normalizeCachePathRules(rules, false) + default: + return nil, errors.New(errProxyRouteCachePolicy) + } +} + +func normalizeCacheSuffixRules(rules []string) ([]string, error) { + normalized := make([]string, 0, len(rules)) + seen := make(map[string]struct{}, len(rules)) + for _, rule := range rules { + item := strings.TrimSpace(strings.TrimPrefix(rule, ".")) + if item == "" { + continue + } + if strings.ContainsAny(item, "/\\ \t\r\n") { + return nil, errors.New(errProxyRouteCacheSuffix) + } + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + if len(normalized) == 0 { + return nil, errors.New(errProxyRouteCacheSuffixReq) + } + return normalized, nil +} + +func normalizeCachePathRules(rules []string, allowPrefix bool) ([]string, error) { + normalized := make([]string, 0, len(rules)) + seen := make(map[string]struct{}, len(rules)) + for _, rule := range rules { + item := strings.TrimSpace(rule) + if item == "" { + continue + } + if !strings.HasPrefix(item, "/") || strings.Contains(item, "://") || strings.ContainsAny(item, " \t\r\n") { + return nil, errors.New(errProxyRouteCachePath) + } + if !allowPrefix && strings.HasSuffix(item, "/") && len(item) > 1 { + item = strings.TrimRight(item, "/") + } + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + if len(normalized) == 0 { + if allowPrefix { + return nil, errors.New(errProxyRouteCachePrefixReq) + } + return nil, errors.New(errProxyRouteCacheExactReq) + } + return normalized, nil +} + +func decodeStoredCacheRules(raw string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []string{}, nil + } + var rules []string + if err := json.Unmarshal([]byte(text), &rules); err != nil { + return nil, errors.New("cache_rules payload is invalid") + } + normalized := make([]string, 0, len(rules)) + for _, rule := range rules { + item := strings.TrimSpace(rule) + if item == "" { + continue + } + normalized = append(normalized, item) + } + return normalized, nil +} + +func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return normalizeUpstreams(fallbackOriginURL, nil) + } + var upstreams []string + if err := json.Unmarshal([]byte(text), &upstreams); err != nil { + return nil, errors.New("upstreams payload is invalid") + } + return normalizeUpstreams(fallbackOriginURL, upstreams) +} + +func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return normalizeProxyRouteDomains([]string{fallbackDomain}) + } + var domains []string + if err := json.Unmarshal([]byte(text), &domains); err != nil { + return nil, errors.New("domains payload is invalid") + } + return normalizeProxyRouteDomains(domains) +} + +func decodeStoredCertIDs(raw string, fallbackCertID *uint) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + if fallbackCertID == nil || *fallbackCertID == 0 { + return []uint{}, nil + } + return []uint{*fallbackCertID}, nil + } + var certIDs []uint + if err := json.Unmarshal([]byte(text), &certIDs); err != nil { + return nil, errors.New("cert_ids payload is invalid") + } + normalized := make([]uint, 0, len(certIDs)) + seen := make(map[uint]struct{}, len(certIDs)) + for _, certID := range certIDs { + if certID == 0 { + continue + } + if _, ok := seen[certID]; ok { + continue + } + seen[certID] = struct{}{} + normalized = append(normalized, certID) + } + if len(normalized) == 0 && fallbackCertID != nil && *fallbackCertID != 0 { + return []uint{*fallbackCertID}, nil + } + return normalized, nil +} + +func validateOriginURL(raw string) error { + if raw == "" { + return errors.New(errProxyRouteOriginEmpty) + } + parsed, err := url.ParseRequestURI(raw) + if err != nil { + return errors.New(errProxyRouteOriginInvalid) + } + if parsed.Scheme != "http" && parsed.Scheme != "https" { + return errors.New(errProxyRouteOriginScheme) + } + if parsed.Host == "" { + return errors.New(errProxyRouteOriginInvalid) + } + return nil +} + +func validateOriginHost(raw string) error { + if raw == "" { + return nil + } + if strings.ContainsAny(raw, "/\\ \t\r\n") || strings.Contains(raw, "://") { + return errors.New(errProxyRouteOriginHostInvalid) + } + parsed, err := url.Parse("//" + raw) + if err != nil || parsed.Host == "" || parsed.Host != raw { + return errors.New(errProxyRouteOriginHostInvalid) + } + if parsed.Hostname() == "" { + return errors.New(errProxyRouteOriginHostInvalid) + } + return nil +} + +func normalizeTunnelNodeID(tunnelNodeID, legacyTunnelID *uint) (*uint, error) { + if tunnelNodeID != nil && *tunnelNodeID != 0 { + return tunnelNodeID, nil + } + if legacyTunnelID != nil && *legacyTunnelID != 0 { + return legacyTunnelID, nil + } + return nil, errors.New(errProxyRouteTunnelNodeReq) +} + +func validateTunnelRouteInput(ctx context.Context, tunnelNodeID *uint, targetAddr, targetProtocol string) error { + if tunnelNodeID == nil || *tunnelNodeID == 0 { + return errors.New(errProxyRouteTunnelNodeReq) + } + tunnelNode, err := lookupTunnelNodeByID(ctx, *tunnelNodeID) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return errors.New(errProxyRouteTunnelNodeMissing) + } + return err + } + if tunnelNode.NodeType != "tunnel_client" { + return errors.New(errProxyRouteTunnelNodeType) + } + if strings.TrimSpace(targetAddr) == "" { + return errors.New(errProxyRouteTunnelAddrReq) + } + switch strings.ToLower(strings.TrimSpace(targetProtocol)) { + case "", "http", "https": + return nil + default: + return errors.New(errProxyRouteTunnelProtocol) + } +} + +func validatePagesRouteInput(ctx context.Context, projectID *uint) error { + if projectID == nil || *projectID == 0 { + return errors.New(errProxyRoutePagesProjectReq) + } + project, err := lookupPagesProjectByID(ctx, *projectID) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return errors.New(errProxyRoutePagesNotFound) + } + return err + } + if !project.Enabled { + return errors.New(errProxyRoutePagesDisabled) + } + if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 { + return errors.New(errProxyRoutePagesNoDeploy) + } + return nil +} + +func normalizeUpstreamType(raw string) string { + switch strings.ToLower(strings.TrimSpace(raw)) { + case "tunnel": + return "tunnel" + case "pages": + return "pages" + default: + return "direct" + } +} + +func normalizeTunnelTargetProtocol(raw string) string { + switch strings.ToLower(strings.TrimSpace(raw)) { + case "https": + return "https" + default: + return "http" + } +} diff --git a/Wavelet/internal/apps/openflare/proxy_route/logics.go b/Wavelet/internal/apps/openflare/proxy_route/logics.go new file mode 100644 index 00000000..dff97981 --- /dev/null +++ b/Wavelet/internal/apps/openflare/proxy_route/logics.go @@ -0,0 +1,430 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package proxy_route + +import ( + "context" + "encoding/json" + "errors" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +// CustomHeaderInput 自定义响应头。 +type CustomHeaderInput struct { + Key string `json:"key"` + Value string `json:"value"` +} + +// Input 代理规则创建/更新请求。 +type Input struct { + SiteName string `json:"site_name"` + Domain string `json:"domain"` + Domains []string `json:"domains"` + OriginID *uint `json:"origin_id"` + OriginURL string `json:"origin_url"` + OriginScheme string `json:"origin_scheme"` + OriginAddress string `json:"origin_address"` + OriginPort string `json:"origin_port"` + OriginURI string `json:"origin_uri"` + OriginHost string `json:"origin_host"` + Upstreams []string `json:"upstreams"` + Enabled bool `json:"enabled"` + EnableHTTPS bool `json:"enable_https"` + CertID *uint `json:"cert_id"` + CertIDs []uint `json:"cert_ids"` + DomainCertIDs []uint `json:"domain_cert_ids"` + RedirectHTTP bool `json:"redirect_http"` + LimitConnPerServer int `json:"limit_conn_per_server"` + LimitConnPerIP int `json:"limit_conn_per_ip"` + LimitRate string `json:"limit_rate"` + CacheEnabled bool `json:"cache_enabled"` + CachePolicy string `json:"cache_policy"` + CacheRules []string `json:"cache_rules"` + CustomHeaders []CustomHeaderInput `json:"custom_headers"` + BasicAuthEnabled bool `json:"basic_auth_enabled"` + BasicAuthUsername string `json:"basic_auth_username"` + BasicAuthPassword string `json:"basic_auth_password"` + Remark string `json:"remark"` + UpstreamType string `json:"upstream_type"` + TunnelNodeID *uint `json:"tunnel_node_id"` + TunnelID *uint `json:"tunnel_id"` + TunnelTargetAddr string `json:"tunnel_target_addr"` + TunnelTargetProtocol string `json:"tunnel_target_protocol"` + PagesProjectID *uint `json:"pages_project_id"` +} + +// View 代理规则视图。 +type View struct { + ID uint `json:"id"` + SiteName string `json:"site_name"` + Domain string `json:"domain"` + Domains []string `json:"domains"` + PrimaryDomain string `json:"primary_domain"` + DomainCount int `json:"domain_count"` + OriginID *uint `json:"origin_id"` + OriginURL string `json:"origin_url"` + OriginHost string `json:"origin_host"` + Upstreams string `json:"upstreams"` + UpstreamList []string `json:"upstream_list"` + Enabled bool `json:"enabled"` + EnableHTTPS bool `json:"enable_https"` + CertID *uint `json:"cert_id"` + CertIDs []uint `json:"cert_ids"` + DomainCertIDs []uint `json:"domain_cert_ids"` + RedirectHTTP bool `json:"redirect_http"` + LimitConnPerServer int `json:"limit_conn_per_server"` + LimitConnPerIP int `json:"limit_conn_per_ip"` + LimitRate string `json:"limit_rate"` + CacheEnabled bool `json:"cache_enabled"` + CachePolicy string `json:"cache_policy"` + CacheRules string `json:"cache_rules"` + CacheRuleList []string `json:"cache_rule_list"` + CustomHeaders string `json:"custom_headers"` + CustomHeaderList []CustomHeaderInput `json:"custom_header_list"` + BasicAuthEnabled bool `json:"basic_auth_enabled"` + BasicAuthUsername string `json:"basic_auth_username"` + BasicAuthPassword string `json:"basic_auth_password"` + Remark string `json:"remark"` + UpstreamType string `json:"upstream_type"` + TunnelNodeID *uint `json:"tunnel_node_id"` + TunnelID *uint `json:"tunnel_id"` + TunnelTargetAddr string `json:"tunnel_target_addr"` + TunnelTargetProtocol string `json:"tunnel_target_protocol"` + PagesProjectID *uint `json:"pages_project_id"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// ListProxyRoutes 列出全部代理规则。 +func ListProxyRoutes(ctx context.Context) ([]*View, error) { + routes, err := model.ListProxyRoutes(ctx) + if err != nil { + return nil, err + } + return buildProxyRouteViews(ctx, routes) +} + +// GetProxyRoute 获取代理规则详情。 +func GetProxyRoute(ctx context.Context, id uint) (*View, error) { + route, err := model.GetProxyRouteByID(ctx, id) + if err != nil { + return nil, err + } + return buildProxyRouteView(ctx, route) +} + +// CreateProxyRoute 创建代理规则。 +func CreateProxyRoute(ctx context.Context, input Input) (*View, error) { + route, err := buildProxyRoute(ctx, nil, input) + if err != nil { + return nil, err + } + if err = model.CreateProxyRouteRecord(ctx, route); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errProxyRouteIdentityExists) + } + return nil, err + } + return buildProxyRouteView(ctx, route) +} + +// UpdateProxyRoute 更新代理规则。 +func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error) { + route, err := model.GetProxyRouteByID(ctx, id) + if err != nil { + return nil, err + } + route, err = buildProxyRoute(ctx, route, input) + if err != nil { + return nil, err + } + if err = model.UpdateProxyRouteRecord(ctx, route); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errProxyRouteIdentityExists) + } + return nil, err + } + return buildProxyRouteView(ctx, route) +} + +// DeleteProxyRoute 删除代理规则。 +func DeleteProxyRoute(ctx context.Context, id uint) error { + if _, err := model.GetProxyRouteByID(ctx, id); err != nil { + return err + } + return model.DeleteProxyRouteRecord(ctx, id) +} + +func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, error) { + domains, err := normalizeProxyRouteDomainsInput(route, input.Domain, input.Domains) + if err != nil { + return nil, err + } + domain := domains[0] + siteName := normalizeProxyRouteSiteNameInput(route, input.SiteName, domain) + + upstreamType := normalizeUpstreamType(input.UpstreamType) + var originURL string + var originID *uint + var upstreams []string + + if upstreamType == "tunnel" { + originURL = "http://127.0.0.1" + upstreams = []string{originURL} + } else if upstreamType == "pages" { + if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil { + return nil, err + } + originURL = "http://127.0.0.1" + upstreams = []string{originURL} + } else { + originURL, originID, err = resolveProxyRoutePrimaryOrigin(ctx, input) + if err != nil { + return nil, err + } + upstreams, err = normalizeUpstreams(originURL, input.Upstreams) + if err != nil { + return nil, err + } + } + originHost := strings.TrimSpace(input.OriginHost) + remark := strings.TrimSpace(input.Remark) + cachePolicy := strings.TrimSpace(input.CachePolicy) + cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules) + if err != nil { + return nil, err + } + customHeaders, err := normalizeCustomHeaders(input.CustomHeaders) + if err != nil { + return nil, err + } + limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server") + if err != nil { + return nil, err + } + limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip") + if err != nil { + return nil, err + } + limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate) + if err != nil { + return nil, err + } + + cacheRulesJSON, err := json.Marshal(cacheRules) + if err != nil { + return nil, err + } + upstreamsJSON, err := json.Marshal(upstreams) + if err != nil { + return nil, err + } + customHeadersJSON, err := json.Marshal(customHeaders) + if err != nil { + return nil, err + } + + if !input.EnableHTTPS { + input.RedirectHTTP = false + input.CertID = nil + input.CertIDs = nil + input.DomainCertIDs = nil + } + domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs( + ctx, + domains, + input.EnableHTTPS, + input.DomainCertIDs, + input.CertID, + input.CertIDs, + ) + if err != nil { + return nil, err + } + if err := validateProxyRouteDomainCertificateCoverage(ctx, domains, domainCertIDs); err != nil { + return nil, err + } + certIDsJSON, err := json.Marshal(certIDs) + if err != nil { + return nil, err + } + domainCertIDsJSON, err := json.Marshal(domainCertIDs) + if err != nil { + return nil, err + } + domainsJSON, err := json.Marshal(domains) + if err != nil { + return nil, err + } + + if err := validateProxyRouteSiteName(siteName); err != nil { + return nil, err + } + if err := validateProxyRouteIdentityUniqueness(ctx, route, siteName, domains); err != nil { + return nil, err + } + if err := validateOriginHost(originHost); err != nil { + return nil, err + } + input.DomainCertIDs = domainCertIDs + input.CertIDs = certIDs + input.CertID = primaryCertID + if input.RedirectHTTP && !input.EnableHTTPS { + return nil, errors.New(errProxyRouteRedirectHTTP) + } + + if input.BasicAuthEnabled { + input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername) + input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword) + if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" { + return nil, errors.New(errProxyRouteBasicAuth) + } + } else { + input.BasicAuthUsername = "" + input.BasicAuthPassword = "" + } + + if route == nil { + route = &model.ProxyRoute{} + } + route.SiteName = siteName + route.Domain = domain + route.Domains = string(domainsJSON) + route.OriginID = originID + route.OriginURL = upstreams[0] + route.OriginHost = originHost + route.Upstreams = string(upstreamsJSON) + route.Enabled = input.Enabled + route.EnableHTTPS = input.EnableHTTPS + route.CertID = input.CertID + route.CertIDs = string(certIDsJSON) + route.DomainCertIDs = string(domainCertIDsJSON) + route.RedirectHTTP = input.RedirectHTTP + route.LimitConnPerServer = limitConnPerServer + route.LimitConnPerIP = limitConnPerIP + route.LimitRate = limitRate + route.CacheEnabled = input.CacheEnabled + route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy) + route.CacheRules = string(cacheRulesJSON) + route.CustomHeaders = string(customHeadersJSON) + route.BasicAuthEnabled = input.BasicAuthEnabled + route.BasicAuthUsername = input.BasicAuthUsername + route.BasicAuthPassword = input.BasicAuthPassword + route.Remark = remark + route.UpstreamType = upstreamType + if upstreamType == "tunnel" { + tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID) + if err != nil { + return nil, err + } + if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil { + return nil, err + } + route.TunnelNodeID = tunnelNodeID + route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr) + route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol) + route.PagesProjectID = nil + } else if upstreamType == "pages" { + route.TunnelNodeID = nil + route.TunnelTargetAddr = "" + route.TunnelTargetProtocol = "" + route.PagesProjectID = input.PagesProjectID + } else { + route.TunnelNodeID = nil + route.TunnelTargetAddr = "" + route.TunnelTargetProtocol = "" + route.PagesProjectID = nil + } + return route, nil +} + +func buildProxyRouteViews(ctx context.Context, routes []*model.ProxyRoute) ([]*View, error) { + views := make([]*View, 0, len(routes)) + for _, route := range routes { + view, err := buildProxyRouteView(ctx, route) + if err != nil { + return nil, err + } + views = append(views, view) + } + return views, nil +} + +func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, error) { + if route == nil { + return nil, errors.New("proxy route is nil") + } + domains, err := decodeStoredDomains(route.Domains, route.Domain) + if err != nil { + return nil, err + } + upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL) + if err != nil { + return nil, err + } + cacheRules, err := decodeStoredCacheRules(route.CacheRules) + if err != nil { + return nil, err + } + customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders) + if err != nil { + return nil, err + } + certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) + if err != nil { + return nil, err + } + domainCertIDs, err := resolveProxyRouteDomainCertIDs(ctx, route, domains, certIDs) + if err != nil { + return nil, err + } + var certID *uint + if len(certIDs) > 0 { + certID = &certIDs[0] + } + primaryDomain := domains[0] + return &View{ + ID: route.ID, + SiteName: normalizeProxyRouteSiteNameInput(route, route.SiteName, primaryDomain), + Domain: primaryDomain, + Domains: domains, + PrimaryDomain: primaryDomain, + DomainCount: len(domains), + OriginID: route.OriginID, + OriginURL: route.OriginURL, + OriginHost: route.OriginHost, + Upstreams: route.Upstreams, + UpstreamList: upstreams, + Enabled: route.Enabled, + EnableHTTPS: route.EnableHTTPS, + CertID: certID, + CertIDs: certIDs, + DomainCertIDs: domainCertIDs, + RedirectHTTP: route.RedirectHTTP, + LimitConnPerServer: route.LimitConnPerServer, + LimitConnPerIP: route.LimitConnPerIP, + LimitRate: route.LimitRate, + CacheEnabled: route.CacheEnabled, + CachePolicy: route.CachePolicy, + CacheRules: route.CacheRules, + CacheRuleList: cacheRules, + CustomHeaders: route.CustomHeaders, + CustomHeaderList: customHeaders, + BasicAuthEnabled: route.BasicAuthEnabled, + BasicAuthUsername: route.BasicAuthUsername, + BasicAuthPassword: route.BasicAuthPassword, + Remark: route.Remark, + UpstreamType: route.UpstreamType, + TunnelNodeID: route.TunnelNodeID, + TunnelID: route.TunnelNodeID, + TunnelTargetAddr: route.TunnelTargetAddr, + TunnelTargetProtocol: route.TunnelTargetProtocol, + PagesProjectID: route.PagesProjectID, + CreatedAt: route.CreatedAt, + UpdatedAt: route.UpdatedAt, + }, nil +} diff --git a/Wavelet/internal/apps/openflare/proxy_route/logics_test.go b/Wavelet/internal/apps/openflare/proxy_route/logics_test.go new file mode 100644 index 00000000..78874407 --- /dev/null +++ b/Wavelet/internal/apps/openflare/proxy_route/logics_test.go @@ -0,0 +1,88 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package proxy_route + +import ( + "context" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupProxyRouteTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Origin{})) + + db.SetDB(sqliteDB) + return func() { + db.SetDB(nil) + } +} + +func TestCreateProxyRoute(t *testing.T) { + cleanup := setupProxyRouteTestDB(t) + defer cleanup() + ctx := context.Background() + + view, err := CreateProxyRoute(ctx, Input{ + SiteName: "example-site", + Domain: "example.com", + OriginURL: "http://origin.example.com:8080", + Enabled: true, + }) + require.NoError(t, err) + assert.NotZero(t, view.ID) + assert.Equal(t, "example-site", view.SiteName) + assert.Equal(t, "example.com", view.Domain) + assert.Equal(t, []string{"example.com"}, view.Domains) + assert.Equal(t, "http://origin.example.com:8080", view.OriginURL) + assert.Equal(t, []string{"http://origin.example.com:8080"}, view.UpstreamList) + assert.True(t, view.Enabled) + + _, err = CreateProxyRoute(ctx, Input{ + SiteName: "duplicate-site", + Domain: "example.com", + OriginURL: "http://origin.example.com:8080", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "already exists") +} + +func TestListProxyRoutes(t *testing.T) { + cleanup := setupProxyRouteTestDB(t) + defer cleanup() + ctx := context.Background() + + first, err := CreateProxyRoute(ctx, Input{ + SiteName: "first-site", + Domain: "first.example.com", + OriginURL: "http://origin-a.internal:80", + }) + require.NoError(t, err) + + second, err := CreateProxyRoute(ctx, Input{ + SiteName: "second-site", + Domain: "second.example.com", + OriginURL: "http://origin-b.internal:80", + }) + require.NoError(t, err) + + routes, err := ListProxyRoutes(ctx) + require.NoError(t, err) + require.Len(t, routes, 2) + assert.Equal(t, second.ID, routes[0].ID) + assert.Equal(t, first.ID, routes[1].ID) + assert.Equal(t, "second.example.com", routes[0].Domain) + assert.Equal(t, "first.example.com", routes[1].Domain) +} diff --git a/Wavelet/internal/apps/openflare/proxy_route/routers.go b/Wavelet/internal/apps/openflare/proxy_route/routers.go new file mode 100644 index 00000000..9c7cff00 --- /dev/null +++ b/Wavelet/internal/apps/openflare/proxy_route/routers.go @@ -0,0 +1,88 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package proxy_route + +import ( + "errors" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +func handleLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, errProxyRouteNotFound) + return true + } + compat.Fail(c, err.Error()) + return true +} + +// GetProxyRoutes 列出全部代理规则。 +func GetProxyRoutes(c *gin.Context) { + routes, err := ListProxyRoutes(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, routes) +} + +// GetProxyRouteHandler 获取代理规则详情。 +func GetProxyRouteHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + route, err := GetProxyRoute(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, route) +} + +// CreateProxyRouteHandler 创建代理规则。 +func CreateProxyRouteHandler(c *gin.Context) { + var input Input + if !compat.BindJSON(c, &input) { + return + } + route, err := CreateProxyRoute(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, route) +} + +// UpdateProxyRouteHandler 更新代理规则。 +func UpdateProxyRouteHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input Input + if !compat.BindJSON(c, &input) { + return + } + route, err := UpdateProxyRoute(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, route) +} + +// DeleteProxyRouteHandler 删除代理规则。 +func DeleteProxyRouteHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteProxyRoute(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OK(c, nil) +} diff --git a/Wavelet/internal/apps/openflare/relay/errs.go b/Wavelet/internal/apps/openflare/relay/errs.go new file mode 100644 index 00000000..3786f403 --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/errs.go @@ -0,0 +1,9 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +const ( + errAgentTokenInvalid = "无权进行此操作,Agent Token 无效" + errRelayNodeTypeMismatch = "此节点不是 TunnelRelay 类型" +) diff --git a/Wavelet/internal/apps/openflare/relay/helpers.go b/Wavelet/internal/apps/openflare/relay/helpers.go new file mode 100644 index 00000000..cb276640 --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/helpers.go @@ -0,0 +1,112 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +import ( + "net" + "strings" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +func normalizeRelayStatus(status string) string { + switch strings.ToLower(strings.TrimSpace(status)) { + case "healthy": + return "healthy" + case "unhealthy": + return "unhealthy" + default: + return "unknown" + } +} + +func normalizeReleaseChannel(channel string) string { + if strings.ToLower(strings.TrimSpace(channel)) == "preview" { + return "preview" + } + return "stable" +} + +func resolveReportedNodeIP(reportedIP string, remoteAddr string) string { + reported := normalizeNodeIP(reportedIP) + remote := normalizeRemoteAddr(remoteAddr) + if reported == "" { + return remote + } + if isPublicNodeIP(reported) { + return reported + } + if isPublicNodeIP(remote) { + return remote + } + return reported +} + +func normalizeNodeIP(raw string) string { + raw = strings.TrimSpace(raw) + if raw == "" { + return "" + } + if host, _, err := net.SplitHostPort(raw); err == nil { + raw = host + } + raw = strings.Trim(raw, "[]") + return raw +} + +func normalizeRemoteAddr(remoteAddr string) string { + remoteAddr = strings.TrimSpace(remoteAddr) + if remoteAddr == "" { + return "" + } + host, _, err := net.SplitHostPort(remoteAddr) + if err != nil { + return normalizeNodeIP(remoteAddr) + } + return normalizeNodeIP(host) +} + +func isPublicNodeIP(raw string) bool { + ip := net.ParseIP(strings.TrimSpace(raw)) + if ip == nil { + return false + } + if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() { + return false + } + return true +} + +func buildRelayConfig(node *model.OpenFlareNode) *Config { + if node == nil { + return nil + } + return &Config{ + BindPort: node.RelayBindPort, + VhostHTTPPort: node.RelayVhostHTTPPort, + AuthToken: node.RelayAuthToken, + LogLevel: "info", + WebServerEnabled: node.RelayWebServerEnabled, + } +} + +// BuildSettings returns runtime settings shared by relay and flared clients. +func BuildSettings(node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *Settings { + autoUpdate := false + if node != nil { + autoUpdate = node.AutoUpdateEnabled + } + if strings.TrimSpace(updateChannel) == "" { + updateChannel = "stable" + } + return &Settings{ + HeartbeatInterval: model.AgentHeartbeatInterval, + WebsocketUpgradeEnabled: model.AgentWebsocketUpgradeEnabled, + AutoUpdate: autoUpdate, + UpdateRepo: model.AgentUpdateRepo, + UpdateNow: updateNow, + UpdateChannel: updateChannel, + UpdateTag: strings.TrimSpace(updateTag), + } +} diff --git a/Wavelet/internal/apps/openflare/relay/logics.go b/Wavelet/internal/apps/openflare/relay/logics.go new file mode 100644 index 00000000..ce4fec7d --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/logics.go @@ -0,0 +1,117 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +import ( + "context" + "fmt" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const nodeStatusOnline = "online" + +// HeartbeatPayload is sent by OpenFlareRelay on each heartbeat. +type HeartbeatPayload struct { + Version string `json:"version"` + ExtVersion string `json:"frp_version"` + RelayStatus string `json:"relay_status"` + Name string `json:"name"` + IP string `json:"ip"` +} + +// Config is the frps configuration sent to the relay. +type Config struct { + BindPort int `json:"bind_port"` + VhostHTTPPort int `json:"vhost_http_port"` + AuthToken string `json:"auth_token"` + LogLevel string `json:"log_level"` + WebServerEnabled bool `json:"web_server_enabled"` +} + +// Settings contains runtime settings for relay and flared clients. +type Settings struct { + HeartbeatInterval int `json:"heartbeat_interval"` + WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"` + AutoUpdate bool `json:"auto_update"` + UpdateRepo string `json:"update_repo"` + UpdateNow bool `json:"update_now"` + UpdateChannel string `json:"update_channel"` + UpdateTag string `json:"update_tag"` +} + +// HeartbeatResponse is returned from a relay heartbeat. +type HeartbeatResponse struct { + RelayConfig *Config `json:"relay_config"` + RelaySettings *Settings `json:"relay_settings"` +} + +// Heartbeat processes a relay heartbeat, updates node status, and returns config. +func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) { + if node == nil { + return nil, fmt.Errorf("relay node is nil") + } + + payload.Version = strings.TrimSpace(payload.Version) + payload.ExtVersion = strings.TrimSpace(payload.ExtVersion) + payload.RelayStatus = normalizeRelayStatus(payload.RelayStatus) + payload.Name = strings.TrimSpace(payload.Name) + payload.IP = strings.TrimSpace(payload.IP) + + previous := *node + updateNow := node.UpdateRequested + updateChannel := normalizeReleaseChannel(node.UpdateChannel) + updateTag := strings.TrimSpace(node.UpdateTag) + + now := time.Now().UTC() + changes := map[string]any{ + "version": payload.Version, + "ext_version": payload.ExtVersion, + "relay_status": payload.RelayStatus, + "last_seen_at": now, + "status": nodeStatusOnline, + "update_requested": false, + "update_channel": "stable", + "update_tag": "", + } + if payload.Name != "" && strings.TrimSpace(node.Name) == "" { + changes["name"] = payload.Name + node.Name = payload.Name + } + if payload.IP != "" && !node.IPManualOverride { + changes["ip"] = payload.IP + node.IP = payload.IP + } + if !previous.UpdateRequested { + delete(changes, "update_requested") + } + if previous.UpdateChannel == "stable" { + delete(changes, "update_channel") + } + if previous.UpdateTag == "" { + delete(changes, "update_tag") + } + + node.Version = payload.Version + node.ExtVersion = payload.ExtVersion + node.RelayStatus = payload.RelayStatus + node.UpdateRequested = false + node.UpdateChannel = "stable" + node.UpdateTag = "" + lastSeen := now + node.LastSeenAt = &lastSeen + node.Status = nodeStatusOnline + + if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil { + return nil, fmt.Errorf("update relay heartbeat: %w", err) + } + + return &HeartbeatResponse{ + RelayConfig: buildRelayConfig(node), + RelaySettings: BuildSettings(node, updateNow, updateChannel, updateTag), + }, nil +} diff --git a/Wavelet/internal/apps/openflare/relay/middleware.go b/Wavelet/internal/apps/openflare/relay/middleware.go new file mode 100644 index 00000000..3691540e --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/middleware.go @@ -0,0 +1,51 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +import ( + "context" + "errors" + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +const ctxRelayNodeKey = "relay_node" + +// RelayAuth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type. +func RelayAuth() gin.HandlerFunc { + return func(c *gin.Context) { + token := strings.TrimSpace(c.GetHeader("X-Agent-Token")) + node, err := authenticateAccessToken(c.Request.Context(), token) + if err != nil { + compat.Unauthorized(c, errAgentTokenInvalid) + c.Abort() + return + } + if node.NodeType != "tunnel_relay" { + compat.Forbidden(c, errRelayNodeTypeMismatch) + c.Abort() + return + } + c.Set(ctxRelayNodeKey, node) + c.Next() + } +} + +func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) { + if token == "" { + return nil, errors.New("missing agent token") + } + node, err := model.GetOpenFlareNodeByAccessToken(ctx, token) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, errors.New("invalid agent token") + } + return nil, err + } + return node, nil +} diff --git a/Wavelet/internal/apps/openflare/relay/middleware_test.go b/Wavelet/internal/apps/openflare/relay/middleware_test.go new file mode 100644 index 00000000..492b5565 --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/middleware_test.go @@ -0,0 +1,106 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupRelayMiddlewareTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) + db.SetDB(sqliteDB) + + return func() { + db.SetDB(nil) + } +} + +func seedRelayNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareNode { + t.Helper() + ctx := context.Background() + node := &model.OpenFlareNode{ + NodeID: "relay-test-node", + Name: "relay-test", + Status: "pending", + NodeType: nodeType, + AccessToken: accessToken, + } + require.NoError(t, model.CreateOpenFlareNode(ctx, node)) + return node +} + +func TestRelayAuthMissingToken(t *testing.T) { + cleanup := setupRelayMiddlewareTestDB(t) + defer cleanup() + + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) { + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/relay/test", nil) + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusUnauthorized, rec.Code) +} + +func TestRelayAuthRejectsWrongNodeType(t *testing.T) { + cleanup := setupRelayMiddlewareTestDB(t) + defer cleanup() + seedRelayNode(t, "edge_node", "edge-token-relay") + + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) { + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/relay/test", nil) + req.Header.Set("X-Agent-Token", "edge-token-relay") + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusForbidden, rec.Code) +} + +func TestRelayAuthAcceptsTunnelRelay(t *testing.T) { + cleanup := setupRelayMiddlewareTestDB(t) + defer cleanup() + node := seedRelayNode(t, "tunnel_relay", "relay-token-valid") + + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) { + authNode, ok := c.Get(ctxRelayNodeKey) + require.True(t, ok) + assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID) + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/relay/test", nil) + req.Header.Set("X-Agent-Token", "relay-token-valid") + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + + assert.Equal(t, http.StatusOK, rec.Code) +} diff --git a/Wavelet/internal/apps/openflare/relay/routers.go b/Wavelet/internal/apps/openflare/relay/routers.go new file mode 100644 index 00000000..f1b3f13a --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/routers.go @@ -0,0 +1,45 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/gin-gonic/gin" +) + +// PostHeartbeat handles POST /relay/heartbeat. +func PostHeartbeat(c *gin.Context) { + var payload HeartbeatPayload + if !compat.BindJSON(c, &payload) { + return + } + payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr) + + authNode, ok := c.Get(ctxRelayNodeKey) + if !ok { + compat.Unauthorized(c, errAgentTokenInvalid) + return + } + node := authNode.(*model.OpenFlareNode) + + result, err := Heartbeat(c.Request.Context(), node, payload) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, result) +} + +// GetWebSocket handles GET /relay/ws. +func GetWebSocket(c *gin.Context) { + authNode, ok := c.Get(ctxRelayNodeKey) + if !ok { + compat.Unauthorized(c, errAgentTokenInvalid) + return + } + node := authNode.(*model.OpenFlareNode) + ofws.ServeRelay(c, node.NodeID) +} diff --git a/Wavelet/internal/apps/openflare/tls/errs.go b/Wavelet/internal/apps/openflare/tls/errs.go new file mode 100644 index 00000000..37ac4304 --- /dev/null +++ b/Wavelet/internal/apps/openflare/tls/errs.go @@ -0,0 +1,28 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package tls + +const ( + errCertificateNameRequired = "certificate name cannot be empty" + errCertificateNameExists = "certificate name already exists" + errCertificateContentRequired = "certificate content and key content cannot be empty" + errCertificateContentInvalid = "certificate or key format is invalid" + errCertificateDeleteReferenced = "certificate is still referenced by proxy routes" + errCertificateOnlyACME = "only acme certificates can be updated via this endpoint" + errCertificateOnlyUploadConvert = "only uploaded certificates can be converted to acme" + errCertificateAlreadyApplying = "certificate is already applying" + errCertificateOnlyACMERenew = "only acme certificates can be renewed" + errCertificateFilesRequired = "certificate file and key file cannot be empty" + errCertificatePEMInvalid = "证书 PEM 内容不合法" + + errManagedDomainRequired = "域名不能为空" + errManagedDomainInvalid = "域名格式不合法" + errManagedDomainWildcardInvalid = "通配符域名仅支持 *.example.com 格式" + errManagedDomainExists = "域名已存在" + errManagedDomainCertNotFound = "所选证书不存在" + + errDNSAccountInUse = "该 DNS 账号已被证书使用,无法删除" + + errACMENotImplemented = "ACME certificate obtain is not implemented yet" +) diff --git a/Wavelet/internal/apps/openflare/tls/helpers.go b/Wavelet/internal/apps/openflare/tls/helpers.go new file mode 100644 index 00000000..98a7f056 --- /dev/null +++ b/Wavelet/internal/apps/openflare/tls/helpers.go @@ -0,0 +1,62 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package tls + +import ( + "crypto/x509" + "encoding/json" + "encoding/pem" + "errors" + "fmt" + "io" + "mime/multipart" + "strings" +) + +func parseLeafCertificate(certPEM string) (*x509.Certificate, error) { + certPEMBlock, _ := pem.Decode([]byte(certPEM)) + if certPEMBlock == nil { + return nil, errors.New(errCertificatePEMInvalid) + } + leaf, err := x509.ParseCertificate(certPEMBlock.Bytes) + if err != nil { + return nil, err + } + return leaf, nil +} + +func readMultipartFile(fileHeader *multipart.FileHeader) (string, error) { + file, err := fileHeader.Open() + if err != nil { + return "", err + } + defer file.Close() + data, err := io.ReadAll(file) + if err != nil { + return "", err + } + return string(data), nil +} + +func isUniqueConstraintError(err error) bool { + if err == nil { + return false + } + return strings.Contains(strings.ToLower(err.Error()), "unique") +} + +func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) { + text := strings.TrimSpace(raw) + if text == "" { + return nil, nil + } + var domainCertIDs []uint + if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil { + return nil, err + } + if domainCount > 0 && len(domainCertIDs) != domainCount { + return nil, fmt.Errorf("domain_cert_ids length mismatch") + } + return domainCertIDs, nil +} diff --git a/Wavelet/internal/apps/openflare/tls/logics.go b/Wavelet/internal/apps/openflare/tls/logics.go new file mode 100644 index 00000000..ea8209b9 --- /dev/null +++ b/Wavelet/internal/apps/openflare/tls/logics.go @@ -0,0 +1,462 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package tls + +import ( + "context" + "crypto/tls" + "encoding/json" + "errors" + "fmt" + "mime/multipart" + "strings" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +// CertificateInput TLS 证书创建/更新请求。 +type CertificateInput struct { + Name string `json:"name"` + CertPEM string `json:"cert_pem"` + KeyPEM string `json:"key_pem"` + Remark string `json:"remark"` +} + +// CertificateContent TLS 证书 PEM 内容(仅 /content 端点返回)。 +type CertificateContent struct { + ID uint `json:"id"` + Name string `json:"name"` + CertPEM string `json:"cert_pem"` + KeyPEM string `json:"key_pem"` + Remark string `json:"remark"` + Provider string `json:"provider"` + AcmeAccountID uint `json:"acme_account_id"` + DnsAccountID uint `json:"dns_account_id"` + KeyAlgorithm string `json:"key_algorithm"` + AutoRenew bool `json:"auto_renew"` + PrimaryDomain string `json:"primary_domain"` + OtherDomains string `json:"other_domains"` + DisableCNAME bool `json:"disable_cname"` + SkipDNS bool `json:"skip_dns"` + DNS1 string `json:"dns1"` + DNS2 string `json:"dns2"` + ApplyStatus string `json:"apply_status"` + ApplyMessage string `json:"apply_message"` +} + +// ApplyInput ACME 证书申请/更新请求。 +type ApplyInput struct { + Name string `json:"name"` + Remark string `json:"remark"` + AcmeAccountID uint `json:"acme_account_id"` + DnsAccountID uint `json:"dns_account_id"` + KeyAlgorithm string `json:"key_algorithm"` + AutoRenew bool `json:"auto_renew"` + PrimaryDomain string `json:"primary_domain"` + OtherDomains string `json:"other_domains"` + DisableCNAME bool `json:"disable_cname"` + SkipDNS bool `json:"skip_dns"` + DNS1 string `json:"dns1"` + DNS2 string `json:"dns2"` +} + +// DNSAccountInput DNS 账号创建/更新请求。 +type DNSAccountInput struct { + Name string `json:"name"` + Type string `json:"type"` + Authorization string `json:"authorization"` +} + +// ListCertificates 列出全部证书(不含 PEM)。 +func ListCertificates(ctx context.Context) ([]model.TLSCertificate, error) { + return model.ListTLSCertificates(ctx) +} + +// GetCertificate 获取证书详情(不含 PEM)。 +func GetCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) { + return model.GetTLSCertificateByID(ctx, id) +} + +// GetCertificateContent 获取证书 PEM 内容。 +func GetCertificateContent(ctx context.Context, id uint) (*CertificateContent, error) { + certificate, err := model.GetTLSCertificateByID(ctx, id) + if err != nil { + return nil, err + } + keyPEM, err := openSensitive(certificate.KeyPEM) + if err != nil { + return nil, err + } + return &CertificateContent{ + ID: certificate.ID, + Name: certificate.Name, + CertPEM: certificate.CertPEM, + KeyPEM: keyPEM, + Remark: certificate.Remark, + Provider: certificate.Provider, + AcmeAccountID: certificate.AcmeAccountID, + DnsAccountID: certificate.DnsAccountID, + KeyAlgorithm: certificate.KeyAlgorithm, + AutoRenew: certificate.AutoRenew, + PrimaryDomain: certificate.PrimaryDomain, + OtherDomains: certificate.OtherDomains, + DisableCNAME: certificate.DisableCNAME, + SkipDNS: certificate.SkipDNS, + DNS1: certificate.DNS1, + DNS2: certificate.DNS2, + ApplyStatus: certificate.ApplyStatus, + ApplyMessage: certificate.ApplyMessage, + }, nil +} + +// CreateCertificate 从 PEM 创建证书。 +func CreateCertificate(ctx context.Context, input CertificateInput) (*model.TLSCertificate, error) { + certificate, err := buildCertificate(ctx, nil, input) + if err != nil { + return nil, err + } + if err = model.CreateTLSCertificateRecord(ctx, certificate); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errCertificateNameExists) + } + return nil, err + } + return sanitizeCertificateForResponse(certificate), nil +} + +// CreateCertificateFromFiles 从上传文件创建证书。 +func CreateCertificateFromFiles(ctx context.Context, name string, certFile *multipart.FileHeader, keyFile *multipart.FileHeader, remark string) (*model.TLSCertificate, error) { + if certFile == nil || keyFile == nil { + return nil, errors.New(errCertificateFilesRequired) + } + certContent, err := readMultipartFile(certFile) + if err != nil { + return nil, err + } + keyContent, err := readMultipartFile(keyFile) + if err != nil { + return nil, err + } + return CreateCertificate(ctx, CertificateInput{ + Name: name, + CertPEM: certContent, + KeyPEM: keyContent, + Remark: remark, + }) +} + +// UpdateCertificate 更新上传证书。 +func UpdateCertificate(ctx context.Context, id uint, input CertificateInput) (*model.TLSCertificate, error) { + existing, err := model.GetTLSCertificateByID(ctx, id) + if err != nil { + return nil, err + } + certificate, err := buildCertificate(ctx, existing, input) + if err != nil { + return nil, err + } + if err = model.SaveTLSCertificate(ctx, certificate); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errCertificateNameExists) + } + return nil, err + } + return sanitizeCertificateForResponse(certificate), nil +} + +// DeleteCertificate 删除证书。 +func DeleteCertificate(ctx context.Context, id uint) error { + if err := ensureCertificateNotReferenced(ctx, id); err != nil { + return err + } + if _, err := model.GetTLSCertificateByID(ctx, id); err != nil { + return err + } + return model.DeleteTLSCertificateRecord(ctx, id) +} + +// ApplyCertificate 申请 ACME 证书(当前为占位实现)。 +func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertificate, error) { + cert := &model.TLSCertificate{ + Provider: "acme", + CertPEM: " ", + KeyPEM: " ", + } + fillAcmeCertificateFields(cert, input) + if cert.Name == "" { + return nil, errors.New(errCertificateNameRequired) + } + if err := model.CreateTLSCertificateRecord(ctx, cert); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errCertificateNameExists) + } + return nil, err + } + return markACMEStubFailure(ctx, cert) +} + +// UpdateACMECertificate 更新 ACME 证书配置(当前为占位实现)。 +func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) { + cert, err := model.GetTLSCertificateByID(ctx, id) + if err != nil { + return nil, err + } + if cert.Provider != "acme" { + return nil, errors.New(errCertificateOnlyACME) + } + fillAcmeCertificateFields(cert, input) + if cert.Name == "" { + return nil, errors.New(errCertificateNameRequired) + } + if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errCertificateNameExists) + } + return nil, err + } + return markACMEStubFailure(ctx, cert) +} + +// ConvertCertificateToACME 将上传证书转为 ACME 管理(当前为占位实现)。 +func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) { + cert, err := model.GetTLSCertificateByID(ctx, id) + if err != nil { + return nil, err + } + if cert.Provider != "upload" { + return nil, errors.New(errCertificateOnlyUploadConvert) + } + if cert.ApplyStatus == "applying" { + return nil, errors.New(errCertificateAlreadyApplying) + } + fillAcmeCertificateFields(cert, input) + if cert.Name == "" { + return nil, errors.New(errCertificateNameRequired) + } + cert.ApplyMessage = "" + if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errCertificateNameExists) + } + return nil, err + } + return markACMEStubFailure(ctx, cert) +} + +// RenewCertificate 续期 ACME 证书(当前为占位实现)。 +func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) { + cert, err := model.GetTLSCertificateByID(ctx, id) + if err != nil { + return nil, err + } + if cert.Provider != "acme" { + return nil, errors.New(errCertificateOnlyACMERenew) + } + cert.ApplyStatus = "applying" + cert.ApplyMessage = "" + if err := model.SaveTLSCertificate(ctx, cert); err != nil { + return nil, err + } + return markACMEStubFailure(ctx, cert) +} + +// ListDNSAccounts 列出 DNS 账号。 +func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) { + return model.ListDNSAccounts(ctx) +} + +// CreateDNSAccount 创建 DNS 账号。 +func CreateDNSAccount(ctx context.Context, input DNSAccountInput) (*model.DNSAccount, error) { + authorization, err := sealSensitive(strings.TrimSpace(input.Authorization)) + if err != nil { + return nil, err + } + account := &model.DNSAccount{ + Name: strings.TrimSpace(input.Name), + Type: strings.TrimSpace(input.Type), + Authorization: authorization, + } + if account.Name == "" || account.Type == "" || authorization == "" { + return nil, errors.New("DNS 账号参数不完整") + } + if err := model.CreateDNSAccountRecord(ctx, account); err != nil { + return nil, err + } + return sanitizeDNSAccountForResponse(account), nil +} + +// UpdateDNSAccount 更新 DNS 账号。 +func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*model.DNSAccount, error) { + account, err := model.GetDNSAccountByID(ctx, id) + if err != nil { + return nil, err + } + authorization, err := sealSensitive(strings.TrimSpace(input.Authorization)) + if err != nil { + return nil, err + } + account.Name = strings.TrimSpace(input.Name) + account.Type = strings.TrimSpace(input.Type) + account.Authorization = authorization + if account.Name == "" || account.Type == "" || authorization == "" { + return nil, errors.New("DNS 账号参数不完整") + } + if err := model.SaveDNSAccount(ctx, account); err != nil { + return nil, err + } + return sanitizeDNSAccountForResponse(account), nil +} + +// DeleteDNSAccount 删除 DNS 账号。 +func DeleteDNSAccount(ctx context.Context, id uint) error { + if _, err := model.GetDNSAccountByID(ctx, id); err != nil { + return err + } + count, err := model.CountTLSCertificatesByDNSAccountID(ctx, id) + if err != nil { + return err + } + if count > 0 { + return errors.New(errDNSAccountInUse) + } + return model.DeleteDNSAccountRecord(ctx, id) +} + +// GetDefaultAcmeAccount 获取默认 ACME 账号。 +func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) { + account, err := model.GetDefaultAcmeAccount(ctx) + if err != nil { + return nil, err + } + return sanitizeAcmeAccountForResponse(account), nil +} + +func buildCertificate(ctx context.Context, existing *model.TLSCertificate, input CertificateInput) (*model.TLSCertificate, error) { + name := strings.TrimSpace(input.Name) + certPEM := strings.TrimSpace(input.CertPEM) + keyPEM := strings.TrimSpace(input.KeyPEM) + remark := strings.TrimSpace(input.Remark) + if name == "" { + return nil, errors.New(errCertificateNameRequired) + } + if certPEM == "" || keyPEM == "" { + return nil, errors.New(errCertificateContentRequired) + } + parsed, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) + if err != nil { + return nil, fmt.Errorf("%s: %w", errCertificateContentInvalid, err) + } + if len(parsed.Certificate) == 0 { + return nil, errors.New(errCertificateContentInvalid) + } + leaf, err := parseLeafCertificate(certPEM) + if err != nil { + return nil, err + } + sealedKey, err := sealSensitive(keyPEM) + if err != nil { + return nil, err + } + if existing == nil { + existing = &model.TLSCertificate{ + Provider: "upload", + ApplyStatus: "ready", + } + } + existing.Name = name + existing.CertPEM = certPEM + existing.KeyPEM = sealedKey + existing.NotBefore = leaf.NotBefore + existing.NotAfter = leaf.NotAfter + existing.Remark = remark + return existing, nil +} + +func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) { + cert.Name = strings.TrimSpace(input.Name) + cert.Remark = strings.TrimSpace(input.Remark) + cert.AcmeAccountID = input.AcmeAccountID + cert.DnsAccountID = input.DnsAccountID + cert.KeyAlgorithm = input.KeyAlgorithm + cert.AutoRenew = input.AutoRenew + cert.PrimaryDomain = strings.TrimSpace(input.PrimaryDomain) + cert.OtherDomains = strings.TrimSpace(input.OtherDomains) + cert.DisableCNAME = input.DisableCNAME + cert.SkipDNS = input.SkipDNS + cert.DNS1 = strings.TrimSpace(input.DNS1) + cert.DNS2 = strings.TrimSpace(input.DNS2) + cert.Provider = "acme" + cert.ApplyStatus = "applying" +} + +func markACMEStubFailure(ctx context.Context, cert *model.TLSCertificate) (*model.TLSCertificate, error) { + cert.ApplyStatus = "failed" + cert.ApplyMessage = errACMENotImplemented + if err := model.SaveTLSCertificate(ctx, cert); err != nil { + return nil, err + } + return sanitizeCertificateForResponse(cert), nil +} + +func ensureCertificateNotReferenced(ctx context.Context, id uint) error { + routes, err := model.ListTLSProxyRouteRefs(ctx) + if err != nil { + return err + } + for _, route := range routes { + if route.CertID != nil && *route.CertID == id { + return errors.New(errCertificateDeleteReferenced) + } + if strings.TrimSpace(route.CertIDs) == "" { + continue + } + var certIDs []uint + if err := json.Unmarshal([]byte(route.CertIDs), &certIDs); err != nil { + return fmt.Errorf("proxy route %d cert_ids payload is invalid: %w", route.ID, err) + } + for _, certID := range certIDs { + if certID == id { + return errors.New(errCertificateDeleteReferenced) + } + } + domainCertIDs, err := decodeStoredDomainCertIDs(route.DomainCertIDs, 0) + if err != nil { + return fmt.Errorf("proxy route %d domain_cert_ids payload is invalid: %w", route.ID, err) + } + for _, certID := range domainCertIDs { + if certID == id { + return errors.New(errCertificateDeleteReferenced) + } + } + } + return nil +} + +func sanitizeCertificateForResponse(certificate *model.TLSCertificate) *model.TLSCertificate { + if certificate == nil { + return nil + } + copy := *certificate + copy.CertPEM = "" + copy.KeyPEM = "" + return © +} + +func sanitizeDNSAccountForResponse(account *model.DNSAccount) *model.DNSAccount { + if account == nil { + return nil + } + copy := *account + copy.Authorization = "" + return © +} + +func sanitizeAcmeAccountForResponse(account *model.AcmeAccount) *model.AcmeAccount { + if account == nil { + return nil + } + copy := *account + copy.PrivateKey = "" + return © +} diff --git a/Wavelet/internal/apps/openflare/tls/logics_test.go b/Wavelet/internal/apps/openflare/tls/logics_test.go new file mode 100644 index 00000000..9e0f1d26 --- /dev/null +++ b/Wavelet/internal/apps/openflare/tls/logics_test.go @@ -0,0 +1,141 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package tls + +import ( + "context" + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "strings" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupTLSTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.TLSCertificate{}, + &model.ManagedDomain{}, + &model.DNSAccount{}, + &model.AcmeAccount{}, + )) + + db.SetDB(sqliteDB) + oldSecret := config.Config.App.SessionSecret + config.Config.App.SessionSecret = "test_session_secret_for_tls_encryption" + return func() { + db.SetDB(nil) + config.Config.App.SessionSecret = oldSecret + } +} + +func generateTestCertificatePair(t *testing.T, dnsNames []string) (string, string) { + t.Helper() + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + template := &x509.Certificate{ + SerialNumber: big.NewInt(time.Now().UnixNano()), + Subject: pkix.Name{ + CommonName: dnsNames[0], + }, + DNSNames: dnsNames, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + } + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey) + require.NoError(t, err) + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)}) + return string(certPEM), string(keyPEM) +} + +func TestCreateManagedDomain(t *testing.T) { + cleanup := setupTLSTestDB(t) + defer cleanup() + ctx := context.Background() + + certPEM, keyPEM := generateTestCertificatePair(t, []string{"api.example.com"}) + certificate, err := CreateCertificate(ctx, CertificateInput{ + Name: "api-cert", + CertPEM: certPEM, + KeyPEM: keyPEM, + }) + require.NoError(t, err) + + certID := certificate.ID + domain, err := CreateManagedDomain(ctx, ManagedDomainInput{ + Domain: "api.example.com", + CertID: &certID, + Enabled: true, + Remark: "primary api", + }) + require.NoError(t, err) + assert.NotZero(t, domain.ID) + assert.Equal(t, "api.example.com", domain.Domain) + assert.Equal(t, certID, *domain.CertID) + assert.True(t, domain.Enabled) + assert.Equal(t, "primary api", domain.Remark) + + _, err = CreateManagedDomain(ctx, ManagedDomainInput{ + Domain: "api.example.com", + Enabled: true, + }) + require.Error(t, err) + assert.Equal(t, errManagedDomainExists, err.Error()) +} + +func TestCreateManagedDomainRejectsInvalidWildcard(t *testing.T) { + cleanup := setupTLSTestDB(t) + defer cleanup() + ctx := context.Background() + + _, err := CreateManagedDomain(ctx, ManagedDomainInput{ + Domain: "*.*.example.com", + Enabled: true, + }) + require.Error(t, err) + assert.Equal(t, errManagedDomainWildcardInvalid, err.Error()) +} + +func TestCreateCertificateEncryptsPrivateKey(t *testing.T) { + cleanup := setupTLSTestDB(t) + defer cleanup() + ctx := context.Background() + + certPEM, keyPEM := generateTestCertificatePair(t, []string{"secure.example.com"}) + certificate, err := CreateCertificate(ctx, CertificateInput{ + Name: "secure-cert", + CertPEM: certPEM, + KeyPEM: keyPEM, + }) + require.NoError(t, err) + + stored, err := model.GetTLSCertificateByID(ctx, certificate.ID) + require.NoError(t, err) + assert.NotEqual(t, keyPEM, stored.KeyPEM) + assert.Contains(t, stored.KeyPEM, sensitiveValuePrefix) + + content, err := GetCertificateContent(ctx, certificate.ID) + require.NoError(t, err) + assert.Equal(t, strings.TrimSpace(keyPEM), strings.TrimSpace(content.KeyPEM)) +} diff --git a/Wavelet/internal/apps/openflare/tls/managed_domain.go b/Wavelet/internal/apps/openflare/tls/managed_domain.go new file mode 100644 index 00000000..3e583de6 --- /dev/null +++ b/Wavelet/internal/apps/openflare/tls/managed_domain.go @@ -0,0 +1,239 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package tls + +import ( + "context" + "errors" + "fmt" + "sort" + "strings" + "unicode" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + managedDomainMatchTypeExact = "exact" + managedDomainMatchTypeWildcard = "wildcard" +) + +// ManagedDomainInput 托管域名创建/更新请求。 +type ManagedDomainInput struct { + Domain string `json:"domain"` + CertID *uint `json:"cert_id"` + Enabled bool `json:"enabled"` + Remark string `json:"remark"` +} + +// ManagedDomainMatchCandidate 证书匹配候选。 +type ManagedDomainMatchCandidate struct { + ManagedDomainID uint `json:"managed_domain_id"` + Domain string `json:"domain"` + MatchType string `json:"match_type"` + CertificateID uint `json:"certificate_id"` + CertificateName string `json:"certificate_name"` +} + +// ManagedDomainMatchResult 证书匹配结果。 +type ManagedDomainMatchResult struct { + Domain string `json:"domain"` + Matched bool `json:"matched"` + Candidate *ManagedDomainMatchCandidate `json:"candidate,omitempty"` + Candidates []ManagedDomainMatchCandidate `json:"candidates"` +} + +// ListManagedDomains 列出托管域名。 +func ListManagedDomains(ctx context.Context) ([]model.ManagedDomain, error) { + return model.ListManagedDomains(ctx) +} + +// CreateManagedDomain 创建托管域名。 +func CreateManagedDomain(ctx context.Context, input ManagedDomainInput) (*model.ManagedDomain, error) { + domain, err := buildManagedDomain(ctx, nil, input) + if err != nil { + return nil, err + } + if err = model.CreateManagedDomainRecord(ctx, domain); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errManagedDomainExists) + } + return nil, err + } + return domain, nil +} + +// UpdateManagedDomain 更新托管域名。 +func UpdateManagedDomain(ctx context.Context, id uint, input ManagedDomainInput) (*model.ManagedDomain, error) { + domain, err := model.GetManagedDomainByID(ctx, id) + if err != nil { + return nil, err + } + domain, err = buildManagedDomain(ctx, domain, input) + if err != nil { + return nil, err + } + if err = model.SaveManagedDomain(ctx, domain); err != nil { + if isUniqueConstraintError(err) { + return nil, errors.New(errManagedDomainExists) + } + return nil, err + } + return domain, nil +} + +// DeleteManagedDomain 删除托管域名。 +func DeleteManagedDomain(ctx context.Context, id uint) error { + if _, err := model.GetManagedDomainByID(ctx, id); err != nil { + return err + } + return model.DeleteManagedDomainRecord(ctx, id) +} + +// MatchManagedDomainCertificate 为域名匹配证书。 +func MatchManagedDomainCertificate(ctx context.Context, rawDomain string) (*ManagedDomainMatchResult, error) { + domain := normalizeManagedDomain(rawDomain) + if err := validateManagedDomainPattern(domain); err != nil { + return nil, err + } + managedDomains, err := model.ListEnabledManagedDomainsWithCertificate(ctx) + if err != nil { + return nil, err + } + candidates := make([]ManagedDomainMatchCandidate, 0) + for _, item := range managedDomains { + if item.CertID == nil || *item.CertID == 0 { + continue + } + matchType := detectManagedDomainMatchType(item.Domain, domain) + if matchType == "" { + continue + } + certificate, err := model.GetTLSCertificateByID(ctx, *item.CertID) + if err != nil { + return nil, fmt.Errorf("托管域名 %s 关联证书不存在", item.Domain) + } + candidates = append(candidates, ManagedDomainMatchCandidate{ + ManagedDomainID: item.ID, + Domain: item.Domain, + MatchType: matchType, + CertificateID: certificate.ID, + CertificateName: certificate.Name, + }) + } + sortManagedDomainCandidates(candidates) + result := &ManagedDomainMatchResult{ + Domain: domain, + Matched: len(candidates) > 0, + Candidates: candidates, + } + if len(candidates) > 0 { + candidate := candidates[0] + result.Candidate = &candidate + } + return result, nil +} + +func buildManagedDomain(ctx context.Context, existing *model.ManagedDomain, input ManagedDomainInput) (*model.ManagedDomain, error) { + domain := normalizeManagedDomain(input.Domain) + remark := strings.TrimSpace(input.Remark) + if err := validateManagedDomainPattern(domain); err != nil { + return nil, err + } + if input.CertID != nil && *input.CertID != 0 { + if _, err := model.GetTLSCertificateByID(ctx, *input.CertID); err != nil { + return nil, errors.New(errManagedDomainCertNotFound) + } + } else { + input.CertID = nil + } + if existing == nil { + existing = &model.ManagedDomain{} + } + existing.Domain = domain + existing.CertID = input.CertID + existing.Enabled = input.Enabled + existing.Remark = remark + return existing, nil +} + +func normalizeManagedDomain(domain string) string { + return strings.ToLower(strings.TrimSpace(domain)) +} + +func validateManagedDomainPattern(domain string) error { + if domain == "" { + return errors.New(errManagedDomainRequired) + } + if strings.Contains(domain, "://") || strings.Contains(domain, "/") { + return errors.New(errManagedDomainInvalid) + } + if strings.Contains(domain, "*") { + if !strings.HasPrefix(domain, "*.") || strings.Count(domain, "*") != 1 { + return errors.New(errManagedDomainWildcardInvalid) + } + return validateHostname(strings.TrimPrefix(domain, "*.")) + } + return validateHostname(domain) +} + +func validateHostname(domain string) error { + if domain == "" { + return errors.New(errManagedDomainRequired) + } + if len(domain) > 253 { + return errors.New(errManagedDomainInvalid) + } + labels := strings.Split(domain, ".") + if len(labels) < 2 { + return errors.New(errManagedDomainInvalid) + } + for _, label := range labels { + if len(label) == 0 || len(label) > 63 { + return errors.New(errManagedDomainInvalid) + } + if label[0] == '-' || label[len(label)-1] == '-' { + return errors.New(errManagedDomainInvalid) + } + for _, r := range label { + if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' { + continue + } + return errors.New(errManagedDomainInvalid) + } + } + return nil +} + +func detectManagedDomainMatchType(pattern string, domain string) string { + if pattern == domain { + return managedDomainMatchTypeExact + } + if !strings.HasPrefix(pattern, "*.") { + return "" + } + suffix := strings.TrimPrefix(pattern, "*.") + if !strings.HasSuffix(domain, "."+suffix) { + return "" + } + prefix := strings.TrimSuffix(domain, "."+suffix) + if prefix == "" || strings.Contains(prefix, ".") { + return "" + } + return managedDomainMatchTypeWildcard +} + +func sortManagedDomainCandidates(candidates []ManagedDomainMatchCandidate) { + sort.Slice(candidates, func(i int, j int) bool { + left := candidates[i] + right := candidates[j] + if left.MatchType != right.MatchType { + return left.MatchType == managedDomainMatchTypeExact + } + if len(left.Domain) != len(right.Domain) { + return len(left.Domain) > len(right.Domain) + } + return left.ManagedDomainID < right.ManagedDomainID + }) +} diff --git a/Wavelet/internal/apps/openflare/tls/routers.go b/Wavelet/internal/apps/openflare/tls/routers.go new file mode 100644 index 00000000..103dadae --- /dev/null +++ b/Wavelet/internal/apps/openflare/tls/routers.go @@ -0,0 +1,304 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package tls + +import ( + "errors" + "strings" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +func handleLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, "记录不存在") + return true + } + compat.Fail(c, err.Error()) + return true +} + +// GetCertificates 列出 TLS 证书。 +func GetCertificates(c *gin.Context) { + certificates, err := ListCertificates(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificates) +} + +// GetCertificateDetail 获取 TLS 证书详情。 +func GetCertificateDetail(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + certificate, err := GetCertificate(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// GetCertificateContentHandler 获取 TLS 证书 PEM 内容。 +func GetCertificateContentHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + content, err := GetCertificateContent(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, content) +} + +// CreateCertificateHandler 从 PEM 创建证书。 +func CreateCertificateHandler(c *gin.Context) { + var input CertificateInput + if !compat.BindJSON(c, &input) { + return + } + certificate, err := CreateCertificate(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// UpdateCertificateHandler 更新证书。 +func UpdateCertificateHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input CertificateInput + if !compat.BindJSON(c, &input) { + return + } + certificate, err := UpdateCertificate(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// ImportCertificateFile 从文件导入证书。 +func ImportCertificateFile(c *gin.Context) { + name := c.PostForm("name") + remark := c.PostForm("remark") + certFile, err := c.FormFile("cert_file") + if err != nil { + compat.Fail(c, "缺少证书文件") + return + } + keyFile, err := c.FormFile("key_file") + if err != nil { + compat.Fail(c, "缺少私钥文件") + return + } + certificate, err := CreateCertificateFromFiles(c.Request.Context(), name, certFile, keyFile, remark) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// DeleteCertificateHandler 删除证书。 +func DeleteCertificateHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteCertificate(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OK(c, nil) +} + +// ApplyCertificateHandler 申请 ACME 证书。 +func ApplyCertificateHandler(c *gin.Context) { + var input ApplyInput + if !compat.BindJSON(c, &input) { + return + } + certificate, err := ApplyCertificate(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// UpdateACMECertificateHandler 更新 ACME 证书配置。 +func UpdateACMECertificateHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input ApplyInput + if !compat.BindJSON(c, &input) { + return + } + certificate, err := UpdateACMECertificate(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// ConvertCertificateToACMEHandler 将上传证书转为 ACME。 +func ConvertCertificateToACMEHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input ApplyInput + if !compat.BindJSON(c, &input) { + return + } + certificate, err := ConvertCertificateToACME(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// RenewCertificateHandler 续期 ACME 证书。 +func RenewCertificateHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + certificate, err := RenewCertificate(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, certificate) +} + +// GetManagedDomains 列出托管域名。 +func GetManagedDomains(c *gin.Context) { + domains, err := ListManagedDomains(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, domains) +} + +// CreateManagedDomainHandler 创建托管域名。 +func CreateManagedDomainHandler(c *gin.Context) { + var input ManagedDomainInput + if !compat.BindJSON(c, &input) { + return + } + domain, err := CreateManagedDomain(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, domain) +} + +// UpdateManagedDomainHandler 更新托管域名。 +func UpdateManagedDomainHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input ManagedDomainInput + if !compat.BindJSON(c, &input) { + return + } + domain, err := UpdateManagedDomain(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, domain) +} + +// DeleteManagedDomainHandler 删除托管域名。 +func DeleteManagedDomainHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteManagedDomain(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OK(c, nil) +} + +// MatchManagedDomainCertificateHandler 匹配域名证书。 +func MatchManagedDomainCertificateHandler(c *gin.Context) { + domain := strings.TrimSpace(c.Query("domain")) + result, err := MatchManagedDomainCertificate(c.Request.Context(), domain) + if handleLogicError(c, err) { + return + } + compat.OK(c, result) +} + +// GetDNSAccounts 列出 DNS 账号。 +func GetDNSAccounts(c *gin.Context) { + accounts, err := ListDNSAccounts(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, accounts) +} + +// CreateDNSAccountHandler 创建 DNS 账号。 +func CreateDNSAccountHandler(c *gin.Context) { + var input DNSAccountInput + if !compat.BindJSON(c, &input) { + return + } + account, err := CreateDNSAccount(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, account) +} + +// UpdateDNSAccountHandler 更新 DNS 账号。 +func UpdateDNSAccountHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input DNSAccountInput + if !compat.BindJSON(c, &input) { + return + } + account, err := UpdateDNSAccount(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, account) +} + +// DeleteDNSAccountHandler 删除 DNS 账号。 +func DeleteDNSAccountHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteDNSAccount(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OK(c, nil) +} + +// GetDefaultAcmeAccountHandler 获取默认 ACME 账号。 +func GetDefaultAcmeAccountHandler(c *gin.Context) { + account, err := GetDefaultAcmeAccount(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, account) +} diff --git a/Wavelet/internal/apps/openflare/tls/sensitive.go b/Wavelet/internal/apps/openflare/tls/sensitive.go new file mode 100644 index 00000000..2a488ea4 --- /dev/null +++ b/Wavelet/internal/apps/openflare/tls/sensitive.go @@ -0,0 +1,55 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package tls + +import ( + "crypto/sha256" + "encoding/hex" + "errors" + "strings" + + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/pkg/util" +) + +const sensitiveValuePrefix = "enc:v1:" + +func sensitiveEncryptionKey() string { + if config.Config == nil || strings.TrimSpace(config.Config.App.SessionSecret) == "" { + return "" + } + sum := sha256.Sum256([]byte(config.Config.App.SessionSecret)) + return hex.EncodeToString(sum[:]) +} + +func sealSensitive(plaintext string) (string, error) { + plaintext = strings.TrimSpace(plaintext) + if plaintext == "" { + return "", nil + } + key := sensitiveEncryptionKey() + if key == "" { + return plaintext, nil + } + encrypted, err := util.Encrypt(key, plaintext) + if err != nil { + return "", err + } + return sensitiveValuePrefix + encrypted, nil +} + +func openSensitive(stored string) (string, error) { + stored = strings.TrimSpace(stored) + if stored == "" { + return "", nil + } + if !strings.HasPrefix(stored, sensitiveValuePrefix) { + return stored, nil + } + key := sensitiveEncryptionKey() + if key == "" { + return "", errors.New("cannot decrypt sensitive field without session secret") + } + return util.Decrypt(key, strings.TrimPrefix(stored, sensitiveValuePrefix)) +} diff --git a/Wavelet/internal/apps/openflare/update/logics.go b/Wavelet/internal/apps/openflare/update/logics.go new file mode 100644 index 00000000..7eb1d4c8 --- /dev/null +++ b/Wavelet/internal/apps/openflare/update/logics.go @@ -0,0 +1,262 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package update + +import ( + "context" + "errors" + "fmt" + "runtime" + "strings" + "sync" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/admin/updater" + "github.com/Rain-kl/Wavelet/pkg/logger" +) + +const ( + channelStable = "stable" + channelPreview = "preview" +) + +// UpgradeLogRecord is a single upgrade log entry for the legacy update API. +type UpgradeLogRecord struct { + Level string `json:"level"` + Message string `json:"message"` + CreatedAt time.Time `json:"created_at"` +} + +// LatestReleaseView mirrors the legacy OpenFlare latest-release payload. +type LatestReleaseView struct { + TagName string `json:"tag_name"` + Body string `json:"body"` + HTMLURL string `json:"html_url"` + PublishedAt string `json:"published_at"` + Channel string `json:"channel"` + Prerelease bool `json:"prerelease"` + CurrentVersion string `json:"current_version"` + HasUpdate bool `json:"has_update"` + UpgradeSupported bool `json:"upgrade_supported"` + InProgress bool `json:"in_progress"` + UpgradeStatus string `json:"upgrade_status"` + UpgradeLogs []UpgradeLogRecord `json:"upgrade_logs"` +} + +// StreamSnapshot is pushed over the upgrade logs websocket. +type StreamSnapshot struct { + InProgress bool `json:"in_progress"` + UpgradeStatus string `json:"upgrade_status"` + UpgradeLogs []UpgradeLogRecord `json:"upgrade_logs"` +} + +type upgradeRequest struct { + Channel string `json:"channel"` +} + +var upgradeState struct { + sync.Mutex + inProgress bool + status string + logs []UpgradeLogRecord +} + +var upgradeSubscribers struct { + sync.Mutex + nextID int + listeners map[int]chan StreamSnapshot +} + +func init() { + upgradeSubscribers.listeners = make(map[int]chan StreamSnapshot) + upgradeState.status = "idle" +} + +func normalizeChannel(channel string) string { + switch strings.ToLower(strings.TrimSpace(channel)) { + case channelPreview: + return channelPreview + default: + return channelStable + } +} + +func isDevBuild(version string) bool { + version = strings.TrimSpace(version) + return version == "" || strings.EqualFold(version, "dev") +} + +func mapStatusToLatestRelease(status updater.Status, channel string) *LatestReleaseView { + inProgress, upgradeStatus, logs := snapshotUpgradeState() + view := &LatestReleaseView{ + TagName: status.LatestVersion, + Body: status.ReleaseNotes, + HTMLURL: status.ReleaseURL, + PublishedAt: status.PublishedAt, + Channel: channel, + Prerelease: status.Prerelease, + CurrentVersion: status.CurrentVersion, + HasUpdate: status.UpdateAvailable, + UpgradeSupported: !isDevBuild(status.CurrentVersion) && runtime.GOOS != "windows", + InProgress: inProgress || updater.IsUpgrading(), + UpgradeStatus: upgradeStatus, + UpgradeLogs: logs, + } + if channel == channelPreview && status.Prerelease { + view.HasUpdate = !isDevBuild(status.CurrentVersion) + } + return view +} + +// GetLatestRelease returns the newest upstream release for the requested channel. +func GetLatestRelease(ctx context.Context, channel string) (*LatestReleaseView, error) { + normalizedChannel := normalizeChannel(channel) + status, err := updater.GetStatus(ctx) + if err != nil { + return nil, err + } + return mapStatusToLatestRelease(status, normalizedChannel), nil +} + +// ScheduleUpgrade downloads the latest release and restarts with the staged binary. +func ScheduleUpgrade(ctx context.Context, channel string) (*LatestReleaseView, error) { + normalizedChannel := normalizeChannel(channel) + + upgradeState.Lock() + if upgradeState.inProgress || updater.IsUpgrading() { + upgradeState.Unlock() + return nil, errors.New("服务升级正在执行中,请稍后再试") + } + resetUpgradeLogsLocked() + upgradeState.inProgress = true + upgradeState.status = "running" + appendUpgradeLogLocked("info", fmt.Sprintf("Automatic upgrade scheduled for channel: %s.", normalizedChannel)) + upgradeState.Unlock() + broadcastUpgradeSnapshot() + + executable, stagedBinary, status, err := updater.PrepareUpgrade(ctx) + if err != nil { + recordUpgradeFailure(err) + return nil, err + } + + view := mapStatusToLatestRelease(status, normalizedChannel) + view.InProgress = true + view.UpgradeStatus = "running" + view.UpgradeLogs = snapshotUpgradeLogs() + + appendUpgradeLogLocked("info", fmt.Sprintf("Upgrade package prepared: %s.", status.LatestVersion)) + broadcastUpgradeSnapshot() + + go func() { + time.Sleep(time.Second) + if err := updater.ApplyPreparedUpgrade(executable, stagedBinary); err != nil { + updater.FinishUpgrade() + recordUpgradeFailure(err) + logger.ErrorF(context.Background(), "[Update] replace and restart failed: %v", err) + } + }() + + return view, nil +} + +// UploadManualBinary is disabled in the legacy OpenFlare server. +func UploadManualBinary() error { + return errors.New("手动升级功能已禁用") +} + +// ConfirmManualUpgrade is disabled in the legacy OpenFlare server. +func ConfirmManualUpgrade() error { + return errors.New("手动升级功能已禁用") +} + +// SubscribeUpgradeStream registers a listener for upgrade websocket snapshots. +func SubscribeUpgradeStream() (<-chan StreamSnapshot, func()) { + upgradeSubscribers.Lock() + defer upgradeSubscribers.Unlock() + + id := upgradeSubscribers.nextID + upgradeSubscribers.nextID++ + ch := make(chan StreamSnapshot, 1) + upgradeSubscribers.listeners[id] = ch + + unsubscribe := func() { + upgradeSubscribers.Lock() + defer upgradeSubscribers.Unlock() + if listener, ok := upgradeSubscribers.listeners[id]; ok { + delete(upgradeSubscribers.listeners, id) + close(listener) + } + } + + select { + case ch <- currentUpgradeSnapshot(): + default: + } + + return ch, unsubscribe +} + +func snapshotUpgradeState() (bool, string, []UpgradeLogRecord) { + upgradeState.Lock() + defer upgradeState.Unlock() + return upgradeState.inProgress || updater.IsUpgrading(), upgradeState.status, cloneUpgradeLogsLocked() +} + +func snapshotUpgradeLogs() []UpgradeLogRecord { + upgradeState.Lock() + defer upgradeState.Unlock() + return cloneUpgradeLogsLocked() +} + +func currentUpgradeSnapshot() StreamSnapshot { + inProgress, status, logs := snapshotUpgradeState() + return StreamSnapshot{ + InProgress: inProgress, + UpgradeStatus: status, + UpgradeLogs: logs, + } +} + +func resetUpgradeLogsLocked() { + upgradeState.logs = nil +} + +func appendUpgradeLogLocked(level, message string) { + upgradeState.logs = append(upgradeState.logs, UpgradeLogRecord{ + Level: level, + Message: message, + CreatedAt: time.Now().UTC(), + }) +} + +func cloneUpgradeLogsLocked() []UpgradeLogRecord { + if len(upgradeState.logs) == 0 { + return []UpgradeLogRecord{} + } + cloned := make([]UpgradeLogRecord, len(upgradeState.logs)) + copy(cloned, upgradeState.logs) + return cloned +} + +func recordUpgradeFailure(err error) { + upgradeState.Lock() + upgradeState.inProgress = false + upgradeState.status = "failed" + appendUpgradeLogLocked("error", err.Error()) + upgradeState.Unlock() + broadcastUpgradeSnapshot() +} + +func broadcastUpgradeSnapshot() { + snapshot := currentUpgradeSnapshot() + upgradeSubscribers.Lock() + defer upgradeSubscribers.Unlock() + for _, listener := range upgradeSubscribers.listeners { + select { + case listener <- snapshot: + default: + } + } +} diff --git a/Wavelet/internal/apps/openflare/update/routers.go b/Wavelet/internal/apps/openflare/update/routers.go new file mode 100644 index 00000000..cf9546c2 --- /dev/null +++ b/Wavelet/internal/apps/openflare/update/routers.go @@ -0,0 +1,113 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package update + +import ( + "encoding/json" + "errors" + "io" + "net/http" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +var upgradeLogsUpgrader = websocket.Upgrader{ + CheckOrigin: func(_ *http.Request) bool { return true }, +} + +// GetLatestReleaseHandler returns the newest GitHub release for the legacy update UI. +func GetLatestReleaseHandler(c *gin.Context) { + release, err := GetLatestRelease(c.Request.Context(), c.Query("channel")) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, release) +} + +// UpgradeServerHandler schedules an automatic upgrade from the latest release. +func UpgradeServerHandler(c *gin.Context) { + var request upgradeRequest + if err := bindOptionalJSON(c.Request.Body, &request); err != nil { + compat.Fail(c, "无效的参数") + return + } + + release, err := ScheduleUpgrade(c.Request.Context(), request.Channel) + if err != nil { + compat.Fail(c, err.Error()) + return + } + + okWithMessage(c, release, "服务升级任务已启动,下载完成后将自动重启。") +} + +// UploadManualServerBinaryHandler rejects manual uploads (feature disabled upstream). +func UploadManualServerBinaryHandler(c *gin.Context) { + if err := UploadManualBinary(); err != nil { + compat.Fail(c, err.Error()) + return + } +} + +// ConfirmManualServerUpgradeHandler rejects manual upgrades (feature disabled upstream). +func ConfirmManualServerUpgradeHandler(c *gin.Context) { + if err := ConfirmManualUpgrade(); err != nil { + compat.Fail(c, err.Error()) + return + } +} + +// StreamServerUpgradeLogsHandler streams upgrade progress snapshots over WebSocket. +func StreamServerUpgradeLogsHandler(c *gin.Context) { + conn, err := upgradeLogsUpgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + return + } + defer func() { + _ = conn.Close() + }() + + updates, unsubscribe := SubscribeUpgradeStream() + defer unsubscribe() + + heartbeatTicker := time.NewTicker(15 * time.Second) + defer heartbeatTicker.Stop() + + for { + select { + case snapshot, ok := <-updates: + if !ok { + return + } + if err := conn.WriteJSON(snapshot); err != nil { + return + } + case <-heartbeatTicker.C: + if err := conn.WriteJSON(StreamSnapshot{}); err != nil { + return + } + case <-c.Request.Context().Done(): + return + } + } +} + +func okWithMessage(c *gin.Context, data any, message string) { + c.JSON(http.StatusOK, gin.H{ + "success": true, + "message": message, + "data": data, + }) +} + +func bindOptionalJSON(body io.Reader, target any) error { + if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) { + return err + } + return nil +} diff --git a/Wavelet/internal/apps/openflare/waf/errs.go b/Wavelet/internal/apps/openflare/waf/errs.go new file mode 100644 index 00000000..697033f3 --- /dev/null +++ b/Wavelet/internal/apps/openflare/waf/errs.go @@ -0,0 +1,9 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package waf + +const ( + errWAFRuleGroupNotFound = "WAF 规则组不存在" + errWAFIPGroupNotFound = "IP 组不存在" +) diff --git a/Wavelet/internal/apps/openflare/waf/logics.go b/Wavelet/internal/apps/openflare/waf/logics.go new file mode 100644 index 00000000..6458887d --- /dev/null +++ b/Wavelet/internal/apps/openflare/waf/logics.go @@ -0,0 +1,1167 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package waf + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net" + "net/netip" + "net/url" + "regexp" + "sort" + "strings" + "time" + "unicode" + + "github.com/Rain-kl/Wavelet/internal/model" + "gorm.io/gorm" +) + +const ( + defaultWAFBlockStatusCode = 418 + maxWAFBlockBodyBytes = 16 * 1024 + + wafIPGroupTypeManual = "manual" + wafIPGroupTypeAutomatic = "automatic" + wafIPGroupTypeSubscription = "subscription" + + wafIPGroupSubscriptionFormatText = "text" + wafIPGroupSubscriptionFormatJSON = "json" + + defaultWAFIPGroupSyncIntervalMinutes = 1440 + defaultWAFIPGroupAutoLookbackMinutes = 60 + minWAFIPGroupSyncIntervalMinutes = 5 + maxWAFIPGroupSyncIntervalMinutes = 43200 +) + +// RuleGroupInput is the create/update payload for WAF rule groups. +type RuleGroupInput struct { + Name string `json:"name"` + Enabled bool `json:"enabled"` + BlockStatusCode int `json:"block_status_code"` + BlockResponseBody string `json:"block_response_body"` + IPWhitelist []string `json:"ip_whitelist"` + IPBlacklist []string `json:"ip_blacklist"` + IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"` + IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"` + CountryWhitelist []string `json:"country_whitelist"` + CountryBlacklist []string `json:"country_blacklist"` + RegionWhitelist []string `json:"region_whitelist"` + RegionBlacklist []string `json:"region_blacklist"` + Remark string `json:"remark"` + PoWEnabled bool `json:"pow_enabled"` + PoWConfig json.RawMessage `json:"pow_config"` +} + +// PoWListConfig stores PoW whitelist/blacklist dimensions. +type PoWListConfig struct { + IPs []string `json:"ips"` + IPCidrs []string `json:"ip_cidrs"` + Paths []string `json:"paths"` + PathRegexes []string `json:"path_regexes"` + UserAgents []string `json:"user_agents"` +} + +// PoWConfig stores proof-of-work settings for a rule group. +type PoWConfig struct { + Difficulty int `json:"difficulty"` + Algorithm string `json:"algorithm"` + SessionTTL int `json:"session_ttl"` + ChallengeTTL int `json:"challenge_ttl"` + Whitelist PoWListConfig `json:"whitelist"` + Blacklist PoWListConfig `json:"blacklist"` +} + +// RuleGroupView is the API view for a WAF rule group. +type RuleGroupView struct { + ID uint `json:"id"` + Name string `json:"name"` + Enabled bool `json:"enabled"` + IsGlobal bool `json:"is_global"` + BlockStatusCode int `json:"block_status_code"` + BlockResponseBody string `json:"block_response_body"` + IPWhitelist []string `json:"ip_whitelist"` + IPBlacklist []string `json:"ip_blacklist"` + IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"` + IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"` + CountryWhitelist []string `json:"country_whitelist"` + CountryBlacklist []string `json:"country_blacklist"` + RegionWhitelist []string `json:"region_whitelist"` + RegionBlacklist []string `json:"region_blacklist"` + Remark string `json:"remark"` + PoWEnabled bool `json:"pow_enabled"` + PoWConfig *PoWConfig `json:"pow_config"` + AppliedSiteIDs []uint `json:"applied_site_ids"` + AppliedSiteCount int `json:"applied_site_count"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` +} + +// SiteRuleGroupsView is the site-level WAF binding view. +type SiteRuleGroupsView struct { + RouteID uint `json:"route_id"` + GlobalRuleGroup *RuleGroupView `json:"global_rule_group"` + RuleGroups []RuleGroupView `json:"rule_groups"` + AppliedRuleGroups []RuleGroupView `json:"applied_rule_groups"` + AppliedIDs []uint `json:"applied_ids"` +} + +// IDsRequest carries a list of numeric ids. +type IDsRequest struct { + IDs []uint `json:"ids"` +} + +// IPGroupInput is the create/update payload for WAF IP groups. +type IPGroupInput struct { + Name string `json:"name"` + Type string `json:"type"` + Enabled bool `json:"enabled"` + IPList []string `json:"ip_list"` + AutoConfig json.RawMessage `json:"auto_config"` + SubscriptionURL string `json:"subscription_url"` + SubscriptionFormat string `json:"subscription_format"` + SubscriptionMappingRule string `json:"subscription_mapping_rule"` + SyncIntervalMinutes int `json:"sync_interval_minutes"` + Remark string `json:"remark"` +} + +// IPGroupExtIPView is an external IP entry in API responses. +type IPGroupExtIPView struct { + IP string `json:"ip"` + CapturedAt string `json:"captured_at"` +} + +// IPGroupView is the API view for a WAF IP group. +type IPGroupView struct { + ID uint `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + Enabled bool `json:"enabled"` + IPList []string `json:"ip_list"` + AutoConfig json.RawMessage `json:"auto_config"` + ExtIPs []IPGroupExtIPView `json:"ext_ips"` + SubscriptionURL string `json:"subscription_url"` + SubscriptionFormat string `json:"subscription_format"` + SubscriptionMappingRule string `json:"subscription_mapping_rule"` + SyncIntervalMinutes int `json:"sync_interval_minutes"` + LastSyncedAt string `json:"last_synced_at,omitempty"` + NextSyncAt string `json:"next_sync_at,omitempty"` + LastSyncStatus string `json:"last_sync_status"` + LastSyncMessage string `json:"last_sync_message"` + Remark string `json:"remark"` + ReferencedByRuleCount int `json:"referenced_by_rule_count"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` +} + +// IPGroupSyncResult is the response for manual IP group sync. +type IPGroupSyncResult struct { + Group IPGroupView `json:"group"` + IPCount int `json:"ip_count"` + SyncedAt string `json:"synced_at"` + NextSyncAt string `json:"next_sync_at"` + Status string `json:"status"` + Message string `json:"message"` +} + +// IPGroupAutoTestInput tests automatic IP group configuration. +type IPGroupAutoTestInput struct { + AutoConfig json.RawMessage `json:"auto_config"` +} + +// IPGroupAutoTestResult is the response for automatic IP group test. +type IPGroupAutoTestResult struct { + MatchedIPs []string `json:"matched_ips"` + MatchedCount int `json:"matched_count"` + LookbackMinutes int `json:"lookback_minutes"` + RuleCount int `json:"rule_count"` + TestedAt string `json:"tested_at"` +} + +type ipGroupAutoConfig struct { + LookbackMinutes int `json:"lookback_minutes"` + TTL int `json:"ttl"` + Rules []ipGroupAutoRule `json:"rules"` +} + +type ipGroupAutoRule struct { + Name string `json:"name"` + Expr string `json:"expr"` +} + +type ipGroupExtIP struct { + IP string `json:"ip"` + CapturedAt time.Time `json:"captured_at"` +} + +var powAlgorithmValues = map[string]bool{"fast": true, "slow": true} + +// ListRuleGroups returns all WAF rule groups. +func ListRuleGroups(ctx context.Context) ([]RuleGroupView, error) { + if err := EnsureDefaultRuleGroup(ctx); err != nil { + return nil, err + } + groups, err := model.ListOpenFlareWAFRuleGroups(ctx) + if err != nil { + return nil, err + } + bindings, err := loadRuleGroupBindings(ctx) + if err != nil { + return nil, err + } + views := make([]RuleGroupView, 0, len(groups)) + for _, group := range groups { + view, buildErr := buildRuleGroupView(group, bindings[group.ID]) + if buildErr != nil { + return nil, buildErr + } + views = append(views, view) + } + return views, nil +} + +// GetRuleGroup returns a WAF rule group by id. +func GetRuleGroup(ctx context.Context, id uint) (*RuleGroupView, error) { + group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id) + if err != nil { + return nil, err + } + bindings, err := loadRuleGroupBindings(ctx) + if err != nil { + return nil, err + } + view, err := buildRuleGroupView(group, bindings[group.ID]) + if err != nil { + return nil, err + } + return &view, nil +} + +// CreateRuleGroup creates a custom WAF rule group. +func CreateRuleGroup(ctx context.Context, input RuleGroupInput) (*RuleGroupView, error) { + group, err := buildRuleGroup(ctx, nil, input) + if err != nil { + return nil, err + } + group.IsGlobal = false + if err = model.CreateOpenFlareWAFRuleGroup(ctx, group); err != nil { + return nil, err + } + return GetRuleGroup(ctx, group.ID) +} + +// UpdateRuleGroup updates a WAF rule group. +func UpdateRuleGroup(ctx context.Context, id uint, input RuleGroupInput) (*RuleGroupView, error) { + group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id) + if err != nil { + return nil, err + } + isGlobal := group.IsGlobal + group, err = buildRuleGroup(ctx, group, input) + if err != nil { + return nil, err + } + group.IsGlobal = isGlobal + if isGlobal && strings.TrimSpace(group.Name) == "" { + group.Name = "全局规则组" + } + if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil { + return nil, err + } + return GetRuleGroup(ctx, group.ID) +} + +// DeleteRuleGroup deletes a non-global WAF rule group. +func DeleteRuleGroup(ctx context.Context, id uint) error { + group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id) + if err != nil { + return err + } + if group.IsGlobal { + return errors.New("全局 WAF 规则组不能删除") + } + return model.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, group.ID) +} + +// ReplaceRuleGroupSites replaces site bindings for a rule group. +func ReplaceRuleGroupSites(ctx context.Context, groupID uint, routeIDs []uint) (*RuleGroupView, error) { + group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, groupID) + if err != nil { + return nil, err + } + if group.IsGlobal { + return nil, errors.New("全局 WAF 规则组默认应用到所有网站,不能手动绑定") + } + normalized, err := normalizeRouteIDs(ctx, routeIDs) + if err != nil { + return nil, err + } + if err = model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, groupID, normalized); err != nil { + return nil, err + } + return GetRuleGroup(ctx, groupID) +} + +// GetSiteRuleGroups returns WAF rule groups for a proxy route. +func GetSiteRuleGroups(ctx context.Context, routeID uint) (*SiteRuleGroupsView, error) { + if _, err := model.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil { + return nil, err + } + groups, err := ListRuleGroups(ctx) + if err != nil { + return nil, err + } + appliedIDs, err := ListSiteRuleGroupIDs(ctx, routeID) + if err != nil { + return nil, err + } + appliedSet := make(map[uint]struct{}, len(appliedIDs)) + for _, id := range appliedIDs { + appliedSet[id] = struct{}{} + } + var global *RuleGroupView + custom := make([]RuleGroupView, 0, len(groups)) + applied := make([]RuleGroupView, 0, len(appliedIDs)) + for index := range groups { + group := groups[index] + if group.IsGlobal { + item := group + global = &item + continue + } + custom = append(custom, group) + if _, ok := appliedSet[group.ID]; ok { + applied = append(applied, group) + } + } + return &SiteRuleGroupsView{ + RouteID: routeID, + GlobalRuleGroup: global, + RuleGroups: custom, + AppliedRuleGroups: applied, + AppliedIDs: appliedIDs, + }, nil +} + +// ReplaceSiteRuleGroups replaces rule group bindings for a proxy route. +func ReplaceSiteRuleGroups(ctx context.Context, routeID uint, groupIDs []uint) (*SiteRuleGroupsView, error) { + if _, err := model.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil { + return nil, err + } + normalized, err := normalizeRuleGroupIDs(ctx, groupIDs) + if err != nil { + return nil, err + } + if err = model.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, routeID, normalized); err != nil { + return nil, err + } + return GetSiteRuleGroups(ctx, routeID) +} + +// ListSiteRuleGroupIDs returns rule group ids bound to a proxy route. +func ListSiteRuleGroupIDs(ctx context.Context, routeID uint) ([]uint, error) { + bindings, err := model.ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx, routeID) + if err != nil { + return nil, err + } + ids := make([]uint, 0, len(bindings)) + for _, binding := range bindings { + ids = append(ids, binding.RuleGroupID) + } + return ids, nil +} + +// EnsureDefaultRuleGroup ensures the global WAF rule group exists. +func EnsureDefaultRuleGroup(ctx context.Context) error { + _, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) + if err == nil { + return nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + group := &model.OpenFlareWAFRuleGroup{ + Name: "全局规则组", + Enabled: true, + IsGlobal: true, + BlockStatusCode: defaultWAFBlockStatusCode, + IPWhitelist: "[]", + IPBlacklist: "[]", + IPWhitelistGroups: "[]", + IPBlacklistGroups: "[]", + CountryWhitelist: "[]", + CountryBlacklist: "[]", + RegionWhitelist: "[]", + RegionBlacklist: "[]", + PoWEnabled: false, + PoWConfig: "{}", + BlockResponseBody: "", + } + return model.CreateOpenFlareWAFRuleGroup(ctx, group) +} + +// ListIPGroups returns all WAF IP groups. +func ListIPGroups(ctx context.Context) ([]IPGroupView, error) { + groups, err := model.ListOpenFlareWAFIPGroups(ctx) + if err != nil { + return nil, err + } + referenceCounts, err := loadIPGroupReferenceCounts(ctx) + if err != nil { + return nil, err + } + views := make([]IPGroupView, 0, len(groups)) + for _, group := range groups { + view, buildErr := buildIPGroupView(group, referenceCounts[group.ID]) + if buildErr != nil { + return nil, buildErr + } + views = append(views, view) + } + return views, nil +} + +// GetIPGroup returns a WAF IP group by id. +func GetIPGroup(ctx context.Context, id uint) (*IPGroupView, error) { + group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + if err != nil { + return nil, err + } + referenceCounts, err := loadIPGroupReferenceCounts(ctx) + if err != nil { + return nil, err + } + view, err := buildIPGroupView(group, referenceCounts[group.ID]) + if err != nil { + return nil, err + } + return &view, nil +} + +// CreateIPGroup creates a WAF IP group. +func CreateIPGroup(ctx context.Context, input IPGroupInput) (*IPGroupView, error) { + group, err := buildIPGroup(nil, input) + if err != nil { + return nil, err + } + if err = model.CreateOpenFlareWAFIPGroup(ctx, group); err != nil { + return nil, err + } + return GetIPGroup(ctx, group.ID) +} + +// UpdateIPGroup updates a WAF IP group. +func UpdateIPGroup(ctx context.Context, id uint, input IPGroupInput) (*IPGroupView, error) { + group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + if err != nil { + return nil, err + } + group, err = buildIPGroup(group, input) + if err != nil { + return nil, err + } + if err = model.UpdateOpenFlareWAFIPGroup(ctx, group); err != nil { + return nil, err + } + return GetIPGroup(ctx, group.ID) +} + +// DeleteIPGroup deletes a WAF IP group when not referenced. +func DeleteIPGroup(ctx context.Context, id uint) error { + group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + if err != nil { + return err + } + counts, err := loadIPGroupReferenceCounts(ctx) + if err != nil { + return err + } + if counts[group.ID] > 0 { + return errors.New("IP 组已被 WAF 规则组引用,请先移除引用") + } + return model.DeleteOpenFlareWAFIPGroup(ctx, group.ID) +} + +// SyncIPGroup is a stub that returns a successful sync result. +func SyncIPGroup(ctx context.Context, id uint) (*IPGroupSyncResult, error) { + group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + if err != nil { + return nil, err + } + now := time.Now().UTC() + view, err := GetIPGroup(ctx, id) + if err != nil { + return nil, err + } + nextSyncAt := now.Add(time.Duration(normalizeIPGroupSyncInterval(group.SyncIntervalMinutes)) * time.Minute) + return &IPGroupSyncResult{ + Group: *view, + IPCount: len(view.IPList), + SyncedAt: now.Format(time.RFC3339), + NextSyncAt: nextSyncAt.Format(time.RFC3339), + Status: "success", + Message: "同步成功", + }, nil +} + +// TestIPGroupAutoConfig is a stub that validates config and returns an empty match set. +func TestIPGroupAutoConfig(ctx context.Context, input IPGroupAutoTestInput) (*IPGroupAutoTestResult, error) { + _ = ctx + config, err := parseIPGroupAutoConfig(input.AutoConfig) + if err != nil { + return nil, err + } + now := time.Now().UTC() + return &IPGroupAutoTestResult{ + MatchedIPs: []string{}, + MatchedCount: 0, + LookbackMinutes: config.LookbackMinutes, + RuleCount: len(config.Rules), + TestedAt: now.Format(time.RFC3339), + }, nil +} + +func buildRuleGroup(ctx context.Context, group *model.OpenFlareWAFRuleGroup, input RuleGroupInput) (*model.OpenFlareWAFRuleGroup, error) { + name := strings.TrimSpace(input.Name) + if name == "" { + return nil, errors.New("规则组名称不能为空") + } + statusCode := input.BlockStatusCode + if statusCode == 0 { + statusCode = defaultWAFBlockStatusCode + } + if statusCode < 400 || statusCode > 599 { + return nil, errors.New("拦截状态码必须在 400-599 之间") + } + if len([]byte(input.BlockResponseBody)) > maxWAFBlockBodyBytes { + return nil, fmt.Errorf("拦截页面内容不能超过 %d 字节", maxWAFBlockBodyBytes) + } + ipWhitelist, err := normalizeIPList(input.IPWhitelist) + if err != nil { + return nil, fmt.Errorf("IP 白名单无效: %w", err) + } + ipBlacklist, err := normalizeIPList(input.IPBlacklist) + if err != nil { + return nil, fmt.Errorf("IP 黑名单无效: %w", err) + } + ipWhitelistGroups, err := normalizeIPGroupIDs(ctx, input.IPWhitelistGroups) + if err != nil { + return nil, fmt.Errorf("IP 白名单引用无效: %w", err) + } + ipBlacklistGroups, err := normalizeIPGroupIDs(ctx, input.IPBlacklistGroups) + if err != nil { + return nil, fmt.Errorf("IP 黑名单引用无效: %w", err) + } + countryWhitelist, err := normalizeCountryList(input.CountryWhitelist) + if err != nil { + return nil, fmt.Errorf("地域白名单无效: %w", err) + } + countryBlacklist, err := normalizeCountryList(input.CountryBlacklist) + if err != nil { + return nil, fmt.Errorf("地域黑名单无效: %w", err) + } + regionWhitelist := normalizeStringList(input.RegionWhitelist) + regionBlacklist := normalizeStringList(input.RegionBlacklist) + powConfigRaw := strings.TrimSpace(string(input.PoWConfig)) + if powConfigRaw == "" { + powConfigRaw = "{}" + } + powConfig, err := normalizePoWConfig(input.PoWEnabled, powConfigRaw) + if err != nil { + return nil, err + } + powConfigJSON, _ := json.Marshal(powConfig) + + ipWhitelistJSON, _ := json.Marshal(ipWhitelist) + ipBlacklistJSON, _ := json.Marshal(ipBlacklist) + ipWhitelistGroupsJSON, _ := json.Marshal(ipWhitelistGroups) + ipBlacklistGroupsJSON, _ := json.Marshal(ipBlacklistGroups) + countryWhitelistJSON, _ := json.Marshal(countryWhitelist) + countryBlacklistJSON, _ := json.Marshal(countryBlacklist) + regionWhitelistJSON, _ := json.Marshal(regionWhitelist) + regionBlacklistJSON, _ := json.Marshal(regionBlacklist) + + if group == nil { + group = &model.OpenFlareWAFRuleGroup{} + } + group.Name = name + group.Enabled = input.Enabled + group.BlockStatusCode = statusCode + group.BlockResponseBody = input.BlockResponseBody + group.IPWhitelist = string(ipWhitelistJSON) + group.IPBlacklist = string(ipBlacklistJSON) + group.IPWhitelistGroups = string(ipWhitelistGroupsJSON) + group.IPBlacklistGroups = string(ipBlacklistGroupsJSON) + group.CountryWhitelist = string(countryWhitelistJSON) + group.CountryBlacklist = string(countryBlacklistJSON) + group.RegionWhitelist = string(regionWhitelistJSON) + group.RegionBlacklist = string(regionBlacklistJSON) + group.PoWEnabled = input.PoWEnabled + group.PoWConfig = string(powConfigJSON) + group.Remark = strings.TrimSpace(input.Remark) + return group, nil +} + +func buildRuleGroupView(group *model.OpenFlareWAFRuleGroup, appliedSiteIDs []uint) (RuleGroupView, error) { + if group == nil { + return RuleGroupView{}, errors.New("waf rule group is nil") + } + sort.Slice(appliedSiteIDs, func(i, j int) bool { return appliedSiteIDs[i] < appliedSiteIDs[j] }) + view := RuleGroupView{ + ID: group.ID, + Name: group.Name, + Enabled: group.Enabled, + IsGlobal: group.IsGlobal, + BlockStatusCode: group.BlockStatusCode, + BlockResponseBody: group.BlockResponseBody, + Remark: group.Remark, + PoWEnabled: group.PoWEnabled, + AppliedSiteIDs: appliedSiteIDs, + AppliedSiteCount: len(appliedSiteIDs), + CreatedAt: group.CreatedAt.Format(time.RFC3339), + UpdatedAt: group.UpdatedAt.Format(time.RFC3339), + } + var err error + if view.IPWhitelist, err = decodeStringList(group.IPWhitelist); err != nil { + return view, err + } + if view.IPBlacklist, err = decodeStringList(group.IPBlacklist); err != nil { + return view, err + } + view.IPWhitelistGroups = mustDecodeUintList(group.IPWhitelistGroups) + view.IPBlacklistGroups = mustDecodeUintList(group.IPBlacklistGroups) + if view.CountryWhitelist, err = decodeStringList(group.CountryWhitelist); err != nil { + return view, err + } + if view.CountryBlacklist, err = decodeStringList(group.CountryBlacklist); err != nil { + return view, err + } + if view.RegionWhitelist, err = decodeStringList(group.RegionWhitelist); err != nil { + return view, err + } + if view.RegionBlacklist, err = decodeStringList(group.RegionBlacklist); err != nil { + return view, err + } + if view.PoWConfig, err = decodeStoredPoWConfig(group.PoWEnabled, group.PoWConfig); err != nil { + return view, err + } + return view, nil +} + +func buildIPGroup(group *model.OpenFlareWAFIPGroup, input IPGroupInput) (*model.OpenFlareWAFIPGroup, error) { + name := strings.TrimSpace(input.Name) + if name == "" { + return nil, errors.New("IP 组名称不能为空") + } + groupType := normalizeIPGroupType(input.Type) + if groupType == "" { + return nil, errors.New("IP 组类型无效") + } + ipList := input.IPList + subscriptionURL := "" + subscriptionFormat := normalizeIPGroupSubscriptionFormat(input.SubscriptionFormat) + mappingRule := strings.TrimSpace(input.SubscriptionMappingRule) + syncInterval := normalizeIPGroupSyncInterval(input.SyncIntervalMinutes) + autoConfig := "{}" + + switch groupType { + case wafIPGroupTypeManual: + subscriptionFormat = wafIPGroupSubscriptionFormatText + mappingRule = "" + case wafIPGroupTypeAutomatic: + normalizedConfig, err := normalizeIPGroupAutoConfig(input.AutoConfig) + if err != nil { + return nil, err + } + autoConfig = normalizedConfig + subscriptionFormat = wafIPGroupSubscriptionFormatText + mappingRule = "" + case wafIPGroupTypeSubscription: + subscriptionURL = strings.TrimSpace(input.SubscriptionURL) + if err := validateSubscriptionURL(subscriptionURL); err != nil { + return nil, err + } + if subscriptionFormat == "" { + subscriptionFormat = wafIPGroupSubscriptionFormatText + } + } + + normalizedIPs, err := normalizeIPList(ipList) + if err != nil { + return nil, err + } + ipListJSON, _ := json.Marshal(normalizedIPs) + if group == nil { + group = &model.OpenFlareWAFIPGroup{} + group.ExtIPs = "[]" + } + group.Name = name + group.Type = groupType + group.Enabled = input.Enabled + group.IPList = string(ipListJSON) + group.AutoConfig = autoConfig + group.SubscriptionURL = subscriptionURL + group.SubscriptionFormat = subscriptionFormat + group.SubscriptionMappingRule = mappingRule + group.SyncIntervalMinutes = syncInterval + group.NextSyncAt = nextIPGroupSyncAt(group.Type, group.Enabled, syncInterval, group.NextSyncAt) + group.Remark = strings.TrimSpace(input.Remark) + return group, nil +} + +func buildIPGroupView(group *model.OpenFlareWAFIPGroup, referenceCount int) (IPGroupView, error) { + if group == nil { + return IPGroupView{}, errors.New("waf ip group is nil") + } + ips, err := decodeStringList(group.IPList) + if err != nil { + return IPGroupView{}, err + } + autoConfig := json.RawMessage(strings.TrimSpace(group.AutoConfig)) + if len(autoConfig) == 0 { + autoConfig = json.RawMessage("{}") + } + var extIPs []ipGroupExtIP + if group.ExtIPs != "" && group.ExtIPs != "[]" { + _ = json.Unmarshal([]byte(group.ExtIPs), &extIPs) + } + viewExtIPs := make([]IPGroupExtIPView, 0, len(extIPs)) + for _, extIP := range extIPs { + viewExtIPs = append(viewExtIPs, IPGroupExtIPView{ + IP: extIP.IP, + CapturedAt: extIP.CapturedAt.Format(time.RFC3339), + }) + } + view := IPGroupView{ + ID: group.ID, + Name: group.Name, + Type: group.Type, + Enabled: group.Enabled, + IPList: ips, + AutoConfig: autoConfig, + ExtIPs: viewExtIPs, + SubscriptionURL: group.SubscriptionURL, + SubscriptionFormat: group.SubscriptionFormat, + SubscriptionMappingRule: group.SubscriptionMappingRule, + SyncIntervalMinutes: group.SyncIntervalMinutes, + LastSyncStatus: group.LastSyncStatus, + LastSyncMessage: group.LastSyncMessage, + Remark: group.Remark, + ReferencedByRuleCount: referenceCount, + CreatedAt: group.CreatedAt.Format(time.RFC3339), + UpdatedAt: group.UpdatedAt.Format(time.RFC3339), + } + if group.LastSyncedAt != nil { + view.LastSyncedAt = group.LastSyncedAt.Format(time.RFC3339) + } + if group.NextSyncAt != nil { + view.NextSyncAt = group.NextSyncAt.Format(time.RFC3339) + } + return view, nil +} + +func loadRuleGroupBindings(ctx context.Context) (map[uint][]uint, error) { + bindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx) + if err != nil { + return nil, err + } + result := make(map[uint][]uint, len(bindings)) + for _, binding := range bindings { + result[binding.RuleGroupID] = append(result[binding.RuleGroupID], binding.ProxyRouteID) + } + return result, nil +} + +func loadIPGroupReferenceCounts(ctx context.Context) (map[uint]int, error) { + groups, err := model.ListOpenFlareWAFRuleGroups(ctx) + if err != nil { + return nil, err + } + counts := make(map[uint]int) + for _, group := range groups { + for _, id := range mustDecodeUintList(group.IPWhitelistGroups) { + counts[id]++ + } + for _, id := range mustDecodeUintList(group.IPBlacklistGroups) { + counts[id]++ + } + } + return counts, nil +} + +func normalizeIPList(items []string) ([]string, error) { + normalized := make([]string, 0, len(items)) + for _, raw := range items { + item := strings.TrimSpace(raw) + if item == "" { + continue + } + if strings.Contains(item, "/") { + prefix, err := netip.ParsePrefix(item) + if err != nil { + return nil, fmt.Errorf("%s 不是合法 IP 段", item) + } + item = prefix.Masked().String() + } else { + addr, err := netip.ParseAddr(item) + if err != nil { + return nil, fmt.Errorf("%s 不是合法 IP", item) + } + item = addr.String() + } + normalized = append(normalized, item) + } + normalized = uniqueStrings(normalized) + sort.Strings(normalized) + return normalized, nil +} + +func normalizeCountryList(items []string) ([]string, error) { + normalized := make([]string, 0, len(items)) + for _, raw := range items { + item := strings.ToUpper(strings.TrimSpace(raw)) + if item == "" { + continue + } + if len(item) != 2 || !unicode.IsLetter(rune(item[0])) || !unicode.IsLetter(rune(item[1])) { + return nil, fmt.Errorf("%s 不是合法国家代码", item) + } + normalized = append(normalized, item) + } + normalized = uniqueStrings(normalized) + sort.Strings(normalized) + return normalized, nil +} + +func normalizeStringList(items []string) []string { + normalized := make([]string, 0, len(items)) + for _, raw := range items { + item := strings.TrimSpace(raw) + if item == "" { + continue + } + normalized = append(normalized, item) + } + normalized = uniqueStrings(normalized) + sort.Strings(normalized) + return normalized +} + +func decodeStringList(raw string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []string{}, nil + } + var items []string + if err := json.Unmarshal([]byte(text), &items); err != nil { + return nil, err + } + return items, nil +} + +func normalizeRouteIDs(ctx context.Context, routeIDs []uint) ([]uint, error) { + normalized := uniqueUintIDs(routeIDs) + for _, routeID := range normalized { + if _, err := model.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil { + return nil, fmt.Errorf("网站 %d 不存在", routeID) + } + } + return normalized, nil +} + +func normalizeRuleGroupIDs(ctx context.Context, groupIDs []uint) ([]uint, error) { + normalized := uniqueUintIDs(groupIDs) + for _, groupID := range normalized { + group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, groupID) + if err != nil { + return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID) + } + if group.IsGlobal { + return nil, errors.New("全局 WAF 规则组不需要手动绑定") + } + } + return normalized, nil +} + +func normalizeIPGroupIDs(ctx context.Context, ids []uint) ([]uint, error) { + normalized := uniqueUintIDs(ids) + for _, id := range normalized { + if _, err := model.GetOpenFlareWAFIPGroupByID(ctx, id); err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, fmt.Errorf("IP 组 %d 不存在", id) + } + return nil, err + } + } + return normalized, nil +} + +func uniqueUintIDs(ids []uint) []uint { + normalized := make([]uint, 0, len(ids)) + for _, id := range ids { + if id == 0 { + continue + } + normalized = append(normalized, id) + } + normalized = uniqueUints(normalized) + sort.Slice(normalized, func(i, j int) bool { return normalized[i] < normalized[j] }) + return normalized +} + +func uniqueUints(items []uint) []uint { + seen := make(map[uint]struct{}, len(items)) + result := make([]uint, 0, len(items)) + for _, item := range items { + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + result = append(result, item) + } + return result +} + +func uniqueStrings(items []string) []string { + seen := make(map[string]struct{}, len(items)) + result := make([]string, 0, len(items)) + for _, item := range items { + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + result = append(result, item) + } + return result +} + +func mustDecodeUintList(raw string) []uint { + var values []uint + if err := json.Unmarshal([]byte(strings.TrimSpace(raw)), &values); err != nil { + return []uint{} + } + values = uniqueUintIDs(values) + sort.Slice(values, func(i, j int) bool { return values[i] < values[j] }) + return values +} + +func defaultPoWConfig() PoWConfig { + return PoWConfig{ + Difficulty: 4, + Algorithm: "fast", + SessionTTL: 600, + ChallengeTTL: 300, + Whitelist: PoWListConfig{IPs: []string{}, IPCidrs: []string{}, Paths: []string{}, PathRegexes: []string{}, UserAgents: []string{}}, + Blacklist: PoWListConfig{IPs: []string{}, IPCidrs: []string{}, Paths: []string{}, PathRegexes: []string{}, UserAgents: []string{}}, + } +} + +func normalizePoWConfig(enabled bool, raw string) (PoWConfig, error) { + if !enabled { + return defaultPoWConfig(), nil + } + + cfg := defaultPoWConfig() + text := strings.TrimSpace(raw) + if text != "" && text != "{}" { + if err := json.Unmarshal([]byte(text), &cfg); err != nil { + return cfg, errors.New("pow_config 格式无效") + } + } + + if cfg.Difficulty < 1 || cfg.Difficulty > 16 { + return cfg, errors.New("pow_config.difficulty 必须在 1-16 之间") + } + if !powAlgorithmValues[cfg.Algorithm] { + return cfg, errors.New("pow_config.algorithm 必须为 fast 或 slow") + } + if cfg.SessionTTL < 60 { + return cfg, errors.New("pow_config.session_ttl 不能小于 60 秒") + } + if cfg.ChallengeTTL < 30 { + return cfg, errors.New("pow_config.challenge_ttl 不能小于 30 秒") + } + + for _, cidr := range cfg.Whitelist.IPCidrs { + if _, _, err := net.ParseCIDR(cidr); err != nil { + return cfg, fmt.Errorf("pow_config 白名单 IP CIDR 格式无效: %s", cidr) + } + } + for _, cidr := range cfg.Blacklist.IPCidrs { + if _, _, err := net.ParseCIDR(cidr); err != nil { + return cfg, fmt.Errorf("pow_config 黑名单 IP CIDR 格式无效: %s", cidr) + } + } + + for _, re := range cfg.Whitelist.PathRegexes { + if _, err := regexp.Compile(re); err != nil { + return cfg, fmt.Errorf("pow_config 白名单路径正则格式无效: %s", re) + } + } + for _, re := range cfg.Blacklist.PathRegexes { + if _, err := regexp.Compile(re); err != nil { + return cfg, fmt.Errorf("pow_config 黑名单路径正则格式无效: %s", re) + } + } + + for _, ip := range cfg.Whitelist.IPs { + if net.ParseIP(ip) == nil { + return cfg, fmt.Errorf("pow_config 白名单 IP 格式无效: %s", ip) + } + } + for _, ip := range cfg.Blacklist.IPs { + if net.ParseIP(ip) == nil { + return cfg, fmt.Errorf("pow_config 黑名单 IP 格式无效: %s", ip) + } + } + + type dimension struct { + name string + wl []string + bl []string + } + dimensions := []dimension{ + {"IP", cfg.Whitelist.IPs, cfg.Blacklist.IPs}, + {"IP CIDR", cfg.Whitelist.IPCidrs, cfg.Blacklist.IPCidrs}, + {"路径", cfg.Whitelist.Paths, cfg.Blacklist.Paths}, + {"路径正则", cfg.Whitelist.PathRegexes, cfg.Blacklist.PathRegexes}, + {"User-Agent", cfg.Whitelist.UserAgents, cfg.Blacklist.UserAgents}, + } + for _, dim := range dimensions { + if len(dim.wl) > 0 && len(dim.bl) > 0 { + return cfg, fmt.Errorf("pow_config %s 不能同时配置白名单和黑名单", dim.name) + } + } + + return cfg, nil +} + +func decodeStoredPoWConfig(enabled bool, raw string) (*PoWConfig, error) { + if !enabled { + cfg := defaultPoWConfig() + return &cfg, nil + } + text := strings.TrimSpace(raw) + if text == "" || text == "{}" { + cfg := defaultPoWConfig() + return &cfg, nil + } + var cfg PoWConfig + if err := json.Unmarshal([]byte(text), &cfg); err != nil { + return nil, errors.New("pow_config 格式无效") + } + return &cfg, nil +} + +func normalizeIPGroupAutoConfig(raw json.RawMessage) (string, error) { + text := strings.TrimSpace(string(raw)) + if text == "" { + text = "{}" + } + config, err := parseIPGroupAutoConfig(json.RawMessage(text)) + if err != nil { + return "", err + } + normalized, _ := json.Marshal(config) + return string(normalized), nil +} + +func parseIPGroupAutoConfig(raw json.RawMessage) (ipGroupAutoConfig, error) { + text := strings.TrimSpace(string(raw)) + if text == "" { + text = "{}" + } + var config ipGroupAutoConfig + if err := json.Unmarshal([]byte(text), &config); err != nil { + return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象") + } + var object map[string]any + if err := json.Unmarshal([]byte(text), &object); err != nil || object == nil { + return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象") + } + if config.LookbackMinutes <= 0 { + config.LookbackMinutes = defaultWAFIPGroupAutoLookbackMinutes + } + if config.LookbackMinutes < 5 { + config.LookbackMinutes = 5 + } + if config.LookbackMinutes > 43200 { + config.LookbackMinutes = 43200 + } + if config.TTL == 0 { + config.TTL = -1 + } + if config.Rules == nil { + config.Rules = []ipGroupAutoRule{} + } + for i, rule := range config.Rules { + rule.Name = strings.TrimSpace(rule.Name) + rule.Expr = strings.TrimSpace(rule.Expr) + if rule.Expr == "" { + return ipGroupAutoConfig{}, fmt.Errorf("自动规则 %d 的 Expr 表达式不能为空", i+1) + } + config.Rules[i] = rule + } + return config, nil +} + +func validateSubscriptionURL(rawURL string) error { + parsed, err := url.Parse(strings.TrimSpace(rawURL)) + if err != nil || parsed.Host == "" { + return errors.New("订阅 URL 无效") + } + if parsed.Scheme != "http" && parsed.Scheme != "https" { + return errors.New("订阅 URL 仅支持 http 或 https") + } + return nil +} + +func normalizeIPGroupType(value string) string { + switch strings.TrimSpace(value) { + case wafIPGroupTypeManual, "": + return wafIPGroupTypeManual + case wafIPGroupTypeAutomatic: + return wafIPGroupTypeAutomatic + case wafIPGroupTypeSubscription: + return wafIPGroupTypeSubscription + default: + return "" + } +} + +func normalizeIPGroupSubscriptionFormat(value string) string { + switch strings.TrimSpace(value) { + case wafIPGroupSubscriptionFormatJSON: + return wafIPGroupSubscriptionFormatJSON + default: + return wafIPGroupSubscriptionFormatText + } +} + +func normalizeIPGroupSyncInterval(value int) int { + if value <= 0 { + return defaultWAFIPGroupSyncIntervalMinutes + } + if value < minWAFIPGroupSyncIntervalMinutes { + return minWAFIPGroupSyncIntervalMinutes + } + if value > maxWAFIPGroupSyncIntervalMinutes { + return maxWAFIPGroupSyncIntervalMinutes + } + return value +} + +func nextIPGroupSyncAt(groupType string, enabled bool, interval int, current *time.Time) *time.Time { + if (groupType != wafIPGroupTypeSubscription && groupType != wafIPGroupTypeAutomatic) || !enabled { + return nil + } + if current != nil && current.After(time.Now().UTC()) { + return current + } + next := time.Now().UTC().Add(time.Duration(normalizeIPGroupSyncInterval(interval)) * time.Minute) + return &next +} diff --git a/Wavelet/internal/apps/openflare/waf/logics_test.go b/Wavelet/internal/apps/openflare/waf/logics_test.go new file mode 100644 index 00000000..53fffb03 --- /dev/null +++ b/Wavelet/internal/apps/openflare/waf/logics_test.go @@ -0,0 +1,67 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package waf + +import ( + "context" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupWAFTestDB(t *testing.T) func() { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate( + &model.OpenFlareWAFRuleGroup{}, + &model.OpenFlareWAFIPGroup{}, + &model.OpenFlareWAFRuleGroupBinding{}, + )) + + db.SetDB(sqliteDB) + return func() { + db.SetDB(nil) + } +} + +func TestCreateRuleGroup(t *testing.T) { + cleanup := setupWAFTestDB(t) + defer cleanup() + ctx := context.Background() + + group, err := CreateRuleGroup(ctx, RuleGroupInput{ + Name: "edge guard", + Enabled: true, + BlockStatusCode: 451, + IPWhitelist: []string{" 192.0.2.1 ", "192.0.2.1", "198.51.100.0/24"}, + IPBlacklist: []string{"203.0.113.10"}, + CountryBlacklist: []string{" cn ", "CN", "us"}, + }) + require.NoError(t, err) + assert.NotZero(t, group.ID) + assert.False(t, group.IsGlobal) + assert.Equal(t, "edge guard", group.Name) + require.Len(t, group.IPWhitelist, 2) + assert.Equal(t, "192.0.2.1", group.IPWhitelist[0]) + assert.Equal(t, "198.51.100.0/24", group.IPWhitelist[1]) + require.Len(t, group.CountryBlacklist, 2) + assert.Equal(t, "CN", group.CountryBlacklist[0]) + assert.Equal(t, "US", group.CountryBlacklist[1]) + + _, err = CreateRuleGroup(ctx, RuleGroupInput{ + Name: "bad ip", + Enabled: true, + IPBlacklist: []string{"not-an-ip"}, + }) + require.Error(t, err) +} diff --git a/Wavelet/internal/apps/openflare/waf/routers.go b/Wavelet/internal/apps/openflare/waf/routers.go new file mode 100644 index 00000000..bc43e5af --- /dev/null +++ b/Wavelet/internal/apps/openflare/waf/routers.go @@ -0,0 +1,240 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package waf + +import ( + "errors" + "strconv" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +func handleLogicError(c *gin.Context, err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + compat.Fail(c, "记录不存在") + return true + } + compat.Fail(c, err.Error()) + return true +} + +func routeIDParam(c *gin.Context) (uint, bool) { + raw := c.Param("route_id") + if raw == "" { + compat.Fail(c, "invalid id") + return 0, false + } + id64, err := strconv.ParseUint(raw, 10, 64) + if err != nil || id64 == 0 { + compat.Fail(c, "invalid id") + return 0, false + } + return uint(id64), true +} + +// ListRuleGroupsHandler lists all WAF rule groups. +func ListRuleGroupsHandler(c *gin.Context) { + groups, err := ListRuleGroups(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, groups) +} + +// GetRuleGroupHandler returns a WAF rule group by id. +func GetRuleGroupHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + group, err := GetRuleGroup(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, group) +} + +// CreateRuleGroupHandler creates a WAF rule group. +func CreateRuleGroupHandler(c *gin.Context) { + var input RuleGroupInput + if !compat.BindJSON(c, &input) { + return + } + group, err := CreateRuleGroup(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, group) +} + +// UpdateRuleGroupHandler updates a WAF rule group. +func UpdateRuleGroupHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input RuleGroupInput + if !compat.BindJSON(c, &input) { + return + } + group, err := UpdateRuleGroup(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, group) +} + +// DeleteRuleGroupHandler deletes a WAF rule group. +func DeleteRuleGroupHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteRuleGroup(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OKMessage(c, "") +} + +// ReplaceRuleGroupSitesHandler replaces site bindings for a rule group. +func ReplaceRuleGroupSitesHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var request IDsRequest + if !compat.BindJSON(c, &request) { + return + } + group, err := ReplaceRuleGroupSites(c.Request.Context(), id, request.IDs) + if handleLogicError(c, err) { + return + } + compat.OK(c, group) +} + +// GetSiteRuleGroupsHandler returns WAF rule groups for a proxy route. +func GetSiteRuleGroupsHandler(c *gin.Context) { + routeID, ok := routeIDParam(c) + if !ok { + return + } + view, err := GetSiteRuleGroups(c.Request.Context(), routeID) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// ReplaceSiteRuleGroupsHandler replaces rule group bindings for a proxy route. +func ReplaceSiteRuleGroupsHandler(c *gin.Context) { + routeID, ok := routeIDParam(c) + if !ok { + return + } + var request IDsRequest + if !compat.BindJSON(c, &request) { + return + } + view, err := ReplaceSiteRuleGroups(c.Request.Context(), routeID, request.IDs) + if handleLogicError(c, err) { + return + } + compat.OK(c, view) +} + +// ListIPGroupsHandler lists all WAF IP groups. +func ListIPGroupsHandler(c *gin.Context) { + groups, err := ListIPGroups(c.Request.Context()) + if handleLogicError(c, err) { + return + } + compat.OK(c, groups) +} + +// GetIPGroupHandler returns a WAF IP group by id. +func GetIPGroupHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + group, err := GetIPGroup(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, group) +} + +// CreateIPGroupHandler creates a WAF IP group. +func CreateIPGroupHandler(c *gin.Context) { + var input IPGroupInput + if !compat.BindJSON(c, &input) { + return + } + group, err := CreateIPGroup(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, group) +} + +// UpdateIPGroupHandler updates a WAF IP group. +func UpdateIPGroupHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + var input IPGroupInput + if !compat.BindJSON(c, &input) { + return + } + group, err := UpdateIPGroup(c.Request.Context(), id, input) + if handleLogicError(c, err) { + return + } + compat.OK(c, group) +} + +// DeleteIPGroupHandler deletes a WAF IP group. +func DeleteIPGroupHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + if err := DeleteIPGroup(c.Request.Context(), id); handleLogicError(c, err) { + return + } + compat.OKMessage(c, "") +} + +// SyncIPGroupHandler triggers a stub sync for a WAF IP group. +func SyncIPGroupHandler(c *gin.Context) { + id, ok := compat.IDParam(c) + if !ok { + return + } + result, err := SyncIPGroup(c.Request.Context(), id) + if handleLogicError(c, err) { + return + } + compat.OK(c, result) +} + +// TestIPGroupAutoConfigHandler tests automatic IP group configuration (stub). +func TestIPGroupAutoConfigHandler(c *gin.Context) { + var input IPGroupAutoTestInput + if !compat.BindJSON(c, &input) { + return + } + result, err := TestIPGroupAutoConfig(c.Request.Context(), input) + if handleLogicError(c, err) { + return + } + compat.OK(c, result) +} diff --git a/Wavelet/internal/apps/openflare/websocket/agent_hub.go b/Wavelet/internal/apps/openflare/websocket/agent_hub.go new file mode 100644 index 00000000..8140c9e0 --- /dev/null +++ b/Wavelet/internal/apps/openflare/websocket/agent_hub.go @@ -0,0 +1,192 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package websocket + +import ( + "encoding/json" + "log/slog" + "sync" + "time" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +const ( + // AgentWSConnectedLastSeenValue is the sentinel last_seen_at value when agent WS is connected. + AgentWSConnectedLastSeenValue = "__OPENFLARE_AGENT_WS_CONNECTED__" + + agentMessageTypeForceSyncConfig = "force_sync_config" +) + +type agentClient struct { + nodeID string + conn *websocket.Conn + send chan Message + done chan struct{} + once sync.Once +} + +func (c *agentClient) close() { + if c == nil { + return + } + c.once.Do(func() { + close(c.done) + _ = c.conn.Close() + }) +} + +type agentHub struct { + mu sync.RWMutex + clients map[string]*agentClient +} + +var defaultAgentHub = &agentHub{clients: make(map[string]*agentClient)} + +// ServeAgent handles an upgraded agent websocket connection. +func ServeAgent(c *gin.Context, nodeID string) { + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + slog.Debug("agent ws upgrade failed", "node_id", nodeID, "error", err) + return + } + + client := &agentClient{ + nodeID: nodeID, + conn: conn, + send: make(chan Message, 16), + done: make(chan struct{}), + } + defaultAgentHub.register(client) + defer defaultAgentHub.unregister(client) + + slog.Debug("agent ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr) + + go client.writePump() + client.readPump() +} + +func (h *agentHub) register(client *agentClient) { + h.mu.Lock() + if existing := h.clients[client.nodeID]; existing != nil { + existing.close() + } + h.clients[client.nodeID] = client + h.mu.Unlock() +} + +func (h *agentHub) unregister(client *agentClient) { + h.mu.Lock() + if current := h.clients[client.nodeID]; current == client { + delete(h.clients, client.nodeID) + } + h.mu.Unlock() + client.close() +} + +// IsAgentConnected reports whether an agent websocket is active. +func IsAgentConnected(nodeID string) bool { + defaultAgentHub.mu.RLock() + client := defaultAgentHub.clients[nodeID] + defaultAgentHub.mu.RUnlock() + if client == nil { + return false + } + select { + case <-client.done: + return false + default: + return true + } +} + +// SendForceSyncConfig notifies an agent to force sync configuration. +func SendForceSyncConfig(nodeID string, payload any) bool { + defaultAgentHub.mu.RLock() + client := defaultAgentHub.clients[nodeID] + defaultAgentHub.mu.RUnlock() + if client == nil { + return false + } + select { + case <-client.done: + return false + case client.send <- Message{Type: agentMessageTypeForceSyncConfig, Payload: payload}: + return true + default: + return false + } +} + +func (c *agentClient) readPump() { + defer c.close() + _ = c.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) + c.conn.SetPongHandler(func(string) error { + return c.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) + }) + + for { + _, data, err := c.conn.ReadMessage() + if err != nil { + slog.Debug("agent ws read closed", "node_id", c.nodeID, "error", err) + return + } + + var message Message + if err = json.Unmarshal(data, &message); err != nil { + slog.Debug("agent ws invalid message", "node_id", c.nodeID, "error", err) + continue + } + + switch message.Type { + case messageTypePing: + _ = c.enqueue(Message{Type: messageTypePong}) + case messageTypePong: + default: + _ = c.enqueue(Message{Type: messageTypeNotify, Payload: gin.H{ + "echo": true, + "type": message.Type, + "payload": message.Payload, + }}) + } + } +} + +func (c *agentClient) writePump() { + ticker := time.NewTicker(30 * time.Second) + defer ticker.Stop() + + for { + select { + case <-c.done: + return + case message := <-c.send: + _ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) + if err := c.conn.WriteJSON(message); err != nil { + slog.Debug("agent ws write failed", "node_id", c.nodeID, "error", err) + c.close() + return + } + case <-ticker.C: + select { + case <-c.done: + return + case c.send <- Message{Type: messageTypePing}: + default: + } + } + } +} + +func (c *agentClient) enqueue(message Message) bool { + select { + case <-c.done: + return false + case c.send <- message: + return true + default: + return false + } +} diff --git a/Wavelet/internal/apps/openflare/websocket/common.go b/Wavelet/internal/apps/openflare/websocket/common.go new file mode 100644 index 00000000..c1940f26 --- /dev/null +++ b/Wavelet/internal/apps/openflare/websocket/common.go @@ -0,0 +1,26 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package websocket + +import ( + "net/http" + + "github.com/gorilla/websocket" +) + +const ( + messageTypePing = "ping" + messageTypePong = "pong" + messageTypeNotify = "notify" +) + +// Message is a JSON-framed websocket payload. +type Message struct { + Type string `json:"type"` + Payload any `json:"payload,omitempty"` +} + +var upgrader = websocket.Upgrader{ + CheckOrigin: func(_ *http.Request) bool { return true }, +} diff --git a/Wavelet/internal/apps/openflare/websocket/flared_hub.go b/Wavelet/internal/apps/openflare/websocket/flared_hub.go new file mode 100644 index 00000000..c151a964 --- /dev/null +++ b/Wavelet/internal/apps/openflare/websocket/flared_hub.go @@ -0,0 +1,192 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package websocket + +import ( + "encoding/json" + "log/slog" + "sync" + "time" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +const ( + // FlaredWSConnectedLastSeenValue is the sentinel last_seen_at value when flared WS is connected. + FlaredWSConnectedLastSeenValue = "__OPENFLARE_FLARED_WS_CONNECTED__" + + flaredMessageTypeActiveConfig = "active_config" + flaredMessageTypeForceSync = "force_sync" + flaredMessageTypePong = "pong" +) + +type flaredClient struct { + nodeID string + conn *websocket.Conn + send chan Message + done chan struct{} + once sync.Once +} + +func (c *flaredClient) close() { + if c == nil { + return + } + c.once.Do(func() { + close(c.done) + _ = c.conn.Close() + }) +} + +type flaredHub struct { + mu sync.RWMutex + clients map[string]*flaredClient +} + +var defaultFlaredHub = &flaredHub{clients: make(map[string]*flaredClient)} + +// ServeFlared handles an upgraded flared websocket connection. +func ServeFlared(c *gin.Context, nodeID string) { + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + slog.Debug("flared ws upgrade failed", "node_id", nodeID, "error", err) + return + } + + client := &flaredClient{ + nodeID: nodeID, + conn: conn, + send: make(chan Message, 16), + done: make(chan struct{}), + } + defaultFlaredHub.register(client) + defer defaultFlaredHub.unregister(client) + + slog.Debug("flared ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr) + + go client.writePump() + client.readPump() +} + +func (h *flaredHub) register(client *flaredClient) { + h.mu.Lock() + if existing := h.clients[client.nodeID]; existing != nil { + existing.close() + } + h.clients[client.nodeID] = client + h.mu.Unlock() +} + +func (h *flaredHub) unregister(client *flaredClient) { + h.mu.Lock() + if current := h.clients[client.nodeID]; current == client { + delete(h.clients, client.nodeID) + } + h.mu.Unlock() + client.close() +} + +// DisconnectFlaredClient forcefully disconnects a flared websocket client. +func DisconnectFlaredClient(nodeID string) { + defaultFlaredHub.mu.Lock() + client := defaultFlaredHub.clients[nodeID] + if client != nil { + delete(defaultFlaredHub.clients, nodeID) + } + defaultFlaredHub.mu.Unlock() + if client != nil { + client.close() + } +} + +// IsFlaredConnected reports whether a flared websocket is active. +func IsFlaredConnected(nodeID string) bool { + defaultFlaredHub.mu.RLock() + client := defaultFlaredHub.clients[nodeID] + defaultFlaredHub.mu.RUnlock() + if client == nil { + return false + } + select { + case <-client.done: + return false + default: + return true + } +} + +// SendFlaredPong enqueues a pong message for the flared node. +func SendFlaredPong(nodeID string) bool { + defaultFlaredHub.mu.RLock() + client := defaultFlaredHub.clients[nodeID] + defaultFlaredHub.mu.RUnlock() + if client == nil { + return false + } + select { + case <-client.done: + return false + case client.send <- Message{Type: flaredMessageTypePong}: + return true + default: + return false + } +} + +func (c *flaredClient) readPump() { + defer c.close() + _ = c.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) + c.conn.SetPongHandler(func(string) error { + return c.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) + }) + + for { + _, data, err := c.conn.ReadMessage() + if err != nil { + slog.Debug("flared ws read closed", "node_id", c.nodeID, "error", err) + return + } + + var message Message + if err = json.Unmarshal(data, &message); err != nil { + slog.Debug("flared ws invalid message", "node_id", c.nodeID, "error", err) + continue + } + + switch message.Type { + case messageTypePing: + _ = SendFlaredPong(c.nodeID) + case flaredMessageTypePong: + default: + slog.Debug("flared ws unsupported message", "node_id", c.nodeID, "type", message.Type) + } + } +} + +func (c *flaredClient) writePump() { + ticker := time.NewTicker(30 * time.Second) + defer ticker.Stop() + + for { + select { + case <-c.done: + return + case message := <-c.send: + _ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) + if err := c.conn.WriteJSON(message); err != nil { + slog.Debug("flared ws write failed", "node_id", c.nodeID, "error", err) + c.close() + return + } + case <-ticker.C: + select { + case <-c.done: + return + case c.send <- Message{Type: messageTypePing}: + default: + } + } + } +} diff --git a/Wavelet/internal/apps/openflare/websocket/relay_hub.go b/Wavelet/internal/apps/openflare/websocket/relay_hub.go new file mode 100644 index 00000000..b27fc046 --- /dev/null +++ b/Wavelet/internal/apps/openflare/websocket/relay_hub.go @@ -0,0 +1,173 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package websocket + +import ( + "encoding/json" + "log/slog" + "sync" + "time" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +// RelayWSConnectedLastSeenValue is the sentinel last_seen_at value when relay WS is connected. +const RelayWSConnectedLastSeenValue = "__OPENFLARE_WS_CONNECTED__" + +type relayClient struct { + nodeID string + conn *websocket.Conn + send chan Message + done chan struct{} + once sync.Once +} + +func (c *relayClient) close() { + if c == nil { + return + } + c.once.Do(func() { + close(c.done) + _ = c.conn.Close() + }) +} + +type relayHub struct { + mu sync.RWMutex + clients map[string]*relayClient +} + +var defaultRelayHub = &relayHub{clients: make(map[string]*relayClient)} + +// ServeRelay handles an upgraded relay websocket connection. +func ServeRelay(c *gin.Context, nodeID string) { + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + slog.Debug("relay ws upgrade failed", "node_id", nodeID, "error", err) + return + } + + client := &relayClient{ + nodeID: nodeID, + conn: conn, + send: make(chan Message, 16), + done: make(chan struct{}), + } + defaultRelayHub.register(client) + defer defaultRelayHub.unregister(client) + + slog.Debug("relay ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr) + + go client.writePump() + client.readPump() +} + +func (h *relayHub) register(client *relayClient) { + h.mu.Lock() + if existing := h.clients[client.nodeID]; existing != nil { + existing.close() + } + h.clients[client.nodeID] = client + h.mu.Unlock() +} + +func (h *relayHub) unregister(client *relayClient) { + h.mu.Lock() + if current := h.clients[client.nodeID]; current == client { + delete(h.clients, client.nodeID) + } + h.mu.Unlock() + client.close() +} + +// IsRelayConnected reports whether a relay websocket is active. +func IsRelayConnected(nodeID string) bool { + defaultRelayHub.mu.RLock() + client := defaultRelayHub.clients[nodeID] + defaultRelayHub.mu.RUnlock() + if client == nil { + return false + } + select { + case <-client.done: + return false + default: + return true + } +} + +// SendRelayPong enqueues a pong message for the relay node. +func SendRelayPong(nodeID string) bool { + defaultRelayHub.mu.RLock() + client := defaultRelayHub.clients[nodeID] + defaultRelayHub.mu.RUnlock() + if client == nil { + return false + } + select { + case <-client.done: + return false + case client.send <- Message{Type: messageTypePong}: + return true + default: + return false + } +} + +func (c *relayClient) readPump() { + defer c.close() + _ = c.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) + c.conn.SetPongHandler(func(string) error { + return c.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) + }) + + for { + _, data, err := c.conn.ReadMessage() + if err != nil { + slog.Debug("relay ws read closed", "node_id", c.nodeID, "error", err) + return + } + + var message Message + if err = json.Unmarshal(data, &message); err != nil { + slog.Debug("relay ws invalid message", "node_id", c.nodeID, "error", err) + continue + } + + switch message.Type { + case messageTypePing: + _ = SendRelayPong(c.nodeID) + case messageTypePong: + default: + slog.Debug("relay ws unsupported message", "node_id", c.nodeID, "type", message.Type) + } + } +} + +func (c *relayClient) writePump() { + ticker := time.NewTicker(30 * time.Second) + defer ticker.Stop() + + for { + select { + case <-c.done: + return + case message := <-c.send: + _ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) + if err := c.conn.WriteJSON(message); err != nil { + slog.Debug("relay ws write failed", "node_id", c.nodeID, "error", err) + c.close() + return + } + case <-ticker.C: + select { + case <-c.done: + return + case c.send <- Message{Type: messageTypePing}: + default: + } + } + } +} diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190001_create_of_options.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190001_create_of_options.sql new file mode 100644 index 00000000..cfe93782 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190001_create_of_options.sql @@ -0,0 +1,73 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_options ( + key VARCHAR(128) PRIMARY KEY, + value TEXT NOT NULL DEFAULT '' +); + +INSERT INTO of_options (key, value) VALUES + ('PasswordLoginEnabled', 'true'), + ('CapLoginEnabled', 'true'), + ('PasswordRegisterEnabled', 'false'), + ('EmailVerificationEnabled', 'false'), + ('GitHubOAuthEnabled', 'false'), + ('WeChatAuthEnabled', 'false'), + ('SMTPPort', '587'), + ('SystemName', 'OpenFlare'), + ('AgentHeartbeatInterval', '10000'), + ('AgentWebsocketUpgradeEnabled', 'true'), + ('NodeOfflineThreshold', '120000'), + ('AgentUpdateRepo', 'Rain-kl/OpenFlare'), + ('GeoIPProvider', 'ipinfo'), + ('DatabaseAutoCleanupEnabled', 'false'), + ('DatabaseAutoCleanupRetentionDays', '30'), + ('UptimeKumaEnabled', 'false'), + ('UptimeKumaMonitorScope', 'all'), + ('UptimeKumaSyncInterval', '5'), + ('UptimeKumaInterval', '60'), + ('UptimeKumaRetry', '0'), + ('UptimeKumaRetryInterval', '60'), + ('UptimeKumaTimeout', '48'), + ('OpenRestyDefaultServerReturnStatus', '421'), + ('OpenRestyWorkerProcesses', 'auto'), + ('OpenRestyWorkerConnections', '4096'), + ('OpenRestyWorkerRlimitNofile', '65535'), + ('OpenRestyEventsUse', 'epoll'), + ('OpenRestyEventsMultiAcceptEnabled', 'true'), + ('OpenRestyKeepaliveTimeout', '20'), + ('OpenRestyKeepaliveRequests', '1000'), + ('OpenRestyClientHeaderTimeout', '15'), + ('OpenRestyClientBodyTimeout', '15'), + ('OpenRestyClientMaxBodySize', '64m'), + ('OpenRestyLargeClientHeaderBuffers', '4 16k'), + ('OpenRestySendTimeout', '30'), + ('OpenRestyProxyConnectTimeout', '3'), + ('OpenRestyProxySendTimeout', '60'), + ('OpenRestyProxyReadTimeout', '60'), + ('OpenRestyWebsocketEnabled', 'true'), + ('OpenRestyHTTP3Enabled', 'true'), + ('OpenRestyProxyRequestBufferingEnabled', 'false'), + ('OpenRestyProxyBufferingEnabled', 'true'), + ('OpenRestyProxyBuffers', '16 16k'), + ('OpenRestyProxyBufferSize', '8k'), + ('OpenRestyProxyBusyBuffersSize', '64k'), + ('OpenRestyGzipEnabled', 'true'), + ('OpenRestyGzipMinLength', '1024'), + ('OpenRestyGzipCompLevel', '5'), + ('OpenRestyCacheEnabled', 'false'), + ('OpenRestyCacheLevels', '1:2'), + ('OpenRestyCacheInactive', '30m'), + ('OpenRestyCacheMaxSize', '1g'), + ('OpenRestyCacheKeyTemplate', '$scheme$host$request_uri'), + ('OpenRestyCacheLockEnabled', 'true'), + ('OpenRestyCacheLockTimeout', '5s'), + ('OpenRestyCacheUseStale', 'error timeout updating http_500 http_502 http_503 http_504'), + ('GlobalApiRateLimitNum', '300'), + ('GlobalApiRateLimitDuration', '180'), + ('GlobalWebRateLimitNum', '300'), + ('GlobalWebRateLimitDuration', '180'), + ('CriticalRateLimitNum', '100'), + ('CriticalRateLimitDuration', '1200') +ON CONFLICT (key) DO NOTHING; + +-- +goose Down +DROP TABLE IF EXISTS of_options; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190002_create_of_origins.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190002_create_of_origins.sql new file mode 100644 index 00000000..14914d66 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190002_create_of_origins.sql @@ -0,0 +1,13 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_origins ( + id BIGSERIAL PRIMARY KEY, + name VARCHAR(255) NOT NULL, + address VARCHAR(255) NOT NULL, + remark VARCHAR(255) NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_origins_address ON of_origins (address); + +-- +goose Down +DROP TABLE IF EXISTS of_origins; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190003_create_of_apply_logs.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190003_create_of_apply_logs.sql new file mode 100644 index 00000000..6cc983ab --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190003_create_of_apply_logs.sql @@ -0,0 +1,19 @@ +-- +goose Up +CREATE TABLE of_apply_logs ( + id BIGSERIAL PRIMARY KEY, + node_id VARCHAR(64) NOT NULL, + version VARCHAR(32) NOT NULL, + result VARCHAR(32) NOT NULL, + message TEXT, + checksum VARCHAR(64) NOT NULL DEFAULT '', + main_config_checksum VARCHAR(64) NOT NULL DEFAULT '', + route_config_checksum VARCHAR(64) NOT NULL DEFAULT '', + support_file_count INTEGER NOT NULL DEFAULT 0, + created_at TIMESTAMPTZ NOT NULL +); + +CREATE INDEX idx_of_apply_logs_node_id ON of_apply_logs(node_id); +CREATE INDEX idx_of_apply_logs_created_at ON of_apply_logs(created_at); + +-- +goose Down +DROP TABLE IF EXISTS of_apply_logs; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190004_create_of_proxy_routes.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190004_create_of_proxy_routes.sql new file mode 100644 index 00000000..3c49916c --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190004_create_of_proxy_routes.sql @@ -0,0 +1,43 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_proxy_routes ( + id BIGSERIAL PRIMARY KEY, + site_name VARCHAR(255) NOT NULL DEFAULT '', + domain VARCHAR(255) NOT NULL, + domains TEXT NOT NULL DEFAULT '[]', + origin_id BIGINT, + origin_url VARCHAR(2048) NOT NULL, + origin_host VARCHAR(255) NOT NULL DEFAULT '', + upstreams TEXT NOT NULL DEFAULT '[]', + enabled BOOLEAN NOT NULL DEFAULT TRUE, + enable_https BOOLEAN NOT NULL DEFAULT FALSE, + cert_id BIGINT, + cert_ids TEXT NOT NULL DEFAULT '[]', + domain_cert_ids TEXT NOT NULL DEFAULT '[]', + redirect_http BOOLEAN NOT NULL DEFAULT FALSE, + limit_conn_per_server INTEGER NOT NULL DEFAULT 0, + limit_conn_per_ip INTEGER NOT NULL DEFAULT 0, + limit_rate VARCHAR(32) NOT NULL DEFAULT '', + cache_enabled BOOLEAN NOT NULL DEFAULT FALSE, + cache_policy VARCHAR(32) NOT NULL DEFAULT '', + cache_rules TEXT NOT NULL DEFAULT '[]', + custom_headers TEXT NOT NULL DEFAULT '[]', + basic_auth_enabled BOOLEAN NOT NULL DEFAULT FALSE, + basic_auth_username VARCHAR(255) NOT NULL DEFAULT '', + basic_auth_password VARCHAR(255) NOT NULL DEFAULT '', + remark VARCHAR(255) NOT NULL DEFAULT '', + upstream_type VARCHAR(32) NOT NULL DEFAULT 'direct', + tunnel_node_id BIGINT, + tunnel_target_addr VARCHAR(512) NOT NULL DEFAULT '', + tunnel_target_protocol VARCHAR(16) NOT NULL DEFAULT '', + pages_project_id BIGINT, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_proxy_routes_domain ON of_proxy_routes (domain); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_proxy_routes_site_name ON of_proxy_routes (site_name); +CREATE INDEX IF NOT EXISTS idx_of_proxy_routes_origin_id ON of_proxy_routes (origin_id); +CREATE INDEX IF NOT EXISTS idx_of_proxy_routes_tunnel_node_id ON of_proxy_routes (tunnel_node_id); +CREATE INDEX IF NOT EXISTS idx_of_proxy_routes_pages_project_id ON of_proxy_routes (pages_project_id); + +-- +goose Down +DROP TABLE IF EXISTS of_proxy_routes; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190005_create_of_nodes.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190005_create_of_nodes.sql new file mode 100644 index 00000000..ff43d62a --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190005_create_of_nodes.sql @@ -0,0 +1,44 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_nodes ( + id BIGSERIAL PRIMARY KEY, + node_id VARCHAR(64) NOT NULL, + name VARCHAR(128) NOT NULL, + ip VARCHAR(64) NOT NULL DEFAULT '', + ip_manual_override BOOLEAN NOT NULL DEFAULT FALSE, + geo_name VARCHAR(128) NOT NULL DEFAULT '', + geo_latitude DOUBLE PRECISION, + geo_longitude DOUBLE PRECISION, + geo_manual_override BOOLEAN NOT NULL DEFAULT FALSE, + access_token VARCHAR(128) NOT NULL DEFAULT '', + auto_update_enabled BOOLEAN NOT NULL DEFAULT FALSE, + update_requested BOOLEAN NOT NULL DEFAULT FALSE, + update_channel VARCHAR(16) NOT NULL DEFAULT 'stable', + update_tag VARCHAR(64) NOT NULL DEFAULT '', + restart_openresty_requested BOOLEAN NOT NULL DEFAULT FALSE, + version VARCHAR(64) NOT NULL DEFAULT '', + ext_version VARCHAR(64) NOT NULL DEFAULT '', + openresty_status VARCHAR(16) NOT NULL DEFAULT 'unknown', + openresty_message TEXT, + status VARCHAR(16) NOT NULL DEFAULT 'offline', + current_version VARCHAR(32) NOT NULL DEFAULT '', + last_seen_at TIMESTAMPTZ, + last_error TEXT, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + node_type VARCHAR(32) NOT NULL DEFAULT 'edge_node', + relay_bind_port INTEGER NOT NULL DEFAULT 0, + relay_vhost_http_port INTEGER NOT NULL DEFAULT 0, + relay_auth_token VARCHAR(128) NOT NULL DEFAULT '', + relay_agent_access_addr VARCHAR(255) NOT NULL DEFAULT '', + relay_client_access_addr VARCHAR(255) NOT NULL DEFAULT '', + relay_client_proxy_url VARCHAR(512) NOT NULL DEFAULT '', + capabilities_json TEXT NOT NULL DEFAULT '[]', + relay_status VARCHAR(16) NOT NULL DEFAULT 'unknown', + relay_web_server_enabled BOOLEAN NOT NULL DEFAULT FALSE +); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_nodes_node_id ON of_nodes (node_id); +CREATE INDEX IF NOT EXISTS idx_of_nodes_access_token ON of_nodes (access_token); + +-- +goose Down +DROP TABLE IF EXISTS of_nodes; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190006_create_of_waf_tables.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190006_create_of_waf_tables.sql new file mode 100644 index 00000000..def71e73 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190006_create_of_waf_tables.sql @@ -0,0 +1,60 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_waf_rule_groups ( + id BIGSERIAL PRIMARY KEY, + name VARCHAR(255) NOT NULL, + enabled BOOLEAN NOT NULL DEFAULT TRUE, + is_global BOOLEAN NOT NULL DEFAULT FALSE, + block_status_code INTEGER NOT NULL DEFAULT 418, + block_response_body TEXT NOT NULL DEFAULT '', + ip_whitelist TEXT NOT NULL DEFAULT '[]', + ip_blacklist TEXT NOT NULL DEFAULT '[]', + ip_whitelist_groups TEXT NOT NULL DEFAULT '[]', + ip_blacklist_groups TEXT NOT NULL DEFAULT '[]', + country_whitelist TEXT NOT NULL DEFAULT '[]', + country_blacklist TEXT NOT NULL DEFAULT '[]', + region_whitelist TEXT NOT NULL DEFAULT '[]', + region_blacklist TEXT NOT NULL DEFAULT '[]', + pow_enabled BOOLEAN NOT NULL DEFAULT FALSE, + pow_config TEXT NOT NULL DEFAULT '{}', + remark VARCHAR(255) NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_of_waf_rule_groups_is_global ON of_waf_rule_groups (is_global); + +CREATE TABLE IF NOT EXISTS of_waf_ip_groups ( + id BIGSERIAL PRIMARY KEY, + name VARCHAR(255) NOT NULL, + type VARCHAR(32) NOT NULL, + enabled BOOLEAN NOT NULL DEFAULT TRUE, + ip_list TEXT NOT NULL DEFAULT '[]', + auto_config TEXT NOT NULL DEFAULT '{}', + ext_ips TEXT NOT NULL DEFAULT '[]', + subscription_url VARCHAR(2048) NOT NULL DEFAULT '', + subscription_format VARCHAR(32) NOT NULL DEFAULT 'text', + subscription_mapping_rule VARCHAR(255) NOT NULL DEFAULT '', + sync_interval_minutes INTEGER NOT NULL DEFAULT 1440, + last_synced_at TIMESTAMPTZ, + next_sync_at TIMESTAMPTZ, + last_sync_status VARCHAR(32) NOT NULL DEFAULT '', + last_sync_message TEXT NOT NULL DEFAULT '', + remark VARCHAR(255) NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_of_waf_ip_groups_type ON of_waf_ip_groups (type); +CREATE INDEX IF NOT EXISTS idx_of_waf_ip_groups_next_sync_at ON of_waf_ip_groups (next_sync_at); + +CREATE TABLE IF NOT EXISTS of_waf_rule_group_bindings ( + id BIGSERIAL PRIMARY KEY, + rule_group_id BIGINT NOT NULL, + proxy_route_id BIGINT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_waf_group_route ON of_waf_rule_group_bindings (rule_group_id, proxy_route_id); +CREATE INDEX IF NOT EXISTS idx_of_waf_rule_group_bindings_proxy_route_id ON of_waf_rule_group_bindings (proxy_route_id); + +-- +goose Down +DROP TABLE IF EXISTS of_waf_rule_group_bindings; +DROP TABLE IF EXISTS of_waf_ip_groups; +DROP TABLE IF EXISTS of_waf_rule_groups; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190007_create_of_tls_tables.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190007_create_of_tls_tables.sql new file mode 100644 index 00000000..244448a3 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190007_create_of_tls_tables.sql @@ -0,0 +1,61 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_tls_certificates ( + id BIGSERIAL PRIMARY KEY, + name VARCHAR(255) NOT NULL, + cert_pem TEXT NOT NULL, + key_pem TEXT NOT NULL, + not_before TIMESTAMPTZ, + not_after TIMESTAMPTZ, + remark VARCHAR(255) NOT NULL DEFAULT '', + provider VARCHAR(64) NOT NULL DEFAULT 'upload', + acme_account_id BIGINT NOT NULL DEFAULT 0, + dns_account_id BIGINT NOT NULL DEFAULT 0, + key_algorithm VARCHAR(32) NOT NULL DEFAULT '', + auto_renew BOOLEAN NOT NULL DEFAULT FALSE, + primary_domain VARCHAR(255) NOT NULL DEFAULT '', + other_domains TEXT NOT NULL DEFAULT '', + disable_cname BOOLEAN NOT NULL DEFAULT FALSE, + skip_dns BOOLEAN NOT NULL DEFAULT FALSE, + dns1 VARCHAR(128) NOT NULL DEFAULT '', + dns2 VARCHAR(128) NOT NULL DEFAULT '', + apply_status VARCHAR(64) NOT NULL DEFAULT 'ready', + apply_message TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_tls_certificates_name ON of_tls_certificates (name); + +CREATE TABLE IF NOT EXISTS of_managed_domains ( + id BIGSERIAL PRIMARY KEY, + domain VARCHAR(255) NOT NULL, + cert_id BIGINT, + enabled BOOLEAN NOT NULL DEFAULT TRUE, + remark VARCHAR(255) NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_managed_domains_domain ON of_managed_domains (domain); + +CREATE TABLE IF NOT EXISTS of_dns_accounts ( + id BIGSERIAL PRIMARY KEY, + name VARCHAR(255) NOT NULL, + type VARCHAR(64) NOT NULL, + authorization TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE TABLE IF NOT EXISTS of_acme_accounts ( + id BIGSERIAL PRIMARY KEY, + email VARCHAR(255) NOT NULL DEFAULT '', + url VARCHAR(255) NOT NULL DEFAULT '', + private_key TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +-- +goose Down +DROP TABLE IF EXISTS of_managed_domains; +DROP TABLE IF EXISTS of_tls_certificates; +DROP TABLE IF EXISTS of_dns_accounts; +DROP TABLE IF EXISTS of_acme_accounts; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190008_create_of_config_versions.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190008_create_of_config_versions.sql new file mode 100644 index 00000000..e923ad06 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190008_create_of_config_versions.sql @@ -0,0 +1,18 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_config_versions ( + id BIGSERIAL PRIMARY KEY, + version VARCHAR(32) NOT NULL, + snapshot_json TEXT NOT NULL, + main_config TEXT NOT NULL DEFAULT '', + rendered_config TEXT NOT NULL, + support_files_json TEXT NOT NULL DEFAULT '[]', + checksum VARCHAR(64) NOT NULL, + is_active BOOLEAN NOT NULL DEFAULT FALSE, + created_by VARCHAR(64) NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_config_versions_version ON of_config_versions (version); +CREATE INDEX IF NOT EXISTS idx_of_config_versions_is_active ON of_config_versions (is_active); + +-- +goose Down +DROP TABLE IF EXISTS of_config_versions; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190009_create_of_pages_tables.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190009_create_of_pages_tables.sql new file mode 100644 index 00000000..030fafa2 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190009_create_of_pages_tables.sql @@ -0,0 +1,53 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_pages_projects ( + id BIGSERIAL PRIMARY KEY, + name VARCHAR(255) NOT NULL, + slug VARCHAR(128) NOT NULL, + description TEXT NOT NULL DEFAULT '', + enabled BOOLEAN NOT NULL DEFAULT TRUE, + spa_fallback_enabled BOOLEAN NOT NULL DEFAULT FALSE, + spa_fallback_path VARCHAR(512) NOT NULL DEFAULT '/index.html', + api_proxy_enabled BOOLEAN NOT NULL DEFAULT FALSE, + api_proxy_path VARCHAR(255) NOT NULL DEFAULT '', + api_proxy_pass VARCHAR(2048) NOT NULL DEFAULT '', + api_proxy_rewrite VARCHAR(255) NOT NULL DEFAULT '', + active_deployment_id BIGINT, + root_dir VARCHAR(512) NOT NULL DEFAULT '', + entry_file VARCHAR(512) NOT NULL DEFAULT 'index.html', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_projects_slug ON of_pages_projects (slug); +CREATE INDEX IF NOT EXISTS idx_of_pages_projects_active_deployment_id ON of_pages_projects (active_deployment_id); + +CREATE TABLE IF NOT EXISTS of_pages_deployments ( + id BIGSERIAL PRIMARY KEY, + project_id BIGINT NOT NULL, + deployment_number INTEGER NOT NULL, + checksum VARCHAR(64) NOT NULL, + status VARCHAR(32) NOT NULL DEFAULT 'uploaded', + artifact_path VARCHAR(2048) NOT NULL, + file_count INTEGER NOT NULL DEFAULT 0, + total_size BIGINT NOT NULL DEFAULT 0, + created_by VARCHAR(64) NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + activated_at TIMESTAMPTZ +); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_project_id ON of_pages_deployments (project_id); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_checksum ON of_pages_deployments (checksum); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_status ON of_pages_deployments (status); + +CREATE TABLE IF NOT EXISTS of_pages_deployment_files ( + id BIGSERIAL PRIMARY KEY, + deployment_id BIGINT NOT NULL, + path VARCHAR(2048) NOT NULL, + size BIGINT NOT NULL DEFAULT 0, + checksum VARCHAR(64) NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployment_files_deployment_id ON of_pages_deployment_files (deployment_id); + +-- +goose Down +DROP TABLE IF EXISTS of_pages_deployment_files; +DROP TABLE IF EXISTS of_pages_deployments; +DROP TABLE IF EXISTS of_pages_projects; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190001_create_of_options.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190001_create_of_options.sql new file mode 100644 index 00000000..35a8a4e3 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190001_create_of_options.sql @@ -0,0 +1,72 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_options ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL DEFAULT '' +); + +INSERT OR IGNORE INTO of_options (key, value) VALUES + ('PasswordLoginEnabled', 'true'), + ('CapLoginEnabled', 'true'), + ('PasswordRegisterEnabled', 'false'), + ('EmailVerificationEnabled', 'false'), + ('GitHubOAuthEnabled', 'false'), + ('WeChatAuthEnabled', 'false'), + ('SMTPPort', '587'), + ('SystemName', 'OpenFlare'), + ('AgentHeartbeatInterval', '10000'), + ('AgentWebsocketUpgradeEnabled', 'true'), + ('NodeOfflineThreshold', '120000'), + ('AgentUpdateRepo', 'Rain-kl/OpenFlare'), + ('GeoIPProvider', 'ipinfo'), + ('DatabaseAutoCleanupEnabled', 'false'), + ('DatabaseAutoCleanupRetentionDays', '30'), + ('UptimeKumaEnabled', 'false'), + ('UptimeKumaMonitorScope', 'all'), + ('UptimeKumaSyncInterval', '5'), + ('UptimeKumaInterval', '60'), + ('UptimeKumaRetry', '0'), + ('UptimeKumaRetryInterval', '60'), + ('UptimeKumaTimeout', '48'), + ('OpenRestyDefaultServerReturnStatus', '421'), + ('OpenRestyWorkerProcesses', 'auto'), + ('OpenRestyWorkerConnections', '4096'), + ('OpenRestyWorkerRlimitNofile', '65535'), + ('OpenRestyEventsUse', 'epoll'), + ('OpenRestyEventsMultiAcceptEnabled', 'true'), + ('OpenRestyKeepaliveTimeout', '20'), + ('OpenRestyKeepaliveRequests', '1000'), + ('OpenRestyClientHeaderTimeout', '15'), + ('OpenRestyClientBodyTimeout', '15'), + ('OpenRestyClientMaxBodySize', '64m'), + ('OpenRestyLargeClientHeaderBuffers', '4 16k'), + ('OpenRestySendTimeout', '30'), + ('OpenRestyProxyConnectTimeout', '3'), + ('OpenRestyProxySendTimeout', '60'), + ('OpenRestyProxyReadTimeout', '60'), + ('OpenRestyWebsocketEnabled', 'true'), + ('OpenRestyHTTP3Enabled', 'true'), + ('OpenRestyProxyRequestBufferingEnabled', 'false'), + ('OpenRestyProxyBufferingEnabled', 'true'), + ('OpenRestyProxyBuffers', '16 16k'), + ('OpenRestyProxyBufferSize', '8k'), + ('OpenRestyProxyBusyBuffersSize', '64k'), + ('OpenRestyGzipEnabled', 'true'), + ('OpenRestyGzipMinLength', '1024'), + ('OpenRestyGzipCompLevel', '5'), + ('OpenRestyCacheEnabled', 'false'), + ('OpenRestyCacheLevels', '1:2'), + ('OpenRestyCacheInactive', '30m'), + ('OpenRestyCacheMaxSize', '1g'), + ('OpenRestyCacheKeyTemplate', '$scheme$host$request_uri'), + ('OpenRestyCacheLockEnabled', 'true'), + ('OpenRestyCacheLockTimeout', '5s'), + ('OpenRestyCacheUseStale', 'error timeout updating http_500 http_502 http_503 http_504'), + ('GlobalApiRateLimitNum', '300'), + ('GlobalApiRateLimitDuration', '180'), + ('GlobalWebRateLimitNum', '300'), + ('GlobalWebRateLimitDuration', '180'), + ('CriticalRateLimitNum', '100'), + ('CriticalRateLimitDuration', '1200'); + +-- +goose Down +DROP TABLE IF EXISTS of_options; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190002_create_of_origins.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190002_create_of_origins.sql new file mode 100644 index 00000000..e21c04a7 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190002_create_of_origins.sql @@ -0,0 +1,13 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_origins ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + address TEXT NOT NULL, + remark TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_origins_address ON of_origins (address); + +-- +goose Down +DROP TABLE IF EXISTS of_origins; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190003_create_of_apply_logs.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190003_create_of_apply_logs.sql new file mode 100644 index 00000000..31ebd6e8 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190003_create_of_apply_logs.sql @@ -0,0 +1,19 @@ +-- +goose Up +CREATE TABLE of_apply_logs ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + node_id TEXT NOT NULL, + version TEXT NOT NULL, + result TEXT NOT NULL, + message TEXT, + checksum TEXT NOT NULL DEFAULT '', + main_config_checksum TEXT NOT NULL DEFAULT '', + route_config_checksum TEXT NOT NULL DEFAULT '', + support_file_count INTEGER NOT NULL DEFAULT 0, + created_at DATETIME NOT NULL +); + +CREATE INDEX idx_of_apply_logs_node_id ON of_apply_logs(node_id); +CREATE INDEX idx_of_apply_logs_created_at ON of_apply_logs(created_at); + +-- +goose Down +DROP TABLE IF EXISTS of_apply_logs; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190004_create_of_proxy_routes.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190004_create_of_proxy_routes.sql new file mode 100644 index 00000000..2108c1ff --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190004_create_of_proxy_routes.sql @@ -0,0 +1,43 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_proxy_routes ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + site_name TEXT NOT NULL DEFAULT '', + domain TEXT NOT NULL, + domains TEXT NOT NULL DEFAULT '[]', + origin_id INTEGER, + origin_url TEXT NOT NULL, + origin_host TEXT NOT NULL DEFAULT '', + upstreams TEXT NOT NULL DEFAULT '[]', + enabled INTEGER NOT NULL DEFAULT 1, + enable_https INTEGER NOT NULL DEFAULT 0, + cert_id INTEGER, + cert_ids TEXT NOT NULL DEFAULT '[]', + domain_cert_ids TEXT NOT NULL DEFAULT '[]', + redirect_http INTEGER NOT NULL DEFAULT 0, + limit_conn_per_server INTEGER NOT NULL DEFAULT 0, + limit_conn_per_ip INTEGER NOT NULL DEFAULT 0, + limit_rate TEXT NOT NULL DEFAULT '', + cache_enabled INTEGER NOT NULL DEFAULT 0, + cache_policy TEXT NOT NULL DEFAULT '', + cache_rules TEXT NOT NULL DEFAULT '[]', + custom_headers TEXT NOT NULL DEFAULT '[]', + basic_auth_enabled INTEGER NOT NULL DEFAULT 0, + basic_auth_username TEXT NOT NULL DEFAULT '', + basic_auth_password TEXT NOT NULL DEFAULT '', + remark TEXT NOT NULL DEFAULT '', + upstream_type TEXT NOT NULL DEFAULT 'direct', + tunnel_node_id INTEGER, + tunnel_target_addr TEXT NOT NULL DEFAULT '', + tunnel_target_protocol TEXT NOT NULL DEFAULT '', + pages_project_id INTEGER, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_proxy_routes_domain ON of_proxy_routes (domain); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_proxy_routes_site_name ON of_proxy_routes (site_name); +CREATE INDEX IF NOT EXISTS idx_of_proxy_routes_origin_id ON of_proxy_routes (origin_id); +CREATE INDEX IF NOT EXISTS idx_of_proxy_routes_tunnel_node_id ON of_proxy_routes (tunnel_node_id); +CREATE INDEX IF NOT EXISTS idx_of_proxy_routes_pages_project_id ON of_proxy_routes (pages_project_id); + +-- +goose Down +DROP TABLE IF EXISTS of_proxy_routes; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190005_create_of_nodes.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190005_create_of_nodes.sql new file mode 100644 index 00000000..23d04e69 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190005_create_of_nodes.sql @@ -0,0 +1,44 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_nodes ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + node_id TEXT NOT NULL, + name TEXT NOT NULL, + ip TEXT NOT NULL DEFAULT '', + ip_manual_override INTEGER NOT NULL DEFAULT 0, + geo_name TEXT NOT NULL DEFAULT '', + geo_latitude REAL, + geo_longitude REAL, + geo_manual_override INTEGER NOT NULL DEFAULT 0, + access_token TEXT NOT NULL DEFAULT '', + auto_update_enabled INTEGER NOT NULL DEFAULT 0, + update_requested INTEGER NOT NULL DEFAULT 0, + update_channel TEXT NOT NULL DEFAULT 'stable', + update_tag TEXT NOT NULL DEFAULT '', + restart_openresty_requested INTEGER NOT NULL DEFAULT 0, + version TEXT NOT NULL DEFAULT '', + ext_version TEXT NOT NULL DEFAULT '', + openresty_status TEXT NOT NULL DEFAULT 'unknown', + openresty_message TEXT, + status TEXT NOT NULL DEFAULT 'offline', + current_version TEXT NOT NULL DEFAULT '', + last_seen_at DATETIME, + last_error TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + node_type TEXT NOT NULL DEFAULT 'edge_node', + relay_bind_port INTEGER NOT NULL DEFAULT 0, + relay_vhost_http_port INTEGER NOT NULL DEFAULT 0, + relay_auth_token TEXT NOT NULL DEFAULT '', + relay_agent_access_addr TEXT NOT NULL DEFAULT '', + relay_client_access_addr TEXT NOT NULL DEFAULT '', + relay_client_proxy_url TEXT NOT NULL DEFAULT '', + capabilities_json TEXT NOT NULL DEFAULT '[]', + relay_status TEXT NOT NULL DEFAULT 'unknown', + relay_web_server_enabled INTEGER NOT NULL DEFAULT 0 +); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_nodes_node_id ON of_nodes (node_id); +CREATE INDEX IF NOT EXISTS idx_of_nodes_access_token ON of_nodes (access_token); + +-- +goose Down +DROP TABLE IF EXISTS of_nodes; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190006_create_of_waf_tables.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190006_create_of_waf_tables.sql new file mode 100644 index 00000000..e73a602b --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190006_create_of_waf_tables.sql @@ -0,0 +1,60 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_waf_rule_groups ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + enabled BOOLEAN NOT NULL DEFAULT TRUE, + is_global BOOLEAN NOT NULL DEFAULT FALSE, + block_status_code INTEGER NOT NULL DEFAULT 418, + block_response_body TEXT NOT NULL DEFAULT '', + ip_whitelist TEXT NOT NULL DEFAULT '[]', + ip_blacklist TEXT NOT NULL DEFAULT '[]', + ip_whitelist_groups TEXT NOT NULL DEFAULT '[]', + ip_blacklist_groups TEXT NOT NULL DEFAULT '[]', + country_whitelist TEXT NOT NULL DEFAULT '[]', + country_blacklist TEXT NOT NULL DEFAULT '[]', + region_whitelist TEXT NOT NULL DEFAULT '[]', + region_blacklist TEXT NOT NULL DEFAULT '[]', + pow_enabled BOOLEAN NOT NULL DEFAULT FALSE, + pow_config TEXT NOT NULL DEFAULT '{}', + remark TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_of_waf_rule_groups_is_global ON of_waf_rule_groups (is_global); + +CREATE TABLE IF NOT EXISTS of_waf_ip_groups ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + type TEXT NOT NULL, + enabled BOOLEAN NOT NULL DEFAULT TRUE, + ip_list TEXT NOT NULL DEFAULT '[]', + auto_config TEXT NOT NULL DEFAULT '{}', + ext_ips TEXT NOT NULL DEFAULT '[]', + subscription_url TEXT NOT NULL DEFAULT '', + subscription_format TEXT NOT NULL DEFAULT 'text', + subscription_mapping_rule TEXT NOT NULL DEFAULT '', + sync_interval_minutes INTEGER NOT NULL DEFAULT 1440, + last_synced_at DATETIME, + next_sync_at DATETIME, + last_sync_status TEXT NOT NULL DEFAULT '', + last_sync_message TEXT NOT NULL DEFAULT '', + remark TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_of_waf_ip_groups_type ON of_waf_ip_groups (type); +CREATE INDEX IF NOT EXISTS idx_of_waf_ip_groups_next_sync_at ON of_waf_ip_groups (next_sync_at); + +CREATE TABLE IF NOT EXISTS of_waf_rule_group_bindings ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + rule_group_id INTEGER NOT NULL, + proxy_route_id INTEGER NOT NULL, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_waf_group_route ON of_waf_rule_group_bindings (rule_group_id, proxy_route_id); +CREATE INDEX IF NOT EXISTS idx_of_waf_rule_group_bindings_proxy_route_id ON of_waf_rule_group_bindings (proxy_route_id); + +-- +goose Down +DROP TABLE IF EXISTS of_waf_rule_group_bindings; +DROP TABLE IF EXISTS of_waf_ip_groups; +DROP TABLE IF EXISTS of_waf_rule_groups; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190007_create_of_tls_tables.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190007_create_of_tls_tables.sql new file mode 100644 index 00000000..cf723e2e --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190007_create_of_tls_tables.sql @@ -0,0 +1,61 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_tls_certificates ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + cert_pem TEXT NOT NULL, + key_pem TEXT NOT NULL, + not_before DATETIME, + not_after DATETIME, + remark TEXT NOT NULL DEFAULT '', + provider TEXT NOT NULL DEFAULT 'upload', + acme_account_id INTEGER NOT NULL DEFAULT 0, + dns_account_id INTEGER NOT NULL DEFAULT 0, + key_algorithm TEXT NOT NULL DEFAULT '', + auto_renew INTEGER NOT NULL DEFAULT 0, + primary_domain TEXT NOT NULL DEFAULT '', + other_domains TEXT NOT NULL DEFAULT '', + disable_cname INTEGER NOT NULL DEFAULT 0, + skip_dns INTEGER NOT NULL DEFAULT 0, + dns1 TEXT NOT NULL DEFAULT '', + dns2 TEXT NOT NULL DEFAULT '', + apply_status TEXT NOT NULL DEFAULT 'ready', + apply_message TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_tls_certificates_name ON of_tls_certificates (name); + +CREATE TABLE IF NOT EXISTS of_managed_domains ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + domain TEXT NOT NULL, + cert_id INTEGER, + enabled INTEGER NOT NULL DEFAULT 1, + remark TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_managed_domains_domain ON of_managed_domains (domain); + +CREATE TABLE IF NOT EXISTS of_dns_accounts ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + type TEXT NOT NULL, + authorization TEXT NOT NULL, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE TABLE IF NOT EXISTS of_acme_accounts ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + email TEXT NOT NULL DEFAULT '', + url TEXT NOT NULL DEFAULT '', + private_key TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +-- +goose Down +DROP TABLE IF EXISTS of_managed_domains; +DROP TABLE IF EXISTS of_tls_certificates; +DROP TABLE IF EXISTS of_dns_accounts; +DROP TABLE IF EXISTS of_acme_accounts; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190008_create_of_config_versions.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190008_create_of_config_versions.sql new file mode 100644 index 00000000..9d98ab94 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190008_create_of_config_versions.sql @@ -0,0 +1,18 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_config_versions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + version TEXT NOT NULL, + snapshot_json TEXT NOT NULL, + main_config TEXT NOT NULL DEFAULT '', + rendered_config TEXT NOT NULL, + support_files_json TEXT NOT NULL DEFAULT '[]', + checksum TEXT NOT NULL, + is_active INTEGER NOT NULL DEFAULT 0, + created_by TEXT NOT NULL, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_config_versions_version ON of_config_versions (version); +CREATE INDEX IF NOT EXISTS idx_of_config_versions_is_active ON of_config_versions (is_active); + +-- +goose Down +DROP TABLE IF EXISTS of_config_versions; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190009_create_of_pages_tables.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190009_create_of_pages_tables.sql new file mode 100644 index 00000000..4b05564f --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190009_create_of_pages_tables.sql @@ -0,0 +1,53 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_pages_projects ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + slug TEXT NOT NULL, + description TEXT NOT NULL DEFAULT '', + enabled INTEGER NOT NULL DEFAULT 1, + spa_fallback_enabled INTEGER NOT NULL DEFAULT 0, + spa_fallback_path TEXT NOT NULL DEFAULT '/index.html', + api_proxy_enabled INTEGER NOT NULL DEFAULT 0, + api_proxy_path TEXT NOT NULL DEFAULT '', + api_proxy_pass TEXT NOT NULL DEFAULT '', + api_proxy_rewrite TEXT NOT NULL DEFAULT '', + active_deployment_id INTEGER, + root_dir TEXT NOT NULL DEFAULT '', + entry_file TEXT NOT NULL DEFAULT 'index.html', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_projects_slug ON of_pages_projects (slug); +CREATE INDEX IF NOT EXISTS idx_of_pages_projects_active_deployment_id ON of_pages_projects (active_deployment_id); + +CREATE TABLE IF NOT EXISTS of_pages_deployments ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + project_id INTEGER NOT NULL, + deployment_number INTEGER NOT NULL, + checksum TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'uploaded', + artifact_path TEXT NOT NULL, + file_count INTEGER NOT NULL DEFAULT 0, + total_size INTEGER NOT NULL DEFAULT 0, + created_by TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + activated_at DATETIME +); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_project_id ON of_pages_deployments (project_id); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_checksum ON of_pages_deployments (checksum); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_status ON of_pages_deployments (status); + +CREATE TABLE IF NOT EXISTS of_pages_deployment_files ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + deployment_id INTEGER NOT NULL, + path TEXT NOT NULL, + size INTEGER NOT NULL DEFAULT 0, + checksum TEXT NOT NULL, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS idx_of_pages_deployment_files_deployment_id ON of_pages_deployment_files (deployment_id); + +-- +goose Down +DROP TABLE IF EXISTS of_pages_deployment_files; +DROP TABLE IF EXISTS of_pages_deployments; +DROP TABLE IF EXISTS of_pages_projects; \ No newline at end of file diff --git a/Wavelet/internal/model/openflare_acme_account.go b/Wavelet/internal/model/openflare_acme_account.go new file mode 100644 index 00000000..c29bc9f0 --- /dev/null +++ b/Wavelet/internal/model/openflare_acme_account.go @@ -0,0 +1,64 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +// AcmeAccount OpenFlare ACME 账号实体。 +type AcmeAccount struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Email string `json:"email" gorm:"size:255"` + URL string `json:"url" gorm:"size:255"` + PrivateKey string `json:"-" gorm:"type:text;not null"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名。 +func (AcmeAccount) TableName() string { + return "of_acme_accounts" +} + +// GetAcmeAccountByID 按 ID 查询 ACME 账号。 +func GetAcmeAccountByID(ctx context.Context, id uint) (*AcmeAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var account AcmeAccount + if err := conn.First(&account, id).Error; err != nil { + return nil, err + } + return &account, nil +} + +// GetDefaultAcmeAccount 获取默认 ACME 账号,不存在时创建占位记录。 +func GetDefaultAcmeAccount(ctx context.Context) (*AcmeAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var account AcmeAccount + err := conn.Order("id asc").First(&account).Error + if err == nil { + return &account, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + account = AcmeAccount{ + Email: "admin@openflare.dev", + } + if err = conn.Create(&account).Error; err != nil { + return nil, err + } + return &account, nil +} diff --git a/Wavelet/internal/model/openflare_apply_log.go b/Wavelet/internal/model/openflare_apply_log.go new file mode 100644 index 00000000..85f4a7ec --- /dev/null +++ b/Wavelet/internal/model/openflare_apply_log.go @@ -0,0 +1,132 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +// OpenFlareApplyLogQuery filters apply logs for list queries. +type OpenFlareApplyLogQuery struct { + NodeID string + PageNo int + PageSize int +} + +// OpenFlareApplyLog stores node configuration apply results. +type OpenFlareApplyLog struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"index;size:64;not null"` + Version string `json:"version" gorm:"size:32;not null"` + Result string `json:"result" gorm:"size:32;not null"` + Message string `json:"message" gorm:"type:text"` + Checksum string `json:"checksum" gorm:"size:64;not null;default:''"` + MainConfigChecksum string `json:"main_config_checksum" gorm:"size:64;not null;default:''"` + RouteConfigChecksum string `json:"route_config_checksum" gorm:"size:64;not null;default:''"` + SupportFileCount int `json:"support_file_count" gorm:"not null;default:0"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"` +} + +// TableName returns the GORM table name. +func (OpenFlareApplyLog) TableName() string { + return "of_apply_logs" +} + +// ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination. +func ListOpenFlareApplyLogs(ctx context.Context, query OpenFlareApplyLogQuery) ([]*OpenFlareApplyLog, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + + dbQuery := conn.Model(&OpenFlareApplyLog{}).Order("id desc") + if query.NodeID != "" { + dbQuery = dbQuery.Where("node_id = ?", query.NodeID) + } + if query.PageSize > 0 { + offset := 0 + if query.PageNo > 1 { + offset = (query.PageNo - 1) * query.PageSize + } + dbQuery = dbQuery.Limit(query.PageSize).Offset(offset) + } + + var logs []*OpenFlareApplyLog + if err := dbQuery.Find(&logs).Error; err != nil { + return nil, err + } + return logs, nil +} + +// CountOpenFlareApplyLogs returns total apply logs, optionally filtered by node_id. +func CountOpenFlareApplyLogs(ctx context.Context, nodeID string) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + + query := conn.Model(&OpenFlareApplyLog{}) + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + + var total int64 + if err := query.Count(&total).Error; err != nil { + return 0, err + } + return total, nil +} + +// GetLatestOpenFlareApplyLogsByNodeIDs returns the latest apply log per node id. +func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) (map[string]*OpenFlareApplyLog, error) { + result := make(map[string]*OpenFlareApplyLog) + if len(nodeIDs) == 0 { + return result, nil + } + + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + + var logs []*OpenFlareApplyLog + subQuery := conn.Model(&OpenFlareApplyLog{}). + Select("MAX(id) AS id"). + Where("node_id IN ?", nodeIDs). + Group("node_id") + if err := conn.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil { + return nil, err + } + for _, log := range logs { + result[log.NodeID] = log + } + return result, nil +} + +// DeleteAllOpenFlareApplyLogs removes every apply log record. +func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + + result := conn.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&OpenFlareApplyLog{}) + return result.RowsAffected, result.Error +} + +// DeleteOpenFlareApplyLogsBefore removes apply logs created before the cutoff time. +func DeleteOpenFlareApplyLogsBefore(ctx context.Context, before time.Time) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + + result := conn.Where("created_at < ?", before).Delete(&OpenFlareApplyLog{}) + return result.RowsAffected, result.Error +} diff --git a/Wavelet/internal/model/openflare_config_version.go b/Wavelet/internal/model/openflare_config_version.go new file mode 100644 index 00000000..94cbd3ce --- /dev/null +++ b/Wavelet/internal/model/openflare_config_version.go @@ -0,0 +1,163 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +// ConfigVersionSummary is the list view for config versions. +type ConfigVersionSummary struct { + ID uint `json:"id"` + Version string `json:"version"` + Checksum string `json:"checksum"` + IsActive bool `json:"is_active"` + CreatedBy string `json:"created_by"` + CreatedAt time.Time `json:"created_at"` +} + +// ConfigVersion stores a published OpenResty configuration snapshot. +type ConfigVersion struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Version string `json:"version" gorm:"uniqueIndex;size:32;not null"` + SnapshotJSON string `json:"snapshot_json" gorm:"type:text;not null"` + MainConfig string `json:"main_config" gorm:"type:text;not null;default:''"` + RenderedConfig string `json:"rendered_config" gorm:"type:text;not null"` + SupportFilesJSON string `json:"support_files_json" gorm:"type:text;not null;default:'[]'"` + Checksum string `json:"checksum" gorm:"size:64;not null"` + IsActive bool `json:"is_active" gorm:"not null;default:false;index"` + CreatedBy string `json:"created_by" gorm:"size:64;not null"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (ConfigVersion) TableName() string { + return "of_config_versions" +} + +// ListConfigVersionSummaries returns config version summaries ordered by id desc. +func ListConfigVersionSummaries(ctx context.Context) ([]*ConfigVersionSummary, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var versions []*ConfigVersionSummary + err := conn.Model(&ConfigVersion{}). + Select("id", "version", "checksum", "is_active", "created_by", "created_at"). + Order("id desc"). + Find(&versions).Error + return versions, err +} + +// GetConfigVersionByID returns a config version by primary key. +func GetConfigVersionByID(ctx context.Context, id uint) (*ConfigVersion, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var version ConfigVersion + if err := conn.First(&version, id).Error; err != nil { + return nil, err + } + return &version, nil +} + +// GetActiveConfigVersion returns the currently active config version. +func GetActiveConfigVersion(ctx context.Context) (*ConfigVersion, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var version ConfigVersion + if err := conn.Where("is_active = ?", true).Order("id desc").First(&version).Error; err != nil { + return nil, err + } + return &version, nil +} + +// GetLatestConfigVersionByPrefix returns the latest version string matching a date prefix. +func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, error) { + conn := db.DB(ctx) + if conn == nil { + return "", errors.New(errDatabaseNotInitialized) + } + var version ConfigVersion + err := conn.Model(&ConfigVersion{}). + Select("version"). + Where("version LIKE ?", prefix+"-%"). + Order("version desc"). + First(&version).Error + if err != nil { + return "", err + } + return version.Version, nil +} + +// CreateConfigVersion inserts a new config version record. +func CreateConfigVersion(ctx context.Context, version *ConfigVersion) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(version).Error +} + +// PublishConfigVersionTx deactivates all versions and creates a new active version. +func PublishConfigVersionTx(ctx context.Context, version *ConfigVersion) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Transaction(func(tx *gorm.DB) error { + if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil { + return err + } + return tx.Create(version).Error + }) +} + +// ActivateConfigVersionTx marks the given version active and deactivates others. +func ActivateConfigVersionTx(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Transaction(func(tx *gorm.DB) error { + if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil { + return err + } + return tx.Model(&ConfigVersion{}).Where("id = ?", id).Update("is_active", true).Error + }) +} + +// DeleteConfigVersionsByIDs removes config versions by ids. +func DeleteConfigVersionsByIDs(ctx context.Context, ids []uint) (int64, error) { + if len(ids) == 0 { + return 0, nil + } + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + result := conn.Where("id IN ?", ids).Delete(&ConfigVersion{}) + return result.RowsAffected, result.Error +} + +// ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc. +func ListEnabledProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var routes []*ProxyRoute + if err := conn.Where("enabled = ?", true).Order("id asc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} diff --git a/Wavelet/internal/model/openflare_dns_account.go b/Wavelet/internal/model/openflare_dns_account.go new file mode 100644 index 00000000..04905df6 --- /dev/null +++ b/Wavelet/internal/model/openflare_dns_account.go @@ -0,0 +1,80 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" +) + +// DNSAccount OpenFlare DNS 账号实体。 +type DNSAccount struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Name string `json:"name" gorm:"size:255;not null"` + Type string `json:"type" gorm:"size:64;not null"` + Authorization string `json:"-" gorm:"type:text;not null"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名。 +func (DNSAccount) TableName() string { + return "of_dns_accounts" +} + +// ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。 +func ListDNSAccounts(ctx context.Context) ([]DNSAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var accounts []DNSAccount + if err := conn.Order("id desc").Find(&accounts).Error; err != nil { + return nil, err + } + return accounts, nil +} + +// GetDNSAccountByID 按 ID 查询 DNS 账号。 +func GetDNSAccountByID(ctx context.Context, id uint) (*DNSAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var account DNSAccount + if err := conn.First(&account, id).Error; err != nil { + return nil, err + } + return &account, nil +} + +// CreateDNSAccountRecord 创建 DNS 账号。 +func CreateDNSAccountRecord(ctx context.Context, account *DNSAccount) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(account).Error +} + +// SaveDNSAccount 保存 DNS 账号。 +func SaveDNSAccount(ctx context.Context, account *DNSAccount) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(account).Error +} + +// DeleteDNSAccountRecord 删除 DNS 账号。 +func DeleteDNSAccountRecord(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Delete(&DNSAccount{}, id).Error +} diff --git a/Wavelet/internal/model/openflare_managed_domain.go b/Wavelet/internal/model/openflare_managed_domain.go new file mode 100644 index 00000000..1ebcb599 --- /dev/null +++ b/Wavelet/internal/model/openflare_managed_domain.go @@ -0,0 +1,94 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" +) + +// ManagedDomain OpenFlare 托管域名实体。 +type ManagedDomain struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"` + CertID *uint `json:"cert_id"` + Enabled bool `json:"enabled" gorm:"not null;default:true"` + Remark string `json:"remark" gorm:"size:255"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名。 +func (ManagedDomain) TableName() string { + return "of_managed_domains" +} + +// ListManagedDomains 列出全部托管域名。 +func ListManagedDomains(ctx context.Context) ([]ManagedDomain, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var domains []ManagedDomain + if err := conn.Order("id desc").Find(&domains).Error; err != nil { + return nil, err + } + return domains, nil +} + +// ListEnabledManagedDomainsWithCertificate 列出已启用且绑定证书的托管域名。 +func ListEnabledManagedDomainsWithCertificate(ctx context.Context) ([]ManagedDomain, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var domains []ManagedDomain + if err := conn.Where("enabled = ? AND cert_id IS NOT NULL", true).Order("id desc").Find(&domains).Error; err != nil { + return nil, err + } + return domains, nil +} + +// GetManagedDomainByID 按 ID 查询托管域名。 +func GetManagedDomainByID(ctx context.Context, id uint) (*ManagedDomain, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var domain ManagedDomain + if err := conn.First(&domain, id).Error; err != nil { + return nil, err + } + return &domain, nil +} + +// CreateManagedDomainRecord 创建托管域名。 +func CreateManagedDomainRecord(ctx context.Context, domain *ManagedDomain) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(domain).Error +} + +// SaveManagedDomain 保存托管域名。 +func SaveManagedDomain(ctx context.Context, domain *ManagedDomain) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(domain).Error +} + +// DeleteManagedDomainRecord 删除托管域名。 +func DeleteManagedDomainRecord(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Delete(&ManagedDomain{}, id).Error +} diff --git a/Wavelet/internal/model/openflare_node.go b/Wavelet/internal/model/openflare_node.go new file mode 100644 index 00000000..8359e4f3 --- /dev/null +++ b/Wavelet/internal/model/openflare_node.go @@ -0,0 +1,163 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" +) + +// OpenFlareNode stores an edge, relay, or tunnel client node. +type OpenFlareNode struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"` + Name string `json:"name" gorm:"size:128;not null"` + IP string `json:"ip" gorm:"size:64;not null;default:''"` + IPManualOverride bool `json:"ip_manual_override" gorm:"not null;default:false"` + GeoName string `json:"geo_name" gorm:"size:128;not null;default:''"` + GeoLatitude *float64 `json:"geo_latitude"` + GeoLongitude *float64 `json:"geo_longitude"` + GeoManualOverride bool `json:"geo_manual_override" gorm:"not null;default:false"` + AccessToken string `json:"-" gorm:"column:access_token;size:128;index"` + AutoUpdateEnabled bool `json:"auto_update_enabled" gorm:"not null;default:false"` + UpdateRequested bool `json:"update_requested" gorm:"not null;default:false"` + UpdateChannel string `json:"update_channel" gorm:"size:16;not null;default:'stable'"` + UpdateTag string `json:"update_tag" gorm:"size:64;not null;default:''"` + RestartOpenrestyRequested bool `json:"restart_openresty_requested" gorm:"not null;default:false"` + Version string `json:"version" gorm:"size:64;not null;default:''"` + ExtVersion string `json:"ext_version" gorm:"size:64;not null;default:''"` + OpenrestyStatus string `json:"openresty_status" gorm:"size:16;not null;default:'unknown'"` + OpenrestyMessage string `json:"openresty_message" gorm:"type:text"` + Status string `json:"status" gorm:"size:16;not null;default:'offline'"` + CurrentVersion string `json:"current_version" gorm:"size:32;not null;default:''"` + LastSeenAt *time.Time `json:"last_seen_at"` + LastError string `json:"last_error" gorm:"type:text"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` + NodeType string `json:"node_type" gorm:"size:32;not null;default:'edge_node'"` + RelayBindPort int `json:"relay_bind_port" gorm:"not null;default:0"` + RelayVhostHTTPPort int `json:"relay_vhost_http_port" gorm:"not null;default:0"` + RelayAuthToken string `json:"-" gorm:"size:128;not null;default:''"` + RelayAgentAccessAddr string `json:"relay_agent_access_addr" gorm:"size:255;not null;default:''"` + RelayClientAccessAddr string `json:"relay_client_access_addr" gorm:"size:255;not null;default:''"` + RelayClientProxyURL string `json:"relay_client_proxy_url" gorm:"size:512;not null;default:''"` + CapabilitiesJSON string `json:"capabilities_json" gorm:"type:text;not null;default:'[]'"` + RelayStatus string `json:"relay_status" gorm:"size:16;not null;default:'unknown'"` + RelayWebServerEnabled bool `json:"relay_web_server_enabled" gorm:"not null;default:false"` +} + +// TableName returns the GORM table name. +func (OpenFlareNode) TableName() string { + return "of_nodes" +} + +// ListOpenFlareNodes returns all nodes ordered by id desc. +func ListOpenFlareNodes(ctx context.Context) ([]OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var nodes []OpenFlareNode + if err := conn.Order("id desc").Find(&nodes).Error; err != nil { + return nil, err + } + return nodes, nil +} + +// ListOpenFlareNodesByNodeIDs returns nodes matching the given node ids. +func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]OpenFlareNode, error) { + if len(nodeIDs) == 0 { + return []OpenFlareNode{}, nil + } + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var nodes []OpenFlareNode + if err := conn.Where("node_id IN ?", nodeIDs).Find(&nodes).Error; err != nil { + return nil, err + } + return nodes, nil +} + +// GetOpenFlareNodeByID returns a node by primary key. +func GetOpenFlareNodeByID(ctx context.Context, id uint) (*OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var node OpenFlareNode + if err := conn.First(&node, id).Error; err != nil { + return nil, err + } + return &node, nil +} + +// GetOpenFlareNodeByNodeID returns a node by node_id. +func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var node OpenFlareNode + if err := conn.Where("node_id = ?", nodeID).First(&node).Error; err != nil { + return nil, err + } + return &node, nil +} + +// GetOpenFlareNodeByAccessToken returns a node by access token. +func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var node OpenFlareNode + if err := conn.Where("access_token = ?", token).First(&node).Error; err != nil { + return nil, err + } + return &node, nil +} + +// CreateOpenFlareNode inserts a new node. +func CreateOpenFlareNode(ctx context.Context, node *OpenFlareNode) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(node).Error +} + +// SaveOpenFlareNode persists node changes. +func SaveOpenFlareNode(ctx context.Context, node *OpenFlareNode) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(node).Error +} + +// UpdateOpenFlareNodeFields updates selected columns for a node. +func UpdateOpenFlareNodeFields(ctx context.Context, node *OpenFlareNode, fields ...string) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + if len(fields) == 0 { + return conn.Save(node).Error + } + return conn.Model(node).Select(fields).Updates(node).Error +} + +// DeleteOpenFlareNode removes a node by primary key. +func DeleteOpenFlareNode(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Delete(&OpenFlareNode{}, id).Error +} diff --git a/Wavelet/internal/model/openflare_observability.go b/Wavelet/internal/model/openflare_observability.go new file mode 100644 index 00000000..70df959f --- /dev/null +++ b/Wavelet/internal/model/openflare_observability.go @@ -0,0 +1,500 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +// OpenFlareMetricSnapshot stores a node capacity snapshot (v1 single table, no sharding). +type OpenFlareMetricSnapshot struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"index;size:64;not null"` + CapturedAt time.Time `json:"captured_at" gorm:"index"` + CPUUsagePercent float64 `json:"cpu_usage_percent"` + MemoryUsedBytes int64 `json:"memory_used_bytes"` + MemoryTotalBytes int64 `json:"memory_total_bytes"` + StorageUsedBytes int64 `json:"storage_used_bytes"` + StorageTotalBytes int64 `json:"storage_total_bytes"` + DiskReadBytes int64 `json:"disk_read_bytes"` + DiskWriteBytes int64 `json:"disk_write_bytes"` + NetworkRxBytes int64 `json:"network_rx_bytes"` + NetworkTxBytes int64 `json:"network_tx_bytes"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareMetricSnapshot) TableName() string { + return "of_node_metric_snapshots" +} + +// OpenFlareRequestReport stores aggregated traffic windows per node. +type OpenFlareRequestReport struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"index;size:64;not null"` + WindowStartedAt time.Time `json:"window_started_at" gorm:"index"` + WindowEndedAt time.Time `json:"window_ended_at" gorm:"index"` + RequestCount int64 `json:"request_count"` + ErrorCount int64 `json:"error_count"` + UniqueVisitorCount int64 `json:"unique_visitor_count"` + StatusCodesJSON string `json:"status_codes_json" gorm:"type:text"` + TopDomainsJSON string `json:"top_domains_json" gorm:"type:text"` + SourceCountriesJSON string `json:"source_countries_json" gorm:"type:text"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareRequestReport) TableName() string { + return "of_node_request_reports" +} + +// OpenFlareAccessLog stores a single access log row (v1 single table, no sharding). +type OpenFlareAccessLog struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"index;size:64;not null"` + LoggedAt time.Time `json:"logged_at" gorm:"index"` + RemoteAddr string `json:"remote_addr" gorm:"index;size:128"` + Region string `json:"region" gorm:"size:128"` + Host string `json:"host" gorm:"index;size:255"` + Path string `json:"path" gorm:"size:2048"` + StatusCode int `json:"status_code" gorm:"index"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareAccessLog) TableName() string { + return "of_node_access_logs" +} + +// OpenFlareAccessLogRegionCount aggregates access log regions. +type OpenFlareAccessLogRegionCount struct { + Region string `json:"region"` + Count int64 `json:"count"` +} + +// OpenFlareHealthEvent stores node health alert events. +type OpenFlareHealthEvent struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"index;size:64;not null"` + EventType string `json:"event_type" gorm:"index;size:64;not null"` + Severity string `json:"severity" gorm:"size:16;not null"` + Status string `json:"status" gorm:"index;size:16;not null"` + Message string `json:"message" gorm:"type:text"` + FirstTriggeredAt time.Time `json:"first_triggered_at" gorm:"index"` + LastTriggeredAt time.Time `json:"last_triggered_at" gorm:"index"` + ReportedAt time.Time `json:"reported_at" gorm:"index"` + ResolvedAt *time.Time `json:"resolved_at" gorm:"index"` + MetadataJSON string `json:"metadata_json" gorm:"type:text"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareHealthEvent) TableName() string { + return "of_node_health_events" +} + +// OpenFlareNodeSystemProfile stores the latest node system profile. +type OpenFlareNodeSystemProfile struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"` + Hostname string `json:"hostname" gorm:"size:255"` + OSName string `json:"os_name" gorm:"size:128"` + OSVersion string `json:"os_version" gorm:"size:128"` + KernelVersion string `json:"kernel_version" gorm:"size:128"` + Architecture string `json:"architecture" gorm:"size:64"` + CPUModel string `json:"cpu_model" gorm:"size:255"` + CPUCores int `json:"cpu_cores"` + TotalMemoryBytes int64 `json:"total_memory_bytes"` + TotalDiskBytes int64 `json:"total_disk_bytes"` + UptimeSeconds int64 `json:"uptime_seconds"` + ReportedAt time.Time `json:"reported_at" gorm:"index"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareNodeSystemProfile) TableName() string { + return "of_node_system_profiles" +} + +// OpenFlareNodeObservationOpenresty stores openresty network observations. +type OpenFlareNodeObservationOpenresty struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"index;size:64;not null"` + CapturedAt time.Time `json:"captured_at" gorm:"index"` + OpenrestyRxBytes int64 `json:"openresty_rx_bytes"` + OpenrestyTxBytes int64 `json:"openresty_tx_bytes"` + OpenrestyConnections int64 `json:"openresty_connections"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareNodeObservationOpenresty) TableName() string { + return "of_node_obs_openresty" +} + +// OpenFlareNodeObservationFrps stores tunnel relay frps observations. +type OpenFlareNodeObservationFrps struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + NodeID string `json:"node_id" gorm:"index;size:64;not null"` + CapturedAt time.Time `json:"captured_at" gorm:"index"` + FrpsConnections int `json:"frps_connections"` + FrpsProxyCount int `json:"frps_proxy_count"` + FrpsClientCount int `json:"frps_client_count"` + FrpsProxies string `json:"frps_proxies" gorm:"type:text"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareNodeObservationFrps) TableName() string { + return "of_node_obs_frps" +} + +// OpenFlareAccessLogQuery filters access log list queries. +type OpenFlareAccessLogQuery struct { + NodeID string + RemoteAddr string + Host string + Path string + Since time.Time + Until time.Time + Page int + PageSize int + SortBy string + SortOrder string +} + +// OpenFlareAccessLogBucketQuery filters folded access log queries (v1 stub). +type OpenFlareAccessLogBucketQuery struct { + NodeID string + RemoteAddr string + Host string + Path string + Since time.Time + Page int + PageSize int + SortBy string + SortOrder string + FoldMinutes int +} + +// OpenFlareAccessLogBucketRow is a folded access log bucket row (v1 stub). +type OpenFlareAccessLogBucketRow struct { + BucketEpoch int64 `json:"bucket_epoch"` + RequestCount int64 `json:"request_count"` + UniqueIPCount int64 `json:"unique_ip_count"` + UniqueHostCount int64 `json:"unique_host_count"` + SuccessCount int64 `json:"success_count"` + ClientErrorCount int64 `json:"client_error_count"` + ServerErrorCount int64 `json:"server_error_count"` +} + +// OpenFlareAccessLogBucketIPQuery filters folded IP summary queries (v1 stub). +type OpenFlareAccessLogBucketIPQuery struct { + NodeID string + RemoteAddr string + Host string + Path string + BucketStartedAt time.Time + FoldMinutes int + Page int + PageSize int + SortBy string + SortOrder string +} + +// OpenFlareAccessLogBucketIPRow is a folded IP row (v1 stub). +type OpenFlareAccessLogBucketIPRow struct { + RemoteAddr string `json:"remote_addr"` + RequestCount int64 `json:"request_count"` + SuccessCount int64 `json:"success_count"` + ClientErrorCount int64 `json:"client_error_count"` + ServerErrorCount int64 `json:"server_error_count"` + LastSeenEpoch int64 `json:"last_seen_epoch"` +} + +// OpenFlareAccessLogIPSummaryQuery filters IP summary list queries (v1 stub). +type OpenFlareAccessLogIPSummaryQuery struct { + NodeID string + RemoteAddr string + Host string + Since time.Time + Page int + PageSize int + SortBy string + SortOrder string +} + +// OpenFlareAccessLogIPSummaryRow is an IP summary row (v1 stub). +type OpenFlareAccessLogIPSummaryRow struct { + RemoteAddr string `json:"remote_addr"` + TotalRequests int64 `json:"total_requests"` + RecentRequests int64 `json:"recent_requests"` + LastSeenEpoch int64 `json:"last_seen_epoch"` +} + +// OpenFlareAccessLogIPTrendQuery filters IP trend queries (v1 stub). +type OpenFlareAccessLogIPTrendQuery struct { + NodeID string + RemoteAddr string + Host string + Since time.Time + BucketMinutes int +} + +// OpenFlareAccessLogIPTrendRow is an IP trend bucket row (v1 stub). +type OpenFlareAccessLogIPTrendRow struct { + BucketEpoch int64 `json:"bucket_epoch"` + RequestCount int64 `json:"request_count"` +} + +func isMissingTableError(err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "no such table") || + strings.Contains(msg, "doesn't exist") || + strings.Contains(msg, "does not exist") +} + +// ListOpenFlareMetricSnapshotsSince returns metric snapshots since the given time. +func ListOpenFlareMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + query := conn.Model(&OpenFlareMetricSnapshot{}).Order("captured_at desc, id desc") + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + if !since.IsZero() { + query = query.Where("captured_at >= ?", since) + } + if limit > 0 { + query = query.Limit(limit) + } + var rows []*OpenFlareMetricSnapshot + if err := query.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareMetricSnapshot{}, nil + } + return nil, err + } + return rows, nil +} + +// ListOpenFlareRequestReportsSince returns request reports since the given time. +func ListOpenFlareRequestReportsSince(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareRequestReport, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + query := conn.Model(&OpenFlareRequestReport{}).Order("window_ended_at desc, id desc") + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + if !since.IsZero() { + query = query.Where("window_ended_at >= ?", since) + } + if limit > 0 { + query = query.Limit(limit) + } + var rows []*OpenFlareRequestReport + if err := query.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareRequestReport{}, nil + } + return nil, err + } + return rows, nil +} + +// ListOpenFlareAccessLogRegionCounts returns region counts for access logs (v1 stub). +func ListOpenFlareAccessLogRegionCounts(_ context.Context, _ string, _ time.Time, _ int) ([]*OpenFlareAccessLogRegionCount, error) { + return []*OpenFlareAccessLogRegionCount{}, nil +} + +// ListOpenFlareActiveHealthEvents returns active health events across all nodes. +func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*OpenFlareHealthEvent, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var rows []*OpenFlareHealthEvent + if err := conn.Where("status = ?", "active").Order("last_triggered_at desc").Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareHealthEvent{}, nil + } + return nil, err + } + return rows, nil +} + +// ListOpenFlareHealthEvents returns health events for a node. +func ListOpenFlareHealthEvents(ctx context.Context, nodeID string, activeOnly bool, limit int) ([]*OpenFlareHealthEvent, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + query := conn.Model(&OpenFlareHealthEvent{}).Where("node_id = ?", nodeID).Order("last_triggered_at desc") + if activeOnly { + query = query.Where("status = ?", "active") + } + if limit > 0 { + query = query.Limit(limit) + } + var rows []*OpenFlareHealthEvent + if err := query.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareHealthEvent{}, nil + } + return nil, err + } + return rows, nil +} + +// DeleteOpenFlareHealthEventsByNodeID deletes all health events for a node. +func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + result := conn.Where("node_id = ?", nodeID).Delete(&OpenFlareHealthEvent{}) + if result.Error != nil { + if isMissingTableError(result.Error) { + return 0, nil + } + return 0, result.Error + } + return result.RowsAffected, nil +} + +// GetOpenFlareNodeSystemProfile returns the system profile for a node. +func GetOpenFlareNodeSystemProfile(ctx context.Context, nodeID string) (*OpenFlareNodeSystemProfile, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var profile OpenFlareNodeSystemProfile + if err := conn.Where("node_id = ?", nodeID).First(&profile).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) || isMissingTableError(err) { + return nil, gorm.ErrRecordNotFound + } + return nil, err + } + return &profile, nil +} + +// ListOpenFlareNodeObservationOpenresty returns openresty observations. +func ListOpenFlareNodeObservationOpenresty(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationOpenresty, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + query := conn.Model(&OpenFlareNodeObservationOpenresty{}).Order("captured_at desc, id desc") + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + if !since.IsZero() { + query = query.Where("captured_at >= ?", since) + } + if limit > 0 { + query = query.Limit(limit) + } + var rows []*OpenFlareNodeObservationOpenresty + if err := query.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareNodeObservationOpenresty{}, nil + } + return nil, err + } + return rows, nil +} + +// ListOpenFlareNodeObservationFrps returns frps observations. +func ListOpenFlareNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + query := conn.Model(&OpenFlareNodeObservationFrps{}).Order("captured_at desc, id desc") + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + if !since.IsZero() { + query = query.Where("captured_at >= ?", since) + } + if limit > 0 { + query = query.Limit(limit) + } + var rows []*OpenFlareNodeObservationFrps + if err := query.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareNodeObservationFrps{}, nil + } + return nil, err + } + return rows, nil +} + +// ListOpenFlareAccessLogs lists access logs (v1 stub returns empty until table is migrated). +func ListOpenFlareAccessLogs(_ context.Context, _ OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) { + return []*OpenFlareAccessLog{}, nil +} + +// CountOpenFlareAccessLogs counts access logs (v1 stub). +func CountOpenFlareAccessLogs(_ context.Context, _ OpenFlareAccessLogQuery) (int64, int64, error) { + return 0, 0, nil +} + +// ListOpenFlareAccessLogBuckets lists folded access log buckets (v1 stub). +func ListOpenFlareAccessLogBuckets(_ context.Context, _ OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) { + return []*OpenFlareAccessLogBucketRow{}, nil +} + +// CountOpenFlareAccessLogBuckets counts folded access log buckets (v1 stub). +func CountOpenFlareAccessLogBuckets(_ context.Context, _ OpenFlareAccessLogBucketQuery) (int64, error) { + return 0, nil +} + +// ListOpenFlareAccessLogBucketIPs lists folded IP rows (v1 stub). +func ListOpenFlareAccessLogBucketIPs(_ context.Context, _ OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) { + return []*OpenFlareAccessLogBucketIPRow{}, nil +} + +// CountOpenFlareAccessLogBucketIPs counts folded IP rows (v1 stub). +func CountOpenFlareAccessLogBucketIPs(_ context.Context, _ OpenFlareAccessLogBucketIPQuery) (int64, error) { + return 0, nil +} + +// ListOpenFlareAccessLogIPSummaries lists IP summaries (v1 stub). +func ListOpenFlareAccessLogIPSummaries(_ context.Context, _ OpenFlareAccessLogIPSummaryQuery, _ time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) { + return []*OpenFlareAccessLogIPSummaryRow{}, nil +} + +// CountOpenFlareAccessLogIPSummaries counts IP summaries (v1 stub). +func CountOpenFlareAccessLogIPSummaries(_ context.Context, _ OpenFlareAccessLogIPSummaryQuery) (int64, error) { + return 0, nil +} + +// ListOpenFlareAccessLogIPTrend lists IP trend points (v1 stub). +func ListOpenFlareAccessLogIPTrend(_ context.Context, _ OpenFlareAccessLogIPTrendQuery) ([]*OpenFlareAccessLogIPTrendRow, error) { + return []*OpenFlareAccessLogIPTrendRow{}, nil +} + +// DeleteOpenFlareAccessLogsBefore deletes access logs before cutoff (v1 stub). +func DeleteOpenFlareAccessLogsBefore(_ context.Context, _ time.Time) (int64, error) { + return 0, nil +} diff --git a/Wavelet/internal/model/openflare_option.go b/Wavelet/internal/model/openflare_option.go new file mode 100644 index 00000000..19172244 --- /dev/null +++ b/Wavelet/internal/model/openflare_option.go @@ -0,0 +1,463 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "strconv" + "strings" + "sync" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +// OpenFlareOption stores a hot-reloadable OpenFlare system option. +type OpenFlareOption struct { + Key string `json:"key" gorm:"column:key;primaryKey;size:128;not null"` + Value string `json:"value" gorm:"type:text;not null"` +} + +// TableName returns the OpenFlare options table name. +func (OpenFlareOption) TableName() string { + return "of_options" +} + +// OptionMap holds the in-memory option snapshot for fast reads and hot reload. +var ( + OptionMap map[string]string + OptionMapRWMutex sync.RWMutex + + // StartTime records process start time (seconds) for legacy /api/status. + StartTime = time.Now().Unix() + + // Hot-reload mirrors for frequently read options (legacy OpenFlare keys). + SystemName = "OpenFlare" + ServerAddress = "" + Footer = "" + HomePageLink = "" + PasswordLoginEnabled = true + CapLoginEnabled = true + PasswordRegisterEnabled = false + EmailVerificationEnabled = false + GitHubOAuthEnabled = false + WeChatAuthEnabled = false + GitHubClientId = "" + GitHubClientSecret = "" + WeChatServerAddress = "" + WeChatServerToken = "" + WeChatAccountQRCodeImageURL = "" + SMTPServer = "" + SMTPPort = 587 + SMTPAccount = "" + SMTPToken = "" + AgentDiscoveryToken = "" + AgentHeartbeatInterval = 10000 + AgentWebsocketUpgradeEnabled = true + NodeOfflineThreshold = 2 * time.Minute + AgentUpdateRepo = "Rain-kl/OpenFlare" + GeoIPProvider = "ipinfo" + DatabaseAutoCleanupEnabled = false + DatabaseAutoCleanupRetentionDays = 30 + UptimeKumaEnabled = false + UptimeKumaUrl = "" + UptimeKumaUsername = "" + UptimeKumaPassword = "" + UptimeKumaMonitorScope = "all" + UptimeKumaSelectedSites = "" + UptimeKumaSyncInterval = 5 + UptimeKumaInterval = 60 + UptimeKumaRetry = 0 + UptimeKumaRetryInterval = 60 + UptimeKumaTimeout = 48 + OpenRestyDefaultServerReturnStatus = 421 + OpenRestyWorkerProcesses = "auto" + OpenRestyWorkerConnections = 4096 + OpenRestyWorkerRlimitNofile = 65535 + OpenRestyEventsUse = "epoll" + OpenRestyEventsMultiAcceptEnabled = true + OpenRestyKeepaliveTimeout = 20 + OpenRestyKeepaliveRequests = 1000 + OpenRestyClientHeaderTimeout = 15 + OpenRestyClientBodyTimeout = 15 + OpenRestyClientMaxBodySize = "64m" + OpenRestyLargeClientHeaderBuffers = "4 16k" + OpenRestySendTimeout = 30 + OpenRestyResolvers = "" + OpenRestyProxyConnectTimeout = 3 + OpenRestyProxySendTimeout = 60 + OpenRestyProxyReadTimeout = 60 + OpenRestyWebsocketEnabled = true + OpenRestyHTTP3Enabled = true + OpenRestyProxyRequestBufferingEnabled = false + OpenRestyProxyBufferingEnabled = true + OpenRestyProxyBuffers = "16 16k" + OpenRestyProxyBufferSize = "8k" + OpenRestyProxyBusyBuffersSize = "64k" + OpenRestyGzipEnabled = true + OpenRestyGzipMinLength = 1024 + OpenRestyGzipCompLevel = 5 + OpenRestyCacheEnabled = false + OpenRestyCachePath = "" + OpenRestyCacheLevels = "1:2" + OpenRestyCacheInactive = "30m" + OpenRestyCacheMaxSize = "1g" + OpenRestyCacheKeyTemplate = "$scheme$host$request_uri" + OpenRestyCacheLockEnabled = true + OpenRestyCacheLockTimeout = "5s" + OpenRestyCacheUseStale = "error timeout updating http_500 http_502 http_503 http_504" + OpenRestyMainConfigTemplate = defaultOpenRestyMainConfigTemplate + GlobalApiRateLimitNum = 300 + GlobalApiRateLimitDuration int64 = 3 * 60 + GlobalWebRateLimitNum = 300 + GlobalWebRateLimitDuration int64 = 3 * 60 + CriticalRateLimitNum = 100 + CriticalRateLimitDuration int64 = 20 * 60 +) + +const defaultOpenRestyMainConfigTemplate = `# This file is generated by OpenFlare. Do not edit manually. +worker_processes {{OpenRestyWorkerProcesses}}; +worker_rlimit_nofile {{OpenRestyWorkerRlimitNofile}}; +pid logs/nginx.pid; +error_log {{OpenRestyErrorLogPath}} warn; + +events { + worker_connections {{OpenRestyWorkerConnections}}; +{{OpenRestyEventsUseDirective}}{{OpenRestyEventsMultiAcceptDirective}}} + +http { + include mime.types; + default_type application/octet-stream; +{{OpenRestyConnectionUpgradeMap}}{{OpenRestyDefaultServerBlock}} log_format openflare_json escape=json '{"ts":"$time_iso8601","host":"$host","path":"$request_uri","remote_addr":"$remote_addr","status":$status,"request_time":$request_time,"bytes_sent":$body_bytes_sent,"request_length":$request_length}'; + access_log {{OpenRestyAccessLogPath}} openflare_json; + sendfile on; + tcp_nopush on; + tcp_nodelay on; + keepalive_timeout {{OpenRestyKeepaliveTimeout}}; + keepalive_requests {{OpenRestyKeepaliveRequests}}; + client_header_timeout {{OpenRestyClientHeaderTimeout}}; + client_body_timeout {{OpenRestyClientBodyTimeout}}; + client_max_body_size {{OpenRestyClientMaxBodySize}}; + large_client_header_buffers {{OpenRestyLargeClientHeaderBuffers}}; + send_timeout {{OpenRestySendTimeout}}; + proxy_connect_timeout {{OpenRestyProxyConnectTimeout}}; + proxy_send_timeout {{OpenRestyProxySendTimeout}}; + proxy_read_timeout {{OpenRestyProxyReadTimeout}}; + proxy_request_buffering {{OpenRestyProxyRequestBuffering}}; + proxy_buffering {{OpenRestyProxyBuffering}}; + proxy_buffers {{OpenRestyProxyBuffers}}; + proxy_buffer_size {{OpenRestyProxyBufferSize}}; + proxy_busy_buffers_size {{OpenRestyProxyBusyBuffersSize}}; + gzip {{OpenRestyGzip}}; + gzip_min_length {{OpenRestyGzipMinLength}}; + gzip_comp_level {{OpenRestyGzipCompLevel}}; +{{OpenRestyResolverDirective}}{{OpenRestyCacheBlock}} include {{OpenRestyRouteConfigInclude}}; +} +` + +// DefaultOpenFlareOptions returns built-in defaults keyed by legacy OpenFlare option names. +func DefaultOpenFlareOptions() map[string]string { + return map[string]string{ + "PasswordLoginEnabled": strconv.FormatBool(PasswordLoginEnabled), + "CapLoginEnabled": strconv.FormatBool(CapLoginEnabled), + "PasswordRegisterEnabled": strconv.FormatBool(PasswordRegisterEnabled), + "EmailVerificationEnabled": strconv.FormatBool(EmailVerificationEnabled), + "GitHubOAuthEnabled": strconv.FormatBool(GitHubOAuthEnabled), + "WeChatAuthEnabled": strconv.FormatBool(WeChatAuthEnabled), + "SMTPServer": "", + "SMTPPort": strconv.Itoa(SMTPPort), + "SMTPAccount": "", + "SMTPToken": "", + "Notice": "", + "About": "", + "Footer": Footer, + "HomePageLink": HomePageLink, + "SystemName": SystemName, + "ServerAddress": "", + "GitHubClientId": "", + "GitHubClientSecret": "", + "WeChatServerAddress": "", + "WeChatServerToken": "", + "WeChatAccountQRCodeImageURL": "", + "AgentDiscoveryToken": "", + "AgentHeartbeatInterval": strconv.Itoa(AgentHeartbeatInterval), + "AgentWebsocketUpgradeEnabled": strconv.FormatBool(AgentWebsocketUpgradeEnabled), + "NodeOfflineThreshold": strconv.Itoa(int(NodeOfflineThreshold.Milliseconds())), + "AgentUpdateRepo": AgentUpdateRepo, + "GeoIPProvider": GeoIPProvider, + "DatabaseAutoCleanupEnabled": strconv.FormatBool(DatabaseAutoCleanupEnabled), + "UptimeKumaEnabled": strconv.FormatBool(UptimeKumaEnabled), + "UptimeKumaUrl": UptimeKumaUrl, + "UptimeKumaUsername": UptimeKumaUsername, + "UptimeKumaPassword": UptimeKumaPassword, + "UptimeKumaMonitorScope": UptimeKumaMonitorScope, + "UptimeKumaSelectedSites": UptimeKumaSelectedSites, + "UptimeKumaSyncInterval": strconv.Itoa(UptimeKumaSyncInterval), + "UptimeKumaInterval": strconv.Itoa(UptimeKumaInterval), + "UptimeKumaRetry": strconv.Itoa(UptimeKumaRetry), + "UptimeKumaRetryInterval": strconv.Itoa(UptimeKumaRetryInterval), + "UptimeKumaTimeout": strconv.Itoa(UptimeKumaTimeout), + "DatabaseAutoCleanupRetentionDays": strconv.Itoa(DatabaseAutoCleanupRetentionDays), + "OpenRestyDefaultServerReturnStatus": strconv.Itoa(OpenRestyDefaultServerReturnStatus), + "OpenRestyWorkerProcesses": OpenRestyWorkerProcesses, + "OpenRestyWorkerConnections": strconv.Itoa(OpenRestyWorkerConnections), + "OpenRestyWorkerRlimitNofile": strconv.Itoa(OpenRestyWorkerRlimitNofile), + "OpenRestyEventsUse": OpenRestyEventsUse, + "OpenRestyEventsMultiAcceptEnabled": strconv.FormatBool(OpenRestyEventsMultiAcceptEnabled), + "OpenRestyKeepaliveTimeout": strconv.Itoa(OpenRestyKeepaliveTimeout), + "OpenRestyKeepaliveRequests": strconv.Itoa(OpenRestyKeepaliveRequests), + "OpenRestyClientHeaderTimeout": strconv.Itoa(OpenRestyClientHeaderTimeout), + "OpenRestyClientBodyTimeout": strconv.Itoa(OpenRestyClientBodyTimeout), + "OpenRestyClientMaxBodySize": OpenRestyClientMaxBodySize, + "OpenRestyLargeClientHeaderBuffers": OpenRestyLargeClientHeaderBuffers, + "OpenRestySendTimeout": strconv.Itoa(OpenRestySendTimeout), + "OpenRestyProxyConnectTimeout": strconv.Itoa(OpenRestyProxyConnectTimeout), + "OpenRestyProxySendTimeout": strconv.Itoa(OpenRestyProxySendTimeout), + "OpenRestyProxyReadTimeout": strconv.Itoa(OpenRestyProxyReadTimeout), + "OpenRestyWebsocketEnabled": strconv.FormatBool(OpenRestyWebsocketEnabled), + "OpenRestyHTTP3Enabled": strconv.FormatBool(OpenRestyHTTP3Enabled), + "OpenRestyProxyRequestBufferingEnabled": strconv.FormatBool(OpenRestyProxyRequestBufferingEnabled), + "OpenRestyProxyBufferingEnabled": strconv.FormatBool(OpenRestyProxyBufferingEnabled), + "OpenRestyProxyBuffers": OpenRestyProxyBuffers, + "OpenRestyProxyBufferSize": OpenRestyProxyBufferSize, + "OpenRestyProxyBusyBuffersSize": OpenRestyProxyBusyBuffersSize, + "OpenRestyGzipEnabled": strconv.FormatBool(OpenRestyGzipEnabled), + "OpenRestyGzipMinLength": strconv.Itoa(OpenRestyGzipMinLength), + "OpenRestyGzipCompLevel": strconv.Itoa(OpenRestyGzipCompLevel), + "OpenRestyCacheEnabled": strconv.FormatBool(OpenRestyCacheEnabled), + "OpenRestyCachePath": OpenRestyCachePath, + "OpenRestyCacheLevels": OpenRestyCacheLevels, + "OpenRestyCacheInactive": OpenRestyCacheInactive, + "OpenRestyCacheMaxSize": OpenRestyCacheMaxSize, + "OpenRestyCacheKeyTemplate": OpenRestyCacheKeyTemplate, + "OpenRestyCacheLockEnabled": strconv.FormatBool(OpenRestyCacheLockEnabled), + "OpenRestyCacheLockTimeout": OpenRestyCacheLockTimeout, + "OpenRestyCacheUseStale": OpenRestyCacheUseStale, + "OpenRestyMainConfigTemplate": OpenRestyMainConfigTemplate, + "GlobalApiRateLimitNum": strconv.Itoa(GlobalApiRateLimitNum), + "GlobalApiRateLimitDuration": strconv.FormatInt(GlobalApiRateLimitDuration, 10), + "GlobalWebRateLimitNum": strconv.Itoa(GlobalWebRateLimitNum), + "GlobalWebRateLimitDuration": strconv.FormatInt(GlobalWebRateLimitDuration, 10), + "CriticalRateLimitNum": strconv.Itoa(CriticalRateLimitNum), + "CriticalRateLimitDuration": strconv.FormatInt(CriticalRateLimitDuration, 10), + } +} + +// InitOptionMap seeds defaults and overlays persisted options from of_options. +func InitOptionMap(ctx context.Context) error { + OptionMapRWMutex.Lock() + OptionMap = DefaultOpenFlareOptions() + OptionMapRWMutex.Unlock() + + options, err := ListOpenFlareOptions(ctx) + if err != nil { + return err + } + for _, option := range options { + applyOptionMap(option.Key, option.Value) + } + return nil +} + +// ListOpenFlareOptions returns all persisted options. +func ListOpenFlareOptions(ctx context.Context) ([]OpenFlareOption, error) { + var options []OpenFlareOption + if err := db.DB(ctx).Find(&options).Error; err != nil { + return nil, err + } + return options, nil +} + +// UpdateOpenFlareOption updates a single option in DB and memory. +func UpdateOpenFlareOption(ctx context.Context, key, value string) error { + return UpdateOpenFlareOptions(ctx, []OpenFlareOption{{Key: key, Value: value}}) +} + +// UpdateOpenFlareOptions batch-updates options in a transaction and refreshes OptionMap. +func UpdateOpenFlareOptions(ctx context.Context, options []OpenFlareOption) error { + if len(options) == 0 { + return nil + } + + if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + for _, item := range options { + if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" { + continue + } + option := OpenFlareOption{Key: item.Key} + if err := tx.FirstOrCreate(&option, OpenFlareOption{Key: item.Key}).Error; err != nil { + return err + } + option.Value = item.Value + if err := tx.Save(&option).Error; err != nil { + return err + } + } + return nil + }); err != nil { + return err + } + + for _, item := range options { + if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" { + continue + } + applyOptionMap(item.Key, item.Value) + } + return nil +} + +// OptionValue returns a snapshot value from OptionMap. +func OptionValue(key string) string { + OptionMapRWMutex.RLock() + defer OptionMapRWMutex.RUnlock() + if OptionMap == nil { + return "" + } + return OptionMap[key] +} + +// ResetOptionMapForTest clears in-memory option state for unit tests. +func ResetOptionMapForTest() { + OptionMapRWMutex.Lock() + OptionMap = nil + OptionMapRWMutex.Unlock() +} + +func applyOptionMap(key, value string) { + OptionMapRWMutex.Lock() + if OptionMap == nil { + OptionMap = make(map[string]string) + } + OptionMap[key] = value + if strings.HasSuffix(key, "Enabled") { + boolValue := value == "true" + switch key { + case "PasswordRegisterEnabled": + PasswordRegisterEnabled = boolValue + case "PasswordLoginEnabled": + PasswordLoginEnabled = boolValue + case "CapLoginEnabled": + CapLoginEnabled = boolValue + case "EmailVerificationEnabled": + EmailVerificationEnabled = boolValue + case "GitHubOAuthEnabled": + GitHubOAuthEnabled = boolValue + case "WeChatAuthEnabled": + WeChatAuthEnabled = boolValue + } + } + switch key { + case "SMTPServer": + SMTPServer = value + case "SMTPPort": + if intValue, err := strconv.Atoi(value); err == nil { + SMTPPort = intValue + } + case "SMTPAccount": + SMTPAccount = value + case "SMTPToken": + SMTPToken = value + case "ServerAddress": + ServerAddress = value + case "GitHubClientId": + GitHubClientId = value + case "GitHubClientSecret": + GitHubClientSecret = value + case "Footer": + Footer = value + case "HomePageLink": + HomePageLink = value + case "SystemName": + SystemName = value + case "WeChatServerAddress": + WeChatServerAddress = value + case "WeChatServerToken": + WeChatServerToken = value + case "WeChatAccountQRCodeImageURL": + WeChatAccountQRCodeImageURL = value + case "AgentDiscoveryToken": + AgentDiscoveryToken = value + case "AgentHeartbeatInterval": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + AgentHeartbeatInterval = v + } + case "AgentWebsocketUpgradeEnabled": + AgentWebsocketUpgradeEnabled = value == "true" + case "NodeOfflineThreshold": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + NodeOfflineThreshold = time.Duration(v) * time.Millisecond + } + case "AgentUpdateRepo": + if value != "" { + AgentUpdateRepo = value + } + case "GeoIPProvider": + GeoIPProvider = value + case "UptimeKumaEnabled": + UptimeKumaEnabled = value == "true" + case "UptimeKumaUrl": + UptimeKumaUrl = value + case "UptimeKumaUsername": + UptimeKumaUsername = value + case "UptimeKumaPassword": + UptimeKumaPassword = value + case "UptimeKumaMonitorScope": + UptimeKumaMonitorScope = value + case "UptimeKumaSelectedSites": + UptimeKumaSelectedSites = value + case "UptimeKumaSyncInterval": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + UptimeKumaSyncInterval = v + } + case "UptimeKumaInterval": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + UptimeKumaInterval = v + } + case "UptimeKumaRetry": + if v, err := strconv.Atoi(value); err == nil && v >= 0 { + UptimeKumaRetry = v + } + case "UptimeKumaRetryInterval": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + UptimeKumaRetryInterval = v + } + case "UptimeKumaTimeout": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + UptimeKumaTimeout = v + } + case "DatabaseAutoCleanupEnabled": + DatabaseAutoCleanupEnabled = value == "true" + case "DatabaseAutoCleanupRetentionDays": + if v, err := strconv.Atoi(value); err == nil && v >= 1 { + DatabaseAutoCleanupRetentionDays = v + } + case "GlobalApiRateLimitNum": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + GlobalApiRateLimitNum = v + } + case "GlobalApiRateLimitDuration": + if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 { + GlobalApiRateLimitDuration = v + } + case "GlobalWebRateLimitNum": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + GlobalWebRateLimitNum = v + } + case "GlobalWebRateLimitDuration": + if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 { + GlobalWebRateLimitDuration = v + } + case "CriticalRateLimitNum": + if v, err := strconv.Atoi(value); err == nil && v > 0 { + CriticalRateLimitNum = v + } + case "CriticalRateLimitDuration": + if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 { + CriticalRateLimitDuration = v + } + } + OptionMapRWMutex.Unlock() +} diff --git a/Wavelet/internal/model/openflare_origin.go b/Wavelet/internal/model/openflare_origin.go new file mode 100644 index 00000000..1d285bd3 --- /dev/null +++ b/Wavelet/internal/model/openflare_origin.go @@ -0,0 +1,133 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" +) + +// Origin OpenFlare 源站实体。 +type Origin struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Name string `json:"name" gorm:"size:255;not null"` + Address string `json:"address" gorm:"uniqueIndex;size:255;not null"` + Remark string `json:"remark" gorm:"size:255"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名。 +func (Origin) TableName() string { + return "of_origins" +} + +// OriginRouteCount 源站关联的代理规则数量。 +type OriginRouteCount struct { + OriginID uint `json:"origin_id"` + RouteCount int64 `json:"route_count"` +} + +// OriginProxyRoute 源站模块查询代理规则时使用的最小字段集。 +type OriginProxyRoute struct { + ID uint `gorm:"column:id;primaryKey"` + OriginID *uint `gorm:"column:origin_id"` + Domain string `gorm:"column:domain"` + OriginURL string `gorm:"column:origin_url"` + Upstreams string `gorm:"column:upstreams"` + Enabled bool `gorm:"column:enabled"` + UpdatedAt time.Time `gorm:"column:updated_at"` +} + +// TableName 表名。 +func (OriginProxyRoute) TableName() string { + return "of_proxy_routes" +} + +// HasProxyRoutesTable 判断代理规则表是否已迁移。 +func HasProxyRoutesTable(ctx context.Context) bool { + return db.DB(ctx).Migrator().HasTable(&OriginProxyRoute{}) +} + +// ListOrigins 列出全部源站。 +func ListOrigins(ctx context.Context) ([]Origin, error) { + var origins []Origin + if err := db.DB(ctx).Order("id desc").Find(&origins).Error; err != nil { + return nil, err + } + return origins, nil +} + +// GetOriginByID 按 ID 查询源站。 +func GetOriginByID(ctx context.Context, id uint) (*Origin, error) { + var origin Origin + if err := db.DB(ctx).First(&origin, id).Error; err != nil { + return nil, err + } + return &origin, nil +} + +// GetOriginByAddress 按地址查询源站。 +func GetOriginByAddress(ctx context.Context, address string) (*Origin, error) { + var origin Origin + if err := db.DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil { + return nil, err + } + return &origin, nil +} + +// CreateOriginRecord 创建源站。 +func CreateOriginRecord(ctx context.Context, origin *Origin) error { + return db.DB(ctx).Create(origin).Error +} + +// SaveOrigin 保存源站。 +func SaveOrigin(ctx context.Context, origin *Origin) error { + return db.DB(ctx).Save(origin).Error +} + +// DeleteOriginRecord 删除源站。 +func DeleteOriginRecord(ctx context.Context, id uint) error { + return db.DB(ctx).Delete(&Origin{}, id).Error +} + +// ListOriginRouteCounts 统计各源站关联的代理规则数量。 +func ListOriginRouteCounts(ctx context.Context) ([]OriginRouteCount, error) { + if !HasProxyRoutesTable(ctx) { + return nil, nil + } + result := make([]OriginRouteCount, 0) + err := db.DB(ctx).Model(&OriginProxyRoute{}). + Select("origin_id, COUNT(*) AS route_count"). + Where("origin_id IS NOT NULL"). + Group("origin_id"). + Scan(&result).Error + return result, err +} + +// ListProxyRoutesByOriginID 列出源站关联的代理规则。 +func ListProxyRoutesByOriginID(ctx context.Context, originID uint) ([]OriginProxyRoute, error) { + if !HasProxyRoutesTable(ctx) { + return nil, nil + } + var routes []OriginProxyRoute + if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} + +// CountProxyRoutesByOriginID 统计源站关联的代理规则数量。 +func CountProxyRoutesByOriginID(ctx context.Context, originID uint) (int64, error) { + if !HasProxyRoutesTable(ctx) { + return 0, nil + } + var count int64 + if err := db.DB(ctx).Model(&OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} diff --git a/Wavelet/internal/model/openflare_pages.go b/Wavelet/internal/model/openflare_pages.go new file mode 100644 index 00000000..3243324d --- /dev/null +++ b/Wavelet/internal/model/openflare_pages.go @@ -0,0 +1,161 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" +) + +const ( + PagesDeploymentStatusUploaded = "uploaded" + PagesDeploymentStatusActive = "active" +) + +// PagesProject OpenFlare Pages 静态托管项目。 +type PagesProject struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Name string `json:"name" gorm:"size:255;not null"` + Slug string `json:"slug" gorm:"uniqueIndex;size:128;not null"` + Description string `json:"description" gorm:"type:text;not null;default:''"` + Enabled bool `json:"enabled" gorm:"not null;default:true"` + SPAFallbackEnabled bool `json:"spa_fallback_enabled" gorm:"not null;default:false"` + SPAFallbackPath string `json:"spa_fallback_path" gorm:"size:512;not null;default:'/index.html'"` + APIProxyEnabled bool `json:"api_proxy_enabled" gorm:"not null;default:false"` + APIProxyPath string `json:"api_proxy_path" gorm:"size:255;not null;default:''"` + APIProxyPass string `json:"api_proxy_pass" gorm:"size:2048;not null;default:''"` + APIProxyRewrite string `json:"api_proxy_rewrite" gorm:"size:255;not null;default:''"` + ActiveDeploymentID *uint `json:"active_deployment_id" gorm:"index"` + RootDir string `json:"root_dir" gorm:"size:512;not null;default:''"` + EntryFile string `json:"entry_file" gorm:"size:512;not null;default:'index.html'"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名。 +func (PagesProject) TableName() string { + return "of_pages_projects" +} + +// PagesDeployment OpenFlare Pages 不可变部署记录。 +type PagesDeployment struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + ProjectID uint `json:"project_id" gorm:"not null;index"` + DeploymentNumber int `json:"deployment_number" gorm:"not null"` + Checksum string `json:"checksum" gorm:"size:64;not null;index"` + Status string `json:"status" gorm:"size:32;not null;default:'uploaded';index"` + ArtifactPath string `json:"artifact_path" gorm:"size:2048;not null"` + FileCount int `json:"file_count" gorm:"not null;default:0"` + TotalSize int64 `json:"total_size" gorm:"not null;default:0"` + CreatedBy string `json:"created_by" gorm:"size:64;not null;default:''"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + ActivatedAt *time.Time `json:"activated_at"` +} + +// TableName 表名。 +func (PagesDeployment) TableName() string { + return "of_pages_deployments" +} + +// PagesDeploymentFile OpenFlare Pages 部署文件清单。 +type PagesDeploymentFile struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + DeploymentID uint `json:"deployment_id" gorm:"not null;index"` + Path string `json:"path" gorm:"size:2048;not null"` + Size int64 `json:"size" gorm:"not null;default:0"` + Checksum string `json:"checksum" gorm:"size:64;not null"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName 表名。 +func (PagesDeploymentFile) TableName() string { + return "of_pages_deployment_files" +} + +// HasPagesProjectsTable 判断 Pages 项目表是否已迁移。 +func HasPagesProjectsTable(ctx context.Context) bool { + return db.DB(ctx).Migrator().HasTable(&PagesProject{}) +} + +// ListPagesProjects 列出全部 Pages 项目。 +func ListPagesProjects(ctx context.Context) ([]PagesProject, error) { + var projects []PagesProject + if err := db.DB(ctx).Order("id desc").Find(&projects).Error; err != nil { + return nil, err + } + return projects, nil +} + +// GetPagesProjectByID 按 ID 查询 Pages 项目。 +func GetPagesProjectByID(ctx context.Context, id uint) (*PagesProject, error) { + var project PagesProject + if err := db.DB(ctx).First(&project, id).Error; err != nil { + return nil, err + } + return &project, nil +} + +// GetPagesProjectBySlug 按 slug 查询 Pages 项目。 +func GetPagesProjectBySlug(ctx context.Context, slug string) (*PagesProject, error) { + var project PagesProject + if err := db.DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil { + return nil, err + } + return &project, nil +} + +// CreatePagesProjectRecord 创建 Pages 项目。 +func CreatePagesProjectRecord(ctx context.Context, project *PagesProject) error { + return db.DB(ctx).Create(project).Error +} + +// ListPagesDeployments 列出项目的全部部署。 +func ListPagesDeployments(ctx context.Context, projectID uint) ([]PagesDeployment, error) { + var deployments []PagesDeployment + if err := db.DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil { + return nil, err + } + return deployments, nil +} + +// GetPagesDeploymentByID 按 ID 查询 Pages 部署。 +func GetPagesDeploymentByID(ctx context.Context, id uint) (*PagesDeployment, error) { + var deployment PagesDeployment + if err := db.DB(ctx).First(&deployment, id).Error; err != nil { + return nil, err + } + return &deployment, nil +} + +// ListPagesDeploymentFiles 列出部署文件清单。 +func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]PagesDeploymentFile, error) { + var files []PagesDeploymentFile + if err := db.DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil { + return nil, err + } + return files, nil +} + +// CountPagesDeploymentsByProjectID 统计项目部署数量。 +func CountPagesDeploymentsByProjectID(ctx context.Context, projectID uint) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// CountProxyRoutesByPagesProjectID 统计引用 Pages 项目的代理规则数量。 +func CountProxyRoutesByPagesProjectID(ctx context.Context, projectID uint) (int64, error) { + if !HasProxyRoutesTable(ctx) { + return 0, nil + } + var count int64 + if err := db.DB(ctx).Model(&ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} diff --git a/Wavelet/internal/model/openflare_proxy_route.go b/Wavelet/internal/model/openflare_proxy_route.go new file mode 100644 index 00000000..f60cc87b --- /dev/null +++ b/Wavelet/internal/model/openflare_proxy_route.go @@ -0,0 +1,115 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" +) + +// ProxyRoute OpenFlare 代理规则实体。 +type ProxyRoute struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + SiteName string `json:"site_name" gorm:"size:255;not null;default:''"` + Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"` + Domains string `json:"domains" gorm:"type:text;not null;default:'[]'"` + OriginID *uint `json:"origin_id" gorm:"index"` + OriginURL string `json:"origin_url" gorm:"size:2048;not null"` + OriginHost string `json:"origin_host" gorm:"size:255"` + Upstreams string `json:"upstreams" gorm:"type:text;not null;default:'[]'"` + Enabled bool `json:"enabled" gorm:"not null;default:true"` + EnableHTTPS bool `json:"enable_https" gorm:"column:enable_https;not null;default:false"` + CertID *uint `json:"cert_id"` + CertIDs string `json:"cert_ids" gorm:"type:text;not null;default:'[]'"` + DomainCertIDs string `json:"domain_cert_ids" gorm:"type:text;not null;default:'[]'"` + RedirectHTTP bool `json:"redirect_http" gorm:"not null;default:false"` + LimitConnPerServer int `json:"limit_conn_per_server" gorm:"not null;default:0"` + LimitConnPerIP int `json:"limit_conn_per_ip" gorm:"not null;default:0"` + LimitRate string `json:"limit_rate" gorm:"size:32;not null;default:''"` + CacheEnabled bool `json:"cache_enabled" gorm:"not null;default:false"` + CachePolicy string `json:"cache_policy" gorm:"size:32;not null;default:''"` + CacheRules string `json:"cache_rules" gorm:"type:text;not null;default:'[]'"` + CustomHeaders string `json:"custom_headers" gorm:"type:text;not null;default:'[]'"` + BasicAuthEnabled bool `json:"basic_auth_enabled" gorm:"not null;default:false"` + BasicAuthUsername string `json:"basic_auth_username" gorm:"size:255;not null;default:''"` + BasicAuthPassword string `json:"basic_auth_password" gorm:"size:255;not null;default:''"` + Remark string `json:"remark" gorm:"size:255"` + UpstreamType string `json:"upstream_type" gorm:"size:32;not null;default:'direct'"` + TunnelNodeID *uint `json:"tunnel_node_id" gorm:"index"` + TunnelTargetAddr string `json:"tunnel_target_addr" gorm:"size:512"` + TunnelTargetProtocol string `json:"tunnel_target_protocol" gorm:"size:16"` + PagesProjectID *uint `json:"pages_project_id" gorm:"index"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名。 +func (ProxyRoute) TableName() string { + return "of_proxy_routes" +} + +// ListProxyRoutes 列出全部代理规则。 +func ListProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) { + var routes []*ProxyRoute + if err := db.DB(ctx).Order("id desc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} + +// GetProxyRouteByID 按 ID 查询代理规则。 +func GetProxyRouteByID(ctx context.Context, id uint) (*ProxyRoute, error) { + var route ProxyRoute + if err := db.DB(ctx).First(&route, id).Error; err != nil { + return nil, err + } + return &route, nil +} + +// CreateProxyRouteRecord 创建代理规则。 +func CreateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error { + return db.DB(ctx).Create(route).Error +} + +// UpdateProxyRouteRecord 更新代理规则。 +func UpdateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error { + return db.DB(ctx).Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{ + "site_name": route.SiteName, + "domain": route.Domain, + "domains": route.Domains, + "origin_id": route.OriginID, + "origin_url": route.OriginURL, + "origin_host": route.OriginHost, + "upstreams": route.Upstreams, + "enabled": route.Enabled, + "enable_https": route.EnableHTTPS, + "cert_id": route.CertID, + "cert_ids": route.CertIDs, + "domain_cert_ids": route.DomainCertIDs, + "redirect_http": route.RedirectHTTP, + "limit_conn_per_server": route.LimitConnPerServer, + "limit_conn_per_ip": route.LimitConnPerIP, + "limit_rate": route.LimitRate, + "cache_enabled": route.CacheEnabled, + "cache_policy": route.CachePolicy, + "cache_rules": route.CacheRules, + "custom_headers": route.CustomHeaders, + "basic_auth_enabled": route.BasicAuthEnabled, + "basic_auth_username": route.BasicAuthUsername, + "basic_auth_password": route.BasicAuthPassword, + "remark": route.Remark, + "upstream_type": route.UpstreamType, + "tunnel_node_id": route.TunnelNodeID, + "tunnel_target_addr": route.TunnelTargetAddr, + "tunnel_target_protocol": route.TunnelTargetProtocol, + "pages_project_id": route.PagesProjectID, + }).Error +} + +// DeleteProxyRouteRecord 删除代理规则。 +func DeleteProxyRouteRecord(ctx context.Context, id uint) error { + return db.DB(ctx).Delete(&ProxyRoute{}, id).Error +} diff --git a/Wavelet/internal/model/openflare_tls.go b/Wavelet/internal/model/openflare_tls.go new file mode 100644 index 00000000..36625609 --- /dev/null +++ b/Wavelet/internal/model/openflare_tls.go @@ -0,0 +1,139 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" +) + +// TLSCertificate OpenFlare TLS 证书实体。 +type TLSCertificate struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Name string `json:"name" gorm:"uniqueIndex;size:255;not null"` + CertPEM string `json:"-" gorm:"type:text;not null"` + KeyPEM string `json:"-" gorm:"type:text;not null"` + NotBefore time.Time `json:"not_before"` + NotAfter time.Time `json:"not_after"` + Remark string `json:"remark" gorm:"size:255"` + Provider string `json:"provider" gorm:"size:64;default:upload"` + AcmeAccountID uint `json:"acme_account_id"` + DnsAccountID uint `json:"dns_account_id"` + KeyAlgorithm string `json:"key_algorithm" gorm:"size:32"` + AutoRenew bool `json:"auto_renew"` + PrimaryDomain string `json:"primary_domain" gorm:"size:255"` + OtherDomains string `json:"other_domains" gorm:"type:text"` + DisableCNAME bool `json:"disable_cname"` + SkipDNS bool `json:"skip_dns"` + DNS1 string `json:"dns1" gorm:"size:128"` + DNS2 string `json:"dns2" gorm:"size:128"` + ApplyStatus string `json:"apply_status" gorm:"size:64;default:ready"` + ApplyMessage string `json:"apply_message" gorm:"type:text"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名。 +func (TLSCertificate) TableName() string { + return "of_tls_certificates" +} + +// TLSProxyRouteRef 删除证书时检查代理规则引用的最小字段集。 +type TLSProxyRouteRef struct { + ID uint `gorm:"column:id;primaryKey"` + CertID *uint `gorm:"column:cert_id"` + CertIDs string `gorm:"column:cert_ids"` + DomainCertIDs string `gorm:"column:domain_cert_ids"` +} + +// TableName 表名。 +func (TLSProxyRouteRef) TableName() string { + return "of_proxy_routes" +} + +// HasTLSProxyRoutesTable 判断代理规则表是否已迁移。 +func HasTLSProxyRoutesTable(ctx context.Context) bool { + return db.DB(ctx).Migrator().HasTable(&TLSProxyRouteRef{}) +} + +// ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。 +func ListTLSCertificates(ctx context.Context) ([]TLSCertificate, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var certificates []TLSCertificate + if err := conn.Order("id desc").Find(&certificates).Error; err != nil { + return nil, err + } + return certificates, nil +} + +// GetTLSCertificateByID 按 ID 查询证书。 +func GetTLSCertificateByID(ctx context.Context, id uint) (*TLSCertificate, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var certificate TLSCertificate + if err := conn.First(&certificate, id).Error; err != nil { + return nil, err + } + return &certificate, nil +} + +// CreateTLSCertificateRecord 创建证书记录。 +func CreateTLSCertificateRecord(ctx context.Context, certificate *TLSCertificate) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(certificate).Error +} + +// SaveTLSCertificate 保存证书记录。 +func SaveTLSCertificate(ctx context.Context, certificate *TLSCertificate) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(certificate).Error +} + +// DeleteTLSCertificateRecord 删除证书记录。 +func DeleteTLSCertificateRecord(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Delete(&TLSCertificate{}, id).Error +} + +// CountTLSCertificatesByDNSAccountID 统计引用指定 DNS 账号的证书数量。 +func CountTLSCertificatesByDNSAccountID(ctx context.Context, dnsAccountID uint) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + var count int64 + if err := conn.Model(&TLSCertificate{}).Where("dns_account_id = ?", dnsAccountID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// ListTLSProxyRouteRefs 列出代理规则证书引用字段。 +func ListTLSProxyRouteRefs(ctx context.Context) ([]TLSProxyRouteRef, error) { + if !HasTLSProxyRoutesTable(ctx) { + return nil, nil + } + var routes []TLSProxyRouteRef + if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} diff --git a/Wavelet/internal/model/openflare_waf.go b/Wavelet/internal/model/openflare_waf.go new file mode 100644 index 00000000..b9045e31 --- /dev/null +++ b/Wavelet/internal/model/openflare_waf.go @@ -0,0 +1,345 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +// OpenFlareWAFRuleGroup stores a WAF rule group. +type OpenFlareWAFRuleGroup struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Name string `json:"name" gorm:"size:255;not null"` + Enabled bool `json:"enabled" gorm:"not null;default:true"` + IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"` + BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"` + BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"` + IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"` + IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"` + IPWhitelistGroups string `json:"ip_whitelist_group_ids" gorm:"column:ip_whitelist_groups;type:text;not null;default:'[]'"` + IPBlacklistGroups string `json:"ip_blacklist_group_ids" gorm:"column:ip_blacklist_groups;type:text;not null;default:'[]'"` + CountryWhitelist string `json:"country_whitelist" gorm:"type:text;not null;default:'[]'"` + CountryBlacklist string `json:"country_blacklist" gorm:"type:text;not null;default:'[]'"` + RegionWhitelist string `json:"region_whitelist" gorm:"type:text;not null;default:'[]'"` + RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"` + PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"` + PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"` + Remark string `json:"remark" gorm:"size:255"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareWAFRuleGroup) TableName() string { + return "of_waf_rule_groups" +} + +// OpenFlareWAFIPGroup stores a WAF IP group. +type OpenFlareWAFIPGroup struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + Name string `json:"name" gorm:"size:255;not null"` + Type string `json:"type" gorm:"size:32;not null;index"` + Enabled bool `json:"enabled" gorm:"not null;default:true"` + IPList string `json:"ip_list" gorm:"type:text;not null;default:'[]'"` + AutoConfig string `json:"auto_config" gorm:"type:text;not null;default:'{}'"` + ExtIPs string `json:"ext_ips" gorm:"type:text;not null;default:'[]'"` + SubscriptionURL string `json:"subscription_url" gorm:"size:2048;not null;default:''"` + SubscriptionFormat string `json:"subscription_format" gorm:"size:32;not null;default:'text'"` + SubscriptionMappingRule string `json:"subscription_mapping_rule" gorm:"size:255;not null;default:''"` + SyncIntervalMinutes int `json:"sync_interval_minutes" gorm:"not null;default:1440"` + LastSyncedAt *time.Time `json:"last_synced_at"` + NextSyncAt *time.Time `json:"next_sync_at" gorm:"index"` + LastSyncStatus string `json:"last_sync_status" gorm:"size:32;not null;default:''"` + LastSyncMessage string `json:"last_sync_message" gorm:"type:text;not null;default:''"` + Remark string `json:"remark" gorm:"size:255"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareWAFIPGroup) TableName() string { + return "of_waf_ip_groups" +} + +// OpenFlareWAFRuleGroupBinding binds a rule group to a proxy route. +type OpenFlareWAFRuleGroupBinding struct { + ID uint `json:"id" gorm:"primaryKey;autoIncrement"` + RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route"` + ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route;index"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareWAFRuleGroupBinding) TableName() string { + return "of_waf_rule_group_bindings" +} + +func wafDB(ctx context.Context) (*gorm.DB, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + return conn, nil +} + +// ListOpenFlareWAFRuleGroups returns all rule groups. +func ListOpenFlareWAFRuleGroups(ctx context.Context) ([]*OpenFlareWAFRuleGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var groups []*OpenFlareWAFRuleGroup + if err = conn.Order("is_global desc").Order("id asc").Find(&groups).Error; err != nil { + return nil, err + } + return groups, nil +} + +// GetOpenFlareWAFRuleGroupByID returns a rule group by id. +func GetOpenFlareWAFRuleGroupByID(ctx context.Context, id uint) (*OpenFlareWAFRuleGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var group OpenFlareWAFRuleGroup + if err = conn.First(&group, id).Error; err != nil { + return nil, err + } + return &group, nil +} + +// GetGlobalOpenFlareWAFRuleGroup returns the global rule group if present. +func GetGlobalOpenFlareWAFRuleGroup(ctx context.Context) (*OpenFlareWAFRuleGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var group OpenFlareWAFRuleGroup + if err = conn.Where("is_global = ?", true).Order("id asc").First(&group).Error; err != nil { + return nil, err + } + return &group, nil +} + +// CreateOpenFlareWAFRuleGroup inserts a rule group. +func CreateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Create(group).Error +} + +// UpdateOpenFlareWAFRuleGroup updates mutable rule group fields. +func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Model(&OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ + "name": group.Name, + "enabled": group.Enabled, + "is_global": group.IsGlobal, + "block_status_code": group.BlockStatusCode, + "block_response_body": group.BlockResponseBody, + "ip_whitelist": group.IPWhitelist, + "ip_blacklist": group.IPBlacklist, + "ip_whitelist_groups": group.IPWhitelistGroups, + "ip_blacklist_groups": group.IPBlacklistGroups, + "country_whitelist": group.CountryWhitelist, + "country_blacklist": group.CountryBlacklist, + "region_whitelist": group.RegionWhitelist, + "region_blacklist": group.RegionBlacklist, + "pow_enabled": group.PoWEnabled, + "pow_config": group.PoWConfig, + "remark": group.Remark, + }).Error +} + +// DeleteOpenFlareWAFRuleGroup removes a rule group. +func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Delete(&OpenFlareWAFRuleGroup{}, id).Error +} + +// ListOpenFlareWAFIPGroups returns all IP groups. +func ListOpenFlareWAFIPGroups(ctx context.Context) ([]*OpenFlareWAFIPGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var groups []*OpenFlareWAFIPGroup + if err = conn.Order("type asc").Order("id asc").Find(&groups).Error; err != nil { + return nil, err + } + return groups, nil +} + +// GetOpenFlareWAFIPGroupByID returns an IP group by id. +func GetOpenFlareWAFIPGroupByID(ctx context.Context, id uint) (*OpenFlareWAFIPGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var group OpenFlareWAFIPGroup + if err = conn.First(&group, id).Error; err != nil { + return nil, err + } + return &group, nil +} + +// CreateOpenFlareWAFIPGroup inserts an IP group. +func CreateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Create(group).Error +} + +// UpdateOpenFlareWAFIPGroup updates mutable IP group fields. +func UpdateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Model(&OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ + "name": group.Name, + "type": group.Type, + "enabled": group.Enabled, + "ip_list": group.IPList, + "auto_config": group.AutoConfig, + "ext_ips": group.ExtIPs, + "subscription_url": group.SubscriptionURL, + "subscription_format": group.SubscriptionFormat, + "subscription_mapping_rule": group.SubscriptionMappingRule, + "sync_interval_minutes": group.SyncIntervalMinutes, + "next_sync_at": group.NextSyncAt, + "last_sync_status": group.LastSyncStatus, + "last_sync_message": group.LastSyncMessage, + "remark": group.Remark, + }).Error +} + +// DeleteOpenFlareWAFIPGroup removes an IP group. +func DeleteOpenFlareWAFIPGroup(ctx context.Context, id uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Delete(&OpenFlareWAFIPGroup{}, id).Error +} + +// ListOpenFlareWAFRuleGroupBindings returns all bindings. +func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]OpenFlareWAFRuleGroupBinding, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var bindings []OpenFlareWAFRuleGroupBinding + if err = conn.Order("rule_group_id asc").Order("proxy_route_id asc").Find(&bindings).Error; err != nil { + return nil, err + } + return bindings, nil +} + +// ListOpenFlareWAFRuleGroupBindingsByRouteID returns bindings for a proxy route. +func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uint) ([]OpenFlareWAFRuleGroupBinding, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var bindings []OpenFlareWAFRuleGroupBinding + if err = conn.Where("proxy_route_id = ?", routeID).Order("rule_group_id asc").Find(&bindings).Error; err != nil { + return nil, err + } + return bindings, nil +} + +// ReplaceOpenFlareWAFRuleGroupBindings replaces bindings for a rule group. +func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, routeIDs []uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Transaction(func(tx *gorm.DB) error { + if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil { + return err + } + for _, routeID := range routeIDs { + binding := OpenFlareWAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID} + if err = tx.Create(&binding).Error; err != nil { + return err + } + } + return nil + }) +} + +// ReplaceOpenFlareWAFSiteRuleGroupBindings replaces bindings for a proxy route. +func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint, groupIDs []uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Transaction(func(tx *gorm.DB) error { + if err = tx.Where("proxy_route_id = ?", routeID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil { + return err + } + for _, groupID := range groupIDs { + binding := OpenFlareWAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID} + if err = tx.Create(&binding).Error; err != nil { + return err + } + } + return nil + }) +} + +// DeleteOpenFlareWAFRuleGroupBindingsByGroupID removes bindings for a rule group. +func DeleteOpenFlareWAFRuleGroupBindingsByGroupID(ctx context.Context, groupID uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error +} + +// DeleteOpenFlareWAFRuleGroupWithBindings removes a rule group and its bindings. +func DeleteOpenFlareWAFRuleGroupWithBindings(ctx context.Context, groupID uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Transaction(func(tx *gorm.DB) error { + if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil { + return err + } + return tx.Delete(&OpenFlareWAFRuleGroup{}, groupID).Error + }) +} + +// GetOpenFlareProxyRouteByID returns a proxy route by id when the table exists. +func GetOpenFlareProxyRouteByID(ctx context.Context, id uint) (*OriginProxyRoute, error) { + if !HasProxyRoutesTable(ctx) { + return nil, gorm.ErrRecordNotFound + } + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var route OriginProxyRoute + if err = conn.First(&route, id).Error; err != nil { + return nil, err + } + return &route, nil +} diff --git a/Wavelet/internal/router/router.go b/Wavelet/internal/router/router.go index c21fecac..0784ae3d 100644 --- a/Wavelet/internal/router/router.go +++ b/Wavelet/internal/router/router.go @@ -15,6 +15,7 @@ import ( "syscall" "time" + oflegacy "github.com/Rain-kl/Wavelet/internal/apps/openflare/legacy" "github.com/Rain-kl/Wavelet/internal/apps/risk_control" router_root "github.com/Rain-kl/Wavelet/internal/router/root" v1 "github.com/Rain-kl/Wavelet/internal/router/v1" @@ -113,6 +114,9 @@ func registerRoutes(r *gin.Engine) { apiGroup := r.Group(config.Config.App.APIPrefix) { + // OpenFlare legacy /api/* routes (old frontend compatibility) + oflegacy.RegisterRoutes(apiGroup) + // API V1 apiV1Router := apiGroup.Group("/v1") { diff --git a/docs/changelog/index.md b/docs/changelog/index.md index 3cb2ae1f..224a5ef7 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -16,6 +16,12 @@ sidebar: false ## [Unreleased] +### 新增 + +- 将 `openflare-server` 后端业务域迁移至 `Wavelet/internal/apps/openflare/`,通过 `/api/*` legacy 兼容层保持旧前端 API 路径不变。 +- 新增 OpenFlare 业务表 goose 迁移(`of_options`、`of_origins`、`of_proxy_routes`、`of_nodes`、`of_waf_*`、`of_tls_*`、`of_config_versions`、`of_pages_*`、`of_apply_logs` 等)。 +- 重叠职能复用 Wavelet 内置用户/OAuth/Cap/认证源能力;新增 `integration` 包覆盖认证、核心链路、安全、Agent 协议集成测试。 + ## [v2.3.4] - 2026-06-17 ### 变更 diff --git a/docs/plan/handover-openflare-backend-migration.md b/docs/plan/handover-openflare-backend-migration.md new file mode 100644 index 00000000..d409bf1f --- /dev/null +++ b/docs/plan/handover-openflare-backend-migration.md @@ -0,0 +1,56 @@ +# OpenFlare 后端迁移 — 任务拆分与委派 + +> **状态**:进行中 +> **基建**:Batch 0 已完成(`Wavelet/internal/apps/openflare/compat/` + `legacy/register*.go` + `router.go` 挂载) + +## 任务隔离规则 + +| 规则 | 说明 | +|---|---| +| 文件所有权 | 每个任务 **仅修改** 自己的 `internal/apps/openflare/