mod inner;
use inner::{Inner, run_inner};
use super::Resolved;
use crate::prompt::Deps;
use crate::prompt::Error;
use crate::prompt::tool::ToolOutcome;
use serde::Deserialize;
use serde_json::Value;
use std::path::Path;
pub(super) const NAME: &str = "multi_tool";
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Envelope {
invocations: Vec<Invocation>,
#[serde(default)]
on_failure: OnFailure,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Invocation {
name: String,
#[serde(default = "empty_input")]
input: Value,
}
fn empty_input() -> Value {
Value::Object(serde_json::Map::new())
}
#[derive(Deserialize, Default, Clone, Copy, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
enum OnFailure {
#[default]
Abort,
RunAll,
}
#[derive(Debug)]
pub(super) enum Fanout {
Outcome(ToolOutcome),
Stopped,
}
const OK: &str = "ok";
const FAILED: &str = "failed";
const DECLINED: &str = "declined";
const SKIPPED: &str = "skipped";
struct Entry {
name: String,
status: &'static str,
text: String,
}
pub(super) fn fan_out(
outer_id: &str,
input: &Value,
step_dir_abs: &Path,
resolved: &Resolved<'_>,
conv_repo: &Path,
conv_id: &str,
deps: &Deps<'_>,
) -> Result<Fanout, Error> {
let envelope = match Envelope::deserialize(input) {
Ok(envelope) => envelope,
Err(err) => return Ok(Fanout::Outcome(malformed(&err))),
};
let total = envelope.invocations.len();
let mut entries: Vec<Entry> = Vec::with_capacity(total);
let mut failed_at: Option<usize> = None;
for (idx, inv) in envelope.invocations.iter().enumerate() {
if envelope.on_failure == OnFailure::Abort
&& let Some(failed) = failed_at
{
entries.push(skipped(inv, failed, total));
continue;
}
let inner = Inner {
outer_id,
k: idx + 1,
inv,
step_dir_abs,
conv_repo,
conv_id,
};
let entry = match run_inner(&inner, resolved, deps)? {
Some(entry) => entry,
None => return Ok(Fanout::Stopped),
};
if entry.status != OK {
failed_at = failed_at.or(Some(idx + 1));
}
entries.push(entry);
}
Ok(Fanout::Outcome(render(&entries)))
}
fn skipped(inv: &Invocation, failed_at: usize, total: usize) -> Entry {
Entry {
name: inv.name.clone(),
status: SKIPPED,
text: format!(
"not run: on_failure \"abort\" ended the envelope after \
[{failed_at}/{total}] failed."
),
}
}
fn render(entries: &[Entry]) -> ToolOutcome {
let total = entries.len();
let ok = entries.iter().filter(|e| e.status == OK).count();
let skip = entries.iter().filter(|e| e.status == SKIPPED).count();
let failed = total - ok - skip;
let mut out = format!("{total} invocations: {ok} ok, {failed} failed, {skip} skipped\n");
for (idx, entry) in entries.iter().enumerate() {
let (k, name, status, text) = (idx + 1, &entry.name, entry.status, &entry.text);
out.push_str(&format!(
"\n=== [{k}/{total}] {name}: {status} ===\n{text}\n"
));
}
ToolOutcome {
content: out.into_bytes(),
is_error: failed > 0,
}
}
fn malformed(err: &serde_json::Error) -> ToolOutcome {
ToolOutcome {
content: format!(
"{NAME}: malformed envelope: {err}. Expected \
{{\"invocations\": [{{\"name\": \"<tool>\", \"input\": {{...}}}}, ...], \
\"on_failure\": \"abort\"|\"run_all\"}}."
)
.into_bytes(),
is_error: true,
}
}