Merge remote-tracking branch 'origin/main' into codex/provider-policy-hardening

This commit is contained in:
elky
2026-09-05 16:21:10 +08:00
1228 changed files with 241230 additions and 31814 deletions
+4 -2
View File
@@ -65,14 +65,14 @@ http.workspace = true
http-body-util = "0.1"
hyper = { version = "1", features = ["client", "server", "http1", "http2"] }
hyper-util = { version = "0.1", features = ["client-legacy", "client-pool", "server-auto", "service", "tokio"] }
ldap3 = { version = "0.11", default-features = false, features = ["sync", "tls-rustls"] }
ldap3 = { version = "0.12.1", default-features = false, features = ["sync", "tls-rustls-ring"] }
libc = "0.2"
md-5 = "0.10"
object_store.workspace = true
parking_lot = "0.12"
percent-encoding.workspace = true
regex.workspace = true
reqwest.workspace = true
rsa = "0.9.10"
rustls.workspace = true
serde.workspace = true
serde_json.workspace = true
@@ -86,6 +86,7 @@ sysinfo = "0.32"
thiserror.workspace = true
tokio.workspace = true
tokio-util.workspace = true
tokio-tungstenite = { version = "0.28", features = ["rustls-tls-webpki-roots"] }
tower = { version = "0.5", features = ["util"] }
tower-http = { version = "0.6", features = ["fs", "compression-gzip", "set-header"] }
tracing.workspace = true
@@ -102,4 +103,5 @@ tikv-jemalloc-sys = { version = "0.6", optional = true }
[dev-dependencies]
aether-test-support.workspace = true
aws-lc-rs.workspace = true
tracing-subscriber.workspace = true
@@ -30,7 +30,7 @@ struct Args {
#[arg(
long,
env = "AETHER_EXECUTION_RUNTIME_UNIX_SOCKET",
default_value = "/tmp/aether-execution-runtime.sock"
default_value = "/tmp/aether-execution-runtime/aether-execution-runtime.sock"
)]
unix_socket: PathBuf,
+9 -8
View File
@@ -1,17 +1,18 @@
pub(crate) use crate::handlers::admin::{
admin_provider_ops_local_action_response, admin_provider_pool_config,
build_internal_control_error_response, create_provider_oauth_catalog_key,
find_duplicate_provider_oauth_key, maybe_build_local_admin_pool_response,
maybe_build_local_admin_response, persist_provider_quota_refresh_state,
provider_oauth_maintenance_endpoint_for_provider, provider_oauth_runtime_endpoint_for_provider,
provider_quota_refresh_endpoint_for_provider, provider_type_supports_quota_refresh,
reconcile_admin_fixed_provider_template_endpoints,
execute_admin_system_import_exclusively, find_duplicate_provider_oauth_key,
maybe_build_local_admin_pool_response, maybe_build_local_admin_response,
persist_provider_quota_refresh_state, provider_oauth_maintenance_endpoint_for_provider,
provider_oauth_runtime_endpoint_for_provider, provider_quota_refresh_endpoint_for_provider,
provider_type_supports_quota_refresh, reconcile_admin_fixed_provider_template_endpoints,
refresh_provider_oauth_account_state_after_update, refresh_provider_pool_quota_locally,
store_admin_provider_ops_balance_cache, update_existing_provider_oauth_catalog_key,
release_admin_system_import_lease, store_admin_provider_ops_balance_cache,
try_acquire_admin_system_import_lease, update_existing_provider_oauth_catalog_key,
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminLocalOAuthRefreshError,
AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult,
AdminStatsTimeRange, AdminStatsUsageFilter, OAUTH_ACCOUNT_BLOCK_PREFIX,
OAUTH_REQUEST_FAILED_PREFIX,
AdminStatsTimeRange, AdminStatsUsageFilter, AdminSystemImportLockError, SystemExportMode,
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REQUEST_FAILED_PREFIX,
};
use crate::handlers::admin::{
@@ -1,3 +1,4 @@
use aether_usage_runtime::decode_internal_report_body_base64;
use base64::Engine as _;
use serde_json::Value;
@@ -43,9 +44,8 @@ pub(crate) fn maybe_normalize_provider_private_sync_report_payload(
}
if let Some(body_base64) = payload.body_base64.as_deref() {
let body_bytes = base64::engine::general_purpose::STANDARD
.decode(body_base64)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let body_bytes =
decode_internal_report_body_base64(body_base64).map_err(GatewayError::Internal)?;
let Some(normalized_bytes) =
normalize_provider_private_stream_bytes(report_context, &body_bytes)?
else {
@@ -11,10 +11,10 @@ use super::{
convert_gemini_chat_response_to_openai_chat, convert_gemini_response_to_openai_responses,
maybe_build_local_core_sync_finalize_response,
};
use crate::ai_serving::GatewayControlDecision;
use crate::ai_serving::{
convert_openai_chat_response_to_openai_responses,
convert_openai_responses_response_to_openai_chat,
convert_openai_responses_response_to_openai_chat, openai_responses_message_item_id,
GatewayControlDecision,
};
use crate::usage::GatewaySyncReportRequest;
@@ -192,7 +192,7 @@ fn aggregates_openai_responses_stream_completed_event_to_final_response() {
"output_text": "Hello",
"output": [{
"type": "message",
"id": "resp_123_msg",
"id": openai_responses_message_item_id("resp_123", 0),
"role": "assistant",
"status": "completed",
"content": [{
@@ -843,7 +843,7 @@ fn converts_claude_cli_response_to_openai_responses_response() {
"output_text": "Hello Claude CLI",
"output": [{
"type": "message",
"id": "msg_cli_123_msg",
"id": openai_responses_message_item_id("msg_cli_123", 0),
"role": "assistant",
"status": "completed",
"content": [{
@@ -907,7 +907,7 @@ fn converts_claude_cli_tool_use_to_openai_responses_function_call() {
"output": [
{
"type": "message",
"id": "msg_cli_tool_123_msg",
"id": openai_responses_message_item_id("msg_cli_tool_123", 0),
"role": "assistant",
"status": "completed",
"content": [{
@@ -977,7 +977,7 @@ fn converts_gemini_cli_response_to_openai_responses_response() {
"output_text": "Hello Gemini CLI",
"output": [{
"type": "message",
"id": "resp_cli_123_msg",
"id": openai_responses_message_item_id("resp_cli_123", 0),
"role": "assistant",
"status": "completed",
"content": [{
@@ -1046,7 +1046,7 @@ fn converts_gemini_cli_function_call_to_openai_responses_function_call() {
"output": [
{
"type": "message",
"id": "resp_cli_tool_123_msg",
"id": openai_responses_message_item_id("resp_cli_tool_123", 0),
"role": "assistant",
"status": "completed",
"content": [{
@@ -1252,7 +1252,7 @@ fn local_finalize_handles_openai_responses_openai_family_sync_response_even_when
"model": "gpt-5",
"output": [{
"type": "message",
"id": "resp_cli_family_123_msg",
"id": openai_responses_message_item_id("resp_cli_family_123", 0),
"role": "assistant",
"status": "completed",
"content": [{
@@ -9,7 +9,7 @@ use aether_ai_serving::{
use aether_dispatch_core::{DispatchSequence, DispatchSequenceItem};
use aether_routing_core::{
rank_vector_for_candidate, CandidateKind, ResolvedRoutingPolicy, RoutingCandidateFacts,
RoutingCandidateTrace, RoutingDecisionTrace,
RoutingCandidateTrace, RoutingDecisionTrace, RoutingExecutionPolicy,
};
use aether_scheduler_core::{
ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate, SchedulerRankingOutcome,
@@ -47,13 +47,12 @@ use crate::cache::{
use crate::clock::current_unix_ms;
use crate::dispatch::refs::dispatch_ref_for_local_candidate;
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
use crate::orchestration::{ExecutionAttemptIdentity, POOL_KEY_RETRY_INDEX_STRIDE};
use crate::scheduler::candidate::is_auth_api_key_concurrency_limit_skip_reason;
use crate::scheduler::config::SchedulerSchedulingMode;
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100;
const AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET: Duration = Duration::from_millis(100);
const AUTH_API_KEY_CONCURRENCY_RETRY_DELAY: Duration = Duration::from_millis(10);
@@ -80,6 +79,13 @@ type DecorateSkippedCandidateFn<'a> = Arc<
pub(crate) trait LocalExecutionAttemptSource<T>: Send {
async fn next_execution_attempt(&mut self) -> Result<Option<T>, GatewayError>;
/// Returns the request-scoped execution behaviour selected by routing.
/// Execution wrappers use this snapshot before consuming the first
/// attempt, avoiding a second lookup against mutable system settings.
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
None
}
async fn drain_execution_attempts(&mut self) -> Result<Vec<T>, GatewayError>;
async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError>;
@@ -481,10 +487,6 @@ where
type ExtraData = Value;
type Error = Infallible;
fn attempt_slot_count(&self, candidate: &Self::Candidate) -> u32 {
local_attempt_slot_count(&candidate.transport)
}
fn build_extra_data(&self, candidate: &Self::Candidate) -> Option<Self::ExtraData> {
available_candidate_extra_data_with_dispatch_ref(candidate, &self.build_extra_data)
}
@@ -1242,9 +1244,7 @@ async fn scheduler_cache_affinity_enabled(
state: PlannerAppState<'_>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> bool {
scheduler_ordering_config_for_routing_policy(state, routing_policy)
.await
.scheduling_mode
scheduler_ordering_config_for_routing_policy(routing_policy).scheduling_mode
== SchedulerSchedulingMode::CacheAffinity
}
@@ -1610,7 +1610,8 @@ async fn persist_available_local_execution_candidate_at_index<F>(
where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
{
let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1);
// Exactly one attempt is materialized per candidate; same-key retries are
// derived lazily by the attempt loop after a failure.
let extra_data = ai_candidate_extra_data_with_ranking(
available_candidate_base_extra_data_with_dispatch_ref(&candidate, build_extra_data),
candidate.ranking.as_ref(),
@@ -1625,53 +1626,34 @@ where
Some(candidate_index),
extra_data,
);
let should_persist = should_persist_available_local_candidate(&candidate);
let mut attempts = Vec::with_capacity(attempt_slots as usize);
let mut owned_candidate = Some(candidate);
let retry_index = effective_retry_index(0, candidate.orchestration.pool_key_index);
let generated_candidate_id = Uuid::new_v4().to_string();
let candidate_id = if should_persist_available_local_candidate(&candidate) {
state
.persist_available_local_candidate(
trace_id,
context.user_id,
context.api_key_id,
&candidate.candidate,
candidate_index,
retry_index,
generated_candidate_id.as_str(),
context.required_capabilities,
extra_data,
current_unix_ms(),
context.error_context,
)
.await
} else {
generated_candidate_id
};
for retry_index in 0..attempt_slots {
let candidate_ref = owned_candidate
.as_ref()
.expect("candidate should remain available until final retry");
let generated_candidate_id = Uuid::new_v4().to_string();
let candidate_id = if should_persist {
state
.persist_available_local_candidate(
trace_id,
context.user_id,
context.api_key_id,
&candidate_ref.candidate,
candidate_index,
effective_retry_index(retry_index, candidate_ref.orchestration.pool_key_index),
generated_candidate_id.as_str(),
context.required_capabilities,
extra_data.clone(),
current_unix_ms(),
context.error_context,
)
.await
} else {
generated_candidate_id
};
let candidate = if retry_index + 1 == attempt_slots {
owned_candidate
.take()
.expect("final retry should consume owned candidate")
} else {
candidate_ref.clone()
};
let retry_index =
effective_retry_index(retry_index, candidate.orchestration.pool_key_index);
attempts.push(LocalExecutionCandidateAttempt {
eligible: candidate,
candidate_index,
retry_index,
candidate_id,
});
}
attempts
vec![LocalExecutionCandidateAttempt {
eligible: candidate,
candidate_index,
retry_index,
candidate_id,
}]
}
fn available_candidate_extra_data_with_dispatch_ref<F>(
@@ -1840,6 +1822,7 @@ fn routing_trace_for_candidate(
CandidateKind::Provider => Some(candidate.key_id.clone()),
CandidateKind::PoolGroup => None,
},
api_format: Some(candidate.endpoint_api_format.clone()),
provider_priority: candidate.provider_priority,
key_priority: candidate
.key_global_priority_for_format
@@ -1923,32 +1906,15 @@ fn build_unpersisted_local_execution_candidate_attempts(
candidate: EligibleLocalExecutionCandidate,
candidate_index: u32,
) -> VecDeque<LocalExecutionCandidateAttempt> {
let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1);
let mut attempts = VecDeque::with_capacity(attempt_slots as usize);
let mut owned_candidate = Some(candidate);
for retry_index in 0..attempt_slots {
let candidate = if retry_index + 1 == attempt_slots {
owned_candidate
.take()
.expect("final retry should consume owned candidate")
} else {
owned_candidate
.as_ref()
.expect("candidate should remain available until final retry")
.clone()
};
let retry_index =
effective_retry_index(retry_index, candidate.orchestration.pool_key_index);
attempts.push_back(LocalExecutionCandidateAttempt {
eligible: candidate,
candidate_index,
retry_index,
candidate_id: Uuid::new_v4().to_string(),
});
}
attempts
// One attempt per candidate; same-key retries are derived lazily by the
// attempt loop after a failure.
let retry_index = effective_retry_index(0, candidate.orchestration.pool_key_index);
VecDeque::from([LocalExecutionCandidateAttempt {
eligible: candidate,
candidate_index,
retry_index,
candidate_id: Uuid::new_v4().to_string(),
}])
}
async fn persist_pool_group_exhaustion_skipped_candidate(
@@ -2276,6 +2242,8 @@ mod tests {
pool_key_index,
pool_key_lease: None,
scheduler_affinity_epoch: None,
// These tests cover persistence shape, not same-key retries.
sticky_key_attempts: Some(1),
},
ranking: None,
}
@@ -2357,16 +2325,11 @@ mod tests {
assert_eq!(stored.len(), 1);
assert_eq!(stored[0].key_id.as_deref(), Some("normal-key"));
assert_eq!(stored[0].candidate_index, 2);
assert_eq!(
stored[0]
.extra_data
.as_ref()
.and_then(|value| value.get("dispatch_ref"))
.and_then(|value| value.get("SingleKey"))
.and_then(|value| value.get("key"))
.and_then(|value| value.get("key_id")),
Some(&json!("normal-key"))
);
assert!(stored[0]
.extra_data
.as_ref()
.and_then(|value| value.get("dispatch_ref"))
.is_none());
}
#[test]
@@ -2514,14 +2477,23 @@ mod tests {
assert!(should_cache_resolved_candidate_page(&cursor));
let fixed_order_app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests([(
"scheduling_mode".to_string(),
json!("fixed_order"),
)]),
);
let fixed_order_app = AppState::new().expect("state should build");
let fixed_order_policy = ResolvedRoutingPolicy {
group_id: Some("routing-group-fixed-order".to_string()),
group_version: Some(1),
selection_source: "test".to_string(),
requested_model: "gpt-5".to_string(),
resolved_model: "gpt-5".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS,
execution_policy: Default::default(),
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
matched_rules: Vec::new(),
};
let mut page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&fixed_order_app),
&model_directive_policy,
@@ -2531,7 +2503,7 @@ mod tests {
true,
None,
&auth_snapshot,
None,
Some(&fixed_order_policy),
None,
None,
false,
@@ -2549,7 +2521,7 @@ mod tests {
auth_snapshot,
client_session_affinity: None,
required_capabilities: None,
routing_policy: None,
routing_policy: Some(fixed_order_policy),
sticky_session_token: None,
request_auth_channel: None,
skipped_user_id: "user-1".to_string(),
@@ -2647,16 +2619,11 @@ mod tests {
);
assert_eq!(stored[1].key_id.as_deref(), Some("normal-key"));
assert_eq!(stored[1].candidate_index, 1);
assert_eq!(
stored[1]
.extra_data
.as_ref()
.and_then(|value| value.get("dispatch_ref"))
.and_then(|value| value.get("SingleKey"))
.and_then(|value| value.get("key"))
.and_then(|value| value.get("key_id")),
Some(&json!("normal-key"))
);
assert!(stored[1]
.extra_data
.as_ref()
.and_then(|value| value.get("dispatch_ref"))
.is_none());
}
#[test]
@@ -2726,7 +2693,7 @@ mod tests {
.as_ref()
.and_then(serde_json::Value::as_object)
.expect("ranking metadata should persist as object extra data");
assert_eq!(extra_data.get("existing"), Some(&json!("value")));
assert!(extra_data.get("existing").is_none());
assert_eq!(
extra_data.get("ranking_mode"),
Some(&json!("CacheAffinity"))
@@ -2739,14 +2706,7 @@ mod tests {
Some(&json!("cached_affinity"))
);
assert_eq!(extra_data.get("demoted_by"), Some(&json!("cross_format")));
assert_eq!(
extra_data
.get("dispatch_ref")
.and_then(|value| value.get("SingleKey"))
.and_then(|value| value.get("key"))
.and_then(|value| value.get("key_id")),
Some(&json!("ranked-key"))
);
assert!(extra_data.get("dispatch_ref").is_none());
}
#[tokio::test]
@@ -3084,7 +3044,7 @@ mod tests {
.as_ref()
.and_then(serde_json::Value::as_object)
.expect("skipped ranking metadata should persist");
assert_eq!(extra_data.get("existing"), Some(&json!("value")));
assert!(extra_data.get("existing").is_none());
assert_eq!(
extra_data.get("ranking_mode"),
Some(&json!("CacheAffinity"))
@@ -278,13 +278,21 @@ mod tests {
assert_eq!(metadata["transport_diagnostics"]["provider_type"], "codex");
assert_eq!(
metadata["transport_diagnostics"]["fingerprint"]["transport_profile"]["profile_id"],
"chrome_136"
metadata["transport_diagnostics"]["key_fingerprint_configured"],
Value::Bool(true)
);
assert_eq!(
metadata["transport_diagnostics"]["key_transport_profile_configured"],
Value::Bool(true)
);
assert_eq!(
metadata["transport_diagnostics"]["resolved_transport_profile_id"],
"chrome_136"
);
assert_eq!(
metadata["transport_diagnostics"]["resolved_transport_profile"]["profile_id"],
"chrome_136"
);
assert_eq!(
metadata["transport_diagnostics"]["request_pair"]["conversion_enabled"],
Value::Bool(true)
@@ -3,17 +3,14 @@ use aether_ai_serving::{
AiCandidateRankingPort, AiRankableCandidateParts, AiRankingContextConfig,
AiRankingSchedulingMode,
};
use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode};
use aether_routing_core::ResolvedRoutingPolicy;
use async_trait::async_trait;
use tokio::sync::Mutex;
use tracing::warn;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_ms;
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
use crate::scheduler::config::{
read_scheduler_ordering_config, SchedulerOrderingConfig, SchedulerSchedulingMode,
};
use crate::scheduler::config::{SchedulerOrderingConfig, SchedulerSchedulingMode};
use aether_scheduler_core::{
matches_affinity_target, ClientSessionAffinity, SchedulerAffinityTarget,
SchedulerMinimalCandidateSelectionCandidate, SchedulerPriorityMode, SchedulerRankableCandidate,
@@ -133,7 +130,7 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
required_capabilities: Option<&serde_json::Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> Vec<EligibleLocalExecutionCandidate> {
let ordering_config = scheduler_ordering_config_for_routing_policy(state, routing_policy).await;
let ordering_config = scheduler_ordering_config_for_routing_policy(routing_policy);
let port = GatewayLocalCandidateRankingPort {
state,
requested_model,
@@ -184,35 +181,24 @@ fn ai_ranking_scheduling_mode(mode: SchedulerSchedulingMode) -> AiRankingSchedul
}
}
pub(crate) async fn scheduler_ordering_config_for_routing_policy(
state: PlannerAppState<'_>,
/// Return the immutable scheduler snapshot carried by a resolved routing
/// policy. A missing policy is a programming error in production request
/// paths; unit tests may use the scheduler default for isolated ranking tests.
pub(crate) fn scheduler_ordering_config_for_routing_policy(
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> SchedulerOrderingConfig {
let system_config = read_scheduler_ordering_config_or_default(state).await;
match routing_policy {
Some(policy) => {
let mut config = scheduler_ordering_config_from_routing_policy(policy);
config.keep_priority_on_conversion |= system_config.keep_priority_on_conversion;
config
Some(policy) => SchedulerOrderingConfig::from_routing_policy(policy),
None => {
#[cfg(test)]
{
SchedulerOrderingConfig::default()
}
#[cfg(not(test))]
{
panic!("resolved routing policy is required before candidate scheduling")
}
}
None => system_config,
}
}
fn scheduler_ordering_config_from_routing_policy(
policy: &ResolvedRoutingPolicy,
) -> SchedulerOrderingConfig {
SchedulerOrderingConfig {
priority_mode: match policy.priority_mode {
RoutingSetPriorityMode::Provider => SchedulerPriorityMode::Provider,
RoutingSetPriorityMode::GlobalKey => SchedulerPriorityMode::GlobalKey,
},
scheduling_mode: match policy.scheduling_mode {
RoutingSchedulingMode::FixedOrder => SchedulerSchedulingMode::FixedOrder,
RoutingSchedulingMode::CacheAffinity => SchedulerSchedulingMode::CacheAffinity,
RoutingSchedulingMode::LoadBalance => SchedulerSchedulingMode::LoadBalance,
},
keep_priority_on_conversion: policy.keep_priority_on_conversion,
}
}
@@ -231,37 +217,32 @@ fn routing_overlaid_candidate(
let overlaid_key_priority = match kind {
LocalExecutionCandidateKind::SingleKey => policy
.ranking_overlay
.key_priority_overrides
.get(candidate.key_id.as_str()),
.key_priority_override_matching_format(candidate.key_id.as_str(), |format| {
crate::ai_serving::api_format_alias_matches(
format,
candidate.endpoint_api_format.as_str(),
)
})
.or_else(|| {
policy
.ranking_overlay
.key_priority_overrides
.get(candidate.key_id.as_str())
.copied()
}),
LocalExecutionCandidateKind::PoolGroup => policy
.ranking_overlay
.pool_priority_overrides
.get(candidate.provider_id.as_str()),
.get(candidate.provider_id.as_str())
.copied(),
};
if let Some(overlaid_key_priority) = overlaid_key_priority.copied() {
if let Some(overlaid_key_priority) = overlaid_key_priority {
overlaid.key_internal_priority = overlaid_key_priority;
overlaid.key_global_priority_for_format = Some(overlaid_key_priority);
}
overlaid
}
async fn read_scheduler_ordering_config_or_default(
state: PlannerAppState<'_>,
) -> SchedulerOrderingConfig {
match read_scheduler_ordering_config(state.app()).await {
Ok(config) => config,
Err(error) => {
warn!(
event_name = "planner_scheduler_ordering_config_load_failed",
log_type = "event",
error = ?error,
"failed to load scheduler ordering config while ranking local execution candidates"
);
SchedulerOrderingConfig::default()
}
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
@@ -270,10 +251,17 @@ mod tests {
use aether_ai_serving::{
ai_ranking_context, build_ai_rankable_candidate, AiRankableCandidateParts,
};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_data::repository::{
provider_catalog::InMemoryProviderCatalogReadRepository,
routing_profiles::InMemoryRoutingGroupRepository,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::routing_profiles::{
CreateRoutingGroupRecord, RoutingGroupWriteRepository,
};
use aether_scheduler_core::{
apply_scheduler_candidate_ranking,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session,
@@ -303,7 +291,11 @@ mod tests {
required_capabilities: Option<&serde_json::Value>,
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
let normalized_client_api_format = client_api_format.trim().to_ascii_lowercase();
let ordering_config = super::read_scheduler_ordering_config_or_default(state).await;
let ordering_config =
crate::scheduler::config::read_system_default_routing_ordering_config(state.app())
.await
.expect("routing strategy should load")
.unwrap_or_default();
let mut candidates = candidates;
let mut rankables = Vec::with_capacity(candidates.len());
let mut ordering_cache = CandidateTransportRankingFactsCache::default();
@@ -378,6 +370,8 @@ mod tests {
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity,
keep_priority_on_conversion: false,
sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS,
execution_policy: Default::default(),
ranking_overlay: aether_routing_core::RankingOverlay::default(),
mutation_plan: Default::default(),
pool_policy_overrides: BTreeMap::new(),
@@ -396,7 +390,7 @@ mod tests {
}
#[tokio::test]
async fn routing_policy_inherits_global_conversion_priority_override() {
async fn routing_policy_ignores_legacy_global_conversion_priority_override() {
let data_state = GatewayDataState::default().with_system_config_values_for_tests([(
"keep_priority_on_conversion".to_string(),
json!(true),
@@ -413,23 +407,24 @@ mod tests {
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS,
execution_policy: Default::default(),
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
matched_rules: Vec::new(),
};
let ordering = super::scheduler_ordering_config_for_routing_policy(
PlannerAppState::new(&state),
Some(&policy),
)
.await;
let ordering = super::scheduler_ordering_config_for_routing_policy(Some(&policy));
assert_eq!(
ordering.scheduling_mode,
crate::scheduler::config::SchedulerSchedulingMode::FixedOrder
);
assert!(ordering.keep_priority_on_conversion);
assert!(
!ordering.keep_priority_on_conversion,
"a resolved routing policy must not inherit the legacy system-config flag"
);
}
#[test]
@@ -447,6 +442,8 @@ mod tests {
priority_mode: aether_routing_core::RoutingSetPriorityMode::GlobalKey,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity,
keep_priority_on_conversion: false,
sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS,
execution_policy: Default::default(),
ranking_overlay: aether_routing_core::RankingOverlay {
pool_priority_overrides: BTreeMap::from([("provider-1".to_string(), 4)]),
key_priority_overrides: BTreeMap::from([("representative-key".to_string(), 1)]),
@@ -570,6 +567,15 @@ mod tests {
api_formats: Option<serde_json::Value>,
allowed_models: Option<serde_json::Value>,
) -> StoredProviderCatalogKey {
let credential_state = AppState::new()
.expect("credential state should build")
.with_data_state_for_tests(
GatewayDataState::disabled()
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
);
let encrypted_api_key = credential_state
.seal_provider_catalog_key_api_key(provider_id, id, "plain-upstream-key")
.expect("api key should encrypt");
StoredProviderCatalogKey::new(
id.to_string(),
provider_id.to_string(),
@@ -581,7 +587,7 @@ mod tests {
.expect("key should build")
.with_transport_fields(
api_formats,
"plain-upstream-key".to_string(),
encrypted_api_key,
None,
None,
Some(json!({"openai:chat": 1})),
@@ -695,7 +701,7 @@ mod tests {
let observed_at_unix_secs = current_unix_secs();
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
)
.with_system_config_values_for_tests(vec![
("provider_priority_mode".to_string(), json!("provider")),
@@ -704,6 +710,7 @@ mod tests {
serde_json::to_value(TunnelAttachmentRecord {
gateway_instance_id: "gateway-b".to_string(),
relay_base_url: "http://gateway-b:8080".to_string(),
tunnel_generation: "test-generation-remote".to_string(),
conn_count: 1,
observed_at_unix_secs,
})
@@ -714,6 +721,7 @@ mod tests {
serde_json::to_value(TunnelAttachmentRecord {
gateway_instance_id: "gateway-a".to_string(),
relay_base_url: "http://gateway-a:8080".to_string(),
tunnel_generation: "test-generation-local".to_string(),
conn_count: 1,
observed_at_unix_secs,
})
@@ -772,7 +780,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -825,7 +833,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
)
.with_system_config_values_for_tests(vec![(
"scheduling_mode".to_string(),
@@ -882,7 +890,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -918,7 +926,8 @@ mod tests {
}
#[tokio::test]
async fn local_execution_ranking_keeps_cross_format_priority_when_global_override_is_enabled() {
async fn local_execution_ranking_keeps_cross_format_priority_when_strategy_override_is_enabled()
{
let provider_catalog = InMemoryProviderCatalogReadRepository::seed(
vec![
sample_provider_with_options("provider-same", false, 10),
@@ -933,14 +942,32 @@ mod tests {
sample_key_for_provider("provider-cross", "key-cross", ""),
],
);
let routing_repository = std::sync::Arc::new(InMemoryRoutingGroupRepository::default());
routing_repository
.create_routing_group(CreateRoutingGroupRecord {
id: "strategy-default".to_string(),
name: "strategy-default".to_string(),
description: None,
enabled: true,
is_system_default: true,
sort_order: 0,
config_json: json!({
"default_policy": {
"keep_priority_on_conversion": true
}
}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
})
.await
.expect("routing strategy should be created");
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
)
.with_system_config_values_for_tests(vec![(
"keep_priority_on_conversion".to_string(),
json!(true),
)]);
.with_routing_group_repository_for_tests(routing_repository);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
@@ -1000,7 +1027,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
)
.with_system_config_values_for_tests(vec![(
"provider_priority_mode".to_string(),
@@ -1066,7 +1093,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1119,7 +1146,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1193,7 +1220,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1273,7 +1300,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1349,7 +1376,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1416,7 +1443,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1499,7 +1526,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1564,7 +1591,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1653,7 +1680,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1739,7 +1766,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1836,7 +1863,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -1941,7 +1968,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -2035,7 +2062,7 @@ mod tests {
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
DEVELOPMENT_ENCRYPTION_KEY,
);
let state = AppState::new()
.expect("state should build")
@@ -23,7 +23,9 @@ use crate::ai_serving::{
use crate::orchestration::LocalExecutionCandidateMetadata;
use crate::stage_metrics::observe_gateway_stage_ms;
use super::candidate_ranking::rank_eligible_local_execution_candidates;
use super::candidate_ranking::{
rank_eligible_local_execution_candidates, scheduler_ordering_config_for_routing_policy,
};
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct EligibleLocalExecutionCandidate {
@@ -378,8 +380,17 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
"candidate_resolution_core",
started_at.elapsed().as_millis() as u64,
);
let sticky_key_attempts = if outcome.eligible_candidates.is_empty() {
None
} else {
Some(
scheduler_ordering_config_for_routing_policy(routing_policy)
.sticky_key_attempts,
)
};
for candidate in &mut outcome.eligible_candidates {
candidate.orchestration.scheduler_affinity_epoch = Some(scheduler_affinity_epoch);
candidate.orchestration.sticky_key_attempts = sticky_key_attempts;
}
(outcome.eligible_candidates, outcome.skipped_candidates)
}
@@ -174,6 +174,9 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
self.ranking_seed,
false,
self.request_operation,
super::candidate_ranking::scheduler_ordering_config_for_routing_policy(
self.routing_policy,
),
)
.await?;
@@ -425,11 +428,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
);
let ordering_config =
super::candidate_ranking::scheduler_ordering_config_for_routing_policy(
state,
routing_policy,
)
.await;
super::candidate_ranking::scheduler_ordering_config_for_routing_policy(routing_policy);
Self {
state,
@@ -1291,6 +1290,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
.then_some(self.client_session_affinity.as_ref())
.flatten(),
self.ranking_seed,
self.ordering_config,
)
.await?;
let skipped_candidates = skipped_candidates
@@ -1473,6 +1473,7 @@ mod tests {
use super::*;
use crate::data::GatewayDataState;
use crate::AppState;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::DataLayerError;
@@ -1884,6 +1885,8 @@ mod tests {
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS,
execution_policy: Default::default(),
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
@@ -1947,6 +1950,8 @@ mod tests {
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS,
execution_policy: Default::default(),
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
@@ -2170,6 +2175,19 @@ mod tests {
None,
)
.expect("endpoint transport should build");
let credential_state = AppState::new()
.expect("credential state should build")
.with_data_state_for_tests(
GatewayDataState::disabled()
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
);
let encrypted_api_key = credential_state
.seal_provider_catalog_key_api_key(
row.provider_id.as_str(),
row.key_id.as_str(),
"plain-upstream-key",
)
.expect("api key should encrypt");
let key = StoredProviderCatalogKey::new(
row.key_id.clone(),
row.provider_id.clone(),
@@ -2181,7 +2199,7 @@ mod tests {
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!([row.endpoint_api_format.clone()])),
"plain-upstream-key".to_string(),
encrypted_api_key,
None,
None,
None,
@@ -2536,7 +2554,7 @@ mod tests {
provider_repository,
candidate_repository,
)
.with_encryption_key_for_tests("development-key");
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let app = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(data_state);
@@ -2656,15 +2674,17 @@ mod tests {
provider_repository,
candidate_repository,
)
.with_encryption_key_for_tests("development-key")
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
// Legacy keys deliberately disagree with the routing policy: the
// resolved policy must be the only source of scheduler ordering.
.with_system_config_values_for_tests([
(
"scheduling_mode".to_string(),
serde_json::json!("fixed_order"),
serde_json::json!("cache_affinity"),
),
(
"keep_priority_on_conversion".to_string(),
serde_json::json!(true),
serde_json::json!(false),
),
]);
let app = AppState::new()
@@ -2681,7 +2701,9 @@ mod tests {
resolved_model: "gpt-5.4-mini".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
keep_priority_on_conversion: true,
sticky_key_attempts: aether_routing_core::DEFAULT_STICKY_KEY_ATTEMPTS,
execution_policy: Default::default(),
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
@@ -11,6 +11,7 @@ use crate::ai_serving::planner::route::{
is_matching_stream_request, resolve_execution_runtime_stream_plan_kind,
};
use crate::ai_serving::{resolve_decision_execution_runtime_auth_context, GatewayControlDecision};
use crate::state::VideoTaskRouteAccess;
use crate::{AiExecutionDecision, AppState, GatewayError};
pub(crate) async fn maybe_build_stream_decision_payload(
@@ -155,16 +156,37 @@ async fn maybe_build_local_video_task_content_stream_decision_payload(
return Ok(None);
}
let _ = state
.hydrate_video_task_for_route(decision.route_family.as_deref(), parts.uri.path())
.await?;
let Some(user_id) = decision
.auth_context
.as_ref()
.filter(|auth_context| auth_context.access_allowed)
.map(|auth_context| auth_context.user_id.trim())
.filter(|value| !value.is_empty())
else {
return Err(crate::video_tasks::not_found_error());
};
if state
.hydrate_video_task_for_route_for_user(
decision.route_family.as_deref(),
parts.uri.path(),
user_id,
)
.await?
!= VideoTaskRouteAccess::Allowed
{
return Err(crate::video_tasks::not_found_error());
}
let Some(action) = state.video_tasks.prepare_openai_content_stream_action(
parts.uri.path(),
parts.uri.query(),
trace_id,
) else {
return Ok(None);
let Some(action) = state
.video_tasks
.prepare_openai_content_stream_action_for_user(
parts.uri.path(),
parts.uri.query(),
trace_id,
user_id,
)
else {
return Err(crate::video_tasks::not_found_error());
};
let crate::video_tasks::LocalVideoTaskContentAction::StreamPlan(plan) = action else {
@@ -16,6 +16,7 @@ use crate::ai_serving::{
build_execution_runtime_auth_context, resolve_execution_runtime_auth_context,
GatewayControlDecision,
};
use crate::state::VideoTaskRouteAccess;
use crate::{AiExecutionDecision, AppState, GatewayError};
pub(crate) async fn maybe_build_sync_decision_payload(
@@ -191,10 +192,6 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload(
return Ok(None);
}
let _ = state
.hydrate_video_task_for_route(decision.route_family.as_deref(), parts.uri.path())
.await?;
let auth_context = resolve_execution_runtime_auth_context(
state,
decision,
@@ -204,16 +201,30 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload(
)
.await?;
let Some(auth_context) = auth_context else {
return Ok(None);
return Err(crate::video_tasks::not_found_error());
};
let Some(follow_up) = state.video_tasks.prepare_follow_up_sync_plan(
if !auth_context.access_allowed || auth_context.user_id.trim().is_empty() {
return Err(crate::video_tasks::not_found_error());
}
if state
.hydrate_video_task_for_route_for_user(
decision.route_family.as_deref(),
parts.uri.path(),
&auth_context.user_id,
)
.await?
!= VideoTaskRouteAccess::Allowed
{
return Err(crate::video_tasks::not_found_error());
}
let Some(follow_up) = state.video_tasks.prepare_follow_up_sync_plan_for_user(
plan_kind,
parts.uri.path(),
Some(body_json),
Some(&auth_context),
trace_id,
) else {
return Ok(None);
return Err(crate::video_tasks::not_found_error());
};
let aether_video_tasks_core::LocalVideoTaskFollowUpPlan {
@@ -236,8 +247,7 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload(
downstream_path = %parts.uri.path(),
provider_api_format = %plan.provider_api_format,
client_api_format = %plan.client_api_format,
upstream_base_url = ?upstream_base_url,
upstream_url = %plan.url,
upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url),
"gateway built local video follow-up sync decision payload"
);
@@ -37,6 +37,10 @@ const ROUTING_GROUP_SELECTION_CACHE_TTL: Duration = Duration::from_secs(30);
const ROUTING_GROUP_SELECTION_CACHE_STALE_TTL: Duration = Duration::from_secs(120);
const CODEX_ACCOUNT_ID_HEADER: &str = "chatgpt-account-id";
const CODEX_FEDRAMP_HEADER: &str = "x-openai-fedramp";
const INVALID_ROUTING_PROVIDER_CONTRACT_MESSAGE: &str =
"routing provider request violates provider contract";
const INVALID_ROUTING_PROVIDER_HEADERS_MESSAGE: &str =
"invalid provider request headers in routing mutation";
#[derive(Debug, Clone)]
pub(crate) struct ResolvedLocalDecisionAuthInput {
@@ -292,6 +296,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m
crate::ai_serving::openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
provider_model.as_str(),
)
})
.unwrap_or_default();
@@ -311,10 +316,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m
)
}
}
.map_err(|violation| GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: format!("routing provider_request violates provider contract: {violation:?}"),
})?;
.map_err(|_| invalid_routing_provider_contract())?;
}
let provider_model = provider_request_body
.get("model")
@@ -649,21 +651,17 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
GatewayRoutingSelectionError::NotFound(explicit_group.unwrap_or_default()),
));
}
None
return Err(routing_selection_error(
GatewayRoutingSelectionError::NoDefault,
));
}
};
let Some((group_id, group_version, group_config_json, selection_source)) = selected_group
else {
input.client_session_affinity = client_session_affinity_from_api_request(
client_api_format,
&parts.headers,
Some(body_json),
);
input.routing_policy = None;
input.routing_trace_seed = None;
input.routing_context = None;
return Ok(());
return Err(routing_selection_error(
GatewayRoutingSelectionError::NoDefault,
));
};
if try_attach_static_default_routing_policy_to_input(
@@ -887,10 +885,36 @@ fn routing_selection_error(error: GatewayRoutingSelectionError) -> GatewayError
GatewayRoutingSelectionError::Repository(message) => {
GatewayError::Internal(format!("routing group repository lookup failed: {message}"))
}
error => GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: error.to_string(),
GatewayRoutingSelectionError::NoDefault => GatewayError::Client {
status: StatusCode::SERVICE_UNAVAILABLE,
message: "no enabled routing strategy is configured for this request".to_string(),
},
GatewayRoutingSelectionError::NotFound(_) => GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: "requested routing group was not found".to_string(),
},
GatewayRoutingSelectionError::Disabled(_) => GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: "requested routing group is not enabled".to_string(),
},
GatewayRoutingSelectionError::Forbidden(_) => GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: "requested routing group is not allowed for this principal".to_string(),
},
}
}
fn invalid_routing_provider_contract() -> GatewayError {
GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: INVALID_ROUTING_PROVIDER_CONTRACT_MESSAGE.to_string(),
}
}
fn invalid_routing_provider_headers() -> GatewayError {
GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: INVALID_ROUTING_PROVIDER_HEADERS_MESSAGE.to_string(),
}
}
@@ -945,14 +969,9 @@ fn btree_headers_to_header_map(
) -> Result<HeaderMap, GatewayError> {
let mut output = HeaderMap::new();
for (name, value) in headers {
let name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: format!("invalid provider request header name in routing mutation: {err}"),
})?;
let value = HeaderValue::from_str(value).map_err(|err| GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: format!("invalid provider request header value in routing mutation: {err}"),
})?;
let name = HeaderName::from_bytes(name.as_bytes())
.map_err(|_| invalid_routing_provider_headers())?;
let value = HeaderValue::from_str(value).map_err(|_| invalid_routing_provider_headers())?;
output.insert(name, value);
}
Ok(output)
@@ -1097,6 +1116,7 @@ fn ensure_report_context_routing_trace(
endpoint_id: decision.endpoint_id.clone().unwrap_or_default(),
model_id,
key_id,
api_format: decision.provider_api_format.clone(),
provider_priority,
key_priority,
},
@@ -1174,6 +1194,50 @@ mod tests {
}
}
#[test]
fn routing_selection_errors_do_not_echo_explicit_group() {
let secret = "private-group?token=Bearer-secret";
for error in [
GatewayRoutingSelectionError::NotFound(secret.to_string()),
GatewayRoutingSelectionError::Disabled(secret.to_string()),
GatewayRoutingSelectionError::Forbidden(secret.to_string()),
] {
let error = routing_selection_error(error);
assert!(matches!(
error,
GatewayError::Client {
status: StatusCode::FORBIDDEN,
ref message,
} if !message.contains(secret)
));
}
}
#[test]
fn routing_provider_errors_do_not_echo_dynamic_details() {
let secret = "https://internal.example/?token=Bearer-secret";
let contract_error = invalid_routing_provider_contract();
let header_error = btree_headers_to_header_map(&BTreeMap::from([(
format!("Authorization: {secret}"),
secret.to_string(),
)]))
.expect_err("invalid header should fail");
for (error, expected_message) in [
(contract_error, INVALID_ROUTING_PROVIDER_CONTRACT_MESSAGE),
(header_error, INVALID_ROUTING_PROVIDER_HEADERS_MESSAGE),
] {
assert!(matches!(
error,
GatewayError::Client {
status: StatusCode::BAD_REQUEST,
ref message,
} if message == expected_message && !message.contains(secret)
));
}
}
#[tokio::test]
async fn explicit_auth_snapshot_override_does_not_fall_back_to_the_planner_cache() {
// AppState::new has no auth snapshot repository. Without the explicit
@@ -1230,6 +1294,7 @@ mod tests {
description: None,
enabled: true,
is_system_default: false,
sort_order: 0,
config_json: json!({}),
version: 1,
created_at: 1,
@@ -140,6 +140,9 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
current_unix_secs(),
false,
spec.operation.map(|operation| operation.as_str()),
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await?;
let outcome = materialize_local_execution_candidates_with_serving(
@@ -246,6 +249,9 @@ pub(crate) async fn build_local_same_format_provider_candidate_attempt_source<'a
current_unix_secs(),
false,
spec.operation.map(|operation| operation.as_str()),
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await?;
@@ -203,6 +203,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
sticky_key_attempts: eligible.orchestration.sticky_key_attempts,
client_requested_stream: body_json
.get("stream")
.and_then(serde_json::Value::as_bool)
@@ -183,6 +183,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
let reasoning_replay_policy = openai_responses_reasoning_replay_policy(
prepared.transport.provider.provider_type.as_str(),
prepared.transport.endpoint.base_url.as_str(),
prepared.mapped_model.as_str(),
);
let redaction = resolve_provider_chat_pii_redaction(
state,
@@ -26,6 +26,7 @@ use super::{
LocalSameFormatProviderCandidateAttemptSource, LocalSameFormatProviderDecisionInput,
LocalSameFormatProviderSpec,
};
use aether_routing_core::RoutingExecutionPolicy;
pub(crate) struct LocalSameFormatProviderSyncAttemptSource<'a> {
state: &'a AppState,
@@ -189,6 +190,13 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
@@ -234,6 +242,13 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncA
impl LocalExecutionAttemptSource<AiStreamAttempt>
for LocalSameFormatProviderStreamAttemptSource<'_>
{
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
@@ -4,7 +4,7 @@ use aether_ai_serving::{
build_ai_execution_report_context,
insert_provider_stream_event_api_format as insert_ai_provider_stream_event_api_format,
provider_stream_event_api_format_for_provider_type as ai_provider_stream_event_api_format_for_provider_type,
AiExecutionReportContextParts, AiRequestOrigin,
AiExecutionReportContextParts, AiRequestOrigin, STICKY_KEY_ATTEMPTS_REPORT_FIELD,
};
use aether_routing_core::ResolvedRoutingPolicy;
use aether_runtime_state::RuntimeLockLease;
@@ -21,7 +21,8 @@ use crate::client_session_affinity::{
};
use crate::orchestration::{
insert_pool_key_lease_report_context_fields, ExecutionAttemptIdentity,
ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD,
ROUTING_EXECUTION_POLICY_REPORT_FIELD, ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD,
SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD,
};
use crate::scheduler::affinity::insert_scheduler_affinity_policy_report_context_field;
@@ -59,6 +60,9 @@ pub(crate) struct LocalExecutionReportContextParts<'a> {
pub(crate) client_session_affinity: Option<&'a ClientSessionAffinity>,
pub(crate) routing_policy: Option<&'a ResolvedRoutingPolicy>,
pub(crate) scheduler_affinity_epoch: Option<u64>,
/// Routing policy sticky-key attempt budget; read back by the attempt
/// loop to derive same-key retries lazily.
pub(crate) sticky_key_attempts: Option<u32>,
pub(crate) client_requested_stream: bool,
pub(crate) upstream_is_stream: bool,
pub(crate) has_envelope: bool,
@@ -72,10 +76,12 @@ pub(crate) fn build_local_execution_report_context(
let RequestOrigin {
client_ip,
user_agent,
forwarded_headers_trusted,
} = parts
.request_origin
.unwrap_or_else(|| request_origin_from_headers(parts.original_headers));
let original_headers = crate::ai_serving::collect_control_headers(parts.original_headers);
let original_headers =
collect_report_context_original_headers(parts.original_headers, forwarded_headers_trusted);
let original_request_body = crate::ai_serving::build_report_context_original_request_echo(
parts.original_request_body_json,
parts.original_request_body_base64,
@@ -102,13 +108,20 @@ pub(crate) fn build_local_execution_report_context(
value,
);
}
if let Some(incoming_tls) =
crate::ai_serving::tls_fingerprint_from_headers(parts.original_headers)
{
merge_incoming_tls_fingerprint(&mut extra_fields, incoming_tls);
if forwarded_headers_trusted {
if let Some(incoming_tls) =
crate::ai_serving::tls_fingerprint_from_headers(parts.original_headers)
{
merge_incoming_tls_fingerprint(&mut extra_fields, incoming_tls);
}
}
insert_pool_key_lease_report_context_fields(&mut extra_fields, parts.pool_key_lease);
insert_scheduler_affinity_policy_report_context_field(&mut extra_fields, parts.routing_policy);
if let Some(policy) = parts.routing_policy {
if let Ok(value) = serde_json::to_value(policy.execution_policy) {
extra_fields.insert(ROUTING_EXECUTION_POLICY_REPORT_FIELD.to_string(), value);
}
}
if let Some(override_policy) = parts
.routing_policy
.and_then(|policy| policy.pool_policy_overrides.get(parts.provider_id))
@@ -124,6 +137,12 @@ pub(crate) fn build_local_execution_report_context(
Value::Number(epoch.into()),
);
}
if let Some(sticky_key_attempts) = parts.sticky_key_attempts {
extra_fields.insert(
STICKY_KEY_ATTEMPTS_REPORT_FIELD.to_string(),
Value::Number(sticky_key_attempts.into()),
);
}
insert_request_path_fields(
&mut extra_fields,
parts.request_path,
@@ -174,6 +193,17 @@ pub(crate) fn build_local_execution_report_context(
})
}
fn collect_report_context_original_headers(
headers: &http::HeaderMap,
forwarded_headers_trusted: bool,
) -> BTreeMap<String, String> {
let mut collected = crate::ai_serving::collect_control_headers(headers);
if !forwarded_headers_trusted {
collected.retain(|name, _| !name.starts_with("x-aether-tls-"));
}
collected
}
fn insert_request_path_fields(
extra_fields: &mut Map<String, Value>,
request_path: Option<&str>,
@@ -243,8 +273,8 @@ mod tests {
use serde_json::{json, Map, Value};
use super::{
build_local_execution_report_context, provider_stream_event_api_format_for_provider_type,
LocalExecutionReportContextParts,
build_local_execution_report_context, collect_report_context_original_headers,
provider_stream_event_api_format_for_provider_type, LocalExecutionReportContextParts,
};
use crate::ai_serving::ExecutionRuntimeAuthContext;
use crate::ai_serving::RequestOrigin;
@@ -274,6 +304,26 @@ mod tests {
);
}
#[test]
fn untrusted_tls_forwarding_headers_are_excluded_from_report_context() {
let mut headers = http::HeaderMap::new();
headers.insert("x-aether-tls-ja3", "spoofed-ja3".parse().unwrap());
headers.insert(http::header::USER_AGENT, "test-client".parse().unwrap());
let untrusted = collect_report_context_original_headers(&headers, false);
assert!(!untrusted.contains_key("x-aether-tls-ja3"));
assert_eq!(
untrusted.get("user-agent").map(String::as_str),
Some("test-client")
);
let trusted = collect_report_context_original_headers(&headers, true);
assert_eq!(
trusted.get("x-aether-tls-ja3").map(String::as_str),
Some("spoofed-ja3")
);
}
#[test]
fn local_execution_report_context_records_request_origin_and_session_affinity() {
let auth_context = ExecutionRuntimeAuthContext {
@@ -324,12 +374,14 @@ mod tests {
request_origin: Some(RequestOrigin {
client_ip: Some("203.0.113.8".to_string()),
user_agent: Some("Claude-Code/1.0".to_string()),
forwarded_headers_trusted: false,
}),
original_request_body_json: Some(&json!({"model": "gpt-5"})),
original_request_body_base64: None,
client_session_affinity: Some(&client_session_affinity),
routing_policy: None,
scheduler_affinity_epoch: None,
sticky_key_attempts: None,
client_requested_stream: false,
upstream_is_stream: false,
has_envelope: false,
@@ -413,6 +465,7 @@ mod tests {
client_session_affinity: None,
routing_policy: None,
scheduler_affinity_epoch: None,
sticky_key_attempts: None,
client_requested_stream: false,
upstream_is_stream: true,
has_envelope: false,
@@ -474,12 +527,17 @@ mod tests {
original_headers: &original_headers,
request_path: None,
request_query_string: None,
request_origin: None,
request_origin: Some(RequestOrigin {
client_ip: None,
user_agent: None,
forwarded_headers_trusted: true,
}),
original_request_body_json: Some(&json!({"model": "gpt-5"})),
original_request_body_base64: None,
client_session_affinity: None,
routing_policy: None,
scheduler_affinity_epoch: None,
sticky_key_attempts: None,
client_requested_stream: false,
upstream_is_stream: false,
has_envelope: false,
@@ -17,6 +17,7 @@ use crate::ai_serving::{
resolve_gemini_files_sync_spec as resolve_sync_spec, LocalGeminiFilesSpec,
};
use crate::{AiExecutionDecision, AppState, GatewayError};
use aether_routing_core::RoutingExecutionPolicy;
use self::decision::maybe_build_local_gemini_files_decision_payload_for_candidate;
use self::support::{
@@ -174,6 +175,13 @@ pub(crate) async fn build_local_gemini_files_stream_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
@@ -212,6 +220,13 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptS
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
@@ -109,6 +109,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
sticky_key_attempts: eligible.orchestration.sticky_key_attempts,
client_requested_stream: spec_metadata.require_streaming,
upstream_is_stream: spec_metadata.require_streaming,
has_envelope: false,
@@ -8,7 +8,10 @@ use crate::ai_serving::transport::{
GeminiFilesRequestBodyError,
};
use crate::ai_serving::GEMINI_FILES_UPLOAD_PLAN_KIND;
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
use crate::ai_serving::{
CandidateFailureDiagnostic, GatewayProviderTransportSnapshot, GEMINI_FILES_DELETE_PLAN_KIND,
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_FILES_GET_PLAN_KIND,
};
use crate::AppState;
use super::support::{
@@ -47,6 +50,26 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
let transport = &attempt.eligible.transport;
let effective_headers = input.effective_headers(&parts.headers);
if matches!(
spec_metadata.decision_kind,
GEMINI_FILES_GET_PLAN_KIND
| GEMINI_FILES_DELETE_PLAN_KIND
| GEMINI_FILES_DOWNLOAD_PLAN_KIND
) && !candidate_matches_owned_gemini_file_mapping(state, parts, input, attempt).await
{
mark_skipped_local_gemini_files_candidate(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"gemini_file_mapping_mismatch",
)
.await;
return None;
}
if let Some(skip_reason) =
gemini_files_transport_unsupported_reason(transport, GEMINI_FILES_CANDIDATE_API_FORMAT)
{
@@ -191,3 +214,64 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
file_name,
})
}
async fn candidate_matches_owned_gemini_file_mapping(
state: &AppState,
parts: &http::request::Parts,
input: &LocalGeminiFilesDecisionInput,
attempt: &LocalGeminiFilesCandidateAttempt,
) -> bool {
let Some(file_name) = normalize_gemini_file_name_from_path(parts.uri.path()) else {
return false;
};
let user_id = input.auth_context.user_id.trim();
if user_id.is_empty() || !state.has_gemini_file_mapping_data_reader() {
return false;
}
let Ok(Some(mapping)) = state
.find_active_gemini_file_mapping_for_owner(
file_name.as_str(),
&attempt.eligible.transport.key.id,
user_id,
crate::clock::current_unix_secs(),
)
.await
else {
return false;
};
mapping.user_id.as_deref().map(str::trim) == Some(user_id)
&& mapping.key_id == attempt.eligible.transport.key.id
}
pub(crate) fn normalize_gemini_file_name_from_path(path: &str) -> Option<String> {
let suffix = path.strip_prefix("/v1beta/files/")?.trim_matches('/');
let suffix = suffix.strip_suffix(":download").unwrap_or(suffix).trim();
let suffix = suffix.strip_prefix("files/").unwrap_or(suffix).trim();
if suffix.is_empty() || suffix.contains('/') {
return None;
}
Some(format!("files/{suffix}"))
}
#[cfg(test)]
mod tests {
use super::normalize_gemini_file_name_from_path;
#[test]
fn normalizes_supported_gemini_file_object_paths() {
assert_eq!(
normalize_gemini_file_name_from_path("/v1beta/files/file-123"),
Some("files/file-123".to_string())
);
assert_eq!(
normalize_gemini_file_name_from_path("/v1beta/files/file-123:download"),
Some("files/file-123".to_string())
);
assert_eq!(
normalize_gemini_file_name_from_path("/v1beta/files/files/abc-123"),
Some("files/abc-123".to_string())
);
assert_eq!(normalize_gemini_file_name_from_path("/v1beta/files"), None);
}
}
@@ -108,6 +108,9 @@ pub(super) async fn materialize_local_gemini_files_candidate_attempts(
Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(),
current_unix_secs(),
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await?;
let outcome = materialize_local_execution_candidates_with_serving(
@@ -181,6 +184,9 @@ pub(super) async fn build_local_gemini_files_candidate_attempt_source<'a>(
Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(),
current_unix_secs(),
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await?;
Ok(build_local_execution_candidate_attempt_source_with_serving(
@@ -19,6 +19,7 @@ use crate::ai_serving::{
resolve_local_image_sync_spec as resolve_sync_spec,
};
use crate::{AiExecutionDecision, AppState, GatewayError};
use aether_routing_core::RoutingExecutionPolicy;
use self::decision::maybe_build_local_openai_image_decision_payload_for_candidate;
use self::support::{
@@ -252,6 +253,13 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
@@ -290,6 +298,13 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptS
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
@@ -122,6 +122,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
sticky_key_attempts: eligible.orchestration.sticky_key_attempts,
client_requested_stream: spec_metadata.require_streaming,
upstream_is_stream,
has_envelope: false,
@@ -9,6 +9,7 @@ use crate::ai_serving::planner::candidate_preparation::{
};
use crate::ai_serving::planner::spec_metadata::local_openai_image_spec_metadata;
use crate::ai_serving::pure::normalize_openai_image_request_with_options;
use crate::ai_serving::transport::antigravity::is_antigravity_provider_transport;
use crate::ai_serving::transport::{
build_grok_browser_headers, build_grok_upstream_url, build_openai_image_headers,
build_openai_image_upstream_url, build_standard_provider_request_headers,
@@ -338,6 +339,25 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport;
let provider_api_format = "gemini:generate_content";
// The gemini:generate_content URL hook rewrites an Antigravity endpoint to
// /v1internal:, and this image path has no v1internal envelope to match it.
// Skip the candidate instead of posting a bare Gemini body that upstream
// would only reject.
if is_antigravity_provider_transport(transport) {
mark_skipped_local_openai_image_candidate(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"transport_unsupported",
)
.await;
return None;
}
let effective_headers = input.effective_headers(&parts.headers);
let prepared_candidate = match prepare_header_authenticated_candidate(
@@ -127,6 +127,9 @@ pub(super) async fn list_local_openai_image_candidate_attempts(
input.client_session_affinity.as_ref(),
current_unix_secs(),
false,
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await
{
@@ -201,6 +204,9 @@ pub(super) async fn build_local_openai_image_candidate_attempt_source<'a>(
input.client_session_affinity.as_ref(),
current_unix_secs(),
false,
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await
{
@@ -16,6 +16,7 @@ use crate::ai_serving::{
LocalVideoCreateSpec,
};
use crate::{AiExecutionDecision, AppState, GatewayError};
use aether_routing_core::RoutingExecutionPolicy;
use self::decision::maybe_build_local_video_create_decision_payload_for_candidate;
use self::support::{
@@ -104,6 +105,13 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
@@ -90,6 +90,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
sticky_key_attempts: eligible.orchestration.sticky_key_attempts,
client_requested_stream: false,
upstream_is_stream: false,
has_envelope: false,
@@ -133,6 +133,9 @@ pub(super) async fn list_local_video_create_candidate_attempts(
input.client_session_affinity.as_ref(),
current_unix_secs(),
false,
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await
{
@@ -190,6 +193,9 @@ pub(super) async fn build_local_video_create_candidate_attempt_source<'a>(
input.client_session_affinity.as_ref(),
current_unix_secs(),
false,
crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy(
input.routing_policy.as_ref(),
),
)
.await
{
@@ -505,7 +505,7 @@ fn projects_uuid_prompt_cache_identity_into_missing_session_headers() {
assert_eq!(headers.get("x-client-request-id"), None);
assert_eq!(
headers.get("user-agent"),
Some(&"codex_cli_rs/0.144.1".to_string())
Some(&"codex_cli_rs/0.153.3".to_string())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert!(!headers.contains_key("version"));
@@ -615,7 +615,7 @@ fn injects_only_codex_client_headers_for_images_requests() {
);
assert_eq!(
headers.get("user-agent"),
Some(&"codex_cli_rs/0.144.1".to_string())
Some(&"codex_cli_rs/0.153.3".to_string())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert!(!headers.contains_key("version"));
@@ -699,7 +699,7 @@ fn preserves_client_context_headers_and_enforces_codex_provider_identity() {
);
assert_eq!(
headers.get("user-agent"),
Some(&"codex_cli_rs/0.144.1".to_string())
Some(&"codex_cli_rs/0.153.3".to_string())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert_eq!(
@@ -763,7 +763,7 @@ fn compact_projects_uuid_prompt_cache_identity_into_session_headers() {
assert_eq!(headers.get("x-client-request-id"), None);
assert_eq!(
headers.get("user-agent"),
Some(&"codex_cli_rs/0.144.1".to_string())
Some(&"codex_cli_rs/0.153.3".to_string())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert!(!headers.contains_key("version"));
@@ -15,11 +15,25 @@ pub(crate) fn is_deepseek_provider(provider_type: &str, base_url: &str) -> bool
host == "deepseek.com" || host.ends_with(".deepseek.com")
}
fn is_deepseek_model(provider_model: &str) -> bool {
let provider_model = provider_model.trim().to_ascii_lowercase();
let leaf = provider_model
.rsplit(['/', ':'])
.next()
.unwrap_or(provider_model.as_str());
leaf == "deepseek" || leaf.starts_with("deepseek-") || leaf.starts_with("deepseek_")
}
fn is_deepseek_upstream(provider_type: &str, base_url: &str, provider_model: &str) -> bool {
is_deepseek_provider(provider_type, base_url) || is_deepseek_model(provider_model)
}
pub(crate) fn openai_responses_reasoning_replay_policy(
provider_type: &str,
base_url: &str,
provider_model: &str,
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
if is_deepseek_provider(provider_type, base_url) {
if is_deepseek_upstream(provider_type, base_url, provider_model) {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
} else {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
@@ -33,7 +47,11 @@ pub(crate) fn apply_deepseek_tool_call_thinking_compat(
provider_api_format: &str,
original_request_body: Option<&Value>,
) {
if !is_deepseek_provider(provider_type, base_url) {
let provider_model = provider_request_body
.get("model")
.and_then(Value::as_str)
.unwrap_or_default();
if !is_deepseek_upstream(provider_type, base_url, provider_model) {
return;
}
@@ -302,11 +320,35 @@ mod tests {
));
assert!(!is_deepseek_provider("custom", "ftp://api.deepseek.com/v1"));
assert_eq!(
openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1"),
openai_responses_reasoning_replay_policy(
"custom",
"https://api.deepseek.com/v1",
"deepseek-v4-flash",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
);
assert_eq!(
openai_responses_reasoning_replay_policy("openai", "https://api.openai.com/v1"),
openai_responses_reasoning_replay_policy(
"openai",
"https://api.openai.com/v1",
"gpt-5.6-sol",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
assert_eq!(
openai_responses_reasoning_replay_policy(
"custom",
"https://api.b.ai/v1",
"deepseek-v4-flash",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
);
assert_eq!(
openai_responses_reasoning_replay_policy(
"custom",
"https://api.b.ai/v1",
"not-deepseek-compatible",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
}
@@ -330,8 +372,11 @@ mod tests {
"input": reasoning_items.clone(),
"future_request_field": {"preserve": true}
});
let replay_policy =
openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1");
let replay_policy = openai_responses_reasoning_replay_policy(
"custom",
"https://api.deepseek.com/v1",
"deepseek-v4-flash",
);
let mut provider_body = crate::ai_serving::build_standard_request_body_with_model_directives_and_request_headers_and_reasoning_replay_policy(
&request,
"openai:responses",
@@ -373,7 +418,11 @@ mod tests {
crate::ai_serving::strip_incompatible_openai_responses_reasoning_items_with_policy(
&mut deepseek,
"openai:responses",
openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1"),
openai_responses_reasoning_replay_policy(
"custom",
"https://api.deepseek.com/v1",
"deepseek-v4-flash",
),
),
0
);
@@ -383,7 +432,11 @@ mod tests {
crate::ai_serving::strip_incompatible_openai_responses_reasoning_items_with_policy(
&mut openai,
"openai:responses",
openai_responses_reasoning_replay_policy("openai", "https://api.openai.com/v1"),
openai_responses_reasoning_replay_policy(
"openai",
"https://api.openai.com/v1",
"gpt-5.6-sol",
),
),
66
);
@@ -417,6 +470,52 @@ mod tests {
assert_eq!(body["messages"][1]["reasoning_content"], "");
}
#[test]
fn custom_relay_deepseek_model_adds_chat_thinking_compat() {
let mut body = json!({
"model": "deepseek-v4-flash",
"messages": [
{"role": "user", "content": "inspect the repository"},
{"role": "assistant", "content": null, "tool_calls": [{
"id": "call_1",
"type": "function",
"function": {"name": "inspect", "arguments": "{}"}
}]},
{"role": "tool", "tool_call_id": "call_1", "content": "done"}
]
});
apply_deepseek_tool_call_thinking_compat(
&mut body,
"custom",
"https://api.b.ai/v1",
"openai:chat",
None,
);
assert_eq!(body["thinking"]["type"], "enabled");
assert_eq!(body["messages"][1]["reasoning_content"], "");
}
#[test]
fn custom_relay_non_deepseek_model_is_not_rewritten() {
let original = json!({
"model": "not-deepseek-compatible",
"messages": [{"role": "assistant", "content": "done"}]
});
let mut body = original.clone();
apply_deepseek_tool_call_thinking_compat(
&mut body,
"custom",
"https://api.b.ai/v1",
"openai:chat",
None,
);
assert_eq!(body, original);
}
#[test]
fn openai_chat_deepseek_honors_disabled_thinking() {
let original = json!({"reasoning_effort": "none"});
@@ -18,6 +18,7 @@ use crate::ai_serving::planner::spec_metadata::{
};
use crate::ai_serving::GatewayControlDecision;
use crate::{AiExecutionDecision, AppState, GatewayError};
use aether_routing_core::RoutingExecutionPolicy;
use super::candidates::{
build_local_standard_candidate_attempt_source, resolve_local_standard_decision_input,
@@ -177,6 +178,13 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
@@ -220,6 +228,13 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
@@ -142,6 +142,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
sticky_key_attempts: eligible.orchestration.sticky_key_attempts,
client_requested_stream: body_json
.get("stream")
.and_then(serde_json::Value::as_bool)
@@ -4,6 +4,10 @@ use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value;
use crate::ai_serving::planner::antigravity::{
build_antigravity_v1internal_provider_request, AntigravityV1InternalRequestError,
AntigravityV1InternalRequestInput, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME,
};
use crate::ai_serving::planner::candidate_preparation::{
prepare_header_authenticated_candidate, prepare_header_authenticated_candidate_from_auth,
OauthPreparationContext,
@@ -26,6 +30,7 @@ use crate::ai_serving::planner::standard::{
openai_provider_request_contract_failure_extra_data, openai_responses_reasoning_replay_policy,
request_body_build_failure_extra_data, request_conversion_failure_extra_data,
};
use crate::ai_serving::transport::antigravity::is_antigravity_provider_transport;
use crate::ai_serving::transport::kiro::{
build_kiro_provider_headers, build_kiro_provider_request_body,
is_kiro_claude_messages_transport, KiroProviderHeadersInput, KiroRequestAuth,
@@ -587,7 +592,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
}
};
crate::ai_serving::hydrate_openai_response_history(
state.runtime_state(),
state,
body_json,
spec_metadata.api_format,
provider_api_format,
@@ -597,6 +602,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
let reasoning_replay_policy = openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
prepared_candidate.mapped_model.as_str(),
);
let redaction = resolve_provider_chat_pii_redaction(
state,
@@ -836,6 +842,29 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
.await);
}
if normalized_provider_api_format == "gemini:generate_content"
&& is_antigravity_provider_transport(transport)
{
return Ok(build_antigravity_cross_format_payload_parts(
state,
parts,
trace_id,
body_json,
input,
attempt,
transport,
spec_metadata.api_format,
provider_api_format,
prepared_candidate.mapped_model,
prepared_candidate.auth_header,
prepared_candidate.auth_value,
provider_request_body,
upstream_is_stream,
redaction.redacted,
)
.await);
}
if normalized_provider_api_format == "gemini:generate_content"
&& is_gemini_cli_provider_transport(transport)
{
@@ -962,6 +991,145 @@ fn apply_transport_request_body_semantics(
)
}
#[allow(clippy::too_many_arguments)]
async fn build_antigravity_cross_format_payload_parts(
state: &AppState,
parts: &http::request::Parts,
trace_id: &str,
original_body_json: &serde_json::Value,
input: &LocalStandardDecisionInput,
attempt: &LocalStandardCandidateAttempt,
transport: &Arc<GatewayProviderTransportSnapshot>,
client_api_format: &str,
provider_api_format: &str,
mapped_model: String,
auth_header: String,
auth_value: String,
gemini_request_body: Value,
upstream_is_stream: bool,
request_redacted: bool,
) -> Option<LocalStandardCandidatePayloadParts> {
let candidate = &attempt.eligible.candidate;
let effective_headers = input.effective_headers(&parts.headers);
let resolved =
match build_antigravity_v1internal_provider_request(AntigravityV1InternalRequestInput {
state,
parts,
transport,
trace_id,
mapped_model: &mapped_model,
provider_api_format,
auth_header: &auth_header,
auth_value: &auth_value,
request_headers: effective_headers,
original_request_body: original_body_json,
gemini_request_body: &gemini_request_body,
upstream_is_stream,
same_format: false,
})
.await
{
Ok(resolved) => resolved,
Err(AntigravityV1InternalRequestError::TransportUnsupported) => {
mark_skipped_local_standard_candidate(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"transport_unsupported",
)
.await;
return None;
}
Err(AntigravityV1InternalRequestError::EnvelopeUnsupported) => {
mark_skipped_local_standard_candidate_with_extra_data(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"provider_request_body_build_failed",
request_body_build_failure_extra_data(
original_body_json,
client_api_format,
provider_api_format,
),
)
.await;
return None;
}
Err(AntigravityV1InternalRequestError::UpstreamUrlUnavailable) => {
mark_skipped_local_standard_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"upstream_url_missing",
CandidateFailureDiagnostic::upstream_url_missing(
client_api_format,
provider_api_format,
"standard_family_antigravity_url",
),
)
.await;
return None;
}
Err(AntigravityV1InternalRequestError::HeaderRulesApplyFailed) => {
mark_skipped_local_standard_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
client_api_format,
provider_api_format,
"standard_family_antigravity_headers",
),
)
.await;
return None;
}
};
let mut provider_request_headers = resolved.headers.headers;
apply_codex_openai_special_headers(
&mut provider_request_headers,
&resolved.body,
effective_headers,
resolved.transport.provider.provider_type.as_str(),
provider_api_format,
Some(trace_id),
resolved.transport.key.decrypted_auth_config.as_deref(),
);
request_identity_response_encoding_when_redacted(
&mut provider_request_headers,
request_redacted,
);
Some(LocalStandardCandidatePayloadParts {
auth_header: resolved.headers.auth_header,
auth_value: resolved.headers.auth_value,
mapped_model,
provider_api_format: provider_api_format.to_string(),
provider_request_body: resolved.body,
provider_request_headers,
upstream_url: resolved.upstream_url,
upstream_is_stream,
envelope_name: Some(ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME),
transport: resolved.transport,
transport_profile: None,
request_redacted,
})
}
#[allow(clippy::too_many_arguments)]
async fn build_gemini_cli_cross_format_payload_parts(
state: &AppState,
@@ -195,6 +195,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
sticky_key_attempts: eligible.orchestration.sticky_key_attempts,
client_requested_stream: body_json
.get("stream")
.and_then(serde_json::Value::as_bool)
@@ -159,6 +159,7 @@ fn finalize_openai_chat_provider_request_body(
openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
mapped_model,
),
)
.err()
@@ -2740,7 +2741,7 @@ mod tests {
.provider_request_headers
.get("x-client-version")
.map(String::as_str),
Some("1.2.3")
Some("4.3.0")
);
assert_eq!(
payload
@@ -2760,7 +2761,7 @@ mod tests {
assert_eq!(payload.provider_request_body["model"], "gemini-2.5-pro");
assert_eq!(
payload.provider_request_body["userAgent"],
"antigravity/cli/1.0.16 (aidev_client; os_type=linux; arch=arm64; auth_method=consumer)"
"vscode/1.X.X (Antigravity/4.3.0)"
);
assert_eq!(payload.provider_request_body["requestType"], "agent");
assert!(payload.provider_request_body.get("contents").is_none());
@@ -1,3 +1,4 @@
use aether_routing_core::RoutingExecutionPolicy;
use async_trait::async_trait;
use std::collections::VecDeque;
use tracing::warn;
@@ -119,6 +120,13 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
let select_started_at = std::time::Instant::now();
let selected = self.next_execution_attempt_with_target_select().await?;
@@ -1,3 +1,4 @@
use aether_routing_core::RoutingExecutionPolicy;
use async_trait::async_trait;
use tracing::warn;
@@ -92,6 +93,13 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
@@ -1,7 +1,6 @@
use std::collections::BTreeMap;
use aether_contracts::RequestBody;
use tracing::debug;
use super::super::{
augment_sync_report_context, build_ai_execution_plan_from_decision,
@@ -10,7 +9,6 @@ use super::super::{
AiStreamAttempt,
};
use crate::ai_serving::planner::common::enforce_provider_body_stream_policy;
use crate::ai_serving::planner::redaction::sanitize_upstream_url_for_log;
use crate::ai_serving::provider_adaptation_requires_eventstream_accept;
use crate::ai_serving::transport::{
build_standard_plan_fallback_headers, build_standard_plan_fallback_openai_chat_url,
@@ -157,21 +155,16 @@ pub(crate) fn build_openai_responses_stream_plan_from_decision(
let Some(auth_pair) = take_ai_upstream_auth_pair(&mut payload) else {
return Ok(None);
};
let (url, url_source) = if let Some(upstream_url) =
take_non_empty_string(&mut payload.upstream_url)
{
(upstream_url, "upstream_url")
let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) {
upstream_url
} else {
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
return Ok(None);
};
(
build_standard_plan_fallback_openai_responses_url(
&upstream_base_url,
parts.uri.query(),
compact,
),
"upstream_base_url",
build_standard_plan_fallback_openai_responses_url(
&upstream_base_url,
parts.uri.query(),
compact,
)
};
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
@@ -238,16 +231,7 @@ pub(crate) fn build_openai_responses_stream_plan_from_decision(
.uri
.query()
.and_then(crate::ai_serving::api::sanitize_request_query_string);
let log_decision_upstream_base_url = payload
.upstream_base_url
.as_deref()
.map(sanitize_upstream_url_for_log);
let log_decision_upstream_url = payload
.upstream_url
.as_deref()
.map(sanitize_upstream_url_for_log);
let log_plan_url = sanitize_upstream_url_for_log(plan.url.as_str());
debug!(
tracing::debug!(
event_name = "local_openai_responses_stream_plan_built",
log_type = "debug",
request_id = %plan.request_id,
@@ -255,12 +239,13 @@ pub(crate) fn build_openai_responses_stream_plan_from_decision(
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
downstream_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query(
parts.uri.path(),
parts.uri.query(),
).unwrap_or_else(|| "/".to_string()),
upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url),
downstream_path = %parts.uri.path(),
downstream_query = ?log_downstream_query,
url_source,
decision_upstream_base_url = ?log_decision_upstream_base_url,
decision_upstream_url = ?log_decision_upstream_url,
plan_url = %log_plan_url,
client_api_format = %plan.client_api_format,
provider_api_format = %plan.provider_api_format,
upstream_is_stream = effective_upstream_is_stream,
@@ -1,7 +1,6 @@
use std::collections::BTreeMap;
use aether_contracts::RequestBody;
use tracing::debug;
use super::super::{
augment_sync_report_context, build_ai_execution_plan_from_decision,
@@ -10,7 +9,6 @@ use super::super::{
AiSyncAttempt,
};
use crate::ai_serving::planner::common::enforce_provider_body_stream_policy;
use crate::ai_serving::planner::redaction::sanitize_upstream_url_for_log;
use crate::ai_serving::transport::{
build_standard_plan_fallback_headers, build_standard_plan_fallback_openai_chat_url,
build_standard_plan_fallback_openai_responses_url, StandardPlanFallbackAcceptPolicy,
@@ -142,21 +140,16 @@ pub(crate) fn build_openai_responses_sync_plan_from_decision(
let Some(auth_pair) = take_ai_upstream_auth_pair(&mut payload) else {
return Ok(None);
};
let (url, url_source) = if let Some(upstream_url) =
take_non_empty_string(&mut payload.upstream_url)
{
(upstream_url, "upstream_url")
let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) {
upstream_url
} else {
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
return Ok(None);
};
(
build_standard_plan_fallback_openai_responses_url(
&upstream_base_url,
parts.uri.query(),
compact,
),
"upstream_base_url",
build_standard_plan_fallback_openai_responses_url(
&upstream_base_url,
parts.uri.query(),
compact,
)
};
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
@@ -205,16 +198,7 @@ pub(crate) fn build_openai_responses_sync_plan_from_decision(
.uri
.query()
.and_then(crate::ai_serving::api::sanitize_request_query_string);
let log_decision_upstream_base_url = payload
.upstream_base_url
.as_deref()
.map(sanitize_upstream_url_for_log);
let log_decision_upstream_url = payload
.upstream_url
.as_deref()
.map(sanitize_upstream_url_for_log);
let log_plan_url = sanitize_upstream_url_for_log(plan.url.as_str());
debug!(
tracing::debug!(
event_name = "local_openai_responses_sync_plan_built",
log_type = "debug",
request_id = %plan.request_id,
@@ -222,12 +206,13 @@ pub(crate) fn build_openai_responses_sync_plan_from_decision(
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
downstream_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query(
parts.uri.path(),
parts.uri.query(),
).unwrap_or_else(|| "/".to_string()),
upstream_origin = %crate::handlers::shared::security_log_url_origin(&plan.url),
downstream_path = %parts.uri.path(),
downstream_query = ?log_downstream_query,
url_source,
decision_upstream_base_url = ?log_decision_upstream_base_url,
decision_upstream_url = ?log_decision_upstream_url,
plan_url = %log_plan_url,
client_api_format = %plan.client_api_format,
provider_api_format = %plan.provider_api_format,
upstream_is_stream = payload.upstream_is_stream,
@@ -3,7 +3,6 @@ use tracing::debug;
use crate::ai_serving::build_request_trace_proxy_value;
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision_with_websocket_mode;
use crate::ai_serving::planner::redaction::sanitize_upstream_url_for_log;
use crate::ai_serving::planner::report_context::{
build_local_execution_report_context, insert_native_client_envelope_name,
insert_provider_stream_event_api_format, LocalExecutionReportContextParts,
@@ -184,6 +183,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
sticky_key_attempts: eligible.orchestration.sticky_key_attempts,
client_requested_stream: body_json
.get("stream")
.and_then(serde_json::Value::as_bool)
@@ -204,12 +204,10 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
&resolved.transport,
);
let log_base_url = sanitize_upstream_url_for_log(resolved.transport.endpoint.base_url.as_str());
let log_request_query = parts
.uri
.query()
.and_then(crate::ai_serving::api::sanitize_request_query_string);
let log_upstream_url = sanitize_upstream_url_for_log(resolved.upstream_url.as_str());
debug!(
event_name = "local_openai_responses_decision_payload_built",
log_type = "debug",
@@ -226,9 +224,12 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
client_api_format = spec_metadata.api_format,
provider_api_format = %resolved.provider_api_format,
request_path = %parts.uri.path(),
request_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query(
parts.uri.path(),
parts.uri.query(),
).unwrap_or_else(|| "/".to_string()),
upstream_origin = %crate::handlers::shared::security_log_url_origin(&resolved.upstream_url),
request_query = ?log_request_query,
upstream_base_url = %log_base_url,
upstream_url = %log_upstream_url,
upstream_is_stream = resolved.upstream_is_stream,
has_envelope = resolved.envelope_name.is_some(),
"gateway built local openai responses decision payload"
@@ -24,7 +24,6 @@ use crate::ai_serving::planner::gemini_cli::{
};
use crate::ai_serving::planner::redaction::{
request_identity_response_encoding_when_redacted, resolve_provider_chat_pii_redaction,
sanitize_upstream_url_for_log,
};
use crate::ai_serving::planner::spec_metadata::local_openai_responses_spec_metadata;
use crate::ai_serving::planner::standard::{
@@ -428,7 +427,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_
}
};
crate::ai_serving::hydrate_openai_response_history(
state.runtime_state(),
state,
body_json,
spec_metadata.api_format,
provider_api_format,
@@ -438,6 +437,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_
let reasoning_replay_policy = openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
mapped_model.as_str(),
);
let redaction = resolve_provider_chat_pii_redaction(
state,
@@ -866,17 +866,10 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_
let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats(spec_metadata.api_format, provider_api_format);
let log_base_url = sanitize_upstream_url_for_log(transport.endpoint.base_url.as_str());
let log_custom_path = transport
.endpoint
.custom_path
.as_deref()
.map(sanitize_upstream_url_for_log);
let log_request_query = parts
.uri
.query()
.and_then(crate::ai_serving::api::sanitize_request_query_string);
let log_upstream_url = sanitize_upstream_url_for_log(upstream_url.as_str());
debug!(
event_name = "local_openai_responses_upstream_url_resolved",
@@ -892,12 +885,14 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_
provider_api_format = %provider_api_format,
execution_strategy = execution_strategy.as_str(),
conversion_mode = conversion_mode.as_str(),
base_url = %log_base_url,
custom_path = ?log_custom_path,
request_path_and_query = %crate::ai_serving::pure::sanitize_request_path_and_query(
parts.uri.path(),
parts.uri.query(),
).unwrap_or_else(|| "/".to_string()),
upstream_origin = %crate::handlers::shared::security_log_url_origin(&upstream_url),
request_path = %parts.uri.path(),
request_query = ?log_request_query,
mapped_model = %mapped_model,
upstream_url = %log_upstream_url,
upstream_is_stream,
"gateway resolved local openai responses upstream url"
);
@@ -2010,8 +2005,6 @@ async fn build_kiro_openai_responses_payload_parts(
};
let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats(client_api_format, provider_api_format);
let log_upstream_url = sanitize_upstream_url_for_log(upstream_url.as_str());
debug!(
event_name = "local_openai_responses_kiro_upstream_url_resolved",
log_type = "debug",
@@ -2026,7 +2019,7 @@ async fn build_kiro_openai_responses_payload_parts(
provider_api_format = %provider_api_format,
execution_strategy = execution_strategy.as_str(),
conversion_mode = conversion_mode.as_str(),
upstream_url = %log_upstream_url,
upstream_origin = %crate::handlers::shared::security_log_url_origin(&upstream_url),
upstream_is_stream,
"gateway resolved local openai responses kiro upstream url"
);
@@ -1005,6 +1005,7 @@ pub(crate) async fn maybe_build_responses_websocket_decision(
reasoning_replay_policy: openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
mapped_model.as_str(),
),
model_directive_patch: input
.model_directive_policy
@@ -1,3 +1,4 @@
use aether_routing_core::RoutingExecutionPolicy;
use async_trait::async_trait;
use tracing::warn;
@@ -161,6 +162,13 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
@@ -204,6 +212,13 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStreamAttemptSource<'_> {
fn routing_execution_policy(&self) -> Option<RoutingExecutionPolicy> {
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
@@ -8,9 +8,13 @@ use crate::constants::{
API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS, API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS,
};
use crate::scheduler::candidate::SchedulerSkippedCandidate;
use crate::scheduler::config::SchedulerOrderingConfig;
use crate::GatewayError;
impl<'a> PlannerAppState<'a> {
/// `ordering_config` is the immutable scheduler snapshot derived from the
/// request's resolved routing policy.
#[allow(clippy::too_many_arguments)]
pub(crate) async fn list_selectable_candidates(
self,
api_format: &str,
@@ -21,6 +25,7 @@ impl<'a> PlannerAppState<'a> {
client_session_affinity: Option<&ClientSessionAffinity>,
now_unix_secs: u64,
enable_model_directives: bool,
ordering_config: SchedulerOrderingConfig,
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, GatewayError> {
crate::scheduler::candidate::list_selectable_candidates(
self.app().data.as_ref(),
@@ -33,10 +38,12 @@ impl<'a> PlannerAppState<'a> {
client_session_affinity,
now_unix_secs,
enable_model_directives,
ordering_config,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn list_selectable_candidates_with_skip_reasons(
self,
api_format: &str,
@@ -47,6 +54,7 @@ impl<'a> PlannerAppState<'a> {
client_session_affinity: Option<&ClientSessionAffinity>,
now_unix_secs: u64,
enable_model_directives: bool,
ordering_config: SchedulerOrderingConfig,
) -> Result<
(
Vec<SchedulerMinimalCandidateSelectionCandidate>,
@@ -64,10 +72,12 @@ impl<'a> PlannerAppState<'a> {
now_unix_secs,
enable_model_directives,
None,
ordering_config,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn list_selectable_candidates_with_skip_reasons_for_request_operation(
self,
api_format: &str,
@@ -79,6 +89,7 @@ impl<'a> PlannerAppState<'a> {
now_unix_secs: u64,
enable_model_directives: bool,
request_operation: Option<&str>,
ordering_config: SchedulerOrderingConfig,
) -> Result<
(
Vec<SchedulerMinimalCandidateSelectionCandidate>,
@@ -103,6 +114,7 @@ impl<'a> PlannerAppState<'a> {
attempt_now_unix_secs,
enable_model_directives,
request_operation,
ordering_config,
)
.await?;
@@ -123,6 +135,7 @@ impl<'a> PlannerAppState<'a> {
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn list_selectable_enumerated_candidates_with_skip_reasons(
self,
api_format: &str,
@@ -132,6 +145,7 @@ impl<'a> PlannerAppState<'a> {
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
now_unix_secs: u64,
ordering_config: SchedulerOrderingConfig,
) -> Result<
(
Vec<SchedulerMinimalCandidateSelectionCandidate>,
@@ -148,10 +162,12 @@ impl<'a> PlannerAppState<'a> {
auth_snapshot,
client_session_affinity,
now_unix_secs,
ordering_config,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn list_selectable_candidates_for_required_capability_without_requested_model(
self,
candidate_api_format: &str,
@@ -160,6 +176,7 @@ impl<'a> PlannerAppState<'a> {
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
now_unix_secs: u64,
ordering_config: SchedulerOrderingConfig,
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, GatewayError> {
let wait_timeout = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS);
let wait_interval = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS.max(1));
@@ -176,6 +193,7 @@ impl<'a> PlannerAppState<'a> {
auth_snapshot,
client_session_affinity,
attempt_now_unix_secs,
ordering_config,
)
.await?;
@@ -177,6 +177,7 @@ pub(crate) use aether_ai_formats::{
api_format_defaults_to_client_error_failover, api_format_defaults_to_non_stream,
api_format_permission_covers, codex_responses_lite_tool_is_client_executed,
intersect_api_format_allowed_lists, is_embedding_api_format, is_rerank_api_format,
normalize_openai_responses_message_item_ids, openai_responses_message_item_id,
openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id,
strip_incompatible_openai_responses_reasoning_items,
strip_incompatible_openai_responses_reasoning_items_with_policy, ApiOperation, ClientSurface,
@@ -2,14 +2,15 @@ use crate::ai_serving::{
hydrate_response_history, normalize_api_format_alias, record_converted_response_history,
response_history_is_loaded, response_history_storage_key, ResponseHistoryRecord,
};
use aether_runtime_state::RuntimeState;
use serde_json::Value;
use tracing::warn;
use crate::GatewayError;
use crate::{AppState, GatewayError};
const RESPONSE_HISTORY_SECRET_PURPOSE: &str = "openai-response-history";
pub(crate) async fn hydrate_openai_response_history(
runtime_state: &RuntimeState,
state: &AppState,
request: &Value,
client_api_format: &str,
provider_api_format: &str,
@@ -33,6 +34,7 @@ pub(crate) async fn hydrate_openai_response_history(
}
let storage_key = response_history_storage_key(previous_response_id, Some(history_scope));
let runtime_state = state.runtime_state();
let payload = runtime_state.kv_get(&storage_key).await.map_err(|error| {
warn!(
event_name = "openai_response_history_read_failed",
@@ -46,8 +48,24 @@ pub(crate) async fn hydrate_openai_response_history(
let Some(payload) = payload else {
return Ok(());
};
let Some(payload) = crate::handlers::shared::open_runtime_secret_payload(
state,
RESPONSE_HISTORY_SECRET_PURPOSE,
&payload,
) else {
let _ = runtime_state.kv_delete(&storage_key).await;
warn!(
event_name = "openai_response_history_decryption_failed",
log_type = "ops",
backend = runtime_state.backend_kind().as_str(),
"gateway rejected undecryptable shared OpenAI response history"
);
return Err(GatewayError::Internal(
"OpenAI response history decryption failed".to_string(),
));
};
if let Err(error) =
hydrate_response_history(previous_response_id, Some(history_scope), &payload)
hydrate_response_history(previous_response_id, Some(history_scope), payload.as_str())
{
let _ = runtime_state.kv_delete(&storage_key).await;
warn!(
@@ -65,11 +83,25 @@ pub(crate) async fn hydrate_openai_response_history(
}
pub(crate) async fn persist_response_history_record(
runtime_state: &RuntimeState,
state: &AppState,
record: ResponseHistoryRecord,
) {
let runtime_state = state.runtime_state();
let Some(sealed_payload) = crate::handlers::shared::seal_runtime_secret_payload(
state,
RESPONSE_HISTORY_SECRET_PURPOSE,
&record.payload,
) else {
warn!(
event_name = "openai_response_history_encryption_unavailable",
log_type = "ops",
backend = runtime_state.backend_kind().as_str(),
"gateway refused to persist unencrypted OpenAI response history"
);
return;
};
if let Err(error) = runtime_state
.kv_set(&record.storage_key, record.payload, Some(record.ttl))
.kv_set(&record.storage_key, sealed_payload, Some(record.ttl))
.await
{
warn!(
@@ -83,7 +115,7 @@ pub(crate) async fn persist_response_history_record(
}
pub(crate) async fn persist_converted_response_history(
runtime_state: &RuntimeState,
state: &AppState,
report_context: &Value,
response: Option<&Value>,
) {
@@ -91,6 +123,126 @@ pub(crate) async fn persist_converted_response_history(
return;
};
if let Some(record) = record_converted_response_history(report_context, response) {
persist_response_history_record(runtime_state, record).await;
persist_response_history_record(state, record).await;
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState};
use serde_json::json;
use sha2::{Digest, Sha256};
use super::{
hydrate_openai_response_history, persist_response_history_record, ResponseHistoryRecord,
};
use crate::{ai_serving::response_history_storage_key, data::GatewayDataState, AppState};
fn response_history_test_state() -> AppState {
AppState::new()
.expect("test state should build")
.with_data_state_for_tests(
GatewayDataState::disabled()
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_runtime_state(Arc::new(RuntimeState::memory(
MemoryRuntimeStateConfig::default(),
)))
}
fn response_history_payload(response_id: &str, scope: &str, marker: &str) -> String {
let expires_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
.saturating_add(3600);
json!({
"version": 1,
"response_id": response_id,
"scope_fingerprint": format!("{:x}", Sha256::digest(scope.trim().as_bytes())),
"expires_at_unix_secs": expires_at_unix_secs,
"transcript": [{"type": "message", "content": marker}],
})
.to_string()
}
#[tokio::test]
async fn response_history_is_encrypted_at_rest_and_hydrates() {
let state = response_history_test_state();
let response_id = "resp_gateway_encrypted_history_v1";
let scope = "response-history-encrypted-scope";
let marker = "private-response-history-marker";
let storage_key = response_history_storage_key(response_id, Some(scope));
let payload = response_history_payload(response_id, scope, marker);
persist_response_history_record(
&state,
ResponseHistoryRecord {
storage_key: storage_key.clone(),
payload,
ttl: Duration::from_secs(6 * 60 * 60),
},
)
.await;
let stored = state
.runtime_kv_get(&storage_key)
.await
.expect("history lookup should succeed")
.expect("history should be persisted");
assert!(crate::handlers::shared::runtime_secret_payload_is_sealed(
&stored
));
assert!(!stored.contains(marker));
hydrate_openai_response_history(
&state,
&json!({"previous_response_id": response_id}),
"openai:responses",
"openai:chat",
scope,
)
.await
.expect("encrypted history should hydrate");
assert!(crate::ai_serving::response_history_is_loaded(
response_id,
Some(scope)
));
}
#[tokio::test]
async fn response_history_reader_rejects_and_deletes_legacy_plaintext() {
let state = response_history_test_state();
let response_id = "resp_gateway_legacy_history_v1";
let scope = "response-history-legacy-scope";
let storage_key = response_history_storage_key(response_id, Some(scope));
let payload = response_history_payload(response_id, scope, "legacy-private-history");
state
.runtime_kv_setex(&storage_key, &payload, 6 * 60 * 60)
.await
.expect("legacy history should store");
let result = hydrate_openai_response_history(
&state,
&json!({"previous_response_id": response_id}),
"openai:responses",
"openai:chat",
scope,
)
.await;
assert!(result.is_err());
assert!(!crate::ai_serving::response_history_is_loaded(
response_id,
Some(scope)
));
assert!(state
.runtime_kv_get(&storage_key)
.await
.expect("history lookup should succeed")
.is_none());
}
}
@@ -1,7 +1,10 @@
use axum::routing::get;
use axum::Router;
use crate::{handlers::proxy::proxy_request, state::AppState};
use crate::{
handlers::{proxy::proxy_request, public::vscodex_ws_proxy},
state::AppState,
};
pub(crate) fn mount_public_support_routes(router: Router<AppState>) -> Router<AppState> {
router
@@ -26,6 +29,7 @@ pub(crate) fn mount_public_support_routes(router: Router<AppState>) -> Router<Ap
.route("/api/capabilities", get(proxy_request))
.route("/api/capabilities/user-configurable", get(proxy_request))
.route("/api/capabilities/model/{*model_path}", get(proxy_request))
.route("/api/vscodex/ws", get(vscodex_ws_proxy))
.route("/install/{*install_path}", get(proxy_request))
.route("/install-tunnel/{*install_path}", get(proxy_request))
.route("/i/{*install_path}", get(proxy_request))
+1 -1
View File
@@ -133,7 +133,7 @@ pub(crate) async fn frontdoor_manifest(State(state): State<AppState>) -> impl In
"internal_gateway": {
"route_groups": INTERNAL_GATEWAY_ROUTE_GROUPS,
"path_prefixes": INTERNAL_GATEWAY_PATH_PREFIXES,
"status": "rust_native_control_plane",
"status": state.internal_gateway_auth_status(),
},
},
"features": {
+254 -3
View File
@@ -1,5 +1,14 @@
use std::net::SocketAddr;
use axum::body::Body;
use axum::extract::{ConnectInfo, Request, State};
use axum::http::{self, HeaderValue, StatusCode};
use axum::middleware::{self, Next};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::Router;
use axum::{Json, Router};
use serde_json::json;
use tracing::warn;
use crate::async_task::{
cancel_video_task, get_video_task_detail, get_video_task_stats, get_video_task_video,
@@ -10,8 +19,18 @@ use crate::hooks::{get_request_audit_bundle, get_request_usage_audit};
use crate::router::metrics;
use crate::state::AppState;
pub(crate) fn mount_operational_routes(router: Router<AppState>) -> Router<AppState> {
router
#[derive(Clone, Copy)]
struct OperationalPermission {
required_permissions: &'static [&'static str],
write: bool,
requires_full_admin_role: bool,
}
pub(crate) fn mount_operational_routes(
router: Router<AppState>,
state: AppState,
) -> Router<AppState> {
let operational = Router::<AppState>::new()
.route("/_gateway/metrics", get(metrics))
.route("/_gateway/async-tasks/video-tasks", get(list_video_tasks))
.route(
@@ -50,4 +69,236 @@ pub(crate) fn mount_operational_routes(router: Router<AppState>) -> Router<AppSt
"/_gateway/audit/request-usage/{request_id}",
get(get_request_usage_audit),
)
.route_layer(middleware::from_fn_with_state(
state,
authorize_operational_request,
));
router.merge(operational)
}
async fn authorize_operational_request(
State(state): State<AppState>,
request: Request,
next: Next,
) -> Response<Body> {
let Some(permission) = operational_permission(request.method(), request.uri().path()) else {
return operational_error_response(
StatusCode::FORBIDDEN,
"operational route permission is not configured",
None,
);
};
let Some(remote_addr) = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|value| value.0)
else {
return operational_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"operational authentication unavailable",
None,
);
};
let headers = request.headers().clone();
let uri = request.uri().clone();
if headers.get_all(http::header::AUTHORIZATION).iter().count() > 1 {
return operational_auth_required_response();
}
match crate::control::resolve_local_admin_session_principal(&state, &headers, &uri).await {
Ok(Some(principal)) => {
if permission.requires_full_admin_role
&& !crate::roles::is_full_admin_role(&principal.user_role)
{
return operational_permission_denied_response(permission.required_permissions[0]);
}
if permission.write && !crate::roles::can_write_admin_console(&principal.user_role) {
return operational_permission_denied_response(permission.required_permissions[0]);
}
}
Ok(None) => {
let client_ip = crate::headers::effective_client_ip(&headers, &remote_addr);
let authenticated = match crate::management_token_auth::authenticate_management_token(
&state, &headers, client_ip,
)
.await
{
Ok(authenticated) => authenticated,
Err(
crate::management_token_auth::ManagementTokenAuthError::Missing
| crate::management_token_auth::ManagementTokenAuthError::Invalid,
) => return operational_auth_required_response(),
Err(crate::management_token_auth::ManagementTokenAuthError::Unavailable) => {
return operational_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"operational authentication unavailable",
None,
)
}
};
if permission.requires_full_admin_role
&& !crate::roles::is_full_admin_role(&authenticated.user.role)
{
return operational_permission_denied_response(permission.required_permissions[0]);
}
if permission.write && !crate::roles::can_write_admin_console(&authenticated.user.role)
{
return operational_permission_denied_response(permission.required_permissions[0]);
}
let missing_permission =
permission
.required_permissions
.iter()
.copied()
.find(|required| {
!management_token_has_operational_permission(
&authenticated.permissions,
required,
)
});
if let Some(required_permission) = missing_permission {
return operational_permission_denied_response(required_permission);
}
let client_ip = client_ip.to_string();
if let Err(err) = state
.record_management_token_usage(&authenticated.token.id, Some(client_ip.as_str()))
.await
{
warn!(
token_id = %authenticated.token.id,
error = ?err,
"gateway failed to record operational management token usage"
);
}
}
Err(err) => {
warn!(error = ?err, "operational admin session authentication failed");
return operational_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"operational authentication unavailable",
None,
);
}
}
let mut response = next.run(request).await;
response.headers_mut().insert(
http::header::CACHE_CONTROL,
HeaderValue::from_static("no-store"),
);
response
}
fn operational_permission(method: &http::Method, path: &str) -> Option<OperationalPermission> {
if path == "/_gateway/metrics" {
return Some(OperationalPermission {
required_permissions: &["admin:monitoring:read"],
write: false,
requires_full_admin_role: false,
});
}
if path.starts_with("/_gateway/async-tasks/video-tasks") {
let write = *method == http::Method::POST && path.ends_with("/cancel");
return Some(OperationalPermission {
required_permissions: if write {
&["admin:video_tasks:write"]
} else {
&["admin:video_tasks:read"]
},
write,
requires_full_admin_role: false,
});
}
if path.starts_with("/_gateway/audit/auth/users/") {
return Some(OperationalPermission {
required_permissions: &["admin:api_keys:read"],
write: false,
requires_full_admin_role: false,
});
}
if path.starts_with("/_gateway/audit/request-audit/") {
return Some(OperationalPermission {
required_permissions: &[
"admin:monitoring:admin",
"admin:usage:read",
"admin:api_keys:read",
],
write: false,
requires_full_admin_role: true,
});
}
if path.starts_with("/_gateway/audit/request-candidates/")
|| path.starts_with("/_gateway/audit/decision-trace/")
{
return Some(OperationalPermission {
required_permissions: &["admin:monitoring:admin"],
write: false,
requires_full_admin_role: true,
});
}
if path.starts_with("/_gateway/audit/") {
return Some(OperationalPermission {
required_permissions: &["admin:usage:read"],
write: false,
requires_full_admin_role: false,
});
}
None
}
fn management_token_has_operational_permission(
permissions: &[String],
required_permission: &str,
) -> bool {
let scope = required_permission
.rsplit_once(':')
.map(|(scope, _)| scope)
.unwrap_or(required_permission);
let admin_permission = format!("{scope}:admin");
permissions
.iter()
.any(|permission| permission == required_permission || permission == &admin_permission)
}
fn operational_auth_required_response() -> Response<Body> {
let mut response = operational_error_response(
StatusCode::UNAUTHORIZED,
"admin authentication required",
None,
);
response.headers_mut().insert(
http::header::WWW_AUTHENTICATE,
HeaderValue::from_static("Bearer"),
);
response
}
fn operational_permission_denied_response(required_permission: &'static str) -> Response<Body> {
operational_error_response(
StatusCode::FORBIDDEN,
"operational permission denied",
Some(required_permission),
)
}
fn operational_error_response(
status: StatusCode,
detail: &'static str,
required_permission: Option<&'static str>,
) -> Response<Body> {
let mut response = (
status,
Json(json!({
"detail": detail,
"required_permission": required_permission,
})),
)
.into_response();
response.headers_mut().insert(
http::header::CACHE_CONTROL,
HeaderValue::from_static("no-store"),
);
response
}
+301 -13
View File
@@ -11,6 +11,7 @@ use crate::constants::*;
use crate::control::GatewayControlDecision;
use crate::control::GatewayLocalAuthRejection;
use crate::headers::should_skip_response_header;
use crate::plan_usage_policy::PlanUsagePolicyRejection;
use crate::rate_limit::FrontdoorUserRpmRejection;
use crate::{insert_header_if_missing, GatewayError};
@@ -52,22 +53,34 @@ pub(crate) fn apply_streaming_response_headers(headers: &mut http::HeaderMap) {
);
}
fn apply_gateway_browser_security_headers(headers: &mut http::HeaderMap) {
// Provider responses are API data, even when an untrusted provider labels
// them as HTML or SVG. Keep a direct navigation to a gateway API route
// from becoming same-origin active content, and prevent referrer leakage
// if a user follows a link rendered from such a response.
headers.insert(
http::header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
headers.insert(
HeaderName::from_static("content-security-policy"),
HeaderValue::from_static(
"default-src 'none'; base-uri 'none'; form-action 'none'; frame-ancestors 'none'; sandbox",
),
);
headers.insert(
HeaderName::from_static("referrer-policy"),
HeaderValue::from_static("no-referrer"),
);
}
pub(crate) fn build_client_response(
upstream_response: reqwest::Response,
trace_id: &str,
control_decision: Option<&GatewayControlDecision>,
) -> Result<Response<Body>, GatewayError> {
let status = upstream_response.status();
let upstream_headers = upstream_response
.headers()
.iter()
.map(|(name, value)| {
(
name.as_str().to_string(),
value.to_str().unwrap_or_default().to_string(),
)
})
.collect::<BTreeMap<_, _>>();
let upstream_headers = collect_safe_response_headers(upstream_response.headers());
let upstream_stream = upstream_response.bytes_stream();
build_client_response_from_parts(
status.as_u16(),
@@ -78,6 +91,40 @@ pub(crate) fn build_client_response(
)
}
fn collect_safe_response_headers(headers: &http::HeaderMap) -> BTreeMap<String, String> {
let connection_declared = aether_http::connection_declared_header_names(
headers
.get_all(http::header::CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok()),
);
headers
.iter()
.filter_map(|(name, value)| {
let normalized = name.as_str().to_ascii_lowercase();
if should_skip_client_response_header(&normalized)
|| connection_declared.contains(&normalized)
{
return None;
}
value
.to_str()
.ok()
.map(|value| (normalized, value.to_string()))
})
.collect()
}
fn should_skip_client_response_header(name: &str) -> bool {
should_skip_response_header(name)
// A provider Location is relative to the provider, not to the gateway.
// Forwarding it lets redirect-following clients bypass the gateway and
// can disclose their gateway Authorization header to another origin.
// Keep Location available inside execution reports, but never expose
// it on the client-facing response boundary.
|| name.eq_ignore_ascii_case(http::header::LOCATION.as_str())
}
pub(crate) fn build_client_response_from_parts(
status_code: u16,
upstream_headers: &BTreeMap<String, String>,
@@ -111,8 +158,17 @@ where
.body(body)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let connection_declared = aether_http::connection_declared_header_names(
upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str()))
.map(|(_, value)| value.as_str()),
);
for (name, value) in upstream_headers {
if should_skip_response_header(name.as_str()) {
if should_skip_client_response_header(name.as_str())
|| connection_declared.contains(&name.to_ascii_lowercase())
{
continue;
}
let header_name = HeaderName::from_bytes(name.as_bytes())
@@ -123,6 +179,7 @@ where
}
mutate_headers(response.headers_mut())?;
apply_streaming_response_headers(response.headers_mut());
apply_gateway_browser_security_headers(response.headers_mut());
insert_header_if_missing(response.headers_mut(), TRACE_ID_HEADER, trace_id)?;
insert_header_if_missing(response.headers_mut(), GATEWAY_HEADER, "rust-phase3b")?;
if let Some(decision) = control_decision {
@@ -258,6 +315,57 @@ pub(crate) fn build_local_user_rpm_limited_response(
)
}
pub(crate) fn build_local_plan_usage_limited_response(
trace_id: &str,
control_decision: Option<&GatewayControlDecision>,
rejection: &PlanUsagePolicyRejection,
) -> Result<Response<Body>, GatewayError> {
let message = "套餐使用限制已达到上限,请稍后重试";
let fallback_payload = json!({
"error": {
"type": "plan_usage_limit_exceeded",
"message": message,
"details": {
"metric": rejection.metric,
"window": rejection.window,
"limit": rejection.limit,
"retry_after": rejection.retry_after,
}
}
});
let payload = build_local_error_payload(
control_decision,
None,
message,
LocalCoreSyncErrorKind::RateLimit,
fallback_payload,
);
let body =
serde_json::to_vec(&payload).map_err(|err| GatewayError::Internal(err.to_string()))?;
let headers = BTreeMap::from([
("content-type".to_string(), "application/json".to_string()),
("Retry-After".to_string(), rejection.retry_after.to_string()),
("X-RateLimit-Limit".to_string(), rejection.limit.to_string()),
("X-RateLimit-Remaining".to_string(), "0".to_string()),
("X-RateLimit-Scope".to_string(), "plan".to_string()),
(
"X-RateLimit-Metric".to_string(),
rejection.metric.to_string(),
),
(
"X-RateLimit-Window".to_string(),
rejection.window.to_string(),
),
]);
build_client_response_from_parts(
StatusCode::TOO_MANY_REQUESTS.as_u16(),
&headers,
Body::from(body),
trace_id,
control_decision,
)
}
pub(crate) fn build_local_http_error_response(
trace_id: &str,
control_decision: Option<&GatewayControlDecision>,
@@ -454,11 +562,13 @@ fn local_error_kind_for_status(status: StatusCode) -> LocalCoreSyncErrorKind {
#[cfg(test)]
mod tests {
use super::{
build_client_response_from_parts, build_local_auth_rejection_response,
build_client_response, build_client_response_from_parts,
build_client_response_from_parts_with_mutator, build_local_auth_rejection_response,
build_local_http_error_response_with_request_path, build_local_overloaded_response,
build_local_user_rpm_limited_response,
build_local_plan_usage_limited_response, build_local_user_rpm_limited_response,
};
use crate::control::{GatewayControlDecision, GatewayLocalAuthRejection};
use crate::plan_usage_policy::PlanUsagePolicyRejection;
use crate::rate_limit::FrontdoorUserRpmRejection;
use axum::body::{to_bytes, Body};
use std::collections::BTreeMap;
@@ -490,6 +600,163 @@ mod tests {
);
}
#[test]
fn upstream_security_headers_are_stripped_before_gateway_headers_are_added() {
let response = build_client_response_from_parts_with_mutator(
200,
&BTreeMap::from([
("set-cookie".to_string(), "session=attacker".to_string()),
(
"x-aether-gateway".to_string(),
"attacker-gateway".to_string(),
),
(
"x-aether-control-action".to_string(),
"attacker-action".to_string(),
),
(
"x-aether-future-control".to_string(),
"attacker-future".to_string(),
),
(
"x-accel-redirect".to_string(),
"/internal/private-file".to_string(),
),
("x-sendfile".to_string(), "/etc/passwd".to_string()),
(
"x-reproxy-url".to_string(),
"http://127.0.0.1:9000/private".to_string(),
),
(
"access-control-allow-origin".to_string(),
"https://attacker.example".to_string(),
),
(
"access-control-allow-credentials".to_string(),
"true".to_string(),
),
("content-length".to_string(), "999999".to_string()),
(
"content-security-policy".to_string(),
"default-src * 'unsafe-inline' 'unsafe-eval'".to_string(),
),
(
"content-security-policy-report-only".to_string(),
"default-src 'none'; report-uri https://attacker.example/csp".to_string(),
),
(
"reporting-endpoints".to_string(),
"attacker=\"https://attacker.example/reports\"".to_string(),
),
("report-to".to_string(), "attacker".to_string()),
(
"nel".to_string(),
"{\"report_to\":\"attacker\"}".to_string(),
),
(
"refresh".to_string(),
"0; url=https://attacker.example".to_string(),
),
("referrer-policy".to_string(), "unsafe-url".to_string()),
("x-content-type-options".to_string(), "invalid".to_string()),
(
"location".to_string(),
"https://provider.example/direct".to_string(),
),
("x-upstream-visible".to_string(), "ok".to_string()),
]),
Body::empty(),
"trace-upstream-header-filter",
None,
|headers| {
headers.insert(
http::HeaderName::from_static("x-aether-control-action"),
http::HeaderValue::from_static("gateway-action"),
);
Ok(())
},
)
.expect("response should build");
assert!(response.headers().get(http::header::SET_COOKIE).is_none());
assert!(response.headers().get("x-aether-future-control").is_none());
assert!(response.headers().get("x-accel-redirect").is_none());
assert!(response.headers().get("x-sendfile").is_none());
assert!(response.headers().get("x-reproxy-url").is_none());
assert!(response
.headers()
.get("access-control-allow-origin")
.is_none());
assert!(response
.headers()
.get("access-control-allow-credentials")
.is_none());
assert!(response
.headers()
.get(http::header::CONTENT_LENGTH)
.is_none());
assert!(response
.headers()
.get("content-security-policy-report-only")
.is_none());
assert!(response.headers().get("reporting-endpoints").is_none());
assert!(response.headers().get("report-to").is_none());
assert!(response.headers().get("nel").is_none());
assert!(response.headers().get("refresh").is_none());
assert!(response.headers().get(http::header::LOCATION).is_none());
assert_eq!(
response.headers()["content-security-policy"],
"default-src 'none'; base-uri 'none'; form-action 'none'; frame-ancestors 'none'; sandbox"
);
assert_eq!(response.headers()["referrer-policy"], "no-referrer");
assert_eq!(
response.headers()[http::header::X_CONTENT_TYPE_OPTIONS],
"nosniff"
);
assert_eq!(response.headers()["x-aether-gateway"], "rust-phase3b");
assert_eq!(
response.headers()["x-aether-control-action"],
"gateway-action"
);
assert_eq!(response.headers()["x-upstream-visible"], "ok");
}
#[tokio::test]
async fn raw_response_collector_honors_all_connection_header_lines() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener");
let addr = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("connection");
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
let mut request = [0_u8; 1024];
let _ = stream.read(&mut request).await.expect("request read");
stream
.write_all(
b"HTTP/1.1 200 OK\r\nConnection: x-first-hop\r\nConnection: x-second-hop\r\nX-First-Hop: first-secret\r\nX-Second-Hop: second-secret\r\nContent-Length: 2\r\n\r\nok",
)
.await
.expect("response write");
});
let upstream = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("client")
.get(format!("http://{addr}/"))
.send()
.await
.expect("upstream response");
let response = build_client_response(upstream, "trace-connection-lines", None)
.expect("client response");
server.await.expect("server");
assert!(response.headers().get("connection").is_none());
assert!(response.headers().get("x-first-hop").is_none());
assert!(response.headers().get("x-second-hop").is_none());
}
fn claude_decision() -> GatewayControlDecision {
GatewayControlDecision::synthetic(
"/v1/messages",
@@ -581,4 +848,25 @@ mod tests {
);
}
}
#[tokio::test]
async fn plan_usage_rejection_exposes_machine_readable_limit_headers() {
let response = build_local_plan_usage_limited_response(
"trace-plan-limit",
None,
&PlanUsagePolicyRejection {
metric: "request_count",
limit: 100.0,
retry_after: 42,
window: "calendar_week",
},
)
.expect("response");
assert_eq!(response.status(), http::StatusCode::TOO_MANY_REQUESTS);
assert_eq!(response.headers()["retry-after"], "42");
assert_eq!(response.headers()["x-ratelimit-scope"], "plan");
assert_eq!(response.headers()["x-ratelimit-window"], "calendar_week");
let payload = response_json(response).await;
assert_eq!(payload["error"]["type"], "plan_usage_limit_exceeded");
}
}
+336 -47
View File
@@ -1,3 +1,4 @@
use std::net::{IpAddr, SocketAddr};
use std::time::{SystemTime, UNIX_EPOCH};
use aether_contracts::ExecutionResult;
@@ -21,7 +22,9 @@ use super::{
};
use crate::{AppState, GatewayError};
pub(crate) use self::cancel::{cancel_video_task_record, CancelVideoTaskError};
pub(crate) use self::cancel::{
cancel_video_task_record, cancel_video_task_record_for_user, CancelVideoTaskError,
};
#[derive(Debug, Deserialize)]
pub(crate) struct ListVideoTasksQuery {
@@ -146,18 +149,21 @@ pub(crate) async fn get_video_task_video(
}
pub(crate) async fn build_video_task_video_response(
state: &AppState,
_state: &AppState,
task_id: &str,
source: VideoTaskVideoSource,
) -> Result<axum::response::Response, GatewayError> {
match source {
VideoTaskVideoSource::Redirect { url } => Ok(Redirect::temporary(&url).into_response()),
VideoTaskVideoSource::Redirect { url } => {
resolve_public_video_target(&url).await?;
Ok(Redirect::temporary(url.as_str()).into_response())
}
VideoTaskVideoSource::Proxy {
url,
header_name,
header_value,
filename,
} => proxy_video_stream(state, task_id, &url, &header_name, &header_value, &filename).await,
} => proxy_video_stream(task_id, &url, &header_name, &header_value, &filename).await,
}
}
@@ -209,25 +215,34 @@ fn video_task_status_name(status: VideoTaskStatus) -> &'static str {
}
async fn proxy_video_stream(
state: &AppState,
task_id: &str,
url: &str,
url: &url::Url,
header_name: &str,
header_value: &str,
filename: &str,
) -> Result<axum::response::Response, GatewayError> {
let response = state
.client
.get(url)
let target = resolve_public_video_target(url).await?;
let client = build_pinned_video_client(&target)?;
let response = client
.get(url.clone())
.header(header_name, header_value)
.send()
.await
.map_err(|err| GatewayError::UpstreamUnavailable {
trace_id: task_id.to_string(),
message: err.to_string(),
message: video_request_failure_message(&err).to_string(),
})?;
if response.status().is_client_error() || response.status().is_server_error() {
if response.status().is_redirection() {
return Err(GatewayError::UpstreamUnavailable {
trace_id: task_id.to_string(),
message: format!(
"video upstream redirect was rejected with HTTP {}",
response.status()
),
});
}
if !response.status().is_success() {
return Err(GatewayError::UpstreamUnavailable {
trace_id: task_id.to_string(),
message: format!("video upstream returned HTTP {}", response.status()),
@@ -235,47 +250,321 @@ async fn proxy_video_stream(
}
let status = response.status();
let content_type = response
.headers()
.get(axum::http::header::CONTENT_TYPE)
.cloned()
.unwrap_or_else(|| axum::http::HeaderValue::from_static("video/mp4"));
let content_length = response
.headers()
.get(axum::http::header::CONTENT_LENGTH)
.cloned();
let cache_control = response
.headers()
.get(axum::http::header::CACHE_CONTROL)
.cloned();
// Do not copy the provider's Content-Length onto a newly wrapped stream.
// Reqwest may decode transfer/content encodings and the provider controls
// the declaration; forwarding a stale value would make the client-facing
// HTTP framing disagree with the bytes produced by this Body. Axum/Hyper
// will select safe framing for the actual stream.
let upstream_headers = response.headers().clone();
let body = Body::from_stream(response.bytes_stream());
let mut outbound = axum::http::Response::builder()
.status(status)
.body(body)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
outbound
.headers_mut()
.insert(axum::http::header::CONTENT_TYPE, content_type);
outbound.headers_mut().insert(
axum::http::header::CONTENT_DISPOSITION,
axum::http::HeaderValue::from_str(&format!("inline; filename=\"{filename}\""))
.map_err(|err| GatewayError::Internal(err.to_string()))?,
);
if let Some(content_length) = content_length {
outbound
.headers_mut()
.insert(axum::http::header::CONTENT_LENGTH, content_length);
}
if let Some(cache_control) = cache_control {
outbound
.headers_mut()
.insert(axum::http::header::CACHE_CONTROL, cache_control);
} else {
outbound.headers_mut().insert(
axum::http::header::CACHE_CONTROL,
axum::http::HeaderValue::from_static("private, max-age=3600"),
);
}
apply_safe_video_response_metadata(outbound.headers_mut(), &upstream_headers, filename)?;
Ok(outbound)
}
fn apply_safe_video_response_metadata(
outbound: &mut axum::http::HeaderMap,
upstream: &axum::http::HeaderMap,
filename: &str,
) -> Result<(), GatewayError> {
let content_type = upstream
.get(axum::http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.and_then(safe_video_content_type)
.unwrap_or_else(|| axum::http::HeaderValue::from_static("application/octet-stream"));
outbound.insert(axum::http::header::CONTENT_TYPE, content_type);
outbound.insert(
axum::http::header::CONTENT_DISPOSITION,
axum::http::HeaderValue::from_str(&format!(
"inline; filename=\"{}\"",
safe_video_filename(filename)
))
.map_err(|err| GatewayError::Internal(err.to_string()))?,
);
outbound.remove(axum::http::header::CONTENT_LENGTH);
outbound.insert(
axum::http::header::CACHE_CONTROL,
axum::http::HeaderValue::from_static("private, no-store"),
);
outbound.insert(
axum::http::header::X_CONTENT_TYPE_OPTIONS,
axum::http::HeaderValue::from_static("nosniff"),
);
Ok(())
}
fn safe_video_content_type(raw_value: &str) -> Option<axum::http::HeaderValue> {
let media_type = raw_value.split(';').next()?.trim().to_ascii_lowercase();
let subtype = media_type.strip_prefix("video/")?;
if subtype.is_empty()
|| !subtype.bytes().all(|byte| {
byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'!' | b'#' | b'$' | b'&' | b'-' | b'^' | b'_' | b'.' | b'+'
)
})
{
return None;
}
axum::http::HeaderValue::from_str(raw_value).ok()
}
fn safe_video_filename(filename: &str) -> String {
let filename = filename
.chars()
.take(255)
.map(|character| {
if character.is_ascii_alphanumeric() || matches!(character, '.' | '-' | '_') {
character
} else {
'_'
}
})
.collect::<String>();
if filename.is_empty() {
"video.mp4".to_string()
} else {
filename
}
}
struct ResolvedVideoTarget {
host: String,
addrs: Vec<SocketAddr>,
}
async fn resolve_public_video_target(url: &url::Url) -> Result<ResolvedVideoTarget, GatewayError> {
if !matches!(url.scheme(), "http" | "https")
|| !url.username().is_empty()
|| url.password().is_some()
{
return Err(video_target_rejected(
"video URL must be an absolute HTTP(S) URL without credentials",
));
}
let port = url
.port_or_known_default()
.ok_or_else(|| video_target_rejected("video URL is missing a port"))?;
let (host, addrs) = match url.host() {
Some(url::Host::Ipv4(ip)) => (ip.to_string(), vec![SocketAddr::new(IpAddr::V4(ip), port)]),
Some(url::Host::Ipv6(ip)) => (ip.to_string(), vec![SocketAddr::new(IpAddr::V6(ip), port)]),
Some(url::Host::Domain(host)) if !host.is_empty() => {
let addrs = aether_http::lookup_host_with_limits(
host,
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| video_target_rejected("video URL DNS resolution failed"))?;
(host.to_string(), addrs)
}
_ => return Err(video_target_rejected("video URL is missing a host")),
};
if addrs.is_empty()
|| addrs
.iter()
.any(|addr| aether_http::is_private_or_reserved_ip(addr.ip()))
{
return Err(video_target_rejected(
"video URL resolves to a private or reserved address",
));
}
Ok(ResolvedVideoTarget { host, addrs })
}
fn build_pinned_video_client(
target: &ResolvedVideoTarget,
) -> Result<reqwest::Client, GatewayError> {
let mut builder = aether_http::apply_http_client_config(
reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none()),
&aether_http::HttpClientConfig {
connect_timeout_ms: Some(10_000),
request_timeout_ms: Some(300_000),
http2_adaptive_window: true,
..aether_http::HttpClientConfig::default()
},
);
if target.host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(&target.host, &target.addrs);
}
builder
.build()
.map_err(|_| GatewayError::Internal("video HTTP client initialization failed".to_string()))
}
fn video_target_rejected(message: &str) -> GatewayError {
GatewayError::Client {
status: axum::http::StatusCode::BAD_GATEWAY,
message: message.to_string(),
}
}
fn video_request_failure_message(error: &reqwest::Error) -> &'static str {
if error.is_timeout() {
"video upstream request timed out"
} else if error.is_connect() {
"video upstream connection failed"
} else if error.is_body() || error.is_decode() {
"video upstream response failed"
} else {
"video upstream request failed"
}
}
#[cfg(test)]
mod tests {
use axum::response::IntoResponse;
use super::{
apply_safe_video_response_metadata, build_video_task_video_response,
resolve_public_video_target, safe_video_content_type, safe_video_filename,
VideoTaskVideoSource,
};
use crate::AppState;
#[tokio::test]
async fn video_redirect_response_accepts_public_target() {
let state = AppState::new().expect("gateway state should build");
let target = "https://8.8.8.8/video.mp4";
let response = build_video_task_video_response(
&state,
"task-public-redirect",
VideoTaskVideoSource::Redirect {
url: url::Url::parse(target).expect("public target should parse"),
},
)
.await
.expect("public redirect should build");
assert_eq!(
response.status(),
axum::http::StatusCode::TEMPORARY_REDIRECT
);
assert_eq!(
response
.headers()
.get(axum::http::header::LOCATION)
.and_then(|value| value.to_str().ok()),
Some(target)
);
}
#[tokio::test]
async fn video_redirect_response_rejects_private_and_reserved_targets() {
let state = AppState::new().expect("gateway state should build");
for raw_url in [
"http://127.0.0.1/video.mp4",
"http://169.254.169.254/latest/meta-data",
"http://10.0.0.1/video.mp4",
"http://[::1]/video.mp4",
] {
let error = build_video_task_video_response(
&state,
"task-rejected-redirect",
VideoTaskVideoSource::Redirect {
url: url::Url::parse(raw_url).expect("target should parse"),
},
)
.await
.expect_err("private or reserved redirect target should be rejected");
assert_eq!(
error.into_response().status(),
axum::http::StatusCode::BAD_GATEWAY,
"unexpected status for {raw_url}"
);
}
}
#[tokio::test]
async fn video_target_resolution_rejects_private_and_reserved_ip_literals() {
for raw_url in [
"http://127.0.0.1/video.mp4",
"http://169.254.169.254/latest/meta-data",
"http://10.0.0.1/video.mp4",
"http://[::1]/video.mp4",
"http://[::ffff:127.0.0.1]/video.mp4",
] {
let url = url::Url::parse(raw_url).unwrap();
assert!(
resolve_public_video_target(&url).await.is_err(),
"target should be rejected: {raw_url}"
);
}
}
#[tokio::test]
async fn video_target_resolution_accepts_public_ip_literals() {
for raw_url in [
"https://8.8.8.8/video.mp4",
"https://[2606:4700:4700::1111]/video.mp4",
] {
let url = url::Url::parse(raw_url).unwrap();
assert!(
resolve_public_video_target(&url).await.is_ok(),
"target should be accepted: {raw_url}"
);
}
}
#[test]
fn video_response_metadata_rejects_active_content_and_sanitizes_filename() {
assert!(safe_video_content_type("video/mp4").is_some());
assert!(safe_video_content_type("video/webm; charset=binary").is_some());
assert!(safe_video_content_type("video/").is_none());
assert!(safe_video_content_type("video/; charset=binary").is_none());
assert!(safe_video_content_type("text/html").is_none());
assert!(safe_video_content_type("video/mp4\r\nx-test: injected").is_none());
assert_eq!(
safe_video_filename("video_123.mp4\"; filename=\"attack.html"),
"video_123.mp4___filename__attack.html"
);
assert_eq!(safe_video_filename(&"x".repeat(1024)).len(), 255);
let mut upstream = axum::http::HeaderMap::new();
upstream.insert(
axum::http::header::CONTENT_TYPE,
axum::http::HeaderValue::from_static("text/html"),
);
upstream.insert(
axum::http::header::CONTENT_LENGTH,
axum::http::HeaderValue::from_static("999999"),
);
let mut outbound = upstream.clone();
apply_safe_video_response_metadata(
&mut outbound,
&upstream,
"video.mp4\"; filename=\"attack.html",
)
.expect("video metadata should build");
assert_eq!(
outbound
.get(axum::http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok()),
Some("application/octet-stream")
);
assert!(outbound.get(axum::http::header::CONTENT_LENGTH).is_none());
assert_eq!(
outbound
.get(axum::http::header::X_CONTENT_TYPE_OPTIONS)
.and_then(|value| value.to_str().ok()),
Some("nosniff")
);
assert_eq!(
outbound
.get(axum::http::header::CONTENT_DISPOSITION)
.and_then(|value| value.to_str().ok()),
Some("inline; filename=\"video.mp4___filename__attack.html\"")
);
}
}
+190 -109
View File
@@ -3,12 +3,13 @@ use aether_data_contracts::repository::video_tasks::{
};
use axum::response::IntoResponse;
use axum::Json;
use serde_json::{json, Map, Value};
use serde_json::json;
use crate::state::VideoTaskRouteAccess;
use crate::{AppState, GatewayError};
use super::super::finalize_video_task_if_terminal;
use super::super::read_video_task_detail;
use super::super::{read_video_task_detail, read_video_task_detail_for_user};
use super::current_unix_secs;
#[derive(Debug)]
@@ -29,7 +30,31 @@ pub(crate) async fn cancel_video_task_record(
state: &AppState,
task_id: &str,
) -> Result<StoredVideoTask, CancelVideoTaskError> {
let Some(task) = read_video_task_detail(state, task_id).await? else {
cancel_video_task_record_inner(state, task_id, None).await
}
pub(crate) async fn cancel_video_task_record_for_user(
state: &AppState,
task_id: &str,
user_id: &str,
) -> Result<StoredVideoTask, CancelVideoTaskError> {
let user_id = user_id.trim();
if user_id.is_empty() {
return Err(CancelVideoTaskError::NotFound);
}
cancel_video_task_record_inner(state, task_id, Some(user_id)).await
}
async fn cancel_video_task_record_inner(
state: &AppState,
task_id: &str,
expected_user_id: Option<&str>,
) -> Result<StoredVideoTask, CancelVideoTaskError> {
let task = match expected_user_id {
Some(user_id) => read_video_task_detail_for_user(state, task_id, user_id).await?,
None => read_video_task_detail(state, task_id).await?,
};
let Some(task) = task else {
return Err(CancelVideoTaskError::NotFound);
};
@@ -45,39 +70,84 @@ pub(crate) async fn cancel_video_task_record(
}
let trace_id = format!("async-task-admin-cancel-{task_id}");
let mut finalize_mutation = None;
if let Some(cancel_plan) = build_video_task_cancel_plan(&task) {
state
.hydrate_video_task_for_route(Some(cancel_plan.route_family), &cancel_plan.request_path)
.await?;
let body_json = json!({});
let follow_up = state.video_tasks.prepare_follow_up_sync_plan(
cancel_plan.plan_kind,
&cancel_plan.request_path,
Some(&body_json),
None,
&trace_id,
);
let follow_up = if let Some(user_id) = expected_user_id {
if state
.hydrate_video_task_for_route_for_user(
Some(cancel_plan.route_family),
&cancel_plan.request_path,
user_id,
)
.await?
!= VideoTaskRouteAccess::Allowed
{
return Err(CancelVideoTaskError::NotFound);
}
state.video_tasks.prepare_follow_up_sync_plan_for_user_id(
cancel_plan.plan_kind,
&cancel_plan.request_path,
Some(&body_json),
user_id,
task.api_key_id.as_deref(),
&trace_id,
)
} else {
state
.hydrate_video_task_for_route(
Some(cancel_plan.route_family),
&cancel_plan.request_path,
)
.await?;
state.video_tasks.prepare_follow_up_sync_plan(
cancel_plan.plan_kind,
&cancel_plan.request_path,
Some(&body_json),
None,
&trace_id,
)
};
if let Some(follow_up) = follow_up {
execute_video_task_cancel_plan(state, &trace_id, follow_up.plan)
.await
.map_err(CancelVideoTaskError::Response)?;
finalize_mutation = Some((
cancel_plan.request_path,
cancel_plan.report_kind.to_string(),
));
} else if expected_user_id.is_none() {
finalize_mutation = Some((
cancel_plan.request_path,
cancel_plan.report_kind.to_string(),
));
}
state
.video_tasks
.apply_finalize_mutation(&cancel_plan.request_path, cancel_plan.report_kind);
}
let request_metadata = build_cancelled_request_metadata(state, &task).await?;
let stored = persist_cancelled_video_task(state, &task, request_metadata)
.await?
.ok_or_else(|| {
CancelVideoTaskError::Gateway(GatewayError::Internal(
let stored = match persist_cancelled_video_task(state, &task).await? {
Some(stored) => stored,
None => {
let current = match expected_user_id {
Some(user_id) => read_video_task_detail_for_user(state, task_id, user_id).await?,
None => read_video_task_detail(state, task_id).await?,
};
let Some(current) = current else {
return Err(CancelVideoTaskError::NotFound);
};
if !current.status.is_active() {
return Err(CancelVideoTaskError::InvalidStatus(current.status));
}
return Err(CancelVideoTaskError::Gateway(GatewayError::Internal(
"video task repository is unavailable".to_string(),
))
})?;
)));
}
};
if let Some((request_path, report_kind)) = finalize_mutation {
state
.video_tasks
.apply_finalize_mutation(&request_path, &report_kind);
}
finalize_video_task_if_terminal(state, &stored).await;
Ok(stored)
}
@@ -91,12 +161,7 @@ struct VideoTaskCancelPlan<'a> {
}
fn build_video_task_cancel_plan(task: &StoredVideoTask) -> Option<VideoTaskCancelPlan<'_>> {
let provider_api_format = task
.provider_api_format
.as_deref()
.or(task.client_api_format.as_deref())
.map(str::trim)
.filter(|value| !value.is_empty())?;
let provider_api_format = task.effective_api_format()?;
match provider_api_format {
"openai:video" => Some(VideoTaskCancelPlan {
@@ -131,99 +196,52 @@ async fn execute_video_task_cancel_plan(
let result =
crate::execution_runtime::execute_execution_runtime_sync_plan(state, Some(trace_id), &plan)
.await
.map_err(|err| {
.map_err(|_| {
GatewayError::UpstreamUnavailable {
trace_id: trace_id.to_string(),
message: format!("{err:?}"),
message: "video cancellation request failed".to_string(),
}
.into_response()
})?;
if result.status_code >= 400 {
let status = axum::http::StatusCode::from_u16(result.status_code)
.unwrap_or(axum::http::StatusCode::BAD_GATEWAY);
let body_json = result
.body
.and_then(|body| body.json_body)
.unwrap_or_else(|| {
json!({
"error": {
"message": result
.error
.as_ref()
.map(|error| error.message.clone())
.unwrap_or_else(|| {
format!("execution runtime returned {}", result.status_code)
}),
}
})
});
return Err((status, Json(body_json)).into_response());
return Err(build_video_task_cancel_upstream_error_response(&result));
}
Ok(())
}
async fn build_cancelled_request_metadata(
state: &AppState,
task: &StoredVideoTask,
) -> Result<Option<Value>, GatewayError> {
let mut metadata = match task.request_metadata.clone() {
Some(Value::Object(object)) => object,
_ => Map::new(),
};
let mut snapshot_value = metadata.get("rust_local_snapshot").cloned();
if snapshot_value.is_none() {
snapshot_value = state
.reconstruct_video_task_snapshot(task)
.await?
.map(|snapshot| {
serde_json::to_value(snapshot)
.map_err(|err| GatewayError::Internal(err.to_string()))
})
.transpose()?;
}
if let Some(snapshot_value_ref) = snapshot_value.as_mut() {
mark_snapshot_value_cancelled(snapshot_value_ref);
metadata.insert(
"rust_owner".to_string(),
Value::String("async_task".to_string()),
);
metadata.insert(
"rust_local_snapshot".to_string(),
snapshot_value_ref.clone(),
);
return Ok(Some(Value::Object(metadata)));
}
Ok(task.request_metadata.clone())
}
fn mark_snapshot_value_cancelled(snapshot_value: &mut Value) {
if let Some(object) = snapshot_value
.get_mut("OpenAi")
.and_then(Value::as_object_mut)
{
object.insert("status".to_string(), Value::String("Cancelled".to_string()));
return;
}
if let Some(object) = snapshot_value
.get_mut("Gemini")
.and_then(Value::as_object_mut)
{
object.insert("status".to_string(), Value::String("Cancelled".to_string()));
}
fn build_video_task_cancel_upstream_error_response(
result: &aether_contracts::ExecutionResult,
) -> axum::response::Response {
let status = axum::http::StatusCode::from_u16(result.status_code)
.unwrap_or(axum::http::StatusCode::BAD_GATEWAY);
tracing::warn!(
event_name = "video_task_cancel_upstream_error",
upstream_status = result.status_code,
"video cancellation upstream response body discarded"
);
(
status,
Json(json!({
"error": {
"message": format!(
"video cancellation upstream returned HTTP {}",
result.status_code
),
}
})),
)
.into_response()
}
async fn persist_cancelled_video_task(
state: &AppState,
task: &StoredVideoTask,
request_metadata: Option<Value>,
) -> Result<Option<StoredVideoTask>, GatewayError> {
let now_unix_secs = current_unix_secs();
state
.data
.upsert_video_task(UpsertVideoTask {
.update_active_video_task(UpsertVideoTask {
id: task.id.clone(),
short_id: task.short_id.clone(),
request_id: task.request_id.clone(),
@@ -240,14 +258,14 @@ async fn persist_cancelled_video_task(
format_converted: task.format_converted,
model: task.model.clone(),
prompt: task.prompt.clone(),
original_request_body: task.original_request_body.clone(),
original_request_body: None,
duration_seconds: task.duration_seconds,
resolution: task.resolution.clone(),
aspect_ratio: task.aspect_ratio.clone(),
size: task.size.clone(),
status: VideoTaskStatus::Cancelled,
progress_percent: task.progress_percent,
progress_message: task.progress_message.clone(),
progress_message: None,
retry_count: task.retry_count,
poll_interval_seconds: task.poll_interval_seconds,
next_poll_at_unix_secs: None,
@@ -258,10 +276,73 @@ async fn persist_cancelled_video_task(
completed_at_unix_secs: Some(now_unix_secs),
updated_at_unix_secs: now_unix_secs,
error_code: task.error_code.clone(),
error_message: task.error_message.clone(),
error_message: None,
video_url: task.video_url.clone(),
request_metadata,
request_metadata: None,
})
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use aether_contracts::{
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionResult, ResponseBody,
};
use axum::body::to_bytes;
use serde_json::json;
use super::build_video_task_cancel_upstream_error_response;
#[tokio::test]
async fn cancellation_upstream_errors_do_not_expose_runtime_payloads() {
let result = ExecutionResult {
request_id: "cancel-secret-request-id".to_string(),
candidate_id: Some("cancel-secret-candidate-id".to_string()),
status_code: 502,
headers: BTreeMap::from([(
"x-internal-secret".to_string(),
"cancel-secret-header".to_string(),
)]),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(json!({
"error": {
"message": "cancel-secret-upstream-body",
}
})),
body_bytes_b64: None,
}),
telemetry: None,
error: Some(ExecutionError {
kind: ExecutionErrorKind::Upstream5xx,
phase: ExecutionPhase::FirstByte,
message: "cancel-secret-runtime-error".to_string(),
upstream_status: Some(502),
retryable: true,
failover_recommended: false,
}),
};
let response = build_video_task_cancel_upstream_error_response(&result);
assert_eq!(response.status(), axum::http::StatusCode::BAD_GATEWAY);
assert!(response.headers().get("x-internal-secret").is_none());
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body should read");
let payload: serde_json::Value =
serde_json::from_slice(&body).expect("response body should parse");
assert_eq!(
payload,
json!({
"error": {
"message": "video cancellation upstream returned HTTP 502",
}
})
);
let body = String::from_utf8(body.to_vec()).expect("response body should be utf-8");
assert!(!body.contains("cancel-secret"));
}
}
+6 -5
View File
@@ -6,13 +6,14 @@ pub(crate) use crate::video_tasks::VideoTaskService;
pub use crate::video_tasks::VideoTaskTruthSourceMode;
pub(crate) use http::{
build_video_task_video_response, cancel_video_task, cancel_video_task_record,
get_video_task_detail, get_video_task_stats, get_video_task_video, list_video_tasks,
CancelVideoTaskError,
cancel_video_task_record_for_user, get_video_task_detail, get_video_task_stats,
get_video_task_video, list_video_tasks, CancelVideoTaskError,
};
pub(crate) use query::{
read_video_task_detail, read_video_task_page, read_video_task_page_summary,
read_video_task_stats, read_video_task_video_source, VideoTaskPageResponse,
VideoTaskStatsResponse, VideoTaskVideoSource,
read_video_task_detail, read_video_task_detail_for_user, read_video_task_page,
read_video_task_page_summary, read_video_task_stats, read_video_task_video_source,
video_task_video_source_from_task, VideoTaskPageResponse, VideoTaskStatsResponse,
VideoTaskVideoSource,
};
pub(crate) use runtime::{
execute_video_task_refresh_plan, finalize_video_task_if_terminal, spawn_video_task_poller,
+263 -5
View File
@@ -26,13 +26,12 @@ pub(crate) struct VideoTaskStatsResponse {
pub(crate) processing_count: u64,
}
#[derive(Debug, Clone)]
pub(crate) enum VideoTaskVideoSource {
Redirect {
url: String,
url: url::Url,
},
Proxy {
url: String,
url: url::Url,
header_name: String,
header_value: String,
filename: String,
@@ -102,6 +101,14 @@ pub(crate) async fn read_video_task_detail(
state.find_video_task_by_id(task_id).await
}
pub(crate) async fn read_video_task_detail_for_user(
state: &AppState,
task_id: &str,
user_id: &str,
) -> Result<Option<StoredVideoTask>, GatewayError> {
state.find_video_task_by_id_for_user(task_id, user_id).await
}
pub(crate) async fn read_video_task_video_source(
state: &AppState,
task_id: &str,
@@ -109,6 +116,13 @@ pub(crate) async fn read_video_task_video_source(
let Some(task) = read_video_task_detail(state, task_id).await? else {
return Ok(None);
};
video_task_video_source_from_task(state, &task).await
}
pub(crate) async fn video_task_video_source_from_task(
state: &AppState,
task: &StoredVideoTask,
) -> Result<Option<VideoTaskVideoSource>, GatewayError> {
let Some(video_url) = task
.video_url
.as_deref()
@@ -119,7 +133,9 @@ pub(crate) async fn read_video_task_video_source(
return Ok(None);
};
if !video_url.contains("generativelanguage.googleapis.com") {
let video_url = parse_video_url(&video_url)?;
if task.effective_api_format() != Some("gemini:video") {
return Ok(Some(VideoTaskVideoSource::Redirect { url: video_url }));
}
@@ -148,6 +164,15 @@ pub(crate) async fn read_video_task_video_source(
));
};
let endpoint_url = parse_video_url(transport.endpoint.base_url.trim()).map_err(|_| {
GatewayError::Internal("provider endpoint URL is invalid for proxied video".to_string())
})?;
if !video_urls_share_origin(&endpoint_url, &video_url) {
return Err(GatewayError::Client {
status: axum::http::StatusCode::BAD_GATEWAY,
message: "video URL origin does not match its provider endpoint".to_string(),
});
}
let api_key = transport.key.decrypted_api_key.trim();
if api_key.is_empty() {
return Err(GatewayError::Internal(
@@ -159,10 +184,34 @@ pub(crate) async fn read_video_task_video_source(
url: video_url,
header_name: "x-goog-api-key".to_string(),
header_value: api_key.to_string(),
filename: format!("video_{task_id}.mp4"),
filename: format!("video_{}.mp4", task.id),
}))
}
fn parse_video_url(raw_url: &str) -> Result<url::Url, GatewayError> {
let url = url::Url::parse(raw_url.trim()).map_err(|_| GatewayError::Client {
status: axum::http::StatusCode::BAD_GATEWAY,
message: "video URL is invalid".to_string(),
})?;
if !matches!(url.scheme(), "http" | "https")
|| url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
{
return Err(GatewayError::Client {
status: axum::http::StatusCode::BAD_GATEWAY,
message: "video URL must be an absolute HTTP(S) URL without credentials".to_string(),
});
}
Ok(url)
}
fn video_urls_share_origin(left: &url::Url, right: &url::Url) -> bool {
left.scheme() == right.scheme()
&& left.host() == right.host()
&& left.port_or_known_default() == right.port_or_known_default()
}
pub(crate) async fn read_video_task_stats(
state: &AppState,
filter: &VideoTaskQueryFilter,
@@ -226,3 +275,212 @@ fn status_key(status: VideoTaskStatus) -> String {
fn start_of_utc_day(now_unix_secs: u64) -> u64 {
now_unix_secs - (now_unix_secs % 86_400)
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::video_tasks::InMemoryVideoTaskRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::video_tasks::{UpsertVideoTask, VideoTaskStatus};
use serde_json::json;
use super::{
parse_video_url, video_task_video_source_from_task, video_urls_share_origin,
VideoTaskVideoSource,
};
use crate::{data::GatewayDataState, AppState};
fn legacy_gemini_video_task() -> aether_data_contracts::repository::video_tasks::StoredVideoTask
{
UpsertVideoTask {
id: "legacy-gemini-task".to_string(),
short_id: Some("legacy-short".to_string()),
request_id: "legacy-request".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("client-key-1".to_string()),
username: None,
api_key_name: None,
external_task_id: Some("operations/upstream-1".to_string()),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("provider-key-1".to_string()),
client_api_format: Some("gemini:video".to_string()),
provider_api_format: None,
format_converted: false,
model: Some("veo-3".to_string()),
prompt: None,
original_request_body: None,
duration_seconds: Some(8),
resolution: Some("720p".to_string()),
aspect_ratio: Some("16:9".to_string()),
size: Some("1280x720".to_string()),
status: VideoTaskStatus::Completed,
progress_percent: 100,
progress_message: None,
retry_count: 0,
poll_interval_seconds: 10,
next_poll_at_unix_secs: None,
poll_count: 1,
max_poll_count: 360,
created_at_unix_ms: 1,
submitted_at_unix_secs: Some(1),
completed_at_unix_secs: Some(2),
updated_at_unix_secs: 2,
error_code: None,
error_message: None,
video_url: Some(
"https://generativelanguage.googleapis.com/v1beta/files/video-1:download?alt=media"
.to_string(),
),
request_metadata: None,
}
.into_stored()
}
fn state_with_gemini_transport() -> AppState {
let state = AppState::new().expect("gateway state should build");
let provider = StoredProviderCatalogProvider::new(
"provider-1".to_string(),
"Gemini".to_string(),
Some("https://ai.google.dev".to_string()),
"gemini".to_string(),
)
.expect("provider should build");
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-1".to_string(),
"provider-1".to_string(),
"gemini:video".to_string(),
None,
None,
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://generativelanguage.googleapis.com".to_string(),
None,
None,
None,
None,
None,
None,
None,
)
.expect("endpoint transport should build");
let encrypted_api_key = state
.seal_provider_catalog_key_api_key(
"provider-1",
"provider-key-1",
"gemini-provider-secret",
)
.expect("provider key should encrypt");
let key = StoredProviderCatalogKey::new(
"provider-key-1".to_string(),
"provider-1".to_string(),
"default".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("provider key should build")
.with_transport_fields(
Some(json!(["gemini:video"])),
encrypted_api_key,
None,
None,
None,
None,
None,
None,
None,
)
.expect("provider key transport should build");
let provider_catalog: Arc<dyn ProviderCatalogReadRepository> = Arc::new(
InMemoryProviderCatalogReadRepository::seed(vec![provider], vec![endpoint], vec![key]),
);
let video_tasks = Arc::new(InMemoryVideoTaskRepository::default());
let data = GatewayDataState::with_video_task_repository_and_provider_transport_for_tests(
video_tasks,
provider_catalog,
DEVELOPMENT_ENCRYPTION_KEY,
);
state.with_data_state_for_tests(data)
}
#[test]
fn video_url_parser_rejects_non_http_and_embedded_credentials() {
for raw_url in [
"file:///etc/passwd",
"data:video/mp4;base64,AAAA",
"https://[email protected]/video.mp4",
"https://user:[email protected]/video.mp4",
"/relative/video.mp4",
] {
assert!(
parse_video_url(raw_url).is_err(),
"URL should be rejected: {raw_url}"
);
}
}
#[test]
fn video_origin_comparison_uses_scheme_host_and_effective_port() {
let base = parse_video_url("https://generativelanguage.googleapis.com/v1beta").unwrap();
for same_origin in [
"https://generativelanguage.googleapis.com/file",
"https://generativelanguage.googleapis.com:443/file",
] {
assert!(video_urls_share_origin(
&base,
&parse_video_url(same_origin).unwrap()
));
}
for different_origin in [
"http://generativelanguage.googleapis.com/file",
"https://generativelanguage.googleapis.com:444/file",
"https://generativelanguage.googleapis.com.evil.test/file",
"https://evil.test/generativelanguage.googleapis.com/file",
] {
assert!(!video_urls_share_origin(
&base,
&parse_video_url(different_origin).unwrap()
));
}
}
#[tokio::test]
async fn legacy_gemini_client_format_uses_authenticated_proxy_source() {
let source = video_task_video_source_from_task(
&state_with_gemini_transport(),
&legacy_gemini_video_task(),
)
.await
.expect("video source should resolve")
.expect("video source should exist");
match source {
VideoTaskVideoSource::Proxy {
url,
header_name,
header_value,
filename,
} => {
assert_eq!(
url.as_str(),
"https://generativelanguage.googleapis.com/v1beta/files/video-1:download?alt=media"
);
assert_eq!(header_name, "x-goog-api-key");
assert_eq!(header_value, "gemini-provider-secret");
assert_eq!(filename, "video_legacy-gemini-task.mp4");
}
VideoTaskVideoSource::Redirect { .. } => {
panic!("legacy Gemini video must not bypass the authenticated proxy")
}
}
}
}
+79 -164
View File
@@ -20,7 +20,7 @@ const VIDEO_TASK_POLL_CLAIM_SECONDS: u64 = 30;
#[derive(Debug, Clone)]
struct VideoTaskRefreshError {
message: String,
category: &'static str,
permanent: bool,
}
@@ -55,7 +55,7 @@ pub(crate) async fn execute_video_task_refresh_plan(
warn!(
event_name = "video_task_refresh_failed",
log_type = "event",
error = %err.message,
error_category = err.category,
permanent = err.permanent,
"gateway video task refresh failed"
);
@@ -79,23 +79,32 @@ async fn poll_video_tasks_once(state: &AppState, batch_size: usize) -> Result<us
let mut refreshed = 0usize;
for (index, task) in tasks.into_iter().enumerate() {
let trace_id = format!("video-task-poller-{index}");
let Some(snapshot) = state.reconstruct_video_task_snapshot(&task).await? else {
continue;
};
let Some(refresh_plan) = state
.video_tasks
.prepare_poll_refresh_plan_for_stored_task(&task, &trace_id)
.prepare_poll_refresh_plan_for_snapshot(snapshot.clone(), &trace_id)
else {
continue;
};
match fetch_video_task_refresh_attempt(state, &refresh_plan).await? {
VideoTaskRefreshAttempt::Success { provider_body } => {
let Some(updated) =
build_successful_poll_update(&task, &provider_body, now_unix_secs)?
let Some(updated) = build_successful_poll_update(
&task,
snapshot.clone(),
&provider_body,
now_unix_secs,
)?
else {
continue;
};
match state.update_active_video_task(updated).await? {
Some(stored) => {
if let Some(snapshot) = LocalVideoTaskSnapshot::from_stored_task(&stored) {
if let Some(snapshot) =
state.reconstruct_video_task_snapshot(&stored).await?
{
state.video_tasks.record_snapshot(snapshot);
}
info!(
@@ -116,7 +125,9 @@ async fn poll_video_tasks_once(state: &AppState, batch_size: usize) -> Result<us
let updated = build_failed_poll_update(&task, &err, now_unix_secs);
match state.update_active_video_task(updated).await? {
Some(stored) => {
if let Some(snapshot) = LocalVideoTaskSnapshot::from_stored_task(&stored) {
if let Some(snapshot) =
state.reconstruct_video_task_snapshot(&stored).await?
{
state.video_tasks.record_snapshot(snapshot);
}
info!(
@@ -190,9 +201,9 @@ async fn fetch_video_task_refresh_attempt(
.await
{
Ok(result) => result,
Err(err) => {
Err(_) => {
return Ok(VideoTaskRefreshAttempt::Error(VideoTaskRefreshError {
message: format!("{err:?}"),
category: "transport_error",
permanent: false,
}));
}
@@ -209,7 +220,7 @@ async fn fetch_video_task_refresh_attempt(
.and_then(|body| body.as_object().cloned())
else {
return Ok(VideoTaskRefreshAttempt::Error(VideoTaskRefreshError {
message: "video task refresh missing json provider body".to_string(),
category: "invalid_provider_response",
permanent: false,
}));
};
@@ -223,20 +234,19 @@ fn classify_refresh_result_error(result: &ExecutionResult) -> VideoTaskRefreshEr
.as_ref()
.and_then(|error| error.upstream_status)
.unwrap_or(result.status_code);
let message = result
.error
.as_ref()
.map(|error| error.message.clone())
.or_else(|| {
result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(|value| value.get("error"))
.and_then(Value::as_str)
.map(str::to_string)
})
.unwrap_or_else(|| format!("upstream returned {status_code}"));
let category = if status_code == 401 {
"authentication_error"
} else if status_code == 403 {
"permission_denied"
} else if status_code == 404 {
"not_found"
} else if status_code == 429 {
"rate_limit"
} else if status_code >= 500 {
"server_error"
} else {
"provider_error"
};
let permanent = result.error.as_ref().map_or(
matches!(status_code, 400 | 401 | 403 | 404 | 422),
|error| match error.kind {
@@ -253,17 +263,18 @@ fn classify_refresh_result_error(result: &ExecutionResult) -> VideoTaskRefreshEr
},
);
VideoTaskRefreshError { message, permanent }
VideoTaskRefreshError {
category,
permanent,
}
}
fn build_successful_poll_update(
task: &StoredVideoTask,
mut snapshot: LocalVideoTaskSnapshot,
provider_body: &Map<String, Value>,
now_unix_secs: u64,
) -> Result<Option<UpsertVideoTask>, GatewayError> {
let Some(mut snapshot) = LocalVideoTaskSnapshot::from_stored_task(task) else {
return Ok(None);
};
snapshot.apply_provider_body(provider_body);
let mut record = snapshot.to_upsert_record();
@@ -283,10 +294,7 @@ fn build_successful_poll_update(
record.format_converted = task.format_converted;
record.model = task.model.clone().or(record.model);
record.prompt = task.prompt.clone().or(record.prompt);
record.original_request_body = task
.original_request_body
.clone()
.or(record.original_request_body);
record.original_request_body = None;
record.duration_seconds = task.duration_seconds.or(record.duration_seconds);
record.resolution = task.resolution.clone().or(record.resolution);
record.aspect_ratio = task.aspect_ratio.clone().or(record.aspect_ratio);
@@ -309,17 +317,11 @@ fn build_successful_poll_update(
if record.status.is_active() && record.poll_count >= record.max_poll_count {
record.status = VideoTaskStatus::Failed;
record.error_code = Some("poll_timeout".to_string());
record.error_message = Some(format!("Task timed out after {} polls", record.poll_count));
record.error_message = None;
record.completed_at_unix_secs = Some(now_unix_secs);
record.next_poll_at_unix_secs = None;
}
record.request_metadata = merge_video_task_request_metadata(
task.request_metadata.clone(),
&snapshot,
Some(provider_body),
None,
)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
record.request_metadata = None;
Ok(Some(record))
}
@@ -332,11 +334,11 @@ fn build_failed_poll_update(
let mut record = stored_task_to_upsert(task);
record.updated_at_unix_secs = now_unix_secs;
record.poll_count = task.poll_count.saturating_add(1);
record.progress_message = Some(format!("Poll error: {}", err.message));
record.progress_message = None;
if err.permanent {
record.status = VideoTaskStatus::Failed;
record.error_code = Some("poll_permanent_error".to_string());
record.error_message = Some(err.message.clone());
record.error_message = None;
record.completed_at_unix_secs = Some(now_unix_secs);
record.next_poll_at_unix_secs = None;
} else {
@@ -348,28 +350,15 @@ fn build_failed_poll_update(
if record.status.is_active() && record.poll_count >= record.max_poll_count {
record.status = VideoTaskStatus::Failed;
record.error_code = Some("poll_timeout".to_string());
record.error_message = Some(format!("Task timed out after {} polls", record.poll_count));
record.error_message = None;
record.completed_at_unix_secs = Some(now_unix_secs);
record.next_poll_at_unix_secs = None;
}
record.request_metadata = LocalVideoTaskSnapshot::from_stored_task(task)
.and_then(|snapshot| {
merge_video_task_request_metadata(
task.request_metadata.clone(),
&snapshot,
None,
Some(err),
)
.ok()
.flatten()
})
.or(task.request_metadata.clone());
record.request_metadata = None;
record
}
fn stored_task_to_upsert(task: &StoredVideoTask) -> UpsertVideoTask {
let snapshot_record =
LocalVideoTaskSnapshot::from_stored_task(task).map(|snapshot| snapshot.to_upsert_record());
UpsertVideoTask {
id: task.id.clone(),
short_id: task.short_id.clone(),
@@ -386,39 +375,15 @@ fn stored_task_to_upsert(task: &StoredVideoTask) -> UpsertVideoTask {
provider_api_format: task.provider_api_format.clone(),
format_converted: task.format_converted,
model: task.model.clone(),
prompt: task.prompt.clone().or_else(|| {
snapshot_record
.as_ref()
.and_then(|record| record.prompt.clone())
}),
original_request_body: task.original_request_body.clone().or_else(|| {
snapshot_record
.as_ref()
.and_then(|record| record.original_request_body.clone())
}),
duration_seconds: task.duration_seconds.or_else(|| {
snapshot_record
.as_ref()
.and_then(|record| record.duration_seconds)
}),
resolution: task.resolution.clone().or_else(|| {
snapshot_record
.as_ref()
.and_then(|record| record.resolution.clone())
}),
aspect_ratio: task.aspect_ratio.clone().or_else(|| {
snapshot_record
.as_ref()
.and_then(|record| record.aspect_ratio.clone())
}),
size: task.size.clone().or_else(|| {
snapshot_record
.as_ref()
.and_then(|record| record.size.clone())
}),
prompt: task.prompt.clone(),
original_request_body: None,
duration_seconds: task.duration_seconds,
resolution: task.resolution.clone(),
aspect_ratio: task.aspect_ratio.clone(),
size: task.size.clone(),
status: task.status,
progress_percent: task.progress_percent,
progress_message: task.progress_message.clone(),
progress_message: None,
retry_count: task.retry_count,
poll_interval_seconds: task.poll_interval_seconds.max(1),
next_poll_at_unix_secs: task.next_poll_at_unix_secs,
@@ -429,9 +394,9 @@ fn stored_task_to_upsert(task: &StoredVideoTask) -> UpsertVideoTask {
completed_at_unix_secs: task.completed_at_unix_secs,
updated_at_unix_secs: task.updated_at_unix_secs,
error_code: task.error_code.clone(),
error_message: task.error_message.clone(),
error_message: None,
video_url: task.video_url.clone(),
request_metadata: task.request_metadata.clone(),
request_metadata: None,
}
}
@@ -443,44 +408,6 @@ fn compute_poll_backoff_seconds(poll_interval_seconds: u32, retry_count: u32) ->
.min(MAX_VIDEO_TASK_POLL_BACKOFF_SECONDS)
}
fn merge_video_task_request_metadata(
existing: Option<Value>,
snapshot: &LocalVideoTaskSnapshot,
provider_body: Option<&Map<String, Value>>,
poll_error: Option<&VideoTaskRefreshError>,
) -> Result<Option<Value>, serde_json::Error> {
let mut metadata = match existing {
Some(Value::Object(object)) => object,
_ => Map::new(),
};
metadata.insert(
"rust_owner".to_string(),
Value::String("async_task".to_string()),
);
metadata.insert(
"rust_local_snapshot".to_string(),
serde_json::to_value(snapshot)?,
);
if let Some(provider_body) = provider_body {
metadata.insert(
"poll_raw_response".to_string(),
Value::Object(provider_body.clone()),
);
metadata.remove("poll_error");
}
if let Some(poll_error) = poll_error {
metadata.insert(
"poll_error".to_string(),
serde_json::json!({
"message": poll_error.message,
"permanent": poll_error.permanent,
"observed_at_unix_secs": now_unix_secs(),
}),
);
}
Ok(Some(Value::Object(metadata)))
}
pub(crate) async fn finalize_video_task_if_terminal(state: &AppState, task: &StoredVideoTask) {
let Some(event) = build_video_task_terminal_usage_event(task) else {
return;
@@ -543,9 +470,9 @@ fn build_video_task_terminal_usage_event(task: &StoredVideoTask) -> Option<Usage
return None;
}
};
let provider_name = LocalVideoTaskSnapshot::from_stored_task(task)
.and_then(|snapshot| snapshot.provider_name().map(str::to_string))
.or_else(|| task.provider_id.clone())
let provider_name = task
.provider_id
.clone()
.unwrap_or_else(|| "unknown".to_string());
let response_time_ms = task
.submitted_at_unix_secs
@@ -580,10 +507,10 @@ fn build_video_task_terminal_usage_event(task: &StoredVideoTask) -> Option<Usage
has_format_conversion: Some(task.format_converted),
is_stream: Some(false),
status_code,
error_message: task.error_message.clone().or(task.error_code.clone()),
error_message: task.error_code.clone(),
response_time_ms,
request_body: task.original_request_body.clone(),
request_metadata: task.request_metadata.clone(),
request_body: None,
request_metadata: None,
..UsageEventData::default()
},
))
@@ -701,48 +628,36 @@ mod tests {
}
#[test]
fn stored_task_to_upsert_restores_sparse_fields_from_snapshot() {
fn stored_task_to_upsert_does_not_restore_sensitive_legacy_snapshot_fields() {
let record = stored_task_to_upsert(&sample_sparse_stored_task());
assert_eq!(record.prompt.as_deref(), Some("hello"));
assert_eq!(
record.original_request_body,
Some(json!({
"prompt": "hello",
"seconds": "4",
"resolution": "720p",
"aspect_ratio": "16:9",
"size": "1280x720"
}))
);
assert_eq!(record.duration_seconds, Some(4));
assert_eq!(record.resolution.as_deref(), Some("720p"));
assert_eq!(record.aspect_ratio.as_deref(), Some("16:9"));
assert_eq!(record.size.as_deref(), Some("1280x720"));
assert!(record.prompt.is_none());
assert!(record.original_request_body.is_none());
assert!(record.duration_seconds.is_none());
assert!(record.resolution.is_none());
assert!(record.aspect_ratio.is_none());
assert!(record.size.is_none());
assert!(record.progress_message.is_none());
assert!(record.error_message.is_none());
assert!(record.request_metadata.is_none());
}
#[test]
fn failed_poll_update_keeps_snapshot_backed_request_body() {
fn failed_poll_update_drops_snapshot_backed_sensitive_fields() {
let record = build_failed_poll_update(
&sample_sparse_stored_task(),
&VideoTaskRefreshError {
message: "temporary failure".to_string(),
category: "transport_error",
permanent: false,
},
100,
);
assert_eq!(
record.original_request_body,
Some(json!({
"prompt": "hello",
"seconds": "4",
"resolution": "720p",
"aspect_ratio": "16:9",
"size": "1280x720"
}))
);
assert_eq!(record.prompt.as_deref(), Some("hello"));
assert_eq!(record.resolution.as_deref(), Some("720p"));
assert!(record.original_request_body.is_none());
assert!(record.prompt.is_none());
assert!(record.resolution.is_none());
assert!(record.progress_message.is_none());
assert!(record.error_message.is_none());
assert!(record.request_metadata.is_none());
}
}
+48 -4
View File
@@ -34,6 +34,7 @@ pub(crate) fn emit_admin_audit(
path_and_query: &str,
control_decision: Option<&GatewayControlDecision>,
) {
let sanitized_path_and_query = sanitize_admin_audit_path(path_and_query);
let Some(decision) = control_decision else {
return;
};
@@ -64,9 +65,10 @@ pub(crate) fn emit_admin_audit(
},
route_kind,
default_target_type(route_family),
path_and_query.to_string(),
sanitized_path_and_query.clone(),
)
};
let target_id = sanitize_admin_audit_target_id(target_id);
let (audit_status, log_level) = classify_admin_audit_response(method, response.status());
if log_level == AdminAuditLogLevel::Info {
@@ -83,7 +85,7 @@ pub(crate) fn emit_admin_audit(
route_family,
route_kind,
method = %method,
path = %path_and_query,
path = %sanitized_path_and_query,
action,
target_type,
target_id = %target_id,
@@ -103,7 +105,7 @@ pub(crate) fn emit_admin_audit(
route_family,
route_kind,
method = %method,
path = %path_and_query,
path = %sanitized_path_and_query,
action,
target_type,
target_id = %target_id,
@@ -112,6 +114,17 @@ pub(crate) fn emit_admin_audit(
}
}
fn sanitize_admin_audit_path(path_and_query: &str) -> String {
crate::middleware::sanitize_access_log_path(path_and_query)
}
fn sanitize_admin_audit_target_id(target_id: String) -> String {
if target_id.trim_start().starts_with('/') {
return sanitize_admin_audit_path(&target_id);
}
target_id
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum AdminAuditLogLevel {
Info,
@@ -151,7 +164,10 @@ fn is_admin_read_method(method: &http::Method) -> bool {
#[cfg(test)]
mod tests {
use super::{classify_admin_audit_response, AdminAuditLogLevel};
use super::{
classify_admin_audit_response, sanitize_admin_audit_path, sanitize_admin_audit_target_id,
AdminAuditLogLevel,
};
use axum::http::{Method, StatusCode};
#[test]
@@ -169,4 +185,32 @@ mod tests {
("failed", AdminAuditLogLevel::Warn)
);
}
#[test]
fn audit_paths_drop_sensitive_query_values() {
assert_eq!(
sanitize_admin_audit_path(
"/api/admin/providers?token=secret&api_key=live-key&limit=25"
),
"/api/admin/providers?limit=25"
);
assert_eq!(
sanitize_admin_audit_path("/install/one-time-secret?view=raw"),
"/install/[redacted]?view=raw"
);
}
#[test]
fn path_shaped_audit_targets_drop_sensitive_query_values() {
assert_eq!(
sanitize_admin_audit_target_id(
"/api/admin/monitoring/trace/request-1?token=secret&limit=25".to_string(),
),
"/api/admin/monitoring/trace/request-1?limit=25"
);
assert_eq!(
sanitize_admin_audit_target_id("resource-id?literal".to_string()),
"resource-id?literal"
);
}
}
+8 -2
View File
@@ -26,7 +26,10 @@ pub(crate) async fn get_request_candidate_trace(
.map_err(|err| GatewayError::Internal(err.to_string()).into_response())?;
match trace {
Some(trace) => Ok(Json(trace)),
Some(mut trace) => {
trace.sanitize_sensitive_diagnostics();
Ok(Json(trace))
}
None => Err((
axum::http::StatusCode::NOT_FOUND,
Json(json!({
@@ -52,7 +55,10 @@ pub(crate) async fn get_decision_trace(
.map_err(|err| GatewayError::Internal(err.to_string()).into_response())?;
match trace {
Some(trace) => Ok(Json(trace)),
Some(mut trace) => {
trace.sanitize_sensitive_diagnostics();
Ok(Json(trace))
}
None => Err((
axum::http::StatusCode::NOT_FOUND,
Json(json!({
+137 -3
View File
@@ -5,7 +5,7 @@ use serde_json::{Map, Value};
use super::schedule::{BackupSchedule, BackupScheduleUnit};
use super::scopes::BackupScope;
#[derive(Debug, Clone, PartialEq, Eq)]
#[derive(Clone, PartialEq, Eq)]
pub(crate) struct S3BackupConfig {
pub(crate) enabled: bool,
pub(crate) scope: BackupScope,
@@ -22,6 +22,28 @@ pub(crate) struct S3BackupConfig {
pub(crate) retention_count: u32,
}
impl fmt::Debug for S3BackupConfig {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
let endpoint_origin = sanitized_endpoint_origin(&self.endpoint);
formatter
.debug_struct("S3BackupConfig")
.field("enabled", &self.enabled)
.field("scope", &self.scope)
.field("endpoint_origin", &endpoint_origin)
.field("region", &self.region)
.field("user_agent", &self.user_agent)
.field("bucket", &self.bucket)
.field("prefix", &self.prefix)
.field("has_access_key_id", &!self.access_key_id.is_empty())
.field("has_secret_access_key", &!self.secret_access_key.is_empty())
.field("path_style", &self.path_style)
.field("compression", &self.compression)
.field("schedule", &self.schedule)
.field("retention_count", &self.retention_count)
.finish()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct BackupConfigError {
message: String,
@@ -84,6 +106,9 @@ impl S3BackupConfig {
"Endpoint(S3 地址)",
enabled,
)?;
if enabled {
validate_s3_endpoint(&endpoint)?;
}
let bucket =
required_or_disabled_string(entries, "backup_s3_bucket", "Bucket(存储桶)", enabled)?;
let access_key_id = required_or_disabled_string(
@@ -99,6 +124,11 @@ impl S3BackupConfig {
enabled,
)?;
let prefix = normalize_s3_prefix(
&optional_string(entries, "backup_s3_prefix")?
.unwrap_or_else(|| "aether/backups/".to_string()),
)?;
Ok(Self {
enabled,
scope,
@@ -108,8 +138,7 @@ impl S3BackupConfig {
user_agent: optional_string(entries, "backup_s3_user_agent")?
.unwrap_or_else(|| "rclone/v1.68.0".to_string()),
bucket,
prefix: optional_string(entries, "backup_s3_prefix")?
.unwrap_or_else(|| "aether/backups/".to_string()),
prefix,
access_key_id,
secret_access_key,
path_style: optional_bool(entries, "backup_s3_path_style")?.unwrap_or(true),
@@ -121,6 +150,48 @@ impl S3BackupConfig {
}
}
fn normalize_s3_prefix(prefix: &str) -> Result<String, BackupConfigError> {
let prefix = prefix.trim().trim_matches('/');
if prefix.is_empty() {
return Ok(String::new());
}
if prefix
.split('/')
.any(|segment| segment.is_empty() || segment == "." || segment == "..")
|| prefix.contains('\\')
{
return Err(BackupConfigError::new(
"Prefix(备份前缀)不能包含空路径段、相对路径段或反斜杠",
));
}
Ok(format!("{prefix}/"))
}
fn validate_s3_endpoint(endpoint: &str) -> Result<(), BackupConfigError> {
let parsed = url::Url::parse(endpoint)
.map_err(|_| BackupConfigError::new("Endpoint(S3 地址)必须是有效的 HTTPS URL"))?;
if parsed.scheme() != "https"
|| parsed.host_str().is_none()
|| !parsed.username().is_empty()
|| parsed.password().is_some()
|| parsed.query().is_some()
|| parsed.fragment().is_some()
{
return Err(BackupConfigError::new(
"Endpoint(S3 地址)必须使用 HTTPS,且不能包含用户凭据、查询参数或片段",
));
}
Ok(())
}
fn sanitized_endpoint_origin(endpoint: &str) -> String {
url::Url::parse(endpoint)
.ok()
.map(|parsed| parsed.origin().ascii_serialization())
.unwrap_or_else(|| "<invalid>".to_string())
}
fn validate_range(label: &str, value: u32, min: u32, max: u32) -> Result<(), BackupConfigError> {
if (min..=max).contains(&value) {
Ok(())
@@ -374,6 +445,69 @@ mod tests {
assert!(err.to_string().contains("Endpoint"));
}
#[test]
fn rejects_insecure_or_credential_bearing_endpoints() {
for endpoint in [
"http://s3.example.com",
"https://user:[email protected]",
"https://s3.example.com?token=secret",
"https://s3.example.com/#fragment",
] {
let entries = serde_json::json!({
"backup_s3_enabled": true,
"backup_s3_endpoint": endpoint,
"backup_s3_bucket": "aether-backups",
"backup_s3_access_key_id": "access",
"backup_s3_secret_access_key": "secret"
});
let error = S3BackupConfig::from_json_map(entries.as_object().unwrap())
.expect_err("unsafe endpoint should fail closed");
assert!(error.to_string().contains("Endpoint"));
}
}
#[test]
fn debug_output_does_not_expose_s3_credentials() {
let entries = serde_json::json!({
"backup_s3_enabled": true,
"backup_s3_endpoint": "https://s3.example.com/path",
"backup_s3_bucket": "aether-backups",
"backup_s3_access_key_id": "access-key-value",
"backup_s3_secret_access_key": "secret-key-value"
});
let config = S3BackupConfig::from_json_map(entries.as_object().unwrap())
.expect("config should parse");
let debug = format!("{config:?}");
assert!(debug.contains("https://s3.example.com"));
assert!(!debug.contains("/path"));
assert!(!debug.contains("access-key-value"));
assert!(!debug.contains("secret-key-value"));
}
#[test]
fn canonicalizes_s3_backup_prefix_once() {
let entries = serde_json::json!({
"backup_s3_enabled": true,
"backup_s3_endpoint": "https://s3.example.com",
"backup_s3_bucket": "aether-backups",
"backup_s3_prefix": "/prod/backups//",
"backup_s3_access_key_id": "access",
"backup_s3_secret_access_key": "secret"
});
let config = S3BackupConfig::from_json_map(entries.as_object().unwrap())
.expect("prefix should be canonicalized");
assert_eq!(config.prefix, "prod/backups/");
for invalid_prefix in ["prod//backups", "prod/../backups", "prod\\backups"] {
let mut entries = entries.clone();
entries["backup_s3_prefix"] = serde_json::json!(invalid_prefix);
assert!(S3BackupConfig::from_json_map(entries.as_object().unwrap()).is_err());
}
}
#[test]
fn applies_default_values_from_system_config_contract() {
let entries = serde_json::json!({
File diff suppressed because it is too large Load Diff
+162
View File
@@ -6,5 +6,167 @@ pub(crate) mod store;
pub(crate) mod task;
pub(crate) mod worker;
pub use executor::{
restore_backup_json, BackupDecryptionKey, BackupRestoreError, BackupRestoreLimits,
RestoredBackupJson, DEFAULT_BACKUP_MAX_ENCRYPTED_BYTES, DEFAULT_BACKUP_MAX_JSON_BYTES,
};
use axum::body::Bytes;
use serde_json::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BackupRestoreScope {
Config,
Users,
Data,
}
impl BackupRestoreScope {
pub const fn as_str(self) -> &'static str {
match self {
Self::Config => "config",
Self::Users => "users",
Self::Data => "data",
}
}
}
#[derive(Debug, thiserror::Error)]
#[error("backup database apply failed: {0}")]
pub struct BackupApplyError(String);
pub async fn apply_restored_backup(
app: &crate::AppState,
restored: RestoredBackupJson,
scope: BackupRestoreScope,
operator_id: Option<&str>,
) -> Result<Result<Value, (http::StatusCode, Value)>, BackupApplyError> {
let (json_bytes, authority) = restored.into_authenticated_parts();
if authority.scope() != scope {
return Err(BackupApplyError(format!(
"authenticated {} backup cannot be applied to {} scope",
authority.scope().as_str(),
scope.as_str(),
)));
}
let request_body = Bytes::from(json_bytes);
let state = crate::admin_api::AdminAppState::new(app);
let result = crate::admin_api::execute_admin_system_import_exclusively(app, async {
match scope {
BackupRestoreScope::Config => {
state
.restore_admin_system_config_backup(&request_body, authority)
.await
}
BackupRestoreScope::Users => {
state
.restore_admin_system_users_backup(&request_body, operator_id, authority)
.await
}
BackupRestoreScope::Data => {
state
.restore_admin_system_data_backup(&request_body, operator_id, authority)
.await
}
}
})
.await
.map_err(|error| {
let message = match error {
crate::admin_api::AdminSystemImportLockError::Conflict => {
"another system import or restore is already running"
}
crate::admin_api::AdminSystemImportLockError::Unavailable => {
"system import coordination is unavailable"
}
crate::admin_api::AdminSystemImportLockError::Lost => {
"system import coordination lease was lost; restore was cancelled and may have partially applied changes"
}
};
BackupApplyError(message.to_string())
})?;
result.map_err(|error| BackupApplyError(error.into_message()))
}
pub(crate) const S3_BACKUP_ENABLED_KEY: &str = "backup_s3_enabled";
pub(crate) const S3_BACKUP_LAST_SLOT_KEY: &str = "backup_s3_last_slot";
#[cfg(test)]
mod tests {
use super::{
apply_restored_backup, BackupDecryptionKey, BackupRestoreLimits, BackupRestoreScope,
RestoredBackupJson,
};
use crate::backup::executor::encrypt_backup_bytes;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use serde_json::json;
fn authenticated_users_backup() -> RestoredBackupJson {
let object_key = "prod/aether-users-backup-20260830-120000.json.zst.aes256gcm";
let compressed = zstd::stream::encode_all(
serde_json::to_vec(&json!({
"version": "1.5",
"exported_at": "2026-08-30T12:00:00Z",
"users": [],
"standalone_keys": [],
}))
.expect("test backup should serialize")
.as_slice(),
0,
)
.expect("test backup should compress");
let (envelope, _) =
encrypt_backup_bytes(DEVELOPMENT_ENCRYPTION_KEY, object_key, &compressed)
.expect("test backup should encrypt");
super::restore_backup_json(
object_key,
&envelope,
&[BackupDecryptionKey::current(DEVELOPMENT_ENCRYPTION_KEY)
.expect("test restore key should build")],
BackupRestoreLimits::default(),
)
.expect("test backup should authenticate")
}
#[tokio::test]
async fn authenticated_backup_cannot_be_applied_to_a_different_scope() {
let restored = authenticated_users_backup();
let error = apply_restored_backup(
&crate::AppState::new().expect("test state should build"),
restored,
BackupRestoreScope::Config,
None,
)
.await
.expect_err("scope mismatch must fail before database access");
assert_eq!(
error.to_string(),
"backup database apply failed: authenticated users backup cannot be applied to config scope"
);
}
#[tokio::test]
async fn authenticated_backup_apply_uses_the_shared_system_import_lock() {
let app = crate::AppState::new().expect("test state should build");
let lock = crate::admin_api::try_acquire_admin_system_import_lease(&app)
.await
.expect("test should acquire the shared import lease");
let error = apply_restored_backup(
&app,
authenticated_users_backup(),
BackupRestoreScope::Users,
None,
)
.await
.expect_err("restore must not interleave with another system import");
assert_eq!(
error.to_string(),
"backup database apply failed: another system import or restore is already running"
);
crate::admin_api::release_admin_system_import_lease(&app, &lock).await;
}
}
+153 -15
View File
@@ -1,5 +1,8 @@
use std::fmt;
const ENCRYPTED_BACKUP_FILE_SUFFIX: &str = ".json.zst.aes256gcm";
const LEGACY_PLAINTEXT_BACKUP_FILE_SUFFIX: &str = ".json.zst";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum BackupScope {
Config,
@@ -52,10 +55,81 @@ impl BackupScope {
}
}
pub(crate) fn from_encrypted_object_key(object_key: &str) -> Option<Self> {
if object_key.is_empty()
|| object_key.starts_with('/')
|| object_key.contains('\0')
|| object_key.contains('\\')
{
return None;
}
let mut segments = object_key.split('/').peekable();
let mut file_name = None;
while let Some(segment) = segments.next() {
if segment.is_empty()
|| segment == "."
|| segment == ".."
|| segment.chars().any(char::is_control)
{
return None;
}
if segments.peek().is_none() {
file_name = Some(segment);
}
}
let file_name = file_name?;
[Self::Config, Self::Users, Self::Data]
.into_iter()
.find(|scope| {
file_name
.strip_prefix(&format!("{}-", scope.file_stem()))
.and_then(|rest| rest.strip_suffix(ENCRYPTED_BACKUP_FILE_SUFFIX))
.is_some_and(is_aether_backup_object_id)
})
}
#[cfg(test)]
pub(crate) fn matching_backup_keys(
self,
prefix: &str,
keys: impl IntoIterator<Item = String>,
) -> Vec<String> {
self.matching_backup_keys_with_suffixes(
prefix,
keys,
&[
ENCRYPTED_BACKUP_FILE_SUFFIX,
LEGACY_PLAINTEXT_BACKUP_FILE_SUFFIX,
],
)
}
pub(crate) fn matching_encrypted_backup_keys(
self,
prefix: &str,
keys: impl IntoIterator<Item = String>,
) -> Vec<String> {
self.matching_backup_keys_with_suffixes(prefix, keys, &[ENCRYPTED_BACKUP_FILE_SUFFIX])
}
pub(crate) fn matching_legacy_plaintext_backup_keys(
self,
prefix: &str,
keys: impl IntoIterator<Item = String>,
) -> Vec<String> {
self.matching_backup_keys_with_suffixes(
prefix,
keys,
&[LEGACY_PLAINTEXT_BACKUP_FILE_SUFFIX],
)
}
fn matching_backup_keys_with_suffixes(
self,
prefix: &str,
keys: impl IntoIterator<Item = String>,
file_suffixes: &[&str],
) -> Vec<String> {
let normalized_prefix = normalized_prefix(prefix);
let expected_prefix = if normalized_prefix.is_empty() {
@@ -64,7 +138,6 @@ impl BackupScope {
format!("{normalized_prefix}/")
};
let file_prefix = format!("{}-", self.file_stem());
let file_suffix = ".json.zst";
keys.into_iter()
.filter(|key| {
@@ -74,20 +147,24 @@ impl BackupScope {
if file_name.contains('/') {
return false;
}
let Some(timestamp) = file_name
.strip_prefix(&file_prefix)
.and_then(|rest| rest.strip_suffix(file_suffix))
else {
let Some(timestamp) = file_name.strip_prefix(&file_prefix).and_then(|rest| {
file_suffixes
.iter()
.find_map(|suffix| rest.strip_suffix(suffix))
}) else {
return false;
};
is_aether_backup_timestamp(timestamp)
is_aether_backup_object_id(timestamp)
})
.collect()
}
fn file_name(self, timestamp: &str) -> String {
format!("{}-{timestamp}.json.zst", self.file_stem())
format!(
"{}-{timestamp}{ENCRYPTED_BACKUP_FILE_SUFFIX}",
self.file_stem()
)
}
}
@@ -110,6 +187,25 @@ fn is_aether_backup_timestamp(timestamp: &str) -> bool {
&& bytes[9..].iter().all(|byte| byte.is_ascii_digit())
}
fn is_aether_backup_object_id(value: &str) -> bool {
if is_aether_backup_timestamp(value) {
return true;
}
let Some((timestamp, collision_digest)) = value.split_once('-').and_then(|(date, rest)| {
let (time, digest) = rest.split_once('-')?;
Some((format!("{date}-{time}"), digest))
}) else {
return false;
};
is_aether_backup_timestamp(&timestamp)
&& collision_digest.len() == 64
&& collision_digest
.bytes()
.all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
}
#[cfg(test)]
mod tests {
use super::BackupScope;
@@ -130,15 +226,15 @@ mod tests {
assert_eq!(
BackupScope::Config.object_key("prod/", "20260524-031500"),
"prod/aether-config-backup-20260524-031500.json.zst"
"prod/aether-config-backup-20260524-031500.json.zst.aes256gcm"
);
assert_eq!(
BackupScope::Users.object_key("prod/", "20260524-031500"),
"prod/aether-users-backup-20260524-031500.json.zst"
"prod/aether-users-backup-20260524-031500.json.zst.aes256gcm"
);
assert_eq!(
BackupScope::Data.object_key("prod/", "20260524-031500"),
"prod/aether-data-backup-20260524-031500.json.zst"
"prod/aether-data-backup-20260524-031500.json.zst.aes256gcm"
);
}
@@ -146,7 +242,7 @@ mod tests {
fn retention_filter_only_matches_same_scope() {
let keys = vec![
"prod/aether-config-backup-20260524-010000.json.zst".to_string(),
"prod/aether-users-backup-20260524-010000.json.zst".to_string(),
"prod/aether-users-backup-20260524-010000.json.zst.aes256gcm".to_string(),
"prod/aether-data-backup-20260524-010000.json.zst".to_string(),
"prod/random.json.zst".to_string(),
];
@@ -155,14 +251,18 @@ mod tests {
assert_eq!(
matched,
vec!["prod/aether-users-backup-20260524-010000.json.zst"]
vec!["prod/aether-users-backup-20260524-010000.json.zst.aes256gcm"]
);
}
#[test]
fn retention_filter_requires_aether_timestamp_format() {
let collision_digest = "a".repeat(64);
let keys = vec![
"prod/aether-users-backup-20260524-010000.json.zst".to_string(),
format!(
"prod/aether-users-backup-20260524-010000-{collision_digest}.json.zst.aes256gcm"
),
"prod/aether-users-backup-foo.json.zst".to_string(),
"prod/aether-users-backup-2026052-010000.json.zst".to_string(),
"prod/aether-users-backup-202605240-010000.json.zst".to_string(),
@@ -171,13 +271,19 @@ mod tests {
"prod/aether-users-backup-20260524010000.json.zst".to_string(),
"prod/aether-users-backup-2026052a-010000.json.zst".to_string(),
"prod/aether-users-backup-20260524-01000x.json.zst".to_string(),
"prod/aether-users-backup-20260524-010000-short.json.zst.aes256gcm".to_string(),
];
let matched = BackupScope::Users.matching_backup_keys("prod/", keys);
assert_eq!(
matched,
vec!["prod/aether-users-backup-20260524-010000.json.zst"]
vec![
"prod/aether-users-backup-20260524-010000.json.zst".to_string(),
format!(
"prod/aether-users-backup-20260524-010000-{collision_digest}.json.zst.aes256gcm"
),
]
);
}
@@ -185,11 +291,11 @@ mod tests {
fn backup_key_prefix_boundaries_are_exact() {
assert_eq!(
BackupScope::Config.object_key("", "20260524-031500"),
"aether-config-backup-20260524-031500.json.zst"
"aether-config-backup-20260524-031500.json.zst.aes256gcm"
);
assert_eq!(
BackupScope::Config.object_key("prod", "20260524-031500"),
"prod/aether-config-backup-20260524-031500.json.zst"
"prod/aether-config-backup-20260524-031500.json.zst.aes256gcm"
);
let keys = vec![
@@ -208,4 +314,36 @@ mod tests {
vec!["prod/aether-config-backup-20260524-010000.json.zst"]
);
}
#[test]
fn encrypted_object_key_parser_binds_scope_and_rejects_path_traversal() {
let collision_digest = "a".repeat(64);
assert_eq!(
BackupScope::from_encrypted_object_key(
"prod/aether-config-backup-20260524-010000.json.zst.aes256gcm"
),
Some(BackupScope::Config)
);
assert_eq!(
BackupScope::from_encrypted_object_key(&format!(
"prod/aether-users-backup-20260524-010000-{collision_digest}.json.zst.aes256gcm"
)),
Some(BackupScope::Users)
);
for key in [
"../aether-data-backup-20260524-010000.json.zst.aes256gcm",
"/aether-data-backup-20260524-010000.json.zst.aes256gcm",
"prod//aether-data-backup-20260524-010000.json.zst.aes256gcm",
"prod/./aether-data-backup-20260524-010000.json.zst.aes256gcm",
"prod\\aether-data-backup-20260524-010000.json.zst.aes256gcm",
"prod/aether-data-backup-invalid.json.zst.aes256gcm",
"prod/unrelated-20260524-010000.json.zst.aes256gcm",
] {
assert_eq!(
BackupScope::from_encrypted_object_key(key),
None,
"unsafe or unrelated key: {key}"
);
}
}
}
+222 -30
View File
@@ -2,23 +2,45 @@ use std::collections::BTreeMap;
use std::fmt;
use std::sync::Arc;
use bytes::Bytes;
use bytes::{Bytes, BytesMut};
use futures_util::TryStreamExt;
use object_store::aws::AmazonS3Builder;
use object_store::path::Path;
use object_store::{ClientOptions, ObjectStore};
use object_store::{ClientOptions, ObjectStore, ObjectStoreExt, PutMode, PutOptions};
use reqwest::header::HeaderValue;
use tokio::sync::RwLock;
use super::config::S3BackupConfig;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum BackupObjectCreateResult {
Created,
AlreadyExists,
}
#[async_trait::async_trait]
pub(crate) trait BackupObjectStore: Send + Sync {
async fn put_object(&self, key: &str, bytes: Bytes) -> Result<(), BackupStoreError>;
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, BackupStoreError>;
async fn put_object_if_absent(
&self,
key: &str,
bytes: Bytes,
) -> Result<BackupObjectCreateResult, BackupStoreError>;
async fn get_object_limited(
&self,
key: &str,
max_bytes: usize,
) -> Result<Bytes, BackupStoreError>;
async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError>;
async fn list_keys_limited(
&self,
prefix: &str,
max_objects: usize,
) -> Result<Vec<String>, BackupStoreError>;
}
#[derive(Debug, Clone, PartialEq, Eq)]
@@ -60,21 +82,72 @@ impl BackupObjectStore for FakeBackupObjectStore {
Ok(())
}
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, BackupStoreError> {
async fn put_object_if_absent(
&self,
key: &str,
bytes: Bytes,
) -> Result<BackupObjectCreateResult, BackupStoreError> {
let mut objects = self.objects.write().await;
if objects.contains_key(key) {
Ok(BackupObjectCreateResult::AlreadyExists)
} else {
objects.insert(key.to_string(), bytes);
Ok(BackupObjectCreateResult::Created)
}
}
async fn get_object_limited(
&self,
key: &str,
max_bytes: usize,
) -> Result<Bytes, BackupStoreError> {
let bytes = self
.objects
.read()
.await
.get(key)
.cloned()
.ok_or_else(|| BackupStoreError::new(format!("backup object `{key}` not found")))?;
if bytes.len() > max_bytes {
return Err(BackupStoreError::new(format!(
"backup object `{key}` exceeds the configured {max_bytes} byte read limit"
)));
}
Ok(bytes)
}
async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> {
self.objects.write().await.remove(key);
Ok(())
}
async fn list_keys_limited(
&self,
prefix: &str,
max_objects: usize,
) -> Result<Vec<String>, BackupStoreError> {
let prefix = directory_list_prefix(prefix);
Ok(self
let keys: Vec<_> = self
.objects
.read()
.await
.keys()
.filter(|key| key.starts_with(&prefix))
.cloned()
.collect())
.collect();
if keys.len() > max_objects {
return Err(BackupStoreError::new(format!(
"backup object listing exceeds the configured {max_objects} object limit"
)));
}
Ok(keys)
}
}
async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> {
self.objects.write().await.remove(key);
Ok(())
#[cfg(test)]
impl FakeBackupObjectStore {
pub(crate) async fn object_bytes(&self, key: &str) -> Option<Bytes> {
self.objects.read().await.get(key).cloned()
}
}
@@ -125,17 +198,68 @@ impl BackupObjectStore for ObjectStoreS3BackupStore {
.map_err(|error| BackupStoreError::object_store("put", key, error))
}
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, BackupStoreError> {
let prefix_path = list_prefix_path(prefix);
let mut keys = self
async fn put_object_if_absent(
&self,
key: &str,
bytes: Bytes,
) -> Result<BackupObjectCreateResult, BackupStoreError> {
let options = PutOptions {
mode: PutMode::Create,
..PutOptions::default()
};
match self
.store
.list(prefix_path.as_ref())
.map_ok(|meta| meta.location.to_string())
.try_collect::<Vec<_>>()
.put_opts(&Path::from(key), bytes.into(), options)
.await
.map_err(|error| BackupStoreError::object_store("list", prefix, error))?;
keys.sort();
Ok(keys)
{
Ok(_) => Ok(BackupObjectCreateResult::Created),
Err(object_store::Error::AlreadyExists { .. }) => {
Ok(BackupObjectCreateResult::AlreadyExists)
}
Err(error) => Err(BackupStoreError::object_store(
"conditional put",
key,
error,
)),
}
}
async fn get_object_limited(
&self,
key: &str,
max_bytes: usize,
) -> Result<Bytes, BackupStoreError> {
let result = self
.store
.get(&Path::from(key))
.await
.map_err(|error| BackupStoreError::object_store("get", key, error))?;
if result.meta.size > u64::try_from(max_bytes).unwrap_or(u64::MAX) {
return Err(BackupStoreError::new(format!(
"backup object `{key}` exceeds the configured {max_bytes} byte read limit"
)));
}
let object_size = result.meta.size;
let mut stream = result.into_stream();
let mut bytes = BytesMut::with_capacity(
usize::try_from(object_size)
.unwrap_or(max_bytes)
.min(max_bytes)
.min(8 * 1024 * 1024),
);
while let Some(chunk) = stream
.try_next()
.await
.map_err(|error| BackupStoreError::object_store("read", key, error))?
{
if bytes.len().saturating_add(chunk.len()) > max_bytes {
return Err(BackupStoreError::new(format!(
"backup object `{key}` exceeds the configured {max_bytes} byte read limit"
)));
}
bytes.extend_from_slice(&chunk);
}
Ok(bytes.freeze())
}
async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> {
@@ -144,6 +268,30 @@ impl BackupObjectStore for ObjectStoreS3BackupStore {
.await
.map_err(|error| BackupStoreError::object_store("delete", key, error))
}
async fn list_keys_limited(
&self,
prefix: &str,
max_objects: usize,
) -> Result<Vec<String>, BackupStoreError> {
let prefix_path = list_prefix_path(prefix);
let mut objects = self.store.list(prefix_path.as_ref());
let mut keys = Vec::new();
while let Some(meta) = objects
.try_next()
.await
.map_err(|error| BackupStoreError::object_store("list", prefix, error))?
{
if keys.len() >= max_objects {
return Err(BackupStoreError::new(format!(
"backup object listing exceeds the configured {max_objects} object limit"
)));
}
keys.push(meta.location.to_string());
}
keys.sort();
Ok(keys)
}
}
fn directory_list_prefix(prefix: &str) -> String {
@@ -166,10 +314,12 @@ fn list_prefix_path(prefix: &str) -> Option<Path> {
#[cfg(test)]
mod tests {
use super::{list_prefix_path, BackupObjectStore, FakeBackupObjectStore};
use super::{
list_prefix_path, BackupObjectCreateResult, BackupObjectStore, FakeBackupObjectStore,
};
#[tokio::test]
async fn fake_backup_object_store_puts_lists_and_deletes() {
async fn fake_backup_object_store_puts_and_lists() {
let store = FakeBackupObjectStore::default();
store
.put_object(
@@ -186,17 +336,59 @@ mod tests {
.await
.unwrap();
let keys = store.list_keys("prod/").await.unwrap();
assert_eq!(keys.len(), 2);
store
.delete_object("prod/aether-data-backup-20260524-010000.json.zst")
.await
.unwrap();
let keys = store.list_keys("prod/").await.unwrap();
let keys = store.list_keys_limited("prod/", 2).await.unwrap();
assert_eq!(
keys,
vec!["prod/aether-data-backup-20260524-020000.json.zst"]
vec![
"prod/aether-data-backup-20260524-010000.json.zst",
"prod/aether-data-backup-20260524-020000.json.zst",
]
);
}
#[tokio::test]
async fn fake_backup_object_store_enforces_read_and_listing_limits() {
let store = FakeBackupObjectStore::default();
store
.put_object("prod/one", bytes::Bytes::from_static(b"1234"))
.await
.unwrap();
store
.put_object("prod/two", bytes::Bytes::from_static(b"5678"))
.await
.unwrap();
assert!(store.get_object_limited("prod/one", 3).await.is_err());
assert_eq!(
store.get_object_limited("prod/one", 4).await.unwrap(),
bytes::Bytes::from_static(b"1234")
);
assert!(store.list_keys_limited("prod/", 1).await.is_err());
assert_eq!(store.list_keys_limited("prod/", 2).await.unwrap().len(), 2);
}
#[tokio::test]
async fn fake_backup_object_store_conditional_put_never_overwrites() {
let store = FakeBackupObjectStore::default();
let key = "prod/aether-data-backup-20260524-010000.json.zst.aes256gcm";
assert_eq!(
store
.put_object_if_absent(key, bytes::Bytes::from_static(b"first"))
.await
.unwrap(),
BackupObjectCreateResult::Created
);
assert_eq!(
store
.put_object_if_absent(key, bytes::Bytes::from_static(b"second"))
.await
.unwrap(),
BackupObjectCreateResult::AlreadyExists
);
assert_eq!(
store.object_bytes(key).await.as_deref(),
Some(b"first".as_slice())
);
}
@@ -218,7 +410,7 @@ mod tests {
.await
.unwrap();
let keys = store.list_keys("prod").await.unwrap();
let keys = store.list_keys_limited("prod", 10).await.unwrap();
assert_eq!(
keys,
+410 -55
View File
@@ -1,4 +1,5 @@
use std::fmt;
use std::future::Future;
use std::time::Duration;
use aether_admin::system::admin_system_config_default_value;
@@ -12,14 +13,15 @@ use chrono::Utc;
use futures_util::FutureExt;
use serde::Serialize;
use serde_json::{json, Map, Value};
use tokio::task::{JoinError, JoinHandle};
use tracing::warn;
use super::config::S3BackupConfig;
use super::executor::{run_backup_with_store, BackupRunResult};
use super::scopes::BackupScope;
use super::store::ObjectStoreS3BackupStore;
use crate::admin_api::AdminAppState;
use crate::handlers::shared::decrypt_catalog_secret_with_fallbacks;
use crate::admin_api::{AdminAppState, SystemExportMode};
use crate::handlers::shared::decrypt_or_migrate_system_config_secret;
use crate::task_runtime::{
append_event_with_logging, build_task_run_id, now_unix_secs, spawn_fire_and_forget,
task_definition, update_run_status, upsert_run_with_logging, TASK_KEY_SYSTEM_S3_BACKUP,
@@ -48,6 +50,9 @@ const S3_BACKUP_CONFIG_KEYS: &[&str] = &[
];
const S3_BACKUP_QUEUED_MESSAGE: &str = "S3 备份任务已提交";
const S3_BACKUP_INTERNAL_ERROR_DETAIL: &str = "S3 备份服务暂时不可用";
const S3_BACKUP_TASK_FAILURE_CODE: &str = "s3_backup_failed";
const S3_BACKUP_SLOT_RECORD_FAILURE_CODE: &str = "s3_backup_slot_record_failed";
const S3_BACKUP_TASK_LOCK_KEY: &str = "task_runtime:lock:system.s3.backup";
const S3_BACKUP_TASK_LOCK_TTL: Duration = Duration::from_secs(60 * 60 * 6);
const S3_BACKUP_TASK_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(60 * 5);
@@ -67,6 +72,16 @@ pub(crate) struct S3BackupTaskError {
detail: String,
}
enum BackupLockRenewalFailure<E> {
Lost,
Backend(E),
}
enum BackupLockRaceOutcome<T> {
BackupCompleted(T),
LeaseLost(Result<(), JoinError>),
}
impl S3BackupTaskError {
fn bad_request(detail: impl Into<String>) -> Self {
Self {
@@ -114,8 +129,12 @@ impl fmt::Display for S3BackupTaskError {
impl std::error::Error for S3BackupTaskError {}
impl From<GatewayError> for S3BackupTaskError {
fn from(error: GatewayError) -> Self {
Self::internal(format!("{error:?}"))
fn from(_error: GatewayError) -> Self {
warn!(
error_category = "dependency_failed",
"S3 backup dependency failed"
);
Self::internal(S3_BACKUP_INTERNAL_ERROR_DETAIL)
}
}
@@ -220,8 +239,6 @@ fn s3_backup_task_payload_json(
) -> Value {
let mut payload = json!({
"scope": config.scope.as_config_value(),
"bucket": config.bucket.clone(),
"prefix": config.prefix.clone(),
"compression": config.compression.clone(),
"trigger": trigger,
});
@@ -259,7 +276,7 @@ fn spawn_s3_backup_worker(
Some(100),
Some("S3 备份任务异常退出".to_string()),
None,
Some("S3 backup task panicked".to_string()),
Some("background_task_panicked".to_string()),
None,
Some(now_unix_secs()),
)
@@ -294,16 +311,67 @@ async fn run_s3_backup_worker_inner(
.await;
append_event_with_logging(&app, &run_id, "running", "S3 backup task started", None).await;
let heartbeat = spawn_s3_backup_task_heartbeat(app.clone(), run_id.clone(), lock);
let result = run_s3_backup_once(&app, &config).await;
heartbeat.abort();
let _ = heartbeat.await;
let heartbeat = spawn_s3_backup_task_heartbeat(app.clone(), run_id.clone(), lock.clone());
let result = match race_backup_with_lock_heartbeat(run_s3_backup_once(&app, &config), heartbeat)
.await
{
BackupLockRaceOutcome::BackupCompleted(result) => {
match require_successful_backup_lock_renewal(
app.runtime_state
.lock_renew(&lock, S3_BACKUP_TASK_LOCK_TTL)
.await,
) {
Ok(()) => result,
Err(BackupLockRenewalFailure::Lost) => {
warn!(
run_id = %run_id,
lock_key = %lock.key,
"S3 backup task lost its distributed lock before publishing completion"
);
Err(S3BackupTaskError::service_unavailable(
"S3 备份任务锁已失效,任务完成状态未发布",
))
}
Err(BackupLockRenewalFailure::Backend(error)) => {
warn!(
run_id = %run_id,
lock_key = %lock.key,
error = %error,
"S3 backup task could not verify its distributed lock before publishing completion"
);
Err(S3BackupTaskError::service_unavailable(
"无法确认 S3 备份任务锁所有权,任务完成状态未发布",
))
}
}
}
BackupLockRaceOutcome::LeaseLost(heartbeat_result) => {
match heartbeat_result {
Ok(()) => warn!(
run_id = %run_id,
"S3 backup task stopped after losing its distributed lock"
),
Err(error) => warn!(
run_id = %run_id,
error = %error,
"S3 backup lock heartbeat task failed"
),
}
Err(S3BackupTaskError::service_unavailable(
"S3 备份任务锁已失效,任务已停止",
))
}
};
match result {
Ok(result) => {
if let Some(slot) = scheduled_backup_slot_to_record(scheduled_slot.as_deref(), true) {
if let Err(error) = record_scheduled_backup_slot(&app, &slot).await {
warn!(error = ?error, run_id = %run_id, "S3 backup slot record failed");
if record_scheduled_backup_slot(&app, &slot).await.is_err() {
warn!(
error_category = "slot_record_failed",
run_id = %run_id,
"S3 backup slot record failed"
);
let _ = update_run_status(
&app,
&run_id,
@@ -311,7 +379,7 @@ async fn run_s3_backup_worker_inner(
Some(100),
Some("S3 备份任务完成,但记录调度时间失败".to_string()),
None,
Some(format!("S3 backup slot record failed: {error:?}")),
Some(S3_BACKUP_SLOT_RECORD_FAILURE_CODE.to_string()),
None,
Some(now_unix_secs()),
)
@@ -321,7 +389,7 @@ async fn run_s3_backup_worker_inner(
&run_id,
"failed",
"S3 backup slot record failed",
Some(json!({ "error": format!("{error:?}") })),
Some(json!({ "error_code": S3_BACKUP_SLOT_RECORD_FAILURE_CODE })),
)
.await;
return;
@@ -349,8 +417,12 @@ async fn run_s3_backup_worker_inner(
)
.await;
}
Err(error) => {
warn!(error = %error, run_id = %run_id, "S3 backup task failed");
Err(_) => {
warn!(
error_category = "backup_execution_failed",
run_id = %run_id,
"S3 backup task failed"
);
let _ = update_run_status(
&app,
&run_id,
@@ -358,7 +430,7 @@ async fn run_s3_backup_worker_inner(
Some(100),
Some("S3 备份任务失败".to_string()),
None,
Some(error.to_string()),
Some(S3_BACKUP_TASK_FAILURE_CODE.to_string()),
None,
Some(now_unix_secs()),
)
@@ -368,13 +440,44 @@ async fn run_s3_backup_worker_inner(
&run_id,
"failed",
"S3 backup task failed",
Some(json!({ "error": error.to_string() })),
Some(json!({ "error_code": S3_BACKUP_TASK_FAILURE_CODE })),
)
.await;
}
}
}
async fn race_backup_with_lock_heartbeat<F, T>(
backup: F,
mut heartbeat: JoinHandle<()>,
) -> BackupLockRaceOutcome<T>
where
F: Future<Output = T>,
{
tokio::pin!(backup);
tokio::select! {
biased;
heartbeat_result = &mut heartbeat => {
BackupLockRaceOutcome::LeaseLost(heartbeat_result)
}
result = &mut backup => {
heartbeat.abort();
let _ = heartbeat.await;
BackupLockRaceOutcome::BackupCompleted(result)
}
}
}
fn require_successful_backup_lock_renewal<E>(
result: Result<bool, E>,
) -> Result<(), BackupLockRenewalFailure<E>> {
match result {
Ok(true) => Ok(()),
Ok(false) => Err(BackupLockRenewalFailure::Lost),
Err(error) => Err(BackupLockRenewalFailure::Backend(error)),
}
}
fn spawn_s3_backup_task_heartbeat(
app: AppState,
run_id: String,
@@ -386,10 +489,30 @@ fn spawn_s3_backup_task_heartbeat(
interval.tick().await;
loop {
interval.tick().await;
let _ = app
.runtime_state
.lock_renew(&lock, S3_BACKUP_TASK_LOCK_TTL)
.await;
match require_successful_backup_lock_renewal(
app.runtime_state
.lock_renew(&lock, S3_BACKUP_TASK_LOCK_TTL)
.await,
) {
Ok(()) => {}
Err(BackupLockRenewalFailure::Lost) => {
warn!(
run_id = %run_id,
lock_key = %lock.key,
"S3 backup task distributed lock is no longer owned"
);
return;
}
Err(BackupLockRenewalFailure::Backend(error)) => {
warn!(
run_id = %run_id,
lock_key = %lock.key,
error = %error,
"S3 backup task distributed lock renewal failed"
);
return;
}
}
let _ = update_run_status(
&app,
&run_id,
@@ -458,9 +581,15 @@ async fn acquire_s3_backup_task_lock(
Ok(None) => Err(S3BackupTaskError::conflict(
"已有 S3 备份任务正在执行,请等待当前任务完成后再试",
)),
Err(error) => Err(S3BackupTaskError::service_unavailable(format!(
"无法获取 S3 备份任务锁:{error}"
))),
Err(_) => {
warn!(
error_category = "lock_acquisition_failed",
"S3 backup task lock acquisition failed"
);
Err(S3BackupTaskError::service_unavailable(
"无法获取 S3 备份任务锁,请稍后重试",
))
}
}
}
@@ -509,25 +638,77 @@ async fn run_s3_backup_once(
app: &AppState,
config: &S3BackupConfig,
) -> Result<BackupRunResult, S3BackupTaskError> {
let admin_state = AdminAppState::new(app);
let payload = match config.scope {
BackupScope::Config => {
admin_state
.build_admin_system_config_export_payload()
.await?
}
BackupScope::Users => {
admin_state
.build_admin_system_users_export_payload()
.await?
}
BackupScope::Data => admin_state.build_admin_system_data_export_payload().await?,
let Some(encryption_secret) = effective_backup_encryption_secret(app) else {
return Err(S3BackupTaskError::service_unavailable(
"S3 备份需要 AETHER_BACKUP_ENCRYPTION_KEY 或可用的数据加密密钥",
));
};
let store = ObjectStoreS3BackupStore::from_config(config)
.map_err(|error| S3BackupTaskError::internal(error.to_string()))?;
run_backup_with_store(config, &store, payload, Utc::now())
let payload = build_s3_backup_payload_exclusively(app, config.scope).await?;
let store = ObjectStoreS3BackupStore::from_config(config).map_err(|_| {
warn!(
error_category = "object_store_initialization_failed",
"S3 backup object store initialization failed"
);
S3BackupTaskError::internal(S3_BACKUP_INTERNAL_ERROR_DETAIL)
})?;
run_backup_with_store(config, &store, payload, Utc::now(), &encryption_secret)
.await
.map_err(|error| S3BackupTaskError::internal(error.to_string()))
.map_err(|_| {
warn!(
error_category = "backup_execution_failed",
"S3 backup execution failed"
);
S3BackupTaskError::internal(S3_BACKUP_INTERNAL_ERROR_DETAIL)
})
}
async fn build_s3_backup_payload_exclusively(
app: &AppState,
scope: BackupScope,
) -> Result<Value, S3BackupTaskError> {
let admin_state = AdminAppState::new(app);
crate::admin_api::execute_admin_system_import_exclusively(app, async {
match scope {
BackupScope::Config => {
admin_state
.build_admin_system_config_export_payload(SystemExportMode::RecoveryBackup)
.await
}
BackupScope::Users => {
admin_state
.build_admin_system_users_export_payload(SystemExportMode::RecoveryBackup)
.await
}
BackupScope::Data => {
admin_state
.build_admin_system_data_export_payload(SystemExportMode::RecoveryBackup)
.await
}
}
})
.await
.map_err(|error| {
warn!(
error_category = "system_import_coordination_failed",
lock_error = ?error,
"S3 backup snapshot could not acquire or retain the system import lock"
);
S3BackupTaskError::service_unavailable(S3_BACKUP_INTERNAL_ERROR_DETAIL)
})?
.map_err(S3BackupTaskError::from)
}
fn effective_backup_encryption_secret(app: &AppState) -> Option<String> {
std::env::var("AETHER_BACKUP_ENCRYPTION_KEY")
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.or_else(|| {
app.encryption_key()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
async fn load_s3_backup_config_for_run(
@@ -551,7 +732,7 @@ pub(crate) async fn load_s3_backup_config_values(
.or_else(|| admin_system_config_default_value(key));
if let Some(value) = value {
let value = if *key == "backup_s3_secret_access_key" {
decrypt_s3_secret_access_key(app, value)?
decrypt_s3_secret_access_key(app, value).await?
} else {
value
};
@@ -561,39 +742,56 @@ pub(crate) async fn load_s3_backup_config_values(
Ok(values)
}
fn decrypt_s3_secret_access_key(app: &AppState, value: Value) -> Result<Value, S3BackupTaskError> {
let Some(ciphertext) = value
async fn decrypt_s3_secret_access_key(
app: &AppState,
value: Value,
) -> Result<Value, S3BackupTaskError> {
let Some(stored_value) = value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(value);
};
let Some(plaintext) = decrypt_catalog_secret_with_fallbacks(app.encryption_key(), ciphertext)
else {
return Err(S3BackupTaskError::bad_request(
let plaintext = decrypt_or_migrate_system_config_secret(
app,
"backup_s3_secret_access_key",
stored_value.to_string(),
)
.await
.map_err(|_| {
S3BackupTaskError::bad_request(
"S3 备份配置无效:Secret Access Key(访问密钥)无法解密,请重新填写",
));
};
)
})?;
Ok(Value::String(plaintext))
}
fn backup_run_result_json(result: &BackupRunResult) -> Value {
json!({
"scope": result.scope.as_config_value(),
"bucket": result.bucket,
"object_key": result.object_key,
"bytes": result.bytes,
"sha256": result.sha256,
"export_version": result.export_version,
"exported_at": result.exported_at,
"compression": result.compression,
"deleted_old_objects": result.deleted_old_objects,
"encryption": result.encryption,
"legacy_encrypted_copies_created": result.legacy_encrypted_copies_created,
"legacy_encrypted_copies_verified": result.legacy_encrypted_copies_verified,
"legacy_plaintext_objects_deleted": result.legacy_plaintext_objects_deleted,
"legacy_plaintext_objects_retained": result.legacy_plaintext_objects_retained,
"retention_cleanup_candidates": result.retention_cleanup_candidates,
"automatic_deletions": result.legacy_plaintext_objects_deleted,
"object_cleanup_mode": "legacy_plaintext_deleted_after_verified_encryption",
"versioned_storage_cleanup_required": result.versioned_storage_cleanup_required,
"versioned_storage_cleanup_notice": "legacy_plaintext_versions_require_external_cleanup",
})
}
#[cfg(test)]
mod tests {
use std::convert::Infallible;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
@@ -603,9 +801,78 @@ mod tests {
};
use crate::data::GatewayDataState;
use crate::handlers::shared::decrypt_system_config_secret;
use crate::state::AppState;
use crate::task_runtime::{now_unix_secs, TASK_KEY_SYSTEM_S3_BACKUP};
#[test]
fn backup_lock_renewal_requires_ownership_and_preserves_backend_errors() {
assert!(matches!(
super::require_successful_backup_lock_renewal::<Infallible>(Ok(true)),
Ok(())
));
assert!(matches!(
super::require_successful_backup_lock_renewal::<Infallible>(Ok(false)),
Err(super::BackupLockRenewalFailure::Lost)
));
assert!(matches!(
super::require_successful_backup_lock_renewal(Err("redis unavailable")),
Err(super::BackupLockRenewalFailure::Backend(
"redis unavailable"
))
));
}
#[tokio::test]
async fn lost_backup_lock_stops_race_without_publishing_backup_result() {
let destructive_stage_reached = Arc::new(AtomicBool::new(false));
let destructive_stage_for_backup = Arc::clone(&destructive_stage_reached);
let backup = async move {
std::future::pending::<()>().await;
destructive_stage_for_backup.store(true, Ordering::Release);
Ok::<(), super::S3BackupTaskError>(())
};
let heartbeat = tokio::spawn(async {});
let outcome = super::race_backup_with_lock_heartbeat(backup, heartbeat).await;
assert!(matches!(
outcome,
super::BackupLockRaceOutcome::LeaseLost(Ok(()))
));
assert!(!destructive_stage_reached.load(Ordering::Acquire));
}
#[tokio::test]
async fn completed_heartbeat_wins_when_backup_completion_is_also_ready() {
let heartbeat = tokio::spawn(async {});
tokio::task::yield_now().await;
let outcome = super::race_backup_with_lock_heartbeat(async { 42_u8 }, heartbeat).await;
assert!(matches!(
outcome,
super::BackupLockRaceOutcome::LeaseLost(Ok(()))
));
}
#[tokio::test]
async fn s3_backup_snapshot_refuses_to_overlap_system_import() {
let app = AppState::new().expect("app state should build");
let lease = crate::admin_api::try_acquire_admin_system_import_lease(&app)
.await
.expect("test should acquire the system import lease");
let error = super::build_s3_backup_payload_exclusively(
&app,
crate::backup::scopes::BackupScope::Config,
)
.await
.expect_err("backup snapshot must not overlap a system import");
crate::admin_api::release_admin_system_import_lease(&app, &lease).await;
assert_eq!(error.status(), axum::http::StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(error.detail(), super::S3_BACKUP_INTERNAL_ERROR_DETAIL);
}
fn valid_s3_backup_config_values() -> Vec<(String, serde_json::Value)> {
vec![
(
@@ -632,6 +899,92 @@ mod tests {
]
}
#[tokio::test]
async fn legacy_plaintext_s3_secret_is_migrated_when_config_loads() {
let plaintext = "legacy-s3-secret-access-key";
let mut entries = valid_s3_backup_config_values();
entries
.iter_mut()
.find(|(key, _)| key == "backup_s3_secret_access_key")
.expect("secret config fixture should exist")
.1 = serde_json::json!(plaintext);
let app = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::disabled()
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
.with_system_config_values_for_tests(entries),
);
let values = super::load_s3_backup_config_values(&app)
.await
.expect("legacy S3 config should load");
assert_eq!(
values.get("backup_s3_secret_access_key"),
Some(&serde_json::json!(plaintext))
);
let stored = app
.read_system_config_json_value_strong("backup_s3_secret_access_key")
.await
.expect("stored S3 secret should read")
.and_then(|value| value.as_str().map(ToOwned::to_owned))
.expect("stored S3 secret should remain a string");
assert_ne!(stored, plaintext);
assert_eq!(
decrypt_system_config_secret(&app, "backup_s3_secret_access_key", &stored)
.expect("migrated S3 secret should decrypt"),
plaintext
);
}
#[tokio::test]
async fn undecryptable_s3_fernet_secret_fails_closed() {
let plaintext = "s3-secret-from-unavailable-key";
let ciphertext = encrypt_python_fernet_plaintext("unavailable-s3-key", plaintext)
.expect("unknown-key fixture should encrypt");
let mut entries = valid_s3_backup_config_values();
entries
.iter_mut()
.find(|(key, _)| key == "backup_s3_secret_access_key")
.expect("secret config fixture should exist")
.1 = serde_json::json!(ciphertext.clone());
let app = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::disabled()
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
.with_system_config_values_for_tests(entries),
);
let error = super::load_s3_backup_config_values(&app)
.await
.expect_err("unknown-key S3 ciphertext must fail closed");
let error_text = error.to_string();
assert!(!error_text.contains(plaintext));
assert!(!error_text.contains(&ciphertext));
assert_eq!(
app.read_system_config_json_value_strong("backup_s3_secret_access_key")
.await
.expect("stored S3 secret should read"),
Some(serde_json::json!(ciphertext))
);
}
#[test]
fn gateway_dependency_errors_are_not_exposed_to_backup_clients() {
let error = super::S3BackupTaskError::from(crate::GatewayError::Internal(
"postgresql://admin:[email protected]/aether".to_string(),
));
assert_eq!(
error.status(),
axum::http::StatusCode::INTERNAL_SERVER_ERROR
);
assert_eq!(error.detail(), super::S3_BACKUP_INTERNAL_ERROR_DETAIL);
assert!(!error.detail().contains("database-secret"));
}
fn stored_s3_backup_run(status: BackgroundTaskStatus) -> StoredBackgroundTaskRun {
let now = now_unix_secs();
StoredBackgroundTaskRun {
@@ -774,7 +1127,9 @@ mod tests {
let payload = super::s3_backup_task_payload_json(&config, "manual", None);
assert!(payload["bucket"].is_string());
assert!(payload.get("bucket").is_none());
assert!(payload.get("prefix").is_none());
assert_eq!(payload["scope"], serde_json::json!("data"));
assert_eq!(payload["trigger"], serde_json::json!("manual"));
assert!(!payload.to_string().contains("secret"));
}
+20 -8
View File
@@ -29,8 +29,11 @@ pub(crate) fn spawn_s3_backup_worker(app: AppState) -> Option<JoinHandle<()>> {
interval.tick().await;
loop {
interval.tick().await;
if let Err(error) = run_s3_backup_schedule_tick(&app, Utc::now()).await {
warn!(error = ?error, "S3 backup schedule tick failed");
if run_s3_backup_schedule_tick(&app, Utc::now()).await.is_err() {
warn!(
error_category = "schedule_tick_failed",
"S3 backup schedule tick failed"
);
}
}
},
@@ -43,15 +46,21 @@ async fn run_s3_backup_schedule_tick(
) -> Result<(), GatewayError> {
let values = match super::task::load_s3_backup_config_values(app).await {
Ok(values) => values,
Err(error) => {
warn!(error = %error, "S3 backup schedule config load failed");
Err(_) => {
warn!(
error_category = "config_load_failed",
"S3 backup schedule config load failed"
);
return Ok(());
}
};
let config = match S3BackupConfig::from_json_map(&values) {
Ok(config) => config,
Err(error) => {
warn!(error = %error, "S3 backup schedule config is invalid");
Err(_) => {
warn!(
error_category = "config_invalid",
"S3 backup schedule config is invalid"
);
return Ok(());
}
};
@@ -68,8 +77,11 @@ async fn run_s3_backup_schedule_tick(
match super::task::start_s3_backup_task_for_schedule(app.clone(), slot).await {
Ok(_) => {}
Err(error) => {
warn!(error = %error, "S3 backup scheduled task submission failed");
Err(_) => {
warn!(
error_category = "task_submission_failed",
"S3 backup scheduled task submission failed"
);
}
}
Ok(())
+349 -46
View File
@@ -1,8 +1,10 @@
use crate::handlers::shared::{
decrypt_catalog_secret_with_fallbacks, system_config_bool, system_config_string,
bark_device_key_binding, canonical_bark_server_url, decrypt_or_migrate_bark_device_key,
system_config_bool, system_config_string,
};
use crate::{AppState, GatewayError};
use serde_json::{json, Value};
use std::net::{IpAddr, SocketAddr};
pub(crate) const BARK_PUSH_ENABLED_KEY: &str = "module.bark_push.enabled";
pub(crate) const BARK_PUSH_DEVICE_KEY_KEY: &str = "module.bark_push.device_key";
@@ -10,8 +12,20 @@ pub(crate) const BARK_PUSH_SERVER_URL_KEY: &str = "module.bark_push.server_url";
pub(crate) const BARK_PUSH_TEMPLATE_KEY: &str = "module.bark_push.template";
const DEFAULT_BARK_API_BASE: &str = "https://api.day.app";
const BARK_ALLOW_HTTP_ENV: &str = "AETHER_BARK_ALLOW_HTTP";
const BARK_ALLOW_PRIVATE_TARGETS_ENV: &str = "AETHER_BARK_ALLOW_PRIVATE_TARGETS";
const MAX_BARK_RESPONSE_BYTES: usize = 64 * 1024;
const BARK_CONNECT_TIMEOUT_MS: u64 = 10_000;
const BARK_REQUEST_TIMEOUT_MS: u64 = 300_000;
const MAX_BARK_SERVER_URL_BYTES: usize = 2 * 1024;
const MAX_BARK_DEVICE_KEY_BYTES: usize = 512;
const MAX_BARK_TEMPLATE_BYTES: usize = 256 * 1024;
const MAX_BARK_TITLE_BYTES: usize = 512;
const MAX_BARK_BODY_BYTES: usize = 2 * 1024 * 1024;
const MAX_BARK_RENDERED_BODY_BYTES: usize = 2 * 1024 * 1024;
const MAX_BARK_RESOLVED_ADDRESSES: usize = 32;
#[derive(Debug, Clone)]
#[derive(Clone)]
pub(crate) struct BarkPushConfig {
pub(crate) enabled: bool,
pub(crate) device_key: Option<String>,
@@ -19,6 +33,21 @@ pub(crate) struct BarkPushConfig {
pub(crate) template: Option<String>,
}
impl std::fmt::Debug for BarkPushConfig {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("BarkPushConfig")
.field("enabled", &self.enabled)
.field(
"device_key",
&self.device_key.as_ref().map(|_| "[REDACTED]"),
)
.field("server_url", &self.server_url)
.field("template", &self.template.as_ref().map(|_| "[REDACTED]"))
.finish()
}
}
pub(crate) async fn bark_push_module_enabled(state: &AppState) -> Result<bool, GatewayError> {
let value = state
.read_system_config_json_value(BARK_PUSH_ENABLED_KEY)
@@ -35,23 +64,34 @@ pub(crate) async fn read_bark_push_config(
state: &AppState,
) -> Result<BarkPushConfig, GatewayError> {
let enabled = bark_push_module_enabled(state).await?;
let device_key = state
.read_system_config_json_value(BARK_PUSH_DEVICE_KEY_KEY)
.await?
.and_then(|value| system_config_string(Some(&value)))
.map(|value| {
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &value).unwrap_or(value)
});
let server_url = state
.read_system_config_json_value(BARK_PUSH_SERVER_URL_KEY)
.await?
.and_then(|value| system_config_string(Some(&value)))
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| DEFAULT_BARK_API_BASE.to_string());
validate_bark_config_field("server_url", &server_url, MAX_BARK_SERVER_URL_BYTES)?;
let server_url = normalized_bark_server_url(&server_url)?;
let binding = bark_device_key_binding(&server_url)
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址不是合法 URL".to_string()))?;
let device_key = state
.read_system_config_json_value(BARK_PUSH_DEVICE_KEY_KEY)
.await?
.and_then(|value| system_config_string(Some(&value)));
let device_key = match device_key {
Some(value) => Some(decrypt_or_migrate_bark_device_key(state, &binding, value).await?),
None => None,
};
if let Some(device_key) = device_key.as_deref() {
validate_bark_config_field("device_key", device_key, MAX_BARK_DEVICE_KEY_BYTES)?;
}
let template = state
.read_system_config_json_value(BARK_PUSH_TEMPLATE_KEY)
.await?
.and_then(|value| system_config_string(Some(&value)));
if let Some(template) = template.as_deref() {
validate_bark_config_field("template", template, MAX_BARK_TEMPLATE_BYTES)?;
}
Ok(BarkPushConfig {
enabled,
@@ -62,7 +102,7 @@ pub(crate) async fn read_bark_push_config(
}
pub(crate) async fn send_bark_push(
state: &AppState,
_state: &AppState,
config: &BarkPushConfig,
title: &str,
markdown_body: &str,
@@ -76,11 +116,13 @@ pub(crate) async fn send_bark_push(
"Bark Device Key 不能为空".to_string(),
));
}
let server_url = normalized_bark_server_url(&config.server_url)?;
let body = render_bark_body(config.template.as_deref(), title, markdown_body);
let response = state
.client
.post(format!("{server_url}/push"))
validate_bark_config_field("device_key", device_key, MAX_BARK_DEVICE_KEY_BYTES)?;
validate_bark_content_field("title", title, MAX_BARK_TITLE_BYTES)?;
validate_bark_content_field("body", markdown_body, MAX_BARK_BODY_BYTES)?;
let (client, push_url) = build_bark_push_client_and_url(&config.server_url).await?;
let body = render_bark_body(config.template.as_deref(), title, markdown_body)?;
let response = client
.post(push_url)
.json(&json!({
"device_key": device_key,
"title": title,
@@ -88,16 +130,14 @@ pub(crate) async fn send_bark_push(
}))
.send()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
.map_err(|err| GatewayError::Internal(bark_request_error_message(&err)))?;
let status = response.status();
let text = response
.text()
let body = aether_http::read_response_bytes_with_limit(response, MAX_BARK_RESPONSE_BYTES)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
.map_err(|err| GatewayError::Internal(bark_response_body_error_message(&err)))?;
let text = String::from_utf8_lossy(&body);
if !status.is_success() {
return Err(GatewayError::Internal(format!(
"Bark 返回 HTTP {status}: {text}"
)));
return Err(GatewayError::Internal(format!("Bark 返回 HTTP {status}")));
}
if let Ok(payload) = serde_json::from_str::<Value>(&text) {
let code_is_ok = payload
@@ -114,53 +154,259 @@ pub(crate) async fn send_bark_push(
})
.unwrap_or(true);
if !code_is_ok {
return Err(GatewayError::Internal(format!("Bark 返回失败: {payload}")));
return Err(GatewayError::Internal("Bark 返回失败".to_string()));
}
}
Ok(())
}
fn normalized_bark_server_url(server_url: &str) -> Result<String, GatewayError> {
let server_url = server_url.trim().trim_end_matches('/');
if server_url.is_empty() {
return Err(GatewayError::Internal(
"Bark 服务器地址不能为空".to_string(),
));
}
if !server_url.starts_with("https://") && !server_url.starts_with("http://") {
return Err(GatewayError::Internal(
"Bark 服务器地址必须以 http:// 或 https:// 开头".to_string(),
));
}
Ok(server_url.to_string())
fn bark_request_error_message(error: &reqwest::Error) -> String {
format!("Bark 请求失败 ({})", bark_reqwest_error_kind(error))
}
fn render_bark_body(template: Option<&str>, title: &str, markdown_body: &str) -> String {
match template {
Some(template) if !template.trim().is_empty() => template
.replace("{title}", title)
.replace("{body}", markdown_body),
_ => markdown_body.to_string(),
fn bark_response_body_error_message(error: &aether_http::ResponseBodyReadError) -> String {
match error {
aether_http::ResponseBodyReadError::TooLarge { max_bytes } => {
format!("Bark 响应超过 {max_bytes} 字节")
}
aether_http::ResponseBodyReadError::Read(error) => {
format!("Bark 响应读取失败 ({})", bark_reqwest_error_kind(error))
}
}
}
fn bark_reqwest_error_kind(error: &reqwest::Error) -> &'static str {
if error.is_timeout() {
"timeout"
} else if error.is_connect() {
"connect"
} else if error.is_request() {
"request"
} else {
"transport"
}
}
fn normalized_bark_server_url(server_url: &str) -> Result<String, GatewayError> {
canonical_bark_server_url(server_url)
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址不是合法 URL".to_string()))
}
async fn build_bark_push_client_and_url(
server_url: &str,
) -> Result<(reqwest::Client, url::Url), GatewayError> {
validate_bark_config_field("server_url", server_url, MAX_BARK_SERVER_URL_BYTES)?;
let normalized = normalized_bark_server_url(server_url)?;
let mut push_url = url::Url::parse(&normalized)
.map_err(|_| GatewayError::Internal("Bark 服务器地址不是合法 URL".to_string()))?;
validate_bark_transport_policy(&push_url, env_flag_enabled(BARK_ALLOW_HTTP_ENV))?;
let host = push_url
.host_str()
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少主机名".to_string()))?
.to_string();
let port = push_url
.port_or_known_default()
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少端口".to_string()))?;
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
tokio::time::timeout(
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
tokio::net::lookup_host((host.as_str(), port)),
)
.await
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析超时".to_string()))?
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析失败".to_string()))?
.take(MAX_BARK_RESOLVED_ADDRESSES)
.collect::<Vec<_>>()
};
let allow_benchmarking_ip = push_url.scheme() == "https"
&& push_url.port_or_known_default() == Some(443)
&& host.eq_ignore_ascii_case("api.day.app");
validate_bark_resolved_addresses(
&addresses,
env_flag_enabled(BARK_ALLOW_PRIVATE_TARGETS_ENV),
allow_benchmarking_ip,
)?;
push_url
.path_segments_mut()
.map_err(|_| GatewayError::Internal("Bark 服务器地址不能作为基础 URL".to_string()))?
.pop_if_empty()
.push("push");
let mut builder = aether_http::apply_http_client_config(
reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none()),
&aether_http::HttpClientConfig {
connect_timeout_ms: Some(BARK_CONNECT_TIMEOUT_MS),
request_timeout_ms: Some(BARK_REQUEST_TIMEOUT_MS),
http2_adaptive_window: true,
..aether_http::HttpClientConfig::default()
},
);
if host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(&host, &addresses);
}
let client = builder
.build()
.map_err(|_| GatewayError::Internal("Bark HTTP 客户端初始化失败".to_string()))?;
Ok((client, push_url))
}
fn validate_bark_transport_policy(url: &url::Url, allow_http: bool) -> Result<(), GatewayError> {
if url.scheme() == "http" && !allow_http {
return Err(GatewayError::Internal(format!(
"Bark 服务器必须使用 HTTPS;如确需明文 HTTP,请显式设置 {BARK_ALLOW_HTTP_ENV}=true"
)));
}
Ok(())
}
fn validate_bark_resolved_addresses(
addresses: &[SocketAddr],
allow_private: bool,
allow_benchmarking_ip: bool,
) -> Result<(), GatewayError> {
if addresses.is_empty() {
return Err(GatewayError::Internal(
"Bark 服务器 DNS 解析未返回地址".to_string(),
));
}
if !allow_private
&& addresses.iter().any(|address| {
aether_http::is_private_or_reserved_ip(address.ip())
&& !(allow_benchmarking_ip
&& aether_http::is_ipv4_benchmarking_fake_ip(address.ip()))
})
{
return Err(GatewayError::Internal(format!(
"Bark 服务器解析到私有或保留地址;如确需内网自建服务,请显式设置 {BARK_ALLOW_PRIVATE_TARGETS_ENV}=true"
)));
}
Ok(())
}
fn env_flag_enabled(key: &str) -> bool {
std::env::var(key).ok().is_some_and(|value| {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
)
})
}
fn validate_bark_config_field(
field: &str,
value: &str,
max_bytes: usize,
) -> Result<(), GatewayError> {
if value.len() > max_bytes || value.bytes().any(|byte| byte == 0) {
return Err(GatewayError::Internal(format!(
"Bark {field} exceeds the allowed size or contains a NUL byte"
)));
}
Ok(())
}
fn validate_bark_content_field(
field: &str,
value: &str,
max_bytes: usize,
) -> Result<(), GatewayError> {
validate_bark_config_field(field, value, max_bytes)
}
fn render_bark_body(
template: Option<&str>,
title: &str,
markdown_body: &str,
) -> Result<String, GatewayError> {
validate_bark_content_field("title", title, MAX_BARK_TITLE_BYTES)?;
validate_bark_content_field("body", markdown_body, MAX_BARK_BODY_BYTES)?;
let template = template
.filter(|value| !value.trim().is_empty())
.unwrap_or("{body}");
validate_bark_config_field("template", template, MAX_BARK_TEMPLATE_BYTES)?;
let mut rendered = String::with_capacity(template.len().min(MAX_BARK_RENDERED_BODY_BYTES));
let mut cursor = 0usize;
while cursor < template.len() {
let remaining = &template[cursor..];
let title_match = remaining.find("{title}");
let body_match = remaining.find("{body}");
let next = match (title_match, body_match) {
(None, None) => {
append_bark_rendered_part(&mut rendered, remaining)?;
cursor = template.len();
continue;
}
(Some(index), None) => (index, "{title}", title),
(None, Some(index)) => (index, "{body}", markdown_body),
(Some(title_index), Some(body_index)) if title_index <= body_index => {
(title_index, "{title}", title)
}
(Some(_), Some(body_index)) => (body_index, "{body}", markdown_body),
};
append_bark_rendered_part(&mut rendered, &remaining[..next.0])?;
append_bark_rendered_part(&mut rendered, next.2)?;
cursor += next.0 + next.1.len();
}
if rendered.is_empty() && template.is_empty() {
return Ok(String::new());
}
Ok(rendered)
}
fn append_bark_rendered_part(output: &mut String, part: &str) -> Result<(), GatewayError> {
let next_len = output
.len()
.checked_add(part.len())
.ok_or_else(|| GatewayError::Internal("Bark rendered body is too large".to_string()))?;
if next_len > MAX_BARK_RENDERED_BODY_BYTES {
return Err(GatewayError::Internal(
"Bark rendered body exceeds the allowed size".to_string(),
));
}
output.push_str(part);
Ok(())
}
#[cfg(test)]
mod tests {
use super::{normalized_bark_server_url, render_bark_body};
use super::{
bark_request_error_message, bark_response_body_error_message, normalized_bark_server_url,
render_bark_body, validate_bark_resolved_addresses, validate_bark_transport_policy,
};
use std::net::SocketAddr;
#[test]
fn bark_body_uses_template_when_provided() {
let rendered = render_bark_body(Some("{title}\n\n{body}"), "告警", "原始正文");
let rendered = render_bark_body(Some("{title}\n\n{body}"), "告警", "原始正文")
.expect("template should render");
assert_eq!(rendered, "告警\n\n原始正文");
}
#[test]
fn bark_body_falls_back_to_markdown_body_for_empty_template() {
assert_eq!(render_bark_body(None, "告警", "原始正文"), "原始正文");
assert_eq!(
render_bark_body(Some(" "), "告警", "原始正文"),
render_bark_body(None, "告警", "原始正文").expect("fallback should render"),
"原始正文"
);
assert_eq!(
render_bark_body(Some(" "), "告警", "原始正文").expect("fallback should render"),
"原始正文"
);
}
#[test]
fn bark_body_rejects_template_expansion_bombs_and_oversized_content() {
let template = "x".repeat(super::MAX_BARK_TEMPLATE_BYTES + 1);
assert!(render_bark_body(Some(&template), "告警", "正文").is_err());
let body = "x".repeat(super::MAX_BARK_BODY_BYTES + 1);
assert!(render_bark_body(None, "告警", &body).is_err());
}
#[test]
@@ -170,4 +416,61 @@ mod tests {
"https://api.day.app"
);
}
#[test]
fn bark_server_url_rejects_credentials_query_and_fragments() {
for invalid in [
"https://[email protected]",
"https://example.com?target=internal",
"https://example.com/#fragment",
] {
assert!(normalized_bark_server_url(invalid).is_err(), "{invalid}");
}
}
#[test]
fn bark_http_transport_requires_explicit_opt_in() {
let url = url::Url::parse("http://bark.example.com").unwrap();
assert!(validate_bark_transport_policy(&url, false).is_err());
assert!(validate_bark_transport_policy(&url, true).is_ok());
}
#[test]
fn bark_private_targets_require_explicit_opt_in() {
let private = [SocketAddr::from(([127, 0, 0, 1], 443))];
assert!(validate_bark_resolved_addresses(&private, false, false).is_err());
assert!(validate_bark_resolved_addresses(&private, true, false).is_ok());
}
#[test]
fn bark_builtin_server_allows_benchmarking_ip_only_with_https_default_port() {
let fake = [SocketAddr::from(([198, 18, 75, 234], 443))];
assert!(validate_bark_resolved_addresses(&fake, false, true).is_ok());
assert!(validate_bark_resolved_addresses(
&[fake[0], SocketAddr::from(([127, 0, 0, 1], 443))],
false,
true,
)
.is_err());
assert!(validate_bark_resolved_addresses(&fake, false, false).is_err());
}
#[tokio::test]
async fn bark_transport_errors_do_not_expose_server_url_or_response_body() {
let secret = "bark-secret-query";
let error = reqwest::Client::new()
.post(format!("ftp://bark.example.test/push?token={secret}"))
.send()
.await
.expect_err("unsupported URL scheme should fail before network I/O");
let message = bark_request_error_message(&error);
assert!(!message.contains(secret));
assert!(!message.contains("bark.example.test"));
let body_error = aether_http::ResponseBodyReadError::Read(error);
let message = bark_response_body_error_message(&body_error);
assert!(!message.contains(secret));
assert!(!message.contains("bark.example.test"));
}
}
File diff suppressed because it is too large Load Diff
@@ -193,6 +193,7 @@ pub(crate) fn parse_probe_url(raw: &str) -> Result<Url, ProbeFailure> {
|| url.password().is_some()
|| url.query().is_some()
|| url.fragment().is_some()
|| (url.scheme() == "ws" && !aether_http::url_has_literal_loopback_host(&url))
{
return Err(ProbeFailure::InvalidEndpoint);
}
@@ -210,6 +211,7 @@ pub(crate) const fn turn_timeout(args: &ProbeArgs) -> Duration {
async fn run_probe(config: &ProbeConfig, started_at: Instant) -> Result<ProbeReport, ProbeFailure> {
let client = wreq::Client::builder()
.no_proxy()
.connect_timeout(config.turn_timeout)
.timeout(config.turn_timeout)
.build()
@@ -434,7 +436,14 @@ mod tests {
#[test]
fn probe_url_rejects_credentials_and_query_strings() {
assert!(parse_probe_url("wss://example.test/v1/responses").is_ok());
assert!(parse_probe_url("ws://localhost:8080/v1/responses").is_ok());
assert!(parse_probe_url("ws://127.42.0.1:8080/v1/responses").is_ok());
assert!(parse_probe_url("ws://[::1]:8080/v1/responses").is_ok());
assert!(parse_probe_url("https://example.test/v1/responses").is_err());
assert!(parse_probe_url("ws://example.test/v1/responses").is_err());
assert!(parse_probe_url("ws://10.0.0.1/v1/responses").is_err());
assert!(parse_probe_url("ws://0.0.0.0:8080/v1/responses").is_err());
assert!(parse_probe_url("ws://[::ffff:127.0.0.1]:8080/v1/responses").is_err());
assert!(parse_probe_url("wss://[email protected]/v1/responses").is_err());
assert!(parse_probe_url("wss://example.test/v1/responses?token=secret").is_err());
}
+1
View File
@@ -422,6 +422,7 @@ mod tests {
local_rejection: None,
allowed_models: None,
ip_rules: None,
verified_api_key_hash: None,
}
}
+9
View File
@@ -173,6 +173,15 @@ impl SystemConfigCache {
self.detach_all_loads();
}
pub(crate) fn invalidate(&self, key: &str) {
let Ok(_mutation) = self.mutation.lock() else {
return;
};
self.generation.fetch_add(1, Ordering::AcqRel);
self.entries.remove(&key.to_string());
self.detach_all_loads();
}
pub(crate) fn insert_if_generation(
&self,
key: String,
+1
View File
@@ -18,6 +18,7 @@ pub(crate) const TUNNEL_AFFINITY_FORWARDED_BY_HEADER: &str =
"x-aether-tunnel-affinity-forwarded-by";
pub(crate) const TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER: &str =
"x-aether-tunnel-affinity-owner-instance-id";
pub(crate) const TUNNEL_AFFINITY_NODE_ID_HEADER: &str = "x-aether-tunnel-affinity-node-id";
pub(crate) const EXECUTION_PATH_PUBLIC_PROXY_PASSTHROUGH: &str = "public_proxy_passthrough";
pub(crate) const EXECUTION_PATH_LOCAL_PROXY_PASSTHROUGH_REMOVED: &str =
"local_proxy_passthrough_removed";
@@ -23,7 +23,7 @@ pub(crate) fn extract_requested_model(
body: &Bytes,
) -> Option<String> {
if decision.route_family.as_deref() == Some("gemini") {
if let Some(model) = extract_gemini_model_from_path(uri.path()) {
if let Some(model) = extract_gemini_requested_model_from_path(uri.path()) {
return Some(model);
}
}
@@ -43,23 +43,46 @@ pub(crate) fn extract_requested_model(
.filter(|value| !value.is_empty())
}
fn extract_gemini_requested_model_from_path(path: &str) -> Option<String> {
let model = extract_gemini_model_from_path(path)?;
Some(
model
.split_once("/operations/")
.map(|(model, _)| model)
.unwrap_or(model.as_str())
.to_string(),
)
}
pub(super) fn extract_request_credentials(
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: &str,
) -> GatewayExtractedCredentials {
extract_request_credentials_with_trusted_auth(headers, uri, auth_endpoint_signature, cfg!(test))
}
pub(super) fn extract_request_credentials_with_trusted_auth(
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: &str,
trusted_auth_verified: bool,
) -> GatewayExtractedCredentials {
let bundle = GatewayCredentialBundle {
authorization_bearer: header_value_str(headers, http::header::AUTHORIZATION.as_str())
.as_deref()
.and_then(extract_bearer_token)
.map(ToOwned::to_owned),
authorization_bearer: unique_header_value_str(
headers,
http::header::AUTHORIZATION.as_str(),
)
.as_deref()
.and_then(extract_bearer_token)
.map(ToOwned::to_owned),
x_api_key: header_value_str(headers, "x-api-key"),
api_key: header_value_str(headers, "api-key"),
x_goog_api_key: header_value_str(headers, "x-goog-api-key"),
query_key: extract_query_api_key(uri),
cookie_header: header_value_str(headers, http::header::COOKIE.as_str()),
};
let trusted_headers = extract_trusted_auth_headers(headers);
let trusted_headers = extract_trusted_auth_headers(headers, trusted_auth_verified);
let trusted_admin_headers = extract_trusted_admin_headers(headers);
let primary = select_primary_credential(auth_endpoint_signature, &bundle);
@@ -71,6 +94,20 @@ pub(super) fn extract_request_credentials(
}
}
fn unique_header_value_str(headers: &http::HeaderMap, key: &str) -> Option<String> {
let mut values = headers.get_all(key).iter();
let value = values.next()?;
if values.next().is_some() {
return None;
}
value
.to_str()
.ok()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(in crate::control) fn resolve_gateway_credential_carrier(
headers: &http::HeaderMap,
uri: &Uri,
@@ -85,6 +122,7 @@ pub(in crate::control) fn resolve_gateway_credential_carrier(
})
}
#[cfg(test)]
fn has_trusted_gateway_marker(headers: &http::HeaderMap) -> bool {
header_value_str(headers, crate::constants::GATEWAY_HEADER)
.unwrap_or_default()
@@ -97,13 +135,32 @@ pub(super) fn build_auth_context_cache_key(
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: &str,
) -> Option<String> {
build_auth_context_cache_key_with_trusted_auth(
headers,
uri,
auth_endpoint_signature,
cfg!(test),
)
}
pub(super) fn build_auth_context_cache_key_with_trusted_auth(
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: &str,
trusted_auth_verified: bool,
) -> Option<String> {
let signature = auth_endpoint_signature.trim();
if signature.is_empty() {
return None;
}
let extracted = extract_request_credentials(headers, uri, signature);
let extracted = extract_request_credentials_with_trusted_auth(
headers,
uri,
signature,
trusted_auth_verified,
);
let trusted_headers = extracted.trusted_headers;
let bundle = extracted.bundle;
if bundle.authorization_bearer.is_none()
@@ -135,7 +192,7 @@ pub(super) fn build_auth_context_cache_key(
})
.unwrap_or_default();
Some(format!(
let raw_cache_identity = format!(
"{signature}\n{}\n{}\n{}\n{}\n{}\n{}\n{}\n{}\n{}\n{}",
bundle.authorization_bearer.unwrap_or_default(),
bundle.x_api_key.unwrap_or_default(),
@@ -147,11 +204,26 @@ pub(super) fn build_auth_context_cache_key(
trusted_api_key_id,
trusted_balance_remaining,
trusted_access_allowed,
))
);
let mut hasher = Sha256::new();
hasher.update(raw_cache_identity.as_bytes());
Some(format!("auth-context:sha256:{:x}", hasher.finalize()))
}
fn extract_trusted_auth_headers(headers: &http::HeaderMap) -> Option<GatewayTrustedAuthHeaders> {
if !has_trusted_gateway_marker(headers) {
fn extract_trusted_auth_headers(
headers: &http::HeaderMap,
trusted_auth_verified: bool,
) -> Option<GatewayTrustedAuthHeaders> {
if !trusted_auth_verified {
return None;
}
#[cfg(test)]
if !header_value_str(headers, crate::constants::GATEWAY_HEADER)
.unwrap_or_default()
.trim()
.to_ascii_lowercase()
.starts_with("rust-phase3")
{
return None;
}
let user_id = header_value_str(headers, crate::constants::TRUSTED_AUTH_USER_ID_HEADER)
@@ -387,7 +459,7 @@ fn extract_bearer_token(value: &str) -> Option<&str> {
return None;
}
let token = token.trim();
if token.is_empty() {
if token.is_empty() || token.chars().any(char::is_whitespace) {
None
} else {
Some(token)
@@ -472,6 +544,46 @@ mod tests {
assert_eq!(requested_model.as_deref(), Some("gpt-5.4"));
}
#[test]
fn extract_requested_model_handles_gemini_generation_and_operation_paths() {
let generation_decision = GatewayControlDecision::synthetic(
"/v1beta/models/gemini-2.5-pro:generateContent",
Some("ai_public".to_string()),
Some("gemini".to_string()),
Some("generate_content".to_string()),
Some("gemini:generate_content".to_string()),
);
let operation_decision = GatewayControlDecision::synthetic(
"/v1beta/models/veo-3/operations/task-123:cancel",
Some("ai_public".to_string()),
Some("gemini".to_string()),
Some("video".to_string()),
Some("gemini:video".to_string()),
);
let headers = http::HeaderMap::new();
assert_eq!(
extract_requested_model(
&generation_decision,
&uri("/v1beta/models/gemini-2.5-pro:generateContent"),
&headers,
&Bytes::new(),
)
.as_deref(),
Some("gemini-2.5-pro")
);
assert_eq!(
extract_requested_model(
&operation_decision,
&uri("/v1beta/models/veo-3/operations/task-123:cancel"),
&headers,
&Bytes::new(),
)
.as_deref(),
Some("veo-3")
);
}
#[test]
fn selects_openai_bearer_as_provider_api_key() {
let mut headers = http::HeaderMap::new();
@@ -491,6 +603,33 @@ mod tests {
);
}
#[test]
fn rejects_duplicate_or_combined_authorization_credentials() {
let mut duplicate = http::HeaderMap::new();
duplicate.append(
http::header::AUTHORIZATION,
"Bearer first-token".parse().unwrap(),
);
duplicate.append(
http::header::AUTHORIZATION,
"Bearer second-token".parse().unwrap(),
);
let extracted =
extract_request_credentials(&duplicate, &uri("/api/admin/system"), "admin:operational");
assert!(extracted.bundle.authorization_bearer.is_none());
assert!(extracted.primary.is_none());
let mut combined = http::HeaderMap::new();
combined.insert(
http::header::AUTHORIZATION,
"Bearer first-token, Bearer second-token".parse().unwrap(),
);
let extracted =
extract_request_credentials(&combined, &uri("/api/admin/system"), "admin:operational");
assert!(extracted.bundle.authorization_bearer.is_none());
assert!(extracted.primary.is_none());
}
#[test]
fn selects_codex_live_bearer_as_provider_api_key() {
let mut headers = http::HeaderMap::new();
@@ -608,7 +747,7 @@ mod tests {
}
#[test]
fn cache_key_includes_cookie_header() {
fn cache_key_hashes_cookie_header_instead_of_retaining_session_secret() {
let mut headers = http::HeaderMap::new();
headers.insert(http::header::COOKIE, "session=abc123".parse().unwrap());
@@ -618,7 +757,8 @@ mod tests {
"internal:session",
)
.expect("cache key should exist");
assert!(cache_key.contains("session=abc123"));
assert!(cache_key.starts_with("auth-context:sha256:"));
assert!(!cache_key.contains("session=abc123"));
}
#[test]
@@ -669,12 +809,10 @@ mod tests {
.expect("trusted cache key should exist");
assert_ne!(first, second);
assert!(first.contains("user-1"));
assert!(first.contains("key-1"));
assert!(first.contains("1.5"));
assert!(first.contains("true"));
assert!(second.contains("user-2"));
assert!(second.contains("false"));
for raw_identity in ["user-1", "key-1", "1.5", "user-2"] {
assert!(!first.contains(raw_identity));
assert!(!second.contains(raw_identity));
}
}
#[test]
+2 -1
View File
@@ -224,7 +224,7 @@ fn wallet_finite_available_usd(
Some(wallet.balance.max(0.0) + wallet.gift_balance.max(0.0))
}
async fn estimate_execution_plan_cost_upper_bound_usd(
pub(crate) async fn estimate_execution_plan_cost_upper_bound_usd(
state: &AppState,
plan: &aether_contracts::ExecutionPlan,
report_context: Option<&serde_json::Value>,
@@ -925,6 +925,7 @@ mod tests {
local_rejection: None,
allowed_models: Some(allowed_models),
ip_rules: None,
verified_api_key_hash: None,
});
decision
}
+8 -5
View File
@@ -7,13 +7,16 @@ mod types;
pub(crate) use credentials::extract_requested_model;
pub(super) use credentials::resolve_gateway_credential_carrier;
pub(crate) use gate::{
execution_plan_balance_capacity_rejection, request_model_local_rejection,
should_buffer_request_for_local_auth, trusted_auth_local_rejection, GatewayLocalAuthRejection,
estimate_execution_plan_cost_upper_bound_usd, execution_plan_balance_capacity_rejection,
request_model_local_rejection, should_buffer_request_for_local_auth,
trusted_auth_local_rejection, GatewayLocalAuthRejection,
};
pub(crate) use resolution::{
refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot,
resolve_execution_runtime_auth_context, GatewayAdminPrincipalContext,
GatewayControlAuthContext,
resolve_execution_runtime_auth_context, resolve_local_admin_session_principal,
GatewayAdminPrincipalContext, GatewayControlAuthContext,
};
pub(super) use resolution::{
resolve_control_decision_auth_with_trusted_auth, ControlDecisionAuthResolution,
};
pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution};
pub(crate) use types::GatewayCredentialCarrier;
+337 -120
View File
@@ -4,11 +4,9 @@ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use axum::http::Uri;
use base64::Engine as _;
use hmac::Mac;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use tracing::{debug, info};
use tracing::{debug, info, warn};
use crate::wallet_runtime::{
local_rejection_from_wallet_access, resolve_wallet_auth_gate_uncached,
@@ -17,7 +15,8 @@ use crate::{AppState, GatewayError};
use super::super::GatewayControlDecision;
use super::credentials::{
build_auth_context_cache_key, current_unix_secs, extract_request_credentials,
build_auth_context_cache_key, build_auth_context_cache_key_with_trusted_auth,
current_unix_secs, extract_request_credentials, extract_request_credentials_with_trusted_auth,
extract_trusted_admin_headers, hash_api_key,
};
use super::gate::GatewayLocalAuthRejection;
@@ -27,6 +26,9 @@ use super::types::{
};
use crate::cache::{AuthContextCacheGeneration, AuthContextInflightRegistration};
use crate::headers::header_value_str;
use crate::local_auth_token::{
decode_local_auth_token, local_auth_token_identity_matches_user, LocalAuthTokenType,
};
const AUTH_CONTEXT_CACHE_TTL: Duration = Duration::from_secs(60);
const AUTH_CONTEXT_CACHE_REFRESH_INTERVAL: Duration = Duration::from_secs(10);
@@ -93,6 +95,30 @@ pub(crate) struct GatewayControlAuthContext {
pub(crate) allowed_models: Option<Vec<String>>,
#[serde(skip)]
pub(crate) ip_rules: Option<Vec<String>>,
/// Credential verifier that established this API-key identity. Long-lived
/// executions use it to prove that a later row with the same IDs is still
/// the record authenticated by the original request.
#[serde(skip)]
pub(crate) verified_api_key_hash: Option<VerifiedApiKeyHash>,
}
#[derive(Clone)]
pub(crate) struct VerifiedApiKeyHash(String);
impl VerifiedApiKeyHash {
fn new(value: String) -> Self {
Self(value)
}
fn as_str(&self) -> &str {
self.0.as_str()
}
}
impl std::fmt::Debug for VerifiedApiKeyHash {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("VerifiedApiKeyHash([REDACTED])")
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
@@ -109,11 +135,30 @@ pub(in super::super) enum ControlDecisionAuthResolution {
}
pub(in super::super) async fn resolve_control_decision_auth(
state: &AppState,
headers: &http::HeaderMap,
uri: &Uri,
trace_id: &str,
decision: GatewayControlDecision,
) -> Result<ControlDecisionAuthResolution, GatewayError> {
resolve_control_decision_auth_with_trusted_auth(
state,
headers,
uri,
trace_id,
decision,
cfg!(test),
)
.await
}
pub(in super::super) async fn resolve_control_decision_auth_with_trusted_auth(
state: &AppState,
headers: &http::HeaderMap,
uri: &Uri,
trace_id: &str,
mut decision: GatewayControlDecision,
trusted_auth_verified: bool,
) -> Result<ControlDecisionAuthResolution, GatewayError> {
if let Some(admin_principal) =
resolve_trusted_admin_principal(headers, decision.auth_endpoint_signature.as_deref())
@@ -132,10 +177,18 @@ pub(in super::super) async fn resolve_control_decision_auth(
decision.admin_principal = Some(admin_principal);
}
let auth_context_cache_key = decision
.auth_endpoint_signature
.as_deref()
.and_then(|signature| build_auth_context_cache_key(headers, uri, signature));
let auth_context_cache_key =
decision
.auth_endpoint_signature
.as_deref()
.and_then(|signature| {
build_auth_context_cache_key_with_trusted_auth(
headers,
uri,
signature,
trusted_auth_verified,
)
});
let mut resolved_auth_context = None;
if let Some(cache_key) = auth_context_cache_key.as_deref() {
@@ -149,6 +202,7 @@ pub(in super::super) async fn resolve_control_decision_auth(
decision.auth_endpoint_signature.as_deref(),
headers,
uri,
trusted_auth_verified,
)
.await?,
);
@@ -168,6 +222,7 @@ pub(in super::super) async fn resolve_control_decision_auth(
uri,
decision.auth_endpoint_signature.as_deref(),
true,
trusted_auth_verified,
)
.await?;
}
@@ -336,7 +391,7 @@ async fn resolve_local_admin_principal(
let Some(access_token) = extracted.bundle.authorization_bearer.as_deref() else {
return Ok(None);
};
let claims = match decode_local_auth_token(access_token, "access") {
let claims = match decode_local_auth_token(access_token, LocalAuthTokenType::Access) {
Ok(claims) => claims,
Err(_) => return Ok(None),
};
@@ -351,6 +406,14 @@ async fn resolve_local_admin_principal(
resolve_local_admin_principal_from_claims(state, headers, uri, &claims).await
}
pub(crate) async fn resolve_local_admin_session_principal(
state: &AppState,
headers: &http::HeaderMap,
uri: &Uri,
) -> Result<Option<GatewayAdminPrincipalContext>, GatewayError> {
resolve_local_admin_principal(state, headers, uri, Some("admin:operational")).await
}
async fn resolve_local_admin_principal_from_claims(
state: &AppState,
headers: &http::HeaderMap,
@@ -373,6 +436,9 @@ async fn resolve_local_admin_principal_from_claims(
if !user.is_active || user.is_deleted || !crate::roles::can_access_admin_console(&user.role) {
return Ok(None);
}
if !local_auth_token_identity_matches_user(claims, &user) {
return Ok(None);
}
let now = chrono::Utc::now();
let Some(session) = state.find_user_session(user_id, session_id).await? else {
@@ -380,6 +446,7 @@ async fn resolve_local_admin_principal_from_claims(
};
if session.is_revoked()
|| session.is_expired(now)
|| session.security_version != user.security_version
|| session.client_device_id != client_device_id
{
return Ok(None);
@@ -431,68 +498,6 @@ fn local_admin_user_agent(headers: &http::HeaderMap) -> Option<String> {
.map(|value| value.chars().take(1000).collect())
}
fn local_auth_secret() -> String {
std::env::var("JWT_SECRET_KEY")
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.unwrap_or_else(|| "aether-rust-dev-jwt-secret".to_string())
}
fn decode_local_auth_token(
token: &str,
expected_type: &str,
) -> Result<serde_json::Map<String, Value>, String> {
let mut parts = token.split('.');
let Some(header_segment) = parts.next() else {
return Err("invalid token".to_string());
};
let Some(payload_segment) = parts.next() else {
return Err("invalid token".to_string());
};
let Some(signature_segment) = parts.next() else {
return Err("invalid token".to_string());
};
if parts.next().is_some() {
return Err("invalid token".to_string());
}
let signing_input = format!("{header_segment}.{payload_segment}");
let signature = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(signature_segment)
.map_err(|_| "invalid token".to_string())?;
let mut mac = hmac::Hmac::<sha2::Sha256>::new_from_slice(local_auth_secret().as_bytes())
.map_err(|_| "invalid token".to_string())?;
mac.update(signing_input.as_bytes());
mac.verify_slice(&signature)
.map_err(|_| "invalid token".to_string())?;
let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(payload_segment)
.map_err(|_| "invalid token".to_string())?;
let payload =
serde_json::from_slice::<Value>(&payload_bytes).map_err(|_| "invalid token".to_string())?;
let payload = payload
.as_object()
.cloned()
.ok_or_else(|| "invalid token".to_string())?;
let actual_type = payload
.get("type")
.and_then(Value::as_str)
.unwrap_or_default();
if actual_type != expected_type {
return Err("invalid token".to_string());
}
let exp = payload
.get("exp")
.and_then(Value::as_i64)
.ok_or_else(|| "invalid token".to_string())?;
if exp <= chrono::Utc::now().timestamp() {
return Err("expired token".to_string());
}
Ok(payload)
}
pub(crate) async fn resolve_execution_runtime_auth_context(
state: &AppState,
decision: &GatewayControlDecision,
@@ -525,6 +530,7 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
Some(auth_endpoint_signature),
headers,
uri,
cfg!(test),
)
.await
.map(Some);
@@ -539,6 +545,7 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
uri,
Some(auth_endpoint_signature),
true,
cfg!(test),
)
.await?
{
@@ -558,6 +565,7 @@ async fn revalidate_cached_auth_context(
auth_endpoint_signature: Option<&str>,
headers: &http::HeaderMap,
uri: &Uri,
trusted_auth_verified: bool,
) -> Result<GatewayControlAuthContext, GatewayError> {
if is_negative_auth_context(&auth_context)
|| !auth_context.access_allowed
@@ -581,6 +589,7 @@ async fn revalidate_cached_auth_context(
uri,
auth_context.clone(),
auth_endpoint_signature,
trusted_auth_verified,
)
.await
{
@@ -622,6 +631,7 @@ async fn revalidate_cached_auth_context(
uri,
auth_context,
auth_endpoint_signature,
trusted_auth_verified,
)
.await;
if refreshed.is_err() {
@@ -639,9 +649,16 @@ async fn resolve_security_fresh_auth_context(
uri: &Uri,
stale: GatewayControlAuthContext,
auth_endpoint_signature: Option<&str>,
trusted_auth_verified: bool,
) -> Result<GatewayControlAuthContext, GatewayError> {
if let Some(refreshed) =
resolve_data_backed_auth_context(state, headers, uri, auth_endpoint_signature).await?
if let Some(refreshed) = resolve_data_backed_auth_context_with_trusted_auth(
state,
headers,
uri,
auth_endpoint_signature,
trusted_auth_verified,
)
.await?
{
return Ok(refreshed);
}
@@ -660,19 +677,27 @@ async fn resolve_data_backed_auth_context_cached(
uri: &Uri,
auth_endpoint_signature: Option<&str>,
cache_negative: bool,
trusted_auth_verified: bool,
) -> Result<Option<GatewayControlAuthContext>, GatewayError> {
let Some(cache_key) = cache_key else {
return resolve_data_backed_auth_context(state, headers, uri, auth_endpoint_signature)
.await;
return resolve_data_backed_auth_context_with_trusted_auth(
state,
headers,
uri,
auth_endpoint_signature,
trusted_auth_verified,
)
.await;
};
loop {
match state.auth_context_cache.register_inflight(cache_key) {
AuthContextInflightRegistration::Leader(guard) => {
let resolved = match resolve_data_backed_auth_context(
let resolved = match resolve_data_backed_auth_context_with_trusted_auth(
state,
headers,
uri,
auth_endpoint_signature,
trusted_auth_verified,
)
.await
{
@@ -708,11 +733,12 @@ async fn resolve_data_backed_auth_context_cached(
}
}
AuthContextInflightRegistration::Bypass => {
return resolve_data_backed_auth_context(
return resolve_data_backed_auth_context_with_trusted_auth(
state,
headers,
uri,
auth_endpoint_signature,
trusted_auth_verified,
)
.await;
}
@@ -768,28 +794,39 @@ pub(crate) async fn refresh_execution_runtime_auth_context_with_snapshot(
return Ok((auth_context, None));
}
let verified_api_key_hash = auth_context.verified_api_key_hash.clone();
let snapshot = {
let _permit = state.acquire_auth_snapshot_load_gate().await?;
state
.data
.read_auth_api_key_snapshot_strong(
&auth_context.user_id,
&auth_context.api_key_id,
current_unix_secs(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
if let Some(key_hash) = verified_api_key_hash.as_ref() {
state
.data
.read_auth_api_key_snapshot_by_key_hash_strong(
key_hash.as_str(),
current_unix_secs(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state
.data
.read_auth_api_key_snapshot_strong(
&auth_context.user_id,
&auth_context.api_key_id,
current_unix_secs(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
}
};
let Some(snapshot) = snapshot else {
let mut denied = auth_context;
denied.access_allowed = false;
denied.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey);
denied.balance_remaining = None;
return Ok((denied, None));
return Ok((deny_refreshed_auth_context(auth_context), None));
};
if snapshot.user_id != auth_context.user_id || snapshot.api_key_id != auth_context.api_key_id {
return Ok((deny_refreshed_auth_context(auth_context), None));
};
let wallet_access = resolve_wallet_auth_gate_uncached(state, &snapshot).await?;
let refreshed = build_data_backed_auth_context(
let mut refreshed = build_data_backed_auth_context(
state,
snapshot.clone(),
auth_endpoint_signature,
@@ -798,9 +835,19 @@ pub(crate) async fn refresh_execution_runtime_auth_context_with_snapshot(
wallet_access,
)
.await;
refreshed.verified_api_key_hash = verified_api_key_hash;
Ok((refreshed, Some(snapshot)))
}
fn deny_refreshed_auth_context(
mut auth_context: GatewayControlAuthContext,
) -> GatewayControlAuthContext {
auth_context.access_allowed = false;
auth_context.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey);
auth_context.balance_remaining = None;
auth_context
}
fn put_cached_auth_context(
state: &AppState,
cache_key: String,
@@ -913,6 +960,23 @@ pub(super) async fn resolve_data_backed_auth_context(
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: Option<&str>,
) -> Result<Option<GatewayControlAuthContext>, GatewayError> {
resolve_data_backed_auth_context_with_trusted_auth(
state,
headers,
uri,
auth_endpoint_signature,
cfg!(test),
)
.await
}
async fn resolve_data_backed_auth_context_with_trusted_auth(
state: &AppState,
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: Option<&str>,
trusted_auth_verified: bool,
) -> Result<Option<GatewayControlAuthContext>, GatewayError> {
let Some(signature) = auth_endpoint_signature
.map(str::trim)
@@ -923,7 +987,12 @@ pub(super) async fn resolve_data_backed_auth_context(
if !state.has_auth_api_key_reader() {
return Ok(None);
}
let extracted = extract_request_credentials(headers, uri, signature);
let extracted = extract_request_credentials_with_trusted_auth(
headers,
uri,
signature,
trusted_auth_verified,
);
let principal = derive_principal_candidate(&extracted);
let now_unix_secs = current_unix_secs();
@@ -955,6 +1024,7 @@ pub(super) async fn resolve_data_backed_auth_context(
local_rejection: Some(GatewayLocalAuthRejection::InvalidApiKey),
allowed_models: None,
ip_rules: None,
verified_api_key_hash: None,
}));
};
@@ -963,17 +1033,17 @@ pub(super) async fn resolve_data_backed_auth_context(
.await;
let wallet_access = resolve_wallet_auth_gate_uncached(state, &snapshot).await?;
Ok(Some(
build_data_backed_auth_context(
state,
snapshot,
signature,
None,
None,
wallet_access,
)
.await,
))
let mut auth_context = build_data_backed_auth_context(
state,
snapshot,
signature,
None,
None,
wallet_access,
)
.await;
auth_context.verified_api_key_hash = Some(VerifiedApiKeyHash::new(key_hash));
Ok(Some(auth_context))
}
Some(GatewayPrincipalCandidate::DeferredBearerToken { raw, carrier }) => {
if let Some(auth_context) = resolve_antigravity_bearer_bridge_auth_context(
@@ -1068,6 +1138,7 @@ async fn resolve_antigravity_bearer_bridge_auth_context(
local_rejection: Some(GatewayLocalAuthRejection::InvalidApiKey),
allowed_models: None,
ip_rules: None,
verified_api_key_hash: None,
}));
};
@@ -1127,6 +1198,7 @@ async fn resolve_trusted_auth_context(
local_rejection: Some(GatewayLocalAuthRejection::InvalidApiKey),
allowed_models: None,
ip_rules: None,
verified_api_key_hash: None,
}));
};
@@ -1158,9 +1230,7 @@ async fn build_data_backed_auth_context(
let invalid_api_key = !snapshot.user_is_active
|| snapshot.user_is_deleted
|| !snapshot.api_key_is_active
|| snapshot
.api_key_expires_at_unix_secs
.is_some_and(|expires_at| expires_at < current_unix_secs());
|| api_key_is_expired(snapshot.api_key_expires_at_unix_secs, current_unix_secs());
let locked_api_key = snapshot.api_key_is_locked && !snapshot.api_key_is_standalone;
let key_access_allowed = header_access_allowed
.map(|value| value && snapshot.currently_usable)
@@ -1225,9 +1295,14 @@ async fn build_data_backed_auth_context(
local_rejection,
allowed_models,
ip_rules: snapshot.api_key_ip_rules,
verified_api_key_hash: None,
}
}
fn api_key_is_expired(expires_at_unix_secs: Option<u64>, now_unix_secs: u64) -> bool {
expires_at_unix_secs.is_some_and(|expires_at| expires_at <= now_unix_secs)
}
fn contains_api_format_or_alias(items: &[String], target: &str) -> bool {
items.iter().any(|item| api_format_matches(item, target))
}
@@ -1282,18 +1357,21 @@ async fn auth_snapshot_allows_requested_provider(
return true;
}
if !state.has_provider_catalog_data_reader() {
return true;
debug!(
"deny requested provider {}: provider catalog is unavailable for allowlist resolution",
requested_provider
);
return false;
}
let providers = match state.list_provider_catalog_providers(true).await {
Ok(value) => value,
Err(err) => {
debug!(
"skip local provider auth gate for requested provider {}: provider catalog lookup failed: {:?}",
requested_provider,
err
warn!(
"deny requested provider {}: provider catalog lookup failed: {:?}",
requested_provider, err
);
return true;
return false;
}
};
@@ -1331,11 +1409,11 @@ async fn auth_snapshot_allows_requested_provider(
{
Ok(value) => value,
Err(err) => {
debug!(
"skip local provider auth gate for requested provider {}: provider endpoint lookup failed: {:?}",
warn!(
"deny requested provider {}: provider endpoint lookup failed: {:?}",
requested_provider, err
);
return true;
return false;
}
};
@@ -1426,7 +1504,8 @@ mod tests {
use std::time::Duration;
use aether_data::repository::auth::{
AuthApiKeyWriteRepository, InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
AuthApiKeyWriteRepository, CreateUserApiKeyRecord, InMemoryAuthApiKeySnapshotRepository,
StoredAuthApiKeySnapshot,
};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::wallet::{
@@ -1441,9 +1520,10 @@ mod tests {
use futures_util::future::join_all;
use super::{
get_cached_auth_context, resolve_control_decision_auth, resolve_data_backed_auth_context,
resolve_execution_runtime_auth_context, ControlDecisionAuthResolution,
GatewayLocalAuthRejection,
api_key_is_expired, get_cached_auth_context,
refresh_execution_runtime_auth_context_with_snapshot, resolve_control_decision_auth,
resolve_data_backed_auth_context, resolve_execution_runtime_auth_context,
ControlDecisionAuthResolution, GatewayLocalAuthRejection,
};
use crate::control::auth::credentials::{build_auth_context_cache_key, hash_api_key};
use crate::control::GatewayControlDecision;
@@ -1481,6 +1561,14 @@ mod tests {
path.parse().expect("uri should parse")
}
#[test]
fn api_key_expiry_is_inclusive_at_the_declared_second() {
assert!(!api_key_is_expired(None, 100));
assert!(!api_key_is_expired(Some(101), 100));
assert!(api_key_is_expired(Some(100), 100));
assert!(api_key_is_expired(Some(99), 100));
}
fn sample_provider(id: &str, name: &str, provider_type: &str) -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
id.to_string(),
@@ -1769,6 +1857,97 @@ mod tests {
assert_eq!(repository.touch_count("key-1"), 1);
}
#[tokio::test]
async fn long_lived_refresh_rejects_same_ids_recreated_with_a_different_credential() {
let old_api_key = "sk-old-websocket-credential";
let new_api_key = "sk-new-websocket-credential";
let old_key_hash = hash_api_key(old_api_key);
let new_key_hash = hash_api_key(new_api_key);
let mut old_snapshot = sample_snapshot("key-stable-id", "user-stable-id");
old_snapshot.user_allowed_api_formats = Some(vec!["openai:responses".to_string()]);
old_snapshot.api_key_allowed_api_formats = Some(vec!["openai:responses".to_string()]);
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(old_key_hash.clone()),
old_snapshot,
)]));
let data = GatewayDataState::with_auth_api_key_repository_for_tests(repository.clone());
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data);
let mut headers = HeaderMap::new();
headers.insert(
http::header::AUTHORIZATION,
format!("Bearer {old_api_key}").parse().unwrap(),
);
let original = resolve_data_backed_auth_context(
&state,
&headers,
&uri("/v1/responses"),
Some("openai:responses"),
)
.await
.expect("initial auth resolution should succeed")
.expect("the old API key should authenticate");
assert!(original.access_allowed);
assert!(original.verified_api_key_hash.is_some());
assert!(
!format!("{original:?}").contains(&old_key_hash),
"the credential verifier must stay redacted from Debug output"
);
assert!(repository
.delete_user_api_key("user-stable-id", "key-stable-id")
.await
.expect("old API key deletion should succeed"));
repository
.create_user_api_key(CreateUserApiKeyRecord {
user_id: "user-stable-id".to_string(),
api_key_id: "key-stable-id".to_string(),
key_hash: new_key_hash,
key_encrypted: None,
name: Some("restored-with-new-secret".to_string()),
allowed_providers: Some(vec!["openai".to_string()]),
allowed_api_formats: Some(vec!["openai:responses".to_string()]),
allowed_models: Some(vec!["gpt-4.1".to_string()]),
ip_rules: None,
rate_limit: 60,
concurrent_limit: Some(5),
force_capabilities: None,
feature_settings: None,
is_active: true,
expires_at_unix_secs: Some(4_102_444_800),
auto_delete_on_expiry: false,
total_requests: 0,
total_tokens: 0,
total_cost_usd: 0.0,
})
.await
.expect("same-ID API key recreation should resolve")
.expect("same-ID API key recreation should persist");
let (refreshed, snapshot) = refresh_execution_runtime_auth_context_with_snapshot(
&state,
original,
Some("openai:responses"),
)
.await
.expect("long-lived auth refresh should resolve");
assert!(!refreshed.access_allowed);
assert_eq!(
refreshed.local_rejection,
Some(GatewayLocalAuthRejection::InvalidApiKey)
);
assert!(snapshot.is_none());
assert_eq!(repository.key_hash_lookup_count(&old_key_hash), 1);
assert_eq!(
repository.snapshot_lookup_count("key-stable-id"),
0,
"a bound long-lived credential must not fall back to identity-only lookup"
);
}
#[tokio::test]
async fn control_auth_context_singleflights_concurrent_cache_misses() {
let api_key = "sk-test-concurrent-auth-miss";
@@ -2396,6 +2575,44 @@ mod tests {
assert_eq!(auth_context.local_rejection, None);
}
#[tokio::test]
async fn data_backed_auth_context_denies_unresolved_provider_id_without_catalog_reader() {
let api_key = "sk-test-provider-no-catalog";
let mut snapshot = sample_snapshot("key-no-catalog", "user-no-catalog");
snapshot.user_allowed_providers = Some(vec!["provider-custom-claude".to_string()]);
snapshot.api_key_allowed_providers = Some(vec!["provider-custom-claude".to_string()]);
snapshot.user_allowed_api_formats = None;
snapshot.api_key_allowed_api_formats = None;
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key(api_key)),
snapshot,
)]));
let data = GatewayDataState::with_auth_api_key_reader_for_tests(repository);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data);
let mut headers = HeaderMap::new();
headers.insert("x-api-key", api_key.parse().unwrap());
let auth_context = resolve_data_backed_auth_context(
&state,
&headers,
&uri("/v1/messages"),
Some("claude:messages"),
)
.await
.expect("resolution should succeed")
.expect("auth context should exist");
assert!(!auth_context.access_allowed);
assert_eq!(
auth_context.local_rejection,
Some(GatewayLocalAuthRejection::ProviderNotAllowed {
provider: "claude".to_string(),
})
);
}
#[tokio::test]
async fn due_antigravity_bearer_refresh_observes_cross_node_allowlist_revocation() {
let raw_bearer = "google-oauth-access-token-revoked-cross-node";
+91 -3
View File
@@ -44,7 +44,7 @@ pub(super) struct GatewayTrustedAdminHeaders {
pub(super) management_token_id: Option<String>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
#[derive(Clone, Default, PartialEq, Eq)]
pub(super) struct GatewayCredentialBundle {
pub(super) authorization_bearer: Option<String>,
pub(super) x_api_key: Option<String>,
@@ -54,7 +54,25 @@ pub(super) struct GatewayCredentialBundle {
pub(super) cookie_header: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
impl std::fmt::Debug for GatewayCredentialBundle {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let redacted = |value: &Option<String>| value.as_ref().map(|_| "[REDACTED]");
formatter
.debug_struct("GatewayCredentialBundle")
.field(
"authorization_bearer",
&redacted(&self.authorization_bearer),
)
.field("x_api_key", &redacted(&self.x_api_key))
.field("api_key", &redacted(&self.api_key))
.field("x_goog_api_key", &redacted(&self.x_goog_api_key))
.field("query_key", &redacted(&self.query_key))
.field("cookie_header", &redacted(&self.cookie_header))
.finish()
}
}
#[derive(Clone, PartialEq, Eq)]
pub(super) enum GatewayPrimaryCredential {
ProviderApiKey {
raw: String,
@@ -70,6 +88,21 @@ pub(super) enum GatewayPrimaryCredential {
},
}
impl std::fmt::Debug for GatewayPrimaryCredential {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let (variant, carrier) = match self {
Self::ProviderApiKey { carrier, .. } => ("ProviderApiKey", carrier),
Self::BearerToken { carrier, .. } => ("BearerToken", carrier),
Self::CookieHeader { carrier, .. } => ("CookieHeader", carrier),
};
formatter
.debug_struct(variant)
.field("raw", &"[REDACTED]")
.field("carrier", carrier)
.finish()
}
}
#[derive(Debug, Clone, PartialEq)]
pub(super) struct GatewayExtractedCredentials {
pub(super) trusted_headers: Option<GatewayTrustedAuthHeaders>,
@@ -78,7 +111,7 @@ pub(super) struct GatewayExtractedCredentials {
pub(super) primary: Option<GatewayPrimaryCredential>,
}
#[derive(Debug, Clone, PartialEq)]
#[derive(Clone, PartialEq)]
pub(super) enum GatewayPrincipalCandidate {
TrustedHeaders(GatewayTrustedAuthHeaders),
ApiKeyHash {
@@ -94,3 +127,58 @@ pub(super) enum GatewayPrincipalCandidate {
carrier: GatewayCredentialCarrier,
},
}
impl std::fmt::Debug for GatewayPrincipalCandidate {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::TrustedHeaders(headers) => formatter
.debug_tuple("TrustedHeaders")
.field(headers)
.finish(),
Self::ApiKeyHash { carrier, .. } => formatter
.debug_struct("ApiKeyHash")
.field("key_hash", &"[REDACTED]")
.field("carrier", carrier)
.finish(),
Self::DeferredBearerToken { carrier, .. } => formatter
.debug_struct("DeferredBearerToken")
.field("raw", &"[REDACTED]")
.field("carrier", carrier)
.finish(),
Self::DeferredCookieHeader { carrier, .. } => formatter
.debug_struct("DeferredCookieHeader")
.field("raw", &"[REDACTED]")
.field("carrier", carrier)
.finish(),
}
}
}
#[cfg(test)]
mod debug_redaction_tests {
use super::{GatewayCredentialBundle, GatewayCredentialCarrier, GatewayPrimaryCredential};
#[test]
fn gateway_credential_debug_output_redacts_raw_authorization_values() {
let bundle = GatewayCredentialBundle {
authorization_bearer: Some("bundle-bearer-canary".to_string()),
api_key: Some("bundle-api-key-canary".to_string()),
cookie_header: Some("bundle-cookie-canary".to_string()),
..GatewayCredentialBundle::default()
};
let primary = GatewayPrimaryCredential::ProviderApiKey {
raw: "primary-api-key-canary".to_string(),
carrier: GatewayCredentialCarrier::ApiKey,
};
let debug = format!("{bundle:?} {primary:?}");
assert!(debug.contains("[REDACTED]"));
for secret in [
"bundle-bearer-canary",
"bundle-api-key-canary",
"bundle-cookie-canary",
"primary-api-key-canary",
] {
assert!(!debug.contains(secret), "debug output leaked {secret}");
}
}
}
File diff suppressed because it is too large Load Diff
+10 -7
View File
@@ -8,9 +8,10 @@ mod public;
mod route;
pub(crate) use auth::{
execution_plan_balance_capacity_rejection, extract_requested_model,
refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot,
request_model_local_rejection, resolve_execution_runtime_auth_context,
estimate_execution_plan_cost_upper_bound_usd, execution_plan_balance_capacity_rejection,
extract_requested_model, refresh_execution_runtime_auth_context,
refresh_execution_runtime_auth_context_with_snapshot, request_model_local_rejection,
resolve_execution_runtime_auth_context, resolve_local_admin_session_principal,
should_buffer_request_for_local_auth, trusted_auth_local_rejection,
GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayCredentialCarrier,
GatewayLocalAuthRejection,
@@ -18,14 +19,16 @@ pub(crate) use auth::{
pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control};
pub(crate) use management_token_permissions::{
all_assignable_management_token_permissions,
audit_admin_read_only_management_token_permissions,
audit_admin_read_only_management_token_permissions, legacy_full_management_token_permissions,
management_token_permission_catalog_payload, management_token_permission_keys_from_value,
management_token_permission_mode_and_summary,
management_token_permissions_cover_all_assignable_permissions,
management_token_permission_mode_and_summary, management_token_principal_has_permission,
management_token_required_permission, normalize_assignable_management_token_permissions,
read_only_management_token_permissions, validate_management_token_admin_route_permission,
};
pub(crate) use public::{resolve_public_request_context, GatewayPublicRequestContext};
pub(crate) use public::{
resolve_public_request_context, resolve_public_request_context_with_trusted_auth,
resolve_public_request_context_without_trusted_auth, GatewayPublicRequestContext,
};
#[cfg(test)]
pub(crate) use route::classify_control_route;
pub(crate) use route::{resolve_control_route, GatewayControlDecision};
+41 -1
View File
@@ -2,7 +2,9 @@ use axum::http::Uri;
use crate::{AppState, GatewayError};
use super::{resolve_control_route, GatewayControlDecision};
use super::{
resolve_control_route, route::resolve_control_route_with_trusted_auth, GatewayControlDecision,
};
pub(crate) type GatewayPublicRequestContext =
aether_gateway_control::PublicRequestContext<GatewayControlDecision>;
@@ -23,3 +25,41 @@ pub(crate) async fn resolve_public_request_context(
control_decision,
))
}
pub(crate) async fn resolve_public_request_context_with_trusted_auth(
state: &AppState,
method: &http::Method,
uri: &Uri,
headers: &http::HeaderMap,
trace_id: &str,
) -> Result<GatewayPublicRequestContext, GatewayError> {
let control_decision =
resolve_control_route_with_trusted_auth(state, method, uri, headers, trace_id, true)
.await?;
Ok(GatewayPublicRequestContext::from_request_parts(
trace_id,
method,
uri,
headers,
control_decision,
))
}
pub(crate) async fn resolve_public_request_context_without_trusted_auth(
state: &AppState,
method: &http::Method,
uri: &Uri,
headers: &http::HeaderMap,
trace_id: &str,
) -> Result<GatewayPublicRequestContext, GatewayError> {
let control_decision =
resolve_control_route_with_trusted_auth(state, method, uri, headers, trace_id, false)
.await?;
Ok(GatewayPublicRequestContext::from_request_parts(
trace_id,
method,
uri,
headers,
control_decision,
))
}
@@ -797,6 +797,18 @@ pub(super) fn classify_admin_operations_family_route(
"admin:users",
false,
))
} else if method == http::Method::DELETE
&& normalized_path_no_trailing.starts_with("/api/admin/users/")
&& normalized_path_no_trailing.contains("/billing/entitlements/")
&& normalized_path_no_trailing.matches('/').count() == 7
{
Some(classified(
"admin_proxy",
"users_manage",
"revoke_user_billing_entitlement",
"admin:users",
false,
))
} else if method == http::Method::GET
&& normalized_path.starts_with("/api/admin/users/")
&& normalized_path.ends_with("/sessions")
@@ -5,19 +5,21 @@ pub(super) fn classify_internal_route(
method: &http::Method,
normalized_path: &str,
) -> Option<ClassifiedRoute> {
if method == http::Method::POST && normalized_path.starts_with("/api/internal/gateway/") {
let route_kind = match normalized_path {
"/api/internal/gateway/resolve" => "resolve",
"/api/internal/gateway/auth-context" => "auth_context",
"/api/internal/gateway/decision-sync" => "decision_sync",
"/api/internal/gateway/decision-stream" => "decision_stream",
"/api/internal/gateway/plan-sync" => "plan_sync",
"/api/internal/gateway/plan-stream" => "plan_stream",
"/api/internal/gateway/report-sync" => "report_sync",
"/api/internal/gateway/report-stream" => "report_stream",
"/api/internal/gateway/finalize-sync" => "finalize_sync",
"/api/internal/gateway/execute-sync" => "execute_sync",
"/api/internal/gateway/execute-stream" => "execute_stream",
if normalized_path == "/api/internal/gateway"
|| normalized_path.starts_with("/api/internal/gateway/")
{
let route_kind = match (method, normalized_path) {
(&http::Method::POST, "/api/internal/gateway/resolve") => "resolve",
(&http::Method::POST, "/api/internal/gateway/auth-context") => "auth_context",
(&http::Method::POST, "/api/internal/gateway/decision-sync") => "decision_sync",
(&http::Method::POST, "/api/internal/gateway/decision-stream") => "decision_stream",
(&http::Method::POST, "/api/internal/gateway/plan-sync") => "plan_sync",
(&http::Method::POST, "/api/internal/gateway/plan-stream") => "plan_stream",
(&http::Method::POST, "/api/internal/gateway/report-sync") => "report_sync",
(&http::Method::POST, "/api/internal/gateway/report-stream") => "report_stream",
(&http::Method::POST, "/api/internal/gateway/finalize-sync") => "finalize_sync",
(&http::Method::POST, "/api/internal/gateway/execute-sync") => "execute_sync",
(&http::Method::POST, "/api/internal/gateway/execute-stream") => "execute_stream",
_ => "unhandled",
};
Some(classified(
+22 -2
View File
@@ -11,7 +11,7 @@ mod oauth;
mod public_support;
use super::auth::{
resolve_control_decision_auth, resolve_gateway_credential_carrier,
resolve_control_decision_auth_with_trusted_auth, resolve_gateway_credential_carrier,
ControlDecisionAuthResolution, GatewayCredentialCarrier,
};
use super::{GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayLocalAuthRejection};
@@ -175,6 +175,17 @@ pub(crate) async fn resolve_control_route(
uri: &Uri,
headers: &http::HeaderMap,
trace_id: &str,
) -> Result<Option<GatewayControlDecision>, GatewayError> {
resolve_control_route_with_trusted_auth(state, method, uri, headers, trace_id, cfg!(test)).await
}
pub(crate) async fn resolve_control_route_with_trusted_auth(
state: &AppState,
method: &http::Method,
uri: &Uri,
headers: &http::HeaderMap,
trace_id: &str,
trusted_auth_verified: bool,
) -> Result<Option<GatewayControlDecision>, GatewayError> {
let Some(mut decision) = classify_control_route(method, uri, headers) else {
return Ok(None);
@@ -185,7 +196,16 @@ pub(crate) async fn resolve_control_route(
crate::system_features::ModelDirectivePolicySnapshot::load(state).await;
}
match resolve_control_decision_auth(state, headers, uri, trace_id, decision).await? {
match resolve_control_decision_auth_with_trusted_auth(
state,
headers,
uri,
trace_id,
decision,
trusted_auth_verified,
)
.await?
{
ControlDecisionAuthResolution::Resolved(decision) => Ok(Some(decision)),
}
}
@@ -285,7 +285,7 @@ pub(super) fn classify_oauth_route(
"admin_proxy",
"provider_oauth_manage",
"batch_import_oauth",
"admin:pool",
"admin:provider_oauth",
false,
))
} else if method == http::Method::POST
@@ -296,7 +296,7 @@ pub(super) fn classify_oauth_route(
"admin_proxy",
"provider_oauth_manage",
"start_batch_import_oauth_task",
"admin:pool",
"admin:provider_oauth",
false,
))
} else if method == http::Method::GET
@@ -307,7 +307,7 @@ pub(super) fn classify_oauth_route(
"admin_proxy",
"provider_oauth_manage",
"get_batch_import_task_status",
"admin:pool",
"admin:provider_oauth",
false,
))
} else if method == http::Method::POST
@@ -197,18 +197,22 @@ pub(super) fn classify_public_support_route(
"public:auth",
false,
))
} else if matches!(method, &http::Method::GET | &http::Method::POST)
// Authentication state-changing endpoints must never be dispatched for GET.
// Besides violating HTTP method semantics, accepting GET here would allow
// browser prefetches/cross-site requests to trigger login, refresh, logout,
// registration, or verification side effects. `/me` is the sole read route.
} else if (method == http::Method::POST
&& matches!(
normalized_path,
"/api/auth/login"
| "/api/auth/refresh"
| "/api/auth/register"
| "/api/auth/me"
| "/api/auth/logout"
| "/api/auth/send-verification-code"
| "/api/auth/verify-email"
| "/api/auth/verification-status"
)
))
|| (method == http::Method::GET && normalized_path == "/api/auth/me")
{
let route_kind = match normalized_path {
"/api/auth/login" => "login",
@@ -520,6 +524,65 @@ pub(super) fn classify_public_support_route(
"aether:ccswitch_usage",
false,
))
} else if method == http::Method::POST
&& matches!(normalized_path, "/api/vscodex/pair" | "/api/vscodex/pair/")
{
Some(classified(
"public_support",
"vscodex",
"pairing_exchange",
"public:vscodex",
false,
))
} else if method == http::Method::GET
&& matches!(
normalized_path,
"/api/users/me/vscodex/devices" | "/api/users/me/vscodex/devices/"
)
{
Some(classified(
"public_support",
"users_me",
"vscodex_devices_list",
"user:self",
false,
))
} else if method == http::Method::POST
&& matches!(
normalized_path,
"/api/users/me/vscodex/pairings" | "/api/users/me/vscodex/pairings/"
)
{
Some(classified(
"public_support",
"users_me",
"vscodex_pairing_create",
"user:self",
false,
))
} else if method == http::Method::POST
&& matches!(
normalized_path,
"/api/users/me/vscodex/ws-tickets" | "/api/users/me/vscodex/ws-tickets/"
)
{
Some(classified(
"public_support",
"users_me",
"vscodex_ws_ticket_create",
"user:self",
false,
))
} else if method == http::Method::DELETE
&& has_single_segment_after_prefix(normalized_path, "/api/users/me/vscodex/devices/")
{
Some(classified(
"public_support",
"users_me",
"vscodex_device_delete",
"user:self",
false,
))
} else if method == http::Method::GET
&& matches!(
normalized_path,
@@ -628,7 +628,7 @@ fn classifies_admin_management_token_write_routes_and_permission_catalog() {
http::Method::POST,
"/api/admin/management-tokens",
"create_token",
"admin:management_tokens:write",
"admin:management_tokens:admin",
),
(
http::Method::PUT,
@@ -640,7 +640,7 @@ fn classifies_admin_management_token_write_routes_and_permission_catalog() {
http::Method::POST,
"/api/admin/management-tokens/token-123/regenerate",
"regenerate_token",
"admin:management_tokens:write",
"admin:management_tokens:admin",
),
];
@@ -68,11 +68,11 @@ fn classifies_admin_provider_oauth_batch_import_task_status_as_admin_proxy_route
);
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("admin:pool")
Some("admin:provider_oauth")
);
assert_eq!(
management_token_required_permission(&http::Method::GET, &decision).as_deref(),
Some("admin:pool:read")
Some("admin:provider_oauth:read")
);
assert!(!decision.is_execution_runtime_candidate());
}
@@ -86,42 +86,42 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() {
"/api/admin/provider-oauth/keys/key-123/complete",
"complete_key_oauth",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
(
http::Method::POST,
"/api/admin/provider-oauth/keys/key-123/refresh",
"refresh_key_oauth",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
(
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-123/complete",
"complete_provider_oauth",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
(
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-123/import-refresh-token",
"import_refresh_token",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
(
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-123/cookie-authorize",
"cookie_authorize",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
(
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-123/cookie-authorize/tasks",
"start_cookie_authorize_task",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
(
http::Method::GET,
@@ -135,7 +135,7 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() {
"/api/admin/provider-oauth/providers/provider-123/agent-identity-import/tasks",
"start_agent_identity_import_task",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
(
http::Method::GET,
@@ -148,22 +148,22 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() {
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-123/batch-import",
"batch_import_oauth",
"admin:pool",
"admin:pool:write",
"admin:provider_oauth",
"admin:provider_oauth:admin",
),
(
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-123/batch-import/tasks",
"start_batch_import_oauth_task",
"admin:pool",
"admin:pool:write",
"admin:provider_oauth",
"admin:provider_oauth:admin",
),
(
http::Method::GET,
"/api/admin/provider-oauth/providers/provider-123/batch-import/tasks/task-123",
"get_batch_import_task_status",
"admin:pool",
"admin:pool:read",
"admin:provider_oauth",
"admin:provider_oauth:read",
),
(
http::Method::POST,
@@ -177,7 +177,7 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() {
"/api/admin/provider-oauth/providers/provider-123/device-poll",
"device_poll",
"admin:provider_oauth",
"admin:provider_oauth:write",
"admin:provider_oauth:admin",
),
] {
let uri: Uri = path.parse().expect("uri should parse");
@@ -103,6 +103,22 @@ fn classifies_admin_user_billing_routes_as_admin_proxy_route() {
Some("admin:users")
);
let revoke_uri: Uri = "/api/admin/users/user-1/billing/entitlements/entitlement-1"
.parse()
.expect("uri should parse");
let revoke = classify_control_route(&http::Method::DELETE, &revoke_uri, &headers)
.expect("route should classify");
assert_eq!(revoke.route_class.as_deref(), Some("admin_proxy"));
assert_eq!(revoke.route_family.as_deref(), Some("users_manage"));
assert_eq!(
revoke.route_kind.as_deref(),
Some("revoke_user_billing_entitlement")
);
assert_eq!(
revoke.auth_endpoint_signature.as_deref(),
Some("admin:users")
);
let context = GatewayPublicRequestContext::from_request_parts(
"trace-user-billing-grant",
&http::Method::POST,
@@ -440,6 +440,26 @@ fn classifies_users_me_routes_as_public_support_route() {
"/api/users/me/available-models",
"available_models",
),
(
http::Method::GET,
"/api/users/me/vscodex/devices",
"vscodex_devices_list",
),
(
http::Method::POST,
"/api/users/me/vscodex/pairings",
"vscodex_pairing_create",
),
(
http::Method::DELETE,
"/api/users/me/vscodex/devices/device-1",
"vscodex_device_delete",
),
(
http::Method::POST,
"/api/users/me/vscodex/ws-tickets",
"vscodex_ws_ticket_create",
),
(
http::Method::PUT,
"/api/users/me/model-capabilities",
@@ -496,6 +516,49 @@ fn classifies_users_me_routes_as_public_support_route() {
}
}
#[test]
fn vscodex_post_routes_buffer_request_body() {
let headers = headers(&[]);
for path in [
"/api/vscodex/pair",
"/api/users/me/vscodex/pairings",
"/api/users/me/vscodex/ws-tickets",
] {
let uri: Uri = path.parse().expect("uri should parse");
let decision = classify_control_route(&http::Method::POST, &uri, &headers)
.expect("route should classify");
let context = GatewayPublicRequestContext::from_request_parts(
"trace-vscodex",
&http::Method::POST,
&uri,
&headers,
Some(decision),
);
assert!(
local_proxy_route_requires_buffered_body(&context),
"{path} should buffer its JSON body"
);
}
}
#[test]
fn classifies_public_vscodex_pairing_exchange() {
let headers = headers(&[]);
let uri: Uri = "/api/vscodex/pair".parse().expect("uri should parse");
let decision =
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
assert_eq!(decision.route_class.as_deref(), Some("public_support"));
assert_eq!(decision.route_family.as_deref(), Some("vscodex"));
assert_eq!(decision.route_kind.as_deref(), Some("pairing_exchange"));
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("public:vscodex")
);
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_ccswitch_usage_as_api_key_public_support_route() {
let headers = headers(&[]);
@@ -884,6 +947,34 @@ fn classifies_auth_routes_as_public_support_route() {
}
}
#[test]
fn does_not_classify_state_changing_auth_routes_for_get() {
for path in [
"/api/auth/login",
"/api/auth/refresh",
"/api/auth/register",
"/api/auth/logout",
"/api/auth/send-verification-code",
"/api/auth/verify-email",
"/api/auth/verification-status",
] {
let headers = headers(&[]);
let uri: Uri = path.parse().expect("uri should parse");
assert!(
classify_control_route(&http::Method::GET, &uri, &headers).is_none(),
"state-changing auth route {path} must not accept GET"
);
}
let headers = headers(&[]);
let uri: Uri = "/api/auth/me".parse().expect("uri should parse");
assert_eq!(
classify_control_route(&http::Method::GET, &uri, &headers)
.and_then(|decision| decision.route_kind),
Some("me".to_string())
);
}
#[test]
fn classifies_oauth_public_providers_route() {
let headers = headers(&[]);
@@ -172,7 +172,7 @@ mod tests {
candidates: vec![DecisionTraceCandidate {
candidate: sample_candidate("req-1"),
provider_name: Some("OpenAI".to_string()),
provider_website: Some("https://openai.com".to_string()),
provider_website: Some("https://openai.com/".to_string()),
provider_type: Some("custom".to_string()),
provider_priority: Some(0),
provider_keep_priority_on_conversion: Some(false),
File diff suppressed because it is too large Load Diff
+275 -47
View File
@@ -1,23 +1,42 @@
use super::{
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, PublicHealthStatusCount, PublicHealthTimelineBucket,
StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint,
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
StoredRequestCandidate, UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate, PublicHealthStatusCount,
PublicHealthTimelineBucket, StoredGeminiFileMapping, StoredGeminiFileMappingListPage,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
};
fn sanitize_request_candidate_rows(
mut candidates: Vec<StoredRequestCandidate>,
) -> Vec<StoredRequestCandidate> {
for candidate in &mut candidates {
candidate.sanitize_sensitive_diagnostics();
}
candidates
}
fn sanitize_request_candidate_row(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
candidate.sanitize_sensitive_diagnostics();
candidate
}
impl GatewayDataState {
pub(crate) async fn list_request_candidates_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_reader {
Some(repository) => repository.list_by_request_id(request_id).await,
Some(repository) => repository
.list_by_request_id(request_id)
.await
.map(sanitize_request_candidate_rows),
None => Ok(Vec::new()),
}
}
@@ -27,7 +46,10 @@ impl GatewayDataState {
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_reader {
Some(repository) => repository.list_attempted_by_request_id(request_id).await,
Some(repository) => repository
.list_attempted_by_request_id(request_id)
.await
.map(sanitize_request_candidate_rows),
None => Ok(Vec::new()),
}
}
@@ -38,7 +60,10 @@ impl GatewayDataState {
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_reader {
Some(repository) => repository.list_by_provider_id(provider_id, limit).await,
Some(repository) => repository
.list_by_provider_id(provider_id, limit)
.await
.map(sanitize_request_candidate_rows),
None => Ok(Vec::new()),
}
}
@@ -48,7 +73,10 @@ impl GatewayDataState {
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_reader {
Some(repository) => repository.list_recent(limit).await,
Some(repository) => repository
.list_recent(limit)
.await
.map(sanitize_request_candidate_rows),
None => Ok(Vec::new()),
}
}
@@ -60,11 +88,10 @@ impl GatewayDataState {
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_reader {
Some(repository) => {
repository
.list_finalized_by_endpoint_ids_since(endpoint_ids, since_unix_secs, limit)
.await
}
Some(repository) => repository
.list_finalized_by_endpoint_ids_since(endpoint_ids, since_unix_secs, limit)
.await
.map(sanitize_request_candidate_rows),
None => Ok(Vec::new()),
}
}
@@ -108,14 +135,19 @@ impl GatewayDataState {
pub(crate) async fn upsert_request_candidate(
&self,
candidate: UpsertRequestCandidateRecord,
mut candidate: UpsertRequestCandidateRecord,
) -> Result<Option<StoredRequestCandidate>, DataLayerError> {
candidate.sanitize_for_persistence();
crate::request_diagnostics::observe_db_operation(
"request_candidate_upsert",
self.database_pool_summary(),
async {
match &self.request_candidate_writer {
Some(repository) => repository.upsert(candidate).await.map(Some),
Some(repository) => repository
.upsert(candidate)
.await
.map(sanitize_request_candidate_row)
.map(Some),
None => Ok(None),
}
},
@@ -170,6 +202,16 @@ impl GatewayDataState {
}
}
pub(crate) async fn upsert_gemini_file_mapping_if_owner_matches(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
match &self.gemini_file_mapping_writer {
Some(repository) => repository.upsert_if_owner_matches(record).await,
None => Ok(None),
}
}
pub(crate) async fn list_gemini_file_mappings(
&self,
query: &GeminiFileMappingListQuery,
@@ -183,6 +225,49 @@ impl GatewayDataState {
}
}
pub(crate) async fn find_gemini_file_mapping_by_file_name(
&self,
file_name: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
match &self.gemini_file_mapping_reader {
Some(repository) => repository.find_by_file_name(file_name).await,
None => Ok(None),
}
}
pub(crate) async fn find_active_gemini_file_mapping_for_user(
&self,
file_name: &str,
user_id: &str,
now_unix_secs: u64,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
match &self.gemini_file_mapping_reader {
Some(repository) => {
repository
.find_active_by_file_name_for_user(file_name, user_id, now_unix_secs)
.await
}
None => Ok(None),
}
}
pub(crate) async fn find_active_gemini_file_mapping_for_owner(
&self,
file_name: &str,
key_id: &str,
user_id: &str,
now_unix_secs: u64,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
match &self.gemini_file_mapping_reader {
Some(repository) => {
repository
.find_active_by_file_name_for_owner(file_name, key_id, user_id, now_unix_secs)
.await
}
None => Ok(None),
}
}
pub(crate) async fn summarize_gemini_file_mappings(
&self,
now_unix_secs: u64,
@@ -208,6 +293,37 @@ impl GatewayDataState {
}
}
pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_user(
&self,
file_name: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
match &self.gemini_file_mapping_writer {
Some(repository) => {
repository
.delete_by_file_name_for_user(file_name, user_id)
.await
}
None => Ok(false),
}
}
pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_owner(
&self,
file_name: &str,
key_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
match &self.gemini_file_mapping_writer {
Some(repository) => {
repository
.delete_by_file_name_for_owner(file_name, key_id, user_id)
.await
}
None => Ok(false),
}
}
pub(crate) async fn delete_gemini_file_mapping_by_id(
&self,
mapping_id: &str,
@@ -357,38 +473,11 @@ impl GatewayDataState {
}
}
pub(crate) async fn update_provider_catalog_key_oauth_credentials(
&self,
key_id: &str,
encrypted_api_key: &str,
encrypted_auth_config: Option<&str>,
expires_at_unix_secs: Option<u64>,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => {
repository
.update_key_oauth_credentials(
key_id,
encrypted_api_key,
encrypted_auth_config,
expires_at_unix_secs,
)
.await
}
None => Ok(false),
}?;
if updated {
self.clear_provider_catalog_cache();
}
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_oauth_runtime_state(
&self,
key_id: &str,
oauth_invalid_at_unix_secs: Option<u64>,
oauth_invalid_reason: Option<&str>,
encrypted_auth_config_update: Option<&str>,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
@@ -398,7 +487,6 @@ impl GatewayDataState {
key_id,
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config_update,
updated_at_unix_secs,
)
.await
@@ -475,6 +563,30 @@ impl GatewayDataState {
Ok(updated)
}
pub(crate) async fn compare_and_swap_provider_catalog_provider_config(
&self,
update: &ProviderCatalogProviderConfigCasUpdate,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.compare_and_swap_provider_config(update).await,
None => Ok(false),
}?;
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn compare_and_swap_provider_catalog_provider_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.compare_and_swap_provider_proxy(update).await,
None => Ok(false),
}?;
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn delete_provider_catalog_provider(
&self,
provider_id: &str,
@@ -543,6 +655,18 @@ impl GatewayDataState {
Ok(updated)
}
pub(crate) async fn compare_and_swap_provider_catalog_endpoint_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.compare_and_swap_endpoint_proxy(update).await,
None => Ok(false),
}?;
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn delete_provider_catalog_endpoint(
&self,
endpoint_id: &str,
@@ -571,6 +695,32 @@ impl GatewayDataState {
Ok(updated)
}
pub(crate) async fn compare_and_swap_provider_catalog_key_proxy(
&self,
update: &ProviderCatalogProxyCasUpdate,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.compare_and_swap_key_proxy(update).await,
None => Ok(false),
}?;
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn compare_and_swap_provider_catalog_key_credentials(
&self,
update: &ProviderCatalogKeyCredentialsCasUpdate,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.compare_and_swap_key_credentials(update).await,
None => Ok(false),
}?;
// Clear on both outcomes: a CAS miss proves the cached credential
// generation was stale and the retry must observe the winning record.
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
@@ -843,3 +993,81 @@ impl GatewayDataState {
Ok(updated)
}
}
#[cfg(test)]
mod request_candidate_security_tests {
use serde_json::json;
use super::{
sanitize_request_candidate_row, sanitize_request_candidate_rows, StoredRequestCandidate,
};
use aether_data_contracts::repository::candidates::RequestCandidateStatus;
fn untrusted_candidate() -> StoredRequestCandidate {
let mut candidate = StoredRequestCandidate::new(
"candidate-untrusted".to_string(),
"request-1".to_string(),
None,
None,
None,
None,
0,
0,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("key-1".to_string()),
RequestCandidateStatus::Failed,
None,
false,
Some(500),
None,
None,
None,
None,
None,
None,
1,
None,
Some(2),
)
.expect("candidate should build");
candidate.skip_reason = Some("Bearer candidate-secret".to_string());
candidate.error_type = Some("candidate-secret".to_string());
candidate.error_message = Some("Bearer candidate-secret".to_string());
candidate.extra_data = Some(json!({
"gateway_execution_runtime": true,
"request_body": {"token": "candidate-secret"}
}));
candidate.required_capabilities = Some(json!({
"vision": 1,
"tenant_secret": "candidate-secret"
}));
candidate
}
fn assert_candidate_is_sanitized(candidate: &StoredRequestCandidate) {
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
assert!(candidate.error_message.is_none());
assert_eq!(
candidate.extra_data,
Some(json!({"gateway_execution_runtime": true}))
);
assert_eq!(
candidate.required_capabilities,
Some(json!({"vision": true}))
);
assert!(!serde_json::to_string(candidate)
.expect("candidate should serialize")
.contains("candidate-secret"));
}
#[test]
fn gateway_candidate_boundary_sanitizes_repository_rows_and_write_results() {
let candidate = sanitize_request_candidate_row(untrusted_candidate());
assert_candidate_is_sanitized(&candidate);
let candidates = sanitize_request_candidate_rows(vec![untrusted_candidate()]);
assert_candidate_is_sanitized(&candidates[0]);
}
}
@@ -895,6 +895,44 @@ impl GatewayDataState {
.await
}
pub(crate) async fn compare_and_set_system_config_string_value(
&self,
key: &str,
expected: &str,
replacement: &str,
) -> Result<bool, DataLayerError> {
if let Some(values) = &self.system_config_values {
let updated = {
let mut values = values.write().expect("system config values lock");
match values.get_mut(key) {
Some(entry) if entry.value.as_str() == Some(expected) => {
entry.value = serde_json::Value::String(replacement.to_string());
entry.updated_at_unix_secs =
Some(current_system_config_updated_at_unix_secs());
true
}
_ => false,
}
};
self.clear_cached_system_config_value(key);
return Ok(updated);
}
let result = match self.backends.as_ref() {
Some(backends) => {
crate::request_diagnostics::observe_db_operation(
"system_config_compare_and_set",
self.database_pool_summary(),
backends.compare_and_set_system_config_string_value(key, expected, replacement),
)
.await
}
None => Ok(false),
};
self.clear_cached_system_config_value(key);
result
}
pub(crate) async fn upsert_system_config_value(
&self,
key: &str,
@@ -11,7 +11,11 @@ use aether_data_contracts::repository::candidates::DecisionTrace;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput};
use aether_data_contracts::repository::proxy_nodes::ProxyNodeTrafficMutation;
use aether_data_contracts::repository::settlement::{
ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement,
UsageSettlementInput,
};
use aether_data_contracts::repository::usage::{
ProxyNodeCounterDelta, StoredRequestUsageAudit, UpsertUsageRecord, UsageWriteRepository,
};
@@ -34,7 +38,7 @@ const LEGACY_REQUEST_LOG_LEVEL_KEY: &str = "request_log_level";
fn usage_request_record_level_from_value(value: Option<&Value>) -> UsageRequestRecordLevel {
let Some(value) = value.and_then(Value::as_str).map(str::trim) else {
return UsageRequestRecordLevel::Full;
return UsageRequestRecordLevel::Basic;
};
if value.eq_ignore_ascii_case("basic")
@@ -45,7 +49,9 @@ fn usage_request_record_level_from_value(value: Option<&Value>) -> UsageRequestR
{
UsageRequestRecordLevel::Basic
} else {
UsageRequestRecordLevel::Full
// Raw HTTP payload capture is disabled at the runtime boundary. The setting remains
// accepted for compatibility, but no longer authorizes collecting request/response data.
UsageRequestRecordLevel::Basic
}
}
@@ -95,6 +101,14 @@ impl StoredVideoTaskReadSide for GatewayDataState {
) -> Result<Option<StoredVideoTask>, DataLayerError> {
GatewayDataState::find_video_task(self, key).await
}
async fn find_stored_video_task_for_user(
&self,
key: VideoTaskLookupKey<'_>,
user_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
GatewayDataState::find_video_task_for_user(self, key, user_id).await
}
}
#[async_trait]
@@ -219,6 +233,13 @@ impl UsageSettlementWriter for GatewayDataState {
GatewayDataState::has_settlement_writer(self)
}
async fn reconcile_usage_policy_cost(
&self,
input: ReconcileUsagePolicyCostInput,
) -> Result<Option<StoredUsagePolicyCostReservation>, DataLayerError> {
GatewayDataState::reconcile_usage_policy_cost(self, input).await
}
async fn settle_usage(
&self,
input: UsageSettlementInput,
@@ -285,10 +306,26 @@ impl aether_usage_runtime::ManualProxyNodeCounter for GatewayDataState {
failed_delta: i64,
latency_ms: Option<i64>,
) -> Result<(), DataLayerError> {
// This API predates incarnation fences and only receives a node id. Read
// the selected node first so every durable path can bind the delta to
// the observed generation. If the node disappeared, fail closed rather
// than allowing a bare id to target a replacement node.
let Some(node) = self.find_proxy_node(node_id).await? else {
return Ok(());
};
if !node.is_manual {
return Ok(());
}
let expected_tunnel_generation = node.tunnel_generation.trim().to_string();
if expected_tunnel_generation.is_empty() {
return Ok(());
}
if let Some(repository) = &self.usage_writer {
let enqueued = repository
.enqueue_proxy_node_counter_delta(ProxyNodeCounterDelta {
node_id: node_id.to_string(),
expected_tunnel_generation: Some(expected_tunnel_generation.clone()),
total_requests_delta: total_delta,
failed_requests_delta: failed_delta,
dns_failures_delta: 0,
@@ -300,14 +337,25 @@ impl aether_usage_runtime::ManualProxyNodeCounter for GatewayDataState {
}
}
match &self.proxy_node_writer {
Some(repository) => {
repository
.increment_manual_node_requests(node_id, total_delta, failed_delta, latency_ms)
.await
}
None => Ok(()),
// The legacy increment method has no generation argument and would
// re-read the current row, which is vulnerable to an id reuse between
// the read above and the write. Use the fenced traffic mutation as the
// only fallback. It intentionally omits the legacy latency-only field;
// preserving counters safely is more important than an unfenced write.
if let Some(repository) = &self.proxy_node_writer {
let _ = repository
.record_traffic(&ProxyNodeTrafficMutation {
node_id: node_id.to_string(),
expected_tunnel_generation: Some(expected_tunnel_generation),
total_requests_delta: total_delta,
failed_requests_delta: failed_delta,
dns_failures_delta: 0,
stream_errors_delta: 0,
})
.await?;
}
let _ = latency_ms;
Ok(())
}
}
@@ -452,6 +500,20 @@ mod tests {
assert_eq!(level, UsageRequestRecordLevel::Basic);
}
#[tokio::test]
async fn usage_runtime_access_disables_full_http_capture() {
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
"request_record_level".to_string(),
json!("full"),
)]);
let level = UsageRuntimeAccess::request_record_level(&state)
.await
.expect("request record level should read");
assert_eq!(level, UsageRequestRecordLevel::Basic);
}
#[tokio::test]
async fn usage_runtime_access_falls_back_to_legacy_request_log_level_alias() {
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
@@ -467,14 +529,14 @@ mod tests {
}
#[tokio::test]
async fn usage_runtime_access_defaults_missing_request_record_level_to_full() {
async fn usage_runtime_access_defaults_missing_request_record_level_to_basic() {
let state = GatewayDataState::disabled();
let level = UsageRuntimeAccess::request_record_level(&state)
.await
.expect("missing request record level should fall back");
assert_eq!(level, UsageRequestRecordLevel::Full);
assert_eq!(level, UsageRequestRecordLevel::Basic);
}
#[tokio::test]
@@ -488,6 +550,20 @@ mod tests {
.await
.expect("body capture policy should read");
assert_eq!(policy.record_level, UsageRequestRecordLevel::Full);
assert_eq!(policy.record_level, UsageRequestRecordLevel::Basic);
}
#[tokio::test]
async fn usage_runtime_access_fails_closed_for_unknown_record_level() {
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
"request_record_level".to_string(),
json!("everything"),
)]);
let level = UsageRuntimeAccess::request_record_level(&state)
.await
.expect("request record level should read");
assert_eq!(level, UsageRequestRecordLevel::Basic);
}
}
+35 -25
View File
@@ -27,8 +27,8 @@ use aether_data::repository::auth::{
StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
};
use aether_data::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,
AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult,
LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
};
use aether_data::repository::gemini_file_mappings::{
GeminiFileMappingListQuery, GeminiFileMappingReadRepository, GeminiFileMappingStats,
@@ -36,9 +36,10 @@ use aether_data::repository::gemini_file_mappings::{
UpsertGeminiFileMappingRecord,
};
use aether_data::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenListPage, StoredManagementTokenWithUser, UpdateManagementTokenRecord,
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
};
use aether_data::repository::oauth_providers::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
@@ -63,21 +64,24 @@ pub(crate) use aether_data::repository::users::{
use aether_data::repository::wallet::{
AdjustWalletBalanceInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery,
AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery,
AdminWalletRefundRequestListQuery, CompleteAdminWalletRefundInput,
CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult,
CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome,
CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome,
CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput,
AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput,
CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput,
CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput,
CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput,
DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput,
FailAdminWalletRefundInput, ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput,
ProcessPaymentCallbackOutcome, RedeemWalletCodeInput, RedeemWalletCodeOutcome,
FailAdminWalletRefundInput, FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome,
StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder,
StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage,
StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger,
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome,
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput,
UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome,
WalletReadRepository, WalletWriteRepository,
};
use aether_data::{
@@ -92,8 +96,9 @@ use aether_data_contracts::repository::background_tasks::{
use aether_data_contracts::repository::billing::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord,
PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository,
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord,
};
use aether_data_contracts::repository::candidate_selection::{
@@ -122,12 +127,14 @@ use aether_data_contracts::repository::pool_scores::{
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogProviderConfigCasUpdate,
ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
@@ -136,7 +143,10 @@ use aether_data_contracts::repository::routing_profiles::{
RoutingGroupReadRepository, RoutingGroupWriteRepository,
};
use aether_data_contracts::repository::settlement::{
SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput,
ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput,
ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput,
ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation,
StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsageSettlementInput,
};
use aether_data_contracts::repository::usage::{
ApiKeyLastUsedDelta, ManagementTokenCounterDelta, PendingUsageCleanupSummary,
@@ -150,9 +160,9 @@ use aether_data_contracts::repository::video_tasks::{
use aether_runtime_state::RuntimeQueueStore;
pub(crate) use self::referrals::{
ReferralAdminStats, ReferralMutationStatus, ReferralRelationshipListQuery,
ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery,
ReferralRewardRecord, ReferralUserDashboard,
ReferralAdminStats, ReferralMutationStatus, ReferralReconciliationSummary,
ReferralRelationshipListQuery, ReferralRelationshipRecord, ReferralRewardConfig,
ReferralRewardListQuery, ReferralRewardRecord, ReferralUserDashboard,
};
#[derive(Clone, Default)]
@@ -4,9 +4,9 @@ use aether_data::DataLayerError;
use super::GatewayDataState;
pub(crate) use aether_data::backend::{
ReferralAdminStats, ReferralMutationStatus, ReferralRelationshipListQuery,
ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery,
ReferralRewardRecord, ReferralUserDashboard,
ReferralAdminStats, ReferralMutationStatus, ReferralReconciliationSummary,
ReferralRelationshipListQuery, ReferralRelationshipRecord, ReferralRewardConfig,
ReferralRewardListQuery, ReferralRewardRecord, ReferralUserDashboard,
};
impl GatewayDataState {
@@ -115,4 +115,13 @@ impl GatewayDataState {
.reverse_referral_rewards_for_order(order_id, amount_usd)
.await
}
pub(crate) async fn reconcile_referral_rewards_once(
&self,
reward_config: Option<ReferralRewardConfig>,
) -> Result<ReferralReconciliationSummary, DataLayerError> {
self.referrals()
.reconcile_referral_rewards_once(reward_config)
.await
}
}
@@ -684,7 +684,7 @@ mod tests {
.unwrap_or_default();
// One Arc is retained by the map and every active request
// owns one through its leader guard or follower state.
if participant_count >= participants + 1 {
if participant_count > participants {
break;
}
tokio::task::yield_now().await;
+272 -19
View File
@@ -7,33 +7,39 @@ use super::{
AdminWalletRefundRequestListQuery, AnnouncementListQuery, AuditLogListQuery,
BackgroundTaskListQuery, BackgroundTaskSummary, BillingModelContextCacheKey,
BillingModelContextCacheState, BillingModelContextInflightState, BillingPlanRecord,
BillingPlanWriteInput, CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
BillingPlanWriteInput, CompareAndSwapPaymentOrderStripeClientSecretInput,
CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
CreateAdminRedeemCodeBatchResult, CreateAnnouncementRecord, CreateManualWalletRechargeInput,
CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput,
CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, DataLayerError,
DatabaseMaintenanceSummary, DecisionTrace, DeleteAdminRedeemCodeBatchInput,
DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput,
GatewayDataState, GatewayProviderTransportSnapshot, LocalVideoTaskReadResponse,
PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, ProcessAdminWalletRefundInput,
ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, RedeemWalletCodeInput,
RedeemWalletCodeOutcome, RequestAuditBundle, RequestCandidateTrace, StoredAdminAuditLogPage,
StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage,
StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage,
StoredAdminWalletLedgerPage, StoredAdminWalletListPage, StoredAdminWalletRefund,
StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
FailWalletRechargeCheckoutInput, GatewayDataState, GatewayProviderTransportSnapshot,
LocalVideoTaskReadResponse, PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord,
PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, ProcessAdminWalletRefundInput,
ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput,
ReconcileUsagePolicyCostInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome,
ReleaseUsagePolicyRequestAdmissionInput, RequestAuditBundle, RequestCandidateTrace,
ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput,
ReserveUsagePolicyRequestOutcome, StoredAdminAuditLogPage, StoredAdminPaymentCallbackPage,
StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCodeBatch,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage,
StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredAnnouncement, StoredAnnouncementPage,
StoredBackgroundTaskEvent, StoredBackgroundTaskRun, StoredBackgroundTaskRunPage,
StoredBillingModelContext, StoredProviderQuotaSnapshot, StoredProviderUsageSummary,
StoredRequestUsageAudit, StoredSuspiciousActivity, StoredUsageSettlement,
StoredUserAuditLogPage, StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary,
StoredVideoTask, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage,
StoredWalletSnapshot, UpdateAnnouncementRecord, UpsertBackgroundTaskEvent,
UpsertBackgroundTaskRun, UpsertUsageRecord, UpsertVideoTask, UsageSettlementInput,
UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord, VideoTaskLookupKey,
VideoTaskModelCount, VideoTaskQueryFilter, VideoTaskStatusCount,
WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult, WalletLookupKey,
WalletMutationOutcome,
StoredRequestUsageAudit, StoredSuspiciousActivity, StoredUsagePolicyCostReservation,
StoredUsagePolicyRequestAdmission, StoredUsageSettlement, StoredUserAuditLogPage,
StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, StoredVideoTask,
StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot,
UpdateAdminWalletRefundGatewayInput, UpdateAnnouncementRecord,
UpdateWalletRechargeCheckoutInput, UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
UpsertUsageRecord, UpsertVideoTask, UsageSettlementInput, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord, VideoTaskLookupKey, VideoTaskModelCount, VideoTaskQueryFilter,
VideoTaskStatusCount, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult,
WalletLookupKey, WalletMutationOutcome,
};
use aether_data_contracts::repository::usage::{
PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest,
@@ -43,7 +49,9 @@ use aether_data_contracts::repository::usage::{
UsageDailyHeatmapQuery,
};
use aether_runtime_state::RuntimeQueueStore;
use aether_video_tasks_core::read_data_backed_video_task_response;
use aether_video_tasks_core::{
read_data_backed_video_task_response, read_data_backed_video_task_response_for_user,
};
use std::time::{Duration, Instant};
use tokio::time::timeout;
@@ -558,6 +566,17 @@ impl GatewayDataState {
}
}
pub(crate) async fn find_video_task_for_user(
&self,
key: VideoTaskLookupKey<'_>,
user_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
match &self.video_task_reader {
Some(repository) => repository.find_for_user(key, user_id).await,
None => Ok(None),
}
}
pub(crate) async fn list_video_task_page(
&self,
filter: &VideoTaskQueryFilter,
@@ -894,6 +913,21 @@ impl GatewayDataState {
}
}
pub(crate) async fn find_wallet_recharge_order_by_order_no(
&self,
user_id: &str,
order_no: &str,
) -> Result<Option<StoredAdminPaymentOrder>, DataLayerError> {
match &self.wallet_reader {
Some(repository) => {
repository
.find_wallet_recharge_order_by_order_no(user_id, order_no)
.await
}
None => Ok(None),
}
}
pub(crate) async fn find_pending_plan_purchase_order_by_user_id(
&self,
user_id: &str,
@@ -909,6 +943,16 @@ impl GatewayDataState {
}
}
pub(crate) async fn find_payment_order_by_order_no(
&self,
order_no: &str,
) -> Result<Option<StoredAdminPaymentOrder>, DataLayerError> {
match &self.wallet_reader {
Some(repository) => repository.find_payment_order_by_order_no(order_no).await,
None => Ok(None),
}
}
pub(crate) async fn find_wallet_refund(
&self,
wallet_id: &str,
@@ -934,6 +978,58 @@ impl GatewayDataState {
}
}
pub(crate) async fn update_wallet_recharge_checkout(
&self,
input: UpdateWalletRechargeCheckoutInput,
) -> Result<Option<WalletMutationOutcome<StoredAdminPaymentOrder>>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.update_wallet_recharge_checkout(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn compare_and_swap_payment_order_stripe_client_secret(
&self,
input: CompareAndSwapPaymentOrderStripeClientSecretInput,
) -> Result<Option<bool>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.compare_and_swap_payment_order_stripe_client_secret(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn fail_wallet_recharge_checkout(
&self,
input: FailWalletRechargeCheckoutInput,
) -> Result<Option<WalletMutationOutcome<StoredAdminPaymentOrder>>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.fail_wallet_recharge_checkout(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn reclaim_wallet_recharge_checkout(
&self,
input: ReclaimWalletRechargeCheckoutInput,
) -> Result<Option<WalletMutationOutcome<StoredAdminPaymentOrder>>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.reclaim_wallet_recharge_checkout(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn create_plan_purchase_order(
&self,
input: CreatePlanPurchaseOrderInput,
@@ -1009,6 +1105,19 @@ impl GatewayDataState {
}
}
pub(crate) async fn update_admin_wallet_refund_gateway(
&self,
input: UpdateAdminWalletRefundGatewayInput,
) -> Result<Option<WalletMutationOutcome<StoredAdminWalletRefund>>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.update_admin_wallet_refund_gateway(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn complete_admin_wallet_refund(
&self,
input: CompleteAdminWalletRefundInput,
@@ -1151,6 +1260,83 @@ impl GatewayDataState {
}
}
pub(crate) async fn reserve_usage_policy_cost(
&self,
input: ReserveUsagePolicyCostInput,
) -> Result<Option<ReserveUsagePolicyCostOutcome>, DataLayerError> {
match &self.settlement_writer {
Some(repository) => repository.reserve_usage_policy_cost(input).await.map(Some),
None => Ok(None),
}
}
pub(crate) async fn reserve_usage_policy_request(
&self,
input: ReserveUsagePolicyRequestInput,
) -> Result<Option<ReserveUsagePolicyRequestOutcome>, DataLayerError> {
match &self.settlement_writer {
Some(repository) => repository
.reserve_usage_policy_request(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn release_usage_policy_request_admission(
&self,
input: ReleaseUsagePolicyRequestAdmissionInput,
) -> Result<Option<StoredUsagePolicyRequestAdmission>, DataLayerError> {
match &self.settlement_writer {
Some(repository) => {
repository
.release_usage_policy_request_admission(input)
.await
}
None => Ok(None),
}
}
pub(crate) async fn cleanup_usage_policy_request_admissions(
&self,
now_unix_secs: u64,
batch_size: usize,
) -> Result<usize, DataLayerError> {
match &self.settlement_writer {
Some(repository) => {
repository
.cleanup_usage_policy_request_admissions(now_unix_secs, batch_size)
.await
}
None => Ok(0),
}
}
pub(crate) async fn reconcile_usage_policy_cost(
&self,
input: ReconcileUsagePolicyCostInput,
) -> Result<Option<StoredUsagePolicyCostReservation>, DataLayerError> {
match &self.settlement_writer {
Some(repository) => repository.reconcile_usage_policy_cost(input).await,
None => Ok(None),
}
}
pub(crate) async fn cleanup_usage_policy_cost_reservations(
&self,
now_unix_secs: u64,
batch_size: usize,
) -> Result<usize, DataLayerError> {
match &self.settlement_writer {
Some(repository) => {
repository
.cleanup_usage_policy_cost_reservations(now_unix_secs, batch_size)
.await
}
None => Ok(0),
}
}
pub(crate) async fn reset_due_provider_quotas(
&self,
now_unix_secs: u64,
@@ -2496,6 +2682,48 @@ impl GatewayDataState {
}
}
pub(crate) async fn find_payment_gateway_config_strong(
&self,
provider: &str,
) -> Result<Option<PaymentGatewayConfigRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.find_payment_gateway_config_strong(provider)
.await
}
None => Ok(None),
}
}
pub(crate) async fn compare_and_swap_payment_gateway_secret(
&self,
update: &PaymentGatewaySecretCasUpdate,
) -> Result<bool, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.compare_and_swap_payment_gateway_secret(update)
.await
}
None => Ok(false),
}
}
pub(crate) async fn compare_and_swap_payment_gateway_config(
&self,
input: &PaymentGatewayConfigCasWriteInput,
) -> Result<AdminBillingMutationOutcome<PaymentGatewayConfigRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.compare_and_swap_payment_gateway_config(input)
.await
}
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn upsert_payment_gateway_config(
&self,
input: &PaymentGatewayConfigWriteInput,
@@ -2578,6 +2806,21 @@ impl GatewayDataState {
}
}
pub(crate) async fn revoke_user_plan_entitlement(
&self,
user_id: &str,
entitlement_id: &str,
) -> Result<AdminBillingMutationOutcome<()>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.revoke_user_plan_entitlement(user_id, entitlement_id)
.await
}
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn find_user_daily_quota_availability(
&self,
user_id: &str,
@@ -2652,6 +2895,16 @@ impl GatewayDataState {
read_data_backed_video_task_response(self, route_family, request_path).await
}
pub(crate) async fn read_video_task_response_for_user(
&self,
route_family: Option<&str>,
request_path: &str,
user_id: &str,
) -> Result<Option<LocalVideoTaskReadResponse>, DataLayerError> {
read_data_backed_video_task_response_for_user(self, route_family, request_path, user_id)
.await
}
pub(crate) async fn find_background_task_run(
&self,
run_id: &str,
@@ -1,10 +1,15 @@
use std::collections::BTreeMap;
use std::sync::{Arc, RwLock};
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
use aether_data_contracts::repository::candidates::RequestCandidateRepository;
use aether_data_contracts::repository::pool_scores::PoolMemberScoreRepository;
use aether_data_contracts::repository::quota::ProviderQuotaRepository;
use aether_data_contracts::repository::routing_profiles::{
StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion,
};
use aether_data_contracts::repository::usage::UsageRepository;
use aether_routing_core::RoutingGroupConfig;
use super::{
AnnouncementReadRepository, AnnouncementWriteRepository, AuthApiKeyReadRepository,
@@ -213,6 +218,21 @@ impl GatewayDataState {
self
}
#[cfg(test)]
pub(crate) fn with_cached_provider_catalog_reader_for_tests<T>(
mut self,
repository: Arc<T>,
) -> Self
where
T: ProviderCatalogReadRepository + 'static,
{
let inner: Arc<dyn ProviderCatalogReadRepository> = repository;
self.provider_catalog_reader = Some(Arc::new(
super::provider_catalog_cache::CachedProviderCatalogReadRepository::new(inner),
));
self
}
#[cfg(test)]
pub(crate) fn with_request_candidate_reader(
mut self,
@@ -877,6 +897,30 @@ impl GatewayDataState {
self
}
#[cfg(test)]
pub(crate) fn with_system_default_routing_group_for_tests(self) -> Self {
let now = 1;
let repository = Arc::new(InMemoryRoutingGroupRepository::seed(
[StoredRoutingGroup {
id: "system-default".to_string(),
name: "system-default".to_string(),
description: Some("test system default routing strategy".to_string()),
enabled: true,
is_system_default: true,
sort_order: 0,
config_json: serde_json::to_value(RoutingGroupConfig::default())
.expect("default routing config should serialize"),
version: 1,
created_at: now,
updated_at: now,
published_at: Some(now),
}],
std::iter::empty::<StoredRoutingGroupBinding>(),
std::iter::empty::<StoredRoutingGroupVersion>(),
));
self.with_routing_group_repository_for_tests(repository)
}
#[cfg(test)]
pub(crate) fn with_auth_api_key_reader(
mut self,
@@ -1746,6 +1790,18 @@ impl GatewayDataState {
}
}
#[cfg(test)]
pub(crate) fn attach_auth_api_key_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
where
T: aether_data::repository::auth::AuthRepository + 'static,
{
let auth_api_key_reader: Arc<dyn AuthApiKeyReadRepository> = repository.clone();
let auth_api_key_writer: Arc<dyn AuthApiKeyWriteRepository> = repository;
self.auth_api_key_reader = Some(auth_api_key_reader);
self.auth_api_key_writer = Some(auth_api_key_writer);
self
}
#[cfg(test)]
pub(crate) fn with_decision_trace_readers_for_tests(
request_candidate_repository: Arc<dyn RequestCandidateReadRepository>,
@@ -2335,6 +2391,36 @@ impl GatewayDataState {
}
}
#[cfg(test)]
pub(crate) fn with_auth_candidate_selection_provider_catalog_request_candidate_and_gemini_file_mapping_repositories_for_tests<
T,
U,
V,
>(
auth_api_key_repository: Arc<dyn AuthApiKeyReadRepository>,
candidate_selection_repository: Arc<dyn MinimalCandidateSelectionReadRepository>,
provider_catalog_repository: Arc<U>,
request_candidate_repository: Arc<T>,
gemini_file_mapping_repository: Arc<V>,
encryption_key: impl Into<String>,
) -> Self
where
T: RequestCandidateRepository + 'static,
U: ProviderCatalogReadRepository + ProviderCatalogWriteRepository + 'static,
V: aether_data::repository::gemini_file_mappings::GeminiFileMappingRepository + 'static,
{
let mut state = Self::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
auth_api_key_repository,
candidate_selection_repository,
provider_catalog_repository,
request_candidate_repository,
encryption_key,
);
state.gemini_file_mapping_reader = Some(gemini_file_mapping_repository.clone());
state.gemini_file_mapping_writer = Some(gemini_file_mapping_repository);
state
}
#[cfg(test)]
pub(crate) fn with_auth_candidate_selection_provider_catalog_request_candidates_for_tests<
T,
@@ -12,10 +12,113 @@ use aether_data_contracts::repository::usage::{
};
use aether_data::repository::auth::AuthApiKeyReadRepository;
use aether_data::repository::management_tokens::{
InMemoryManagementTokenRepository, ManagementTokenReadRepository,
ManagementTokenWriteRepository, StoredManagementToken, StoredManagementTokenUserSummary,
StoredManagementTokenWithUser,
};
use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode};
use aether_data::repository::proxy_nodes::{ProxyNodeReadRepository, ProxyNodeWriteRepository};
use aether_data::repository::users::{
InMemoryUserReadRepository, StoredUserAuthRecord, UserReadRepository,
};
use sha2::{Digest, Sha256};
use super::{GatewayDataConfig, GatewayDataState};
impl GatewayDataState {
pub(crate) fn with_tunnel_management_auth_for_testkit(
node_id: &str,
tunnel_generation: &str,
raw_token: &str,
encryption_key: impl Into<String>,
) -> Result<Self, aether_data::DataLayerError> {
const TOKEN_ID: &str = "token-tunnel-harness";
const USER_ID: &str = "user-tunnel-harness";
let node = StoredProxyNode::new(
node_id.to_string(),
"tunnel harness node".to_string(),
"127.0.0.1".to_string(),
0,
false,
"offline".to_string(),
30,
0,
0,
0,
0,
0,
true,
false,
0,
)?
.with_tunnel_generation(tunnel_generation.to_string());
let proxy_repository = Arc::new(InMemoryProxyNodeRepository::seed([node]));
let user_summary = StoredManagementTokenUserSummary::new(
USER_ID.to_string(),
Some("[email protected]".to_string()),
"tunnel_harness_admin".to_string(),
"admin".to_string(),
)?;
let token = StoredManagementToken::new(
TOKEN_ID.to_string(),
USER_ID.to_string(),
"tunnel harness token".to_string(),
)?
.with_permissions(Some(serde_json::json!(["admin:proxy_nodes:admin"])));
let token_hash = format!("{:x}", Sha256::digest(raw_token.as_bytes()));
let token_repository = Arc::new(InMemoryManagementTokenRepository::seed_with_hashes(
[StoredManagementTokenWithUser::new(token, user_summary)],
[(token_hash, TOKEN_ID.to_string())],
));
let token_reader: Arc<dyn ManagementTokenReadRepository> = token_repository.clone();
let token_writer: Arc<dyn ManagementTokenWriteRepository> = token_repository;
let user = StoredUserAuthRecord::new(
USER_ID.to_string(),
Some("[email protected]".to_string()),
true,
"tunnel_harness_admin".to_string(),
None,
"admin".to_string(),
"local".to_string(),
None,
None,
None,
true,
false,
None,
None,
)?;
let user_reader: Arc<dyn UserReadRepository> =
Arc::new(InMemoryUserReadRepository::seed_auth_users([user]));
let mut state =
Self::with_proxy_node_repository_for_testkit(proxy_repository, encryption_key);
state.management_token_reader = Some(token_reader);
state.management_token_writer = Some(token_writer);
state.user_reader = Some(user_reader);
Ok(state)
}
pub(crate) fn with_proxy_node_repository_for_testkit<T>(
repository: Arc<T>,
encryption_key: impl Into<String>,
) -> Self
where
T: ProxyNodeReadRepository + ProxyNodeWriteRepository + 'static,
{
let proxy_node_reader: Arc<dyn ProxyNodeReadRepository> = repository.clone();
let proxy_node_writer: Arc<dyn ProxyNodeWriteRepository> = repository;
let mut state = Self::disabled();
state.config = GatewayDataConfig::disabled().with_encryption_key(encryption_key);
state.proxy_node_reader = Some(proxy_node_reader);
state.proxy_node_writer = Some(proxy_node_writer);
state
}
pub(crate) fn with_openai_chat_pressure_repositories_for_testkit<T, U, V>(
auth_api_key_repository: Arc<dyn AuthApiKeyReadRepository>,
candidate_selection_repository: Arc<dyn MinimalCandidateSelectionReadRepository>,
+1 -1
View File
@@ -312,7 +312,7 @@ async fn data_state_checks_user_uniqueness_through_user_reader() {
Some("[email protected]".to_string()),
true,
"admin".to_string(),
Some(format!("$2b$12${}", "a".repeat(53))),
Some("$2b$12$4qL4tdcsFwVaDTw5Ck3xzu8GpNdre56DiNR6Dnw7t6gCXaEnqAe7G".to_string()),
"admin".to_string(),
"local".to_string(),
None,

Some files were not shown because too many files have changed in this diff Show More