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
@@ -0,0 +1,249 @@
use std::sync::{Arc, OnceLock};
use std::time::{SystemTime, UNIX_EPOCH};
use aether_provider_transport::CodexFingerprintConvergenceContext;
use http::{request::Parts, HeaderMap};
use serde_json::Value;
use uuid::Uuid;
use crate::client_session_affinity::codex_request_signals_from_request;
#[derive(Debug, Clone)]
pub(crate) struct CodexFingerprintContextSlot(Arc<OnceLock<CodexFingerprintConvergenceContext>>);
impl Default for CodexFingerprintContextSlot {
fn default() -> Self {
Self(Arc::new(OnceLock::new()))
}
}
impl CodexFingerprintContextSlot {
fn resolve(
&self,
headers: &HeaderMap,
body_json: &Value,
) -> CodexFingerprintConvergenceContext {
self.0
.get_or_init(|| {
build_codex_fingerprint_context(headers, body_json, Uuid::now_v7().to_string())
})
.clone()
}
}
pub(crate) fn resolve_codex_fingerprint_context(
parts: &Parts,
body_json: &Value,
) -> CodexFingerprintConvergenceContext {
if let Some(context) = parts
.extensions
.get::<CodexFingerprintConvergenceContext>()
.cloned()
{
return context;
}
if let Some(slot) = parts.extensions.get::<CodexFingerprintContextSlot>() {
return slot.resolve(&parts.headers, body_json);
}
build_codex_fingerprint_context(&parts.headers, body_json, Uuid::now_v7().to_string())
}
pub(crate) fn install_codex_fingerprint_context_slot(parts: &mut Parts) {
if parts
.extensions
.get::<CodexFingerprintConvergenceContext>()
.is_none()
&& parts
.extensions
.get::<CodexFingerprintContextSlot>()
.is_none()
{
parts
.extensions
.insert(CodexFingerprintContextSlot::default());
}
}
pub(crate) fn ensure_codex_fingerprint_context(
parts: &mut Parts,
body_json: &Value,
) -> CodexFingerprintConvergenceContext {
let context = resolve_codex_fingerprint_context(parts, body_json);
if parts
.extensions
.get::<CodexFingerprintConvergenceContext>()
.is_none()
{
parts.extensions.remove::<CodexFingerprintContextSlot>();
parts.extensions.insert(context.clone());
}
context
}
pub(crate) fn attach_codex_logical_turn_context(
parts: &mut Parts,
body_json: &Value,
logical_turn_id: &str,
) -> CodexFingerprintConvergenceContext {
let context =
build_codex_fingerprint_context(&parts.headers, body_json, logical_turn_id.to_string());
parts.extensions.remove::<CodexFingerprintContextSlot>();
parts.extensions.insert(context.clone());
context
}
pub(crate) fn restore_codex_logical_turn_context(
parts: &mut Parts,
context: &CodexFingerprintConvergenceContext,
) {
parts.extensions.remove::<CodexFingerprintContextSlot>();
parts.extensions.insert(context.clone());
}
fn build_codex_fingerprint_context(
headers: &HeaderMap,
body_json: &Value,
logical_turn_id: String,
) -> CodexFingerprintConvergenceContext {
let signals = codex_request_signals_from_request(headers, Some(body_json));
let mut context =
CodexFingerprintConvergenceContext::new(logical_turn_id, current_unix_millis());
if let Some(turn_id) = signals.turn_id {
context = context.with_original_turn_id(turn_id);
}
if let Some(session_id) = signals.thread_id.or(signals.session_id) {
context = context.with_original_client_session_id(session_id);
}
if let Some(prompt_cache_key) = signals.prompt_cache_key {
context = context.with_original_prompt_cache_key(prompt_cache_key);
}
context
}
fn current_unix_millis() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis()
.try_into()
.unwrap_or(u64::MAX)
}
#[cfg(test)]
mod tests {
use http::HeaderValue;
use serde_json::json;
use super::*;
#[test]
fn request_signals_are_captured_once_for_the_logical_turn() {
let request = http::Request::builder()
.header("thread-id", "header-thread")
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
let body = json!({
"prompt_cache_key": "client-cache",
"client_metadata": {
"turn_id": "client-turn",
"thread_id": "body-thread"
}
});
let context = attach_codex_logical_turn_context(&mut parts, &body, "logical-turn");
assert_eq!(context.logical_turn_id(), "logical-turn");
assert_eq!(context.original_turn_id(), Some("client-turn"));
assert_eq!(context.original_client_session_id(), Some("header-thread"));
assert_eq!(context.original_prompt_cache_key(), Some("client-cache"));
assert_eq!(
parts.extensions.get::<CodexFingerprintConvergenceContext>(),
Some(&context)
);
}
#[test]
fn restored_context_wins_over_retry_request_signals() {
let original = CodexFingerprintConvergenceContext::new("logical-turn", 1234)
.with_original_turn_id("original-turn")
.with_original_client_session_id("original-thread")
.with_original_prompt_cache_key("original-cache");
let request = http::Request::builder()
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
parts
.headers
.insert("thread-id", HeaderValue::from_static("retry-thread"));
restore_codex_logical_turn_context(&mut parts, &original);
let resolved = resolve_codex_fingerprint_context(
&parts,
&json!({
"prompt_cache_key": "retry-cache",
"client_metadata": {"turn_id": "retry-turn"}
}),
);
assert_eq!(resolved, original);
assert_eq!(resolved.turn_started_at_unix_ms(), 1234);
}
#[test]
fn generated_context_is_persisted_for_http_replanning() {
let request = http::Request::builder()
.header("session-id", "client-session")
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
let body = json!({
"prompt_cache_key": "client-cache",
"client_metadata": {"turn_id": "client-turn"}
});
let first = ensure_codex_fingerprint_context(&mut parts, &body);
let second = resolve_codex_fingerprint_context(
&parts,
&json!({
"prompt_cache_key": "retry-cache",
"client_metadata": {"turn_id": "retry-turn"}
}),
);
assert_eq!(second, first);
assert_eq!(second.original_turn_id(), Some("client-turn"));
assert_eq!(second.original_prompt_cache_key(), Some("client-cache"));
}
#[test]
fn installed_slot_reuses_context_across_cloned_parts() {
let request = http::Request::builder()
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
install_codex_fingerprint_context_slot(&mut parts);
let cloned_parts = parts.clone();
let first = resolve_codex_fingerprint_context(
&parts,
&json!({
"prompt_cache_key": "first-cache",
"client_metadata": {"turn_id": "first-turn"}
}),
);
let second = resolve_codex_fingerprint_context(
&cloned_parts,
&json!({
"prompt_cache_key": "second-cache",
"client_metadata": {"turn_id": "second-turn"}
}),
);
assert_eq!(second, first);
assert_eq!(second.original_turn_id(), Some("first-turn"));
assert_eq!(second.original_prompt_cache_key(), Some("first-cache"));
}
}
@@ -1,5 +1,6 @@
mod adaptation;
pub(crate) mod api;
pub(crate) mod codex_context;
mod finalize;
mod planner;
mod pure;
@@ -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,
@@ -60,6 +60,7 @@ pub(crate) mod windsurf {
pub(crate) use aether_provider_transport::{
append_transport_diagnostics_to_value, apply_codex_oauth_fingerprint_convergence,
apply_codex_oauth_fingerprint_convergence_with_context,
apply_local_auth_config_header_overrides, apply_local_body_rules,
apply_local_body_rules_with_request_headers, apply_local_header_rules,
apply_local_header_rules_with_request_headers, apply_standard_provider_request_body_rules,