Improve gateway scheduling and runtime admission

This commit is contained in:
elky
2026-06-24 01:53:45 +08:00
parent cf0af8fa1e
commit d336d1a7fa
87 changed files with 9671 additions and 804 deletions
@@ -23,6 +23,7 @@ use super::principal::derive_principal_candidate;
use super::types::{
GatewayCredentialCarrier, GatewayPrincipalCandidate, GatewayTrustedAuthHeaders,
};
use crate::cache::AuthContextInflightRegistration;
use crate::headers::header_value_str;
const AUTH_CONTEXT_CACHE_TTL: Duration = Duration::from_secs(60);
@@ -137,14 +138,15 @@ pub(in super::super) async fn resolve_control_decision_auth(
if let Some(cache_key) = auth_context_cache_key.as_deref() {
if let Some(auth_context) = get_cached_auth_context(state, cache_key) {
resolved_auth_context = if auth_context_cache_refresh_on_hit() {
let refreshed = refresh_execution_runtime_auth_context(
state,
auth_context,
decision.auth_endpoint_signature.as_deref(),
Some(
refresh_cached_auth_context_or_reuse(
state,
cache_key,
auth_context,
decision.auth_endpoint_signature.as_deref(),
)
.await?,
)
.await?;
put_cached_auth_context(state, cache_key.to_string(), refreshed.clone());
Some(refreshed)
} else {
Some(auth_context)
};
@@ -152,11 +154,13 @@ pub(in super::super) async fn resolve_control_decision_auth(
}
if resolved_auth_context.is_none() {
resolved_auth_context = resolve_data_backed_auth_context(
resolved_auth_context = resolve_data_backed_auth_context_cached(
state,
auth_context_cache_key.as_deref(),
headers,
uri,
decision.auth_endpoint_signature.as_deref(),
true,
)
.await?;
if let (Some(cache_key), Some(auth_context)) = (
@@ -501,14 +505,15 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
if !auth_context_cache_refresh_on_hit() {
return Ok(Some(auth_context));
}
return Ok(Some(
refresh_execution_runtime_auth_context(
state,
auth_context,
decision.auth_endpoint_signature.as_deref(),
)
.await?,
));
return refresh_decision_auth_context_on_hit(
state,
headers,
uri,
decision.auth_endpoint_signature.as_deref(),
auth_context,
)
.await
.map(Some);
}
let Some(auth_endpoint_signature) = decision.auth_endpoint_signature.as_deref() else {
@@ -524,18 +529,25 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
return Ok(Some(auth_context));
}
let refreshed = refresh_execution_runtime_auth_context(
let refreshed = refresh_cached_auth_context_or_reuse(
state,
&cache_key,
auth_context,
Some(auth_endpoint_signature),
)
.await?;
put_cached_auth_context(state, cache_key, refreshed.clone());
return Ok(Some(refreshed));
}
if let Some(auth_context) =
resolve_data_backed_auth_context(state, headers, uri, Some(auth_endpoint_signature)).await?
if let Some(auth_context) = resolve_data_backed_auth_context_cached(
state,
Some(cache_key.as_str()),
headers,
uri,
Some(auth_endpoint_signature),
true,
)
.await?
{
if auth_context.user_id.is_empty() || auth_context.api_key_id.is_empty() {
return Ok(None);
@@ -547,6 +559,109 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
Ok(None)
}
async fn refresh_decision_auth_context_on_hit(
state: &AppState,
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: Option<&str>,
auth_context: GatewayControlAuthContext,
) -> Result<GatewayControlAuthContext, GatewayError> {
let Some(auth_endpoint_signature) = auth_endpoint_signature else {
return Ok(auth_context);
};
let Some(cache_key) = build_auth_context_cache_key(headers, uri, auth_endpoint_signature)
else {
return refresh_execution_runtime_auth_context(
state,
auth_context,
Some(auth_endpoint_signature),
)
.await;
};
refresh_cached_auth_context_or_reuse(
state,
&cache_key,
auth_context,
Some(auth_endpoint_signature),
)
.await
}
async fn refresh_cached_auth_context_or_reuse(
state: &AppState,
cache_key: &str,
auth_context: GatewayControlAuthContext,
auth_endpoint_signature: Option<&str>,
) -> Result<GatewayControlAuthContext, GatewayError> {
match state.auth_context_cache.register_inflight(cache_key) {
AuthContextInflightRegistration::Leader(_guard) => {
let refreshed = refresh_execution_runtime_auth_context(
state,
auth_context,
auth_endpoint_signature,
)
.await?;
put_cached_auth_context(state, cache_key.to_string(), refreshed.clone());
Ok(refreshed)
}
AuthContextInflightRegistration::Follower => Ok(auth_context),
AuthContextInflightRegistration::Bypass => {
refresh_execution_runtime_auth_context(state, auth_context, auth_endpoint_signature)
.await
}
}
}
async fn resolve_data_backed_auth_context_cached(
state: &AppState,
cache_key: Option<&str>,
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: Option<&str>,
cache_negative: 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;
};
loop {
let notified = state.auth_context_cache.notified();
match state.auth_context_cache.register_inflight(cache_key) {
AuthContextInflightRegistration::Leader(_guard) => {
let resolved =
resolve_data_backed_auth_context(state, headers, uri, auth_endpoint_signature)
.await?;
if let Some(auth_context) = resolved.as_ref() {
if cache_negative
|| (!auth_context.user_id.is_empty() && !auth_context.api_key_id.is_empty())
{
put_cached_auth_context(state, cache_key.to_string(), auth_context.clone());
}
}
return Ok(resolved);
}
AuthContextInflightRegistration::Follower => {
notified.await;
if let Some(auth_context) = get_cached_auth_context(state, cache_key) {
return Ok(Some(auth_context));
}
if !cache_negative {
return Ok(None);
}
}
AuthContextInflightRegistration::Bypass => {
return resolve_data_backed_auth_context(
state,
headers,
uri,
auth_endpoint_signature,
)
.await;
}
}
}
}
pub(crate) async fn refresh_execution_runtime_auth_context(
state: &AppState,
auth_context: GatewayControlAuthContext,
@@ -1176,6 +1291,7 @@ fn get_cached_auth_context(state: &AppState, cache_key: &str) -> Option<GatewayC
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use aether_data::repository::auth::{
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
@@ -1188,6 +1304,7 @@ mod tests {
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use axum::http::{HeaderMap, Uri};
use futures_util::future::join_all;
use super::{
get_cached_auth_context, resolve_control_decision_auth, resolve_data_backed_auth_context,
@@ -1347,6 +1464,61 @@ mod tests {
assert_eq!(repository.touch_count("key-1"), 1);
}
#[tokio::test]
async fn control_auth_context_singleflights_concurrent_cache_misses() {
let api_key = "sk-test-concurrent-auth-miss";
let repository = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key(api_key)),
sample_snapshot("key-concurrent-auth-miss", "user-concurrent-auth-miss"),
)])
.with_lookup_delay_for_tests(Duration::from_millis(20)),
);
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 {api_key}").parse().unwrap(),
);
let request_uri = uri("/v1/chat/completions");
let tasks = (0..32).map(|index| {
let decision = GatewayControlDecision::synthetic(
"/v1/chat/completions",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
);
let trace_id = format!("trace-concurrent-auth-miss-{index}");
let state = &state;
let headers = &headers;
let request_uri = &request_uri;
async move {
resolve_control_decision_auth(state, headers, request_uri, &trace_id, decision)
.await
}
});
for result in join_all(tasks).await {
let ControlDecisionAuthResolution::Resolved(decision) =
result.expect("auth resolution should succeed");
let auth_context = decision
.auth_context
.expect("auth context should be resolved");
assert_eq!(auth_context.user_id, "user-concurrent-auth-miss");
assert_eq!(auth_context.api_key_id, "key-concurrent-auth-miss");
}
assert_eq!(
repository.key_hash_lookup_count(&hash_api_key(api_key)),
1,
"concurrent cache misses for one auth context should only load one snapshot"
);
}
#[tokio::test]
async fn data_backed_auth_context_marks_wallet_denial_as_not_allowed() {
let api_key = "sk-test-empty-wallet";
@@ -1496,6 +1668,85 @@ mod tests {
assert!(!second.access_allowed);
}
#[tokio::test]
async fn execution_runtime_auth_context_singleflights_concurrent_cache_refreshes() {
let api_key = "sk-test-runtime-auth-refresh";
let auth_repository = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key(api_key)),
sample_snapshot("key-runtime-auth-refresh", "user-runtime-auth-refresh"),
)])
.with_lookup_delay_for_tests(Duration::from_millis(20)),
);
let data =
GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository.clone());
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data);
let decision = GatewayControlDecision::synthetic(
"/v1/chat/completions",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
);
let mut headers = HeaderMap::new();
headers.insert("x-api-key", api_key.parse().unwrap());
let request_uri = uri("/v1/chat/completions");
let first = resolve_execution_runtime_auth_context(
&state,
&decision,
&headers,
&request_uri,
"trace-runtime-auth-refresh-prime",
)
.await
.expect("resolution should succeed")
.expect("auth context should exist");
assert_eq!(first.api_key_id, "key-runtime-auth-refresh");
assert_eq!(
auth_repository.key_hash_lookup_count(&hash_api_key(api_key)),
1
);
let tasks = (0..32).map(|index| {
let trace_id = format!("trace-runtime-auth-refresh-{index}");
let state = &state;
let decision = &decision;
let headers = &headers;
let request_uri = &request_uri;
async move {
resolve_execution_runtime_auth_context(
state,
decision,
headers,
request_uri,
&trace_id,
)
.await
}
});
for result in join_all(tasks).await {
let auth_context = result
.expect("resolution should succeed")
.expect("auth context should exist");
assert_eq!(auth_context.user_id, "user-runtime-auth-refresh");
assert_eq!(auth_context.api_key_id, "key-runtime-auth-refresh");
}
assert_eq!(
auth_repository.key_hash_lookup_count(&hash_api_key(api_key)),
1,
"cache refreshes should reuse the existing auth context under concurrent pressure"
);
assert_eq!(
auth_repository.snapshot_lookup_count("key-runtime-auth-refresh"),
1,
"only one cached auth context refresh should read the snapshot by user/key id"
);
}
#[tokio::test]
async fn data_backed_auth_context_allows_provider_id_for_matching_provider_type() {
let api_key = "sk-test-provider-id";