mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 18:37:46 +08:00
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.
This commit is contained in:
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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<dyn MinimalCandidateSelectionReadRepository> =
|
||||
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(),
|
||||
|
||||
@@ -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<Stri
|
||||
|
||||
fn routing_group_selection_cache_key(
|
||||
explicit_group: Option<&str>,
|
||||
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::<Vec<_>>()
|
||||
.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<LocalRequestedModelDecisionInput, GatewayError> {
|
||||
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();
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -231,8 +231,32 @@ pub(crate) async fn estimate_execution_plan_cost_upper_bound_usd(
|
||||
report_context: Option<&serde_json::Value>,
|
||||
) -> Result<Option<f64>, 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<Option<f64>, 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<i64>,
|
||||
requested_processing_tier: Option<&str>,
|
||||
cache_ttl_minutes: Option<i64>,
|
||||
use_base_cost: bool,
|
||||
) -> Result<Option<f64>, 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(
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -7733,6 +7733,7 @@ impl<'a> AdminAppState<'a> {
|
||||
feature_settings: key
|
||||
.contains_key("feature_settings")
|
||||
.then(|| feature_settings.clone()),
|
||||
routing_group_selection: None,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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<Option<Option<String>>, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
Option::<String>::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<Option<String>>,
|
||||
) -> Result<Option<Option<String>>, Response<Body>> {
|
||||
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<Option<Value>>,
|
||||
routing_group_patch: Option<Option<String>>,
|
||||
) -> Result<Option<Option<Value>>, 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<Option<Value>>,
|
||||
) -> Result<Option<Option<Value>>, 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<String, String> {
|
||||
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<String, String>,
|
||||
) -> Map<String, Value> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
@@ -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<Body>) -> Response<Body> {
|
||||
response.headers_mut().insert(
|
||||
http::header::CACHE_CONTROL,
|
||||
@@ -43,6 +51,8 @@ struct UsersMeCreateApiKeyRequest {
|
||||
concurrent_limit: Option<i32>,
|
||||
#[serde(default)]
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
routing_group_id: Option<String>,
|
||||
#[serde(default, alias = "allowed_ips")]
|
||||
ip_rules: Option<Vec<String>>,
|
||||
}
|
||||
@@ -57,6 +67,8 @@ struct UsersMeUpdateApiKeyRequest {
|
||||
concurrent_limit: Option<i32>,
|
||||
#[serde(default, deserialize_with = "deserialize_optional_json_patch")]
|
||||
feature_settings: Option<Option<serde_json::Value>>,
|
||||
#[serde(default, deserialize_with = "deserialize_routing_group_patch")]
|
||||
routing_group_id: Option<Option<String>>,
|
||||
#[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<String, String>,
|
||||
) -> 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<String, String>,
|
||||
) -> 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<String, String> {
|
||||
@@ -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::<Vec<_>>(),
|
||||
)
|
||||
@@ -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::<UsersMeCreateApiKeyRequest>(request.clone()).is_err());
|
||||
assert!(serde_json::from_value::<UsersMeUpdateApiKeyRequest>(request).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_ip_rules_trims_ip_and_cidr_values() {
|
||||
let values = normalize_users_me_ip_rules(Some(vec![
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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<Body> {
|
||||
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::<RoutingGroupConfig>(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::<Vec<_>>();
|
||||
Json(json!({"total": items.len(), "items": items})).into_response()
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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": []})] {
|
||||
|
||||
@@ -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<bool, GatewayRoutingSelectionError> {
|
||||
// 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: &[],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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")),
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
@@ -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();
|
||||
}
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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]
|
||||
|
||||
+154
@@ -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);
|
||||
@@ -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 {
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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<String> = 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();
|
||||
}
|
||||
|
||||
@@ -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"#;
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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)
|
||||
|
||||
+20
-2
@@ -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,
|
||||
|
||||
+20
-2
@@ -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,
|
||||
|
||||
@@ -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<Option<serde_json::Value>>,
|
||||
/// Self-service updates merge routing selection separately against the
|
||||
/// current stored settings. `None` retains administrative replacement semantics.
|
||||
pub routing_group_selection: Option<UpdateApiKeyRoutingGroupSelection>,
|
||||
}
|
||||
|
||||
#[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<Option<String>>,
|
||||
}
|
||||
|
||||
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<serde_json::Value>>,
|
||||
) -> Option<serde_json::Value> {
|
||||
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 {
|
||||
|
||||
@@ -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<f64>,
|
||||
pub finalized_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
@@ -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::<UsageSettlementInput>(
|
||||
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());
|
||||
|
||||
@@ -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<String, f64>,
|
||||
multiplier: f64,
|
||||
}
|
||||
|
||||
impl BillingMultiplierSnapshot {
|
||||
pub fn from_factors(factors: BTreeMap<String, f64>) -> Result<Self, DataLayerError> {
|
||||
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::<f64>()
|
||||
};
|
||||
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<f64, DataLayerError> {
|
||||
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<Option<BillingMultiplierSnapshot>, 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());
|
||||
}
|
||||
}
|
||||
@@ -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<Valu
|
||||
sanitize_usage_request_metadata_object(value?.as_object()?)
|
||||
}
|
||||
|
||||
/// Keep the request's first captured billing snapshot and reservation owner across retries.
|
||||
pub fn preserve_usage_routing_group_snapshot(
|
||||
incoming: Option<Value>,
|
||||
previous: Option<&Value>,
|
||||
) -> Option<Value> {
|
||||
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<String, Value>) -> Option<Value> {
|
||||
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<String, Value>) -> 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<String, Value>) -> 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<String, Value>) -> 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<String, Value>,
|
||||
target: &mut Map<String, Value>,
|
||||
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<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
||||
let Some(value) = source
|
||||
.get(key)
|
||||
@@ -1263,7 +1364,111 @@ fn safe_version_value(value: &Value) -> Option<String> {
|
||||
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() {
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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<f64> {
|
||||
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<f64> {
|
||||
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<bool> {
|
||||
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();
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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 =
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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>, 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;
|
||||
}
|
||||
@@ -115,6 +115,7 @@ WHERE version=20260919000000;
|
||||
20260923000000,
|
||||
20261001000000,
|
||||
20261004000000,
|
||||
20261007000000,
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -136,14 +136,16 @@ fn apply_allocations(
|
||||
}
|
||||
fn decimal_sum(
|
||||
rows: &[&StoredRequestUsageAudit],
|
||||
value: impl Fn(&StoredRequestUsageAudit) -> f64,
|
||||
value: impl Fn(&StoredRequestUsageAudit) -> Option<f64>,
|
||||
) -> Option<String> {
|
||||
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::<Vec<_>>();
|
||||
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 {
|
||||
|
||||
@@ -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),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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::<RoutingGroupConfig>(encoded).unwrap(),
|
||||
config
|
||||
);
|
||||
}
|
||||
for invalid in [
|
||||
serde_json::json!(null),
|
||||
serde_json::json!("true"),
|
||||
serde_json::json!(1),
|
||||
] {
|
||||
assert!(
|
||||
serde_json::from_value::<RoutingGroupConfig>(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::<RoutingGroupConfig>(
|
||||
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::<RoutingGroupConfig>(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<String>,
|
||||
#[serde(default)]
|
||||
pub allowed_keys: Vec<String>,
|
||||
/// 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<String, bool>,
|
||||
#[serde(default)]
|
||||
pub provider_priority_overrides: BTreeMap<String, i32>,
|
||||
#[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<String>,
|
||||
/// The default policy is global for the selected strategy group. Model
|
||||
@@ -238,6 +330,23 @@ pub struct RoutingGroupConfig {
|
||||
pub rules: Vec<RoutingRule>,
|
||||
}
|
||||
|
||||
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,
|
||||
|
||||
@@ -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<String>,
|
||||
/// Display name captured alongside the selected group by the gateway.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_version: Option<i64>,
|
||||
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::<Vec<_>>();
|
||||
// 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::<RoutingGroupConfig>(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
|
||||
|
||||
@@ -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<String>,
|
||||
#[serde(default)]
|
||||
|
||||
@@ -55,11 +55,15 @@ pub struct RoutingRuntimeFacts {
|
||||
pub priority_mode: Option<RoutingSetPriorityMode>,
|
||||
}
|
||||
|
||||
#[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<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_version: Option<i64>,
|
||||
pub selection_source: String,
|
||||
#[serde(default)]
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<Value>) -> 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(
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -66,6 +66,10 @@ impl UsageSettlementWriter for ReuseStore {
|
||||
&self,
|
||||
input: ReconcileUsagePolicyCostInput,
|
||||
) -> Result<Option<StoredUsagePolicyCostReservation>, 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
|
||||
|
||||
@@ -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<T>(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
|
||||
|
||||
@@ -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
|
||||
|
||||
+29
-3
@@ -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<string, boolean> | 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<ApiKey> {
|
||||
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<ApiKey> {
|
||||
const response = await apiClient.post<ApiKey>('/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<ApiKey & { message: string }> {
|
||||
const response = await apiClient.put<ApiKey & { message: string }>(
|
||||
`/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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 // 客户端从请求进入网关到完成的总耗时
|
||||
|
||||
@@ -19,6 +19,10 @@ const props = withDefaults(defineProps<{
|
||||
collisionPadding: 0,
|
||||
ariaLabel: undefined,
|
||||
})
|
||||
|
||||
const emit = defineEmits<{
|
||||
openAutoFocus: [event: Event]
|
||||
}>()
|
||||
</script>
|
||||
|
||||
<template>
|
||||
@@ -39,6 +43,7 @@ const props = withDefaults(defineProps<{
|
||||
:align-offset="props.alignOffset"
|
||||
:collision-padding="props.collisionPadding"
|
||||
:aria-label="props.ariaLabel"
|
||||
@open-auto-focus="emit('openAutoFocus', $event)"
|
||||
>
|
||||
<slot />
|
||||
</PopoverContent>
|
||||
|
||||
@@ -28,6 +28,12 @@
|
||||
<SelectValue :placeholder="loading ? '正在加载分组' : '选择策略分组'" />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem
|
||||
v-if="isNewDraft"
|
||||
value="new"
|
||||
>
|
||||
新建策略
|
||||
</SelectItem>
|
||||
<SelectItem
|
||||
v-for="group in groups"
|
||||
:key="group.id"
|
||||
@@ -42,7 +48,7 @@
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-8"
|
||||
:disabled="busy"
|
||||
:disabled="busy || isNewDraft"
|
||||
title="新建分组"
|
||||
aria-label="新建策略"
|
||||
@click="openCreate"
|
||||
@@ -85,7 +91,7 @@
|
||||
class="h-8 w-8"
|
||||
:class="{ 'text-primary': draftDirty }"
|
||||
:disabled="!canSaveDraft"
|
||||
:title="saving ? '正在保存…' : !routingSchedulingValid ? '请先选择适用模型' : draftDirty ? '保存修改' : '已保存'"
|
||||
:title="saving ? '正在保存…' : billingMultiplierError ?? (!routingSchedulingValid ? '请先选择适用模型' : draftDirty ? '保存修改' : '已保存')"
|
||||
aria-label="保存调度"
|
||||
:aria-busy="saving"
|
||||
@click="saveDraft"
|
||||
@@ -101,12 +107,12 @@
|
||||
<div
|
||||
v-if="draft"
|
||||
ref="groupMetadata"
|
||||
class="min-w-0 border-b border-border/50 p-3"
|
||||
class="min-w-0 space-y-2 border-b border-border/50 p-3"
|
||||
aria-label="分组信息"
|
||||
:inert="busy"
|
||||
>
|
||||
<div class="flex min-w-0 items-center gap-2">
|
||||
<label class="min-w-0 flex-1">
|
||||
<label class="block min-w-0 flex-1">
|
||||
<span class="sr-only">策略名称</span>
|
||||
<Input
|
||||
v-model="draft.name"
|
||||
@@ -126,6 +132,40 @@
|
||||
/>
|
||||
</label>
|
||||
</div>
|
||||
<div class="flex min-w-0 items-center justify-between gap-3">
|
||||
<label class="flex min-w-0 items-center gap-2 text-xs">
|
||||
<span class="shrink-0">分组倍率</span>
|
||||
<Input
|
||||
:model-value="billingMultiplierInput"
|
||||
type="number"
|
||||
min="0"
|
||||
step="any"
|
||||
size="sm"
|
||||
class="w-24 min-w-0"
|
||||
aria-label="分组倍率"
|
||||
:aria-invalid="Boolean(billingMultiplierError)"
|
||||
:aria-describedby="billingMultiplierError ? 'group-billing-multiplier-error' : undefined"
|
||||
:disabled="busy"
|
||||
@update:model-value="updateBillingMultiplier"
|
||||
/>
|
||||
<span class="shrink-0 text-muted-foreground">倍</span>
|
||||
</label>
|
||||
<label class="flex shrink-0 items-center gap-1 text-xs">
|
||||
<span>用户可见</span>
|
||||
<Switch
|
||||
v-model="draft.config_json.user_visible"
|
||||
:disabled="busy"
|
||||
aria-label="用户可见"
|
||||
/>
|
||||
</label>
|
||||
</div>
|
||||
<p
|
||||
v-if="billingMultiplierError"
|
||||
id="group-billing-multiplier-error"
|
||||
class="text-xs text-destructive"
|
||||
>
|
||||
{{ billingMultiplierError }}
|
||||
</p>
|
||||
</div>
|
||||
<RoutingSchedulingPolicyEditor
|
||||
v-if="draft"
|
||||
@@ -353,49 +393,6 @@
|
||||
<slot />
|
||||
</div>
|
||||
</div>
|
||||
<Dialog
|
||||
:model-value="createDialogOpen"
|
||||
title="新建策略分组"
|
||||
description="创建独立分组,创建成功后切换到新分组。"
|
||||
size="md"
|
||||
:persistent="busy"
|
||||
@update:model-value="closeCreate"
|
||||
>
|
||||
<div
|
||||
class="space-y-4"
|
||||
:inert="busy"
|
||||
>
|
||||
<label class="block space-y-1.5 text-sm"><span>分组名称</span><Input
|
||||
v-model="createForm.name"
|
||||
aria-label="新分组名称"
|
||||
placeholder="例如:日常使用"
|
||||
/></label>
|
||||
<label class="flex items-center justify-between gap-3 text-sm"><span>启用分组</span><Switch
|
||||
v-model="createForm.enabled"
|
||||
aria-label="启用新分组"
|
||||
/></label>
|
||||
<p class="text-xs text-muted-foreground">
|
||||
使用默认调度配置;创建后可在提供商目录中调整成员和顺序。
|
||||
</p>
|
||||
</div>
|
||||
<template #footer>
|
||||
<Button
|
||||
:disabled="busy || !createForm.name.trim()"
|
||||
aria-label="创建策略分组"
|
||||
@click="createGroup"
|
||||
>
|
||||
{{ creating ? '创建中…' : '创建分组' }}
|
||||
</Button>
|
||||
<Button
|
||||
variant="outline"
|
||||
:disabled="busy"
|
||||
aria-label="取消新建分组"
|
||||
@click="closeCreate(false)"
|
||||
>
|
||||
取消
|
||||
</Button>
|
||||
</template>
|
||||
</Dialog>
|
||||
<AlertDialog
|
||||
v-model="deleteDialogOpen"
|
||||
type="destructive"
|
||||
@@ -409,16 +406,17 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, onBeforeUnmount, onMounted, ref, watch } from 'vue'
|
||||
import { computed, nextTick, onBeforeUnmount, onMounted, ref, watch } from 'vue'
|
||||
import { onBeforeRouteLeave, onBeforeRouteUpdate, useRoute, useRouter, type RouteLocationNormalized } from 'vue-router'
|
||||
import { ChevronRight, Plus, Save, Star, Trash2 } from 'lucide-vue-next'
|
||||
import { Button, Card, Dialog, Input, Select, SelectContent, SelectItem, SelectTrigger, SelectValue, Switch } from '@/components/ui'
|
||||
import { Button, Card, Input, Select, SelectContent, SelectItem, SelectTrigger, SelectValue, Switch } from '@/components/ui'
|
||||
import { AlertDialog } from '@/components/common'
|
||||
import HelpHint from '@/components/common/HelpHint.vue'
|
||||
import {
|
||||
DEFAULT_STICKY_KEY_ATTEMPTS,
|
||||
createEmptyRoutingGroupConfig,
|
||||
normalizeStickyKeyAttempts,
|
||||
parseBillingMultiplier,
|
||||
type RoutingModelPolicy,
|
||||
type RoutingPriorityMode,
|
||||
type RoutingSchedulingMode,
|
||||
@@ -484,10 +482,8 @@ const loadingError = ref<string | null>(null)
|
||||
const saving = ref(false)
|
||||
const saveConflict = ref(false)
|
||||
const deleting = ref(false)
|
||||
const creating = ref(false)
|
||||
const busy = computed(() => loading.value || saving.value || deleting.value || creating.value)
|
||||
const createDialogOpen = ref(false)
|
||||
const createForm = ref({ name: '', enabled: true })
|
||||
const busy = computed(() => loading.value || saving.value || deleting.value)
|
||||
const billingMultiplierInput = ref('1')
|
||||
const draftGeneration = ref(0)
|
||||
const groupMetadata = ref<HTMLElement | null>(null)
|
||||
const advancedOpen = ref(false)
|
||||
@@ -498,7 +494,8 @@ let discardConfirmation: Promise<boolean> | null = null
|
||||
|
||||
const routeGroupId = computed(() => queryToString(route.query.group))
|
||||
const defaultGroupId = computed(() => groups.value.find(group => group.is_system_default)?.id ?? groups.value[0]?.id ?? null)
|
||||
const selectedValue = computed(() => draft.value?.id ?? '')
|
||||
const isNewDraft = computed(() => draft.value != null && !draft.value.id)
|
||||
const selectedValue = computed(() => isNewDraft.value ? 'new' : draft.value?.id ?? '')
|
||||
const priorityMode = 'provider' as const
|
||||
const schedulingMode = computed(() => activePolicy.value?.schedulingMode ?? draft.value?.config_json.default_policy.scheduling_mode ?? 'cache_affinity')
|
||||
const emptyMessage = computed(() => loading.value ? '正在加载调度策略' : loadingError.value ?? (groups.value.length ? '未找到调度策略' : '还没有调度策略'))
|
||||
@@ -507,8 +504,9 @@ const stickyKeyAttempts = computed(() => draft.value?.config_json.default_policy
|
||||
const cfHeartbeat = computed(() => draft.value?.config_json.default_policy.enable_cf_heartbeat ?? false)
|
||||
const cyberContinueFailover = computed(() => draft.value?.config_json.default_policy.cyber_continue_failover ?? false)
|
||||
const cancelOnClientDisconnect = computed(() => draft.value?.config_json.default_policy.cancel_on_client_disconnect ?? false)
|
||||
const draftDirty = computed(() => draft.value != null && (routingFailoverPending.value || savedDraftSnapshot.value !== draftSnapshotValue(draft.value)))
|
||||
const canSaveDraft = computed(() => Boolean(draft.value) && !busy.value && draftDirty.value && routingSchedulingValid.value)
|
||||
const billingMultiplierError = computed(() => parseBillingMultiplier(billingMultiplierInput.value) == null ? '分组倍率必须是大于或等于 0 的有效数字' : null)
|
||||
const draftDirty = computed(() => draft.value != null && (Boolean(billingMultiplierError.value) || routingFailoverPending.value || savedDraftSnapshot.value !== draftSnapshotValue(draft.value)))
|
||||
const canSaveDraft = computed(() => Boolean(draft.value) && !busy.value && draftDirty.value && routingSchedulingValid.value && !billingMultiplierError.value)
|
||||
|
||||
function queryToString(value: unknown): string | null {
|
||||
if (Array.isArray(value)) return typeof value[0] === 'string' ? value[0] : null
|
||||
@@ -545,22 +543,31 @@ function resetEditors(): void {
|
||||
}
|
||||
|
||||
function selectGroup(group: RoutingGroupRecord, preserveSelection = false): void {
|
||||
const selection = preserveSelection && draft.value?.id === group.id ? activePolicy.value : null
|
||||
const selection = preserveSelection ? activePolicy.value : null
|
||||
resetEditors()
|
||||
initialSchedulingSelection.value = selection
|
||||
? { id: selection.id, scope: selection.scope, modelNames: [...selection.modelNames] }
|
||||
: null
|
||||
draft.value = { id: group.id, version: group.version, name: group.name, enabled: group.enabled, is_system_default: group.is_system_default, config_json: cloneConfig(group.config_json) }
|
||||
billingMultiplierInput.value = String(draft.value.config_json.billing_multiplier)
|
||||
savedDraftSnapshot.value = draftSnapshotValue(draft.value)
|
||||
}
|
||||
|
||||
function openCreate(): void {
|
||||
if (busy.value) return
|
||||
createForm.value = { name: '', enabled: true }
|
||||
createDialogOpen.value = true
|
||||
if (busy.value || isNewDraft.value) return
|
||||
openGroup('new')
|
||||
}
|
||||
function closeCreate(value: boolean): void { if (!busy.value) createDialogOpen.value = value }
|
||||
|
||||
function startNewDraft(): void {
|
||||
resetEditors()
|
||||
draft.value = { version: 0, name: '', enabled: true, is_system_default: groups.value.length === 0, config_json: createEmptyRoutingGroupConfig() }
|
||||
billingMultiplierInput.value = '1'
|
||||
savedDraftSnapshot.value = null
|
||||
void nextTick(() => {
|
||||
groupMetadata.value?.scrollIntoView?.({ block: 'nearest' })
|
||||
groupMetadata.value?.querySelector<HTMLInputElement>('[aria-label="策略名称"]')?.focus()
|
||||
})
|
||||
}
|
||||
|
||||
function clearDraft(): void {
|
||||
resetEditors()
|
||||
@@ -571,10 +578,7 @@ function clearDraft(): void {
|
||||
function syncRouteState(): void {
|
||||
if (loading.value) return
|
||||
if (routeGroupId.value === 'new') {
|
||||
if (!draft.value) { const group = groups.value.find(item => item.id === defaultGroupId.value); if (group) selectGroup(group) }
|
||||
openCreate()
|
||||
internalNavigation = true
|
||||
void router.replace({ name: 'ProviderManagement', query: { ...route.query, view: undefined, group: draft.value?.id } }).finally(() => { internalNavigation = false })
|
||||
if (!isNewDraft.value) startNewDraft()
|
||||
return
|
||||
}
|
||||
const group = groups.value.find(item => item.id === (routeGroupId.value ?? defaultGroupId.value))
|
||||
@@ -603,7 +607,7 @@ async function confirmDiscard(): Promise<boolean> {
|
||||
async function guardNavigation(to: RouteLocationNormalized): Promise<boolean> {
|
||||
if (internalNavigation) return true
|
||||
const targetGroup = queryToString(to.query.group) ?? defaultGroupId.value
|
||||
const staysOnDraft = to.name === 'ProviderManagement' && (targetGroup === selectedValue.value || targetGroup === 'new')
|
||||
const staysOnDraft = to.name === 'ProviderManagement' && targetGroup === selectedValue.value
|
||||
if (staysOnDraft) return true
|
||||
if (busy.value) {
|
||||
showError('正在保存调度设置,请稍候再切换')
|
||||
@@ -625,6 +629,13 @@ function updateDraftConfig(value: RoutingGroupConfig): void {
|
||||
if (draft.value) draft.value.config_json = normalizeProviderSchedulingConfig(value)
|
||||
}
|
||||
|
||||
function updateBillingMultiplier(value: string | number): void {
|
||||
if (!draft.value || busy.value) return
|
||||
billingMultiplierInput.value = String(value)
|
||||
const parsed = parseBillingMultiplier(value)
|
||||
if (parsed != null) draft.value.config_json.billing_multiplier = parsed
|
||||
}
|
||||
|
||||
function updatePriorityPolicy(policy: RoutingModelPolicy): void {
|
||||
if (busy.value) return
|
||||
routingSchedulingPolicyEditor.value?.updateSelectedPolicy(policy)
|
||||
@@ -700,7 +711,7 @@ async function loadGlobalModels(options: { cacheTtlMs?: number } = {}): Promise<
|
||||
}
|
||||
|
||||
async function saveDraft(): Promise<boolean> {
|
||||
if (!draft.value?.id || busy.value) return false
|
||||
if (!draft.value || busy.value) return false
|
||||
const name = draft.value.name.trim()
|
||||
if (!name) {
|
||||
groupMetadata.value?.scrollIntoView?.({ block: 'nearest' })
|
||||
@@ -708,6 +719,12 @@ async function saveDraft(): Promise<boolean> {
|
||||
showError('策略名称不能为空')
|
||||
return false
|
||||
}
|
||||
if (billingMultiplierError.value) {
|
||||
groupMetadata.value?.scrollIntoView?.({ block: 'nearest' })
|
||||
groupMetadata.value?.querySelector<HTMLInputElement>('[aria-label="分组倍率"]')?.focus()
|
||||
showError(billingMultiplierError.value)
|
||||
return false
|
||||
}
|
||||
if (routingFailoverPolicyEditor.value && !routingFailoverPolicyEditor.value.commitJsonDrafts()) { failoverOpen.value = true; return false }
|
||||
const failoverError = validateRoutingFailoverPolicy(draft.value.config_json.default_policy)
|
||||
if (failoverError) { failoverOpen.value = true; showError(failoverError); return false }
|
||||
@@ -715,19 +732,33 @@ async function saveDraft(): Promise<boolean> {
|
||||
const targetGroupId = draft.value.id
|
||||
const submittedGeneration = draftGeneration.value
|
||||
const submittedSnapshot = draftSnapshotValue(draft.value)
|
||||
const payload = { name, enabled: draft.value.enabled, is_system_default: draft.value.is_system_default, expected_version: draft.value.version, config_json: cloneConfig(draft.value.config_json) }
|
||||
const payload = { name, enabled: draft.value.enabled, is_system_default: draft.value.is_system_default, config_json: cloneConfig(draft.value.config_json) }
|
||||
const expectedVersion = draft.value.version
|
||||
saving.value = true
|
||||
try {
|
||||
const saved = await updateRoutingGroup(targetGroupId, payload)
|
||||
const saved = targetGroupId
|
||||
? await updateRoutingGroup(targetGroupId, { ...payload, expected_version: expectedVersion })
|
||||
: await createRoutingGroup({ ...payload, sort_order: groups.value.length })
|
||||
const unchanged = draftGeneration.value === submittedGeneration && draft.value?.id === targetGroupId && draftSnapshotValue(draft.value) === submittedSnapshot
|
||||
replaceGroup(saved, unchanged, true)
|
||||
success('调度策略已保存')
|
||||
if (!targetGroupId) {
|
||||
// Retain any newer edits while attaching the server ID, so retries update the created group.
|
||||
if (!unchanged && draft.value && draftGeneration.value === submittedGeneration && !draft.value.id) {
|
||||
draft.value.id = saved.id
|
||||
draft.value.version = saved.version
|
||||
savedDraftSnapshot.value = draftSnapshotValue({ ...saved, config_json: cloneConfig(saved.config_json) })
|
||||
}
|
||||
internalNavigation = true
|
||||
try { await router.replace({ name: 'ProviderManagement', query: { ...route.query, view: undefined, group: saved.id } }) }
|
||||
finally { internalNavigation = false }
|
||||
}
|
||||
success(targetGroupId ? '调度策略已保存' : '策略分组已创建')
|
||||
emit('saved')
|
||||
return unchanged
|
||||
} catch (err) {
|
||||
const status = (err as { response?: { status?: number } })?.response?.status
|
||||
saveConflict.value = status === 409
|
||||
showError(status === 409 ? '此分组已在其他操作中更新。当前修改已保留,请重新加载最新分组后再编辑。' : parseApiError(err, '保存调度策略失败'))
|
||||
saveConflict.value = Boolean(targetGroupId) && status === 409
|
||||
showError(saveConflict.value ? '此分组已在其他操作中更新。当前修改已保留,请重新加载最新分组后再编辑。' : parseApiError(err, targetGroupId ? '保存调度策略失败' : '创建策略分组失败'))
|
||||
log.error('保存调度策略失败:', err)
|
||||
return false
|
||||
} finally { saving.value = false }
|
||||
@@ -738,22 +769,6 @@ async function ensureSaved(): Promise<boolean> {
|
||||
const approved = await confirm({ title: '先保存当前分组', message: '当前分组有未保存的修改。保存后继续?', confirmText: '保存并继续', cancelText: '继续编辑', variant: 'question' })
|
||||
return approved && await saveDraft()
|
||||
}
|
||||
async function createGroup(): Promise<void> {
|
||||
if (busy.value || !createForm.value.name.trim()) return
|
||||
if (draftDirty.value && !await ensureSaved()) return
|
||||
creating.value = true
|
||||
try {
|
||||
const saved = await createRoutingGroup({ name: createForm.value.name.trim(), enabled: createForm.value.enabled, is_system_default: groups.value.length === 0, sort_order: groups.value.length, config_json: createEmptyRoutingGroupConfig() })
|
||||
replaceGroup(saved, true)
|
||||
createDialogOpen.value = false
|
||||
internalNavigation = true
|
||||
try { await router.replace({ name: 'ProviderManagement', query: { ...route.query, view: undefined, group: saved.id } }) }
|
||||
finally { internalNavigation = false }
|
||||
success('策略分组已创建')
|
||||
emit('saved')
|
||||
} catch (err) { showError(parseApiError(err, '创建策略分组失败')); log.error('创建策略分组失败:', err) }
|
||||
finally { creating.value = false }
|
||||
}
|
||||
|
||||
async function confirmDeleteDraft(): Promise<void> {
|
||||
if (!draft.value?.id || busy.value) return
|
||||
|
||||
+244
-19
@@ -2,7 +2,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import { createApp, defineComponent, h, nextTick, ref, type App } from 'vue'
|
||||
import { createMemoryHistory, createRouter, RouterView, type LocationQueryRaw } from 'vue-router'
|
||||
import ProviderSchedulingView from '../ProviderSchedulingView.vue'
|
||||
import { createEmptyRoutingGroupConfig, type RoutingGroupConfig } from '@/features/routing/utils/routingPolicy'
|
||||
import { createEmptyRoutingGroupConfig, setDefaultProviderPriorityOverrides, type RoutingGroupConfig } from '@/features/routing/utils/routingPolicy'
|
||||
import type { RoutingGroupRecord, RoutingGroupUpdateRequest } from '@/api/routing-profiles'
|
||||
|
||||
const routingApi = vi.hoisted(() => ({ listRoutingGroups: vi.fn(), updateRoutingGroup: vi.fn(), createRoutingGroup: vi.fn(), deleteRoutingGroup: vi.fn() }))
|
||||
@@ -101,6 +101,12 @@ async function editName(root: HTMLElement, name: string) {
|
||||
input.dispatchEvent(new Event('input', { bubbles: true }))
|
||||
await nextTick()
|
||||
}
|
||||
async function editMultiplier(root: HTMLElement, value: string, label = '分组倍率') {
|
||||
const input = element<HTMLInputElement>(root, `[aria-label="${label}"]`)
|
||||
input.value = value
|
||||
input.dispatchEvent(new Event('input', { bubbles: true }))
|
||||
await nextTick()
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
@@ -116,6 +122,116 @@ afterEach(() => {
|
||||
})
|
||||
|
||||
describe('ProviderSchedulingView workspace navigation', () => {
|
||||
it('defaults legacy groups to private and saves visibility without losing scheduling or rankings', async () => {
|
||||
const legacyConfig = createEmptyRoutingGroupConfig()
|
||||
Reflect.deleteProperty(legacyConfig, 'user_visible')
|
||||
const { root, workspace } = await mountWorkspace({}, [group('default', { is_system_default: true, config_json: legacyConfig })])
|
||||
expect(button(root, '用户可见').getAttribute('aria-checked')).toBe('false')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
button(root, '用户可见').click()
|
||||
await nextTick()
|
||||
expect(button(root, '用户可见').getAttribute('aria-checked')).toBe('true')
|
||||
expect(button(root, '启用策略').getAttribute('aria-checked')).toBe('true')
|
||||
expect(button(root, '保存调度').disabled).toBe(false)
|
||||
expect(routingApi.updateRoutingGroup).not.toHaveBeenCalled()
|
||||
const config = setDefaultProviderPriorityOverrides(
|
||||
JSON.parse(JSON.stringify(contextChange.mock.lastCall?.[0].config)),
|
||||
{ 'provider-a': 4, 'provider-b': 1 },
|
||||
)
|
||||
config.default_policy.scheduling_mode = 'fixed_order'
|
||||
config.disabled_providers = ['provider-c']
|
||||
workspace.value?.updateDraftConfig(config)
|
||||
await nextTick()
|
||||
expect(button(root, '用户可见').getAttribute('aria-checked')).toBe('true')
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(routingApi.updateRoutingGroup).toHaveBeenLastCalledWith('default', expect.objectContaining({
|
||||
enabled: true,
|
||||
config_json: { ...config, user_visible: true },
|
||||
}))
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
button(root, '用户可见').click()
|
||||
await nextTick()
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(routingApi.updateRoutingGroup).toHaveBeenLastCalledWith('default', expect.objectContaining({
|
||||
config_json: { ...config, user_visible: false },
|
||||
}))
|
||||
expect(button(root, '用户可见').getAttribute('aria-checked')).toBe('false')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
})
|
||||
|
||||
it('defaults old groups to a multiplier of one and saves decimal or zero values with scheduling changes', async () => {
|
||||
const legacyConfig = createEmptyRoutingGroupConfig()
|
||||
Reflect.deleteProperty(legacyConfig, 'billing_multiplier')
|
||||
const { root, workspace } = await mountWorkspace({}, [group('default', { is_system_default: true, config_json: legacyConfig })])
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('1')
|
||||
await editMultiplier(root, '1.25')
|
||||
const config = JSON.parse(JSON.stringify(contextChange.mock.lastCall?.[0].config))
|
||||
config.disabled_providers = ['provider-a']
|
||||
workspace.value?.updateDraftConfig(config)
|
||||
await nextTick()
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(routingApi.updateRoutingGroup).toHaveBeenLastCalledWith('default', expect.objectContaining({ config_json: expect.objectContaining({ billing_multiplier: 1.25, disabled_providers: ['provider-a'] }) }))
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('1.25')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
await editMultiplier(root, '0')
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(routingApi.updateRoutingGroup).toHaveBeenLastCalledWith('default', expect.objectContaining({ config_json: expect.objectContaining({ billing_multiplier: 0 }) }))
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('0')
|
||||
})
|
||||
|
||||
it('keeps invalid multiplier input unsaved across scheduling changes and cancelled navigation', async () => {
|
||||
const { root, router, workspace } = await mountWorkspace()
|
||||
await editMultiplier(root, '2')
|
||||
await editMultiplier(root, '')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('')
|
||||
expect(element(root, '[aria-label="分组倍率"]').getAttribute('aria-invalid')).toBe('true')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
const config = JSON.parse(JSON.stringify(contextChange.mock.lastCall?.[0].config))
|
||||
config.disabled_providers = ['provider-a']
|
||||
workspace.value?.updateDraftConfig(config)
|
||||
await nextTick()
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('')
|
||||
expect(contextChange.mock.lastCall?.[0].config.billing_multiplier).toBe(2)
|
||||
confirm.mockResolvedValueOnce(true)
|
||||
expect(await workspace.value?.ensureSaved()).toBe(false)
|
||||
expect(routingApi.updateRoutingGroup).not.toHaveBeenCalled()
|
||||
expect(toast.error).toHaveBeenCalledWith('分组倍率必须是大于或等于 0 的有效数字')
|
||||
await chooseGroup(root, 'last')
|
||||
expect(router.currentRoute.value.query.group).toBeUndefined()
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('')
|
||||
await editMultiplier(root, '-1')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
confirm.mockResolvedValue(true)
|
||||
await chooseGroup(root, 'last')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('1')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
})
|
||||
|
||||
it('validates new group multipliers without replacing blank input and accepts zero', async () => {
|
||||
const { root, router } = await mountWorkspace({ group: 'new' })
|
||||
expect(router.currentRoute.value.query.group).toBe('new')
|
||||
expect(selector(root).textContent?.trim()).toBe('新建策略')
|
||||
expect(document.querySelector('[role="dialog"]')).toBeNull()
|
||||
await editName(root, '免费分组')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('1')
|
||||
await editMultiplier(root, '')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
button(root, '保存调度').click()
|
||||
expect(routingApi.createRoutingGroup).not.toHaveBeenCalled()
|
||||
await editMultiplier(root, '0')
|
||||
expect(button(root, '保存调度').disabled).toBe(false)
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(routingApi.createRoutingGroup).toHaveBeenCalledWith(expect.objectContaining({ name: '免费分组', config_json: expect.objectContaining({ billing_multiplier: 0 }) }))
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('0')
|
||||
expect(router.currentRoute.value.query.group).toBe('created')
|
||||
})
|
||||
|
||||
it('opens the system default immediately and forwards provider inspection and refreshes', async () => {
|
||||
const { root, providerRevision } = await mountWorkspace()
|
||||
expect(selector(root).textContent?.trim()).toBe('default · 默认')
|
||||
@@ -245,42 +361,151 @@ describe('ProviderSchedulingView workspace navigation', () => {
|
||||
expect(router.currentRoute.value.query.group).toBeUndefined()
|
||||
})
|
||||
|
||||
it('opens creation independently and cancellation preserves the selected group and draft', async () => {
|
||||
it('asks before replacing existing edits with a new inline group draft', async () => {
|
||||
const { root, router } = await mountWorkspace()
|
||||
await editName(root, '保留的分组草稿')
|
||||
button(root, '新建策略').click()
|
||||
await nextTick()
|
||||
await flush()
|
||||
expect(selector(root).textContent?.trim()).toBe('default · 默认')
|
||||
expect(router.currentRoute.value.query.group).toBeUndefined()
|
||||
expect(document.querySelector('[role="dialog"][aria-label="新建策略分组"]')).not.toBeNull()
|
||||
button(root, '取消新建分组').click()
|
||||
await nextTick()
|
||||
expect(document.querySelector('[role="dialog"]')).toBeNull()
|
||||
expect(button(root, '保存调度').disabled).toBe(false)
|
||||
expect(contextChange.mock.lastCall?.[0].groupName).toBe('保留的分组草稿')
|
||||
expect(routingApi.createRoutingGroup).not.toHaveBeenCalled()
|
||||
expect(confirm).not.toHaveBeenCalled()
|
||||
expect(confirm).toHaveBeenCalledOnce()
|
||||
confirm.mockResolvedValue(true)
|
||||
button(root, '新建策略').click()
|
||||
await flush()
|
||||
expect(router.currentRoute.value.query.group).toBe('new')
|
||||
expect(selector(root).textContent?.trim()).toBe('新建策略')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="策略名称"]').value).toBe('')
|
||||
expect(routingApi.createRoutingGroup).not.toHaveBeenCalled()
|
||||
expect(routingApi.updateRoutingGroup).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('creates in a dialog and switches only after the request succeeds', async () => {
|
||||
const { root, router } = await mountWorkspace({ group: 'new' })
|
||||
expect(selector(root).textContent?.trim()).toBe('default · 默认')
|
||||
expect(router.currentRoute.value.query.group).toBe('default')
|
||||
const name = element<HTMLInputElement>(root, '[aria-label="新分组名称"]')
|
||||
name.value = '新策略'
|
||||
name.dispatchEvent(new Event('input', { bubbles: true }))
|
||||
it('starts an inline draft and creates the complete edited configuration only on header save', async () => {
|
||||
const { root, router, workspace } = await mountWorkspace()
|
||||
button(root, '新建策略').click()
|
||||
await flush()
|
||||
expect(router.currentRoute.value.query.group).toBe('new')
|
||||
expect(selector(root).textContent?.trim()).toBe('新建策略')
|
||||
expect(document.querySelector('[role="dialog"]')).toBeNull()
|
||||
expect(root.querySelector('button[aria-label="删除策略"]')).toBeNull()
|
||||
expect(root.querySelector('[data-testid="provider-directory"]')).not.toBeNull()
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('1')
|
||||
expect(button(root, '启用策略').getAttribute('aria-checked')).toBe('true')
|
||||
expect(button(root, '用户可见').getAttribute('aria-checked')).toBe('false')
|
||||
expect(button(root, '设为系统默认').getAttribute('aria-pressed')).toBe('false')
|
||||
expect(routingApi.createRoutingGroup).not.toHaveBeenCalled()
|
||||
await editName(root, '新策略')
|
||||
await editMultiplier(root, '1.25')
|
||||
button(root, '启用策略').click()
|
||||
button(root, '用户可见').click()
|
||||
await nextTick()
|
||||
const config = JSON.parse(JSON.stringify(contextChange.mock.lastCall?.[0].config)) as RoutingGroupConfig
|
||||
config.disabled_providers = ['provider-a']
|
||||
config.default_policy.scheduling_mode = 'fixed_order'
|
||||
config.default_policy.sticky_key_attempts = 3
|
||||
workspace.value?.updateDraftConfig(config)
|
||||
await nextTick()
|
||||
button(root, '新建策略').click()
|
||||
await flush()
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="策略名称"]').value).toBe('新策略')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('1.25')
|
||||
expect(button(root, '启用策略').getAttribute('aria-checked')).toBe('false')
|
||||
expect(button(root, '用户可见').getAttribute('aria-checked')).toBe('true')
|
||||
expect(config.user_visible).toBe(true)
|
||||
expect(contextChange.mock.lastCall?.[0].config).toEqual(config)
|
||||
expect(routingApi.createRoutingGroup).not.toHaveBeenCalled()
|
||||
let finish!: (group: RoutingGroupRecord) => void
|
||||
routingApi.createRoutingGroup.mockReturnValue(new Promise(resolve => { finish = resolve }))
|
||||
button(root, '创建策略分组').click()
|
||||
button(root, '保存调度').click()
|
||||
await nextTick()
|
||||
expect(selector(root).textContent?.trim()).toBe('default · 默认')
|
||||
finish(group('created', { name: '新策略' }))
|
||||
expect(router.currentRoute.value.query.group).toBe('new')
|
||||
expect(selector(root).textContent?.trim()).toBe('新建策略')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
button(root, '保存调度').click()
|
||||
expect(routingApi.createRoutingGroup).toHaveBeenCalledOnce()
|
||||
expect(routingApi.createRoutingGroup).toHaveBeenCalledWith(expect.objectContaining({ name: '新策略', enabled: false, is_system_default: false, config_json: config }))
|
||||
finish(group('created', { name: '新策略', enabled: false, config_json: config }))
|
||||
await flush()
|
||||
expect(routingApi.createRoutingGroup).toHaveBeenCalledWith(expect.objectContaining({ name: '新策略', config_json: expect.objectContaining({ disabled_providers: [] }) }))
|
||||
expect(router.currentRoute.value.query.group).toBe('created')
|
||||
expect(selector(root).textContent?.trim()).toBe('新策略')
|
||||
expect(selector(root).textContent?.trim()).toBe('新策略 · 停用')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
expect(confirm).not.toHaveBeenCalled()
|
||||
expect(routingApi.updateRoutingGroup).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it.each([false, true])('opens a new draft from its deep link with first-group default=%s', async firstGroup => {
|
||||
const { root, router } = await mountWorkspace({ group: 'new' }, firstGroup ? [] : [group('existing')])
|
||||
expect(router.currentRoute.value.query.group).toBe('new')
|
||||
expect(selector(root).textContent?.trim()).toBe('新建策略')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="策略名称"]').value).toBe('')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('1')
|
||||
expect(button(root, '启用策略').getAttribute('aria-checked')).toBe('true')
|
||||
expect(button(root, '用户可见').getAttribute('aria-checked')).toBe('false')
|
||||
expect(button(root, '设为系统默认').getAttribute('aria-pressed')).toBe(String(firstGroup))
|
||||
expect(root.querySelector('button[aria-label="删除策略"]')).toBeNull()
|
||||
expect(document.querySelector('[role="dialog"]')).toBeNull()
|
||||
expect(routingApi.createRoutingGroup).not.toHaveBeenCalled()
|
||||
await editName(root, '深链新建')
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(routingApi.createRoutingGroup).toHaveBeenCalledWith(expect.objectContaining({
|
||||
name: '深链新建',
|
||||
enabled: true,
|
||||
is_system_default: firstGroup,
|
||||
config_json: expect.objectContaining({ billing_multiplier: 1, user_visible: false }),
|
||||
}))
|
||||
expect(router.currentRoute.value.query.group).toBe('created')
|
||||
})
|
||||
|
||||
it.each([500, 409])('keeps a failed new draft intact and retries the same create payload after status %s', async status => {
|
||||
const { root, router } = await mountWorkspace({ group: 'new' })
|
||||
await editName(root, '重试创建')
|
||||
await editMultiplier(root, '0.75')
|
||||
routingApi.createRoutingGroup.mockRejectedValueOnce({ response: { status } })
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(router.currentRoute.value.query.group).toBe('new')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="策略名称"]').value).toBe('重试创建')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('0.75')
|
||||
expect(button(root, '保存调度').disabled).toBe(false)
|
||||
expect(toast.error).toHaveBeenCalled()
|
||||
expect(root.querySelector('button[aria-label="重新加载分组"]')).toBeNull()
|
||||
const failedPayload = routingApi.createRoutingGroup.mock.calls[0][0]
|
||||
button(root, '保存调度').click()
|
||||
await flush()
|
||||
expect(routingApi.createRoutingGroup).toHaveBeenCalledTimes(2)
|
||||
expect(routingApi.createRoutingGroup).toHaveBeenLastCalledWith(failedPayload)
|
||||
expect(router.currentRoute.value.query.group).toBe('created')
|
||||
expect(button(root, '保存调度').disabled).toBe(true)
|
||||
expect(routingApi.updateRoutingGroup).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it.each(['switch', 'leave'] as const)('confirms discarding a new draft before %s navigation', async navigation => {
|
||||
const { root, router } = await mountWorkspace({ group: 'new' })
|
||||
await editName(root, '保留新建草稿')
|
||||
await editMultiplier(root, '2')
|
||||
const navigate = () => navigation === 'switch'
|
||||
? chooseGroup(root, 'last')
|
||||
: router.push({ name: 'Other' })
|
||||
await navigate()
|
||||
expect(confirm).toHaveBeenCalledOnce()
|
||||
expect(router.currentRoute.value.name).toBe('ProviderManagement')
|
||||
expect(router.currentRoute.value.query.group).toBe('new')
|
||||
expect(selector(root).textContent?.trim()).toBe('新建策略')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="策略名称"]').value).toBe('保留新建草稿')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="分组倍率"]').value).toBe('2')
|
||||
confirm.mockResolvedValue(true)
|
||||
await navigate()
|
||||
expect(router.currentRoute.value.name).toBe(navigation === 'switch' ? 'ProviderManagement' : 'Other')
|
||||
if (navigation === 'switch') {
|
||||
expect(router.currentRoute.value.query.group).toBe('last')
|
||||
expect(element<HTMLInputElement>(root, '[aria-label="策略名称"]').value).toBe('last')
|
||||
}
|
||||
expect(routingApi.createRoutingGroup).not.toHaveBeenCalled()
|
||||
expect(routingApi.updateRoutingGroup).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('saves strategy metadata with scheduling and updates the default marker', async () => {
|
||||
|
||||
@@ -150,16 +150,17 @@ describe('RoutingFailoverPolicyEditor', () => {
|
||||
expect(editor.value?.commitJsonDrafts()).toBe(false)
|
||||
})
|
||||
|
||||
it('edits independent global budgets and documents sticky retry exclusion', async () => {
|
||||
it('edits independent global transfer budgets', async () => {
|
||||
const { root, policy } = mountEditor()
|
||||
expect(root.textContent).toContain('首次尝试和粘性同 Key 重试不计入')
|
||||
expect(root.textContent).toContain('不会中断已开始的调用')
|
||||
const count = control<HTMLInputElement>(root, '全局最大转移次数')
|
||||
count.value = '4'
|
||||
count.dispatchEvent(new Event('input', { bubbles: true }))
|
||||
await nextTick()
|
||||
expect(policy.value.max_transfer_count).toBe(4)
|
||||
expect(policy.value.max_transfer_timeout_seconds).toBe(0)
|
||||
await input(control<HTMLInputElement>(root, '全局最大转移时间'), '15')
|
||||
expect(policy.value.max_transfer_count).toBe(4)
|
||||
expect(policy.value.max_transfer_timeout_seconds).toBe(15)
|
||||
})
|
||||
|
||||
it('adds regex and status-only rules and reports invalid drafts', async () => {
|
||||
|
||||
@@ -29,7 +29,7 @@ function mountEditor(options: { modelIds?: string[], keyMode?: boolean } = {}) {
|
||||
initial.model_policies = [{
|
||||
...getDefaultModelPolicy(initial),
|
||||
provider_priority_overrides: { outside: 99 },
|
||||
key_priority_overrides_by_format: { 'openai:chat': { outside: 99 }, 'claude:chat': { 'claude-key': 45 } },
|
||||
key_priority_overrides_by_format: { 'openai:chat': { outside: 99 }, 'claude:messages': { 'claude-key': 45 } },
|
||||
pool_priority_overrides: { 'other-pool': 88 },
|
||||
}]
|
||||
const config = ref(initial)
|
||||
@@ -183,6 +183,20 @@ describe('RoutingPriorityPolicyEditor ordering', () => {
|
||||
expect(rows.every(row => !row.textContent?.includes('停用'))).toBe(true)
|
||||
})
|
||||
|
||||
it('shows the selected model policy membership while retaining legacy exclusions as defaults', async () => {
|
||||
const { root, config } = mountEditor()
|
||||
config.value.disabled_providers = ['A', 'C']
|
||||
config.value.model_policies[0].provider_enabled_overrides = { A: true, B: false }
|
||||
await vi.waitFor(() => expect(rowNames(root)).toHaveLength(4))
|
||||
const rows = [...root.querySelectorAll<HTMLElement>('[draggable="true"]')]
|
||||
expect(rows.find(row => row.textContent?.includes('提供商 A'))?.textContent).not.toContain('本组禁用')
|
||||
expect(rows.find(row => row.textContent?.includes('提供商 B'))?.textContent).toContain('本组禁用')
|
||||
expect(rows.find(row => row.textContent?.includes('提供商 C'))?.textContent).toContain('本组禁用')
|
||||
await click(root, '置顶 提供商 C')
|
||||
expect(getDefaultModelPolicy(config.value).provider_enabled_overrides).toEqual({ A: true, B: false })
|
||||
expect(config.value.disabled_providers).toEqual(['A', 'C'])
|
||||
})
|
||||
|
||||
it('emits provider inspection and refreshes health and keys without altering the draft', async () => {
|
||||
const { root, config, revision, inspect } = mountEditor()
|
||||
await vi.waitFor(() => expect(rowNames(root)).toHaveLength(4))
|
||||
|
||||
@@ -83,8 +83,8 @@ async function clickText(root: HTMLElement, text: string) {
|
||||
}
|
||||
|
||||
async function select(root: HTMLElement, model: string) {
|
||||
await openModels(root)
|
||||
control<HTMLInputElement>(root, `选择模型 ${model}`).click()
|
||||
const picker = await openModels(root)
|
||||
control<HTMLInputElement>(picker, `选择模型 ${model}`).click()
|
||||
await flush()
|
||||
}
|
||||
|
||||
@@ -98,9 +98,19 @@ async function openModels(root: HTMLElement) {
|
||||
if (control(root, '全部模型').getAttribute('aria-pressed') === 'true') {
|
||||
await clickText(root, '区分模型')
|
||||
}
|
||||
if (root.querySelector('[aria-label="全局模型选择列表"]')) return
|
||||
const activeCard = root.querySelector('[aria-label^="选择调度配置 "][aria-pressed="true"]')?.closest('section')
|
||||
const edit = activeCard?.querySelector<HTMLButtonElement>('[aria-label="编辑模型"]')
|
||||
if (edit) {
|
||||
if (!document.querySelector('[aria-label="编辑适用模型"]')) {
|
||||
edit.click()
|
||||
await flush()
|
||||
}
|
||||
return control(document.body, '编辑适用模型')
|
||||
}
|
||||
if (root.querySelector('[aria-label="全局模型选择列表"]')) return root
|
||||
control<HTMLButtonElement>(root, '选择适用模型').click()
|
||||
await flush()
|
||||
return root
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
@@ -208,7 +218,7 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
const { root, config, selection } = mountEditor(writeSchedulingPolicies(initial, [selected, fallback]), 'config-only')
|
||||
control<HTMLButtonElement>(root, '选择调度配置 2').click()
|
||||
await flush()
|
||||
expect(root.querySelector('[aria-label="选择适用模型"]')).toBeNull()
|
||||
expect(control(root, '调度配置 2').querySelector('[aria-label="编辑模型"]')).toBeNull()
|
||||
expect(control(root, '当前配置的适用模型').textContent).toContain('未单独指定的模型')
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ scope: 'all', modelNames: [], schedulingMode: 'cache_affinity' }))
|
||||
await clickText(root, '负载均衡')
|
||||
@@ -216,8 +226,8 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
expect(getModelScheduling(config.value, 'model-b').scheduling_mode).toBe('load_balance')
|
||||
control<HTMLButtonElement>(root, '选择调度配置 1').click()
|
||||
await flush()
|
||||
expect(root.querySelectorAll('[aria-label="选择适用模型"]')).toHaveLength(1)
|
||||
expect(control(root, '选择适用模型').textContent).toContain('模型 A')
|
||||
expect(root.querySelectorAll('[aria-label="编辑模型"]')).toHaveLength(1)
|
||||
expect(control(root, '已配置模型').textContent).toContain('模型 A')
|
||||
})
|
||||
|
||||
it('expands only the selected model editor and keeps every shared ranking attached to its configuration', async () => {
|
||||
@@ -225,17 +235,33 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
const first = { ...createSchedulingPolicy(initial), models: ['model-a', 'model-c'], schedulingMode: 'fixed_order' as const }
|
||||
const second = { ...createSchedulingPolicy(initial), models: ['model-b'], schedulingMode: 'load_balance' as const }
|
||||
const { root, config, selection, editor } = mountEditor(writeSchedulingPolicies(initial, [first, second]), 'config-only')
|
||||
expect(root.querySelectorAll('[aria-label="选择适用模型"]')).toHaveLength(1)
|
||||
expect(root.querySelector('button[aria-label="选择适用模型"]')).toBeNull()
|
||||
expect(control(root, '调度配置 1').contains(control(root, '全局模型选择列表'))).toBe(true)
|
||||
expect(control(root, '调度配置 2').querySelector('[aria-label="全局模型选择列表"]')).toBeNull()
|
||||
expect(root.querySelectorAll('[aria-label="编辑模型"]')).toHaveLength(2)
|
||||
for (const index of [1, 2]) {
|
||||
const card = control(root, `调度配置 ${index}`)
|
||||
const edit = control<HTMLButtonElement>(card, '编辑模型')
|
||||
expect(edit.textContent?.trim()).toBe('')
|
||||
expect(edit.querySelector('svg')).not.toBeNull()
|
||||
expect(edit.compareDocumentPosition(control(card, `删除调度配置 ${index}`)) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy()
|
||||
}
|
||||
expect(control(root, '调度配置 2').querySelector('[aria-label="当前配置的适用模型"]')).toBeNull()
|
||||
expect(document.querySelector('[aria-label="全局模型选择列表"]')).toBeNull()
|
||||
expect(control(root, '已配置模型').querySelector('[title="model-a"]')?.textContent).toBe('模型 A')
|
||||
expect(control(root, '已配置模型').querySelector('[title="model-c"]')?.textContent).toBe('模型 C')
|
||||
expect(control(root, '选择调度配置 1').textContent).toContain('模型 A +1')
|
||||
expect(root.querySelector('[aria-label="调整排序"]')).toBeNull()
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ id: first.id, modelNames: ['model-a', 'model-c'], schedulingMode: 'fixed_order' }))
|
||||
control<HTMLButtonElement>(root, '选择调度配置 2').click()
|
||||
const secondEdit = control<HTMLButtonElement>(control(root, '调度配置 2'), '编辑模型')
|
||||
secondEdit.click()
|
||||
await flush()
|
||||
expect(control(root, '调度配置 2').contains(control(root, '全局模型选择列表'))).toBe(true)
|
||||
expect(control(root, '调度配置 1').querySelector('[aria-label="全局模型选择列表"]')).toBeNull()
|
||||
expect(control(root, '选择调度配置 2').getAttribute('aria-expanded')).toBe('true')
|
||||
expect(control(root, '选择调度配置 2').getAttribute('aria-pressed')).toBe('true')
|
||||
expect(control(root, '调度配置 1').querySelector('[aria-label="当前配置的适用模型"]')).toBeNull()
|
||||
expect(control<HTMLInputElement>(control(document.body, '编辑适用模型'), '选择模型 model-b').checked).toBe(true)
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ id: second.id, modelNames: ['model-b'], schedulingMode: 'load_balance' }))
|
||||
control<HTMLButtonElement>(control(document.body, '编辑适用模型'), '完成选择').click()
|
||||
await flush()
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
await vi.waitFor(() => expect(document.activeElement).toBe(secondEdit))
|
||||
const policy = { ...getDefaultModelPolicy(initial), provider_priority_overrides: { provider: 6 } }
|
||||
editor.value!.updateSelectedPolicy(policy)
|
||||
await flush()
|
||||
@@ -247,7 +273,7 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
await flush()
|
||||
expect(getModelPolicy(config.value, 'model-a').provider_priority_overrides).toEqual({ shared: 3 })
|
||||
expect(getModelPolicy(config.value, 'model-c').provider_priority_overrides).toEqual({ shared: 3 })
|
||||
expect(root.querySelectorAll('[aria-label="选择适用模型"]')).toHaveLength(1)
|
||||
expect(root.querySelectorAll('[aria-label="编辑模型"]')).toHaveLength(2)
|
||||
control<HTMLButtonElement>(root, '删除调度配置 1').click()
|
||||
await flush()
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ id: second.id, modelNames: ['model-b'] }))
|
||||
@@ -271,12 +297,14 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
const { root, selection, config, editor } = mountEditor(writeSchedulingPolicies(initial, [first, second]), 'config-only')
|
||||
const firstButton = control<HTMLButtonElement>(root, '选择调度配置 1')
|
||||
expect(firstButton.getAttribute('aria-expanded')).toBe('true')
|
||||
await openModels(root)
|
||||
selection.mockClear()
|
||||
firstButton.click()
|
||||
await flush()
|
||||
expect(firstButton.getAttribute('aria-expanded')).toBe('false')
|
||||
expect(firstButton.getAttribute('aria-pressed')).toBe('true')
|
||||
expect(root.querySelector('[aria-label="全局模型选择列表"]')).toBeNull()
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
expect(root.querySelectorAll('[aria-label="编辑模型"]')).toHaveLength(2)
|
||||
expect(root.querySelector('[aria-label="调度策略"]')).toBeNull()
|
||||
expect(selection).not.toHaveBeenCalled()
|
||||
editor.value!.updateSelectedPolicy({ ...getDefaultModelPolicy(initial), provider_priority_overrides: { provider: 8 } })
|
||||
@@ -289,7 +317,7 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
firstButton.click()
|
||||
await flush()
|
||||
expect(firstButton.getAttribute('aria-expanded')).toBe('true')
|
||||
expect(control(root, '调度配置 1').contains(control(root, '全局模型选择列表'))).toBe(true)
|
||||
expect(control(root, '调度配置 1').contains(control(root, '编辑模型'))).toBe(true)
|
||||
expect(control(root, '调度配置 1').contains(control(root, '调度策略'))).toBe(true)
|
||||
firstButton.click()
|
||||
await flush()
|
||||
@@ -302,8 +330,9 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
await flush()
|
||||
expect(control(root, '选择调度配置 3').getAttribute('aria-expanded')).toBe('true')
|
||||
expect(control(root, '选择调度配置 3').getAttribute('aria-pressed')).toBe('true')
|
||||
expect(control(root, '调度配置 3').contains(control(root, '搜索全局模型'))).toBe(true)
|
||||
expect(root.querySelectorAll('[aria-label="全局模型选择列表"]')).toHaveLength(1)
|
||||
expect(control(control(root, '调度配置 3'), '编辑模型')).toBeTruthy()
|
||||
expect(control(root, '当前配置的适用模型').textContent).toContain('请选择适用模型')
|
||||
expect(document.querySelector('[aria-label="全局模型选择列表"]')).toBeNull()
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ policy: null, modelNames: [] }))
|
||||
})
|
||||
|
||||
@@ -316,7 +345,7 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
id: 'previous-generated-id', scope: 'selected', modelNames: ['model-c', 'model-b'],
|
||||
})
|
||||
expect(control(root, '选择调度配置 2').getAttribute('aria-pressed')).toBe('true')
|
||||
expect(control(root, '调度配置 2').contains(control(root, '全局模型选择列表'))).toBe(true)
|
||||
expect(control(control(root, '调度配置 2'), '编辑模型')).toBeTruthy()
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ modelNames: ['model-b', 'model-c'] }))
|
||||
|
||||
const legacy = writeSchedulingPolicies(initial, [first, createSchedulingPolicy(initial, 'all')])
|
||||
@@ -335,12 +364,12 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
const secondCard = control<HTMLElement>(root, '调度配置 2')
|
||||
expect(firstCard.contains(control(root, '调度策略'))).toBe(true)
|
||||
expect(secondCard.querySelector('[aria-label="调度策略"]')).toBeNull()
|
||||
expect(control(firstCard, '全局模型选择列表').compareDocumentPosition(control(firstCard, '调度策略')) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy()
|
||||
expect(control(firstCard, '当前配置的适用模型').compareDocumentPosition(control(firstCard, '调度策略')) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy()
|
||||
control<HTMLButtonElement>(secondCard, '选择调度配置 2').click()
|
||||
await flush()
|
||||
expect(firstCard.querySelector('[aria-label="调度策略"]')).toBeNull()
|
||||
const secondStrategy = control(secondCard, '调度策略')
|
||||
expect(control(secondCard, '全局模型选择列表').compareDocumentPosition(secondStrategy) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy()
|
||||
expect(control(secondCard, '当前配置的适用模型').compareDocumentPosition(secondStrategy) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy()
|
||||
expect([...secondStrategy.querySelectorAll('button')].find(button => button.textContent?.trim() === '负载均衡')?.getAttribute('aria-pressed')).toBe('true')
|
||||
await clickText(secondStrategy, '缓存亲和')
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ id: second.id, modelNames: ['model-b'], priorityMode: 'provider' }))
|
||||
@@ -352,44 +381,86 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
expect(getModelScheduling(config.value, 'model-b').scheduling_mode).toBe('cache_affinity')
|
||||
control<HTMLButtonElement>(secondCard, '选择调度配置 2').click()
|
||||
await flush()
|
||||
control<HTMLInputElement>(root, '选择模型 model-c').click()
|
||||
await flush()
|
||||
await select(root, 'model-c')
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ id: second.id, modelNames: ['model-b', 'model-c'] }))
|
||||
})
|
||||
|
||||
it('searches and adds models inside the active card without carrying its search or edits to another configuration', async () => {
|
||||
it('keeps live model edits when finishing or escaping and isolates the next configuration picker', async () => {
|
||||
const initial = createEmptyRoutingGroupConfig()
|
||||
const first = { ...createSchedulingPolicy(initial), models: ['model-a'], schedulingMode: 'fixed_order' as const }
|
||||
const second = { ...createSchedulingPolicy(initial), models: ['model-b'], schedulingMode: 'load_balance' as const }
|
||||
const { root, config } = mountEditor(writeSchedulingPolicies(initial, [first, second]), 'config-only', {}, true)
|
||||
const firstCard = control(root, '调度配置 1')
|
||||
const search = control<HTMLInputElement>(firstCard, '搜索全局模型')
|
||||
const { root, config, selection } = mountEditor(writeSchedulingPolicies(initial, [first, second]), 'config-only', {}, true)
|
||||
const picker = await openModels(root)
|
||||
expect(picker.querySelector('[aria-label="清空已选"]')).toBeNull()
|
||||
expect(picker.textContent).not.toMatch(/已选\s*\d/)
|
||||
const search = control<HTMLInputElement>(picker, '搜索全局模型')
|
||||
search.value = '模型 C'
|
||||
search.dispatchEvent(new Event('input', { bubbles: true }))
|
||||
await flush()
|
||||
control<HTMLInputElement>(firstCard, '选择模型 model-c').click()
|
||||
control<HTMLInputElement>(picker, '选择模型 model-c').click()
|
||||
await flush()
|
||||
expect(getModelScheduling(config.value, 'model-c').scheduling_mode).toBe('fixed_order')
|
||||
expect(getModelScheduling(config.value, 'model-b').scheduling_mode).toBe('load_balance')
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ id: first.id, modelNames: ['model-a', 'model-c'] }))
|
||||
control<HTMLButtonElement>(picker, '完成选择').click()
|
||||
await flush()
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
expect(readSchedulingPolicies(config.value)[0].models).toEqual(['model-a', 'model-c'])
|
||||
expect(control(root, '已配置模型').textContent).toContain('模型 C')
|
||||
const reopened = await openModels(root)
|
||||
expect(control<HTMLInputElement>(reopened, '选择模型 model-c').checked).toBe(true)
|
||||
control<HTMLButtonElement>(root, '选择调度配置 2').click()
|
||||
await flush()
|
||||
expect(root.querySelectorAll('[aria-label="搜索全局模型"]')).toHaveLength(1)
|
||||
const nextSearch = control<HTMLInputElement>(root, '搜索全局模型')
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
expect(selection).toHaveBeenLastCalledWith(expect.objectContaining({ id: second.id, modelNames: ['model-b'] }))
|
||||
const nextPicker = await openModels(root)
|
||||
expect(document.querySelectorAll('[aria-label="搜索全局模型"]')).toHaveLength(1)
|
||||
const nextSearch = control<HTMLInputElement>(nextPicker, '搜索全局模型')
|
||||
expect(nextSearch.value).toBe('')
|
||||
expect(control<HTMLInputElement>(root, '选择模型 model-b').checked).toBe(true)
|
||||
expect(control<HTMLInputElement>(nextPicker, '选择模型 model-b').checked).toBe(true)
|
||||
nextSearch.value = 'model-c'
|
||||
nextSearch.dispatchEvent(new Event('input', { bubbles: true }))
|
||||
await flush()
|
||||
expect(control<HTMLInputElement>(root, '选择模型 model-c').disabled).toBe(true)
|
||||
expect(control(root, '全局模型选择列表').textContent).toContain('已用于配置 1')
|
||||
expect(control<HTMLInputElement>(nextPicker, '选择模型 model-c').disabled).toBe(true)
|
||||
expect(control(nextPicker, '全局模型选择列表').textContent).toContain('已用于配置 1')
|
||||
nextSearch.dispatchEvent(new KeyboardEvent('keydown', { key: 'Escape', bubbles: true }))
|
||||
await flush()
|
||||
expect(root.querySelector('[aria-label="全局模型选择列表"]')).not.toBeNull()
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
expect(readSchedulingPolicies(config.value).map(entry => entry.models)).toEqual([['model-a', 'model-c'], ['model-b']])
|
||||
expect(control(root, '选择调度配置 2').getAttribute('aria-pressed')).toBe('true')
|
||||
await vi.waitFor(() => expect(document.activeElement).toBe(control(control(root, '调度配置 2'), '编辑模型')))
|
||||
const add = control<HTMLButtonElement>(root, '添加调度配置')
|
||||
for (const card of root.querySelectorAll('section[aria-label^="调度配置 "]')) {
|
||||
expect(card.compareDocumentPosition(add) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy()
|
||||
}
|
||||
expect(root.querySelector('[aria-label="完成选择"]')).toBeNull()
|
||||
expect(document.querySelector('[aria-label="完成选择"]')).toBeNull()
|
||||
})
|
||||
|
||||
it('keeps live model edits when Escape or saving closes the popover', async () => {
|
||||
const initial = createEmptyRoutingGroupConfig()
|
||||
const first = { ...createSchedulingPolicy(initial), models: ['model-a'] }
|
||||
const { root, config, disabled, selection } = mountEditor(writeSchedulingPolicies(initial, [first]), 'config-only', {}, true)
|
||||
await select(root, 'model-c')
|
||||
const saved = JSON.stringify(config.value)
|
||||
control(document.body, '搜索全局模型').dispatchEvent(new KeyboardEvent('keydown', { key: 'Escape', bubbles: true }))
|
||||
await flush()
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
expect(JSON.stringify(config.value)).toBe(saved)
|
||||
await openModels(root)
|
||||
selection.mockClear()
|
||||
disabled.value = true
|
||||
await flush()
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
expect(control<HTMLButtonElement>(root, '编辑模型').disabled).toBe(true)
|
||||
expect(control(root, '选择调度配置 1').getAttribute('aria-pressed')).toBe('true')
|
||||
expect(JSON.stringify(config.value)).toBe(saved)
|
||||
expect(selection).not.toHaveBeenCalled()
|
||||
disabled.value = false
|
||||
await flush()
|
||||
expect(document.querySelector('[aria-label="编辑适用模型"]')).toBeNull()
|
||||
const reopened = await openModels(root)
|
||||
expect(control<HTMLInputElement>(reopened, '选择模型 model-a').checked).toBe(true)
|
||||
expect(control<HTMLInputElement>(reopened, '选择模型 model-c').checked).toBe(true)
|
||||
})
|
||||
|
||||
it('keeps a strategy chosen before model selection and lets unfinished configurations choose their strategy', async () => {
|
||||
@@ -402,8 +473,7 @@ describe('RoutingSchedulingPolicyEditor', () => {
|
||||
control<HTMLButtonElement>(root, '添加调度配置').click()
|
||||
await flush()
|
||||
await clickText(root, '负载均衡')
|
||||
control<HTMLInputElement>(root, '选择模型 model-b').click()
|
||||
await flush()
|
||||
await select(root, 'model-b')
|
||||
expect(getModelScheduling(config.value, 'model-b').scheduling_mode).toBe('load_balance')
|
||||
expect(getModelScheduling(config.value, 'model-a').scheduling_mode).toBe('fixed_order')
|
||||
expect(root.querySelectorAll('[aria-label="调度策略"]')).toHaveLength(1)
|
||||
|
||||
@@ -6,9 +6,11 @@ import {
|
||||
createEmptyRoutingGroupConfig,
|
||||
getDefaultModelPolicy,
|
||||
getModelScheduling,
|
||||
isRoutingProviderEnabled,
|
||||
modelSchedulingRuleId,
|
||||
normalizeRoutingGroupConfig,
|
||||
normalizeStickyKeyAttempts,
|
||||
parseBillingMultiplier,
|
||||
resolveModelKeyPriorityOverride,
|
||||
setDefaultPoolPriorityOverrides,
|
||||
setDefaultProviderPriorityOverrides,
|
||||
@@ -25,6 +27,27 @@ describe('routingPolicy', () => {
|
||||
expect(config.default_policy.priority_mode).toBe('provider')
|
||||
expect(config.default_policy.scheduling_mode).toBe('cache_affinity')
|
||||
expect(config.default_policy.cancel_on_client_disconnect).toBe(false)
|
||||
expect(config.billing_multiplier).toBe(1)
|
||||
expect(config.user_visible).toBe(false)
|
||||
})
|
||||
|
||||
it('keeps new and legacy groups private unless user visibility is explicitly true', () => {
|
||||
expect(createEmptyRoutingGroupConfig().user_visible).toBe(false)
|
||||
expect(normalizeRoutingGroupConfig({ user_visible: true }).user_visible).toBe(true)
|
||||
for (const user_visible of [undefined, null, false, 'true', 1]) {
|
||||
expect(normalizeRoutingGroupConfig({ user_visible } as unknown as Parameters<typeof normalizeRoutingGroupConfig>[0]).user_visible).toBe(false)
|
||||
}
|
||||
})
|
||||
|
||||
it('preserves a nonnegative billing multiplier while defaulting legacy configs to one', () => {
|
||||
expect(createEmptyRoutingGroupConfig().billing_multiplier).toBe(1)
|
||||
expect(normalizeRoutingGroupConfig({ billing_multiplier: 0 }).billing_multiplier).toBe(0)
|
||||
expect(normalizeRoutingGroupConfig({ billing_multiplier: 1.25 }).billing_multiplier).toBe(1.25)
|
||||
expect(parseBillingMultiplier('0')).toBe(0)
|
||||
expect(parseBillingMultiplier('1.25')).toBe(1.25)
|
||||
for (const value of ['', ' ', '-0.1', 'Infinity', '1e309', Number.NaN, Infinity, null, undefined, true]) {
|
||||
expect(parseBillingMultiplier(value)).toBeNull()
|
||||
}
|
||||
})
|
||||
|
||||
it('preserves cancellation policy across model scheduling edits', () => {
|
||||
@@ -59,14 +82,43 @@ describe('routingPolicy', () => {
|
||||
expect(next.model_policies[0].allowed_providers).toEqual(['provider-a'])
|
||||
})
|
||||
|
||||
it('normalizes model membership independently and preserves explicit false overrides', () => {
|
||||
const overrides = { enabled: true, disabled: false }
|
||||
const config = normalizeRoutingGroupConfig({ model_policies: [
|
||||
{ ...createEmptyModelPolicy('model-a'), provider_enabled_overrides: overrides },
|
||||
{ model: 'legacy-model' } as ReturnType<typeof createEmptyModelPolicy>,
|
||||
{ ...createEmptyModelPolicy('invalid-model'), provider_enabled_overrides: { '': true, valid: false, string: 'false' } as unknown as Record<string, boolean> },
|
||||
] })
|
||||
expect(config.model_policies[0].provider_enabled_overrides).toEqual(overrides)
|
||||
expect(config.model_policies[0].provider_enabled_overrides).not.toBe(overrides)
|
||||
expect(config.model_policies[1].provider_enabled_overrides).toEqual({})
|
||||
expect(config.model_policies[2].provider_enabled_overrides).toEqual({ valid: false })
|
||||
expect(createEmptyModelPolicy().provider_enabled_overrides).toEqual({})
|
||||
})
|
||||
|
||||
it('resolves model membership before default membership and legacy group exclusions', () => {
|
||||
const config = createEmptyRoutingGroupConfig()
|
||||
config.disabled_providers = ['provider-a', 'provider-b']
|
||||
config.model_policies = [{ ...createEmptyModelPolicy('*'), provider_enabled_overrides: { 'provider-a': true, 'provider-c': false } }]
|
||||
const selected = { ...createEmptyModelPolicy('model-a'), provider_enabled_overrides: { 'provider-b': true, 'provider-c': true, 'provider-d': false } }
|
||||
expect(isRoutingProviderEnabled(config, 'provider-a', selected)).toBe(true)
|
||||
expect(isRoutingProviderEnabled(config, 'provider-b', selected)).toBe(true)
|
||||
expect(isRoutingProviderEnabled(config, 'provider-c', selected)).toBe(true)
|
||||
expect(isRoutingProviderEnabled(config, 'provider-d', selected)).toBe(false)
|
||||
expect(isRoutingProviderEnabled(config, 'provider-b', createEmptyModelPolicy('model-b'))).toBe(false)
|
||||
expect(isRoutingProviderEnabled(config, 'provider-c')).toBe(false)
|
||||
expect(isRoutingProviderEnabled(config, 'provider-d')).toBe(true)
|
||||
})
|
||||
|
||||
it('stores default priority overrides on the wildcard model policy', () => {
|
||||
const config = upsertModelPolicy(createEmptyRoutingGroupConfig(), createEmptyModelPolicy('gpt-5'))
|
||||
const config = upsertModelPolicy({ ...createEmptyRoutingGroupConfig(), user_visible: true }, createEmptyModelPolicy('gpt-5'))
|
||||
const next = setDefaultProviderPriorityOverrides(config, {
|
||||
'provider-a': 0,
|
||||
'provider-b': 2,
|
||||
})
|
||||
|
||||
const policy = getDefaultModelPolicy(next)
|
||||
expect(next.user_visible).toBe(true)
|
||||
expect(policy.model).toBe(DEFAULT_ROUTING_POLICY_MODEL)
|
||||
expect(next.model_policies.map(item => item.model)).toEqual([DEFAULT_ROUTING_POLICY_MODEL, 'gpt-5'])
|
||||
expect(policy.provider_priority_overrides).toEqual({
|
||||
|
||||
@@ -5,6 +5,7 @@ import {
|
||||
getDefaultModelPolicy,
|
||||
getModelPolicy,
|
||||
getModelScheduling,
|
||||
isRoutingProviderEnabled,
|
||||
modelSchedulingRuleId,
|
||||
setDefaultProviderPriorityOverrides,
|
||||
setModelKeyPriorityOverridesForFormat,
|
||||
@@ -22,6 +23,23 @@ import {
|
||||
} from '../utils/schedulingPolicies'
|
||||
|
||||
describe('strategy-scoped scheduling policies', () => {
|
||||
it.each([false, true])('preserves user visibility %s and billing multipliers through scheduling changes and provider editor projections', userVisible => {
|
||||
for (const multiplier of [0, 1.25]) {
|
||||
const config = { ...createEmptyRoutingGroupConfig(), billing_multiplier: multiplier, user_visible: userVisible }
|
||||
const entry = createSchedulingPolicy(config, 'selected')
|
||||
entry.models = ['gpt-5']
|
||||
entry.schedulingMode = 'fixed_order'
|
||||
entry.policy.provider_priority_overrides = { provider: 2 }
|
||||
const updated = writeSchedulingPolicies(config, [entry])
|
||||
expect(updated.billing_multiplier).toBe(multiplier)
|
||||
expect(updated.user_visible).toBe(userVisible)
|
||||
expect(normalizeProviderSchedulingConfig(updated).billing_multiplier).toBe(multiplier)
|
||||
expect(normalizeProviderSchedulingConfig(updated).user_visible).toBe(userVisible)
|
||||
expect(schedulingPolicyEditorConfig(updated, readSchedulingPolicies(updated)[0]).billing_multiplier).toBe(multiplier)
|
||||
expect(schedulingPolicyEditorConfig(updated, readSchedulingPolicies(updated)[0]).user_visible).toBe(userVisible)
|
||||
}
|
||||
})
|
||||
|
||||
it('normalizes legacy Key scheduling without mutating unrelated rules, scopes, or historical priorities', () => {
|
||||
const config = createEmptyRoutingGroupConfig()
|
||||
config.default_policy.priority_mode = 'global_key'
|
||||
@@ -94,6 +112,7 @@ describe('strategy-scoped scheduling policies', () => {
|
||||
entry.policy = {
|
||||
...entry.policy,
|
||||
allowed_providers: ['provider-a'],
|
||||
provider_enabled_overrides: { 'provider-a': false, 'provider-b': true },
|
||||
provider_priority_overrides: { 'provider-a': 2 },
|
||||
key_priority_overrides_by_format: { 'openai:chat': { 'key-a': 1 } },
|
||||
pool_priority_overrides: { 'pool-a': 3 },
|
||||
@@ -126,6 +145,45 @@ describe('strategy-scoped scheduling policies', () => {
|
||||
expect(config.disabled_providers).toEqual(['disabled-provider'])
|
||||
})
|
||||
|
||||
it('preserves independent model membership through serialization without changing group defaults', () => {
|
||||
const config = createEmptyRoutingGroupConfig()
|
||||
config.disabled_providers = ['legacy-disabled']
|
||||
const first = { ...createSchedulingPolicy(config), models: ['model-a', 'model-b'] }
|
||||
first.policy.provider_enabled_overrides = { provider: false, 'legacy-disabled': true }
|
||||
const second = { ...createSchedulingPolicy(config), models: ['model-c'] }
|
||||
const fallback = createSchedulingPolicy(config, 'all')
|
||||
fallback.policy.provider_enabled_overrides = { 'default-disabled': false }
|
||||
const saved = writeSchedulingPolicies(config, [first, second, fallback])
|
||||
const reloaded = readSchedulingPolicies(JSON.parse(JSON.stringify(saved)))
|
||||
expect(reloaded).toHaveLength(3)
|
||||
expect(reloaded[0].policy.provider_enabled_overrides).toEqual(first.policy.provider_enabled_overrides)
|
||||
expect(reloaded[1].policy.provider_enabled_overrides).toEqual({})
|
||||
expect(saved.disabled_providers).toEqual(['legacy-disabled'])
|
||||
expect(getDefaultModelPolicy(saved).provider_enabled_overrides).toEqual({ 'default-disabled': false })
|
||||
for (const model of ['model-a', 'model-b']) {
|
||||
expect(getModelPolicy(saved, model).provider_enabled_overrides).toEqual(first.policy.provider_enabled_overrides)
|
||||
}
|
||||
const firstEditor = schedulingPolicyEditorConfig(saved, reloaded[0])
|
||||
const secondEditor = schedulingPolicyEditorConfig(saved, reloaded[1])
|
||||
expect(isRoutingProviderEnabled(firstEditor, 'provider')).toBe(false)
|
||||
expect(isRoutingProviderEnabled(secondEditor, 'provider')).toBe(true)
|
||||
expect(isRoutingProviderEnabled(firstEditor, 'legacy-disabled')).toBe(true)
|
||||
expect(isRoutingProviderEnabled(secondEditor, 'legacy-disabled')).toBe(false)
|
||||
expect(isRoutingProviderEnabled(firstEditor, 'default-disabled')).toBe(false)
|
||||
expect(getDefaultModelPolicy(firstEditor).provider_enabled_overrides).toEqual(first.policy.provider_enabled_overrides)
|
||||
getDefaultModelPolicy(firstEditor).provider_enabled_overrides.provider = true
|
||||
expect(getModelPolicy(saved, 'model-a').provider_enabled_overrides.provider).toBe(false)
|
||||
})
|
||||
|
||||
it('does not merge legacy models with different membership into one configuration', () => {
|
||||
const config = createEmptyRoutingGroupConfig()
|
||||
config.model_policies = [
|
||||
{ ...createEmptyModelPolicy('model-a'), provider_enabled_overrides: { provider: false } },
|
||||
{ ...createEmptyModelPolicy('model-b'), provider_enabled_overrides: { provider: true } },
|
||||
]
|
||||
expect(readSchedulingPolicies(config).map(entry => entry.models)).toEqual([['model-a'], ['model-b']])
|
||||
})
|
||||
|
||||
it('retains separate strategies even when their settings are identical', () => {
|
||||
const config = createEmptyRoutingGroupConfig()
|
||||
const first = { ...createSchedulingPolicy(config), models: ['model-a'] }
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
<template>
|
||||
<Popover v-model:open="open">
|
||||
<PopoverTrigger as-child>
|
||||
<Button
|
||||
type="button"
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-7 w-7 shrink-0 text-muted-foreground"
|
||||
:disabled="disabled"
|
||||
aria-label="编辑模型"
|
||||
title="编辑模型"
|
||||
>
|
||||
<Pencil class="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
</PopoverTrigger>
|
||||
<PopoverContent
|
||||
align="end"
|
||||
:side-offset="6"
|
||||
:collision-padding="12"
|
||||
class="w-[min(22rem,calc(100vw-1.5rem))] overflow-hidden p-0"
|
||||
aria-label="编辑适用模型"
|
||||
@open-auto-focus.prevent="focusSearch"
|
||||
>
|
||||
<div class="flex items-center justify-between gap-2 px-3 pt-2 text-xs font-medium">
|
||||
<span>适用模型</span>
|
||||
<Button
|
||||
type="button"
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-7 w-7 text-muted-foreground"
|
||||
aria-label="关闭模型选择"
|
||||
@click="open = false"
|
||||
>
|
||||
<X class="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
</div>
|
||||
<RoutingModelSelector
|
||||
ref="selector"
|
||||
inline
|
||||
compact
|
||||
narrow
|
||||
:model-value="modelValue"
|
||||
:models="models"
|
||||
:assigned-models="assignedModels"
|
||||
:loading="loading"
|
||||
:error="error"
|
||||
:disabled="disabled"
|
||||
@update:model-value="emit('update:modelValue', $event)"
|
||||
@reload="emit('reload')"
|
||||
@close="open = false"
|
||||
/>
|
||||
<div class="flex justify-end border-t border-border/50 px-2 py-1.5">
|
||||
<Button
|
||||
type="button"
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
class="h-7 px-2 text-xs"
|
||||
aria-label="完成选择"
|
||||
@click="open = false"
|
||||
>
|
||||
完成
|
||||
</Button>
|
||||
</div>
|
||||
</PopoverContent>
|
||||
</Popover>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, nextTick, ref, watch } from 'vue'
|
||||
import { Pencil, X } from 'lucide-vue-next'
|
||||
import { Button, Popover, PopoverContent, PopoverTrigger } from '@/components/ui'
|
||||
import type { GlobalModelResponse } from '@/api/global-models'
|
||||
import RoutingModelSelector from './RoutingModelSelector.vue'
|
||||
|
||||
const props = defineProps<{
|
||||
modelValue: string[]
|
||||
models: GlobalModelResponse[]
|
||||
assignedModels: Record<string, number>
|
||||
loading?: boolean
|
||||
error?: string | null
|
||||
disabled?: boolean
|
||||
open: boolean
|
||||
}>()
|
||||
|
||||
const emit = defineEmits<{
|
||||
'update:modelValue': [models: string[]]
|
||||
reload: []
|
||||
'update:open': [open: boolean]
|
||||
}>()
|
||||
|
||||
const open = computed({
|
||||
get: () => props.open,
|
||||
set: value => emit('update:open', value),
|
||||
})
|
||||
const selector = ref<InstanceType<typeof RoutingModelSelector> | null>(null)
|
||||
|
||||
watch(() => props.disabled, disabled => { if (disabled) open.value = false })
|
||||
|
||||
async function focusSearch(): Promise<void> {
|
||||
await nextTick()
|
||||
selector.value?.focusSearch()
|
||||
}
|
||||
</script>
|
||||
@@ -217,6 +217,7 @@ const props = defineProps<{
|
||||
const emit = defineEmits<{
|
||||
'update:modelValue': [models: string[]]
|
||||
reload: []
|
||||
close: []
|
||||
}>()
|
||||
|
||||
const trigger = ref<HTMLButtonElement | null>(null)
|
||||
@@ -255,15 +256,24 @@ async function openModels(): Promise<void> {
|
||||
if (props.disabled) return
|
||||
open.value = true
|
||||
await nextTick()
|
||||
searchInput.value?.inputRef?.focus({ preventScroll: true })
|
||||
focusSearch()
|
||||
}
|
||||
|
||||
function closeModels(): void {
|
||||
if (props.inline) return
|
||||
if (props.inline) {
|
||||
emit('close')
|
||||
return
|
||||
}
|
||||
open.value = false
|
||||
trigger.value?.focus({ preventScroll: true })
|
||||
}
|
||||
|
||||
function focusSearch(): void {
|
||||
searchInput.value?.inputRef?.focus({ preventScroll: true })
|
||||
}
|
||||
|
||||
defineExpose({ focusSearch })
|
||||
|
||||
function modelLabel(name: string): string {
|
||||
return props.models.find(model => model.name === name)?.display_name || name
|
||||
}
|
||||
|
||||
@@ -167,7 +167,7 @@
|
||||
class="rounded bg-muted px-1.5 py-0.5 text-[10px] text-muted-foreground"
|
||||
>停用</span>
|
||||
<span
|
||||
v-if="config.disabled_providers.includes(row.id)"
|
||||
v-if="!isRoutingProviderEnabled(config, row.id, targetModelPolicy)"
|
||||
class="rounded bg-muted px-1.5 py-0.5 text-[10px] text-muted-foreground"
|
||||
>本组禁用</span>
|
||||
</div>
|
||||
@@ -214,6 +214,7 @@ import {
|
||||
DEFAULT_ROUTING_POLICY_MODEL,
|
||||
getDefaultModelPolicy,
|
||||
getModelPolicy,
|
||||
isRoutingProviderEnabled,
|
||||
setModelProviderPriorityOverrides,
|
||||
type RoutingDefaultPolicy,
|
||||
type RoutingGroupConfig,
|
||||
|
||||
@@ -132,19 +132,19 @@
|
||||
<div
|
||||
role="group"
|
||||
aria-label="模型调度配置"
|
||||
class="space-y-1.5"
|
||||
class="divide-y divide-border/60 border-b border-border/60"
|
||||
>
|
||||
<section
|
||||
v-for="(entry, index) in entries"
|
||||
:key="entry.id"
|
||||
class="min-w-0 overflow-hidden rounded-md border transition-colors"
|
||||
:class="selectedEntryId === entry.id ? 'border-primary/40 bg-primary/5' : 'border-border/60 bg-background'"
|
||||
class="min-w-0"
|
||||
:aria-label="`调度配置 ${index + 1}`"
|
||||
>
|
||||
<div class="flex min-w-0 items-center">
|
||||
<button
|
||||
type="button"
|
||||
class="flex min-h-8 min-w-0 flex-1 items-center gap-2 rounded-md px-2 text-left text-xs hover:bg-muted/40 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-inset focus-visible:ring-ring disabled:opacity-50"
|
||||
class="flex min-h-10 min-w-0 flex-1 items-center gap-2 rounded-md px-1 text-left text-xs hover:bg-muted/40 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-inset focus-visible:ring-ring disabled:opacity-50"
|
||||
:class="selectedEntryId === entry.id ? 'text-primary' : 'text-foreground'"
|
||||
:aria-label="`选择调度配置 ${index + 1}`"
|
||||
:aria-pressed="selectedEntryId === entry.id"
|
||||
:aria-expanded="expandedId === entry.id"
|
||||
@@ -152,7 +152,7 @@
|
||||
@click="toggleEntry(entry.id)"
|
||||
>
|
||||
<ChevronRight
|
||||
class="h-3.5 w-3.5 shrink-0 text-muted-foreground transition-transform"
|
||||
class="h-3.5 w-3.5 shrink-0 transition-transform"
|
||||
:class="expandedId === entry.id ? 'rotate-90' : ''"
|
||||
/>
|
||||
<span
|
||||
@@ -161,6 +161,19 @@
|
||||
>{{ scopeSummary(entry) }}</span>
|
||||
<span class="shrink-0 text-muted-foreground">{{ schedulingModeLabel(entry.schedulingMode) }}</span>
|
||||
</button>
|
||||
<RoutingModelSelectionPopover
|
||||
v-if="entry.scope === 'selected'"
|
||||
:open="editingModelsId === entry.id"
|
||||
:model-value="entry.models"
|
||||
:models="globalModels"
|
||||
:assigned-models="otherModelOwners(entry.id)"
|
||||
:loading="loadingModels"
|
||||
:error="modelsError"
|
||||
:disabled="disabled"
|
||||
@update:open="open => setModelEditorOpen(entry.id, open)"
|
||||
@update:model-value="models => updateEntry(entry.id, { models })"
|
||||
@reload="emit('reload-models')"
|
||||
/>
|
||||
<Button
|
||||
v-if="entries.length > 1"
|
||||
type="button"
|
||||
@@ -178,32 +191,40 @@
|
||||
v-if="selectedEntryId === entry.id && expandedId === entry.id"
|
||||
role="region"
|
||||
aria-label="当前配置的适用模型"
|
||||
class="min-w-0 border-t border-border/50"
|
||||
class="min-w-0 px-1 pb-2"
|
||||
>
|
||||
<RoutingModelSelector
|
||||
<span
|
||||
v-if="entry.scope === 'selected'"
|
||||
compact
|
||||
inline
|
||||
:narrow="sidebar"
|
||||
:model-value="entry.models"
|
||||
:models="globalModels"
|
||||
:assigned-models="otherModelOwners(entry.id)"
|
||||
:loading="loadingModels"
|
||||
:error="modelsError"
|
||||
:disabled="disabled"
|
||||
@update:model-value="models => updateEntry(entry.id, { models })"
|
||||
@reload="emit('reload-models')"
|
||||
/>
|
||||
class="mb-1.5 block text-xs font-medium text-muted-foreground"
|
||||
>适用模型</span>
|
||||
<div
|
||||
v-if="entry.models.length"
|
||||
class="flex min-w-0 flex-wrap gap-1.5 py-1"
|
||||
aria-label="已配置模型"
|
||||
>
|
||||
<span
|
||||
v-for="name in entry.models"
|
||||
:key="name"
|
||||
class="max-w-full break-words rounded-md bg-muted/70 px-2 py-1 text-xs text-foreground [overflow-wrap:anywhere]"
|
||||
:title="name"
|
||||
>{{ modelDisplayName(name) }}</span>
|
||||
</div>
|
||||
<p
|
||||
v-else-if="entry.scope === 'selected'"
|
||||
class="py-2 text-xs leading-5 text-muted-foreground"
|
||||
>
|
||||
请选择适用模型
|
||||
</p>
|
||||
<p
|
||||
v-else
|
||||
class="px-2 py-2 text-xs leading-5 text-muted-foreground"
|
||||
class="py-2 text-xs leading-5 text-muted-foreground"
|
||||
>
|
||||
此默认配置适用于未单独指定的模型,新增模型也会自动使用。
|
||||
</p>
|
||||
</div>
|
||||
<div
|
||||
v-if="selectedEntryId === entry.id && expandedId === entry.id"
|
||||
class="min-w-0 space-y-1.5 border-t border-border/50 p-2"
|
||||
class="min-w-0 space-y-1.5 px-1 pb-3 pt-1"
|
||||
>
|
||||
<div class="flex h-6 items-center gap-1 text-xs font-medium text-muted-foreground">
|
||||
<span>调度策略</span>
|
||||
@@ -421,6 +442,7 @@ import HelpHint from '@/components/common/HelpHint.vue'
|
||||
import type { GlobalModelResponse } from '@/api/global-models'
|
||||
import RoutingPriorityPolicyEditor from './RoutingPriorityPolicyEditor.vue'
|
||||
import RoutingModelSelector from './RoutingModelSelector.vue'
|
||||
import RoutingModelSelectionPopover from './RoutingModelSelectionPopover.vue'
|
||||
import { getDefaultModelPolicy, normalizeRoutingGroupConfig, type RoutingGroupConfig, type RoutingModelPolicy, type RoutingPriorityMode, type RoutingSchedulingMode } from '../utils/routingPolicy'
|
||||
import {
|
||||
createSchedulingPolicy,
|
||||
@@ -490,6 +512,7 @@ const fallbackScheduling = {
|
||||
scheduling_mode: props.config.default_policy.scheduling_mode,
|
||||
}
|
||||
const expandedId = ref<string | null>(initialSelectedEntry?.id ?? null)
|
||||
const editingModelsId = ref<string | null>(null)
|
||||
const validationError = computed(() => validateSchedulingPolicies(entries.value))
|
||||
const hasAllModels = computed(() => entries.value.some(entry => entry.scope === 'all'))
|
||||
const assignedModels = computed(() => new Set(entries.value.filter(entry => entry.scope === 'selected').flatMap(entry => entry.models)))
|
||||
@@ -535,6 +558,7 @@ function emitSelection(): void {
|
||||
|
||||
function toggleEntry(id: string): void {
|
||||
if (props.disabled) return
|
||||
editingModelsId.value = null
|
||||
if (layout.value === 'config-only') {
|
||||
if (selectedEntryId.value === id) expandedId.value = expandedId.value === id ? null : id
|
||||
else selectEntry(id)
|
||||
@@ -548,11 +572,22 @@ function toggleEntry(id: string): void {
|
||||
|
||||
function selectEntry(id: string): void {
|
||||
if (props.disabled) return
|
||||
editingModelsId.value = null
|
||||
selectedEntryId.value = id
|
||||
expandedId.value = id
|
||||
emitSelection()
|
||||
}
|
||||
|
||||
function setModelEditorOpen(id: string, open: boolean): void {
|
||||
if (!open) {
|
||||
if (editingModelsId.value === id) editingModelsId.value = null
|
||||
return
|
||||
}
|
||||
if (props.disabled) return
|
||||
selectEntry(id)
|
||||
editingModelsId.value = id
|
||||
}
|
||||
|
||||
function updateSelectedPolicy(policy: RoutingModelPolicy): void {
|
||||
const entry = selectedEntry.value
|
||||
if (!entry || props.disabled || entry.scope === 'selected' && !entry.models.length) return
|
||||
@@ -571,10 +606,19 @@ function schedulingModeLabel(mode: RoutingSchedulingMode): string {
|
||||
function scopeSummary(entry: SchedulingPolicy): string {
|
||||
if (entry.scope === 'all') return '默认配置'
|
||||
if (entry.models.length === 0) return '请选择适用模型'
|
||||
if (layout.value === 'config-only') {
|
||||
const first = entry.models[0] ?? ''
|
||||
const label = props.globalModels.find(model => model.name === first)?.display_name || first
|
||||
return label + (entry.models.length > 1 ? ` +${entry.models.length - 1}` : '')
|
||||
}
|
||||
const labels = entry.models.slice(0, 2).map(name => props.globalModels.find(model => model.name === name)?.display_name || name)
|
||||
return labels.join('、') + (entry.models.length > 2 ? ` 等 ${entry.models.length} 个模型` : '')
|
||||
}
|
||||
|
||||
function modelDisplayName(name: string): string {
|
||||
return props.globalModels.find(model => model.name === name)?.display_name || name
|
||||
}
|
||||
|
||||
function otherModelOwners(entryId: string): Record<string, number> {
|
||||
return Object.fromEntries(entries.value.flatMap((entry, index) => entry.id !== entryId && entry.scope === 'selected'
|
||||
? entry.models.map(model => [model, index + 1])
|
||||
@@ -595,6 +639,7 @@ function publish(): void {
|
||||
|
||||
function setScopeMode(scope: SchedulingPolicy['scope']): void {
|
||||
if (props.disabled || scopeMode.value === scope) return
|
||||
editingModelsId.value = null
|
||||
if (scopeMode.value === 'all') allModelsDraft = entries.value
|
||||
else selectedModelsDraft = entries.value
|
||||
|
||||
@@ -636,6 +681,7 @@ function updateEntry(id: string, patch: Partial<SchedulingPolicy>): void {
|
||||
|
||||
function addEntry(): void {
|
||||
if (!canAddEntry.value) return
|
||||
editingModelsId.value = null
|
||||
const entry = createSchedulingPolicy(props.config)
|
||||
entries.value.push(entry)
|
||||
selectedEntryId.value = entry.id
|
||||
@@ -645,6 +691,7 @@ function addEntry(): void {
|
||||
|
||||
function removeEntry(id: string): void {
|
||||
if (props.disabled || entries.value.length === 1) return
|
||||
editingModelsId.value = null
|
||||
entries.value = entries.value.filter(entry => entry.id !== id)
|
||||
if (entries.value.every(entry => entry.scope === 'all')) {
|
||||
scopeMode.value = 'all'
|
||||
|
||||
@@ -33,6 +33,7 @@ export interface RoutingModelPolicy {
|
||||
model: string
|
||||
allowed_providers: string[]
|
||||
allowed_keys: string[]
|
||||
provider_enabled_overrides: Record<string, boolean>
|
||||
provider_priority_overrides: Record<string, number>
|
||||
key_priority_overrides: Record<string, number>
|
||||
/** api_format -> key_id -> priority;同一 Key 在不同 API 格式下可独立排序 */
|
||||
@@ -65,6 +66,8 @@ export interface RoutingSetSchedulingAction {
|
||||
}
|
||||
|
||||
export interface RoutingGroupConfig {
|
||||
billing_multiplier: number
|
||||
user_visible: boolean
|
||||
disabled_providers: string[]
|
||||
default_policy: RoutingDefaultPolicy
|
||||
model_policies: RoutingModelPolicy[]
|
||||
@@ -77,6 +80,8 @@ export const SCHEDULING_POLICY_RULE_PREFIX = 'ui_scheduling_policy:'
|
||||
|
||||
export function createEmptyRoutingGroupConfig(): RoutingGroupConfig {
|
||||
return {
|
||||
billing_multiplier: 1,
|
||||
user_visible: false,
|
||||
default_policy: {
|
||||
...normalizeRoutingFailoverPolicy(),
|
||||
priority_mode: 'provider',
|
||||
@@ -93,6 +98,13 @@ export function createEmptyRoutingGroupConfig(): RoutingGroupConfig {
|
||||
}
|
||||
}
|
||||
|
||||
export function parseBillingMultiplier(value: unknown): number | null {
|
||||
if (typeof value !== 'number' && typeof value !== 'string') return null
|
||||
if (typeof value === 'string' && !value.trim()) return null
|
||||
const parsed = Number(value)
|
||||
return Number.isFinite(parsed) && parsed >= 0 ? parsed : null
|
||||
}
|
||||
|
||||
export function normalizeStickyKeyAttempts(value: unknown): number {
|
||||
const parsed = Math.trunc(Number(value))
|
||||
if (!Number.isFinite(parsed) || parsed < 0) return DEFAULT_STICKY_KEY_ATTEMPTS
|
||||
@@ -104,6 +116,7 @@ export function createEmptyModelPolicy(model = ''): RoutingModelPolicy {
|
||||
model,
|
||||
allowed_providers: [],
|
||||
allowed_keys: [],
|
||||
provider_enabled_overrides: {},
|
||||
provider_priority_overrides: {},
|
||||
key_priority_overrides: {},
|
||||
key_priority_overrides_by_format: {},
|
||||
@@ -125,6 +138,8 @@ export function normalizeRoutingGroupConfig(value: Partial<RoutingGroupConfig> |
|
||||
} = rawDefaultPolicy
|
||||
|
||||
return {
|
||||
billing_multiplier: parseBillingMultiplier(value?.billing_multiplier) ?? base.billing_multiplier,
|
||||
user_visible: value?.user_visible === true,
|
||||
default_policy: {
|
||||
...base.default_policy,
|
||||
...defaultPolicyWithoutLegacyHeartbeat,
|
||||
@@ -145,6 +160,8 @@ export function normalizeRoutingGroupConfig(value: Partial<RoutingGroupConfig> |
|
||||
...policy,
|
||||
allowed_providers: Array.isArray(policy.allowed_providers) ? [...policy.allowed_providers] : [],
|
||||
allowed_keys: Array.isArray(policy.allowed_keys) ? [...policy.allowed_keys] : [],
|
||||
provider_enabled_overrides: Object.fromEntries(Object.entries(policy.provider_enabled_overrides ?? {})
|
||||
.filter(([id, enabled]) => id.length > 0 && typeof enabled === 'boolean')),
|
||||
provider_priority_overrides: { ...(policy.provider_priority_overrides ?? {}) },
|
||||
key_priority_overrides: { ...(policy.key_priority_overrides ?? {}) },
|
||||
key_priority_overrides_by_format: normalizeKeyPriorityOverridesByFormat(
|
||||
@@ -199,6 +216,16 @@ export function getModelPolicy(config: RoutingGroupConfig, model: string): Routi
|
||||
?? createEmptyModelPolicy(normalizedModel)
|
||||
}
|
||||
|
||||
export function isRoutingProviderEnabled(
|
||||
config: RoutingGroupConfig,
|
||||
providerId: string,
|
||||
policy?: RoutingModelPolicy | null,
|
||||
): boolean {
|
||||
return policy?.provider_enabled_overrides?.[providerId]
|
||||
?? getDefaultModelPolicy(config).provider_enabled_overrides[providerId]
|
||||
?? !config.disabled_providers.includes(providerId)
|
||||
}
|
||||
|
||||
export function upsertDefaultModelPolicy(
|
||||
config: RoutingGroupConfig,
|
||||
patch: Partial<Omit<RoutingModelPolicy, 'model'>>,
|
||||
|
||||
@@ -18,7 +18,9 @@ export interface RoutingCandidateTrace {
|
||||
|
||||
export interface RoutingDecisionTrace {
|
||||
group_id?: string | null
|
||||
group_name?: string | null
|
||||
group_version?: number | null
|
||||
billing_multiplier?: number | null
|
||||
selection_source: string
|
||||
selected_rules: string[]
|
||||
original_model: string
|
||||
|
||||
@@ -2,6 +2,7 @@ import {
|
||||
DEFAULT_ROUTING_POLICY_MODEL,
|
||||
SCHEDULING_POLICY_RULE_PREFIX,
|
||||
createEmptyModelPolicy,
|
||||
getDefaultModelPolicy,
|
||||
getModelPolicy,
|
||||
getModelScheduling,
|
||||
isGeneratedModelSchedulingRule,
|
||||
@@ -202,8 +203,16 @@ export function writeSchedulingPolicies(config: RoutingGroupConfig, entries: Sch
|
||||
}
|
||||
|
||||
export function schedulingPolicyEditorConfig(config: RoutingGroupConfig, entry: SchedulingPolicy): RoutingGroupConfig {
|
||||
// Project inherited membership for display without copying it into the editable policy.
|
||||
const disabledProviders = new Set(config.disabled_providers)
|
||||
for (const [providerId, enabled] of Object.entries(getDefaultModelPolicy(config).provider_enabled_overrides)) {
|
||||
if (enabled) disabledProviders.delete(providerId)
|
||||
else disabledProviders.add(providerId)
|
||||
}
|
||||
return normalizeRoutingGroupConfig({
|
||||
disabled_providers: config.disabled_providers,
|
||||
billing_multiplier: config.billing_multiplier,
|
||||
user_visible: config.user_visible,
|
||||
disabled_providers: [...disabledProviders],
|
||||
default_policy: {
|
||||
...config.default_policy,
|
||||
priority_mode: 'provider',
|
||||
|
||||
@@ -996,6 +996,10 @@ const emit = defineEmits<{
|
||||
cacheReadInputTokens?: number | null
|
||||
cost?: number | null
|
||||
actualCost?: number | null
|
||||
billingMultiplier?: number | null
|
||||
billingCost?: number | null
|
||||
routingGroupId?: string | null
|
||||
routingGroupName?: string | null
|
||||
responseTimeMs?: number | null
|
||||
firstByteTimeMs?: number | null
|
||||
isStream?: boolean | null
|
||||
@@ -1319,6 +1323,10 @@ function emitDetailRequestState(nextDetail: RequestDetail) {
|
||||
cacheReadInputTokens: nextDetail.cache_read_input_tokens ?? null,
|
||||
cost: detailTotalCost(nextDetail),
|
||||
actualCost: nextDetail.actual_cost ?? null,
|
||||
billingMultiplier: nextDetail.billing_multiplier,
|
||||
billingCost: nextDetail.billing_cost,
|
||||
routingGroupId: nextDetail.routing_group_id ?? null,
|
||||
routingGroupName: nextDetail.routing_group_name ?? null,
|
||||
responseTimeMs: nextDetail.response_time_ms ?? undefined,
|
||||
firstByteTimeMs: nextDetail.first_byte_time_ms ?? null,
|
||||
isStream: nextDetail.is_stream ?? null,
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
<template>
|
||||
<div
|
||||
class="flex flex-col items-end gap-0.5"
|
||||
:class="compact ? 'text-[10px]' : 'text-xs'"
|
||||
>
|
||||
<template v-if="record.usage_available !== false && record.usage_pricing_available !== false">
|
||||
<span
|
||||
data-usage-cost="base"
|
||||
class="text-primary"
|
||||
:class="compact ? 'text-sm font-semibold leading-5' : 'font-medium'"
|
||||
>{{ formatCurrency(record.cost || 0) }}</span>
|
||||
<span
|
||||
v-if="billing.cost !== null"
|
||||
data-usage-cost="routing-group"
|
||||
class="whitespace-nowrap text-muted-foreground"
|
||||
title="实际扣费"
|
||||
>{{ formatCurrency(billing.cost) }}</span>
|
||||
<span
|
||||
v-if="showKeyCost"
|
||||
data-usage-cost="provider-key"
|
||||
class="text-muted-foreground"
|
||||
:title="`提供商 Key 成本(Key 倍率 ${record.rate_multiplier}×)`"
|
||||
>{{ formatCurrency(record.actual_cost ?? 0) }}</span>
|
||||
</template>
|
||||
<span
|
||||
v-else-if="record.usage_available === false"
|
||||
data-usage-unavailable="cost"
|
||||
class="text-muted-foreground"
|
||||
:class="compact ? 'text-sm font-medium leading-5' : ''"
|
||||
title="上游未提供可验证的 token/费用用量"
|
||||
>不可用</span>
|
||||
<span
|
||||
v-else
|
||||
data-usage-unpriced="cost"
|
||||
class="text-muted-foreground"
|
||||
:class="compact ? 'text-sm font-medium leading-5' : ''"
|
||||
title="token 用量可验证,但当前计价规则不支持该音频用量分项"
|
||||
>未计价</span>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed } from 'vue'
|
||||
import type { UsageRecord } from '../types'
|
||||
import { formatCurrency } from '@/utils/format'
|
||||
import { resolveUsageBilling } from '../utils/usageBilling'
|
||||
|
||||
const props = defineProps<{
|
||||
record: UsageRecord
|
||||
showActualCost: boolean
|
||||
compact?: boolean
|
||||
}>()
|
||||
|
||||
const billing = computed(() => resolveUsageBilling(props.record))
|
||||
const showKeyCost = computed(() => props.showActualCost
|
||||
&& typeof props.record.actual_cost === 'number' && Number.isFinite(props.record.actual_cost)
|
||||
&& typeof props.record.rate_multiplier === 'number' && Number.isFinite(props.record.rate_multiplier)
|
||||
&& props.record.rate_multiplier >= 0 && props.record.rate_multiplier !== 1)
|
||||
</script>
|
||||
@@ -0,0 +1,24 @@
|
||||
<template>
|
||||
<div class="flex min-w-0 flex-col gap-0.5">
|
||||
<span
|
||||
data-usage-provider="routing-group"
|
||||
class="truncate text-foreground"
|
||||
:title="groupLabel"
|
||||
>{{ groupLabel }}</span>
|
||||
<span
|
||||
data-usage-provider="provider-key"
|
||||
class="truncate text-muted-foreground"
|
||||
:title="providerKeyLabel"
|
||||
>{{ providerKeyLabel }}</span>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed } from 'vue'
|
||||
import type { UsageRecord } from '../types'
|
||||
|
||||
const props = defineProps<{ record: UsageRecord }>()
|
||||
const groupLabel = computed(() => props.record.routing_group_name?.trim()
|
||||
|| props.record.routing_group_id?.trim() || '未记录分组')
|
||||
const providerKeyLabel = computed(() => `${props.record.provider?.trim() || '-'} · ${props.record.provider_key_name?.trim() || '-'}`)
|
||||
</script>
|
||||
@@ -293,28 +293,12 @@
|
||||
</Badge>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex flex-col items-end flex-shrink-0">
|
||||
<span
|
||||
v-if="record.usage_available !== false && record.usage_pricing_available !== false"
|
||||
class="text-sm text-primary font-semibold leading-5"
|
||||
>{{ formatCurrency(record.cost || 0) }}</span>
|
||||
<span
|
||||
v-else-if="record.usage_available === false"
|
||||
data-usage-unavailable="cost"
|
||||
class="text-sm text-muted-foreground font-medium leading-5"
|
||||
title="上游未提供可验证的 token/费用用量"
|
||||
>不可用</span>
|
||||
<span
|
||||
v-else
|
||||
data-usage-unpriced="cost"
|
||||
class="text-sm text-muted-foreground font-medium leading-5"
|
||||
title="token 用量可验证,但当前计价规则不支持该音频用量分项"
|
||||
>未计价</span>
|
||||
<span
|
||||
v-if="record.usage_available !== false && record.usage_pricing_available !== false && showActualCost && record.actual_cost !== undefined && record.rate_multiplier && record.rate_multiplier !== 1.0"
|
||||
class="text-[10px] text-muted-foreground"
|
||||
>{{ formatCurrency(record.actual_cost) }}</span>
|
||||
</div>
|
||||
<UsageCostDisplay
|
||||
:record="record"
|
||||
:show-actual-cost="showActualCost"
|
||||
compact
|
||||
class="shrink-0"
|
||||
/>
|
||||
</div>
|
||||
|
||||
<!-- 第二行:时间 + API格式 -->
|
||||
@@ -331,19 +315,22 @@
|
||||
</template>
|
||||
</div>
|
||||
|
||||
<!-- 第三行:用户 + 提供商 -->
|
||||
<!-- 用户与上游提供商信息 -->
|
||||
<div
|
||||
v-if="isAdmin"
|
||||
class="mt-1 flex min-w-0 items-center gap-1.5 text-[10px] leading-3.5 text-muted-foreground"
|
||||
class="mt-1 min-w-0 truncate text-[10px] leading-3.5 text-muted-foreground"
|
||||
:title="formatRecordUserSegment(record)"
|
||||
>
|
||||
<span
|
||||
class="min-w-0 truncate"
|
||||
:title="formatRecordUserProviderLine(record)"
|
||||
>
|
||||
{{ formatRecordUserSegment(record) }}
|
||||
</span>
|
||||
<span class="shrink-0 text-muted-foreground/40">·</span>
|
||||
<span class="min-w-0 truncate">{{ formatRecordProviderSegment(record) }}</span>
|
||||
{{ formatRecordUserSegment(record) }}
|
||||
</div>
|
||||
<div
|
||||
v-if="isAdmin"
|
||||
class="mt-1 flex min-w-0 items-center gap-1.5 text-[10px] leading-3.5"
|
||||
>
|
||||
<UsageProviderDisplay
|
||||
:record="record"
|
||||
class="flex-1"
|
||||
/>
|
||||
<!-- 手机与桌面保持相同的标记优先级:发生故障转移时优先显示转移标记。 -->
|
||||
<Shuffle
|
||||
v-if="record.has_fallback"
|
||||
@@ -774,20 +761,10 @@
|
||||
class="py-4 w-[16%]"
|
||||
>
|
||||
<div class="flex min-w-0 items-center gap-1">
|
||||
<div class="flex min-w-0 flex-col text-xs gap-0.5">
|
||||
<span class="truncate">{{ record.provider }}</span>
|
||||
<span
|
||||
v-if="record.provider_key_name"
|
||||
class="text-muted-foreground truncate"
|
||||
:title="record.provider_key_name"
|
||||
>
|
||||
{{ record.provider_key_name }}
|
||||
<span
|
||||
v-if="record.rate_multiplier && record.rate_multiplier !== 1.0"
|
||||
class="text-foreground/60"
|
||||
>({{ record.rate_multiplier }}x)</span>
|
||||
</span>
|
||||
</div>
|
||||
<UsageProviderDisplay
|
||||
:record="record"
|
||||
class="text-xs"
|
||||
/>
|
||||
<Shuffle
|
||||
v-if="record.has_fallback"
|
||||
data-usage-attempt-marker="fallback"
|
||||
@@ -974,34 +951,10 @@
|
||||
v-if="isColumnVisible('cost')"
|
||||
class="text-right py-4 w-[6%]"
|
||||
>
|
||||
<div
|
||||
v-if="record.usage_available !== false && record.usage_pricing_available !== false"
|
||||
class="flex flex-col items-end text-xs gap-0.5"
|
||||
>
|
||||
<span class="text-primary font-medium">{{ formatCurrency(record.cost || 0) }}</span>
|
||||
<span
|
||||
v-if="showActualCost && record.actual_cost !== undefined && record.rate_multiplier && record.rate_multiplier !== 1.0"
|
||||
class="text-muted-foreground"
|
||||
>
|
||||
{{ formatCurrency(record.actual_cost) }}
|
||||
</span>
|
||||
</div>
|
||||
<div
|
||||
v-else-if="record.usage_available === false"
|
||||
data-usage-unavailable="cost"
|
||||
class="text-xs text-muted-foreground"
|
||||
title="上游未提供可验证的 token/费用用量"
|
||||
>
|
||||
不可用
|
||||
</div>
|
||||
<div
|
||||
v-else
|
||||
data-usage-unpriced="cost"
|
||||
class="text-xs text-muted-foreground"
|
||||
title="token 用量可验证,但当前计价规则不支持该音频用量分项"
|
||||
>
|
||||
未计价
|
||||
</div>
|
||||
<UsageCostDisplay
|
||||
:record="record"
|
||||
:show-actual-cost="showActualCost"
|
||||
/>
|
||||
</TableCell>
|
||||
<TableCell
|
||||
v-if="isColumnVisible('performance')"
|
||||
@@ -1110,7 +1063,7 @@ import {
|
||||
TableFilterMenu,
|
||||
} from '@/components/ui'
|
||||
import { Ban, EyeOff, RefreshCcw, Search, Shuffle } from 'lucide-vue-next'
|
||||
import { formatTokens, formatCurrency } from '@/utils/format'
|
||||
import { formatTokens } from '@/utils/format'
|
||||
import { getCacheCreationTokens, getCacheReadTokens, getEffectiveInputTokens } from '../token-normalization'
|
||||
import {
|
||||
formatOutputRate,
|
||||
@@ -1140,6 +1093,8 @@ import type { MultiSelectOption } from '@/components/common/MultiSelect.vue'
|
||||
import ElapsedTimeText from './ElapsedTimeText.vue'
|
||||
import ServerUserSelector from './ServerUserSelector.vue'
|
||||
import UsageModelDisplay from './UsageModelDisplay.vue'
|
||||
import UsageCostDisplay from './UsageCostDisplay.vue'
|
||||
import UsageProviderDisplay from './UsageProviderDisplay.vue'
|
||||
|
||||
export interface UserOption {
|
||||
id: string
|
||||
@@ -1466,18 +1421,10 @@ function getRecordUserName(record: UsageRecord): string {
|
||||
return record.username || record.user_email || (record.user_id ? `User ${record.user_id}` : '已删除用户')
|
||||
}
|
||||
|
||||
function formatRecordUserProviderLine(record: UsageRecord): string {
|
||||
return `${formatRecordUserSegment(record)} · ${formatRecordProviderSegment(record)}`
|
||||
}
|
||||
|
||||
function formatRecordUserSegment(record: UsageRecord): string {
|
||||
return `${getRecordUserName(record)} / ${record.api_key?.name || '-'}`
|
||||
}
|
||||
|
||||
function formatRecordProviderSegment(record: UsageRecord): string {
|
||||
return `${record.provider || '-'} / ${record.provider_key_name || '-'}`
|
||||
}
|
||||
|
||||
watch(() => props.filterSearch, (value) => {
|
||||
if (value !== localSearch.value) {
|
||||
cancelPendingSearchEmit()
|
||||
|
||||
@@ -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<boolean>
|
||||
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(),
|
||||
|
||||
@@ -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<HTMLElement>('[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<HTMLElement>('[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()])
|
||||
|
||||
|
||||
@@ -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> = {}): 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<UsageRecord>) {
|
||||
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<UsageRecord>) {
|
||||
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 })
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user