From 911c7f88754d2fea48641ff16960633bf1284d29 Mon Sep 17 00:00:00 2001 From: elky Date: Wed, 7 Oct 2026 14:29:48 +0800 Subject: [PATCH] feat: add selectable routing groups and composite billing Support per-model provider enablement and compact model editing. Capture request-time billing factors, charge customer costs separately, and preserve historical statistics without backfills. --- .../planner/candidate_materialization.rs | 2 + .../ai_serving/planner/candidate_ranking.rs | 6 + .../ai_serving/planner/candidate_source.rs | 113 ++++++ .../src/ai_serving/planner/decision_input.rs | 236 +++++++++++- .../src/ai_serving/planner/report_context.rs | 89 ++++- apps/aether-gateway/src/control/auth/gate.rs | 217 ++++++++++- .../src/control/route/public_support.rs | 2 + .../src/control/tests/public_support.rs | 5 + .../src/dispatch/pool_scheduler.rs | 2 + .../handlers/admin/request/system/import.rs | 1 + .../src/handlers/admin/routing/mod.rs | 3 +- .../admin/users/api_keys/responses/update.rs | 1 + .../src/handlers/public/support/user_me.rs | 3 + .../public/support/user_me_api_key_routing.rs | 351 ++++++++++++++++++ .../public/support/user_me_api_keys.rs | 200 ++++++++-- .../handlers/public/support/user_me_routes.rs | 17 +- .../public/support/user_me_routing_groups.rs | 58 +++ .../handlers/public/support/user_me_usage.rs | 65 ++++ apps/aether-gateway/src/routing/resolver.rs | 124 ++++++- apps/aether-gateway/src/routing/selection.rs | 275 +++++++++++++- apps/aether-gateway/src/routing/trace.rs | 2 + .../scheduler/candidate/tests/selection.rs | 13 +- .../src/tests/frontdoor/internal.rs | 7 + .../src/tests/frontdoor/public_support.rs | 3 + .../public_support/api_key_routing.rs | 192 ++++++++++ .../public_support/routing_groups.rs | 135 +++++++ .../tests/architecture/sql_and_data.rs | 2 +- .../aether-admin/src/observability/usage.rs | 107 +++++- crates/aether-ai/formats/src/codex_profile.rs | 5 +- ...00000_separate_customer_billing_amount.sql | 154 ++++++++ .../aether-data/adapters/postgres/src/auth.rs | 157 +++++++- .../adapters/postgres/src/settlement.rs | 79 ++++ .../postgres/src/usage/analytics_tests.rs | 121 ++++++ .../adapters/postgres/src/usage/dashboard.rs | 16 +- .../src/usage/dashboard_history_tests.rs | 23 ++ .../adapters/postgres/src/usage/mod.rs | 40 +- .../src/usage/queries/dashboard_history.sql | 8 +- .../list_recent_usage_audits_prefix.sql | 22 +- .../queries/list_usage_audits_prefix.sql | 22 +- .../contracts/src/repository/auth.rs | 39 ++ .../src/repository/settlement/types.rs | 60 ++- .../repository/usage/billing_multiplier.rs | 169 +++++++++ .../src/repository/usage/metadata_policy.rs | 213 ++++++++++- .../contracts/src/repository/usage/mod.rs | 8 +- .../contracts/src/repository/usage/types.rs | 79 ++++ .../contracts/src/repository/wallet/types.rs | 1 + .../postgres/001_types_and_tables.sql | 1 + .../postgres/190_overview_analytics.sql | 75 +++- .../generated/postgres/baseline/007_stats.sql | 1 + .../runtime/schema/logical/007_stats.toml | 6 + .../src/backend/stats/postgres_daily/mod.rs | 6 + .../src/backend/stats/postgres_daily/sql.rs | 12 + .../runtime/src/lifecycle/migrate/tests.rs | 2 + .../migrate/tests/customer_billing_upgrade.rs | 114 ++++++ .../migrate/tests/legacy_overview_upgrade.rs | 1 + .../migrate/tests/overview_fact_metadata.rs | 2 +- .../runtime/src/repository/auth/memory.rs | 147 +++++++- .../runtime/src/repository/auth/mod.rs | 4 +- .../src/repository/settlement/memory.rs | 120 ++++++ .../runtime/src/repository/usage/memory.rs | 26 +- .../src/repository/usage/memory/analytics.rs | 12 +- .../usage/memory/dashboard_summary.rs | 10 +- .../src/repository/usage/memory/tests.rs | 109 ++++++ crates/aether-routing-core/src/model.rs | 115 +++++- crates/aether-routing-core/src/policy.rs | 173 ++++++++- crates/aether-routing-core/src/ranking.rs | 3 +- crates/aether-routing-core/src/trace.rs | 6 +- crates/aether-routing-core/src/validation.rs | 5 + .../bin/usage_settlement_hotspot_baseline.rs | 1 + .../runtime/src/request_metadata.rs | 68 +++- crates/aether-usage/runtime/src/runtime.rs | 23 +- crates/aether-usage/runtime/src/settlement.rs | 127 ++++++- .../runtime/src/settlement_reuse_tests.rs | 272 +++++++++++++- crates/aether-usage/runtime/src/worker.rs | 118 +++--- frontend/src/api/dashboard.ts | 4 + frontend/src/api/me.ts | 32 +- frontend/src/api/usage.ts | 12 + frontend/src/api/usageRecords.ts | 6 + .../components/ui/popover/PopoverContent.vue | 5 + .../components/ProviderSchedulingView.vue | 191 +++++----- .../ProviderSchedulingView.navigation.spec.ts | 263 ++++++++++++- .../RoutingFailoverPolicyEditor.spec.ts | 7 +- .../RoutingPriorityPolicyEditor.spec.ts | 16 +- .../RoutingSchedulingPolicyEditor.spec.ts | 144 +++++-- .../routing/__tests__/routingPolicy.spec.ts | 54 ++- .../__tests__/schedulingPolicies.spec.ts | 58 +++ .../RoutingModelSelectionPopover.vue | 103 +++++ .../components/RoutingModelSelector.vue | 14 +- .../RoutingPriorityPolicyEditor.vue | 3 +- .../RoutingSchedulingPolicyEditor.vue | 89 +++-- .../features/routing/utils/routingPolicy.ts | 27 ++ .../features/routing/utils/routingTrace.ts | 2 + .../routing/utils/schedulingPolicies.ts | 11 +- .../usage/components/RequestDetailDrawer.vue | 8 + .../usage/components/UsageCostDisplay.vue | 59 +++ .../usage/components/UsageProviderDisplay.vue | 24 ++ .../usage/components/UsageRecordsTable.vue | 113 ++---- .../RequestDetailDrawer.pricing.spec.ts | 42 +++ .../__tests__/UsageRecordsTable.spec.ts | 94 +++++ .../__tests__/useUsageData.spec.ts | 42 +++ .../usage/composables/useUsageData.ts | 5 +- .../utils/__tests__/usageBilling.spec.ts | 98 +++++ .../src/features/usage/utils/usageBilling.ts | 68 ++++ frontend/src/mocks/handler.ts | 83 ++++- .../src/views/admin/ProviderManagement.vue | 29 +- .../__tests__/ProviderManagement.spec.ts | 116 +++++- .../ProviderSchedulingView.failover.spec.ts | 9 +- frontend/src/views/shared/Usage.vue | 15 +- frontend/src/views/user/MyApiKeys.vue | 127 ++++++- .../user/__tests__/MyApiKeys.ccswitch.spec.ts | 104 ++++++ 110 files changed, 6524 insertions(+), 559 deletions(-) create mode 100644 apps/aether-gateway/src/handlers/public/support/user_me_api_key_routing.rs create mode 100644 apps/aether-gateway/src/handlers/public/support/user_me_routing_groups.rs create mode 100644 apps/aether-gateway/src/tests/frontdoor/public_support/api_key_routing.rs create mode 100644 apps/aether-gateway/src/tests/frontdoor/public_support/routing_groups.rs create mode 100644 crates/aether-data/adapters/postgres/migrations/20261007000000_separate_customer_billing_amount.sql create mode 100644 crates/aether-data/contracts/src/repository/usage/billing_multiplier.rs create mode 100644 crates/aether-data/runtime/src/lifecycle/migrate/tests/customer_billing_upgrade.rs create mode 100644 frontend/src/features/routing/components/RoutingModelSelectionPopover.vue create mode 100644 frontend/src/features/usage/components/UsageCostDisplay.vue create mode 100644 frontend/src/features/usage/components/UsageProviderDisplay.vue create mode 100644 frontend/src/features/usage/utils/__tests__/usageBilling.spec.ts create mode 100644 frontend/src/features/usage/utils/usageBilling.ts diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs index 0369f5118..c1bed4ab2 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs @@ -2579,6 +2579,8 @@ mod tests { let fixed_order_app = AppState::new().expect("state should build"); let fixed_order_policy = ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("routing-group-fixed-order".to_string()), group_version: Some(1), selection_source: "test".to_string(), diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs index a3840edc8..35865786a 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs @@ -362,6 +362,8 @@ mod tests { candidate.key_internal_priority = 3; candidate.key_global_priority_for_format = Some(2); let policy = aether_routing_core::ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("group-1".to_string()), group_version: Some(1), selection_source: "system_default".to_string(), @@ -399,6 +401,8 @@ mod tests { .expect("state should build") .with_data_state_for_tests(data_state); let policy = aether_routing_core::ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("group-1".to_string()), group_version: Some(1), selection_source: "system_default".to_string(), @@ -434,6 +438,8 @@ mod tests { candidate.key_internal_priority = 3; candidate.key_global_priority_for_format = Some(2); let policy = aether_routing_core::ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("group-1".to_string()), group_version: Some(1), selection_source: "system_default".to_string(), diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs index 50479935a..c3ae22dfc 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs @@ -1970,6 +1970,113 @@ mod tests { .is_none()); } + #[tokio::test] + async fn model_provider_enablement_filters_candidate_pages_without_affecting_other_models() { + let mut rows = Vec::new(); + for model in ["model-a", "model-b", "model-c"] { + for (provider, priority) in [ + ("provider-legacy-disabled", 0), + ("provider-model-disabled", 1), + ("provider-other", 2), + ("provider-inactive", 3), + ] { + let mut row = standard_candidate_row(provider, "openai:chat", priority); + row.global_model_id = format!("global-{model}"); + row.global_model_name = model.into(); + row.model_provider_model_name = model.into(); + row.model_id = format!("{provider}-{model}"); + row.provider_is_active = provider != "provider-inactive"; + rows.push(row); + } + } + let repository: Arc = + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)); + let app = AppState::new().unwrap().with_data_state_for_tests( + GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository), + ); + let auth = unrestricted_auth_snapshot(); + let directives = crate::system_features::ModelDirectivePolicySnapshot::load(&app).await; + let config = serde_json::from_value(serde_json::json!({ + "disabled_providers": ["provider-legacy-disabled"], + "model_policies": [ + { "model": "model-a", "provider_enabled_overrides": { + "provider-model-disabled": false, "provider-inactive": true + } }, + { "model": "model-b", "provider_enabled_overrides": { + "provider-legacy-disabled": true, "provider-inactive": true + } } + ], + "rules": [{ "id": "legacy-allowlist", "actions": [{ + "type": "restrict_providers", "provider_ids": [ + "provider-legacy-disabled", "provider-model-disabled", "provider-other", "provider-inactive" + ] + }] }] + })).unwrap(); + // Revisit A after B to exercise candidate caches shared by the app. + for (model, expected) in [ + ("model-a", vec!["provider-other"]), + ( + "model-b", + vec![ + "provider-legacy-disabled", + "provider-model-disabled", + "provider-other", + ], + ), + ("model-c", vec!["provider-model-disabled", "provider-other"]), + ("model-a", vec!["provider-other"]), + ] { + let policy = aether_routing_core::resolve_routing_policy( + &config, + aether_routing_core::RoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + selection_source: "test", + requested_model: model, + resolved_model: model, + api_format: "openai:chat", + user_id: None, + api_key_id: None, + headers: &serde_json::json!({}), + body: &serde_json::json!({}), + phase: aether_routing_core::RoutingRulePhase::ClientRequest, + }, + ) + .unwrap(); + let mut cursor = LocalCandidatePreselectionPageCursor::new( + PlannerAppState::new(&app), + &directives, + "openai:chat", + model, + None, + false, + None, + &auth, + Some(&policy), + None, + None, + true, + LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, + true, + None, + ) + .await; + let mut providers = Vec::new(); + while let Some(page) = cursor.next_page().await.unwrap() { + providers.extend( + page.candidates + .into_iter() + .map(|candidate| candidate.provider_id), + ); + } + providers.sort(); + assert_eq!( + providers, expected, + "provider enablement must remain isolated for {model}" + ); + } + } + #[tokio::test] async fn routing_policy_collects_candidate_pages_before_final_ranking() { let rows = (0..300) @@ -1992,6 +2099,8 @@ mod tests { let model_directive_policy = crate::system_features::ModelDirectivePolicySnapshot::load(&app).await; let routing_policy = ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("routing-group-1".to_string()), group_version: Some(1), selection_source: "test".to_string(), @@ -2057,6 +2166,8 @@ mod tests { let model_directive_policy = crate::system_features::ModelDirectivePolicySnapshot::load(&app).await; let routing_policy = ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("routing-group-fallback".to_string()), group_version: Some(1), selection_source: "test".to_string(), @@ -2899,6 +3010,8 @@ mod tests { let model_directive_policy = crate::system_features::ModelDirectivePolicySnapshot::load(&app).await; let routing_policy = ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("routing-group-codex-first".to_string()), group_version: Some(1), selection_source: "test".to_string(), diff --git a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs index e29a98f08..c24aa1888 100644 --- a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs +++ b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs @@ -555,24 +555,42 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( input.provider_outbound_context = Some(crate::ai_serving::codex_context::resolve_codex_fingerprint_context(parts, body_json)); let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER); + let preferred_group = if explicit_group.is_none() && !input.auth_context.api_key_is_standalone { + state + .read_auth_api_key_feature_settings( + &input.auth_context.user_id, + &input.auth_context.api_key_id, + false, + ) + .await? + .as_ref() + .and_then(|settings| settings.get("routing_group_id")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_owned) + } else { + None + }; let selected_group = match state.routing_group_read_repository() { Some(repository) => { // Explicit non-default groups are authorized against principal // bindings, so both selection and its cache key must retain the // caller context. Only the implicit no-binding system-default // path is global and can skip the membership lookup. - let principal_context_required = if explicit_group.is_some() { - true - } else { - repository - .has_any_routing_group_binding() - .await - .map_err(|error| { - routing_selection_error(GatewayRoutingSelectionError::Repository( - error.to_string(), - )) - })? - }; + let principal_context_required = + if explicit_group.is_some() || preferred_group.is_some() { + true + } else { + repository + .has_any_routing_group_binding() + .await + .map_err(|error| { + routing_selection_error(GatewayRoutingSelectionError::Repository( + error.to_string(), + )) + })? + }; let user_group_ids = if principal_context_required { let user_groups_lookup_started_at = std::time::Instant::now(); let user_groups = state @@ -595,6 +613,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( principal_context_required.then(|| input.auth_context.api_key_id.clone()); let selection_cache_key = routing_group_selection_cache_key( explicit_group.as_deref(), + preferred_group.as_deref(), selection_user_id.as_deref(), selection_api_key_id.as_deref(), &user_group_ids, @@ -612,6 +631,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( repository.as_ref(), GatewayRoutingSelectionInput { explicit_group: explicit_group.as_deref(), + preferred_group: preferred_group.as_deref(), user_id: selection_user_id.as_deref(), api_key_id: selection_api_key_id.as_deref(), user_group_ids: &user_group_ids, @@ -628,6 +648,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( || { let repository = repository.clone(); let explicit_group = explicit_group.clone(); + let preferred_group = preferred_group.clone(); let user_id = selection_user_id.clone(); let api_key_id = selection_api_key_id.clone(); let user_group_ids = user_group_ids.clone(); @@ -637,6 +658,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( repository.as_ref(), GatewayRoutingSelectionInput { explicit_group: explicit_group.as_deref(), + preferred_group: preferred_group.as_deref(), user_id: user_id.as_deref(), api_key_id: api_key_id.as_deref(), user_group_ids: &user_group_ids, @@ -662,6 +684,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( selection.group.map(|group| { ( Some(group.id), + group.name, Some(group.version), group.config_json, selection.source, @@ -669,13 +692,14 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( }) } None => { - if explicit_group + if let Some(requested_group) = explicit_group + .or(preferred_group) .as_deref() .map(str::trim) - .is_some_and(|value| !value.is_empty()) + .filter(|value| !value.is_empty()) { return Err(routing_selection_error( - GatewayRoutingSelectionError::NotFound(explicit_group.unwrap_or_default()), + GatewayRoutingSelectionError::NotFound(requested_group.to_string()), )); } return Err(routing_selection_error( @@ -684,7 +708,8 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( } }; - let Some((group_id, group_version, group_config_json, selection_source)) = selected_group + let Some((group_id, group_name, group_version, group_config_json, selection_source)) = + selected_group else { return Err(routing_selection_error( GatewayRoutingSelectionError::NoDefault, @@ -701,6 +726,12 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( &group_config_json, selection_source.as_str(), )? { + if let Some(policy) = input.routing_policy.as_mut() { + policy.group_name = Some(group_name.clone()); + } + if let Some(trace) = input.routing_trace_seed.as_mut() { + trace.group_name = Some(group_name); + } return Ok(()); } @@ -786,6 +817,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( final_policy_resolve_started_at.elapsed().as_millis() as u64, ); final_policy.mutation_plan = policy.mutation_plan.clone(); + final_policy.group_name = Some(group_name); input.routing_trace_seed = Some(build_routing_trace_seed(&final_policy, client_api_format)); input.routing_policy = Some(final_policy); input.routing_context = Some(LocalRoutingRequestContext { @@ -966,6 +998,7 @@ fn routing_header_value_str(headers: &http::HeaderMap, key: &str) -> Option, + preferred_group: Option<&str>, user_id: Option<&str>, api_key_id: Option<&str>, user_group_ids: &[String], @@ -976,8 +1009,9 @@ fn routing_group_selection_cache_key( .collect::>() .join(","); format!( - "v1|explicit={}|user={}|api_key={}|groups={}", + "v2|explicit={}|preferred={}|user={}|api_key={}|groups={}", escape_cache_key_part(explicit_group.unwrap_or_default()), + escape_cache_key_part(preferred_group.unwrap_or_default()), escape_cache_key_part(user_id.unwrap_or_default()), escape_cache_key_part(api_key_id.unwrap_or_default()), groups @@ -1175,10 +1209,13 @@ mod tests { use std::sync::Arc; use super::*; + use aether_data::repository::auth::{ + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, + }; use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; use aether_data_contracts::repository::routing_profiles::{ CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, RoutingGroupBindingSubject, - RoutingGroupWriteRepository, + RoutingGroupWriteRepository, UpdateRoutingGroupRecord, }; use aether_provider_transport::snapshot::{ GatewayProviderTransportEndpoint, GatewayProviderTransportKey, @@ -1189,12 +1226,14 @@ mod tests { fn explicit_routing_selection_cache_key_is_principal_specific() { let first = routing_group_selection_cache_key( Some("private"), + None, Some("user-1"), Some("key-1"), &["team-1".to_string()], ); let second = routing_group_selection_cache_key( Some("private"), + None, Some("user-2"), Some("key-2"), &["team-2".to_string()], @@ -1311,6 +1350,160 @@ mod tests { assert!(matches!(error, GatewayError::Internal(message) if message.contains("identity"))); } + #[tokio::test] + async fn api_key_routing_selection_applies_at_planner_and_invalidates_after_changes() { + let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed( + ["api-key-1", "api-key-2"].map(|key_id| { + ( + None, + StoredAuthApiKeySnapshot::new( + "user-1".into(), + "alice".into(), + None, + "user".into(), + "local".into(), + true, + false, + None, + None, + None, + key_id.into(), + Some(key_id.into()), + true, + false, + false, + None, + None, + None, + None, + None, + None, + ) + .unwrap(), + ) + }), + )); + let groups = Arc::new(InMemoryRoutingGroupRepository::default()); + for (id, visible, is_default, multiplier) in [ + ("default", false, true, 1.0), + ("discount", true, false, 0.5), + ("premium", true, false, 2.0), + ] { + groups.create_routing_group(CreateRoutingGroupRecord { + id: id.into(), name: format!("{id}-name"), description: None, + enabled: true, is_system_default: is_default, sort_order: 0, + config_json: json!({ "user_visible": visible, "billing_multiplier": multiplier }), + version: 1, created_at: 1, updated_at: 1, published_at: None, + }).await.unwrap(); + } + let state = AppState::new().unwrap().with_data_state_for_tests( + crate::data::GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository) + .with_routing_group_repository_for_tests(groups.clone()), + ); + for (key_id, group_id) in [("api-key-1", "discount"), ("api-key-2", "premium")] { + assert!(state + .set_user_api_key_feature_settings( + "user-1", + key_id, + Some(json!({ "routing_group_id": group_id })) + ) + .await + .unwrap() + .is_some()); + } + let (parts, _) = http::Request::builder().body(()).unwrap().into_parts(); + let (header_parts, _) = http::Request::builder() + .header(ROUTING_GROUP_HEADER, "premium") + .body(()) + .unwrap() + .into_parts(); + + async fn attach( + state: &AppState, + parts: &http::request::Parts, + key_id: &str, + ) -> Result { + let mut input = sample_decision_input(); + input.auth_context.api_key_id = key_id.into(); + input.auth_snapshot.api_key_id = key_id.into(); + attach_routing_policy_to_local_requested_model_input( + state, + parts, + &mut input, + &json!({ "model": "gpt-5" }), + "openai:chat", + ) + .await?; + Ok(input) + } + + // Revisit the first key after the second to exercise both cached choices. + for (key_id, group_id, multiplier) in [ + ("api-key-1", "discount", 0.5), + ("api-key-2", "premium", 2.0), + ("api-key-1", "discount", 0.5), + ] { + let input = attach(&state, &parts, key_id).await.unwrap(); + let policy = input.routing_policy.as_ref().unwrap(); + assert_eq!(policy.group_id.as_deref(), Some(group_id)); + assert_eq!(policy.selection_source, "api_key_selection"); + assert_eq!(policy.billing_multiplier, multiplier); + assert_eq!( + input + .routing_trace_seed + .as_ref() + .unwrap() + .billing_multiplier, + Some(multiplier) + ); + } + let header = attach(&state, &header_parts, "api-key-1").await.unwrap(); + let policy = header.routing_policy.unwrap(); + assert_eq!(policy.group_id.as_deref(), Some("premium")); + assert_eq!(policy.selection_source, "explicit_header"); + + groups + .update_routing_group( + "discount", + UpdateRoutingGroupRecord { + config_json: Some(json!({ "user_visible": false, "billing_multiplier": 0.5 })), + ..Default::default() + }, + ) + .await + .unwrap(); + state.invalidate_provider_routing_caches(); + assert!(matches!( + attach(&state, &parts, "api-key-1").await, + Err(GatewayError::Client { + status: StatusCode::FORBIDDEN, + .. + }) + )); + let header = attach(&state, &header_parts, "api-key-1").await.unwrap(); + assert_eq!( + header.routing_policy.unwrap().group_id.as_deref(), + Some("premium") + ); + + assert!(state + .set_user_api_key_feature_settings("user-1", "api-key-1", None) + .await + .unwrap() + .is_some()); + let cleared = attach(&state, &parts, "api-key-1").await.unwrap(); + let policy = cleared.routing_policy.unwrap(); + assert_eq!(policy.group_id.as_deref(), Some("default")); + assert_eq!(policy.selection_source, "system_default"); + assert_eq!(policy.billing_multiplier, 1.0); + // Clearing one key's preference must not disturb the other key's selection. + let other = attach(&state, &parts, "api-key-2").await.unwrap(); + assert_eq!( + other.routing_policy.unwrap().group_id.as_deref(), + Some("premium") + ); + } + #[tokio::test] async fn explicit_routing_attachment_authorizes_and_caches_per_principal() { let repository = Arc::new(InMemoryRoutingGroupRepository::default()); @@ -1322,7 +1515,7 @@ mod tests { enabled: true, is_system_default: false, sort_order: 0, - config_json: json!({}), + config_json: json!({"billing_multiplier": 0.5}), version: 1, created_at: 1, updated_at: 1, @@ -1368,6 +1561,11 @@ mod tests { .as_ref() .expect("explicit selection should attach routing policy"); assert_eq!(policy.group_id.as_deref(), Some("private-group")); + assert_eq!(policy.group_name.as_deref(), Some("private")); + assert_eq!(policy.billing_multiplier, 0.5); + let trace = allowed.routing_trace_seed.as_ref().unwrap(); + assert_eq!(trace.group_name.as_deref(), Some("private")); + assert_eq!(trace.billing_multiplier, Some(0.5)); assert_eq!(policy.selection_source, "explicit_header"); let mut denied = sample_decision_input(); diff --git a/apps/aether-gateway/src/ai_serving/planner/report_context.rs b/apps/aether-gateway/src/ai_serving/planner/report_context.rs index bd621e27b..cf28bbbd5 100644 --- a/apps/aether-gateway/src/ai_serving/planner/report_context.rs +++ b/apps/aether-gateway/src/ai_serving/planner/report_context.rs @@ -6,6 +6,11 @@ use aether_ai_serving::{ provider_stream_event_api_format_for_provider_type as ai_provider_stream_event_api_format_for_provider_type, AiExecutionReportContextParts, AiRequestOrigin, STICKY_KEY_ATTEMPTS_REPORT_FIELD, }; +use aether_data_contracts::repository::usage::{ + BillingMultiplierSnapshot, BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, +}; use aether_routing_core::ResolvedRoutingPolicy; use aether_runtime_state::RuntimeLockLease; use aether_scheduler_core::{ClientSessionAffinity, SchedulerRankingOutcome}; @@ -87,6 +92,46 @@ pub(crate) fn build_local_execution_report_context( parts.original_request_body_base64, ); let mut extra_fields = parts.extra_fields; + // Always overwrite caller-supplied extras with the planner's immutable policy snapshot. + let billing_multiplier = parts + .routing_policy + .map(|policy| policy.billing_multiplier) + .filter(|value| value.is_finite() && *value >= 0.0) + .unwrap_or(1.0); + extra_fields.insert( + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY.to_string(), + Value::from(billing_multiplier), + ); + extra_fields.insert( + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY.to_string(), + serde_json::to_value( + BillingMultiplierSnapshot::from_factors(BTreeMap::from([( + "routing_group".to_string(), + billing_multiplier, + )])) + .expect("validated routing multiplier must produce a billing snapshot"), + ) + .expect("validated billing snapshot must serialize"), + ); + for (field, value) in [ + ( + ROUTING_GROUP_ID_METADATA_KEY, + parts + .routing_policy + .and_then(|policy| policy.group_id.as_deref()), + ), + ( + ROUTING_GROUP_NAME_METADATA_KEY, + parts + .routing_policy + .and_then(|policy| policy.group_name.as_deref()), + ), + ] { + extra_fields.remove(field); + if let Some(value) = value { + extra_fields.insert(field.to_string(), Value::String(value.to_string())); + } + } if let Some(value) = parts .client_session_affinity .and_then(client_session_affinity_report_context_value) @@ -341,6 +386,27 @@ mod tests { Some("codex".to_string()), Some("account=account-1;session=session-1".to_string()), ); + let mut routing_policy = aether_routing_core::resolve_routing_policy( + &aether_routing_core::RoutingGroupConfig { + billing_multiplier: 0.25, + ..Default::default() + }, + aether_routing_core::RoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(7), + selection_source: "system_default", + requested_model: "gpt-5", + resolved_model: "gpt-5", + api_format: "openai:chat", + user_id: None, + api_key_id: None, + headers: &json!({}), + body: &json!({}), + phase: aether_routing_core::RoutingRulePhase::ClientRequest, + }, + ) + .expect("routing policy should resolve"); + routing_policy.group_name = Some("请求时的分组".to_string()); let report_context = build_local_execution_report_context(LocalExecutionReportContextParts { @@ -379,16 +445,35 @@ mod tests { original_request_body_json: Some(&json!({"model": "gpt-5"})), original_request_body_base64: None, client_session_affinity: Some(&client_session_affinity), - routing_policy: None, + routing_policy: Some(&routing_policy), scheduler_affinity_epoch: None, sticky_key_attempts: None, client_requested_stream: false, upstream_is_stream: false, has_envelope: false, needs_conversion: false, - extra_fields: Map::new(), + extra_fields: Map::from_iter([ + ( + "billing_multiplier_snapshot".to_string(), + json!({ + "version": 1, "factors": {"routing_group": 99.0}, "multiplier": 99.0 + }), + ), + ("routing_group_billing_multiplier".to_string(), json!(99)), + ("routing_group_id".to_string(), json!("forged-group")), + ("routing_group_name".to_string(), json!("forged-name")), + ]), }); + assert_eq!(report_context["routing_group_billing_multiplier"], 0.25); + assert_eq!( + report_context["billing_multiplier_snapshot"], + json!({ + "version": 1, "factors": {"routing_group": 0.25}, "multiplier": 0.25 + }) + ); + assert_eq!(report_context["routing_group_id"], "group-1"); + assert_eq!(report_context["routing_group_name"], "请求时的分组"); assert_eq!( report_context["client_ip"], Value::String("203.0.113.8".to_string()) diff --git a/apps/aether-gateway/src/control/auth/gate.rs b/apps/aether-gateway/src/control/auth/gate.rs index 467b29e74..6e79bf6e6 100644 --- a/apps/aether-gateway/src/control/auth/gate.rs +++ b/apps/aether-gateway/src/control/auth/gate.rs @@ -231,8 +231,32 @@ pub(crate) async fn estimate_execution_plan_cost_upper_bound_usd( report_context: Option<&serde_json::Value>, ) -> Result, GatewayError> { let started_at = std::time::Instant::now(); - let result = - estimate_execution_plan_cost_upper_bound_usd_inner(state, plan, report_context).await; + let result = async { + let multiplier_snapshot = + aether_data_contracts::repository::usage::billing_multiplier_snapshot(report_context) + .map_err(|error| GatewayError::Internal(error.to_string()))?; + let estimate = estimate_execution_plan_cost_upper_bound_usd_inner( + state, + plan, + report_context, + multiplier_snapshot.is_some(), + ) + .await?; + let Some(snapshot) = multiplier_snapshot else { + return Ok(estimate); + }; + // Cache the unmultiplied base estimate so different request snapshots + // cannot reuse one another's charge. Pricing validation still runs for + // a zero multiplier, even when the request has no finite token bound. + if snapshot.multiplier() == 0.0 { + return Ok(Some(0.0)); + } + estimate + .map(|cost| snapshot.cost(cost)) + .transpose() + .map_err(|error| GatewayError::Internal(error.to_string())) + } + .await; observe_gateway_stage_ms( "auth_capacity_cost_estimate", started_at.elapsed().as_millis() as u64, @@ -244,6 +268,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner( state: &AppState, plan: &aether_contracts::ExecutionPlan, report_context: Option<&serde_json::Value>, + use_base_cost: bool, ) -> Result, GatewayError> { let api_format = crate::ai_serving::normalize_api_format_alias(&plan.provider_api_format); let body_json = plan.body.json_body.as_ref(); @@ -311,7 +336,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner( if model_id.is_none() && global_model_name.is_none() { return Ok(None); } - let cache_key = execution_plan_cost_upper_bound_cache_key( + let mut cache_key = execution_plan_cost_upper_bound_cache_key( plan, model_id, global_model_name, @@ -321,6 +346,11 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner( requested_processing_tier.as_deref(), cache_ttl_minutes, ); + if use_base_cost { + // Legacy requests cache provider Key cost; new requests cache base cost. + // These values must never share a cache entry for the same provider Key. + cache_key.insert_str(0, "base\x1f"); + } let ttl = state.frontdoor_runtime_guards.auth_capacity_cache_ttl; if ttl.is_zero() { let _permit = state.acquire_auth_snapshot_load_gate().await?; @@ -335,6 +365,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner( max_output_tokens, requested_processing_tier.as_deref(), cache_ttl_minutes, + use_base_cost, ) .await; } @@ -353,6 +384,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner( max_output_tokens, requested_processing_tier.as_deref(), cache_ttl_minutes, + use_base_cost, ) .await }) @@ -371,6 +403,7 @@ async fn calculate_execution_plan_cost_upper_bound( max_output_tokens: Option, requested_processing_tier: Option<&str>, cache_ttl_minutes: Option, + use_base_cost: bool, ) -> Result, GatewayError> { let context = load_execution_plan_billing_context(state, plan, model_id, global_model_name).await?; @@ -383,11 +416,13 @@ async fn calculate_execution_plan_cost_upper_bound( estimate.requested_processing_tier = requested_processing_tier.map(ToOwned::to_owned); estimate.cache_ttl_minutes = cache_ttl_minutes; estimate.max_output_tokens = max_output_tokens; + let mut pricing = aether_billing::BillingModelPricingSnapshot::from(context); + if use_base_cost { + pricing.provider_billing_type = None; + pricing.provider_api_key_rate_multipliers = None; + } aether_billing::BillingService::new() - .estimate_authorization_cost_upper_bound( - &aether_billing::BillingModelPricingSnapshot::from(context), - &estimate, - ) + .estimate_authorization_cost_upper_bound(&pricing, &estimate) .map_err(|err| GatewayError::Internal(err.to_string())) } @@ -860,10 +895,10 @@ mod tests { use serde_json::json; use super::{ - available_balance_capacity_usd, execution_plan_balance_capacity_rejection, - execution_plan_cost_upper_bound_cache_key, max_output_tokens_from_request, - openai_request_input_is_self_contained, output_choice_count_upper_bound, - request_model_local_rejection, GatewayLocalAuthRejection, + available_balance_capacity_usd, estimate_execution_plan_cost_upper_bound_usd, + execution_plan_balance_capacity_rejection, execution_plan_cost_upper_bound_cache_key, + max_output_tokens_from_request, openai_request_input_is_self_contained, + output_choice_count_upper_bound, request_model_local_rejection, GatewayLocalAuthRejection, }; use crate::control::{GatewayControlAuthContext, GatewayControlDecision}; use crate::data::GatewayDataState; @@ -2098,6 +2133,166 @@ mod tests { assert_eq!(estimate, 6.5); } + #[tokio::test] + async fn charge_estimate_and_capacity_use_request_multiplier_without_key_cost_or_cache_leaks() { + let context = billing_context_with_pricing( + Some(json!({"tiers": [{ + "up_to": null, + "input_price_per_1m": 0.0, + "output_price_per_1m": 10.0 + }]})), + None, + Some(json!({"openai:chat": 2.0})), + None, + ); + let mut state = state_with_quota_and_wallet(quota_availability(15.0, false), context); + Arc::make_mut(&mut state.frontdoor_runtime_guards).auth_capacity_cache_ttl = + Duration::from_secs(60); + let decision = decision_with_allowed_models(vec!["gpt-5".to_string()]); + let plan = execution_plan( + json!({"model": "gpt-5", "messages": [], "max_tokens": 1_000_000}), + "openai:chat", + ); + let legacy = billing_report_context(); + let mut discounted = legacy.clone(); + discounted["billing_multiplier_snapshot"] = json!({ + "version": 1, + "factors": {"routing_group": 2.0, "promotion": 0.25}, + "multiplier": 0.5 + }); + let mut marked_up = legacy.clone(); + marked_up["billing_multiplier_snapshot"] = json!({ + "version": 1, + "factors": {"routing_group": 3.0}, + "multiplier": 3.0 + }); + let mut legacy_group_snapshot = legacy.clone(); + legacy_group_snapshot["routing_group_billing_multiplier"] = json!(1.0); + + // Reuse the same cache for legacy Key cost, independent request + // multipliers, and the old group-only snapshot representation. + for (report_context, expected) in [ + (&legacy, 20.0), + (&discounted, 5.0), + (&marked_up, 30.0), + (&legacy_group_snapshot, 10.0), + (&discounted, 5.0), + (&legacy, 20.0), + ] { + assert_eq!( + estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(report_context)) + .await + .expect("charge estimate should resolve"), + Some(expected) + ); + } + assert_eq!( + execution_plan_balance_capacity_rejection(&state, &decision, &plan, Some(&discounted)) + .await + .expect("discounted request capacity should resolve"), + None + ); + assert_eq!( + execution_plan_balance_capacity_rejection(&state, &decision, &plan, Some(&marked_up)) + .await + .expect("marked-up request capacity should resolve"), + Some(GatewayLocalAuthRejection::BalanceDenied { + remaining: Some(15.0) + }) + ); + } + + #[tokio::test] + async fn charge_estimate_uses_base_price_when_provider_is_free_tier() { + let context = billing_context_with_pricing( + Some(json!({"tiers": [{ + "up_to": null, + "input_price_per_1m": 0.0, + "output_price_per_1m": 10.0 + }]})), + None, + Some(json!({"openai:chat": 0.0})), + Some("free_tier"), + ); + let state = state_with_quota_and_wallet(quota_availability(15.0, false), context); + let plan = execution_plan( + json!({"model": "gpt-5", "messages": [], "max_tokens": 1_000_000}), + "openai:chat", + ); + let mut report_context = billing_report_context(); + assert_eq!( + estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context)) + .await + .expect("legacy free-tier estimate should resolve"), + Some(0.0) + ); + report_context["routing_group_billing_multiplier"] = json!(0.5); + assert_eq!( + estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context)) + .await + .expect("charge estimate should use the model base price"), + Some(5.0) + ); + } + + #[tokio::test] + async fn zero_charge_multiplier_bounds_unknown_cost_but_still_rejects_invalid_pricing() { + let context = billing_context_with_pricing( + Some(json!({"tiers": [{ + "up_to": null, + "input_price_per_1m": 0.0, + "output_price_per_1m": 10.0 + }]})), + None, + None, + None, + ); + let state = state_with_quota_and_wallet(quota_availability(0.0, false), context); + let plan = execution_plan(json!({"model": "gpt-5", "messages": []}), "openai:chat"); + let mut report_context = billing_report_context(); + assert_eq!( + estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context)) + .await + .expect("an unspecified output limit has no finite estimate"), + None + ); + report_context["billing_multiplier_snapshot"] = json!({ + "version": 1, + "factors": {"routing_group": 0.0}, + "multiplier": 0.0 + }); + assert_eq!( + estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context)) + .await + .expect("zero multiplier should bound the charge"), + Some(0.0) + ); + let invalid_context = billing_context_with_pricing( + Some(json!({ + "tiers": [{"up_to": null, "input_price_per_1m": 1.0}], + "processing_tiers": { + "priority": {"tiers": [{}], "price_multiplier": -1.0} + } + })), + None, + None, + None, + ); + let invalid_state = + state_with_quota_and_wallet(quota_availability(0.0, false), invalid_context); + let invalid_plan = execution_plan( + json!({"model": "gpt-5", "messages": [], "service_tier": "priority"}), + "openai:chat", + ); + assert!(estimate_execution_plan_cost_upper_bound_usd( + &invalid_state, + &invalid_plan, + Some(&report_context) + ) + .await + .is_err()); + } + #[test] fn daily_quota_estimate_treats_free_tier_as_zero_cost() { let context = billing_context_with_pricing( diff --git a/apps/aether-gateway/src/control/route/public_support.rs b/apps/aether-gateway/src/control/route/public_support.rs index 2ada334de..62767c522 100644 --- a/apps/aether-gateway/src/control/route/public_support.rs +++ b/apps/aether-gateway/src/control/route/public_support.rs @@ -637,6 +637,7 @@ pub(super) fn classify_public_support_route( | "/api/users/me/usage/interval-timeline" | "/api/users/me/usage/heatmap" | "/api/users/me/providers" + | "/api/users/me/routing-groups" | "/api/users/me/available-models" | "/api/users/me/client-config" | "/api/users/me/endpoint-status" @@ -654,6 +655,7 @@ pub(super) fn classify_public_support_route( "/api/users/me/usage/interval-timeline" => "usage_interval_timeline", "/api/users/me/usage/heatmap" => "usage_heatmap", "/api/users/me/providers" => "providers", + "/api/users/me/routing-groups" => "routing_groups", "/api/users/me/available-models" => "available_models", "/api/users/me/client-config" => "client_config", "/api/users/me/endpoint-status" => "endpoint_status", diff --git a/apps/aether-gateway/src/control/tests/public_support.rs b/apps/aether-gateway/src/control/tests/public_support.rs index b758c658a..1df7c7d52 100644 --- a/apps/aether-gateway/src/control/tests/public_support.rs +++ b/apps/aether-gateway/src/control/tests/public_support.rs @@ -461,6 +461,11 @@ fn classifies_users_me_routes_as_public_support_route() { "/api/users/me/available-models", "available_models", ), + ( + http::Method::GET, + "/api/users/me/routing-groups", + "routing_groups", + ), ( http::Method::GET, "/api/users/me/vscodex/devices", diff --git a/apps/aether-gateway/src/dispatch/pool_scheduler.rs b/apps/aether-gateway/src/dispatch/pool_scheduler.rs index dcb60a8ee..28bd8ae78 100644 --- a/apps/aether-gateway/src/dispatch/pool_scheduler.rs +++ b/apps/aether-gateway/src/dispatch/pool_scheduler.rs @@ -5524,6 +5524,8 @@ mod tests { key_ids: [&str; N], ) -> ResolvedRoutingPolicy { ResolvedRoutingPolicy { + billing_multiplier: 1.0, + group_name: None, group_id: Some("routing-group-1".to_string()), group_version: Some(1), selection_source: "test".to_string(), diff --git a/apps/aether-gateway/src/handlers/admin/request/system/import.rs b/apps/aether-gateway/src/handlers/admin/request/system/import.rs index a63e08ea8..b80092864 100644 --- a/apps/aether-gateway/src/handlers/admin/request/system/import.rs +++ b/apps/aether-gateway/src/handlers/admin/request/system/import.rs @@ -7733,6 +7733,7 @@ impl<'a> AdminAppState<'a> { feature_settings: key .contains_key("feature_settings") .then(|| feature_settings.clone()), + routing_group_selection: None, }, ) .await?; diff --git a/apps/aether-gateway/src/handlers/admin/routing/mod.rs b/apps/aether-gateway/src/handlers/admin/routing/mod.rs index 44a9e53e7..d00595790 100644 --- a/apps/aether-gateway/src/handlers/admin/routing/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/routing/mod.rs @@ -411,7 +411,7 @@ async fn dry_run_routing_group( let headers_json = payload.headers.unwrap_or_else(|| json!({})); let mut header_map = header_map_from_value(&headers_json)?; let mut body = payload.body.unwrap_or_else(|| json!({})); - let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput { + let mut policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput { group_id: Some(group.id.as_str()), group_version: Some(group.version), group_config_json: &group.config_json, @@ -425,6 +425,7 @@ async fn dry_run_routing_group( body: &body, phase: payload.phase.unwrap_or(RoutingRulePhase::ClientRequest), })?; + policy.group_name = Some(group.name.clone()); let patch_summary = patch_summary(&policy.mutation_plan); apply_routing_mutation_plan(&mut body, &mut header_map, &policy.mutation_plan)?; let mut trace = build_routing_trace_seed(&policy, api_format); diff --git a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs index 08668a6c8..8d47344c0 100644 --- a/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs +++ b/apps/aether-gateway/src/handlers/admin/users/api_keys/responses/update.rs @@ -125,6 +125,7 @@ pub(crate) async fn build_admin_update_user_api_key_response( concurrent_limit_present, ip_rules, feature_settings, + routing_group_selection: None, }) .await? else { diff --git a/apps/aether-gateway/src/handlers/public/support/user_me.rs b/apps/aether-gateway/src/handlers/public/support/user_me.rs index 54d196599..3255a9a67 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me.rs @@ -32,6 +32,9 @@ use user_me_usage::*; #[path = "user_me_catalog.rs"] mod user_me_catalog; use user_me_catalog::*; +#[path = "user_me_routing_groups.rs"] +mod user_me_routing_groups; +use user_me_routing_groups::*; #[path = "user_me_preferences.rs"] mod user_me_preferences; use user_me_preferences::*; diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_api_key_routing.rs b/apps/aether-gateway/src/handlers/public/support/user_me_api_key_routing.rs new file mode 100644 index 000000000..d6ebee026 --- /dev/null +++ b/apps/aether-gateway/src/handlers/public/support/user_me_api_key_routing.rs @@ -0,0 +1,351 @@ +use std::collections::BTreeMap; + +use aether_data_contracts::repository::routing_profiles::RoutingGroupLookupKey; +use axum::{body::Body, http::StatusCode, response::Response}; +use serde::Deserialize; +use serde_json::{Map, Value}; + +use super::{build_auth_error_response, normalize_feature_settings, AppState}; +use crate::routing::selection::routing_group_is_user_visible; + +const ROUTING_GROUP_ID: &str = "routing_group_id"; + +pub(super) fn deserialize_routing_group_patch<'de, D>( + deserializer: D, +) -> Result>, D::Error> +where + D: serde::Deserializer<'de>, +{ + Option::::deserialize(deserializer).map(Some) +} + +pub(super) fn api_key_routing_group_id(settings: Option<&Value>) -> Option<&str> { + settings + .and_then(|value| value.get(ROUTING_GROUP_ID)) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +pub(super) async fn validate_routing_group_patch( + state: &AppState, + current: Option<&Value>, + requested: Option>, +) -> Result>, Response> { + let Some(Some(requested)) = requested else { + return Ok(requested); + }; + let id = requested.trim(); + if id.is_empty() || id.len() > 128 { + return Err(build_auth_error_response( + StatusCode::BAD_REQUEST, + "routing_group_id 必须是有效的策略分组 ID;跟随默认请传 null", + false, + )); + } + // A group can become hidden or disabled after selection. An unrelated edit + // (including a form resubmitting its unchanged selection) must remain valid. + if api_key_routing_group_id(current) == Some(id) { + return Ok(Some(Some(id.to_string()))); + } + if !state.has_routing_group_data_reader() { + return Err(build_auth_error_response( + StatusCode::SERVICE_UNAVAILABLE, + "策略分组目录暂不可用", + false, + )); + } + let group = state + .find_routing_group(RoutingGroupLookupKey::Id(id)) + .await + .map_err(|error| { + build_auth_error_response( + StatusCode::INTERNAL_SERVER_ERROR, + format!("user API key routing group lookup failed: {error:?}"), + false, + ) + })?; + if !group + .as_ref() + .is_some_and(|group| group.enabled && routing_group_is_user_visible(group)) + { + return Err(build_auth_error_response( + StatusCode::BAD_REQUEST, + "所选策略分组不存在或当前不可选", + false, + )); + } + Ok(Some(Some(id.to_string()))) +} + +/// Compose the initial create record. Existing-key updates must use the atomic +/// repository routing selection patch instead of merging a pre-read snapshot. +pub(super) fn merge_api_key_feature_settings( + current: Option<&Value>, + incoming: Option>, + routing_group_patch: Option>, +) -> Result>, String> { + if incoming.is_none() && routing_group_patch.is_none() { + return Ok(None); + } + let mut settings = normalize_feature_settings(incoming.unwrap_or_else(|| current.cloned()))? + .and_then(|value| value.as_object().cloned()) + .unwrap_or_default(); + settings.remove(ROUTING_GROUP_ID); + settings.remove("routing_group_name"); + let group_id = routing_group_patch + .unwrap_or_else(|| api_key_routing_group_id(current).map(ToOwned::to_owned)); + if let Some(group_id) = group_id { + settings.insert(ROUTING_GROUP_ID.to_string(), Value::String(group_id)); + } + Ok(Some( + (!settings.is_empty()).then_some(Value::Object(settings)), + )) +} + +pub(super) fn normalize_api_key_feature_settings_patch( + value: Option>, +) -> Result>, String> { + let Some(value) = value else { + return Ok(None); + }; + let Some(Value::Object(mut settings)) = normalize_feature_settings(value)? else { + return Ok(Some(None)); + }; + settings.remove(ROUTING_GROUP_ID); + settings.remove("routing_group_name"); + Ok(Some( + (!settings.is_empty()).then_some(Value::Object(settings)), + )) +} + +pub(super) async fn routing_group_names( + state: &AppState, + needed: bool, +) -> BTreeMap { + if !needed || !state.has_routing_group_data_reader() { + return BTreeMap::new(); + } + match state.list_routing_groups().await { + Ok(groups) => groups + .into_iter() + .map(|group| (group.id, group.name)) + .collect(), + Err(error) => { + tracing::warn!(?error, "API key routing group names unavailable"); + BTreeMap::new() + } + } +} + +pub(super) fn routing_group_payload_fields( + settings: Option<&Value>, + names: &BTreeMap, +) -> Map { + let id = api_key_routing_group_id(settings); + Map::from_iter([ + (ROUTING_GROUP_ID.to_string(), serde_json::json!(id)), + ( + "routing_group_name".to_string(), + serde_json::json!(id.and_then(|id| names.get(id))), + ), + ]) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; + use aether_data_contracts::repository::routing_profiles::StoredRoutingGroup; + use serde_json::json; + + use super::*; + + fn group(id: &str, enabled: bool, visible: Value) -> StoredRoutingGroup { + StoredRoutingGroup { + id: id.to_string(), + name: format!("{id} current name"), + description: None, + enabled, + is_system_default: false, + sort_order: 0, + config_json: json!({"user_visible": visible}), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + } + } + + fn state() -> AppState { + let repository = Arc::new(InMemoryRoutingGroupRepository::seed( + [ + group("public", true, json!(true)), + group("hidden", true, json!(false)), + group("disabled", false, json!(true)), + group("legacy", true, Value::Null), + group("malformed", true, json!("true")), + ], + [], + [], + )); + AppState::new().unwrap().with_data_state_for_tests( + crate::data::GatewayDataState::disabled() + .with_routing_group_repository_for_tests(repository), + ) + } + + #[tokio::test] + async fn only_enabled_user_visible_groups_can_be_newly_selected() { + let state = state(); + for id in [ + "", + " ", + "missing", + "hidden", + "disabled", + "legacy", + "malformed", + ] { + assert_eq!( + validate_routing_group_patch(&state, None, Some(Some(id.to_string()))) + .await + .unwrap_err() + .status(), + StatusCode::BAD_REQUEST, + "id={id}" + ); + } + assert_eq!( + validate_routing_group_patch(&state, None, Some(Some(" public ".to_string()))) + .await + .unwrap(), + Some(Some("public".to_string())) + ); + let unavailable = AppState::new() + .unwrap() + .with_data_state_for_tests(crate::data::GatewayDataState::disabled()); + assert_eq!( + validate_routing_group_patch(&unavailable, None, Some(Some("public".into()))) + .await + .unwrap_err() + .status(), + StatusCode::SERVICE_UNAVAILABLE + ); + } + + #[tokio::test] + async fn unchanged_or_cleared_choices_do_not_require_current_visibility_or_catalog_access() { + let state = AppState::new() + .unwrap() + .with_data_state_for_tests(crate::data::GatewayDataState::disabled()); + let current = json!({"routing_group_id": "hidden"}); + for patch in [None, Some(None), Some(Some("hidden".to_string()))] { + assert_eq!( + validate_routing_group_patch(&state, Some(¤t), patch.clone()) + .await + .unwrap(), + patch + ); + } + } + + #[test] + fn feature_patch_does_not_carry_a_pre_read_routing_selection() { + assert_eq!( + normalize_api_key_feature_settings_patch(None).unwrap(), + None + ); + assert_eq!( + normalize_api_key_feature_settings_patch(Some(None)).unwrap(), + Some(None) + ); + let patch = normalize_api_key_feature_settings_patch(Some(Some(json!({ + "routing_group_id": "stale-or-forged", + "routing_group_name": "stale name", + "chat_pii_redaction": {"enabled": false}, + })))) + .unwrap() + .flatten() + .unwrap(); + assert!(patch.get("routing_group_id").is_none()); + assert!(patch.get("routing_group_name").is_none()); + assert_eq!(patch["chat_pii_redaction"]["enabled"], false); + } + + #[test] + fn feature_updates_cannot_inject_replace_or_clear_a_routing_choice() { + let current = + json!({"routing_group_id": "saved", "chat_pii_redaction": {"enabled": false}}); + let injection = json!({"routing_group_id": "hidden", "routing_group_name": "forged", "chat_pii_redaction": {"enabled": true}}); + let created = merge_api_key_feature_settings(None, Some(Some(injection.clone())), None) + .unwrap() + .flatten() + .unwrap(); + assert!(created.get("routing_group_id").is_none()); + assert!(created.get("routing_group_name").is_none()); + for feature_patch in [ + Some(injection), + Some(json!({"routing_group_id": null})), + None, + ] { + let updated = merge_api_key_feature_settings(Some(¤t), Some(feature_patch), None) + .unwrap() + .flatten() + .unwrap(); + assert_eq!(updated["routing_group_id"], "saved"); + assert!(updated.get("routing_group_name").is_none()); + } + let untouched = merge_api_key_feature_settings(Some(¤t), None, None).unwrap(); + assert_eq!( + untouched, None, + "name/rate/IP-only updates must leave settings untouched" + ); + } + + #[test] + fn validated_top_level_selection_and_clear_preserve_other_feature_settings() { + let current = json!({"routing_group_id": "saved", "chat_pii_redaction": {"enabled": false}, "another_setting": 7}); + let selected = + merge_api_key_feature_settings(Some(¤t), None, Some(Some("public".to_string()))) + .unwrap() + .flatten() + .unwrap(); + assert_eq!(selected["routing_group_id"], "public"); + assert_eq!(selected["another_setting"], 7); + assert_eq!(selected["chat_pii_redaction"]["enabled"], false); + let cleared = merge_api_key_feature_settings(Some(&selected), None, Some(None)) + .unwrap() + .flatten() + .unwrap(); + assert!(cleared.get("routing_group_id").is_none()); + assert_eq!(cleared["another_setting"], 7); + let replacement = merge_api_key_feature_settings( + None, + Some(Some(json!({"routing_group_id": "hidden"}))), + Some(Some("public".to_string())), + ) + .unwrap() + .flatten() + .unwrap(); + assert_eq!(replacement, json!({"routing_group_id": "public"})); + } + + #[tokio::test] + async fn names_are_resolved_from_the_catalog_without_persisting_a_name_snapshot() { + let state = state(); + let names = routing_group_names(&state, true).await; + let current = json!({"routing_group_id": "hidden", "routing_group_name": "stale"}); + let fields = routing_group_payload_fields(Some(¤t), &names); + assert_eq!(fields["routing_group_id"], "hidden"); + assert_eq!(fields["routing_group_name"], "hidden current name"); + let missing = + routing_group_payload_fields(Some(&json!({"routing_group_id": "deleted"})), &names); + assert_eq!(missing["routing_group_id"], "deleted"); + assert_eq!(missing["routing_group_name"], Value::Null); + let default = routing_group_payload_fields(None, &names); + assert_eq!(default["routing_group_id"], Value::Null); + assert_eq!(default["routing_group_name"], Value::Null); + } +} diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs b/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs index baec48c2f..b042752d3 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_api_keys.rs @@ -26,6 +26,14 @@ use super::{ const USERS_ME_API_KEY_WRITE_UNAVAILABLE_DETAIL: &str = "用户 API 密钥写入暂不可用"; +#[path = "user_me_api_key_routing.rs"] +mod routing_selection; +use routing_selection::{ + api_key_routing_group_id, deserialize_routing_group_patch, merge_api_key_feature_settings, + normalize_api_key_feature_settings_patch, routing_group_names, routing_group_payload_fields, + validate_routing_group_patch, +}; + fn users_me_api_key_secret_response(mut response: Response) -> Response { response.headers_mut().insert( http::header::CACHE_CONTROL, @@ -43,6 +51,8 @@ struct UsersMeCreateApiKeyRequest { concurrent_limit: Option, #[serde(default)] feature_settings: Option, + #[serde(default)] + routing_group_id: Option, #[serde(default, alias = "allowed_ips")] ip_rules: Option>, } @@ -57,6 +67,8 @@ struct UsersMeUpdateApiKeyRequest { concurrent_limit: Option, #[serde(default, deserialize_with = "deserialize_optional_json_patch")] feature_settings: Option>, + #[serde(default, deserialize_with = "deserialize_routing_group_patch")] + routing_group_id: Option>, #[serde( default, alias = "allowed_ips", @@ -170,8 +182,9 @@ fn build_users_me_api_key_list_payload( state: &AppState, record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, is_locked: bool, + group_names: &BTreeMap, ) -> serde_json::Value { - json!({ + let mut payload = json!({ "id": record.api_key_id, "name": record.name, "key_display": users_me_masked_api_key_display(state, record), @@ -187,15 +200,24 @@ fn build_users_me_api_key_list_payload( "ip_rules": record.ip_rules, "force_capabilities": record.force_capabilities, "feature_settings": record.feature_settings, - }) + }); + payload + .as_object_mut() + .expect("API key payload is an object") + .extend(routing_group_payload_fields( + record.feature_settings.as_ref(), + group_names, + )); + payload } fn build_users_me_api_key_detail_payload( state: &AppState, record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord, is_locked: bool, + group_names: &BTreeMap, ) -> serde_json::Value { - json!({ + let mut payload = json!({ "id": record.api_key_id, "name": record.name, "key_display": users_me_masked_api_key_display(state, record), @@ -210,7 +232,15 @@ fn build_users_me_api_key_detail_payload( "last_used_at": format_users_me_optional_unix_secs_iso8601(record.last_used_at_unix_secs), "expires_at": format_users_me_optional_unix_secs_iso8601(record.expires_at_unix_secs), "created_at": format_users_me_optional_unix_secs_iso8601(record.created_at_unix_secs), - }) + }); + payload + .as_object_mut() + .expect("API key payload is an object") + .extend(routing_group_payload_fields( + record.feature_settings.as_ref(), + group_names, + )); + payload } fn normalize_users_me_required_api_key_name(value: &str) -> Result { @@ -336,6 +366,13 @@ pub(super) async fn handle_users_me_api_keys_get( }; records.retain(|record| !record.is_standalone); records.sort_by(|left, right| left.api_key_id.cmp(&right.api_key_id)); + let group_names = routing_group_names( + state, + records + .iter() + .any(|record| api_key_routing_group_id(record.feature_settings.as_ref()).is_some()), + ) + .await; let snapshot_ids = records .iter() @@ -367,7 +404,7 @@ pub(super) async fn handle_users_me_api_keys_get( .get(&record.api_key_id) .map(|snapshot| snapshot.api_key_is_locked) .unwrap_or(false); - build_users_me_api_key_list_payload(state, record, is_locked) + build_users_me_api_key_list_payload(state, record, is_locked, &group_names) }) .collect::>(), ) @@ -461,8 +498,16 @@ pub(super) async fn handle_users_me_api_key_detail_get( } }; + let group_names = routing_group_names( + state, + api_key_routing_group_id(record.feature_settings.as_ref()).is_some(), + ) + .await; Json(build_users_me_api_key_detail_payload( - state, &record, is_locked, + state, + &record, + is_locked, + &group_names, )) .into_response() } @@ -567,8 +612,17 @@ pub(super) async fn handle_users_me_api_key_create( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } }; - let feature_settings = match normalize_feature_settings(payload.feature_settings) { - Ok(value) => value, + let routing_group_patch = + match validate_routing_group_patch(state, None, Some(payload.routing_group_id)).await { + Ok(value) => value, + Err(response) => return response, + }; + let feature_settings = match merge_api_key_feature_settings( + None, + Some(payload.feature_settings), + routing_group_patch, + ) { + Ok(value) => value.flatten(), Err(detail) => { return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } @@ -631,26 +685,36 @@ pub(super) async fn handle_users_me_api_key_create( return build_users_me_api_key_writer_unavailable_response(); }; - users_me_api_key_secret_response( - Json(json!({ - "id": created.api_key_id, - "name": created.name, - "key": plaintext_key, - "key_display": users_me_masked_api_key_display(state, &created), - "is_active": created.is_active, - "is_locked": false, - "rate_limit": created.rate_limit, - "concurrent_limit": created.concurrent_limit, - "ip_rules": created.ip_rules, - "feature_settings": created.feature_settings, - "last_used_at": format_users_me_optional_unix_secs_iso8601(created.last_used_at_unix_secs), - "created_at": format_users_me_optional_unix_secs_iso8601(created.created_at_unix_secs), - "total_requests": created.total_requests, - "total_cost_usd": created.total_cost_usd, - "message": "API密钥创建成功", - })) - .into_response(), + let group_names = routing_group_names( + state, + api_key_routing_group_id(created.feature_settings.as_ref()).is_some(), ) + .await; + let mut payload = json!({ + "id": created.api_key_id, + "name": created.name, + "key": plaintext_key, + "key_display": users_me_masked_api_key_display(state, &created), + "is_active": created.is_active, + "is_locked": false, + "rate_limit": created.rate_limit, + "concurrent_limit": created.concurrent_limit, + "ip_rules": created.ip_rules, + "feature_settings": created.feature_settings, + "last_used_at": format_users_me_optional_unix_secs_iso8601(created.last_used_at_unix_secs), + "created_at": format_users_me_optional_unix_secs_iso8601(created.created_at_unix_secs), + "total_requests": created.total_requests, + "total_cost_usd": created.total_cost_usd, + "message": "API密钥创建成功", + }); + payload + .as_object_mut() + .expect("API key payload is an object") + .extend(routing_group_payload_fields( + created.feature_settings.as_ref(), + &group_names, + )); + users_me_api_key_secret_response(Json(payload).into_response()) } pub(super) async fn handle_users_me_api_key_update( @@ -715,14 +779,39 @@ pub(super) async fn handle_users_me_api_key_update( return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); } }; - let feature_settings = match payload.feature_settings { - Some(value) => match normalize_feature_settings(value) { - Ok(value) => Some(value), - Err(detail) => { - return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false); + let current_features = if matches!(payload.routing_group_id.as_ref(), Some(Some(_))) { + match state + .read_auth_api_key_feature_settings(&auth.user.id, &snapshot.api_key_id, false) + .await + { + Ok(value) => value, + Err(error) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("user API key feature settings lookup failed: {error:?}"), + false, + ) } - }, - None => None, + } + } else { + None + }; + let routing_group_patch = match validate_routing_group_patch( + state, + current_features.as_ref(), + payload.routing_group_id, + ) + .await + { + Ok(value) => value, + Err(response) => return response, + }; + let feature_settings = match normalize_api_key_feature_settings_patch(payload.feature_settings) + { + Ok(value) => value, + Err(detail) => { + return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false) + } }; let ip_rules = match payload.ip_rules { Some(value) => match normalize_users_me_ip_rules(value) { @@ -752,6 +841,11 @@ pub(super) async fn handle_users_me_api_key_update( concurrent_limit_present, ip_rules, feature_settings, + routing_group_selection: Some( + aether_data::repository::auth::UpdateApiKeyRoutingGroupSelection { + group_id: routing_group_patch, + }, + ), }, ) .await @@ -768,8 +862,17 @@ pub(super) async fn handle_users_me_api_key_update( return build_users_me_api_key_mutation_conflict_response(); }; - let mut payload = - build_users_me_api_key_detail_payload(state, &updated, snapshot.api_key_is_locked); + let group_names = routing_group_names( + state, + api_key_routing_group_id(updated.feature_settings.as_ref()).is_some(), + ) + .await; + let mut payload = build_users_me_api_key_detail_payload( + state, + &updated, + snapshot.api_key_is_locked, + &group_names, + ); payload["message"] = json!("API密钥已更新"); Json(payload).into_response() } @@ -1095,7 +1198,8 @@ mod tests { use axum::{response::IntoResponse, Json}; use super::{ - normalize_users_me_ip_rules, users_me_api_key_secret_response, UsersMeUpdateApiKeyRequest, + normalize_users_me_ip_rules, users_me_api_key_secret_response, UsersMeCreateApiKeyRequest, + UsersMeUpdateApiKeyRequest, }; use serde_json::json; @@ -1110,6 +1214,28 @@ mod tests { ); } + #[test] + fn routing_group_selection_patch_distinguishes_missing_clear_and_valid_string() { + let unchanged: UsersMeUpdateApiKeyRequest = + serde_json::from_value(json!({"name": "renamed"})).unwrap(); + assert_eq!(unchanged.routing_group_id, None); + let cleared: UsersMeUpdateApiKeyRequest = + serde_json::from_value(json!({"routing_group_id": null})).unwrap(); + assert_eq!(cleared.routing_group_id, Some(None)); + let selected: UsersMeUpdateApiKeyRequest = + serde_json::from_value(json!({"routing_group_id": "group-1"})).unwrap(); + assert_eq!(selected.routing_group_id, Some(Some("group-1".to_string()))); + let created: UsersMeCreateApiKeyRequest = + serde_json::from_value(json!({"name": "created", "routing_group_id": "group-1"})) + .unwrap(); + assert_eq!(created.routing_group_id.as_deref(), Some("group-1")); + for invalid in [json!(true), json!(3), json!([]), json!({})] { + let request = json!({"name": "invalid", "routing_group_id": invalid}); + assert!(serde_json::from_value::(request.clone()).is_err()); + assert!(serde_json::from_value::(request).is_err()); + } + } + #[test] fn normalize_ip_rules_trims_ip_and_cidr_values() { let values = normalize_users_me_ip_rules(Some(vec![ diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs b/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs index 42b930491..b5945df35 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_routes.rs @@ -17,12 +17,12 @@ use super::{ handle_users_me_model_capabilities_get, handle_users_me_model_capabilities_put, handle_users_me_password_patch, handle_users_me_preferences_get, handle_users_me_preferences_put, handle_users_me_providers_get, handle_users_me_referral_get, - handle_users_me_sessions_get, handle_users_me_update_session, handle_users_me_usage_active_get, - handle_users_me_usage_get, handle_users_me_usage_heatmap_get, - handle_users_me_usage_interval_timeline_get, handle_users_me_vscodex_request, - users_me_api_key_capabilities_path_matches, users_me_api_key_detail_path_matches, - users_me_api_key_install_sessions_path_matches, users_me_api_key_providers_path_matches, - users_me_management_token_detail_path_matches, + handle_users_me_routing_groups_get, handle_users_me_sessions_get, + handle_users_me_update_session, handle_users_me_usage_active_get, handle_users_me_usage_get, + handle_users_me_usage_heatmap_get, handle_users_me_usage_interval_timeline_get, + handle_users_me_vscodex_request, users_me_api_key_capabilities_path_matches, + users_me_api_key_detail_path_matches, users_me_api_key_install_sessions_path_matches, + users_me_api_key_providers_path_matches, users_me_management_token_detail_path_matches, users_me_management_token_regenerate_path_matches, users_me_management_token_toggle_path_matches, users_me_management_tokens_root, users_me_session_detail_path_matches, AppState, GatewayPublicRequestContext, @@ -231,6 +231,11 @@ pub(crate) async fn maybe_build_local_users_me_response( Some("providers") if request_context.request_path == "/api/users/me/providers" => { Some(handle_users_me_providers_get(state, request_context, headers).await) } + Some("routing_groups") + if request_context.request_path == "/api/users/me/routing-groups" => + { + Some(handle_users_me_routing_groups_get(state, request_context, headers).await) + } Some("preferences") if request_context.request_path == "/api/users/me/preferences" => { Some(handle_users_me_preferences_get(state, request_context, headers).await) } diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_routing_groups.rs b/apps/aether-gateway/src/handlers/public/support/user_me_routing_groups.rs new file mode 100644 index 000000000..b0336cd56 --- /dev/null +++ b/apps/aether-gateway/src/handlers/public/support/user_me_routing_groups.rs @@ -0,0 +1,58 @@ +use aether_routing_core::RoutingGroupConfig; +use axum::{ + body::Body, + http, + response::{IntoResponse, Response}, + Json, +}; +use serde_json::json; + +use super::{ + build_auth_error_response, resolve_authenticated_local_user, AppState, + GatewayPublicRequestContext, +}; +use crate::routing::selection::routing_group_is_user_visible; + +pub(super) async fn handle_users_me_routing_groups_get( + state: &AppState, + request_context: &GatewayPublicRequestContext, + headers: &http::HeaderMap, +) -> Response { + if let Err(response) = resolve_authenticated_local_user(state, request_context, headers).await { + return response; + } + if !state.has_routing_group_data_reader() { + return build_auth_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "策略分组目录暂不可用", + false, + ); + } + let groups = match state.list_routing_groups().await { + Ok(groups) => groups, + Err(error) => { + return build_auth_error_response( + http::StatusCode::INTERNAL_SERVER_ERROR, + format!("user routing group lookup failed: {error:?}"), + false, + ) + } + }; + let items = groups + .into_iter() + .filter(|group| group.enabled && routing_group_is_user_visible(group)) + .filter_map(|group| { + let config = serde_json::from_value::(group.config_json).ok()?; + if !config.billing_multiplier.is_finite() || config.billing_multiplier < 0.0 { + return None; + } + Some(json!({ + "id": group.id, + "name": group.name, + "billing_multiplier": config.billing_multiplier, + "is_default": group.is_system_default, + })) + }) + .collect::>(); + Json(json!({"total": items.len(), "items": items})).into_response() +} diff --git a/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs b/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs index 15a421748..83e59de2e 100644 --- a/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs +++ b/apps/aether-gateway/src/handlers/public/support/user_me_usage.rs @@ -561,6 +561,8 @@ fn build_users_me_usage_record_payload( let cache_read_price_per_1m = item.settlement_cache_read_price_per_1m(); let cache_creation_input_tokens = users_me_usage_cache_creation_tokens(item); let rate_multiplier = item.settlement_rate_multiplier(); + let billing_multiplier = item.billing_multiplier(); + let billing_cost = item.billing_cost().map(|cost| round_to(cost, 6)); let client_is_stream = users_me_usage_client_is_stream(item); let upstream_is_stream = users_me_usage_upstream_is_stream(item); let mut payload = json!({ @@ -576,6 +578,8 @@ fn build_users_me_usage_record_payload( "output_tokens": item.output_tokens, "total_tokens": item.total_tokens, "cost": round_to(item.total_cost_usd, 6), + "billing_multiplier": billing_multiplier, + "billing_cost": billing_cost, "response_time_ms": item.response_time_ms, "first_byte_time_ms": item.first_byte_time_ms, "is_stream": item.is_stream, @@ -614,6 +618,8 @@ fn build_users_me_usage_record_payload( ), }); payload["end_to_end_time_ms"] = json!(users_me_usage_metadata_u64(item, "end_to_end_time_ms")); + payload["routing_group_id"] = json!(item.routing_group_id()); + payload["routing_group_name"] = json!(item.routing_group_name()); payload["end_to_end_first_byte_time_ms"] = json!(users_me_usage_metadata_u64( item, "end_to_end_first_byte_time_ms" @@ -643,6 +649,8 @@ fn build_users_me_usage_record_payload( fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_json::Value { let cache_creation_input_tokens = users_me_usage_cache_creation_tokens(item); + let billing_multiplier = item.billing_multiplier(); + let billing_cost = item.billing_cost().map(|cost| round_to(cost, 6)); let client_is_stream = users_me_usage_client_is_stream(item); let upstream_is_stream = users_me_usage_upstream_is_stream(item); let mut payload = json!({ @@ -659,6 +667,8 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_ "cost": round_to(item.total_cost_usd, 6), "actual_cost": round_to(item.actual_total_cost_usd, 6), "rate_multiplier": item.settlement_rate_multiplier(), + "billing_multiplier": billing_multiplier, + "billing_cost": billing_cost, "response_time_ms": item.response_time_ms, "first_byte_time_ms": item.first_byte_time_ms, "updated_at": unix_secs_to_rfc3339(item.updated_at_unix_secs), @@ -685,6 +695,8 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_ "response_model": item.provider_response_model(), "has_fallback": item.has_fallback(), }); + payload["routing_group_id"] = json!(item.routing_group_id()); + payload["routing_group_name"] = json!(item.routing_group_name()); payload["end_to_end_time_ms"] = json!(users_me_usage_metadata_u64(item, "end_to_end_time_ms")); payload["end_to_end_first_byte_time_ms"] = json!(users_me_usage_metadata_u64( item, @@ -1867,6 +1879,59 @@ mod tests { assert_eq!(payload["cache_creation_ephemeral_1h_input_tokens"], 6); } + #[test] + fn user_usage_payloads_preserve_routing_group_snapshot_and_precise_display_cost() { + for (metadata, multiplier, cost, group_name) in [ + (None, 1.0, json!(0.0), serde_json::Value::Null), + ( + Some(json!({"routing_group_billing_multiplier": 0.0})), + 0.0, + json!(0.0), + serde_json::Value::Null, + ), + ( + Some(json!({ + "routing_group_billing_multiplier": 2.5, + "routing_group_id": "group-1", + "routing_group_name": "请求时的分组", + "rate_multiplier": 0.5 + })), + 2.5, + json!(0.000004), + json!("请求时的分组"), + ), + ( + Some(json!({ + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 2.5, "user_group": 2.0}, "multiplier": 5.0}, + "routing_group_billing_multiplier": 2.5, + "routing_group_id": "group-1", + "routing_group_name": "请求时的分组", + "rate_multiplier": 0.5 + })), + 5.0, + json!(0.000007), + json!("请求时的分组"), + ), + ] { + let item = StoredRequestUsageAudit { + total_cost_usd: 0.00000149, + actual_total_cost_usd: 0.0000002, + request_metadata: metadata, + ..sample_usage("completed") + }; + let record = build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false); + let active = build_users_me_usage_active_payload(&item); + for payload in [&record, &active] { + assert_eq!(payload["billing_multiplier"], multiplier); + assert_eq!(payload["billing_cost"], cost); + assert_eq!(payload["routing_group_name"], group_name); + assert_eq!(payload["cost"], 0.000001); + } + assert!(record.get("actual_cost").is_none()); + assert!(record.get("rate_multiplier").is_none()); + } + } + #[test] fn user_usage_payloads_expose_response_model_separately_from_mapping() { let item = StoredRequestUsageAudit { diff --git a/apps/aether-gateway/src/routing/resolver.rs b/apps/aether-gateway/src/routing/resolver.rs index 547df881c..2469f37f6 100644 --- a/apps/aether-gateway/src/routing/resolver.rs +++ b/apps/aether-gateway/src/routing/resolver.rs @@ -84,10 +84,19 @@ pub(crate) fn resolve_gateway_static_default_routing_policy( let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else { return Ok(None); }; + let billing_multiplier = match input.group_config_json.get("billing_multiplier") { + Some(value) => value + .as_f64() + .filter(|value| value.is_finite() && *value >= 0.0) + .ok_or_else(invalid_routing_group_config)?, + None => 1.0, + }; crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy.clone()); Ok(Some(ResolvedRoutingPolicy { + billing_multiplier, group_id: input.group_id.map(str::to_string), + group_name: None, group_version: input.group_version, selection_source: input.selection_source.to_string(), requested_model: input.requested_model.to_string(), @@ -216,6 +225,7 @@ mod tests { #[test] fn static_default_policy_matches_full_resolver_without_body_context() { let config = json!({ + "billing_multiplier": 2.5, "default_policy": { "priority_mode": "global_key", "scheduling_mode": "load_balance", @@ -262,6 +272,7 @@ mod tests { .expect("full policy should resolve"); assert_eq!(static_policy, full_policy); + assert_eq!(static_policy.billing_multiplier, 2.5); assert_eq!(static_policy.execution_policy.max_transfer_count, 3); assert_eq!( static_policy.execution_policy.max_transfer_timeout_seconds, @@ -288,6 +299,50 @@ mod tests { assert!(static_policy.matched_rules.is_empty()); } + #[test] + fn static_default_billing_multiplier_defaults_to_one_and_validates_input() { + for config in [ + json!({}), + json!({"billing_multiplier": 0.0}), + json!({"billing_multiplier": 0.25}), + ] { + let policy = + resolve_gateway_static_default_routing_policy(GatewayStaticRoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + group_config_json: &config, + selection_source: "system_default", + requested_model: "model-a", + resolved_model: "model-a", + }) + .unwrap() + .unwrap(); + let expected = config + .get("billing_multiplier") + .and_then(Value::as_f64) + .unwrap_or(1.0); + assert_eq!(policy.billing_multiplier, expected); + assert_eq!( + crate::routing::build_routing_trace_seed(&policy, "openai:chat").billing_multiplier, + Some(expected) + ); + } + for value in [json!(-1), json!(null), json!("2"), json!(true)] { + let config = json!({"billing_multiplier": value}); + assert!(resolve_gateway_static_default_routing_policy( + GatewayStaticRoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + group_config_json: &config, + selection_source: "system_default", + requested_model: "model-a", + resolved_model: "model-a", + }, + ) + .is_err()); + } + } + #[test] fn dynamic_routing_config_is_not_static_default() { let config = json!({ @@ -323,17 +378,16 @@ mod tests { "model_policies": [], "rules": [] }); - let static_policy = resolve_gateway_static_default_routing_policy( - GatewayStaticRoutingPolicyInput { + let static_policy = + resolve_gateway_static_default_routing_policy(GatewayStaticRoutingPolicyInput { group_id: Some("group-1"), group_version: Some(1), group_config_json: &config, selection_source: "system_default", requested_model: model, resolved_model: model, - }, - ) - .expect("provider exclusions should defer to the full resolver"); + }) + .expect("provider exclusions should defer to the full resolver"); assert!(static_policy.is_none()); let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput { @@ -356,6 +410,66 @@ mod tests { } } + #[test] + fn model_provider_enablement_requires_full_resolver_and_valid_boolean_values() { + for (model, enabled) in [("model-a", false), ("model-b", true)] { + let config = json!({ "model_policies": [{ + "model": "model-a", "provider_enabled_overrides": { "provider-1": false } + }] }); + assert!(resolve_gateway_static_default_routing_policy( + GatewayStaticRoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + group_config_json: &config, + selection_source: "system_default", + requested_model: model, + resolved_model: model, + } + ) + .unwrap() + .is_none()); + let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + group_config_json: &config, + selection_source: "system_default", + requested_model: model, + resolved_model: model, + api_format: "openai:chat", + user_id: None, + api_key_id: None, + headers: &json!({}), + body: &json!({}), + phase: RoutingRulePhase::ClientRequest, + }) + .unwrap(); + assert_eq!( + policy.ranking_overlay.provider_allowed("provider-1"), + enabled + ); + } + for invalid in [json!("false"), json!(0), Value::Null] { + let config = json!({ "model_policies": [{ + "model": "model-a", "provider_enabled_overrides": { "provider-1": invalid } + }] }); + assert!(resolve_gateway_routing_policy(GatewayRoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + group_config_json: &config, + selection_source: "system_default", + requested_model: "model-a", + resolved_model: "model-a", + api_format: "openai:chat", + user_id: None, + api_key_id: None, + headers: &json!({}), + body: &json!({}), + phase: RoutingRulePhase::ClientRequest, + }) + .is_err()); + } + } + #[test] fn empty_provider_exclusions_preserve_static_default_fast_path() { for config in [json!({}), json!({"disabled_providers": []})] { diff --git a/apps/aether-gateway/src/routing/selection.rs b/apps/aether-gateway/src/routing/selection.rs index 2d3d2e953..1ce4f934a 100644 --- a/apps/aether-gateway/src/routing/selection.rs +++ b/apps/aether-gateway/src/routing/selection.rs @@ -23,6 +23,9 @@ pub(crate) enum GatewayRoutingSelectionError { #[derive(Debug, Clone, Default)] pub(crate) struct GatewayRoutingSelectionInput<'a> { pub explicit_group: Option<&'a str>, + /// A public group selected on the user's API key. Explicit request headers + /// take precedence, while an unavailable saved choice must fail closed. + pub preferred_group: Option<&'a str>, pub user_id: Option<&'a str>, pub api_key_id: Option<&'a str>, pub user_group_ids: &'a [String], @@ -69,6 +72,28 @@ pub(crate) async fn select_gateway_routing_group( }); } + if let Some(preferred) = input + .preferred_group + .map(str::trim) + .filter(|value| !value.is_empty()) + { + let group = repository + .find_routing_group(RoutingGroupLookupKey::Id(preferred)) + .await + .map_err(repository_selection_error)? + .ok_or_else(|| GatewayRoutingSelectionError::NotFound(preferred.to_string()))?; + if !group.enabled { + return Err(GatewayRoutingSelectionError::Disabled(group.id)); + } + if !has_authenticated_principal(&input) || !routing_group_is_user_visible(&group) { + return Err(GatewayRoutingSelectionError::Forbidden(group.id)); + } + return Ok(GatewayRoutingGroupSelection { + group: Some(group), + source: "api_key_selection".to_string(), + }); + } + // When there are no bindings at all, no principal-specific lookup can // produce a group. The data-state repository answers this with a cached // existence query, so the common "routing configured but unused" case @@ -128,6 +153,11 @@ async fn explicit_group_allowed( group: &StoredRoutingGroup, input: &GatewayRoutingSelectionInput<'_>, ) -> Result { + // Public selection is opt-in. Turning it off does not revoke existing + // administrator-granted bindings or the system-default compatibility path. + if has_authenticated_principal(input) && routing_group_is_user_visible(group) { + return Ok(true); + } if group.is_system_default { return Ok(true); } @@ -147,6 +177,22 @@ async fn explicit_group_allowed( Ok(false) } +pub(crate) fn routing_group_is_user_visible(group: &StoredRoutingGroup) -> bool { + group + .config_json + .get("user_visible") + .and_then(serde_json::Value::as_bool) + == Some(true) +} + +fn has_authenticated_principal(input: &GatewayRoutingSelectionInput<'_>) -> bool { + input + .user_id + .into_iter() + .chain(input.api_key_id) + .any(|id| !id.trim().is_empty()) +} + fn repository_selection_error(error: impl std::fmt::Display) -> GatewayRoutingSelectionError { GatewayRoutingSelectionError::Repository(error.to_string()) } @@ -231,6 +277,225 @@ mod tests { } } + #[tokio::test] + async fn public_groups_allow_authenticated_selection_but_visibility_is_opt_in() { + let repository = InMemoryRoutingGroupRepository::default(); + for (id, config, enabled) in [ + ("public", json!({ "user_visible": true }), true), + ("private", json!({ "user_visible": false }), true), + ("legacy", json!({}), true), + ("malformed", json!({ "user_visible": "true" }), true), + ("disabled", json!({ "user_visible": true }), false), + ] { + repository + .create_routing_group(CreateRoutingGroupRecord { + id: id.into(), + name: format!("{id}-name"), + description: None, + enabled, + is_system_default: false, + sort_order: 0, + config_json: config, + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .unwrap(); + } + for (user_id, api_key_id) in [(Some("user-1"), None), (None, Some("key-1"))] { + for explicit in ["public", "public-name"] { + let selected = select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + explicit_group: Some(explicit), + user_id, + api_key_id, + user_group_ids: &[], + ..Default::default() + }, + ) + .await + .unwrap(); + assert_eq!(selected.group.unwrap().id, "public"); + } + } + for id in ["private", "legacy", "malformed"] { + assert_eq!( + select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + explicit_group: Some(id), + user_id: Some("user-1"), + ..Default::default() + } + ) + .await + .unwrap_err(), + GatewayRoutingSelectionError::Forbidden(id.into()) + ); + } + assert_eq!( + select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + explicit_group: Some("disabled"), + user_id: Some("user-1"), + ..Default::default() + } + ) + .await + .unwrap_err(), + GatewayRoutingSelectionError::Disabled("disabled".into()) + ); + assert_eq!( + select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + explicit_group: Some("public"), + ..Default::default() + } + ) + .await + .unwrap_err(), + GatewayRoutingSelectionError::Forbidden("public".into()) + ); + } + + #[tokio::test] + async fn api_key_selected_public_group_precedes_bindings_and_header_precedes_saved_choice() { + let repository = InMemoryRoutingGroupRepository::default(); + for (id, visible) in [ + ("selected", true), + ("header", true), + ("private-default", false), + ] { + repository + .create_routing_group(CreateRoutingGroupRecord { + id: id.into(), + name: id.into(), + description: None, + enabled: true, + is_system_default: false, + sort_order: 0, + config_json: json!({ "user_visible": visible }), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .unwrap(); + } + repository + .create_routing_group_binding(CreateRoutingGroupBindingRecord { + id: "admin-default".into(), + group_id: "private-default".into(), + subject_type: RoutingGroupBindingSubject::ApiKey, + subject_id: "key-1".into(), + is_default: true, + allow_explicit_select: true, + created_at: 1, + updated_at: 1, + }) + .await + .unwrap(); + + for (explicit, preferred, expected, source) in [ + (None, Some("selected"), "selected", "api_key_selection"), + ( + Some("header"), + Some("selected"), + "header", + "explicit_header", + ), + // An explicit authorized request overrides an invalid saved choice. + (Some("header"), Some("missing"), "header", "explicit_header"), + ( + Some("private-default"), + Some("selected"), + "private-default", + "explicit_header", + ), + (None, None, "private-default", "api_key_default"), + ] { + let selection = select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + explicit_group: explicit, + preferred_group: preferred, + api_key_id: Some("key-1"), + user_id: Some("user-1"), + user_group_ids: &[], + }, + ) + .await + .unwrap(); + assert_eq!(selection.group.unwrap().id, expected); + assert_eq!(selection.source, source); + } + } + + #[tokio::test] + async fn unavailable_api_key_group_selection_fails_closed_without_default_fallback() { + let repository = InMemoryRoutingGroupRepository::default(); + for (id, visible, enabled, is_default) in [ + ("private-default", false, true, true), + ("disabled", true, false, false), + ] { + repository + .create_routing_group(CreateRoutingGroupRecord { + id: id.into(), + name: format!("{id}-name"), + description: None, + enabled, + is_system_default: is_default, + sort_order: 0, + config_json: json!({ "user_visible": visible }), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .unwrap(); + } + for (id, error) in [ + ( + "private-default", + GatewayRoutingSelectionError::Forbidden("private-default".into()), + ), + ( + "disabled", + GatewayRoutingSelectionError::Disabled("disabled".into()), + ), + ( + "missing", + GatewayRoutingSelectionError::NotFound("missing".into()), + ), + // Saved selections are stable IDs, not mutable group names. + ( + "private-default-name", + GatewayRoutingSelectionError::NotFound("private-default-name".into()), + ), + ] { + assert_eq!( + select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + preferred_group: Some(id), + user_id: Some("user-1"), + ..Default::default() + } + ) + .await + .unwrap_err(), + error + ); + } + } + #[tokio::test] async fn selects_api_key_default_binding() { let repository = InMemoryRoutingGroupRepository::default(); @@ -268,6 +533,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: None, + preferred_group: None, user_id: None, api_key_id: Some("api-key-1"), user_group_ids: &[], @@ -304,6 +570,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: None, + preferred_group: None, user_id: Some("user-1"), api_key_id: Some("api-key-1"), user_group_ids: &["user-group-1".to_string()], @@ -327,7 +594,7 @@ mod tests { enabled: true, is_system_default: false, sort_order: 0, - config_json: json!({}), + config_json: json!({ "user_visible": false }), version: 1, created_at: 1, updated_at: 1, @@ -353,6 +620,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: Some("private-group"), + preferred_group: None, user_id: Some("user-1"), api_key_id: Some("api-key-1"), user_group_ids: &["team-1".to_string()], @@ -373,6 +641,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: Some("missing"), + preferred_group: None, user_id: Some("user-1"), api_key_id: Some("api-key-1"), user_group_ids: &[], @@ -397,6 +666,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: Some("group-name"), + preferred_group: None, user_id: Some("user-1"), api_key_id: Some("api-key-1"), user_group_ids: &[], @@ -423,6 +693,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: None, + preferred_group: None, user_id: Some("user-1"), api_key_id: Some("api-key-1"), user_group_ids: &[], @@ -463,6 +734,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: Some("disabled-group"), + preferred_group: None, user_id: Some("user-1"), api_key_id: Some("api-key-1"), user_group_ids: &[], @@ -514,6 +786,7 @@ mod tests { &repository, GatewayRoutingSelectionInput { explicit_group: Some("private-group"), + preferred_group: None, user_id: Some("user-1"), api_key_id: Some("api-key-1"), user_group_ids: &[], diff --git a/apps/aether-gateway/src/routing/trace.rs b/apps/aether-gateway/src/routing/trace.rs index 16270361b..a1835ad8c 100644 --- a/apps/aether-gateway/src/routing/trace.rs +++ b/apps/aether-gateway/src/routing/trace.rs @@ -5,7 +5,9 @@ pub(crate) fn build_routing_trace_seed( client_api_format: &str, ) -> RoutingDecisionTrace { RoutingDecisionTrace { + billing_multiplier: Some(policy.billing_multiplier), group_id: policy.group_id.clone(), + group_name: policy.group_name.clone(), group_version: policy.group_version, selection_source: policy.selection_source.clone(), selected_rules: policy diff --git a/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs b/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs index 6efd0a329..3a1179e5f 100644 --- a/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs +++ b/apps/aether-gateway/src/scheduler/candidate/tests/selection.rs @@ -370,7 +370,7 @@ async fn selects_by_provider_priority_when_priority_mode_is_provider() { } #[tokio::test] -async fn selects_by_global_key_priority_when_priority_mode_is_global_key() { +async fn legacy_global_key_routing_default_selects_by_provider_priority() { let mut provider_first = sample_row(); provider_first.provider_id = "provider-a".to_string(); provider_first.provider_name = "provider-a".to_string(); @@ -402,6 +402,13 @@ async fn selects_by_global_key_priority_when_priority_mode_is_global_key() { ) .await; + // Legacy configuration stays readable, but the routing-to-scheduler + // boundary normalizes it to the supported provider ordering mode. + assert_eq!( + ordering_config(&state).await.priority_mode, + aether_scheduler_core::SchedulerPriorityMode::Provider + ); + let selected = select_candidate( state.data.as_ref(), &state, @@ -415,8 +422,8 @@ async fn selects_by_global_key_priority_when_priority_mode_is_global_key() { .expect("selection should succeed") .expect("candidate should exist"); - assert_eq!(selected.provider_id, "provider-b"); - assert_eq!(selected.key_id, "key-b"); + assert_eq!(selected.provider_id, "provider-a"); + assert_eq!(selected.key_id, "key-a"); } #[tokio::test] diff --git a/apps/aether-gateway/src/tests/frontdoor/internal.rs b/apps/aether-gateway/src/tests/frontdoor/internal.rs index bce477a65..6780ecaa7 100644 --- a/apps/aether-gateway/src/tests/frontdoor/internal.rs +++ b/apps/aether-gateway/src/tests/frontdoor/internal.rs @@ -1038,6 +1038,13 @@ async fn gateway_rejects_internal_gateway_report_with_tampered_protected_context ("endpoint_id", json!("endpoint-unrelated-victim")), ("key_id", json!("key-unrelated-victim")), ("client_api_format", json!("gemini:video")), + ("routing_group_billing_multiplier", json!(0.0)), + ( + "billing_multiplier_snapshot", + json!({"version": 1, "factors": {"routing_group": 0.0}, "multiplier": 0.0}), + ), + ("routing_group_id", json!("unrelated-group")), + ("routing_group_name", json!("forged-group")), ("task_id", json!("task-unrelated-victim")), ("local_task_id", json!("local-task-unrelated-victim")), ("local_short_id", json!("short-unrelated-victim")), diff --git a/apps/aether-gateway/src/tests/frontdoor/public_support.rs b/apps/aether-gateway/src/tests/frontdoor/public_support.rs index ec7c5b551..d3dc5424b 100644 --- a/apps/aether-gateway/src/tests/frontdoor/public_support.rs +++ b/apps/aether-gateway/src/tests/frontdoor/public_support.rs @@ -52,10 +52,13 @@ const TEST_EMAIL_VERIFICATION_TOKEN: &str = #[path = "public_support/announcement_user_list.rs"] mod announcement_user_list; +mod api_key_routing; #[path = "public_support/auth_cookie.rs"] mod auth_cookie; #[path = "public_support/dashboard.rs"] mod dashboard; +#[path = "public_support/routing_groups.rs"] +mod routing_groups; #[path = "public_support/vscodex.rs"] mod vscodex; diff --git a/apps/aether-gateway/src/tests/frontdoor/public_support/api_key_routing.rs b/apps/aether-gateway/src/tests/frontdoor/public_support/api_key_routing.rs new file mode 100644 index 000000000..4aa9fecc0 --- /dev/null +++ b/apps/aether-gateway/src/tests/frontdoor/public_support/api_key_routing.rs @@ -0,0 +1,192 @@ +use super::*; +use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; +use aether_data_contracts::repository::routing_profiles::{ + CreateRoutingGroupRecord, RoutingGroupWriteRepository, UpdateRoutingGroupRecord, +}; + +#[tokio::test] +async fn users_me_api_key_routing_selection_round_trips_and_cannot_bypass_visibility() { + let groups = Arc::new(InMemoryRoutingGroupRepository::default()); + for (id, enabled, visible) in [ + ("public", true, true), + ("hidden", true, false), + ("disabled", false, true), + ] { + groups + .create_routing_group(CreateRoutingGroupRecord { + id: id.into(), + name: format!("{id} strategy"), + description: None, + enabled, + is_system_default: false, + sort_order: 0, + config_json: json!({"user_visible": visible}), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .unwrap(); + } + let now = Utc::now(); + let mut user = sample_auth_user(now); + user.role = "user".into(); + let token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".into(), json!(user.id)), + ("role".into(), json!(user.role)), + ( + "created_at".into(), + json!(user.created_at.map(|date| date.to_rfc3339())), + ), + ("session_id".into(), json!("session-api-key-routing")), + ]), + now + chrono::Duration::hours(1), + ); + let users = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user])); + let keys = Arc::new(InMemoryAuthApiKeySnapshotRepository::default()); + let (url, upstream_hits, gateway, upstream) = start_auth_gateway_with_builder(|| { + let data = GatewayDataState::with_auth_api_key_repository_for_tests(keys) + .with_user_reader(users) + .with_routing_group_repository_for_tests(groups.clone()) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY); + AppState::new() + .unwrap() + .with_data_state_for_tests(data) + .with_auth_sessions_for_tests([sample_auth_session( + "user-auth-1", + "session-api-key-routing", + "device-api-key-routing", + "refresh-api-key-routing", + now, + )]) + }) + .await; + let client = reqwest::Client::new(); + let endpoint = format!("{url}/api/users/me/api-keys"); + let request = |method: reqwest::Method, path: &str| { + client + .request(method, path) + .bearer_auth(&token) + .header("x-client-device-id", "device-api-key-routing") + .header("user-agent", "AetherTest/1.0") + }; + let response = request(reqwest::Method::POST, &endpoint) + .json(&json!({ + "name": "selected key", "routing_group_id": "public", + "feature_settings": {"chat_pii_redaction": {"enabled": true}}, + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let created: serde_json::Value = response.json().await.unwrap(); + assert_eq!(created["routing_group_id"], "public"); + assert_eq!(created["routing_group_name"], "public strategy"); + assert_eq!(created["feature_settings"]["routing_group_id"], "public"); + assert!(created["feature_settings"] + .get("routing_group_name") + .is_none()); + let detail_url = format!("{endpoint}/{}", created["id"].as_str().unwrap()); + + let list: serde_json::Value = request(reqwest::Method::GET, &endpoint) + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(list[0]["routing_group_id"], "public"); + assert_eq!(list[0]["routing_group_name"], "public strategy"); + let detail: serde_json::Value = request(reqwest::Method::GET, &detail_url) + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(detail["routing_group_id"], "public"); + assert_eq!(detail["routing_group_name"], "public strategy"); + + let response = request(reqwest::Method::PUT, &detail_url) + .json(&json!({ + "name": "renamed key", "feature_settings": { + "chat_pii_redaction": {"enabled": false}, "routing_group_id": "hidden", + }, + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let edited: serde_json::Value = response.json().await.unwrap(); + assert_eq!(edited["name"], "renamed key"); + assert_eq!(edited["routing_group_id"], "public"); + assert_eq!( + edited["feature_settings"]["chat_pii_redaction"]["enabled"], + false + ); + for id in ["hidden", "disabled", "missing"] { + let rejected = request(reqwest::Method::PUT, &detail_url) + .json(&json!({ + "name": "must not change", "routing_group_id": id, + })) + .send() + .await + .unwrap(); + assert_eq!(rejected.status(), StatusCode::BAD_REQUEST, "id={id}"); + } + let unchanged: serde_json::Value = request(reqwest::Method::GET, &detail_url) + .send() + .await + .unwrap() + .json() + .await + .unwrap(); + assert_eq!(unchanged["name"], "renamed key"); + assert_eq!(unchanged["routing_group_id"], "public"); + + groups + .update_routing_group( + "public", + UpdateRoutingGroupRecord { + name: Some("renamed strategy".into()), + config_json: Some(json!({"user_visible": false})), + ..Default::default() + }, + ) + .await + .unwrap(); + let response = request(reqwest::Method::PUT, &detail_url) + .json(&json!({ + "name": "existing hidden choice", "routing_group_id": "public", + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let renamed: serde_json::Value = response.json().await.unwrap(); + assert_eq!(renamed["routing_group_id"], "public"); + assert_eq!(renamed["routing_group_name"], "renamed strategy"); + let response = request(reqwest::Method::PUT, &detail_url) + .json(&json!({"routing_group_id": null})) + .send() + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let cleared: serde_json::Value = response.json().await.unwrap(); + assert_eq!(cleared["routing_group_id"], serde_json::Value::Null); + assert_eq!(cleared["routing_group_name"], serde_json::Value::Null); + assert!(cleared["feature_settings"] + .get("routing_group_id") + .is_none()); + assert_eq!( + cleared["feature_settings"]["chat_pii_redaction"]["enabled"], + false + ); + + assert_eq!(*upstream_hits.lock().unwrap(), 0); + gateway.abort(); + upstream.abort(); +} diff --git a/apps/aether-gateway/src/tests/frontdoor/public_support/routing_groups.rs b/apps/aether-gateway/src/tests/frontdoor/public_support/routing_groups.rs new file mode 100644 index 000000000..28288ca56 --- /dev/null +++ b/apps/aether-gateway/src/tests/frontdoor/public_support/routing_groups.rs @@ -0,0 +1,135 @@ +use super::*; +use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; +use aether_data_contracts::repository::routing_profiles::{ + CreateRoutingGroupRecord, RoutingGroupWriteRepository, UpdateRoutingGroupRecord, +}; + +#[tokio::test] +async fn users_me_routing_groups_requires_auth_and_exposes_only_public_enabled_summaries() { + let repository = Arc::new(InMemoryRoutingGroupRepository::default()); + for (id, config, enabled, is_default) in [ + ( + "discount", + json!({ "user_visible": true, "billing_multiplier": 0.5, + "disabled_providers": ["secret-provider"] }), + true, + false, + ), + ("regular", json!({ "user_visible": true }), true, false), + ( + "free", + json!({ "user_visible": true, "billing_multiplier": 0.0 }), + true, + false, + ), + ( + "private-default", + json!({ "user_visible": false }), + true, + true, + ), + ("legacy-private", json!({}), true, false), + ("disabled", json!({ "user_visible": true }), false, false), + ] { + repository + .create_routing_group(CreateRoutingGroupRecord { + id: id.into(), + name: format!("{id}-name"), + description: Some("private-description".into()), + enabled, + is_system_default: is_default, + sort_order: 0, + config_json: config, + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + }) + .await + .unwrap(); + } + let now = Utc::now(); + let mut user = sample_auth_user(now); + user.role = "user".into(); + let token = build_test_auth_token( + "access", + serde_json::Map::from_iter([ + ("user_id".into(), json!(user.id)), + ("role".into(), json!(user.role)), + ( + "created_at".into(), + json!(user.created_at.map(|date| date.to_rfc3339())), + ), + ("session_id".into(), json!("session-routing-groups")), + ]), + now + chrono::Duration::hours(1), + ); + let users = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user])); + let (url, upstream_hits, gateway, upstream) = start_auth_gateway_with_builder(|| { + let data = GatewayDataState::with_user_reader_for_tests(users) + .with_routing_group_repository_for_tests(repository.clone()); + AppState::new() + .unwrap() + .with_data_state_for_tests(data) + .with_auth_sessions_for_tests([sample_auth_session( + "user-auth-1", + "session-routing-groups", + "device-routing-groups", + "refresh-placeholder", + now, + )]) + }) + .await; + let client = reqwest::Client::new(); + let endpoint = format!("{url}/api/users/me/routing-groups"); + let unauthorized = client.get(&endpoint).send().await.unwrap(); + assert_eq!(unauthorized.status(), StatusCode::UNAUTHORIZED); + let request = || { + client + .get(&endpoint) + .bearer_auth(&token) + .header("x-client-device-id", "device-routing-groups") + .header("user-agent", "AetherTest/1.0") + }; + let response = request().send().await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let payload: serde_json::Value = response.json().await.unwrap(); + assert_eq!(payload["total"], 3); + let items = payload["items"].as_array().unwrap(); + assert_eq!(items.len(), 3); + for (id, multiplier) in [("discount", 0.5), ("regular", 1.0), ("free", 0.0)] { + assert_eq!( + items.iter().find(|item| item["id"] == id).unwrap(), + &json!({ + "id": id, "name": format!("{id}-name"), + "billing_multiplier": multiplier, "is_default": false, + }) + ); + } + assert!(!payload.to_string().contains("private")); + assert!(!payload.to_string().contains("secret-provider")); + repository + .update_routing_group( + "private-default", + UpdateRoutingGroupRecord { + config_json: Some(json!({ "user_visible": true })), + ..Default::default() + }, + ) + .await + .unwrap(); + let payload: serde_json::Value = request().send().await.unwrap().json().await.unwrap(); + assert_eq!(payload["total"], 4); + assert_eq!( + payload["items"] + .as_array() + .unwrap() + .iter() + .find(|item| item["id"] == "private-default") + .unwrap()["is_default"], + true + ); + assert_eq!(*upstream_hits.lock().unwrap(), 0); + gateway.abort(); + upstream.abort(); +} diff --git a/apps/aether-gateway/tests/architecture/sql_and_data.rs b/apps/aether-gateway/tests/architecture/sql_and_data.rs index b68729259..7fa199888 100644 --- a/apps/aether-gateway/tests/architecture/sql_and_data.rs +++ b/apps/aether-gateway/tests/architecture/sql_and_data.rs @@ -1449,7 +1449,7 @@ fn usage_repositories_are_owned_by_contracts_and_driver_adapters() { "sql" ) .len(), - 27, + 29, "all PostgreSQL usage SQL fragments should be owned by the adapter crate" ); } diff --git a/crates/aether-admin/src/observability/usage.rs b/crates/aether-admin/src/observability/usage.rs index 8b8465508..f839e3ccf 100644 --- a/crates/aether-admin/src/observability/usage.rs +++ b/crates/aether-admin/src/observability/usage.rs @@ -1295,6 +1295,8 @@ fn admin_usage_active_request_json( let cache_creation_input_tokens = admin_usage_cache_creation_tokens(item); let client_is_stream = admin_usage_client_is_stream(item); let upstream_is_stream = admin_usage_upstream_is_stream(item); + let billing_multiplier = item.billing_multiplier(); + let billing_cost = item.billing_cost().map(|cost| round_to(cost, 6)); let mut value = json!({ "id": item.id, "status": item.status, @@ -1334,6 +1336,11 @@ fn admin_usage_active_request_json( "request_path_and_query": admin_usage_metadata_string(item, "request_path_and_query"), "has_fallback": admin_usage_has_fallback(item), }); + value["billing_multiplier"] = json!(billing_multiplier); + value["billing_cost"] = json!(billing_cost); + value["routing_group_id"] = json!(item.routing_group_id()); + value["routing_group_name"] = json!(item.routing_group_name()); + value["rate_multiplier"] = json!(item.settlement_rate_multiplier()); value["end_to_end_time_ms"] = json!(admin_usage_metadata_u64(item, "end_to_end_time_ms")); value["end_to_end_first_byte_time_ms"] = json!(admin_usage_metadata_u64( item, @@ -1395,6 +1402,8 @@ pub fn admin_usage_record_json( .unwrap_or_else(|| "已删除用户".to_string()); let client_is_stream = admin_usage_client_is_stream(item); let upstream_is_stream = admin_usage_upstream_is_stream(item); + let billing_multiplier = item.billing_multiplier(); + let billing_cost = item.billing_cost().map(|cost| round_to(cost, 6)); let mut payload = json!({ "id": item.id, @@ -1443,6 +1452,10 @@ pub fn admin_usage_record_json( "provider_key_name": provider_key_name, "model_version": Value::Null, }); + payload["billing_multiplier"] = json!(billing_multiplier); + payload["billing_cost"] = json!(billing_cost); + payload["routing_group_id"] = json!(item.routing_group_id()); + payload["routing_group_name"] = json!(item.routing_group_name()); let object = payload .as_object_mut() .expect("admin usage record payload should be an object"); @@ -2723,7 +2736,7 @@ pub fn build_admin_usage_replay_plan_response( mod tests { use std::collections::BTreeMap; - use serde_json::json; + use serde_json::{json, Value}; use super::{ admin_usage_active_request_json, admin_usage_client_is_stream, admin_usage_has_body_value, @@ -2848,6 +2861,98 @@ mod tests { assert_eq!(record["client_is_stream"], false); } + #[test] + fn admin_usage_payloads_preserve_routing_group_snapshot_and_precise_display_cost() { + for (metadata, multiplier, cost, group_name) in [ + (None, 1.0, json!(0.0), Value::Null), + ( + Some(json!({"routing_group_billing_multiplier": 0.0})), + 0.0, + json!(0.0), + Value::Null, + ), + ( + Some(json!({ + "routing_group_billing_multiplier": 2.5, + "routing_group_id": "group-1", + "routing_group_name": "请求时的分组", + "rate_multiplier": 0.5 + })), + 2.5, + json!(0.000004), + json!("请求时的分组"), + ), + ( + Some(json!({ + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 2.5, "user_group": 2.0}, "multiplier": 5.0}, + "routing_group_billing_multiplier": 2.5, + "routing_group_id": "group-1", + "routing_group_name": "请求时的分组", + "rate_multiplier": 0.5 + })), + 5.0, + json!(0.000007), + json!("请求时的分组"), + ), + ] { + let item = StoredRequestUsageAudit { + total_cost_usd: 0.00000149, + actual_total_cost_usd: 0.0000002, + request_metadata: metadata, + ..sample_usage("completed", Some(200), None) + }; + let record = admin_usage_record_json( + &item, + &BTreeMap::new(), + &BTreeMap::new(), + false, + false, + None, + ); + let active = admin_usage_active_request_json(&item, None, None, None); + let detail = build_admin_usage_detail_payload( + &item, + &BTreeMap::new(), + &BTreeMap::new(), + false, + false, + None, + false, + None, + &BTreeMap::new(), + ); + for payload in [&record, &active, &detail] { + assert_eq!(payload["billing_multiplier"], multiplier); + assert_eq!(payload["billing_cost"], cost); + assert_eq!(payload["routing_group_name"], group_name); + assert_eq!( + payload["routing_group_id"], + if group_name.is_null() { + Value::Null + } else { + json!("group-1") + } + ); + assert_eq!( + payload["rate_multiplier"], + if group_name.is_null() { + Value::Null + } else { + json!(0.5) + } + ); + assert_eq!(payload["cost"], 0.000001); + assert_eq!(payload["actual_cost"], 0.0); + } + } + let item = StoredRequestUsageAudit { + total_cost_usd: f64::MAX, + request_metadata: Some(json!({"routing_group_billing_multiplier": 2.0})), + ..sample_usage("completed", Some(200), None) + }; + assert!(admin_usage_active_request_json(&item, None, None, None)["billing_cost"].is_null()); + } + #[test] fn admin_usage_payloads_expose_response_model_separately_from_mapping() { let item = StoredRequestUsageAudit { diff --git a/crates/aether-ai/formats/src/codex_profile.rs b/crates/aether-ai/formats/src/codex_profile.rs index ddd16b3cc..44592c9b4 100644 --- a/crates/aether-ai/formats/src/codex_profile.rs +++ b/crates/aether-ai/formats/src/codex_profile.rs @@ -109,7 +109,10 @@ mod tests { assert_eq!(profile.originator, "codex_cli_rs"); assert!(profile.user_agent.starts_with("codex_cli_rs/0.200.1 (")); assert!(profile.user_agent.ends_with(") unknown")); - assert!(profile.user_agent.contains(std::env::consts::ARCH)); + let architecture = super::OS_INFO + .architecture() + .unwrap_or(std::env::consts::ARCH); + assert!(profile.user_agent.contains(&format!("; {architecture})"))); } #[test] diff --git a/crates/aether-data/adapters/postgres/migrations/20261007000000_separate_customer_billing_amount.sql b/crates/aether-data/adapters/postgres/migrations/20261007000000_separate_customer_billing_amount.sql new file mode 100644 index 000000000..890482f0e --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20261007000000_separate_customer_billing_amount.sql @@ -0,0 +1,154 @@ +-- Customer charges use the immutable request-time factor snapshot. Provider +-- procurement cost remains in actual_total_cost_usd for legacy reporting. +CREATE OR REPLACE FUNCTION public.usage_customer_billable_amount( + metadata jsonb, base_cost numeric, legacy_cost numeric +) RETURNS numeric LANGUAGE plpgsql IMMUTABLE PARALLEL SAFE AS $$ +DECLARE factor jsonb; multiplier numeric; amount numeric; + factor_name text; factor_value jsonb; factor_number double precision; + expected_multiplier double precision := 1.0; factor_count integer := 0; + has_zero boolean := false; +BEGIN + IF metadata ? 'billing_multiplier_snapshot' THEN + IF jsonb_typeof(metadata->'billing_multiplier_snapshot') <> 'object' + OR metadata #> '{billing_multiplier_snapshot,version}' IS DISTINCT FROM '1'::jsonb + OR jsonb_typeof(metadata #> '{billing_multiplier_snapshot,factors}') IS DISTINCT FROM 'object' + THEN RETURN NULL; END IF; + factor := metadata #> '{billing_multiplier_snapshot,multiplier}'; + FOR factor_name, factor_value IN + SELECT key, value FROM jsonb_each(metadata #> '{billing_multiplier_snapshot,factors}') ORDER BY key COLLATE "C" + LOOP + factor_count := factor_count + 1; + IF factor_count > 16 OR factor_name = '' OR length(factor_name) > 64 + OR factor_name !~ '^[A-Za-z0-9_]+$' + OR jsonb_typeof(factor_value) IS DISTINCT FROM 'number' + THEN RETURN NULL; END IF; + factor_number := factor_value::text::double precision; + IF factor_number < 0 OR factor_number > 1.7976931348623157e308::double precision + THEN RETURN NULL; END IF; + has_zero := has_zero OR factor_number = 0; + END LOOP; + -- Rust short-circuits zero before multiplying any of the other factors. + IF has_zero THEN expected_multiplier := 0; + ELSE + FOR factor_name, factor_value IN + SELECT key, value FROM jsonb_each(metadata #> '{billing_multiplier_snapshot,factors}') ORDER BY key COLLATE "C" + LOOP + factor_number := factor_value::text::double precision; + BEGIN + expected_multiplier := expected_multiplier * factor_number; + EXCEPTION WHEN numeric_value_out_of_range THEN + -- PostgreSQL raises on float underflow; Rust rounds that product to 0. + IF expected_multiplier::numeric * factor_number::numeric > 1.7976931348623157e308::numeric + THEN RETURN NULL; END IF; + expected_multiplier := 0; + END; + END LOOP; + END IF; + ELSIF metadata ? 'routing_group_billing_multiplier' THEN + factor := metadata->'routing_group_billing_multiplier'; + expected_multiplier := NULL; + ELSE + RETURN CASE WHEN legacy_cost NOT IN ('NaN'::numeric,'Infinity'::numeric,'-Infinity'::numeric) + THEN round(legacy_cost,8) END; + END IF; + IF jsonb_typeof(factor) IS DISTINCT FROM 'number' THEN RETURN NULL; END IF; + multiplier := factor::text::numeric; + factor_number := factor::text::double precision; + IF factor_number < 0 + OR factor_number > 1.7976931348623157e308::double precision + OR (expected_multiplier IS NOT NULL AND factor_number <> expected_multiplier) + OR multiplier < 0 OR multiplier > 1.7976931348623157e308::numeric + OR base_cost IS NULL OR base_cost < 0 + OR base_cost IN ('NaN'::numeric,'Infinity'::numeric,'-Infinity'::numeric) + THEN RETURN NULL; END IF; + amount := base_cost * multiplier; + IF amount > 1.7976931348623157e308::numeric THEN RETURN NULL; END IF; + RETURN round(amount,8); +EXCEPTION WHEN numeric_value_out_of_range OR invalid_text_representation THEN + -- Corrupt captured pricing must not abort an entire analytics query. + RETURN NULL; +END $$; + +CREATE OR REPLACE VIEW public.usage_analytics_facts_v1 AS +SELECT u.request_id, COALESCE(u.id, u.request_id) AS id, u.created_at, + CASE WHEN identity.owner_id IS NOT NULL AND identity.is_standalone=false THEN identity.owner_id END AS actor_user_id, + identity.owner_id AS credential_owner_id, + CASE WHEN identity.owner_id IS NULL THEN 'unknown' WHEN identity.is_standalone THEN 'standalone' + WHEN NOT identity.is_standalone THEN 'employee' ELSE 'unknown' END AS attribution_kind, + CASE WHEN identity.owner_id IS NULL THEN 'unknown' WHEN identity.is_standalone THEN 'standalone_key' + WHEN NOT identity.is_standalone THEN 'user_account' ELSE 'unknown' END AS attribution_source, + COALESCE(a.record_kind, 'request') AS record_kind, a.parent_request_id, + u.api_key_id, u.model, u.target_model, u.provider_id, u.provider_name, + u.api_format, u.endpoint_kind, u.request_type, u.is_stream, u.has_format_conversion, + u.status, u.status_code, u.error_category, u.failure_origin, u.failure_stage, u.failure_reason, + u.failure_schema_version, u.response_time_ms, u.first_byte_time_ms, + COALESCE(s.billing_status, u.billing_status) AS settlement_status, + COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb AS usage_available, + COALESCE(metadata.value->'usage_pricing_available', 'true'::jsonb) <> 'false'::jsonb + AND (s.billing_total_cost_usd IS NOT NULL OR COALESCE(s.billing_status, u.billing_status) = 'settled') AS pricing_available, + CASE WHEN COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb + THEN b.input_tokens END AS input_tokens, + CASE WHEN COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb + THEN b.output_tokens END AS output_tokens, + CASE WHEN COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb + THEN b.total_tokens END AS total_tokens, + CASE WHEN COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb + THEN b.cache_read_input_tokens END AS cache_read_input_tokens, + CASE WHEN COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb + THEN b.cache_creation_input_tokens END AS cache_creation_input_tokens, + CASE WHEN COALESCE(metadata.value->'usage_pricing_available', 'true'::jsonb) <> 'false'::jsonb + AND (s.billing_total_cost_usd IS NOT NULL OR COALESCE(s.billing_status, u.billing_status) = 'settled') + THEN round(COALESCE(s.billing_total_cost_usd::numeric, u.total_cost_usd::numeric), 8) END AS rated_amount, + CASE WHEN COALESCE(metadata.value->'usage_pricing_available', 'true'::jsonb) <> 'false'::jsonb + AND (s.billing_actual_total_cost_usd IS NOT NULL OR COALESCE(s.billing_status, u.billing_status) = 'settled') + THEN public.usage_customer_billable_amount(metadata.value, + COALESCE(s.billing_total_cost_usd::numeric, u.total_cost_usd::numeric), + COALESCE(s.billing_actual_total_cost_usd::numeric, u.actual_total_cost_usd::numeric)) END AS billable_amount, + s.quota_covered_amount_usd AS quota_covered_amount, + s.wallet_consumed_amount_usd AS wallet_consumed_amount, + s.wallet_debit_amount_usd AS wallet_debit_amount, + s.wallet_recharge_debit_usd AS wallet_recharge_debit_amount, + s.wallet_gift_debit_usd AS wallet_gift_debit_amount, + s.wallet_overdraft_usd AS wallet_overdraft_amount, + s.allocation_status, s.finalized_at AS settled_at, + CASE WHEN s.billing_total_cost_usd IS NOT NULL THEN 'settlement_snapshot' ELSE 'legacy_float' END AS amount_source, + b.upstream_is_stream, + CASE WHEN metadata.value #>> '{analytics_measurement,source}' IN ('reported','estimated','mixed') + THEN metadata.value #>> '{analytics_measurement,source}' ELSE 'unknown' END AS token_source, + CASE WHEN COALESCE(metadata.value->'usage_available','true'::jsonb) <> 'false'::jsonb + AND COALESCE(metadata.value->'usage_pricing_available','true'::jsonb) <> 'false'::jsonb + AND s.input_price_per_1m IS NOT NULL AND s.billing_cache_read_cost_usd IS NOT NULL + THEN round(s.input_price_per_1m::numeric * b.cache_read_input_tokens::numeric / 1000000,8) END AS cache_estimated_full_cost_amount, + CASE WHEN COALESCE(metadata.value->'usage_available','true'::jsonb) <> 'false'::jsonb + AND COALESCE(metadata.value->'usage_pricing_available','true'::jsonb) <> 'false'::jsonb + AND s.input_price_per_1m IS NOT NULL AND s.billing_cache_read_cost_usd IS NOT NULL + THEN round(s.billing_cache_read_cost_usd::numeric,8) END AS cache_read_cost_amount, + CASE WHEN COALESCE(metadata.value->'usage_available','true'::jsonb) <> 'false'::jsonb + AND COALESCE(metadata.value->'usage_pricing_available','true'::jsonb) <> 'false'::jsonb + AND s.input_price_per_1m IS NOT NULL AND s.billing_cache_creation_cost_usd IS NOT NULL + THEN round(s.billing_cache_creation_cost_usd::numeric,8) END AS cache_creation_cost_amount +FROM public.usage u +-- OFFSET 0 keeps this projection from being flattened: large metadata is +-- detoasted and parsed once per request, rather than once per metric expression. +CROSS JOIN LATERAL (SELECT u.request_metadata::jsonb AS value OFFSET 0) metadata +LEFT JOIN public.usage_settlement_snapshots s USING (request_id) +LEFT JOIN public.usage_attribution_snapshots a USING (request_id) +JOIN public.usage_billing_facts b USING (request_id) +LEFT JOIN public.api_keys k ON k.id=u.api_key_id +CROSS JOIN LATERAL ( + SELECT CASE WHEN a.request_id IS NOT NULL THEN a.credential_owner_id + WHEN EXISTS (SELECT 1 FROM public.users WHERE id=u.user_id AND NOT is_deleted) THEN u.user_id END AS owner_id, + COALESCE(k.is_standalone, + CASE WHEN jsonb_typeof(metadata.value #> '{analytics_attribution,is_standalone}')='boolean' + THEN (metadata.value #>> '{analytics_attribution,is_standalone}')::boolean END, + CASE WHEN jsonb_typeof(metadata.value->'api_key_is_standalone')='boolean' + THEN (metadata.value->>'api_key_is_standalone')::boolean END, + CASE WHEN a.attribution_source='user_account' THEN false + WHEN a.attribution_source='standalone_key' THEN true END, + CASE WHEN u.api_key_id IS NULL THEN false END) AS is_standalone +) identity; + +-- Do not backfill existing rows or scan historical usage during the upgrade. +-- Historical daily totals retain their legacy charge through the read fallback; +-- normal daily aggregation writes billing_cost for newly aggregated days. +ALTER TABLE public.stats_daily ADD COLUMN IF NOT EXISTS billing_cost numeric(20,8); diff --git a/crates/aether-data/adapters/postgres/src/auth.rs b/crates/aether-data/adapters/postgres/src/auth.rs index 573dcdeea..7fbbb4d81 100644 --- a/crates/aether-data/adapters/postgres/src/auth.rs +++ b/crates/aether-data/adapters/postgres/src/auth.rs @@ -561,7 +561,18 @@ SET rate_limit = CASE WHEN $7 THEN $8 ELSE rate_limit END, concurrent_limit = CASE WHEN $9 THEN $10 ELSE concurrent_limit END, ip_rules = CASE WHEN $11 THEN $12::jsonb ELSE ip_rules END, - feature_settings = CASE WHEN $13 THEN $14::jsonb ELSE feature_settings END, + feature_settings = CASE WHEN $16 THEN + NULLIF( + (COALESCE(CASE WHEN $13 THEN $14::jsonb ELSE feature_settings END, '{}'::jsonb) + - 'routing_group_id' - 'routing_group_name') + || CASE WHEN $17 THEN + CASE WHEN $18::text IS NULL THEN '{}'::jsonb + ELSE jsonb_build_object('routing_group_id', $18::text) END + WHEN jsonb_typeof(feature_settings->'routing_group_id') = 'string' THEN + jsonb_build_object('routing_group_id', feature_settings->'routing_group_id') + ELSE '{}'::jsonb END, + '{}'::jsonb) + ELSE CASE WHEN $13 THEN $14::jsonb ELSE feature_settings END END, updated_at = NOW() WHERE user_id = $1 AND id = $2 @@ -1357,6 +1368,18 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository { .bind(record.feature_settings.is_some()) .bind(feature_settings) .bind(false) + .bind(record.routing_group_selection.is_some()) + .bind( + record + .routing_group_selection + .as_ref() + .is_some_and(|patch| patch.group_id.is_some()), + ) + .bind( + record + .routing_group_selection + .and_then(|patch| patch.group_id.flatten()), + ) .fetch_optional(&self.pool) .await .map_postgres_err()?; @@ -1418,6 +1441,18 @@ WHERE id = $2 .bind(record.feature_settings.is_some()) .bind(feature_settings) .bind(true) + .bind(record.routing_group_selection.is_some()) + .bind( + record + .routing_group_selection + .as_ref() + .is_some_and(|patch| patch.group_id.is_some()), + ) + .bind( + record + .routing_group_selection + .and_then(|patch| patch.group_id.flatten()), + ) .fetch_optional(&self.pool) .await .map_postgres_err()?; @@ -2143,12 +2178,126 @@ mod tests { .contains("key_encrypted = CASE WHEN $3 THEN $4 ELSE key_encrypted END")); assert!(UPDATE_USER_API_KEY_BASIC_SQL .contains("ip_rules = CASE WHEN $11 THEN $12::jsonb ELSE ip_rules END")); - assert!(UPDATE_USER_API_KEY_BASIC_SQL.contains( - "feature_settings = CASE WHEN $13 THEN $14::jsonb ELSE feature_settings END" - )); + assert!(UPDATE_USER_API_KEY_BASIC_SQL + .contains("CASE WHEN $13 THEN $14::jsonb ELSE feature_settings END")); assert!(UPDATE_USER_API_KEY_BASIC_SQL.contains("AND ($15 = FALSE OR is_locked = FALSE)")); } + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL; uses only a temporary table"] + async fn live_api_key_routing_patch_preserves_concurrent_feature_edits() { + use aether_data_contracts::repository::auth::{ + AuthApiKeyWriteRepository, UpdateApiKeyRoutingGroupSelection, + UpdateUserApiKeyBasicRecord, + }; + use serde_json::json; + + let pool = sqlx::postgres::PgPoolOptions::new() + .max_connections(1) + .connect(&std::env::var("AETHER_TEST_DATABASE_URL").unwrap()) + .await + .unwrap(); + // The repository's complete production UPDATE runs against a session-local table. + sqlx::raw_sql( + r#" +CREATE TEMP TABLE api_keys ( + id text PRIMARY KEY, user_id text, key_hash text, key_encrypted text, name text, + allowed_providers json, allowed_api_formats json, allowed_models json, + ip_rules jsonb, rate_limit integer, concurrent_limit integer, + force_capabilities json, feature_settings jsonb, is_active boolean DEFAULT true, + is_locked boolean DEFAULT false, is_standalone boolean DEFAULT false, + expires_at timestamptz, auto_delete_on_expiry boolean DEFAULT false, + total_requests bigint DEFAULT 0, total_tokens bigint DEFAULT 0, + total_cost_usd numeric DEFAULT 0, last_used_at timestamptz, + created_at timestamptz DEFAULT NOW(), updated_at timestamptz DEFAULT NOW() +); +INSERT INTO api_keys (id,user_id,key_hash,name,feature_settings) +VALUES ('key-1','user-1','hash-1','key','{"routing_group_id":"a","pii":false}'); +"#, + ) + .execute(&pool) + .await + .unwrap(); + let repository = SqlxAuthApiKeySnapshotReadRepository::new(pool.clone()); + let patch = |features, group_id| UpdateUserApiKeyBasicRecord { + user_id: "user-1".into(), + api_key_id: "key-1".into(), + key_encrypted: None, + key_encrypted_present: false, + name: None, + name_present: false, + rate_limit: None, + rate_limit_present: false, + concurrent_limit: None, + concurrent_limit_present: false, + ip_rules: None, + feature_settings: features, + routing_group_selection: Some(UpdateApiKeyRoutingGroupSelection { group_id }), + }; + // Prepared before the group change: stale or injected group fields must not win. + let stale_feature_edit = + patch(Some(Some(json!({"routing_group_id":"a","pii":true}))), None); + repository + .update_user_api_key_basic_if_unlocked(patch(None, Some(Some("b".into())))) + .await + .unwrap() + .unwrap(); + let edited = repository + .update_user_api_key_basic_if_unlocked(stale_feature_edit) + .await + .unwrap() + .unwrap(); + assert_eq!( + edited.feature_settings, + Some(json!({"routing_group_id":"b","pii":true})) + ); + let changed = repository + .update_user_api_key_basic_if_unlocked(patch(None, Some(Some("c".into())))) + .await + .unwrap() + .unwrap(); + assert_eq!( + changed.feature_settings, + Some(json!({"routing_group_id":"c","pii":true})) + ); + let cleared_features = repository + .update_user_api_key_basic_if_unlocked(patch(Some(None), None)) + .await + .unwrap() + .unwrap(); + assert_eq!( + cleared_features.feature_settings, + Some(json!({"routing_group_id":"c"})) + ); + let cleared_group = repository + .update_user_api_key_basic_if_unlocked(patch(None, Some(None))) + .await + .unwrap() + .unwrap(); + assert_eq!(cleared_group.feature_settings, None); + let mut admin = patch(Some(Some(json!({"admin":true}))), None); + admin.routing_group_selection = None; + assert_eq!( + repository + .update_user_api_key_basic(admin) + .await + .unwrap() + .unwrap() + .feature_settings, + Some(json!({"admin":true})) + ); + sqlx::query("UPDATE api_keys SET is_locked=true") + .execute(&pool) + .await + .unwrap(); + assert!(repository + .update_user_api_key_basic_if_unlocked(patch(None, Some(Some("d".into())))) + .await + .unwrap() + .is_none()); + pool.close().await; + } + #[tokio::test] async fn repository_constructs_from_lazy_pool() { let factory = PostgresPoolFactory::new(PostgresPoolConfig { diff --git a/crates/aether-data/adapters/postgres/src/settlement.rs b/crates/aether-data/adapters/postgres/src/settlement.rs index a03c2271b..e31704918 100644 --- a/crates/aether-data/adapters/postgres/src/settlement.rs +++ b/crates/aether-data/adapters/postgres/src/settlement.rs @@ -1509,6 +1509,85 @@ mod tests { (pool, schema) } + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] + async fn live_composite_billing_settlement_preserves_provider_cost_and_is_idempotent() { + use super::*; + + let (pool, schema) = isolated_settlement_test_pool().await; + let result = AssertUnwindSafe(async { + for table in ["wallets", "usage", "usage_settlement_snapshots", "usage_counter_deltas"] { + sqlx::query(&format!("CREATE TABLE {table} (LIKE public.{table} INCLUDING ALL)")) + .execute(&pool).await.expect("settlement fixture table should be created"); + } + let repository = SqlxSettlementRepository::new(pool.clone()); + for (scenario, charge, quota_covered) in [ + ("wallet", 20.0, 0.0), + ("quota_and_wallet", 20.0, 7.0), + ("zero_charge", 0.0, 0.0), + ] { + sqlx::query("INSERT INTO users (id, username, email_verified) VALUES ($1, $1, false)") + .bind(scenario).execute(&pool).await.expect("user should insert"); + sqlx::query("INSERT INTO wallets (id, user_id, balance, gift_balance, total_consumed, limit_mode, created_at, updated_at) VALUES ($1, $1, 100, 0, 0, 'finite', NOW(), NOW())") + .bind(scenario).execute(&pool).await.expect("wallet should insert"); + // A zero-charge request must leave an active quota untouched too. + if scenario != "wallet" { + let grant = serde_json::json!([{ + "type": "daily_quota", "daily_quota_usd": 7.0, + "reset_timezone": "UTC", "allow_wallet_overage": true, + }]); + sqlx::query("INSERT INTO billing_plans (id, title, price_amount, duration_unit, duration_value, entitlements_json, created_at, updated_at) VALUES ($1, $1, 10, 'month', 1, $2, NOW(), NOW())") + .bind(scenario).bind(&grant).execute(&pool).await.expect("plan should insert"); + sqlx::query("INSERT INTO user_plan_entitlements (id, user_id, plan_id, payment_order_id, starts_at, expires_at, entitlements_snapshot, created_at, updated_at) VALUES ($1, $1, $1, $1, NOW() - INTERVAL '1 hour', NOW() + INTERVAL '1 day', $2, NOW(), NOW())") + .bind(scenario).bind(&grant).execute(&pool).await.expect("entitlement should insert"); + } + let multiplier = charge / 10.0; + let metadata = serde_json::json!({"billing_multiplier_snapshot": { + "version": 1, "factors": {"routing_group": multiplier}, "multiplier": multiplier, + }}); + sqlx::query("INSERT INTO usage (id, request_id, user_id, provider_id, provider_name, model, status, billing_status, total_cost_usd, actual_total_cost_usd, request_metadata) VALUES ($1, $1, $1, 'provider', 'Provider', 'model', 'completed', 'pending', 10, 5, $2)") + .bind(scenario).bind(metadata).execute(&pool).await.expect("usage should insert"); + let input = UsageSettlementInput { + request_id: scenario.to_string(), user_id: Some(scenario.to_string()), + api_key_id: None, api_key_is_standalone: false, + provider_id: Some("provider".to_string()), + status: "completed".to_string(), billing_status: "pending".to_string(), + total_cost_usd: 10.0, actual_total_cost_usd: 5.0, + billing_cost_usd: Some(charge), finalized_at_unix_secs: None, + }; + let settled = repository.settle_usage(input.clone()).await.unwrap().unwrap(); + assert_eq!(settled.billing_status, "settled", "{scenario}"); + assert_eq!(settled.wallet_balance_before, Some(100.0)); + assert_eq!(settled.wallet_balance_after, Some(100.0 - (charge - quota_covered))); + assert_eq!(repository.settle_usage(input).await.unwrap(), Some(settled), "replayed {scenario}"); + + let wallet: (f64, f64) = sqlx::query_as("SELECT (balance + gift_balance)::double precision, total_consumed::double precision FROM wallets WHERE id = $1") + .bind(scenario).fetch_one(&pool).await.unwrap(); + assert_eq!(wallet, (100.0 - (charge - quota_covered), charge - quota_covered), "{scenario}"); + let quota: (i64, f64) = sqlx::query_as("SELECT COUNT(*), COALESCE(SUM(amount_usd), 0)::double precision FROM entitlement_usage_ledgers WHERE request_id = $1") + .bind(scenario).fetch_one(&pool).await.unwrap(); + assert_eq!(quota, (if quota_covered > 0.0 { 1 } else { 0 }, quota_covered), "{scenario}"); + let allocation: (f64, f64, f64, String) = sqlx::query_as("SELECT quota_covered_amount_usd::double precision, wallet_consumed_amount_usd::double precision, wallet_debit_amount_usd::double precision, allocation_status FROM usage_settlement_snapshots WHERE request_id = $1") + .bind(scenario).fetch_one(&pool).await.unwrap(); + assert_eq!(allocation, (quota_covered, charge - quota_covered, charge - quota_covered, "complete".to_string()), "{scenario}"); + let costs: (f64, f64) = sqlx::query_as("SELECT total_cost_usd::double precision, actual_total_cost_usd::double precision FROM usage WHERE request_id = $1") + .bind(scenario).fetch_one(&pool).await.unwrap(); + assert_eq!(costs, (10.0, 5.0), "base and upstream cost must remain unchanged"); + let provider_cost: (i64, f64) = sqlx::query_as("SELECT COUNT(*), COALESCE(SUM(total_cost_usd_delta), 0)::double precision FROM usage_counter_deltas WHERE request_id = $1 AND kind = 'provider_monthly' AND target_id = 'provider'") + .bind(scenario).fetch_one(&pool).await.unwrap(); + assert_eq!(provider_cost, (1, 5.0), "upstream cost must be recorded once even for a zero-charge request"); + } + }).catch_unwind().await; + sqlx::query(&format!("DROP SCHEMA {schema} CASCADE")) + .execute(&pool) + .await + .expect("isolated settlement schema should be removed"); + pool.close().await; + if let Err(panic) = result { + std::panic::resume_unwind(panic); + } + } + #[tokio::test] #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] async fn live_usage_policy_window_aggregates_preserve_exact_admission_and_idempotency() { diff --git a/crates/aether-data/adapters/postgres/src/usage/analytics_tests.rs b/crates/aether-data/adapters/postgres/src/usage/analytics_tests.rs index 785fef557..9d0328016 100644 --- a/crates/aether-data/adapters/postgres/src/usage/analytics_tests.rs +++ b/crates/aether-data/adapters/postgres/src/usage/analytics_tests.rs @@ -682,6 +682,7 @@ async fn live_overview_settlement_allocations_preserve_unlimited_and_finite_wall billing_status: "pending".into(), total_cost_usd: cost, actual_total_cost_usd: cost, + billing_cost_usd: None, finalized_at_unix_secs: None, }; assert_eq!( @@ -1266,6 +1267,19 @@ async fn live_overview_dashboard_total_matches_canonical_settlement_and_legacy_t serde_json::json!({}), 1002, ), + ( + "billing-snapshot", + "openai:chat", + 120, + serde_json::json!({ + "billing_multiplier_snapshot": { + "version": 1, + "factors": {"routing_group": 2.0, "user_group": 0.75}, + "multiplier": 1.5 + } + }), + 120, + ), ( "unavailable", "openai:chat", @@ -1340,7 +1354,114 @@ async fn live_overview_dashboard_total_matches_canonical_settlement_and_legacy_t .await .unwrap(); assert_eq!(total.total_tokens, expected_tokens, "{case}"); + if case == "billing-snapshot" { + assert_eq!(total.billable_amount.as_deref(), Some("0.37500000")); + } assert_dashboard_total_matches_canonical(&total, &canonical); } tx.rollback().await.unwrap(); } + +#[tokio::test] +#[ignore = "requires migrated isolated AETHER_TEST_DATABASE_URL"] +async fn live_customer_billing_amount_matches_canonical_and_dashboard_facts() { + let pool = sqlx::PgPool::connect(&std::env::var("AETHER_TEST_DATABASE_URL").unwrap()) + .await + .unwrap(); + let mut tx = pool.begin().await.unwrap(); + let start = Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(); + let composite = serde_json::json!({ + "billing_multiplier_snapshot": { + "version": 1, "factors": {"routing_group": 2.0, "user_group": 0.75}, "multiplier": 1.5 + }, + "routing_group_billing_multiplier": 99.0, + "rate_multiplier": 0.25 + }); + for (case, metadata, expected) in [ + ("legacy", serde_json::json!({}), Some("0.50000000")), + ("composite", composite.clone(), Some("3.00000000")), + ("settlement-base", composite, Some("6.00000000")), + ( + "free", + serde_json::json!({"routing_group_billing_multiplier": 0}), + Some("0.00000000"), + ), + ( + "null", + serde_json::json!({"billing_multiplier_snapshot": null, "routing_group_billing_multiplier": 1}), + None, + ), + ( + "negative-factor", + serde_json::json!({"billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": -1}, "multiplier": 1}}), + None, + ), + ( + "mismatch", + serde_json::json!({"billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 2}, "multiplier": 1}}), + None, + ), + ( + "negative-legacy", + serde_json::json!({"routing_group_billing_multiplier": -1}), + None, + ), + ( + "string-legacy", + serde_json::json!({"routing_group_billing_multiplier": "1"}), + None, + ), + ( + "zero-before-overflow", + serde_json::json!({"billing_multiplier_snapshot": {"version": 1, "factors": {"a": 1e308, "b": 1e308, "z": 0}, "multiplier": 0}}), + Some("0.00000000"), + ), + ( + "overflow", + serde_json::json!({"billing_multiplier_snapshot": {"version": 1, "factors": {"a": 1e308, "b": 1e308}, "multiplier": 1}}), + None, + ), + ( + "bad-key", + serde_json::json!({"billing_multiplier_snapshot": {"version": 1, "factors": {"routing-group": 2}, "multiplier": 2}}), + None, + ), + ] { + let request = uuid::Uuid::new_v4().to_string(); + sqlx::query("INSERT INTO usage(id,request_id,model,provider_name,status,billing_status,total_cost_usd,actual_total_cost_usd,created_at,request_metadata) VALUES($1,$1,$1,'billing-test','completed','settled',2,0.5,$2,$3)") + .bind(&request).bind(start).bind(metadata).execute(&mut *tx).await.unwrap(); + if case == "settlement-base" { + sqlx::query("INSERT INTO usage_settlement_snapshots(request_id,billing_status,billing_total_cost_usd,billing_actual_total_cost_usd) VALUES($1,'settled',4,0.25)") + .bind(&request).execute(&mut *tx).await.unwrap(); + } + let amount: Option = sqlx::query_scalar( + "SELECT billable_amount::text FROM usage_analytics_facts_v1 WHERE request_id=$1", + ) + .bind(&request) + .fetch_one(&mut *tx) + .await + .unwrap(); + assert_eq!(amount.as_deref(), expected, "{case}"); + let query = UsageAnalyticsQuery { + from_unix_ms: start.timestamp_millis() as u64, + to_unix_ms: (start + chrono::Duration::hours(1)).timestamp_millis() as u64, + model: Some(request), + ..Default::default() + }; + let canonical = super::analytics::read_analytics_metrics(&mut tx, &query, false) + .await + .unwrap(); + let inline = super::dashboard::read_dashboard_total_metrics(&mut tx, &query, false) + .await + .unwrap(); + assert_dashboard_total_matches_canonical(&inline, &canonical); + if let Some(expected) = expected { + assert_eq!( + canonical.billable_amount.as_deref(), + Some(expected), + "{case}" + ); + } + } + tx.rollback().await.unwrap(); +} diff --git a/crates/aether-data/adapters/postgres/src/usage/dashboard.rs b/crates/aether-data/adapters/postgres/src/usage/dashboard.rs index 0c9e260db..59a5e9b55 100644 --- a/crates/aether-data/adapters/postgres/src/usage/dashboard.rs +++ b/crates/aether-data/adapters/postgres/src/usage/dashboard.rs @@ -17,10 +17,10 @@ SELECT u.created_at, u.api_key_id, u.model, u.provider_id, u.api_format, u.endpo u.request_type, u.status, u.is_stream, u.has_format_conversion, u.failure_origin, 'request'::text AS record_kind, COALESCE(s.billing_status, u.billing_status) AS settlement_status, - COALESCE(availability.usage_available, 'true'::jsonb) <> 'false'::jsonb AS usage_available, - COALESCE(availability.usage_pricing_available, 'true'::jsonb) <> 'false'::jsonb + COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb AS usage_available, + COALESCE(metadata.value->'usage_pricing_available', 'true'::jsonb) <> 'false'::jsonb AND (s.billing_total_cost_usd IS NOT NULL OR COALESCE(s.billing_status, u.billing_status) = 'settled') AS pricing_available, - CASE WHEN COALESCE(availability.usage_available, 'true'::jsonb) <> 'false'::jsonb THEN + CASE WHEN COALESCE(metadata.value->'usage_available', 'true'::jsonb) <> 'false'::jsonb THEN GREATEST( COALESCE( CASE @@ -89,15 +89,15 @@ SELECT u.created_at, u.api_key_id, u.model, u.provider_id, u.api_format, u.endpo ), 0 )::bigint END AS total_tokens, - CASE WHEN COALESCE(availability.usage_pricing_available, 'true'::jsonb) <> 'false'::jsonb + CASE WHEN COALESCE(metadata.value->'usage_pricing_available', 'true'::jsonb) <> 'false'::jsonb AND (s.billing_actual_total_cost_usd IS NOT NULL OR COALESCE(s.billing_status, u.billing_status) = 'settled') - THEN round(COALESCE(s.billing_actual_total_cost_usd::numeric, u.actual_total_cost_usd::numeric), 8) END AS billable_amount, + THEN public.usage_customer_billable_amount(metadata.value, + COALESCE(s.billing_total_cost_usd::numeric, u.total_cost_usd::numeric), + COALESCE(s.billing_actual_total_cost_usd::numeric, u.actual_total_cost_usd::numeric)) END AS billable_amount, s.allocation_status FROM public.usage u LEFT JOIN public.usage_settlement_snapshots s USING (request_id) -CROSS JOIN LATERAL json_to_record( - CASE WHEN json_typeof(u.request_metadata)='object' THEN u.request_metadata ELSE '{}'::json END -) AS availability(usage_available jsonb, usage_pricing_available jsonb) +CROSS JOIN LATERAL (SELECT u.request_metadata::jsonb AS value OFFSET 0) metadata WHERE NOT EXISTS (SELECT 1 FROM public.usage_attribution_snapshots a WHERE a.request_id=u.request_id AND a.record_kind='session') ) AS usage_analytics_facts_v1"#; diff --git a/crates/aether-data/adapters/postgres/src/usage/dashboard_history_tests.rs b/crates/aether-data/adapters/postgres/src/usage/dashboard_history_tests.rs index 5eb66a89d..78b8b7610 100644 --- a/crates/aether-data/adapters/postgres/src/usage/dashboard_history_tests.rs +++ b/crates/aether-data/adapters/postgres/src/usage/dashboard_history_tests.rs @@ -134,6 +134,29 @@ async fn live_dashboard_restores_legacy_history_without_replaying_or_double_coun assert_eq!(advanced.activity_days, restored.activity_days); assert_eq!(advanced.active_days, restored.active_days); + // New daily rollups retain customer charges independently after detail + // expires; older NULL daily charges retain their original legacy cost. + sqlx::query("UPDATE stats_daily SET billing_cost=1.5 WHERE id='recent'") + .execute(&pool).await.unwrap(); + sqlx::query("DELETE FROM usage WHERE request_id='overlap'") + .execute(&pool).await.unwrap(); + let billed_history = repo.query_dashboard_summary(&query).await.unwrap(); + assert_eq!(billed_history.total.billable_amount.as_deref(), Some("123456791.59691357")); + assert_eq!(billed_history.total.request_count, restored.total.request_count); + assert_eq!(billed_history.today, restored.today); + let provider_cost: String = sqlx::query_scalar("SELECT actual_total_cost::text FROM stats_daily WHERE id='recent'") + .fetch_one(&pool).await.unwrap(); + assert_eq!(provider_cost, "0.30000003"); + + // The pre-activation live prefix applies the same composite snapshot + // to its finalized base amount, independently of procurement cost. + sqlx::query("UPDATE usage SET request_metadata=$1 WHERE request_id='before-shared'") + .bind(serde_json::json!({"billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 2, "user_group": 0.75}, "multiplier": 1.5}})) + .execute(&pool).await.unwrap(); + let billed_prefix = repo.query_dashboard_summary(&query).await.unwrap(); + assert_eq!(billed_prefix.total.billable_amount.as_deref(), Some("123456791.65864196")); + assert_eq!(billed_prefix.today.billable_amount.as_deref(), Some("0.93518517")); + // A summary cutoff without legacy daily history must leave the normal // future-only projection and its requested calendar unchanged. sqlx::query("DELETE FROM stats_daily").execute(&pool).await.unwrap(); diff --git a/crates/aether-data/adapters/postgres/src/usage/mod.rs b/crates/aether-data/adapters/postgres/src/usage/mod.rs index 1349778f1..4a2456f23 100644 --- a/crates/aether-data/adapters/postgres/src/usage/mod.rs +++ b/crates/aether-data/adapters/postgres/src/usage/mod.rs @@ -42,18 +42,21 @@ use crate::{ PostgresTransactionRunner, }; use aether_data_contracts::repository::usage::{ - api_key_usage_contribution, model_usage_contribution, provider_api_key_usage_contribution, - sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence, - sanitize_usage_request_metadata, usage_can_recover_terminal_failure, - usage_error_category_for_status_code, usage_lifecycle_update_allowed, ApiKeyUsageDelta, - ModelUsageDelta, PendingUsageCleanupSummary, ProviderApiKeyUsageContribution, - ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, StoredProviderApiKeyUsageSummary, - StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredRequestUsageAudit, - StoredUsageDailySummary, UpsertUsageRecord, UsageAuditListQuery, UsageCounterFlushSummary, - UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, - UsageReadRepository, UsageWriteRepository, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, - PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, - REQUESTED_REASONING_EFFORT_METADATA_KEY, + api_key_usage_contribution, model_usage_contribution, preserve_usage_routing_group_snapshot, + provider_api_key_usage_contribution, sanitize_usage_capture_controls_for_persistence, + sanitize_usage_for_persistence, sanitize_usage_request_metadata, + usage_can_recover_terminal_failure, usage_error_category_for_status_code, + usage_lifecycle_update_allowed, ApiKeyUsageDelta, ModelUsageDelta, PendingUsageCleanupSummary, + ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, + StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary, + StoredProviderUsageSummary, StoredRequestUsageAudit, StoredUsageDailySummary, + UpsertUsageRecord, UsageAuditListQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot, + UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, UsageReadRepository, + UsageWriteRepository, BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, + PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, + PROVIDER_SERVICE_TIER_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, }; use aether_data_contracts::DataLayerError; @@ -8762,6 +8765,15 @@ ORDER BY "usage".user_id ASC ); request_metadata_json = json_bind_text(request_metadata_value.as_ref())?; } + if capture_update_allowed { + request_metadata_value = preserve_usage_routing_group_snapshot( + request_metadata_value, + previous_usage + .as_ref() + .and_then(|stored| stored.request_metadata.as_ref()), + ); + request_metadata_json = json_bind_text(request_metadata_value.as_ref())?; + } let _row = sqlx::query(UPSERT_SQL) .bind(Uuid::new_v4().to_string()) .bind(&usage.request_id) @@ -12554,6 +12566,10 @@ fn retain_previous_request_audit_metadata( "request_path", "request_query_string", "request_path_and_query", + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, + ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, ] { if let Some(value) = previous_metadata.get(key) { retained.insert(key.to_string(), value.clone()); diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/dashboard_history.sql b/crates/aether-data/adapters/postgres/src/usage/queries/dashboard_history.sql index 3eb9a991f..ab49956cb 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/dashboard_history.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/dashboard_history.sql @@ -11,7 +11,7 @@ WITH daily AS ( CASE WHEN effective_input_tokens=0 AND input_tokens>0 THEN input_tokens ELSE effective_input_tokens END + cache_creation_tokens + cache_read_tokens ELSE total_input_context END AS cache_input_tokens, - actual_total_cost::numeric AS billable_amount + COALESCE(billing_cost,actual_total_cost::numeric) AS billable_amount FROM stats_daily ), facts AS MATERIALIZED ( SELECT (day AT TIME ZONE 'UTC')::date AS day, request_count, @@ -22,7 +22,11 @@ WITH daily AS ( SELECT (b.created_at AT TIME ZONE 'UTC')::date, 1::bigint, b.input_tokens, b.output_tokens, b.total_tokens, b.cache_creation_input_tokens, b.cache_read_input_tokens, b.total_input_context, - COALESCE(s.billing_actual_total_cost_usd::numeric,u.actual_total_cost_usd::numeric) + CASE WHEN COALESCE(u.request_metadata::jsonb->'usage_pricing_available','true'::jsonb)<>'false'::jsonb + AND (s.billing_actual_total_cost_usd IS NOT NULL OR COALESCE(s.billing_status,u.billing_status)='settled') + THEN public.usage_customer_billable_amount(u.request_metadata::jsonb, + COALESCE(s.billing_total_cost_usd::numeric,u.total_cost_usd::numeric), + COALESCE(s.billing_actual_total_cost_usd::numeric,u.actual_total_cost_usd::numeric)) END FROM usage_billing_facts b JOIN usage u USING (request_id) LEFT JOIN usage_settlement_snapshots s USING (request_id) diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql b/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql index d6d8a736d..e9346d73b 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/list_recent_usage_audits_prefix.sql @@ -176,6 +176,10 @@ SELECT NULL::bytea AS client_response_body_compressed, CASE WHEN NULLIF(BTRIM("usage".request_metadata->>'client_ip'), '') IS NOT NULL + OR NULLIF(BTRIM("usage".request_metadata->>'routing_group_id'), '') IS NOT NULL + OR NULLIF(BTRIM("usage".request_metadata->>'routing_group_name'), '') IS NOT NULL + OR json_typeof("usage".request_metadata->'routing_group_billing_multiplier') = 'number' + OR "usage".request_metadata->'billing_multiplier_snapshot' IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'user_agent'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'request_path'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'request_path_and_query'), '') IS NOT NULL @@ -192,7 +196,17 @@ SELECT OR ("usage".request_metadata->>'usage_pricing_available') IN ('true', 'false') OR json_typeof("usage".request_metadata->'live_session') = 'object' OR json_typeof("usage".request_metadata->'realtime_session') = 'object' - THEN jsonb_strip_nulls(jsonb_build_object( + THEN (jsonb_strip_nulls(jsonb_build_object( + 'routing_group_id', + NULLIF(BTRIM("usage".request_metadata->>'routing_group_id'), ''), + 'routing_group_name', + NULLIF(BTRIM("usage".request_metadata->>'routing_group_name'), ''), + 'routing_group_billing_multiplier', + CASE + WHEN json_typeof("usage".request_metadata->'routing_group_billing_multiplier') = 'number' + THEN "usage".request_metadata->'routing_group_billing_multiplier' + ELSE NULL + END, 'client_ip', NULLIF(BTRIM("usage".request_metadata->>'client_ip'), ''), 'user_agent', @@ -255,7 +269,11 @@ SELECT THEN "usage".request_metadata->'realtime_session' ELSE NULL END - ))::json + )) || CASE + WHEN "usage".request_metadata->'billing_multiplier_snapshot' IS NOT NULL + THEN jsonb_build_object('billing_multiplier_snapshot', "usage".request_metadata->'billing_multiplier_snapshot') + ELSE '{}'::jsonb + END)::json ELSE NULL::json END AS request_metadata, NULL::varchar AS http_request_body_ref, diff --git a/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql b/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql index d6d8a736d..e9346d73b 100644 --- a/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql +++ b/crates/aether-data/adapters/postgres/src/usage/queries/list_usage_audits_prefix.sql @@ -176,6 +176,10 @@ SELECT NULL::bytea AS client_response_body_compressed, CASE WHEN NULLIF(BTRIM("usage".request_metadata->>'client_ip'), '') IS NOT NULL + OR NULLIF(BTRIM("usage".request_metadata->>'routing_group_id'), '') IS NOT NULL + OR NULLIF(BTRIM("usage".request_metadata->>'routing_group_name'), '') IS NOT NULL + OR json_typeof("usage".request_metadata->'routing_group_billing_multiplier') = 'number' + OR "usage".request_metadata->'billing_multiplier_snapshot' IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'user_agent'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'request_path'), '') IS NOT NULL OR NULLIF(BTRIM("usage".request_metadata->>'request_path_and_query'), '') IS NOT NULL @@ -192,7 +196,17 @@ SELECT OR ("usage".request_metadata->>'usage_pricing_available') IN ('true', 'false') OR json_typeof("usage".request_metadata->'live_session') = 'object' OR json_typeof("usage".request_metadata->'realtime_session') = 'object' - THEN jsonb_strip_nulls(jsonb_build_object( + THEN (jsonb_strip_nulls(jsonb_build_object( + 'routing_group_id', + NULLIF(BTRIM("usage".request_metadata->>'routing_group_id'), ''), + 'routing_group_name', + NULLIF(BTRIM("usage".request_metadata->>'routing_group_name'), ''), + 'routing_group_billing_multiplier', + CASE + WHEN json_typeof("usage".request_metadata->'routing_group_billing_multiplier') = 'number' + THEN "usage".request_metadata->'routing_group_billing_multiplier' + ELSE NULL + END, 'client_ip', NULLIF(BTRIM("usage".request_metadata->>'client_ip'), ''), 'user_agent', @@ -255,7 +269,11 @@ SELECT THEN "usage".request_metadata->'realtime_session' ELSE NULL END - ))::json + )) || CASE + WHEN "usage".request_metadata->'billing_multiplier_snapshot' IS NOT NULL + THEN jsonb_build_object('billing_multiplier_snapshot', "usage".request_metadata->'billing_multiplier_snapshot') + ELSE '{}'::jsonb + END)::json ELSE NULL::json END AS request_metadata, NULL::varchar AS http_request_body_ref, diff --git a/crates/aether-data/contracts/src/repository/auth.rs b/crates/aether-data/contracts/src/repository/auth.rs index 3fedc5c0b..f9fba6d0e 100644 --- a/crates/aether-data/contracts/src/repository/auth.rs +++ b/crates/aether-data/contracts/src/repository/auth.rs @@ -577,6 +577,45 @@ pub struct UpdateUserApiKeyBasicRecord { /// unchanged. Keeping this patch in the basic mutation record lets repositories apply the /// complete user-key update in one atomic write. pub feature_settings: Option>, + /// Self-service updates merge routing selection separately against the + /// current stored settings. `None` retains administrative replacement semantics. + pub routing_group_selection: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct UpdateApiKeyRoutingGroupSelection { + /// `None` preserves the latest stored group; `Some(None)` follows the + /// default; `Some(Some(id))` selects the validated public group. + pub group_id: Option>, +} + +impl UpdateApiKeyRoutingGroupSelection { + /// Repositories must call this while holding the same write lock as the + /// surrounding API key mutation, so unrelated edits cannot restore a stale + /// group choice or a stale feature-settings object. + pub fn merge_feature_settings( + &self, + current: Option<&serde_json::Value>, + replacement: Option>, + ) -> Option { + let group_id = match &self.group_id { + None => current + .and_then(|value| value.get("routing_group_id")) + .and_then(serde_json::Value::as_str) + .map(|id| serde_json::Value::String(id.to_string())), + Some(group_id) => group_id.clone().map(serde_json::Value::String), + }; + let mut settings = replacement + .unwrap_or_else(|| current.cloned()) + .and_then(|value| value.as_object().cloned()) + .unwrap_or_default(); + settings.remove("routing_group_id"); + settings.remove("routing_group_name"); + if let Some(group_id) = group_id { + settings.insert("routing_group_id".to_string(), group_id); + } + (!settings.is_empty()).then_some(serde_json::Value::Object(settings)) + } } impl std::fmt::Debug for UpdateUserApiKeyBasicRecord { diff --git a/crates/aether-data/contracts/src/repository/settlement/types.rs b/crates/aether-data/contracts/src/repository/settlement/types.rs index 4d9c6a0e6..865b8742f 100644 --- a/crates/aether-data/contracts/src/repository/settlement/types.rs +++ b/crates/aether-data/contracts/src/repository/settlement/types.rs @@ -346,6 +346,10 @@ pub struct UsageSettlementInput { pub billing_status: String, pub total_cost_usd: f64, pub actual_total_cost_usd: f64, + /// Customer charge after all captured billing factors, independent of upstream cost. + /// Missing values retain the legacy charge based on `actual_total_cost_usd`. + #[serde(default)] + pub billing_cost_usd: Option, pub finalized_at_unix_secs: Option, } @@ -366,6 +370,14 @@ impl UsageSettlementInput { "settlement cost must be finite".to_string(), )); } + if self + .billing_cost_usd + .is_some_and(|value| !value.is_finite() || value < 0.0) + { + return Err(crate::DataLayerError::InvalidInput( + "settlement billing_cost_usd must be finite and non-negative".to_string(), + )); + } Ok(()) } } @@ -511,14 +523,17 @@ pub fn settlement_billing_status_for_usage_status(status: &str) -> &'static str } pub fn settlement_billable_cost_usd(input: &UsageSettlementInput) -> f64 { - input.actual_total_cost_usd.max(0.0) + input + .billing_cost_usd + .unwrap_or(input.actual_total_cost_usd) + .max(0.0) } #[cfg(test)] mod tests { use super::{ - validate_wallet_settlement_values, ReconcileUsagePolicyCostInput, - ReserveUsagePolicyCostInput, ReserveUsagePolicyRequestInput, + settlement_billable_cost_usd, validate_wallet_settlement_values, + ReconcileUsagePolicyCostInput, ReserveUsagePolicyCostInput, ReserveUsagePolicyRequestInput, UsagePolicyCostReservationState, UsagePolicyCostWindow, UsagePolicyRequestWindow, UsageSettlementInput, }; @@ -535,11 +550,50 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 0.1, actual_total_cost_usd: 0.1, + billing_cost_usd: None, finalized_at_unix_secs: None, }; assert!(input.validate().is_err()); } + #[test] + fn explicit_customer_charge_is_validated_independently_of_upstream_cost() { + let mut input: UsageSettlementInput = serde_json::from_value(serde_json::json!({ + "request_id": "billing-charge", + "user_id": "user-1", + "api_key_id": null, + "provider_id": "provider-1", + "status": "completed", + "billing_status": "pending", + "total_cost_usd": 2.0, + "actual_total_cost_usd": 0.5, + "finalized_at_unix_secs": null, + })) + .expect("legacy settlement input should deserialize"); + assert_eq!(input.billing_cost_usd, None); + assert_eq!(settlement_billable_cost_usd(&input), 0.5); + assert!(input.validate().is_ok()); + + for charge in [3.0, 0.0] { + input.billing_cost_usd = Some(charge); + assert!(input.validate().is_ok()); + assert_eq!(settlement_billable_cost_usd(&input), charge); + assert_eq!(input.actual_total_cost_usd, 0.5); + assert_eq!( + serde_json::from_value::( + serde_json::to_value(&input).unwrap() + ) + .unwrap(), + input + ); + } + + for charge in [-0.01, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + input.billing_cost_usd = Some(charge); + assert!(input.validate().is_err()); + } + } + #[test] fn wallet_settlement_values_reject_corruption_and_overflow() { assert!(validate_wallet_settlement_values(-3.0, 0.0, 12.0, 1.0).is_ok()); diff --git a/crates/aether-data/contracts/src/repository/usage/billing_multiplier.rs b/crates/aether-data/contracts/src/repository/usage/billing_multiplier.rs new file mode 100644 index 000000000..0d8d89f26 --- /dev/null +++ b/crates/aether-data/contracts/src/repository/usage/billing_multiplier.rs @@ -0,0 +1,169 @@ +use std::collections::BTreeMap; + +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +use crate::DataLayerError; + +use super::ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY; + +pub const BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY: &str = "billing_multiplier_snapshot"; + +/// Immutable customer pricing factors. Provider Key rates belong to upstream cost, +/// not this snapshot. Add future factors (for example `user_group`) at admission. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct BillingMultiplierSnapshot { + version: u32, + factors: BTreeMap, + multiplier: f64, +} + +impl BillingMultiplierSnapshot { + pub fn from_factors(factors: BTreeMap) -> Result { + if factors.len() > 16 + || factors.iter().any(|(name, value)| { + name.is_empty() + || name.len() > 64 + || !name + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'_') + || !value.is_finite() + || *value < 0.0 + }) + { + return Err(invalid_snapshot()); + } + let multiplier = if factors.values().any(|value| *value == 0.0) { + 0.0 + } else { + factors.values().product::() + }; + if !multiplier.is_finite() { + return Err(invalid_snapshot()); + } + Ok(Self { + version: 1, + factors, + multiplier, + }) + } + + pub fn validate(&self) -> Result<(), DataLayerError> { + let expected = Self::from_factors(self.factors.clone())?; + if self.version != 1 || self.multiplier != expected.multiplier { + return Err(invalid_snapshot()); + } + Ok(()) + } + + pub fn multiplier(&self) -> f64 { + self.multiplier + } + + pub fn cost(&self, base_cost: f64) -> Result { + self.validate()?; + let cost = base_cost * self.multiplier; + if !base_cost.is_finite() || base_cost < 0.0 || !cost.is_finite() { + return Err(DataLayerError::InvalidInput( + "customer billing cost must be finite and non-negative".to_string(), + )); + } + // Match wallet storage and usage-policy cost units (eight decimals). + // Scaling a finite large amount must not introduce infinity by itself. + let scaled = cost * 100_000_000.0; + Ok(if scaled.is_finite() { + scaled.round() / 100_000_000.0 + } else { + cost + }) + } +} + +fn invalid_snapshot() -> DataLayerError { + DataLayerError::InvalidInput("invalid billing multiplier snapshot".to_string()) +} + +/// None preserves legacy charging. A malformed captured snapshot is an error, +/// never an instruction to silently charge a different rate. +pub fn billing_multiplier_snapshot( + metadata: Option<&Value>, +) -> Result, DataLayerError> { + let Some(metadata) = metadata.and_then(Value::as_object) else { + return Ok(None); + }; + if let Some(value) = metadata.get(BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY) { + let snapshot: BillingMultiplierSnapshot = + serde_json::from_value(value.clone()).map_err(|_| invalid_snapshot())?; + snapshot.validate()?; + return Ok(Some(snapshot)); + } + if let Some(value) = metadata.get(ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY) { + let multiplier = value.as_f64().ok_or_else(invalid_snapshot)?; + return BillingMultiplierSnapshot::from_factors(BTreeMap::from([( + "routing_group".to_string(), + multiplier, + )])) + .map(Some); + } + Ok(None) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn composes_customer_factors_without_provider_cost_and_freezes_them() { + let snapshot = BillingMultiplierSnapshot::from_factors(BTreeMap::from([ + ("routing_group".to_string(), 2.0), + ("user_group".to_string(), 0.75), + ])) + .unwrap(); + assert_eq!(snapshot.multiplier(), 1.5); + assert_eq!(snapshot.cost(10.0).unwrap(), 15.0); + assert_eq!(snapshot.cost(0.123456789).unwrap(), 0.18518518); + let metadata = json!({"billing_multiplier_snapshot": snapshot, "routing_group_billing_multiplier": 99, "rate_multiplier": 0.1}); + assert_eq!( + billing_multiplier_snapshot(Some(&metadata)).unwrap(), + Some(snapshot) + ); + assert_eq!(billing_multiplier_snapshot(None).unwrap(), None); + } + + #[test] + fn rejects_corrupt_overflowing_snapshots_and_accepts_zero_rates() { + for factors in [ + BTreeMap::from([("routing_group".into(), -1.0)]), + BTreeMap::from([("routing_group".into(), f64::INFINITY)]), + BTreeMap::from([ + ("routing_group".into(), f64::MAX), + ("user_group".into(), 2.0), + ]), + ] { + assert!(BillingMultiplierSnapshot::from_factors(factors).is_err()); + } + let zero = BillingMultiplierSnapshot::from_factors(BTreeMap::from([ + ("routing_group".into(), 0.0), + ("user_group".into(), 2.0), + ])) + .unwrap(); + assert_eq!(zero.cost(10.0).unwrap(), 0.0); + for invalid in [ + Value::Null, + json!({"version": 2, "factors": {}, "multiplier": 1}), + json!({"version": 1, "factors": {"routing_group": 2}, "multiplier": 1}), + ] { + assert!(billing_multiplier_snapshot(Some( + &json!({"billing_multiplier_snapshot": invalid}) + )) + .is_err()); + } + let doubled = BillingMultiplierSnapshot::from_factors(BTreeMap::from([( + "routing_group".into(), + 2.0, + )])) + .unwrap(); + assert!(doubled.cost(f64::MAX).is_err()); + } +} diff --git a/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs b/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs index f69b16e14..41f4cce75 100644 --- a/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs +++ b/crates/aether-data/contracts/src/repository/usage/metadata_policy.rs @@ -8,14 +8,17 @@ use serde_json::{Map, Value}; use crate::repository::candidates::sanitize_request_candidate_skip_reason; use super::{ - normalize_provider_response_model, LIVE_SESSION_METADATA_KEY, + billing_multiplier_snapshot, normalize_provider_response_model, + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, LIVE_SESSION_METADATA_KEY, PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_RESPONSE_MODEL_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, REALTIME_SESSION_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, - USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY, - WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY, + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY, + USAGE_PRICING_AVAILABLE_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY, + WEBSOCKET_TRANSPORT_METADATA_KEY, }; const UPSTREAM_IS_STREAM_KEY: &str = "upstream_is_stream"; @@ -43,8 +46,68 @@ pub fn sanitize_usage_request_metadata_ref(value: Option<&Value>) -> Option, + previous: Option<&Value>, +) -> Option { + let Some(previous) = previous.and_then(Value::as_object) else { + return incoming; + }; + let mut snapshot = Map::from_iter( + [ + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, + ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, + PLAN_USAGE_RESERVATION_TOKEN_KEY, + ] + .into_iter() + .filter_map(|key| { + previous + .get(key) + .map(|value| (key.to_string(), value.clone())) + }), + ); + if !snapshot.contains_key(BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY) + && snapshot.contains_key(ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY) + { + let captured = billing_multiplier_snapshot(Some(&Value::Object(snapshot.clone()))) + .ok() + .flatten() + .and_then(|snapshot| serde_json::to_value(snapshot).ok()) + .unwrap_or(Value::Null); + snapshot.insert( + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY.to_string(), + captured, + ); + } + let Some(Value::Object(snapshot)) = sanitize_usage_request_metadata_object(&snapshot) else { + return incoming; + }; + let mut metadata = incoming + .and_then(|value| value.as_object().cloned()) + .unwrap_or_default(); + metadata.extend(snapshot); + Some(Value::Object(metadata)) +} + pub fn sanitize_usage_request_metadata_object(source: &Map) -> Option { let mut target = Map::new(); + if let Some(snapshot) = source.get(BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY) { + let metadata = serde_json::json!({BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY: snapshot}); + let snapshot = match billing_multiplier_snapshot(Some(&metadata)) { + Ok(Some(snapshot)) => serde_json::to_value(snapshot) + .expect("validated billing multiplier snapshot must serialize"), + // Preserve an invalid marker so malformed financial input cannot silently + // fall back to legacy billing after metadata projection. + _ => Value::Null, + }; + target.insert( + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY.to_string(), + snapshot, + ); + } if let Some(source) = source .get("analytics_measurement") .and_then(|value| value.get("source")) @@ -81,6 +144,8 @@ pub fn sanitize_usage_request_metadata_object(source: &Map) -> Op } insert_token(source, &mut target, "trace_id", 128); + insert_token(source, &mut target, ROUTING_GROUP_ID_METADATA_KEY, 128); + insert_bounded_text(source, &mut target, ROUTING_GROUP_NAME_METADATA_KEY, 256); insert_ip_address(source, &mut target, "client_ip"); insert_client_family(source, &mut target); for key in [ @@ -181,6 +246,7 @@ pub fn sanitize_usage_request_metadata_object(source: &Map) -> Op for key in [ "rate_multiplier", + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, "input_price_per_1m", "output_price_per_1m", "cache_creation_price_per_1m", @@ -189,6 +255,20 @@ pub fn sanitize_usage_request_metadata_object(source: &Map) -> Op ] { insert_nonnegative_number(source, &mut target, key); } + // An invalid legacy routing factor must remain a financial tombstone. Dropping it + // would make a subsequent reader silently fall back to the historical provider charge. + if source + .get(ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY) + .is_some_and(|value| { + !value + .as_f64() + .is_some_and(|value| value.is_finite() && value >= 0.0) + }) + { + target + .entry(BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY.to_string()) + .or_insert(Value::Null); + } let billing_snapshot = source .get("billing_snapshot") @@ -1156,6 +1236,27 @@ fn insert_token( target.insert(key.to_string(), Value::String(value.to_string())); } +fn insert_bounded_text( + source: &Map, + target: &mut Map, + key: &str, + max_len: usize, +) { + let Some(value) = source + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| { + !value.is_empty() + && value.chars().count() <= max_len + && !value.chars().any(char::is_control) + }) + else { + return; + }; + target.insert(key.to_string(), Value::String(value.to_string())); +} + fn insert_dimension_token(source: &Map, target: &mut Map, key: &str) { let Some(value) = source .get(key) @@ -1263,7 +1364,111 @@ fn safe_version_value(value: &Value) -> Option { mod tests { use serde_json::json; - use super::{sanitize_usage_request_metadata, sanitize_usage_request_metadata_ref}; + use super::{ + billing_multiplier_snapshot, preserve_usage_routing_group_snapshot, + sanitize_usage_request_metadata, sanitize_usage_request_metadata_ref, + }; + + #[test] + fn billing_multiplier_snapshot_projection_preserves_invalid_marker_and_immutable_factors() { + for snapshot in [ + serde_json::Value::Null, + json!({"version": 1, "factors": {"routing_group": 2.0}, "multiplier": 1.0}), + json!({"version": 99, "factors": {}, "multiplier": 1.0}), + ] { + let projected = sanitize_usage_request_metadata(Some(json!({ + "billing_multiplier_snapshot": snapshot, + "routing_group_billing_multiplier": 0.5, + }))) + .unwrap(); + assert_eq!( + projected.get("billing_multiplier_snapshot"), + Some(&serde_json::Value::Null) + ); + assert!(billing_multiplier_snapshot(Some(&projected)).is_err()); + let preserved = preserve_usage_routing_group_snapshot( + Some(json!({"billing_multiplier_snapshot": {"version": 1, "factors": {}, "multiplier": 1.0}})), + Some(&projected), + ).unwrap(); + assert!(billing_multiplier_snapshot(Some(&preserved)).is_err()); + } + let legacy = + json!({"routing_group_billing_multiplier": 0.25, "routing_group_name": "历史分组"}); + let preserved = preserve_usage_routing_group_snapshot( + Some(json!({"billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 99.0}, "multiplier": 99.0}})), + Some(&legacy), + ).unwrap(); + assert_eq!( + billing_multiplier_snapshot(Some(&preserved)) + .unwrap() + .unwrap() + .multiplier(), + 0.25 + ); + assert_eq!(preserved["routing_group_name"], "历史分组"); + } + + #[test] + fn billing_multiplier_snapshot_projection_rejects_malformed_legacy_factors() { + for factor in [serde_json::Value::Null, json!(-1), json!("2"), json!({})] { + let projected = sanitize_usage_request_metadata(Some(json!({ + "routing_group_billing_multiplier": factor, + }))) + .expect("invalid financial input must retain a tombstone"); + assert_eq!( + projected["billing_multiplier_snapshot"], + serde_json::Value::Null + ); + assert!(billing_multiplier_snapshot(Some(&projected)).is_err()); + assert_eq!( + sanitize_usage_request_metadata(Some(projected.clone())), + Some(projected) + ); + } + let generic = sanitize_usage_request_metadata(Some(json!({ + "routing_group_billing_multiplier": -1, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 2}, "multiplier": 2}, + }))) + .unwrap(); + assert_eq!( + billing_multiplier_snapshot(Some(&generic)) + .unwrap() + .unwrap() + .multiplier(), + 2.0 + ); + } + + #[test] + fn billing_multiplier_snapshot_preserves_the_original_reservation_owner() { + let token_a = "550e8400-e29b-41d4-a716-446655440001"; + let token_b = "550e8400-e29b-41d4-a716-446655440002"; + let incoming = json!({ + "plan_usage_reservation_token": token_b, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 3}, "multiplier": 3}, + }); + let previous = json!({ + "plan_usage_reservation_token": token_a, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 0.5}, "multiplier": 0.5}, + }); + let preserved = + preserve_usage_routing_group_snapshot(Some(incoming.clone()), Some(&previous)).unwrap(); + assert_eq!(preserved["plan_usage_reservation_token"], token_a); + assert_eq!( + billing_multiplier_snapshot(Some(&preserved)) + .unwrap() + .unwrap() + .multiplier(), + 0.5 + ); + + for empty in [json!({}), json!({"plan_usage_reservation_token": " "})] { + let preserved = + preserve_usage_routing_group_snapshot(Some(incoming.clone()), Some(&empty)) + .unwrap(); + assert_eq!(preserved["plan_usage_reservation_token"], token_b); + } + } #[test] fn account_attribution_preserves_key_flag_without_custom_identity_or_purpose() { diff --git a/crates/aether-data/contracts/src/repository/usage/mod.rs b/crates/aether-data/contracts/src/repository/usage/mod.rs index b34dcba94..014169010 100644 --- a/crates/aether-data/contracts/src/repository/usage/mod.rs +++ b/crates/aether-data/contracts/src/repository/usage/mod.rs @@ -2,6 +2,7 @@ mod analytics; #[cfg(test)] mod analytics_tests; mod attribution; +mod billing_multiplier; mod capture_memory; mod compression; mod dashboard_summary; @@ -12,6 +13,7 @@ mod types; pub use analytics::*; pub use attribution::*; +pub use billing_multiplier::*; #[doc(hidden)] pub use capture_memory::{ mark_usage_capture_memory_omitted, usage_json_heap_estimate, UsageCaptureMemoryBudget, @@ -60,6 +62,8 @@ pub use types::{ PROVIDER_RESPONSE_MODEL_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, REALTIME_SESSION_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, - USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY, - WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY, + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY, + USAGE_PRICING_AVAILABLE_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY, + WEBSOCKET_TRANSPORT_METADATA_KEY, }; diff --git a/crates/aether-data/contracts/src/repository/usage/types.rs b/crates/aether-data/contracts/src/repository/usage/types.rs index 60cd0ed3c..04f4307c9 100644 --- a/crates/aether-data/contracts/src/repository/usage/types.rs +++ b/crates/aether-data/contracts/src/repository/usage/types.rs @@ -11,6 +11,10 @@ pub const PROVIDER_RESPONSE_MODEL_METADATA_KEY: &str = "provider_response_model" pub const PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY: &str = "provider_cache_ttl_minutes"; pub const ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY: &str = "routing_candidate_skip_reason"; pub const ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY: &str = "routing_failure_diagnostic"; +/// Immutable routing-group multiplier captured when the request is planned. +pub const ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY: &str = "routing_group_billing_multiplier"; +pub const ROUTING_GROUP_ID_METADATA_KEY: &str = "routing_group_id"; +pub const ROUTING_GROUP_NAME_METADATA_KEY: &str = "routing_group_name"; pub const WEBSOCKET_MODE_METADATA_KEY: &str = "websocket_mode"; pub const WEBSOCKET_TRANSPORT_METADATA_KEY: &str = "websocket_transport"; pub const PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY: &str = "plan_usage_reservation_deferred"; @@ -828,6 +832,47 @@ impl StoredRequestUsageAudit { self.request_metadata_number("rate_multiplier") } + /// Historical requests without a captured multiplier retain the original 1x rate. + pub fn routing_group_billing_multiplier(&self) -> f64 { + self.request_metadata_number(ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY) + .filter(|value| value.is_finite() && *value >= 0.0) + .unwrap_or(1.0) + } + + /// Routing factor projection retained for callers inspecting this individual factor. + pub fn routing_group_billing_cost(&self) -> Option { + let cost = self.total_cost_usd * self.routing_group_billing_multiplier(); + cost.is_finite().then_some(cost) + } + + pub fn billing_multiplier(&self) -> f64 { + super::billing_multiplier_snapshot(self.request_metadata.as_ref()) + .ok() + .flatten() + .map(|snapshot| snapshot.multiplier()) + .unwrap_or(1.0) + } + + /// Customer charge is independent of upstream Key cost. Legacy rows keep their + /// original charge; no current configuration is consulted for historical usage. + pub fn billing_cost(&self) -> Option { + match super::billing_multiplier_snapshot(self.request_metadata.as_ref()).ok()? { + Some(snapshot) => snapshot.cost(self.total_cost_usd).ok(), + None => self + .actual_total_cost_usd + .is_finite() + .then_some(self.actual_total_cost_usd.max(0.0)), + } + } + + pub fn routing_group_id(&self) -> Option<&str> { + self.request_metadata_string(ROUTING_GROUP_ID_METADATA_KEY) + } + + pub fn routing_group_name(&self) -> Option<&str> { + self.request_metadata_string(ROUTING_GROUP_NAME_METADATA_KEY) + } + pub fn settlement_is_free_tier(&self) -> Option { self.request_metadata_bool("is_free_tier") } @@ -3407,6 +3452,40 @@ mod tests { assert!(record.validate().is_err()); } + #[test] + fn routing_group_snapshot_defaults_legacy_multiplier_without_inventing_a_group() { + let mut usage = sample_usage(); + usage.total_cost_usd = 4.0; + assert_eq!(usage.routing_group_billing_multiplier(), 1.0); + assert_eq!(usage.routing_group_billing_cost(), Some(4.0)); + assert_eq!(usage.routing_group_id(), None); + assert_eq!(usage.routing_group_name(), None); + for (value, multiplier, cost) in [ + (json!(0), 0.0, 0.0), + (json!(0.25), 0.25, 1.0), + (json!(2.5), 2.5, 10.0), + (json!(-2), 1.0, 4.0), + (json!("Infinity"), 1.0, 4.0), + (json!(f64::INFINITY), 1.0, 4.0), + (json!(f64::NAN), 1.0, 4.0), + ] { + usage.request_metadata = Some(json!({ + "routing_group_billing_multiplier": value, + "routing_group_id": "group-recorded", + "routing_group_name": "请求时的分组", + "rate_multiplier": 0.75 + })); + assert_eq!(usage.routing_group_billing_multiplier(), multiplier); + assert_eq!(usage.routing_group_billing_cost(), Some(cost)); + assert_eq!(usage.routing_group_id(), Some("group-recorded")); + assert_eq!(usage.routing_group_name(), Some("请求时的分组")); + assert_eq!(usage.settlement_rate_multiplier(), Some(0.75)); + } + usage.request_metadata = Some(json!({"routing_group_billing_multiplier": 2.0})); + usage.total_cost_usd = f64::MAX; + assert_eq!(usage.routing_group_billing_cost(), None); + } + #[test] fn settlement_accessors_prefer_typed_metadata() { let mut usage = sample_usage(); diff --git a/crates/aether-data/contracts/src/repository/wallet/types.rs b/crates/aether-data/contracts/src/repository/wallet/types.rs index ce0d28d67..df7aaac25 100644 --- a/crates/aether-data/contracts/src/repository/wallet/types.rs +++ b/crates/aether-data/contracts/src/repository/wallet/types.rs @@ -3286,6 +3286,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 0.1, actual_total_cost_usd: 0.1, + billing_cost_usd: None, finalized_at_unix_secs: None, }; assert!(input.validate().is_err()); diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql index 0db6c2360..ea8c52c1e 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/001_types_and_tables.sql @@ -901,6 +901,7 @@ CREATE TABLE IF NOT EXISTS public.stats_daily ( cache_read_tokens bigint DEFAULT '0'::bigint NOT NULL, total_cost numeric(20,8) DEFAULT '0'::double precision NOT NULL, actual_total_cost numeric(20,8) DEFAULT '0'::double precision NOT NULL, + billing_cost numeric(20,8), input_cost numeric(20,8) DEFAULT '0'::double precision NOT NULL, output_cost numeric(20,8) DEFAULT '0'::double precision NOT NULL, cache_creation_cost numeric(20,8) DEFAULT '0'::double precision NOT NULL, diff --git a/crates/aether-data/runtime/schema/bootstrap/postgres/190_overview_analytics.sql b/crates/aether-data/runtime/schema/bootstrap/postgres/190_overview_analytics.sql index 7db83dce9..594e83313 100644 --- a/crates/aether-data/runtime/schema/bootstrap/postgres/190_overview_analytics.sql +++ b/crates/aether-data/runtime/schema/bootstrap/postgres/190_overview_analytics.sql @@ -201,6 +201,77 @@ DROP TRIGGER IF EXISTS overview_usage_delete_attribution ON public.usage; CREATE TRIGGER overview_usage_delete_attribution BEFORE DELETE ON public.usage FOR EACH ROW EXECUTE FUNCTION public.overview_delete_attribution(); +-- Customer charges use the immutable request-time factor snapshot. Provider +-- procurement cost remains in actual_total_cost_usd for legacy reporting. +CREATE OR REPLACE FUNCTION public.usage_customer_billable_amount( + metadata jsonb, base_cost numeric, legacy_cost numeric +) RETURNS numeric LANGUAGE plpgsql IMMUTABLE PARALLEL SAFE AS $$ +DECLARE factor jsonb; multiplier numeric; amount numeric; + factor_name text; factor_value jsonb; factor_number double precision; + expected_multiplier double precision := 1.0; factor_count integer := 0; + has_zero boolean := false; +BEGIN + IF metadata ? 'billing_multiplier_snapshot' THEN + IF jsonb_typeof(metadata->'billing_multiplier_snapshot') <> 'object' + OR metadata #> '{billing_multiplier_snapshot,version}' IS DISTINCT FROM '1'::jsonb + OR jsonb_typeof(metadata #> '{billing_multiplier_snapshot,factors}') IS DISTINCT FROM 'object' + THEN RETURN NULL; END IF; + factor := metadata #> '{billing_multiplier_snapshot,multiplier}'; + FOR factor_name, factor_value IN + SELECT key, value FROM jsonb_each(metadata #> '{billing_multiplier_snapshot,factors}') ORDER BY key COLLATE "C" + LOOP + factor_count := factor_count + 1; + IF factor_count > 16 OR factor_name = '' OR length(factor_name) > 64 + OR factor_name !~ '^[A-Za-z0-9_]+$' + OR jsonb_typeof(factor_value) IS DISTINCT FROM 'number' + THEN RETURN NULL; END IF; + factor_number := factor_value::text::double precision; + IF factor_number < 0 OR factor_number > 1.7976931348623157e308::double precision + THEN RETURN NULL; END IF; + has_zero := has_zero OR factor_number = 0; + END LOOP; + -- Rust short-circuits zero before multiplying any of the other factors. + IF has_zero THEN expected_multiplier := 0; + ELSE + FOR factor_name, factor_value IN + SELECT key, value FROM jsonb_each(metadata #> '{billing_multiplier_snapshot,factors}') ORDER BY key COLLATE "C" + LOOP + factor_number := factor_value::text::double precision; + BEGIN + expected_multiplier := expected_multiplier * factor_number; + EXCEPTION WHEN numeric_value_out_of_range THEN + -- PostgreSQL raises on float underflow; Rust rounds that product to 0. + IF expected_multiplier::numeric * factor_number::numeric > 1.7976931348623157e308::numeric + THEN RETURN NULL; END IF; + expected_multiplier := 0; + END; + END LOOP; + END IF; + ELSIF metadata ? 'routing_group_billing_multiplier' THEN + factor := metadata->'routing_group_billing_multiplier'; + expected_multiplier := NULL; + ELSE + RETURN CASE WHEN legacy_cost NOT IN ('NaN'::numeric,'Infinity'::numeric,'-Infinity'::numeric) + THEN round(legacy_cost,8) END; + END IF; + IF jsonb_typeof(factor) IS DISTINCT FROM 'number' THEN RETURN NULL; END IF; + multiplier := factor::text::numeric; + factor_number := factor::text::double precision; + IF factor_number < 0 + OR factor_number > 1.7976931348623157e308::double precision + OR (expected_multiplier IS NOT NULL AND factor_number <> expected_multiplier) + OR multiplier < 0 OR multiplier > 1.7976931348623157e308::numeric + OR base_cost IS NULL OR base_cost < 0 + OR base_cost IN ('NaN'::numeric,'Infinity'::numeric,'-Infinity'::numeric) + THEN RETURN NULL; END IF; + amount := base_cost * multiplier; + IF amount > 1.7976931348623157e308::numeric THEN RETURN NULL; END IF; + RETURN round(amount,8); +EXCEPTION WHEN numeric_value_out_of_range OR invalid_text_representation THEN + -- Corrupt captured pricing must not abort an entire analytics query. + RETURN NULL; +END $$; + CREATE OR REPLACE VIEW public.usage_analytics_facts_v1 AS SELECT u.request_id, COALESCE(u.id, u.request_id) AS id, u.created_at, CASE WHEN identity.owner_id IS NOT NULL AND identity.is_standalone=false THEN identity.owner_id END AS actor_user_id, @@ -233,7 +304,9 @@ SELECT u.request_id, COALESCE(u.id, u.request_id) AS id, u.created_at, THEN round(COALESCE(s.billing_total_cost_usd::numeric, u.total_cost_usd::numeric), 8) END AS rated_amount, CASE WHEN COALESCE(metadata.value->'usage_pricing_available', 'true'::jsonb) <> 'false'::jsonb AND (s.billing_actual_total_cost_usd IS NOT NULL OR COALESCE(s.billing_status, u.billing_status) = 'settled') - THEN round(COALESCE(s.billing_actual_total_cost_usd::numeric, u.actual_total_cost_usd::numeric), 8) END AS billable_amount, + THEN public.usage_customer_billable_amount(metadata.value, + COALESCE(s.billing_total_cost_usd::numeric, u.total_cost_usd::numeric), + COALESCE(s.billing_actual_total_cost_usd::numeric, u.actual_total_cost_usd::numeric)) END AS billable_amount, s.quota_covered_amount_usd AS quota_covered_amount, s.wallet_consumed_amount_usd AS wallet_consumed_amount, s.wallet_debit_amount_usd AS wallet_debit_amount, diff --git a/crates/aether-data/runtime/schema/generated/postgres/baseline/007_stats.sql b/crates/aether-data/runtime/schema/generated/postgres/baseline/007_stats.sql index 2319906e5..5f2cd95fd 100644 --- a/crates/aether-data/runtime/schema/generated/postgres/baseline/007_stats.sql +++ b/crates/aether-data/runtime/schema/generated/postgres/baseline/007_stats.sql @@ -172,6 +172,7 @@ CREATE TABLE IF NOT EXISTS public.stats_daily ( cache_read_tokens bigint DEFAULT 0 NOT NULL, total_cost double precision DEFAULT 0 NOT NULL, actual_total_cost double precision DEFAULT 0 NOT NULL, + billing_cost numeric(20,8), input_cost double precision DEFAULT 0 NOT NULL, output_cost double precision DEFAULT 0 NOT NULL, cache_creation_cost double precision DEFAULT 0 NOT NULL, diff --git a/crates/aether-data/runtime/schema/logical/007_stats.toml b/crates/aether-data/runtime/schema/logical/007_stats.toml index 70b7ccfe6..dbc2ce36d 100644 --- a/crates/aether-data/runtime/schema/logical/007_stats.toml +++ b/crates/aether-data/runtime/schema/logical/007_stats.toml @@ -499,6 +499,12 @@ name = "actual_total_cost" type = "float64" default = 0 +[[table.stats_daily.columns]] +name = "billing_cost" +type = "decimal_money" +nullable = true +driver.postgres.type = "numeric(20,8)" + [[table.stats_daily.columns]] name = "input_cost" type = "float64" diff --git a/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs b/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs index 479c20582..5601db12b 100644 --- a/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs +++ b/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs @@ -159,6 +159,12 @@ async fn perform_stats_aggregation_for_day( .execute(&mut *tx) .await?; + sqlx::query(UPDATE_STATS_DAILY_BILLING_COST_SQL) + .bind(day_start_utc) + .bind(day_end_utc) + .execute(&mut *tx) + .await?; + let model_rows = upsert_stats_daily_model_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?; let provider_rows = diff --git a/crates/aether-data/runtime/src/backend/stats/postgres_daily/sql.rs b/crates/aether-data/runtime/src/backend/stats/postgres_daily/sql.rs index cf233480b..a7a158165 100644 --- a/crates/aether-data/runtime/src/backend/stats/postgres_daily/sql.rs +++ b/crates/aether-data/runtime/src/backend/stats/postgres_daily/sql.rs @@ -2404,6 +2404,18 @@ WHERE created_at >= $1 AND provider_name NOT IN ('unknown', 'pending') "#; +// Keep customer consumption separate from the upstream procurement-cost rollup. +pub(super) const UPDATE_STATS_DAILY_BILLING_COST_SQL: &str = r#" +UPDATE stats_daily SET billing_cost=( + SELECT round(COALESCE(sum(billable_amount),0),8) + FROM usage_analytics_facts_v1 + WHERE created_at >= $1 AND created_at < $2 + AND status NOT IN ('pending','streaming') + AND provider_name NOT IN ('unknown','pending') +) +WHERE date=$1 +"#; + #[cfg(test)] mod tests { use super::{ diff --git a/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs b/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs index 380ec6481..a59054565 100644 --- a/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs +++ b/crates/aether-data/runtime/src/lifecycle/migrate/tests.rs @@ -27,6 +27,7 @@ use crate::lifecycle::bootstrap::postgres::{ EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION, EMPTY_DATABASE_SNAPSHOT_SQL, }; +mod customer_billing_upgrade; mod dashboard_user_anonymization; mod legacy_overview_upgrade; mod migration_deadlines; @@ -1597,6 +1598,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() { 20260923000000, 20261001000000, 20261004000000, + 20261007000000, ] ); } diff --git a/crates/aether-data/runtime/src/lifecycle/migrate/tests/customer_billing_upgrade.rs b/crates/aether-data/runtime/src/lifecycle/migrate/tests/customer_billing_upgrade.rs new file mode 100644 index 000000000..7b29f53a2 --- /dev/null +++ b/crates/aether-data/runtime/src/lifecycle/migrate/tests/customer_billing_upgrade.rs @@ -0,0 +1,114 @@ +use super::*; + +const BILLING_VERSION: i64 = 20261007000000; + +#[tokio::test] +async fn customer_billing_upgrade_preserves_history_and_aggregates_new_days() { + let Some(server) = ManagedPostgresServer::try_start().await.unwrap() else { + return; + }; + let mut connection = PgConnection::connect(server.database_url()).await.unwrap(); + connection.ensure_migrations_table().await.unwrap(); + for migration in POSTGRES_MIGRATOR + .iter() + .filter(|migration| migration.version < BILLING_VERSION) + { + connection.apply(migration).await.unwrap(); + } + let pool = PgPool::connect(server.database_url()).await.unwrap(); + sqlx::raw_sql( + r#" +INSERT INTO stats_daily(id,date,total_requests,actual_total_cost,is_complete) +VALUES ('history','2026-07-17 00:00:00+00',1,0.5,true); +INSERT INTO usage(id,request_id,model,provider_name,status,billing_status, + total_cost_usd,actual_total_cost_usd,created_at,request_metadata) +VALUES ('history','history','m','p','completed','settled',2,0.5, + '2026-07-17 12:00:00+00','{"routing_group_billing_multiplier":2}'); +"#, + ) + .execute(&pool) + .await + .unwrap(); + let history_before: serde_json::Value = + query_scalar("SELECT to_jsonb(d) FROM stats_daily d WHERE id='history'") + .fetch_one(&pool) + .await + .unwrap(); + let usage_before: serde_json::Value = + query_scalar("SELECT to_jsonb(u) FROM usage u WHERE request_id='history'") + .fetch_one(&pool) + .await + .unwrap(); + let migration = POSTGRES_MIGRATOR + .iter() + .find(|migration| migration.version == BILLING_VERSION) + .unwrap(); + connection.apply(migration).await.unwrap(); + + // Even retained requests with captured factors must not rewrite old daily totals. + let history_after: serde_json::Value = + query_scalar("SELECT to_jsonb(d) - 'billing_cost' FROM stats_daily d WHERE id='history'") + .fetch_one(&pool) + .await + .unwrap(); + assert_eq!(history_after, history_before); + let usage_after: serde_json::Value = + query_scalar("SELECT to_jsonb(u) FROM usage u WHERE request_id='history'") + .fetch_one(&pool) + .await + .unwrap(); + assert_eq!(usage_after, usage_before); + let legacy_cost: (Option, String) = sqlx::query_as( + "SELECT billing_cost::text, COALESCE(billing_cost,actual_total_cost::numeric)::text FROM stats_daily WHERE id='history'", + ) + .fetch_one(&pool) + .await + .unwrap(); + assert_eq!(legacy_cost, (None, "0.50000000".to_string())); + + sqlx::raw_sql( + r#" +INSERT INTO usage(id,request_id,model,provider_name,status,billing_status, + total_cost_usd,actual_total_cost_usd,created_at,request_metadata) +VALUES ('new-billed','new-billed','m','p','completed','settled',4,1, + '2026-07-18 12:00:00+00', + '{"billing_multiplier_snapshot":{"version":1,"factors":{"routing_group":2,"user_group":0.75},"multiplier":1.5}}'), + ('new-legacy','new-legacy','m','p','completed','settled',2,0.5, + '2026-07-18 13:00:00+00','{}'); +"#, + ) + .execute(&pool) + .await + .unwrap(); + let new_day = historical_stats_day() + chrono::Duration::days(1); + let backend = postgres_backend(server.database_url()); + let summary = backend + .aggregate_stats_daily(&crate::StatsDailyAggregationInput { + target_day_utc: new_day, + aggregated_at: new_day + chrono::Duration::days(1), + }) + .await + .unwrap() + .unwrap(); + assert_eq!(summary.day_start_utc, new_day); + assert_eq!(summary.total_requests, 2); + let new_costs: (String, String) = sqlx::query_as( + "SELECT billing_cost::text, actual_total_cost::text FROM stats_daily WHERE date=$1", + ) + .bind(new_day) + .fetch_one(&pool) + .await + .unwrap(); + assert_eq!( + new_costs, + ("6.50000000".to_string(), "1.50000000".to_string()) + ); + assert!(query_scalar::<_, bool>( + "SELECT billing_cost IS NULL FROM stats_daily WHERE id='history'", + ) + .fetch_one(&pool) + .await + .unwrap()); + backend.pool().close().await; + pool.close().await; +} diff --git a/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs b/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs index 843c20b85..86da747be 100644 --- a/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs +++ b/crates/aether-data/runtime/src/lifecycle/migrate/tests/legacy_overview_upgrade.rs @@ -115,6 +115,7 @@ WHERE version=20260919000000; 20260923000000, 20261001000000, 20261004000000, + 20261007000000, ] ); assert_eq!( diff --git a/crates/aether-data/runtime/src/lifecycle/migrate/tests/overview_fact_metadata.rs b/crates/aether-data/runtime/src/lifecycle/migrate/tests/overview_fact_metadata.rs index ca8552f24..3fcedcab3 100644 --- a/crates/aether-data/runtime/src/lifecycle/migrate/tests/overview_fact_metadata.rs +++ b/crates/aether-data/runtime/src/lifecycle/migrate/tests/overview_fact_metadata.rs @@ -168,7 +168,7 @@ VALUES('employee-key','owner',repeat('e',64),false), let bootstrap = include_str!("../../../../schema/bootstrap/postgres/190_overview_analytics.sql"); let view_start = bootstrap - .find("CREATE OR REPLACE VIEW public.usage_analytics_facts_v1 AS") + .find("CREATE OR REPLACE FUNCTION public.usage_customer_billable_amount(") .unwrap(); sqlx::raw_sql(&bootstrap[view_start..]) .execute(&pool) diff --git a/crates/aether-data/runtime/src/repository/auth/memory.rs b/crates/aether-data/runtime/src/repository/auth/memory.rs index cb4775ec5..950041ecd 100644 --- a/crates/aether-data/runtime/src/repository/auth/memory.rs +++ b/crates/aether-data/runtime/src/repository/auth/memory.rs @@ -1024,8 +1024,13 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { export.ip_rules = ip_rules; } } - if let Some(feature_settings) = record.feature_settings { - if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + if let Some(selection) = record.routing_group_selection { + export.feature_settings = selection.merge_feature_settings( + export.feature_settings.as_ref(), + record.feature_settings, + ); + } else if let Some(feature_settings) = record.feature_settings { export.feature_settings = match feature_settings { Some(serde_json::Value::Null) | None => None, Some(value) => Some(value), @@ -1089,8 +1094,13 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository { export.ip_rules = ip_rules; } } - if let Some(feature_settings) = record.feature_settings { - if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) { + if let Some(selection) = record.routing_group_selection { + export.feature_settings = selection.merge_feature_settings( + export.feature_settings.as_ref(), + record.feature_settings, + ); + } else if let Some(feature_settings) = record.feature_settings { export.feature_settings = match feature_settings { Some(serde_json::Value::Null) | None => None, Some(value) => Some(value), @@ -2036,6 +2046,7 @@ mod tests { concurrent_limit_present: false, ip_rules: None, feature_settings: Some(Some(serde_json::json!({"must_not_change": true}))), + routing_group_selection: None, }) .await .expect("locked basic update should resolve") @@ -2103,6 +2114,7 @@ mod tests { concurrent_limit_present: false, ip_rules: None, feature_settings: Some(Some(serde_json::json!({"admin": true}))), + routing_group_selection: None, }) .await .expect("administrator update should resolve") @@ -2363,6 +2375,130 @@ mod tests { ); } + #[tokio::test] + async fn user_key_feature_and_group_updates_merge_against_the_locked_current_record() { + use super::super::UpdateApiKeyRoutingGroupSelection; + + fn patch() -> UpdateUserApiKeyBasicRecord { + UpdateUserApiKeyBasicRecord { + user_id: "user-1".into(), + api_key_id: "key-1".into(), + key_encrypted: None, + key_encrypted_present: false, + name: None, + name_present: false, + rate_limit: None, + rate_limit_present: false, + concurrent_limit: None, + concurrent_limit_present: false, + ip_rules: None, + feature_settings: None, + routing_group_selection: Some(UpdateApiKeyRoutingGroupSelection { group_id: None }), + } + } + + // Both repository entry points must apply the same merge, with the + // self-service entry point additionally fencing locked keys. + for require_unlocked in [false, true] { + let repository = InMemoryAuthApiKeySnapshotRepository::seed([( + None, + sample_snapshot("key-1", "user-1"), + )]); + repository.set_user_api_key_feature_settings("user-1", "key-1", Some(serde_json::json!({ + "routing_group_id": "group-a", "routing_group_name": "stale-name", "pii": {"enabled": false}, + }))).await.unwrap().unwrap(); + async fn apply( + repository: &InMemoryAuthApiKeySnapshotRepository, + record: UpdateUserApiKeyBasicRecord, + require_unlocked: bool, + ) -> StoredAuthApiKeyExportRecord { + if require_unlocked { + repository + .update_user_api_key_basic_if_unlocked(record) + .await + .unwrap() + .unwrap() + } else { + repository + .update_user_api_key_basic(record) + .await + .unwrap() + .unwrap() + } + } + + // The PII request is prepared while A is selected, but another + // request selects B before that prepared replacement is committed. + let mut prepared_pii = patch(); + prepared_pii.feature_settings = Some(Some(serde_json::json!({ + "pii": {"enabled": true}, "routing_group_id": "group-a", "routing_group_name": "injected-name", + }))); + let mut select_b = patch(); + select_b.routing_group_selection.as_mut().unwrap().group_id = + Some(Some("group-b".into())); + apply(&repository, select_b, require_unlocked).await; + let merged = apply(&repository, prepared_pii, require_unlocked).await; + assert_eq!( + merged.feature_settings, + Some(serde_json::json!({ + "pii": {"enabled": true}, "routing_group_id": "group-b", + })) + ); + + // Conversely a group-only request prepared before a feature change + // must preserve the latest feature object when it reaches storage. + let mut prepared_group = patch(); + prepared_group + .routing_group_selection + .as_mut() + .unwrap() + .group_id = Some(Some("group-c".into())); + let mut latest_pii = patch(); + latest_pii.feature_settings = Some(Some( + serde_json::json!({ "pii": { "enabled": false, "mode": "strict" } }), + )); + apply(&repository, latest_pii, require_unlocked).await; + let merged = apply(&repository, prepared_group, require_unlocked).await; + assert_eq!( + merged.feature_settings, + Some(serde_json::json!({ + "pii": {"enabled": false, "mode": "strict"}, "routing_group_id": "group-c", + })) + ); + + let mut clear_features = patch(); + clear_features.feature_settings = Some(None); + let cleared = apply(&repository, clear_features, require_unlocked).await; + assert_eq!( + cleared.feature_settings, + Some(serde_json::json!({"routing_group_id": "group-c"})) + ); + let mut clear_group = patch(); + clear_group + .routing_group_selection + .as_mut() + .unwrap() + .group_id = Some(None); + assert!(apply(&repository, clear_group, require_unlocked) + .await + .feature_settings + .is_none()); + + // Administrative callers can still replace the complete document. + let mut admin = patch(); + admin.routing_group_selection = None; + admin.feature_settings = Some(Some( + serde_json::json!({ "routing_group_id": "admin-group", "admin": true }), + )); + assert_eq!( + apply(&repository, admin, require_unlocked) + .await + .feature_settings, + Some(serde_json::json!({ "routing_group_id": "admin-group", "admin": true })) + ); + } + } + #[tokio::test] async fn update_user_api_key_basic_updates_concurrent_limit() { let repository = InMemoryAuthApiKeySnapshotRepository::seed(vec![( @@ -2384,6 +2520,7 @@ mod tests { concurrent_limit_present: true, ip_rules: None, feature_settings: None, + routing_group_selection: None, }) .await .expect("update should succeed") @@ -2419,6 +2556,7 @@ mod tests { concurrent_limit_present: true, ip_rules: None, feature_settings: None, + routing_group_selection: None, }) .await .expect("nullable values should clear") @@ -2441,6 +2579,7 @@ mod tests { concurrent_limit_present: false, ip_rules: None, feature_settings: None, + routing_group_selection: None, }) .await .expect("zero rate limit should persist") diff --git a/crates/aether-data/runtime/src/repository/auth/mod.rs b/crates/aether-data/runtime/src/repository/auth/mod.rs index 3bb4e469a..9fde95aca 100644 --- a/crates/aether-data/runtime/src/repository/auth/mod.rs +++ b/crates/aether-data/runtime/src/repository/auth/mod.rs @@ -6,8 +6,8 @@ pub use aether_data_contracts::repository::auth::{ AuthApiKeyLookupKey, AuthApiKeyReadRepository, AuthApiKeyWriteRepository, AuthRepository, CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, ResolvedAuthApiKeySnapshot, ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery, - StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, - UpdateUserApiKeyBasicRecord, + StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, UpdateApiKeyRoutingGroupSelection, + UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord, }; #[cfg(feature = "postgres")] pub use aether_data_postgres::SqlxAuthApiKeySnapshotReadRepository; diff --git a/crates/aether-data/runtime/src/repository/settlement/memory.rs b/crates/aether-data/runtime/src/repository/settlement/memory.rs index 33032c6b2..0ebfcde52 100644 --- a/crates/aether-data/runtime/src/repository/settlement/memory.rs +++ b/crates/aether-data/runtime/src/repository/settlement/memory.rs @@ -1086,6 +1086,116 @@ mod tests { .expect("wallet should build") } + fn group_billed_input(request_id: &str) -> UsageSettlementInput { + UsageSettlementInput { + request_id: request_id.to_string(), + user_id: Some("user-1".to_string()), + api_key_id: Some("key-1".to_string()), + api_key_is_standalone: false, + provider_id: Some("provider-1".to_string()), + status: "completed".to_string(), + billing_status: "pending".to_string(), + total_cost_usd: 2.0, + actual_total_cost_usd: 0.5, + billing_cost_usd: Some(3.0), + finalized_at_unix_secs: Some(200), + } + } + + #[tokio::test] + async fn group_customer_charge_debits_user_wallet_once_without_inflating_provider_cost() { + let repository = + InMemorySettlementRepository::seed(vec![sample_user_wallet("user-wallet", "user-1")]); + let input = group_billed_input("group-billed-user"); + let first = repository + .settle_usage(input.clone()) + .await + .unwrap() + .unwrap(); + assert_eq!(first.wallet_id.as_deref(), Some("user-wallet")); + assert_eq!(first.wallet_balance_before, Some(12.0)); + assert_eq!(first.wallet_balance_after, Some(9.0)); + assert_eq!(first.provider_monthly_used_usd, Some(0.5)); + assert_eq!(repository.settle_usage(input).await.unwrap(), Some(first)); + + repository.wallets.with_mut(|wallets| { + let wallet = &wallets["user-wallet"]; + assert_eq!(wallet.balance, 7.0); + assert_eq!(wallet.gift_balance, 2.0); + assert_eq!(wallet.total_consumed, 3.0); + }); + assert_eq!( + repository.provider_monthly_used.read().unwrap()["provider-1"], + 0.5 + ); + } + + #[tokio::test] + async fn zero_group_customer_charge_keeps_wallet_unchanged_and_records_provider_cost() { + let repository = + InMemorySettlementRepository::seed(vec![sample_user_wallet("user-wallet", "user-1")]); + let mut input = group_billed_input("group-billed-free"); + input.billing_cost_usd = Some(0.0); + let first = repository + .settle_usage(input.clone()) + .await + .unwrap() + .unwrap(); + assert_eq!(first.billing_status, "settled"); + assert_eq!(first.wallet_balance_before, Some(12.0)); + assert_eq!(first.wallet_balance_after, Some(12.0)); + assert_eq!(first.provider_monthly_used_usd, Some(0.5)); + assert_eq!(repository.settle_usage(input).await.unwrap(), Some(first)); + repository.wallets.with_mut(|wallets| { + let wallet = &wallets["user-wallet"]; + assert_eq!(wallet.balance, 10.0); + assert_eq!(wallet.gift_balance, 2.0); + assert_eq!(wallet.total_consumed, 0.0); + }); + assert_eq!( + repository.provider_monthly_used.read().unwrap()["provider-1"], + 0.5 + ); + } + + #[tokio::test] + async fn standalone_key_wallet_pays_group_customer_charge_without_debiting_owner() { + let repository = InMemorySettlementRepository::seed(vec![ + sample_wallet(), + sample_user_wallet("owner-wallet", "user-1"), + ]); + let mut input = group_billed_input("group-billed-standalone"); + input.api_key_is_standalone = true; + let settlement = repository.settle_usage(input).await.unwrap().unwrap(); + assert_eq!(settlement.wallet_id.as_deref(), Some("wallet-1")); + assert_eq!(settlement.wallet_balance_after, Some(9.0)); + assert_eq!(settlement.provider_monthly_used_usd, Some(0.5)); + repository.wallets.with_mut(|wallets| { + assert_eq!(wallets["wallet-1"].balance, 7.0); + assert_eq!(wallets["wallet-1"].total_consumed, 3.0); + assert_eq!(wallets["owner-wallet"].balance, 10.0); + assert_eq!(wallets["owner-wallet"].gift_balance, 2.0); + assert_eq!(wallets["owner-wallet"].total_consumed, 0.0); + }); + } + + #[tokio::test] + async fn invalid_customer_charge_rejects_settlement_before_mutating_financial_state() { + for charge in [-0.01, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]); + let mut input = group_billed_input("group-billed-invalid"); + input.billing_cost_usd = Some(charge); + assert!(repository.settle_usage(input).await.is_err()); + repository.wallets.with_mut(|wallets| { + assert_eq!(wallets["wallet-1"].balance, 10.0); + assert_eq!(wallets["wallet-1"].gift_balance, 2.0); + assert_eq!(wallets["wallet-1"].total_consumed, 0.0); + }); + assert!(repository.provider_monthly_used.read().unwrap().is_empty()); + assert!(repository.settlements.read().unwrap().is_empty()); + } + } + #[tokio::test] async fn settles_usage_against_wallet_and_provider_quota() { let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]); @@ -1100,6 +1210,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 3.0, actual_total_cost_usd: 6.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }) .await @@ -1127,6 +1238,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 3.0, actual_total_cost_usd: 6.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }) .await @@ -1152,6 +1264,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 3.0, actual_total_cost_usd: 6.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }) .await @@ -1181,6 +1294,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 3.0, actual_total_cost_usd: 1.5, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }) .await @@ -1207,6 +1321,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 3.0, actual_total_cost_usd: 15.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }) .await @@ -1234,6 +1349,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 3.0, actual_total_cost_usd: 6.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }; @@ -1263,6 +1379,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 3.0, actual_total_cost_usd: 6.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }; let mut tasks = Vec::new(); @@ -1303,6 +1420,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 1.0, actual_total_cost_usd: 1.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(200), }) .await; @@ -1334,6 +1452,7 @@ mod tests { billing_status: "pending".to_string(), total_cost_usd: 2.0, actual_total_cost_usd: 1.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(250), }) .await @@ -1351,6 +1470,7 @@ mod tests { billing_status: "settled".to_string(), total_cost_usd: 2.0, actual_total_cost_usd: 1.0, + billing_cost_usd: None, finalized_at_unix_secs: Some(250), }) .await diff --git a/crates/aether-data/runtime/src/repository/usage/memory.rs b/crates/aether-data/runtime/src/repository/usage/memory.rs index f117a6d4a..3d1de26ab 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory.rs @@ -4,9 +4,9 @@ use std::sync::RwLock; use aether_ai_formats::UPSTREAM_IS_STREAM_KEY; use aether_data_contracts::repository::usage::{ - canonical_usage_body_ref_for, parse_usage_body_ref, sanitize_usage_request_metadata, - usage_body_ref, StoredUsageAuditAggregation, StoredUsageAuditSummary, - StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary, + canonical_usage_body_ref_for, parse_usage_body_ref, preserve_usage_routing_group_snapshot, + sanitize_usage_request_metadata, usage_body_ref, StoredUsageAuditAggregation, + StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary, @@ -22,9 +22,11 @@ use aether_data_contracts::repository::usage::{ UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageSettledCostSummaryQuery, - UsageTimeSeriesGranularity, UsageTimeSeriesQuery, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, - PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, - REQUESTED_REASONING_EFFORT_METADATA_KEY, + UsageTimeSeriesGranularity, UsageTimeSeriesQuery, BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, + PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, + PROVIDER_SERVICE_TIER_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, }; use async_trait::async_trait; use chrono::Utc; @@ -3035,6 +3037,10 @@ fn retain_previous_request_audit_metadata( "request_path", "request_query_string", "request_path_and_query", + ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, + ROUTING_GROUP_ID_METADATA_KEY, + ROUTING_GROUP_NAME_METADATA_KEY, ] { if let Some(value) = metadata.get(key) { retained.insert(key.to_string(), value.clone()); @@ -3213,7 +3219,13 @@ impl UsageWriteRepository for InMemoryUsageReadRepository { .and_then(|existing| existing.request_metadata.clone()) } }); - let request_metadata = sanitize_memory_request_metadata(request_metadata); + let request_metadata = + sanitize_memory_request_metadata(preserve_usage_routing_group_snapshot( + request_metadata, + existing + .as_ref() + .and_then(|stored| stored.request_metadata.as_ref()), + )); let (request_body, request_body_ref, request_body_state) = merge_usage_body_capture( capture_usage.request_body.take(), capture_usage.request_body_ref.take(), diff --git a/crates/aether-data/runtime/src/repository/usage/memory/analytics.rs b/crates/aether-data/runtime/src/repository/usage/memory/analytics.rs index 72414013d..a709e3b7f 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory/analytics.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory/analytics.rs @@ -136,14 +136,16 @@ fn apply_allocations( } fn decimal_sum( rows: &[&StoredRequestUsageAudit], - value: impl Fn(&StoredRequestUsageAudit) -> f64, + value: impl Fn(&StoredRequestUsageAudit) -> Option, ) -> Option { let amounts = rows .iter() .filter(|row| { available(row, USAGE_PRICING_AVAILABLE_METADATA_KEY) && row.billing_status == "settled" }) - .map(|row| (value(row) * 100_000_000.0).round() as i128) + .filter_map(|row| value(row)) + .filter(|amount| amount.is_finite()) + .map(|amount| (amount * 100_000_000.0).round() as i128) .collect::>(); if amounts.is_empty() { None @@ -270,8 +272,8 @@ fn metrics( metrics.first_byte_p90_ms = first_percentile(0.9); metrics.first_byte_p99_ms = first_percentile(0.99); metrics.usage_active_users = users.len() as u64; - metrics.rated_amount = decimal_sum(rows, |row| row.total_cost_usd); - metrics.billable_amount = decimal_sum(rows, |row| row.actual_total_cost_usd); + metrics.rated_amount = decimal_sum(rows, |row| Some(row.total_cost_usd)); + metrics.billable_amount = decimal_sum(rows, |row| row.billing_cost()); metrics } @@ -281,7 +283,7 @@ fn dashboard_total_metrics( ) -> UsageAnalyticsMetrics { let mut metrics = UsageAnalyticsMetrics { request_count: rows.len() as u64, - billable_amount: decimal_sum(rows, |row| row.actual_total_cost_usd), + billable_amount: decimal_sum(rows, |row| row.billing_cost()), ..Default::default() }; for row in rows { diff --git a/crates/aether-data/runtime/src/repository/usage/memory/dashboard_summary.rs b/crates/aether-data/runtime/src/repository/usage/memory/dashboard_summary.rs index e94b66fd6..d084556d6 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory/dashboard_summary.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory/dashboard_summary.rs @@ -52,8 +52,10 @@ impl DashboardProjection { != Some(false) }; let usage = available(USAGE_AVAILABLE_METADATA_KEY); - let priced = - available(USAGE_PRICING_AVAILABLE_METADATA_KEY) && row.billing_status == "settled"; + let billing_cost = row.billing_cost(); + let priced = available(USAGE_PRICING_AVAILABLE_METADATA_KEY) + && row.billing_status == "settled" + && billing_cost.is_some(); let stream = row .request_metadata .as_ref() @@ -105,7 +107,9 @@ impl DashboardProjection { actor: analytics::actor(row, keys).map(str::to_owned), metrics, billable_units: priced - .then(|| (row.actual_total_cost_usd * 100_000_000.0).round() as i128), + .then_some(billing_cost) + .flatten() + .map(|cost| (cost * 100_000_000.0).round() as i128), }, ); } diff --git a/crates/aether-data/runtime/src/repository/usage/memory/tests.rs b/crates/aether-data/runtime/src/repository/usage/memory/tests.rs index 3d8c0a2a9..bb2d94955 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory/tests.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory/tests.rs @@ -22,6 +22,63 @@ use aether_data_contracts::repository::usage::{ }; use serde_json::json; +#[tokio::test] +async fn customer_billing_statistics_use_frozen_factors_and_preserve_legacy_provider_cost() { + use aether_data_contracts::repository::usage::*; + let now = chrono::Utc::now(); + let at = now - chrono::Duration::seconds(10); + let mut billed = sample_usage("customer-billed", at.timestamp()); + billed.total_cost_usd = 2.0; + billed.actual_total_cost_usd = 0.5; + billed.request_metadata = Some(json!({ + "billing_multiplier_snapshot": { + "version": 1, + "factors": {"routing_group": 2.0, "user_group": 0.75}, + "multiplier": 1.5 + }, + "routing_group_billing_multiplier": 99.0, + "rate_multiplier": 0.25 + })); + let mut legacy = sample_usage("customer-legacy", at.timestamp()); + legacy.total_cost_usd = 2.0; + legacy.actual_total_cost_usd = 0.5; + let mut free = sample_usage("customer-free", at.timestamp()); + free.total_cost_usd = 2.0; + free.actual_total_cost_usd = 0.5; + free.request_metadata = Some(json!({"routing_group_billing_multiplier": 0.0})); + let mut invalid = sample_usage("customer-invalid", at.timestamp()); + invalid.total_cost_usd = 999.0; + invalid.actual_total_cost_usd = 999.0; + invalid.request_metadata = Some(json!({"billing_multiplier_snapshot": null})); + let repo = InMemoryUsageReadRepository::seed([billed, legacy, free, invalid]) + .with_dashboard_stats_since(at - chrono::Duration::seconds(1)); + let overview = repo + .query_usage_analytics(&UsageAnalyticsQuery { + from_unix_ms: (at - chrono::Duration::seconds(1)).timestamp_millis() as u64, + to_unix_ms: now.timestamp_millis() as u64, + timezone: "UTC".into(), + limit: 1, + ..Default::default() + }) + .await + .unwrap(); + assert_eq!( + overview.summary.billable_amount.as_deref(), + Some("3.50000000") + ); + let query = UsageDashboardAnalyticsQuery { + timezone: "UTC".into(), + }; + let analytics = repo.query_dashboard_analytics(&query).await.unwrap(); + assert_eq!( + analytics.total.summary.billable_amount.as_deref(), + Some("3.50000000") + ); + let summary = repo.query_dashboard_summary(&query).await.unwrap(); + assert_eq!(summary.total.billable_amount.as_deref(), Some("3.50000000")); + assert_eq!(summary.total.pricing_available_count, 3); +} + #[tokio::test] async fn overview_model_performance_merges_provider_samples_without_pagination() { use aether_data_contracts::repository::usage::*; @@ -617,6 +674,58 @@ fn sample_upsert_usage_record(request_id: &str) -> UpsertUsageRecord { } } +#[tokio::test] +async fn upsert_preserves_routing_group_snapshot_across_terminal_metadata_replacement() { + for terminal_metadata in [ + None, + Some(json!({"rate_multiplier": 0.5, "billing_snapshot": {"status": "complete"}})), + Some(json!({ + "routing_group_billing_multiplier": 99.0, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 3.0}, "multiplier": 3.0}, + "routing_group_id": "changed-group", + "routing_group_name": "changed-group-name", + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440002", + "rate_multiplier": 0.5 + })), + ] { + let repository = InMemoryUsageReadRepository::default(); + let mut pending = sample_upsert_usage_record("req-group-snapshot"); + pending.request_metadata = Some(json!({ + "routing_group_billing_multiplier": 0.25, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 0.25, "user_group": 2.0}, "multiplier": 0.5}, + "routing_group_id": "group-original", + "routing_group_name": "请求时的分组", + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440001" + })); + repository + .upsert(pending) + .await + .expect("pending usage should persist"); + let mut terminal = sample_upsert_usage_record("req-group-snapshot"); + terminal.status = "completed".to_string(); + terminal.request_metadata = terminal_metadata; + terminal.updated_at_unix_secs += 1; + let stored = repository + .upsert(terminal) + .await + .expect("terminal usage should persist"); + assert_eq!(stored.routing_group_billing_multiplier(), 0.25); + assert_eq!(stored.billing_multiplier(), 0.5); + assert_eq!( + stored.request_metadata.as_ref().unwrap()["billing_multiplier_snapshot"], + json!({ + "version": 1, "factors": {"routing_group": 0.25, "user_group": 2.0}, "multiplier": 0.5 + }) + ); + assert_eq!(stored.routing_group_id(), Some("group-original")); + assert_eq!(stored.routing_group_name(), Some("请求时的分组")); + assert_eq!( + stored.request_metadata.as_ref().unwrap()["plan_usage_reservation_token"], + "550e8400-e29b-41d4-a716-446655440001" + ); + } +} + #[tokio::test] async fn upsert_preserves_full_http_captures_across_lifecycle_updates() { let repository = InMemoryUsageReadRepository::default(); diff --git a/crates/aether-routing-core/src/model.rs b/crates/aether-routing-core/src/model.rs index 88e866765..0f8c3952c 100644 --- a/crates/aether-routing-core/src/model.rs +++ b/crates/aether-routing-core/src/model.rs @@ -132,6 +132,86 @@ fn is_false(value: &bool) -> bool { mod execution_policy_tests { use super::*; + #[test] + fn group_visibility_is_opt_in_and_round_trips_without_losing_policy() { + let legacy: RoutingGroupConfig = serde_json::from_str("{}").unwrap(); + assert!(!legacy.user_visible); + assert!(!RoutingGroupConfig::default().user_visible); + for user_visible in [false, true] { + let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({ + "user_visible": user_visible, + "billing_multiplier": 0.5, + "disabled_providers": ["private-provider"], + "default_policy": { "scheduling_mode": "fixed_order" } + })) + .unwrap(); + let encoded = serde_json::to_value(&config).unwrap(); + assert_eq!(encoded["user_visible"], user_visible); + assert_eq!(config.billing_multiplier, 0.5); + assert_eq!(config.disabled_providers, ["private-provider"]); + assert_eq!( + serde_json::from_value::(encoded).unwrap(), + config + ); + } + for invalid in [ + serde_json::json!(null), + serde_json::json!("true"), + serde_json::json!(1), + ] { + assert!( + serde_json::from_value::(serde_json::json!({ + "user_visible": invalid + })) + .is_err() + ); + } + } + + #[test] + fn group_billing_multiplier_defaults_to_one_and_rejects_invalid_values() { + let legacy: RoutingGroupConfig = serde_json::from_str("{}").unwrap(); + assert_eq!(legacy.billing_multiplier, 1.0); + assert_eq!(RoutingGroupConfig::default().billing_multiplier, 1.0); + for multiplier in [0.0, 0.25, 1.0, 2.5] { + let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({ + "billing_multiplier": multiplier + })) + .unwrap(); + crate::validate_routing_group_config(&config).unwrap(); + assert_eq!(config.billing_multiplier, multiplier); + assert_eq!( + serde_json::from_value::( + serde_json::to_value(&config).unwrap() + ) + .unwrap(), + config + ); + } + for multiplier in [-1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + let config = RoutingGroupConfig { + billing_multiplier: multiplier, + ..RoutingGroupConfig::default() + }; + assert!(matches!( + crate::validate_routing_group_config(&config), + Err(crate::RoutingValidationError::InvalidBillingMultiplier) + )); + } + for value in [ + serde_json::json!(null), + serde_json::json!("2"), + serde_json::json!(false), + ] { + assert!( + serde_json::from_value::(serde_json::json!({ + "billing_multiplier": value + })) + .is_err() + ); + } + } + #[test] fn routing_failover_configuration_round_trips_and_validates() { let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({ @@ -188,6 +268,11 @@ pub struct RoutingModelPolicy { pub allowed_providers: Vec, #[serde(default)] pub allowed_keys: Vec, + /// Per-model provider enablement. A `false` value adds a provider to this + /// model's exclusions and `true` removes an inherited exclusion, including + /// one from the legacy group-wide `disabled_providers` baseline. + #[serde(default)] + pub provider_enabled_overrides: BTreeMap, #[serde(default)] pub provider_priority_overrides: BTreeMap, #[serde(default)] @@ -222,10 +307,17 @@ pub struct RoutingRule { pub stop_processing: bool, } -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct RoutingGroupConfig { - /// Providers excluded from every model in this group, including providers - /// otherwise selected by model policies or routing rules. + /// Whether authenticated users can discover and explicitly select this + /// group. Private bindings and automatic defaults remain independent. + #[serde(default)] + pub user_visible: bool, + /// Group-wide billing multiplier, snapshotted when a request is routed. + #[serde(default = "default_billing_multiplier")] + pub billing_multiplier: f64, + /// Legacy provider exclusion baseline for the group. Explicit per-model + /// enablement overrides may change it; allowlists and rules cannot. #[serde(default)] pub disabled_providers: Vec, /// The default policy is global for the selected strategy group. Model @@ -238,6 +330,23 @@ pub struct RoutingGroupConfig { pub rules: Vec, } +pub(crate) fn default_billing_multiplier() -> f64 { + 1.0 +} + +impl Default for RoutingGroupConfig { + fn default() -> Self { + Self { + user_visible: false, + billing_multiplier: default_billing_multiplier(), + disabled_providers: Vec::new(), + default_policy: RoutingDefaultPolicy::default(), + model_policies: Vec::new(), + rules: Vec::new(), + } + } +} + #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct RoutingGroupRecord { pub id: String, diff --git a/crates/aether-routing-core/src/policy.rs b/crates/aether-routing-core/src/policy.rs index 1cec1bff6..a3621e88e 100644 --- a/crates/aether-routing-core/src/policy.rs +++ b/crates/aether-routing-core/src/policy.rs @@ -47,10 +47,15 @@ pub struct MatchedRoutingRule { pub phase: RoutingRulePhase, } -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct ResolvedRoutingPolicy { + #[serde(default = "crate::model::default_billing_multiplier")] + pub billing_multiplier: f64, #[serde(default, skip_serializing_if = "Option::is_none")] pub group_id: Option, + /// Display name captured alongside the selected group by the gateway. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub group_name: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub group_version: Option, pub selection_source: String, @@ -80,7 +85,9 @@ pub fn resolve_routing_policy( .map_err(|error| RoutingPolicyError::InvalidConfig(error.to_string()))?; let mut policy = ResolvedRoutingPolicy { + billing_multiplier: config.billing_multiplier, group_id: input.group_id.map(str::to_string), + group_name: None, group_version: input.group_version, selection_source: input.selection_source.to_string(), requested_model: input.requested_model.to_string(), @@ -105,6 +112,11 @@ pub fn resolve_routing_policy( { apply_model_policy(&mut policy, model_policy); } + for model_policy in + matching_provider_enablement_policies(config, input.requested_model, input.resolved_model) + { + apply_provider_enable_overrides(&mut policy, model_policy); + } let condition_context = RoutingConditionContext { model: input.requested_model, @@ -284,6 +296,57 @@ fn matching_model_policies<'a>( .collect() } +fn matching_provider_enablement_policies<'a>( + config: &'a RoutingGroupConfig, + requested_model: &str, + resolved_model: &str, +) -> Vec<&'a RoutingModelPolicy> { + let mut matches = matching_model_policies(config, requested_model, resolved_model) + .into_iter() + .filter(|policy| !policy.provider_enabled_overrides.is_empty()) + .collect::>(); + // Enablement is a layered exception map: broad defaults first, then + // prefixes, then exact model entries. Other model-policy fields retain + // their historical configured-order merge semantics. + matches.sort_by_key(|policy| model_pattern_specificity(&policy.model)); + matches +} + +fn apply_provider_enable_overrides( + policy: &mut ResolvedRoutingPolicy, + model_policy: &RoutingModelPolicy, +) { + for (provider_id, enabled) in &model_policy.provider_enabled_overrides { + if *enabled { + policy + .ranking_overlay + .disabled_providers + .retain(|disabled| disabled != provider_id); + } else if !policy + .ranking_overlay + .disabled_providers + .iter() + .any(|disabled| disabled == provider_id) + { + policy + .ranking_overlay + .disabled_providers + .push(provider_id.clone()); + } + } +} + +fn model_pattern_specificity(pattern: &str) -> (u8, usize) { + let pattern = pattern.trim(); + if pattern == "*" { + (0, 0) + } else if let Some(prefix) = pattern.strip_suffix('*') { + (1, prefix.len()) + } else { + (2, 0) + } +} + fn model_allowed(patterns: &[String], requested_model: &str) -> bool { patterns.is_empty() || patterns @@ -405,7 +468,7 @@ mod tests { } #[test] - fn group_disabled_providers_apply_to_every_model_and_cannot_be_reenabled() { + fn legacy_group_exclusions_cannot_be_bypassed_by_allowlists_or_rule_actions() { let config: RoutingGroupConfig = serde_json::from_value(json!({ "disabled_providers": ["provider-disabled"], "model_policies": [{ @@ -478,6 +541,106 @@ mod tests { } } + #[test] + fn model_provider_enablement_is_scoped_and_specific_overrides_win() { + let mut config: RoutingGroupConfig = serde_json::from_value(json!({ + "disabled_providers": ["provider-root"], + "model_policies": [ + { + "model": "*", + "provider_enabled_overrides": { + "provider-model": false, + "provider-specific": false + } + }, + { + "model": "model-*", + "provider_enabled_overrides": { + "provider-model": true, + "provider-specific": true + } + }, + { + "model": "model-exact", + "provider_enabled_overrides": { + "provider-specific": false, + "provider-exact": true, + "provider-root": true + } + } + ] + })) + .unwrap(); + + // Persisted order need not put broad defaults first. Only the new + // enablement map follows specificity; priorities retain their old order. + config.model_policies.reverse(); + config.model_policies[0] + .provider_priority_overrides + .insert("provider-model".into(), 1); + config.model_policies[2] + .provider_priority_overrides + .insert("provider-model".into(), 9); + + let for_model = |model: &str| { + resolve_routing_policy( + &config, + RoutingPolicyInput { + group_id: Some("group-1"), + group_version: Some(1), + selection_source: "test", + requested_model: model, + resolved_model: model, + api_format: "openai:chat", + user_id: None, + api_key_id: None, + headers: &json!({}), + body: &json!({}), + phase: RoutingRulePhase::ClientRequest, + }, + ) + .unwrap() + }; + + let exact = for_model("model-exact"); + assert!(exact.ranking_overlay.provider_allowed("provider-model")); + assert_eq!( + exact.ranking_overlay.provider_priority_overrides["provider-model"], + 9 + ); + assert!(!exact.ranking_overlay.provider_allowed("provider-specific")); + assert!(exact.ranking_overlay.provider_allowed("provider-exact")); + assert!(exact.ranking_overlay.provider_allowed("provider-root")); + + let wildcard_prefix = for_model("model-other"); + assert!(wildcard_prefix + .ranking_overlay + .provider_allowed("provider-model")); + assert!(wildcard_prefix + .ranking_overlay + .provider_allowed("provider-specific")); + assert!(!wildcard_prefix + .ranking_overlay + .provider_allowed("provider-root")); + + let unrelated = for_model("other-model"); + assert!(!unrelated.ranking_overlay.provider_allowed("provider-model")); + assert!(!unrelated + .ranking_overlay + .provider_allowed("provider-specific")); + assert!(!unrelated.ranking_overlay.provider_allowed("provider-root")); + + let encoded = serde_json::to_value(&config).unwrap(); + assert_eq!( + encoded["model_policies"][2]["provider_enabled_overrides"]["provider-model"], + false + ); + assert_eq!( + serde_json::from_value::(encoded).unwrap(), + config + ); + } + #[test] fn all_model_scheduling_and_rankings_apply_to_future_models() { let config: RoutingGroupConfig = serde_json::from_value(json!({ @@ -599,6 +762,8 @@ mod tests { #[test] fn resolves_model_policy_and_matching_rule() { let config = RoutingGroupConfig { + user_visible: false, + billing_multiplier: 1.0, disabled_providers: vec![], default_policy: RoutingDefaultPolicy::default(), model_policies: vec![RoutingModelPolicy { @@ -668,6 +833,8 @@ mod tests { #[test] fn default_policy_applies_to_models_without_an_override() { let config = RoutingGroupConfig { + user_visible: false, + billing_multiplier: 2.5, disabled_providers: vec![], default_policy: RoutingDefaultPolicy { priority_mode: RoutingSetPriorityMode::GlobalKey, @@ -707,6 +874,7 @@ mod tests { assert_eq!(special.scheduling_mode, RoutingSchedulingMode::LoadBalance); assert!(special.keep_priority_on_conversion); assert_eq!(special.sticky_key_attempts, 3); + assert_eq!(special.billing_multiplier, 2.5); assert_eq!( special.ranking_overlay.allowed_providers, vec!["provider-special"] @@ -741,6 +909,7 @@ mod tests { assert_eq!(ordinary.scheduling_mode, RoutingSchedulingMode::LoadBalance); assert!(ordinary.keep_priority_on_conversion); assert_eq!(ordinary.sticky_key_attempts, 3); + assert_eq!(ordinary.billing_multiplier, 2.5); assert!(ordinary.ranking_overlay.allowed_providers.is_empty()); assert!(ordinary.ranking_overlay.allowed_keys.is_empty()); assert!(ordinary diff --git a/crates/aether-routing-core/src/ranking.rs b/crates/aether-routing-core/src/ranking.rs index a86190dcc..92bfc7298 100644 --- a/crates/aether-routing-core/src/ranking.rs +++ b/crates/aether-routing-core/src/ranking.rs @@ -13,7 +13,8 @@ pub enum CandidateKind { #[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] pub struct RankingOverlay { - /// Group-wide exclusions take precedence over every provider allowlist. + /// Effective provider exclusions after model overrides. These take + /// precedence over every provider allowlist. #[serde(default)] pub disabled_providers: Vec, #[serde(default)] diff --git a/crates/aether-routing-core/src/trace.rs b/crates/aether-routing-core/src/trace.rs index 573c61851..5c8280454 100644 --- a/crates/aether-routing-core/src/trace.rs +++ b/crates/aether-routing-core/src/trace.rs @@ -55,11 +55,15 @@ pub struct RoutingRuntimeFacts { pub priority_mode: Option, } -#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] pub struct RoutingDecisionTrace { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub billing_multiplier: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub group_id: Option, #[serde(default, skip_serializing_if = "Option::is_none")] + pub group_name: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] pub group_version: Option, pub selection_source: String, #[serde(default)] diff --git a/crates/aether-routing-core/src/validation.rs b/crates/aether-routing-core/src/validation.rs index a4c3fb8d7..80dd62055 100644 --- a/crates/aether-routing-core/src/validation.rs +++ b/crates/aether-routing-core/src/validation.rs @@ -28,6 +28,8 @@ const ROUTING_POOL_PRESETS: &[&str] = &[ #[derive(Debug, Error, Clone, PartialEq, Eq)] pub enum RoutingValidationError { + #[error("routing group billing multiplier must be a non-negative finite number")] + InvalidBillingMultiplier, #[error("routing failover rules are invalid: {0}")] InvalidFailoverRules(String), #[error("routing rule id is empty")] @@ -71,6 +73,9 @@ pub enum RoutingValidationError { pub fn validate_routing_group_config( config: &RoutingGroupConfig, ) -> Result<(), RoutingValidationError> { + if !config.billing_multiplier.is_finite() || config.billing_multiplier < 0.0 { + return Err(RoutingValidationError::InvalidBillingMultiplier); + } crate::validate_routing_failover_rules(&config.default_policy.execution_policy.failover_rules) .map_err(RoutingValidationError::InvalidFailoverRules)?; let mut rule_ids = BTreeSet::new(); diff --git a/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs b/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs index 894e14ef7..6ec77fa56 100644 --- a/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs +++ b/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs @@ -527,6 +527,7 @@ fn settlement_input(index: usize) -> UsageSettlementInput { billing_status: "pending".to_string(), total_cost_usd: 0.0, actual_total_cost_usd: COST_PER_REQUEST_USD, + billing_cost_usd: None, finalized_at_unix_secs: Some(now_unix_secs().saturating_add(index as u64)), } } diff --git a/crates/aether-usage/runtime/src/request_metadata.rs b/crates/aether-usage/runtime/src/request_metadata.rs index 8763b494d..8a6ff9428 100644 --- a/crates/aether-usage/runtime/src/request_metadata.rs +++ b/crates/aether-usage/runtime/src/request_metadata.rs @@ -8,9 +8,11 @@ use aether_data_contracts::repository::usage::{ sanitize_usage_request_metadata_object as project_usage_request_metadata_object, sanitize_usage_request_metadata_ref as project_usage_request_metadata_ref, usage_body_capture_is_authoritative, UsageBodyCaptureState, - PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, - PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_RESPONSE_MODEL_METADATA_KEY, - PROVIDER_SERVICE_TIER_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, + BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, + PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, + PROVIDER_RESPONSE_MODEL_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY, + REQUESTED_REASONING_EFFORT_METADATA_KEY, ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, + ROUTING_GROUP_ID_METADATA_KEY, ROUTING_GROUP_NAME_METADATA_KEY, }; use serde_json::{Map, Value}; @@ -112,6 +114,10 @@ pub(crate) fn retain_first_byte_request_metadata(value: Option) -> Option | "model_id" | "global_model_id" | "global_model_name" + | ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY + | BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY + | ROUTING_GROUP_ID_METADATA_KEY + | ROUTING_GROUP_NAME_METADATA_KEY ) }); (!metadata.is_empty()).then_some(Value::Object(metadata)) @@ -542,6 +548,10 @@ mod tests { "request_path": "/v1/chat/completions", "upstream_is_stream": true, "proxy": {"mode": "manual", "node_id": "proxy-1"}, + "routing_group_billing_multiplier": 0.25, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 0.25, "user_group": 2.0}, "multiplier": 0.5}, + "routing_group_id": "group-1", + "routing_group_name": "默认调度策略", "billing_snapshot": {"dimensions": [1, 2, 3]}, "settlement_snapshot": {"status": "pending"}, "stage_timings_ms": {"planning": 12} @@ -555,7 +565,11 @@ mod tests { "client_ip": "203.0.113.8", "request_path": "/v1/chat/completions", "request_path_and_query": "/v1/chat/completions", - "upstream_is_stream": true + "upstream_is_stream": true, + "routing_group_billing_multiplier": 0.25, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 0.25, "user_group": 2.0}, "multiplier": 0.5}, + "routing_group_id": "group-1", + "routing_group_name": "默认调度策略" }) ); } @@ -644,6 +658,52 @@ mod tests { .is_none()); } + #[test] + fn routing_group_snapshot_survives_seed_and_sanitization() { + for multiplier in [0.0, 0.25, 1.0, 2.5] { + let context = json!({ + "routing_group_billing_multiplier": multiplier, + "billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": multiplier, "user_group": 2.0}, "multiplier": multiplier * 2.0}, + "routing_group_id": "group-1", + "routing_group_name": "请求时的分组", + "rate_multiplier": 0.75, + "routing_trace": {"untrusted": true} + }); + let metadata = build_usage_request_metadata_seed(&sample_plan(), context.as_object()) + .expect("group snapshot should survive projection"); + assert_eq!(metadata["routing_group_billing_multiplier"], multiplier); + assert_eq!( + metadata["billing_multiplier_snapshot"], + context["billing_multiplier_snapshot"] + ); + assert_eq!(metadata["routing_group_id"], "group-1"); + assert_eq!(metadata["routing_group_name"], "请求时的分组"); + assert_eq!(metadata["rate_multiplier"], 0.75); + assert!(metadata.get("routing_trace").is_none()); + assert_eq!( + sanitize_usage_request_metadata(Some(metadata.clone())), + Some(metadata) + ); + } + for multiplier in [json!(-1), json!("Infinity"), json!(f64::NAN)] { + let context = json!({"routing_group_billing_multiplier": multiplier}); + let metadata = build_usage_request_metadata_seed(&sample_plan(), context.as_object()) + .expect( + "invalid pricing must retain a marker instead of falling back to legacy billing", + ); + assert_eq!( + metadata.get("billing_multiplier_snapshot"), + Some(&Value::Null) + ); + assert!( + aether_data_contracts::repository::usage::billing_multiplier_snapshot(Some( + &metadata + )) + .is_err() + ); + } + } + #[test] fn builds_seed_from_context_and_allowlisted_metadata_only() { let metadata = build_usage_request_metadata_seed( diff --git a/crates/aether-usage/runtime/src/runtime.rs b/crates/aether-usage/runtime/src/runtime.rs index 48c1354f9..1fe525738 100644 --- a/crates/aether-usage/runtime/src/runtime.rs +++ b/crates/aether-usage/runtime/src/runtime.rs @@ -26,9 +26,7 @@ use crate::request_metadata::{ request_body_derived_facts_action, retain_first_byte_request_metadata, RequestBodyDerivedFactsAction, }; -use crate::settlement::{ - reconcile_usage_policy_cost_for_event_with_result, settle_usage_with_reconciled_cost, -}; +use crate::settlement::settle_usage_after_upsert; use crate::shutdown::{UsageBackgroundTasks, UsageShutdownState}; use crate::worker::{ build_usage_queue_worker_with_record_gate, UsageWorkerControl, UsageWorkerObservation, @@ -5156,21 +5154,6 @@ impl UsageRuntime { where T: UsageRuntimeAccess, { - let reconciled = match reconcile_usage_policy_cost_for_event_with_result(data, event).await - { - Ok(reconciled) => reconciled, - Err(err) => { - warn!( - event_name = "usage_event_cost_reconciliation_failed", - log_type = "event", - usage_event_type = ?event.event_type, - request_id = %event.request_id, - error = %err, - "usage runtime failed to reconcile plan cost before direct usage upsert" - ); - return false; - } - }; match build_upsert_usage_record_from_event(event) { Ok(record) => match catch_usage_writer_panic( "direct usage upsert", @@ -5179,9 +5162,7 @@ impl UsageRuntime { .await { Ok(Some(stored)) => { - if let Err(err) = - settle_usage_with_reconciled_cost(data, &stored, reconciled).await - { + if let Err(err) = settle_usage_after_upsert(data, &stored, event).await { warn!( event_name = "usage_terminal_settlement_failed", log_type = "event", diff --git a/crates/aether-usage/runtime/src/settlement.rs b/crates/aether-usage/runtime/src/settlement.rs index 65d92b190..d54d3a9bb 100644 --- a/crates/aether-usage/runtime/src/settlement.rs +++ b/crates/aether-usage/runtime/src/settlement.rs @@ -7,7 +7,7 @@ use aether_data_contracts::repository::settlement::{ }; use aether_data_contracts::repository::usage::PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY; use aether_data_contracts::repository::usage::{ - cancelled_request_fee_is_billable, StoredRequestUsageAudit, + billing_multiplier_snapshot, cancelled_request_fee_is_billable, StoredRequestUsageAudit, }; use aether_data_contracts::{DataLayerError, DataLayerError::InvalidInput}; use async_trait::async_trait; @@ -72,16 +72,25 @@ pub(crate) async fn reconcile_usage_policy_cost_for_event_with_result( return Ok(None); }; let actual_cost_units = if terminal_state == UsagePolicyCostReservationState::Finalized { - let actual_cost_usd = event.data.actual_total_cost_usd.ok_or_else(|| { + let snapshot = billing_multiplier_snapshot(event.data.request_metadata.as_ref())?; + let cost = if snapshot.is_some() { + event.data.total_cost_usd + } else { + event.data.actual_total_cost_usd + }; + let actual_cost_usd = cost.ok_or_else(|| { InvalidInput( "completed usage event with a plan reservation token is missing actual cost" .to_string(), ) })?; - nonnegative_usd_to_usage_policy_cost_units(finite_cost(actual_cost_usd)?.max(0.0)) - .ok_or_else(|| { - InvalidInput("usage policy settlement cost exceeds the supported range".to_string()) - })? + let actual_cost_usd = match snapshot { + Some(snapshot) => snapshot.cost(actual_cost_usd)?, + None => finite_cost(actual_cost_usd)?.max(0.0), + }; + nonnegative_usd_to_usage_policy_cost_units(actual_cost_usd).ok_or_else(|| { + InvalidInput("usage policy settlement cost exceeds the supported range".to_string()) + })? } else { 0 }; @@ -115,6 +124,34 @@ pub async fn settle_usage_if_needed( settle_usage_with_reconciled_cost(writer, usage, None).await } +pub(crate) async fn settle_usage_after_upsert( + writer: &dyn UsageSettlementWriter, + usage: &StoredRequestUsageAudit, + event: &UsageEvent, +) -> Result<(), DataLayerError> { + // Different admissions can share a client request id. Finalize that event's + // own server-issued reservation without borrowing the other admission's rate. + if event_usage_policy_reservation_token(event).is_some() + && event_usage_policy_reservation_token(event) != usage_policy_reservation_token(usage) + && !plan_usage_reservation_reconciliation_is_deferred(event.data.request_metadata.as_ref()) + { + let billable = event.event_type == UsageEventType::Completed + || (event.event_type == UsageEventType::Cancelled + && cancelled_request_fee_is_billable(event.data.request_metadata.as_ref())); + if billable + && billing_multiplier_snapshot(event.data.request_metadata.as_ref())?.is_none() + && billing_multiplier_snapshot(usage.request_metadata.as_ref())?.is_some() + { + return Err(InvalidInput( + "colliding usage admission is missing its own billing multiplier snapshot" + .to_string(), + )); + } + reconcile_usage_policy_cost_for_event(writer, event).await?; + } + settle_usage_if_needed(writer, usage).await +} + pub(crate) async fn settle_usage_with_reconciled_cost( writer: &dyn UsageSettlementWriter, usage: &StoredRequestUsageAudit, @@ -127,6 +164,11 @@ pub(crate) async fn settle_usage_with_reconciled_cost( return Ok(()); } + let billing_cost_usd = match billing_multiplier_snapshot(usage.request_metadata.as_ref())? { + Some(snapshot) => snapshot.cost(usage.total_cost_usd)?, + None => finite_cost(usage.actual_total_cost_usd)?.max(0.0), + }; + let finalized_at_unix_secs = usage .finalized_at_unix_secs .or(Some(usage.updated_at_unix_secs)); @@ -148,14 +190,14 @@ pub(crate) async fn settle_usage_with_reconciled_cost( { ( UsagePolicyCostReservationState::Finalized, - nonnegative_usd_to_usage_policy_cost_units( - finite_cost(usage.actual_total_cost_usd)?.max(0.0), - ) - .ok_or_else(|| { - InvalidInput( - "usage policy settlement cost exceeds the supported range".to_string(), - ) - })?, + nonnegative_usd_to_usage_policy_cost_units(billing_cost_usd).ok_or_else( + || { + InvalidInput( + "usage policy settlement cost exceeds the supported range" + .to_string(), + ) + }, + )?, ) } else { (UsagePolicyCostReservationState::Released, 0) @@ -194,6 +236,7 @@ pub(crate) async fn settle_usage_with_reconciled_cost( billing_status: usage.billing_status.clone(), total_cost_usd: finite_cost(usage.total_cost_usd)?, actual_total_cost_usd: finite_cost(usage.actual_total_cost_usd)?, + billing_cost_usd: Some(billing_cost_usd), finalized_at_unix_secs, }; let _ = writer.settle_usage(input).await?; @@ -449,6 +492,62 @@ mod tests { ); } + #[tokio::test] + async fn composite_billing_rate_charges_customer_without_changing_provider_cost() { + for (group_rate, user_rate, expected_cost) in [(2.0, 0.75, 1.875), (0.0, 3.0, 0.0)] { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let mut usage = sample_usage(); + let snapshot = + aether_data_contracts::repository::usage::BillingMultiplierSnapshot::from_factors( + std::collections::BTreeMap::from([ + ("routing_group".to_string(), group_rate), + ("user_group".to_string(), user_rate), + ]), + ) + .unwrap(); + usage.request_metadata.as_mut().unwrap()["billing_multiplier_snapshot"] = + json!(snapshot); + settle_usage_if_needed(&writer, &usage).await.unwrap(); + let inputs = writer.inputs.lock().unwrap(); + assert_eq!(inputs[0].billing_cost_usd, Some(expected_cost)); + assert_eq!(inputs[0].total_cost_usd, 1.25); + assert_eq!(inputs[0].actual_total_cost_usd, 0.75); + assert_eq!( + writer.reconciliations.lock().unwrap()[0].actual_cost_units, + (expected_cost * 100_000_000.0).round() as u64 + ); + } + } + + #[tokio::test] + async fn corrupt_or_overflowing_billing_rate_never_changes_wallet_or_reservation() { + for (base, snapshot) in [ + (1.25, serde_json::Value::Null), + ( + 1.25, + json!({"version":1,"factors":{"routing_group":2},"multiplier":1}), + ), + ( + f64::MAX, + json!({"version":1,"factors":{"routing_group":2},"multiplier":2}), + ), + ] { + let writer = TestSettlementWriter { + has_writer: true, + ..Default::default() + }; + let mut usage = sample_usage(); + usage.total_cost_usd = base; + usage.request_metadata.as_mut().unwrap()["billing_multiplier_snapshot"] = snapshot; + assert!(settle_usage_if_needed(&writer, &usage).await.is_err()); + assert!(writer.inputs.lock().unwrap().is_empty()); + assert!(writer.reconciliations.lock().unwrap().is_empty()); + } + } + #[tokio::test] async fn releases_pending_cancelled_usage_without_wallet_settlement() { let writer = TestSettlementWriter { diff --git a/crates/aether-usage/runtime/src/settlement_reuse_tests.rs b/crates/aether-usage/runtime/src/settlement_reuse_tests.rs index 34fc4a60b..9330bbb2f 100644 --- a/crates/aether-usage/runtime/src/settlement_reuse_tests.rs +++ b/crates/aether-usage/runtime/src/settlement_reuse_tests.rs @@ -66,6 +66,10 @@ impl UsageSettlementWriter for ReuseStore { &self, input: ReconcileUsagePolicyCostInput, ) -> Result, DataLayerError> { + assert!( + self.upserts.load(Ordering::Relaxed) > 0, + "a durable usage row must exist before a reservation is finalized" + ); input.validate()?; self.reconciliations.lock().unwrap().push(input.clone()); tokio::task::yield_now().await; @@ -185,7 +189,7 @@ async fn write(store: &ReuseStore, event: UsageEvent, direct: bool) { } #[tokio::test] -async fn worker_and_direct_writes_reuse_confirmed_reservation_and_still_settle_wallet() { +async fn worker_and_direct_writes_persist_before_reconciling_and_settling_wallet() { for direct in [false, true] { let store = ReuseStore::default(); write(&store, event(), direct).await; @@ -198,11 +202,12 @@ async fn worker_and_direct_writes_reuse_confirmed_reservation_and_still_settle_w assert_eq!(settlements.len(), 1); assert_eq!(settlements[0].request_id, "req-1"); assert_eq!(settlements[0].actual_total_cost_usd, 0.75); + assert_eq!(settlements[0].billing_cost_usd, Some(0.75)); } } #[tokio::test] -async fn missing_or_different_reconciliation_results_keep_stored_usage_reconciliation() { +async fn missing_or_different_reconciliation_results_do_not_repeat_stored_usage_reconciliation() { let changes: [fn(&mut StoredUsagePolicyCostReservation); 9] = [ |row| row.request_id = "other-request".to_string(), |row| row.subject_id = "other-user".to_string(), @@ -223,7 +228,7 @@ async fn missing_or_different_reconciliation_results_keep_stored_usage_reconcili ..Default::default() }; write(&store, event(), direct).await; - assert_eq!(store.reconciliations.lock().unwrap().len(), 2); + assert_eq!(store.reconciliations.lock().unwrap().len(), 1); assert_eq!(store.settlements.lock().unwrap().len(), 1); } } @@ -252,19 +257,56 @@ async fn changed_stored_usage_is_reconciled_using_its_own_identity_cost_and_term }; write(&store, event(), direct).await; let reconciliations = store.reconciliations.lock().unwrap(); - assert_eq!(reconciliations.len(), 2); - assert_eq!(reconciliations[1].request_id, stored.request_id); + let stored_token = stored.request_metadata.as_ref().unwrap() + ["plan_usage_reservation_token"] + .as_str() + .unwrap(); assert_eq!( - reconciliations[1].subject_id, + reconciliations.len(), + if stored_token == RESERVATION_TOKEN { + 1 + } else { + 2 + } + ); + if stored_token != RESERVATION_TOKEN { + assert_eq!(reconciliations[0].reservation_token, RESERVATION_TOKEN); + assert_eq!(reconciliations[0].actual_cost_units, 75_000_000); + } + let reconciliation = reconciliations.last().unwrap(); + assert_eq!(reconciliation.request_id, stored.request_id); + assert_eq!( + reconciliation.subject_id, stored.user_id.as_ref().unwrap().as_str() ); assert_eq!( - reconciliations[1].reservation_token, + reconciliation.reservation_token, stored.request_metadata.as_ref().unwrap()["plan_usage_reservation_token"] .as_str() .unwrap() ); - assert_ne!(reconciliations[0], reconciliations[1]); + assert_eq!( + reconciliation.actual_cost_units, + if stored.status == "failed" { + 0 + } else { + (stored.actual_total_cost_usd * 100_000_000.0) as u64 + } + ); + assert_eq!( + reconciliation.terminal_state, + if stored.status == "failed" { + UsagePolicyCostReservationState::Released + } else { + UsagePolicyCostReservationState::Finalized + } + ); + assert_eq!( + reconciliation.finalized_at_unix_secs, + stored + .finalized_at_unix_secs + .unwrap_or(stored.updated_at_unix_secs) + ); let settlements = store.settlements.lock().unwrap(); assert_eq!(settlements.len(), 1); assert_eq!( @@ -276,6 +318,185 @@ async fn changed_stored_usage_is_reconciled_using_its_own_identity_cost_and_term } } +#[tokio::test] +async fn worker_and_direct_settle_customer_multiplier_snapshot_without_changing_provider_cost() { + for direct in [false, true] { + for (group_multiplier, promotion_multiplier, expected) in [ + (0.0, 1.0, 0.0), + (1.0, 1.0, 1.25), + (2.0, 0.25, 0.625), + (3.0, 1.0, 3.75), + ] { + let store = ReuseStore::default(); + let mut event = event(); + event.data.request_metadata.as_mut().unwrap()["billing_multiplier_snapshot"] = json!({ + "version": 1, + "factors": {"routing_group": group_multiplier, "promotion": promotion_multiplier}, + "multiplier": group_multiplier * promotion_multiplier + }); + write(&store, event, direct).await; + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), 1, "direct={direct}"); + assert_eq!( + reconciliations[0].actual_cost_units, + (expected * 100_000_000.0) as u64 + ); + let settlements = store.settlements.lock().unwrap(); + assert_eq!(settlements.len(), 1); + assert_eq!(settlements[0].billing_cost_usd, Some(expected)); + assert_eq!(settlements[0].total_cost_usd, 1.25); + assert_eq!(settlements[0].actual_total_cost_usd, 0.75); + } + } +} + +#[tokio::test] +async fn sparse_terminal_event_uses_persisted_multiplier_and_cost_for_both_ledger_and_wallet() { + for direct in [false, true] { + let mut stored = sample_usage(); + stored.total_cost_usd = 2.0; + stored.actual_total_cost_usd = 0.25; + stored.request_metadata.as_mut().unwrap()["billing_multiplier_snapshot"] = json!({ + "version": 1, + "factors": {"routing_group": 3.0, "promotion": 0.5}, + "multiplier": 1.5 + }); + let expected = stored.billing_cost().unwrap(); + assert_eq!(expected, 3.0); + let store = ReuseStore { + stored_override: Some(stored), + ..Default::default() + }; + // Sparse asynchronous completion has neither the captured multiplier + // nor authoritative charges. Persistence restores the original snapshot. + let mut terminal = event(); + terminal.data.total_cost_usd = None; + terminal.data.actual_total_cost_usd = None; + write(&store, terminal, direct).await; + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), 1, "direct={direct}"); + assert_eq!(reconciliations[0].actual_cost_units, 300_000_000); + let settlements = store.settlements.lock().unwrap(); + assert_eq!(settlements.len(), 1); + assert_eq!(settlements[0].billing_cost_usd, Some(expected)); + assert_eq!(settlements[0].total_cost_usd, 2.0); + assert_eq!(settlements[0].actual_total_cost_usd, 0.25); + } +} + +#[tokio::test] +async fn colliding_request_id_reconciles_new_token_with_its_own_snapshot_after_upsert() { + for direct in [false, true] { + let mut stored = sample_usage(); + stored.total_cost_usd = 2.0; + stored.request_metadata = Some(json!({ + "plan_usage_reservation_token": "previous-token", + "billing_multiplier_snapshot": { + "version": 1, + "factors": {"routing_group": 0.5}, + "multiplier": 0.5 + } + })); + let store = ReuseStore { + stored_override: Some(stored), + ..Default::default() + }; + let mut terminal = event(); + terminal.data.request_metadata.as_mut().unwrap()["billing_multiplier_snapshot"] = json!({ + "version": 1, + "factors": {"routing_group": 3.0}, + "multiplier": 3.0 + }); + write(&store, terminal, direct).await; + assert_eq!(store.upserts.load(Ordering::Relaxed), 1); + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), 2, "direct={direct}"); + assert_eq!(reconciliations[0].reservation_token, RESERVATION_TOKEN); + assert_eq!(reconciliations[0].actual_cost_units, 375_000_000); + assert_eq!(reconciliations[1].reservation_token, "previous-token"); + assert_eq!(reconciliations[1].actual_cost_units, 100_000_000); + let settlements = store.settlements.lock().unwrap(); + assert_eq!(settlements.len(), 1); + assert_eq!(settlements[0].billing_cost_usd, Some(1.0)); + assert_eq!(settlements[0].actual_total_cost_usd, 0.75); + } +} + +#[tokio::test] +async fn sparse_colliding_token_cannot_borrow_another_requests_multiplier() { + for direct in [false, true] { + for (event_type, billable_cancel) in [ + (UsageEventType::Completed, false), + (UsageEventType::Cancelled, true), + ] { + let mut stored = sample_usage(); + stored.request_metadata = Some(json!({ + "plan_usage_reservation_token": "previous-token", + "billing_multiplier_snapshot": { + "version": 1, + "factors": {"routing_group": 0.0}, + "multiplier": 0.0 + } + })); + let store = ReuseStore { + stored_override: Some(stored), + ..Default::default() + }; + let mut terminal = event(); + terminal.event_type = event_type; + terminal.data.request_metadata.as_mut().unwrap()["cancelled_request_fee"] = + json!(billable_cancel); + if direct { + write(&store, terminal, true).await; + } else { + assert!(write_event_record(&store, &terminal).await.is_err()); + } + assert_eq!(store.upserts.load(Ordering::Relaxed), 1); + assert!( + store.reconciliations.lock().unwrap().is_empty(), + "direct={direct}" + ); + assert!(store.settlements.lock().unwrap().is_empty()); + } + } +} + +#[tokio::test] +async fn failed_colliding_token_is_released_without_borrowing_snapshot_or_charge() { + for direct in [false, true] { + let mut stored = sample_usage(); + stored.billing_status = "settled".to_string(); + stored.request_metadata = Some(json!({ + "plan_usage_reservation_token": "previous-token", + "billing_multiplier_snapshot": { + "version": 1, + "factors": {"routing_group": 2.0}, + "multiplier": 2.0 + } + })); + let store = ReuseStore { + stored_override: Some(stored), + ..Default::default() + }; + let mut terminal = event(); + terminal.event_type = UsageEventType::Failed; + terminal.data.total_cost_usd = None; + terminal.data.actual_total_cost_usd = None; + write(&store, terminal, direct).await; + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), 2, "direct={direct}"); + assert_eq!(reconciliations[0].reservation_token, RESERVATION_TOKEN); + assert_eq!( + reconciliations[0].terminal_state, + UsagePolicyCostReservationState::Released + ); + assert_eq!(reconciliations[0].actual_cost_units, 0); + assert_eq!(reconciliations[1].reservation_token, "previous-token"); + assert_eq!(reconciliations[1].actual_cost_units, 250_000_000); + assert!(store.settlements.lock().unwrap().is_empty()); + } +} + #[tokio::test] async fn cancellation_release_billable_cancellation_and_zero_cost_preserve_settlement_rules() { for direct in [false, true] { @@ -329,7 +550,7 @@ async fn cancellation_release_billable_cancellation_and_zero_cost_preserve_settl } #[tokio::test] -async fn reconciliation_failure_stops_both_writes_before_upsert_and_wallet_settlement() { +async fn reconciliation_failure_keeps_durable_usage_but_stops_wallet_settlement() { for direct in [false, true] { let store = ReuseStore { response: ReconcileResponse::Error, @@ -341,13 +562,13 @@ async fn reconciliation_failure_stops_both_writes_before_upsert_and_wallet_settl assert!(write_event_record(&store, &event()).await.is_err()); } assert_eq!(store.reconciliations.lock().unwrap().len(), 1); - assert_eq!(store.upserts.load(Ordering::Relaxed), 0); + assert_eq!(store.upserts.load(Ordering::Relaxed), 1); assert!(store.settlements.lock().unwrap().is_empty()); } } #[tokio::test] -async fn retry_after_upsert_failure_reconciles_again_before_settling() { +async fn retry_after_upsert_failure_reconciles_only_the_successfully_persisted_usage() { for direct in [false, true] { let store = ReuseStore { fail_next_upsert: AtomicBool::new(true), @@ -358,9 +579,10 @@ async fn retry_after_upsert_failure_reconciles_again_before_settling() { } else { assert!(write_event_record(&store, &event()).await.is_err()); } + assert!(store.reconciliations.lock().unwrap().is_empty()); assert!(store.settlements.lock().unwrap().is_empty()); write(&store, event(), direct).await; - assert_eq!(store.reconciliations.lock().unwrap().len(), 2); + assert_eq!(store.reconciliations.lock().unwrap().len(), 1); assert_eq!(store.upserts.load(Ordering::Relaxed), 2); assert_eq!(store.settlements.lock().unwrap().len(), 1); } @@ -481,11 +703,17 @@ async fn concurrent_duplicate_delivery_debits_real_memory_wallet_only_once() { let store = store.clone(); let runtime = runtime.clone(); tasks.spawn(async move { + let mut terminal = event(); + terminal.data.request_metadata.as_mut().unwrap()["billing_multiplier_snapshot"] = json!({ + "version": 1, + "factors": {"routing_group": 3.0, "promotion": 0.5}, + "multiplier": 1.5 + }); if index % 2 == 0 { - write_event_record(store.as_ref(), &event()).await.unwrap(); + write_event_record(store.as_ref(), &terminal).await.unwrap(); } else { runtime - .record_terminal_event_direct(store.as_ref(), event()) + .record_terminal_event_direct(store.as_ref(), terminal) .await; } }); @@ -495,13 +723,25 @@ async fn concurrent_duplicate_delivery_debits_real_memory_wallet_only_once() { } assert_eq!(store.reconciliations.lock().unwrap().len(), 32); assert_eq!(store.settlements.lock().unwrap().len(), 32); + assert!(store + .reconciliations + .lock() + .unwrap() + .iter() + .all(|input| input.actual_cost_units == 187_500_000)); + assert!(store + .settlements + .lock() + .unwrap() + .iter() + .all(|input| input.billing_cost_usd == Some(1.875) && input.actual_total_cost_usd == 0.75)); let wallet = wallets .find(WalletLookupKey::UserId("user-1")) .await .unwrap() .unwrap(); - assert_eq!(wallet.balance + wallet.gift_balance, 11.25); - assert_eq!(wallet.total_consumed, 0.75); + assert_eq!(wallet.balance + wallet.gift_balance, 10.125); + assert_eq!(wallet.total_consumed, 1.875); assert!(matches!( store .repository diff --git a/crates/aether-usage/runtime/src/worker.rs b/crates/aether-usage/runtime/src/worker.rs index dc98b6bbc..3af764258 100644 --- a/crates/aether-usage/runtime/src/worker.rs +++ b/crates/aether-usage/runtime/src/worker.rs @@ -16,9 +16,7 @@ use crate::queue::UsageDeadLetterOutcome; use crate::runtime::{ UsageBillingEventEnricher, UsageRuntimeAccess, UsageWorkerRecordConcurrencyGate, }; -use crate::settlement::{ - reconcile_usage_policy_cost_for_event_with_result, settle_usage_with_reconciled_cost, -}; +use crate::settlement::settle_usage_after_upsert; use crate::{ build_upsert_usage_record_from_event, UsageEvent, UsageEventType, UsageQueue, UsageRuntimeConfig, UsageSettlementWriter, @@ -796,10 +794,11 @@ pub async fn write_event_record(data: &T, event: &UsageEvent) -> Result<(), D where T: UsageRecordWriter + UsageSettlementWriter + Send + Sync, { - let reconciled = reconcile_usage_policy_cost_for_event_with_result(data, event).await?; let record = build_upsert_usage_record_from_event(event)?; if let Some(stored) = data.upsert_usage_record(record).await? { - settle_usage_with_reconciled_cost(data, &stored, reconciled).await?; + // Sparse terminal events may omit pricing factors. Only the stored request + // snapshot is authoritative before finalizing an immutable cost reservation. + settle_usage_after_upsert(data, &stored, event).await?; } // Manual proxy traffic is counted at the actual transport-attempt boundary. Usage events are // replayable, so emitting that side effect here would count normal requests and reclaims twice. @@ -1240,49 +1239,49 @@ mod tests { .lock() .expect("records lock") .push(record.clone()); - Ok(Some( - StoredRequestUsageAudit::new( - "usage-1".to_string(), - record.request_id, - record.user_id, - record.api_key_id, - record.username, - record.api_key_name, - record.provider_name, - record.model, - record.target_model, - record.provider_id, - record.provider_endpoint_id, - record.provider_api_key_id, - record.request_type, - record.api_format, - record.api_family, - record.endpoint_kind, - record.endpoint_api_format, - record.provider_api_family, - record.provider_endpoint_kind, - record.has_format_conversion.unwrap_or(false), - record.is_stream.unwrap_or(false), - record.input_tokens.unwrap_or_default() as i32, - record.output_tokens.unwrap_or_default() as i32, - record.total_tokens.unwrap_or_default() as i32, - record.total_cost_usd.unwrap_or_default(), - record.actual_total_cost_usd.unwrap_or_default(), - record.status_code.map(i32::from), - record.error_message, - record.error_category, - record.response_time_ms.map(|value| value as i32), - record.first_byte_time_ms.map(|value| value as i32), - record.status, - record.billing_status, - record - .created_at_unix_ms - .unwrap_or(record.updated_at_unix_secs) as i64, - record.updated_at_unix_secs as i64, - record.finalized_at_unix_secs.map(|value| value as i64), - ) - .expect("stored usage should build"), - )) + let mut stored = StoredRequestUsageAudit::new( + "usage-1".to_string(), + record.request_id, + record.user_id, + record.api_key_id, + record.username, + record.api_key_name, + record.provider_name, + record.model, + record.target_model, + record.provider_id, + record.provider_endpoint_id, + record.provider_api_key_id, + record.request_type, + record.api_format, + record.api_family, + record.endpoint_kind, + record.endpoint_api_format, + record.provider_api_family, + record.provider_endpoint_kind, + record.has_format_conversion.unwrap_or(false), + record.is_stream.unwrap_or(false), + record.input_tokens.unwrap_or_default() as i32, + record.output_tokens.unwrap_or_default() as i32, + record.total_tokens.unwrap_or_default() as i32, + record.total_cost_usd.unwrap_or_default(), + record.actual_total_cost_usd.unwrap_or_default(), + record.status_code.map(i32::from), + record.error_message, + record.error_category, + record.response_time_ms.map(|value| value as i32), + record.first_byte_time_ms.map(|value| value as i32), + record.status, + record.billing_status, + record + .created_at_unix_ms + .unwrap_or(record.updated_at_unix_secs) as i64, + record.updated_at_unix_secs as i64, + record.finalized_at_unix_secs.map(|value| value as i64), + ) + .expect("stored usage should build"); + stored.request_metadata = record.request_metadata; + Ok(Some(stored)) } } @@ -1510,7 +1509,7 @@ mod tests { } #[tokio::test] - async fn same_request_id_terminal_events_reconcile_each_reservation_token_before_upsert() { + async fn same_request_id_terminal_events_reconcile_each_persisted_reservation_token() { let store = TestUsageStore::default(); let mut first = sample_event(); first.request_id = "shared-client-trace".to_string(); @@ -1619,7 +1618,12 @@ mod tests { worker.queue.ensure_consumer_group().await.expect("group"); let mut event = sample_event(); event.data.request_metadata = Some(serde_json::json!({ - "plan_usage_reservation_token": "pricing-retry-reservation" + "plan_usage_reservation_token": "pricing-retry-reservation", + "billing_multiplier_snapshot": { + "version": 1, + "factors": {"routing_group": 3.0, "promotion": 0.5}, + "multiplier": 1.5 + } })); worker.queue.enqueue(&event).await.expect("enqueue"); let batch = worker @@ -1676,17 +1680,27 @@ mod tests { assert_eq!(records[0].total_cost_usd, Some(0.456)); assert_eq!(records[0].actual_total_cost_usd, Some(0.123)); assert_eq!(records[0].total_tokens, Some(10)); + assert_eq!( + records[0].request_metadata.as_ref().unwrap()["billing_multiplier_snapshot"] + ["multiplier"], + 1.5 + ); } { let reconciliations = store.reconciliations.lock().expect("reconciliations lock"); assert_eq!(reconciliations.len(), 1); - assert_eq!(reconciliations[0].actual_cost_units, 12_300_000); + assert_eq!(reconciliations[0].actual_cost_units, 68_400_000); assert_eq!( reconciliations[0].reservation_token, "pricing-retry-reservation" ); } - assert_eq!(store.settlements.lock().expect("settlements lock").len(), 1); + { + let settlements = store.settlements.lock().expect("settlements lock"); + assert_eq!(settlements.len(), 1); + assert_eq!(settlements[0].billing_cost_usd, Some(0.456 * 1.5)); + assert_eq!(settlements[0].actual_total_cost_usd, 0.123); + } assert_eq!( store.enrich_calls.lock().expect("enrich calls lock").len(), 2 diff --git a/frontend/src/api/dashboard.ts b/frontend/src/api/dashboard.ts index 4d36ba1f8..30d74b9fa 100644 --- a/frontend/src/api/dashboard.ts +++ b/frontend/src/api/dashboard.ts @@ -250,6 +250,10 @@ export interface RequestDetail { output_cost?: number total_cost?: number actual_cost?: number + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + billing_cost?: number | null cache_creation_cost?: number cache_read_cost?: number image_output_cost?: number diff --git a/frontend/src/api/me.ts b/frontend/src/api/me.ts index b1509c05f..f0f65e93e 100644 --- a/frontend/src/api/me.ts +++ b/frontend/src/api/me.ts @@ -27,6 +27,13 @@ export interface Profile { feature_settings?: FeatureSettingsMap | null } +export interface UserRoutingGroup { + id: string + name: string + billing_multiplier: number + is_default: boolean +} + export interface UserPreferences { avatar_url?: string bio?: string @@ -66,7 +73,11 @@ export interface UsageRecordDetail { output_tokens: number total_tokens: number cost: number // 官方费率 - actual_cost?: number // 倍率消耗(仅管理员可见) + actual_cost?: number // 提供商 Key 成本(仅管理员可见);旧记录也用于兼容历史扣费 + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + billing_cost?: number | null rate_multiplier?: number // 成本倍率(仅管理员可见) response_time_ms?: number | null first_byte_time_ms?: number | null @@ -196,6 +207,8 @@ export interface ApiKey { allowed_providers?: ProviderConfig[] force_capabilities?: Record | null // 强制能力配置 feature_settings?: FeatureSettingsMap | null + routing_group_id?: string | null + routing_group_name?: string | null } export type InstallTargetCli = 'claude_code' | 'codex_cli' | 'gemini_cli' @@ -278,7 +291,12 @@ export const meApi = { return response.data }, - async createApiKey(data: { name: string; rate_limit?: number | null; concurrent_limit?: number | null; ip_rules?: string[] | null; feature_settings?: FeatureSettingsMap | null }): Promise { + async getRoutingGroups(): Promise<{ items: UserRoutingGroup[]; total: number }> { + const response = await apiClient.get<{ items: UserRoutingGroup[]; total: number }>('/api/users/me/routing-groups') + return response.data + }, + + async createApiKey(data: { name: string; rate_limit?: number | null; concurrent_limit?: number | null; ip_rules?: string[] | null; feature_settings?: FeatureSettingsMap | null; routing_group_id?: string | null }): Promise { const response = await apiClient.post('/api/users/me/api-keys', data) return response.data }, @@ -318,7 +336,7 @@ export const meApi = { async updateApiKey( keyId: string, - data: { name?: string; rate_limit?: number | null; concurrent_limit?: number | null; ip_rules?: string[] | null; feature_settings?: FeatureSettingsMap | null | undefined } + data: { name?: string; rate_limit?: number | null; concurrent_limit?: number | null; ip_rules?: string[] | null; feature_settings?: FeatureSettingsMap | null | undefined; routing_group_id?: string | null } ): Promise { const response = await apiClient.put( `/api/users/me/api-keys/${keyId}`, @@ -370,6 +388,10 @@ export const meApi = { cost: number actual_cost?: number | null rate_multiplier?: number | null + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + billing_cost?: number | null response_time_ms: number | null first_byte_time_ms: number | null end_to_end_time_ms?: number | null @@ -417,6 +439,10 @@ export const meApi = { cost: number actual_cost?: number | null rate_multiplier?: number | null + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + billing_cost?: number | null response_time_ms: number | null first_byte_time_ms: number | null end_to_end_time_ms?: number | null diff --git a/frontend/src/api/usage.ts b/frontend/src/api/usage.ts index e6bd58359..9fe22ad37 100644 --- a/frontend/src/api/usage.ts +++ b/frontend/src/api/usage.ts @@ -30,6 +30,10 @@ export interface UsageRecord { cache_read_input_tokens?: number total_tokens: number cost?: number + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + billing_cost?: number | null response_time?: number response_time_ms?: number | null first_byte_time_ms?: number | null @@ -617,6 +621,10 @@ export const usageApi = { cost: number actual_cost?: number | null rate_multiplier?: number | null + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + billing_cost?: number | null response_time_ms: number | null first_byte_time_ms: number | null end_to_end_time_ms?: number | null @@ -689,6 +697,10 @@ export const usageApi = { cost: number actual_cost?: number | null rate_multiplier?: number | null + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + billing_cost?: number | null response_time_ms: number | null first_byte_time_ms: number | null end_to_end_time_ms?: number | null diff --git a/frontend/src/api/usageRecords.ts b/frontend/src/api/usageRecords.ts index 24afc8c38..59e1d76d3 100644 --- a/frontend/src/api/usageRecords.ts +++ b/frontend/src/api/usageRecords.ts @@ -39,6 +39,12 @@ export interface UsageRecord { total_tokens: number cost: number actual_cost?: number + /** Combined customer billing multiplier captured for this request. */ + billing_multiplier?: number | null + routing_group_id?: string | null + routing_group_name?: string | null + /** Captured customer charge. Null means unavailable and must not be recalculated. */ + billing_cost?: number | null response_time_ms?: number | null first_byte_time_ms?: number | null // 首字时间 (TTFB) end_to_end_time_ms?: number | null // 客户端从请求进入网关到完成的总耗时 diff --git a/frontend/src/components/ui/popover/PopoverContent.vue b/frontend/src/components/ui/popover/PopoverContent.vue index 94dabff4d..a56cef60f 100644 --- a/frontend/src/components/ui/popover/PopoverContent.vue +++ b/frontend/src/components/ui/popover/PopoverContent.vue @@ -19,6 +19,10 @@ const props = withDefaults(defineProps<{ collisionPadding: 0, ariaLabel: undefined, }) + +const emit = defineEmits<{ + openAutoFocus: [event: Event] +}>() - +
- - {{ formatRecordUserSegment(record) }} - - · - {{ formatRecordProviderSegment(record) }} + {{ formatRecordUserSegment(record) }} +
+
+
-
- {{ record.provider }} - - {{ record.provider_key_name }} - ({{ record.rate_multiplier }}x) - -
+ -
- {{ formatCurrency(record.cost || 0) }} - - {{ formatCurrency(record.actual_cost) }} - -
-
- 不可用 -
-
- 未计价 -
+ props.filterSearch, (value) => { if (value !== localSearch.value) { cancelPendingSearchEmit() diff --git a/frontend/src/features/usage/components/__tests__/RequestDetailDrawer.pricing.spec.ts b/frontend/src/features/usage/components/__tests__/RequestDetailDrawer.pricing.spec.ts index 53685b3f4..2f80d7202 100644 --- a/frontend/src/features/usage/components/__tests__/RequestDetailDrawer.pricing.spec.ts +++ b/frontend/src/features/usage/components/__tests__/RequestDetailDrawer.pricing.spec.ts @@ -113,6 +113,48 @@ function buildFastTierDetail(): RequestDetail { } describe('RequestDetailDrawer settlement pricing', () => { + it.each([ + { billingMultiplier: 0, billingCost: 0 }, + { billingMultiplier: undefined, billingCost: undefined }, + { billingMultiplier: 2, billingCost: null }, + ])('preserves zero, missing, and explicitly unavailable billing facts in list updates: %o', async ({ billingMultiplier, billingCost }) => { + apiMocks.getRequestDetail.mockResolvedValue({ + ...buildEmbeddingDetail(), + billing_multiplier: billingMultiplier, + billing_cost: billingCost, + routing_group_id: 'group-1', + routing_group_name: '历史分组', + actual_cost: 0.000005, + } satisfies RequestDetail) + const updates = vi.fn() + let isOpen!: Ref + const root = document.createElement('div') + document.body.appendChild(root) + const app = createApp({ + setup() { + isOpen = ref(false) + return () => h(RequestDetailDrawer, { + isOpen: isOpen.value, + requestId: 'usage-embedding-1', + onRequestState: updates, + }) + }, + }) + app.mount(root) + mountedApps.push({ app, root }) + isOpen.value = true + await nextTick() + await vi.waitFor(() => expect(updates).toHaveBeenCalledWith(expect.objectContaining({ + id: 'usage-embedding-1', + cost: 0.00001, + actualCost: 0.000005, + billingMultiplier, + billingCost, + routingGroupId: 'group-1', + routingGroupName: '历史分组', + }))) + }) + it('labels an unmetered OpenAI Live WebSocket detail without rendering zero usage as billing', async () => { apiMocks.getRequestDetail.mockResolvedValue({ ...buildEmbeddingDetail(), diff --git a/frontend/src/features/usage/components/__tests__/UsageRecordsTable.spec.ts b/frontend/src/features/usage/components/__tests__/UsageRecordsTable.spec.ts index 6d2687e30..8be725809 100644 --- a/frontend/src/features/usage/components/__tests__/UsageRecordsTable.spec.ts +++ b/frontend/src/features/usage/components/__tests__/UsageRecordsTable.spec.ts @@ -188,6 +188,100 @@ afterEach(() => { }) describe('UsageRecordsTable', () => { + it.each([ + [{ routing_group_name: '生产策略', routing_group_id: 'group-1' }, '生产策略'], + [{ routing_group_name: ' ', routing_group_id: 'group-history' }, 'group-history'], + [{}, '未记录分组'], + ])('shows group then the provider Key in desktop and mobile layouts', (group, expected) => { + const root = mountUsageRecordsTable([buildRecord({ + ...group, + provider: '上游 A', + provider_key_name: '供应商 Key', + api_key: { id: 'user-key', name: '用户 API Key', display: 'sk-user' }, + })]) + expect([...root.querySelectorAll('[data-usage-provider="routing-group"]')].map(element => element.textContent?.trim())) + .toEqual([expected, expected]) + const providerLines = [...root.querySelectorAll('[data-usage-provider="provider-key"]')] + expect(providerLines.map(element => element.textContent?.trim())).toEqual(['上游 A · 供应商 Key', '上游 A · 供应商 Key']) + for (const providerLine of providerLines) { + expect(providerLine.previousElementSibling?.getAttribute('data-usage-provider')).toBe('routing-group') + expect(providerLine.title).toBe('上游 A · 供应商 Key') + } + }) + + it('does not substitute a user API Key when the provider Key is unavailable', () => { + const root = mountUsageRecordsTable([buildRecord({ + provider: '上游 A', + api_key: { id: 'user-key', name: '用户 API Key', display: 'sk-user' }, + })]) + expect([...root.querySelectorAll('[data-usage-provider="provider-key"]')].map(element => element.textContent?.trim())) + .toEqual(['上游 A · -', '上游 A · -']) + }) + + it.each([ + [0, '$0.00'], + [1, '$10.00'], + [0.5, '$5.00'], + [2, '$20.00'], + [undefined, '$3.00'], + ])('shows customer charges in both layouts for multiplier %s, including legacy charges', (multiplier, expected) => { + const root = mountUsageRecordsTable([buildRecord({ + cost: 10, + actual_cost: 3, + rate_multiplier: 0.3, + billing_multiplier: multiplier as number | undefined, + })], { isAdmin: false }) + const subtitles = [...root.querySelectorAll('[data-usage-cost="routing-group"]')] + expect(subtitles).toHaveLength(2) + expect(subtitles.map(element => element.textContent?.trim())).toEqual([expected, expected]) + expect(subtitles.every(element => element.getAttribute('title') === '实际扣费')).toBe(true) + expect([...root.querySelectorAll('[data-usage-cost="base"]')].map(element => element.textContent?.trim())).toEqual(['$10.00', '$10.00']) + expect(root.querySelector('[data-usage-cost="provider-key"]')).toBeNull() + expect(root.querySelector('[data-usage-provider]')).toBeNull() + }) + + it('uses the historical customer charge independently of the administrator-only provider Key cost', () => { + const root = mountUsageRecordsTable([buildRecord({ + cost: 10, + actual_cost: 3, + rate_multiplier: 0.3, + billing_multiplier: 2, + billing_cost: 19.99, + })], { showActualCost: true }) + const subtitles = [...root.querySelectorAll('[data-usage-cost="routing-group"]')] + expect(subtitles.map(element => element.textContent?.trim())).toEqual(['$19.99', '$19.99']) + const keyCosts = [...root.querySelectorAll('[data-usage-cost="provider-key"]')] + expect(keyCosts.map(element => element.textContent?.trim())).toEqual(['$3.00', '$3.00']) + expect(keyCosts.every(element => element.title.includes('提供商 Key 成本'))).toBe(true) + }) + + it('does not replace an unavailable customer charge with the base or provider Key cost', () => { + const root = mountUsageRecordsTable([buildRecord({ + cost: 10, + actual_cost: 3, + rate_multiplier: 0.3, + billing_multiplier: 2, + billing_cost: null, + })], { showActualCost: true }) + expect(root.querySelector('[data-usage-cost="routing-group"]')).toBeNull() + expect([...root.querySelectorAll('[data-usage-cost="base"]')].map(element => element.textContent?.trim())) + .toEqual(['$10.00', '$10.00']) + expect([...root.querySelectorAll('[data-usage-cost="provider-key"]')].map(element => element.textContent?.trim())) + .toEqual(['$3.00', '$3.00']) + }) + + it.each(['usage_available', 'usage_pricing_available'] as const)('hides all numeric costs when %s is false', field => { + const root = mountUsageRecordsTable([buildRecord({ + [field]: false, + actual_cost: 3, + rate_multiplier: 0.3, + billing_multiplier: 2, + billing_cost: 0.02, + })], { showActualCost: true }) + expect(root.querySelector('[data-usage-cost]')).toBeNull() + expect(root.querySelectorAll(field === 'usage_available' ? '[data-usage-unavailable="cost"]' : '[data-usage-unpriced="cost"]')).toHaveLength(2) + }) + it('shows output TPS after the request completes', () => { const root = mountUsageRecordsTable([buildRecord()]) diff --git a/frontend/src/features/usage/composables/__tests__/useUsageData.spec.ts b/frontend/src/features/usage/composables/__tests__/useUsageData.spec.ts index 0aadc5cc7..6c91ba6e3 100644 --- a/frontend/src/features/usage/composables/__tests__/useUsageData.spec.ts +++ b/frontend/src/features/usage/composables/__tests__/useUsageData.spec.ts @@ -49,6 +49,7 @@ vi.mock('@/utils/logger', () => ({ import { useUsageData } from '../useUsageData' import type { UsageRecord } from '../../types' +import { resolveUsageBilling } from '../../utils/usageBilling' function buildUsageRecord(overrides: Partial = {}): UsageRecord { return { @@ -522,6 +523,47 @@ describe('useUsageData', () => { }) }) + it('preserves billing snapshots across sparse refreshes but accepts newer zero or unavailable amounts', async () => { + const { loadRecords, currentRecords } = useUsageData({ isAdminPage: ref(true) }) + async function refresh(overrides: Partial) { + getAllUsageRecordsMock.mockResolvedValueOnce({ + records: [buildUsageRecord(overrides)], total: 1, limit: 20, offset: 0, + }) + await loadRecords({ page: 1, pageSize: 20 }) + } + await refresh({ updated_at: '2026-05-01T00:00:02Z', billing_multiplier: 2, billing_cost: 0.02, routing_group_id: 'group-1', routing_group_name: '历史分组' }) + await refresh({ updated_at: '2026-05-01T00:00:03Z' }) + expect(currentRecords.value[0]).toMatchObject({ billing_multiplier: 2, billing_cost: 0.02, routing_group_id: 'group-1', routing_group_name: '历史分组' }) + await refresh({ updated_at: '2026-05-01T00:00:01Z', billing_multiplier: 0, billing_cost: 0, routing_group_name: '过时分组' }) + expect(currentRecords.value[0]).toMatchObject({ billing_multiplier: 2, billing_cost: 0.02, routing_group_id: 'group-1', routing_group_name: '历史分组' }) + await refresh({ updated_at: '2026-05-01T00:00:04Z', billing_multiplier: 0, billing_cost: 0 }) + expect(currentRecords.value[0]).toMatchObject({ cost: 0.01, billing_multiplier: 0, billing_cost: 0 }) + await refresh({ updated_at: '2026-05-01T00:00:05Z', billing_multiplier: null, billing_cost: null }) + expect(currentRecords.value[0]).toMatchObject({ billing_multiplier: 0, billing_cost: null }) + await refresh({ updated_at: '2026-05-01T00:00:06Z', cost: 0.02, billing_multiplier: 2 }) + expect(resolveUsageBilling(currentRecords.value[0])).toEqual({ multiplier: 2, cost: null }) + }) + + it('recalculates a completed group cost when a mixed-version list omits the amount and resets costs for a changed group', async () => { + const { loadRecords, currentRecords } = useUsageData({ isAdminPage: ref(true) }) + async function refresh(overrides: Partial) { + getAllUsageRecordsMock.mockResolvedValueOnce({ + records: [buildUsageRecord(overrides)], total: 1, limit: 20, offset: 0, + }) + await loadRecords({ page: 1, pageSize: 20 }) + } + await refresh({ status: 'pending', cost: 0, routing_group_id: 'g1', billing_multiplier: 2, billing_cost: 0 }) + await refresh({ status: 'completed', cost: 3, routing_group_id: 'g1', billing_multiplier: 2 }) + expect(resolveUsageBilling(currentRecords.value[0])).toEqual({ cost: 6, multiplier: 2 }) + await refresh({ status: 'completed', cost: 3, routing_group_id: 'g1', billing_multiplier: 2, billing_cost: 5.99999999 }) + // A sparse list zero that the base-cost merger rejects must not erase the precise captured amount. + await refresh({ status: 'completed', cost: 0, routing_group_id: 'g1' }) + expect(currentRecords.value[0].billing_cost).toBe(5.99999999) + await refresh({ status: 'completed', cost: 3, routing_group_id: 'g2' }) + expect(resolveUsageBilling(currentRecords.value[0])).toEqual({ cost: null, multiplier: 1 }) + expect(currentRecords.value[0].billing_cost).toBeNull() + }) + it('refreshes exact admin record totals after an estimated first page', async () => { const isAdminPage = ref(true) const { loadRecords, totalRecords } = useUsageData({ isAdminPage }) diff --git a/frontend/src/features/usage/composables/useUsageData.ts b/frontend/src/features/usage/composables/useUsageData.ts index c4842c382..e729161c9 100644 --- a/frontend/src/features/usage/composables/useUsageData.ts +++ b/frontend/src/features/usage/composables/useUsageData.ts @@ -11,6 +11,7 @@ import type { EnhancedModelStatsItem } from '../types' import { createDefaultStats } from '../types' +import { mergeUsageBillingSnapshot } from '../utils/usageBilling' import { log } from '@/utils/logger' import { getErrorStatus } from '@/types/api-error' import { isUsageProviderVisible, normalizeUsageProviderStats } from '../utils/providerStats' @@ -610,8 +611,10 @@ export function useUsageData(options: UseUsageDataOptions) { { preferNext: nextTimingIsAuthoritative }, ) + const mergedCost = mergeSparseRecordMetric(existing.cost, record.cost) ?? record.cost return { ...record, + ...mergeUsageBillingSnapshot(existing, { ...record, cost: mergedCost }, statusProgressed), // 保留详情抽屉/活跃轮询已经拿到的完整指标,避免列表刷新用 0 或空值回退。 status: mergedStatus, provider: statusProgressed @@ -634,7 +637,7 @@ export function useUsageData(options: UseUsageDataOptions) { record.cache_creation_ephemeral_1h_input_tokens ) ?? record.cache_creation_ephemeral_1h_input_tokens, cache_read_input_tokens: mergeSparseRecordMetric(existing.cache_read_input_tokens, record.cache_read_input_tokens) ?? record.cache_read_input_tokens, - cost: mergeSparseRecordMetric(existing.cost, record.cost) ?? record.cost, + cost: mergedCost, actual_cost: mergeSparseRecordMetric(existing.actual_cost, record.actual_cost) ?? record.actual_cost, response_time_ms: responseTiming.response_time_ms, first_byte_time_ms: mergeUsageRecordFirstByteTimeMs( diff --git a/frontend/src/features/usage/utils/__tests__/usageBilling.spec.ts b/frontend/src/features/usage/utils/__tests__/usageBilling.spec.ts new file mode 100644 index 000000000..6d748112e --- /dev/null +++ b/frontend/src/features/usage/utils/__tests__/usageBilling.spec.ts @@ -0,0 +1,98 @@ +import { describe, expect, it } from 'vitest' +import { mergeUsageBillingSnapshot, resolveUsageBilling } from '../usageBilling' + +describe('request billing snapshots', () => { + it.each([undefined, null, -1, NaN, Infinity])('uses the historical default for invalid or absent multiplier %s', multiplier => { + expect(resolveUsageBilling({ cost: 2, actual_cost: 0.5, billing_multiplier: multiplier })) + .toEqual({ cost: 0.5, multiplier: 1 }) + }) + + it('does not invent a historical charge when actual_cost is unavailable', () => { + expect(resolveUsageBilling({ cost: 2 })).toEqual({ cost: null, multiplier: 1 }) + expect(resolveUsageBilling({ cost: 2, actual_cost: 0 })).toEqual({ cost: 0, multiplier: 1 }) + }) + + it.each([null, NaN, Infinity, -1])('does not recalculate an explicit unavailable billing cost %s', billingCost => { + expect(resolveUsageBilling({ cost: 2, actual_cost: 0.5, billing_multiplier: 3, billing_cost: billingCost })) + .toEqual({ cost: null, multiplier: 3 }) + }) + + it('keeps an overflowed or unavailable base-cost product unknown', () => { + expect(resolveUsageBilling({ cost: Number.MAX_VALUE, billing_multiplier: 2 })) + .toEqual({ cost: null, multiplier: 2 }) + expect(resolveUsageBilling({ cost: NaN, billing_multiplier: 0 })) + .toEqual({ cost: null, multiplier: 0 }) + }) + + it('preserves a captured zero cost instead of recalculating it', () => { + expect(resolveUsageBilling({ cost: 2, billing_multiplier: 3, billing_cost: 0 })) + .toEqual({ cost: 0, multiplier: 3 }) + }) + + it('keeps a snapshot across sparse or stale updates and accepts authoritative zeroes', () => { + const existing = { billing_multiplier: 2, billing_cost: 4 } + expect(mergeUsageBillingSnapshot(existing, {})).toEqual(existing) + expect(mergeUsageBillingSnapshot(existing, { billing_multiplier: null })).toEqual(existing) + const free = { billing_multiplier: 0, billing_cost: 0 } + expect(mergeUsageBillingSnapshot(existing, free, false)).toEqual(existing) + expect(mergeUsageBillingSnapshot(existing, free)).toEqual(free) + }) + + it.each([null, NaN, Infinity, -1])('retains an authoritative unavailable amount %s across sparse refreshes', billingCost => { + const unavailable = mergeUsageBillingSnapshot( + { cost: 2, billing_multiplier: 2, billing_cost: 4 }, + { billing_cost: billingCost }, + ) + expect(unavailable.billing_cost).toBeNull() + const refreshed = mergeUsageBillingSnapshot({ ...unavailable, cost: 2 }, { cost: 3, billing_multiplier: 4 }) + expect(resolveUsageBilling({ ...refreshed, cost: 3 })).toEqual({ cost: null, multiplier: 4 }) + expect(mergeUsageBillingSnapshot(refreshed, { billing_cost: 12 }).billing_cost).toBe(12) + }) + + it('keeps historical names through sparse updates without attaching one group name to another ID', () => { + const existing = { routing_group_id: 'g1', routing_group_name: '历史分组', billing_multiplier: 2, billing_cost: 4 } + expect(mergeUsageBillingSnapshot(existing, {})).toEqual(existing) + expect(mergeUsageBillingSnapshot(existing, { routing_group_id: 'g2', routing_group_name: '新分组' }, false)).toEqual(existing) + const changedGroup = mergeUsageBillingSnapshot(existing, { routing_group_id: 'g2' }) + expect(changedGroup).toEqual({ + routing_group_id: 'g2', routing_group_name: undefined, + billing_multiplier: undefined, billing_cost: null, + }) + expect(resolveUsageBilling({ ...changedGroup, cost: 2, actual_cost: 0.5 })).toEqual({ cost: null, multiplier: 1 }) + const replacementSnapshot = mergeUsageBillingSnapshot(existing, { + routing_group_id: 'g2', billing_multiplier: 3, billing_cost: 6, + }) + expect(resolveUsageBilling({ ...replacementSnapshot, cost: 2 })).toEqual({ cost: 6, multiplier: 3 }) + }) + + it('recalculates fallback cost when a new multiplier arrives without its paired amount', () => { + const next = mergeUsageBillingSnapshot( + { billing_multiplier: 2, billing_cost: 4 }, + { billing_multiplier: 0 }, + ) + expect(resolveUsageBilling({ ...next, cost: 2 })).toEqual({ cost: 0, multiplier: 0 }) + const previouslyMissingMultiplier = mergeUsageBillingSnapshot( + { billing_cost: 4 }, + { billing_multiplier: 0 }, + ) + expect(resolveUsageBilling({ ...previouslyMissingMultiplier, cost: 2 })).toEqual({ cost: 0, multiplier: 0 }) + }) + + it('does not reuse a paired amount when another billing factor changes within the same routing group', () => { + const updated = mergeUsageBillingSnapshot( + { cost: 2, routing_group_id: 'g1', billing_multiplier: 2, billing_cost: 4 }, + { routing_group_id: 'g1', billing_multiplier: 6 }, + ) + expect(resolveUsageBilling({ ...updated, cost: 2 })).toEqual({ cost: 12, multiplier: 6 }) + }) + + it('invalidates a pending amount when a trusted base cost advances without a paired amount', () => { + const pending = { cost: 0, routing_group_id: 'g1', billing_multiplier: 2, billing_cost: 0 } + const completed = mergeUsageBillingSnapshot(pending, { cost: 3, routing_group_id: 'g1', billing_multiplier: 2 }) + expect(resolveUsageBilling({ ...completed, cost: 3 })).toEqual({ cost: 6, multiplier: 2 }) + const stale = mergeUsageBillingSnapshot(pending, { cost: 3 }, false) + expect(stale.billing_cost).toBe(0) + const paired = mergeUsageBillingSnapshot(pending, { cost: 3, billing_cost: 5.99999999 }) + expect(paired.billing_cost).toBe(5.99999999) + }) +}) diff --git a/frontend/src/features/usage/utils/usageBilling.ts b/frontend/src/features/usage/utils/usageBilling.ts new file mode 100644 index 000000000..79c116c2a --- /dev/null +++ b/frontend/src/features/usage/utils/usageBilling.ts @@ -0,0 +1,68 @@ +import type { UsageRecord } from '@/api/usageRecords' + +export type UsageBillingSnapshot = Pick< + UsageRecord, + 'billing_multiplier' | 'billing_cost' | 'routing_group_id' | 'routing_group_name' +> + +function nonnegativeFinite(value: unknown): number | undefined { + return typeof value === 'number' && Number.isFinite(value) && value >= 0 ? value : undefined +} + +export function resolveUsageBilling(record: UsageBillingSnapshot & { cost: number, actual_cost?: number | null }) { + const capturedMultiplier = nonnegativeFinite(record.billing_multiplier) + const multiplier = capturedMultiplier ?? 1 + + // A present billing_cost is authoritative, including null (for example when + // the server rejected an overflowed product). Never replace it with a guess. + if (record.billing_cost !== undefined) { + const capturedCost = nonnegativeFinite(record.billing_cost) + return { multiplier, cost: capturedCost ?? null } + } + + if (capturedMultiplier !== undefined) { + const baseCost = nonnegativeFinite(record.cost) + const derivedCost = baseCost === undefined ? null : baseCost * multiplier + return { multiplier, cost: derivedCost !== null && Number.isFinite(derivedCost) ? derivedCost : null } + } + + // Before group snapshots were introduced, actual_cost was the only captured + // charge. Falling back to the catalogue/base cost would misstate old records. + const historicalCost = nonnegativeFinite(record.actual_cost) + return { multiplier, cost: historicalCost ?? null } +} + +/** Request snapshots survive sparse updates; an accepted zero is a real value. */ +export function mergeUsageBillingSnapshot( + existing: UsageBillingSnapshot & { cost?: number | null }, + next: UsageBillingSnapshot & { cost?: number | null }, + acceptNext = true, +): UsageBillingSnapshot { + const groupId = (acceptNext ? next.routing_group_id?.trim() : undefined) || existing.routing_group_id + const sameGroup = !groupId || !existing.routing_group_id || groupId === existing.routing_group_id + const nextMultiplier = acceptNext ? nonnegativeFinite(next.billing_multiplier) : undefined + const existingMultiplier = sameGroup ? nonnegativeFinite(existing.billing_multiplier) : undefined + const multiplierChanged = nextMultiplier !== undefined && nextMultiplier !== (existingMultiplier ?? 1) + const nextBaseCost = acceptNext ? nonnegativeFinite(next.cost) : undefined + const baseCostChanged = nextBaseCost !== undefined && nextBaseCost !== nonnegativeFinite(existing.cost) + const nextCapturedCost = acceptNext && next.billing_cost !== undefined + ? (nonnegativeFinite(next.billing_cost) ?? null) + : undefined + const existingCapturedCost = existing.billing_cost !== undefined + ? (nonnegativeFinite(existing.billing_cost) ?? null) + : undefined + + return { + routing_group_id: groupId, + routing_group_name: (acceptNext ? next.routing_group_name?.trim() : undefined) + || (sameGroup ? existing.routing_group_name : undefined), + billing_multiplier: nextMultiplier ?? existingMultiplier, + billing_cost: nextCapturedCost !== undefined + ? nextCapturedCost + : !sameGroup && nextMultiplier === undefined + ? null + : sameGroup && (existingCapturedCost === null || (!multiplierChanged && !baseCostChanged)) + ? existingCapturedCost + : undefined, + } +} diff --git a/frontend/src/mocks/handler.ts b/frontend/src/mocks/handler.ts index 7c10c8136..586834e76 100644 --- a/frontend/src/mocks/handler.ts +++ b/frontend/src/mocks/handler.ts @@ -727,6 +727,7 @@ interface MockManagedUserApiKey { is_locked: boolean is_standalone: false feature_settings?: Record | null + routing_group_id?: string | null rate_limit?: number | null concurrent_limit?: number | null ip_rules?: string[] | null @@ -735,6 +736,30 @@ interface MockManagedUserApiKey { force_capabilities?: Record | null } +const mockSelfUserApiKeys: MockManagedUserApiKey[] = MOCK_USER_API_KEYS.map((key, index) => ({ + ...key, + fullKey: `sk-ae-demo-user-${index + 1}`, + is_locked: false, + is_standalone: false, + routing_group_id: null, +})) + +function mockSelfUserApiKeyPayload(key: MockManagedUserApiKey) { + return { + ...publicMockManagedUserApiKey(key), + routing_group_id: key.routing_group_id ?? null, + routing_group_name: MOCK_ROUTING_GROUPS.find(group => group.id === key.routing_group_id)?.name ?? null, + } +} + +function mockSelectableRoutingGroupId(value: unknown): string | null { + if (value == null) return null + const group = MOCK_ROUTING_GROUPS.find(group => group.id === value + && group.enabled && group.config_json.user_visible === true) + if (!group) throw { response: createMockResponse({ detail: '该策略分组不可选择' }, 403) } + return group.id +} + const mockManagedUserApiKeysByUserId = new Map([ [MOCK_NORMAL_USER.id ?? '', MOCK_USER_API_KEYS.map((key, index) => ({ ...key, @@ -1418,7 +1443,19 @@ const mockHandlers: Record Promise { await delay() - return createMockResponse(MOCK_USER_API_KEYS) + return createMockResponse(mockSelfUserApiKeys.map(mockSelfUserApiKeyPayload)) + }, + + 'GET /api/users/me/routing-groups': async () => { + await delay() + const items = MOCK_ROUTING_GROUPS.filter(group => group.enabled && group.config_json.user_visible === true) + .map(group => ({ + id: group.id, + name: group.name, + billing_multiplier: group.config_json.billing_multiplier ?? 1, + is_default: group.is_system_default, + })) + return createMockResponse({ items, total: items.length }) }, 'GET /api/users/me/client-config': async () => { @@ -1435,18 +1472,25 @@ const mockHandlers: Record Promise { await delay() const body = JSON.parse(config.data || '{}') - const newKey = { + const newKey: MockManagedUserApiKey = { id: `key-demo-${Date.now()}`, - key: `sk-aether-demo-${Math.random().toString(36).substring(2, 15)}`, + fullKey: `sk-aether-demo-${Math.random().toString(36).substring(2, 15)}`, key_display: 'sk-ae...demo', name: body.name || '新密钥(演示)', created_at: new Date().toISOString(), is_active: true, + is_locked: false, is_standalone: false, + routing_group_id: mockSelectableRoutingGroupId(body.routing_group_id), + feature_settings: body.feature_settings ?? null, + rate_limit: body.rate_limit ?? 0, + concurrent_limit: body.concurrent_limit ?? null, + ip_rules: body.ip_rules ?? null, total_requests: 0, total_cost_usd: 0 } - return createMockResponse(newKey) + mockSelfUserApiKeys.unshift(newKey) + return createMockResponse({ ...mockSelfUserApiKeyPayload(newKey), key: newKey.fullKey }) }, 'GET /api/users/me/usage': async () => { @@ -3990,13 +4034,40 @@ registerDynamicRoute('DELETE', '/api/admin/api-keys/:keyId', async (_config, par return createMockResponse({ message: '删除成功(演示模式)' }) }) +registerDynamicRoute('GET', '/api/users/me/api-keys/:keyId', async (config, params) => { + await delay() + const key = mockSelfUserApiKeys.find(key => key.id === params.keyId) + if (!key) throw { response: createMockResponse({ detail: 'API Key 不存在' }, 404) } + return createMockResponse(config.params?.include_key + ? { key: key.fullKey } + : mockSelfUserApiKeyPayload(key)) +}) + +registerDynamicRoute('PUT', '/api/users/me/api-keys/:keyId', async (config, params) => { + await delay() + const key = mockSelfUserApiKeys.find(key => key.id === params.keyId) + if (!key) throw { response: createMockResponse({ detail: 'API Key 不存在' }, 404) } + const body = mockRequestObject(config) + const groupId = 'routing_group_id' in body && body.routing_group_id !== key.routing_group_id + ? mockSelectableRoutingGroupId(body.routing_group_id) + : key.routing_group_id + if (typeof body.name === 'string') key.name = body.name + if (typeof body.rate_limit === 'number') key.rate_limit = body.rate_limit + if (typeof body.concurrent_limit === 'number') key.concurrent_limit = body.concurrent_limit + if ('ip_rules' in body) key.ip_rules = body.ip_rules as string[] | null + if ('feature_settings' in body) key.feature_settings = body.feature_settings as Record | null + key.routing_group_id = groupId + return createMockResponse({ ...mockSelfUserApiKeyPayload(key), message: 'API密钥已更新' }) +}) + // 用户 API Key 删除 registerDynamicRoute('DELETE', '/api/users/me/api-keys/:keyId', async (_config, params) => { await delay() - const key = MOCK_USER_API_KEYS.find(k => k.id === params.keyId) - if (!key) { + const index = mockSelfUserApiKeys.findIndex(key => key.id === params.keyId) + if (index < 0) { throw { response: createMockResponse({ detail: 'API Key 不存在' }, 404) } } + mockSelfUserApiKeys.splice(index, 1) return createMockResponse({ message: '删除成功(演示模式)' }) }) diff --git a/frontend/src/views/admin/ProviderManagement.vue b/frontend/src/views/admin/ProviderManagement.vue index e12afc71b..1a10bf2e7 100644 --- a/frontend/src/views/admin/ProviderManagement.vue +++ b/frontend/src/views/admin/ProviderManagement.vue @@ -160,7 +160,7 @@ :priority="getGroupPriority(provider)" :edit-context="priorityEditContext" :enabled="isGroupEnabled(provider.id)" - :disabled="schedulingBusy" + :disabled="priorityEditingDisabled" :priority-disabled="priorityEditingDisabled" :show-priority="false" @update:priority="setGroupPriority(provider.id, $event)" @@ -170,7 +170,7 @@ @@ -213,7 +213,7 @@ :priority="getGroupPriority(provider)" :edit-context="priorityEditContext" :enabled="isGroupEnabled(provider.id)" - :disabled="schedulingBusy" + :disabled="priorityEditingDisabled" :priority-disabled="priorityEditingDisabled" @update:priority="setGroupPriority(provider.id, $event)" /> @@ -222,7 +222,7 @@ @@ -325,7 +325,7 @@ import ProviderGroupControls from '@/features/providers/components/ProviderGroup import ProviderGroupToggleButton from '@/features/providers/components/ProviderGroupToggleButton.vue' import ProviderPriorityInput from '@/features/providers/components/ProviderPriorityInput.vue' import { providerGroupPriority, sortGroupProviders, moveGroupProvider } from '@/features/providers/utils/groupPriority' -import { getDefaultModelPolicy, normalizeRoutingGroupConfig, type RoutingModelPolicy, type RoutingPriorityMode, type RoutingSchedulingMode, type RoutingGroupConfig } from '@/features/routing/utils/routingPolicy' +import { getDefaultModelPolicy, isRoutingProviderEnabled, normalizeRoutingGroupConfig, type RoutingModelPolicy, type RoutingPriorityMode, type RoutingSchedulingMode, type RoutingGroupConfig } from '@/features/routing/utils/routingPolicy' import ProviderDragHandle from '@/features/providers/components/ProviderDragHandle.vue' import ProviderDeleteProgressCard from '@/features/providers/components/ProviderDeleteProgressCard.vue' import ProviderEmptyState from '@/features/providers/components/ProviderEmptyState.vue' @@ -595,15 +595,16 @@ function getGroupPriority(provider: ProviderWithEndpointsSummary) { return providerGroupPriority(priorityConfig.value, provider) } function isGroupEnabled(providerId: string) { - return !schedulingContext.value.config?.disabled_providers?.includes(providerId) + const config = schedulingContext.value.config + return config ? isRoutingProviderEnabled(config, providerId, selectedPolicy.value?.policy) : true } function setGroupEnabled(providerId: string, enabled: boolean) { - const config = schedulingContext.value.config - if (!config || schedulingBusy.value) return - const disabled = new Set(config.disabled_providers ?? []) - if (enabled) disabled.delete(providerId) - else disabled.add(providerId) - schedulingWorkspace.value?.updateDraftConfig({ ...config, disabled_providers: [...disabled] }) + const policy = selectedPolicy.value?.policy + if (!policy || priorityEditingDisabled.value) return + schedulingWorkspace.value?.updatePriorityPolicy({ + ...policy, + provider_enabled_overrides: { ...policy.provider_enabled_overrides, [providerId]: enabled }, + }) } function setGroupPriority(providerId: string, priority: number) { const policy = selectedPolicy.value?.policy @@ -748,9 +749,9 @@ function handleRowClick(event: MouseEvent, providerId: string) { // 打开添加提供商对话框 function openAddProviderDialog() { - const { groupId, groupName } = schedulingContext.value + const { groupId, groupName, config } = schedulingContext.value if (!groupId || schedulingBusy.value) { - showInfo(legacyT('请先创建或选择策略分组')) + showInfo(legacyT(!groupId && config ? '请先保存新分组,再添加提供商' : '请先创建或选择策略分组')) return } providerCreationGroup.value = { id: groupId, name: groupName } diff --git a/frontend/src/views/admin/__tests__/ProviderManagement.spec.ts b/frontend/src/views/admin/__tests__/ProviderManagement.spec.ts index 5c83bf99d..e42c061f0 100644 --- a/frontend/src/views/admin/__tests__/ProviderManagement.spec.ts +++ b/frontend/src/views/admin/__tests__/ProviderManagement.spec.ts @@ -24,6 +24,8 @@ const apiMocks = vi.hoisted(() => ({ updateProvider: vi.fn(), })) +const toastMocks = vi.hoisted(() => ({ success: vi.fn(), error: vi.fn(), info: vi.fn() })) + vi.mock('@/api/endpoints', async (importOriginal) => ({ ...await importOriginal(), ...apiMocks, @@ -34,7 +36,7 @@ vi.mock('@/composables/useConfirm', () => ({ })) vi.mock('@/composables/useToast', () => ({ - useToast: () => ({ success: vi.fn(), error: vi.fn(), info: vi.fn() }), + useToast: () => toastMocks, })) vi.mock('@/features/providers/composables/useProviderBalance', () => ({ @@ -417,10 +419,11 @@ describe('ProviderManagement group directory', () => { expect(row.textContent).toContain('本组启用') expect(row.textContent).not.toContain('本组禁用') expect(row.textContent).toContain('全局停用') - expect(workspace.groups['group-b']!.disabled_providers).toEqual([]) + expect(workspace.groups['group-b']!.disabled_providers).toEqual(['provider-4']) + expect(getDefaultModelPolicy(workspace.groups['group-b']!).provider_enabled_overrides).toEqual({ 'provider-4': true }) expect(workspace.groups['group-a']!.disabled_providers).toEqual([]) - expect(workspace.updateDraftConfig).toHaveBeenCalledExactlyOnceWith(workspace.groups['group-b']) - expect(workspace.updatePriorityPolicy).not.toHaveBeenCalled() + expect(workspace.updateDraftConfig).not.toHaveBeenCalled() + expect(workspace.updatePriorityPolicy).toHaveBeenCalledOnce() expect(providers[3]!.is_active).toBe(false) expect(apiMocks.updateProvider).not.toHaveBeenCalled() expect(apiMocks.getProvidersSummary).toHaveBeenCalledTimes(1) @@ -476,7 +479,8 @@ describe('ProviderManagement group directory', () => { await settle() expect(providerOrder(root)).toEqual(['provider-1', 'provider-2', 'provider-3', 'provider-4']) expect(groupDisabled.textContent).toContain('本组启用') - expect(workspace.groups['group-a']!.disabled_providers).toEqual([]) + expect(workspace.groups['group-a']!.disabled_providers).toEqual(['provider-2']) + expect(getDefaultModelPolicy(workspace.groups['group-a']!).provider_enabled_overrides).toEqual({ 'provider-2': true }) globallyDisabled.querySelector('[aria-label="Provider 1 本组启用"]')!.click() await settle() @@ -527,6 +531,46 @@ describe('ProviderManagement group directory', () => { expect(providerOrder(root)).toEqual(['provider-3', 'provider-1', 'provider-2', 'provider-4']) }) + it.each(['table', 'mobile list'] as const)('keeps model configuration membership independent in the %s', async layout => { + const providers = mockSortableProviders() + const config = createEmptyRoutingGroupConfig() + workspace.groups['group-a'] = writeSchedulingPolicies(config, [ + { ...createSchedulingPolicy(config), models: ['Model One', 'Model Two'] }, + { ...createSchedulingPolicy(config), models: ['Model Three'] }, + createSchedulingPolicy(config, 'all'), + ]) + const root = await mountView() + const row = [...root.querySelectorAll('[data-provider-sort-id="provider-1"]')] + .find(element => layout === 'table' ? element.closest('table') : !element.closest('table'))! + const toggle = row.querySelector('[aria-label="Provider 1 本组启用"]')! + toggle.click() + await settle() + expect(toggle.getAttribute('aria-pressed')).toBe('false') + expect(row.textContent).toContain('本组禁用') + for (const model of ['Model One', 'Model Two']) { + expect(getModelPolicy(workspace.groups['group-a']!, model).provider_enabled_overrides).toEqual({ 'provider-1': false }) + } + expect(getModelPolicy(workspace.groups['group-a']!, 'Model Three').provider_enabled_overrides).toEqual({}) + expect(getDefaultModelPolicy(workspace.groups['group-a']!).provider_enabled_overrides).toEqual({}) + expect(workspace.groups['group-a']!.disabled_providers).toEqual([]) + + root.querySelector('[data-select-policy="1"]')!.click() + await settle() + expect(toggle.getAttribute('aria-pressed')).toBe('true') + expect(row.textContent).toContain('本组启用') + expect(row.textContent).not.toContain('本组禁用') + root.querySelector('[data-select-policy="0"]')!.click() + await settle() + expect(toggle.getAttribute('aria-pressed')).toBe('false') + toggle.click() + await settle() + expect(getModelPolicy(workspace.groups['group-a']!, 'Model One').provider_enabled_overrides).toEqual({ 'provider-1': true }) + expect(getModelPolicy(workspace.groups['group-a']!, 'Model Three').provider_enabled_overrides).toEqual({}) + expect(workspace.updateDraftConfig).not.toHaveBeenCalled() + expect(providers[0]!.is_active).toBe(true) + expect(apiMocks.updateProvider).not.toHaveBeenCalled() + }) + it('drags the shared model configuration without splitting its models', async () => { mockSortableProviders() const config = createEmptyRoutingGroupConfig() @@ -585,11 +629,12 @@ describe('ProviderManagement group directory', () => { expect(toggle.disabled).toBe(false) toggle.click() await settle() - expect(workspace.groups['group-a']!.disabled_providers).toContain('provider-1') + expect(getDefaultModelPolicy(workspace.groups['group-a']!).provider_enabled_overrides['provider-1']).toBe(false) + expect(workspace.groups['group-a']!.disabled_providers).toEqual([]) expect(workspace.updatePriorityPolicy).toHaveBeenCalled() }) - it('disables ranking before a configuration is selected without blocking provider creation', async () => { + it('disables ranking and membership before a configuration is selected without blocking provider creation', async () => { mockSortableProviders() workspace.selectionReady = false const root = await mountView() @@ -599,10 +644,10 @@ describe('ProviderManagement group directory', () => { const handle = root.querySelector('[data-provider-drag-handle]') expect(handle == null || handle.disabled).toBe(true) const toggle = providerElements(root)[0]!.querySelector('[aria-label="Provider 1 本组启用"]')! - expect(toggle.disabled).toBe(false) + expect(toggle.disabled).toBe(true) toggle.click() await settle() - expect(workspace.groups['group-a']!.disabled_providers).toContain('provider-1') + expect(workspace.groups['group-a']!.disabled_providers).toEqual([]) findButton(root, '新增提供商').click() await settle() expect(root.querySelector('[data-routing-group-id="group-a"]')).not.toBeNull() @@ -631,9 +676,10 @@ describe('ProviderManagement group directory', () => { expect(toggle.disabled).toBe(false) toggle.click() await settle() - expect(workspace.groups['group-a']!.disabled_providers).toContain('provider-1') + expect(getModelPolicy(workspace.groups['group-a']!, 'Missing Model').provider_enabled_overrides).toEqual({ 'provider-1': false }) + expect(workspace.groups['group-a']!.disabled_providers).toEqual([]) expect(providerOrder(root)).toEqual(['provider-1', 'provider-2', 'provider-3', 'provider-4']) - expect(workspace.updatePriorityPolicy).toHaveBeenCalledOnce() + expect(workspace.updatePriorityPolicy).toHaveBeenCalledTimes(2) }) it('collects every API page before applying group priority', async () => { @@ -663,6 +709,54 @@ describe('ProviderManagement group directory', () => { findButton(root, '新增提供商').click() await settle() expect(root.querySelector('[data-save-edited-provider]')).toBeNull() + expect(toastMocks.info).toHaveBeenCalledWith('请先保存新分组,再添加提供商') + expect(workspace.ensureSaved).not.toHaveBeenCalled() + }) + + it('edits provider priority and membership in an unsaved group without saving it', async () => { + const providers = mockSortableProviders() + workspace.groups.new = { + ...createEmptyRoutingGroupConfig(), + disabled_providers: ['provider-4'], + } + const savedGroups = { + 'group-a': structuredClone(workspace.groups['group-a']), + 'group-b': structuredClone(workspace.groups['group-b']), + } + const root = await mountView('/admin/providers?group=new') + expect(root.querySelector('[data-scheduling-group="new"]')).not.toBeNull() + expect(providerOrder(root)).toEqual(['provider-1', 'provider-2', 'provider-3', 'provider-4']) + + const input = await openPriorityInput(root, 'Provider 4') + expect(input.disabled).toBe(false) + input.value = '0' + input.dispatchEvent(new Event('input', { bubbles: true })) + input.dispatchEvent(new KeyboardEvent('keydown', { key: 'Enter', bubbles: true })) + await settle() + + expect(getDefaultModelPolicy(workspace.groups.new!).provider_priority_overrides['provider-4']).toBe(0) + expect(providerOrder(root)).toEqual(['provider-4', 'provider-1', 'provider-2', 'provider-3']) + const row = providerElements(root)[0]! + const toggle = row.querySelector('[aria-label="Provider 4 本组启用"]')! + expect(toggle.disabled).toBe(false) + expect(toggle.getAttribute('aria-pressed')).toBe('false') + toggle.click() + await settle() + + expect(toggle.getAttribute('aria-pressed')).toBe('true') + expect(workspace.groups.new!.disabled_providers).toEqual(['provider-4']) + expect(getDefaultModelPolicy(workspace.groups.new!).provider_enabled_overrides).toEqual({ 'provider-4': true }) + expect(getDefaultModelPolicy(workspace.groups.new!).provider_priority_overrides['provider-4']).toBe(0) + expect(workspace.updatePriorityPolicy).toHaveBeenCalledTimes(2) + expect(workspace.updateDraftConfig).not.toHaveBeenCalled() + expect(workspace.groups['group-a']).toEqual(savedGroups['group-a']) + expect(workspace.groups['group-b']).toEqual(savedGroups['group-b']) + expect(workspace.ensureSaved).not.toHaveBeenCalled() + expect(workspace.refreshGroups).not.toHaveBeenCalled() + expect(apiMocks.updateProvider).not.toHaveBeenCalled() + expect(providers.map(provider => provider.provider_priority)).toEqual([10, 20, 30, 40]) + expect(mountedRouter!.currentRoute.value.query.group).toBe('new') + expect(root.querySelector('[data-provider-detail]')).toBeNull() }) it('opens details and applies edited snapshots for providers outside the loaded directory', async () => { diff --git a/frontend/src/views/admin/__tests__/ProviderSchedulingView.failover.spec.ts b/frontend/src/views/admin/__tests__/ProviderSchedulingView.failover.spec.ts index 334c5165e..6ea68d295 100644 --- a/frontend/src/views/admin/__tests__/ProviderSchedulingView.failover.spec.ts +++ b/frontend/src/views/admin/__tests__/ProviderSchedulingView.failover.spec.ts @@ -195,15 +195,22 @@ describe('ProviderSchedulingView failover persistence', () => { button(root, '区分模型').click() await flush() expect(button(root, '保存调度').disabled).toBe(true) + const modelPicker = root.querySelector('[aria-label="选择适用模型"]') + ?? root.querySelector('[aria-label="编辑模型"]') + if (!modelPicker) throw new Error('Missing model picker') + modelPicker.click() + await flush() element(document.body, 'input[aria-label="选择模型 model-a"]').click() await flush() - element(root, 'input[aria-label="选择模型 model-a"]').click() + element(document.body, 'input[aria-label="选择模型 model-a"]').click() await flush() expect(button(root, '保存调度').disabled).toBe(true) for (const model of ['model-a', 'model-b']) { element(document.body, `input[aria-label="选择模型 ${model}"]`).click() await nextTick() } + button(document.body, '完成选择').click() + await flush() byText('负载均衡').click() await nextTick() expect(button(root, '保存调度').disabled).toBe(false) diff --git a/frontend/src/views/shared/Usage.vue b/frontend/src/views/shared/Usage.vue index df6b4b49c..c70638d26 100644 --- a/frontend/src/views/shared/Usage.vue +++ b/frontend/src/views/shared/Usage.vue @@ -170,6 +170,7 @@