use std::collections::{HashMap, HashSet};
use crate::error::{Error, Result};
use crate::handlers::Handlers;
use crate::process::HarnessOptions;
use crate::protocol::{
InitializeConversationResponse, InputEvent, OutputEventEvent, StepUpdate, ToolConfirmation,
TrajectoryStateUpdate, TrajectoryStateUpdateState, UsageMetadata,
};
use crate::steps::{Step, StepAssembler};
use crate::RawClient;
#[derive(Debug)]
pub struct Client {
raw: RawClient,
handlers: Handlers,
assembler: StepAssembler,
usage: Option<UsageMetadata>,
trajectory_usage: HashMap<String, UsageMetadata>,
}
impl Client {
pub async fn launch(options: HarnessOptions) -> Result<Self> {
Self::launch_with(options, Handlers::new()).await
}
pub async fn launch_with(options: HarnessOptions, handlers: Handlers) -> Result<Self> {
let raw = RawClient::launch(options).await?;
let mut assembler = StepAssembler::new(raw.cascade_id().map(str::to_string));
let initialize = raw.initialize_response().clone();
for update in initialize.history.iter().cloned() {
assembler.ingest(update);
}
Ok(Self {
usage: initialize.cumulative_usage.clone(),
trajectory_usage: initialize
.trajectory_usage
.iter()
.filter_map(|e| Some((e.trajectory_id.clone()?, e.usage.clone()?)))
.collect(),
raw,
handlers,
assembler,
})
}
pub fn cascade_id(&self) -> Option<&str> {
self.raw.cascade_id()
}
pub fn initialize_response(&self) -> &InitializeConversationResponse {
self.raw.initialize_response()
}
pub fn usage(&self) -> Option<&UsageMetadata> {
self.usage.as_ref()
}
pub fn trajectory_usage(&self) -> &HashMap<String, UsageMetadata> {
&self.trajectory_usage
}
pub fn raw(&mut self) -> &mut RawClient {
&mut self.raw
}
pub async fn send(&mut self, prompt: impl Into<String>) -> Result<Turn<'_>> {
self.send_event(InputEvent::user(prompt)).await
}
pub async fn send_event(&mut self, event: InputEvent) -> Result<Turn<'_>> {
self.raw.send(&event).await?;
Ok(Turn {
client: self,
finished: false,
failure: None,
answered: HashSet::new(),
})
}
pub async fn cancel(&mut self) -> Result<()> {
self.raw.send(&InputEvent::halt()).await
}
pub async fn shutdown(self) -> Result<()> {
self.raw.shutdown().await
}
fn record_usage(&mut self, update: crate::protocol::UsageUpdate) {
if let Some(total) = update.total {
self.usage = Some(total);
}
for entry in update.agents {
if let (Some(id), Some(usage)) = (entry.trajectory_id, entry.usage) {
self.trajectory_usage.insert(id, usage);
}
}
}
}
#[derive(Debug)]
pub struct Turn<'a> {
client: &'a mut Client,
finished: bool,
failure: Option<String>,
answered: HashSet<(String, u32, &'static str)>,
}
impl Turn<'_> {
pub async fn next_step(&mut self) -> Result<Option<Step>> {
loop {
if self.finished {
return match self.failure.take() {
Some(message) => Err(Error::Turn { message }),
None => Ok(None),
};
}
let Some(event) = self.client.raw.next_event().await? else {
self.finished = true;
continue;
};
match event.into_event() {
Some(OutputEventEvent::StepUpdate(update)) => {
self.answer_in_band_requests(&update).await?;
return Ok(Some(self.client.assembler.ingest(update)));
}
Some(OutputEventEvent::ToolCall(call)) => {
let response = self.client.handlers.call_tool(call).await;
self.client
.raw
.send(&InputEvent::tool_response(response))
.await?;
}
Some(OutputEventEvent::CallHookRequest(request)) => {
let response = self.client.handlers.call_hook(request).await;
self.client
.raw
.send(&InputEvent::hook_response(response))
.await?;
}
Some(OutputEventEvent::PolicyDecisionRequest(request)) => {
let response = self.client.handlers.call_policy(request).await;
self.client
.raw
.send(&InputEvent::policy_response(response))
.await?;
}
Some(OutputEventEvent::UsageUpdate(update)) => self.client.record_usage(update),
Some(OutputEventEvent::TrajectoryStateUpdate(update)) => {
self.observe_trajectory(&update)
}
Some(OutputEventEvent::SessionEndResponse(_)) => self.finished = true,
Some(OutputEventEvent::InitializeConversationResponse(_)) | None => {}
}
}
}
pub async fn collect_text(&mut self) -> Result<String> {
let mut out = String::new();
let mut seen = HashSet::new();
while let Some(step) = self.next_step().await? {
if step.is_final() && seen.insert(step.id()) {
if let Some(text) = step.user_facing_text() {
out.push_str(text);
}
}
}
Ok(out)
}
pub fn is_finished(&self) -> bool {
self.finished
}
async fn answer_in_band_requests(&mut self, update: &StepUpdate) -> Result<()> {
let trajectory_id = update.trajectory_id.clone().unwrap_or_default();
let step_index = update.step_index.unwrap_or_default();
if let Some(request) = update.questions_request.clone() {
if self
.answered
.insert((trajectory_id.clone(), step_index, "questions"))
{
let mut response = self.client.handlers.call_questions(request).await;
response.trajectory_id = Some(trajectory_id.clone());
response.step_index = Some(step_index);
self.client
.raw
.send(&InputEvent {
question_response: Some(response),
..Default::default()
})
.await?;
}
}
if update.tool_confirmation_request.is_some()
&& self
.answered
.insert((trajectory_id.clone(), step_index, "confirm"))
{
let accepted = self.client.handlers.call_confirm(update.clone()).await;
self.client
.raw
.send(&InputEvent {
tool_confirmation: Some(ToolConfirmation {
trajectory_id: Some(trajectory_id),
step_index: Some(step_index),
accepted: Some(accepted),
}),
..Default::default()
})
.await?;
}
Ok(())
}
fn observe_trajectory(&mut self, update: &TrajectoryStateUpdate) {
let id = update.trajectory_id.as_deref().unwrap_or_default();
let error = update.error.clone().filter(|e| !e.is_empty());
if !self.client.assembler.is_main(id) {
if let Some(error) = error {
log::info!("subagent trajectory {id} failed: {error}");
}
return;
}
match update.state {
Some(TrajectoryStateUpdateState::FullyIdle) => {
self.finished = true;
self.failure = error;
}
Some(TrajectoryStateUpdateState::Cancelled) => {
self.finished = true;
self.failure = Some(error.unwrap_or_else(|| "turn cancelled".into()));
}
_ => {}
}
}
}