mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
feat(gateway): harden provider request execution
Preserve exact request payloads and model client surface and API operation explicitly. Add Anthropic compatibility profiles, bounded stream commitment, and scoped OAuth retry behavior across provider transports.
This commit is contained in:
@@ -17,8 +17,8 @@ use aether_contracts::{
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
StoredProviderCatalogKey,
|
||||
ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_runtime_state::RuntimeLockLease;
|
||||
use base64::{engine::general_purpose::STANDARD, Engine as _};
|
||||
@@ -42,6 +42,12 @@ const OAUTH_EXPIRED_PREFIX: &str = "[OAUTH_EXPIRED] ";
|
||||
const OAUTH_REFRESH_FAILED_PREFIX: &str = "[REFRESH_FAILED] ";
|
||||
const OAUTH_REQUEST_FAILED_PREFIX: &str = "[REQUEST_FAILED] ";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
struct ProviderTransportCredentialFence {
|
||||
encrypted_auth_config: String,
|
||||
credential: ProviderCatalogKeyOAuthCredentialFence,
|
||||
}
|
||||
|
||||
struct GatewayLocalOAuthHttpExecutor<'a> {
|
||||
state: &'a AppState,
|
||||
}
|
||||
@@ -173,32 +179,6 @@ fn local_oauth_transport_context_allows_reload(
|
||||
initial, current,
|
||||
);
|
||||
}
|
||||
if initial
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
&& initial.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
{
|
||||
let initial_config = initial
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok());
|
||||
let current_config = current
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok());
|
||||
return current
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
&& current.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
&& initial_config == current_config
|
||||
&& initial.key.decrypted_api_key == current.key.decrypted_api_key;
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
@@ -1380,7 +1360,7 @@ impl AppState {
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
let expected_auth_config = if current_transport
|
||||
let expected_credential_fence = if current_transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
@@ -1388,10 +1368,10 @@ impl AppState {
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
{
|
||||
match self
|
||||
.capture_provider_transport_auth_config_fence(¤t_transport)
|
||||
.capture_provider_transport_credential_fence(¤t_transport)
|
||||
.await?
|
||||
{
|
||||
Some(ciphertext) => Some(ciphertext),
|
||||
Some(fence) => Some(fence),
|
||||
None => {
|
||||
let Some(reloaded) = self
|
||||
.read_provider_transport_snapshot_uncached(
|
||||
@@ -1485,7 +1465,7 @@ impl AppState {
|
||||
.persist_local_oauth_refresh_entry(
|
||||
¤t_transport,
|
||||
&refreshed_entry,
|
||||
expected_auth_config.as_deref(),
|
||||
expected_credential_fence.as_ref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -1533,9 +1513,9 @@ impl AppState {
|
||||
let lock_owner = format!("aether-gateway-admin-{}", std::process::id());
|
||||
let initial_transport = transport.clone();
|
||||
let mut current_transport = transport.clone();
|
||||
current_transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
let expected_refresh_fingerprint =
|
||||
provider_transport::codex_agent_identity_refresh_fingerprint(¤t_transport, None);
|
||||
let expected_refresh_fingerprint = self
|
||||
.oauth_refresh
|
||||
.refresh_fingerprint_for_transport(&initial_transport);
|
||||
let executor = GatewayLocalOAuthHttpExecutor { state: self };
|
||||
let transport_refresh_token_fingerprint = oauth_auth_config_refresh_token_fingerprint(
|
||||
current_transport.key.decrypted_auth_config.as_deref(),
|
||||
@@ -1560,8 +1540,8 @@ impl AppState {
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
let expected_auth_config = match self
|
||||
.capture_provider_transport_auth_config_fence(¤t_transport)
|
||||
let expected_credential_fence = match self
|
||||
.capture_provider_transport_credential_fence(¤t_transport)
|
||||
.await
|
||||
.map_err(
|
||||
|err| provider_transport::LocalOAuthRefreshError::InvalidResponse {
|
||||
@@ -1569,7 +1549,7 @@ impl AppState {
|
||||
message: format!("{err:?}"),
|
||||
},
|
||||
)? {
|
||||
Some(ciphertext) => Some(ciphertext),
|
||||
Some(fence) => Some(fence),
|
||||
None if current_transport.key.decrypted_auth_config.is_some() => {
|
||||
let Some(reloaded) = self
|
||||
.read_provider_transport_snapshot_uncached(
|
||||
@@ -1588,7 +1568,6 @@ impl AppState {
|
||||
return Ok(None);
|
||||
};
|
||||
current_transport = reloaded;
|
||||
current_transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
continue;
|
||||
}
|
||||
None => None,
|
||||
@@ -1621,7 +1600,6 @@ impl AppState {
|
||||
continue;
|
||||
};
|
||||
current_transport = reloaded_transport;
|
||||
current_transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -1684,7 +1662,7 @@ impl AppState {
|
||||
.persist_local_oauth_refresh_entry(
|
||||
¤t_transport,
|
||||
&refreshed_entry,
|
||||
expected_auth_config.as_deref(),
|
||||
expected_credential_fence.as_ref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -1772,6 +1750,16 @@ impl AppState {
|
||||
&self,
|
||||
transport: &provider_transport::GatewayProviderTransportSnapshot,
|
||||
) -> Result<Option<String>, GatewayError> {
|
||||
Ok(self
|
||||
.capture_provider_transport_credential_fence(transport)
|
||||
.await?
|
||||
.map(|fence| fence.encrypted_auth_config))
|
||||
}
|
||||
|
||||
async fn capture_provider_transport_credential_fence(
|
||||
&self,
|
||||
transport: &provider_transport::GatewayProviderTransportSnapshot,
|
||||
) -> Result<Option<ProviderTransportCredentialFence>, GatewayError> {
|
||||
let key_id = transport.key.id.trim();
|
||||
let stored = self
|
||||
.data
|
||||
@@ -1780,11 +1768,51 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
.into_iter()
|
||||
.next();
|
||||
let Some(ciphertext) = stored.and_then(|key| key.encrypted_auth_config) else {
|
||||
let Some(stored) = stored else {
|
||||
return Ok(None);
|
||||
};
|
||||
if stored.provider_id != transport.provider.id
|
||||
|| stored.auth_type != transport.key.auth_type
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
let provider = self
|
||||
.data
|
||||
.list_provider_catalog_providers_by_ids(std::slice::from_ref(&stored.provider_id))
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
.into_iter()
|
||||
.next();
|
||||
let Some(provider) = provider else {
|
||||
return Ok(None);
|
||||
};
|
||||
if provider.provider_type != transport.provider.provider_type {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let stored_api_key = match stored.encrypted_api_key.as_deref() {
|
||||
Some(ciphertext) => Some(
|
||||
decrypt_catalog_secret_with_fallbacks(self.data.encryption_key(), ciphertext)
|
||||
.ok_or_else(|| {
|
||||
GatewayError::Internal(
|
||||
"provider api_key could not be verified for runtime fencing"
|
||||
.to_string(),
|
||||
)
|
||||
})?,
|
||||
),
|
||||
None => None,
|
||||
};
|
||||
let transport_api_key = (!transport.key.decrypted_api_key.is_empty())
|
||||
.then_some(transport.key.decrypted_api_key.as_str());
|
||||
if stored_api_key.as_deref() != transport_api_key {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(ciphertext) = stored.encrypted_auth_config.as_deref() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let plaintext =
|
||||
decrypt_catalog_secret_with_fallbacks(self.data.encryption_key(), ciphertext.as_str())
|
||||
decrypt_catalog_secret_with_fallbacks(self.data.encryption_key(), ciphertext)
|
||||
.ok_or_else(|| {
|
||||
GatewayError::Internal(
|
||||
"provider auth_config could not be verified for runtime fencing"
|
||||
@@ -1810,7 +1838,15 @@ impl AppState {
|
||||
if config != transport_config {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(Some(ciphertext))
|
||||
Ok(Some(ProviderTransportCredentialFence {
|
||||
encrypted_auth_config: ciphertext.to_string(),
|
||||
credential: ProviderCatalogKeyOAuthCredentialFence {
|
||||
encrypted_api_key: stored.encrypted_api_key,
|
||||
auth_type: stored.auth_type,
|
||||
provider_id: stored.provider_id,
|
||||
provider_type: provider.provider_type,
|
||||
},
|
||||
}))
|
||||
}
|
||||
|
||||
pub(crate) async fn mark_provider_catalog_key_oauth_invalid(
|
||||
@@ -1883,18 +1919,23 @@ impl AppState {
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn mark_provider_catalog_key_oauth_invalid_fenced(
|
||||
pub(crate) async fn mark_provider_transport_oauth_invalid_fenced(
|
||||
&self,
|
||||
key_id: &str,
|
||||
provider_type: &str,
|
||||
transport: &provider_transport::GatewayProviderTransportSnapshot,
|
||||
invalid_reason: &str,
|
||||
expected_encrypted_auth_config: &str,
|
||||
) -> Result<bool, GatewayError> {
|
||||
let invalid_reason = invalid_reason.trim();
|
||||
let expected_encrypted_auth_config = expected_encrypted_auth_config.trim();
|
||||
if invalid_reason.is_empty() || expected_encrypted_auth_config.is_empty() {
|
||||
let key_id = transport.key.id.trim();
|
||||
let provider_type = transport.provider.provider_type.as_str();
|
||||
if invalid_reason.is_empty() || key_id.is_empty() {
|
||||
return Ok(false);
|
||||
}
|
||||
let Some(expected_credential_fence) = self
|
||||
.capture_provider_transport_credential_fence(transport)
|
||||
.await?
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
|
||||
let Some(mut latest_key) = self
|
||||
.data
|
||||
@@ -1906,7 +1947,12 @@ impl AppState {
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if latest_key.encrypted_auth_config.as_deref() != Some(expected_encrypted_auth_config)
|
||||
if latest_key.encrypted_auth_config.as_deref()
|
||||
!= Some(expected_credential_fence.encrypted_auth_config.as_str())
|
||||
|| latest_key.encrypted_api_key
|
||||
!= expected_credential_fence.credential.encrypted_api_key
|
||||
|| latest_key.auth_type != expected_credential_fence.credential.auth_type
|
||||
|| latest_key.provider_id != expected_credential_fence.credential.provider_id
|
||||
|| !provider_key_is_oauth_managed(&latest_key, provider_type)
|
||||
{
|
||||
return Ok(false);
|
||||
@@ -1940,9 +1986,10 @@ impl AppState {
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: Some(
|
||||
expected_encrypted_auth_config.to_string(),
|
||||
expected_credential_fence.encrypted_auth_config.clone(),
|
||||
),
|
||||
encrypted_auth_config: expected_encrypted_auth_config.to_string(),
|
||||
expected_credential: Some(expected_credential_fence.credential),
|
||||
encrypted_auth_config: expected_credential_fence.encrypted_auth_config,
|
||||
encrypted_api_key_update: None,
|
||||
expires_at_unix_secs_update: None,
|
||||
oauth_invalid_at_unix_secs: latest_key.oauth_invalid_at_unix_secs,
|
||||
@@ -1965,7 +2012,7 @@ impl AppState {
|
||||
&self,
|
||||
transport: &provider_transport::GatewayProviderTransportSnapshot,
|
||||
entry: &provider_transport::CachedOAuthEntry,
|
||||
expected_auth_config: Option<&str>,
|
||||
expected_credential_fence: Option<&ProviderTransportCredentialFence>,
|
||||
) -> Result<(), GatewayError> {
|
||||
let key_id = transport.key.id.trim();
|
||||
if key_id.is_empty() {
|
||||
@@ -1973,6 +2020,19 @@ impl AppState {
|
||||
}
|
||||
|
||||
if local_oauth_refresh_entry_should_stay_memory_only(transport, entry) {
|
||||
let expected_credential_fence = expected_credential_fence.ok_or_else(|| {
|
||||
GatewayError::Internal(
|
||||
"memory-only OAuth refresh has no starting credential fence".to_string(),
|
||||
)
|
||||
})?;
|
||||
let current_credential_fence = self
|
||||
.capture_provider_transport_credential_fence(transport)
|
||||
.await?;
|
||||
if current_credential_fence.as_ref() != Some(expected_credential_fence) {
|
||||
return Err(GatewayError::Internal(
|
||||
"OAuth credential changed while memory-only refresh was in flight".to_string(),
|
||||
));
|
||||
}
|
||||
tracing::info!(
|
||||
key_id = %key_id,
|
||||
provider_id = %transport.provider.id,
|
||||
@@ -2034,18 +2094,22 @@ impl AppState {
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
let expected_credential_fence = expected_credential_fence.ok_or_else(|| {
|
||||
GatewayError::Internal(
|
||||
"Agent Identity task registration has no starting credential fence".to_string(),
|
||||
)
|
||||
})?;
|
||||
let expected_encrypted_auth_config =
|
||||
expected_auth_config.map(str::to_string).ok_or_else(|| {
|
||||
GatewayError::Internal(
|
||||
"Agent Identity task registration has no starting auth_config fence"
|
||||
.to_string(),
|
||||
)
|
||||
})?;
|
||||
expected_credential_fence.encrypted_auth_config.clone();
|
||||
if latest_key.encrypted_auth_config.as_deref()
|
||||
!= Some(expected_encrypted_auth_config.as_str())
|
||||
|| latest_key.encrypted_api_key
|
||||
!= expected_credential_fence.credential.encrypted_api_key
|
||||
|| latest_key.auth_type != expected_credential_fence.credential.auth_type
|
||||
|| latest_key.provider_id != expected_credential_fence.credential.provider_id
|
||||
{
|
||||
return Err(GatewayError::Internal(
|
||||
"Agent Identity auth_config changed while task registration was in flight"
|
||||
"Agent Identity credential changed while task registration was in flight"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
@@ -2093,6 +2157,7 @@ impl AppState {
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: Some(expected_encrypted_auth_config),
|
||||
expected_credential: Some(expected_credential_fence.credential.clone()),
|
||||
encrypted_auth_config,
|
||||
encrypted_api_key_update: None,
|
||||
expires_at_unix_secs_update: None,
|
||||
@@ -2147,17 +2212,13 @@ impl AppState {
|
||||
.map(|value| encrypt_python_fernet_plaintext(encryption_key, value.as_str()))
|
||||
.transpose()
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let requires_fenced_persistence = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
&& transport.key.auth_type.trim().eq_ignore_ascii_case("oauth");
|
||||
let requires_fenced_persistence =
|
||||
provider_transport::supports_local_oauth_request_auth_resolution(transport);
|
||||
if requires_fenced_persistence
|
||||
&& (expected_auth_config.is_none() || encrypted_auth_config.is_none())
|
||||
&& (expected_credential_fence.is_none() || encrypted_auth_config.is_none())
|
||||
{
|
||||
return Err(GatewayError::Internal(
|
||||
"Codex OAuth refresh persistence is missing its auth_config fence".to_string(),
|
||||
"OAuth refresh persistence is missing its credential fence".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
@@ -2172,7 +2233,13 @@ impl AppState {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let observed_encrypted_auth_config = latest_key.encrypted_auth_config.clone();
|
||||
let observed_credential_matches = expected_credential_fence.is_none_or(|expected| {
|
||||
latest_key.encrypted_auth_config.as_deref()
|
||||
== Some(expected.encrypted_auth_config.as_str())
|
||||
&& latest_key.encrypted_api_key == expected.credential.encrypted_api_key
|
||||
&& latest_key.auth_type == expected.credential.auth_type
|
||||
&& latest_key.provider_id == expected.credential.provider_id
|
||||
});
|
||||
latest_key.encrypted_api_key = Some(encrypted_api_key.clone());
|
||||
latest_key.encrypted_auth_config = encrypted_auth_config.clone();
|
||||
latest_key.expires_at_unix_secs = entry.expires_at_unix_secs;
|
||||
@@ -2191,17 +2258,20 @@ impl AppState {
|
||||
latest_key.status_snapshot =
|
||||
sync_provider_key_oauth_status_snapshot(current_status_snapshot, &latest_key);
|
||||
let used_fenced_persistence =
|
||||
expected_auth_config.is_some() && encrypted_auth_config.is_some();
|
||||
let updated = if let (Some(expected_auth_config), Some(encrypted_auth_config)) =
|
||||
(expected_auth_config, encrypted_auth_config.as_deref())
|
||||
expected_credential_fence.is_some() && encrypted_auth_config.is_some();
|
||||
let updated = if let (Some(expected_credential_fence), Some(encrypted_auth_config)) =
|
||||
(expected_credential_fence, encrypted_auth_config.as_deref())
|
||||
{
|
||||
if observed_encrypted_auth_config.as_deref() != Some(expected_auth_config) {
|
||||
if !observed_credential_matches {
|
||||
false
|
||||
} else {
|
||||
self.compare_and_update_provider_catalog_key_oauth_runtime_state(
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: Some(expected_auth_config.to_string()),
|
||||
expected_encrypted_auth_config: Some(
|
||||
expected_credential_fence.encrypted_auth_config.clone(),
|
||||
),
|
||||
expected_credential: Some(expected_credential_fence.credential.clone()),
|
||||
encrypted_auth_config: encrypted_auth_config.to_string(),
|
||||
encrypted_api_key_update: Some(encrypted_api_key.clone()),
|
||||
expires_at_unix_secs_update: Some(entry.expires_at_unix_secs),
|
||||
@@ -2250,7 +2320,7 @@ impl AppState {
|
||||
};
|
||||
if !updated && (requires_fenced_persistence || used_fenced_persistence) {
|
||||
return Err(GatewayError::Internal(
|
||||
"Codex OAuth credential changed during refresh persistence".to_string(),
|
||||
"OAuth credential changed during refresh persistence".to_string(),
|
||||
));
|
||||
}
|
||||
let metadata_refresh_token_fingerprint =
|
||||
@@ -2295,13 +2365,13 @@ impl AppState {
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
let expected_auth_config = if transport_has_auth_config {
|
||||
self.capture_provider_transport_auth_config_fence(transport)
|
||||
let expected_credential_fence = if transport_has_auth_config {
|
||||
self.capture_provider_transport_credential_fence(transport)
|
||||
.await?
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if transport_has_auth_config && expected_auth_config.is_none() {
|
||||
if transport_has_auth_config && expected_credential_fence.is_none() {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
@@ -2320,10 +2390,13 @@ impl AppState {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
if expected_auth_config
|
||||
.as_deref()
|
||||
.is_some_and(|expected| latest_key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
if expected_credential_fence.as_ref().is_some_and(|expected| {
|
||||
latest_key.encrypted_auth_config.as_deref()
|
||||
!= Some(expected.encrypted_auth_config.as_str())
|
||||
|| latest_key.encrypted_api_key != expected.credential.encrypted_api_key
|
||||
|| latest_key.auth_type != expected.credential.auth_type
|
||||
|| latest_key.provider_id != expected.credential.provider_id
|
||||
}) {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
@@ -2356,13 +2429,18 @@ impl AppState {
|
||||
latest_key.status_snapshot =
|
||||
sync_provider_key_oauth_status_snapshot(current_status_snapshot, &latest_key);
|
||||
|
||||
if let Some(expected_auth_config) = expected_auth_config.as_ref() {
|
||||
if let Some(expected_credential_fence) = expected_credential_fence.as_ref() {
|
||||
updated = self
|
||||
.compare_and_update_provider_catalog_key_oauth_runtime_state(
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: Some(expected_auth_config.clone()),
|
||||
encrypted_auth_config: expected_auth_config.clone(),
|
||||
expected_encrypted_auth_config: Some(
|
||||
expected_credential_fence.encrypted_auth_config.clone(),
|
||||
),
|
||||
expected_credential: Some(expected_credential_fence.credential.clone()),
|
||||
encrypted_auth_config: expected_credential_fence
|
||||
.encrypted_auth_config
|
||||
.clone(),
|
||||
encrypted_api_key_update: None,
|
||||
expires_at_unix_secs_update: None,
|
||||
oauth_invalid_at_unix_secs: latest_key.oauth_invalid_at_unix_secs,
|
||||
@@ -2407,11 +2485,12 @@ impl AppState {
|
||||
// Codex credentials are replaceable under a stable key id. Without a
|
||||
// conditional delete, refresh failure handling must retain them after
|
||||
// writing the generation-fenced marker.
|
||||
let auto_removed = if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
let auto_removed = if expected_credential_fence.is_none()
|
||||
&& !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
&& admin_provider_quota_pure::provider_auto_remove_banned_keys(
|
||||
transport.provider.config.as_ref(),
|
||||
)
|
||||
@@ -2909,6 +2988,60 @@ mod tests {
|
||||
(state, repository, encrypted_auth_config)
|
||||
}
|
||||
|
||||
fn vertex_service_account_state(
|
||||
) -> (AppState, Arc<InMemoryProviderCatalogReadRepository>, String) {
|
||||
let mut provider = sample_provider();
|
||||
provider.provider_type = "vertex_ai".to_string();
|
||||
let mut endpoint = sample_endpoint();
|
||||
endpoint.api_format = "gemini:generate_content".to_string();
|
||||
endpoint.api_family = Some("gemini".to_string());
|
||||
endpoint.base_url = "https://aiplatform.googleapis.com".to_string();
|
||||
let auth_config = json!({
|
||||
"client_email": "[email protected]",
|
||||
"private_key": "TEST-PRIVATE-KEY",
|
||||
"project_id": "demo-project"
|
||||
});
|
||||
let encrypted_auth_config =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &auth_config.to_string())
|
||||
.expect("Vertex auth config should encrypt");
|
||||
let encrypted_api_key =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__")
|
||||
.expect("Vertex placeholder should encrypt");
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"Vertex service account".to_string(),
|
||||
"service_account".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("Vertex key should build")
|
||||
.with_transport_fields(
|
||||
Some(json!(["gemini:generate_content"])),
|
||||
encrypted_api_key,
|
||||
Some(encrypted_auth_config.clone()),
|
||||
None,
|
||||
Some(json!({"gemini:generate_content": 1})),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("Vertex key transport should build");
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![endpoint],
|
||||
vec![key],
|
||||
));
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(repository.clone())
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
(state, repository, encrypted_auth_config)
|
||||
}
|
||||
|
||||
fn state_with_global_format_conversion(enabled: bool) -> AppState {
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
@@ -3769,6 +3902,65 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn memory_only_vertex_refresh_rejects_replaced_credential_generation() {
|
||||
let (state, repository, initial_encrypted_auth_config) = vertex_service_account_state();
|
||||
let transport = state
|
||||
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
|
||||
.await
|
||||
.expect("Vertex transport should load")
|
||||
.expect("Vertex transport should exist");
|
||||
let expected_credential_fence = state
|
||||
.capture_provider_transport_credential_fence(&transport)
|
||||
.await
|
||||
.expect("Vertex fence should load")
|
||||
.expect("Vertex fence should match");
|
||||
let entry = crate::provider_transport::CachedOAuthEntry {
|
||||
provider_type: "vertex_ai".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer memory-only-token".to_string(),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
source_fingerprint: Some("service-account-generation".to_string()),
|
||||
};
|
||||
|
||||
state
|
||||
.persist_local_oauth_refresh_entry(&transport, &entry, Some(&expected_credential_fence))
|
||||
.await
|
||||
.expect("unchanged Vertex credential should accept memory-only token");
|
||||
let unchanged = repository
|
||||
.list_keys_by_ids(&["key-1".to_string()])
|
||||
.await
|
||||
.expect("Vertex key should load")
|
||||
.pop()
|
||||
.expect("Vertex key should exist");
|
||||
assert_eq!(
|
||||
unchanged.encrypted_auth_config.as_deref(),
|
||||
Some(initial_encrypted_auth_config.as_str())
|
||||
);
|
||||
|
||||
let mut replacement = unchanged;
|
||||
replacement.encrypted_api_key = Some(
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "admin-replacement")
|
||||
.expect("replacement credential should encrypt"),
|
||||
);
|
||||
repository
|
||||
.update_key(&replacement)
|
||||
.await
|
||||
.expect("replacement credential should persist");
|
||||
|
||||
assert!(
|
||||
state
|
||||
.persist_local_oauth_refresh_entry(
|
||||
&transport,
|
||||
&entry,
|
||||
Some(&expected_credential_fence),
|
||||
)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn failed_refresh_persistence_discards_provisional_auth_and_cache_entry() {
|
||||
let mut resolution = Some(crate::provider_transport::LocalOAuthResolution {
|
||||
@@ -3789,6 +3981,7 @@ mod tests {
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
local_refresh_guard: None,
|
||||
});
|
||||
|
||||
super::discard_failed_local_oauth_refresh_resolution(&mut resolution);
|
||||
@@ -3860,13 +4053,18 @@ mod tests {
|
||||
"email": "[email protected]",
|
||||
"expires_at": 4102444800_u64
|
||||
});
|
||||
let (state, repository, expected_auth_config) =
|
||||
let (state, repository, _expected_auth_config) =
|
||||
codex_oauth_state(&initial_config, "access-old");
|
||||
let transport = state
|
||||
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
|
||||
.await
|
||||
.expect("transport should load")
|
||||
.expect("transport should exist");
|
||||
let expected_credential_fence = state
|
||||
.capture_provider_transport_credential_fence(&transport)
|
||||
.await
|
||||
.expect("credential fence should load")
|
||||
.expect("credential fence should match the initial transport");
|
||||
|
||||
let replacement_config = json!({
|
||||
"provider_type": "codex",
|
||||
@@ -3914,7 +4112,7 @@ mod tests {
|
||||
.persist_local_oauth_refresh_entry(
|
||||
&transport,
|
||||
&refreshed_entry,
|
||||
Some(expected_auth_config.as_str()),
|
||||
Some(&expected_credential_fence),
|
||||
)
|
||||
.await
|
||||
.is_err());
|
||||
@@ -3935,4 +4133,108 @@ mod tests {
|
||||
);
|
||||
assert_eq!(stored.expires_at_unix_secs, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_refresh_failure_does_not_mark_access_token_only_replacement() {
|
||||
let initial_config = json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "refresh-stable",
|
||||
"expires_at": 4102444800_u64
|
||||
});
|
||||
let (state, repository, _) = codex_oauth_state(&initial_config, "access-old");
|
||||
let stale_transport = state
|
||||
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
|
||||
.await
|
||||
.expect("transport should load")
|
||||
.expect("transport should exist");
|
||||
|
||||
let replacement_api_key =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "access-admin")
|
||||
.expect("replacement api key should encrypt");
|
||||
let mut replaced = repository
|
||||
.list_keys_by_ids(&["key-1".to_string()])
|
||||
.await
|
||||
.expect("key should load")
|
||||
.pop()
|
||||
.expect("key should exist");
|
||||
replaced.encrypted_api_key = Some(replacement_api_key.clone());
|
||||
repository
|
||||
.update_key(&replaced)
|
||||
.await
|
||||
.expect("access token replacement should persist");
|
||||
|
||||
assert!(!state
|
||||
.persist_local_oauth_refresh_failure_state(
|
||||
&stale_transport,
|
||||
401,
|
||||
r#"{"error":"invalid_token"}"#,
|
||||
true,
|
||||
)
|
||||
.await
|
||||
.expect("stale failure should be ignored"));
|
||||
|
||||
let stored = repository
|
||||
.list_keys_by_ids(&["key-1".to_string()])
|
||||
.await
|
||||
.expect("key should reload")
|
||||
.pop()
|
||||
.expect("replacement should remain");
|
||||
assert_eq!(
|
||||
stored.encrypted_api_key.as_deref(),
|
||||
Some(replacement_api_key.as_str())
|
||||
);
|
||||
assert!(stored.oauth_invalid_at_unix_secs.is_none());
|
||||
assert!(stored.oauth_invalid_reason.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_request_invalidation_does_not_mark_access_token_only_replacement() {
|
||||
let initial_config = json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "refresh-stable",
|
||||
"expires_at": 4102444800_u64
|
||||
});
|
||||
let (state, repository, _) = codex_oauth_state(&initial_config, "access-old");
|
||||
let stale_transport = state
|
||||
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
|
||||
.await
|
||||
.expect("transport should load")
|
||||
.expect("transport should exist");
|
||||
|
||||
let replacement_api_key =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "access-admin")
|
||||
.expect("replacement api key should encrypt");
|
||||
let mut replaced = repository
|
||||
.list_keys_by_ids(&["key-1".to_string()])
|
||||
.await
|
||||
.expect("key should load")
|
||||
.pop()
|
||||
.expect("key should exist");
|
||||
replaced.encrypted_api_key = Some(replacement_api_key.clone());
|
||||
repository
|
||||
.update_key(&replaced)
|
||||
.await
|
||||
.expect("access token replacement should persist");
|
||||
|
||||
assert!(!state
|
||||
.mark_provider_transport_oauth_invalid_fenced(
|
||||
&stale_transport,
|
||||
"[OAUTH_EXPIRED] stale request",
|
||||
)
|
||||
.await
|
||||
.expect("stale invalidation should be ignored"));
|
||||
|
||||
let stored = repository
|
||||
.list_keys_by_ids(&["key-1".to_string()])
|
||||
.await
|
||||
.expect("key should reload")
|
||||
.pop()
|
||||
.expect("replacement should remain");
|
||||
assert_eq!(
|
||||
stored.encrypted_api_key.as_deref(),
|
||||
Some(replacement_api_key.as_str())
|
||||
);
|
||||
assert!(stored.oauth_invalid_at_unix_secs.is_none());
|
||||
assert!(stored.oauth_invalid_reason.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user