From e3bfd9ca6d2d3c9d26fa5cad1d21162138f5e05d Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 18 Jun 2026 16:51:06 +0800 Subject: [PATCH] migrate --- .../apps/openflare/agent/access_log_region.go | 95 +++ .../internal/apps/openflare/agent/config.go | 9 +- .../apps/openflare/agent/config_test.go | 46 ++ Wavelet/internal/apps/openflare/agent/errs.go | 1 - .../internal/apps/openflare/agent/helpers.go | 53 +- .../apps/openflare/agent/helpers_test.go | 131 +++ .../internal/apps/openflare/agent/logics.go | 25 +- .../apps/openflare/agent/observability.go | 71 +- .../internal/apps/openflare/agent/routers.go | 36 +- .../apps/openflare/agent/waf_ip_group.go | 200 +++++ .../apps/openflare/agent/waf_ip_group_test.go | 158 ++++ Wavelet/internal/apps/openflare/auth/errs.go | 1 + .../apps/openflare/auth/oauth_shortcuts.go | 353 ++++++++ .../internal/apps/openflare/flared/logics.go | 3 + .../apps/openflare/flared/observability.go | 50 ++ .../openflare/flared/observability_test.go | 71 ++ .../internal/apps/openflare/geoip/lookup.go | 5 + .../integration/agent_protocol_test.go | 5 + .../openflare/integration/auth_option_test.go | 15 + .../apps/openflare/legacy/auth_shortcuts.go | 52 ++ .../apps/openflare/legacy/middleware.go | 5 - .../apps/openflare/legacy/register.go | 6 +- .../apps/openflare/legacy/register_auth.go | 12 +- .../internal/apps/openflare/option/routers.go | 4 +- ...00cea377a6ceffd4c72c76c77910686be8dc4e.zip | Bin 0 -> 142 bytes Wavelet/internal/apps/openflare/pages/errs.go | 35 +- .../internal/apps/openflare/pages/logics.go | 72 ++ .../apps/openflare/pages/logics_test.go | 84 ++ .../internal/apps/openflare/relay/logics.go | 34 +- .../apps/openflare/relay/logics_test.go | 174 ++++ .../apps/openflare/relay/observability.go | 69 ++ ...dd_of_node_access_logs_composite_index.sql | 5 + .../202606190012_create_of_node_obs_frpc.sql | 15 + ...dd_of_node_access_logs_composite_index.sql | 5 + .../202606190012_create_of_node_obs_frpc.sql | 15 + .../internal/model/openflare_access_log.go | 757 ++++++++++++++++++ .../model/openflare_access_log_test.go | 144 ++++ .../internal/model/openflare_observability.go | 96 +-- Wavelet/internal/model/openflare_waf.go | 16 + docs/changelog/index.md | 7 + 40 files changed, 2819 insertions(+), 116 deletions(-) create mode 100644 Wavelet/internal/apps/openflare/agent/access_log_region.go create mode 100644 Wavelet/internal/apps/openflare/agent/config_test.go create mode 100644 Wavelet/internal/apps/openflare/agent/helpers_test.go create mode 100644 Wavelet/internal/apps/openflare/agent/waf_ip_group.go create mode 100644 Wavelet/internal/apps/openflare/agent/waf_ip_group_test.go create mode 100644 Wavelet/internal/apps/openflare/auth/oauth_shortcuts.go create mode 100644 Wavelet/internal/apps/openflare/flared/observability.go create mode 100644 Wavelet/internal/apps/openflare/flared/observability_test.go create mode 100644 Wavelet/internal/apps/openflare/legacy/auth_shortcuts.go create mode 100644 Wavelet/internal/apps/openflare/pages/data/pages/artifacts/published-site/1d0b1001941b126350515d2eaa00cea377a6ceffd4c72c76c77910686be8dc4e.zip create mode 100644 Wavelet/internal/apps/openflare/relay/logics_test.go create mode 100644 Wavelet/internal/apps/openflare/relay/observability.go create mode 100644 Wavelet/internal/db/migrator/goose/postgres/202606190011_add_of_node_access_logs_composite_index.sql create mode 100644 Wavelet/internal/db/migrator/goose/postgres/202606190012_create_of_node_obs_frpc.sql create mode 100644 Wavelet/internal/db/migrator/goose/sqlite/202606190011_add_of_node_access_logs_composite_index.sql create mode 100644 Wavelet/internal/db/migrator/goose/sqlite/202606190012_create_of_node_obs_frpc.sql create mode 100644 Wavelet/internal/model/openflare_access_log.go create mode 100644 Wavelet/internal/model/openflare_access_log_test.go diff --git a/Wavelet/internal/apps/openflare/agent/access_log_region.go b/Wavelet/internal/apps/openflare/agent/access_log_region.go new file mode 100644 index 00000000..f3823535 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/access_log_region.go @@ -0,0 +1,95 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "log/slog" + "net" + "strings" + + pkggeoip "github.com/rain-kl/openflare/pkg/geoip" +) + +var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) { + return pkggeoip.NewMaxMindGeoIPService() +} + +type accessLogRegionResolver struct { + provider pkggeoip.GeoIPService + cache map[string]string +} + +func newAccessLogRegionResolver() (*accessLogRegionResolver, error) { + provider, err := accessLogGeoProviderFactory() + if err != nil { + return nil, err + } + return &accessLogRegionResolver{ + provider: provider, + cache: make(map[string]string), + }, nil +} + +func (r *accessLogRegionResolver) Close() { + if r == nil || r.provider == nil { + return + } + if err := r.provider.Close(); err != nil { + slog.Warn("close access log geo provider failed", "error", err) + } +} + +func (r *accessLogRegionResolver) Resolve(rawIP string) string { + if r == nil || r.provider == nil { + return "" + } + normalizedIP := normalizeAccessLogIP(rawIP) + if normalizedIP == "" { + return "" + } + if cached, ok := r.cache[normalizedIP]; ok { + return cached + } + + info, err := r.provider.GetGeoInfo(net.ParseIP(normalizedIP)) + if err != nil || info == nil { + r.cache[normalizedIP] = "" + return "" + } + + region := strings.TrimSpace(info.Name) + if region == "" { + region = strings.TrimSpace(info.ISOCode) + } + r.cache[normalizedIP] = region + return region +} + +func normalizeAccessLogIP(raw string) string { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "" + } + + if ip := net.ParseIP(trimmed); ip != nil { + return ip.String() + } + + trimmed = strings.TrimPrefix(trimmed, "[") + trimmed = strings.TrimSuffix(trimmed, "]") + if ip := net.ParseIP(trimmed); ip != nil { + return ip.String() + } + + host, _, err := net.SplitHostPort(strings.TrimSpace(raw)) + if err != nil { + return "" + } + host = strings.TrimPrefix(host, "[") + host = strings.TrimSuffix(host, "]") + if ip := net.ParseIP(host); ip != nil { + return ip.String() + } + return "" +} diff --git a/Wavelet/internal/apps/openflare/agent/config.go b/Wavelet/internal/apps/openflare/agent/config.go index e8470356..122697f9 100644 --- a/Wavelet/internal/apps/openflare/agent/config.go +++ b/Wavelet/internal/apps/openflare/agent/config.go @@ -11,6 +11,7 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/db" + openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty" "gorm.io/gorm" ) @@ -89,8 +90,12 @@ func sourceSupportFiles(files []SupportFile) []SupportFile { } func isRuntimeGeneratedSupportFile(path string) bool { - path = strings.TrimSpace(path) - return strings.HasPrefix(path, "runtime/") + switch strings.TrimSpace(path) { + case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName: + return true + default: + return false + } } func isActiveConfigNotFound(err error) bool { diff --git a/Wavelet/internal/apps/openflare/agent/config_test.go b/Wavelet/internal/apps/openflare/agent/config_test.go new file mode 100644 index 00000000..2a4f2c8e --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/config_test.go @@ -0,0 +1,46 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "testing" + + openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty" +) + +func TestIsRuntimeGeneratedSupportFile(t *testing.T) { + tests := []struct { + path string + want bool + }{ + {path: "pow_config.json", want: true}, + {path: "waf_config.json", want: true}, + {path: openrestyrender.SourceConfigFileName, want: true}, + {path: "runtime/custom.json", want: false}, + {path: "certs/example.pem", want: false}, + } + for _, tc := range tests { + if got := isRuntimeGeneratedSupportFile(tc.path); got != tc.want { + t.Fatalf("isRuntimeGeneratedSupportFile(%q) = %v, want %v", tc.path, got, tc.want) + } + } +} + +func TestSourceSupportFilesFiltersRuntimeGeneratedFiles(t *testing.T) { + files := []SupportFile{ + {Path: "certs/example.pem", Content: "pem"}, + {Path: "pow_config.json", Content: "{}"}, + {Path: "waf_config.json", Content: "{}"}, + {Path: openrestyrender.SourceConfigFileName, Content: "{}"}, + {Path: "routes/extra.json", Content: "{}"}, + } + + filtered := sourceSupportFiles(files) + if len(filtered) != 2 { + t.Fatalf("expected 2 support files, got %d: %+v", len(filtered), filtered) + } + if filtered[0].Path != "certs/example.pem" || filtered[1].Path != "routes/extra.json" { + t.Fatalf("unexpected filtered files: %+v", filtered) + } +} diff --git a/Wavelet/internal/apps/openflare/agent/errs.go b/Wavelet/internal/apps/openflare/agent/errs.go index efb2a961..ed553aab 100644 --- a/Wavelet/internal/apps/openflare/agent/errs.go +++ b/Wavelet/internal/apps/openflare/agent/errs.go @@ -17,5 +17,4 @@ const ( 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 index 54dd0138..85b15875 100644 --- a/Wavelet/internal/apps/openflare/agent/helpers.go +++ b/Wavelet/internal/apps/openflare/agent/helpers.go @@ -11,6 +11,7 @@ import ( "strings" "time" + ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip" "github.com/Rain-kl/Wavelet/internal/model" ) @@ -97,6 +98,41 @@ func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveNa now := time.Now() node.LastSeenAt = &now node.LastError = truncateForDatabase(payload.LastError, 16000) + if !node.GeoManualOverride { + applyGeoInfoFromIP(node, node.IP) + } +} + +func applyGeoInfoFromIP(node *model.OpenFlareNode, rawIP string) { + if node == nil { + return + } + node.GeoName = "" + node.GeoLatitude = nil + node.GeoLongitude = nil + ip := net.ParseIP(strings.TrimSpace(rawIP)) + if ip == nil { + return + } + info, err := ofgeoip.GeoInfoFromIP(ip) + if err != nil || info == nil { + return + } + if strings.TrimSpace(info.Name) != "" { + node.GeoName = strings.TrimSpace(info.Name) + } + if info.Latitude != nil && info.Longitude != nil { + node.GeoLatitude = cloneCoordinate(info.Latitude) + node.GeoLongitude = cloneCoordinate(info.Longitude) + } +} + +func cloneCoordinate(value *float64) *float64 { + if value == nil { + return nil + } + cloned := *value + return &cloned } func truncateForDatabase(value string, max int) string { @@ -199,6 +235,7 @@ func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenF } appendIfChanged("name", previous.Name, current.Name) appendIfChanged("ip", previous.IP, current.IP) + appendIfChanged("geo_name", previous.GeoName, current.GeoName) appendIfChanged("version", previous.Version, current.Version) appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion) appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus) @@ -210,12 +247,25 @@ func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenF appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel) appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag) appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested) + if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) { + changes["geo_latitude"] = current.GeoLatitude + } + if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) { + changes["geo_longitude"] = current.GeoLongitude + } if !lastSeenAtEqual(previous.LastSeenAt, current.LastSeenAt) { changes["last_seen_at"] = current.LastSeenAt } return changes } +func coordinatesEqual(before *float64, after *float64) bool { + if before == nil || after == nil { + return before == after + } + return *before == *after +} + func lastSeenAtEqual(before *time.Time, after *time.Time) bool { if before == nil || after == nil { return before == after @@ -241,7 +291,8 @@ func isUniqueConstraintError(err error) bool { return strings.Contains(strings.ToLower(err.Error()), "unique") } -func refreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) { +// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations. +func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) { if node == nil { return } diff --git a/Wavelet/internal/apps/openflare/agent/helpers_test.go b/Wavelet/internal/apps/openflare/agent/helpers_test.go new file mode 100644 index 00000000..a7e9ec23 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/helpers_test.go @@ -0,0 +1,131 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "net" + "testing" + + "github.com/Rain-kl/Wavelet/internal/model" + pkggeoip "github.com/rain-kl/openflare/pkg/geoip" +) + +type fakeGeoIPProvider struct { + info *pkggeoip.GeoInfo +} + +func (f *fakeGeoIPProvider) Name() string { return "fake-geoip" } + +func (f *fakeGeoIPProvider) GetGeoInfo(ip net.IP) (*pkggeoip.GeoInfo, error) { + return f.info, nil +} + +func (f *fakeGeoIPProvider) UpdateDatabase() error { return nil } + +func (f *fakeGeoIPProvider) Close() error { return nil } + +func withFakeGeoIPProvider(t *testing.T, info *pkggeoip.GeoInfo) { + t.Helper() + previous := pkggeoip.CurrentProvider + pkggeoip.CurrentProvider = &fakeGeoIPProvider{info: info} + t.Cleanup(func() { + pkggeoip.CurrentProvider = previous + }) +} + +func geoipFloat(value float64) *float64 { + return &value +} + +func TestApplyGeoInfoFromIP(t *testing.T) { + latitude := 31.2304 + longitude := 121.4737 + withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{ + Name: "Shanghai", + Latitude: geoipFloat(latitude), + Longitude: geoipFloat(longitude), + }) + + node := &model.OpenFlareNode{IP: "203.0.113.10"} + applyGeoInfoFromIP(node, node.IP) + + if node.GeoName != "Shanghai" { + t.Fatalf("expected geo_name Shanghai, got %q", node.GeoName) + } + if node.GeoLatitude == nil || *node.GeoLatitude != latitude { + t.Fatalf("unexpected geo_latitude: %+v", node.GeoLatitude) + } + if node.GeoLongitude == nil || *node.GeoLongitude != longitude { + t.Fatalf("unexpected geo_longitude: %+v", node.GeoLongitude) + } +} + +func TestApplyGeoInfoFromIPSkipsInvalidIP(t *testing.T) { + withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{Name: "Should Not Apply"}) + + node := &model.OpenFlareNode{ + IP: "203.0.113.10", + GeoName: "Existing", + GeoLatitude: geoipFloat(1), + GeoLongitude: geoipFloat(2), + } + applyGeoInfoFromIP(node, "not-an-ip") + + if node.GeoName != "" || node.GeoLatitude != nil || node.GeoLongitude != nil { + t.Fatalf("expected geo fields to be cleared on invalid IP, got %+v", node) + } +} + +func TestApplyNodeRuntimeRespectsGeoManualOverride(t *testing.T) { + withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{ + Name: "Shanghai", + Latitude: geoipFloat(31.2304), + Longitude: geoipFloat(121.4737), + }) + + node := &model.OpenFlareNode{ + GeoManualOverride: true, + GeoName: "Manual", + GeoLatitude: geoipFloat(10), + GeoLongitude: geoipFloat(20), + } + applyNodeRuntime(node, NodePayload{ + IP: "203.0.113.10", + Version: "1.0.0", + }, true) + + if node.GeoName != "Manual" { + t.Fatalf("expected manual geo_name to be preserved, got %q", node.GeoName) + } + if node.GeoLatitude == nil || *node.GeoLatitude != 10 { + t.Fatalf("expected manual geo_latitude to be preserved, got %+v", node.GeoLatitude) + } +} + +func TestCollectHeartbeatChangesTracksGeoFields(t *testing.T) { + before := &model.OpenFlareNode{ + IP: "10.0.0.1", + GeoName: "Old Region", + } + after := &model.OpenFlareNode{ + IP: "203.0.113.10", + GeoName: "New Region", + GeoLatitude: geoipFloat(31.2304), + GeoLongitude: geoipFloat(121.4737), + } + + changes := collectHeartbeatChanges(before, after) + if changes["ip"] != after.IP { + t.Fatalf("expected ip change, got %+v", changes) + } + if changes["geo_name"] != after.GeoName { + t.Fatalf("expected geo_name change, got %+v", changes) + } + if changes["geo_latitude"] != after.GeoLatitude { + t.Fatalf("expected geo_latitude change, got %+v", changes) + } + if changes["geo_longitude"] != after.GeoLongitude { + t.Fatalf("expected geo_longitude change, got %+v", changes) + } +} diff --git a/Wavelet/internal/apps/openflare/agent/logics.go b/Wavelet/internal/apps/openflare/agent/logics.go index fefd0eee..2f20610e 100644 --- a/Wavelet/internal/apps/openflare/agent/logics.go +++ b/Wavelet/internal/apps/openflare/agent/logics.go @@ -28,7 +28,7 @@ func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, if err := model.SaveOpenFlareNode(ctx, authNode); err != nil { return nil, err } - refreshAccessTokenCache(ctx, authNode) + RefreshAccessTokenCache(ctx, authNode) return &RegistrationResponse{ NodeID: authNode.NodeID, AccessToken: authNode.AccessToken, @@ -74,7 +74,7 @@ func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*Registrat } return nil, err } - refreshAccessTokenCache(ctx, record) + RefreshAccessTokenCache(ctx, record) return &RegistrationResponse{ NodeID: record.NodeID, AccessToken: record.AccessToken, @@ -116,24 +116,29 @@ func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload N } } - refreshAccessTokenCache(ctx, authNode) + RefreshAccessTokenCache(ctx, authNode) reportedAt := time.Now() if authNode.LastSeenAt != nil { reportedAt = *authNode.LastSeenAt } - persistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt) + PersistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt) activeConfig, err := getActiveConfigMeta(ctx) if err != nil && !isActiveConfigNotFound(err) { return nil, err } + wafIPGroups, err := ChangedWAFIPGroupsForAgent(ctx, nil, payload.WAFIPGroupChecksums) + if err != nil { + return nil, err + } + return &HeartbeatResponse{ Node: authNode, AgentSettings: buildAgentSettings(authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow), ActiveConfig: activeConfig, - WAFIPGroups: nil, + WAFIPGroups: wafIPGroups, }, nil } @@ -149,9 +154,13 @@ func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) { 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 +// SyncWAFIPGroups returns WAF IP groups whose checksums differ from the agent state. +func SyncWAFIPGroups(ctx context.Context, input WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) { + groups, err := ChangedWAFIPGroupsForAgent(ctx, input.IDs, input.Checksums) + if err != nil { + return nil, err + } + return &WAFIPGroupSyncResult{Groups: groups}, nil } // ReportApplyLog records an agent apply result. diff --git a/Wavelet/internal/apps/openflare/agent/observability.go b/Wavelet/internal/apps/openflare/agent/observability.go index 0ef76935..fd9f3745 100644 --- a/Wavelet/internal/apps/openflare/agent/observability.go +++ b/Wavelet/internal/apps/openflare/agent/observability.go @@ -7,6 +7,7 @@ import ( "context" "encoding/json" "errors" + "log/slog" "strings" "time" @@ -18,12 +19,14 @@ import ( ) const ( - healthEventStatusActive = "active" - healthEventStatusResolved = "resolved" - healthSeverityInfo = "info" - healthSeverityWarning = "warning" - healthSeverityCritical = "critical" - accessLogPathMaxLength = 100 + healthEventStatusActive = "active" + healthEventStatusResolved = "resolved" + healthSeverityInfo = "info" + healthSeverityWarning = "warning" + healthSeverityCritical = "critical" + nodeAccessLogRetentionDays = 90 + nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour + accessLogPathMaxLength = 100 ) // NodeSystemProfile is the agent-reported system profile. @@ -102,7 +105,8 @@ type NodeHealthEvent struct { Metadata map[string]string `json:"metadata"` } -func persistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) { +// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events. +func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) { if strings.TrimSpace(nodeID) == "" { return } @@ -275,15 +279,29 @@ func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *NodeTrafficRep } func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, reportedAt time.Time) error { + if len(logs) == 0 { + return nil + } + resolver, err := newAccessLogRegionResolver() + if err != nil { + slog.Warn("initialize access log geo resolver failed", "node_id", nodeID, "error", err) + } + if resolver != nil { + defer resolver.Close() + } for _, item := range logs { record := &model.OpenFlareAccessLog{ NodeID: nodeID, LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt), RemoteAddr: strings.TrimSpace(item.RemoteAddr), + Region: "", Host: strings.TrimSpace(item.Host), Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength), StatusCode: item.StatusCode, } + if resolver != nil { + record.Region = resolver.Resolve(record.RemoteAddr) + } exists, err := accessLogExists(tx, record) if err != nil { return err @@ -295,16 +313,27 @@ func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, rep return err } } - return nil + _, err = deleteAccessLogsByNodeBefore(tx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow)) + return err } func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error { + return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil) +} + +// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types. +func ReconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error { activeTypes := make(map[string]NodeHealthEvent, len(events)) for _, event := range events { eventType := normalizeHealthEventType(event.EventType) if eventType == "" { continue } + if len(managedEventTypes) > 0 { + if _, ok := managedEventTypes[eventType]; !ok { + continue + } + } event.EventType = eventType event.Severity = normalizeHealthSeverity(event.Severity) if event.TriggeredAtUnix <= 0 { @@ -314,7 +343,21 @@ func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEv } var activeEvents []*model.OpenFlareHealthEvent - if err := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive).Find(&activeEvents).Error; err != nil { + query := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive) + if len(managedEventTypes) > 0 { + scopedTypes := make([]string, 0, len(managedEventTypes)) + for eventType := range managedEventTypes { + eventType = normalizeHealthEventType(eventType) + if eventType != "" { + scopedTypes = append(scopedTypes, eventType) + } + } + if len(scopedTypes) == 0 { + return nil + } + query = query.Where("event_type IN ?", scopedTypes) + } + if err := query.Find(&activeEvents).Error; err != nil { return err } @@ -391,6 +434,11 @@ func requestReportExists(tx *gorm.DB, nodeID string, windowStartedAt, windowEnde return count > 0, nil } +func deleteAccessLogsByNodeBefore(tx *gorm.DB, nodeID string, before time.Time) (int64, error) { + result := tx.Where("node_id = ? AND logged_at < ?", nodeID, before).Delete(&model.OpenFlareAccessLog{}) + return result.RowsAffected, result.Error +} + func accessLogExists(tx *gorm.DB, record *model.OpenFlareAccessLog) (bool, error) { var count int64 if err := tx.Model(&model.OpenFlareAccessLog{}). @@ -438,6 +486,11 @@ func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time { return time.Unix(unixSeconds, 0).UTC() } +// MarshalJSON serializes a value for database JSON columns. +func MarshalJSON(value any) string { + return marshalJSON(value) +} + func marshalJSON(value any) string { if value == nil { return "" diff --git a/Wavelet/internal/apps/openflare/agent/routers.go b/Wavelet/internal/apps/openflare/agent/routers.go index 5e2836d4..135a7d55 100644 --- a/Wavelet/internal/apps/openflare/agent/routers.go +++ b/Wavelet/internal/apps/openflare/agent/routers.go @@ -5,8 +5,10 @@ package agent import ( "net/http" + "strconv" "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/pages" "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" "github.com/gin-gonic/gin" ) @@ -98,7 +100,7 @@ func GetActiveConfigHandler(c *gin.Context) { compat.OK(c, config) } -// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent (stub). +// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent. func SyncWAFIPGroupsHandler(c *gin.Context) { var input WAFIPGroupSyncInput if !compat.BindJSON(c, &input) { @@ -129,13 +131,33 @@ func ReportApplyLogHandler(c *gin.Context) { compat.OK(c, log) } -// DownloadPagesPackageHandler is a stub until Pages agent packaging is migrated. +// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent. func DownloadPagesPackageHandler(c *gin.Context) { - c.JSON(http.StatusNotFound, compat.Envelope{ - Success: false, - Message: errPagesPackageNotFound, - Data: nil, - }) + deploymentID, ok := pagesDeploymentIDParam(c) + if !ok { + return + } + filePath, fileName, err := pages.GetDeploymentPackagePath(c.Request.Context(), deploymentID) + if err != nil { + compat.Fail(c, err.Error()) + return + } + c.Header("Content-Disposition", "attachment; filename="+fileName) + c.File(filePath) +} + +func pagesDeploymentIDParam(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 } // AgentWebSocketHandler upgrades an authenticated agent websocket connection. diff --git a/Wavelet/internal/apps/openflare/agent/waf_ip_group.go b/Wavelet/internal/apps/openflare/agent/waf_ip_group.go new file mode 100644 index 00000000..cc80fa11 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/waf_ip_group.go @@ -0,0 +1,200 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "sort" + "strings" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +type snapshotWAFRuleGroupRef struct { + IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"` + IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"` +} + +type snapshotWAFSection struct { + RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"` +} + +type activeConfigSnapshot struct { + WAF snapshotWAFSection `json:"waf"` +} + +// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state. +func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) { + targetIDs := uniqueUintIDs(ids) + if len(targetIDs) == 0 { + activeIDs, err := activeConfigWAFIPGroupIDs(ctx) + if err != nil { + return nil, err + } + targetIDs = activeIDs + } + if len(targetIDs) == 0 { + return []WAFIPGroup{}, nil + } + groups, err := buildAgentWAFIPGroups(ctx, targetIDs) + if err != nil { + return nil, err + } + changed := make([]WAFIPGroup, 0, len(groups)) + for _, group := range groups { + if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum { + continue + } + changed = append(changed, group) + } + return changed, nil +} + +func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) { + ids = uniqueUintIDs(ids) + if len(ids) == 0 { + return []WAFIPGroup{}, nil + } + sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) + groups, err := model.ListOpenFlareWAFIPGroupsByIDs(ctx, ids) + if err != nil { + return nil, err + } + groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups)) + for _, group := range groups { + groupByID[group.ID] = group + } + result := make([]WAFIPGroup, 0, len(ids)) + for _, id := range ids { + group := groupByID[id] + if group == nil { + continue + } + agentGroup, err := buildAgentWAFIPGroup(group) + if err != nil { + return nil, err + } + result = append(result, agentGroup) + } + return result, nil +} + +func buildAgentWAFIPGroup(group *model.OpenFlareWAFIPGroup) (WAFIPGroup, error) { + if group == nil { + return WAFIPGroup{}, errors.New("IP 组不存在") + } + ips, err := decodeWAFIPGroupStringList(group.IPList) + if err != nil { + return WAFIPGroup{}, err + } + if !group.Enabled { + ips = []string{} + } + agentGroup := WAFIPGroup{ + ID: group.ID, + Name: group.Name, + Type: group.Type, + Enabled: group.Enabled, + IPList: ips, + } + agentGroup.Checksum = checksumAgentWAFIPGroup(agentGroup) + return agentGroup, nil +} + +func checksumAgentWAFIPGroup(group WAFIPGroup) string { + payload := struct { + ID uint `json:"id"` + Enabled bool `json:"enabled"` + IPList []string `json:"ip_list"` + }{ + ID: group.ID, + Enabled: group.Enabled, + IPList: append([]string{}, group.IPList...), + } + sort.Strings(payload.IPList) + data, _ := json.Marshal(payload) + sum := sha256.Sum256(data) + return hex.EncodeToString(sum[:]) +} + +func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) { + version, err := loadActiveConfigVersion(ctx) + if err != nil { + if isActiveConfigNotFound(err) { + return []uint{}, nil + } + return nil, err + } + snapshot, err := parseActiveConfigSnapshot(version.SnapshotJSON) + if err != nil { + return nil, err + } + idSet := make(map[uint]struct{}) + for _, group := range snapshot.WAF.RuleGroups { + for _, id := range group.IPWhitelistGroups { + if id > 0 { + idSet[id] = struct{}{} + } + } + for _, id := range group.IPBlacklistGroups { + if id > 0 { + idSet[id] = struct{}{} + } + } + } + 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] }) + return ids, nil +} + +func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) { + text := strings.TrimSpace(snapshotJSON) + if text == "" { + return &activeConfigSnapshot{}, nil + } + var snapshot activeConfigSnapshot + if err := json.Unmarshal([]byte(text), &snapshot); err != nil { + return nil, err + } + if snapshot.WAF.RuleGroups == nil { + snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{} + } + return &snapshot, nil +} + +func decodeWAFIPGroupStringList(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 uniqueUintIDs(ids []uint) []uint { + normalized := make([]uint, 0, len(ids)) + seen := make(map[uint]struct{}, len(ids)) + for _, id := range ids { + if id == 0 { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + normalized = append(normalized, id) + } + return normalized +} diff --git a/Wavelet/internal/apps/openflare/agent/waf_ip_group_test.go b/Wavelet/internal/apps/openflare/agent/waf_ip_group_test.go new file mode 100644 index 00000000..b2b02176 --- /dev/null +++ b/Wavelet/internal/apps/openflare/agent/waf_ip_group_test.go @@ -0,0 +1,158 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "encoding/json" + "strconv" + "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 setupWAFIPGroupTestDB(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.OpenFlareWAFIPGroup{}, + &configVersionRecord{}, + )) + + db.SetDB(sqliteDB) + return func() { + db.SetDB(nil) + } +} + +func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) { + t.Helper() + + snapshot := map[string]any{ + "routes": []any{}, + "waf": map[string]any{ + "rule_groups": []map[string]any{ + { + "id": 1, + "name": "agent refs", + "enabled": true, + "ip_blacklist_group_ids": []uint{ipGroupID}, + }, + }, + "bindings": []any{}, + }, + } + snapshotJSON, err := json.Marshal(snapshot) + require.NoError(t, err) + + require.NoError(t, db.DB(ctx).Create(&configVersionRecord{ + Version: "20260618-001", + SnapshotJSON: string(snapshotJSON), + Checksum: "test-checksum", + IsActive: true, + }).Error) +} + +func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) { + cleanup := setupWAFIPGroupTestDB(t) + defer cleanup() + + ctx := context.Background() + ipGroup := &model.OpenFlareWAFIPGroup{ + Name: "agent runtime group", + Type: "manual", + Enabled: true, + IPList: `["203.0.113.44"]`, + } + require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) + + groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) + require.NoError(t, err) + require.Len(t, groups, 1) + assert.Equal(t, ipGroup.ID, groups[0].ID) + assert.Equal(t, "203.0.113.44", groups[0].IPList[0]) + assert.NotEmpty(t, groups[0].Checksum) + + groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10) + same, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum}) + require.NoError(t, err) + assert.Empty(t, same) + + ipGroup.IPList = `["203.0.113.45"]` + require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup)) + + delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum}) + require.NoError(t, err) + require.Len(t, delta, 1) + assert.Equal(t, ipGroup.ID, delta[0].ID) + assert.Equal(t, "203.0.113.45", delta[0].IPList[0]) + assert.NotEqual(t, groups[0].Checksum, delta[0].Checksum) +} + +func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) { + cleanup := setupWAFIPGroupTestDB(t) + defer cleanup() + + ctx := context.Background() + ipGroup := &model.OpenFlareWAFIPGroup{ + Name: "sync group", + Type: "manual", + Enabled: true, + IPList: `["198.51.100.10"]`, + } + require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) + + result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{ + IDs: []uint{ipGroup.ID}, + Checksums: map[string]string{}, + }) + require.NoError(t, err) + require.Len(t, result.Groups, 1) + assert.Equal(t, ipGroup.ID, result.Groups[0].ID) + assert.Equal(t, "198.51.100.10", result.Groups[0].IPList[0]) + + result, err = SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{ + IDs: []uint{ipGroup.ID}, + Checksums: map[string]string{ + strconv.FormatUint(uint64(ipGroup.ID), 10): result.Groups[0].Checksum, + }, + }) + require.NoError(t, err) + assert.Empty(t, result.Groups) +} + +func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) { + cleanup := setupWAFIPGroupTestDB(t) + defer cleanup() + + ctx := context.Background() + ipGroup := &model.OpenFlareWAFIPGroup{ + Name: "disabled group", + Type: "manual", + Enabled: true, + IPList: `["203.0.113.10"]`, + } + require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + ipGroup.Enabled = false + require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup)) + seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) + + groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) + require.NoError(t, err) + require.Len(t, groups, 1) + assert.False(t, groups[0].Enabled) + assert.Empty(t, groups[0].IPList) + assert.NotEmpty(t, groups[0].Checksum) +} diff --git a/Wavelet/internal/apps/openflare/auth/errs.go b/Wavelet/internal/apps/openflare/auth/errs.go index 1a8ace6f..5f431144 100644 --- a/Wavelet/internal/apps/openflare/auth/errs.go +++ b/Wavelet/internal/apps/openflare/auth/errs.go @@ -5,6 +5,7 @@ package auth const ( errInvalidParams = "无效的参数" + errUnauthorized = "无权进行此操作,未登录或 token 无效" errPasswordLoginDisabled = "管理员关闭了密码登录" errUsernameOrPasswordWrong = "用户名或密码错误" errBannedAccount = "用户已被封禁" diff --git a/Wavelet/internal/apps/openflare/auth/oauth_shortcuts.go b/Wavelet/internal/apps/openflare/auth/oauth_shortcuts.go new file mode 100644 index 00000000..b92e29ed --- /dev/null +++ b/Wavelet/internal/apps/openflare/auth/oauth_shortcuts.go @@ -0,0 +1,353 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "strings" + "time" + + "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/pkg/logger" + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +const ( + errGitHubOAuthDisabled = "管理员未开启通过 GitHub 登录以及注册" + errWeChatOAuthDisabled = "管理员未开启通过微信登录以及注册" + errRegistrationClosed = "管理员关闭了新用户注册" + errGitHubAlreadyBound = "该 GitHub 账户已被绑定" + errWeChatAlreadyBound = "该微信账号已被绑定" +) + +type githubOAuthResponse struct { + AccessToken string `json:"access_token"` +} + +type githubUser struct { + Login string `json:"login"` + Name string `json:"name"` + Email string `json:"email"` +} + +type wechatLoginResponse struct { + Success bool `json:"success"` + Message string `json:"message"` + Data string `json:"data"` +} + +// GitHubOAuth handles the legacy GET /oauth/github shortcut. +func GitHubOAuth(ctx context.Context, c *gin.Context, code string) (LegacyUser, error) { + if current := currentUserFromLegacyToken(ctx, c); current != nil { + if err := GitHubBind(ctx, c, current, code); err != nil { + return LegacyUser{}, err + } + return LegacyUser{}, nil + } + + if !model.GitHubOAuthEnabled { + return LegacyUser{}, errors.New(errGitHubOAuthDisabled) + } + + githubUser, err := getGitHubUserInfoByCode(code) + if err != nil { + return LegacyUser{}, err + } + + user, err := findUserByShortcutBinding(ctx, githubUser.Login, "github", "GitHub") + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return LegacyUser{}, errors.New(errRegistrationClosed) + } + return LegacyUser{}, err + } + if !user.IsActive { + return LegacyUser{}, errors.New(errBannedAccount) + } + return finishLegacyLogin(ctx, c, user) +} + +// GitHubBind binds a GitHub account to the current user. +func GitHubBind(ctx context.Context, c *gin.Context, current *model.User, code string) error { + if current == nil { + return errors.New(errUnauthorized) + } + if !model.GitHubOAuthEnabled { + return errors.New(errGitHubOAuthDisabled) + } + + githubUser, err := getGitHubUserInfoByCode(code) + if err != nil { + return err + } + if _, err := findUserByShortcutBinding(ctx, githubUser.Login, "github", "GitHub"); err == nil { + return errors.New(errGitHubAlreadyBound) + } else if !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + + return bindShortcutExternalAccount(ctx, current.ID, githubUser.Login, githubUser.Login, githubUser.Email, "github", "GitHub") +} + +// WeChatOAuth handles the legacy GET /oauth/wechat shortcut. +func WeChatOAuth(ctx context.Context, c *gin.Context, code string) (LegacyUser, error) { + if !model.WeChatAuthEnabled { + return LegacyUser{}, errors.New(errWeChatOAuthDisabled) + } + + wechatID, err := getWeChatIDByCode(code) + if err != nil { + return LegacyUser{}, err + } + + user, err := findUserByShortcutBinding(ctx, wechatID, "wechat", "WeChat") + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return LegacyUser{}, errors.New(errRegistrationClosed) + } + return LegacyUser{}, err + } + if !user.IsActive { + return LegacyUser{}, errors.New(errBannedAccount) + } + return finishLegacyLogin(ctx, c, user) +} + +// WeChatBind binds a WeChat account to the current user. +func WeChatBind(ctx context.Context, userID uint64, code string) error { + if userID == 0 { + return errors.New(errUnauthorized) + } + if !model.WeChatAuthEnabled { + return errors.New(errWeChatOAuthDisabled) + } + + wechatID, err := getWeChatIDByCode(code) + if err != nil { + return err + } + if _, err := findUserByShortcutBinding(ctx, wechatID, "wechat", "WeChat"); err == nil { + return errors.New(errWeChatAlreadyBound) + } else if !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + + return bindShortcutExternalAccount(ctx, userID, wechatID, wechatID, "", "wechat", "WeChat") +} + +// EmailBind binds a verified email address to the current user. +func EmailBind(ctx context.Context, userID uint64, email, code string) error { + email = strings.TrimSpace(email) + code = strings.TrimSpace(code) + if userID == 0 { + return errors.New(errUnauthorized) + } + if email == "" || code == "" { + return errors.New(errInvalidParams) + } + if !verifyEmailCode(ctx, email, "register", code) { + return errors.New(errEmailCodeInvalid) + } + + var user model.User + if err := db.DB(ctx).Where("id = ?", userID).First(&user).Error; err != nil { + return errors.New(errUserNotFound) + } + user.Email = email + return db.DB(ctx).Model(&user).Update("email", email).Error +} + +func finishLegacyLogin(ctx context.Context, c *gin.Context, user *model.User) (LegacyUser, error) { + if user == nil { + return LegacyUser{}, errors.New(errUserNotFound) + } + 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 shortcut 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 +} + +func getGitHubUserInfoByCode(code string) (*githubUser, error) { + code = strings.TrimSpace(code) + if code == "" { + return nil, errors.New(errInvalidParams) + } + values := map[string]string{ + "client_id": model.GitHubClientId, + "client_secret": model.GitHubClientSecret, + "code": code, + } + jsonData, err := json.Marshal(values) + if err != nil { + return nil, err + } + + client := http.Client{Timeout: 5 * time.Second} + req, err := http.NewRequest(http.MethodPost, "https://github.com/login/oauth/access_token", bytes.NewBuffer(jsonData)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + + res, err := client.Do(req) + if err != nil { + slog.Error("github oauth access token request failed", "error", err) + return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") + } + defer res.Body.Close() + + var oauthResponse githubOAuthResponse + if err := json.NewDecoder(res.Body).Decode(&oauthResponse); err != nil { + return nil, err + } + if strings.TrimSpace(oauthResponse.AccessToken) == "" { + return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") + } + + req, err = http.NewRequest(http.MethodGet, "https://api.github.com/user", nil) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", oauthResponse.AccessToken)) + + res2, err := client.Do(req) + if err != nil { + slog.Error("github user info request failed", "error", err) + return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") + } + defer res2.Body.Close() + + var ghUser githubUser + if err := json.NewDecoder(res2.Body).Decode(&ghUser); err != nil { + return nil, err + } + if strings.TrimSpace(ghUser.Login) == "" { + return nil, errors.New("返回值非法,用户字段为空,请稍后重试!") + } + return &ghUser, nil +} + +func getWeChatIDByCode(code string) (string, error) { + code = strings.TrimSpace(code) + if code == "" { + return "", errors.New(errInvalidParams) + } + serverAddress := strings.TrimRight(strings.TrimSpace(model.WeChatServerAddress), "/") + if serverAddress == "" { + return "", errors.New(errWeChatOAuthDisabled) + } + + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("%s/api/wechat/user?code=%s", serverAddress, code), nil) + if err != nil { + return "", err + } + req.Header.Set("Authorization", model.WeChatServerToken) + + client := http.Client{Timeout: 5 * time.Second} + httpResponse, err := client.Do(req) + if err != nil { + return "", err + } + defer func(body io.ReadCloser) { + if closeErr := body.Close(); closeErr != nil { + slog.Error("failed to close wechat response body", "error", closeErr) + } + }(httpResponse.Body) + + var res wechatLoginResponse + if err := json.NewDecoder(httpResponse.Body).Decode(&res); err != nil { + return "", err + } + if !res.Success { + if strings.TrimSpace(res.Message) == "" { + return "", errors.New(errInvalidParams) + } + return "", errors.New(res.Message) + } + if strings.TrimSpace(res.Data) == "" { + return "", errors.New(errEmailCodeInvalid) + } + return strings.TrimSpace(res.Data), nil +} + +func findUserByShortcutBinding(ctx context.Context, externalID string, sourceNames ...string) (*model.User, error) { + externalID = strings.TrimSpace(externalID) + if externalID == "" { + return nil, gorm.ErrRecordNotFound + } + + query := db.DB(ctx). + Table("w_external_accounts AS ea"). + Select("u.*"). + Joins("JOIN w_users u ON u.id = ea.user_id"). + Where("ea.external_id = ?", externalID) + if len(sourceNames) > 0 { + lowered := make([]string, 0, len(sourceNames)) + for _, name := range sourceNames { + trimmed := strings.ToLower(strings.TrimSpace(name)) + if trimmed != "" { + lowered = append(lowered, trimmed) + } + } + if len(lowered) > 0 { + query = query. + Joins("JOIN w_auth_sources s ON s.id = ea.auth_source_id"). + Where("LOWER(s.name) IN ?", lowered) + } + } + + var user model.User + if err := query.First(&user).Error; err != nil { + return nil, err + } + return &user, nil +} + +func bindShortcutExternalAccount(ctx context.Context, userID uint64, externalID, externalUsername, email string, sourceNames ...string) error { + source, err := resolveShortcutAuthSource(ctx, sourceNames...) + if err != nil { + return err + } + return model.BindExternalAccount(ctx, &model.ExternalAccount{ + AuthSourceID: source.ID, + UserID: userID, + ExternalID: strings.TrimSpace(externalID), + ExternalUsername: strings.TrimSpace(externalUsername), + Email: strings.TrimSpace(email), + }) +} + +func resolveShortcutAuthSource(ctx context.Context, sourceNames ...string) (*model.AuthSource, error) { + for _, name := range sourceNames { + source, err := model.GetAuthSourceByName(ctx, name) + if err == nil { + return source, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + } + return nil, errors.New("认证源不存在") +} diff --git a/Wavelet/internal/apps/openflare/flared/logics.go b/Wavelet/internal/apps/openflare/flared/logics.go index 1f4f9272..83b328c2 100644 --- a/Wavelet/internal/apps/openflare/flared/logics.go +++ b/Wavelet/internal/apps/openflare/flared/logics.go @@ -10,6 +10,7 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" "github.com/Rain-kl/Wavelet/internal/apps/openflare/relay" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" @@ -143,6 +144,8 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil { return nil, fmt.Errorf("update flared heartbeat: %w", err) } + agent.RefreshAccessTokenCache(ctx, node) + persistFlaredObservability(ctx, node.NodeID, payload, now) activeConfig, err := getActiveConfigMeta(ctx) if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { diff --git a/Wavelet/internal/apps/openflare/flared/observability.go b/Wavelet/internal/apps/openflare/flared/observability.go new file mode 100644 index 00000000..499dd214 --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/observability.go @@ -0,0 +1,50 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +import ( + "context" + "fmt" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" + "github.com/Rain-kl/Wavelet/internal/db" + "go.uber.org/zap" +) + +const flaredRuntimeUnhealthyEventType = "flared_runtime_unhealthy" + +func persistFlaredObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) { + connected := make([]string, 0, len(payload.ConnectedRelays)) + for _, relay := range payload.ConnectedRelays { + connected = append(connected, fmt.Sprintf("%s:%s", relay.RelayNodeID, relay.Status)) + } + managedTypes := map[string]struct{}{ + flaredRuntimeUnhealthyEventType: {}, + } + var events []agent.NodeHealthEvent + if payload.TunnelStatus == "unhealthy" { + events = append(events, agent.NodeHealthEvent{ + EventType: flaredRuntimeUnhealthyEventType, + Severity: "critical", + Message: "openflared runtime is not healthy", + TriggeredAtUnix: reportedAt.Unix(), + Metadata: map[string]string{ + "tunnel_status": payload.TunnelStatus, + "client_version": payload.ClientVersion, + "current_version": payload.CurrentVersion, + "current_checksum": payload.CurrentChecksum, + "connected_relays": strings.Join(connected, ","), + }, + }) + } + conn := db.DB(ctx) + if conn == nil { + return + } + if err := agent.ReconcileScopedNodeHealthEvents(conn, nodeID, events, reportedAt, managedTypes); err != nil { + zap.L().Error("persist flared health events failed", zap.String("node_id", nodeID), zap.Error(err)) + } +} diff --git a/Wavelet/internal/apps/openflare/flared/observability_test.go b/Wavelet/internal/apps/openflare/flared/observability_test.go new file mode 100644 index 00000000..ca477a33 --- /dev/null +++ b/Wavelet/internal/apps/openflare/flared/observability_test.go @@ -0,0 +1,71 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package flared + +import ( + "context" + "testing" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" + "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 setupFlaredObservabilityTestDB(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.OpenFlareHealthEvent{}, + )) + + db.SetDB(sqliteDB) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + + return func() { + db.SetDB(nil) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + } +} + +func TestHeartbeatFlaredEmitsHealthEventOnUnhealthy(t *testing.T) { + cleanup := setupFlaredObservabilityTestDB(t) + defer cleanup() + + ctx := context.Background() + node := &model.OpenFlareNode{ + NodeID: "node-flared-unhealthy", + Name: "flared-unhealthy", + AccessToken: "tunnel-token-unhealthy", + Status: "pending", + NodeType: "tunnel_client", + } + require.NoError(t, db.DB(ctx).Create(node).Error) + + _, err := Heartbeat(ctx, node, HeartbeatPayload{ + ClientVersion: "v0.2.0", + FrpVersion: "0.61.0", + TunnelStatus: "unhealthy", + CurrentVersion: "v1", + CurrentChecksum: "checksum-1", + }) + require.NoError(t, err) + + events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 20) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, flaredRuntimeUnhealthyEventType, events[0].EventType) + assert.Equal(t, "active", events[0].Status) +} diff --git a/Wavelet/internal/apps/openflare/geoip/lookup.go b/Wavelet/internal/apps/openflare/geoip/lookup.go index ab253f8c..dcd892bb 100644 --- a/Wavelet/internal/apps/openflare/geoip/lookup.go +++ b/Wavelet/internal/apps/openflare/geoip/lookup.go @@ -34,6 +34,11 @@ func IsValidProvider(provider string) bool { return pkggeoip.IsValidProvider(provider) } +// GeoInfoFromIP resolves geographic information using the configured default provider. +func GeoInfoFromIP(ip net.IP) (*pkggeoip.GeoInfo, error) { + return pkggeoip.GetGeoInfo(ip) +} + // Lookup resolves geographic information for rawIP using the given provider. func Lookup(provider, rawIP string) (*LookupView, error) { trimmedProvider := strings.TrimSpace(provider) diff --git a/Wavelet/internal/apps/openflare/integration/agent_protocol_test.go b/Wavelet/internal/apps/openflare/integration/agent_protocol_test.go index d5a4b783..c17bf2e8 100644 --- a/Wavelet/internal/apps/openflare/integration/agent_protocol_test.go +++ b/Wavelet/internal/apps/openflare/integration/agent_protocol_test.go @@ -46,6 +46,11 @@ func setupProtocolTestEnv(t *testing.T) (*gin.Engine, func()) { &model.OpenFlareNode{}, &model.OpenFlareOption{}, &model.OpenFlareApplyLog{}, + &model.OpenFlareNodeSystemProfile{}, + &model.OpenFlareMetricSnapshot{}, + &model.OpenFlareHealthEvent{}, + &model.OpenFlareNodeObservationFrps{}, + &model.OpenFlareNodeObservationFrpc{}, &configVersionRecord{}, )) diff --git a/Wavelet/internal/apps/openflare/integration/auth_option_test.go b/Wavelet/internal/apps/openflare/integration/auth_option_test.go index 119b0cec..60f9202a 100644 --- a/Wavelet/internal/apps/openflare/integration/auth_option_test.go +++ b/Wavelet/internal/apps/openflare/integration/auth_option_test.go @@ -184,6 +184,21 @@ func TestGETOptionRequiresRootAuth(t *testing.T) { }) } +func TestGETNodesWithOpenFlareToken(t *testing.T) { + dbConn, r := setupAuthOptionIntegration(t) + require.NoError(t, dbConn.AutoMigrate(&model.OpenFlareNode{})) + seedUser(t, dbConn, "admin", "password123", true) + rootToken := loginAndGetToken(t, r, "admin", "password123") + + w := performJSONRequest(t, r, http.MethodGet, "/api/nodes/", 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) diff --git a/Wavelet/internal/apps/openflare/legacy/auth_shortcuts.go b/Wavelet/internal/apps/openflare/legacy/auth_shortcuts.go new file mode 100644 index 00000000..e9bd7a82 --- /dev/null +++ b/Wavelet/internal/apps/openflare/legacy/auth_shortcuts.go @@ -0,0 +1,52 @@ +// 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" +) + +// GitHubOAuth handles GET /oauth/github for legacy GitHub login or bind. +func GitHubOAuth(c *gin.Context) { + user, err := ofauth.GitHubOAuth(c.Request.Context(), c, c.Query("code")) + if err != nil { + compat.Fail(c, err.Error()) + return + } + if user.ID == 0 { + compat.OKMessage(c, "bind") + return + } + compat.OK(c, user) +} + +// WeChatOAuth handles GET /oauth/wechat for legacy WeChat login. +func WeChatOAuth(c *gin.Context) { + user, err := ofauth.WeChatOAuth(c.Request.Context(), c, c.Query("code")) + if err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OK(c, user) +} + +// WeChatBind handles GET /oauth/wechat/bind for legacy WeChat account binding. +func WeChatBind(c *gin.Context) { + if err := ofauth.WeChatBind(c.Request.Context(), callerUserID(c), c.Query("code")); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} + +// EmailBind handles GET /oauth/email/bind for legacy email binding. +func EmailBind(c *gin.Context) { + if err := ofauth.EmailBind(c.Request.Context(), callerUserID(c), c.Query("email"), c.Query("code")); err != nil { + compat.Fail(c, err.Error()) + return + } + compat.OKMessage(c, "") +} diff --git a/Wavelet/internal/apps/openflare/legacy/middleware.go b/Wavelet/internal/apps/openflare/legacy/middleware.go index f0c1dff7..20c9d173 100644 --- a/Wavelet/internal/apps/openflare/legacy/middleware.go +++ b/Wavelet/internal/apps/openflare/legacy/middleware.go @@ -9,11 +9,6 @@ import ( "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() diff --git a/Wavelet/internal/apps/openflare/legacy/register.go b/Wavelet/internal/apps/openflare/legacy/register.go index 68523ca8..e88de2dd 100644 --- a/Wavelet/internal/apps/openflare/legacy/register.go +++ b/Wavelet/internal/apps/openflare/legacy/register.go @@ -4,10 +4,14 @@ // Package legacy registers OpenFlare /api/* compatibility routes for the old frontend. package legacy -import "github.com/gin-gonic/gin" +import ( + "github.com/Rain-kl/Wavelet/internal/apps/openflare/compat" + "github.com/gin-gonic/gin" +) // RegisterRoutes mounts all OpenFlare legacy API routes under the /api group. func RegisterRoutes(apiGroup *gin.RouterGroup) { + apiGroup.Use(compat.BridgeOpenFlareToken()) registerAuthRoutes(apiGroup) registerOptionRoutes(apiGroup) registerOriginRoutes(apiGroup) diff --git a/Wavelet/internal/apps/openflare/legacy/register_auth.go b/Wavelet/internal/apps/openflare/legacy/register_auth.go index 61dd49b1..7fe17444 100644 --- a/Wavelet/internal/apps/openflare/legacy/register_auth.go +++ b/Wavelet/internal/apps/openflare/legacy/register_auth.go @@ -16,12 +16,16 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) { oauthGroup := apiGroup.Group("/oauth") { + oauthGroup.GET("/github", GitHubOAuth) + oauthGroup.GET("/wechat", WeChatOAuth) + oauthGroup.GET("/wechat/bind", compat.BridgeOpenFlareToken(), compat.UserAuth(), WeChatBind) + oauthGroup.GET("/email/bind", compat.BridgeOpenFlareToken(), compat.UserAuth(), EmailBind) 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.Use(compat.UserAuth()) { externalAccounts.GET("/", ListExternalAccounts) externalAccounts.POST("/:id/delete", DeleteExternalAccount) @@ -41,7 +45,7 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) { userGroup.GET("/logout", Logout) selfGroup := userGroup.Group("/") - selfGroup.Use(bridgeOpenFlareToken(), compat.UserAuth()) + selfGroup.Use(compat.UserAuth()) { selfGroup.GET("/self", GetSelf) selfGroup.POST("/self/update", UpdateSelf) @@ -50,7 +54,7 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) { } adminGroup := userGroup.Group("/") - adminGroup.Use(bridgeOpenFlareToken(), compat.AdminAuth()) + adminGroup.Use(compat.AdminAuth()) { adminGroup.GET("/", GetAllUsers) adminGroup.GET("/search", SearchUsers) @@ -63,7 +67,7 @@ func registerAuthRoutes(apiGroup *gin.RouterGroup) { } authSourceGroup := apiGroup.Group("/auth-sources") - authSourceGroup.Use(bridgeOpenFlareToken(), compat.RootAuth()) + authSourceGroup.Use(compat.RootAuth()) { authSourceGroup.GET("/", ListAuthSources) authSourceGroup.POST("/", CreateAuthSource) diff --git a/Wavelet/internal/apps/openflare/option/routers.go b/Wavelet/internal/apps/openflare/option/routers.go index 89f00f11..b2b2a24e 100644 --- a/Wavelet/internal/apps/openflare/option/routers.go +++ b/Wavelet/internal/apps/openflare/option/routers.go @@ -20,7 +20,7 @@ func RegisterRoutes(apiGroup *gin.RouterGroup) { apiGroup.GET("/about", getAboutHandler) optionRoute := apiGroup.Group("/option") - optionRoute.Use(compat.BridgeOpenFlareToken(), compat.RootAuth()) + optionRoute.Use(compat.RootAuth()) { optionRoute.GET("/", listOptionsHandler) optionRoute.POST("/update", updateOptionHandler) @@ -30,7 +30,7 @@ func RegisterRoutes(apiGroup *gin.RouterGroup) { } uptimeKumaRoute := apiGroup.Group("/uptimekuma") - uptimeKumaRoute.Use(compat.BridgeOpenFlareToken(), compat.RootAuth()) + uptimeKumaRoute.Use(compat.RootAuth()) { uptimeKumaRoute.POST("/sync", syncUptimeKumaHandler) } diff --git a/Wavelet/internal/apps/openflare/pages/data/pages/artifacts/published-site/1d0b1001941b126350515d2eaa00cea377a6ceffd4c72c76c77910686be8dc4e.zip b/Wavelet/internal/apps/openflare/pages/data/pages/artifacts/published-site/1d0b1001941b126350515d2eaa00cea377a6ceffd4c72c76c77910686be8dc4e.zip new file mode 100644 index 0000000000000000000000000000000000000000..75e34bb3d412d1181deb5fd01c11424adb87f239 GIT binary patch literal 142 zcmWIWW@Zs#-~d8&zy%b@%u7kF(90;v%{g_RjfH{X|Nj7Qb`JNucPc^ZnSeOJn~_O` e0bv5N9LNMzfG{t>o0SbD#|VTLK-vMsVE_R3B@`F{ literal 0 HcmV?d00001 diff --git a/Wavelet/internal/apps/openflare/pages/errs.go b/Wavelet/internal/apps/openflare/pages/errs.go index 4bbc6f3f..80e05f01 100644 --- a/Wavelet/internal/apps/openflare/pages/errs.go +++ b/Wavelet/internal/apps/openflare/pages/errs.go @@ -4,20 +4,23 @@ 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" + 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" + errPagesPackagePathEmpty = "Pages 部署包路径为空" + errPagesPackageNotInActiveConfig = "Pages 部署尚未进入激活配置" + errPagesInvalidSnapshotFormat = "配置快照格式无效" ) diff --git a/Wavelet/internal/apps/openflare/pages/logics.go b/Wavelet/internal/apps/openflare/pages/logics.go index ba1850f3..846667c1 100644 --- a/Wavelet/internal/apps/openflare/pages/logics.go +++ b/Wavelet/internal/apps/openflare/pages/logics.go @@ -5,6 +5,7 @@ package pages import ( "context" + "encoding/json" "errors" "fmt" "mime/multipart" @@ -341,6 +342,77 @@ func ActivateDeployment(ctx context.Context, projectID uint, deploymentID uint) return GetProject(ctx, project.ID) } +// GetDeploymentPackagePath returns the on-disk artifact path and download filename for an agent package request. +func GetDeploymentPackagePath(ctx context.Context, deploymentID uint) (string, string, error) { + deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) + if err != nil { + return "", "", err + } + if err = ensureDeploymentInActiveSnapshot(ctx, deployment.ID); err != nil { + return "", "", err + } + if strings.TrimSpace(deployment.ArtifactPath) == "" { + return "", "", errors.New(errPagesPackagePathEmpty) + } + if _, err = os.Stat(deployment.ArtifactPath); err != nil { + return "", "", fmt.Errorf("Pages 部署包不存在: %w", err) + } + return deployment.ArtifactPath, fmt.Sprintf("pages-deployment-%d.zip", deployment.ID), nil +} + +func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) error { + version, err := model.GetActiveConfigVersion(ctx) + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return errors.New(errPagesPackageNotInActiveConfig) + } + return err + } + routes, err := parseSnapshotRoutes(version.SnapshotJSON) + if err != nil { + return err + } + for _, route := range routes { + if route.UpstreamType != "pages" || route.PagesDeployment == nil { + continue + } + if route.PagesDeployment.DeploymentID == deploymentID { + return nil + } + } + return errors.New(errPagesPackageNotInActiveConfig) +} + +type snapshotPagesDeployment struct { + DeploymentID uint `json:"deployment_id"` +} + +type snapshotRouteRef struct { + UpstreamType string `json:"upstream_type"` + PagesDeployment *snapshotPagesDeployment `json:"pages_deployment"` +} + +func parseSnapshotRoutes(snapshotJSON string) ([]snapshotRouteRef, error) { + text := strings.TrimSpace(snapshotJSON) + if text == "" { + return []snapshotRouteRef{}, nil + } + if strings.HasPrefix(text, "[") { + var routes []snapshotRouteRef + if err := json.Unmarshal([]byte(text), &routes); err != nil { + return nil, errors.New(errPagesInvalidSnapshotFormat) + } + return routes, nil + } + var snapshot struct { + Routes []snapshotRouteRef `json:"routes"` + } + if err := json.Unmarshal([]byte(text), &snapshot); err != nil { + return nil, errors.New(errPagesInvalidSnapshotFormat) + } + return snapshot.Routes, nil +} + // DeleteDeployment 删除 Pages 部署。 func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error { project, err := model.GetPagesProjectByID(ctx, projectID) diff --git a/Wavelet/internal/apps/openflare/pages/logics_test.go b/Wavelet/internal/apps/openflare/pages/logics_test.go index 77a0413c..f9791993 100644 --- a/Wavelet/internal/apps/openflare/pages/logics_test.go +++ b/Wavelet/internal/apps/openflare/pages/logics_test.go @@ -4,7 +4,13 @@ package pages import ( + "archive/zip" + "bytes" "context" + "fmt" + "mime/multipart" + "net/http/httptest" + "strconv" "testing" "github.com/Rain-kl/Wavelet/internal/db" @@ -26,6 +32,7 @@ func setupPagesTestDB(t *testing.T) func() { &model.PagesProject{}, &model.PagesDeployment{}, &model.PagesDeploymentFile{}, + &model.ConfigVersion{}, )) db.SetDB(sqliteDB) @@ -82,3 +89,80 @@ func TestCreateProjectRejectsUnsafeFallbackPath(t *testing.T) { require.Error(t, err) assert.Contains(t, err.Error(), "回退路径") } + +func TestGetDeploymentPackagePathRequiresActiveConfigSnapshot(t *testing.T) { + cleanup := setupPagesTestDB(t) + defer cleanup() + ctx := context.Background() + + project, err := CreateProject(ctx, Input{ + Name: "Published Site", + Slug: "published-site", + Enabled: true, + }) + require.NoError(t, err) + + deployment, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "site.zip", testPagesZip(t, map[string]string{ + "index.html": "ok", + })), "root") + require.NoError(t, err) + + _, err = ActivateDeployment(ctx, project.ID, deployment.ID) + require.NoError(t, err) + + _, _, err = GetDeploymentPackagePath(ctx, deployment.ID) + require.Error(t, err) + assert.Contains(t, err.Error(), "激活配置") + + require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{ + Version: "v2026-001", + SnapshotJSON: fmt.Sprintf(`{"routes":[{"upstream_type":"pages","pages_deployment":{"deployment_id":%d}}]}`, deployment.ID), + MainConfig: "", + RenderedConfig: "", + SupportFilesJSON: "[]", + Checksum: "test-checksum", + IsActive: true, + CreatedBy: "test", + }).Error) + + filePath, fileName, err := GetDeploymentPackagePath(ctx, deployment.ID) + require.NoError(t, err) + assert.NotEmpty(t, filePath) + assert.Equal(t, "pages-deployment-"+strconv.FormatUint(uint64(deployment.ID), 10)+".zip", fileName) +} + +func testPagesZip(t *testing.T, files map[string]string) []byte { + t.Helper() + + var buffer bytes.Buffer + writer := zip.NewWriter(&buffer) + for name, content := range files { + file, err := writer.Create(name) + require.NoError(t, err) + _, err = file.Write([]byte(content)) + require.NoError(t, err) + } + require.NoError(t, writer.Close()) + return buffer.Bytes() +} + +func testPagesMultipartFile(t *testing.T, fileName string, content []byte) *multipart.FileHeader { + t.Helper() + + var body bytes.Buffer + writer := multipart.NewWriter(&body) + part, err := writer.CreateFormFile("package", fileName) + require.NoError(t, err) + _, err = part.Write(content) + require.NoError(t, err) + require.NoError(t, writer.Close()) + + req := httptest.NewRequest("POST", "/", &body) + req.Header.Set("Content-Type", writer.FormDataContentType()) + require.NoError(t, req.ParseMultipartForm(int64(len(content))+1024)) + + file, header, err := req.FormFile("package") + require.NoError(t, err) + file.Close() + return header +} diff --git a/Wavelet/internal/apps/openflare/relay/logics.go b/Wavelet/internal/apps/openflare/relay/logics.go index ce4fec7d..79c5e3dd 100644 --- a/Wavelet/internal/apps/openflare/relay/logics.go +++ b/Wavelet/internal/apps/openflare/relay/logics.go @@ -9,19 +9,38 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" ) const nodeStatusOnline = "online" +// ProxyStat describes a single frps proxy reported by the relay. +type ProxyStat 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"` +} + // 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"` + Version string `json:"version"` + ExtVersion string `json:"frp_version"` + RelayStatus string `json:"relay_status"` + FrpsConnCount int `json:"frps_connections"` + FrpsProxyCount int `json:"frps_proxy_count"` + FrpsClientCount int `json:"frps_client_count"` + FrpsProxies []ProxyStat `json:"frps_proxies,omitempty"` + Name string `json:"name"` + IP string `json:"ip"` + Profile *agent.NodeSystemProfile `json:"profile,omitempty"` + Snapshot *agent.NodeMetricSnapshot `json:"snapshot,omitempty"` + HealthEvents []agent.NodeHealthEvent `json:"health_events,omitempty"` } // Config is the frps configuration sent to the relay. @@ -109,6 +128,11 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil { return nil, fmt.Errorf("update relay heartbeat: %w", err) } + if err := reconcileRelayHealthEvents(ctx, node.NodeID, payload.RelayStatus, now); err != nil { + return nil, fmt.Errorf("reconcile relay health events: %w", err) + } + agent.RefreshAccessTokenCache(ctx, node) + persistRelayHeartbeatObservability(ctx, node.NodeID, payload, now) return &HeartbeatResponse{ RelayConfig: buildRelayConfig(node), diff --git a/Wavelet/internal/apps/openflare/relay/logics_test.go b/Wavelet/internal/apps/openflare/relay/logics_test.go new file mode 100644 index 00000000..eb0f6caa --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/logics_test.go @@ -0,0 +1,174 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +import ( + "context" + "encoding/json" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" + "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 setupRelayTestDB(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.OpenFlareNodeSystemProfile{}, + &model.OpenFlareMetricSnapshot{}, + &model.OpenFlareHealthEvent{}, + &model.OpenFlareNodeObservationFrps{}, + )) + + db.SetDB(sqliteDB) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + + return func() { + db.SetDB(nil) + option.ResetInitializationForTest() + agent.ResetAuthCacheForTest() + } +} + +func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) { + cleanup := setupRelayTestDB(t) + defer cleanup() + + ctx := context.Background() + now := time.Now().UTC().Truncate(time.Second) + + node := &model.OpenFlareNode{ + NodeID: "node-relay-observe", + Name: "relay-1", + AccessToken: "relay-token", + Status: "pending", + NodeType: "tunnel_relay", + RelayStatus: "unknown", + } + require.NoError(t, db.DB(ctx).Create(node).Error) + + proxies := []ProxyStat{ + { + Name: "proxy-a", + Type: "http", + Status: "online", + ClientVersion: "0.61.0", + ClientAddr: "10.0.0.2:12345", + }, + } + _, err := Heartbeat(ctx, node, HeartbeatPayload{ + Version: "v0.1.0", + ExtVersion: "0.61.0", + RelayStatus: "healthy", + FrpsConnCount: 7, + FrpsProxyCount: 3, + FrpsClientCount: 2, + FrpsProxies: proxies, + Name: "relay-runtime", + IP: "203.0.113.9", + Profile: &agent.NodeSystemProfile{ + Hostname: "relay-runtime", + OSName: "Ubuntu", + OSVersion: "24.04", + Architecture: "amd64", + CPUCores: 4, + ReportedAtUnix: now.Unix(), + }, + Snapshot: &agent.NodeMetricSnapshot{ + CapturedAtUnix: now.Unix(), + CPUUsagePercent: 12.5, + NetworkRxBytes: 1024, + NetworkTxBytes: 2048, + }, + HealthEvents: []agent.NodeHealthEvent{}, + }) + require.NoError(t, err) + + var stored model.OpenFlareNode + require.NoError(t, db.DB(ctx).Where("node_id = ?", node.NodeID).First(&stored).Error) + assert.Equal(t, "online", stored.Status) + assert.Equal(t, "healthy", stored.RelayStatus) + assert.Equal(t, "203.0.113.9", stored.IP) + assert.Equal(t, "v0.1.0", stored.Version) + assert.Equal(t, "0.61.0", stored.ExtVersion) + + profile, err := model.GetOpenFlareNodeSystemProfile(ctx, node.NodeID) + require.NoError(t, err) + assert.Equal(t, "relay-runtime", profile.Hostname) + assert.Equal(t, "Ubuntu", profile.OSName) + + snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10) + require.NoError(t, err) + require.Len(t, snapshots, 1) + assert.Equal(t, 12.5, snapshots[0].CPUUsagePercent) + + frpsObs, err := model.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1) + require.NoError(t, err) + require.Len(t, frpsObs, 1) + assert.Equal(t, 7, frpsObs[0].FrpsConnections) + assert.Equal(t, 3, frpsObs[0].FrpsProxyCount) + assert.Equal(t, 2, frpsObs[0].FrpsClientCount) + + var decoded []ProxyStat + require.NoError(t, json.Unmarshal([]byte(frpsObs[0].FrpsProxies), &decoded)) + require.Len(t, decoded, 1) + assert.Equal(t, "proxy-a", decoded[0].Name) + assert.Equal(t, "online", decoded[0].Status) +} + +func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) { + cleanup := setupRelayTestDB(t) + defer cleanup() + + ctx := context.Background() + node := &model.OpenFlareNode{ + NodeID: "node-relay-unhealthy", + Name: "relay-unhealthy", + AccessToken: "relay-token-unhealthy", + Status: "pending", + NodeType: "tunnel_relay", + RelayStatus: "healthy", + } + require.NoError(t, db.DB(ctx).Create(node).Error) + + _, err := Heartbeat(ctx, node, HeartbeatPayload{ + Version: "v0.1.0", + ExtVersion: "0.61.0", + RelayStatus: "unhealthy", + }) + require.NoError(t, err) + + events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, true, 10) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, relayFrpsUnhealthyEventType, events[0].EventType) + assert.Equal(t, "active", events[0].Status) + + _, err = Heartbeat(ctx, node, HeartbeatPayload{ + Version: "v0.1.0", + ExtVersion: "0.61.0", + RelayStatus: "healthy", + }) + require.NoError(t, err) + + events, err = model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 10) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, "resolved", events[0].Status) +} diff --git a/Wavelet/internal/apps/openflare/relay/observability.go b/Wavelet/internal/apps/openflare/relay/observability.go new file mode 100644 index 00000000..2c188514 --- /dev/null +++ b/Wavelet/internal/apps/openflare/relay/observability.go @@ -0,0 +1,69 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package relay + +import ( + "context" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "go.uber.org/zap" + "gorm.io/gorm" +) + +const relayFrpsUnhealthyEventType = "frps_unhealthy" + +func reconcileRelayHealthEvents(ctx context.Context, nodeID string, relayStatus string, reportedAt time.Time) error { + if relayStatus == "unknown" { + return nil + } + managedTypes := map[string]struct{}{ + relayFrpsUnhealthyEventType: {}, + } + events := []agent.NodeHealthEvent{} + if relayStatus == "unhealthy" { + events = append(events, agent.NodeHealthEvent{ + EventType: relayFrpsUnhealthyEventType, + Severity: "critical", + Message: "frps runtime is not healthy", + TriggeredAtUnix: reportedAt.Unix(), + Metadata: map[string]string{ + "relay_status": relayStatus, + }, + }) + } + conn := db.DB(ctx) + if conn == nil { + return nil + } + return conn.Transaction(func(tx *gorm.DB) error { + return agent.ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, managedTypes) + }) +} + +func persistRelayHeartbeatObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) { + agent.PersistHeartbeatObservability(ctx, nodeID, agent.NodePayload{ + Profile: payload.Profile, + Snapshot: payload.Snapshot, + HealthEvents: payload.HealthEvents, + }, reportedAt) + + conn := db.DB(ctx) + if conn == nil { + return + } + frpsObs := &model.OpenFlareNodeObservationFrps{ + NodeID: nodeID, + CapturedAt: reportedAt, + FrpsConnections: payload.FrpsConnCount, + FrpsProxyCount: payload.FrpsProxyCount, + FrpsClientCount: payload.FrpsClientCount, + FrpsProxies: agent.MarshalJSON(payload.FrpsProxies), + } + if err := conn.Create(frpsObs).Error; err != nil { + zap.L().Error("persist relay frps observation failed", zap.String("node_id", nodeID), zap.Error(err)) + } +} diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190011_add_of_node_access_logs_composite_index.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190011_add_of_node_access_logs_composite_index.sql new file mode 100644 index 00000000..a478ae9d --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190011_add_of_node_access_logs_composite_index.sql @@ -0,0 +1,5 @@ +-- +goose Up +CREATE INDEX IF NOT EXISTS idx_of_node_access_logs_node_id_logged_at ON of_node_access_logs (node_id, logged_at); + +-- +goose Down +DROP INDEX IF EXISTS idx_of_node_access_logs_node_id_logged_at; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/postgres/202606190012_create_of_node_obs_frpc.sql b/Wavelet/internal/db/migrator/goose/postgres/202606190012_create_of_node_obs_frpc.sql new file mode 100644 index 00000000..a98437ac --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/postgres/202606190012_create_of_node_obs_frpc.sql @@ -0,0 +1,15 @@ +-- +goose Up +CREATE TABLE of_node_obs_frpc ( + id BIGSERIAL PRIMARY KEY, + node_id VARCHAR(64) NOT NULL, + captured_at TIMESTAMPTZ NOT NULL, + tunnel_status VARCHAR(16) NOT NULL DEFAULT '', + connected_relays_count INTEGER NOT NULL DEFAULT 0, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX idx_of_node_obs_frpc_node_id ON of_node_obs_frpc (node_id); +CREATE INDEX idx_of_node_obs_frpc_captured_at ON of_node_obs_frpc (captured_at); + +-- +goose Down +DROP TABLE IF EXISTS of_node_obs_frpc; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190011_add_of_node_access_logs_composite_index.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190011_add_of_node_access_logs_composite_index.sql new file mode 100644 index 00000000..a478ae9d --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190011_add_of_node_access_logs_composite_index.sql @@ -0,0 +1,5 @@ +-- +goose Up +CREATE INDEX IF NOT EXISTS idx_of_node_access_logs_node_id_logged_at ON of_node_access_logs (node_id, logged_at); + +-- +goose Down +DROP INDEX IF EXISTS idx_of_node_access_logs_node_id_logged_at; \ No newline at end of file diff --git a/Wavelet/internal/db/migrator/goose/sqlite/202606190012_create_of_node_obs_frpc.sql b/Wavelet/internal/db/migrator/goose/sqlite/202606190012_create_of_node_obs_frpc.sql new file mode 100644 index 00000000..033e6f59 --- /dev/null +++ b/Wavelet/internal/db/migrator/goose/sqlite/202606190012_create_of_node_obs_frpc.sql @@ -0,0 +1,15 @@ +-- +goose Up +CREATE TABLE IF NOT EXISTS of_node_obs_frpc ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + node_id TEXT NOT NULL, + captured_at DATETIME NOT NULL, + tunnel_status TEXT NOT NULL DEFAULT '', + connected_relays_count INTEGER NOT NULL DEFAULT 0, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX IF NOT EXISTS idx_of_node_obs_frpc_node_id ON of_node_obs_frpc (node_id); +CREATE INDEX IF NOT EXISTS idx_of_node_obs_frpc_captured_at ON of_node_obs_frpc (captured_at); + +-- +goose Down +DROP TABLE IF EXISTS of_node_obs_frpc; \ No newline at end of file diff --git a/Wavelet/internal/model/openflare_access_log.go b/Wavelet/internal/model/openflare_access_log.go new file mode 100644 index 00000000..abcd8af2 --- /dev/null +++ b/Wavelet/internal/model/openflare_access_log.go @@ -0,0 +1,757 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "errors" + "fmt" + "sort" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "gorm.io/gorm" +) + +const openFlareAccessLogTable = "of_node_access_logs" + +type openFlareAccessLogBucketAggregateRow struct { + BucketEpoch int64 `gorm:"column:bucket_epoch"` + RequestCount int64 `gorm:"column:request_count"` + SuccessCount int64 `gorm:"column:success_count"` + ClientErrorCount int64 `gorm:"column:client_error_count"` + ServerErrorCount int64 `gorm:"column:server_error_count"` +} + +type openFlareAccessLogBucketDimensionRow struct { + BucketEpoch int64 `gorm:"column:bucket_epoch"` + Value string `gorm:"column:value"` +} + +type openFlareAccessLogIPAggregateRow struct { + RemoteAddr string `gorm:"column:remote_addr"` + RequestCount int64 `gorm:"column:request_count"` + SuccessCount int64 `gorm:"column:success_count"` + ClientErrorCount int64 `gorm:"column:client_error_count"` + ServerErrorCount int64 `gorm:"column:server_error_count"` + LastSeenEpoch int64 `gorm:"column:last_seen_epoch"` +} + +type openFlareAccessLogIPSummaryRow struct { + RemoteAddr string `gorm:"column:remote_addr"` + TotalRequests int64 `gorm:"column:total_requests"` + RecentRequests int64 `gorm:"column:recent_requests"` + LastSeenEpoch int64 `gorm:"column:last_seen_epoch"` +} + +type openFlareAccessLogIPTrendRow struct { + BucketEpoch int64 `gorm:"column:bucket_epoch"` + RequestCount int64 `gorm:"column:request_count"` +} + +// ListOpenFlareAccessLogs lists access logs matching the query. +func ListOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + tx := applyOpenFlareAccessLogFilters(conn.Model(&OpenFlareAccessLog{}), query) + tx = tx.Order(openFlareAccessLogOrderClause(query.SortBy, query.SortOrder)) + if query.PageSize > 0 { + if query.Page < 0 { + query.Page = 0 + } + tx = tx.Offset(query.Page * query.PageSize).Limit(query.PageSize) + } + var rows []*OpenFlareAccessLog + if err := tx.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareAccessLog{}, nil + } + return nil, err + } + return rows, nil +} + +// CountOpenFlareAccessLogs counts access logs and distinct IPs matching the query. +func CountOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, 0, errors.New(errDatabaseNotInitialized) + } + totalRecords, err := countOpenFlareAccessLogRecords(conn, query) + if err != nil { + if isMissingTableError(err) { + return 0, 0, nil + } + return 0, 0, err + } + totalIPs, err := countDistinctOpenFlareAccessLogIPs(conn, query) + if err != nil { + if isMissingTableError(err) { + return 0, 0, nil + } + return 0, 0, err + } + return totalRecords, totalIPs, nil +} + +// ListOpenFlareAccessLogRegionCounts returns region counts for access logs. +func ListOpenFlareAccessLogRegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + filter := OpenFlareAccessLogQuery{ + NodeID: nodeID, + Since: since, + } + clause, args := buildOpenFlareAccessLogFilterClause(filter) + sql := fmt.Sprintf(` +SELECT TRIM(region) AS region, COUNT(*) AS count +FROM %s +WHERE %s AND TRIM(region) <> '' +GROUP BY TRIM(region) +ORDER BY count DESC, region ASC`, openFlareAccessLogTable, clause) + if limit > 0 { + sql += fmt.Sprintf(" LIMIT %d", limit) + } + var rows []*OpenFlareAccessLogRegionCount + if err := conn.Raw(sql, args...).Scan(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareAccessLogRegionCount{}, nil + } + return nil, err + } + return rows, nil +} + +// ListOpenFlareAccessLogBuckets lists folded access log buckets. +func ListOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) { + rows, err := buildOpenFlareAccessLogBucketRows(ctx, query) + if err != nil { + return nil, err + } + start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize) + if start >= len(rows) { + return []*OpenFlareAccessLogBucketRow{}, nil + } + return rows[start:end], nil +} + +// CountOpenFlareAccessLogBuckets counts folded access log buckets. +func CountOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) (int64, error) { + rows, err := buildOpenFlareAccessLogBucketRows(ctx, query) + if err != nil { + return 0, err + } + return int64(len(rows)), nil +} + +// ListOpenFlareAccessLogBucketIPs lists folded IP rows for a bucket window. +func ListOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) { + rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query) + if err != nil { + return nil, err + } + start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize) + if start >= len(rows) { + return []*OpenFlareAccessLogBucketIPRow{}, nil + } + return rows[start:end], nil +} + +// CountOpenFlareAccessLogBucketIPs counts folded IP rows for a bucket window. +func CountOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) (int64, error) { + rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query) + if err != nil { + return 0, err + } + return int64(len(rows)), nil +} + +// ListOpenFlareAccessLogIPSummaries lists IP summaries. +func ListOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) { + rows, err := buildOpenFlareAccessLogIPSummaryRows(ctx, query, recentSince) + if err != nil { + return nil, err + } + start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize) + if start >= len(rows) { + return []*OpenFlareAccessLogIPSummaryRow{}, nil + } + return rows[start:end], nil +} + +// CountOpenFlareAccessLogIPSummaries counts IP summaries. +func CountOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery) (int64, error) { + rows, err := buildOpenFlareAccessLogIPSummaryRows(ctx, query, time.Time{}) + if err != nil { + return 0, err + } + return int64(len(rows)), nil +} + +// ListOpenFlareAccessLogIPTrend lists IP trend points. +func ListOpenFlareAccessLogIPTrend(ctx context.Context, query OpenFlareAccessLogIPTrendQuery) ([]*OpenFlareAccessLogIPTrendRow, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + remoteAddr := strings.TrimSpace(query.RemoteAddr) + if remoteAddr == "" { + return []*OpenFlareAccessLogIPTrendRow{}, nil + } + filter := OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: remoteAddr, + Host: query.Host, + Since: query.Since, + } + clause, args := buildOpenFlareAccessLogFilterClause(filter) + bucketSeconds := int64(query.BucketMinutes * 60) + if bucketSeconds <= 0 { + bucketSeconds = 1800 + } + bucketExpr := openFlareAccessLogBucketEpochExpr(openFlareAccessLogDialect(conn), bucketSeconds) + queryClause := combineOpenFlareAccessLogSQLClauses(clause, "TRIM(remote_addr) = ?") + queryArgs := append(append([]any{}, args...), remoteAddr) + sql := fmt.Sprintf(` +SELECT + %s AS bucket_epoch, + COUNT(*) AS request_count +FROM %s +WHERE %s +GROUP BY bucket_epoch +ORDER BY bucket_epoch ASC`, bucketExpr, openFlareAccessLogTable, queryClause) + var rows []*OpenFlareAccessLogIPTrendRow + if err := conn.Raw(sql, queryArgs...).Scan(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareAccessLogIPTrendRow{}, nil + } + return nil, err + } + return rows, nil +} + +// DeleteOpenFlareAccessLogsBefore deletes access logs older than cutoff. +func DeleteOpenFlareAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + result := conn.Where("logged_at < ?", cutoff).Delete(&OpenFlareAccessLog{}) + if result.Error != nil { + if isMissingTableError(result.Error) { + return 0, nil + } + return 0, result.Error + } + return result.RowsAffected, nil +} + +func buildOpenFlareAccessLogBucketRows(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + filter := openFlareAccessLogQueryFromBucket(query) + clause, args := buildOpenFlareAccessLogFilterClause(filter) + bucketSeconds := int64(query.FoldMinutes * 60) + if bucketSeconds <= 0 { + bucketSeconds = 180 + } + bucketExpr := openFlareAccessLogBucketEpochExpr(openFlareAccessLogDialect(conn), bucketSeconds) + + type bucketAccumulator struct { + requestCount int64 + uniqueIPs map[string]struct{} + uniqueHosts map[string]struct{} + successCount int64 + clientErrorCount int64 + serverErrorCount int64 + } + accumulators := make(map[int64]*bucketAccumulator) + + var partials []openFlareAccessLogBucketAggregateRow + sql := fmt.Sprintf(` +SELECT + %s AS bucket_epoch, + COUNT(*) AS request_count, + SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count, + SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count, + SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count +FROM %s +WHERE %s +GROUP BY bucket_epoch`, bucketExpr, openFlareAccessLogTable, clause) + if err := conn.Raw(sql, args...).Scan(&partials).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareAccessLogBucketRow{}, nil + } + return nil, err + } + for _, partial := range partials { + accumulator := accumulators[partial.BucketEpoch] + if accumulator == nil { + accumulator = &bucketAccumulator{ + uniqueIPs: make(map[string]struct{}), + uniqueHosts: make(map[string]struct{}), + } + accumulators[partial.BucketEpoch] = accumulator + } + accumulator.requestCount += partial.RequestCount + accumulator.successCount += partial.SuccessCount + accumulator.clientErrorCount += partial.ClientErrorCount + accumulator.serverErrorCount += partial.ServerErrorCount + } + + for _, column := range []string{"remote_addr", "host"} { + dimensions, err := queryOpenFlareAccessLogBucketDimensionRows(conn, clause, args, column, bucketExpr) + if err != nil { + return nil, err + } + for _, item := range dimensions { + accumulator := accumulators[item.BucketEpoch] + if accumulator == nil { + accumulator = &bucketAccumulator{ + uniqueIPs: make(map[string]struct{}), + uniqueHosts: make(map[string]struct{}), + } + accumulators[item.BucketEpoch] = accumulator + } + trimmed := strings.TrimSpace(item.Value) + if trimmed == "" { + continue + } + switch column { + case "remote_addr": + accumulator.uniqueIPs[trimmed] = struct{}{} + case "host": + accumulator.uniqueHosts[trimmed] = struct{}{} + } + } + } + + rows := make([]*OpenFlareAccessLogBucketRow, 0, len(accumulators)) + for bucketEpoch, accumulator := range accumulators { + rows = append(rows, &OpenFlareAccessLogBucketRow{ + BucketEpoch: bucketEpoch, + RequestCount: accumulator.requestCount, + UniqueIPCount: int64(len(accumulator.uniqueIPs)), + UniqueHostCount: int64(len(accumulator.uniqueHosts)), + SuccessCount: accumulator.successCount, + ClientErrorCount: accumulator.clientErrorCount, + ServerErrorCount: accumulator.serverErrorCount, + }) + } + sortOpenFlareAccessLogBucketRows(rows, query.SortBy, query.SortOrder) + return rows, nil +} + +func queryOpenFlareAccessLogBucketDimensionRows(conn *gorm.DB, clause string, args []any, column string, bucketExpr string) ([]openFlareAccessLogBucketDimensionRow, error) { + var rows []openFlareAccessLogBucketDimensionRow + sql := fmt.Sprintf(` +SELECT + %s AS bucket_epoch, + TRIM(%s) AS value +FROM %s +WHERE %s AND TRIM(%s) <> '' +GROUP BY bucket_epoch, TRIM(%s)`, bucketExpr, column, openFlareAccessLogTable, clause, column, column) + if err := conn.Raw(sql, args...).Scan(&rows).Error; err != nil { + if isMissingTableError(err) { + return []openFlareAccessLogBucketDimensionRow{}, nil + } + return nil, err + } + return rows, nil +} + +func buildOpenFlareAccessLogBucketIPRows(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) { + if query.BucketStartedAt.IsZero() { + return []*OpenFlareAccessLogBucketIPRow{}, nil + } + foldMinutes := query.FoldMinutes + if foldMinutes <= 0 { + foldMinutes = 3 + } + bucketStartedAt := query.BucketStartedAt.UTC() + filter := OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Path: query.Path, + Since: bucketStartedAt, + Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute), + } + rows, err := queryOpenFlareAccessLogIPAggregateRows(ctx, filter, false) + if err != nil { + return nil, err + } + sortOpenFlareAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder) + return rows, nil +} + +func buildOpenFlareAccessLogIPSummaryRows(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + filter := OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Since: query.Since, + } + clause, args := buildOpenFlareAccessLogFilterClause(filter) + lastSeenExpr := openFlareAccessLogEpochExpr(openFlareAccessLogDialect(conn)) + recentClause := "0" + queryArgs := make([]any, 0, len(args)+1) + if !recentSince.IsZero() { + recentClause = "CASE WHEN logged_at >= ? THEN 1 ELSE 0 END" + queryArgs = append(queryArgs, recentSince) + } + queryArgs = append(queryArgs, args...) + sql := fmt.Sprintf(` +SELECT + TRIM(remote_addr) AS remote_addr, + COUNT(*) AS total_requests, + SUM(%s) AS recent_requests, + MAX(%s) AS last_seen_epoch +FROM %s +WHERE %s AND TRIM(remote_addr) <> '' +GROUP BY TRIM(remote_addr)`, recentClause, lastSeenExpr, openFlareAccessLogTable, clause) + var partials []openFlareAccessLogIPSummaryRow + if err := conn.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareAccessLogIPSummaryRow{}, nil + } + return nil, err + } + rows := make([]*OpenFlareAccessLogIPSummaryRow, 0, len(partials)) + for _, partial := range partials { + remoteAddr := strings.TrimSpace(partial.RemoteAddr) + if remoteAddr == "" { + continue + } + rows = append(rows, &OpenFlareAccessLogIPSummaryRow{ + RemoteAddr: remoteAddr, + TotalRequests: partial.TotalRequests, + RecentRequests: partial.RecentRequests, + LastSeenEpoch: partial.LastSeenEpoch, + }) + } + sortOpenFlareAccessLogIPSummaryRows(rows, query.SortBy, query.SortOrder) + return rows, nil +} + +func queryOpenFlareAccessLogIPAggregateRows(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]*OpenFlareAccessLogBucketIPRow, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + clause, args := buildOpenFlareAccessLogFilterClause(filter) + lastSeenExpr := openFlareAccessLogEpochExpr(openFlareAccessLogDialect(conn)) + queryClause := clause + queryArgs := append([]any{}, args...) + if exactRemoteAddr { + trimmed := strings.TrimSpace(filter.RemoteAddr) + if trimmed == "" { + return []*OpenFlareAccessLogBucketIPRow{}, nil + } + queryClause = combineOpenFlareAccessLogSQLClauses(queryClause, "TRIM(remote_addr) = ?") + queryArgs = append(queryArgs, trimmed) + } + sql := fmt.Sprintf(` +SELECT + TRIM(remote_addr) AS remote_addr, + COUNT(*) AS request_count, + SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count, + SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count, + SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count, + MAX(%s) AS last_seen_epoch +FROM %s +WHERE %s AND TRIM(remote_addr) <> '' +GROUP BY TRIM(remote_addr)`, lastSeenExpr, openFlareAccessLogTable, queryClause) + var partials []openFlareAccessLogIPAggregateRow + if err := conn.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareAccessLogBucketIPRow{}, nil + } + return nil, err + } + rows := make([]*OpenFlareAccessLogBucketIPRow, 0, len(partials)) + for _, partial := range partials { + remoteAddr := strings.TrimSpace(partial.RemoteAddr) + if remoteAddr == "" { + continue + } + rows = append(rows, &OpenFlareAccessLogBucketIPRow{ + RemoteAddr: remoteAddr, + RequestCount: partial.RequestCount, + SuccessCount: partial.SuccessCount, + ClientErrorCount: partial.ClientErrorCount, + ServerErrorCount: partial.ServerErrorCount, + LastSeenEpoch: partial.LastSeenEpoch, + }) + } + return rows, nil +} + +func openFlareAccessLogQueryFromBucket(query OpenFlareAccessLogBucketQuery) OpenFlareAccessLogQuery { + return OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Path: query.Path, + Since: query.Since, + } +} + +func buildOpenFlareAccessLogFilterClause(query OpenFlareAccessLogQuery) (string, []any) { + parts := make([]string, 0, 6) + args := make([]any, 0, 6) + if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" { + parts = append(parts, "node_id = ?") + args = append(args, trimmed) + } + if trimmed := strings.TrimSpace(query.RemoteAddr); trimmed != "" { + parts = append(parts, "remote_addr LIKE ?") + args = append(args, trimmed+"%") + } + if trimmed := strings.TrimSpace(query.Host); trimmed != "" { + parts = append(parts, "host LIKE ?") + args = append(args, trimmed+"%") + } + if trimmed := strings.TrimSpace(query.Path); trimmed != "" { + parts = append(parts, "path LIKE ?") + args = append(args, trimmed+"%") + } + if !query.Since.IsZero() { + parts = append(parts, "logged_at >= ?") + args = append(args, query.Since) + } + if !query.Until.IsZero() { + parts = append(parts, "logged_at < ?") + args = append(args, query.Until) + } + if len(parts) == 0 { + return "TRUE", nil + } + return strings.Join(parts, " AND "), args +} + +func applyOpenFlareAccessLogFilters(tx *gorm.DB, query OpenFlareAccessLogQuery) *gorm.DB { + clause, args := buildOpenFlareAccessLogFilterClause(query) + if clause == "TRUE" { + return tx + } + return tx.Where(clause, args...) +} + +func countOpenFlareAccessLogRecords(conn *gorm.DB, query OpenFlareAccessLogQuery) (int64, error) { + var count int64 + if err := applyOpenFlareAccessLogFilters(conn.Model(&OpenFlareAccessLog{}), query).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +func countDistinctOpenFlareAccessLogIPs(conn *gorm.DB, query OpenFlareAccessLogQuery) (int64, error) { + clause, args := buildOpenFlareAccessLogFilterClause(query) + sql := fmt.Sprintf(` +SELECT COUNT(*) FROM ( + SELECT TRIM(remote_addr) AS remote_addr + FROM %s + WHERE %s AND remote_addr <> '' + GROUP BY TRIM(remote_addr) +) AS ips`, openFlareAccessLogTable, clause) + var total int64 + if err := conn.Raw(sql, args...).Scan(&total).Error; err != nil { + return 0, err + } + return total, nil +} + +func openFlareAccessLogDialect(conn *gorm.DB) string { + if conn == nil || conn.Dialector == nil { + return "sqlite" + } + switch conn.Dialector.Name() { + case "postgres": + return "postgres" + default: + return "sqlite" + } +} + +func openFlareAccessLogBucketEpochExpr(dialect string, bucketSeconds int64) string { + switch dialect { + case "postgres": + return fmt.Sprintf("FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC') / %d) * %d", bucketSeconds, bucketSeconds) + default: + return fmt.Sprintf("(CAST(strftime('%%s', logged_at) AS INTEGER) / %d) * %d", bucketSeconds, bucketSeconds) + } +} + +func openFlareAccessLogEpochExpr(dialect string) string { + switch dialect { + case "postgres": + return "FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC'))::bigint" + default: + return "CAST((julianday(logged_at) - 2440587.5) * 86400 AS INTEGER)" + } +} + +func combineOpenFlareAccessLogSQLClauses(left string, right string) string { + if strings.TrimSpace(left) == "" || left == "TRUE" { + return right + } + return left + " AND " + right +} + +func openFlareAccessLogOrderClause(sortBy string, sortOrder string) string { + direction := "DESC" + if openFlareAccessLogNormalizeSortOrder(sortOrder) == "asc" { + direction = "ASC" + } + column := "logged_at" + switch strings.TrimSpace(sortBy) { + case "status_code": + column = "status_code" + case "remote_addr": + column = "remote_addr" + case "host": + column = "host" + case "path": + column = "path" + } + if column == "logged_at" { + return column + " " + direction + ", id " + direction + } + return column + " " + direction + ", logged_at " + direction + ", id " + direction +} + +func sortOpenFlareAccessLogBucketIPRows(items []*OpenFlareAccessLogBucketIPRow, sortBy string, sortOrder string) { + desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != "asc" + sort.Slice(items, func(i int, j int) bool { + left := items[i] + right := items[j] + if left == nil || right == nil { + return left != nil + } + var compare int + switch strings.TrimSpace(sortBy) { + case "last_seen_at": + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + case "remote_addr": + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + default: + compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount) + } + if compare == 0 { + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + } + if compare == 0 { + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + } + if desc { + return compare > 0 + } + return compare < 0 + }) +} + +func sortOpenFlareAccessLogBucketRows(items []*OpenFlareAccessLogBucketRow, sortBy string, sortOrder string) { + desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != "asc" + sort.Slice(items, func(i int, j int) bool { + left := items[i] + right := items[j] + if left == nil || right == nil { + return left != nil + } + var compare int + switch strings.TrimSpace(sortBy) { + case "request_count": + compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount) + default: + compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch) + } + if compare == 0 { + compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch) + } + if desc { + return compare > 0 + } + return compare < 0 + }) +} + +func sortOpenFlareAccessLogIPSummaryRows(items []*OpenFlareAccessLogIPSummaryRow, sortBy string, sortOrder string) { + desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != "asc" + sort.Slice(items, func(i int, j int) bool { + left := items[i] + right := items[j] + if left == nil || right == nil { + return left != nil + } + var compare int + switch strings.TrimSpace(sortBy) { + case "recent_requests": + compare = openFlareAccessLogCompareInt64(left.RecentRequests, right.RecentRequests) + case "last_seen_at": + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + case "remote_addr": + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + default: + compare = openFlareAccessLogCompareInt64(left.TotalRequests, right.TotalRequests) + } + if compare == 0 { + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + } + if compare == 0 { + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + } + if desc { + return compare > 0 + } + return compare < 0 + }) +} + +func openFlareAccessLogPaginateBounds(total int, page int, pageSize int) (int, int) { + if page < 0 { + page = 0 + } + if pageSize <= 0 { + return 0, total + } + start := page * pageSize + if start > total { + start = total + } + end := start + pageSize + if end > total { + end = total + } + return start, end +} + +func openFlareAccessLogNormalizeSortOrder(sortOrder string) string { + if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") { + return "asc" + } + return "desc" +} + +func openFlareAccessLogCompareInt64(left int64, right int64) int { + switch { + case left > right: + return 1 + case left < right: + return -1 + default: + return 0 + } +} diff --git a/Wavelet/internal/model/openflare_access_log_test.go b/Wavelet/internal/model/openflare_access_log_test.go new file mode 100644 index 00000000..e0c50549 --- /dev/null +++ b/Wavelet/internal/model/openflare_access_log_test.go @@ -0,0 +1,144 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func()) { + t.Helper() + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + require.NoError(t, err) + require.NoError(t, sqliteDB.AutoMigrate(&OpenFlareAccessLog{})) + db.SetDB(sqliteDB) + return context.Background(), func() { + db.SetDB(nil) + } +} + +func seedOpenFlareAccessLogs(t *testing.T, ctx context.Context, now time.Time) { + t.Helper() + records := []*OpenFlareAccessLog{ + {NodeID: "node-a", LoggedAt: now.Add(-5 * time.Minute), RemoteAddr: "1.1.1.1", Region: "US", Host: "a.example.com", Path: "/alpha", StatusCode: 200}, + {NodeID: "node-a", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "2.2.2.2", Region: "US", Host: "a.example.com", Path: "/beta", StatusCode: 404}, + {NodeID: "node-b", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "1.1.1.1", Region: "EU", Host: "b.example.com", Path: "/gamma", StatusCode: 502}, + {NodeID: "node-b", LoggedAt: now.Add(-2 * time.Minute), RemoteAddr: "3.3.3.3", Region: "EU", Host: "b.example.com", Path: "/delta", StatusCode: 200}, + {NodeID: "node-b", LoggedAt: now.Add(-1 * time.Minute), RemoteAddr: "", Region: "", Host: "b.example.com", Path: "/empty-ip", StatusCode: 200}, + } + for index, record := range records { + require.NoError(t, db.DB(ctx).Create(record).Error, "seed access log %d", index) + } +} + +func TestListOpenFlareAccessLogsPaginated(t *testing.T) { + ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t) + defer cleanup() + + now := time.Now().UTC() + for index := range 15 { + record := &OpenFlareAccessLog{ + NodeID: "node-page", + LoggedAt: now.Add(-time.Duration(index) * time.Minute), + RemoteAddr: fmt.Sprintf("203.0.113.%d", (index%5)+1), + Host: "example.com", + Path: fmt.Sprintf("/path-%02d", index), + StatusCode: 200, + } + require.NoError(t, db.DB(ctx).Create(record).Error) + } + + query := OpenFlareAccessLogQuery{ + NodeID: "node-page", + Since: now.Add(-24 * time.Hour), + Page: 1, + PageSize: 5, + SortBy: "logged_at", + SortOrder: "desc", + } + page, err := ListOpenFlareAccessLogs(ctx, query) + require.NoError(t, err) + require.Len(t, page, 5) + assert.Equal(t, "/path-05", page[0].Path) + assert.Equal(t, "/path-09", page[4].Path) +} + +func TestCountOpenFlareAccessLogs(t *testing.T) { + ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t) + defer cleanup() + + now := time.Now().UTC() + seedOpenFlareAccessLogs(t, ctx, now) + + query := OpenFlareAccessLogQuery{ + Since: now.Add(-10 * time.Minute), + } + totalRecords, totalIPs, err := CountOpenFlareAccessLogs(ctx, query) + require.NoError(t, err) + assert.Equal(t, int64(5), totalRecords) + assert.Equal(t, int64(3), totalIPs) +} + +func TestListOpenFlareAccessLogsFiltersAndSort(t *testing.T) { + ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t) + defer cleanup() + + now := time.Now().UTC() + seedOpenFlareAccessLogs(t, ctx, now) + + query := OpenFlareAccessLogQuery{ + NodeID: "node-a", + Since: now.Add(-10 * time.Minute), + SortBy: "status_code", + SortOrder: "desc", + } + rows, err := ListOpenFlareAccessLogs(ctx, query) + require.NoError(t, err) + require.Len(t, rows, 2) + assert.Equal(t, 404, rows[0].StatusCode) + assert.Equal(t, 200, rows[1].StatusCode) +} + +func TestListOpenFlareAccessLogsMissingTableGraceful(t *testing.T) { + ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t) + defer cleanup() + require.NoError(t, db.DB(ctx).Migrator().DropTable(&OpenFlareAccessLog{})) + + query := OpenFlareAccessLogQuery{Since: time.Now().UTC().Add(-time.Hour)} + rows, err := ListOpenFlareAccessLogs(ctx, query) + require.NoError(t, err) + assert.Empty(t, rows) + + totalRecords, totalIPs, err := CountOpenFlareAccessLogs(ctx, query) + require.NoError(t, err) + assert.Zero(t, totalRecords) + assert.Zero(t, totalIPs) +} + +func TestDeleteOpenFlareAccessLogsBefore(t *testing.T) { + ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t) + defer cleanup() + + now := time.Now().UTC() + seedOpenFlareAccessLogs(t, ctx, now) + + deleted, err := DeleteOpenFlareAccessLogsBefore(ctx, now.Add(-2*time.Minute)) + require.NoError(t, err) + assert.Equal(t, int64(3), deleted) + + totalRecords, _, err := CountOpenFlareAccessLogs(ctx, OpenFlareAccessLogQuery{Since: now.Add(-10 * time.Minute)}) + require.NoError(t, err) + assert.Equal(t, int64(2), totalRecords) +} diff --git a/Wavelet/internal/model/openflare_observability.go b/Wavelet/internal/model/openflare_observability.go index 70df959f..c49963a1 100644 --- a/Wavelet/internal/model/openflare_observability.go +++ b/Wavelet/internal/model/openflare_observability.go @@ -141,6 +141,21 @@ func (OpenFlareNodeObservationOpenresty) TableName() string { return "of_node_obs_openresty" } +// OpenFlareNodeObservationFrpc stores tunnel client frpc observations. +type OpenFlareNodeObservationFrpc 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"` + TunnelStatus string `json:"tunnel_status" gorm:"size:16"` + ConnectedRelaysCount int `json:"connected_relays_count"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName returns the GORM table name. +func (OpenFlareNodeObservationFrpc) TableName() string { + return "of_node_obs_frpc" +} + // OpenFlareNodeObservationFrps stores tunnel relay frps observations. type OpenFlareNodeObservationFrps struct { ID uint `json:"id" gorm:"primaryKey;autoIncrement"` @@ -321,11 +336,6 @@ func ListOpenFlareRequestReportsSince(ctx context.Context, nodeID string, since 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) @@ -423,6 +433,32 @@ func ListOpenFlareNodeObservationOpenresty(ctx context.Context, nodeID string, s return rows, nil } +// ListOpenFlareNodeObservationFrpc returns frpc observations. +func ListOpenFlareNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + query := conn.Model(&OpenFlareNodeObservationFrpc{}).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 []*OpenFlareNodeObservationFrpc + if err := query.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*OpenFlareNodeObservationFrpc{}, 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) @@ -448,53 +484,3 @@ func ListOpenFlareNodeObservationFrps(ctx context.Context, nodeID string, since } 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_waf.go b/Wavelet/internal/model/openflare_waf.go index b9045e31..3d0dd3f1 100644 --- a/Wavelet/internal/model/openflare_waf.go +++ b/Wavelet/internal/model/openflare_waf.go @@ -184,6 +184,22 @@ func ListOpenFlareWAFIPGroups(ctx context.Context) ([]*OpenFlareWAFIPGroup, erro return groups, nil } +// ListOpenFlareWAFIPGroupsByIDs returns IP groups for the given ids. +func ListOpenFlareWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*OpenFlareWAFIPGroup, error) { + if len(ids) == 0 { + return []*OpenFlareWAFIPGroup{}, nil + } + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var groups []*OpenFlareWAFIPGroup + if err = conn.Where("id IN ?", ids).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) diff --git a/docs/changelog/index.md b/docs/changelog/index.md index d06035fb..ed7aba8a 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -27,6 +27,13 @@ sidebar: false - Agent heartbeat 恢复可观测性数据持久化(系统画像、指标快照、流量报表、健康事件等)。 - Wavelet 默认数据库名由 `wavelet` 调整为 `openflare`(PostgreSQL、ClickHouse、SQLite 后备路径同步更新)。 - Wavelet 默认 PostgreSQL `application_name` 由 `wavelet-server` 调整为 `openflare-server`,Redis 键前缀由 `wavelet:` 调整为 `openflare:`。 +- 实装 Agent WAF IP 组同步(heartbeat `waf_ip_groups` 增量下发与 `/api/agent/waf/ip-groups/sync`)。 +- 实装 Pages Agent 部署包下载(`/api/agent/pages/deployments/:deployment_id/package` 二进制响应)。 +- 全局 `OpenFlare-Token` 桥接至 legacy `/api/*` 管理端路由。 +- 实装访问日志单表查询层(列表、折叠、IP 汇总/趋势、地域统计)及 `(node_id, logged_at)` 复合索引。 +- 扩展 Relay/Flared heartbeat 载荷与可观测性持久化(frps 观测、健康事件);新增 `of_node_obs_frpc` 单表。 +- Agent heartbeat 恢复 Geo 自动更新、访问日志地域解析与 90 天保留清理;对齐 config `support_files` 过滤规则。 +- 补全 OAuth 快捷路由(`/api/oauth/github`、`/api/oauth/wechat`、`/api/oauth/wechat/bind`、`/api/oauth/email/bind`)。 ## [v2.3.4] - 2026-06-17