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:
elky
2026-07-27 09:36:31 +08:00
parent 79b70f7b5c
commit 531cf11025
152 changed files with 13984 additions and 2075 deletions
@@ -54,6 +54,223 @@ impl LocalFailoverClassification {
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum FailureRetryAction {
Stop,
SameCredential,
NextCandidate,
NextCredential,
NextEndpoint,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum FailureScope {
None,
Credential,
CredentialModel,
Endpoint,
Provider,
}
impl FailureScope {
pub(crate) const fn affects_credential(self) -> bool {
matches!(self, Self::Credential | Self::CredentialModel)
}
pub(crate) const fn allows_key_wide_effects(self) -> bool {
matches!(self, Self::None | Self::Credential)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum FailureTokenAction {
None,
ForceRefresh,
#[allow(dead_code)]
Quarantine,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct FailureDisposition {
pub(crate) retry_action: FailureRetryAction,
pub(crate) failure_scope: FailureScope,
pub(crate) token_action: FailureTokenAction,
pub(crate) preserve_upstream_error: bool,
}
impl FailureDisposition {
const fn new(
retry_action: FailureRetryAction,
failure_scope: FailureScope,
token_action: FailureTokenAction,
preserve_upstream_error: bool,
) -> Self {
Self {
retry_action,
failure_scope,
token_action,
preserve_upstream_error,
}
}
}
pub(crate) const fn failure_disposition_from_local_classification(
classification: LocalFailoverClassification,
status_code: u16,
) -> FailureDisposition {
match classification {
LocalFailoverClassification::StopStatusCode
| LocalFailoverClassification::StopErrorPattern
| LocalFailoverClassification::StopExecutionError
| LocalFailoverClassification::StopCyberPolicy => FailureDisposition::new(
FailureRetryAction::Stop,
FailureScope::None,
FailureTokenAction::None,
true,
),
LocalFailoverClassification::UseDefault => FailureDisposition::new(
FailureRetryAction::Stop,
FailureScope::None,
FailureTokenAction::None,
status_code >= 400,
),
LocalFailoverClassification::RetrySuccessPattern => FailureDisposition::new(
FailureRetryAction::NextCandidate,
FailureScope::None,
FailureTokenAction::None,
false,
),
LocalFailoverClassification::RetryStatusCode
| LocalFailoverClassification::RetryUpstreamFailure => FailureDisposition::new(
FailureRetryAction::NextCandidate,
FailureScope::None,
FailureTokenAction::None,
false,
),
}
}
pub(crate) const fn classify_anthropic_failure_disposition(
classification: LocalFailoverClassification,
status_code: u16,
) -> FailureDisposition {
if matches!(
classification,
LocalFailoverClassification::StopStatusCode
| LocalFailoverClassification::StopErrorPattern
| LocalFailoverClassification::StopExecutionError
| LocalFailoverClassification::StopCyberPolicy
) {
let generic = failure_disposition_from_local_classification(classification, status_code);
return match status_code {
401 => FailureDisposition::new(
generic.retry_action,
FailureScope::Credential,
FailureTokenAction::ForceRefresh,
generic.preserve_upstream_error,
),
403 => FailureDisposition::new(
generic.retry_action,
FailureScope::Credential,
FailureTokenAction::None,
generic.preserve_upstream_error,
),
404 => FailureDisposition::new(
generic.retry_action,
FailureScope::Endpoint,
FailureTokenAction::None,
generic.preserve_upstream_error,
),
429 => FailureDisposition::new(
generic.retry_action,
FailureScope::CredentialModel,
FailureTokenAction::None,
generic.preserve_upstream_error,
),
529 => FailureDisposition::new(
generic.retry_action,
FailureScope::Provider,
FailureTokenAction::None,
generic.preserve_upstream_error,
),
500..=599 => FailureDisposition::new(
generic.retry_action,
FailureScope::Endpoint,
FailureTokenAction::None,
generic.preserve_upstream_error,
),
_ => generic,
};
}
match status_code {
400 => FailureDisposition::new(
FailureRetryAction::Stop,
FailureScope::None,
FailureTokenAction::None,
true,
),
401 => FailureDisposition::new(
FailureRetryAction::NextCredential,
FailureScope::Credential,
FailureTokenAction::ForceRefresh,
true,
),
403 => FailureDisposition::new(
FailureRetryAction::NextCredential,
FailureScope::Credential,
FailureTokenAction::None,
true,
),
404 => FailureDisposition::new(
FailureRetryAction::NextEndpoint,
FailureScope::Endpoint,
FailureTokenAction::None,
true,
),
413 => FailureDisposition::new(
FailureRetryAction::Stop,
FailureScope::None,
FailureTokenAction::None,
true,
),
429 => FailureDisposition::new(
FailureRetryAction::NextCredential,
FailureScope::CredentialModel,
FailureTokenAction::None,
true,
),
529 => FailureDisposition::new(
FailureRetryAction::NextEndpoint,
FailureScope::Provider,
FailureTokenAction::None,
true,
),
500..=599 => FailureDisposition::new(
FailureRetryAction::NextEndpoint,
FailureScope::Endpoint,
FailureTokenAction::None,
true,
),
_ => failure_disposition_from_local_classification(classification, status_code),
}
}
pub(crate) fn classify_failure_disposition(
provider_api_format: &str,
classification: LocalFailoverClassification,
status_code: u16,
) -> FailureDisposition {
if provider_api_format
.trim()
.eq_ignore_ascii_case("claude:messages")
{
classify_anthropic_failure_disposition(classification, status_code)
} else {
failure_disposition_from_local_classification(classification, status_code)
}
}
pub(crate) fn classify_local_failover(
policy: &LocalFailoverPolicy,
input: LocalFailoverInput<'_>,
@@ -267,7 +484,11 @@ fn local_failover_regex_rule_matches(
mod tests {
use std::collections::BTreeSet;
use super::{classify_local_failover, LocalFailoverClassification, LocalFailoverInput};
use super::{
classify_anthropic_failure_disposition, classify_local_failover,
failure_disposition_from_local_classification, FailureDisposition, FailureRetryAction,
FailureScope, FailureTokenAction, LocalFailoverClassification, LocalFailoverInput,
};
use crate::orchestration::{LocalFailoverPolicy, LocalFailoverRegexRule};
#[test]
@@ -544,4 +765,138 @@ mod tests {
LocalFailoverClassification::UseDefault
);
}
#[test]
fn legacy_classification_preserves_candidate_by_candidate_retry() {
assert_eq!(
failure_disposition_from_local_classification(
LocalFailoverClassification::RetryUpstreamFailure,
429,
),
FailureDisposition {
retry_action: FailureRetryAction::NextCandidate,
failure_scope: FailureScope::None,
token_action: FailureTokenAction::None,
preserve_upstream_error: false,
}
);
assert_eq!(
failure_disposition_from_local_classification(
LocalFailoverClassification::StopErrorPattern,
400,
)
.retry_action,
FailureRetryAction::Stop
);
}
#[test]
fn anthropic_bad_request_stops_and_preserves_upstream_error() {
let disposition = classify_anthropic_failure_disposition(
LocalFailoverClassification::RetryUpstreamFailure,
400,
);
assert_eq!(disposition.retry_action, FailureRetryAction::Stop);
assert_eq!(disposition.failure_scope, FailureScope::None);
assert_eq!(disposition.token_action, FailureTokenAction::None);
assert!(disposition.preserve_upstream_error);
}
#[test]
fn anthropic_auth_failures_refresh_then_rotate_only_when_needed() {
let unauthorized = classify_anthropic_failure_disposition(
LocalFailoverClassification::RetryUpstreamFailure,
401,
);
assert_eq!(
unauthorized.retry_action,
FailureRetryAction::NextCredential
);
assert_eq!(unauthorized.failure_scope, FailureScope::Credential);
assert_eq!(unauthorized.token_action, FailureTokenAction::ForceRefresh);
let forbidden = classify_anthropic_failure_disposition(
LocalFailoverClassification::RetryUpstreamFailure,
403,
);
assert_eq!(forbidden.retry_action, FailureRetryAction::NextCredential);
assert_eq!(forbidden.failure_scope, FailureScope::Credential);
assert_eq!(forbidden.token_action, FailureTokenAction::None);
}
#[test]
fn anthropic_rate_limit_rotates_with_credential_model_scope() {
let disposition = classify_anthropic_failure_disposition(
LocalFailoverClassification::RetryUpstreamFailure,
429,
);
assert_eq!(disposition.retry_action, FailureRetryAction::NextCredential);
assert_eq!(disposition.failure_scope, FailureScope::CredentialModel);
assert!(disposition.failure_scope.affects_credential());
assert!(!disposition.failure_scope.allows_key_wide_effects());
assert!(disposition.preserve_upstream_error);
}
#[test]
fn anthropic_overload_moves_endpoint_without_credential_penalty() {
let disposition = classify_anthropic_failure_disposition(
LocalFailoverClassification::RetryUpstreamFailure,
529,
);
assert_eq!(disposition.retry_action, FailureRetryAction::NextEndpoint);
assert_eq!(disposition.failure_scope, FailureScope::Provider);
assert!(!disposition.failure_scope.affects_credential());
assert!(!disposition.failure_scope.allows_key_wide_effects());
assert_eq!(disposition.token_action, FailureTokenAction::None);
assert!(disposition.preserve_upstream_error);
}
#[test]
fn anthropic_not_found_moves_endpoint_and_oversize_stops() {
let not_found = classify_anthropic_failure_disposition(
LocalFailoverClassification::RetryUpstreamFailure,
404,
);
assert_eq!(not_found.retry_action, FailureRetryAction::NextEndpoint);
assert_eq!(not_found.failure_scope, FailureScope::Endpoint);
assert!(not_found.preserve_upstream_error);
let oversized = classify_anthropic_failure_disposition(
LocalFailoverClassification::RetryUpstreamFailure,
413,
);
assert_eq!(oversized.retry_action, FailureRetryAction::Stop);
assert_eq!(oversized.failure_scope, FailureScope::None);
assert!(oversized.preserve_upstream_error);
}
#[test]
fn only_unscoped_and_credential_failures_allow_key_wide_effects() {
assert!(FailureScope::None.allows_key_wide_effects());
assert!(FailureScope::Credential.allows_key_wide_effects());
assert!(!FailureScope::CredentialModel.allows_key_wide_effects());
assert!(!FailureScope::Endpoint.allows_key_wide_effects());
assert!(!FailureScope::Provider.allows_key_wide_effects());
}
#[test]
fn anthropic_explicit_stop_keeps_failure_resource_scope() {
let auth = classify_anthropic_failure_disposition(
LocalFailoverClassification::StopStatusCode,
401,
);
assert_eq!(auth.retry_action, FailureRetryAction::Stop);
assert_eq!(auth.failure_scope, FailureScope::Credential);
assert_eq!(auth.token_action, FailureTokenAction::ForceRefresh);
let overloaded = classify_anthropic_failure_disposition(
LocalFailoverClassification::StopStatusCode,
529,
);
assert_eq!(overloaded.retry_action, FailureRetryAction::Stop);
assert_eq!(overloaded.failure_scope, FailureScope::Provider);
}
}
+468 -50
View File
@@ -26,9 +26,10 @@ use tokio::sync::Mutex as TokioMutex;
use tracing::warn;
use super::{
local_failover_error_message, project_local_adaptive_rate_limit,
classify_failure_disposition, local_failover_error_message, project_local_adaptive_rate_limit,
project_local_adaptive_success, project_local_failure_health, project_local_key_circuit_closed,
project_local_key_circuit_failure, project_local_success_health, LocalFailoverClassification,
project_local_key_circuit_failure, project_local_success_health, FailureScope,
LocalFailoverClassification,
};
use crate::ai_serving::extract_pool_sticky_session_token;
use crate::client_session_affinity::{
@@ -613,7 +614,8 @@ async fn record_attempt_failure_effect(
context: LocalExecutionEffectContext<'_>,
effect: LocalAttemptFailureEffect,
) {
if !local_candidate_failure_should_invalidate_affinity(
if !local_candidate_failure_should_invalidate_affinity_for_provider(
&context.plan.provider_api_format,
effect.classification,
effect.status_code,
) {
@@ -683,6 +685,13 @@ async fn record_adaptive_rate_limit_effect(
context: LocalExecutionEffectContext<'_>,
effect: LocalAdaptiveRateLimitEffect<'_>,
) {
if !local_candidate_failure_should_apply_key_effects(
&context.plan.provider_api_format,
effect.classification,
effect.status_code,
) {
return;
}
let Some(auth_config_fence) =
capture_local_execution_auth_config_fence(state, context.plan).await
else {
@@ -964,6 +973,13 @@ async fn record_health_failure_effect(
context: LocalExecutionEffectContext<'_>,
effect: LocalHealthFailureEffect,
) {
if !local_candidate_failure_should_apply_key_effects(
&context.plan.provider_api_format,
effect.classification,
effect.status_code,
) {
return;
}
let api_format = context.plan.provider_api_format.trim();
if api_format.is_empty() {
return;
@@ -1231,6 +1247,13 @@ async fn record_pool_error_effect(
context: LocalExecutionEffectContext<'_>,
effect: LocalPoolErrorEffect<'_>,
) {
if !local_candidate_failure_should_apply_key_effects(
&context.plan.provider_api_format,
effect.classification,
effect.status_code,
) {
return;
}
let terminal_error_reason =
admin_provider_pool_key_terminal_error_reason(effect.status_code, effect.error_body);
if terminal_error_reason.is_none()
@@ -1379,13 +1402,7 @@ async fn record_oauth_invalidation_effect(
if !transport.key.auth_type.trim().eq_ignore_ascii_case("oauth") {
return;
}
if transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("codex")
&& !execution_plan_bearer_matches_transport(plan, &transport)
{
if !execution_plan_bearer_matches_transport(plan, &transport) {
return;
}
@@ -1397,28 +1414,8 @@ async fn record_oauth_invalidation_effect(
return;
};
let expected_auth_config = match state
.capture_provider_transport_auth_config_fence(&transport)
.await
{
Ok(Some(ciphertext)) => ciphertext,
Ok(None) => return,
Err(err) => {
warn!(
"gateway orchestration effects: failed to capture oauth invalidation fence for provider {} endpoint {} key {}: {:?}",
plan.provider_id, plan.endpoint_id, plan.key_id, err
);
return;
}
};
if let Err(err) = state
.mark_provider_catalog_key_oauth_invalid_fenced(
&plan.key_id,
transport.provider.provider_type.as_str(),
invalid_reason.as_str(),
expected_auth_config.as_str(),
)
.mark_provider_transport_oauth_invalid_fenced(&transport, invalid_reason.as_str())
.await
{
warn!(
@@ -1444,16 +1441,20 @@ fn execution_plan_bearer_matches_transport(
plan: &ExecutionPlan,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
) -> bool {
let current_token = transport.key.decrypted_api_key.trim();
!current_token.is_empty()
&& plan.headers.iter().any(|(name, value)| {
name.eq_ignore_ascii_case("authorization")
&& value
.trim()
.strip_prefix("Bearer ")
.map(str::trim)
.is_some_and(|token| token == current_token)
})
let Some(plan_token) = execution_plan_authorization(plan).and_then(bearer_access_token) else {
return false;
};
crate::provider_transport::resolve_local_generic_oauth_transport_authorization(transport)
.as_deref()
.and_then(bearer_access_token)
.is_some_and(|current_token| current_token == plan_token)
}
fn bearer_access_token(authorization: &str) -> Option<&str> {
let mut parts = authorization.split_ascii_whitespace();
let scheme = parts.next()?;
let token = parts.next()?;
(scheme.eq_ignore_ascii_case("bearer") && parts.next().is_none()).then_some(token)
}
fn resolve_local_oauth_invalid_reason(
@@ -1467,6 +1468,12 @@ fn resolve_local_oauth_invalid_reason(
status_code,
upstream_message.as_deref(),
),
_ if super::oauth_status_may_be_invalid(status_code, response_text) => Some(format!(
"[OAUTH_EXPIRED] {}",
upstream_message
.as_deref()
.unwrap_or("OAuth access token was rejected")
)),
_ => None,
}
}
@@ -1492,6 +1499,46 @@ fn local_candidate_failure_should_invalidate_affinity(
}
}
fn local_candidate_failure_should_invalidate_affinity_for_provider(
provider_api_format: &str,
classification: LocalFailoverClassification,
status_code: u16,
) -> bool {
if !local_candidate_failure_should_invalidate_affinity(classification, status_code) {
return false;
}
if !provider_api_format
.trim()
.eq_ignore_ascii_case("claude:messages")
{
return true;
}
let disposition =
classify_failure_disposition(provider_api_format, classification, status_code);
!(disposition.retry_action == crate::orchestration::FailureRetryAction::Stop
&& disposition.failure_scope == FailureScope::None)
}
fn local_candidate_failure_should_apply_key_effects(
provider_api_format: &str,
classification: LocalFailoverClassification,
status_code: u16,
) -> bool {
if !provider_api_format
.trim()
.eq_ignore_ascii_case("claude:messages")
{
return true;
}
matches!(
classify_failure_disposition(provider_api_format, classification, status_code)
.failure_scope,
FailureScope::Credential
)
}
fn local_candidate_failure_should_record_pool_error(
classification: LocalFailoverClassification,
status_code: u16,
@@ -1708,12 +1755,13 @@ mod tests {
use serde_json::{json, Value};
use super::{
apply_local_execution_effect, local_candidate_failure_should_record_pool_error,
pool_score_feedback_gate_allows, pool_score_hard_state_for_status,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect,
ProviderKeyEffectLockPool,
apply_local_execution_effect, execution_plan_bearer_matches_transport,
local_candidate_failure_should_apply_key_effects,
local_candidate_failure_should_record_pool_error, pool_score_feedback_gate_allows,
pool_score_hard_state_for_status, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect,
LocalAttemptFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext,
LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect,
LocalPoolErrorEffect, ProviderKeyEffectLockPool,
};
use crate::data::{GatewayDataConfig, GatewayDataState};
use crate::orchestration::LocalFailoverClassification;
@@ -1760,6 +1808,13 @@ mod tests {
}
}
fn sample_claude_plan() -> ExecutionPlan {
let mut plan = sample_plan();
plan.provider_name = Some("anthropic".to_string());
plan.provider_api_format = "claude:messages".to_string();
plan
}
#[test]
fn pool_score_feedback_gate_suppresses_repeated_success_writes() {
super::POOL_SCORE_FEEDBACK_GATE.clear();
@@ -1864,7 +1919,7 @@ mod tests {
url: "https://chatgpt.com/backend-api/codex".to_string(),
headers: BTreeMap::from([(
"authorization".to_string(),
"Bearer __placeholder__".to_string(),
"Bearer codex-access-token".to_string(),
)]),
content_type: Some("application/json".to_string()),
content_encoding: None,
@@ -1959,8 +2014,8 @@ mod tests {
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["openai:responses"])),
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__")
.expect("placeholder api key should encrypt"),
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "codex-access-token")
.expect("access token should encrypt"),
Some(encrypted_auth_config),
None,
Some(serde_json::json!({"openai:responses": 1})),
@@ -1975,6 +2030,10 @@ mod tests {
fn sample_codex_agent_identity_key() -> StoredProviderCatalogKey {
let mut key = sample_codex_key();
key.name = "Agent Identity".to_string();
key.encrypted_api_key = Some(
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__")
.expect("placeholder api key should encrypt"),
);
key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
@@ -2014,6 +2073,36 @@ mod tests {
)
}
fn claude_code_oauth_state() -> AppState {
let mut provider = sample_codex_provider();
provider.name = "claude_code".to_string();
provider.provider_type = "claude_code".to_string();
let mut endpoint = sample_codex_endpoint();
endpoint.api_format = "claude:messages".to_string();
endpoint.api_family = Some("claude".to_string());
endpoint.base_url = "https://api.anthropic.com".to_string();
let mut key = sample_codex_key();
key.api_formats = Some(json!(["claude:messages"]));
key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"provider_type":"claude_code","refresh_token":"rt-claude-local-123"}"#,
)
.expect("Claude Code auth config should encrypt"),
);
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
}
fn codex_state_with_redis(redis_url: &str, redis_key_prefix: &str) -> AppState {
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_codex_provider()],
@@ -2724,6 +2813,179 @@ mod tests {
));
}
#[test]
fn anthropic_non_credential_failures_do_not_apply_key_wide_effects() {
assert!(!local_candidate_failure_should_apply_key_effects(
"claude:messages",
LocalFailoverClassification::RetryUpstreamFailure,
529,
));
assert!(!local_candidate_failure_should_apply_key_effects(
"claude:messages",
LocalFailoverClassification::RetryUpstreamFailure,
429,
));
assert!(!local_candidate_failure_should_apply_key_effects(
"claude:messages",
LocalFailoverClassification::RetryUpstreamFailure,
503,
));
assert!(local_candidate_failure_should_apply_key_effects(
"claude:messages",
LocalFailoverClassification::RetryUpstreamFailure,
401,
));
assert!(local_candidate_failure_should_apply_key_effects(
"claude:messages",
LocalFailoverClassification::RetryUpstreamFailure,
403,
));
assert!(!local_candidate_failure_should_apply_key_effects(
"claude:messages",
LocalFailoverClassification::RetryUpstreamFailure,
400,
));
assert!(local_candidate_failure_should_apply_key_effects(
"openai:chat",
LocalFailoverClassification::RetryUpstreamFailure,
529,
));
assert!(local_candidate_failure_should_apply_key_effects(
"openai:chat",
LocalFailoverClassification::RetryUpstreamFailure,
429,
));
assert!(local_candidate_failure_should_apply_key_effects(
"openai:chat",
LocalFailoverClassification::RetryUpstreamFailure,
503,
));
}
#[tokio::test]
async fn anthropic_non_credential_failures_preserve_key_wide_state() {
for status_code in [400, 429, 503, 529] {
let mut key = sample_adaptive_key();
let circuit = json!({
"openai:chat": {
"open": true,
"reason": "existing-state"
}
});
key.circuit_breaker_by_format = Some(circuit.clone());
let expected_adaptive_state = ProviderCatalogKeyAdaptiveState::from(&key);
let expected_health = key.health_by_format.clone();
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_pool_health_provider()],
vec![sample_health_endpoint()],
vec![key],
));
let state = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
);
let plan = sample_claude_plan();
let report_context = json!({
"api_key_id": "api-key-1",
"client_api_format": "claude:messages",
"model": "claude-sonnet-4-5",
});
let cache_key = build_scheduler_affinity_cache_key_for_api_key_id(
"api-key-1",
"claude:messages",
"claude-sonnet-4-5",
)
.expect("scheduler affinity cache key should build");
let target = SchedulerAffinityTarget {
provider_id: plan.provider_id.clone(),
endpoint_id: plan.endpoint_id.clone(),
key_id: plan.key_id.clone(),
};
state.remember_scheduler_affinity_target(
&cache_key,
target.clone(),
SCHEDULER_AFFINITY_TTL,
16,
);
let headers = BTreeMap::from([("Retry-After".to_string(), "120".to_string())]);
let context = LocalExecutionEffectContext {
plan: &plan,
report_context: Some(&report_context),
};
let classification = LocalFailoverClassification::RetryUpstreamFailure;
apply_local_execution_effect(
&state,
context,
LocalExecutionEffect::AttemptFailure(LocalAttemptFailureEffect {
status_code,
classification,
}),
)
.await;
apply_local_execution_effect(
&state,
context,
LocalExecutionEffect::AdaptiveRateLimit(LocalAdaptiveRateLimitEffect {
status_code,
classification,
headers: Some(&headers),
}),
)
.await;
apply_local_execution_effect(
&state,
context,
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
status_code,
classification,
}),
)
.await;
apply_local_execution_effect(
&state,
context,
LocalExecutionEffect::PoolError(LocalPoolErrorEffect {
status_code,
classification,
headers: &headers,
error_body: Some(r#"{"error":{"message":"temporarily unavailable"}}"#),
}),
)
.await;
let stored_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
assert_eq!(
ProviderCatalogKeyAdaptiveState::from(&stored_key),
expected_adaptive_state,
"Anthropic status {status_code} must not update key-wide adaptive state"
);
assert_eq!(
stored_key.health_by_format, expected_health,
"Anthropic status {status_code} must not update key-wide health"
);
assert_eq!(
stored_key.circuit_breaker_by_format,
Some(circuit),
"Anthropic status {status_code} must not clear pool key state"
);
let expected_affinity = (status_code == 400).then_some(target);
assert_eq!(
state.read_scheduler_affinity_target(&cache_key, SCHEDULER_AFFINITY_TTL),
expected_affinity,
"Anthropic status {status_code} must invalidate only retryable target affinity"
);
}
}
#[test]
fn terminal_pool_account_errors_project_pool_hard_state() {
assert_eq!(
@@ -2790,6 +3052,162 @@ mod tests {
assert_eq!(stored_key.circuit_breaker_by_format, None);
}
#[tokio::test]
async fn oauth_bearer_generation_match_supports_generic_auth_config_token() {
let state = codex_state();
let mut transport = state
.read_provider_transport_snapshot(
"provider-codex-cli-local-1",
"endpoint-codex-cli-local-1",
"key-codex-cli-local-1",
)
.await
.expect("transport should load")
.expect("transport should exist");
transport.provider.provider_type = "claude_code".to_string();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config =
Some(json!({"accessToken": "current-access-token"}).to_string());
let mut plan = sample_codex_plan();
plan.headers.insert(
"authorization".to_string(),
"Bearer current-access-token".to_string(),
);
assert!(execution_plan_bearer_matches_transport(&plan, &transport));
plan.headers.insert(
"authorization".to_string(),
"Bearer stale-access-token".to_string(),
);
assert!(!execution_plan_bearer_matches_transport(&plan, &transport));
transport.key.decrypted_api_key = "replacement-access-token".to_string();
plan.headers.insert(
"authorization".to_string(),
"Bearer current-access-token".to_string(),
);
assert!(!execution_plan_bearer_matches_transport(&plan, &transport));
plan.headers.insert(
"authorization".to_string(),
"Bearer replacement-access-token".to_string(),
);
assert!(execution_plan_bearer_matches_transport(&plan, &transport));
transport.key.decrypted_auth_config = Some(
json!({
"accessToken": "current-access-token",
"request": {
"extraHeaders": {
"Authorization": "Bearer nested-override-token"
}
}
})
.to_string(),
);
plan.headers.insert(
"authorization".to_string(),
"Bearer replacement-access-token".to_string(),
);
assert!(!execution_plan_bearer_matches_transport(&plan, &transport));
plan.headers.insert(
"authorization".to_string(),
"Bearer nested-override-token".to_string(),
);
assert!(execution_plan_bearer_matches_transport(&plan, &transport));
}
#[tokio::test]
async fn oauth_invalidation_marks_claude_code_authentication_failures_only() {
let state = claude_code_oauth_state();
let mut plan = sample_codex_plan();
plan.provider_name = Some("claude_code".to_string());
plan.provider_api_format = "claude:messages".to_string();
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
status_code: 403,
response_text: Some(
r#"{"type":"error","error":{"type":"permission_error","message":"insufficient scope"}}"#,
),
}),
)
.await;
let unmarked = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
assert!(unmarked.oauth_invalid_at_unix_secs.is_none());
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
status_code: 403,
response_text: Some(
r#"{"type":"error","error":{"type":"authentication_error","message":"invalid access token"}}"#,
),
}),
)
.await;
let marked = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
assert!(marked.oauth_invalid_at_unix_secs.is_some());
assert_eq!(
marked.oauth_invalid_reason.as_deref(),
Some("[OAUTH_EXPIRED] invalid access token")
);
}
#[tokio::test]
async fn oauth_invalidation_marks_claude_code_unauthorized_without_body() {
let state = claude_code_oauth_state();
let mut plan = sample_codex_plan();
plan.provider_name = Some("claude_code".to_string());
plan.provider_api_format = "claude:messages".to_string();
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
status_code: 401,
response_text: None,
}),
)
.await;
let stored_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
assert!(stored_key.oauth_invalid_at_unix_secs.is_some());
assert_eq!(
stored_key.oauth_invalid_reason.as_deref(),
Some("[OAUTH_EXPIRED] OAuth access token was rejected")
);
}
#[tokio::test]
async fn oauth_invalidation_marks_codex_key_expired() {
let state = codex_state();
+13 -5
View File
@@ -9,6 +9,7 @@ mod attempt;
mod classifier;
mod effects;
mod health;
mod oauth_error;
mod policy;
mod recovery;
mod report_effects;
@@ -24,8 +25,10 @@ pub(crate) use self::attempt::{
LocalExecutionCandidateMetadata, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD,
};
pub(crate) use self::classifier::{
classify_local_failover, local_failover_error_message, LocalFailoverClassification,
LocalFailoverInput,
classify_anthropic_failure_disposition, classify_failure_disposition, classify_local_failover,
failure_disposition_from_local_classification, local_failover_error_message,
FailureDisposition, FailureRetryAction, FailureScope, FailureTokenAction,
LocalFailoverClassification, LocalFailoverInput,
};
pub(crate) use self::effects::{
apply_local_execution_effect, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect,
@@ -37,6 +40,9 @@ pub(crate) use self::health::{
project_local_failure_health, project_local_key_circuit_closed,
project_local_key_circuit_failure, project_local_success_health,
};
pub(crate) use self::oauth_error::{
oauth_status_may_be_invalid, oauth_status_proves_access_token_invalid,
};
pub(crate) use self::policy::{
append_local_failover_policy_to_value, codex_cyber_flag_passthrough_enabled,
cyber_continue_failover_enabled, local_failover_policy_from_report_context,
@@ -44,8 +50,8 @@ pub(crate) use self::policy::{
LocalFailoverRegexRule, CYBER_CONTINUE_FAILOVER_CONFIG_KEY,
};
pub(crate) use self::recovery::{
analyze_local_failover, recover_local_failover_decision, LocalFailoverAnalysis,
LocalFailoverDecision,
analyze_local_failover, apply_provider_failure_disposition, recover_local_failover_decision,
LocalFailoverAnalysis, LocalFailoverDecision,
};
#[cfg(test)]
pub(crate) use self::report_effects::clear_local_report_effect_caches_for_tests;
@@ -65,7 +71,9 @@ pub(crate) async fn resolve_local_failover_analysis_for_attempt(
}
let policy = resolve_local_failover_policy(state, plan, report_context).await;
analyze_local_failover(&policy, LocalFailoverInput::new(status_code, response_text))
let analysis =
analyze_local_failover(&policy, LocalFailoverInput::new(status_code, response_text));
apply_provider_failure_disposition(&plan.provider_api_format, status_code, analysis)
}
pub(crate) async fn resolve_local_failover_decision_for_attempt(
@@ -0,0 +1,113 @@
pub(crate) fn oauth_status_may_be_invalid(status_code: u16, response_text: Option<&str>) -> bool {
if status_code == 401 {
return true;
}
if status_code != 403 {
return false;
}
let Some(response_text) = response_text else {
return false;
};
if let Ok(body) = serde_json::from_str::<serde_json::Value>(response_text) {
let error_type = body
.get("error")
.and_then(|error| error.get("type"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| {
body.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty() && !value.eq_ignore_ascii_case("error"))
});
if let Some(error_type) = error_type {
return is_oauth_invalid_error_taxonomy(error_type);
}
let error_code = body
.get("error")
.and_then(|error| error.get("code"))
.or_else(|| body.get("code"))
.or_else(|| body.get("error").filter(|error| error.is_string()))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if error_code.is_some_and(is_oauth_invalid_error_taxonomy) {
return true;
}
return response_has_oauth_invalid_phrase(response_text);
}
response_has_oauth_invalid_phrase(response_text)
}
pub(crate) fn oauth_status_proves_access_token_invalid(
status_code: u16,
response_text: Option<&str>,
) -> bool {
if status_code == 401 {
return true;
}
if status_code != 403 {
return false;
}
response_text.is_some_and(response_has_oauth_invalid_phrase)
}
fn is_oauth_invalid_error_taxonomy(value: &str) -> bool {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"authentication_error"
| "invalid_authentication_token"
| "invalid_token"
| "oauth_token_invalid"
| "token_invalid"
| "token_expired"
| "unauthenticated"
| "biscuit_baker_service_auth_credential_error_status"
)
}
fn response_has_oauth_invalid_phrase(response_text: &str) -> bool {
let response_text = response_text.to_ascii_lowercase();
if [
"oauth_token_invalid",
"invalid_token",
"biscuit_baker_service_auth_credential_error_status",
]
.iter()
.any(|taxonomy| contains_ascii_taxonomy_token(&response_text, taxonomy))
{
return true;
}
[
"oauth token is invalid",
"oauth token is expired",
"oauth token has expired",
"invalid access token",
"access token invalid",
"access token expired",
"expired access token",
"authentication token has been invalidated",
"token has been invalidated",
"personal access token owner is inactive",
"security token included in the request is expired",
]
.iter()
.any(|needle| response_text.contains(needle))
}
fn contains_ascii_taxonomy_token(text: &str, taxonomy: &str) -> bool {
text.match_indices(taxonomy).any(|(start, matched)| {
let end = start + matched.len();
let is_identifier_byte = |byte: u8| byte.is_ascii_alphanumeric() || byte == b'_';
let has_left_boundary = start == 0 || !is_identifier_byte(text.as_bytes()[start - 1]);
let has_right_boundary = end == text.len() || !is_identifier_byte(text.as_bytes()[end]);
has_left_boundary && has_right_boundary
})
}
@@ -1,4 +1,7 @@
use super::classifier::{classify_local_failover, LocalFailoverClassification, LocalFailoverInput};
use super::classifier::{
classify_failure_disposition, classify_local_failover, FailureRetryAction,
LocalFailoverClassification, LocalFailoverInput,
};
use super::LocalFailoverPolicy;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -44,6 +47,37 @@ pub(crate) fn analyze_local_failover(
}
}
pub(crate) fn apply_provider_failure_disposition(
provider_api_format: &str,
status_code: u16,
analysis: LocalFailoverAnalysis,
) -> LocalFailoverAnalysis {
if status_code < 400
&& matches!(
analysis.classification,
LocalFailoverClassification::UseDefault
)
{
return analysis;
}
let disposition =
classify_failure_disposition(provider_api_format, analysis.classification, status_code);
let decision = match disposition.retry_action {
FailureRetryAction::Stop | FailureRetryAction::SameCredential => {
LocalFailoverDecision::StopLocalFailover
}
FailureRetryAction::NextCandidate
| FailureRetryAction::NextCredential
| FailureRetryAction::NextEndpoint => LocalFailoverDecision::RetryNextCandidate,
};
LocalFailoverAnalysis {
classification: analysis.classification,
decision,
}
}
pub(crate) fn recover_local_failover_decision(
policy: &LocalFailoverPolicy,
input: LocalFailoverInput<'_>,
@@ -70,7 +104,10 @@ const fn decision_from_classification(
#[cfg(test)]
mod tests {
use super::{analyze_local_failover, recover_local_failover_decision, LocalFailoverDecision};
use super::{
analyze_local_failover, apply_provider_failure_disposition,
recover_local_failover_decision, LocalFailoverAnalysis, LocalFailoverDecision,
};
use crate::orchestration::{
LocalFailoverClassification, LocalFailoverInput, LocalFailoverPolicy,
};
@@ -161,4 +198,44 @@ mod tests {
LocalFailoverClassification::StopCyberPolicy
);
}
#[test]
fn anthropic_failure_disposition_controls_candidate_retry() {
let policy = LocalFailoverPolicy::default();
for status_code in [400, 413] {
let analysis = analyze_local_failover(
&policy,
LocalFailoverInput::new(status_code, Some(r#"{"error":{"message":"failed"}}"#)),
);
assert_eq!(
apply_provider_failure_disposition("claude:messages", status_code, analysis,)
.decision,
LocalFailoverDecision::StopLocalFailover,
"Anthropic status {status_code} must not blindly rotate credentials"
);
}
for status_code in [401, 403, 404, 429, 529] {
let analysis = analyze_local_failover(
&policy,
LocalFailoverInput::new(status_code, Some(r#"{"error":{"message":"failed"}}"#)),
);
assert_eq!(
apply_provider_failure_disposition("claude:messages", status_code, analysis,)
.decision,
LocalFailoverDecision::RetryNextCandidate,
"Anthropic status {status_code} should continue candidate failover"
);
}
}
#[test]
fn provider_failure_disposition_preserves_non_failure_default() {
let analysis = LocalFailoverAnalysis::use_default();
assert_eq!(
apply_provider_failure_disposition("claude:messages", 200, analysis).decision,
LocalFailoverDecision::UseDefault
);
}
}