use crate::context::{self, SkillMeta};
use crate::providers::InferenceProvider;
use crate::reasoning::{
build_chat_request_messages, initial_prev_resp_id, reasoning_artifact_tokens,
warn_on_missing_reasoning_artifacts,
};
use crate::sessions::{AssistantResponse, RequestContext, SessionCommand, SessionState};
use crate::tools::ToolOutput;
use crate::tools::context::ToolContext;
use crate::tools::load_tools::{LoadToolsArgs, apply_load_tools};
use crate::tools::set_working_dir::{SetWorkingDirArgs, resolve_working_dir_path};
use crate::tools::unload_tools::{UnloadToolsArgs, apply_unload_tools};
use choreo_ai_protocols::openai::{ChatRequestMessage, ChatToolDefinition};
use choreo_ai_protocols::{
ChatToolCall, ChatTurnRequest, ChatTurnResult, StreamEvent, ToolResultItem,
model_reasoning_capability,
};
use choreo_proto::{
AssistantToolCallRecord, DaemonMessage, OutputStream, ReasoningProducer, SessionEvent,
SessionStatus,
};
use std::collections::{HashMap, HashSet};
use std::io;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use std::time::Instant;
use tracing::{debug, warn};
mod system_content;
mod tool_execution;
pub(crate) use system_content::*;
pub(crate) use tool_execution::*;
fn resolve_reasoning_effort(
client: &InferenceProvider,
model: &str,
session_id: u64,
turn_iter: u32,
configured_effort: &str,
) -> String {
if configured_effort == "off" {
return configured_effort.to_string();
}
let slug = client.provider_slug();
let capability = model_reasoning_capability(slug, model);
if capability.available_effort_levels.is_empty() {
warn!(
session_id, turn = turn_iter, model,
effort = %configured_effort,
"model does not support reasoning, disabling",
);
"off".to_string()
} else if !capability
.available_effort_levels
.iter()
.any(|l| l == configured_effort)
{
warn!(
session_id, turn = turn_iter, model,
effort = %configured_effort,
valid = ?capability.available_effort_levels,
"reasoning effort '{}' not in model's capability set, disabling",
configured_effort,
);
"off".to_string()
} else {
configured_effort.to_string()
}
}
fn estimate_prompt_tokens(
model: &str,
messages: &[ChatRequestMessage],
tools: &[ChatToolDefinition],
) -> (Option<&'static tiktoken::CoreBpe>, u32) {
let encoding =
tiktoken::encoding_for_model(model).or_else(|| tiktoken::get_encoding("cl100k_base"));
let estimated = if let Some(enc) = &encoding {
let content_tokens: u32 = messages
.iter()
.filter_map(|m| m.content.as_deref())
.map(|text| {
let n = enc.count(text);
#[allow(clippy::cast_possible_truncation)]
{
n as u32
}
})
.sum();
let image_tokens: u32 = messages
.iter()
.map(|m| {
#[allow(clippy::cast_possible_truncation)]
(m.images.len() as u32).saturating_mul(IMAGE_TOKEN_ESTIMATE)
})
.sum();
let tool_call_tokens: u32 = messages
.iter()
.filter_map(|m| m.tool_calls.as_ref())
.flat_map(|calls| calls.iter())
.map(|tc| {
#[allow(clippy::cast_possible_truncation)]
let tc_tokens = (enc.count(&tc.id) + enc.count(&tc.kind)) as u32
+ (enc.count(&tc.function.name) + enc.count(&tc.function.arguments)) as u32;
tc_tokens
})
.sum();
let tool_def_tokens: u32 = tools
.iter()
.filter_map(|def| match serde_json::to_string(def) {
Ok(s) => {
#[allow(clippy::cast_possible_truncation)]
Some(enc.count(&s) as u32)
}
Err(e) => {
warn!(error = %e, "failed to serialize tool definition for token estimation");
None
}
})
.sum();
let artifact_tokens: u32 = messages
.iter()
.filter_map(|m| m.reasoning_artifact.as_ref())
.map(|artifact| reasoning_artifact_tokens(enc, artifact))
.sum();
content_tokens + tool_call_tokens + tool_def_tokens + artifact_tokens + image_tokens
} else {
tracing::warn!("no tiktoken encoding available for {model}");
0
};
(encoding, estimated)
}
fn sort_by_call_order<T>(
tool_calls: &[AssistantToolCallRecord],
items: &mut [T],
call_id_of: impl Fn(&T) -> &str,
) {
let order: HashMap<&str, usize> = tool_calls
.iter()
.enumerate()
.map(|(i, tc)| (tc.call_id.as_str(), i))
.collect();
if order.is_empty() {
return;
}
items.sort_by_key(|item| order.get(call_id_of(item)).copied().unwrap_or(usize::MAX));
}
enum PendingConfigChange {
LoadTools(Vec<String>),
UnloadTools(Vec<String>),
SetWorkingDir(Option<PathBuf>),
}
fn is_session_config_tool(name: &str) -> bool {
matches!(name, "load_tools" | "unload_tools" | "set_working_dir")
}
fn concurrent_tool_status_label(tools: &[ChatToolCall]) -> String {
if tools.len() == 1 {
tools.first().map(|t| t.name.clone()).unwrap_or_default()
} else {
"(parallel)".into()
}
}
fn pending_config_change(
tool_call: &ChatToolCall,
output: &ToolOutput,
base_working_dir: Option<&Path>,
) -> Option<PendingConfigChange> {
if !is_session_config_tool(&tool_call.name) {
return None;
}
match tool_call.name.as_str() {
"load_tools" => {
let Ok(args) = serde_json::from_str::<LoadToolsArgs>(&tool_call.arguments_json) else {
warn!(
tool_call_id = %tool_call.id,
"load_tools: could not parse args to mirror onto worker config",
);
return None;
};
Some(PendingConfigChange::LoadTools(args.groups))
}
"unload_tools" => {
let Ok(args) = serde_json::from_str::<UnloadToolsArgs>(&tool_call.arguments_json)
else {
warn!(
tool_call_id = %tool_call.id,
"unload_tools: could not parse args to mirror onto worker config",
);
return None;
};
Some(PendingConfigChange::UnloadTools(args.groups))
}
"set_working_dir" => {
if let Some(path) = output
.result_json
.as_ref()
.and_then(|v| v.get("path"))
.and_then(|v| v.as_str())
{
return Some(PendingConfigChange::SetWorkingDir(Some(PathBuf::from(
path,
))));
}
let Ok(args) = serde_json::from_str::<SetWorkingDirArgs>(&tool_call.arguments_json)
else {
warn!(
tool_call_id = %tool_call.id,
"set_working_dir: could not parse args to mirror onto worker config",
);
return Some(PendingConfigChange::SetWorkingDir(None));
};
let path = resolve_working_dir_path(&args.path, base_working_dir).ok();
Some(PendingConfigChange::SetWorkingDir(path))
}
_ => None,
}
}
fn apply_pending_config_change(
session: &mut SessionState,
change: &PendingConfigChange,
protected: &HashSet<String>,
) {
match change {
PendingConfigChange::LoadTools(groups) => {
apply_load_tools(&mut session.config.active_tool_groups, groups);
debug!(groups = ?groups, "mirrored load_tools onto worker session config");
}
PendingConfigChange::UnloadTools(groups) => {
apply_unload_tools(&mut session.config.active_tool_groups, groups, protected);
debug!(groups = ?groups, "mirrored unload_tools onto worker session config");
}
PendingConfigChange::SetWorkingDir(path) => {
if let Some(path) = path {
session.config.working_dir = Some(path.clone());
}
session.discovered_skills = None;
debug!(path = ?path, "mirrored set_working_dir onto worker session config");
}
}
}
const MAX_TRUNCATION_RECOVERIES: u32 = 3;
const TRUNCATION_RECOVERY_INSTRUCTION: &str = "Your previous response was \
truncated at the model's output-token limit while a tool call was being \
generated, so that call was discarded and did not run. Continue the task, \
but produce much smaller outputs: split large writes into several smaller \
tool calls (write a file in sections and append, or use multiple smaller \
files), and avoid emitting very large tool arguments in a single call.";
pub(crate) fn run_agent_loop(
client: &InferenceProvider,
session: &mut SessionState,
model: &str,
request_id: u32,
cancel_rx: &crossbeam_channel::Receiver<()>,
ctx: &RequestContext,
user_text: Option<&str>,
) -> io::Result<bool> {
let max_turns = ctx.max_turns;
let limited = max_turns > 0;
let provider_slug = client.provider_slug();
let mut prev_resp_id = initial_prev_resp_id(session, provider_slug, model);
let mut tool_results: Vec<ToolResultItem> = Vec::new();
let mut known_hint_paths: Vec<PathBuf> = Vec::new();
let mut pending_hints: Vec<String> = Vec::new();
let mut pending_user_text: Option<String> = None;
let mut truncation_recoveries: u32 = 0;
warn_on_missing_reasoning_artifacts(session, ctx.session_id, provider_slug, model);
if session.discovered_skills.is_none() {
session.discovered_skills = Some(context::discover_skills_ambient(
session.config.working_dir.as_deref(),
));
}
let mut turn_iter: u32 = 0;
let request_first_turn_id = session.next_turn_id;
loop {
if limited && turn_iter >= max_turns {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("tool loop exceeded {max_turns} iterations"),
));
}
debug!(
session_id = ctx.session_id,
turn = turn_iter,
"agent loop turn"
);
let configured = session.config.reasoning_effort.as_deref().unwrap_or("off");
let thinking_effort =
resolve_reasoning_effort(client, model, ctx.session_id, turn_iter, configured);
crate::metrics::record_turn(model);
let tools = ctx
.tool_registry
.available_definitions(&session.config.active_tool_groups);
if is_cancelled_once(cancel_rx) {
return Ok(true);
}
let turn_user_text = if turn_iter == 0 {
user_text.map(std::string::ToString::to_string)
} else {
pending_user_text.take()
};
let (current_turn_id, _) = session.start_turn(turn_user_text);
broadcast_turn_appended(&ctx.cmd_tx, session, ctx.session_id, current_turn_id);
if ctx
.cmd_tx
.send(SessionCommand::StatusChanged(SessionStatus::Inference))
.is_err()
{
return Ok(false);
}
let system_content = {
let skills: &[SkillMeta] = session.discovered_skills.as_deref().unwrap_or_default();
build_system_content(
&SystemContentParams {
working_dir: session.config.working_dir.as_deref(),
context_config: &session.config.context_config,
skills,
loaded_skill_bodies: &session.loaded_skill_bodies,
tool_registry: &ctx.tool_registry,
pending_hints: &pending_hints,
session_title: session.config.title.as_deref(),
},
&mut session.context_cache,
)
};
pending_hints.clear();
let messages = build_chat_request_messages(
session,
Some(&system_content),
provider_slug,
model,
Some(request_first_turn_id),
);
let (encoding, estimated_prompt_tokens) = estimate_prompt_tokens(model, &messages, &tools);
let _ = ctx
.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Session {
session_id: Some(ctx.session_id),
event: SessionEvent::Started {
request_id,
turn_id: current_turn_id,
estimated_prompt_tokens,
},
}));
let mut retry_cb: Option<choreo_ai_protocols::openai::RetryCallback> = Some(Box::new({
let cmd_tx = ctx.cmd_tx.clone();
move |attempt, max_attempts, delay| {
let _ = cmd_tx.send(SessionCommand::StatusChanged(SessionStatus::Retrying {
attempt,
max_attempts,
#[allow(clippy::cast_possible_truncation)]
delay_ms: delay.as_millis() as u64,
}));
}
}));
let mut output_token_count: u32 = 0;
let oc_session_id = ctx.session_id.to_string();
let oc_request_id = request_id.to_string();
match client.chat_completion_turn_streaming(
ChatTurnRequest {
model,
messages: &messages,
tools: &tools,
thinking_effort,
on_retry: &mut retry_cb,
cancel_rx: Some(cancel_rx),
previous_response_id: prev_resp_id.as_deref(),
tool_results: &tool_results,
programmatic_tool_calling: client.supports_programmatic_tool_calling(model),
session_id: oc_session_id,
request_id: oc_request_id,
},
&mut |event| {
match event {
StreamEvent::Answer(text) => {
if let Some(enc) = &encoding {
let n = enc.count(&text);
#[allow(clippy::cast_possible_truncation)]
{
output_token_count += n as u32;
}
}
let _ =
ctx.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Session {
session_id: Some(ctx.session_id),
event: SessionEvent::OutputChunk {
request_id,
stream: OutputStream::Answer,
data: text.into_bytes(),
},
}));
let _ =
ctx.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Session {
session_id: Some(ctx.session_id),
event: SessionEvent::LiveOutputTokenCount {
request_id,
output_tokens: output_token_count,
},
}));
}
StreamEvent::Reasoning(text) => {
if let Some(enc) = &encoding {
let n = enc.count(&text);
#[allow(clippy::cast_possible_truncation)]
{
output_token_count += n as u32;
}
}
let _ =
ctx.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Session {
session_id: Some(ctx.session_id),
event: SessionEvent::OutputChunk {
request_id,
stream: OutputStream::Reasoning,
data: text.into_bytes(),
},
}));
let _ =
ctx.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Session {
session_id: Some(ctx.session_id),
event: SessionEvent::LiveOutputTokenCount {
request_id,
output_tokens: output_token_count,
},
}));
}
_ => {}
}
Ok(())
},
) {
Ok(ChatTurnResult::FinalText(final_text)) => {
debug!(
session_id = ctx.session_id,
turn = turn_iter,
response_len = final_text.content.len(),
reasoning = final_text.reasoning.as_deref().unwrap_or_default(),
"model returned final text",
);
let token_usage = final_text.usage;
accumulate_token_usage(session, token_usage.as_ref(), turn_iter, ctx);
broadcast_token_usage(ctx, session);
let mut assistant_text = final_text.content;
if final_text.truncated {
warn!(
session_id = ctx.session_id,
turn = turn_iter,
"provider cut the final answer at the output length limit"
);
assistant_text.push_str("\n\n⚠ response truncated (length limit)");
}
let producer = ReasoningProducer {
provider_slug: provider_slug.to_string(),
model: model.to_string(),
};
session.set_assistant_response(
current_turn_id,
AssistantResponse {
text: Some(assistant_text),
reasoning: final_text.reasoning,
token_usage,
reasoning_artifact: final_text.reasoning_artifact.clone(),
reasoning_producer: Some(producer.clone()),
..Default::default()
},
);
session
.config
.last_response_id
.clone_from(&final_text.response_id);
session.config.last_response_id_producer = Some(producer);
finalize_and_broadcast_turn(session, ctx, current_turn_id)?;
tool_results.clear();
return Ok(false);
}
Ok(ChatTurnResult::ToolUse(tool_use)) => {
let token_usage = tool_use.usage;
accumulate_token_usage(session, token_usage.as_ref(), turn_iter, ctx);
broadcast_token_usage(ctx, session);
let tool_call_records: Vec<AssistantToolCallRecord> = tool_use
.tool_calls
.iter()
.map(|tc| AssistantToolCallRecord {
call_id: tc.id.clone(),
name: tc.name.clone(),
arguments_json: tc.arguments_json.clone(),
})
.collect();
let description_by_call: HashMap<String, String> = tool_use
.tool_calls
.iter()
.map(|tc| (tc.id.clone(), ctx.tool_registry.describe_invocation(tc)))
.collect();
let invocation_descriptions: Vec<String> = tool_call_records
.iter()
.map(|tc| {
description_by_call
.get(&tc.call_id)
.cloned()
.unwrap_or_default()
})
.collect();
let producer = ReasoningProducer {
provider_slug: provider_slug.to_string(),
model: model.to_string(),
};
session.set_assistant_response(
current_turn_id,
AssistantResponse {
text: tool_use.content.clone(),
reasoning: tool_use.reasoning.clone(),
tool_calls: tool_call_records.clone(),
token_usage,
reasoning_artifact: tool_use.reasoning_artifact.clone(),
reasoning_producer: Some(producer.clone()),
},
);
session.seed_tool_results(
current_turn_id,
&tool_call_records,
&invocation_descriptions,
);
broadcast_turn_appended(&ctx.cmd_tx, session, ctx.session_id, current_turn_id);
prev_resp_id.clone_from(&tool_use.response_id);
session.config.last_response_id.clone_from(&prev_resp_id);
session.config.last_response_id_producer = Some(producer);
tool_results.clear();
let (mutators, concurrent): (Vec<_>, Vec<_>) = tool_use
.tool_calls
.into_iter()
.partition(|tc| is_session_config_tool(&tc.name));
let turn_base_working_dir = session.config.working_dir.clone();
let mut pending_config_changes: Vec<PendingConfigChange> = Vec::new();
let mut cancelled = false;
let mut executed_tool_calls: HashSet<String> = HashSet::new();
for tool_call in mutators {
if is_cancelled_once(cancel_rx) {
cancelled = true;
break;
}
let invocation_description = description_by_call
.get(&tool_call.id)
.cloned()
.unwrap_or_default();
if let Err(e) =
ctx.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Session {
session_id: Some(ctx.session_id),
event: SessionEvent::ToolCallStarted {
request_id,
call_id: tool_call.id.clone(),
tool_name: tool_call.name.clone(),
arguments_json: tool_call.arguments_json.clone(),
invocation_description: invocation_description.clone(),
},
}))
{
warn!(%request_id, call_id = %tool_call.id, error = %e, "failed to broadcast ToolCallStarted");
}
let tool_timeout =
determine_tool_timeout(&tool_call.name, &tool_call.arguments_json)
.unwrap_or(Duration::from_secs(60));
if ctx
.cmd_tx
.send(SessionCommand::StatusChanged(SessionStatus::ToolCall(
tool_call.name.clone(),
)))
.is_err()
{
return Ok(false);
}
debug!(
session_id = ctx.session_id,
turn = turn_iter,
tool_name = %tool_call.name,
tool_call_id = %tool_call.id,
args_preview = %tool_call
.arguments_json
.get(..tool_call.arguments_json.len().min(200))
.unwrap_or(&tool_call.arguments_json),
"executing tool (serial)",
);
let turn_working_dir = session.config.working_dir.clone();
let (mut output, tool_cancelled, image) = execute_tool_with_timeout(
&tool_call,
ctx.substrate_credential.as_ref(),
turn_working_dir.as_deref(),
tool_timeout,
request_id,
ctx.session_id,
session,
cancel_rx,
ctx,
&invocation_description,
);
if tool_cancelled {
cancelled = true;
}
record_tool_completion(
request_id,
session,
&tool_call,
&mut output,
image,
ctx,
current_turn_id,
&mut tool_results,
&mut known_hint_paths,
&mut pending_hints,
);
executed_tool_calls.insert(tool_call.id.clone());
if !output.is_error
&& let Some(change) = pending_config_change(
&tool_call,
&output,
turn_base_working_dir.as_deref(),
)
{
pending_config_changes.push(change);
}
if cancelled {
break;
}
}
if !cancelled && !concurrent.is_empty() {
for tc in &concurrent {
let invocation_description =
description_by_call.get(&tc.id).cloned().unwrap_or_default();
if let Err(e) =
ctx.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Session {
session_id: Some(ctx.session_id),
event: SessionEvent::ToolCallStarted {
request_id,
call_id: tc.id.clone(),
tool_name: tc.name.clone(),
arguments_json: tc.arguments_json.clone(),
invocation_description,
},
}))
{
warn!(%request_id, call_id = %tc.id, error = %e, "failed to broadcast ToolCallStarted");
}
}
if ctx
.cmd_tx
.send(SessionCommand::StatusChanged(SessionStatus::ToolCall(
concurrent_tool_status_label(&concurrent),
)))
.is_err()
{
return Ok(false);
}
debug!(
session_id = ctx.session_id,
turn = turn_iter,
count = concurrent.len(),
"dispatching {} tools concurrently",
concurrent.len(),
);
let cancel_flag = Arc::new(AtomicBool::new(false));
let tool_ctx = ToolContext {
session_id: ctx.session_id,
db: Arc::clone(&ctx.db),
daemon_tx: ctx.daemon_tx.clone(),
active_tool_groups: session.config.active_tool_groups.clone(),
reasoning_effort: session.config.reasoning_effort.clone(),
selected_model: session.config.selected_model.clone(),
working_dir: session.config.working_dir.clone(),
cancelled: Arc::clone(&cancel_flag),
account_name: session.config.account_name.clone(),
discovered_skills: session
.discovered_skills
.clone()
.map(std::sync::Arc::new),
};
let cmd_tx = ctx.cmd_tx.clone();
let reg = Arc::clone(&ctx.tool_registry);
let (batch_tx, batch_rx) = crossbeam_channel::unbounded::<ToolHandle>();
let mut call_infos: Vec<CallInfo> = Vec::with_capacity(concurrent.len());
for tool_call in concurrent {
let timeout =
determine_tool_timeout(&tool_call.name, &tool_call.arguments_json);
let invocation_description = description_by_call
.get(&tool_call.id)
.cloned()
.unwrap_or_default();
let started_at = Instant::now();
let call_id = tool_call.id.clone();
let tool_name = tool_call.name.clone();
let arguments_json = tool_call.arguments_json.clone();
let kill_tx = spawn_single_tool(SpawnToolArgs {
tool_call,
timeout,
request_id,
session_id: ctx.session_id,
registry: Arc::clone(®),
cmd_tx: cmd_tx.clone(),
x_credentials: ctx.substrate_credential.clone(),
working_dir: session.config.working_dir.clone(),
ctx: tool_ctx.clone(),
invocation_description: invocation_description.clone(),
started_at,
result_tx: batch_tx.clone(),
});
call_infos.push(CallInfo {
call_id,
tool_name,
arguments_json,
invocation_description,
started_at,
kill_tx,
});
}
drop(batch_tx);
let batch_size = call_infos.len();
let mut process_tool_handle =
|ToolHandle {
tool_call,
mut output,
image,
started_at,
}: ToolHandle| {
let elapsed = started_at.elapsed();
debug!(
session_id = ctx.session_id,
turn = turn_iter,
tool_name = %tool_call.name,
elapsed_ms = elapsed.as_millis(),
result_len = output.content.len(),
is_error = output.is_error,
"tool finished (concurrent)",
);
record_tool_completion(
request_id,
session,
&tool_call,
&mut output,
image,
ctx,
current_turn_id,
&mut tool_results,
&mut known_hint_paths,
&mut pending_hints,
);
executed_tool_calls.insert(tool_call.id.clone());
};
let mut delivered: HashSet<String> = HashSet::with_capacity(batch_size);
while delivered.len() < batch_size {
let (cancelled_now, handle_msg) = crossbeam_channel::select_biased! {
recv(cancel_rx) -> _ => (true, None),
recv(batch_rx) -> msg => (false, Some(msg)),
};
if cancelled_now {
cancel_flag.store(true, Ordering::Relaxed);
cancelled = true;
for info in &call_infos {
let _ = info.kill_tx.send(());
}
while let Ok(handle) = batch_rx.try_recv() {
delivered.insert(handle.tool_call.id.clone());
process_tool_handle(handle);
}
while delivered.len() < batch_size {
if let Ok(handle) = batch_rx.recv() {
delivered.insert(handle.tool_call.id.clone());
process_tool_handle(handle);
} else {
warn!(
session_id = ctx.session_id,
request_id,
delivered = delivered.len(),
expected = batch_size,
"concurrent tool batch ended early after cancel; synthesizing missing tool results",
);
for info in missing_calls(&call_infos, &delivered) {
process_tool_handle(panic_tool_handle(info));
}
break;
}
}
break;
}
if let Some(msg) = handle_msg {
if let Ok(handle) = msg {
delivered.insert(handle.tool_call.id.clone());
process_tool_handle(handle);
} else {
warn!(
session_id = ctx.session_id,
request_id,
delivered = delivered.len(),
expected = batch_size,
"concurrent tool batch ended early; synthesizing missing tool results",
);
for info in missing_calls(&call_infos, &delivered) {
process_tool_handle(panic_tool_handle(info));
}
break;
}
}
}
sort_by_call_order(&tool_call_records, &mut tool_results, |r| {
r.call_id.as_str()
});
}
for change in &pending_config_changes {
apply_pending_config_change(
session,
change,
ctx.tool_registry.protected_groups(),
);
}
if cancelled {
session.mark_unexecuted_tool_results(current_turn_id, &executed_tool_calls);
broadcast_turn_appended(&ctx.cmd_tx, session, ctx.session_id, current_turn_id);
return Ok(true);
}
}
Ok(_) => {
warn!("provider returned an unhandled ChatTurnResult variant");
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"provider returned an unhandled turn result variant",
));
}
Err(choreo_proto::InferenceError::Cancelled) => {
return Ok(true);
}
Err(e) => match e {
choreo_proto::InferenceError::TruncatedToolCall { discarded } => {
truncation_recoveries = truncation_recoveries.saturating_add(1);
let names = discarded
.iter()
.map(|d| d.name.as_str())
.filter(|name| !name.is_empty())
.collect::<Vec<_>>()
.join(", ");
let subject = if names.is_empty() {
"A tool call".to_string()
} else {
format!("The {names} tool call")
};
tracing::warn!(
session_id = ctx.session_id,
request_id,
recovery = truncation_recoveries,
max_recoveries = MAX_TRUNCATION_RECOVERIES,
tools = %names,
"truncated tool call: discarding the partial call and retrying in smaller steps",
);
session.set_assistant_response(
current_turn_id,
AssistantResponse {
text: Some(format!(
"[{subject} was cut off at the model's output-token \
limit before its arguments finished, so it did not run.]"
)),
..Default::default()
},
);
finalize_and_broadcast_turn(session, ctx, current_turn_id)?;
tool_results.clear();
prev_resp_id = None;
session.config.last_response_id = None;
session.config.last_response_id_producer = None;
if truncation_recoveries >= MAX_TRUNCATION_RECOVERIES {
tracing::warn!(
session_id = ctx.session_id,
request_id,
recovery = truncation_recoveries,
"output-token truncation persisted after the recovery budget; \
ending the request",
);
return Ok(false);
}
pending_user_text = Some(TRUNCATION_RECOVERY_INSTRUCTION.to_string());
}
other => {
session.set_turn_error(current_turn_id, other.to_string());
tracing::debug!(
session_id = ctx.session_id,
turn_id = current_turn_id,
error = %other,
"failure marked on turn; finalize will deliver the error turn to clients via TurnAppended",
);
if let Err(persist_err) =
finalize_and_broadcast_turn(session, ctx, current_turn_id)
{
warn!(
session_id = ctx.session_id,
turn_id = current_turn_id,
error = %persist_err,
"failed to persist the failed turn; the inference error is still reported",
);
}
return Err(other.into());
}
},
}
turn_iter += 1;
}
}
pub const IMAGE_TOKEN_ESTIMATE: u32 = 1000;
#[cfg(test)]
#[serial_test::serial(catalog)]
mod tests;