feat(codex): stabilize identity across retries

This commit is contained in:
elky
2026-09-01 15:33:40 +08:00
parent b538aa2d66
commit a39048ecce
19 changed files with 1731 additions and 156 deletions
@@ -2,6 +2,7 @@ use std::collections::BTreeMap;
use std::time::Duration;
use aether_ai_serving::{run_ai_authenticated_decision_input, AiAuthenticatedDecisionInputPort};
use aether_provider_transport::CodexFingerprintConvergenceContext;
use aether_routing_core::{
rank_vector_for_candidate, CandidateKind, ResolvedRoutingPolicy, RoutingCandidateFacts,
RoutingCandidateTrace, RoutingDecisionTrace, RoutingPoolExpansionTrace, RoutingRulePhase,
@@ -55,7 +56,7 @@ pub(crate) struct LocalRequestedModelDecisionInput {
pub(crate) client_surface: Option<ClientSurface>,
pub(crate) gateway_credential_carrier: Option<GatewayCredentialCarrier>,
pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
pub(crate) original_client_session_id: Option<String>,
pub(crate) codex_fingerprint_context: Option<CodexFingerprintConvergenceContext>,
pub(crate) routing_policy: Option<ResolvedRoutingPolicy>,
pub(crate) routing_trace_seed: Option<RoutingDecisionTrace>,
pub(crate) routing_context: Option<LocalRoutingRequestContext>,
@@ -376,13 +377,25 @@ fn apply_codex_oauth_fingerprint_convergence_to_decision(
else {
return;
};
crate::ai_serving::transport::apply_codex_oauth_fingerprint_convergence(
transport,
provider_api_format,
input.original_client_session_id.as_deref(),
&mut decision.provider_request_headers,
provider_request_body,
);
let Some(context) = input.codex_fingerprint_context.as_ref() else {
return;
};
let applied =
crate::ai_serving::transport::apply_codex_oauth_fingerprint_convergence_with_context(
transport,
provider_api_format,
context,
&mut decision.provider_request_headers,
provider_request_body,
);
if applied {
decision.prompt_cache_key = provider_request_body
.get("prompt_cache_key")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
}
}
struct GatewayAuthenticatedDecisionInputPort<'a> {
@@ -471,7 +484,7 @@ pub(crate) fn build_local_requested_model_decision_input(
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -486,7 +499,8 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
body_json: &Value,
client_api_format: &str,
) -> Result<(), GatewayError> {
input.original_client_session_id = original_client_session_id_from_headers(&parts.headers);
input.codex_fingerprint_context =
Some(crate::ai_serving::codex_context::resolve_codex_fingerprint_context(parts, body_json));
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
let selected_group = match state.routing_group_read_repository() {
Some(repository) => {
@@ -737,12 +751,6 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
Ok(())
}
fn original_client_session_id_from_headers(headers: &HeaderMap) -> Option<String> {
routing_header_value_str(headers, "session-id")
.or_else(|| routing_header_value_str(headers, "session_id"))
.or_else(|| routing_header_value_str(headers, "x-session-id"))
}
fn try_attach_static_default_routing_policy_to_input(
input: &mut LocalRequestedModelDecisionInput,
parts: &http::request::Parts,
@@ -1106,38 +1114,6 @@ mod tests {
GatewayProviderTransportProvider,
};
#[test]
fn original_client_session_id_accepts_live_header_as_fallback() {
let headers = HeaderMap::from_iter([(
HeaderName::from_static("x-session-id"),
HeaderValue::from_static("live-thread-1"),
)]);
assert_eq!(
original_client_session_id_from_headers(&headers).as_deref(),
Some("live-thread-1")
);
}
#[test]
fn original_client_session_id_prefers_responses_headers_over_live_fallback() {
let headers = HeaderMap::from_iter([
(
HeaderName::from_static("session-id"),
HeaderValue::from_static("responses-session"),
),
(
HeaderName::from_static("x-session-id"),
HeaderValue::from_static("live-thread"),
),
]);
assert_eq!(
original_client_session_id_from_headers(&headers).as_deref(),
Some("responses-session")
);
}
#[test]
fn explicit_routing_selection_cache_key_is_principal_specific() {
let first = routing_group_selection_cache_key(
@@ -1350,7 +1326,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
model_directive_policy: Default::default(),
@@ -1599,7 +1575,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
model_directive_policy: Default::default(),
@@ -1669,7 +1645,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -1754,7 +1730,13 @@ mod tests {
});
let mut with_mutation = sample_decision_input();
for input in [&mut no_context, &mut empty_mutation, &mut with_mutation] {
input.original_client_session_id = Some("client-session-1".to_string());
input.codex_fingerprint_context = Some(
CodexFingerprintConvergenceContext::new(
uuid::Uuid::new_v4().to_string(),
1_756_668_000_000,
)
.with_original_client_session_id("client-session-1".to_string()),
);
}
let mut stable_identity = None;
@@ -1801,6 +1783,10 @@ mod tests {
.provider_request_body
.as_ref()
.expect("request body");
assert_eq!(
decision.prompt_cache_key.as_deref(),
body.get("prompt_cache_key").and_then(Value::as_str)
);
assert_eq!(
body["prompt_cache_key"],
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
@@ -2005,6 +1991,7 @@ mod tests {
let body = decision.provider_request_body.as_ref().expect("body");
assert!(body.get("prompt_cache_key").is_none());
assert!(decision.prompt_cache_key.is_none());
assert!(body.get("client_metadata").is_none());
assert!(!decision.provider_request_headers.contains_key("session-id"));
assert!(!decision.provider_request_headers.contains_key("thread-id"));
@@ -377,7 +377,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -2182,7 +2182,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,