mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
feat: add selectable routing groups and composite billing
Support per-model provider enablement and compact model editing. Capture request-time billing factors, charge customer costs separately, and preserve historical statistics without backfills.
This commit is contained in:
@@ -2579,6 +2579,8 @@ mod tests {
|
||||
|
||||
let fixed_order_app = AppState::new().expect("state should build");
|
||||
let fixed_order_policy = ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("routing-group-fixed-order".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "test".to_string(),
|
||||
|
||||
@@ -362,6 +362,8 @@ mod tests {
|
||||
candidate.key_internal_priority = 3;
|
||||
candidate.key_global_priority_for_format = Some(2);
|
||||
let policy = aether_routing_core::ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "system_default".to_string(),
|
||||
@@ -399,6 +401,8 @@ mod tests {
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let policy = aether_routing_core::ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "system_default".to_string(),
|
||||
@@ -434,6 +438,8 @@ mod tests {
|
||||
candidate.key_internal_priority = 3;
|
||||
candidate.key_global_priority_for_format = Some(2);
|
||||
let policy = aether_routing_core::ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "system_default".to_string(),
|
||||
|
||||
@@ -1970,6 +1970,113 @@ mod tests {
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn model_provider_enablement_filters_candidate_pages_without_affecting_other_models() {
|
||||
let mut rows = Vec::new();
|
||||
for model in ["model-a", "model-b", "model-c"] {
|
||||
for (provider, priority) in [
|
||||
("provider-legacy-disabled", 0),
|
||||
("provider-model-disabled", 1),
|
||||
("provider-other", 2),
|
||||
("provider-inactive", 3),
|
||||
] {
|
||||
let mut row = standard_candidate_row(provider, "openai:chat", priority);
|
||||
row.global_model_id = format!("global-{model}");
|
||||
row.global_model_name = model.into();
|
||||
row.model_provider_model_name = model.into();
|
||||
row.model_id = format!("{provider}-{model}");
|
||||
row.provider_is_active = provider != "provider-inactive";
|
||||
rows.push(row);
|
||||
}
|
||||
}
|
||||
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows));
|
||||
let app = AppState::new().unwrap().with_data_state_for_tests(
|
||||
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository),
|
||||
);
|
||||
let auth = unrestricted_auth_snapshot();
|
||||
let directives = crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let config = serde_json::from_value(serde_json::json!({
|
||||
"disabled_providers": ["provider-legacy-disabled"],
|
||||
"model_policies": [
|
||||
{ "model": "model-a", "provider_enabled_overrides": {
|
||||
"provider-model-disabled": false, "provider-inactive": true
|
||||
} },
|
||||
{ "model": "model-b", "provider_enabled_overrides": {
|
||||
"provider-legacy-disabled": true, "provider-inactive": true
|
||||
} }
|
||||
],
|
||||
"rules": [{ "id": "legacy-allowlist", "actions": [{
|
||||
"type": "restrict_providers", "provider_ids": [
|
||||
"provider-legacy-disabled", "provider-model-disabled", "provider-other", "provider-inactive"
|
||||
]
|
||||
}] }]
|
||||
})).unwrap();
|
||||
// Revisit A after B to exercise candidate caches shared by the app.
|
||||
for (model, expected) in [
|
||||
("model-a", vec!["provider-other"]),
|
||||
(
|
||||
"model-b",
|
||||
vec![
|
||||
"provider-legacy-disabled",
|
||||
"provider-model-disabled",
|
||||
"provider-other",
|
||||
],
|
||||
),
|
||||
("model-c", vec!["provider-model-disabled", "provider-other"]),
|
||||
("model-a", vec!["provider-other"]),
|
||||
] {
|
||||
let policy = aether_routing_core::resolve_routing_policy(
|
||||
&config,
|
||||
aether_routing_core::RoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
selection_source: "test",
|
||||
requested_model: model,
|
||||
resolved_model: model,
|
||||
api_format: "openai:chat",
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
headers: &serde_json::json!({}),
|
||||
body: &serde_json::json!({}),
|
||||
phase: aether_routing_core::RoutingRulePhase::ClientRequest,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let mut cursor = LocalCandidatePreselectionPageCursor::new(
|
||||
PlannerAppState::new(&app),
|
||||
&directives,
|
||||
"openai:chat",
|
||||
model,
|
||||
None,
|
||||
false,
|
||||
None,
|
||||
&auth,
|
||||
Some(&policy),
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
true,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
let mut providers = Vec::new();
|
||||
while let Some(page) = cursor.next_page().await.unwrap() {
|
||||
providers.extend(
|
||||
page.candidates
|
||||
.into_iter()
|
||||
.map(|candidate| candidate.provider_id),
|
||||
);
|
||||
}
|
||||
providers.sort();
|
||||
assert_eq!(
|
||||
providers, expected,
|
||||
"provider enablement must remain isolated for {model}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn routing_policy_collects_candidate_pages_before_final_ranking() {
|
||||
let rows = (0..300)
|
||||
@@ -1992,6 +2099,8 @@ mod tests {
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let routing_policy = ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("routing-group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "test".to_string(),
|
||||
@@ -2057,6 +2166,8 @@ mod tests {
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let routing_policy = ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("routing-group-fallback".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "test".to_string(),
|
||||
@@ -2899,6 +3010,8 @@ mod tests {
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let routing_policy = ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("routing-group-codex-first".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "test".to_string(),
|
||||
|
||||
@@ -555,24 +555,42 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
input.provider_outbound_context =
|
||||
Some(crate::ai_serving::codex_context::resolve_codex_fingerprint_context(parts, body_json));
|
||||
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
|
||||
let preferred_group = if explicit_group.is_none() && !input.auth_context.api_key_is_standalone {
|
||||
state
|
||||
.read_auth_api_key_feature_settings(
|
||||
&input.auth_context.user_id,
|
||||
&input.auth_context.api_key_id,
|
||||
false,
|
||||
)
|
||||
.await?
|
||||
.as_ref()
|
||||
.and_then(|settings| settings.get("routing_group_id"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_owned)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let selected_group = match state.routing_group_read_repository() {
|
||||
Some(repository) => {
|
||||
// Explicit non-default groups are authorized against principal
|
||||
// bindings, so both selection and its cache key must retain the
|
||||
// caller context. Only the implicit no-binding system-default
|
||||
// path is global and can skip the membership lookup.
|
||||
let principal_context_required = if explicit_group.is_some() {
|
||||
true
|
||||
} else {
|
||||
repository
|
||||
.has_any_routing_group_binding()
|
||||
.await
|
||||
.map_err(|error| {
|
||||
routing_selection_error(GatewayRoutingSelectionError::Repository(
|
||||
error.to_string(),
|
||||
))
|
||||
})?
|
||||
};
|
||||
let principal_context_required =
|
||||
if explicit_group.is_some() || preferred_group.is_some() {
|
||||
true
|
||||
} else {
|
||||
repository
|
||||
.has_any_routing_group_binding()
|
||||
.await
|
||||
.map_err(|error| {
|
||||
routing_selection_error(GatewayRoutingSelectionError::Repository(
|
||||
error.to_string(),
|
||||
))
|
||||
})?
|
||||
};
|
||||
let user_group_ids = if principal_context_required {
|
||||
let user_groups_lookup_started_at = std::time::Instant::now();
|
||||
let user_groups = state
|
||||
@@ -595,6 +613,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
principal_context_required.then(|| input.auth_context.api_key_id.clone());
|
||||
let selection_cache_key = routing_group_selection_cache_key(
|
||||
explicit_group.as_deref(),
|
||||
preferred_group.as_deref(),
|
||||
selection_user_id.as_deref(),
|
||||
selection_api_key_id.as_deref(),
|
||||
&user_group_ids,
|
||||
@@ -612,6 +631,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
repository.as_ref(),
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: explicit_group.as_deref(),
|
||||
preferred_group: preferred_group.as_deref(),
|
||||
user_id: selection_user_id.as_deref(),
|
||||
api_key_id: selection_api_key_id.as_deref(),
|
||||
user_group_ids: &user_group_ids,
|
||||
@@ -628,6 +648,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
|| {
|
||||
let repository = repository.clone();
|
||||
let explicit_group = explicit_group.clone();
|
||||
let preferred_group = preferred_group.clone();
|
||||
let user_id = selection_user_id.clone();
|
||||
let api_key_id = selection_api_key_id.clone();
|
||||
let user_group_ids = user_group_ids.clone();
|
||||
@@ -637,6 +658,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
repository.as_ref(),
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: explicit_group.as_deref(),
|
||||
preferred_group: preferred_group.as_deref(),
|
||||
user_id: user_id.as_deref(),
|
||||
api_key_id: api_key_id.as_deref(),
|
||||
user_group_ids: &user_group_ids,
|
||||
@@ -662,6 +684,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
selection.group.map(|group| {
|
||||
(
|
||||
Some(group.id),
|
||||
group.name,
|
||||
Some(group.version),
|
||||
group.config_json,
|
||||
selection.source,
|
||||
@@ -669,13 +692,14 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
})
|
||||
}
|
||||
None => {
|
||||
if explicit_group
|
||||
if let Some(requested_group) = explicit_group
|
||||
.or(preferred_group)
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return Err(routing_selection_error(
|
||||
GatewayRoutingSelectionError::NotFound(explicit_group.unwrap_or_default()),
|
||||
GatewayRoutingSelectionError::NotFound(requested_group.to_string()),
|
||||
));
|
||||
}
|
||||
return Err(routing_selection_error(
|
||||
@@ -684,7 +708,8 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
}
|
||||
};
|
||||
|
||||
let Some((group_id, group_version, group_config_json, selection_source)) = selected_group
|
||||
let Some((group_id, group_name, group_version, group_config_json, selection_source)) =
|
||||
selected_group
|
||||
else {
|
||||
return Err(routing_selection_error(
|
||||
GatewayRoutingSelectionError::NoDefault,
|
||||
@@ -701,6 +726,12 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
&group_config_json,
|
||||
selection_source.as_str(),
|
||||
)? {
|
||||
if let Some(policy) = input.routing_policy.as_mut() {
|
||||
policy.group_name = Some(group_name.clone());
|
||||
}
|
||||
if let Some(trace) = input.routing_trace_seed.as_mut() {
|
||||
trace.group_name = Some(group_name);
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
@@ -786,6 +817,7 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
final_policy_resolve_started_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
final_policy.mutation_plan = policy.mutation_plan.clone();
|
||||
final_policy.group_name = Some(group_name);
|
||||
input.routing_trace_seed = Some(build_routing_trace_seed(&final_policy, client_api_format));
|
||||
input.routing_policy = Some(final_policy);
|
||||
input.routing_context = Some(LocalRoutingRequestContext {
|
||||
@@ -966,6 +998,7 @@ fn routing_header_value_str(headers: &http::HeaderMap, key: &str) -> Option<Stri
|
||||
|
||||
fn routing_group_selection_cache_key(
|
||||
explicit_group: Option<&str>,
|
||||
preferred_group: Option<&str>,
|
||||
user_id: Option<&str>,
|
||||
api_key_id: Option<&str>,
|
||||
user_group_ids: &[String],
|
||||
@@ -976,8 +1009,9 @@ fn routing_group_selection_cache_key(
|
||||
.collect::<Vec<_>>()
|
||||
.join(",");
|
||||
format!(
|
||||
"v1|explicit={}|user={}|api_key={}|groups={}",
|
||||
"v2|explicit={}|preferred={}|user={}|api_key={}|groups={}",
|
||||
escape_cache_key_part(explicit_group.unwrap_or_default()),
|
||||
escape_cache_key_part(preferred_group.unwrap_or_default()),
|
||||
escape_cache_key_part(user_id.unwrap_or_default()),
|
||||
escape_cache_key_part(api_key_id.unwrap_or_default()),
|
||||
groups
|
||||
@@ -1175,10 +1209,13 @@ mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::*;
|
||||
use aether_data::repository::auth::{
|
||||
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
|
||||
};
|
||||
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
|
||||
use aether_data_contracts::repository::routing_profiles::{
|
||||
CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, RoutingGroupBindingSubject,
|
||||
RoutingGroupWriteRepository,
|
||||
RoutingGroupWriteRepository, UpdateRoutingGroupRecord,
|
||||
};
|
||||
use aether_provider_transport::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -1189,12 +1226,14 @@ mod tests {
|
||||
fn explicit_routing_selection_cache_key_is_principal_specific() {
|
||||
let first = routing_group_selection_cache_key(
|
||||
Some("private"),
|
||||
None,
|
||||
Some("user-1"),
|
||||
Some("key-1"),
|
||||
&["team-1".to_string()],
|
||||
);
|
||||
let second = routing_group_selection_cache_key(
|
||||
Some("private"),
|
||||
None,
|
||||
Some("user-2"),
|
||||
Some("key-2"),
|
||||
&["team-2".to_string()],
|
||||
@@ -1311,6 +1350,160 @@ mod tests {
|
||||
assert!(matches!(error, GatewayError::Internal(message) if message.contains("identity")));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn api_key_routing_selection_applies_at_planner_and_invalidates_after_changes() {
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(
|
||||
["api-key-1", "api-key-2"].map(|key_id| {
|
||||
(
|
||||
None,
|
||||
StoredAuthApiKeySnapshot::new(
|
||||
"user-1".into(),
|
||||
"alice".into(),
|
||||
None,
|
||||
"user".into(),
|
||||
"local".into(),
|
||||
true,
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
key_id.into(),
|
||||
Some(key_id.into()),
|
||||
true,
|
||||
false,
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}),
|
||||
));
|
||||
let groups = Arc::new(InMemoryRoutingGroupRepository::default());
|
||||
for (id, visible, is_default, multiplier) in [
|
||||
("default", false, true, 1.0),
|
||||
("discount", true, false, 0.5),
|
||||
("premium", true, false, 2.0),
|
||||
] {
|
||||
groups.create_routing_group(CreateRoutingGroupRecord {
|
||||
id: id.into(), name: format!("{id}-name"), description: None,
|
||||
enabled: true, is_system_default: is_default, sort_order: 0,
|
||||
config_json: json!({ "user_visible": visible, "billing_multiplier": multiplier }),
|
||||
version: 1, created_at: 1, updated_at: 1, published_at: None,
|
||||
}).await.unwrap();
|
||||
}
|
||||
let state = AppState::new().unwrap().with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository)
|
||||
.with_routing_group_repository_for_tests(groups.clone()),
|
||||
);
|
||||
for (key_id, group_id) in [("api-key-1", "discount"), ("api-key-2", "premium")] {
|
||||
assert!(state
|
||||
.set_user_api_key_feature_settings(
|
||||
"user-1",
|
||||
key_id,
|
||||
Some(json!({ "routing_group_id": group_id }))
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some());
|
||||
}
|
||||
let (parts, _) = http::Request::builder().body(()).unwrap().into_parts();
|
||||
let (header_parts, _) = http::Request::builder()
|
||||
.header(ROUTING_GROUP_HEADER, "premium")
|
||||
.body(())
|
||||
.unwrap()
|
||||
.into_parts();
|
||||
|
||||
async fn attach(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
key_id: &str,
|
||||
) -> Result<LocalRequestedModelDecisionInput, GatewayError> {
|
||||
let mut input = sample_decision_input();
|
||||
input.auth_context.api_key_id = key_id.into();
|
||||
input.auth_snapshot.api_key_id = key_id.into();
|
||||
attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
&json!({ "model": "gpt-5" }),
|
||||
"openai:chat",
|
||||
)
|
||||
.await?;
|
||||
Ok(input)
|
||||
}
|
||||
|
||||
// Revisit the first key after the second to exercise both cached choices.
|
||||
for (key_id, group_id, multiplier) in [
|
||||
("api-key-1", "discount", 0.5),
|
||||
("api-key-2", "premium", 2.0),
|
||||
("api-key-1", "discount", 0.5),
|
||||
] {
|
||||
let input = attach(&state, &parts, key_id).await.unwrap();
|
||||
let policy = input.routing_policy.as_ref().unwrap();
|
||||
assert_eq!(policy.group_id.as_deref(), Some(group_id));
|
||||
assert_eq!(policy.selection_source, "api_key_selection");
|
||||
assert_eq!(policy.billing_multiplier, multiplier);
|
||||
assert_eq!(
|
||||
input
|
||||
.routing_trace_seed
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.billing_multiplier,
|
||||
Some(multiplier)
|
||||
);
|
||||
}
|
||||
let header = attach(&state, &header_parts, "api-key-1").await.unwrap();
|
||||
let policy = header.routing_policy.unwrap();
|
||||
assert_eq!(policy.group_id.as_deref(), Some("premium"));
|
||||
assert_eq!(policy.selection_source, "explicit_header");
|
||||
|
||||
groups
|
||||
.update_routing_group(
|
||||
"discount",
|
||||
UpdateRoutingGroupRecord {
|
||||
config_json: Some(json!({ "user_visible": false, "billing_multiplier": 0.5 })),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
state.invalidate_provider_routing_caches();
|
||||
assert!(matches!(
|
||||
attach(&state, &parts, "api-key-1").await,
|
||||
Err(GatewayError::Client {
|
||||
status: StatusCode::FORBIDDEN,
|
||||
..
|
||||
})
|
||||
));
|
||||
let header = attach(&state, &header_parts, "api-key-1").await.unwrap();
|
||||
assert_eq!(
|
||||
header.routing_policy.unwrap().group_id.as_deref(),
|
||||
Some("premium")
|
||||
);
|
||||
|
||||
assert!(state
|
||||
.set_user_api_key_feature_settings("user-1", "api-key-1", None)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some());
|
||||
let cleared = attach(&state, &parts, "api-key-1").await.unwrap();
|
||||
let policy = cleared.routing_policy.unwrap();
|
||||
assert_eq!(policy.group_id.as_deref(), Some("default"));
|
||||
assert_eq!(policy.selection_source, "system_default");
|
||||
assert_eq!(policy.billing_multiplier, 1.0);
|
||||
// Clearing one key's preference must not disturb the other key's selection.
|
||||
let other = attach(&state, &parts, "api-key-2").await.unwrap();
|
||||
assert_eq!(
|
||||
other.routing_policy.unwrap().group_id.as_deref(),
|
||||
Some("premium")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn explicit_routing_attachment_authorizes_and_caches_per_principal() {
|
||||
let repository = Arc::new(InMemoryRoutingGroupRepository::default());
|
||||
@@ -1322,7 +1515,7 @@ mod tests {
|
||||
enabled: true,
|
||||
is_system_default: false,
|
||||
sort_order: 0,
|
||||
config_json: json!({}),
|
||||
config_json: json!({"billing_multiplier": 0.5}),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
@@ -1368,6 +1561,11 @@ mod tests {
|
||||
.as_ref()
|
||||
.expect("explicit selection should attach routing policy");
|
||||
assert_eq!(policy.group_id.as_deref(), Some("private-group"));
|
||||
assert_eq!(policy.group_name.as_deref(), Some("private"));
|
||||
assert_eq!(policy.billing_multiplier, 0.5);
|
||||
let trace = allowed.routing_trace_seed.as_ref().unwrap();
|
||||
assert_eq!(trace.group_name.as_deref(), Some("private"));
|
||||
assert_eq!(trace.billing_multiplier, Some(0.5));
|
||||
assert_eq!(policy.selection_source, "explicit_header");
|
||||
|
||||
let mut denied = sample_decision_input();
|
||||
|
||||
@@ -6,6 +6,11 @@ use aether_ai_serving::{
|
||||
provider_stream_event_api_format_for_provider_type as ai_provider_stream_event_api_format_for_provider_type,
|
||||
AiExecutionReportContextParts, AiRequestOrigin, STICKY_KEY_ATTEMPTS_REPORT_FIELD,
|
||||
};
|
||||
use aether_data_contracts::repository::usage::{
|
||||
BillingMultiplierSnapshot, BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY,
|
||||
ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY, ROUTING_GROUP_ID_METADATA_KEY,
|
||||
ROUTING_GROUP_NAME_METADATA_KEY,
|
||||
};
|
||||
use aether_routing_core::ResolvedRoutingPolicy;
|
||||
use aether_runtime_state::RuntimeLockLease;
|
||||
use aether_scheduler_core::{ClientSessionAffinity, SchedulerRankingOutcome};
|
||||
@@ -87,6 +92,46 @@ pub(crate) fn build_local_execution_report_context(
|
||||
parts.original_request_body_base64,
|
||||
);
|
||||
let mut extra_fields = parts.extra_fields;
|
||||
// Always overwrite caller-supplied extras with the planner's immutable policy snapshot.
|
||||
let billing_multiplier = parts
|
||||
.routing_policy
|
||||
.map(|policy| policy.billing_multiplier)
|
||||
.filter(|value| value.is_finite() && *value >= 0.0)
|
||||
.unwrap_or(1.0);
|
||||
extra_fields.insert(
|
||||
ROUTING_GROUP_BILLING_MULTIPLIER_METADATA_KEY.to_string(),
|
||||
Value::from(billing_multiplier),
|
||||
);
|
||||
extra_fields.insert(
|
||||
BILLING_MULTIPLIER_SNAPSHOT_METADATA_KEY.to_string(),
|
||||
serde_json::to_value(
|
||||
BillingMultiplierSnapshot::from_factors(BTreeMap::from([(
|
||||
"routing_group".to_string(),
|
||||
billing_multiplier,
|
||||
)]))
|
||||
.expect("validated routing multiplier must produce a billing snapshot"),
|
||||
)
|
||||
.expect("validated billing snapshot must serialize"),
|
||||
);
|
||||
for (field, value) in [
|
||||
(
|
||||
ROUTING_GROUP_ID_METADATA_KEY,
|
||||
parts
|
||||
.routing_policy
|
||||
.and_then(|policy| policy.group_id.as_deref()),
|
||||
),
|
||||
(
|
||||
ROUTING_GROUP_NAME_METADATA_KEY,
|
||||
parts
|
||||
.routing_policy
|
||||
.and_then(|policy| policy.group_name.as_deref()),
|
||||
),
|
||||
] {
|
||||
extra_fields.remove(field);
|
||||
if let Some(value) = value {
|
||||
extra_fields.insert(field.to_string(), Value::String(value.to_string()));
|
||||
}
|
||||
}
|
||||
if let Some(value) = parts
|
||||
.client_session_affinity
|
||||
.and_then(client_session_affinity_report_context_value)
|
||||
@@ -341,6 +386,27 @@ mod tests {
|
||||
Some("codex".to_string()),
|
||||
Some("account=account-1;session=session-1".to_string()),
|
||||
);
|
||||
let mut routing_policy = aether_routing_core::resolve_routing_policy(
|
||||
&aether_routing_core::RoutingGroupConfig {
|
||||
billing_multiplier: 0.25,
|
||||
..Default::default()
|
||||
},
|
||||
aether_routing_core::RoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(7),
|
||||
selection_source: "system_default",
|
||||
requested_model: "gpt-5",
|
||||
resolved_model: "gpt-5",
|
||||
api_format: "openai:chat",
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
headers: &json!({}),
|
||||
body: &json!({}),
|
||||
phase: aether_routing_core::RoutingRulePhase::ClientRequest,
|
||||
},
|
||||
)
|
||||
.expect("routing policy should resolve");
|
||||
routing_policy.group_name = Some("请求时的分组".to_string());
|
||||
|
||||
let report_context =
|
||||
build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
@@ -379,16 +445,35 @@ mod tests {
|
||||
original_request_body_json: Some(&json!({"model": "gpt-5"})),
|
||||
original_request_body_base64: None,
|
||||
client_session_affinity: Some(&client_session_affinity),
|
||||
routing_policy: None,
|
||||
routing_policy: Some(&routing_policy),
|
||||
scheduler_affinity_epoch: None,
|
||||
sticky_key_attempts: None,
|
||||
client_requested_stream: false,
|
||||
upstream_is_stream: false,
|
||||
has_envelope: false,
|
||||
needs_conversion: false,
|
||||
extra_fields: Map::new(),
|
||||
extra_fields: Map::from_iter([
|
||||
(
|
||||
"billing_multiplier_snapshot".to_string(),
|
||||
json!({
|
||||
"version": 1, "factors": {"routing_group": 99.0}, "multiplier": 99.0
|
||||
}),
|
||||
),
|
||||
("routing_group_billing_multiplier".to_string(), json!(99)),
|
||||
("routing_group_id".to_string(), json!("forged-group")),
|
||||
("routing_group_name".to_string(), json!("forged-name")),
|
||||
]),
|
||||
});
|
||||
|
||||
assert_eq!(report_context["routing_group_billing_multiplier"], 0.25);
|
||||
assert_eq!(
|
||||
report_context["billing_multiplier_snapshot"],
|
||||
json!({
|
||||
"version": 1, "factors": {"routing_group": 0.25}, "multiplier": 0.25
|
||||
})
|
||||
);
|
||||
assert_eq!(report_context["routing_group_id"], "group-1");
|
||||
assert_eq!(report_context["routing_group_name"], "请求时的分组");
|
||||
assert_eq!(
|
||||
report_context["client_ip"],
|
||||
Value::String("203.0.113.8".to_string())
|
||||
|
||||
@@ -231,8 +231,32 @@ pub(crate) async fn estimate_execution_plan_cost_upper_bound_usd(
|
||||
report_context: Option<&serde_json::Value>,
|
||||
) -> Result<Option<f64>, GatewayError> {
|
||||
let started_at = std::time::Instant::now();
|
||||
let result =
|
||||
estimate_execution_plan_cost_upper_bound_usd_inner(state, plan, report_context).await;
|
||||
let result = async {
|
||||
let multiplier_snapshot =
|
||||
aether_data_contracts::repository::usage::billing_multiplier_snapshot(report_context)
|
||||
.map_err(|error| GatewayError::Internal(error.to_string()))?;
|
||||
let estimate = estimate_execution_plan_cost_upper_bound_usd_inner(
|
||||
state,
|
||||
plan,
|
||||
report_context,
|
||||
multiplier_snapshot.is_some(),
|
||||
)
|
||||
.await?;
|
||||
let Some(snapshot) = multiplier_snapshot else {
|
||||
return Ok(estimate);
|
||||
};
|
||||
// Cache the unmultiplied base estimate so different request snapshots
|
||||
// cannot reuse one another's charge. Pricing validation still runs for
|
||||
// a zero multiplier, even when the request has no finite token bound.
|
||||
if snapshot.multiplier() == 0.0 {
|
||||
return Ok(Some(0.0));
|
||||
}
|
||||
estimate
|
||||
.map(|cost| snapshot.cost(cost))
|
||||
.transpose()
|
||||
.map_err(|error| GatewayError::Internal(error.to_string()))
|
||||
}
|
||||
.await;
|
||||
observe_gateway_stage_ms(
|
||||
"auth_capacity_cost_estimate",
|
||||
started_at.elapsed().as_millis() as u64,
|
||||
@@ -244,6 +268,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner(
|
||||
state: &AppState,
|
||||
plan: &aether_contracts::ExecutionPlan,
|
||||
report_context: Option<&serde_json::Value>,
|
||||
use_base_cost: bool,
|
||||
) -> Result<Option<f64>, GatewayError> {
|
||||
let api_format = crate::ai_serving::normalize_api_format_alias(&plan.provider_api_format);
|
||||
let body_json = plan.body.json_body.as_ref();
|
||||
@@ -311,7 +336,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner(
|
||||
if model_id.is_none() && global_model_name.is_none() {
|
||||
return Ok(None);
|
||||
}
|
||||
let cache_key = execution_plan_cost_upper_bound_cache_key(
|
||||
let mut cache_key = execution_plan_cost_upper_bound_cache_key(
|
||||
plan,
|
||||
model_id,
|
||||
global_model_name,
|
||||
@@ -321,6 +346,11 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner(
|
||||
requested_processing_tier.as_deref(),
|
||||
cache_ttl_minutes,
|
||||
);
|
||||
if use_base_cost {
|
||||
// Legacy requests cache provider Key cost; new requests cache base cost.
|
||||
// These values must never share a cache entry for the same provider Key.
|
||||
cache_key.insert_str(0, "base\x1f");
|
||||
}
|
||||
let ttl = state.frontdoor_runtime_guards.auth_capacity_cache_ttl;
|
||||
if ttl.is_zero() {
|
||||
let _permit = state.acquire_auth_snapshot_load_gate().await?;
|
||||
@@ -335,6 +365,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner(
|
||||
max_output_tokens,
|
||||
requested_processing_tier.as_deref(),
|
||||
cache_ttl_minutes,
|
||||
use_base_cost,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -353,6 +384,7 @@ async fn estimate_execution_plan_cost_upper_bound_usd_inner(
|
||||
max_output_tokens,
|
||||
requested_processing_tier.as_deref(),
|
||||
cache_ttl_minutes,
|
||||
use_base_cost,
|
||||
)
|
||||
.await
|
||||
})
|
||||
@@ -371,6 +403,7 @@ async fn calculate_execution_plan_cost_upper_bound(
|
||||
max_output_tokens: Option<i64>,
|
||||
requested_processing_tier: Option<&str>,
|
||||
cache_ttl_minutes: Option<i64>,
|
||||
use_base_cost: bool,
|
||||
) -> Result<Option<f64>, GatewayError> {
|
||||
let context =
|
||||
load_execution_plan_billing_context(state, plan, model_id, global_model_name).await?;
|
||||
@@ -383,11 +416,13 @@ async fn calculate_execution_plan_cost_upper_bound(
|
||||
estimate.requested_processing_tier = requested_processing_tier.map(ToOwned::to_owned);
|
||||
estimate.cache_ttl_minutes = cache_ttl_minutes;
|
||||
estimate.max_output_tokens = max_output_tokens;
|
||||
let mut pricing = aether_billing::BillingModelPricingSnapshot::from(context);
|
||||
if use_base_cost {
|
||||
pricing.provider_billing_type = None;
|
||||
pricing.provider_api_key_rate_multipliers = None;
|
||||
}
|
||||
aether_billing::BillingService::new()
|
||||
.estimate_authorization_cost_upper_bound(
|
||||
&aether_billing::BillingModelPricingSnapshot::from(context),
|
||||
&estimate,
|
||||
)
|
||||
.estimate_authorization_cost_upper_bound(&pricing, &estimate)
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
@@ -860,10 +895,10 @@ mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
available_balance_capacity_usd, execution_plan_balance_capacity_rejection,
|
||||
execution_plan_cost_upper_bound_cache_key, max_output_tokens_from_request,
|
||||
openai_request_input_is_self_contained, output_choice_count_upper_bound,
|
||||
request_model_local_rejection, GatewayLocalAuthRejection,
|
||||
available_balance_capacity_usd, estimate_execution_plan_cost_upper_bound_usd,
|
||||
execution_plan_balance_capacity_rejection, execution_plan_cost_upper_bound_cache_key,
|
||||
max_output_tokens_from_request, openai_request_input_is_self_contained,
|
||||
output_choice_count_upper_bound, request_model_local_rejection, GatewayLocalAuthRejection,
|
||||
};
|
||||
use crate::control::{GatewayControlAuthContext, GatewayControlDecision};
|
||||
use crate::data::GatewayDataState;
|
||||
@@ -2098,6 +2133,166 @@ mod tests {
|
||||
assert_eq!(estimate, 6.5);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn charge_estimate_and_capacity_use_request_multiplier_without_key_cost_or_cache_leaks() {
|
||||
let context = billing_context_with_pricing(
|
||||
Some(json!({"tiers": [{
|
||||
"up_to": null,
|
||||
"input_price_per_1m": 0.0,
|
||||
"output_price_per_1m": 10.0
|
||||
}]})),
|
||||
None,
|
||||
Some(json!({"openai:chat": 2.0})),
|
||||
None,
|
||||
);
|
||||
let mut state = state_with_quota_and_wallet(quota_availability(15.0, false), context);
|
||||
Arc::make_mut(&mut state.frontdoor_runtime_guards).auth_capacity_cache_ttl =
|
||||
Duration::from_secs(60);
|
||||
let decision = decision_with_allowed_models(vec!["gpt-5".to_string()]);
|
||||
let plan = execution_plan(
|
||||
json!({"model": "gpt-5", "messages": [], "max_tokens": 1_000_000}),
|
||||
"openai:chat",
|
||||
);
|
||||
let legacy = billing_report_context();
|
||||
let mut discounted = legacy.clone();
|
||||
discounted["billing_multiplier_snapshot"] = json!({
|
||||
"version": 1,
|
||||
"factors": {"routing_group": 2.0, "promotion": 0.25},
|
||||
"multiplier": 0.5
|
||||
});
|
||||
let mut marked_up = legacy.clone();
|
||||
marked_up["billing_multiplier_snapshot"] = json!({
|
||||
"version": 1,
|
||||
"factors": {"routing_group": 3.0},
|
||||
"multiplier": 3.0
|
||||
});
|
||||
let mut legacy_group_snapshot = legacy.clone();
|
||||
legacy_group_snapshot["routing_group_billing_multiplier"] = json!(1.0);
|
||||
|
||||
// Reuse the same cache for legacy Key cost, independent request
|
||||
// multipliers, and the old group-only snapshot representation.
|
||||
for (report_context, expected) in [
|
||||
(&legacy, 20.0),
|
||||
(&discounted, 5.0),
|
||||
(&marked_up, 30.0),
|
||||
(&legacy_group_snapshot, 10.0),
|
||||
(&discounted, 5.0),
|
||||
(&legacy, 20.0),
|
||||
] {
|
||||
assert_eq!(
|
||||
estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(report_context))
|
||||
.await
|
||||
.expect("charge estimate should resolve"),
|
||||
Some(expected)
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
execution_plan_balance_capacity_rejection(&state, &decision, &plan, Some(&discounted))
|
||||
.await
|
||||
.expect("discounted request capacity should resolve"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
execution_plan_balance_capacity_rejection(&state, &decision, &plan, Some(&marked_up))
|
||||
.await
|
||||
.expect("marked-up request capacity should resolve"),
|
||||
Some(GatewayLocalAuthRejection::BalanceDenied {
|
||||
remaining: Some(15.0)
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn charge_estimate_uses_base_price_when_provider_is_free_tier() {
|
||||
let context = billing_context_with_pricing(
|
||||
Some(json!({"tiers": [{
|
||||
"up_to": null,
|
||||
"input_price_per_1m": 0.0,
|
||||
"output_price_per_1m": 10.0
|
||||
}]})),
|
||||
None,
|
||||
Some(json!({"openai:chat": 0.0})),
|
||||
Some("free_tier"),
|
||||
);
|
||||
let state = state_with_quota_and_wallet(quota_availability(15.0, false), context);
|
||||
let plan = execution_plan(
|
||||
json!({"model": "gpt-5", "messages": [], "max_tokens": 1_000_000}),
|
||||
"openai:chat",
|
||||
);
|
||||
let mut report_context = billing_report_context();
|
||||
assert_eq!(
|
||||
estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context))
|
||||
.await
|
||||
.expect("legacy free-tier estimate should resolve"),
|
||||
Some(0.0)
|
||||
);
|
||||
report_context["routing_group_billing_multiplier"] = json!(0.5);
|
||||
assert_eq!(
|
||||
estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context))
|
||||
.await
|
||||
.expect("charge estimate should use the model base price"),
|
||||
Some(5.0)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn zero_charge_multiplier_bounds_unknown_cost_but_still_rejects_invalid_pricing() {
|
||||
let context = billing_context_with_pricing(
|
||||
Some(json!({"tiers": [{
|
||||
"up_to": null,
|
||||
"input_price_per_1m": 0.0,
|
||||
"output_price_per_1m": 10.0
|
||||
}]})),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
let state = state_with_quota_and_wallet(quota_availability(0.0, false), context);
|
||||
let plan = execution_plan(json!({"model": "gpt-5", "messages": []}), "openai:chat");
|
||||
let mut report_context = billing_report_context();
|
||||
assert_eq!(
|
||||
estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context))
|
||||
.await
|
||||
.expect("an unspecified output limit has no finite estimate"),
|
||||
None
|
||||
);
|
||||
report_context["billing_multiplier_snapshot"] = json!({
|
||||
"version": 1,
|
||||
"factors": {"routing_group": 0.0},
|
||||
"multiplier": 0.0
|
||||
});
|
||||
assert_eq!(
|
||||
estimate_execution_plan_cost_upper_bound_usd(&state, &plan, Some(&report_context))
|
||||
.await
|
||||
.expect("zero multiplier should bound the charge"),
|
||||
Some(0.0)
|
||||
);
|
||||
let invalid_context = billing_context_with_pricing(
|
||||
Some(json!({
|
||||
"tiers": [{"up_to": null, "input_price_per_1m": 1.0}],
|
||||
"processing_tiers": {
|
||||
"priority": {"tiers": [{}], "price_multiplier": -1.0}
|
||||
}
|
||||
})),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
let invalid_state =
|
||||
state_with_quota_and_wallet(quota_availability(0.0, false), invalid_context);
|
||||
let invalid_plan = execution_plan(
|
||||
json!({"model": "gpt-5", "messages": [], "service_tier": "priority"}),
|
||||
"openai:chat",
|
||||
);
|
||||
assert!(estimate_execution_plan_cost_upper_bound_usd(
|
||||
&invalid_state,
|
||||
&invalid_plan,
|
||||
Some(&report_context)
|
||||
)
|
||||
.await
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn daily_quota_estimate_treats_free_tier_as_zero_cost() {
|
||||
let context = billing_context_with_pricing(
|
||||
|
||||
@@ -637,6 +637,7 @@ pub(super) fn classify_public_support_route(
|
||||
| "/api/users/me/usage/interval-timeline"
|
||||
| "/api/users/me/usage/heatmap"
|
||||
| "/api/users/me/providers"
|
||||
| "/api/users/me/routing-groups"
|
||||
| "/api/users/me/available-models"
|
||||
| "/api/users/me/client-config"
|
||||
| "/api/users/me/endpoint-status"
|
||||
@@ -654,6 +655,7 @@ pub(super) fn classify_public_support_route(
|
||||
"/api/users/me/usage/interval-timeline" => "usage_interval_timeline",
|
||||
"/api/users/me/usage/heatmap" => "usage_heatmap",
|
||||
"/api/users/me/providers" => "providers",
|
||||
"/api/users/me/routing-groups" => "routing_groups",
|
||||
"/api/users/me/available-models" => "available_models",
|
||||
"/api/users/me/client-config" => "client_config",
|
||||
"/api/users/me/endpoint-status" => "endpoint_status",
|
||||
|
||||
@@ -461,6 +461,11 @@ fn classifies_users_me_routes_as_public_support_route() {
|
||||
"/api/users/me/available-models",
|
||||
"available_models",
|
||||
),
|
||||
(
|
||||
http::Method::GET,
|
||||
"/api/users/me/routing-groups",
|
||||
"routing_groups",
|
||||
),
|
||||
(
|
||||
http::Method::GET,
|
||||
"/api/users/me/vscodex/devices",
|
||||
|
||||
@@ -5524,6 +5524,8 @@ mod tests {
|
||||
key_ids: [&str; N],
|
||||
) -> ResolvedRoutingPolicy {
|
||||
ResolvedRoutingPolicy {
|
||||
billing_multiplier: 1.0,
|
||||
group_name: None,
|
||||
group_id: Some("routing-group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "test".to_string(),
|
||||
|
||||
@@ -7733,6 +7733,7 @@ impl<'a> AdminAppState<'a> {
|
||||
feature_settings: key
|
||||
.contains_key("feature_settings")
|
||||
.then(|| feature_settings.clone()),
|
||||
routing_group_selection: None,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
|
||||
@@ -411,7 +411,7 @@ async fn dry_run_routing_group(
|
||||
let headers_json = payload.headers.unwrap_or_else(|| json!({}));
|
||||
let mut header_map = header_map_from_value(&headers_json)?;
|
||||
let mut body = payload.body.unwrap_or_else(|| json!({}));
|
||||
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
let mut policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
group_id: Some(group.id.as_str()),
|
||||
group_version: Some(group.version),
|
||||
group_config_json: &group.config_json,
|
||||
@@ -425,6 +425,7 @@ async fn dry_run_routing_group(
|
||||
body: &body,
|
||||
phase: payload.phase.unwrap_or(RoutingRulePhase::ClientRequest),
|
||||
})?;
|
||||
policy.group_name = Some(group.name.clone());
|
||||
let patch_summary = patch_summary(&policy.mutation_plan);
|
||||
apply_routing_mutation_plan(&mut body, &mut header_map, &policy.mutation_plan)?;
|
||||
let mut trace = build_routing_trace_seed(&policy, api_format);
|
||||
|
||||
@@ -125,6 +125,7 @@ pub(crate) async fn build_admin_update_user_api_key_response(
|
||||
concurrent_limit_present,
|
||||
ip_rules,
|
||||
feature_settings,
|
||||
routing_group_selection: None,
|
||||
})
|
||||
.await?
|
||||
else {
|
||||
|
||||
@@ -32,6 +32,9 @@ use user_me_usage::*;
|
||||
#[path = "user_me_catalog.rs"]
|
||||
mod user_me_catalog;
|
||||
use user_me_catalog::*;
|
||||
#[path = "user_me_routing_groups.rs"]
|
||||
mod user_me_routing_groups;
|
||||
use user_me_routing_groups::*;
|
||||
#[path = "user_me_preferences.rs"]
|
||||
mod user_me_preferences;
|
||||
use user_me_preferences::*;
|
||||
|
||||
@@ -0,0 +1,351 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_data_contracts::repository::routing_profiles::RoutingGroupLookupKey;
|
||||
use axum::{body::Body, http::StatusCode, response::Response};
|
||||
use serde::Deserialize;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::{build_auth_error_response, normalize_feature_settings, AppState};
|
||||
use crate::routing::selection::routing_group_is_user_visible;
|
||||
|
||||
const ROUTING_GROUP_ID: &str = "routing_group_id";
|
||||
|
||||
pub(super) fn deserialize_routing_group_patch<'de, D>(
|
||||
deserializer: D,
|
||||
) -> Result<Option<Option<String>>, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
Option::<String>::deserialize(deserializer).map(Some)
|
||||
}
|
||||
|
||||
pub(super) fn api_key_routing_group_id(settings: Option<&Value>) -> Option<&str> {
|
||||
settings
|
||||
.and_then(|value| value.get(ROUTING_GROUP_ID))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub(super) async fn validate_routing_group_patch(
|
||||
state: &AppState,
|
||||
current: Option<&Value>,
|
||||
requested: Option<Option<String>>,
|
||||
) -> Result<Option<Option<String>>, Response<Body>> {
|
||||
let Some(Some(requested)) = requested else {
|
||||
return Ok(requested);
|
||||
};
|
||||
let id = requested.trim();
|
||||
if id.is_empty() || id.len() > 128 {
|
||||
return Err(build_auth_error_response(
|
||||
StatusCode::BAD_REQUEST,
|
||||
"routing_group_id 必须是有效的策略分组 ID;跟随默认请传 null",
|
||||
false,
|
||||
));
|
||||
}
|
||||
// A group can become hidden or disabled after selection. An unrelated edit
|
||||
// (including a form resubmitting its unchanged selection) must remain valid.
|
||||
if api_key_routing_group_id(current) == Some(id) {
|
||||
return Ok(Some(Some(id.to_string())));
|
||||
}
|
||||
if !state.has_routing_group_data_reader() {
|
||||
return Err(build_auth_error_response(
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"策略分组目录暂不可用",
|
||||
false,
|
||||
));
|
||||
}
|
||||
let group = state
|
||||
.find_routing_group(RoutingGroupLookupKey::Id(id))
|
||||
.await
|
||||
.map_err(|error| {
|
||||
build_auth_error_response(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("user API key routing group lookup failed: {error:?}"),
|
||||
false,
|
||||
)
|
||||
})?;
|
||||
if !group
|
||||
.as_ref()
|
||||
.is_some_and(|group| group.enabled && routing_group_is_user_visible(group))
|
||||
{
|
||||
return Err(build_auth_error_response(
|
||||
StatusCode::BAD_REQUEST,
|
||||
"所选策略分组不存在或当前不可选",
|
||||
false,
|
||||
));
|
||||
}
|
||||
Ok(Some(Some(id.to_string())))
|
||||
}
|
||||
|
||||
/// Compose the initial create record. Existing-key updates must use the atomic
|
||||
/// repository routing selection patch instead of merging a pre-read snapshot.
|
||||
pub(super) fn merge_api_key_feature_settings(
|
||||
current: Option<&Value>,
|
||||
incoming: Option<Option<Value>>,
|
||||
routing_group_patch: Option<Option<String>>,
|
||||
) -> Result<Option<Option<Value>>, String> {
|
||||
if incoming.is_none() && routing_group_patch.is_none() {
|
||||
return Ok(None);
|
||||
}
|
||||
let mut settings = normalize_feature_settings(incoming.unwrap_or_else(|| current.cloned()))?
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default();
|
||||
settings.remove(ROUTING_GROUP_ID);
|
||||
settings.remove("routing_group_name");
|
||||
let group_id = routing_group_patch
|
||||
.unwrap_or_else(|| api_key_routing_group_id(current).map(ToOwned::to_owned));
|
||||
if let Some(group_id) = group_id {
|
||||
settings.insert(ROUTING_GROUP_ID.to_string(), Value::String(group_id));
|
||||
}
|
||||
Ok(Some(
|
||||
(!settings.is_empty()).then_some(Value::Object(settings)),
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) fn normalize_api_key_feature_settings_patch(
|
||||
value: Option<Option<Value>>,
|
||||
) -> Result<Option<Option<Value>>, String> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(Value::Object(mut settings)) = normalize_feature_settings(value)? else {
|
||||
return Ok(Some(None));
|
||||
};
|
||||
settings.remove(ROUTING_GROUP_ID);
|
||||
settings.remove("routing_group_name");
|
||||
Ok(Some(
|
||||
(!settings.is_empty()).then_some(Value::Object(settings)),
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) async fn routing_group_names(
|
||||
state: &AppState,
|
||||
needed: bool,
|
||||
) -> BTreeMap<String, String> {
|
||||
if !needed || !state.has_routing_group_data_reader() {
|
||||
return BTreeMap::new();
|
||||
}
|
||||
match state.list_routing_groups().await {
|
||||
Ok(groups) => groups
|
||||
.into_iter()
|
||||
.map(|group| (group.id, group.name))
|
||||
.collect(),
|
||||
Err(error) => {
|
||||
tracing::warn!(?error, "API key routing group names unavailable");
|
||||
BTreeMap::new()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn routing_group_payload_fields(
|
||||
settings: Option<&Value>,
|
||||
names: &BTreeMap<String, String>,
|
||||
) -> Map<String, Value> {
|
||||
let id = api_key_routing_group_id(settings);
|
||||
Map::from_iter([
|
||||
(ROUTING_GROUP_ID.to_string(), serde_json::json!(id)),
|
||||
(
|
||||
"routing_group_name".to_string(),
|
||||
serde_json::json!(id.and_then(|id| names.get(id))),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
|
||||
use aether_data_contracts::repository::routing_profiles::StoredRoutingGroup;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn group(id: &str, enabled: bool, visible: Value) -> StoredRoutingGroup {
|
||||
StoredRoutingGroup {
|
||||
id: id.to_string(),
|
||||
name: format!("{id} current name"),
|
||||
description: None,
|
||||
enabled,
|
||||
is_system_default: false,
|
||||
sort_order: 0,
|
||||
config_json: json!({"user_visible": visible}),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn state() -> AppState {
|
||||
let repository = Arc::new(InMemoryRoutingGroupRepository::seed(
|
||||
[
|
||||
group("public", true, json!(true)),
|
||||
group("hidden", true, json!(false)),
|
||||
group("disabled", false, json!(true)),
|
||||
group("legacy", true, Value::Null),
|
||||
group("malformed", true, json!("true")),
|
||||
],
|
||||
[],
|
||||
[],
|
||||
));
|
||||
AppState::new().unwrap().with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::disabled()
|
||||
.with_routing_group_repository_for_tests(repository),
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn only_enabled_user_visible_groups_can_be_newly_selected() {
|
||||
let state = state();
|
||||
for id in [
|
||||
"",
|
||||
" ",
|
||||
"missing",
|
||||
"hidden",
|
||||
"disabled",
|
||||
"legacy",
|
||||
"malformed",
|
||||
] {
|
||||
assert_eq!(
|
||||
validate_routing_group_patch(&state, None, Some(Some(id.to_string())))
|
||||
.await
|
||||
.unwrap_err()
|
||||
.status(),
|
||||
StatusCode::BAD_REQUEST,
|
||||
"id={id}"
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
validate_routing_group_patch(&state, None, Some(Some(" public ".to_string())))
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(Some("public".to_string()))
|
||||
);
|
||||
let unavailable = AppState::new()
|
||||
.unwrap()
|
||||
.with_data_state_for_tests(crate::data::GatewayDataState::disabled());
|
||||
assert_eq!(
|
||||
validate_routing_group_patch(&unavailable, None, Some(Some("public".into())))
|
||||
.await
|
||||
.unwrap_err()
|
||||
.status(),
|
||||
StatusCode::SERVICE_UNAVAILABLE
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unchanged_or_cleared_choices_do_not_require_current_visibility_or_catalog_access() {
|
||||
let state = AppState::new()
|
||||
.unwrap()
|
||||
.with_data_state_for_tests(crate::data::GatewayDataState::disabled());
|
||||
let current = json!({"routing_group_id": "hidden"});
|
||||
for patch in [None, Some(None), Some(Some("hidden".to_string()))] {
|
||||
assert_eq!(
|
||||
validate_routing_group_patch(&state, Some(¤t), patch.clone())
|
||||
.await
|
||||
.unwrap(),
|
||||
patch
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn feature_patch_does_not_carry_a_pre_read_routing_selection() {
|
||||
assert_eq!(
|
||||
normalize_api_key_feature_settings_patch(None).unwrap(),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_api_key_feature_settings_patch(Some(None)).unwrap(),
|
||||
Some(None)
|
||||
);
|
||||
let patch = normalize_api_key_feature_settings_patch(Some(Some(json!({
|
||||
"routing_group_id": "stale-or-forged",
|
||||
"routing_group_name": "stale name",
|
||||
"chat_pii_redaction": {"enabled": false},
|
||||
}))))
|
||||
.unwrap()
|
||||
.flatten()
|
||||
.unwrap();
|
||||
assert!(patch.get("routing_group_id").is_none());
|
||||
assert!(patch.get("routing_group_name").is_none());
|
||||
assert_eq!(patch["chat_pii_redaction"]["enabled"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn feature_updates_cannot_inject_replace_or_clear_a_routing_choice() {
|
||||
let current =
|
||||
json!({"routing_group_id": "saved", "chat_pii_redaction": {"enabled": false}});
|
||||
let injection = json!({"routing_group_id": "hidden", "routing_group_name": "forged", "chat_pii_redaction": {"enabled": true}});
|
||||
let created = merge_api_key_feature_settings(None, Some(Some(injection.clone())), None)
|
||||
.unwrap()
|
||||
.flatten()
|
||||
.unwrap();
|
||||
assert!(created.get("routing_group_id").is_none());
|
||||
assert!(created.get("routing_group_name").is_none());
|
||||
for feature_patch in [
|
||||
Some(injection),
|
||||
Some(json!({"routing_group_id": null})),
|
||||
None,
|
||||
] {
|
||||
let updated = merge_api_key_feature_settings(Some(¤t), Some(feature_patch), None)
|
||||
.unwrap()
|
||||
.flatten()
|
||||
.unwrap();
|
||||
assert_eq!(updated["routing_group_id"], "saved");
|
||||
assert!(updated.get("routing_group_name").is_none());
|
||||
}
|
||||
let untouched = merge_api_key_feature_settings(Some(¤t), None, None).unwrap();
|
||||
assert_eq!(
|
||||
untouched, None,
|
||||
"name/rate/IP-only updates must leave settings untouched"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validated_top_level_selection_and_clear_preserve_other_feature_settings() {
|
||||
let current = json!({"routing_group_id": "saved", "chat_pii_redaction": {"enabled": false}, "another_setting": 7});
|
||||
let selected =
|
||||
merge_api_key_feature_settings(Some(¤t), None, Some(Some("public".to_string())))
|
||||
.unwrap()
|
||||
.flatten()
|
||||
.unwrap();
|
||||
assert_eq!(selected["routing_group_id"], "public");
|
||||
assert_eq!(selected["another_setting"], 7);
|
||||
assert_eq!(selected["chat_pii_redaction"]["enabled"], false);
|
||||
let cleared = merge_api_key_feature_settings(Some(&selected), None, Some(None))
|
||||
.unwrap()
|
||||
.flatten()
|
||||
.unwrap();
|
||||
assert!(cleared.get("routing_group_id").is_none());
|
||||
assert_eq!(cleared["another_setting"], 7);
|
||||
let replacement = merge_api_key_feature_settings(
|
||||
None,
|
||||
Some(Some(json!({"routing_group_id": "hidden"}))),
|
||||
Some(Some("public".to_string())),
|
||||
)
|
||||
.unwrap()
|
||||
.flatten()
|
||||
.unwrap();
|
||||
assert_eq!(replacement, json!({"routing_group_id": "public"}));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn names_are_resolved_from_the_catalog_without_persisting_a_name_snapshot() {
|
||||
let state = state();
|
||||
let names = routing_group_names(&state, true).await;
|
||||
let current = json!({"routing_group_id": "hidden", "routing_group_name": "stale"});
|
||||
let fields = routing_group_payload_fields(Some(¤t), &names);
|
||||
assert_eq!(fields["routing_group_id"], "hidden");
|
||||
assert_eq!(fields["routing_group_name"], "hidden current name");
|
||||
let missing =
|
||||
routing_group_payload_fields(Some(&json!({"routing_group_id": "deleted"})), &names);
|
||||
assert_eq!(missing["routing_group_id"], "deleted");
|
||||
assert_eq!(missing["routing_group_name"], Value::Null);
|
||||
let default = routing_group_payload_fields(None, &names);
|
||||
assert_eq!(default["routing_group_id"], Value::Null);
|
||||
assert_eq!(default["routing_group_name"], Value::Null);
|
||||
}
|
||||
}
|
||||
@@ -26,6 +26,14 @@ use super::{
|
||||
|
||||
const USERS_ME_API_KEY_WRITE_UNAVAILABLE_DETAIL: &str = "用户 API 密钥写入暂不可用";
|
||||
|
||||
#[path = "user_me_api_key_routing.rs"]
|
||||
mod routing_selection;
|
||||
use routing_selection::{
|
||||
api_key_routing_group_id, deserialize_routing_group_patch, merge_api_key_feature_settings,
|
||||
normalize_api_key_feature_settings_patch, routing_group_names, routing_group_payload_fields,
|
||||
validate_routing_group_patch,
|
||||
};
|
||||
|
||||
fn users_me_api_key_secret_response(mut response: Response<Body>) -> Response<Body> {
|
||||
response.headers_mut().insert(
|
||||
http::header::CACHE_CONTROL,
|
||||
@@ -43,6 +51,8 @@ struct UsersMeCreateApiKeyRequest {
|
||||
concurrent_limit: Option<i32>,
|
||||
#[serde(default)]
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
routing_group_id: Option<String>,
|
||||
#[serde(default, alias = "allowed_ips")]
|
||||
ip_rules: Option<Vec<String>>,
|
||||
}
|
||||
@@ -57,6 +67,8 @@ struct UsersMeUpdateApiKeyRequest {
|
||||
concurrent_limit: Option<i32>,
|
||||
#[serde(default, deserialize_with = "deserialize_optional_json_patch")]
|
||||
feature_settings: Option<Option<serde_json::Value>>,
|
||||
#[serde(default, deserialize_with = "deserialize_routing_group_patch")]
|
||||
routing_group_id: Option<Option<String>>,
|
||||
#[serde(
|
||||
default,
|
||||
alias = "allowed_ips",
|
||||
@@ -170,8 +182,9 @@ fn build_users_me_api_key_list_payload(
|
||||
state: &AppState,
|
||||
record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord,
|
||||
is_locked: bool,
|
||||
group_names: &BTreeMap<String, String>,
|
||||
) -> serde_json::Value {
|
||||
json!({
|
||||
let mut payload = json!({
|
||||
"id": record.api_key_id,
|
||||
"name": record.name,
|
||||
"key_display": users_me_masked_api_key_display(state, record),
|
||||
@@ -187,15 +200,24 @@ fn build_users_me_api_key_list_payload(
|
||||
"ip_rules": record.ip_rules,
|
||||
"force_capabilities": record.force_capabilities,
|
||||
"feature_settings": record.feature_settings,
|
||||
})
|
||||
});
|
||||
payload
|
||||
.as_object_mut()
|
||||
.expect("API key payload is an object")
|
||||
.extend(routing_group_payload_fields(
|
||||
record.feature_settings.as_ref(),
|
||||
group_names,
|
||||
));
|
||||
payload
|
||||
}
|
||||
|
||||
fn build_users_me_api_key_detail_payload(
|
||||
state: &AppState,
|
||||
record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord,
|
||||
is_locked: bool,
|
||||
group_names: &BTreeMap<String, String>,
|
||||
) -> serde_json::Value {
|
||||
json!({
|
||||
let mut payload = json!({
|
||||
"id": record.api_key_id,
|
||||
"name": record.name,
|
||||
"key_display": users_me_masked_api_key_display(state, record),
|
||||
@@ -210,7 +232,15 @@ fn build_users_me_api_key_detail_payload(
|
||||
"last_used_at": format_users_me_optional_unix_secs_iso8601(record.last_used_at_unix_secs),
|
||||
"expires_at": format_users_me_optional_unix_secs_iso8601(record.expires_at_unix_secs),
|
||||
"created_at": format_users_me_optional_unix_secs_iso8601(record.created_at_unix_secs),
|
||||
})
|
||||
});
|
||||
payload
|
||||
.as_object_mut()
|
||||
.expect("API key payload is an object")
|
||||
.extend(routing_group_payload_fields(
|
||||
record.feature_settings.as_ref(),
|
||||
group_names,
|
||||
));
|
||||
payload
|
||||
}
|
||||
|
||||
fn normalize_users_me_required_api_key_name(value: &str) -> Result<String, String> {
|
||||
@@ -336,6 +366,13 @@ pub(super) async fn handle_users_me_api_keys_get(
|
||||
};
|
||||
records.retain(|record| !record.is_standalone);
|
||||
records.sort_by(|left, right| left.api_key_id.cmp(&right.api_key_id));
|
||||
let group_names = routing_group_names(
|
||||
state,
|
||||
records
|
||||
.iter()
|
||||
.any(|record| api_key_routing_group_id(record.feature_settings.as_ref()).is_some()),
|
||||
)
|
||||
.await;
|
||||
|
||||
let snapshot_ids = records
|
||||
.iter()
|
||||
@@ -367,7 +404,7 @@ pub(super) async fn handle_users_me_api_keys_get(
|
||||
.get(&record.api_key_id)
|
||||
.map(|snapshot| snapshot.api_key_is_locked)
|
||||
.unwrap_or(false);
|
||||
build_users_me_api_key_list_payload(state, record, is_locked)
|
||||
build_users_me_api_key_list_payload(state, record, is_locked, &group_names)
|
||||
})
|
||||
.collect::<Vec<_>>(),
|
||||
)
|
||||
@@ -461,8 +498,16 @@ pub(super) async fn handle_users_me_api_key_detail_get(
|
||||
}
|
||||
};
|
||||
|
||||
let group_names = routing_group_names(
|
||||
state,
|
||||
api_key_routing_group_id(record.feature_settings.as_ref()).is_some(),
|
||||
)
|
||||
.await;
|
||||
Json(build_users_me_api_key_detail_payload(
|
||||
state, &record, is_locked,
|
||||
state,
|
||||
&record,
|
||||
is_locked,
|
||||
&group_names,
|
||||
))
|
||||
.into_response()
|
||||
}
|
||||
@@ -567,8 +612,17 @@ pub(super) async fn handle_users_me_api_key_create(
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
}
|
||||
};
|
||||
let feature_settings = match normalize_feature_settings(payload.feature_settings) {
|
||||
Ok(value) => value,
|
||||
let routing_group_patch =
|
||||
match validate_routing_group_patch(state, None, Some(payload.routing_group_id)).await {
|
||||
Ok(value) => value,
|
||||
Err(response) => return response,
|
||||
};
|
||||
let feature_settings = match merge_api_key_feature_settings(
|
||||
None,
|
||||
Some(payload.feature_settings),
|
||||
routing_group_patch,
|
||||
) {
|
||||
Ok(value) => value.flatten(),
|
||||
Err(detail) => {
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
}
|
||||
@@ -631,26 +685,36 @@ pub(super) async fn handle_users_me_api_key_create(
|
||||
return build_users_me_api_key_writer_unavailable_response();
|
||||
};
|
||||
|
||||
users_me_api_key_secret_response(
|
||||
Json(json!({
|
||||
"id": created.api_key_id,
|
||||
"name": created.name,
|
||||
"key": plaintext_key,
|
||||
"key_display": users_me_masked_api_key_display(state, &created),
|
||||
"is_active": created.is_active,
|
||||
"is_locked": false,
|
||||
"rate_limit": created.rate_limit,
|
||||
"concurrent_limit": created.concurrent_limit,
|
||||
"ip_rules": created.ip_rules,
|
||||
"feature_settings": created.feature_settings,
|
||||
"last_used_at": format_users_me_optional_unix_secs_iso8601(created.last_used_at_unix_secs),
|
||||
"created_at": format_users_me_optional_unix_secs_iso8601(created.created_at_unix_secs),
|
||||
"total_requests": created.total_requests,
|
||||
"total_cost_usd": created.total_cost_usd,
|
||||
"message": "API密钥创建成功",
|
||||
}))
|
||||
.into_response(),
|
||||
let group_names = routing_group_names(
|
||||
state,
|
||||
api_key_routing_group_id(created.feature_settings.as_ref()).is_some(),
|
||||
)
|
||||
.await;
|
||||
let mut payload = json!({
|
||||
"id": created.api_key_id,
|
||||
"name": created.name,
|
||||
"key": plaintext_key,
|
||||
"key_display": users_me_masked_api_key_display(state, &created),
|
||||
"is_active": created.is_active,
|
||||
"is_locked": false,
|
||||
"rate_limit": created.rate_limit,
|
||||
"concurrent_limit": created.concurrent_limit,
|
||||
"ip_rules": created.ip_rules,
|
||||
"feature_settings": created.feature_settings,
|
||||
"last_used_at": format_users_me_optional_unix_secs_iso8601(created.last_used_at_unix_secs),
|
||||
"created_at": format_users_me_optional_unix_secs_iso8601(created.created_at_unix_secs),
|
||||
"total_requests": created.total_requests,
|
||||
"total_cost_usd": created.total_cost_usd,
|
||||
"message": "API密钥创建成功",
|
||||
});
|
||||
payload
|
||||
.as_object_mut()
|
||||
.expect("API key payload is an object")
|
||||
.extend(routing_group_payload_fields(
|
||||
created.feature_settings.as_ref(),
|
||||
&group_names,
|
||||
));
|
||||
users_me_api_key_secret_response(Json(payload).into_response())
|
||||
}
|
||||
|
||||
pub(super) async fn handle_users_me_api_key_update(
|
||||
@@ -715,14 +779,39 @@ pub(super) async fn handle_users_me_api_key_update(
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
}
|
||||
};
|
||||
let feature_settings = match payload.feature_settings {
|
||||
Some(value) => match normalize_feature_settings(value) {
|
||||
Ok(value) => Some(value),
|
||||
Err(detail) => {
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
let current_features = if matches!(payload.routing_group_id.as_ref(), Some(Some(_))) {
|
||||
match state
|
||||
.read_auth_api_key_feature_settings(&auth.user.id, &snapshot.api_key_id, false)
|
||||
.await
|
||||
{
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("user API key feature settings lookup failed: {error:?}"),
|
||||
false,
|
||||
)
|
||||
}
|
||||
},
|
||||
None => None,
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let routing_group_patch = match validate_routing_group_patch(
|
||||
state,
|
||||
current_features.as_ref(),
|
||||
payload.routing_group_id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(value) => value,
|
||||
Err(response) => return response,
|
||||
};
|
||||
let feature_settings = match normalize_api_key_feature_settings_patch(payload.feature_settings)
|
||||
{
|
||||
Ok(value) => value,
|
||||
Err(detail) => {
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false)
|
||||
}
|
||||
};
|
||||
let ip_rules = match payload.ip_rules {
|
||||
Some(value) => match normalize_users_me_ip_rules(value) {
|
||||
@@ -752,6 +841,11 @@ pub(super) async fn handle_users_me_api_key_update(
|
||||
concurrent_limit_present,
|
||||
ip_rules,
|
||||
feature_settings,
|
||||
routing_group_selection: Some(
|
||||
aether_data::repository::auth::UpdateApiKeyRoutingGroupSelection {
|
||||
group_id: routing_group_patch,
|
||||
},
|
||||
),
|
||||
},
|
||||
)
|
||||
.await
|
||||
@@ -768,8 +862,17 @@ pub(super) async fn handle_users_me_api_key_update(
|
||||
return build_users_me_api_key_mutation_conflict_response();
|
||||
};
|
||||
|
||||
let mut payload =
|
||||
build_users_me_api_key_detail_payload(state, &updated, snapshot.api_key_is_locked);
|
||||
let group_names = routing_group_names(
|
||||
state,
|
||||
api_key_routing_group_id(updated.feature_settings.as_ref()).is_some(),
|
||||
)
|
||||
.await;
|
||||
let mut payload = build_users_me_api_key_detail_payload(
|
||||
state,
|
||||
&updated,
|
||||
snapshot.api_key_is_locked,
|
||||
&group_names,
|
||||
);
|
||||
payload["message"] = json!("API密钥已更新");
|
||||
Json(payload).into_response()
|
||||
}
|
||||
@@ -1095,7 +1198,8 @@ mod tests {
|
||||
use axum::{response::IntoResponse, Json};
|
||||
|
||||
use super::{
|
||||
normalize_users_me_ip_rules, users_me_api_key_secret_response, UsersMeUpdateApiKeyRequest,
|
||||
normalize_users_me_ip_rules, users_me_api_key_secret_response, UsersMeCreateApiKeyRequest,
|
||||
UsersMeUpdateApiKeyRequest,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
@@ -1110,6 +1214,28 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routing_group_selection_patch_distinguishes_missing_clear_and_valid_string() {
|
||||
let unchanged: UsersMeUpdateApiKeyRequest =
|
||||
serde_json::from_value(json!({"name": "renamed"})).unwrap();
|
||||
assert_eq!(unchanged.routing_group_id, None);
|
||||
let cleared: UsersMeUpdateApiKeyRequest =
|
||||
serde_json::from_value(json!({"routing_group_id": null})).unwrap();
|
||||
assert_eq!(cleared.routing_group_id, Some(None));
|
||||
let selected: UsersMeUpdateApiKeyRequest =
|
||||
serde_json::from_value(json!({"routing_group_id": "group-1"})).unwrap();
|
||||
assert_eq!(selected.routing_group_id, Some(Some("group-1".to_string())));
|
||||
let created: UsersMeCreateApiKeyRequest =
|
||||
serde_json::from_value(json!({"name": "created", "routing_group_id": "group-1"}))
|
||||
.unwrap();
|
||||
assert_eq!(created.routing_group_id.as_deref(), Some("group-1"));
|
||||
for invalid in [json!(true), json!(3), json!([]), json!({})] {
|
||||
let request = json!({"name": "invalid", "routing_group_id": invalid});
|
||||
assert!(serde_json::from_value::<UsersMeCreateApiKeyRequest>(request.clone()).is_err());
|
||||
assert!(serde_json::from_value::<UsersMeUpdateApiKeyRequest>(request).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_ip_rules_trims_ip_and_cidr_values() {
|
||||
let values = normalize_users_me_ip_rules(Some(vec![
|
||||
|
||||
@@ -17,12 +17,12 @@ use super::{
|
||||
handle_users_me_model_capabilities_get, handle_users_me_model_capabilities_put,
|
||||
handle_users_me_password_patch, handle_users_me_preferences_get,
|
||||
handle_users_me_preferences_put, handle_users_me_providers_get, handle_users_me_referral_get,
|
||||
handle_users_me_sessions_get, handle_users_me_update_session, handle_users_me_usage_active_get,
|
||||
handle_users_me_usage_get, handle_users_me_usage_heatmap_get,
|
||||
handle_users_me_usage_interval_timeline_get, handle_users_me_vscodex_request,
|
||||
users_me_api_key_capabilities_path_matches, users_me_api_key_detail_path_matches,
|
||||
users_me_api_key_install_sessions_path_matches, users_me_api_key_providers_path_matches,
|
||||
users_me_management_token_detail_path_matches,
|
||||
handle_users_me_routing_groups_get, handle_users_me_sessions_get,
|
||||
handle_users_me_update_session, handle_users_me_usage_active_get, handle_users_me_usage_get,
|
||||
handle_users_me_usage_heatmap_get, handle_users_me_usage_interval_timeline_get,
|
||||
handle_users_me_vscodex_request, users_me_api_key_capabilities_path_matches,
|
||||
users_me_api_key_detail_path_matches, users_me_api_key_install_sessions_path_matches,
|
||||
users_me_api_key_providers_path_matches, users_me_management_token_detail_path_matches,
|
||||
users_me_management_token_regenerate_path_matches,
|
||||
users_me_management_token_toggle_path_matches, users_me_management_tokens_root,
|
||||
users_me_session_detail_path_matches, AppState, GatewayPublicRequestContext,
|
||||
@@ -231,6 +231,11 @@ pub(crate) async fn maybe_build_local_users_me_response(
|
||||
Some("providers") if request_context.request_path == "/api/users/me/providers" => {
|
||||
Some(handle_users_me_providers_get(state, request_context, headers).await)
|
||||
}
|
||||
Some("routing_groups")
|
||||
if request_context.request_path == "/api/users/me/routing-groups" =>
|
||||
{
|
||||
Some(handle_users_me_routing_groups_get(state, request_context, headers).await)
|
||||
}
|
||||
Some("preferences") if request_context.request_path == "/api/users/me/preferences" => {
|
||||
Some(handle_users_me_preferences_get(state, request_context, headers).await)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
use aether_routing_core::RoutingGroupConfig;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_auth_error_response, resolve_authenticated_local_user, AppState,
|
||||
GatewayPublicRequestContext,
|
||||
};
|
||||
use crate::routing::selection::routing_group_is_user_visible;
|
||||
|
||||
pub(super) async fn handle_users_me_routing_groups_get(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
headers: &http::HeaderMap,
|
||||
) -> Response<Body> {
|
||||
if let Err(response) = resolve_authenticated_local_user(state, request_context, headers).await {
|
||||
return response;
|
||||
}
|
||||
if !state.has_routing_group_data_reader() {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"策略分组目录暂不可用",
|
||||
false,
|
||||
);
|
||||
}
|
||||
let groups = match state.list_routing_groups().await {
|
||||
Ok(groups) => groups,
|
||||
Err(error) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("user routing group lookup failed: {error:?}"),
|
||||
false,
|
||||
)
|
||||
}
|
||||
};
|
||||
let items = groups
|
||||
.into_iter()
|
||||
.filter(|group| group.enabled && routing_group_is_user_visible(group))
|
||||
.filter_map(|group| {
|
||||
let config = serde_json::from_value::<RoutingGroupConfig>(group.config_json).ok()?;
|
||||
if !config.billing_multiplier.is_finite() || config.billing_multiplier < 0.0 {
|
||||
return None;
|
||||
}
|
||||
Some(json!({
|
||||
"id": group.id,
|
||||
"name": group.name,
|
||||
"billing_multiplier": config.billing_multiplier,
|
||||
"is_default": group.is_system_default,
|
||||
}))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
Json(json!({"total": items.len(), "items": items})).into_response()
|
||||
}
|
||||
@@ -561,6 +561,8 @@ fn build_users_me_usage_record_payload(
|
||||
let cache_read_price_per_1m = item.settlement_cache_read_price_per_1m();
|
||||
let cache_creation_input_tokens = users_me_usage_cache_creation_tokens(item);
|
||||
let rate_multiplier = item.settlement_rate_multiplier();
|
||||
let billing_multiplier = item.billing_multiplier();
|
||||
let billing_cost = item.billing_cost().map(|cost| round_to(cost, 6));
|
||||
let client_is_stream = users_me_usage_client_is_stream(item);
|
||||
let upstream_is_stream = users_me_usage_upstream_is_stream(item);
|
||||
let mut payload = json!({
|
||||
@@ -576,6 +578,8 @@ fn build_users_me_usage_record_payload(
|
||||
"output_tokens": item.output_tokens,
|
||||
"total_tokens": item.total_tokens,
|
||||
"cost": round_to(item.total_cost_usd, 6),
|
||||
"billing_multiplier": billing_multiplier,
|
||||
"billing_cost": billing_cost,
|
||||
"response_time_ms": item.response_time_ms,
|
||||
"first_byte_time_ms": item.first_byte_time_ms,
|
||||
"is_stream": item.is_stream,
|
||||
@@ -614,6 +618,8 @@ fn build_users_me_usage_record_payload(
|
||||
),
|
||||
});
|
||||
payload["end_to_end_time_ms"] = json!(users_me_usage_metadata_u64(item, "end_to_end_time_ms"));
|
||||
payload["routing_group_id"] = json!(item.routing_group_id());
|
||||
payload["routing_group_name"] = json!(item.routing_group_name());
|
||||
payload["end_to_end_first_byte_time_ms"] = json!(users_me_usage_metadata_u64(
|
||||
item,
|
||||
"end_to_end_first_byte_time_ms"
|
||||
@@ -643,6 +649,8 @@ fn build_users_me_usage_record_payload(
|
||||
|
||||
fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_json::Value {
|
||||
let cache_creation_input_tokens = users_me_usage_cache_creation_tokens(item);
|
||||
let billing_multiplier = item.billing_multiplier();
|
||||
let billing_cost = item.billing_cost().map(|cost| round_to(cost, 6));
|
||||
let client_is_stream = users_me_usage_client_is_stream(item);
|
||||
let upstream_is_stream = users_me_usage_upstream_is_stream(item);
|
||||
let mut payload = json!({
|
||||
@@ -659,6 +667,8 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
|
||||
"cost": round_to(item.total_cost_usd, 6),
|
||||
"actual_cost": round_to(item.actual_total_cost_usd, 6),
|
||||
"rate_multiplier": item.settlement_rate_multiplier(),
|
||||
"billing_multiplier": billing_multiplier,
|
||||
"billing_cost": billing_cost,
|
||||
"response_time_ms": item.response_time_ms,
|
||||
"first_byte_time_ms": item.first_byte_time_ms,
|
||||
"updated_at": unix_secs_to_rfc3339(item.updated_at_unix_secs),
|
||||
@@ -685,6 +695,8 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
|
||||
"response_model": item.provider_response_model(),
|
||||
"has_fallback": item.has_fallback(),
|
||||
});
|
||||
payload["routing_group_id"] = json!(item.routing_group_id());
|
||||
payload["routing_group_name"] = json!(item.routing_group_name());
|
||||
payload["end_to_end_time_ms"] = json!(users_me_usage_metadata_u64(item, "end_to_end_time_ms"));
|
||||
payload["end_to_end_first_byte_time_ms"] = json!(users_me_usage_metadata_u64(
|
||||
item,
|
||||
@@ -1867,6 +1879,59 @@ mod tests {
|
||||
assert_eq!(payload["cache_creation_ephemeral_1h_input_tokens"], 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_usage_payloads_preserve_routing_group_snapshot_and_precise_display_cost() {
|
||||
for (metadata, multiplier, cost, group_name) in [
|
||||
(None, 1.0, json!(0.0), serde_json::Value::Null),
|
||||
(
|
||||
Some(json!({"routing_group_billing_multiplier": 0.0})),
|
||||
0.0,
|
||||
json!(0.0),
|
||||
serde_json::Value::Null,
|
||||
),
|
||||
(
|
||||
Some(json!({
|
||||
"routing_group_billing_multiplier": 2.5,
|
||||
"routing_group_id": "group-1",
|
||||
"routing_group_name": "请求时的分组",
|
||||
"rate_multiplier": 0.5
|
||||
})),
|
||||
2.5,
|
||||
json!(0.000004),
|
||||
json!("请求时的分组"),
|
||||
),
|
||||
(
|
||||
Some(json!({
|
||||
"billing_multiplier_snapshot": {"version": 1, "factors": {"routing_group": 2.5, "user_group": 2.0}, "multiplier": 5.0},
|
||||
"routing_group_billing_multiplier": 2.5,
|
||||
"routing_group_id": "group-1",
|
||||
"routing_group_name": "请求时的分组",
|
||||
"rate_multiplier": 0.5
|
||||
})),
|
||||
5.0,
|
||||
json!(0.000007),
|
||||
json!("请求时的分组"),
|
||||
),
|
||||
] {
|
||||
let item = StoredRequestUsageAudit {
|
||||
total_cost_usd: 0.00000149,
|
||||
actual_total_cost_usd: 0.0000002,
|
||||
request_metadata: metadata,
|
||||
..sample_usage("completed")
|
||||
};
|
||||
let record = build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false);
|
||||
let active = build_users_me_usage_active_payload(&item);
|
||||
for payload in [&record, &active] {
|
||||
assert_eq!(payload["billing_multiplier"], multiplier);
|
||||
assert_eq!(payload["billing_cost"], cost);
|
||||
assert_eq!(payload["routing_group_name"], group_name);
|
||||
assert_eq!(payload["cost"], 0.000001);
|
||||
}
|
||||
assert!(record.get("actual_cost").is_none());
|
||||
assert!(record.get("rate_multiplier").is_none());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_usage_payloads_expose_response_model_separately_from_mapping() {
|
||||
let item = StoredRequestUsageAudit {
|
||||
|
||||
@@ -84,10 +84,19 @@ pub(crate) fn resolve_gateway_static_default_routing_policy(
|
||||
let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let billing_multiplier = match input.group_config_json.get("billing_multiplier") {
|
||||
Some(value) => value
|
||||
.as_f64()
|
||||
.filter(|value| value.is_finite() && *value >= 0.0)
|
||||
.ok_or_else(invalid_routing_group_config)?,
|
||||
None => 1.0,
|
||||
};
|
||||
crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy.clone());
|
||||
|
||||
Ok(Some(ResolvedRoutingPolicy {
|
||||
billing_multiplier,
|
||||
group_id: input.group_id.map(str::to_string),
|
||||
group_name: None,
|
||||
group_version: input.group_version,
|
||||
selection_source: input.selection_source.to_string(),
|
||||
requested_model: input.requested_model.to_string(),
|
||||
@@ -216,6 +225,7 @@ mod tests {
|
||||
#[test]
|
||||
fn static_default_policy_matches_full_resolver_without_body_context() {
|
||||
let config = json!({
|
||||
"billing_multiplier": 2.5,
|
||||
"default_policy": {
|
||||
"priority_mode": "global_key",
|
||||
"scheduling_mode": "load_balance",
|
||||
@@ -262,6 +272,7 @@ mod tests {
|
||||
.expect("full policy should resolve");
|
||||
|
||||
assert_eq!(static_policy, full_policy);
|
||||
assert_eq!(static_policy.billing_multiplier, 2.5);
|
||||
assert_eq!(static_policy.execution_policy.max_transfer_count, 3);
|
||||
assert_eq!(
|
||||
static_policy.execution_policy.max_transfer_timeout_seconds,
|
||||
@@ -288,6 +299,50 @@ mod tests {
|
||||
assert!(static_policy.matched_rules.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn static_default_billing_multiplier_defaults_to_one_and_validates_input() {
|
||||
for config in [
|
||||
json!({}),
|
||||
json!({"billing_multiplier": 0.0}),
|
||||
json!({"billing_multiplier": 0.25}),
|
||||
] {
|
||||
let policy =
|
||||
resolve_gateway_static_default_routing_policy(GatewayStaticRoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
group_config_json: &config,
|
||||
selection_source: "system_default",
|
||||
requested_model: "model-a",
|
||||
resolved_model: "model-a",
|
||||
})
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let expected = config
|
||||
.get("billing_multiplier")
|
||||
.and_then(Value::as_f64)
|
||||
.unwrap_or(1.0);
|
||||
assert_eq!(policy.billing_multiplier, expected);
|
||||
assert_eq!(
|
||||
crate::routing::build_routing_trace_seed(&policy, "openai:chat").billing_multiplier,
|
||||
Some(expected)
|
||||
);
|
||||
}
|
||||
for value in [json!(-1), json!(null), json!("2"), json!(true)] {
|
||||
let config = json!({"billing_multiplier": value});
|
||||
assert!(resolve_gateway_static_default_routing_policy(
|
||||
GatewayStaticRoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
group_config_json: &config,
|
||||
selection_source: "system_default",
|
||||
requested_model: "model-a",
|
||||
resolved_model: "model-a",
|
||||
},
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_routing_config_is_not_static_default() {
|
||||
let config = json!({
|
||||
@@ -323,17 +378,16 @@ mod tests {
|
||||
"model_policies": [],
|
||||
"rules": []
|
||||
});
|
||||
let static_policy = resolve_gateway_static_default_routing_policy(
|
||||
GatewayStaticRoutingPolicyInput {
|
||||
let static_policy =
|
||||
resolve_gateway_static_default_routing_policy(GatewayStaticRoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
group_config_json: &config,
|
||||
selection_source: "system_default",
|
||||
requested_model: model,
|
||||
resolved_model: model,
|
||||
},
|
||||
)
|
||||
.expect("provider exclusions should defer to the full resolver");
|
||||
})
|
||||
.expect("provider exclusions should defer to the full resolver");
|
||||
assert!(static_policy.is_none());
|
||||
|
||||
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
@@ -356,6 +410,66 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_provider_enablement_requires_full_resolver_and_valid_boolean_values() {
|
||||
for (model, enabled) in [("model-a", false), ("model-b", true)] {
|
||||
let config = json!({ "model_policies": [{
|
||||
"model": "model-a", "provider_enabled_overrides": { "provider-1": false }
|
||||
}] });
|
||||
assert!(resolve_gateway_static_default_routing_policy(
|
||||
GatewayStaticRoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
group_config_json: &config,
|
||||
selection_source: "system_default",
|
||||
requested_model: model,
|
||||
resolved_model: model,
|
||||
}
|
||||
)
|
||||
.unwrap()
|
||||
.is_none());
|
||||
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
group_config_json: &config,
|
||||
selection_source: "system_default",
|
||||
requested_model: model,
|
||||
resolved_model: model,
|
||||
api_format: "openai:chat",
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
headers: &json!({}),
|
||||
body: &json!({}),
|
||||
phase: RoutingRulePhase::ClientRequest,
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
policy.ranking_overlay.provider_allowed("provider-1"),
|
||||
enabled
|
||||
);
|
||||
}
|
||||
for invalid in [json!("false"), json!(0), Value::Null] {
|
||||
let config = json!({ "model_policies": [{
|
||||
"model": "model-a", "provider_enabled_overrides": { "provider-1": invalid }
|
||||
}] });
|
||||
assert!(resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
group_config_json: &config,
|
||||
selection_source: "system_default",
|
||||
requested_model: "model-a",
|
||||
resolved_model: "model-a",
|
||||
api_format: "openai:chat",
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
headers: &json!({}),
|
||||
body: &json!({}),
|
||||
phase: RoutingRulePhase::ClientRequest,
|
||||
})
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_provider_exclusions_preserve_static_default_fast_path() {
|
||||
for config in [json!({}), json!({"disabled_providers": []})] {
|
||||
|
||||
@@ -23,6 +23,9 @@ pub(crate) enum GatewayRoutingSelectionError {
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub(crate) struct GatewayRoutingSelectionInput<'a> {
|
||||
pub explicit_group: Option<&'a str>,
|
||||
/// A public group selected on the user's API key. Explicit request headers
|
||||
/// take precedence, while an unavailable saved choice must fail closed.
|
||||
pub preferred_group: Option<&'a str>,
|
||||
pub user_id: Option<&'a str>,
|
||||
pub api_key_id: Option<&'a str>,
|
||||
pub user_group_ids: &'a [String],
|
||||
@@ -69,6 +72,28 @@ pub(crate) async fn select_gateway_routing_group(
|
||||
});
|
||||
}
|
||||
|
||||
if let Some(preferred) = input
|
||||
.preferred_group
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let group = repository
|
||||
.find_routing_group(RoutingGroupLookupKey::Id(preferred))
|
||||
.await
|
||||
.map_err(repository_selection_error)?
|
||||
.ok_or_else(|| GatewayRoutingSelectionError::NotFound(preferred.to_string()))?;
|
||||
if !group.enabled {
|
||||
return Err(GatewayRoutingSelectionError::Disabled(group.id));
|
||||
}
|
||||
if !has_authenticated_principal(&input) || !routing_group_is_user_visible(&group) {
|
||||
return Err(GatewayRoutingSelectionError::Forbidden(group.id));
|
||||
}
|
||||
return Ok(GatewayRoutingGroupSelection {
|
||||
group: Some(group),
|
||||
source: "api_key_selection".to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
// When there are no bindings at all, no principal-specific lookup can
|
||||
// produce a group. The data-state repository answers this with a cached
|
||||
// existence query, so the common "routing configured but unused" case
|
||||
@@ -128,6 +153,11 @@ async fn explicit_group_allowed(
|
||||
group: &StoredRoutingGroup,
|
||||
input: &GatewayRoutingSelectionInput<'_>,
|
||||
) -> Result<bool, GatewayRoutingSelectionError> {
|
||||
// Public selection is opt-in. Turning it off does not revoke existing
|
||||
// administrator-granted bindings or the system-default compatibility path.
|
||||
if has_authenticated_principal(input) && routing_group_is_user_visible(group) {
|
||||
return Ok(true);
|
||||
}
|
||||
if group.is_system_default {
|
||||
return Ok(true);
|
||||
}
|
||||
@@ -147,6 +177,22 @@ async fn explicit_group_allowed(
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
pub(crate) fn routing_group_is_user_visible(group: &StoredRoutingGroup) -> bool {
|
||||
group
|
||||
.config_json
|
||||
.get("user_visible")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
== Some(true)
|
||||
}
|
||||
|
||||
fn has_authenticated_principal(input: &GatewayRoutingSelectionInput<'_>) -> bool {
|
||||
input
|
||||
.user_id
|
||||
.into_iter()
|
||||
.chain(input.api_key_id)
|
||||
.any(|id| !id.trim().is_empty())
|
||||
}
|
||||
|
||||
fn repository_selection_error(error: impl std::fmt::Display) -> GatewayRoutingSelectionError {
|
||||
GatewayRoutingSelectionError::Repository(error.to_string())
|
||||
}
|
||||
@@ -231,6 +277,225 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn public_groups_allow_authenticated_selection_but_visibility_is_opt_in() {
|
||||
let repository = InMemoryRoutingGroupRepository::default();
|
||||
for (id, config, enabled) in [
|
||||
("public", json!({ "user_visible": true }), true),
|
||||
("private", json!({ "user_visible": false }), true),
|
||||
("legacy", json!({}), true),
|
||||
("malformed", json!({ "user_visible": "true" }), true),
|
||||
("disabled", json!({ "user_visible": true }), false),
|
||||
] {
|
||||
repository
|
||||
.create_routing_group(CreateRoutingGroupRecord {
|
||||
id: id.into(),
|
||||
name: format!("{id}-name"),
|
||||
description: None,
|
||||
enabled,
|
||||
is_system_default: false,
|
||||
sort_order: 0,
|
||||
config_json: config,
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
for (user_id, api_key_id) in [(Some("user-1"), None), (None, Some("key-1"))] {
|
||||
for explicit in ["public", "public-name"] {
|
||||
let selected = select_gateway_routing_group(
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some(explicit),
|
||||
user_id,
|
||||
api_key_id,
|
||||
user_group_ids: &[],
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(selected.group.unwrap().id, "public");
|
||||
}
|
||||
}
|
||||
for id in ["private", "legacy", "malformed"] {
|
||||
assert_eq!(
|
||||
select_gateway_routing_group(
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some(id),
|
||||
user_id: Some("user-1"),
|
||||
..Default::default()
|
||||
}
|
||||
)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
GatewayRoutingSelectionError::Forbidden(id.into())
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
select_gateway_routing_group(
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some("disabled"),
|
||||
user_id: Some("user-1"),
|
||||
..Default::default()
|
||||
}
|
||||
)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
GatewayRoutingSelectionError::Disabled("disabled".into())
|
||||
);
|
||||
assert_eq!(
|
||||
select_gateway_routing_group(
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some("public"),
|
||||
..Default::default()
|
||||
}
|
||||
)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
GatewayRoutingSelectionError::Forbidden("public".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn api_key_selected_public_group_precedes_bindings_and_header_precedes_saved_choice() {
|
||||
let repository = InMemoryRoutingGroupRepository::default();
|
||||
for (id, visible) in [
|
||||
("selected", true),
|
||||
("header", true),
|
||||
("private-default", false),
|
||||
] {
|
||||
repository
|
||||
.create_routing_group(CreateRoutingGroupRecord {
|
||||
id: id.into(),
|
||||
name: id.into(),
|
||||
description: None,
|
||||
enabled: true,
|
||||
is_system_default: false,
|
||||
sort_order: 0,
|
||||
config_json: json!({ "user_visible": visible }),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
repository
|
||||
.create_routing_group_binding(CreateRoutingGroupBindingRecord {
|
||||
id: "admin-default".into(),
|
||||
group_id: "private-default".into(),
|
||||
subject_type: RoutingGroupBindingSubject::ApiKey,
|
||||
subject_id: "key-1".into(),
|
||||
is_default: true,
|
||||
allow_explicit_select: true,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
for (explicit, preferred, expected, source) in [
|
||||
(None, Some("selected"), "selected", "api_key_selection"),
|
||||
(
|
||||
Some("header"),
|
||||
Some("selected"),
|
||||
"header",
|
||||
"explicit_header",
|
||||
),
|
||||
// An explicit authorized request overrides an invalid saved choice.
|
||||
(Some("header"), Some("missing"), "header", "explicit_header"),
|
||||
(
|
||||
Some("private-default"),
|
||||
Some("selected"),
|
||||
"private-default",
|
||||
"explicit_header",
|
||||
),
|
||||
(None, None, "private-default", "api_key_default"),
|
||||
] {
|
||||
let selection = select_gateway_routing_group(
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: explicit,
|
||||
preferred_group: preferred,
|
||||
api_key_id: Some("key-1"),
|
||||
user_id: Some("user-1"),
|
||||
user_group_ids: &[],
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(selection.group.unwrap().id, expected);
|
||||
assert_eq!(selection.source, source);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unavailable_api_key_group_selection_fails_closed_without_default_fallback() {
|
||||
let repository = InMemoryRoutingGroupRepository::default();
|
||||
for (id, visible, enabled, is_default) in [
|
||||
("private-default", false, true, true),
|
||||
("disabled", true, false, false),
|
||||
] {
|
||||
repository
|
||||
.create_routing_group(CreateRoutingGroupRecord {
|
||||
id: id.into(),
|
||||
name: format!("{id}-name"),
|
||||
description: None,
|
||||
enabled,
|
||||
is_system_default: is_default,
|
||||
sort_order: 0,
|
||||
config_json: json!({ "user_visible": visible }),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
for (id, error) in [
|
||||
(
|
||||
"private-default",
|
||||
GatewayRoutingSelectionError::Forbidden("private-default".into()),
|
||||
),
|
||||
(
|
||||
"disabled",
|
||||
GatewayRoutingSelectionError::Disabled("disabled".into()),
|
||||
),
|
||||
(
|
||||
"missing",
|
||||
GatewayRoutingSelectionError::NotFound("missing".into()),
|
||||
),
|
||||
// Saved selections are stable IDs, not mutable group names.
|
||||
(
|
||||
"private-default-name",
|
||||
GatewayRoutingSelectionError::NotFound("private-default-name".into()),
|
||||
),
|
||||
] {
|
||||
assert_eq!(
|
||||
select_gateway_routing_group(
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
preferred_group: Some(id),
|
||||
user_id: Some("user-1"),
|
||||
..Default::default()
|
||||
}
|
||||
)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
error
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn selects_api_key_default_binding() {
|
||||
let repository = InMemoryRoutingGroupRepository::default();
|
||||
@@ -268,6 +533,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: None,
|
||||
preferred_group: None,
|
||||
user_id: None,
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &[],
|
||||
@@ -304,6 +570,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: None,
|
||||
preferred_group: None,
|
||||
user_id: Some("user-1"),
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &["user-group-1".to_string()],
|
||||
@@ -327,7 +594,7 @@ mod tests {
|
||||
enabled: true,
|
||||
is_system_default: false,
|
||||
sort_order: 0,
|
||||
config_json: json!({}),
|
||||
config_json: json!({ "user_visible": false }),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
@@ -353,6 +620,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some("private-group"),
|
||||
preferred_group: None,
|
||||
user_id: Some("user-1"),
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &["team-1".to_string()],
|
||||
@@ -373,6 +641,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some("missing"),
|
||||
preferred_group: None,
|
||||
user_id: Some("user-1"),
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &[],
|
||||
@@ -397,6 +666,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some("group-name"),
|
||||
preferred_group: None,
|
||||
user_id: Some("user-1"),
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &[],
|
||||
@@ -423,6 +693,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: None,
|
||||
preferred_group: None,
|
||||
user_id: Some("user-1"),
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &[],
|
||||
@@ -463,6 +734,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some("disabled-group"),
|
||||
preferred_group: None,
|
||||
user_id: Some("user-1"),
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &[],
|
||||
@@ -514,6 +786,7 @@ mod tests {
|
||||
&repository,
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: Some("private-group"),
|
||||
preferred_group: None,
|
||||
user_id: Some("user-1"),
|
||||
api_key_id: Some("api-key-1"),
|
||||
user_group_ids: &[],
|
||||
|
||||
@@ -5,7 +5,9 @@ pub(crate) fn build_routing_trace_seed(
|
||||
client_api_format: &str,
|
||||
) -> RoutingDecisionTrace {
|
||||
RoutingDecisionTrace {
|
||||
billing_multiplier: Some(policy.billing_multiplier),
|
||||
group_id: policy.group_id.clone(),
|
||||
group_name: policy.group_name.clone(),
|
||||
group_version: policy.group_version,
|
||||
selection_source: policy.selection_source.clone(),
|
||||
selected_rules: policy
|
||||
|
||||
@@ -370,7 +370,7 @@ async fn selects_by_provider_priority_when_priority_mode_is_provider() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn selects_by_global_key_priority_when_priority_mode_is_global_key() {
|
||||
async fn legacy_global_key_routing_default_selects_by_provider_priority() {
|
||||
let mut provider_first = sample_row();
|
||||
provider_first.provider_id = "provider-a".to_string();
|
||||
provider_first.provider_name = "provider-a".to_string();
|
||||
@@ -402,6 +402,13 @@ async fn selects_by_global_key_priority_when_priority_mode_is_global_key() {
|
||||
)
|
||||
.await;
|
||||
|
||||
// Legacy configuration stays readable, but the routing-to-scheduler
|
||||
// boundary normalizes it to the supported provider ordering mode.
|
||||
assert_eq!(
|
||||
ordering_config(&state).await.priority_mode,
|
||||
aether_scheduler_core::SchedulerPriorityMode::Provider
|
||||
);
|
||||
|
||||
let selected = select_candidate(
|
||||
state.data.as_ref(),
|
||||
&state,
|
||||
@@ -415,8 +422,8 @@ async fn selects_by_global_key_priority_when_priority_mode_is_global_key() {
|
||||
.expect("selection should succeed")
|
||||
.expect("candidate should exist");
|
||||
|
||||
assert_eq!(selected.provider_id, "provider-b");
|
||||
assert_eq!(selected.key_id, "key-b");
|
||||
assert_eq!(selected.provider_id, "provider-a");
|
||||
assert_eq!(selected.key_id, "key-a");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -1038,6 +1038,13 @@ async fn gateway_rejects_internal_gateway_report_with_tampered_protected_context
|
||||
("endpoint_id", json!("endpoint-unrelated-victim")),
|
||||
("key_id", json!("key-unrelated-victim")),
|
||||
("client_api_format", json!("gemini:video")),
|
||||
("routing_group_billing_multiplier", json!(0.0)),
|
||||
(
|
||||
"billing_multiplier_snapshot",
|
||||
json!({"version": 1, "factors": {"routing_group": 0.0}, "multiplier": 0.0}),
|
||||
),
|
||||
("routing_group_id", json!("unrelated-group")),
|
||||
("routing_group_name", json!("forged-group")),
|
||||
("task_id", json!("task-unrelated-victim")),
|
||||
("local_task_id", json!("local-task-unrelated-victim")),
|
||||
("local_short_id", json!("short-unrelated-victim")),
|
||||
|
||||
@@ -52,10 +52,13 @@ const TEST_EMAIL_VERIFICATION_TOKEN: &str =
|
||||
|
||||
#[path = "public_support/announcement_user_list.rs"]
|
||||
mod announcement_user_list;
|
||||
mod api_key_routing;
|
||||
#[path = "public_support/auth_cookie.rs"]
|
||||
mod auth_cookie;
|
||||
#[path = "public_support/dashboard.rs"]
|
||||
mod dashboard;
|
||||
#[path = "public_support/routing_groups.rs"]
|
||||
mod routing_groups;
|
||||
#[path = "public_support/vscodex.rs"]
|
||||
mod vscodex;
|
||||
|
||||
|
||||
@@ -0,0 +1,192 @@
|
||||
use super::*;
|
||||
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
|
||||
use aether_data_contracts::repository::routing_profiles::{
|
||||
CreateRoutingGroupRecord, RoutingGroupWriteRepository, UpdateRoutingGroupRecord,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn users_me_api_key_routing_selection_round_trips_and_cannot_bypass_visibility() {
|
||||
let groups = Arc::new(InMemoryRoutingGroupRepository::default());
|
||||
for (id, enabled, visible) in [
|
||||
("public", true, true),
|
||||
("hidden", true, false),
|
||||
("disabled", false, true),
|
||||
] {
|
||||
groups
|
||||
.create_routing_group(CreateRoutingGroupRecord {
|
||||
id: id.into(),
|
||||
name: format!("{id} strategy"),
|
||||
description: None,
|
||||
enabled,
|
||||
is_system_default: false,
|
||||
sort_order: 0,
|
||||
config_json: json!({"user_visible": visible}),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let now = Utc::now();
|
||||
let mut user = sample_auth_user(now);
|
||||
user.role = "user".into();
|
||||
let token = build_test_auth_token(
|
||||
"access",
|
||||
serde_json::Map::from_iter([
|
||||
("user_id".into(), json!(user.id)),
|
||||
("role".into(), json!(user.role)),
|
||||
(
|
||||
"created_at".into(),
|
||||
json!(user.created_at.map(|date| date.to_rfc3339())),
|
||||
),
|
||||
("session_id".into(), json!("session-api-key-routing")),
|
||||
]),
|
||||
now + chrono::Duration::hours(1),
|
||||
);
|
||||
let users = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]));
|
||||
let keys = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
|
||||
let (url, upstream_hits, gateway, upstream) = start_auth_gateway_with_builder(|| {
|
||||
let data = GatewayDataState::with_auth_api_key_repository_for_tests(keys)
|
||||
.with_user_reader(users)
|
||||
.with_routing_group_repository_for_tests(groups.clone())
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
|
||||
AppState::new()
|
||||
.unwrap()
|
||||
.with_data_state_for_tests(data)
|
||||
.with_auth_sessions_for_tests([sample_auth_session(
|
||||
"user-auth-1",
|
||||
"session-api-key-routing",
|
||||
"device-api-key-routing",
|
||||
"refresh-api-key-routing",
|
||||
now,
|
||||
)])
|
||||
})
|
||||
.await;
|
||||
let client = reqwest::Client::new();
|
||||
let endpoint = format!("{url}/api/users/me/api-keys");
|
||||
let request = |method: reqwest::Method, path: &str| {
|
||||
client
|
||||
.request(method, path)
|
||||
.bearer_auth(&token)
|
||||
.header("x-client-device-id", "device-api-key-routing")
|
||||
.header("user-agent", "AetherTest/1.0")
|
||||
};
|
||||
let response = request(reqwest::Method::POST, &endpoint)
|
||||
.json(&json!({
|
||||
"name": "selected key", "routing_group_id": "public",
|
||||
"feature_settings": {"chat_pii_redaction": {"enabled": true}},
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let created: serde_json::Value = response.json().await.unwrap();
|
||||
assert_eq!(created["routing_group_id"], "public");
|
||||
assert_eq!(created["routing_group_name"], "public strategy");
|
||||
assert_eq!(created["feature_settings"]["routing_group_id"], "public");
|
||||
assert!(created["feature_settings"]
|
||||
.get("routing_group_name")
|
||||
.is_none());
|
||||
let detail_url = format!("{endpoint}/{}", created["id"].as_str().unwrap());
|
||||
|
||||
let list: serde_json::Value = request(reqwest::Method::GET, &endpoint)
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(list[0]["routing_group_id"], "public");
|
||||
assert_eq!(list[0]["routing_group_name"], "public strategy");
|
||||
let detail: serde_json::Value = request(reqwest::Method::GET, &detail_url)
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(detail["routing_group_id"], "public");
|
||||
assert_eq!(detail["routing_group_name"], "public strategy");
|
||||
|
||||
let response = request(reqwest::Method::PUT, &detail_url)
|
||||
.json(&json!({
|
||||
"name": "renamed key", "feature_settings": {
|
||||
"chat_pii_redaction": {"enabled": false}, "routing_group_id": "hidden",
|
||||
},
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let edited: serde_json::Value = response.json().await.unwrap();
|
||||
assert_eq!(edited["name"], "renamed key");
|
||||
assert_eq!(edited["routing_group_id"], "public");
|
||||
assert_eq!(
|
||||
edited["feature_settings"]["chat_pii_redaction"]["enabled"],
|
||||
false
|
||||
);
|
||||
for id in ["hidden", "disabled", "missing"] {
|
||||
let rejected = request(reqwest::Method::PUT, &detail_url)
|
||||
.json(&json!({
|
||||
"name": "must not change", "routing_group_id": id,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(rejected.status(), StatusCode::BAD_REQUEST, "id={id}");
|
||||
}
|
||||
let unchanged: serde_json::Value = request(reqwest::Method::GET, &detail_url)
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(unchanged["name"], "renamed key");
|
||||
assert_eq!(unchanged["routing_group_id"], "public");
|
||||
|
||||
groups
|
||||
.update_routing_group(
|
||||
"public",
|
||||
UpdateRoutingGroupRecord {
|
||||
name: Some("renamed strategy".into()),
|
||||
config_json: Some(json!({"user_visible": false})),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let response = request(reqwest::Method::PUT, &detail_url)
|
||||
.json(&json!({
|
||||
"name": "existing hidden choice", "routing_group_id": "public",
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let renamed: serde_json::Value = response.json().await.unwrap();
|
||||
assert_eq!(renamed["routing_group_id"], "public");
|
||||
assert_eq!(renamed["routing_group_name"], "renamed strategy");
|
||||
let response = request(reqwest::Method::PUT, &detail_url)
|
||||
.json(&json!({"routing_group_id": null}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let cleared: serde_json::Value = response.json().await.unwrap();
|
||||
assert_eq!(cleared["routing_group_id"], serde_json::Value::Null);
|
||||
assert_eq!(cleared["routing_group_name"], serde_json::Value::Null);
|
||||
assert!(cleared["feature_settings"]
|
||||
.get("routing_group_id")
|
||||
.is_none());
|
||||
assert_eq!(
|
||||
cleared["feature_settings"]["chat_pii_redaction"]["enabled"],
|
||||
false
|
||||
);
|
||||
|
||||
assert_eq!(*upstream_hits.lock().unwrap(), 0);
|
||||
gateway.abort();
|
||||
upstream.abort();
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
use super::*;
|
||||
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
|
||||
use aether_data_contracts::repository::routing_profiles::{
|
||||
CreateRoutingGroupRecord, RoutingGroupWriteRepository, UpdateRoutingGroupRecord,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn users_me_routing_groups_requires_auth_and_exposes_only_public_enabled_summaries() {
|
||||
let repository = Arc::new(InMemoryRoutingGroupRepository::default());
|
||||
for (id, config, enabled, is_default) in [
|
||||
(
|
||||
"discount",
|
||||
json!({ "user_visible": true, "billing_multiplier": 0.5,
|
||||
"disabled_providers": ["secret-provider"] }),
|
||||
true,
|
||||
false,
|
||||
),
|
||||
("regular", json!({ "user_visible": true }), true, false),
|
||||
(
|
||||
"free",
|
||||
json!({ "user_visible": true, "billing_multiplier": 0.0 }),
|
||||
true,
|
||||
false,
|
||||
),
|
||||
(
|
||||
"private-default",
|
||||
json!({ "user_visible": false }),
|
||||
true,
|
||||
true,
|
||||
),
|
||||
("legacy-private", json!({}), true, false),
|
||||
("disabled", json!({ "user_visible": true }), false, false),
|
||||
] {
|
||||
repository
|
||||
.create_routing_group(CreateRoutingGroupRecord {
|
||||
id: id.into(),
|
||||
name: format!("{id}-name"),
|
||||
description: Some("private-description".into()),
|
||||
enabled,
|
||||
is_system_default: is_default,
|
||||
sort_order: 0,
|
||||
config_json: config,
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let now = Utc::now();
|
||||
let mut user = sample_auth_user(now);
|
||||
user.role = "user".into();
|
||||
let token = build_test_auth_token(
|
||||
"access",
|
||||
serde_json::Map::from_iter([
|
||||
("user_id".into(), json!(user.id)),
|
||||
("role".into(), json!(user.role)),
|
||||
(
|
||||
"created_at".into(),
|
||||
json!(user.created_at.map(|date| date.to_rfc3339())),
|
||||
),
|
||||
("session_id".into(), json!("session-routing-groups")),
|
||||
]),
|
||||
now + chrono::Duration::hours(1),
|
||||
);
|
||||
let users = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]));
|
||||
let (url, upstream_hits, gateway, upstream) = start_auth_gateway_with_builder(|| {
|
||||
let data = GatewayDataState::with_user_reader_for_tests(users)
|
||||
.with_routing_group_repository_for_tests(repository.clone());
|
||||
AppState::new()
|
||||
.unwrap()
|
||||
.with_data_state_for_tests(data)
|
||||
.with_auth_sessions_for_tests([sample_auth_session(
|
||||
"user-auth-1",
|
||||
"session-routing-groups",
|
||||
"device-routing-groups",
|
||||
"refresh-placeholder",
|
||||
now,
|
||||
)])
|
||||
})
|
||||
.await;
|
||||
let client = reqwest::Client::new();
|
||||
let endpoint = format!("{url}/api/users/me/routing-groups");
|
||||
let unauthorized = client.get(&endpoint).send().await.unwrap();
|
||||
assert_eq!(unauthorized.status(), StatusCode::UNAUTHORIZED);
|
||||
let request = || {
|
||||
client
|
||||
.get(&endpoint)
|
||||
.bearer_auth(&token)
|
||||
.header("x-client-device-id", "device-routing-groups")
|
||||
.header("user-agent", "AetherTest/1.0")
|
||||
};
|
||||
let response = request().send().await.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.unwrap();
|
||||
assert_eq!(payload["total"], 3);
|
||||
let items = payload["items"].as_array().unwrap();
|
||||
assert_eq!(items.len(), 3);
|
||||
for (id, multiplier) in [("discount", 0.5), ("regular", 1.0), ("free", 0.0)] {
|
||||
assert_eq!(
|
||||
items.iter().find(|item| item["id"] == id).unwrap(),
|
||||
&json!({
|
||||
"id": id, "name": format!("{id}-name"),
|
||||
"billing_multiplier": multiplier, "is_default": false,
|
||||
})
|
||||
);
|
||||
}
|
||||
assert!(!payload.to_string().contains("private"));
|
||||
assert!(!payload.to_string().contains("secret-provider"));
|
||||
repository
|
||||
.update_routing_group(
|
||||
"private-default",
|
||||
UpdateRoutingGroupRecord {
|
||||
config_json: Some(json!({ "user_visible": true })),
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let payload: serde_json::Value = request().send().await.unwrap().json().await.unwrap();
|
||||
assert_eq!(payload["total"], 4);
|
||||
assert_eq!(
|
||||
payload["items"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.find(|item| item["id"] == "private-default")
|
||||
.unwrap()["is_default"],
|
||||
true
|
||||
);
|
||||
assert_eq!(*upstream_hits.lock().unwrap(), 0);
|
||||
gateway.abort();
|
||||
upstream.abort();
|
||||
}
|
||||
@@ -1449,7 +1449,7 @@ fn usage_repositories_are_owned_by_contracts_and_driver_adapters() {
|
||||
"sql"
|
||||
)
|
||||
.len(),
|
||||
27,
|
||||
29,
|
||||
"all PostgreSQL usage SQL fragments should be owned by the adapter crate"
|
||||
);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user