From a9789f28aedb7a057509f8346edaf5f551fa08f2 Mon Sep 17 00:00:00 2001 From: Henry Park <16583448+henrypark133@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:15:40 +0000 Subject: [PATCH 1/7] feat: maintain transactional inference and service spend counters --- crates/api/tests/e2e_all/main.rs | 1 + crates/api/tests/e2e_all/spend_counters.rs | 218 ++++++++++++++++++ .../sql/V0081__add_spend_counters.sql | 19 ++ crates/database/src/repositories/mod.rs | 1 + .../organization_service_usage.rs | 22 +- .../src/repositories/organization_usage.rs | 14 +- .../src/repositories/spend_counters.rs | 31 +++ 7 files changed, 297 insertions(+), 9 deletions(-) create mode 100644 crates/api/tests/e2e_all/spend_counters.rs create mode 100644 crates/database/src/migrations/sql/V0081__add_spend_counters.sql create mode 100644 crates/database/src/repositories/spend_counters.rs diff --git a/crates/api/tests/e2e_all/main.rs b/crates/api/tests/e2e_all/main.rs index d9784dd57..2df566665 100644 --- a/crates/api/tests/e2e_all/main.rs +++ b/crates/api/tests/e2e_all/main.rs @@ -87,6 +87,7 @@ mod score; mod serving_provider; mod session_logout; mod signature_verification; +mod spend_counters; mod usage_chat_completions; mod usage_provider_attribution; mod usage_recording; diff --git a/crates/api/tests/e2e_all/spend_counters.rs b/crates/api/tests/e2e_all/spend_counters.rs new file mode 100644 index 000000000..94eadcd83 --- /dev/null +++ b/crates/api/tests/e2e_all/spend_counters.rs @@ -0,0 +1,218 @@ +use crate::common::{ + create_api_key_in_workspace, create_org, list_workspaces, setup_test_server_with_database, + E2E_QWEN_MODEL_NAME, +}; +use database::models::RecordUsageRequest; +use database::repositories::{ + OrganizationServiceUsageRepository, OrganizationUsageRepository, RecordServiceUsageRequest, +}; +use services::usage::InferenceType; +use uuid::Uuid; + +async fn fixture() -> ( + axum_test::TestServer, + std::sync::Arc, + RecordUsageRequest, + Uuid, +) { + let (server, database) = setup_test_server_with_database().await; + let org = create_org(&server).await; + let workspace = list_workspaces(&server, org.id.clone()) + .await + .into_iter() + .next() + .expect("organization has a default workspace"); + let api_key = create_api_key_in_workspace( + &server, + workspace.id.clone(), + "spend-counter-key".to_string(), + ) + .await; + let client = database.pool().get().await.expect("database connection"); + let model = client + .query_one( + "SELECT id, model_name FROM models WHERE model_name = $1", + &[&E2E_QWEN_MODEL_NAME], + ) + .await + .expect("shared model fixture"); + let model_id: Uuid = model.get("id"); + let model_name: String = model.get("model_name"); + drop(client); + + let request = RecordUsageRequest { + organization_id: org.id.parse().expect("organization UUID"), + workspace_id: workspace.id.parse().expect("workspace UUID"), + api_key_id: api_key.id.parse().expect("API key UUID"), + model_id, + model_name, + input_tokens: 1, + output_tokens: 1, + input_cost: 1_000, + output_cost: 2_000, + total_cost: 3_000, + inference_type: InferenceType::ChatCompletion.as_str().to_string(), + ttft_ms: None, + avg_itl_ms: None, + inference_id: Some(Uuid::new_v4()), + provider_request_id: Some("spend-counter-test".to_string()), + stop_reason: None, + response_id: None, + image_count: None, + cache_read_tokens: 0, + cache_write_tokens: 0, + billing_details: None, + service_tier: None, + context_band: None, + served_provider_tier: None, + served_provider_type: None, + served_via_fallback: false, + }; + ( + server, + database, + request, + api_key.id.parse().expect("API key UUID"), + ) +} + +#[tokio::test] +async fn spend_counters_track_inference_and_service_separately_on_one_key() -> anyhow::Result<()> { + let (_server, database, inference, api_key_id) = fixture().await; + let inference_repository = OrganizationUsageRepository::new(database.pool().clone()); + let service_repository = OrganizationServiceUsageRepository::new(database.pool().clone()); + let service_id = Uuid::new_v4(); + let client = database.pool().get().await?; + client + .execute( + "INSERT INTO services (id, service_name, display_name, unit, cost_per_unit) VALUES ($1, $2, $2, 'request', 5000)", + &[&service_id, &format!("spend-counter-service-{service_id}")], + ) + .await?; + drop(client); + + let service = RecordServiceUsageRequest { + organization_id: inference.organization_id, + workspace_id: inference.workspace_id, + api_key_id, + service_id, + quantity: 1, + total_cost: 5_000, + inference_id: Some(Uuid::new_v4()), + }; + service_repository.record_usage(&service).await?; + service_repository.record_usage(&service).await?; + let mut second_service = service.clone(); + second_service.inference_id = Some(Uuid::new_v4()); + second_service.total_cost = 2_000; + service_repository.record_usage(&second_service).await?; + inference_repository.record_usage(inference.clone()).await?; + let mut second_inference = inference.clone(); + second_inference.inference_id = Some(Uuid::new_v4()); + second_inference.provider_request_id = Some("spend-counter-second-inference".to_string()); + second_inference.total_cost = 1_500; + second_inference.input_cost = 1_500; + second_inference.output_cost = 0; + inference_repository.record_usage(second_inference).await?; + + let client = database.pool().get().await?; + let counters = client + .query_one( + "SELECT inference_spent, service_spent FROM api_key_spend WHERE api_key_id = $1", + &[&api_key_id], + ) + .await?; + assert_eq!(counters.get::<_, i64>("inference_spent"), 4_500); + assert_eq!(counters.get::<_, i64>("service_spent"), 7_000); + let balance = client + .query_one( + "SELECT inference_spent, service_spent FROM organization_balance WHERE organization_id = $1", + &[&inference.organization_id], + ) + .await?; + assert_eq!(balance.get::<_, i64>("inference_spent"), 4_500); + assert_eq!(balance.get::<_, i64>("service_spent"), 7_000); + Ok(()) +} + +#[tokio::test] +async fn spend_counters_ignore_duplicates_and_failed_posts() -> anyhow::Result<()> { + let (_server, database, inference, api_key_id) = fixture().await; + let repository = OrganizationUsageRepository::new(database.pool().clone()); + repository.record_usage(inference.clone()).await?; + repository.record_usage(inference.clone()).await?; + + let client = database.pool().get().await?; + client + .execute( + "UPDATE api_key_spend SET inference_spent = $2 WHERE api_key_id = $1", + &[&api_key_id, &i64::MAX], + ) + .await?; + drop(client); + + let overflow_id = Uuid::new_v4(); + let mut overflow = inference.clone(); + overflow.inference_id = Some(overflow_id); + overflow.provider_request_id = Some("spend-counter-overflow".to_string()); + overflow.total_cost = 1; + overflow.input_cost = 1; + overflow.output_cost = 0; + let error = repository + .record_usage(overflow) + .await + .expect_err("a counter overflow must roll back the usage transaction"); + let services::common::RepositoryError::DatabaseError(cause) = error + .downcast::() + .expect("repository error") + else { + panic!("expected a database overflow error"); + }; + assert_eq!( + cause.to_string(), + "Database error (22003): bigint out of range" + ); + + let client = database.pool().get().await?; + let key = client + .query_one( + "SELECT inference_spent, service_spent FROM api_key_spend WHERE api_key_id = $1", + &[&api_key_id], + ) + .await?; + assert_eq!(key.get::<_, i64>("inference_spent"), i64::MAX); + assert_eq!(key.get::<_, i64>("service_spent"), 0); + let usage_count: i64 = client + .query_one( + "SELECT COUNT(*)::BIGINT FROM organization_usage_log WHERE inference_id = $1", + &[&inference.inference_id], + ) + .await? + .get(0); + assert_eq!(usage_count, 1); + let overflow_count: i64 = client + .query_one( + "SELECT COUNT(*)::BIGINT FROM organization_usage_log WHERE inference_id = $1", + &[&overflow_id], + ) + .await? + .get(0); + assert_eq!(overflow_count, 0); + let balance = client + .query_one( + "SELECT total_spent, inference_spent, service_spent, unresolved_unfunded_amount FROM organization_balance WHERE organization_id = $1", + &[&inference.organization_id], + ) + .await?; + assert_eq!(balance.get::<_, i64>("total_spent"), 3_000); + assert_eq!(balance.get::<_, i64>("inference_spent"), 3_000); + assert_eq!(balance.get::<_, i64>("service_spent"), 0); + assert_eq!(balance.get::<_, i64>("unresolved_unfunded_amount"), 3_000); + client + .execute( + "UPDATE api_key_spend SET inference_spent = $2 WHERE api_key_id = $1", + &[&api_key_id, &3_000i64], + ) + .await?; + Ok(()) +} diff --git a/crates/database/src/migrations/sql/V0081__add_spend_counters.sql b/crates/database/src/migrations/sql/V0081__add_spend_counters.sql new file mode 100644 index 000000000..93162c19f --- /dev/null +++ b/crates/database/src/migrations/sql/V0081__add_spend_counters.sql @@ -0,0 +1,19 @@ +-- Persist split inference/service spend totals for bounded analytics reads. +-- All values are integer nano-dollars (scale 9, USD). +-- Existing rows are initialized to zero; an out-of-band backfill must run before +-- these columns are used as lifetime totals. + +ALTER TABLE organization_balance + ADD COLUMN inference_spent BIGINT NOT NULL DEFAULT 0, + ADD COLUMN service_spent BIGINT NOT NULL DEFAULT 0; + +CREATE TABLE api_key_spend ( + api_key_id UUID PRIMARY KEY REFERENCES api_keys(id) ON DELETE CASCADE, + inference_spent BIGINT NOT NULL DEFAULT 0, + service_spent BIGINT NOT NULL DEFAULT 0, + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +COMMENT ON TABLE api_key_spend IS 'Cached per-API-key inference and service spending totals in nano-dollars'; +COMMENT ON COLUMN api_key_spend.inference_spent IS 'Cumulative inference spending in nano-dollars (scale 9, USD)'; +COMMENT ON COLUMN api_key_spend.service_spent IS 'Cumulative service spending in nano-dollars (scale 9, USD)'; diff --git a/crates/database/src/repositories/mod.rs b/crates/database/src/repositories/mod.rs index fa76a8de3..6f95ae1db 100644 --- a/crates/database/src/repositories/mod.rs +++ b/crates/database/src/repositories/mod.rs @@ -34,6 +34,7 @@ pub mod retry; pub mod service; pub mod service_usage_repository_impl; pub mod session; +pub(crate) mod spend_counters; pub mod usage_repository_impl; pub mod user; pub mod utils; diff --git a/crates/database/src/repositories/organization_service_usage.rs b/crates/database/src/repositories/organization_service_usage.rs index 8784676ff..8f4cfe4a9 100644 --- a/crates/database/src/repositories/organization_service_usage.rs +++ b/crates/database/src/repositories/organization_service_usage.rs @@ -4,6 +4,7 @@ use crate::repositories::credit_allocation::{ allocate_usage, load_allocations, lock_organization_accounting, CreditAllocationPolicy, UsageAllocationParent, }; +use crate::repositories::spend_counters::increment_api_key_spend; use crate::repositories::utils::map_db_error; use crate::retry_db; use anyhow::{Context, Result}; @@ -363,22 +364,27 @@ impl OrganizationServiceUsageRepository { .execute( r#" INSERT INTO organization_balance ( - organization_id, total_spent, last_usage_at, total_requests, total_tokens, updated_at - ) VALUES ($1, $2, $3, 0, 0, $4) + organization_id, total_spent, inference_spent, service_spent, + last_usage_at, total_requests, total_tokens, updated_at + ) VALUES ($1, $2, 0, $2, $3, 0, 0, $4) ON CONFLICT (organization_id) DO UPDATE SET total_spent = organization_balance.total_spent + $2, + service_spent = organization_balance.service_spent + $2, last_usage_at = $3, updated_at = $4 "#, - &[ - &request.organization_id, - &request.total_cost, - &now, - &now, - ], + &[&request.organization_id, &request.total_cost, &now, &now], ) .await .map_err(map_db_error)?; + increment_api_key_spend( + &transaction, + request.api_key_id, + 0, + request.total_cost, + now, + ) + .await?; transaction.commit().await.map_err(map_db_error)?; (r, Some(allocation.allocations)) diff --git a/crates/database/src/repositories/organization_usage.rs b/crates/database/src/repositories/organization_usage.rs index 3ec9cbc18..03194285c 100644 --- a/crates/database/src/repositories/organization_usage.rs +++ b/crates/database/src/repositories/organization_usage.rs @@ -7,6 +7,7 @@ use crate::repositories::credit_allocation::{ allocate_usage, load_allocations, lock_organization_accounting, CreditAllocationPolicy, UsageAllocationParent, }; +use crate::repositories::spend_counters::increment_api_key_spend; use crate::repositories::utils::map_db_error; use crate::retry_db; use anyhow::{Context, Result}; @@ -200,13 +201,16 @@ impl OrganizationUsageRepository { INSERT INTO organization_balance ( organization_id, total_spent, + inference_spent, + service_spent, last_usage_at, total_requests, total_tokens, updated_at - ) VALUES ($1, $2, $3, 1, $4, $5) + ) VALUES ($1, $2, $2, 0, $3, 1, $4, $5) ON CONFLICT (organization_id) DO UPDATE SET total_spent = organization_balance.total_spent + $2, + inference_spent = organization_balance.inference_spent + $2, total_requests = organization_balance.total_requests + 1, total_tokens = organization_balance.total_tokens + $4, last_usage_at = $3, @@ -222,6 +226,14 @@ impl OrganizationUsageRepository { ) .await .map_err(map_db_error)?; + increment_api_key_spend( + &transaction, + request.api_key_id, + request.total_cost, + 0, + now, + ) + .await?; transaction.commit().await.map_err(map_db_error)?; (row, true, Some(allocation.allocations)) diff --git a/crates/database/src/repositories/spend_counters.rs b/crates/database/src/repositories/spend_counters.rs new file mode 100644 index 000000000..79f364e4d --- /dev/null +++ b/crates/database/src/repositories/spend_counters.rs @@ -0,0 +1,31 @@ +use crate::repositories::utils::map_db_error; +use chrono::{DateTime, Utc}; +use services::common::RepositoryError; +use tokio_postgres::Transaction; +use uuid::Uuid; + +/// Add one usage charge to the per-key split counters inside its posting transaction. +pub async fn increment_api_key_spend( + transaction: &Transaction<'_>, + api_key_id: Uuid, + inference_spent: i64, + service_spent: i64, + updated_at: DateTime, +) -> Result<(), RepositoryError> { + transaction + .execute( + r#" + INSERT INTO api_key_spend ( + api_key_id, inference_spent, service_spent, updated_at + ) VALUES ($1, $2, $3, $4) + ON CONFLICT (api_key_id) DO UPDATE SET + inference_spent = api_key_spend.inference_spent + $2, + service_spent = api_key_spend.service_spent + $3, + updated_at = $4 + "#, + &[&api_key_id, &inference_spent, &service_spent, &updated_at], + ) + .await + .map_err(map_db_error)?; + Ok(()) +} From 4cd053eeb0c48762806a90bed4dc8e8ec716b60f Mon Sep 17 00:00:00 2001 From: Henry Park <16583448+henrypark133@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:49:40 +0000 Subject: [PATCH 2/7] feat: reconcile historical spend counters before reader rollout --- .config/nextest.toml | 4 + .github/workflows/test.yml | 2 +- Dockerfile | 3 +- .../src/bin/backfill-spend-counters.rs | 128 ++++ crates/database/src/lib.rs | 4 + .../V0082__add_spend_counters_readiness.sql | 12 + .../database/src/spend_counters_backfill.rs | 399 +++++++++++ .../database/tests/spend_counter_backfill.rs | 660 ++++++++++++++++++ 8 files changed, 1210 insertions(+), 2 deletions(-) create mode 100644 crates/database/src/bin/backfill-spend-counters.rs create mode 100644 crates/database/src/migrations/sql/V0082__add_spend_counters_readiness.sql create mode 100644 crates/database/src/spend_counters_backfill.rs create mode 100644 crates/database/tests/spend_counter_backfill.rs diff --git a/.config/nextest.toml b/.config/nextest.toml index 41628550e..95747b942 100644 --- a/.config/nextest.toml +++ b/.config/nextest.toml @@ -47,6 +47,10 @@ threads-required = "num-test-threads" filter = "package(api) & binary(e2e_all)" test-group = "e2e-db" +[[profile.default.overrides]] +filter = "package(database) & binary(spend_counter_backfill)" +test-group = "e2e-db" + [[profile.default.scripts]] filter = "package(api) & binary(e2e_all)" setup = "e2e-db-bootstrap" diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 82e9a44ef..0f45afc51 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -199,7 +199,7 @@ jobs: cache-key: e2e - name: Run e2e tests - run: cargo nextest run --test e2e_all + run: cargo nextest run --test e2e_all --test spend_counter_backfill env: POSTGRES_PRIMARY_APP_ID: ${{ secrets.POSTGRES_PRIMARY_APP_ID }} DATABASE_HOST: localhost diff --git a/Dockerfile b/Dockerfile index cdcd4d847..cf00fcb4a 100644 --- a/Dockerfile +++ b/Dockerfile @@ -40,7 +40,7 @@ COPY crates/ ./crates/ COPY .cargo/ ./.cargo/ # Build the application in release mode -RUN cargo build --release --locked --bin api +RUN cargo build --release --locked --bin api --bin backfill-spend-counters # Runtime stage @@ -86,6 +86,7 @@ WORKDIR /app # Copy the built binary COPY --from=builder /app/target/release/api /app/api +COPY --from=builder /app/target/release/backfill-spend-counters /app/backfill-spend-counters # Copy the migration SQL files RUN mkdir -p /app/crates/database/src/migrations/sql diff --git a/crates/database/src/bin/backfill-spend-counters.rs b/crates/database/src/bin/backfill-spend-counters.rs new file mode 100644 index 000000000..8eebffc71 --- /dev/null +++ b/crates/database/src/bin/backfill-spend-counters.rs @@ -0,0 +1,128 @@ +use anyhow::{bail, Context, Result}; +use database::spend_counters_backfill::{ + incomplete_organizations, PreparedSpendBackfill, SpendBackfillOutcome, +}; +use database::Database; +use std::time::Duration; +use uuid::Uuid; + +const BATCH_SIZE: i64 = 100; +const DEFAULT_TIMEOUT_SECONDS: u64 = 300; + +#[derive(Debug)] +struct Options { + organization_id: Option, + statement_timeout: Duration, +} + +#[tokio::main] +async fn main() -> Result<()> { + let options = match parse_args(std::env::args().skip(1).collect())? { + Some(options) => options, + None => return Ok(()), + }; + tracing_subscriber::fmt() + .compact() + .with_target(false) + .with_thread_ids(false) + .with_thread_names(false) + .init(); + + let config = config::DatabaseConfig::from_env() + .map_err(|error| anyhow::anyhow!("failed to load database configuration: {error}"))?; + let database = Database::from_config(&config) + .await + .context("connecting to database")?; + tracing::info!( + statement_timeout_seconds = options.statement_timeout.as_secs(), + "starting spend counter backfill; all writers must maintain counters and raw history must not be rewritten" + ); + + if let Some(organization_id) = options.organization_id { + reconcile_one(&database, organization_id, options.statement_timeout).await?; + } else { + let mut after = None; + loop { + let ids = incomplete_organizations(database.pool(), after, BATCH_SIZE).await?; + if ids.is_empty() { + break; + } + after = ids.last().copied(); + for organization_id in ids { + reconcile_one(&database, organization_id, options.statement_timeout).await?; + } + } + database::ensure_spend_counters_ready(database.pool()).await?; + } + tracing::info!("spend counter backfill complete"); + Ok(()) +} + +async fn reconcile_one( + database: &Database, + organization_id: Uuid, + statement_timeout: Duration, +) -> Result<()> { + let Some(prepared) = + PreparedSpendBackfill::prepare(database.pool(), organization_id, statement_timeout).await? + else { + tracing::info!(%organization_id, "spend counters already ready"); + return Ok(()); + }; + match prepared.apply().await? { + SpendBackfillOutcome::Applied { key_count } => { + tracing::info!(%organization_id, key_count, "reconciled spend counters"); + } + SpendBackfillOutcome::AlreadyComplete => { + tracing::info!(%organization_id, "spend counters completed by another worker"); + } + } + Ok(()) +} + +fn parse_args(args: Vec) -> Result> { + if args.iter().any(|arg| arg == "--help" || arg == "-h") { + print_help(); + return Ok(None); + } + let mut organization_id = None; + let mut timeout_seconds = DEFAULT_TIMEOUT_SECONDS; + let mut args = args.into_iter(); + while let Some(arg) = args.next() { + match arg.as_str() { + "--organization" => { + let value = args.next().context("--organization needs a UUID")?; + organization_id = Some( + value + .parse() + .with_context(|| format!("invalid organization UUID: {value}"))?, + ); + } + "--statement-timeout-seconds" => { + let value = args + .next() + .context("--statement-timeout-seconds needs an integer")?; + timeout_seconds = value + .parse() + .with_context(|| format!("invalid timeout seconds: {value}"))?; + if timeout_seconds == 0 { + bail!("--statement-timeout-seconds must be positive"); + } + } + unknown => bail!("unknown argument {unknown}; use --help"), + } + } + Ok(Some(Options { + organization_id, + statement_timeout: Duration::from_secs(timeout_seconds), + })) +} + +fn print_help() { + println!( + "Usage: backfill-spend-counters [--organization UUID] [--statement-timeout-seconds N]\n\n\ +Reconciles incomplete organization spend counters. Deploy counter writers (#1116)\n\ +to every process and drain all older writers first. Do not rewrite or delete raw usage history\n\ +while this process runs. The command does not run migrations." + ); +} diff --git a/crates/database/src/lib.rs b/crates/database/src/lib.rs index 503970ce2..93af5ce41 100644 --- a/crates/database/src/lib.rs +++ b/crates/database/src/lib.rs @@ -8,6 +8,7 @@ pub mod patroni_discovery; pub mod pool; pub mod repositories; pub mod shutdown_coordinator; +pub mod spend_counters_backfill; mod usage_reporting_indexes; pub use constants::*; @@ -21,6 +22,9 @@ pub use repositories::{ SessionRepository, UserRepository, }; pub use shutdown_coordinator::{ShutdownCoordinator, ShutdownStage, ShutdownStageResult}; +pub use spend_counters_backfill::{ + ensure_spend_counters_ready, PreparedSpendBackfill, SpendBackfillOutcome, +}; pub use usage_reporting_indexes::ensure_usage_reporting_indexes; use anyhow::Result; diff --git a/crates/database/src/migrations/sql/V0082__add_spend_counters_readiness.sql b/crates/database/src/migrations/sql/V0082__add_spend_counters_readiness.sql new file mode 100644 index 000000000..bf8c1e26f --- /dev/null +++ b/crates/database/src/migrations/sql/V0082__add_spend_counters_readiness.sql @@ -0,0 +1,12 @@ +-- Mark whether an organization's split spend counters include its historical logs. +-- Existing organizations remain NULL until the backfill operator completes them. +-- New balance rows are complete by construction after all usage writers maintain spend counters. + +ALTER TABLE organization_balance + ADD COLUMN spend_counters_ready_at TIMESTAMPTZ; + +ALTER TABLE organization_balance + ALTER COLUMN spend_counters_ready_at SET DEFAULT NOW(); + +COMMENT ON COLUMN organization_balance.spend_counters_ready_at IS + 'Non-NULL after the spend counter historical reconciliation completed for this organization'; diff --git a/crates/database/src/spend_counters_backfill.rs b/crates/database/src/spend_counters_backfill.rs new file mode 100644 index 000000000..9ee2311d1 --- /dev/null +++ b/crates/database/src/spend_counters_backfill.rs @@ -0,0 +1,399 @@ +use crate::repositories::credit_allocation::lock_organization_accounting; +use crate::DbPool; +use anyhow::{bail, Context, Result}; +use chrono::{DateTime, Utc}; +use std::time::Duration; +use tokio_postgres::{IsolationLevel, Transaction}; +use uuid::Uuid; + +const LOCK_TIMEOUT: Duration = Duration::from_secs(5); +const APPLY_TIMEOUT: Duration = Duration::from_secs(5); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct KeyCorrection { + api_key_id: Uuid, + inference_spent: i64, + service_spent: i64, +} + +/// A repeatable-read snapshot of one incomplete organization's correction. +/// +/// The fields are intentionally private: callers can only obtain corrections +/// from the database snapshot and apply them through the guarded method below. +#[derive(Debug)] +pub struct PreparedSpendBackfill { + pool: DbPool, + organization_id: Uuid, + key_corrections: Vec, + inference_spent: i64, + service_spent: i64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpendBackfillOutcome { + Applied { key_count: usize }, + AlreadyComplete, +} + +impl PreparedSpendBackfill { + /// Capture one organization's raw/counter delta without holding its + /// accounting lock. The snapshot transaction commits before apply. + pub async fn prepare( + pool: &DbPool, + organization_id: Uuid, + statement_timeout: Duration, + ) -> Result> { + validate_timeout(statement_timeout)?; + let mut client = pool + .get() + .await + .context("acquiring database connection for spend snapshot")?; + let transaction = client + .build_transaction() + .isolation_level(IsolationLevel::RepeatableRead) + .read_only(true) + .start() + .await + .context("starting repeatable-read spend snapshot")?; + set_timeout(&transaction, "statement_timeout", statement_timeout).await?; + + let balance = transaction + .query_opt( + "SELECT spend_counters_ready_at, inference_spent, service_spent + FROM organization_balance WHERE organization_id = $1", + &[&organization_id], + ) + .await + .context("reading organization spend readiness")?; + let Some(balance) = balance else { + if transaction + .query_opt( + "SELECT id FROM organizations WHERE id = $1", + &[&organization_id], + ) + .await + .context("checking organization existence")? + .is_some() + { + bail!("organization {organization_id} has no organization_balance row"); + } + bail!("organization {organization_id} does not exist"); + }; + let ready_at: Option> = balance.get("spend_counters_ready_at"); + if ready_at.is_some() { + transaction + .commit() + .await + .context("committing already-complete spend snapshot")?; + return Ok(None); + } + let counted_inference: i64 = balance.get("inference_spent"); + let counted_service: i64 = balance.get("service_spent"); + + let rows = transaction + .query( + r#" + WITH raw_by_key AS ( + SELECT api_key_id, + SUM(total_cost) AS inference_spent, + 0::NUMERIC AS service_spent + FROM organization_usage_log + WHERE organization_id = $1 + GROUP BY api_key_id + UNION ALL + SELECT api_key_id, + 0::NUMERIC AS inference_spent, + SUM(total_cost) AS service_spent + FROM organization_service_usage_log + WHERE organization_id = $1 + GROUP BY api_key_id + ), + raw AS ( + SELECT api_key_id, + SUM(inference_spent) AS inference_spent, + SUM(service_spent) AS service_spent + FROM raw_by_key + GROUP BY api_key_id + ), + counted AS ( + SELECT spend.api_key_id, spend.inference_spent, spend.service_spent + FROM api_key_spend AS spend + JOIN api_keys AS key ON key.id = spend.api_key_id + JOIN workspaces AS workspace ON workspace.id = key.workspace_id + WHERE workspace.organization_id = $1 + ) + SELECT COALESCE(raw.api_key_id, counted.api_key_id) AS api_key_id, + COALESCE(raw.inference_spent, 0)::TEXT AS raw_inference_spent, + COALESCE(raw.service_spent, 0)::TEXT AS raw_service_spent, + COALESCE(counted.inference_spent, 0) AS counted_inference_spent, + COALESCE(counted.service_spent, 0) AS counted_service_spent + FROM raw + FULL OUTER JOIN counted ON counted.api_key_id = raw.api_key_id + ORDER BY api_key_id + "#, + &[&organization_id], + ) + .await + .context("reading raw and counted spend by API key")?; + + let mut raw_inference = 0_i64; + let mut raw_service = 0_i64; + let mut key_corrections = Vec::with_capacity(rows.len()); + for row in rows { + let api_key_id: Uuid = row.get("api_key_id"); + let raw_key_inference = parse_total(&row.get::<_, String>("raw_inference_spent"))?; + let raw_key_service = parse_total(&row.get::<_, String>("raw_service_spent"))?; + raw_inference = raw_inference + .checked_add(raw_key_inference) + .context("raw inference spend exceeds BIGINT")?; + raw_service = raw_service + .checked_add(raw_key_service) + .context("raw service spend exceeds BIGINT")?; + let inference_correction = checked_correction( + raw_key_inference, + row.get("counted_inference_spent"), + api_key_id, + "inference", + )?; + let service_correction = checked_correction( + raw_key_service, + row.get("counted_service_spent"), + api_key_id, + "service", + )?; + if inference_correction != 0 || service_correction != 0 { + key_corrections.push(KeyCorrection { + api_key_id, + inference_spent: inference_correction, + service_spent: service_correction, + }); + } + } + let inference_spent = checked_correction( + raw_inference, + counted_inference, + organization_id, + "organization inference", + )?; + let service_spent = checked_correction( + raw_service, + counted_service, + organization_id, + "organization service", + )?; + + transaction + .commit() + .await + .context("committing spend snapshot before accounting lock")?; + Ok(Some(Self { + pool: pool.clone(), + organization_id, + key_corrections, + inference_spent, + service_spent, + })) + } + + /// Apply the prepared delta under the existing organization accounting lock. + /// The completion marker and every counter correction commit atomically. + pub async fn apply(self) -> Result { + let mut client = self + .pool + .get() + .await + .context("acquiring database connection for spend apply")?; + let transaction = client + .transaction() + .await + .context("starting spend apply transaction")?; + set_timeout(&transaction, "statement_timeout", APPLY_TIMEOUT).await?; + set_timeout(&transaction, "lock_timeout", LOCK_TIMEOUT).await?; + lock_organization_accounting(&transaction, self.organization_id) + .await + .map_err(|error| anyhow::anyhow!(error))?; + + let ready_at: Option> = transaction + .query_one( + "SELECT spend_counters_ready_at FROM organization_balance WHERE organization_id = $1", + &[&self.organization_id], + ) + .await + .context("rechecking organization spend readiness")? + .get(0); + if ready_at.is_some() { + transaction + .commit() + .await + .context("committing already-complete spend apply")?; + return Ok(SpendBackfillOutcome::AlreadyComplete); + } + + if !self.key_corrections.is_empty() { + let key_ids: Vec = self + .key_corrections + .iter() + .map(|correction| correction.api_key_id) + .collect(); + let inference: Vec = self + .key_corrections + .iter() + .map(|correction| correction.inference_spent) + .collect(); + let service: Vec = self + .key_corrections + .iter() + .map(|correction| correction.service_spent) + .collect(); + transaction + .execute( + r#" + INSERT INTO api_key_spend ( + api_key_id, inference_spent, service_spent, updated_at + ) + SELECT key_id, inference, service, NOW() + FROM UNNEST($1::UUID[], $2::BIGINT[], $3::BIGINT[]) + AS correction(key_id, inference, service) + ON CONFLICT (api_key_id) DO UPDATE SET + inference_spent = api_key_spend.inference_spent + EXCLUDED.inference_spent, + service_spent = api_key_spend.service_spent + EXCLUDED.service_spent, + updated_at = NOW() + "#, + &[&key_ids, &inference, &service], + ) + .await + .context("applying API-key spend corrections")?; + } + let updated = transaction + .execute( + r#" + UPDATE organization_balance + SET inference_spent = inference_spent + $2, + service_spent = service_spent + $3, + spend_counters_ready_at = NOW(), + updated_at = NOW() + WHERE organization_id = $1 AND spend_counters_ready_at IS NULL + "#, + &[ + &self.organization_id, + &self.inference_spent, + &self.service_spent, + ], + ) + .await + .context("marking organization spend counters ready")?; + if updated != 1 { + bail!( + "organization {} disappeared before spend readiness update", + self.organization_id + ); + } + let key_count = self.key_corrections.len(); + transaction + .commit() + .await + .context("committing spend counter corrections")?; + Ok(SpendBackfillOutcome::Applied { key_count }) + } +} + +/// Fail unless every organization, including inactive ones, has a balance row +/// with a completed historical spend snapshot. +pub async fn ensure_spend_counters_ready(pool: &DbPool) -> Result<()> { + let client = pool + .get() + .await + .context("acquiring database connection for spend readiness check")?; + let row = client + .query_one( + r#" + SELECT + COUNT(*) FILTER (WHERE balance.organization_id IS NULL)::BIGINT, + COUNT(*) FILTER ( + WHERE balance.organization_id IS NOT NULL + AND balance.spend_counters_ready_at IS NULL + )::BIGINT + FROM organizations AS organization + LEFT JOIN organization_balance AS balance + ON balance.organization_id = organization.id + "#, + &[], + ) + .await + .context("checking spend counter readiness")?; + let missing_balances: i64 = row.get(0); + let incomplete: i64 = row.get(1); + if missing_balances != 0 || incomplete != 0 { + bail!( + "spend counters are incomplete: {missing_balances} organizations lack balances, {incomplete} remain unreconciled" + ); + } + Ok(()) +} + +pub async fn incomplete_organizations( + pool: &DbPool, + after: Option, + limit: i64, +) -> Result> { + let client = pool + .get() + .await + .context("acquiring database connection for spend organization list")?; + let rows = client + .query( + r#" + SELECT organization.id + FROM organizations AS organization + LEFT JOIN organization_balance AS balance + ON balance.organization_id = organization.id + WHERE ($1::UUID IS NULL OR organization.id > $1) + AND (balance.organization_id IS NULL OR balance.spend_counters_ready_at IS NULL) + ORDER BY organization.id + LIMIT $2 + "#, + &[&after, &limit], + ) + .await + .context("listing incomplete spend organizations")?; + Ok(rows.into_iter().map(|row| row.get(0)).collect()) +} + +async fn set_timeout( + transaction: &Transaction<'_>, + setting: &str, + timeout: Duration, +) -> Result<()> { + let value = format!("{}ms", timeout.as_millis()); + transaction + .query_one("SELECT set_config($1, $2, true)", &[&setting, &value]) + .await + .with_context(|| format!("setting {setting}"))?; + Ok(()) +} + +fn validate_timeout(timeout: Duration) -> Result<()> { + if timeout.as_millis() == 0 || timeout.as_millis() > i32::MAX as u128 { + bail!( + "spend snapshot timeout must be between 1ms and {}ms", + i32::MAX + ); + } + Ok(()) +} + +fn parse_total(value: &str) -> Result { + value + .parse::() + .with_context(|| format!("spend total {value} exceeds BIGINT")) +} + +fn checked_correction(raw: i64, counted: i64, id: Uuid, kind: &str) -> Result { + let correction = raw + .checked_sub(counted) + .with_context(|| format!("{kind} correction overflow for {id}"))?; + if correction < 0 { + bail!("{kind} counter exceeds raw history for {id}: raw={raw}, counted={counted}"); + } + Ok(correction) +} diff --git a/crates/database/tests/spend_counter_backfill.rs b/crates/database/tests/spend_counter_backfill.rs new file mode 100644 index 000000000..a5e47c6b2 --- /dev/null +++ b/crates/database/tests/spend_counter_backfill.rs @@ -0,0 +1,660 @@ +use chrono::Utc; +use database::models::RecordUsageRequest; +use database::repositories::{ + OrganizationServiceUsageRepository, OrganizationUsageRepository, RecordServiceUsageRequest, +}; +use database::{ensure_spend_counters_ready, migrations, DbPool, PreparedSpendBackfill}; +use deadpool::Runtime; +use deadpool_postgres::{Config, PoolConfig, Timeouts}; +use services::usage::InferenceType; +use std::time::Duration; +use tokio_postgres::NoTls; +use uuid::Uuid; + +struct TestDatabase { + pool: DbPool, + admin: DbPool, + database_name: String, +} + +struct Fixture { + organization_id: Uuid, + mixed_key: Uuid, + service_key: Uuid, + deleted_key: Uuid, + empty_key: Uuid, + model_id: Uuid, + service_id: Uuid, + workspace_id: Uuid, +} + +#[tokio::test] +async fn backfill_reconciles_snapshot_delta_and_is_idempotent() -> anyhow::Result<()> { + let database = test_database().await?; + let fixture = fixture(&database.pool).await?; + + // Historical rows exceed the C1 counters already present for the mixed key. + insert_inference(&database.pool, &fixture, fixture.mixed_key, 100).await?; + insert_service(&database.pool, &fixture, fixture.mixed_key, 40).await?; + insert_service(&database.pool, &fixture, fixture.service_key, 60).await?; + insert_service(&database.pool, &fixture, fixture.deleted_key, 10).await?; + let client = database.pool.get().await?; + client + .execute( + "INSERT INTO api_key_spend (api_key_id, inference_spent, service_spent) + VALUES ($1, 20, 10)", + &[&fixture.mixed_key], + ) + .await?; + client + .execute( + "UPDATE organization_balance SET total_spent = 777, inference_spent = 20, service_spent = 10, + spend_counters_ready_at = NULL WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + drop(client); + + // Prepare twice from the same snapshot, then post through C1 before applying. + let first = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await? + .expect("fixture organization is incomplete"); + let second = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await? + .expect("second snapshot is also incomplete"); + let inference_repository = OrganizationUsageRepository::new(database.pool.clone()); + let service_repository = OrganizationServiceUsageRepository::new(database.pool.clone()); + inference_repository + .record_usage(inference_request(&fixture, 5)) + .await?; + service_repository + .record_usage(&RecordServiceUsageRequest { + organization_id: fixture.organization_id, + workspace_id: fixture.workspace_id, + api_key_id: fixture.mixed_key, + service_id: fixture.service_id, + quantity: 1, + total_cost: 7, + inference_id: Some(Uuid::new_v4()), + }) + .await?; + + let (first_result, second_result) = tokio::join!(first.apply(), second.apply()); + let outcomes = [first_result?, second_result?]; + assert!(outcomes.iter().any(|outcome| matches!( + outcome, + database::SpendBackfillOutcome::Applied { key_count: 3 } + ))); + assert!(outcomes.contains(&database::SpendBackfillOutcome::AlreadyComplete)); + assert!(PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30) + ) + .await? + .is_none()); + + let client = database.pool.get().await?; + let row = client + .query_one( + "SELECT total_spent, inference_spent, service_spent, spend_counters_ready_at + FROM organization_balance WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + assert_eq!(row.get::<_, i64>("total_spent"), 789); + assert_eq!(row.get::<_, i64>("inference_spent"), 105); + assert_eq!(row.get::<_, i64>("service_spent"), 117); + assert!(row + .get::<_, Option>>("spend_counters_ready_at") + .is_some()); + let rows = client + .query( + "SELECT api_key_id, inference_spent, service_spent FROM api_key_spend + WHERE api_key_id = ANY($1) ORDER BY api_key_id", + &[&vec![ + fixture.mixed_key, + fixture.service_key, + fixture.deleted_key, + ]], + ) + .await?; + assert_eq!(rows.len(), 3); + let totals: Vec<(Uuid, i64, i64)> = rows + .iter() + .map(|row| (row.get(0), row.get(1), row.get(2))) + .collect(); + assert!(totals.contains(&(fixture.mixed_key, 105, 47))); + assert!(totals.contains(&(fixture.service_key, 0, 60))); + assert!(totals.contains(&(fixture.deleted_key, 0, 10))); + assert_eq!( + client + .query_one( + "SELECT inference_spent, service_spent FROM api_key_spend WHERE api_key_id = $1", + &[&fixture.mixed_key], + ) + .await? + .get::<_, i64>(0), + 105 + ); + drop(client); + drop(inference_repository); + drop(service_repository); + database.cleanup().await?; + Ok(()) +} + +#[tokio::test] +async fn backfill_rejects_negative_drift_and_rolls_back_overflow() -> anyhow::Result<()> { + let database = test_database().await?; + let fixture = fixture(&database.pool).await?; + let client = database.pool.get().await?; + client + .execute( + "INSERT INTO api_key_spend (api_key_id, inference_spent) + VALUES ($1, 1)", + &[&fixture.empty_key], + ) + .await?; + client + .execute( + "UPDATE organization_balance SET spend_counters_ready_at = NULL + WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + drop(client); + let error = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await + .expect_err("counter-only spend must reject negative correction"); + assert!(error.to_string().contains("counter exceeds raw history")); + + let client = database.pool.get().await?; + client + .execute( + "DELETE FROM api_key_spend WHERE api_key_id = $1", + &[&fixture.empty_key], + ) + .await?; + client + .execute( + "INSERT INTO organization_usage_log ( + id, organization_id, workspace_id, api_key_id, model_id, model_name, + input_tokens, output_tokens, total_tokens, input_cost, output_cost, + total_cost, inference_type, inference_id + ) VALUES ($1, $2, $3, $4, $5, 'counter-test', 1, 0, 1, 1, 0, 1, + 'chat_completion', $1)", + &[ + &Uuid::new_v4(), + &fixture.organization_id, + &fixture.workspace_id, + &fixture.empty_key, + &fixture.model_id, + ], + ) + .await?; + client + .execute( + "UPDATE organization_balance SET spend_counters_ready_at = NULL + WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + drop(client); + let prepared = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await? + .expect("overflow fixture is incomplete"); + let client = database.pool.get().await?; + client + .execute( + "INSERT INTO api_key_spend (api_key_id, inference_spent) + VALUES ($1, $2)", + &[&fixture.empty_key, &i64::MAX], + ) + .await?; + drop(client); + let error = prepared + .apply() + .await + .expect_err("counter overflow must fail during the apply transaction"); + let pg_error = error + .chain() + .find_map(|cause| cause.downcast_ref::()) + .expect("overflow should preserve the PostgreSQL error cause"); + assert_eq!( + pg_error.code(), + Some(&tokio_postgres::error::SqlState::NUMERIC_VALUE_OUT_OF_RANGE) + ); + let client = database.pool.get().await?; + let row = client + .query_one( + "SELECT inference_spent, spend_counters_ready_at + FROM organization_balance WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + assert!(row.get::<_, Option>>(1).is_none()); + assert_eq!( + client + .query_one( + "SELECT inference_spent FROM api_key_spend WHERE api_key_id = $1", + &[&fixture.empty_key], + ) + .await? + .get::<_, i64>(0), + i64::MAX + ); + assert_eq!(row.get::<_, i64>(0), 0); + client + .execute( + "DELETE FROM api_key_spend WHERE api_key_id = $1", + &[&fixture.empty_key], + ) + .await?; + drop(client); + + // Fail after the key upsert to prove the entire transaction rolls back. + let prepared = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await? + .expect("failed apply must leave the organization retryable"); + let client = database.pool.get().await?; + client + .execute( + "UPDATE organization_balance SET inference_spent = $2 WHERE organization_id = $1", + &[&fixture.organization_id, &i64::MAX], + ) + .await?; + drop(client); + let error = prepared + .apply() + .await + .expect_err("balance overflow must roll back key corrections"); + let database_error = error + .chain() + .find_map(|cause| cause.downcast_ref::()) + .expect("preserved PostgreSQL error"); + assert_eq!( + database_error.code(), + Some(&tokio_postgres::error::SqlState::NUMERIC_VALUE_OUT_OF_RANGE) + ); + let client = database.pool.get().await?; + assert!(client + .query_opt( + "SELECT 1 FROM api_key_spend WHERE api_key_id = $1", + &[&fixture.empty_key] + ) + .await? + .is_none()); + let balance = client.query_one( + "SELECT inference_spent, spend_counters_ready_at FROM organization_balance WHERE organization_id = $1", + &[&fixture.organization_id], + ).await?; + assert_eq!(balance.get::<_, i64>(0), i64::MAX); + assert!(balance.get::<_, Option>>(1).is_none()); + drop(client); + database.cleanup().await?; + Ok(()) +} + +#[tokio::test] +async fn spend_readiness_checks_missing_and_inactive_organizations() -> anyhow::Result<()> { + let database = test_database().await?; + let fixture = fixture(&database.pool).await?; + let client = database.pool.get().await?; + client + .execute( + "UPDATE organization_balance SET spend_counters_ready_at = NULL + WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + drop(client); + let prepared = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await? + .expect("empty historical organization is incomplete"); + assert_eq!( + prepared.apply().await?, + database::SpendBackfillOutcome::Applied { key_count: 0 } + ); + ensure_spend_counters_ready(&database.pool).await?; + + let mut client = database.pool.get().await?; + let migration_tx = client.build_transaction().start().await?; + migration_tx + .batch_execute( + "CREATE TEMP TABLE organization_balance ( + organization_id UUID PRIMARY KEY, + inference_spent BIGINT NOT NULL DEFAULT 0, + service_spent BIGINT NOT NULL DEFAULT 0 + ); + INSERT INTO organization_balance (organization_id) VALUES (uuid_generate_v4());", + ) + .await?; + migration_tx + .batch_execute(include_str!( + "../src/migrations/sql/V0082__add_spend_counters_readiness.sql" + )) + .await?; + let legacy = migration_tx + .query_one( + "SELECT spend_counters_ready_at FROM organization_balance LIMIT 1", + &[], + ) + .await?; + assert!(legacy.get::<_, Option>>(0).is_none()); + let future = migration_tx + .query_one( + "INSERT INTO organization_balance (organization_id) + VALUES (uuid_generate_v4()) + RETURNING spend_counters_ready_at", + &[], + ) + .await?; + assert!(future.get::<_, Option>>(0).is_some()); + migration_tx.rollback().await?; + + client + .execute( + "UPDATE organizations SET is_active = false WHERE id = $1", + &[&fixture.organization_id], + ) + .await?; + ensure_spend_counters_ready(&database.pool).await?; + client + .execute( + "UPDATE organization_balance SET spend_counters_ready_at = NULL + WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + assert!(ensure_spend_counters_ready(&database.pool) + .await + .expect_err("inactive incomplete organizations must block readers") + .to_string() + .contains("remain unreconciled")); + client + .execute( + "DELETE FROM organization_balance WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + assert!(ensure_spend_counters_ready(&database.pool) + .await + .expect_err("missing balances must be distinguished") + .to_string() + .contains("lack balances")); + drop(client); + database.cleanup().await?; + Ok(()) +} + +async fn test_database() -> anyhow::Result { + let admin_config = pool_config(None); + let admin: DbPool = admin_config + .create_pool(Some(Runtime::Tokio1), NoTls)? + .into(); + let database_name = format!("spend_counter_backfill_{}", Uuid::new_v4().simple()); + admin + .get() + .await? + .batch_execute(&format!("CREATE DATABASE {database_name}")) + .await?; + let scoped_config = pool_config(Some(&database_name)); + let pool: DbPool = scoped_config + .create_pool(Some(Runtime::Tokio1), NoTls)? + .into(); + migrations::run(&pool).await?; + Ok(TestDatabase { + pool, + admin, + database_name, + }) +} + +fn pool_config(database_name: Option<&str>) -> Config { + let mut config = Config::new(); + config.host = Some( + std::env::var("PGHOST") + .or_else(|_| std::env::var("DATABASE_HOST")) + .unwrap_or_else(|_| "localhost".to_string()), + ); + config.port = Some( + std::env::var("PGPORT") + .or_else(|_| std::env::var("DATABASE_PORT")) + .unwrap_or_else(|_| "5432".to_string()) + .parse() + .expect("database port must be numeric"), + ); + config.dbname = Some(database_name.map(ToString::to_string).unwrap_or_else(|| { + std::env::var("PGDATABASE") + .or_else(|_| std::env::var("DATABASE_NAME")) + .unwrap_or_else(|_| "platform_api".to_string()) + })); + config.user = Some( + std::env::var("PGUSER") + .or_else(|_| std::env::var("DATABASE_USERNAME")) + .unwrap_or_else(|_| "postgres".to_string()), + ); + config.password = Some( + std::env::var("PGPASSWORD") + .or_else(|_| std::env::var("DATABASE_PASSWORD")) + .unwrap_or_else(|_| "postgres".to_string()), + ); + config.pool = Some(PoolConfig { + max_size: 4, + timeouts: Timeouts { + wait: Some(Duration::from_secs(10)), + create: Some(Duration::from_secs(10)), + recycle: Some(Duration::from_secs(10)), + }, + ..Default::default() + }); + config +} + +async fn fixture(pool: &DbPool) -> anyhow::Result { + let client = pool.get().await?; + let now = Utc::now(); + let suffix = Uuid::new_v4().simple().to_string(); + let user_id = Uuid::new_v4(); + let organization_id = Uuid::new_v4(); + let workspace_id = Uuid::new_v4(); + let mixed_key = Uuid::new_v4(); + let service_key = Uuid::new_v4(); + let deleted_key = Uuid::new_v4(); + let empty_key = Uuid::new_v4(); + let model_id = Uuid::new_v4(); + let service_id = Uuid::new_v4(); + client + .execute( + "INSERT INTO users (id, email, username, auth_provider, provider_user_id) + VALUES ($1, $2, $3, 'test', $4)", + &[ + &user_id, + &format!("{suffix}@example.test"), + &suffix, + &format!("provider-{suffix}"), + ], + ) + .await?; + client + .execute( + "INSERT INTO organizations (id, name, created_at, updated_at) + VALUES ($1, $2, $3, $3)", + &[&organization_id, &format!("org-{suffix}"), &now], + ) + .await?; + client + .execute( + "INSERT INTO workspaces (id, name, organization_id, created_by_user_id) + VALUES ($1, 'workspace', $2, $3)", + &[&workspace_id, &organization_id, &user_id], + ) + .await?; + for (id, name) in [ + (mixed_key, "mixed"), + (service_key, "service"), + (deleted_key, "deleted"), + (empty_key, "empty"), + ] { + client + .execute( + "INSERT INTO api_keys + (id, key_hash, key_prefix, name, workspace_id, created_by_user_id) + VALUES ($1, $2, 'test', $3, $4, $5)", + &[&id, &format!("hash-{id}"), &name, &workspace_id, &user_id], + ) + .await?; + } + client + .execute( + "UPDATE api_keys SET deleted_at = NOW(), is_active = false WHERE id = $1", + &[&deleted_key], + ) + .await?; + client + .execute( + "INSERT INTO models (id, model_name, model_display_name, model_description) + VALUES ($1, 'counter-test', 'counter-test', 'counter-test')", + &[&model_id], + ) + .await?; + client + .execute( + "INSERT INTO services + (id, service_name, display_name, unit, cost_per_unit) + VALUES ($1, $2, $2, 'request', 1)", + &[&service_id, &format!("service-{suffix}")], + ) + .await?; + Ok(Fixture { + organization_id, + mixed_key, + service_key, + deleted_key, + empty_key, + model_id, + service_id, + workspace_id, + }) +} + +fn inference_request(fixture: &Fixture, total_cost: i64) -> RecordUsageRequest { + RecordUsageRequest { + organization_id: fixture.organization_id, + workspace_id: fixture.workspace_id, + api_key_id: fixture.mixed_key, + model_id: fixture.model_id, + model_name: "counter-test".to_string(), + input_tokens: 1, + output_tokens: 1, + input_cost: total_cost, + output_cost: 0, + total_cost, + inference_type: InferenceType::ChatCompletion.as_str().to_string(), + ttft_ms: None, + avg_itl_ms: None, + inference_id: Some(Uuid::new_v4()), + provider_request_id: Some(format!("counter-{total_cost}-{}", Uuid::new_v4())), + stop_reason: None, + response_id: None, + image_count: None, + cache_read_tokens: 0, + cache_write_tokens: 0, + billing_details: None, + service_tier: None, + context_band: None, + served_provider_tier: None, + served_provider_type: None, + served_via_fallback: false, + } +} + +async fn insert_inference( + pool: &DbPool, + fixture: &Fixture, + api_key_id: Uuid, + cost: i64, +) -> anyhow::Result<()> { + let client = pool.get().await?; + let id = Uuid::new_v4(); + client + .execute( + "INSERT INTO organization_usage_log ( + id, organization_id, workspace_id, api_key_id, model_id, model_name, + input_tokens, output_tokens, total_tokens, input_cost, output_cost, + total_cost, inference_type, inference_id + ) VALUES ($1, $2, $3, $4, $5, 'counter-test', 1, 0, 1, $6, 0, $6, + 'chat_completion', $1)", + &[ + &id, + &fixture.organization_id, + &fixture.workspace_id, + &api_key_id, + &fixture.model_id, + &cost, + ], + ) + .await?; + Ok(()) +} + +async fn insert_service( + pool: &DbPool, + fixture: &Fixture, + api_key_id: Uuid, + cost: i64, +) -> anyhow::Result<()> { + let client = pool.get().await?; + client + .execute( + "INSERT INTO organization_service_usage_log + (organization_id, workspace_id, api_key_id, service_id, quantity, total_cost) + VALUES ($1, $2, $3, $4, 1, $5)", + &[ + &fixture.organization_id, + &fixture.workspace_id, + &api_key_id, + &fixture.service_id, + &cost, + ], + ) + .await?; + Ok(()) +} + +impl TestDatabase { + async fn cleanup(self) -> anyhow::Result<()> { + drop(self.pool); + self.admin + .get() + .await? + .batch_execute(&format!("DROP DATABASE {}", self.database_name)) + .await?; + Ok(()) + } +} From 5d129663a3965a7cc1a606e6866c5b3352d25125 Mon Sep 17 00:00:00 2001 From: Henry Park <16583448+henrypark133@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:53:15 +0000 Subject: [PATCH 3/7] perf: replace lifetime spending scans with ready counters --- .config/nextest.toml | 2 +- .github/workflows/test.yml | 2 +- crates/api/src/main.rs | 3 + crates/api/tests/e2e_all/api_keys.rs | 126 +++---- crates/database/src/repositories/analytics.rs | 12 +- crates/database/src/repositories/api_key.rs | 21 +- .../src/repositories/organization_usage.rs | 9 +- crates/database/tests/counter_readers.rs | 312 ++++++++++++++++++ crates/services/src/admin/analytics.rs | 7 +- 9 files changed, 400 insertions(+), 94 deletions(-) create mode 100644 crates/database/tests/counter_readers.rs diff --git a/.config/nextest.toml b/.config/nextest.toml index 95747b942..56aed4286 100644 --- a/.config/nextest.toml +++ b/.config/nextest.toml @@ -48,7 +48,7 @@ filter = "package(api) & binary(e2e_all)" test-group = "e2e-db" [[profile.default.overrides]] -filter = "package(database) & binary(spend_counter_backfill)" +filter = "package(database) & (binary(spend_counter_backfill) | binary(counter_readers))" test-group = "e2e-db" [[profile.default.scripts]] diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 0f45afc51..1ded81424 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -199,7 +199,7 @@ jobs: cache-key: e2e - name: Run e2e tests - run: cargo nextest run --test e2e_all --test spend_counter_backfill + run: cargo nextest run --test e2e_all --test spend_counter_backfill --test counter_readers env: POSTGRES_PRIMARY_APP_ID: ${{ secrets.POSTGRES_PRIMARY_APP_ID }} DATABASE_HOST: localhost diff --git a/crates/api/src/main.rs b/crates/api/src/main.rs index 55ffec788..678cfc553 100644 --- a/crates/api/src/main.rs +++ b/crates/api/src/main.rs @@ -27,6 +27,9 @@ async fn main() { .await .expect("Usage reporting index prerequisites are not satisfied"); } + database::ensure_spend_counters_ready(database.pool()) + .await + .expect("Spend counters are incomplete; run the spend-counter backfill before starting Cloud API"); let auth_components = init_auth_services(database.clone(), &config); // Initialize OpenTelemetry pipeline diff --git a/crates/api/tests/e2e_all/api_keys.rs b/crates/api/tests/e2e_all/api_keys.rs index 91051415f..4e72cee80 100644 --- a/crates/api/tests/e2e_all/api_keys.rs +++ b/crates/api/tests/e2e_all/api_keys.rs @@ -1,4 +1,9 @@ use crate::common::*; +use database::models::RecordUsageRequest; +use database::repositories::{ + OrganizationServiceUsageRepository, OrganizationUsageRepository, RecordServiceUsageRequest, +}; +use services::usage::InferenceType; // ============================================ // API Key Creation and Management Tests @@ -128,76 +133,73 @@ async fn test_list_workspace_api_keys_orders_by_usage() { .unwrap() .get(0); - for (api_key_id, total_cost) in [ - (service_spend_key_id, 100_000_000_i64), - (inference_spend_key_id, 300_000_000_i64), + drop(client); + let inference_repository = OrganizationUsageRepository::new(database.pool().clone()); + let mut inference = RecordUsageRequest { + organization_id, + workspace_id, + api_key_id: service_spend_key_id, + model_id, + model_name: model_name.clone(), + input_tokens: 10, + output_tokens: 10, + input_cost: 1, + output_cost: 1, + total_cost: 100_000_000, + inference_type: InferenceType::ChatCompletion.as_str().to_string(), + ttft_ms: None, + avg_itl_ms: None, + inference_id: Some(uuid::Uuid::new_v4()), + provider_request_id: None, + stop_reason: None, + response_id: None, + image_count: None, + cache_read_tokens: 0, + cache_write_tokens: 0, + billing_details: None, + service_tier: None, + context_band: None, + served_provider_tier: None, + served_provider_type: None, + served_via_fallback: false, + }; + inference_repository + .record_usage(inference.clone()) + .await + .unwrap(); + inference.api_key_id = inference_spend_key_id; + inference.total_cost = 300_000_000; + inference.inference_id = Some(uuid::Uuid::new_v4()); + inference_repository.record_usage(inference).await.unwrap(); + OrganizationServiceUsageRepository::new(database.pool().clone()) + .record_usage(&RecordServiceUsageRequest { + organization_id, + workspace_id, + api_key_id: service_spend_key_id, + service_id, + quantity: 1, + total_cost: 400_000_000, + inference_id: None, + }) + .await + .unwrap(); + + let client = database.pool().get().await.unwrap(); + + for (api_key_id, hours) in [ + (service_spend_key_id, 3_i64), + (inference_spend_key_id, 2_i64), + (unused_key_id, 1_i64), ] { client .execute( - r#" - INSERT INTO organization_usage_log ( - id, organization_id, workspace_id, api_key_id, - model_id, model_name, input_tokens, output_tokens, - total_tokens, input_cost, output_cost, total_cost, - inference_type, created_at - ) VALUES ($1, $2, $3, $4, $5, $6, 10, 10, 20, 1, 1, $7, - 'chat_completion', NOW()) - "#, - &[ - &uuid::Uuid::new_v4(), - &organization_id, - &workspace_id, - &api_key_id, - &model_id, - &model_name, - &total_cost, - ], + "UPDATE api_keys SET created_at = NOW() + ($2::BIGINT * INTERVAL '1 hour') WHERE id = $1", + &[&api_key_id, &hours], ) .await .unwrap(); } - client - .execute( - r#" - INSERT INTO organization_service_usage_log ( - id, organization_id, workspace_id, api_key_id, - service_id, quantity, total_cost, inference_id, created_at - ) VALUES ($1, $2, $3, $4, $5, 1, 400000000, NULL, NOW()) - "#, - &[ - &uuid::Uuid::new_v4(), - &organization_id, - &workspace_id, - &service_spend_key_id, - &service_id, - ], - ) - .await - .unwrap(); - - client - .execute( - "UPDATE api_keys SET created_at = NOW() + INTERVAL '3 hours' WHERE id = $1", - &[&service_spend_key_id], - ) - .await - .unwrap(); - client - .execute( - "UPDATE api_keys SET created_at = NOW() + INTERVAL '2 hours' WHERE id = $1", - &[&inference_spend_key_id], - ) - .await - .unwrap(); - client - .execute( - "UPDATE api_keys SET created_at = NOW() + INTERVAL '1 hour' WHERE id = $1", - &[&unused_key_id], - ) - .await - .unwrap(); - let default_response = server .get(format!("/v1/workspaces/{}/api-keys?limit=3", workspace.id).as_str()) .add_header("Authorization", format!("Bearer {}", get_session_id())) diff --git a/crates/database/src/repositories/analytics.rs b/crates/database/src/repositories/analytics.rs index a13e1dd22..898b799ec 100644 --- a/crates/database/src/repositories/analytics.rs +++ b/crates/database/src/repositories/analytics.rs @@ -545,16 +545,16 @@ impl AnalyticsRepository for PgAnalyticsRepository { let paying_org_count: i64 = limits_row.get(2); let granted_org_count: i64 = limits_row.get(3); - // All-time consumed cost. `total` (from the cached balance) is ALL usage - // (inference + services); the inference/service splits come from their logs - // and reconcile to the total. + // All-time consumed cost. Keep the legacy total independently; historical + // adjustments can differ from the balance splits. let consumed_row = client .query_one( r#" SELECT - (SELECT COALESCE(SUM(total_spent), 0) FROM organization_balance)::bigint as total_nano, - (SELECT COALESCE(SUM(total_cost), 0) FROM organization_usage_log)::bigint as inference_nano, - (SELECT COALESCE(SUM(total_cost), 0) FROM organization_service_usage_log)::bigint as service_nano + COALESCE(SUM(total_spent), 0)::bigint as total_nano, + COALESCE(SUM(inference_spent), 0)::bigint as inference_nano, + COALESCE(SUM(service_spent), 0)::bigint as service_nano + FROM organization_balance "#, &[], ) diff --git a/crates/database/src/repositories/api_key.rs b/crates/database/src/repositories/api_key.rs index d587eede9..fc175b3db 100644 --- a/crates/database/src/repositories/api_key.rs +++ b/crates/database/src/repositories/api_key.rs @@ -253,8 +253,8 @@ impl ApiKeyRepository { Ok(row.get::<_, i64>("count")) } - /// List API keys for a workspace with usage data - /// This is the primary method to list API keys, using an efficient JOIN query + /// List API keys for a workspace with usage data. + /// This is the primary method to list API keys, using the spend counter JOIN. pub async fn list_by_workspace_paginated( &self, workspace_id: Uuid, @@ -305,22 +305,11 @@ impl ApiKeyRepository { ak.deleted_at, ak.spend_limit, ( - COALESCE(inference_usage.total_cost, 0) - + COALESCE(service_usage.total_cost, 0) + COALESCE(spend.inference_spent, 0) + + COALESCE(spend.service_spent, 0) )::BIGINT as usage FROM api_keys ak - LEFT JOIN ( - SELECT api_key_id, COALESCE(SUM(total_cost), 0)::BIGINT AS total_cost - FROM organization_usage_log - WHERE workspace_id = $1 - GROUP BY api_key_id - ) inference_usage ON ak.id = inference_usage.api_key_id - LEFT JOIN ( - SELECT api_key_id, COALESCE(SUM(total_cost), 0)::BIGINT AS total_cost - FROM organization_service_usage_log - WHERE workspace_id = $1 - GROUP BY api_key_id - ) service_usage ON ak.id = service_usage.api_key_id + LEFT JOIN api_key_spend spend ON ak.id = spend.api_key_id WHERE ak.workspace_id = $1 AND ak.deleted_at IS NULL ORDER BY {order_by_column} {order_dir}{tie_breaker} LIMIT $2 OFFSET $3 diff --git a/crates/database/src/repositories/organization_usage.rs b/crates/database/src/repositories/organization_usage.rs index 03194285c..fecc7e99d 100644 --- a/crates/database/src/repositories/organization_usage.rs +++ b/crates/database/src/repositories/organization_usage.rs @@ -56,7 +56,7 @@ impl OrganizationUsageRepository { } } - /// Get total spend for a specific API key + /// Get inference spend for a specific API key for admission-limit checks. pub async fn get_api_key_spend(&self, api_key_id: Uuid) -> Result { let row = retry_db!("get_api_key_spend", { let client = self @@ -69,9 +69,10 @@ impl OrganizationUsageRepository { client .query_one( r#" - SELECT COALESCE(SUM(total_cost), 0)::BIGINT as total_spend - FROM organization_usage_log - WHERE api_key_id = $1 + SELECT COALESCE( + (SELECT inference_spent FROM api_key_spend WHERE api_key_id = $1), + 0 + )::BIGINT as total_spend "#, &[&api_key_id], ) diff --git a/crates/database/tests/counter_readers.rs b/crates/database/tests/counter_readers.rs new file mode 100644 index 000000000..5135e22dd --- /dev/null +++ b/crates/database/tests/counter_readers.rs @@ -0,0 +1,312 @@ +use database::repositories::{ + ApiKeyRepository, OrganizationUsageRepository, PgAnalyticsRepository, +}; +use database::{ensure_spend_counters_ready, DbPool}; +use deadpool::Runtime; +use deadpool_postgres::{Config, PoolConfig}; +use services::admin::AnalyticsRepository; +use services::workspace::ports::{ApiKeyOrderBy, ApiKeyOrderDirection}; +use tokio_postgres::NoTls; +use uuid::Uuid; + +fn env_value(primary: &str, fallback: &str, default: &str) -> String { + std::env::var(primary) + .or_else(|_| std::env::var(fallback)) + .unwrap_or_else(|_| default.to_string()) +} + +fn pool_config() -> Config { + let mut config = Config::new(); + config.host = Some(env_value("PGHOST", "DATABASE_HOST", "localhost")); + config.port = Some( + env_value("PGPORT", "DATABASE_PORT", "5432") + .parse() + .expect("database port must be numeric"), + ); + config.dbname = Some(env_value( + "PGDATABASE", + "TEST_DATABASE_NAME", + "platform_api", + )); + config.user = Some(env_value("PGUSER", "DATABASE_USERNAME", "postgres")); + config.password = Some(env_value("PGPASSWORD", "DATABASE_PASSWORD", "postgres")); + config.pool = Some(PoolConfig::new(4)); + config +} + +async fn new_pool(config: &Config) -> anyhow::Result { + Ok(DbPool::new( + config.create_pool(Some(Runtime::Tokio1), NoTls)?, + )) +} + +async fn scoped_pool() -> anyhow::Result<(DbPool, DbPool, String)> { + let admin_pool = new_pool(&pool_config()).await?; + let admin = admin_pool.get().await?; + let schema = format!("counter_readers_{}", Uuid::new_v4().simple()); + admin + .batch_execute( + format!( + r#" + CREATE SCHEMA {schema}; + CREATE TABLE {schema}.api_keys ( + id UUID PRIMARY KEY, + key_hash TEXT NOT NULL, + key_prefix TEXT NOT NULL, + name TEXT NOT NULL, + workspace_id UUID NOT NULL, + created_by_user_id UUID NOT NULL, + created_at TIMESTAMPTZ NOT NULL, + expires_at TIMESTAMPTZ, + last_used_at TIMESTAMPTZ, + is_active BOOLEAN NOT NULL, + deleted_at TIMESTAMPTZ, + spend_limit BIGINT + ); + CREATE TABLE {schema}.api_key_spend ( + api_key_id UUID PRIMARY KEY, + inference_spent BIGINT NOT NULL, + service_spent BIGINT NOT NULL, + updated_at TIMESTAMPTZ NOT NULL + ); + CREATE TABLE {schema}.organizations ( + id UUID PRIMARY KEY, + is_active BOOLEAN NOT NULL + ); + CREATE TABLE {schema}.organization_balance ( + organization_id UUID PRIMARY KEY, + total_spent BIGINT NOT NULL, + inference_spent BIGINT NOT NULL, + service_spent BIGINT NOT NULL, + spend_counters_ready_at TIMESTAMPTZ DEFAULT NOW() + ); + CREATE TABLE {schema}.organization_limits_history ( + organization_id UUID NOT NULL, + spend_limit BIGINT NOT NULL, + credit_type TEXT NOT NULL, + effective_until TIMESTAMPTZ, + source TEXT + ); + "# + ) + .as_str(), + ) + .await?; + drop(admin); + + let mut scoped_config = pool_config(); + scoped_config.options = Some(format!("-c search_path={schema},pg_catalog")); + let scoped_pool = new_pool(&scoped_config).await?; + Ok((admin_pool, scoped_pool, schema)) +} + +async fn drop_schema( + admin_pool: DbPool, + scoped_pool: DbPool, + schema: String, +) -> anyhow::Result<()> { + drop(scoped_pool); + let admin = admin_pool.get().await?; + admin + .batch_execute(format!("DROP SCHEMA {schema} CASCADE").as_str()) + .await?; + Ok(()) +} + +async fn insert_key( + pool: &DbPool, + workspace_id: Uuid, + key_id: Uuid, + name: &str, + inference_spent: i64, + service_spent: i64, + deleted: bool, +) -> anyhow::Result<()> { + let client = pool.get().await?; + let now = chrono::Utc::now(); + let deleted_at = deleted.then_some(now); + client + .execute( + "INSERT INTO api_keys (id, key_hash, key_prefix, name, workspace_id, created_by_user_id, created_at, is_active, deleted_at) VALUES ($1, $2, $3, $4, $5, $6, $7, true, $8)", + &[ + &key_id, + &format!("hash-{key_id}"), + &"sk-test", + &name, + &workspace_id, + &Uuid::new_v4(), + &now, + &deleted_at, + ], + ) + .await?; + client + .execute( + "INSERT INTO api_key_spend (api_key_id, inference_spent, service_spent, updated_at) VALUES ($1, $2, $3, $4)", + &[&key_id, &inference_spent, &service_spent, &now], + ) + .await?; + Ok(()) +} + +#[tokio::test] +async fn key_list_reads_counter_cohort_with_stable_pagination() -> anyhow::Result<()> { + let (admin_pool, pool, schema) = scoped_pool().await?; + let workspace_id = Uuid::new_v4(); + let zero = Uuid::new_v4(); + let inference = Uuid::new_v4(); + let service = Uuid::new_v4(); + let mixed = Uuid::new_v4(); + let deleted = Uuid::new_v4(); + let other_workspace = Uuid::new_v4(); + let other_workspace_key = Uuid::new_v4(); + insert_key(&pool, workspace_id, zero, "zero", 0, 0, false).await?; + insert_key(&pool, workspace_id, inference, "inference", 30, 0, false).await?; + insert_key(&pool, workspace_id, service, "service", 0, 40, false).await?; + insert_key(&pool, workspace_id, mixed, "mixed", 50, 70, false).await?; + insert_key(&pool, workspace_id, deleted, "deleted", 900, 900, true).await?; + insert_key( + &pool, + other_workspace, + other_workspace_key, + "other-workspace", + 9_000, + 9_000, + false, + ) + .await?; + pool.get() + .await? + .execute("DELETE FROM api_key_spend WHERE api_key_id = $1", &[&zero]) + .await?; + + // The isolated schema deliberately has no raw usage tables. + ensure_spend_counters_ready(&pool).await?; + let repository = ApiKeyRepository::new(pool.clone()); + let first_page = repository + .list_by_workspace_paginated( + workspace_id, + 3, + 0, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Desc), + ) + .await?; + assert_eq!( + first_page.iter().map(|key| key.id).collect::>(), + vec![mixed, service, inference] + ); + assert_eq!( + first_page.iter().map(|key| key.usage).collect::>(), + vec![120, 40, 30] + ); + + let second_page = repository + .list_by_workspace_paginated( + workspace_id, + 3, + 3, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Desc), + ) + .await?; + assert_eq!( + second_page.iter().map(|key| key.id).collect::>(), + vec![zero] + ); + assert_eq!(second_page[0].usage, 0); + assert!(!second_page.iter().any(|key| key.id == deleted)); + + let ascending = repository + .list_by_workspace_paginated( + workspace_id, + 2, + 0, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Asc), + ) + .await?; + assert_eq!( + ascending.iter().map(|key| key.id).collect::>(), + vec![zero, inference] + ); + let past_end = repository + .list_by_workspace_paginated( + workspace_id, + 2, + 99, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Desc), + ) + .await?; + assert!(past_end.is_empty()); + assert!(!first_page.iter().any(|key| key.id == other_workspace_key)); + + drop_schema(admin_pool, pool, schema).await +} + +#[tokio::test] +async fn admission_spend_reads_inference_counter_and_excludes_service_counter() -> anyhow::Result<()> +{ + let (admin_pool, pool, schema) = scoped_pool().await?; + let inference_key = Uuid::new_v4(); + let service_key = Uuid::new_v4(); + let now = chrono::Utc::now(); + let client = pool.get().await?; + for (key_id, inference_spent, service_spent) in [ + (inference_key, 321_i64, 0_i64), + (service_key, 0_i64, 654_i64), + ] { + client + .execute( + "INSERT INTO api_key_spend (api_key_id, inference_spent, service_spent, updated_at) VALUES ($1, $2, $3, $4)", + &[&key_id, &inference_spent, &service_spent, &now], + ) + .await?; + } + drop(client); + + ensure_spend_counters_ready(&pool).await?; + let repository = OrganizationUsageRepository::new(pool.clone()); + assert_eq!(repository.get_api_key_spend(inference_key).await?, 321); + assert_eq!(repository.get_api_key_spend(service_key).await?, 0); + assert_eq!(repository.get_api_key_spend(Uuid::new_v4()).await?, 0); + + drop_schema(admin_pool, pool, schema).await +} + +#[tokio::test] +async fn billing_summary_reads_balance_splits_and_preserves_legacy_total() -> anyhow::Result<()> { + let (admin_pool, pool, schema) = scoped_pool().await?; + let organization_id = Uuid::new_v4(); + let client = pool.get().await?; + client + .execute( + "INSERT INTO organizations (id, is_active) VALUES ($1, true)", + &[&organization_id], + ) + .await?; + client + .execute( + "INSERT INTO organization_limits_history (organization_id, spend_limit, credit_type, source) VALUES ($1, 0, 'payment', 'test')", + &[&organization_id], + ) + .await?; + client + .execute( + "INSERT INTO organization_balance (organization_id, total_spent, inference_spent, service_spent) VALUES ($1, 123, 35, 20)", + &[&organization_id], + ) + .await?; + drop(client); + + ensure_spend_counters_ready(&pool).await?; + let summary = PgAnalyticsRepository::new(pool.clone()) + .get_billing_summary() + .await?; + assert_eq!(summary.total_consumed_usd, 123e-9); + assert_eq!(summary.inference_consumed_usd, 35e-9); + assert_eq!(summary.service_consumed_usd, 20e-9); + + drop_schema(admin_pool, pool, schema).await +} diff --git a/crates/services/src/admin/analytics.rs b/crates/services/src/admin/analytics.rs index 209d27312..4f9a7ad5f 100644 --- a/crates/services/src/admin/analytics.rs +++ b/crates/services/src/admin/analytics.rs @@ -220,12 +220,11 @@ pub struct BillingSummary { /// Sum of active grant-type spend limits (caps), USD pub active_grant_credit_limit_usd: f64, /// All-time consumed cost across all orgs, USD — **all usage** (from - /// organization_balance: inference + services). `inference_consumed_usd + - /// service_consumed_usd` reconcile to this. + /// organization_balance: inference + services). pub total_consumed_usd: f64, - /// All-time inference consumed cost, USD (organization_usage_log) + /// All-time inference consumed cost, USD (organization_balance.inference_spent) pub inference_consumed_usd: f64, - /// All-time service consumed cost, USD (organization_service_usage_log, e.g. web_search) + /// All-time service consumed cost, USD (organization_balance.service_spent, e.g. web_search) pub service_consumed_usd: f64, pub paying_org_count: i64, pub granted_org_count: i64, From eacfd851e698071ec9991a66cfdc0c683c760430 Mon Sep 17 00:00:00 2001 From: Henry Park <16583448+henrypark133@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:54:34 +0000 Subject: [PATCH 4/7] test: isolate billing totals from concurrent fixtures --- .config/nextest.toml | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/.config/nextest.toml b/.config/nextest.toml index 56aed4286..868806d7a 100644 --- a/.config/nextest.toml +++ b/.config/nextest.toml @@ -43,6 +43,12 @@ filter = 'package(api) & binary(e2e_all) & test(/^admin_pricing_changes::test_(r test-group = "e2e-db" threads-required = "num-test-threads" +[[profile.default.overrides]] +filter = 'package(api) & binary(e2e_all) & test(/^admin_analytics::test_admin_platform_billing_summary$/)' +test-group = "e2e-db" +# Exact before/after platform totals require exclusive access to the shared DB. +threads-required = "num-test-threads" + [[profile.default.overrides]] filter = "package(api) & binary(e2e_all)" test-group = "e2e-db" From 2997c87e6962791562debc936f53f5207746eab1 Mon Sep 17 00:00:00 2001 From: Henry Park <16583448+henrypark133@users.noreply.github.com> Date: Tue, 22 Sep 2026 19:19:14 +0000 Subject: [PATCH 5/7] Address PR review feedback (#1116) - Combine balance and key counter writes into one SQL round trip - Cover service overflow rollback and retain exact SQLSTATE checks --- crates/api/tests/e2e_all/spend_counters.rs | 165 ++++++++++++++++-- crates/database/src/repositories/mod.rs | 1 - .../organization_service_usage.rs | 45 +++-- .../src/repositories/organization_usage.rs | 57 +++--- .../src/repositories/spend_counters.rs | 31 ---- 5 files changed, 213 insertions(+), 86 deletions(-) delete mode 100644 crates/database/src/repositories/spend_counters.rs diff --git a/crates/api/tests/e2e_all/spend_counters.rs b/crates/api/tests/e2e_all/spend_counters.rs index 94eadcd83..3a82f198c 100644 --- a/crates/api/tests/e2e_all/spend_counters.rs +++ b/crates/api/tests/e2e_all/spend_counters.rs @@ -28,6 +28,7 @@ async fn fixture() -> ( "spend-counter-key".to_string(), ) .await; + let api_key_id: Uuid = api_key.id.parse().expect("API key UUID"); let client = database.pool().get().await.expect("database connection"); let model = client .query_one( @@ -43,7 +44,7 @@ async fn fixture() -> ( let request = RecordUsageRequest { organization_id: org.id.parse().expect("organization UUID"), workspace_id: workspace.id.parse().expect("workspace UUID"), - api_key_id: api_key.id.parse().expect("API key UUID"), + api_key_id, model_id, model_name, input_tokens: 1, @@ -68,12 +69,7 @@ async fn fixture() -> ( served_provider_type: None, served_via_fallback: false, }; - ( - server, - database, - request, - api_key.id.parse().expect("API key UUID"), - ) + (server, database, request, api_key_id) } #[tokio::test] @@ -139,6 +135,7 @@ async fn spend_counters_track_inference_and_service_separately_on_one_key() -> a async fn spend_counters_ignore_duplicates_and_failed_posts() -> anyhow::Result<()> { let (_server, database, inference, api_key_id) = fixture().await; let repository = OrganizationUsageRepository::new(database.pool().clone()); + let service_repository = OrganizationServiceUsageRepository::new(database.pool().clone()); repository.record_usage(inference.clone()).await?; repository.record_usage(inference.clone()).await?; @@ -168,9 +165,11 @@ async fn spend_counters_ignore_duplicates_and_failed_posts() -> anyhow::Result<( else { panic!("expected a database overflow error"); }; - assert_eq!( - cause.to_string(), - "Database error (22003): bigint out of range" + // map_db_error retains SQLSTATE in this prefix; PostgreSQL diagnostic wording + // is not the failure-kind contract checked by this rollback test. + assert!( + cause.to_string().starts_with("Database error (22003): "), + "expected SQLSTATE 22003, got: {cause}" ); let client = database.pool().get().await?; @@ -210,9 +209,155 @@ async fn spend_counters_ignore_duplicates_and_failed_posts() -> anyhow::Result<( assert_eq!(balance.get::<_, i64>("unresolved_unfunded_amount"), 3_000); client .execute( + // Restore the counter after forcing it to MAX; leaving MAX in the + // shared database can overflow aggregate spend reads. "UPDATE api_key_spend SET inference_spent = $2 WHERE api_key_id = $1", &[&api_key_id, &3_000i64], ) .await?; + + let service_id = Uuid::new_v4(); + client + .execute( + "INSERT INTO services (id, service_name, display_name, unit, cost_per_unit) VALUES ($1, $2, $2, 'request', 1)", + &[&service_id, &format!("spend-counter-overflow-service-{service_id}")], + ) + .await?; + + let service_key_overflow_id = Uuid::new_v4(); + client + .execute( + "UPDATE api_key_spend SET service_spent = $2 WHERE api_key_id = $1", + &[&api_key_id, &i64::MAX], + ) + .await?; + drop(client); + + let service_request = RecordServiceUsageRequest { + organization_id: inference.organization_id, + workspace_id: inference.workspace_id, + api_key_id, + service_id, + quantity: 1, + total_cost: 1, + inference_id: Some(service_key_overflow_id), + }; + let error = service_repository + .record_usage(&service_request) + .await + .expect_err("a per-key service counter overflow must roll back the service transaction"); + let services::common::RepositoryError::DatabaseError(cause) = error + .downcast::() + .expect("repository error") + else { + panic!("expected a database overflow error"); + }; + assert!( + cause.to_string().starts_with("Database error (22003): "), + "expected SQLSTATE 22003, got: {cause}" + ); + + let client = database.pool().get().await?; + let key = client + .query_one( + "SELECT inference_spent, service_spent FROM api_key_spend WHERE api_key_id = $1", + &[&api_key_id], + ) + .await?; + assert_eq!(key.get::<_, i64>("inference_spent"), 3_000); + assert_eq!(key.get::<_, i64>("service_spent"), i64::MAX); + let service_count: i64 = client + .query_one( + "SELECT COUNT(*)::BIGINT FROM organization_service_usage_log WHERE inference_id = $1", + &[&service_key_overflow_id], + ) + .await? + .get(0); + assert_eq!(service_count, 0); + let balance = client + .query_one( + "SELECT total_spent, inference_spent, service_spent, unresolved_unfunded_amount FROM organization_balance WHERE organization_id = $1", + &[&inference.organization_id], + ) + .await?; + assert_eq!(balance.get::<_, i64>("total_spent"), 3_000); + assert_eq!(balance.get::<_, i64>("inference_spent"), 3_000); + assert_eq!(balance.get::<_, i64>("service_spent"), 0); + assert_eq!(balance.get::<_, i64>("unresolved_unfunded_amount"), 3_000); + + let organization_overflow_id = Uuid::new_v4(); + client + .execute( + "UPDATE api_key_spend SET service_spent = 0 WHERE api_key_id = $1", + &[&api_key_id], + ) + .await?; + client + .execute( + "UPDATE organization_balance SET service_spent = $2 WHERE organization_id = $1", + &[&inference.organization_id, &i64::MAX], + ) + .await?; + drop(client); + + let mut organization_overflow_request = service_request.clone(); + organization_overflow_request.inference_id = Some(organization_overflow_id); + let error = service_repository + .record_usage(&organization_overflow_request) + .await + .expect_err( + "an organization service counter overflow must roll back the service transaction", + ); + let services::common::RepositoryError::DatabaseError(cause) = error + .downcast::() + .expect("repository error") + else { + panic!("expected a database overflow error"); + }; + assert!( + cause.to_string().starts_with("Database error (22003): "), + "expected SQLSTATE 22003, got: {cause}" + ); + + let client = database.pool().get().await?; + let key = client + .query_one( + "SELECT inference_spent, service_spent FROM api_key_spend WHERE api_key_id = $1", + &[&api_key_id], + ) + .await?; + assert_eq!(key.get::<_, i64>("inference_spent"), 3_000); + assert_eq!(key.get::<_, i64>("service_spent"), 0); + let service_count: i64 = client + .query_one( + "SELECT COUNT(*)::BIGINT FROM organization_service_usage_log WHERE inference_id = $1", + &[&organization_overflow_id], + ) + .await? + .get(0); + assert_eq!(service_count, 0); + let balance = client + .query_one( + "SELECT total_spent, inference_spent, service_spent, unresolved_unfunded_amount FROM organization_balance WHERE organization_id = $1", + &[&inference.organization_id], + ) + .await?; + assert_eq!(balance.get::<_, i64>("total_spent"), 3_000); + assert_eq!(balance.get::<_, i64>("inference_spent"), 3_000); + assert_eq!(balance.get::<_, i64>("service_spent"), i64::MAX); + assert_eq!(balance.get::<_, i64>("unresolved_unfunded_amount"), 3_000); + + client + .execute( + "UPDATE api_key_spend SET service_spent = 0 WHERE api_key_id = $1", + &[&api_key_id], + ) + .await?; + client + .execute( + "UPDATE organization_balance SET service_spent = 0 WHERE organization_id = $1", + &[&inference.organization_id], + ) + .await?; Ok(()) } diff --git a/crates/database/src/repositories/mod.rs b/crates/database/src/repositories/mod.rs index 6f95ae1db..fa76a8de3 100644 --- a/crates/database/src/repositories/mod.rs +++ b/crates/database/src/repositories/mod.rs @@ -34,7 +34,6 @@ pub mod retry; pub mod service; pub mod service_usage_repository_impl; pub mod session; -pub(crate) mod spend_counters; pub mod usage_repository_impl; pub mod user; pub mod utils; diff --git a/crates/database/src/repositories/organization_service_usage.rs b/crates/database/src/repositories/organization_service_usage.rs index 8f4cfe4a9..9438ebd88 100644 --- a/crates/database/src/repositories/organization_service_usage.rs +++ b/crates/database/src/repositories/organization_service_usage.rs @@ -4,7 +4,6 @@ use crate::repositories::credit_allocation::{ allocate_usage, load_allocations, lock_organization_accounting, CreditAllocationPolicy, UsageAllocationParent, }; -use crate::repositories::spend_counters::increment_api_key_spend; use crate::repositories::utils::map_db_error; use crate::retry_db; use anyhow::{Context, Result}; @@ -363,28 +362,38 @@ impl OrganizationServiceUsageRepository { transaction .execute( r#" - INSERT INTO organization_balance ( - organization_id, total_spent, inference_spent, service_spent, - last_usage_at, total_requests, total_tokens, updated_at - ) VALUES ($1, $2, 0, $2, $3, 0, 0, $4) - ON CONFLICT (organization_id) DO UPDATE SET - total_spent = organization_balance.total_spent + $2, - service_spent = organization_balance.service_spent + $2, - last_usage_at = $3, + WITH balance_upsert AS ( + INSERT INTO organization_balance ( + organization_id, total_spent, inference_spent, service_spent, + last_usage_at, total_requests, total_tokens, updated_at + ) VALUES ($1, $2, 0, $2, $3, 0, 0, $4) + ON CONFLICT (organization_id) DO UPDATE SET + total_spent = organization_balance.total_spent + $2, + service_spent = organization_balance.service_spent + $2, + last_usage_at = $3, + updated_at = $4 + RETURNING organization_id + ) + INSERT INTO api_key_spend ( + api_key_id, inference_spent, service_spent, updated_at + ) + SELECT $5, 0, $2, $4 + FROM balance_upsert + WHERE TRUE + ON CONFLICT (api_key_id) DO UPDATE SET + service_spent = api_key_spend.service_spent + $2, updated_at = $4 "#, - &[&request.organization_id, &request.total_cost, &now, &now], + &[ + &request.organization_id, + &request.total_cost, + &now, + &now, + &request.api_key_id, + ], ) .await .map_err(map_db_error)?; - increment_api_key_spend( - &transaction, - request.api_key_id, - 0, - request.total_cost, - now, - ) - .await?; transaction.commit().await.map_err(map_db_error)?; (r, Some(allocation.allocations)) diff --git a/crates/database/src/repositories/organization_usage.rs b/crates/database/src/repositories/organization_usage.rs index 03194285c..ea043d75f 100644 --- a/crates/database/src/repositories/organization_usage.rs +++ b/crates/database/src/repositories/organization_usage.rs @@ -7,7 +7,6 @@ use crate::repositories::credit_allocation::{ allocate_usage, load_allocations, lock_organization_accounting, CreditAllocationPolicy, UsageAllocationParent, }; -use crate::repositories::spend_counters::increment_api_key_spend; use crate::repositories::utils::map_db_error; use crate::retry_db; use anyhow::{Context, Result}; @@ -194,26 +193,39 @@ impl OrganizationUsageRepository { ) .await .map_err(map_db_error)?; - // New insert succeeded — update organization balance + // New insert succeeded — update organization balance and the + // per-key spend counter in one statement. transaction .execute( r#" - INSERT INTO organization_balance ( - organization_id, - total_spent, - inference_spent, - service_spent, - last_usage_at, - total_requests, - total_tokens, - updated_at - ) VALUES ($1, $2, $2, 0, $3, 1, $4, $5) - ON CONFLICT (organization_id) DO UPDATE SET - total_spent = organization_balance.total_spent + $2, - inference_spent = organization_balance.inference_spent + $2, - total_requests = organization_balance.total_requests + 1, - total_tokens = organization_balance.total_tokens + $4, - last_usage_at = $3, + WITH balance_upsert AS ( + INSERT INTO organization_balance ( + organization_id, + total_spent, + inference_spent, + service_spent, + last_usage_at, + total_requests, + total_tokens, + updated_at + ) VALUES ($1, $2, $2, 0, $3, 1, $4, $5) + ON CONFLICT (organization_id) DO UPDATE SET + total_spent = organization_balance.total_spent + $2, + inference_spent = organization_balance.inference_spent + $2, + total_requests = organization_balance.total_requests + 1, + total_tokens = organization_balance.total_tokens + $4, + last_usage_at = $3, + updated_at = $5 + RETURNING organization_id + ) + INSERT INTO api_key_spend ( + api_key_id, inference_spent, service_spent, updated_at + ) + SELECT $6, $2, 0, $5 + FROM balance_upsert + WHERE TRUE + ON CONFLICT (api_key_id) DO UPDATE SET + inference_spent = api_key_spend.inference_spent + $2, updated_at = $5 "#, &[ @@ -222,18 +234,11 @@ impl OrganizationUsageRepository { &now, &(total_tokens as i64), &now, + &request.api_key_id, ], ) .await .map_err(map_db_error)?; - increment_api_key_spend( - &transaction, - request.api_key_id, - request.total_cost, - 0, - now, - ) - .await?; transaction.commit().await.map_err(map_db_error)?; (row, true, Some(allocation.allocations)) diff --git a/crates/database/src/repositories/spend_counters.rs b/crates/database/src/repositories/spend_counters.rs deleted file mode 100644 index 79f364e4d..000000000 --- a/crates/database/src/repositories/spend_counters.rs +++ /dev/null @@ -1,31 +0,0 @@ -use crate::repositories::utils::map_db_error; -use chrono::{DateTime, Utc}; -use services::common::RepositoryError; -use tokio_postgres::Transaction; -use uuid::Uuid; - -/// Add one usage charge to the per-key split counters inside its posting transaction. -pub async fn increment_api_key_spend( - transaction: &Transaction<'_>, - api_key_id: Uuid, - inference_spent: i64, - service_spent: i64, - updated_at: DateTime, -) -> Result<(), RepositoryError> { - transaction - .execute( - r#" - INSERT INTO api_key_spend ( - api_key_id, inference_spent, service_spent, updated_at - ) VALUES ($1, $2, $3, $4) - ON CONFLICT (api_key_id) DO UPDATE SET - inference_spent = api_key_spend.inference_spent + $2, - service_spent = api_key_spend.service_spent + $3, - updated_at = $4 - "#, - &[&api_key_id, &inference_spent, &service_spent, &updated_at], - ) - .await - .map_err(map_db_error)?; - Ok(()) -} From 0fb7a71aa752fd17c66501230651769d89b41be5 Mon Sep 17 00:00:00 2001 From: Henry Park <16583448+henrypark133@users.noreply.github.com> Date: Tue, 22 Sep 2026 19:19:17 +0000 Subject: [PATCH 6/7] Address PR review feedback (#1119) - Cover CLI batching, anomalies, timeout bounds, and argument validation - Clarify fixed apply deadlines and clean test databases after failures --- .../src/bin/backfill-spend-counters.rs | 99 +++++- .../database/src/spend_counters_backfill.rs | 1 + .../database/tests/spend_counter_backfill.rs | 214 +++++++------ .../tests/spend_counter_backfill/coverage.rs | 281 ++++++++++++++++++ 4 files changed, 497 insertions(+), 98 deletions(-) create mode 100644 crates/database/tests/spend_counter_backfill/coverage.rs diff --git a/crates/database/src/bin/backfill-spend-counters.rs b/crates/database/src/bin/backfill-spend-counters.rs index 8eebffc71..a780fa8f0 100644 --- a/crates/database/src/bin/backfill-spend-counters.rs +++ b/crates/database/src/bin/backfill-spend-counters.rs @@ -105,8 +105,12 @@ fn parse_args(args: Vec) -> Result> { timeout_seconds = value .parse() .with_context(|| format!("invalid timeout seconds: {value}"))?; - if timeout_seconds == 0 { - bail!("--statement-timeout-seconds must be positive"); + let timeout = Duration::from_secs(timeout_seconds); + if timeout.as_millis() == 0 || timeout.as_millis() > i32::MAX as u128 { + bail!( + "--statement-timeout-seconds must produce a timeout between 1ms and {}ms", + i32::MAX + ); } } unknown => bail!("unknown argument {unknown}; use --help"), @@ -123,6 +127,95 @@ fn print_help() { "Usage: backfill-spend-counters [--organization UUID] [--statement-timeout-seconds N]\n\n\ Reconciles incomplete organization spend counters. Deploy counter writers (#1116)\n\ to every process and drain all older writers first. Do not rewrite or delete raw usage history\n\ -while this process runs. The command does not run migrations." +while this process runs. The statement timeout applies to the repeatable-read snapshot;\n\ +apply and accounting-lock statements remain capped at 5 seconds. The command does not run migrations." ); } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_args_defaults() { + let options = parse_args(Vec::new()).unwrap().unwrap(); + assert_eq!(options.organization_id, None); + assert_eq!(options.statement_timeout, Duration::from_secs(300)); + } + + #[test] + fn parse_args_accepts_organization_uuid() { + let organization_id = Uuid::new_v4(); + let options = parse_args(vec![ + "--organization".to_string(), + organization_id.to_string(), + ]) + .unwrap() + .unwrap(); + assert_eq!(options.organization_id, Some(organization_id)); + } + + #[test] + fn parse_args_rejects_missing_organization_uuid() { + let error = parse_args(vec!["--organization".to_string()]).unwrap_err(); + assert!(error.to_string().contains("--organization needs a UUID")); + } + + #[test] + fn parse_args_rejects_invalid_organization_uuid() { + let error = + parse_args(vec!["--organization".to_string(), "not-a-uuid".to_string()]).unwrap_err(); + assert!(error.to_string().contains("invalid organization UUID")); + } + + #[test] + fn parse_args_rejects_missing_statement_timeout() { + let error = parse_args(vec!["--statement-timeout-seconds".to_string()]).unwrap_err(); + assert!(error + .to_string() + .contains("--statement-timeout-seconds needs an integer")); + } + + #[test] + fn parse_args_rejects_invalid_statement_timeout() { + let error = parse_args(vec![ + "--statement-timeout-seconds".to_string(), + "not-a-number".to_string(), + ]) + .unwrap_err(); + assert!(error.to_string().contains("invalid timeout seconds")); + } + + #[test] + fn parse_args_rejects_zero_statement_timeout() { + let error = parse_args(vec![ + "--statement-timeout-seconds".to_string(), + "0".to_string(), + ]) + .unwrap_err(); + assert!(error.to_string().contains("must produce a timeout")); + } + + #[test] + fn parse_args_rejects_oversized_statement_timeout() { + let seconds = (i32::MAX as u64 / 1000) + 1; + let error = parse_args(vec![ + "--statement-timeout-seconds".to_string(), + seconds.to_string(), + ]) + .unwrap_err(); + assert!(error.to_string().contains("must produce a timeout")); + } + + #[test] + fn parse_args_help_returns_none() { + assert!(parse_args(vec!["--help".to_string()]).unwrap().is_none()); + assert!(parse_args(vec!["-h".to_string()]).unwrap().is_none()); + } + + #[test] + fn parse_args_rejects_unknown_flag() { + let error = parse_args(vec!["--unknown".to_string()]).unwrap_err(); + assert!(error.to_string().contains("unknown argument --unknown")); + } +} diff --git a/crates/database/src/spend_counters_backfill.rs b/crates/database/src/spend_counters_backfill.rs index 9ee2311d1..c63f355d0 100644 --- a/crates/database/src/spend_counters_backfill.rs +++ b/crates/database/src/spend_counters_backfill.rs @@ -7,6 +7,7 @@ use tokio_postgres::{IsolationLevel, Transaction}; use uuid::Uuid; const LOCK_TIMEOUT: Duration = Duration::from_secs(5); +// ponytail: each apply statement is capped at 5s; larger units require a resumable protocol. const APPLY_TIMEOUT: Duration = Duration::from_secs(5); #[derive(Debug, Clone, Copy, PartialEq, Eq)] diff --git a/crates/database/tests/spend_counter_backfill.rs b/crates/database/tests/spend_counter_backfill.rs index a5e47c6b2..4d95387b3 100644 --- a/crates/database/tests/spend_counter_backfill.rs +++ b/crates/database/tests/spend_counter_backfill.rs @@ -7,10 +7,14 @@ use database::{ensure_spend_counters_ready, migrations, DbPool, PreparedSpendBac use deadpool::Runtime; use deadpool_postgres::{Config, PoolConfig, Timeouts}; use services::usage::InferenceType; +use std::future::Future; use std::time::Duration; use tokio_postgres::NoTls; use uuid::Uuid; +#[path = "spend_counter_backfill/coverage.rs"] +mod coverage; + struct TestDatabase { pool: DbPool, admin: DbPool, @@ -30,7 +34,7 @@ struct Fixture { #[tokio::test] async fn backfill_reconciles_snapshot_delta_and_is_idempotent() -> anyhow::Result<()> { - let database = test_database().await?; + run_backfill_test(|database| async move { let fixture = fixture(&database.pool).await?; // Historical rows exceed the C1 counters already present for the mixed key. @@ -148,13 +152,14 @@ async fn backfill_reconciles_snapshot_delta_and_is_idempotent() -> anyhow::Resul drop(client); drop(inference_repository); drop(service_repository); - database.cleanup().await?; Ok(()) + }) + .await } #[tokio::test] async fn backfill_rejects_negative_drift_and_rolls_back_overflow() -> anyhow::Result<()> { - let database = test_database().await?; + run_backfill_test(|database| async move { let fixture = fixture(&database.pool).await?; let client = database.pool.get().await?; client @@ -312,103 +317,134 @@ async fn backfill_rejects_negative_drift_and_rolls_back_overflow() -> anyhow::Re assert_eq!(balance.get::<_, i64>(0), i64::MAX); assert!(balance.get::<_, Option>>(1).is_none()); drop(client); - database.cleanup().await?; Ok(()) + }) + .await } #[tokio::test] async fn spend_readiness_checks_missing_and_inactive_organizations() -> anyhow::Result<()> { - let database = test_database().await?; - let fixture = fixture(&database.pool).await?; - let client = database.pool.get().await?; - client - .execute( - "UPDATE organization_balance SET spend_counters_ready_at = NULL + run_backfill_test(|database| async move { + let fixture = fixture(&database.pool).await?; + let client = database.pool.get().await?; + client + .execute( + "UPDATE organization_balance SET spend_counters_ready_at = NULL WHERE organization_id = $1", - &[&fixture.organization_id], + &[&fixture.organization_id], + ) + .await?; + drop(client); + let prepared = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), ) - .await?; - drop(client); - let prepared = PreparedSpendBackfill::prepare( - &database.pool, - fixture.organization_id, - Duration::from_secs(30), - ) - .await? - .expect("empty historical organization is incomplete"); - assert_eq!( - prepared.apply().await?, - database::SpendBackfillOutcome::Applied { key_count: 0 } - ); - ensure_spend_counters_ready(&database.pool).await?; + .await? + .expect("empty historical organization is incomplete"); + assert_eq!( + prepared.apply().await?, + database::SpendBackfillOutcome::Applied { key_count: 0 } + ); + ensure_spend_counters_ready(&database.pool).await?; - let mut client = database.pool.get().await?; - let migration_tx = client.build_transaction().start().await?; - migration_tx - .batch_execute( - "CREATE TEMP TABLE organization_balance ( + let mut client = database.pool.get().await?; + let migration_tx = client.build_transaction().start().await?; + migration_tx + .batch_execute( + "CREATE TEMP TABLE organization_balance ( organization_id UUID PRIMARY KEY, inference_spent BIGINT NOT NULL DEFAULT 0, service_spent BIGINT NOT NULL DEFAULT 0 ); INSERT INTO organization_balance (organization_id) VALUES (uuid_generate_v4());", - ) - .await?; - migration_tx - .batch_execute(include_str!( - "../src/migrations/sql/V0082__add_spend_counters_readiness.sql" - )) - .await?; - let legacy = migration_tx - .query_one( - "SELECT spend_counters_ready_at FROM organization_balance LIMIT 1", - &[], - ) - .await?; - assert!(legacy.get::<_, Option>>(0).is_none()); - let future = migration_tx - .query_one( - "INSERT INTO organization_balance (organization_id) + ) + .await?; + migration_tx + .batch_execute(include_str!( + "../src/migrations/sql/V0082__add_spend_counters_readiness.sql" + )) + .await?; + let legacy = migration_tx + .query_one( + "SELECT spend_counters_ready_at FROM organization_balance LIMIT 1", + &[], + ) + .await?; + assert!(legacy.get::<_, Option>>(0).is_none()); + let future = migration_tx + .query_one( + "INSERT INTO organization_balance (organization_id) VALUES (uuid_generate_v4()) RETURNING spend_counters_ready_at", - &[], - ) - .await?; - assert!(future.get::<_, Option>>(0).is_some()); - migration_tx.rollback().await?; + &[], + ) + .await?; + assert!(future.get::<_, Option>>(0).is_some()); + migration_tx.rollback().await?; - client - .execute( - "UPDATE organizations SET is_active = false WHERE id = $1", - &[&fixture.organization_id], - ) - .await?; - ensure_spend_counters_ready(&database.pool).await?; - client - .execute( - "UPDATE organization_balance SET spend_counters_ready_at = NULL + client + .execute( + "UPDATE organizations SET is_active = false WHERE id = $1", + &[&fixture.organization_id], + ) + .await?; + ensure_spend_counters_ready(&database.pool).await?; + client + .execute( + "UPDATE organization_balance SET spend_counters_ready_at = NULL WHERE organization_id = $1", - &[&fixture.organization_id], - ) - .await?; - assert!(ensure_spend_counters_ready(&database.pool) - .await - .expect_err("inactive incomplete organizations must block readers") - .to_string() - .contains("remain unreconciled")); - client - .execute( - "DELETE FROM organization_balance WHERE organization_id = $1", - &[&fixture.organization_id], - ) - .await?; - assert!(ensure_spend_counters_ready(&database.pool) - .await - .expect_err("missing balances must be distinguished") - .to_string() - .contains("lack balances")); - drop(client); - database.cleanup().await?; + &[&fixture.organization_id], + ) + .await?; + assert!(ensure_spend_counters_ready(&database.pool) + .await + .expect_err("inactive incomplete organizations must block readers") + .to_string() + .contains("remain unreconciled")); + client + .execute( + "DELETE FROM organization_balance WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + assert!(ensure_spend_counters_ready(&database.pool) + .await + .expect_err("missing balances must be distinguished") + .to_string() + .contains("lack balances")); + drop(client); + Ok(()) + }) + .await +} + +async fn run_backfill_test(test: F) -> anyhow::Result<()> +where + F: FnOnce(TestDatabase) -> Fut + Send + 'static, + Fut: Future> + Send + 'static, +{ + let database = test_database().await?; + let admin = database.admin.clone(); + let database_name = database.database_name.clone(); + let result = tokio::spawn(test(database)).await; + let cleanup: anyhow::Result<()> = async { + let client = admin.get().await?; + client + .batch_execute(&format!( + "DROP DATABASE IF EXISTS {} WITH (FORCE)", + database_name + )) + .await?; + Ok(()) + } + .await; + let test_result = match result { + Ok(result) => result, + Err(error) => Err(anyhow::Error::new(error).context("backfill test task failed")), + }; + test_result?; + cleanup?; Ok(()) } @@ -646,15 +682,3 @@ async fn insert_service( .await?; Ok(()) } - -impl TestDatabase { - async fn cleanup(self) -> anyhow::Result<()> { - drop(self.pool); - self.admin - .get() - .await? - .batch_execute(&format!("DROP DATABASE {}", self.database_name)) - .await?; - Ok(()) - } -} diff --git a/crates/database/tests/spend_counter_backfill/coverage.rs b/crates/database/tests/spend_counter_backfill/coverage.rs new file mode 100644 index 000000000..0c95f5803 --- /dev/null +++ b/crates/database/tests/spend_counter_backfill/coverage.rs @@ -0,0 +1,281 @@ +use super::*; +use database::spend_counters_backfill::incomplete_organizations; +use std::process::{Output, Stdio}; +use std::time::{Duration, Instant}; +use tokio::process::Command; +use tokio::time::timeout; + +fn stable_id(value: u128) -> Uuid { + Uuid::from_u128(value) +} + +async fn insert_empty_organization(pool: &DbPool, organization_id: Uuid) -> anyhow::Result<()> { + let client = pool.get().await?; + client + .execute( + "INSERT INTO organizations (id, name, created_at, updated_at) + VALUES ($1, $2, NOW(), NOW())", + &[ + &organization_id, + &format!("backfill-test-{organization_id}"), + ], + ) + .await?; + Ok(()) +} + +async fn set_incomplete(pool: &DbPool, organization_id: Uuid) -> anyhow::Result<()> { + let client = pool.get().await?; + client + .execute( + "UPDATE organization_balance + SET spend_counters_ready_at = NULL + WHERE organization_id = $1", + &[&organization_id], + ) + .await?; + Ok(()) +} + +#[tokio::test] +async fn incomplete_organizations_orders_and_pages_ready_incomplete_and_missing_rows( +) -> anyhow::Result<()> { + run_backfill_test(|database| async move { + let ready = stable_id(1); + let incomplete = stable_id(2); + let missing = stable_id(3); + for organization_id in [ready, incomplete, missing] { + insert_empty_organization(&database.pool, organization_id).await?; + } + set_incomplete(&database.pool, incomplete).await?; + let client = database.pool.get().await?; + client + .execute( + "DELETE FROM organization_balance WHERE organization_id = $1", + &[&missing], + ) + .await?; + drop(client); + + let first = incomplete_organizations(&database.pool, None, 1).await?; + assert_eq!(first, vec![incomplete]); + let second = incomplete_organizations(&database.pool, first.last().copied(), 1).await?; + assert_eq!(second, vec![missing]); + let third = incomplete_organizations(&database.pool, second.last().copied(), 1).await?; + assert!(third.is_empty()); + Ok(()) + }) + .await +} + +#[tokio::test] +async fn prepare_fails_for_organization_without_balance_row() -> anyhow::Result<()> { + run_backfill_test(|database| async move { + let fixture = fixture(&database.pool).await?; + let client = database.pool.get().await?; + client + .execute( + "DELETE FROM organization_balance WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await?; + drop(client); + + let error = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await + .expect_err("missing balance row must fail preparation"); + assert_eq!( + error.to_string(), + format!( + "organization {} has no organization_balance row", + fixture.organization_id + ) + ); + Ok(()) + }) + .await +} + +#[tokio::test] +async fn prepare_rejects_zero_submillisecond_and_oversized_timeouts() -> anyhow::Result<()> { + run_backfill_test(|database| async move { + let organization_id = Uuid::new_v4(); + let timeouts = [ + Duration::ZERO, + Duration::from_nanos(1), + Duration::from_millis(i32::MAX as u64 + 1), + ]; + for statement_timeout in timeouts { + let error = + PreparedSpendBackfill::prepare(&database.pool, organization_id, statement_timeout) + .await + .expect_err("invalid timeout must fail before acquiring a connection"); + assert!(error.to_string().contains("between 1ms and 2147483647ms")); + } + Ok(()) + }) + .await +} + +#[tokio::test] +async fn apply_times_out_when_accounting_lock_is_held() -> anyhow::Result<()> { + run_backfill_test(|database| async move { + let fixture = fixture(&database.pool).await?; + set_incomplete(&database.pool, fixture.organization_id).await?; + let prepared = PreparedSpendBackfill::prepare( + &database.pool, + fixture.organization_id, + Duration::from_secs(30), + ) + .await? + .expect("fixture organization is incomplete"); + + let mut lock_client = database.pool.get().await?; + let lock_transaction = lock_client.transaction().await?; + lock_transaction + .query_one( + "SELECT id FROM organizations WHERE id = $1 FOR UPDATE", + &[&fixture.organization_id], + ) + .await?; + + let started = Instant::now(); + let result = timeout(Duration::from_secs(15), prepared.apply()) + .await + .expect("apply must finish within the outer 15 second bound"); + let elapsed = started.elapsed(); + let error = result.expect_err("the held accounting lock must time out"); + assert!(elapsed >= Duration::from_secs(4), "apply returned too early: {elapsed:?}"); + assert!(elapsed < Duration::from_secs(15), "apply exceeded outer bound: {elapsed:?}"); + // Both deadlines are 5s. map_db_error maps statement cancellation (57014) + // to QueryTimeout, while a lock timeout retains SQLSTATE 55P03. + let repository_error = error + .downcast_ref::() + .expect("accounting lock failures retain their repository error type"); + match repository_error { + services::common::RepositoryError::QueryTimeout => {}, + services::common::RepositoryError::DatabaseError(cause) => assert!( + cause.to_string().starts_with("Database error (55P03): "), + "expected lock timeout SQLSTATE 55P03, got: {cause}" + ), + other => panic!("unexpected accounting lock failure: {other}"), + } + + let client = database.pool.get().await?; + let ready_at = client + .query_one( + "SELECT spend_counters_ready_at FROM organization_balance WHERE organization_id = $1", + &[&fixture.organization_id], + ) + .await? + .get::<_, Option>>(0); + assert!(ready_at.is_none(), "failed apply must leave readiness unchanged"); + drop(client); + lock_transaction.rollback().await?; + Ok(()) + }) + .await +} + +fn cli_command(database_name: &str) -> Command { + let config = pool_config(None); + let host = config.host.unwrap_or_else(|| "localhost".to_string()); + let port = config.port.unwrap_or(5432).to_string(); + let user = config.user.unwrap_or_else(|| "postgres".to_string()); + let password = config.password.unwrap_or_else(|| "postgres".to_string()); + let mut command = Command::new(env!("CARGO_BIN_EXE_backfill-spend-counters")); + command + .env("DATABASE_CONNECTION_MODE", "patroni") + .env("POSTGRES_PRIMARY_APP_ID", "postgres-test") + .env("GATEWAY_SUBDOMAIN", "localhost") + .env("DATABASE_HOST", host) + .env("DATABASE_PORT", port) + .env("DATABASE_NAME", database_name) + .env("DATABASE_USERNAME", user) + .env("DATABASE_PASSWORD", password) + .env("DATABASE_TLS_ENABLED", "false") + .env("DATABASE_MAX_CONNECTIONS", "4") + .env("DATABASE_REFRESH_INTERVAL", "30") + .stdin(std::process::Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()); + command +} + +async fn run_cli(database_name: &str, args: &[&str]) -> anyhow::Result { + let mut command = cli_command(database_name); + command.args(args).kill_on_drop(true); + Ok(timeout(Duration::from_secs(60), command.output()).await??) +} + +#[tokio::test] +async fn cli_reconciles_more_than_one_batch_of_empty_organizations() -> anyhow::Result<()> { + run_backfill_test(|database| async move { + for value in 1..=101 { + let organization_id = stable_id(10_000 + value); + insert_empty_organization(&database.pool, organization_id).await?; + set_incomplete(&database.pool, organization_id).await?; + } + + let output = run_cli(&database.database_name, &[]).await?; + assert!( + output.status.success(), + "CLI failed: {}", + String::from_utf8_lossy(&output.stderr) + ); + let client = database.pool.get().await?; + let ready_count: i64 = client + .query_one( + "SELECT COUNT(*) FROM organization_balance + WHERE spend_counters_ready_at IS NOT NULL", + &[], + ) + .await? + .get(0); + assert_eq!(ready_count, 101); + Ok(()) + }) + .await +} + +#[tokio::test] +async fn cli_stops_on_anomaly_and_preserves_fail_fast_exit() -> anyhow::Result<()> { + run_backfill_test(|database| async move { + let anomalous = stable_id(20_000); + let remaining = stable_id(20_001); + insert_empty_organization(&database.pool, anomalous).await?; + insert_empty_organization(&database.pool, remaining).await?; + set_incomplete(&database.pool, remaining).await?; + let client = database.pool.get().await?; + client + .execute( + "DELETE FROM organization_balance WHERE organization_id = $1", + &[&anomalous], + ) + .await?; + drop(client); + + let output = run_cli(&database.database_name, &[]).await?; + assert!(!output.status.success(), "anomalous organization must fail the CLI"); + assert!( + String::from_utf8_lossy(&output.stderr).contains("no organization_balance row"), + "CLI stderr did not identify the anomaly: {}", + String::from_utf8_lossy(&output.stderr) + ); + let client = database.pool.get().await?; + let ready_at = client + .query_one( + "SELECT spend_counters_ready_at FROM organization_balance WHERE organization_id = $1", + &[&remaining], + ) + .await? + .get::<_, Option>>(0); + assert!(ready_at.is_none(), "fail-fast CLI must not reconcile later organizations"); + Ok(()) + }) + .await +} From 47cdf1a10765a10d1f14f670ff21bc6adcf9926d Mon Sep 17 00:00:00 2001 From: Henry Park <16583448+henrypark133@users.noreply.github.com> Date: Tue, 22 Sep 2026 19:19:19 +0000 Subject: [PATCH 7/7] Address PR review feedback (#1120) - Reuse shared database test configuration for standalone reader tests - Clean isolated schemas after test errors and assertion panics --- crates/api/tests/e2e_all/api_keys.rs | 17 +- .../src/repositories/organization_usage.rs | 8 +- crates/database/tests/counter_readers.rs | 230 +++++++++--------- crates/database/tests/support/mod.rs | 24 +- crates/services/src/admin/analytics.rs | 6 +- crates/services/src/usage/ports.rs | 2 +- 6 files changed, 148 insertions(+), 139 deletions(-) diff --git a/crates/api/tests/e2e_all/api_keys.rs b/crates/api/tests/e2e_all/api_keys.rs index 4e72cee80..0dd97b4c2 100644 --- a/crates/api/tests/e2e_all/api_keys.rs +++ b/crates/api/tests/e2e_all/api_keys.rs @@ -135,12 +135,12 @@ async fn test_list_workspace_api_keys_orders_by_usage() { drop(client); let inference_repository = OrganizationUsageRepository::new(database.pool().clone()); - let mut inference = RecordUsageRequest { + let inference = RecordUsageRequest { organization_id, workspace_id, api_key_id: service_spend_key_id, model_id, - model_name: model_name.clone(), + model_name, input_tokens: 10, output_tokens: 10, input_cost: 1, @@ -167,10 +167,15 @@ async fn test_list_workspace_api_keys_orders_by_usage() { .record_usage(inference.clone()) .await .unwrap(); - inference.api_key_id = inference_spend_key_id; - inference.total_cost = 300_000_000; - inference.inference_id = Some(uuid::Uuid::new_v4()); - inference_repository.record_usage(inference).await.unwrap(); + inference_repository + .record_usage(RecordUsageRequest { + api_key_id: inference_spend_key_id, + total_cost: 300_000_000, + inference_id: Some(uuid::Uuid::new_v4()), + ..inference + }) + .await + .unwrap(); OrganizationServiceUsageRepository::new(database.pool().clone()) .record_usage(&RecordServiceUsageRequest { organization_id, diff --git a/crates/database/src/repositories/organization_usage.rs b/crates/database/src/repositories/organization_usage.rs index fecc7e99d..92ee6d173 100644 --- a/crates/database/src/repositories/organization_usage.rs +++ b/crates/database/src/repositories/organization_usage.rs @@ -56,7 +56,7 @@ impl OrganizationUsageRepository { } } - /// Get inference spend for a specific API key for admission-limit checks. + /// Get inference-only spend for a specific API key for admission-limit checks. pub async fn get_api_key_spend(&self, api_key_id: Uuid) -> Result { let row = retry_db!("get_api_key_spend", { let client = self @@ -72,7 +72,7 @@ impl OrganizationUsageRepository { SELECT COALESCE( (SELECT inference_spent FROM api_key_spend WHERE api_key_id = $1), 0 - )::BIGINT as total_spend + )::BIGINT as inference_spend "#, &[&api_key_id], ) @@ -80,8 +80,8 @@ impl OrganizationUsageRepository { .map_err(map_db_error) })?; - let total_spend: i64 = row.get("total_spend"); - Ok(total_spend) + let inference_spend: i64 = row.get("inference_spend"); + Ok(inference_spend) } /// Record usage and update balance atomically. diff --git a/crates/database/tests/counter_readers.rs b/crates/database/tests/counter_readers.rs index 5135e22dd..363aaf417 100644 --- a/crates/database/tests/counter_readers.rs +++ b/crates/database/tests/counter_readers.rs @@ -1,3 +1,6 @@ +#[allow(dead_code)] +mod support; + use database::repositories::{ ApiKeyRepository, OrganizationUsageRepository, PgAnalyticsRepository, }; @@ -6,35 +9,13 @@ use deadpool::Runtime; use deadpool_postgres::{Config, PoolConfig}; use services::admin::AnalyticsRepository; use services::workspace::ports::{ApiKeyOrderBy, ApiKeyOrderDirection}; +use support::pool_config; use tokio_postgres::NoTls; use uuid::Uuid; -fn env_value(primary: &str, fallback: &str, default: &str) -> String { - std::env::var(primary) - .or_else(|_| std::env::var(fallback)) - .unwrap_or_else(|_| default.to_string()) -} - -fn pool_config() -> Config { - let mut config = Config::new(); - config.host = Some(env_value("PGHOST", "DATABASE_HOST", "localhost")); - config.port = Some( - env_value("PGPORT", "DATABASE_PORT", "5432") - .parse() - .expect("database port must be numeric"), - ); - config.dbname = Some(env_value( - "PGDATABASE", - "TEST_DATABASE_NAME", - "platform_api", - )); - config.user = Some(env_value("PGUSER", "DATABASE_USERNAME", "postgres")); - config.password = Some(env_value("PGPASSWORD", "DATABASE_PASSWORD", "postgres")); - config.pool = Some(PoolConfig::new(4)); - config -} - async fn new_pool(config: &Config) -> anyhow::Result { + let mut config = config.clone(); + config.pool = Some(PoolConfig::new(4)); Ok(DbPool::new( config.create_pool(Some(Runtime::Tokio1), NoTls)?, )) @@ -113,6 +94,19 @@ async fn drop_schema( Ok(()) } +// A spawned task lets cleanup run after either a returned error or an assertion panic. +async fn with_scoped_pool(test: F) -> anyhow::Result<()> +where + F: FnOnce(DbPool) -> Fut, + Fut: std::future::Future> + Send + 'static, +{ + let (admin_pool, pool, schema) = scoped_pool().await?; + let result = tokio::spawn(test(pool.clone())).await; + let cleanup = drop_schema(admin_pool, pool, schema).await; + result??; + cleanup +} + async fn insert_key( pool: &DbPool, workspace_id: Uuid, @@ -151,104 +145,106 @@ async fn insert_key( #[tokio::test] async fn key_list_reads_counter_cohort_with_stable_pagination() -> anyhow::Result<()> { - let (admin_pool, pool, schema) = scoped_pool().await?; - let workspace_id = Uuid::new_v4(); - let zero = Uuid::new_v4(); - let inference = Uuid::new_v4(); - let service = Uuid::new_v4(); - let mixed = Uuid::new_v4(); - let deleted = Uuid::new_v4(); - let other_workspace = Uuid::new_v4(); - let other_workspace_key = Uuid::new_v4(); - insert_key(&pool, workspace_id, zero, "zero", 0, 0, false).await?; - insert_key(&pool, workspace_id, inference, "inference", 30, 0, false).await?; - insert_key(&pool, workspace_id, service, "service", 0, 40, false).await?; - insert_key(&pool, workspace_id, mixed, "mixed", 50, 70, false).await?; - insert_key(&pool, workspace_id, deleted, "deleted", 900, 900, true).await?; - insert_key( - &pool, - other_workspace, - other_workspace_key, - "other-workspace", - 9_000, - 9_000, - false, - ) - .await?; - pool.get() - .await? - .execute("DELETE FROM api_key_spend WHERE api_key_id = $1", &[&zero]) - .await?; - - // The isolated schema deliberately has no raw usage tables. - ensure_spend_counters_ready(&pool).await?; - let repository = ApiKeyRepository::new(pool.clone()); - let first_page = repository - .list_by_workspace_paginated( - workspace_id, - 3, - 0, - Some(ApiKeyOrderBy::Usage), - Some(ApiKeyOrderDirection::Desc), + with_scoped_pool(|pool| async move { + let workspace_id = Uuid::new_v4(); + let zero = Uuid::new_v4(); + let inference = Uuid::new_v4(); + let service = Uuid::new_v4(); + let mixed = Uuid::new_v4(); + let deleted = Uuid::new_v4(); + let other_workspace = Uuid::new_v4(); + let other_workspace_key = Uuid::new_v4(); + insert_key(&pool, workspace_id, zero, "zero", 0, 0, false).await?; + insert_key(&pool, workspace_id, inference, "inference", 30, 0, false).await?; + insert_key(&pool, workspace_id, service, "service", 0, 40, false).await?; + insert_key(&pool, workspace_id, mixed, "mixed", 50, 70, false).await?; + insert_key(&pool, workspace_id, deleted, "deleted", 900, 900, true).await?; + insert_key( + &pool, + other_workspace, + other_workspace_key, + "other-workspace", + 9_000, + 9_000, + false, ) .await?; - assert_eq!( - first_page.iter().map(|key| key.id).collect::>(), - vec![mixed, service, inference] - ); - assert_eq!( - first_page.iter().map(|key| key.usage).collect::>(), - vec![120, 40, 30] - ); + pool.get() + .await? + .execute("DELETE FROM api_key_spend WHERE api_key_id = $1", &[&zero]) + .await?; - let second_page = repository - .list_by_workspace_paginated( - workspace_id, - 3, - 3, - Some(ApiKeyOrderBy::Usage), - Some(ApiKeyOrderDirection::Desc), - ) - .await?; - assert_eq!( - second_page.iter().map(|key| key.id).collect::>(), - vec![zero] - ); - assert_eq!(second_page[0].usage, 0); - assert!(!second_page.iter().any(|key| key.id == deleted)); + // The isolated schema deliberately has no raw usage tables. + ensure_spend_counters_ready(&pool).await?; + let repository = ApiKeyRepository::new(pool.clone()); + let first_page = repository + .list_by_workspace_paginated( + workspace_id, + 3, + 0, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Desc), + ) + .await?; + assert_eq!( + first_page.iter().map(|key| key.id).collect::>(), + vec![mixed, service, inference] + ); + assert_eq!( + first_page.iter().map(|key| key.usage).collect::>(), + vec![120, 40, 30] + ); - let ascending = repository - .list_by_workspace_paginated( - workspace_id, - 2, - 0, - Some(ApiKeyOrderBy::Usage), - Some(ApiKeyOrderDirection::Asc), - ) - .await?; - assert_eq!( - ascending.iter().map(|key| key.id).collect::>(), - vec![zero, inference] - ); - let past_end = repository - .list_by_workspace_paginated( - workspace_id, - 2, - 99, - Some(ApiKeyOrderBy::Usage), - Some(ApiKeyOrderDirection::Desc), - ) - .await?; - assert!(past_end.is_empty()); - assert!(!first_page.iter().any(|key| key.id == other_workspace_key)); + let second_page = repository + .list_by_workspace_paginated( + workspace_id, + 3, + 3, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Desc), + ) + .await?; + assert_eq!( + second_page.iter().map(|key| key.id).collect::>(), + vec![zero] + ); + assert_eq!(second_page[0].usage, 0); + assert!(!second_page.iter().any(|key| key.id == deleted)); - drop_schema(admin_pool, pool, schema).await + let ascending = repository + .list_by_workspace_paginated( + workspace_id, + 2, + 0, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Asc), + ) + .await?; + assert_eq!( + ascending.iter().map(|key| key.id).collect::>(), + vec![zero, inference] + ); + let past_end = repository + .list_by_workspace_paginated( + workspace_id, + 2, + 99, + Some(ApiKeyOrderBy::Usage), + Some(ApiKeyOrderDirection::Desc), + ) + .await?; + assert!(past_end.is_empty()); + assert!(!first_page.iter().any(|key| key.id == other_workspace_key)); + + Ok(()) + }) + .await } #[tokio::test] async fn admission_spend_reads_inference_counter_and_excludes_service_counter() -> anyhow::Result<()> { - let (admin_pool, pool, schema) = scoped_pool().await?; + with_scoped_pool(|pool| async move { let inference_key = Uuid::new_v4(); let service_key = Uuid::new_v4(); let now = chrono::Utc::now(); @@ -272,12 +268,13 @@ async fn admission_spend_reads_inference_counter_and_excludes_service_counter() assert_eq!(repository.get_api_key_spend(service_key).await?, 0); assert_eq!(repository.get_api_key_spend(Uuid::new_v4()).await?, 0); - drop_schema(admin_pool, pool, schema).await + Ok(()) + }).await } #[tokio::test] async fn billing_summary_reads_balance_splits_and_preserves_legacy_total() -> anyhow::Result<()> { - let (admin_pool, pool, schema) = scoped_pool().await?; + with_scoped_pool(|pool| async move { let organization_id = Uuid::new_v4(); let client = pool.get().await?; client @@ -308,5 +305,6 @@ async fn billing_summary_reads_balance_splits_and_preserves_legacy_total() -> an assert_eq!(summary.inference_consumed_usd, 35e-9); assert_eq!(summary.service_consumed_usd, 20e-9); - drop_schema(admin_pool, pool, schema).await + Ok(()) + }).await } diff --git a/crates/database/tests/support/mod.rs b/crates/database/tests/support/mod.rs index cc3090c6d..5b4445ca6 100644 --- a/crates/database/tests/support/mod.rs +++ b/crates/database/tests/support/mod.rs @@ -223,18 +223,22 @@ pub async fn cleanup_usage_fixtures( Ok(()) } -fn pool_config() -> Config { +fn env_value(primary: &str, fallback: &str, default: &str) -> String { + std::env::var(primary) + .or_else(|_| std::env::var(fallback)) + .unwrap_or_else(|_| default.to_string()) +} + +pub fn pool_config() -> Config { let mut config = Config::new(); - config.host = Some(std::env::var("PGHOST").unwrap_or_else(|_| "localhost".to_string())); + config.host = Some(env_value("PGHOST", "DATABASE_HOST", "localhost")); config.port = Some( - std::env::var("PGPORT") - .ok() - .and_then(|value| value.parse::().ok()) - .unwrap_or(5432), + env_value("PGPORT", "DATABASE_PORT", "5432") + .parse() + .expect("database port must be numeric"), ); - config.dbname = - Some(std::env::var("PGDATABASE").unwrap_or_else(|_| "platform_api".to_string())); - config.user = Some(std::env::var("PGUSER").unwrap_or_else(|_| "postgres".to_string())); - config.password = Some(std::env::var("PGPASSWORD").unwrap_or_else(|_| "postgres".to_string())); + config.dbname = Some(env_value("PGDATABASE", "DATABASE_NAME", "platform_api")); + config.user = Some(env_value("PGUSER", "DATABASE_USERNAME", "postgres")); + config.password = Some(env_value("PGPASSWORD", "DATABASE_PASSWORD", "postgres")); config } diff --git a/crates/services/src/admin/analytics.rs b/crates/services/src/admin/analytics.rs index 4f9a7ad5f..2ac012205 100644 --- a/crates/services/src/admin/analytics.rs +++ b/crates/services/src/admin/analytics.rs @@ -219,8 +219,10 @@ pub struct BillingSummary { pub active_paid_credit_limit_usd: f64, /// Sum of active grant-type spend limits (caps), USD pub active_grant_credit_limit_usd: f64, - /// All-time consumed cost across all orgs, USD — **all usage** (from - /// organization_balance: inference + services). + /// All-time consumed cost across all orgs, USD — legacy + /// `organization_balance.total_spent` total, kept independent of the + /// inference/service splits (historical adjustments can diverge from + /// inference + services). pub total_consumed_usd: f64, /// All-time inference consumed cost, USD (organization_balance.inference_spent) pub inference_consumed_usd: f64, diff --git a/crates/services/src/usage/ports.rs b/crates/services/src/usage/ports.rs index 210b8a7ba..94ca516c3 100644 --- a/crates/services/src/usage/ports.rs +++ b/crates/services/src/usage/ports.rs @@ -366,7 +366,7 @@ pub trait UsageRepository: Send + Sync { offset: Option, ) -> anyhow::Result<(Vec, i64)>; - /// Get total spend for a specific API key + /// Get inference-only spend for a specific API key async fn get_api_key_spend(&self, api_key_id: Uuid) -> anyhow::Result; /// Get costs by inference IDs (for HuggingFace billing integration)