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:
elky
2026-10-07 14:49:57 +08:00
parent 310098a853
commit 911c7f8875
110 changed files with 6524 additions and 559 deletions
@@ -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();