use std::collections::BTreeMap;
use std::sync::atomic::AtomicU32;
use crate::cancel;
use crate::client::{CompletionResult, GatewayClient, Message, ToolSchema};
use crate::debug::{DebugCapture, DebugEvent};
use crate::dialects::{ToolDialect, ToolDialectRegistry};
use crate::lua::ToolCallCounts;
use crate::model::CompletionOptions;
use crate::observe::{Observer, detail};
use crate::tools::{ToolId, ToolRegistry};
use crate::untrusted;
use crate::{Error, Result};
use super::support::advance_turn;
pub(crate) struct SectionProgress<'a> {
pub(crate) execution: &'a str,
pub(crate) observer: &'a dyn Observer,
pub(crate) section: &'a str,
pub(crate) turns: &'a AtomicU32,
pub(crate) debug: Option<&'a dyn DebugCapture>,
pub(crate) completion_options: &'a CompletionOptions,
}
#[derive(Debug, Clone, Copy)]
pub(crate) enum ProseMode {
SingleShot,
Loop { max_tool_iterations: usize },
}
#[derive(Debug, Clone)]
pub(crate) struct ProseInferenceResult {
pub text: Option<String>,
pub finish_reason: Option<String>,
}
#[expect(
clippy::too_many_arguments,
reason = "counts and global_aliases extend the loop's borrowed context for per-VM call tracking"
)]
pub(crate) async fn run_tool_loop(
client: &GatewayClient,
schemas: &[ToolSchema],
dispatch: &BTreeMap<String, ToolId>,
registry: &ToolRegistry<'_>,
prose: String,
max_tool_iterations: usize,
progress: SectionProgress<'_>,
counts: Option<&ToolCallCounts>,
global_aliases: Option<&BTreeMap<String, ToolId>>,
) -> Result<(String, Option<String>)> {
let mut conversation = Vec::new();
let outcome = run_prose_inference(
client,
schemas,
dispatch,
registry,
&mut conversation,
prose,
ProseMode::Loop {
max_tool_iterations,
},
progress,
counts,
global_aliases,
)
.await?;
match outcome.text {
Some(text) => Ok((text, outcome.finish_reason)),
None => Err(Error::ToolLoopExhausted),
}
}
#[expect(
clippy::too_many_arguments,
clippy::too_many_lines,
reason = "counts and global_aliases extend the loop's borrowed context for per-VM call tracking"
)]
pub(crate) async fn run_prose_inference(
client: &GatewayClient,
schemas: &[ToolSchema],
dispatch: &BTreeMap<String, ToolId>,
registry: &ToolRegistry<'_>,
conversation: &mut Vec<Message>,
prose: String,
mode: ProseMode,
progress: SectionProgress<'_>,
counts: Option<&ToolCallCounts>,
global_aliases: Option<&BTreeMap<String, ToolId>>,
) -> Result<ProseInferenceResult> {
let SectionProgress {
execution,
observer,
section,
turns,
debug,
completion_options,
} = progress;
let dialect_registry = ToolDialectRegistry::builtin();
let dialect: &dyn ToolDialect = dialect_registry
.get(completion_options.tool_dialect)
.ok_or(Error::UnknownDialect(completion_options.tool_dialect))?;
conversation.push(Message::user(prose));
let tool_arg = if schemas.is_empty() {
None
} else {
Some(schemas)
};
let max_tool_iterations = match mode {
ProseMode::SingleShot => 1,
ProseMode::Loop {
max_tool_iterations,
} => max_tool_iterations,
};
for _ in 0..max_tool_iterations {
let completion = tokio::select! {
biased;
() = cancel::wait_cancelled() => Err(Error::Interrupted),
result = client.complete(conversation, tool_arg, completion_options) => result.map_err(Error::from),
};
if let Err(Error::Interrupted) = &completion {
return Err(Error::Interrupted);
}
if completion.is_err() {
observer.observe(execution, section, detail::MODEL_TURN_FAILED);
}
let completion = completion?;
let turn = advance_turn(turns);
if let Some(capture) = debug {
capture.on_event(
execution,
section,
turn,
DebugEvent::Request {
body: completion.request_body,
},
);
capture.on_event(
execution,
section,
turn,
DebugEvent::Response {
body: completion.response_body.clone(),
finish_reason: completion.finish_reason.clone(),
reasoning_content: completion.reasoning_content.clone(),
},
);
}
observer.observe(execution, section, detail::MODEL_TURN_COMPLETED);
match completion.result {
CompletionResult::Text(text) => {
if completion.finish_reason.as_deref() == Some("length") {
observer.observe(execution, section, detail::MODEL_TURN_TRUNCATED);
}
return Ok(ProseInferenceResult {
text: Some(text),
finish_reason: completion.finish_reason,
});
}
CompletionResult::ToolCalls(calls) => {
let finish_reason = completion.finish_reason.clone();
let mut results: Vec<crate::dialects::FramedToolResult> =
Vec::with_capacity(calls.len());
for call in &calls {
let Some(id) = dispatch.get(&call.name) else {
observer.observe(execution, section, detail::TOOL_CALL_FAILED);
let global_exists =
global_aliases.is_some_and(|g| g.contains_key(&call.name));
let in_scope: Vec<String> = dispatch.keys().cloned().collect();
return Err(Error::OutOfScopeToolCall {
name: call.name.clone(),
global_exists,
in_scope,
});
};
let Some(tool) = registry.get(id) else {
observer.observe(execution, section, detail::TOOL_CALL_FAILED);
return Err(Error::UnknownScopedTool(call.name.clone()));
};
if let Some(counts) = counts {
counts.increment(&call.name)?;
}
let call_result = tokio::select! {
biased;
() = cancel::wait_cancelled() => {
observer.observe(execution, section, detail::TOOL_CALL_FAILED);
return Err(Error::Interrupted);
}
result = tool.call(call.arguments.clone()) => result,
};
observer.observe(
execution,
section,
if call_result.is_ok() {
detail::TOOL_CALL_SUCCEEDED
} else {
detail::TOOL_CALL_FAILED
},
);
let output = call_result.map_err(Error::tool)?;
let result = match output.trust() {
crate::tools::OutputTrust::Untrusted => untrusted::wrap(output.text()),
crate::tools::OutputTrust::Trusted => output.text().to_owned(),
};
results.push(crate::dialects::FramedToolResult::new(
call.id.clone(),
result,
));
}
dialect.echo_tool_results(conversation, &calls, &results)?;
if matches!(mode, ProseMode::SingleShot) {
return Ok(ProseInferenceResult {
text: None,
finish_reason,
});
}
}
}
}
match mode {
ProseMode::SingleShot => Ok(ProseInferenceResult {
text: None,
finish_reason: None,
}),
ProseMode::Loop { .. } => Err(Error::ToolLoopExhausted),
}
}