mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
Improve gateway scheduling and runtime admission
This commit is contained in:
@@ -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";
|
||||
|
||||
Reference in New Issue
Block a user