diff --git a/.squawk.toml b/.squawk.toml index 8fd3371e8..744243838 100644 --- a/.squawk.toml +++ b/.squawk.toml @@ -1,18 +1,15 @@ # squawk — Postgres migration-safety linter for go/internal/store/migrations/. # Gate config for tools/sql-migration-gate/moon.yml. # -# WHY THESE EXCLUSIONS: the sole migration (0001_init.sql) is a single collapsed -# schema-bootstrap that runs ONCE against an empty, pre-live database inside a -# transaction (see the migration runner + advisory lock). squawk's default rule -# set is calibrated for INCREMENTAL migrations against a large LIVE production -# table, where a full-table rewrite or a non-CONCURRENT index build takes a -# blocking lock that stalls live traffic. On an empty pre-live DB that hazard -# does not exist, so the rules below are accepted here. They stay OFF only for -# this bootstrap posture — a future incremental migration against live data -# would want them back (revisit this list when the first post-live migration -# lands). The gate stays GREEN on the current file and RED on genuinely unsafe -# NEW DDL (e.g. adding a NOT NULL column with no default), which these -# exclusions do not silence. +# WHY THESE EXCLUSIONS: every migration so far (0001_init.sql, then the +# append-only files after it) creates tables that are empty when it runs, inside +# one transaction (see the migration runner + advisory lock). squawk's default +# rule set is calibrated for changes to a large LIVE table, where a full-table +# rewrite or a non-CONCURRENT index build takes a blocking lock that stalls +# traffic. On a new empty table that hazard does not exist, so the rules below +# are accepted. Revisit this list when a migration first alters or indexes a +# populated table. The gate stays RED on genuinely unsafe NEW DDL (e.g. adding +# a NOT NULL column with no default), which these exclusions do not silence. pg_version = "16.0" diff --git a/docs/self-host.md b/docs/self-host.md index eeca7129c..ca09ed6af 100644 --- a/docs/self-host.md +++ b/docs/self-host.md @@ -169,6 +169,12 @@ With `--database-external` the stack only connects to the `--database` DSN you name; it never starts, stops, or owns that instance's lifecycle. The flag is the opt-out switch and `--database` (or `$COMPASS_DATABASE_DSN`) carries the DSN. +The server keeps each raw token-usage event (one row per upstream model call) +for 90 days by default, and deletes older events once a day. Set the window with +`--usage-event-retention` (or `$COMPASS_USAGE_EVENT_RETENTION`) as a Go duration +such as `720h`; `0` keeps every event. The hourly and daily usage totals built +from those events are always kept. + ## Secrets Compass keeps its secret *values* in your configured `secretspec` provider, not diff --git a/go/cmd/compass-server/main.go b/go/cmd/compass-server/main.go index 7bb8ced20..e83caaa86 100644 --- a/go/cmd/compass-server/main.go +++ b/go/cmd/compass-server/main.go @@ -21,6 +21,7 @@ import ( "time" "github.com/RigelBuild/compass/go/internal/otel" + "github.com/RigelBuild/compass/go/internal/usage" "github.com/RigelBuild/compass/go/server" ) @@ -220,6 +221,12 @@ func buildServeConfig(args []string) (server.ServeConfig, bool, error) { return server.ServeConfig{}, false, err } + usageRetention, err := resolveUsageEventRetention( + firstNonEmpty(*f.usageRetention, os.Getenv("COMPASS_USAGE_EVENT_RETENTION"))) + if err != nil { + return server.ServeConfig{}, false, err + } + return server.ServeConfig{ SocketPath: socketPath, Version: version, @@ -239,6 +246,7 @@ func buildServeConfig(args []string) (server.ServeConfig, bool, error) { // read one source, so no --otel-endpoint flag. Empty = tracing off. OtelEndpoint: os.Getenv("OTEL_EXPORTER_OTLP_ENDPOINT"), TranscriptSafetyValveCapBytes: positiveCap(*f.transcriptSafetyValveCapBytes, os.Getenv("COMPASS_TRANSCRIPT_SAFETY_VALVE_CAP_BYTES")), + UsageEventRetention: usageRetention, }, false, nil } @@ -281,6 +289,7 @@ type serveFlags struct { adminHandle *string corsAllowedOrigin *string publicURL *string + usageRetention *string } // registerServeFlags declares the core compass-server flags on the given FlagSet @@ -350,6 +359,11 @@ func registerServeFlags(fs *flag.FlagSet) serveFlags { "Linear webhooks must set it."), transcriptSafetyValveCapBytes: fs.Int("transcript-safety-valve-cap-bytes", 0, "Hot-tail safety-valve cap in bytes. Defaults to $COMPASS_TRANSCRIPT_SAFETY_VALVE_CAP_BYTES."), + usageRetention: fs.String("usage-event-retention", "", + "How long raw token-usage events are kept before the daily prune "+ + "deletes them, as a Go duration (e.g. 720h). The hourly and daily "+ + "usage totals are kept. Falls back to $COMPASS_USAGE_EVENT_RETENTION, "+ + "then 2160h (90 days). 0 disables the prune."), } } @@ -392,6 +406,22 @@ func resolveNetworkDoor(listen, tlsCert, tlsKey string) (string, *server.TLSConf } } +// resolveUsageEventRetention parses the retention window (flag, then env). Empty +// keeps the default; 0 is kept, because it is the operator's opt-out. +func resolveUsageEventRetention(v string) (time.Duration, error) { + if v == "" { + return usage.DefaultRetention, nil + } + d, err := time.ParseDuration(v) + if err != nil { + return 0, fmt.Errorf("invalid --usage-event-retention %q: %w", v, err) + } + if d < 0 { + return 0, fmt.Errorf("invalid --usage-event-retention %q: it must not be negative; 0 disables the prune", v) + } + return d, nil +} + // forgeFlags holds the RIG-1810/RIG-2883 forge CLI flag pointers, registered as // a group so run() stays short (they mirror the S3 flag set's precedence). type forgeFlags struct { diff --git a/go/cmd/compass-server/main_test.go b/go/cmd/compass-server/main_test.go index cd4eaaf25..e15d3c5df 100644 --- a/go/cmd/compass-server/main_test.go +++ b/go/cmd/compass-server/main_test.go @@ -12,6 +12,7 @@ import ( "flag" "strings" "testing" + "time" "github.com/RigelBuild/compass/go/server" ) @@ -411,3 +412,45 @@ func TestBuildServeConfigBadFlagIsUsageError(t *testing.T) { t.Fatalf("a bad flag must not read as ErrHelp (that is a clean help exit): %v", err) } } + +// TestBuildServeConfigUsageEventRetention pins the retention window's +// flag-then-env precedence, its 90-day default, 0 as the opt-out, and bad input. +func TestBuildServeConfigUsageEventRetention(t *testing.T) { + for _, tc := range []struct { + name, flag, env string + want time.Duration + wantErr bool + }{ + {name: "unset_keeps_90_days", want: 90 * 24 * time.Hour}, + {name: "flag_sets_the_window", flag: "720h", want: 720 * time.Hour}, + {name: "env_is_the_fallback", env: "48h", want: 48 * time.Hour}, + {name: "flag_beats_env", flag: "24h", env: "48h", want: 24 * time.Hour}, + {name: "flag_0_disables_the_sweep", flag: "0", env: "48h", want: 0}, + {name: "env_0_disables_the_sweep", env: "0", want: 0}, + {name: "day_unit_is_rejected", flag: "90d", wantErr: true}, + {name: "bad_env_is_rejected", env: "soon", wantErr: true}, + {name: "negative_is_rejected", flag: "-24h", wantErr: true}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Setenv("COMPASS_USAGE_EVENT_RETENTION", tc.env) + t.Setenv("COMPASS_NATS_URL", "nats://127.0.0.1:4222") // required since the event fabric landed + args := []string{"--database", "postgres://x/db", "--socket", "/tmp/x.sock"} + if tc.flag != "" { + args = append(args, "--usage-event-retention", tc.flag) + } + cfg, _, err := buildServeConfig(args) + if tc.wantErr { + if err == nil || !strings.Contains(err.Error(), "--usage-event-retention") { + t.Fatalf("buildServeConfig = %v, want an error naming --usage-event-retention", err) + } + return + } + if err != nil { + t.Fatalf("buildServeConfig = %v, want nil", err) + } + if cfg.UsageEventRetention != tc.want { + t.Errorf("UsageEventRetention = %v, want %v", cfg.UsageEventRetention, tc.want) + } + }) + } +} diff --git a/go/internal/store/db/models.go b/go/internal/store/db/models.go index ead7b0858..c8372bb83 100644 --- a/go/internal/store/db/models.go +++ b/go/internal/store/db/models.go @@ -330,6 +330,69 @@ type Token struct { RevokedAt pgtype.Timestamptz } +type TokenUsageEvent struct { + TenantID string + ID string + OccurredAt pgtype.Timestamptz + AgentAccountID string + OwnerUserID string + SessionID string + RequestID string + Provider string + Model string + CredentialID string + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheWriteTokens int64 + TotalTokens int64 + CostMicroUsd int64 + RateVersion string + Outcome string + CreatedAt pgtype.Timestamptz +} + +type TokenUsagePruneHorizon struct { + Singleton bool + Horizon pgtype.Timestamptz + CreatedAt pgtype.Timestamptz + UpdatedAt pgtype.Timestamptz +} + +type TokenUsageRollupsDaily struct { + TenantID string + BucketStart pgtype.Timestamptz + OwnerUserID string + AgentAccountID string + Provider string + Model string + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheWriteTokens int64 + TotalTokens int64 + CostMicroUsd int64 + CreatedAt pgtype.Timestamptz + UpdatedAt pgtype.Timestamptz +} + +type TokenUsageRollupsHourly struct { + TenantID string + BucketStart pgtype.Timestamptz + OwnerUserID string + AgentAccountID string + Provider string + Model string + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheWriteTokens int64 + TotalTokens int64 + CostMicroUsd int64 + CreatedAt pgtype.Timestamptz + UpdatedAt pgtype.Timestamptz +} + type Topic struct { ID string ChannelID string diff --git a/go/internal/store/db/querier.go b/go/internal/store/db/querier.go index f3a8473be..fdd2a2d59 100644 --- a/go/internal/store/db/querier.go +++ b/go/internal/store/db/querier.go @@ -26,6 +26,9 @@ type Querier interface { // (someone else advanced first), NOT an error — the wrapper reports it as // advanced=false. AdvanceForgeDeliveredRevisionCAS(ctx context.Context, arg AdvanceForgeDeliveredRevisionCASParams) (int64, error) + // AdvanceTokenUsagePruneHorizon commits before the prune deletes anything, and + // waits for a rebuild that holds the old horizon. It only moves forward. + AdvanceTokenUsagePruneHorizon(ctx context.Context, cutoff pgtype.Timestamptz) (int64, error) AgentForContainer(ctx context.Context, containerName string) (string, error) // Presence-component read queries (sqlc adoption T4, RIG-3034). These replace the // const-hoisted SQL in internal/store/presence_reads.go (it was never in the @@ -50,6 +53,10 @@ type Querier interface { // Feeds isAgentWorkspaceVisible: membership on the agent's home channel. AgentWorkspaceVisible(ctx context.Context, arg AgentWorkspaceVisibleParams) (bool, error) AgentsByOwner(ctx context.Context, ownerUserID string) ([]AgentsByOwnerRow, error) + // AppendTokenUsageEvents inserts a batch and adds only the events it inserted + // to both rollups, so a replayed id counts once. The arrays hold one element + // per event, and no id repeats. + AppendTokenUsageEvents(ctx context.Context, arg AppendTokenUsageEventsParams) error AuthoredArtifactByCoordinate(ctx context.Context, arg AuthoredArtifactByCoordinateParams) (ForgeAuthoredArtifact, error) AuthoredArtifactByRequestID(ctx context.Context, arg AuthoredArtifactByRequestIDParams) (ForgeAuthoredArtifact, error) // Agent-transcript queries (sqlc adoption T5, RIG-3034). These replace the @@ -130,6 +137,12 @@ type Querier interface { // A DELETE ... RETURNING takes no ORDER BY, so the Store method sorts the // returned slice by session id to keep a sweep pass deterministic and diffable. DeleteSessionBindingsForRunner(ctx context.Context, runnerID string) ([]DeleteSessionBindingsForRunnerRow, error) + // DeleteTokenUsageEventsBefore deletes old events of the tx's tenant only, + // because row-level security scopes it. The prune runs it once per tenant. + DeleteTokenUsageEventsBefore(ctx context.Context, cutoff pgtype.Timestamptz) (int64, error) + // DeleteTokenUsageRollupsFrom drops both rollups from the horizon on. Older + // rollups can hold pruned events, so they stay. + DeleteTokenUsageRollupsFrom(ctx context.Context, horizon pgtype.Timestamptz) error DeleteTopic(ctx context.Context, id string) error // Agent-forge-subscription / artifact-cursor queries (sqlc adoption T6, // RIG-3034). These replace the inline SQL literals in @@ -332,6 +345,9 @@ type Querier interface { ListForgeNotifyTargets(ctx context.Context, arg ListForgeNotifyTargetsParams) ([]ListForgeNotifyTargetsRow, error) ListIssues(ctx context.Context) ([]ListIssuesRow, error) ListMessages(ctx context.Context, arg ListMessagesParams) ([]ListMessagesRow, error) + // ListTenantIDs lists every tenant. tenants has no row-level security, so the + // app role sees them all without the BYPASSRLS system role. + ListTenantIDs(ctx context.Context) ([]string, error) // Topic-domain queries (sqlc adoption T4, RIG-3034). These replace the inline // SQL literals in internal/store/topics.go; the hand-written Store methods keep // their signatures, the UpdateTopic tx orchestration, the rename/merge resolution @@ -429,6 +445,11 @@ type Querier interface { // the displaced value is a READ, so the PK cannot cover it. This lock is taken // before the read, so it does. LockSessionBindingAccount(ctx context.Context, arg LockSessionBindingAccountParams) error + // Token-usage queries (Plane A). The tenant tx sets the GUC that RLS reads and + // that tenant_id defaults to, so no statement here names a tenant. + // LockTokenUsage serializes one tenant's appends and rebuilds. A rebuild beside + // an append would count the append's events twice or lose them. + LockTokenUsage(ctx context.Context) error MarkMentionsRouted(ctx context.Context, arg MarkMentionsRoutedParams) error MergeTopicLastSeq(ctx context.Context, arg MergeTopicLastSeqParams) error MessageByID(ctx context.Context, id string) (MessageByIDRow, error) @@ -517,6 +538,9 @@ type Querier interface { ResolveVisibleGlobalHandles(ctx context.Context, arg ResolveVisibleGlobalHandlesParams) ([]ResolveVisibleGlobalHandlesRow, error) ReviveTopic(ctx context.Context, id string) error RevokeToken(ctx context.Context, hash []byte) (int64, error) + // RollUpTokenUsageFrom rebuilds both rollups from the events at or after the + // horizon. + RollUpTokenUsageFrom(ctx context.Context, horizon pgtype.Timestamptz) error SafetyValveSegments(ctx context.Context, arg SafetyValveSegmentsParams) ([]SafetyValveSegmentsRow, error) // Scaffold-only query proving sqlc generation works end to end (T1). // @@ -594,6 +618,13 @@ type Querier interface { SweepChannels(ctx context.Context, accountID string) ([]string, error) TenantIDBySlug(ctx context.Context, slug string) (string, error) TokenHashExists(ctx context.Context, hash []byte) (bool, error) + // TokenUsagePruneHorizon reads the prune horizon and holds it until the tx + // ends, so a prune cannot delete the events a rebuild is about to count. + TokenUsagePruneHorizon(ctx context.Context) (pgtype.Timestamptz, error) + // TokenUsageSeries sums one granularity's rollup rows per bucket. granularity + // is the usage.Granularity value: 1 is hourly, 2 is daily. An empty agent list + // or provider means no filter. + TokenUsageSeries(ctx context.Context, arg TokenUsageSeriesParams) ([]TokenUsageSeriesRow, error) // Authorization-probe queries (sqlc adoption T6, RIG-3034). These replace the // inline SQL literals in internal/store/authz.go; the hand-written helpers keep // their signatures and the not-found/forbidden merge, wrapping these EXISTS diff --git a/go/internal/store/db/tenant.sql.go b/go/internal/store/db/tenant.sql.go index 9d10d1456..d1dc2e36e 100644 --- a/go/internal/store/db/tenant.sql.go +++ b/go/internal/store/db/tenant.sql.go @@ -35,6 +35,32 @@ func (q *Queries) InsertTenant(ctx context.Context, arg InsertTenantParams) erro return err } +const listTenantIDs = `-- name: ListTenantIDs :many +SELECT id FROM tenants ORDER BY id +` + +// ListTenantIDs lists every tenant. tenants has no row-level security, so the +// app role sees them all without the BYPASSRLS system role. +func (q *Queries) ListTenantIDs(ctx context.Context) ([]string, error) { + rows, err := q.db.Query(ctx, listTenantIDs) + if err != nil { + return nil, err + } + defer rows.Close() + var items []string + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return nil, err + } + items = append(items, id) + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + const tenantIDBySlug = `-- name: TenantIDBySlug :one SELECT id FROM tenants WHERE slug = $1 ` diff --git a/go/internal/store/db/token_usage.sql.go b/go/internal/store/db/token_usage.sql.go new file mode 100644 index 000000000..25efb5798 --- /dev/null +++ b/go/internal/store/db/token_usage.sql.go @@ -0,0 +1,292 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: token_usage.sql + +package db + +import ( + "context" + + "github.com/jackc/pgx/v5/pgtype" +) + +const advanceTokenUsagePruneHorizon = `-- name: AdvanceTokenUsagePruneHorizon :execrows +UPDATE token_usage_prune_horizon SET horizon = GREATEST(horizon, $1::timestamptz) +` + +// AdvanceTokenUsagePruneHorizon commits before the prune deletes anything, and +// waits for a rebuild that holds the old horizon. It only moves forward. +func (q *Queries) AdvanceTokenUsagePruneHorizon(ctx context.Context, cutoff pgtype.Timestamptz) (int64, error) { + result, err := q.db.Exec(ctx, advanceTokenUsagePruneHorizon, cutoff) + if err != nil { + return 0, err + } + return result.RowsAffected(), nil +} + +const appendTokenUsageEvents = `-- name: AppendTokenUsageEvents :exec +WITH inserted AS ( + INSERT INTO token_usage_events ( + id, occurred_at, agent_account_id, owner_user_id, session_id, request_id, provider, + model, credential_id, input_tokens, output_tokens, cache_read_tokens, + cache_write_tokens, total_tokens, cost_micro_usd, rate_version, outcome) + SELECT unnest($1::text[]), unnest($2::timestamptz[]), + unnest($3::text[]), unnest($4::text[]), + unnest($5::text[]), unnest($6::text[]), + unnest($7::text[]), unnest($8::text[]), + unnest($9::text[]), unnest($10::bigint[]), + unnest($11::bigint[]), unnest($12::bigint[]), + unnest($13::bigint[]), unnest($14::bigint[]), + unnest($15::bigint[]), unnest($16::text[]), + unnest($17::text[]) + ON CONFLICT (tenant_id, id) DO NOTHING + RETURNING occurred_at, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd +), hourly AS ( + INSERT INTO token_usage_rollups_hourly ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) + SELECT date_trunc('hour', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM inserted + GROUP BY 1, 2, 3, 4, 5 + ON CONFLICT (tenant_id, bucket_start, owner_user_id, agent_account_id, provider, model) + DO UPDATE SET + input_tokens = token_usage_rollups_hourly.input_tokens + EXCLUDED.input_tokens, + output_tokens = token_usage_rollups_hourly.output_tokens + EXCLUDED.output_tokens, + cache_read_tokens = token_usage_rollups_hourly.cache_read_tokens + EXCLUDED.cache_read_tokens, + cache_write_tokens = token_usage_rollups_hourly.cache_write_tokens + EXCLUDED.cache_write_tokens, + total_tokens = token_usage_rollups_hourly.total_tokens + EXCLUDED.total_tokens, + cost_micro_usd = token_usage_rollups_hourly.cost_micro_usd + EXCLUDED.cost_micro_usd +) +INSERT INTO token_usage_rollups_daily ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) +SELECT date_trunc('day', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM inserted + GROUP BY 1, 2, 3, 4, 5 +ON CONFLICT (tenant_id, bucket_start, owner_user_id, agent_account_id, provider, model) +DO UPDATE SET + input_tokens = token_usage_rollups_daily.input_tokens + EXCLUDED.input_tokens, + output_tokens = token_usage_rollups_daily.output_tokens + EXCLUDED.output_tokens, + cache_read_tokens = token_usage_rollups_daily.cache_read_tokens + EXCLUDED.cache_read_tokens, + cache_write_tokens = token_usage_rollups_daily.cache_write_tokens + EXCLUDED.cache_write_tokens, + total_tokens = token_usage_rollups_daily.total_tokens + EXCLUDED.total_tokens, + cost_micro_usd = token_usage_rollups_daily.cost_micro_usd + EXCLUDED.cost_micro_usd +` + +type AppendTokenUsageEventsParams struct { + Ids []string + OccurredAt []pgtype.Timestamptz + AgentAccountIds []string + OwnerUserIds []string + SessionIds []string + RequestIds []string + Providers []string + Models []string + CredentialIds []string + InputTokens []int64 + OutputTokens []int64 + CacheReadTokens []int64 + CacheWriteTokens []int64 + TotalTokens []int64 + CostMicroUsd []int64 + RateVersions []string + Outcomes []string +} + +// AppendTokenUsageEvents inserts a batch and adds only the events it inserted +// to both rollups, so a replayed id counts once. The arrays hold one element +// per event, and no id repeats. +func (q *Queries) AppendTokenUsageEvents(ctx context.Context, arg AppendTokenUsageEventsParams) error { + _, err := q.db.Exec(ctx, appendTokenUsageEvents, + arg.Ids, + arg.OccurredAt, + arg.AgentAccountIds, + arg.OwnerUserIds, + arg.SessionIds, + arg.RequestIds, + arg.Providers, + arg.Models, + arg.CredentialIds, + arg.InputTokens, + arg.OutputTokens, + arg.CacheReadTokens, + arg.CacheWriteTokens, + arg.TotalTokens, + arg.CostMicroUsd, + arg.RateVersions, + arg.Outcomes, + ) + return err +} + +const deleteTokenUsageEventsBefore = `-- name: DeleteTokenUsageEventsBefore :execrows +DELETE FROM token_usage_events WHERE occurred_at < $1::timestamptz +` + +// DeleteTokenUsageEventsBefore deletes old events of the tx's tenant only, +// because row-level security scopes it. The prune runs it once per tenant. +func (q *Queries) DeleteTokenUsageEventsBefore(ctx context.Context, cutoff pgtype.Timestamptz) (int64, error) { + result, err := q.db.Exec(ctx, deleteTokenUsageEventsBefore, cutoff) + if err != nil { + return 0, err + } + return result.RowsAffected(), nil +} + +const deleteTokenUsageRollupsFrom = `-- name: DeleteTokenUsageRollupsFrom :exec +WITH hourly AS ( + DELETE FROM token_usage_rollups_hourly WHERE bucket_start >= $1::timestamptz +) +DELETE FROM token_usage_rollups_daily WHERE bucket_start >= $1::timestamptz +` + +// DeleteTokenUsageRollupsFrom drops both rollups from the horizon on. Older +// rollups can hold pruned events, so they stay. +func (q *Queries) DeleteTokenUsageRollupsFrom(ctx context.Context, horizon pgtype.Timestamptz) error { + _, err := q.db.Exec(ctx, deleteTokenUsageRollupsFrom, horizon) + return err +} + +const lockTokenUsage = `-- name: LockTokenUsage :exec + +SELECT pg_advisory_xact_lock(hashtext('token_usage:' || current_setting('compass.tenant_id', TRUE))) +` + +// Token-usage queries (Plane A). The tenant tx sets the GUC that RLS reads and +// that tenant_id defaults to, so no statement here names a tenant. +// LockTokenUsage serializes one tenant's appends and rebuilds. A rebuild beside +// an append would count the append's events twice or lose them. +func (q *Queries) LockTokenUsage(ctx context.Context) error { + _, err := q.db.Exec(ctx, lockTokenUsage) + return err +} + +const rollUpTokenUsageFrom = `-- name: RollUpTokenUsageFrom :exec +WITH hourly AS ( + INSERT INTO token_usage_rollups_hourly ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) + SELECT date_trunc('hour', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM token_usage_events + WHERE occurred_at >= $1::timestamptz + GROUP BY 1, 2, 3, 4, 5 +) +INSERT INTO token_usage_rollups_daily ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) +SELECT date_trunc('day', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM token_usage_events + WHERE occurred_at >= $1::timestamptz + GROUP BY 1, 2, 3, 4, 5 +` + +// RollUpTokenUsageFrom rebuilds both rollups from the events at or after the +// horizon. +func (q *Queries) RollUpTokenUsageFrom(ctx context.Context, horizon pgtype.Timestamptz) error { + _, err := q.db.Exec(ctx, rollUpTokenUsageFrom, horizon) + return err +} + +const tokenUsagePruneHorizon = `-- name: TokenUsagePruneHorizon :one +SELECT horizon FROM token_usage_prune_horizon FOR SHARE +` + +// TokenUsagePruneHorizon reads the prune horizon and holds it until the tx +// ends, so a prune cannot delete the events a rebuild is about to count. +func (q *Queries) TokenUsagePruneHorizon(ctx context.Context) (pgtype.Timestamptz, error) { + row := q.db.QueryRow(ctx, tokenUsagePruneHorizon) + var horizon pgtype.Timestamptz + err := row.Scan(&horizon) + return horizon, err +} + +const tokenUsageSeries = `-- name: TokenUsageSeries :many +SELECT r.bucket_start::timestamptz AS bucket_start, + sum(r.input_tokens)::bigint AS input_tokens, + sum(r.output_tokens)::bigint AS output_tokens, + sum(r.cache_read_tokens)::bigint AS cache_read_tokens, + sum(r.cache_write_tokens)::bigint AS cache_write_tokens, + sum(r.total_tokens)::bigint AS total_tokens, + sum(r.cost_micro_usd)::bigint AS cost_micro_usd + FROM (SELECT bucket_start, agent_account_id, provider, input_tokens, output_tokens, + cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd + FROM token_usage_rollups_hourly + WHERE $1::integer = 1 + UNION ALL + SELECT bucket_start, agent_account_id, provider, input_tokens, output_tokens, + cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd + FROM token_usage_rollups_daily + WHERE $1::integer = 2) AS r + WHERE r.bucket_start >= $2::timestamptz + AND r.bucket_start < $3::timestamptz + AND (coalesce(cardinality($4::text[]), 0) = 0 + OR r.agent_account_id = ANY ($4::text[])) + AND ($5::text = '' OR r.provider = $5::text) + GROUP BY r.bucket_start + ORDER BY r.bucket_start +` + +type TokenUsageSeriesParams struct { + Granularity int32 + StartAt pgtype.Timestamptz + EndAt pgtype.Timestamptz + AgentAccountIds []string + Provider string +} + +type TokenUsageSeriesRow struct { + BucketStart pgtype.Timestamptz + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheWriteTokens int64 + TotalTokens int64 + CostMicroUsd int64 +} + +// TokenUsageSeries sums one granularity's rollup rows per bucket. granularity +// is the usage.Granularity value: 1 is hourly, 2 is daily. An empty agent list +// or provider means no filter. +func (q *Queries) TokenUsageSeries(ctx context.Context, arg TokenUsageSeriesParams) ([]TokenUsageSeriesRow, error) { + rows, err := q.db.Query(ctx, tokenUsageSeries, + arg.Granularity, + arg.StartAt, + arg.EndAt, + arg.AgentAccountIds, + arg.Provider, + ) + if err != nil { + return nil, err + } + defer rows.Close() + var items []TokenUsageSeriesRow + for rows.Next() { + var i TokenUsageSeriesRow + if err := rows.Scan( + &i.BucketStart, + &i.InputTokens, + &i.OutputTokens, + &i.CacheReadTokens, + &i.CacheWriteTokens, + &i.TotalTokens, + &i.CostMicroUsd, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} diff --git a/go/internal/store/migrate_upgrade_pgtest_test.go b/go/internal/store/migrate_upgrade_pgtest_test.go new file mode 100644 index 000000000..4688115fe --- /dev/null +++ b/go/internal/store/migrate_upgrade_pgtest_test.go @@ -0,0 +1,168 @@ +//go:build pgtest + +package store + +import ( + "context" + "testing" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/RigelBuild/compass/go/internal/pgtest" + "github.com/RigelBuild/compass/go/internal/store/db" +) + +// TestOpenUpgradesV1DatabaseToTokenUsage builds a database that stopped at v1, +// as a deployed one would, and proves Open migrates it forward. A fresh-schema +// test cannot catch a later migration that only works on an empty database. +func TestOpenUpgradesV1DatabaseToTokenUsage(t *testing.T) { + ctx := t.Context() + dsn := pgtest.RequireDSN(t) + const tenant TenantID = "upgrade-tenant" + + applyV1Only(t, dsn) + // The tenant predates the upgrade; the token-usage foreign keys must accept it. + seed, err := pgxpool.New(ctx, dsn) + if err != nil { + t.Fatalf("connect to seed: %v", err) + } + err = db.New(seed).InsertTenant(ctx, db.InsertTenantParams{ + ID: string(tenant), Slug: string(tenant), DisplayName: string(tenant), + CreatedAtUnixMs: time.Now().UnixMilli(), + }) + seed.Close() + if err != nil { + t.Fatalf("seed tenant at v1: %v", err) + } + + s := openStore(t, dsn) + + var version int + if err := s.pool.QueryRow(ctx, "SELECT max(version) FROM schema_migrations").Scan(&version); err != nil { + t.Fatalf("read schema version: %v", err) + } + migs, err := loadMigrations() + if err != nil { + t.Fatalf("load migrations: %v", err) + } + if want := migs[len(migs)-1].version; version != want { + t.Fatalf("schema version after upgrade = %d, want %d", version, want) + } + + for _, tbl := range []string{"token_usage_events", "token_usage_rollups_hourly", "token_usage_rollups_daily"} { + var enabled, forced bool + err := s.pool.QueryRow(ctx, + `SELECT relrowsecurity, relforcerowsecurity FROM pg_class + WHERE oid = to_regclass(format('%I.%I', current_schema(), $1::text))`, tbl, + ).Scan(&enabled, &forced) + if err != nil { + t.Fatalf("%s: read RLS flags (table missing?): %v", tbl, err) + } + if !enabled || !forced { + t.Errorf("%s: relrowsecurity=%t relforcerowsecurity=%t, want both true", tbl, enabled, forced) + } + } + + var horizon pgtype.Timestamptz + if err := s.pool.QueryRow(ctx, "SELECT horizon FROM token_usage_prune_horizon").Scan(&horizon); err != nil { + t.Fatalf("read prune horizon row: %v", err) + } + if horizon.InfinityModifier != pgtype.NegativeInfinity { + t.Errorf("prune horizon = %+v, want -infinity", horizon) + } + + at := time.Date(2026, time.March, 10, 5, 30, 0, 0, time.UTC) + tctx := WithTenant(ctx, tenant) + err = s.WithTx(tctx, func(tx pgx.Tx) error { + return db.New(tx).AppendTokenUsageEvents(tctx, db.AppendTokenUsageEventsParams{ + Ids: []string{"ev-1"}, + OccurredAt: []pgtype.Timestamptz{{Time: at, Valid: true}}, + AgentAccountIds: []string{"agent-1"}, + OwnerUserIds: []string{"user-1"}, + SessionIds: []string{"session-1"}, + RequestIds: []string{"request-1"}, + Providers: []string{"anthropic"}, + Models: []string{"model-1"}, + CredentialIds: []string{"cred-1"}, + InputTokens: []int64{7}, + OutputTokens: []int64{3}, + CacheReadTokens: []int64{0}, + CacheWriteTokens: []int64{0}, + TotalTokens: []int64{10}, + CostMicroUsd: []int64{42}, + RateVersions: []string{"v1"}, + Outcomes: []string{"ok"}, + }) + }) + if err != nil { + t.Fatalf("append token usage after upgrade: %v", err) + } + + daily := db.TokenUsageSeriesParams{ + Granularity: 2, + StartAt: pgtype.Timestamptz{InfinityModifier: pgtype.NegativeInfinity, Valid: true}, + EndAt: pgtype.Timestamptz{InfinityModifier: pgtype.Infinity, Valid: true}, + } + if got := readSeries(t, s, tctx, daily); len(got) != 1 || got[0].TotalTokens != 10 || got[0].CostMicroUsd != 42 { + t.Errorf("upgraded tenant daily series = %+v, want one bucket of 10 tokens / 42 micro-USD", got) + } + // The bootstrap tenant Open seeds must not see another tenant's usage. + if got := readSeries(t, s, WithTenant(ctx, s.bootstrapTenantID), daily); len(got) != 0 { + t.Errorf("bootstrap tenant daily series = %+v, want none (RLS leak)", got) + } +} + +// applyV1Only migrates the empty schema at dsn to v1 through the runner's own +// steps, leaving every later migration pending for Open. +func applyV1Only(t *testing.T, dsn string) { + t.Helper() + ctx := t.Context() + migs, err := loadMigrations() + if err != nil { + t.Fatalf("load migrations: %v", err) + } + pool, err := pgxpool.New(ctx, dsn) + if err != nil { + t.Fatalf("connect for v1: %v", err) + } + defer pool.Close() + conn, err := pool.Acquire(ctx) + if err != nil { + t.Fatalf("acquire for v1: %v", err) + } + defer conn.Release() + // 0001 edits cluster-global roles, so it must hold the same lock migrate + // holds, or a parallel package's Open can race it. + if _, err := conn.Exec(ctx, "SELECT pg_advisory_lock($1)", migrationLockKey); err != nil { + t.Fatalf("acquire migration lock: %v", err) + } + defer func() { + if _, err := conn.Exec(context.WithoutCancel(ctx), "SELECT pg_advisory_unlock($1)", migrationLockKey); err != nil { + t.Errorf("release migration lock: %v", err) + } + }() + if err := ensureMigrationsTable(ctx, conn); err != nil { + t.Fatal(err) + } + if err := applyMigration(ctx, conn, migs[0]); err != nil { + t.Fatal(err) + } +} + +// readSeries reads one tenant's token-usage series in a tenant tx. +func readSeries(t *testing.T, s *Store, ctx context.Context, p db.TokenUsageSeriesParams) []db.TokenUsageSeriesRow { + t.Helper() + var rows []db.TokenUsageSeriesRow + err := s.WithTx(ctx, func(tx pgx.Tx) error { + var err error + rows, err = db.New(tx).TokenUsageSeries(ctx, p) + return err + }) + if err != nil { + t.Fatalf("read token usage series: %v", err) + } + return rows +} diff --git a/go/internal/store/migrations/0003_token_usage.sql b/go/internal/store/migrations/0003_token_usage.sql new file mode 100644 index 000000000..80b457533 --- /dev/null +++ b/go/internal/store/migrations/0003_token_usage.sql @@ -0,0 +1,142 @@ +-- 0003_token_usage: the Plane-A token-usage store — the raw event log, its +-- hourly and daily rollups, and the global prune horizon. +-- +-- Migrations are append-only: 0001_init.sql is frozen, so this file cannot add +-- the new tables to 0001's tenant_tables / updated_at_tables arrays. It repeats +-- those loops' bodies for its own tables instead, with the same policy text, +-- grants, and trigger. 0001 already created compass_app, compass_system, and +-- set_updated_at(). + +-- token_usage_events: the append-only raw log of upstream model calls the LLM +-- gateway reports, idempotent on the server-assigned id. Retention deletes old +-- rows; the rollups keep their sums. No FK to accounts: a log row outlives its +-- account. +CREATE TABLE token_usage_events ( + tenant_id TEXT NOT NULL DEFAULT current_setting('compass.tenant_id', TRUE) REFERENCES tenants (id) ON DELETE RESTRICT, + id TEXT NOT NULL, + occurred_at TIMESTAMPTZ NOT NULL, + agent_account_id TEXT NOT NULL, + owner_user_id TEXT NOT NULL, + session_id TEXT NOT NULL, + request_id TEXT NOT NULL, + provider TEXT NOT NULL, + model TEXT NOT NULL, + credential_id TEXT NOT NULL, + input_tokens BIGINT NOT NULL, + output_tokens BIGINT NOT NULL, + cache_read_tokens BIGINT NOT NULL, + cache_write_tokens BIGINT NOT NULL, + total_tokens BIGINT NOT NULL, + cost_micro_usd BIGINT NOT NULL, + rate_version TEXT NOT NULL, + outcome TEXT NOT NULL CHECK (outcome IN ('ok', 'error', 'aborted')), + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (tenant_id, id) +); + +-- The rebuild re-rolls one tenant's events from the prune horizon on. +CREATE INDEX token_usage_events_occurred_at_idx ON token_usage_events (tenant_id, occurred_at); + +-- token_usage_rollups_hourly / _daily: the per-bucket sums of the events, keyed +-- like the in-memory reference. bucket_start is the UTC-aligned bucket start. +-- Rows outlive the events they sum, so a prune never touches them. bucket_start +-- follows tenant_id in the key because the series read and the rebuild range +-- over it. +CREATE TABLE token_usage_rollups_hourly ( + tenant_id TEXT NOT NULL DEFAULT current_setting('compass.tenant_id', TRUE) REFERENCES tenants (id) ON DELETE RESTRICT, + bucket_start TIMESTAMPTZ NOT NULL, + owner_user_id TEXT NOT NULL, + agent_account_id TEXT NOT NULL, + provider TEXT NOT NULL, + model TEXT NOT NULL, + input_tokens BIGINT NOT NULL, + output_tokens BIGINT NOT NULL, + cache_read_tokens BIGINT NOT NULL, + cache_write_tokens BIGINT NOT NULL, + total_tokens BIGINT NOT NULL, + cost_micro_usd BIGINT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (tenant_id, bucket_start, owner_user_id, agent_account_id, provider, model) +); + +CREATE TABLE token_usage_rollups_daily ( + tenant_id TEXT NOT NULL DEFAULT current_setting('compass.tenant_id', TRUE) REFERENCES tenants (id) ON DELETE RESTRICT, + bucket_start TIMESTAMPTZ NOT NULL, + owner_user_id TEXT NOT NULL, + agent_account_id TEXT NOT NULL, + provider TEXT NOT NULL, + model TEXT NOT NULL, + input_tokens BIGINT NOT NULL, + output_tokens BIGINT NOT NULL, + cache_read_tokens BIGINT NOT NULL, + cache_write_tokens BIGINT NOT NULL, + total_tokens BIGINT NOT NULL, + cost_micro_usd BIGINT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (tenant_id, bucket_start, owner_user_id, agent_account_id, provider, model) +); + +-- token_usage_prune_horizon: the one global row holding the UTC day the latest +-- prune cut at. Rollups before it may count pruned events, so a rebuild keeps +-- them. Not tenant-scoped, because a prune spans every tenant. It only moves +-- forward, and '-infinity' means no prune has run. +CREATE TABLE token_usage_prune_horizon ( + singleton BOOLEAN PRIMARY KEY DEFAULT TRUE CHECK (singleton), + horizon TIMESTAMPTZ NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +INSERT INTO token_usage_prune_horizon (horizon) VALUES ('-infinity'); + +-- The table grant 0001's schema-wide GRANT gives every table it created. Named +-- per table here, because re-running 0001's ALL TABLES grant would hand +-- server_key_state its revoked DELETE back. +GRANT SELECT, INSERT, UPDATE, DELETE + ON token_usage_events, token_usage_rollups_hourly, token_usage_rollups_daily, + token_usage_prune_horizon + TO compass_app, compass_system; + +-- token_usage_prune_horizon holds the one row the migration inserts. Re-inserting +-- it at '-infinity' would let a rebuild drop pruned-day rollups; UPDATE stays. +REVOKE INSERT, DELETE ON token_usage_prune_horizon FROM compass_app, compass_system; + +-- 0001's tenant_tables loop body, for this migration's tenant-owned tables. +DO $$ +DECLARE + t text; + tenant_tables text[] := ARRAY[ + 'token_usage_events', 'token_usage_rollups_hourly', 'token_usage_rollups_daily' + ]; +BEGIN + FOREACH t IN ARRAY tenant_tables LOOP + EXECUTE format('ALTER TABLE %I ENABLE ROW LEVEL SECURITY', t); + EXECUTE format('ALTER TABLE %I FORCE ROW LEVEL SECURITY', t); + EXECUTE format($f$ + CREATE POLICY tenant_isolation ON %I + USING ((SELECT current_setting('compass.tenant_id', TRUE)) <> '' + AND tenant_id = (SELECT current_setting('compass.tenant_id', TRUE))) + WITH CHECK ((SELECT current_setting('compass.tenant_id', TRUE)) <> '' + AND tenant_id = (SELECT current_setting('compass.tenant_id', TRUE))) + $f$, t); + END LOOP; +END $$; + +-- 0001's updated_at_tables loop body, for this migration's updated_at tables. +DO $$ +DECLARE + t text; + updated_at_tables text[] := ARRAY[ + 'token_usage_rollups_hourly', + 'token_usage_rollups_daily', + 'token_usage_prune_horizon' + ]; +BEGIN + FOREACH t IN ARRAY updated_at_tables LOOP + EXECUTE format( + 'CREATE TRIGGER set_updated_at BEFORE UPDATE ON %I + FOR EACH ROW EXECUTE FUNCTION set_updated_at()', t); + END LOOP; +END $$; diff --git a/go/internal/store/queries/tenant.sql b/go/internal/store/queries/tenant.sql index 7cab59318..57bf4323b 100644 --- a/go/internal/store/queries/tenant.sql +++ b/go/internal/store/queries/tenant.sql @@ -8,3 +8,8 @@ INSERT INTO tenants (id, slug, display_name, created_at_unix_ms) VALUES ($1, $2, -- name: TenantIDBySlug :one SELECT id FROM tenants WHERE slug = $1; + +-- ListTenantIDs lists every tenant. tenants has no row-level security, so the +-- app role sees them all without the BYPASSRLS system role. +-- name: ListTenantIDs :many +SELECT id FROM tenants ORDER BY id; diff --git a/go/internal/store/queries/token_usage.sql b/go/internal/store/queries/token_usage.sql new file mode 100644 index 000000000..e4ff362e7 --- /dev/null +++ b/go/internal/store/queries/token_usage.sql @@ -0,0 +1,138 @@ +-- Token-usage queries (Plane A). The tenant tx sets the GUC that RLS reads and +-- that tenant_id defaults to, so no statement here names a tenant. + +-- LockTokenUsage serializes one tenant's appends and rebuilds. A rebuild beside +-- an append would count the append's events twice or lose them. +-- name: LockTokenUsage :exec +SELECT pg_advisory_xact_lock(hashtext('token_usage:' || current_setting('compass.tenant_id', TRUE))); + +-- AppendTokenUsageEvents inserts a batch and adds only the events it inserted +-- to both rollups, so a replayed id counts once. The arrays hold one element +-- per event, and no id repeats. +-- name: AppendTokenUsageEvents :exec +WITH inserted AS ( + INSERT INTO token_usage_events ( + id, occurred_at, agent_account_id, owner_user_id, session_id, request_id, provider, + model, credential_id, input_tokens, output_tokens, cache_read_tokens, + cache_write_tokens, total_tokens, cost_micro_usd, rate_version, outcome) + SELECT unnest(@ids::text[]), unnest(@occurred_at::timestamptz[]), + unnest(@agent_account_ids::text[]), unnest(@owner_user_ids::text[]), + unnest(@session_ids::text[]), unnest(@request_ids::text[]), + unnest(@providers::text[]), unnest(@models::text[]), + unnest(@credential_ids::text[]), unnest(@input_tokens::bigint[]), + unnest(@output_tokens::bigint[]), unnest(@cache_read_tokens::bigint[]), + unnest(@cache_write_tokens::bigint[]), unnest(@total_tokens::bigint[]), + unnest(@cost_micro_usd::bigint[]), unnest(@rate_versions::text[]), + unnest(@outcomes::text[]) + ON CONFLICT (tenant_id, id) DO NOTHING + RETURNING occurred_at, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd +), hourly AS ( + INSERT INTO token_usage_rollups_hourly ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) + SELECT date_trunc('hour', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM inserted + GROUP BY 1, 2, 3, 4, 5 + ON CONFLICT (tenant_id, bucket_start, owner_user_id, agent_account_id, provider, model) + DO UPDATE SET + input_tokens = token_usage_rollups_hourly.input_tokens + EXCLUDED.input_tokens, + output_tokens = token_usage_rollups_hourly.output_tokens + EXCLUDED.output_tokens, + cache_read_tokens = token_usage_rollups_hourly.cache_read_tokens + EXCLUDED.cache_read_tokens, + cache_write_tokens = token_usage_rollups_hourly.cache_write_tokens + EXCLUDED.cache_write_tokens, + total_tokens = token_usage_rollups_hourly.total_tokens + EXCLUDED.total_tokens, + cost_micro_usd = token_usage_rollups_hourly.cost_micro_usd + EXCLUDED.cost_micro_usd +) +INSERT INTO token_usage_rollups_daily ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) +SELECT date_trunc('day', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM inserted + GROUP BY 1, 2, 3, 4, 5 +ON CONFLICT (tenant_id, bucket_start, owner_user_id, agent_account_id, provider, model) +DO UPDATE SET + input_tokens = token_usage_rollups_daily.input_tokens + EXCLUDED.input_tokens, + output_tokens = token_usage_rollups_daily.output_tokens + EXCLUDED.output_tokens, + cache_read_tokens = token_usage_rollups_daily.cache_read_tokens + EXCLUDED.cache_read_tokens, + cache_write_tokens = token_usage_rollups_daily.cache_write_tokens + EXCLUDED.cache_write_tokens, + total_tokens = token_usage_rollups_daily.total_tokens + EXCLUDED.total_tokens, + cost_micro_usd = token_usage_rollups_daily.cost_micro_usd + EXCLUDED.cost_micro_usd; + +-- TokenUsageSeries sums one granularity's rollup rows per bucket. granularity +-- is the usage.Granularity value: 1 is hourly, 2 is daily. An empty agent list +-- or provider means no filter. +-- name: TokenUsageSeries :many +SELECT r.bucket_start::timestamptz AS bucket_start, + sum(r.input_tokens)::bigint AS input_tokens, + sum(r.output_tokens)::bigint AS output_tokens, + sum(r.cache_read_tokens)::bigint AS cache_read_tokens, + sum(r.cache_write_tokens)::bigint AS cache_write_tokens, + sum(r.total_tokens)::bigint AS total_tokens, + sum(r.cost_micro_usd)::bigint AS cost_micro_usd + FROM (SELECT bucket_start, agent_account_id, provider, input_tokens, output_tokens, + cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd + FROM token_usage_rollups_hourly + WHERE @granularity::integer = 1 + UNION ALL + SELECT bucket_start, agent_account_id, provider, input_tokens, output_tokens, + cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd + FROM token_usage_rollups_daily + WHERE @granularity::integer = 2) AS r + WHERE r.bucket_start >= @start_at::timestamptz + AND r.bucket_start < @end_at::timestamptz + AND (coalesce(cardinality(@agent_account_ids::text[]), 0) = 0 + OR r.agent_account_id = ANY (@agent_account_ids::text[])) + AND (@provider::text = '' OR r.provider = @provider::text) + GROUP BY r.bucket_start + ORDER BY r.bucket_start; + +-- TokenUsagePruneHorizon reads the prune horizon and holds it until the tx +-- ends, so a prune cannot delete the events a rebuild is about to count. +-- name: TokenUsagePruneHorizon :one +SELECT horizon FROM token_usage_prune_horizon FOR SHARE; + +-- DeleteTokenUsageRollupsFrom drops both rollups from the horizon on. Older +-- rollups can hold pruned events, so they stay. +-- name: DeleteTokenUsageRollupsFrom :exec +WITH hourly AS ( + DELETE FROM token_usage_rollups_hourly WHERE bucket_start >= @horizon::timestamptz +) +DELETE FROM token_usage_rollups_daily WHERE bucket_start >= @horizon::timestamptz; + +-- RollUpTokenUsageFrom rebuilds both rollups from the events at or after the +-- horizon. +-- name: RollUpTokenUsageFrom :exec +WITH hourly AS ( + INSERT INTO token_usage_rollups_hourly ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) + SELECT date_trunc('hour', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM token_usage_events + WHERE occurred_at >= @horizon::timestamptz + GROUP BY 1, 2, 3, 4, 5 +) +INSERT INTO token_usage_rollups_daily ( + bucket_start, owner_user_id, agent_account_id, provider, model, input_tokens, + output_tokens, cache_read_tokens, cache_write_tokens, total_tokens, cost_micro_usd) +SELECT date_trunc('day', occurred_at, 'UTC'), owner_user_id, agent_account_id, provider, + model, sum(input_tokens), sum(output_tokens), sum(cache_read_tokens), + sum(cache_write_tokens), sum(total_tokens), sum(cost_micro_usd) + FROM token_usage_events + WHERE occurred_at >= @horizon::timestamptz + GROUP BY 1, 2, 3, 4, 5; + +-- AdvanceTokenUsagePruneHorizon commits before the prune deletes anything, and +-- waits for a rebuild that holds the old horizon. It only moves forward. +-- name: AdvanceTokenUsagePruneHorizon :execrows +UPDATE token_usage_prune_horizon SET horizon = GREATEST(horizon, @cutoff::timestamptz); + +-- DeleteTokenUsageEventsBefore deletes old events of the tx's tenant only, +-- because row-level security scopes it. The prune runs it once per tenant. +-- name: DeleteTokenUsageEventsBefore :execrows +DELETE FROM token_usage_events WHERE occurred_at < @cutoff::timestamptz; diff --git a/go/internal/store/rls_pgtest_test.go b/go/internal/store/rls_pgtest_test.go index d3777b829..6e1392419 100644 --- a/go/internal/store/rls_pgtest_test.go +++ b/go/internal/store/rls_pgtest_test.go @@ -546,10 +546,12 @@ func TestRLSCatalogEnabledAndForced(t *testing.T) { // NOT mere bookkeeping: the bucket-A loop at the end of this test iterates // this map and asserts each listed table has RLS DISABLED, so removing the // entry would let an accidental `ENABLE ROW LEVEL SECURITY` on either table - // pass unnoticed. + // pass unnoticed. token_usage_prune_horizon is global because one prune + // spans every tenant. bucketA := map[string]bool{ "tenants": true, "tokens": true, "agent_config_bundle": true, "server_secrets": true, "server_key_state": true, + "token_usage_prune_horizon": true, } // Enumerate every table in the current (per-test) schema that carries a @@ -609,6 +611,7 @@ func TestRLSCatalogEnabledAndForced(t *testing.T) { "linear_agent_sessions", "issues", "forge_repo_subscriptions", "forge_artifact_cursors", "forge_state_transitions", + "token_usage_events", "token_usage_rollups_hourly", "token_usage_rollups_daily", } for _, tbl := range tenantOwned { if !enumerated[tbl] { diff --git a/go/internal/store/server_secrets_pgtest_test.go b/go/internal/store/server_secrets_pgtest_test.go index 7103dc191..75e77cb7e 100644 --- a/go/internal/store/server_secrets_pgtest_test.go +++ b/go/internal/store/server_secrets_pgtest_test.go @@ -9,6 +9,18 @@ import ( "testing" ) +// hasTablePrivilege reports whether role holds priv on tbl in the test schema. +func hasTablePrivilege(t *testing.T, s *Store, role, tbl, priv string) bool { + t.Helper() + var ok bool + if err := s.pool.QueryRow(t.Context(), + `SELECT has_table_privilege($1, (current_schema()||'.'||$2)::regclass, $3)`, + role, tbl, priv).Scan(&ok); err != nil { + t.Fatal(err) + } + return ok +} + func TestT0ServerSecretsShape(t *testing.T) { s := newTestStore(t) ctx := context.Background() @@ -44,10 +56,11 @@ func TestT0ServerSecretsShape(t *testing.T) { } // The load-bearing half: grants are NOT inherited from 0001's snapshot, so - // each table needs its own. Both tables are asserted with their OWN expected + // each table needs its own. Each table is asserted with their OWN expected // privilege set, and server_key_state's withheld DELETE is asserted ABSENT — // that omission is a deliberate least-privilege choice (the tripwire digest - // must not be droppable), so it is pinned, not left to chance. + // must not be droppable), so it is pinned, not left to chance. The usage + // prune horizon withholds INSERT and DELETE for the reason in 0003_token_usage.sql. for _, tc := range []struct { tbl string granted []string @@ -55,27 +68,16 @@ func TestT0ServerSecretsShape(t *testing.T) { }{ {"server_secrets", []string{"SELECT", "INSERT", "UPDATE", "DELETE"}, nil}, {"server_key_state", []string{"SELECT", "INSERT", "UPDATE"}, []string{"DELETE"}}, + {"token_usage_prune_horizon", []string{"SELECT", "UPDATE"}, []string{"INSERT", "DELETE"}}, } { for _, role := range []string{"compass_app", "compass_system"} { for _, priv := range tc.granted { - var ok bool - if err := s.pool.QueryRow(ctx, - `SELECT has_table_privilege($1, (current_schema()||'.'||$2)::regclass, $3)`, - role, tc.tbl, priv).Scan(&ok); err != nil { - t.Fatal(err) - } - if !ok { + if !hasTablePrivilege(t, s, role, tc.tbl, priv) { t.Fatalf("%s: %s lacks %s", tc.tbl, role, priv) } } for _, priv := range tc.denied { - var ok bool - if err := s.pool.QueryRow(ctx, - `SELECT has_table_privilege($1, (current_schema()||'.'||$2)::regclass, $3)`, - role, tc.tbl, priv).Scan(&ok); err != nil { - t.Fatal(err) - } - if ok { + if hasTablePrivilege(t, s, role, tc.tbl, priv) { t.Fatalf("%s: %s has %s, which is deliberately withheld", tc.tbl, role, priv) } } diff --git a/go/internal/store/updated_at_pgtest_test.go b/go/internal/store/updated_at_pgtest_test.go index 27d605ce2..30f546f74 100644 --- a/go/internal/store/updated_at_pgtest_test.go +++ b/go/internal/store/updated_at_pgtest_test.go @@ -4,10 +4,10 @@ package store // The set_updated_at trigger convention (RIG-3495), proven against real // Postgres. updated_at is maintained by ONE mechanism — the BEFORE UPDATE -// trigger 0001_init.sql installs on every table carrying the column — and no -// query file sets it by hand. These tests are the enforcement of that: they -// fail if the trigger is missing, which is exactly what happens if someone -// re-adds a table without an updated_at_tables entry, or drops the block. +// trigger each migration installs on every table it creates with the column — +// and no query file sets it by hand. These tests are the enforcement of that: +// they fail if the trigger is missing, which is exactly what happens if a +// migration adds such a table without its trigger, or drops the block. // // Three properties, each a distinct failure mode: // @@ -152,7 +152,7 @@ func TestSecretsUpdatedAtIsLive(t *testing.T) { // TestUpdatedAtTriggerCatalogFloor is the catalog guard the three behavioural // tests above cannot be: they each pin ONE named table, so a FUTURE table that -// declares updated_at and forgets its updated_at_tables entry is silently +// declares updated_at and forgets its trigger is silently // untriggered — the exact rot RIG-3495 exists to prevent, reintroduced by // omission rather than by edit. This enumerates the live catalog instead of // trusting a hand-maintained list, the same self-auditing posture as @@ -190,7 +190,7 @@ func TestUpdatedAtTriggerCatalogFloor(t *testing.T) { if err := rows.Scan(&tbl); err != nil { t.Fatalf("scan catalog row: %v", err) } - t.Errorf("%s: declares updated_at but has no set_updated_at trigger — add it to updated_at_tables in 0001_init.sql, or the column can only ever equal created_at and every reader of it is reading a lie", tbl) + t.Errorf("%s: declares updated_at but has no set_updated_at trigger — create it in the migration that adds the table, or the column can only ever equal created_at and every reader of it is reading a lie", tbl) } if err := rows.Err(); err != nil { t.Fatalf("iterate catalog rows: %v", err) diff --git a/go/internal/usage/export_test.go b/go/internal/usage/export_test.go new file mode 100644 index 000000000..0b2858563 --- /dev/null +++ b/go/internal/usage/export_test.go @@ -0,0 +1,26 @@ +package usage + +import ( + "context" + + "github.com/jackc/pgx/v5" + + "github.com/RigelBuild/compass/go/internal/store" +) + +// ClearMemoryRollups drops every rollup of a NewMemory store. The contract +// suite runs in usage_test, so it reaches the unexported state through here. +func ClearMemoryRollups(s Store) { s.(*Memory).clearRollups() } + +// ClearPostgresRollups drops every tenant's rollups from a NewPostgres store. +// It runs as the system role, because a tenant sees only its own rows. +func ClearPostgresRollups(ctx context.Context, s Store) error { + ctx = store.WithSystemRole(ctx) + return s.(*Postgres).st.WithTx(ctx, func(tx pgx.Tx) error { + if _, err := tx.Exec(ctx, "DELETE FROM token_usage_rollups_hourly"); err != nil { + return err + } + _, err := tx.Exec(ctx, "DELETE FROM token_usage_rollups_daily") + return err + }) +} diff --git a/go/internal/usage/memory.go b/go/internal/usage/memory.go new file mode 100644 index 000000000..021af15a1 --- /dev/null +++ b/go/internal/usage/memory.go @@ -0,0 +1,168 @@ +package usage + +import ( + "cmp" + "context" + "math" + "slices" + "sync" + + "github.com/RigelBuild/compass/go/internal/store" +) + +// rollupGranularities are the widths every event is rolled up at. +var rollupGranularities = [...]Granularity{GranularityHour, GranularityDay} + +// rollupKey is one rollup row's identity, the same key the Postgres tables use. +type rollupKey struct { + granularity Granularity + startUnixMs int64 + ownerUserID string + agentAccountID string + provider string + model string +} + +// memoryTenant is one tenant's raw events and rollups. +type memoryTenant struct { + events map[string]TokenUsageEvent + rollups map[rollupKey]Bucket +} + +// Memory is the in-memory reference Store. It keeps the tenant from ctx, or "" +// when ctx carries none. +type Memory struct { + mu sync.Mutex + tenants map[store.TenantID]*memoryTenant + // horizonUnixMs is the UTC day the latest prune cut at. Rollups before it + // may count pruned events, so a rebuild must keep them. + horizonUnixMs int64 +} + +var _ Store = (*Memory)(nil) + +// NewMemory returns an empty in-memory Store. +func NewMemory() *Memory { + return &Memory{tenants: map[store.TenantID]*memoryTenant{}, horizonUnixMs: math.MinInt64} +} + +// AppendTokenUsage implements Store. +func (m *Memory) AppendTokenUsage(ctx context.Context, events []TokenUsageEvent) error { + if err := ValidateEvents(events); err != nil { + return err + } + m.mu.Lock() + defer m.mu.Unlock() + t := m.tenant(ctx) + for i := range events { + if _, seen := t.events[events[i].ID]; seen { + continue + } + t.events[events[i].ID] = events[i] + t.rollUp(&events[i]) + } + return nil +} + +// rollUp adds e to every rollup bucket that holds it. +func (t *memoryTenant) rollUp(e *TokenUsageEvent) { + for _, g := range rollupGranularities { + k := rollupKey{ + granularity: g, + startUnixMs: g.BucketStart(e.OccurredAtUnixMs), + ownerUserID: e.OwnerUserID, + agentAccountID: e.AgentAccountID, + provider: e.Provider, + model: e.Model, + } + b := t.rollups[k].plus(e.usage()) + b.StartUnixMs = k.startUnixMs + t.rollups[k] = b + } +} + +// TokenUsageSeries implements Store. +func (m *Memory) TokenUsageSeries(ctx context.Context, q SeriesQuery) ([]Bucket, error) { + if err := q.Validate(); err != nil { + return nil, err + } + m.mu.Lock() + defer m.mu.Unlock() + sums := map[int64]Bucket{} + for k, b := range m.tenant(ctx).rollups { + if k.granularity != q.Granularity || k.startUnixMs < q.StartUnixMs || k.startUnixMs >= q.EndUnixMs { + continue + } + if len(q.AgentAccountIDs) > 0 && !slices.Contains(q.AgentAccountIDs, k.agentAccountID) { + continue + } + if q.Provider != "" && k.provider != q.Provider { + continue + } + sums[k.startUnixMs] = sums[k.startUnixMs].plus(b) + } + series := make([]Bucket, 0, len(sums)) + for start, b := range sums { + b.StartUnixMs = start + series = append(series, b) + } + slices.SortFunc(series, func(a, b Bucket) int { return cmp.Compare(a.StartUnixMs, b.StartUnixMs) }) + return series, nil +} + +// RebuildTokenUsageRollups implements Store. +func (m *Memory) RebuildTokenUsageRollups(ctx context.Context) error { + m.mu.Lock() + defer m.mu.Unlock() + t := m.tenant(ctx) + for k := range t.rollups { + if k.startUnixMs >= m.horizonUnixMs { + delete(t.rollups, k) + } + } + for id := range t.events { + if e := t.events[id]; e.OccurredAtUnixMs >= m.horizonUnixMs { + t.rollUp(&e) + } + } + return nil +} + +// PruneTokenUsageBefore implements Store. +func (m *Memory) PruneTokenUsageBefore(_ context.Context, beforeUnixMs int64) (int64, error) { + cutoff := GranularityDay.BucketStart(beforeUnixMs) + m.mu.Lock() + defer m.mu.Unlock() + var n int64 + for _, t := range m.tenants { + for id, e := range t.events { + if e.OccurredAtUnixMs < cutoff { + delete(t.events, id) + n++ + } + } + } + // An earlier cutoff must not pull the horizon back over pruned days. + m.horizonUnixMs = max(m.horizonUnixMs, cutoff) + return n, nil +} + +// clearRollups drops every tenant's rollups and keeps the raw events. +func (m *Memory) clearRollups() { + m.mu.Lock() + defer m.mu.Unlock() + for _, t := range m.tenants { + clear(t.rollups) + } +} + +// tenant returns ctx's tenant state, creating it. The caller holds m.mu. +func (m *Memory) tenant(ctx context.Context) *memoryTenant { + id, _ := store.TenantFromContext(ctx) + t, ok := m.tenants[id] + if !ok { + t = &memoryTenant{events: map[string]TokenUsageEvent{}, rollups: map[rollupKey]Bucket{}} + m.tenants[id] = t + } + return t +} diff --git a/go/internal/usage/memory_test.go b/go/internal/usage/memory_test.go new file mode 100644 index 000000000..7a2856c22 --- /dev/null +++ b/go/internal/usage/memory_test.go @@ -0,0 +1,22 @@ +package usage_test + +import ( + "context" + "fmt" + "testing" + + "github.com/RigelBuild/compass/go/internal/store" + "github.com/RigelBuild/compass/go/internal/usage" + "github.com/RigelBuild/compass/go/internal/usage/usagetest" +) + +func TestMemory(t *testing.T) { + usagetest.Run(t, usagetest.Harness{ + New: func(*testing.T) usage.Store { return usage.NewMemory() }, + Ctx: func(t *testing.T, tenant int) context.Context { + t.Helper() + return store.WithTenant(t.Context(), store.TenantID(fmt.Sprint("t", tenant))) + }, + CorruptRollups: func(_ *testing.T, s usage.Store) { usage.ClearMemoryRollups(s) }, + }) +} diff --git a/go/internal/usage/postgres.go b/go/internal/usage/postgres.go new file mode 100644 index 000000000..620a5c2a2 --- /dev/null +++ b/go/internal/usage/postgres.go @@ -0,0 +1,225 @@ +package usage + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/RigelBuild/compass/go/internal/store" + "github.com/RigelBuild/compass/go/internal/store/db" +) + +// Postgres is the Store backed by the store of record. Each call but the prune +// runs in one transaction that row-level security scopes to ctx's tenant. +type Postgres struct { + st *store.Store +} + +var _ Store = (*Postgres)(nil) + +// NewPostgres returns the Store that st backs. +func NewPostgres(st *store.Store) *Postgres { + return &Postgres{st: st} +} + +// AppendTokenUsage implements Store. +func (p *Postgres) AppendTokenUsage(ctx context.Context, events []TokenUsageEvent) error { + if err := ValidateEvents(events); err != nil { + return err + } + if len(events) == 0 { + return nil + } + params := appendParams(events) + err := p.st.WithTx(ctx, func(tx pgx.Tx) error { + q := db.New(tx) + if err := q.LockTokenUsage(ctx); err != nil { + return err + } + return q.AppendTokenUsageEvents(ctx, params) + }) + if err != nil { + return fmt.Errorf("usage: append token usage: %w", err) + } + return nil +} + +// TokenUsageSeries implements Store. +func (p *Postgres) TokenUsageSeries(ctx context.Context, q SeriesQuery) ([]Bucket, error) { + if err := q.Validate(); err != nil { + return nil, err + } + var rows []db.TokenUsageSeriesRow + err := p.st.WithTx(ctx, func(tx pgx.Tx) error { + var err error + rows, err = db.New(tx).TokenUsageSeries(ctx, db.TokenUsageSeriesParams{ + Granularity: int32(q.Granularity), + StartAt: bound(q.StartUnixMs), + EndAt: bound(q.EndUnixMs), + AgentAccountIds: q.AgentAccountIDs, + Provider: q.Provider, + }) + return err + }) + if err != nil { + return nil, fmt.Errorf("usage: read token usage series: %w", err) + } + series := make([]Bucket, len(rows)) + for i := range rows { + r := &rows[i] + series[i] = Bucket{ + StartUnixMs: r.BucketStart.Time.UnixMilli(), + InputTokens: r.InputTokens, + OutputTokens: r.OutputTokens, + CacheReadTokens: r.CacheReadTokens, + CacheWriteTokens: r.CacheWriteTokens, + TotalTokens: r.TotalTokens, + CostMicroUSD: r.CostMicroUsd, + } + } + return series, nil +} + +// RebuildTokenUsageRollups implements Store. +func (p *Postgres) RebuildTokenUsageRollups(ctx context.Context) error { + err := p.st.WithTx(ctx, func(tx pgx.Tx) error { + q := db.New(tx) + if err := q.LockTokenUsage(ctx); err != nil { + return err + } + horizon, err := q.TokenUsagePruneHorizon(ctx) + if err != nil { + return err + } + if err := q.DeleteTokenUsageRollupsFrom(ctx, horizon); err != nil { + return err + } + return q.RollUpTokenUsageFrom(ctx, horizon) + }) + if err != nil { + return fmt.Errorf("usage: rebuild token usage rollups: %w", err) + } + return nil +} + +// PruneTokenUsageBefore implements Store. The BYPASSRLS system role is only for +// the named background loops, so each tenant's events go in that tenant's tx. +func (p *Postgres) PruneTokenUsageBefore(ctx context.Context, beforeUnixMs int64) (int64, error) { + cutoff := bound(GranularityDay.BucketStart(beforeUnixMs)) + // The horizon commits before any delete. Its update waits out a rebuild that + // holds the old horizon, and every later rebuild reads the new one. + var tenants []string + err := p.st.WithTx(ctx, func(tx pgx.Tx) error { + q := db.New(tx) + n, err := q.AdvanceTokenUsagePruneHorizon(ctx, cutoff) + if err != nil { + return err + } + // Without the row the rebuild also fails, so the prune must not delete events. + if n == 0 { + return errors.New("prune horizon row is missing") + } + tenants, err = q.ListTenantIDs(ctx) + return err + }) + if err != nil { + return 0, fmt.Errorf("usage: advance token usage prune horizon: %w", err) + } + var deleted int64 + for _, id := range tenants { + n, err := p.pruneTenant(store.WithTenant(ctx, store.TenantID(id)), cutoff) + if err != nil { + // The earlier tenants' deletes have committed, so the count keeps them. + return deleted, fmt.Errorf("usage: prune token usage of tenant %s: %w", id, err) + } + deleted += n + } + return deleted, nil +} + +// pruneTenant deletes ctx's tenant's events before cutoff. Row-level security is +// what keeps the delete inside that tenant. +func (p *Postgres) pruneTenant(ctx context.Context, cutoff pgtype.Timestamptz) (int64, error) { + var n int64 + err := p.st.WithTx(ctx, func(tx pgx.Tx) error { + var err error + n, err = db.New(tx).DeleteTokenUsageEventsBefore(ctx, cutoff) + return err + }) + if err != nil { + return 0, err + } + return n, nil +} + +// appendParams lays the batch out one column per slice, the shape the append +// statement binds. A repeated id keeps its first event, as in the reference. +func appendParams(events []TokenUsageEvent) db.AppendTokenUsageEventsParams { + n := len(events) + p := db.AppendTokenUsageEventsParams{ + Ids: make([]string, 0, n), + OccurredAt: make([]pgtype.Timestamptz, 0, n), + AgentAccountIds: make([]string, 0, n), + OwnerUserIds: make([]string, 0, n), + SessionIds: make([]string, 0, n), + RequestIds: make([]string, 0, n), + Providers: make([]string, 0, n), + Models: make([]string, 0, n), + CredentialIds: make([]string, 0, n), + InputTokens: make([]int64, 0, n), + OutputTokens: make([]int64, 0, n), + CacheReadTokens: make([]int64, 0, n), + CacheWriteTokens: make([]int64, 0, n), + TotalTokens: make([]int64, 0, n), + CostMicroUsd: make([]int64, 0, n), + RateVersions: make([]string, 0, n), + Outcomes: make([]string, 0, n), + } + seen := make(map[string]struct{}, n) + for i := range events { + e := &events[i] + if _, dup := seen[e.ID]; dup { + continue + } + seen[e.ID] = struct{}{} + p.Ids = append(p.Ids, e.ID) + p.OccurredAt = append(p.OccurredAt, timestamptz(e.OccurredAtUnixMs)) + p.AgentAccountIds = append(p.AgentAccountIds, e.AgentAccountID) + p.OwnerUserIds = append(p.OwnerUserIds, e.OwnerUserID) + p.SessionIds = append(p.SessionIds, e.SessionID) + p.RequestIds = append(p.RequestIds, e.RequestID) + p.Providers = append(p.Providers, e.Provider) + p.Models = append(p.Models, e.Model) + p.CredentialIds = append(p.CredentialIds, e.CredentialID) + p.InputTokens = append(p.InputTokens, e.InputTokens) + p.OutputTokens = append(p.OutputTokens, e.OutputTokens) + p.CacheReadTokens = append(p.CacheReadTokens, e.CacheReadTokens) + p.CacheWriteTokens = append(p.CacheWriteTokens, e.CacheWriteTokens) + p.TotalTokens = append(p.TotalTokens, e.TotalTokens) + p.CostMicroUsd = append(p.CostMicroUsd, e.CostMicroUSD) + p.RateVersions = append(p.RateVersions, e.RateVersion) + p.Outcomes = append(p.Outcomes, e.Outcome) + } + return p +} + +// timestamptz converts unix milliseconds to the Postgres parameter type. +func timestamptz(unixMs int64) pgtype.Timestamptz { + return pgtype.Timestamptz{Time: time.UnixMilli(unixMs), Valid: true} +} + +// bound converts a comparison bound. pgx wraps an out-of-range time silently, +// so such a bound becomes the infinity that compares the same way. +func bound(unixMs int64) pgtype.Timestamptz { + switch { + case unixMs < minTimestamptzMs: + return pgtype.Timestamptz{InfinityModifier: pgtype.NegativeInfinity, Valid: true} + case unixMs > maxTimestamptzMs: + return pgtype.Timestamptz{InfinityModifier: pgtype.Infinity, Valid: true} + } + return timestamptz(unixMs) +} diff --git a/go/internal/usage/postgres_pgtest_test.go b/go/internal/usage/postgres_pgtest_test.go new file mode 100644 index 000000000..e23ef8a2f --- /dev/null +++ b/go/internal/usage/postgres_pgtest_test.go @@ -0,0 +1,280 @@ +//go:build pgtest && unix + +package usage_test + +import ( + "context" + "slices" + "sync" + "testing" + "time" + + "github.com/jackc/pgx/v5" + + "github.com/RigelBuild/compass/go/internal/pgtest" + "github.com/RigelBuild/compass/go/internal/store" + "github.com/RigelBuild/compass/go/internal/store/db" + "github.com/RigelBuild/compass/go/internal/usage" + "github.com/RigelBuild/compass/go/internal/usage/usagetest" +) + +// pgTenants are seeded into every fresh schema, so a fixed id names each one. +var pgTenants = [...]store.TenantID{"usage-tenant-0", "usage-tenant-1"} + +func TestPostgres(t *testing.T) { + usagetest.Run(t, usagetest.Harness{ + New: newPostgres, + Ctx: func(t *testing.T, tenant int) context.Context { + t.Helper() + return store.WithTenant(t.Context(), pgTenants[tenant]) + }, + CorruptRollups: func(t *testing.T, s usage.Store) { + t.Helper() + if err := usage.ClearPostgresRollups(t.Context(), s); err != nil { + t.Fatalf("clear rollups: %v", err) + } + }, + }) +} + +// newPostgres opens a store on a fresh schema and seeds both tenants. +func newPostgres(t *testing.T) usage.Store { + t.Helper() + _, s := openPostgres(t) + return s +} + +// openPostgres is newPostgres that also returns the store of record, so a test +// can hold its own tx beside the usage store's. +func openPostgres(t *testing.T) (*store.Store, usage.Store) { + t.Helper() + st, err := store.Open(t.Context(), pgtest.RequireDSN(t)) + if err != nil { + t.Fatalf("store.Open: %v", err) + } + t.Cleanup(st.Close) + err = st.WithTx(t.Context(), func(tx pgx.Tx) error { + for _, id := range pgTenants { + if err := db.New(tx).InsertTenant(t.Context(), db.InsertTenantParams{ + ID: string(id), + Slug: string(id), + DisplayName: string(id), + CreatedAtUnixMs: time.Now().UnixMilli(), + }); err != nil { + return err + } + } + return nil + }) + if err != nil { + t.Fatalf("seed tenants: %v", err) + } + return st, usage.NewPostgres(st) +} + +// A rebuild reads the horizon FOR SHARE, so a prune must not advance it, and +// so must not delete events the rebuild would miss, until the rebuild commits. +func TestPostgresPruneWaitsForHeldHorizon(t *testing.T) { + st, s := openPostgres(t) + ctx := store.WithTenant(t.Context(), pgTenants[0]) + day0 := time.Date(2026, time.March, 10, 0, 0, 0, 0, time.UTC) + if err := s.AppendTokenUsage(ctx, []usage.TokenUsageEvent{ + pgEvent("old-1", day0.Add(time.Hour), 1), + pgEvent("old-2", day0.Add(2*time.Hour), 2), + }); err != nil { + t.Fatalf("AppendTokenUsage: %v", err) + } + + pid, commit := holdTx(t, st, ctx, func(tx pgx.Tx) error { + _, err := db.New(tx).TokenUsagePruneHorizon(ctx) + return err + }) + type result struct { + n int64 + err error + } + done := make(chan result, 1) + go func() { + n, err := s.PruneTokenUsageBefore(t.Context(), day0.Add(48*time.Hour).UnixMilli()) + done <- result{n, err} + }() + + waitBlockedBy(t, st, pid) + if n := countEvents(t, st, ctx); n != 2 { + t.Fatalf("with the horizon held, %d events remain, want 2: the prune did not wait", n) + } + + commit() + r := <-done + if r.err != nil || r.n != 2 { + t.Fatalf("PruneTokenUsageBefore = %d, %v; want 2, nil", r.n, r.err) + } + if n := countEvents(t, st, ctx); n != 0 { + t.Fatalf("after the prune, %d events remain, want 0", n) + } +} + +// A rebuild beside an append would count the append twice or lose it, so both +// wait for whoever holds the tenant's usage lock. +func TestPostgresWritersWaitForUsageLock(t *testing.T) { + day0 := time.Date(2026, time.March, 10, 0, 0, 0, 0, time.UTC) + q := usage.SeriesQuery{ + Granularity: usage.GranularityDay, + StartUnixMs: day0.UnixMilli(), + EndUnixMs: day0.Add(24 * time.Hour).UnixMilli(), + } + for _, tc := range []struct { + name string + // prepare runs before the lock is taken; write runs while it is held. + prepare func(ctx context.Context, s usage.Store) error + write func(ctx context.Context, s usage.Store) error + wantInput int64 + }{ + {"append", func(context.Context, usage.Store) error { return nil }, + func(ctx context.Context, s usage.Store) error { + return s.AppendTokenUsage(ctx, []usage.TokenUsageEvent{pgEvent("e2", day0, 2)}) + }, 3}, + {"rebuild", usage.ClearPostgresRollups, + func(ctx context.Context, s usage.Store) error { return s.RebuildTokenUsageRollups(ctx) }, 1}, + } { + t.Run(tc.name, func(t *testing.T) { + st, s := openPostgres(t) + ctx := store.WithTenant(t.Context(), pgTenants[0]) + if err := s.AppendTokenUsage(ctx, []usage.TokenUsageEvent{pgEvent("e1", day0, 1)}); err != nil { + t.Fatalf("AppendTokenUsage: %v", err) + } + if err := tc.prepare(ctx, s); err != nil { + t.Fatalf("prepare: %v", err) + } + before := series(t, s, ctx, q) + + pid, commit := holdTx(t, st, ctx, func(tx pgx.Tx) error { + return db.New(tx).LockTokenUsage(ctx) + }) + done := make(chan error, 1) + go func() { done <- tc.write(ctx, s) }() + + waitBlockedBy(t, st, pid) + if got := series(t, s, ctx, q); !slices.Equal(got, before) { + t.Fatalf("with the usage lock held, rollups moved: got %+v, want %+v", got, before) + } + + commit() + if err := <-done; err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + got := series(t, s, ctx, q) + if slices.Equal(got, before) { + t.Fatalf("the %s left the series unchanged, so the held-lock check proved nothing", tc.name) + } + if len(got) != 1 || got[0].InputTokens != tc.wantInput { + t.Fatalf("after the %s, series = %+v, want one bucket of %d input tokens", tc.name, got, tc.wantInput) + } + }) + } +} + +// pgEvent is a valid event with n input tokens. +func pgEvent(id string, at time.Time, n int64) usage.TokenUsageEvent { + return usage.TokenUsageEvent{ + ID: id, + OccurredAtUnixMs: at.UnixMilli(), + AgentAccountID: "a1", + OwnerUserID: "owner-1", + Provider: "anthropic", + Model: "model-1", + InputTokens: n, + TotalTokens: n, + Outcome: usage.OutcomeOK, + } +} + +// holdTx runs fn in a tx scoped by ctx and holds that tx open until commit +// runs. It returns the tx's backend pid, which a waiter reports as its blocker. +func holdTx(t *testing.T, st *store.Store, ctx context.Context, fn func(pgx.Tx) error) (pid int, commit func()) { + t.Helper() + ready := make(chan struct{}) + release := make(chan struct{}) + done := make(chan error, 1) + go func() { + done <- st.WithTx(ctx, func(tx pgx.Tx) error { + if err := fn(tx); err != nil { + return err + } + if err := tx.QueryRow(ctx, "SELECT pg_backend_pid()").Scan(&pid); err != nil { + return err + } + close(ready) + <-release + return nil + }) + }() + select { + case <-ready: + case err := <-done: + t.Fatalf("hold tx: %v", err) + } + var once sync.Once + // A failed test still releases the tx, so the pool can close. + t.Cleanup(func() { once.Do(func() { close(release); <-done }) }) + return pid, func() { + once.Do(func() { + close(release) + if err := <-done; err != nil { + t.Fatalf("commit held tx: %v", err) + } + }) + } +} + +// waitBlockedBy returns once some backend waits on pid. Postgres signals no +// lock wait, so this polls the catalog. On timeout it returns anyway, and the +// caller's assertion reports the waiter that never waited. +func waitBlockedBy(t *testing.T, st *store.Store, pid int) { + t.Helper() + ctx := store.WithSystemRole(t.Context()) + deadline := time.After(5 * time.Second) + tick := time.NewTicker(10 * time.Millisecond) + defer tick.Stop() + for { + var blocked bool + err := st.WithTx(ctx, func(tx pgx.Tx) error { + return tx.QueryRow(ctx, + `SELECT EXISTS (SELECT 1 FROM pg_stat_activity WHERE $1 = ANY(pg_blocking_pids(pid)))`, + pid).Scan(&blocked) + }) + if err != nil { + t.Fatalf("poll pg_blocking_pids: %v", err) + } + if blocked { + return + } + select { + case <-deadline: + return + case <-tick.C: + } + } +} + +// countEvents counts ctx's tenant's raw events, which the rollups outlive. +func countEvents(t *testing.T, st *store.Store, ctx context.Context) int { + t.Helper() + var n int + err := st.WithTx(ctx, func(tx pgx.Tx) error { + return tx.QueryRow(ctx, "SELECT count(*) FROM token_usage_events").Scan(&n) + }) + if err != nil { + t.Fatalf("count events: %v", err) + } + return n +} + +func series(t *testing.T, s usage.Store, ctx context.Context, q usage.SeriesQuery) []usage.Bucket { + t.Helper() + got, err := s.TokenUsageSeries(ctx, q) + if err != nil { + t.Fatalf("TokenUsageSeries: %v", err) + } + return got +} diff --git a/go/internal/usage/retention.go b/go/internal/usage/retention.go new file mode 100644 index 000000000..16d498aa8 --- /dev/null +++ b/go/internal/usage/retention.go @@ -0,0 +1,80 @@ +package usage + +import ( + "context" + "log/slog" + "time" +) + +// DefaultRetention is the server's default window for raw events. The rollups +// outlive it, so the usage charts keep their history. +const DefaultRetention = 90 * 24 * time.Hour + +// sweepInterval is the time between prunes. The cutoff moves one UTC day at a +// time, so a faster sweep would only rescan the same events. +const sweepInterval = 24 * time.Hour + +// pruner is the part of Store the sweeper drives. +type pruner interface { + PruneTokenUsageBefore(ctx context.Context, beforeUnixMs int64) (int64, error) +} + +// RetentionConfig configures a RetentionSweeper. +type RetentionConfig struct { + // Retention is how long raw events are kept. 0 or less keeps every event, + // because the sweep does not run. + Retention time.Duration + // Log is the sweep logger; nil uses slog.Default(). + Log *slog.Logger +} + +// RetentionSweeper deletes the raw events older than the retention window. +type RetentionSweeper struct { + store pruner + retention time.Duration + log *slog.Logger +} + +// NewRetentionSweeper returns a sweeper that prunes s. +func NewRetentionSweeper(s pruner, cfg RetentionConfig) *RetentionSweeper { + log := cfg.Log + if log == nil { + log = slog.Default() + } + return &RetentionSweeper{store: s, retention: cfg.Retention, log: log} +} + +// Run prunes at once and then on every sweepInterval tick, until ctx ends. It +// never returns a prune error: in the serve group that would stop the server. +func (w *RetentionSweeper) Run(ctx context.Context) error { + if w.retention <= 0 { + // The serve group treats a nil return as a clean exit, not a failure. + w.log.InfoContext(ctx, "usage retention: raw token-usage prune disabled") + return nil + } + w.sweep(ctx) + t := time.NewTicker(sweepInterval) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return nil + case <-t.C: + w.sweep(ctx) + } + } +} + +// sweep runs one prune. A failed prune is logged, and the next tick retries it. +func (w *RetentionSweeper) sweep(ctx context.Context) { + deleted, err := w.store.PruneTokenUsageBefore(ctx, time.Now().Add(-w.retention).UnixMilli()) + switch { + case ctx.Err() != nil: + // Shutdown interrupted the prune; the next start sweeps again. + case err != nil: + w.log.ErrorContext(ctx, "usage retention: prune raw token-usage events", "error", err) + case deleted > 0: + w.log.InfoContext(ctx, "usage retention: pruned raw token-usage events", + "deleted", deleted, "retention", w.retention) + } +} diff --git a/go/internal/usage/retention_test.go b/go/internal/usage/retention_test.go new file mode 100644 index 000000000..34a12dd48 --- /dev/null +++ b/go/internal/usage/retention_test.go @@ -0,0 +1,91 @@ +package usage_test + +import ( + "context" + "errors" + "log/slog" + "testing" + "testing/synctest" + "time" + + "github.com/RigelBuild/compass/go/internal/usage" +) + +// chanPruner hands each prune cutoff to the test. Blocking on the receive is +// what lets the synctest clock run forward to the next tick. +type chanPruner struct { + cutoffs chan int64 + err error +} + +func (p *chanPruner) PruneTokenUsageBefore(ctx context.Context, beforeUnixMs int64) (int64, error) { + select { + case p.cutoffs <- beforeUnixMs: + return 0, p.err + case <-ctx.Done(): + return 0, ctx.Err() + } +} + +func TestRetentionSweeperPrunesAtStartAndDaily(t *testing.T) { + const day = 24 * time.Hour + tests := []struct { + name string + retention time.Duration // the configured window + window time.Duration // the window the cutoffs must use + err error // what every prune returns + }{ + {name: "configured_window", retention: 7 * day, window: 7 * day}, + {name: "failed_prune_is_retried_next_day", retention: day, window: day, err: errors.New("database down")}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + p := &chanPruner{cutoffs: make(chan int64), err: tt.err} + w := usage.NewRetentionSweeper(p, usage.RetentionConfig{ + Retention: tt.retention, + Log: slog.New(slog.DiscardHandler), + }) + ctx, cancel := context.WithCancel(t.Context()) + errc := make(chan error, 1) + start := time.Now() + go func() { errc <- w.Run(ctx) }() + + for _, at := range []time.Time{start, start.Add(day)} { + got := <-p.cutoffs + if now := time.Now(); !now.Equal(at) { + t.Fatalf("prune ran at %v, want %v", now, at) + } + if want := at.Add(-tt.window).UnixMilli(); got != want { + t.Fatalf("prune cutoff = %d, want %d (%v before %v)", got, want, tt.window, at) + } + } + cancel() + if err := <-errc; err != nil { + t.Fatalf("Run = %v, want nil after cancel", err) + } + }) + }) + } +} + +func TestRetentionSweeperZeroRetentionPrunesNothing(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + // A prune blocks on the unread channel, so Run returns only if it never prunes. + p := &chanPruner{cutoffs: make(chan int64)} + w := usage.NewRetentionSweeper(p, usage.RetentionConfig{Log: slog.New(slog.DiscardHandler)}) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + errc := make(chan error, 1) + go func() { errc <- w.Run(ctx) }() + synctest.Wait() + select { + case err := <-errc: + if err != nil { + t.Fatalf("Run = %v, want nil", err) + } + default: + t.Fatal("Run with a zero retention window is still running, want it to return without pruning") + } + }) +} diff --git a/go/internal/usage/usage.go b/go/internal/usage/usage.go new file mode 100644 index 000000000..d425b7ec7 --- /dev/null +++ b/go/internal/usage/usage.go @@ -0,0 +1,197 @@ +// Package usage is the Plane-A token-usage contract: the event the LLM gateway +// emits, the Store seam every backend satisfies, and the in-memory reference. +package usage + +import ( + "context" + "errors" + "fmt" + "math" + "strconv" + "time" +) + +// ErrInvalidArgument marks a malformed event batch or query. It is returned +// before any write, so a rejected batch leaves no trace. +var ErrInvalidArgument = errors.New("usage: invalid argument") + +// The Outcome values. Failed and aborted calls still carry the partial usage +// the provider billed, so they are recorded and rolled up like "ok". +const ( + OutcomeOK = "ok" + OutcomeError = "error" + OutcomeAborted = "aborted" +) + +// TokenUsageEvent is one completed upstream model call. Money is integer +// micro-USD so that no float reaches a store. +type TokenUsageEvent struct { + ID string // server-assigned UUID + OccurredAtUnixMs int64 + AgentAccountID string // the calling agent (attribution spine) + OwnerUserID string // rollup key only; the store resolves tenancy, never from this field + SessionID string // runner session id, when known + RequestID string // gateway request UUID + Provider string // e.g. "anthropic" + Model string // resolved model id + CredentialID string // which pool account served it + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheWriteTokens int64 + TotalTokens int64 + CostMicroUSD int64 // total; computed gateway-side from the pricing table + RateVersion string // pricing-table version applied at write time + Outcome string // "ok" | "error" | "aborted" (partial usage still recorded) +} + +// Granularity is a rollup bucket width. Backends persist these values, so +// never renumber them. +type Granularity int32 + +// The rollup granularities. Unspecified is the zero value a query must not use. +const ( + GranularityUnspecified Granularity = iota + GranularityHour + GranularityDay +) + +const ( + hourMs = int64(time.Hour / time.Millisecond) + dayMs = 24 * hourMs +) + +// The unix milliseconds pgx can send as a timestamptz. Postgres starts at +// 4714-11-24 BC, and pgx overflows int64 microseconds past MaxInt64/1000. +const ( + minTimestamptzMs = -210_866_803_200_000 + maxTimestamptzMs = math.MaxInt64 / 1000 +) + +// SeriesQuery selects the rollup buckets whose start lies in +// [StartUnixMs, EndUnixMs). +type SeriesQuery struct { + Granularity Granularity + StartUnixMs int64 + EndUnixMs int64 + // AgentAccountIDs narrows the read to these agents; empty means all. It is + // a set so that a subtree read can resolve the tree once and sum once. + AgentAccountIDs []string + // Provider narrows the read to one provider; empty means all. + Provider string +} + +// Bucket is one bucket's usage, summed over every rollup row a query matched. +type Bucket struct { + StartUnixMs int64 + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheWriteTokens int64 + TotalTokens int64 + CostMicroUSD int64 +} + +// Store is the Plane-A usage store seam: an append-only event write plus rollup +// reads. Each call but PruneTokenUsageBefore is scoped to ctx's tenant, which +// the backend resolves. +type Store interface { + // AppendTokenUsage records a batch atomically. It is idempotent on event ID + // for events still inside the retention window, so a retried batch never + // double-counts. A pruned event's ID is forgotten and counts again. + AppendTokenUsage(ctx context.Context, events []TokenUsageEvent) error + // TokenUsageSeries returns the non-empty buckets in the window, oldest first. + TokenUsageSeries(ctx context.Context, query SeriesQuery) ([]Bucket, error) + // RebuildTokenUsageRollups recomputes every rollup from the prune horizon on. + // Older rollups outlive their pruned events, so it keeps them. + RebuildTokenUsageRollups(ctx context.Context) error + // PruneTokenUsageBefore deletes every tenant's raw events before the UTC day + // that holds beforeUnixMs and keeps the rollups. It returns the delete count. + PruneTokenUsageBefore(ctx context.Context, beforeUnixMs int64) (int64, error) +} + +// BucketStart returns the start of the UTC-aligned bucket that holds unixMs, +// or unixMs unchanged for a granularity no rollup is kept at. +func (g Granularity) BucketStart(unixMs int64) int64 { + width := g.widthMs() + if width == 0 { + return unixMs + } + return unixMs - unixMs%width +} + +// widthMs is the bucket width, or 0 for a granularity no rollup is kept at. +func (g Granularity) widthMs() int64 { + switch g { + case GranularityHour: + return hourMs + case GranularityDay: + return dayMs + case GranularityUnspecified: + return 0 + } + return 0 +} + +// Validate rejects a granularity that no rollup is kept at. Without the check, +// the read would return an empty series and hide the error. +func (q *SeriesQuery) Validate() error { + if q.Granularity.widthMs() == 0 { + return fmt.Errorf("%w: granularity %d has no rollup", ErrInvalidArgument, q.Granularity) + } + return nil +} + +// ValidateEvents rejects a batch that holds an event the rollups cannot key or +// sum. Backends call it before they write, so a bad batch writes nothing. +func ValidateEvents(events []TokenUsageEvent) error { + for i := range events { + if problem := events[i].problem(); problem != "" { + return fmt.Errorf("%w: event %d (%q): %s", ErrInvalidArgument, i, events[i].ID, problem) + } + } + return nil +} + +// problem names the first defect in e, or returns "" if e is valid. +func (e *TokenUsageEvent) problem() string { + switch { + case e.ID == "": + return "empty id" + case e.OccurredAtUnixMs <= 0: + return "non-positive occurred_at" + case e.OccurredAtUnixMs > maxTimestamptzMs: + return "occurred_at past the last storable timestamp" + case e.AgentAccountID == "", e.OwnerUserID == "", e.Provider == "", e.Model == "": + return "empty rollup key (agent, owner, provider, model)" + case e.InputTokens < 0, e.OutputTokens < 0, e.CacheReadTokens < 0, e.CacheWriteTokens < 0, + e.TotalTokens < 0, e.CostMicroUSD < 0: + return "negative token count or cost" + case e.Outcome != OutcomeOK && e.Outcome != OutcomeError && e.Outcome != OutcomeAborted: + return "unknown outcome " + strconv.Quote(e.Outcome) + } + return "" +} + +// usage is the event's contribution to one rollup bucket. +func (e *TokenUsageEvent) usage() Bucket { + return Bucket{ + InputTokens: e.InputTokens, + OutputTokens: e.OutputTokens, + CacheReadTokens: e.CacheReadTokens, + CacheWriteTokens: e.CacheWriteTokens, + TotalTokens: e.TotalTokens, + CostMicroUSD: e.CostMicroUSD, + } +} + +// plus adds o's sums to b and keeps b's StartUnixMs. +func (b Bucket) plus(o Bucket) Bucket { + b.InputTokens += o.InputTokens + b.OutputTokens += o.OutputTokens + b.CacheReadTokens += o.CacheReadTokens + b.CacheWriteTokens += o.CacheWriteTokens + b.TotalTokens += o.TotalTokens + b.CostMicroUSD += o.CostMicroUSD + return b +} diff --git a/go/internal/usage/usagetest/usagetest.go b/go/internal/usage/usagetest/usagetest.go new file mode 100644 index 000000000..0a1be1a41 --- /dev/null +++ b/go/internal/usage/usagetest/usagetest.go @@ -0,0 +1,397 @@ +// Package usagetest is the contract suite every usage.Store backend must pass. +package usagetest + +import ( + "context" + "errors" + "math" + "slices" + "testing" + "time" + + "github.com/RigelBuild/compass/go/internal/usage" +) + +// Harness adapts one backend to the suite. +type Harness struct { + // New returns an empty store. Each subtest calls it once. + New func(t *testing.T) usage.Store + // Ctx returns a context scoped to tenant 0 or tenant 1 of the store. + Ctx func(t *testing.T, tenant int) context.Context + // CorruptRollups deletes every rollup of s and keeps the raw events. + CorruptRollups func(t *testing.T, s usage.Store) +} + +// Run runs the contract suite against the backend h adapts. +func Run(t *testing.T, h Harness) { + t.Helper() + t.Run("reappending_the_same_ids_does_not_double_count", reappendDoesNotDoubleCount(h)) + t.Run("mixed_batch_counts_only_new_ids", mixedBatchCountsOnlyNewIDs(h)) + t.Run("buckets_split_on_utc_hour_and_day_boundaries", bucketsSplitOnUTCBoundaries(h)) + t.Run("agent_and_provider_filters_narrow_the_sum", filtersNarrowTheSum(h)) + t.Run("empty_range_returns_no_buckets", emptyRangeReturnsNoBuckets(h)) + t.Run("unbounded_range_returns_every_bucket", unboundedRangeReturnsEveryBucket(h)) + t.Run("series_is_ordered_by_bucket_start", seriesIsOrdered(h)) + t.Run("rebuild_restores_cleared_rollups", rebuildRestoresClearedRollups(h)) + t.Run("prune_drops_old_events_and_keeps_rollups", pruneKeepsRollups(h)) + t.Run("prune_accepts_extreme_cutoffs", pruneAcceptsExtremeCutoffs(h)) + t.Run("tenants_are_isolated", tenantsAreIsolated(h)) + t.Run("invalid_events_are_rejected_and_write_nothing", invalidEventsWriteNothing(h)) + t.Run("invalid_queries_are_rejected", invalidQueriesAreRejected(h)) +} + +// day0 is a UTC midnight, so the fixtures sit on known bucket boundaries. +var day0 = time.Date(2026, time.March, 10, 0, 0, 0, 0, time.UTC) + +// at returns day0 plus d in unix milliseconds. +func at(d time.Duration) int64 { return day0.Add(d).UnixMilli() } + +const ( + hour = time.Hour + day = 24 * time.Hour +) + +// event builds a valid event whose sums are distinct multiples of n, so a +// wrong field in a rollup shows up as a wrong sum. +func event(id string, occurredAtUnixMs int64, agent, provider string, n int64) usage.TokenUsageEvent { + return usage.TokenUsageEvent{ + ID: id, + OccurredAtUnixMs: occurredAtUnixMs, + AgentAccountID: agent, + OwnerUserID: "owner-1", + SessionID: "session-1", + RequestID: "request-" + id, + Provider: provider, + Model: "model-1", + CredentialID: "cred-1", + InputTokens: n, + OutputTokens: 2 * n, + CacheReadTokens: 3 * n, + CacheWriteTokens: 4 * n, + TotalTokens: 10 * n, + CostMicroUSD: 100 * n, + RateVersion: "v1", + Outcome: usage.OutcomeOK, + } +} + +// bucket is the expected rollup for events whose n values sum to n. +func bucket(startUnixMs, n int64) usage.Bucket { + return usage.Bucket{ + StartUnixMs: startUnixMs, + InputTokens: n, + OutputTokens: 2 * n, + CacheReadTokens: 3 * n, + CacheWriteTokens: 4 * n, + TotalTokens: 10 * n, + CostMicroUSD: 100 * n, + } +} + +// query is a query at granularity g over the whole fixture range. +func query(g usage.Granularity) usage.SeriesQuery { + return usage.SeriesQuery{Granularity: g, StartUnixMs: at(-30 * day), EndUnixMs: at(30 * day)} +} + +func mustAppend(t *testing.T, ctx context.Context, s usage.Store, events ...usage.TokenUsageEvent) { + t.Helper() + if err := s.AppendTokenUsage(ctx, events); err != nil { + t.Fatalf("AppendTokenUsage: %v", err) + } +} + +func mustRebuild(t *testing.T, ctx context.Context, s usage.Store) { + t.Helper() + if err := s.RebuildTokenUsageRollups(ctx); err != nil { + t.Fatalf("RebuildTokenUsageRollups: %v", err) + } +} + +func mustPrune(t *testing.T, ctx context.Context, s usage.Store, beforeUnixMs, wantDeleted int64) { + t.Helper() + n, err := s.PruneTokenUsageBefore(ctx, beforeUnixMs) + if err != nil { + t.Fatalf("PruneTokenUsageBefore: %v", err) + } + if n != wantDeleted { + t.Fatalf("PruneTokenUsageBefore deleted %d events, want %d", n, wantDeleted) + } +} + +func wantSeries(t *testing.T, ctx context.Context, s usage.Store, q usage.SeriesQuery, want ...usage.Bucket) { + t.Helper() + got, err := s.TokenUsageSeries(ctx, q) + if err != nil { + t.Fatalf("TokenUsageSeries(%+v): %v", q, err) + } + if !slices.Equal(got, want) { + t.Fatalf("TokenUsageSeries(%+v)\n got %+v\n want %+v", q, got, want) + } +} + +func reappendDoesNotDoubleCount(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + batch := []usage.TokenUsageEvent{ + event("e1", at(1*hour), "a1", "anthropic", 1), + event("e2", at(2*hour), "a1", "anthropic", 2), + } + mustAppend(t, ctx, s, batch...) + mustAppend(t, ctx, s, batch...) + wantSeries(t, ctx, s, query(usage.GranularityHour), bucket(at(1*hour), 1), bucket(at(2*hour), 2)) + wantSeries(t, ctx, s, query(usage.GranularityDay), bucket(at(0), 3)) + } +} + +func mixedBatchCountsOnlyNewIDs(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + mustAppend(t, ctx, s, event("e1", at(1*hour), "a1", "anthropic", 1)) + // The repeated e1 carries other sums: the first write wins. + mustAppend(t, ctx, s, + event("e1", at(1*hour), "a1", "anthropic", 50), + event("e2", at(1*hour), "a1", "anthropic", 2), + event("e3", at(1*hour), "a1", "anthropic", 4), + event("e3", at(1*hour), "a1", "anthropic", 4), + ) + wantSeries(t, ctx, s, query(usage.GranularityHour), bucket(at(1*hour), 7)) + } +} + +func bucketsSplitOnUTCBoundaries(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + mustAppend(t, ctx, s, + event("e1", at(hour-time.Millisecond), "a1", "anthropic", 1), + event("e2", at(hour), "a1", "anthropic", 2), + event("e3", at(day-time.Millisecond), "a1", "anthropic", 4), + event("e4", at(day), "a1", "anthropic", 8), + ) + wantSeries(t, ctx, s, query(usage.GranularityHour), + bucket(at(0), 1), bucket(at(hour), 2), bucket(at(23*hour), 4), bucket(at(day), 8)) + wantSeries(t, ctx, s, query(usage.GranularityDay), bucket(at(0), 7), bucket(at(day), 8)) + + // The window is half-open on bucket start: [01:00, 23:00) keeps only the 01:00 bucket. + wantSeries(t, ctx, s, + usage.SeriesQuery{Granularity: usage.GranularityHour, StartUnixMs: at(hour), EndUnixMs: at(23 * hour)}, + bucket(at(hour), 2)) + // A window that starts inside a bucket excludes it. + wantSeries(t, ctx, s, + usage.SeriesQuery{Granularity: usage.GranularityDay, StartUnixMs: at(time.Millisecond), EndUnixMs: at(30 * day)}, + bucket(at(day), 8)) + } +} + +func filtersNarrowTheSum(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + mustAppend(t, ctx, s, + event("e1", at(hour), "a1", "anthropic", 1), + event("e2", at(hour), "a2", "anthropic", 2), + event("e3", at(hour), "a3", "openai", 4), + event("e4", at(hour), "a1", "openai", 8), + ) + q := query(usage.GranularityHour) + wantSeries(t, ctx, s, q, bucket(at(hour), 15)) + + q.AgentAccountIDs = []string{"a1", "a3"} + wantSeries(t, ctx, s, q, bucket(at(hour), 13)) + + q.Provider = "openai" + wantSeries(t, ctx, s, q, bucket(at(hour), 12)) + + q.AgentAccountIDs = nil + q.Provider = "anthropic" + wantSeries(t, ctx, s, q, bucket(at(hour), 3)) + + q.AgentAccountIDs = []string{"a2"} + q.Provider = "openai" + wantSeries(t, ctx, s, q) + } +} + +func emptyRangeReturnsNoBuckets(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + wantSeries(t, ctx, s, query(usage.GranularityHour)) + + mustAppend(t, ctx, s, event("e1", at(hour), "a1", "anthropic", 1)) + wantSeries(t, ctx, s, usage.SeriesQuery{Granularity: usage.GranularityHour, StartUnixMs: at(hour), EndUnixMs: at(hour)}) + wantSeries(t, ctx, s, usage.SeriesQuery{Granularity: usage.GranularityHour, StartUnixMs: at(2 * day), EndUnixMs: at(3 * day)}) + // Control: the event is stored, so only the windows above are empty. + wantSeries(t, ctx, s, usage.SeriesQuery{Granularity: usage.GranularityHour, StartUnixMs: at(hour), EndUnixMs: at(2 * hour)}, + bucket(at(hour), 1)) + } +} + +// The extreme bounds are past what a timestamptz holds, so a backend must clamp +// them rather than fail or wrap. +func unboundedRangeReturnsEveryBucket(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + mustAppend(t, ctx, s, + event("e1", at(-day), "a1", "anthropic", 1), + event("e2", at(hour), "a1", "anthropic", 2), + event("e3", at(2*day), "a1", "anthropic", 4), + ) + wantSeries(t, ctx, s, + usage.SeriesQuery{Granularity: usage.GranularityHour, StartUnixMs: math.MinInt64, EndUnixMs: math.MaxInt64}, + bucket(at(-day), 1), bucket(at(hour), 2), bucket(at(2*day), 4)) + wantSeries(t, ctx, s, + usage.SeriesQuery{Granularity: usage.GranularityDay, StartUnixMs: math.MinInt64, EndUnixMs: math.MaxInt64}, + bucket(at(-day), 1), bucket(at(0), 2), bucket(at(2*day), 4)) + } +} + +func seriesIsOrdered(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + mustAppend(t, ctx, s, + event("e1", at(5*hour), "a1", "anthropic", 1), + event("e2", at(2*day), "a1", "anthropic", 2), + event("e3", at(0), "a1", "anthropic", 4), + ) + mustAppend(t, ctx, s, event("e4", at(-day), "a1", "anthropic", 8)) + wantSeries(t, ctx, s, query(usage.GranularityHour), + bucket(at(-day), 8), bucket(at(0), 4), bucket(at(5*hour), 1), bucket(at(2*day), 2)) + wantSeries(t, ctx, s, query(usage.GranularityDay), + bucket(at(-day), 8), bucket(at(0), 5), bucket(at(2*day), 2)) + } +} + +func rebuildRestoresClearedRollups(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + mustAppend(t, ctx, s, + event("e1", at(hour), "a1", "anthropic", 1), + event("e2", at(hour), "a2", "openai", 2), + event("e3", at(day+hour), "a1", "anthropic", 4), + ) + hourly := []usage.Bucket{bucket(at(hour), 3), bucket(at(day+hour), 4)} + daily := []usage.Bucket{bucket(at(0), 3), bucket(at(day), 4)} + wantSeries(t, ctx, s, query(usage.GranularityHour), hourly...) + wantSeries(t, ctx, s, query(usage.GranularityDay), daily...) + + h.CorruptRollups(t, s) + wantSeries(t, ctx, s, query(usage.GranularityHour)) + + mustRebuild(t, ctx, s) + wantSeries(t, ctx, s, query(usage.GranularityHour), hourly...) + wantSeries(t, ctx, s, query(usage.GranularityDay), daily...) + } +} + +func pruneKeepsRollups(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx0, ctx1 := h.New(t), h.Ctx(t, 0), h.Ctx(t, 1) + mustAppend(t, ctx0, s, + event("t0-old", at(hour), "a1", "anthropic", 1), + event("t0-new", at(day), "a1", "anthropic", 2), + ) + mustAppend(t, ctx1, s, + event("t1-old", at(23*hour), "a1", "anthropic", 4), + event("t1-new", at(day+hour), "a1", "anthropic", 8), + ) + hourly := []usage.Bucket{bucket(at(hour), 1), bucket(at(day), 2)} + wantSeries(t, ctx0, s, query(usage.GranularityHour), hourly...) + + // The cutoff rounds down to the UTC day, so a mid-day-1 cutoff keeps day 1. + // The prune spans tenants: it deletes one old event from each. + mustPrune(t, ctx0, s, at(day+12*hour), 2) + wantSeries(t, ctx0, s, query(usage.GranularityHour), hourly...) + + // A rebuild keeps the rollups older than the prune horizon. + mustRebuild(t, ctx0, s) + wantSeries(t, ctx0, s, query(usage.GranularityHour), hourly...) + + // With the rollups gone, a rebuild can only count the events that remain. + h.CorruptRollups(t, s) + mustRebuild(t, ctx0, s) + wantSeries(t, ctx0, s, query(usage.GranularityHour), bucket(at(day), 2)) + wantSeries(t, ctx0, s, query(usage.GranularityDay), bucket(at(day), 2)) + + mustPrune(t, ctx0, s, at(day+12*hour), 0) + } +} + +func pruneAcceptsExtremeCutoffs(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + mustAppend(t, ctx, s, + event("e1", at(hour), "a1", "anthropic", 1), + event("e2", at(day), "a1", "anthropic", 2), + ) + mustPrune(t, ctx, s, math.MinInt64, 0) + mustPrune(t, ctx, s, math.MaxInt64, 2) + // Every rollup is now older than the horizon, so a rebuild keeps them all. + mustRebuild(t, ctx, s) + wantSeries(t, ctx, s, query(usage.GranularityHour), bucket(at(hour), 1), bucket(at(day), 2)) + } +} + +func tenantsAreIsolated(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx0, ctx1 := h.New(t), h.Ctx(t, 0), h.Ctx(t, 1) + mustAppend(t, ctx0, s, event("t0-e1", at(hour), "a1", "anthropic", 1)) + wantSeries(t, ctx1, s, query(usage.GranularityHour)) + + mustAppend(t, ctx1, s, event("t1-e1", at(hour), "a1", "anthropic", 2)) + wantSeries(t, ctx0, s, query(usage.GranularityHour), bucket(at(hour), 1)) + wantSeries(t, ctx1, s, query(usage.GranularityHour), bucket(at(hour), 2)) + + // A rebuild in one tenant leaves the other tenant's rollups as they were. + mustRebuild(t, ctx1, s) + wantSeries(t, ctx0, s, query(usage.GranularityHour), bucket(at(hour), 1)) + wantSeries(t, ctx1, s, query(usage.GranularityHour), bucket(at(hour), 2)) + + // Event IDs are unique per tenant, so each tenant counts "shared" once. + mustAppend(t, ctx0, s, event("shared", at(hour), "a1", "anthropic", 4)) + mustAppend(t, ctx1, s, event("shared", at(hour), "a1", "anthropic", 8)) + wantSeries(t, ctx0, s, query(usage.GranularityHour), bucket(at(hour), 5)) + wantSeries(t, ctx1, s, query(usage.GranularityHour), bucket(at(hour), 10)) + } +} + +func invalidEventsWriteNothing(h Harness) func(*testing.T) { + return func(t *testing.T) { + valid := event("ok", at(hour), "a1", "anthropic", 1) + for name, mutate := range map[string]func(*usage.TokenUsageEvent){ + "empty_id": func(e *usage.TokenUsageEvent) { e.ID = "" }, + "zero_time": func(e *usage.TokenUsageEvent) { e.OccurredAtUnixMs = 0 }, + "time_past_max": func(e *usage.TokenUsageEvent) { e.OccurredAtUnixMs = math.MaxInt64 }, + "empty_agent": func(e *usage.TokenUsageEvent) { e.AgentAccountID = "" }, + "empty_owner": func(e *usage.TokenUsageEvent) { e.OwnerUserID = "" }, + "empty_provider": func(e *usage.TokenUsageEvent) { e.Provider = "" }, + "empty_model": func(e *usage.TokenUsageEvent) { e.Model = "" }, + "negative_tokens": func(e *usage.TokenUsageEvent) { e.OutputTokens = -1 }, + "negative_cost": func(e *usage.TokenUsageEvent) { e.CostMicroUSD = -1 }, + "unknown_outcome": func(e *usage.TokenUsageEvent) { e.Outcome = "timeout" }, + } { + t.Run(name, func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + bad := event("bad", at(hour), "a1", "anthropic", 2) + mutate(&bad) + err := s.AppendTokenUsage(ctx, []usage.TokenUsageEvent{valid, bad}) + if !errors.Is(err, usage.ErrInvalidArgument) { + t.Fatalf("AppendTokenUsage error = %v, want ErrInvalidArgument", err) + } + wantSeries(t, ctx, s, query(usage.GranularityHour)) + + // The valid event was not recorded, so appending it now counts it once. + mustAppend(t, ctx, s, valid) + wantSeries(t, ctx, s, query(usage.GranularityHour), bucket(at(hour), 1)) + }) + } + } +} + +func invalidQueriesAreRejected(h Harness) func(*testing.T) { + return func(t *testing.T) { + s, ctx := h.New(t), h.Ctx(t, 0) + for _, g := range []usage.Granularity{usage.GranularityUnspecified, usage.GranularityDay + 1, -1} { + if _, err := s.TokenUsageSeries(ctx, query(g)); !errors.Is(err, usage.ErrInvalidArgument) { + t.Fatalf("TokenUsageSeries(granularity %d) error = %v, want ErrInvalidArgument", g, err) + } + } + } +} diff --git a/go/server/serve.go b/go/server/serve.go index 0c30dd069..524ae404d 100644 --- a/go/server/serve.go +++ b/go/server/serve.go @@ -138,6 +138,9 @@ type ServeConfig struct { // points the whole deployment's custody — including the master key — at a // managed store, per the record's A2 KMS-by-provider-URI custody note. SecretProvider string + // UsageEventRetention is how long raw token-usage events are kept; 0 turns + // the daily prune off. The CLI defaults it to 90 days. + UsageEventRetention time.Duration } // ForgeConfig configures the board webhook-ingestion lane (RIG-2883) and the @@ -904,6 +907,8 @@ func Serve(ctx context.Context, cfg ServeConfig) error { // projection on the comms bus, both on the serve group rooted on gctx // (cancels at shutdown; presence also ends when drainDoors closes the bus). startCommsConsumers(gctx, g, commsBus, fab, st, hub, hubLog) + // A daily prune bounds the raw usage log; the rollups keep their sums. + startUsageRetention(gctx, g, st, cfg.UsageEventRetention, hubLog) // Drain member of the same group: wake on gctx cancellation, then hand off to // drainDoors. A drain that overruns (a handler still wedged mid-replay) // surfaces as the error rather than a false clean shutdown; a real serve diff --git a/go/server/sinks.go b/go/server/sinks.go index 2d078dea7..2279a47df 100644 --- a/go/server/sinks.go +++ b/go/server/sinks.go @@ -9,6 +9,7 @@ package server import ( "context" "log/slog" + "time" "golang.org/x/sync/errgroup" @@ -21,6 +22,7 @@ import ( "github.com/RigelBuild/compass/go/internal/presence" "github.com/RigelBuild/compass/go/internal/runnerhub" "github.com/RigelBuild/compass/go/internal/store" + "github.com/RigelBuild/compass/go/internal/usage" ) // The lifecycle sink is the Bridge board (internal/board): a session lifecycle @@ -182,6 +184,13 @@ func startForgeIngestLanes(gctx context.Context, g *errgroup.Group, board *board } } +// startUsageRetention starts the token-usage retention sweeper on the serve +// group. The raw log grows with activity; the rollups keep the history. +func startUsageRetention(gctx context.Context, g *errgroup.Group, st *store.Store, retention time.Duration, log *slog.Logger) { + w := usage.NewRetentionSweeper(usage.NewPostgres(st), usage.RetentionConfig{Retention: retention, Log: log}) + g.Go(func() error { return w.Run(gctx) }) +} + // logFrameDiagnostics emits the hub's frame-loss snapshot as one line. Serve // calls it on shutdown, so every run states plainly how many relayed frames // never reached their surface.