Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions usage-service/internal/collector/auth_snapshot.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
44 changes: 42 additions & 2 deletions usage-service/internal/collector/collector.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package collector

import (
"context"
"encoding/json"
"errors"
"net"
"strings"
Expand Down Expand Up @@ -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++
})
Expand All @@ -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
Expand Down
63 changes: 63 additions & 0 deletions usage-service/internal/collector/collector_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
}
Loading