diff --git a/usage-service/internal/collector/auth_snapshot.go b/usage-service/internal/collector/auth_snapshot.go index addf05a62..2fbc3e835 100644 --- a/usage-service/internal/collector/auth_snapshot.go +++ b/usage-service/internal/collector/auth_snapshot.go @@ -39,6 +39,18 @@ func newAuthSnapshotResolver() *authSnapshotResolver { } } +func (r *authSnapshotResolver) clear() { + if r == nil { + return + } + r.mu.Lock() + defer r.mu.Unlock() + r.baseURL = "" + r.managementKey = "" + r.expiresAt = time.Time{} + r.snapshots = nil +} + func (r *authSnapshotResolver) lookup(ctx context.Context, cfg RuntimeConfig, authIndices map[string]struct{}) map[string]authSnapshot { if r == nil || len(authIndices) == 0 { return nil diff --git a/usage-service/internal/collector/collector.go b/usage-service/internal/collector/collector.go index b2c2dcc2e..a96b96dee 100644 --- a/usage-service/internal/collector/collector.go +++ b/usage-service/internal/collector/collector.go @@ -2,6 +2,7 @@ package collector import ( "context" + "encoding/json" "errors" "net" "strings" @@ -385,9 +386,24 @@ func (m *Manager) processItems(ctx context.Context, cfg RuntimeConfig, items []s }) events := make([]usage.Event, 0, len(items)) for _, item := range items { - event, err := usage.NormalizeRaw([]byte(item)) + payload := strings.TrimSpace(item) + if payload == "" { + continue + } + control, enabled := classifyUsageControlPayload(payload) + switch control { + case usageControlSupportRefresh: + continue + case usageControlRefresh: + if enabled && m.snapshotResolver != nil { + m.snapshotResolver.clear() + } + continue + } + + event, err := usage.NormalizeRaw([]byte(payload)) if err != nil { - _ = m.store.AddDeadLetter(ctx, item, err) + _ = m.store.AddDeadLetter(ctx, payload, err) m.setStatus(func(status *Status) { status.DeadLetters++ }) @@ -410,6 +426,30 @@ func (m *Manager) processItems(ctx context.Context, cfg RuntimeConfig, items []s return nil } +type usageControlPayload string + +const ( + usageControlNone usageControlPayload = "" + usageControlSupportRefresh usageControlPayload = "support_refresh" + usageControlRefresh usageControlPayload = "refresh" +) + +func classifyUsageControlPayload(payload string) (usageControlPayload, bool) { + var record map[string]bool + if err := json.Unmarshal([]byte(payload), &record); err != nil || len(record) != 1 { + return usageControlNone, false + } + for key, enabled := range record { + switch key { + case "support_refresh": + return usageControlSupportRefresh, enabled + case "refresh": + return usageControlRefresh, enabled + } + } + return usageControlNone, false +} + func (m *Manager) enrichAccountSnapshots(ctx context.Context, cfg RuntimeConfig, events []usage.Event) { if len(events) == 0 || m.snapshotResolver == nil { return diff --git a/usage-service/internal/collector/collector_test.go b/usage-service/internal/collector/collector_test.go index 80590b339..a6d19e583 100644 --- a/usage-service/internal/collector/collector_test.go +++ b/usage-service/internal/collector/collector_test.go @@ -268,3 +268,66 @@ func TestManagerConsumesSubscribeStream(t *testing.T) { t.Fatalf("total inserted = %d, want 1", status.TotalInserted) } } + +func TestManagerSkipsUsageControlPayloadsAndRefreshesSnapshots(t *testing.T) { + db := newTestStore(t) + manager := NewManager(testConfig(t, "subscribe"), db) + manager.snapshotResolver.baseURL = "http://127.0.0.1:1455" + manager.snapshotResolver.managementKey = "management-key" + manager.snapshotResolver.expiresAt = time.Now().Add(time.Minute) + manager.snapshotResolver.snapshots = map[string]authSnapshot{ + "auth-1": { + Account: "alice@example.com", + Label: "Alice", + FileName: "alice.json", + Provider: "codex", + CapturedAtMS: time.Now().UnixMilli(), + }, + } + + ctx := context.Background() + cfg := RuntimeConfig{ + CPAUpstreamURL: "http://127.0.0.1:1455", + ManagementKey: "management-key", + } + + if err := manager.processItems(ctx, cfg, []string{ + " ", + `{"support_refresh":true}`, + }); err != nil { + t.Fatalf("process support refresh: %v", err) + } + + events, deadLetters, err := db.Counts(ctx) + if err != nil { + t.Fatalf("counts: %v", err) + } + if events != 0 { + t.Fatalf("events = %d, want 0", events) + } + if deadLetters != 0 { + t.Fatalf("dead letters = %d, want 0", deadLetters) + } + if len(manager.snapshotResolver.snapshots) != 1 { + t.Fatalf("support_refresh cleared snapshots, want cache retained") + } + + if err := manager.processItems(ctx, cfg, []string{`{"refresh":true}`}); err != nil { + t.Fatalf("process refresh: %v", err) + } + + events, deadLetters, err = db.Counts(ctx) + if err != nil { + t.Fatalf("counts after refresh: %v", err) + } + if events != 0 { + t.Fatalf("events after refresh = %d, want 0", events) + } + if deadLetters != 0 { + t.Fatalf("dead letters after refresh = %d, want 0", deadLetters) + } + if manager.snapshotResolver.baseURL != "" || manager.snapshotResolver.managementKey != "" || + !manager.snapshotResolver.expiresAt.IsZero() || manager.snapshotResolver.snapshots != nil { + t.Fatalf("refresh did not clear snapshot cache") + } +}