use super::{HttpRequest, HttpTransport, ProviderEvent, transport::HttpResponseMetadata};
use crate::cancellation::AgentCancellation;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use std::{
cell::RefCell,
io::Write,
sync::{Arc, Mutex, atomic::AtomicU64},
time::Instant,
};
pub(crate) const TURN_STATE_HEADER: &str = "x-codex-turn-state";
const MAX_TURN_STATE_BYTES: usize = 4096;
pub(crate) fn valid_turn_state(value: &str) -> bool {
!value.trim().is_empty()
&& value.len() <= MAX_TURN_STATE_BYTES
&& value.bytes().all(|byte| (0x20..=0x7e).contains(&byte))
}
#[derive(Clone)]
pub(crate) struct CodexTurnContext(Arc<Mutex<TurnState>>);
impl Default for CodexTurnContext {
fn default() -> Self {
Self(Arc::new(Mutex::new(TurnState {
id: uuid::Uuid::new_v4().to_string(),
salt: *uuid::Uuid::new_v4().as_bytes(),
owner: None,
routing: None,
previous: None,
attempts: 0,
})))
}
}
impl std::fmt::Debug for CodexTurnContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("CodexTurnContext(<private>)")
}
}
impl PartialEq for CodexTurnContext {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for CodexTurnContext {}
struct TurnState {
id: String,
salt: [u8; 16],
owner: Option<[u8; 32]>,
routing: Option<String>,
previous: Option<Fingerprint>,
attempts: u64,
}
struct Fingerprint {
input_count: usize,
input: [u8; 32],
configuration: [u8; 32],
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub(crate) struct CacheDiagnostics {
pub(crate) turn_id: String,
pub(crate) attempt_ordinal: u64,
pub(crate) started_at: String,
pub(crate) elapsed_ms: u64,
pub(crate) input_items: usize,
pub(crate) previous_input_items: Option<usize>,
pub(crate) previous_prefix_unchanged: Option<bool>,
pub(crate) input_fingerprint: String,
pub(crate) configuration_fingerprint: String,
pub(crate) instructions_fingerprint: String,
pub(crate) tools_fingerprint: String,
pub(crate) affinity_fingerprint: String,
pub(crate) session_affinity_sent: bool,
pub(crate) turn_state_sent: bool,
pub(crate) turn_state_received: bool,
}
struct HashWriter<'a>(&'a mut Sha256);
impl Write for HashWriter<'_> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.update(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
fn add_json(hash: &mut Sha256, value: &Value) {
serde_json::to_writer(HashWriter(hash), value).expect("JSON value digest");
hash.update([0]);
}
fn digest(salt: &[u8; 16], value: &Value) -> String {
let mut hash = Sha256::new();
hash.update(salt);
add_json(&mut hash, value);
crate::hex::lower_hex(hash.finalize())
}
impl CodexTurnContext {
fn prepare(&self, request: &mut HttpRequest) -> CacheDiagnostics {
let mut state = self.0.lock().unwrap_or_else(|error| error.into_inner());
let mut owner = Sha256::new();
for value in [
request.headers.get("chatgpt-account-id"),
request.headers.get("session-id"),
] {
owner.update(value.map_or("", String::as_str));
owner.update([0]);
}
add_json(&mut owner, &request.body["model"]);
let owner: [u8; 32] = owner.finalize().into();
if state.owner != Some(owner) {
state.routing = None;
state.previous = None;
state.owner = Some(owner);
}
request.headers.remove(TURN_STATE_HEADER);
if let Some(value) = &state.routing {
request
.headers
.insert(TURN_STATE_HEADER.to_string(), value.clone());
}
let input = request.body["input"]
.as_array()
.map_or(&[][..], Vec::as_slice);
let mut configuration = Sha256::new();
configuration.update(state.salt);
if let Some(body) = request.body.as_object() {
for (key, value) in body {
if key != "input" {
add_json(&mut configuration, &Value::String(key.clone()));
add_json(&mut configuration, value);
}
}
}
let configuration: [u8; 32] = configuration.finalize().into();
let mut hash = Sha256::new();
hash.update(state.salt);
let mut prefix = None;
for (index, item) in input.iter().enumerate() {
if state
.previous
.as_ref()
.is_some_and(|last| last.input_count == index)
{
prefix = Some(hash.clone().finalize().into());
}
add_json(&mut hash, item);
}
let input_hash: [u8; 32] = hash.finalize().into();
if state
.previous
.as_ref()
.is_some_and(|last| last.input_count == input.len())
{
prefix = Some(input_hash);
}
state.attempts = state.attempts.saturating_add(1);
let diagnostic = CacheDiagnostics {
turn_id: state.id.clone(),
attempt_ordinal: state.attempts,
started_at: chrono::Utc::now().to_rfc3339(),
elapsed_ms: 0,
input_items: input.len(),
previous_input_items: state.previous.as_ref().map(|last| last.input_count),
previous_prefix_unchanged: state
.previous
.as_ref()
.map(|last| prefix == Some(last.input) && configuration == last.configuration),
input_fingerprint: crate::hex::lower_hex(input_hash),
configuration_fingerprint: crate::hex::lower_hex(configuration),
instructions_fingerprint: digest(&state.salt, &request.body["instructions"]),
tools_fingerprint: digest(&state.salt, &request.body["tools"]),
affinity_fingerprint: digest(&state.salt, &request.body["prompt_cache_key"]),
session_affinity_sent: request.headers.contains_key("session-id"),
turn_state_sent: state.routing.is_some(),
turn_state_received: false,
};
state.previous = Some(Fingerprint {
input_count: input.len(),
input: input_hash,
configuration,
});
diagnostic
}
fn observe(&self, metadata: &HttpResponseMetadata) -> bool {
let Some(value) = metadata
.codex_turn_state
.as_deref()
.filter(|value| valid_turn_state(value))
else {
return false;
};
let mut state = self.0.lock().unwrap_or_else(|error| error.into_inner());
if state.routing.is_none() {
state.routing = Some(value.to_string());
}
true
}
}
pub(super) struct CodexTurnTransport<'a, T> {
pub(super) inner: &'a T,
pub(super) context: CodexTurnContext,
pub(super) diagnostic: RefCell<Option<CacheDiagnostics>>,
}
impl<T> CodexTurnTransport<'_, T> {
pub(super) fn annotate(&self, event: &mut ProviderEvent) {
if let ProviderEvent::ResponseIdentity(identity) = event {
identity.cache = self.diagnostic.borrow().clone().map(Box::new);
}
}
}
impl<T: HttpTransport> HttpTransport for CodexTurnTransport<'_, T> {
fn stream_json_cancellable_with_semantic_deadline(
&self,
request: HttpRequest,
cancellation: &AgentCancellation,
deadline: &AtomicU64,
on_chunk: &mut dyn FnMut(&str) -> anyhow::Result<()>,
) -> anyhow::Result<()> {
self.stream_json_cancellable_with_response_metadata(
request,
cancellation,
deadline,
&mut |_| {},
on_chunk,
)
}
fn stream_json_cancellable_with_response_metadata(
&self,
mut request: HttpRequest,
cancellation: &AgentCancellation,
deadline: &AtomicU64,
on_metadata: &mut dyn FnMut(HttpResponseMetadata),
on_chunk: &mut dyn FnMut(&str) -> anyhow::Result<()>,
) -> anyhow::Result<()> {
cancellation.check()?;
let mut diagnostic = self.context.prepare(&mut request);
let start = Instant::now();
let result = self.inner.stream_json_cancellable_with_response_metadata(
request,
cancellation,
deadline,
&mut |metadata| {
diagnostic.turn_state_received |= self.context.observe(&metadata);
on_metadata(metadata);
},
on_chunk,
);
diagnostic.elapsed_ms = start.elapsed().as_millis().min(u128::from(u64::MAX)) as u64;
*self.diagnostic.borrow_mut() = Some(diagnostic);
result
}
}