mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 03:09:50 +08:00
feat: add selectable routing groups and composite billing
Support per-model provider enablement and compact model editing. Capture request-time billing factors, charge customer costs separately, and preserve historical statistics without backfills.
This commit is contained in:
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user