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())
+206 -11
View File
@@ -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(&current), 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(&current), 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(&current), 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(&current), 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(&current), &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 {
+119 -5
View File
@@ -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": []})] {
+274 -1
View File
@@ -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: &[],
+2
View File
@@ -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"
);
}