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
@@ -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())