use std::sync::Arc;
use crate::Codex;
use crate::command::exec::{ExecCommand, ExecResumeCommand};
use crate::error::{Error, Result};
use crate::types::{JsonLineEvent, QueryResult, TokenUsage};
#[derive(Debug, Clone)]
pub struct TurnRecord {
pub result: QueryResult,
}
impl TurnRecord {
#[must_use]
pub fn events(&self) -> &[JsonLineEvent] {
&self.result.events
}
#[must_use]
pub fn usage(&self) -> Option<TokenUsage> {
self.result.usage
}
}
pub struct Session {
codex: Arc<Codex>,
thread_id: Option<String>,
history: Vec<TurnRecord>,
budget: Option<crate::budget::TokenBudget>,
}
impl Session {
pub fn new(codex: Arc<Codex>) -> Self {
Self {
codex,
thread_id: None,
history: Vec::new(),
budget: None,
}
}
pub fn resume(codex: Arc<Codex>, thread_id: impl Into<String>) -> Self {
Self {
codex,
thread_id: Some(thread_id.into()),
history: Vec::new(),
budget: None,
}
}
#[must_use]
pub fn with_budget(mut self, budget: crate::budget::TokenBudget) -> Self {
self.budget = Some(budget);
self
}
#[must_use]
pub fn budget(&self) -> Option<&crate::budget::TokenBudget> {
self.budget.as_ref()
}
pub async fn send(&mut self, prompt: impl Into<String>) -> Result<Vec<JsonLineEvent>> {
let prompt = prompt.into();
match &self.thread_id {
None => {
let cmd = ExecCommand::new(&prompt);
self.run_exec(cmd).await
}
Some(id) => {
let cmd = ExecResumeCommand::new()
.session_id(id.clone())
.prompt(prompt);
self.run_resume(cmd).await
}
}
}
pub async fn execute(&mut self, cmd: ExecCommand) -> Result<Vec<JsonLineEvent>> {
self.run_exec(cmd).await
}
pub async fn execute_resume(&mut self, cmd: ExecResumeCommand) -> Result<Vec<JsonLineEvent>> {
self.run_resume(cmd).await
}
pub async fn stream<F>(
&mut self,
prompt: impl Into<String>,
handler: F,
) -> Result<Vec<JsonLineEvent>>
where
F: FnMut(JsonLineEvent),
{
let prompt = prompt.into();
match &self.thread_id {
None => {
self.stream_execute(ExecCommand::new(&prompt), handler)
.await
}
Some(id) => {
let cmd = ExecResumeCommand::new()
.session_id(id.clone())
.prompt(prompt);
self.stream_execute_resume(cmd, handler).await
}
}
}
pub async fn stream_execute<F>(
&mut self,
cmd: ExecCommand,
mut handler: F,
) -> Result<Vec<JsonLineEvent>>
where
F: FnMut(JsonLineEvent),
{
self.check_budget()?;
let codex = Arc::clone(&self.codex);
let mut collected = Vec::new();
let outcome = cmd
.stream(&codex, |event| {
collected.push(event.clone());
handler(event);
})
.await;
self.finish_stream(collected, outcome)
}
pub async fn stream_execute_resume<F>(
&mut self,
cmd: ExecResumeCommand,
mut handler: F,
) -> Result<Vec<JsonLineEvent>>
where
F: FnMut(JsonLineEvent),
{
self.check_budget()?;
let codex = Arc::clone(&self.codex);
let mut collected = Vec::new();
let outcome = cmd
.stream(&codex, |event| {
collected.push(event.clone());
handler(event);
})
.await;
self.finish_stream(collected, outcome)
}
fn finish_stream(
&mut self,
collected: Vec<JsonLineEvent>,
outcome: Result<()>,
) -> Result<Vec<JsonLineEvent>> {
match outcome {
Ok(()) => Ok(self.record_turn(collected)),
Err(e) => {
self.capture_thread_id(&collected);
Err(e)
}
}
}
#[must_use]
pub fn id(&self) -> Option<&str> {
self.thread_id.as_deref()
}
#[must_use]
pub fn total_turns(&self) -> usize {
self.history.len()
}
#[must_use]
pub fn history(&self) -> &[TurnRecord] {
&self.history
}
#[must_use]
pub fn last_result(&self) -> Option<&QueryResult> {
self.history.last().map(|turn| &turn.result)
}
#[must_use]
pub fn total_tokens(&self) -> u64 {
self.history
.iter()
.filter_map(|turn| turn.usage().and_then(|u| u.total()))
.sum()
}
#[must_use]
pub fn turns_missing_usage(&self) -> usize {
self.history
.iter()
.filter(|turn| turn.usage().and_then(|u| u.total()).is_none())
.count()
}
fn record_turn(&mut self, events: Vec<JsonLineEvent>) -> Vec<JsonLineEvent> {
self.capture_thread_id(&events);
let result = QueryResult::from_events(events);
let events = result.events.clone();
if let Some(budget) = &self.budget {
budget.record(result.usage.and_then(|usage| usage.total()));
}
self.history.push(TurnRecord { result });
events
}
fn check_budget(&self) -> Result<()> {
match &self.budget {
Some(budget) => budget.check(),
None => Ok(()),
}
}
async fn run_exec(&mut self, cmd: ExecCommand) -> Result<Vec<JsonLineEvent>> {
self.check_budget()?;
match cmd.execute_json_lines(&self.codex).await {
Ok(events) => Ok(self.record_turn(events)),
Err(Error::CommandFailed {
stdout,
stderr,
exit_code,
command,
working_dir,
}) => {
self.try_capture_thread_id_from_stdout(&stdout);
Err(Error::CommandFailed {
stdout,
stderr,
exit_code,
command,
working_dir,
})
}
Err(e) => Err(e),
}
}
async fn run_resume(&mut self, cmd: ExecResumeCommand) -> Result<Vec<JsonLineEvent>> {
self.check_budget()?;
match cmd.execute_json_lines(&self.codex).await {
Ok(events) => Ok(self.record_turn(events)),
Err(Error::CommandFailed {
stdout,
stderr,
exit_code,
command,
working_dir,
}) => {
self.try_capture_thread_id_from_stdout(&stdout);
Err(Error::CommandFailed {
stdout,
stderr,
exit_code,
command,
working_dir,
})
}
Err(e) => Err(e),
}
}
fn capture_thread_id(&mut self, events: &[JsonLineEvent]) {
if let Some(id) = events.iter().find_map(|e| e.thread_id()) {
self.thread_id = Some(id.to_string());
}
}
fn try_capture_thread_id_from_stdout(&mut self, stdout: &str) {
for line in stdout.lines() {
if let Ok(event) = serde_json::from_str::<JsonLineEvent>(line)
&& let Some(id) = event.thread_id()
{
self.thread_id = Some(id.to_string());
return;
}
}
}
}
impl std::fmt::Debug for Session {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Session")
.field("thread_id", &self.thread_id)
.field("total_turns", &self.history.len())
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn test_codex() -> Arc<Codex> {
Arc::new(Codex::builder().binary("/usr/bin/false").build().unwrap())
}
#[test]
fn new_session_has_no_state() {
let session = Session::new(test_codex());
assert!(session.id().is_none());
assert_eq!(session.total_turns(), 0);
assert!(session.history().is_empty());
}
#[test]
fn resume_session_has_thread_id() {
let session = Session::resume(test_codex(), "thread_abc");
assert_eq!(session.id(), Some("thread_abc"));
assert_eq!(session.total_turns(), 0);
}
#[test]
fn capture_thread_id_from_events() {
let mut session = Session::new(test_codex());
let events: Vec<JsonLineEvent> = vec![
serde_json::from_str(r#"{"type":"message.created","role":"assistant"}"#).unwrap(),
serde_json::from_str(
r#"{"type":"thread.started","thread_id":"thread_xyz","session_id":"sess_1"}"#,
)
.unwrap(),
];
session.capture_thread_id(&events);
assert_eq!(session.id(), Some("thread_xyz"));
}
#[test]
fn capture_thread_id_noop_when_absent() {
let mut session = Session::new(test_codex());
let events: Vec<JsonLineEvent> =
vec![serde_json::from_str(r#"{"type":"message.created"}"#).unwrap()];
session.capture_thread_id(&events);
assert!(session.id().is_none());
}
#[test]
fn try_capture_thread_id_from_stdout_parses_json() {
let mut session = Session::new(test_codex());
let stdout = r#"{"type":"thread.started","thread_id":"thread_err"}
{"type":"error","message":"something went wrong"}"#;
session.try_capture_thread_id_from_stdout(stdout);
assert_eq!(session.id(), Some("thread_err"));
}
#[test]
fn try_capture_thread_id_from_stdout_ignores_garbage() {
let mut session = Session::new(test_codex());
session.try_capture_thread_id_from_stdout("not json\nalso not json");
assert!(session.id().is_none());
}
#[test]
fn debug_impl() {
let session = Session::resume(test_codex(), "thread_dbg");
let debug = format!("{session:?}");
assert!(debug.contains("thread_dbg"));
assert!(debug.contains("total_turns: 0"));
}
#[cfg(unix)]
fn streaming_session() -> Session {
let script = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fake-codex.sh");
Session::new(Arc::new(
Codex::builder()
.binary("/bin/bash")
.arg(script.to_str().unwrap())
.build()
.expect("bash must exist"),
))
}
fn turn(json: &[&str]) -> Vec<JsonLineEvent> {
json.iter()
.map(|line| serde_json::from_str(line).unwrap())
.collect()
}
fn completed(total: u64) -> Vec<JsonLineEvent> {
turn(&[&format!(
r#"{{"type":"turn.completed","usage":{{"total_tokens":{total}}}}}"#
)])
}
#[test]
fn record_turn_captures_usage_and_thread_id() {
let mut session = Session::new(test_codex());
session.record_turn(turn(&[
r#"{"type":"thread.started","thread_id":"thread_1"}"#,
r#"{"type":"item.completed","item":{"item_type":"agent_message","text":"hi"}}"#,
r#"{"type":"turn.completed","usage":{"total_tokens":42}}"#,
]));
assert_eq!(session.id(), Some("thread_1"));
assert_eq!(session.total_turns(), 1);
assert_eq!(session.last_result().unwrap().result, "hi");
assert_eq!(session.total_tokens(), 42);
assert_eq!(session.turns_missing_usage(), 0);
}
#[test]
fn total_tokens_sums_across_turns() {
let mut session = Session::new(test_codex());
for total in [10, 20, 30] {
session.record_turn(completed(total));
}
assert_eq!(session.total_turns(), 3);
assert_eq!(session.total_tokens(), 60);
assert_eq!(session.turns_missing_usage(), 0);
}
#[test]
fn unreported_usage_is_counted_not_hidden() {
let mut session = Session::new(test_codex());
session.record_turn(completed(40));
session.record_turn(turn(&[r#"{"type":"turn.completed"}"#]));
assert_eq!(session.total_tokens(), 40);
assert_eq!(session.turns_missing_usage(), 1);
assert_eq!(session.total_turns(), 2);
}
#[test]
fn zero_usage_and_unreported_usage_are_distinguishable() {
let mut reported = Session::new(test_codex());
reported.record_turn(completed(0));
let mut unreported = Session::new(test_codex());
unreported.record_turn(turn(&[r#"{"type":"turn.completed"}"#]));
assert_eq!(reported.total_tokens(), unreported.total_tokens());
assert_eq!(reported.turns_missing_usage(), 0);
assert_eq!(unreported.turns_missing_usage(), 1);
}
#[test]
fn turn_record_exposes_events_and_usage() {
let mut session = Session::new(test_codex());
session.record_turn(turn(&[
r#"{"type":"turn.started"}"#,
r#"{"type":"turn.completed","usage":{"total_tokens":5}}"#,
]));
let record = &session.history()[0];
assert_eq!(record.events().len(), 2);
assert_eq!(record.usage().unwrap().total(), Some(5));
}
#[test]
fn last_result_is_none_before_any_turn() {
let session = Session::new(test_codex());
assert!(session.last_result().is_none());
assert_eq!(session.total_tokens(), 0);
assert_eq!(session.turns_missing_usage(), 0);
}
#[tokio::test]
#[cfg(unix)]
async fn stream_delivers_events_and_records_the_turn() {
let mut session = streaming_session();
let mut seen = Vec::new();
let events = session
.stream("test prompt", |event| seen.push(event.event_type.clone()))
.await
.unwrap();
assert!(
seen.contains(&"turn.completed".to_string()),
"saw: {seen:?}"
);
assert_eq!(events.len(), seen.len(), "handler and return value agree");
assert_eq!(session.total_turns(), 1);
assert_eq!(session.id(), Some("thread_test"));
assert_eq!(session.total_tokens(), 165);
assert_eq!(session.last_result().unwrap().result, "hello");
}
#[tokio::test]
#[cfg(unix)]
async fn streaming_turns_accumulate_like_buffered_ones() {
let mut session = streaming_session();
session.stream("first", |_| {}).await.unwrap();
session.stream("second", |_| {}).await.unwrap();
assert_eq!(session.total_turns(), 2);
assert_eq!(session.total_tokens(), 330);
assert_eq!(session.turns_missing_usage(), 0);
}
#[test]
fn a_session_without_a_budget_is_unaffected() {
let mut session = Session::new(test_codex());
session.record_turn(completed(500));
assert!(session.budget().is_none());
assert_eq!(session.total_tokens(), 500);
}
#[test]
fn turns_accumulate_into_the_attached_budget() {
let budget = crate::budget::TokenBudget::builder()
.max_tokens(1000)
.build();
let mut session = Session::new(test_codex()).with_budget(budget.clone());
session.record_turn(completed(300));
session.record_turn(completed(250));
assert_eq!(budget.total_tokens(), 550);
assert_eq!(session.total_tokens(), 550);
}
#[tokio::test]
#[cfg(unix)]
async fn a_turn_past_the_ceiling_is_refused_before_it_runs() {
let budget = crate::budget::TokenBudget::builder()
.max_tokens(100)
.build();
let mut session = streaming_session().with_budget(budget.clone());
session.send("first").await.unwrap();
assert!(budget.total_tokens() >= 100);
let refused = session.send("second").await;
assert!(
matches!(refused, Err(Error::TokenBudgetExceeded { .. })),
"expected the second turn to be refused, got: {refused:?}"
);
assert_eq!(session.total_turns(), 1);
}
#[test]
fn an_all_zero_usage_turn_does_not_advance_the_budget() {
let budget = crate::budget::TokenBudget::builder()
.max_tokens(100)
.build();
let mut session = Session::new(test_codex()).with_budget(budget.clone());
for _ in 0..10 {
session.record_turn(completed(0));
}
assert_eq!(budget.total_tokens(), 0);
assert_eq!(
budget.turns_missing_usage(),
0,
"zero is reported, not missing"
);
assert!(budget.check().is_ok(), "a review-only session never stops");
}
#[test]
fn a_turn_with_no_usage_is_counted_as_unmeasured() {
let budget = crate::budget::TokenBudget::builder()
.max_tokens(100)
.build();
let mut session = Session::new(test_codex()).with_budget(budget.clone());
session.record_turn(turn(&[r#"{"type":"turn.completed"}"#]));
assert_eq!(budget.total_tokens(), 0);
assert_eq!(budget.turns_missing_usage(), 1);
assert_eq!(session.turns_missing_usage(), 1);
}
}