use std::{
io,
path::Path,
sync::{Arc, Mutex},
time::Duration,
};
use basis::{
AllowAll, Approver, Bound, CancellationToken, DenyAll, Event, EventSink, ModelSelector,
RunOutcome, RunSpec, Runtime, RuntimeBuilder, ShellAccess, TurnOptions, Workspace,
WorkspaceBuilder, provider,
};
use serde_json::Value;
use tokio::time::{self, Instant};
use crate::{
Error,
approve::Approve,
data_dir::{AgentPaths, DataDir, valid_task_handle},
events::EventLog,
inbox,
live::DriveContext,
lock,
state::{
MAX_RESULT_BYTES, MAX_TASKS, MessageReply, PendingTerminal, TaskMeta, bounded_text,
cancel_requested, load_meta, now_ms, read_terminal, request_cancel, save_meta,
write_terminal,
},
};
pub const POLL: Duration = Duration::from_millis(100);
#[derive(Debug, Clone, PartialEq)]
pub enum WaitOutcome {
Terminal(Value),
TimedOut { attached: bool },
}
pub(crate) async fn wait_for_terminal(
data: &DataDir,
task: &str,
timeout: Duration,
ctx: &DriveContext,
) -> Result<WaitOutcome, String> {
let deadline = Instant::now().checked_add(timeout);
loop {
if let Some(terminal) = poll_once(data.clone(), task.to_string(), ctx.clone()).await? {
return Ok(WaitOutcome::Terminal(terminal));
}
if deadline.is_some_and(|deadline| Instant::now() >= deadline) {
return Ok(WaitOutcome::TimedOut {
attached: is_attached(data, task).await?,
});
}
time::sleep(POLL).await;
}
}
async fn poll_once(
data: DataDir,
task: String,
ctx: DriveContext,
) -> Result<Option<Value>, String> {
let handle = tokio::runtime::Handle::current();
tokio::task::spawn_blocking(move || -> Result<Option<Value>, String> {
let paths = resolve(&data, &task)?;
if let Some(terminal) = read_terminal(&paths)? {
return Ok(Some(terminal));
}
match try_attach(&paths)? {
Some(guard) => handle.block_on(drive(&data, &task, guard, &ctx)),
None => Ok(None),
}
})
.await
.unwrap_or_else(|error| Err(format!("poll task: {error}")))
}
async fn is_attached(data: &DataDir, task: &str) -> Result<bool, String> {
let data = data.clone();
let task = task.to_string();
tokio::task::spawn_blocking(move || -> Result<bool, String> {
let paths = resolve(&data, &task)?;
Ok(lock::is_held(&paths.attach_lock()))
})
.await
.unwrap_or_else(|error| Err(format!("check attach lock: {error}")))
}
pub(crate) async fn wait_for_message(
data: &DataDir,
task: &str,
message_id: &str,
timeout: Duration,
prompt_host: Option<Arc<dyn crate::approve::PromptHost>>,
) -> Result<WaitOutcome, String> {
let deadline = Instant::now().checked_add(timeout);
let ctx = DriveContext::new(None, prompt_host);
loop {
match poll_message_once(
data.clone(),
task.to_string(),
message_id.to_string(),
ctx.clone(),
)
.await?
{
MessagePoll::Resolved(payload) => return Ok(WaitOutcome::Terminal(payload)),
MessagePoll::Drove => continue,
MessagePoll::Idle => {}
}
if deadline.is_some_and(|deadline| Instant::now() >= deadline) {
return Ok(WaitOutcome::TimedOut {
attached: is_attached(data, task).await?,
});
}
time::sleep(POLL).await;
}
}
enum MessagePoll {
Resolved(Value),
Drove,
Idle,
}
async fn poll_message_once(
data: DataDir,
task: String,
message_id: String,
ctx: DriveContext,
) -> Result<MessagePoll, String> {
let handle = tokio::runtime::Handle::current();
tokio::task::spawn_blocking(move || -> Result<MessagePoll, String> {
let paths = resolve(&data, &task)?;
let messages = inbox::load(&paths)?;
let terminal = read_terminal(&paths)?;
if let Some(payload) =
inbox::message_payload_for_dispatch(&task, &messages, &message_id, terminal.as_ref())?
{
return Ok(MessagePoll::Resolved(payload));
}
if terminal.is_none()
&& let Some(guard) = try_attach(&paths)?
{
let _ = handle.block_on(drive(&data, &task, guard, &ctx))?;
return Ok(MessagePoll::Drove);
}
Ok(MessagePoll::Idle)
})
.await
.unwrap_or_else(|error| Err(format!("poll message: {error}")))
}
pub(crate) fn resolve(data: &DataDir, task: &str) -> Result<AgentPaths, String> {
data.agent_dir(task)
.filter(AgentPaths::exists)
.ok_or_else(|| format!("no task directory for {task}"))
}
pub(crate) fn try_attach(paths: &AgentPaths) -> Result<Option<lock::Lock>, String> {
lock::try_exclusive(&paths.attach_lock())
.map_err(|error| format!("acquire task attach lock: {error}"))
}
pub(crate) fn cancel_tree(data: &DataDir, task: &str) -> Result<(), String> {
let mut queue = vec![task.to_string()];
let mut visited = 0_usize;
while let Some(current) = queue.pop() {
visited += 1;
if visited > MAX_TASKS {
break;
}
let Some(paths) = data.agent_dir(¤t).filter(AgentPaths::exists) else {
continue;
};
if read_terminal(&paths)?.is_some() {
continue;
}
if !cancel_requested(&paths) {
request_cancel(&paths, Some(task))?;
}
queue.extend(children_of(data, ¤t)?);
}
Ok(())
}
fn children_of(data: &DataDir, task: &str) -> Result<Vec<String>, String> {
let Some((key, _)) = valid_task_handle(task) else {
return Ok(Vec::new());
};
let agents = data.agents_dir(key);
let entries = match std::fs::read_dir(&agents) {
Ok(entries) => entries,
Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(Vec::new()),
Err(error) => return Err(format!("scan workspace agents: {error}")),
};
let mut children = Vec::new();
for entry in entries {
let entry = entry.map_err(|error| format!("scan workspace agents: {error}"))?;
let id = entry.file_name().to_string_lossy().into_owned();
let handle = format!("{key}/{id}");
let Some(paths) = data.agent_dir(&handle) else {
continue;
};
let Ok(meta) = load_meta(&paths) else {
continue;
};
if !meta.detached && meta.parent.as_deref() == Some(task) {
children.push(handle);
}
}
Ok(children)
}
pub(crate) async fn drive(
data: &DataDir,
task: &str,
mut guard: lock::Lock,
ctx: &DriveContext,
) -> Result<Option<Value>, String> {
let paths = resolve(data, task)?;
if let Some(terminal) = read_terminal(&paths)? {
return Ok(Some(terminal));
}
guard.write_fingerprint();
let mut meta = load_meta(&paths)?;
if meta.pending_terminal.is_none() {
match existing_conversation(&meta) {
Some(agent_id) => {
let (key, _) = valid_task_handle(task)
.ok_or_else(|| format!("malformed task handle {task}"))?;
match try_conversation(data, key, &agent_id)? {
Some(_conversation) => {
run_model(data, task, &paths, &mut meta, ctx).await?;
}
None => return Ok(None),
}
}
None => run_model(data, task, &paths, &mut meta, ctx).await?,
}
}
Ok(Some(settle(data, &paths, &mut meta, ctx).await?))
}
fn existing_conversation(meta: &TaskMeta) -> Option<String> {
if meta.agent_id.is_empty() {
meta.continues.clone()
} else {
Some(meta.agent_id.clone())
}
}
fn try_conversation(
data: &DataDir,
key: &str,
agent_id: &str,
) -> Result<Option<lock::Lock>, String> {
let path = data
.conversation_lock(key, agent_id)
.map_err(|error| format!("prepare conversation lock: {error}"))?;
lock::try_exclusive(&path).map_err(|error| format!("acquire conversation lock: {error}"))
}
async fn run_model(
data: &DataDir,
task: &str,
paths: &AgentPaths,
meta: &mut TaskMeta,
ctx: &DriveContext,
) -> Result<(), String> {
if meta.deadline_passed() {
return record_pending(
paths,
meta,
PendingTerminal::Failed {
error: "task deadline elapsed before the next turn".to_string(),
},
Some(Bound::Deadline),
);
}
if cancel_requested(paths) {
return record_pending(paths, meta, PendingTerminal::Cancelled, None);
}
inbox::revert_in_flight(paths)?;
let events = match EventLog::open(paths) {
Ok(log) => Arc::new(Mutex::new(log)),
Err(error) => {
return record_failure(
paths,
meta,
format!("open task event journal: {error}"),
None,
);
}
};
let runtime = match task_runtime(data, task, meta) {
Ok(runtime) => runtime,
Err(error) => return record_failure(paths, meta, error, None),
};
let (builder, spec) = run_parts(meta);
let workspace = match builder.with_runtime_builder(runtime).open().await {
Ok(workspace) => Arc::new(workspace),
Err(error) => return record_failure(paths, meta, error.to_string(), None),
};
let reattached = !meta.agent_id.is_empty();
let existing = existing_conversation(meta);
let prepared = match existing.as_deref() {
Some(agent_id) => workspace.resume(agent_id, spec),
None => workspace.prepare(spec),
};
let mut run = match prepared {
Ok(run) => run.with_workspace(workspace),
Err(error) => return record_failure(paths, meta, error.to_string(), None),
};
if !reattached {
meta.agent_id = run.agent_id().to_string();
meta.answered_before = run.answered_turns();
meta.updated_ms = now_ms();
save_meta(paths, meta)?;
}
let mut initial_done = false;
let mut last_result = String::new();
let mut last_stopped_by: Option<Bound> = None;
if reattached {
initial_done = run.answered_turns() > meta.answered_before;
if let Some(message) = run
.history()
.iter()
.rev()
.find(|message| matches!(message.role, mentra::Role::Assistant))
{
last_result = message.text();
}
}
let cancellation = CancellationToken::default();
loop {
if cancel_requested(paths) {
return record_pending(paths, meta, PendingTerminal::Cancelled, None);
}
let remaining = remaining_deadline(meta.deadline_at_ms);
if remaining.as_ref().is_some_and(Duration::is_zero) {
return record_pending(
paths,
meta,
PendingTerminal::Failed {
error: "task deadline elapsed before the next turn".to_string(),
},
Some(Bound::Deadline),
);
}
let message = if initial_done {
inbox::start_next(paths)?
} else {
None
};
if initial_done && message.is_none() {
let (result, truncated) = bounded_text(last_result, MAX_RESULT_BYTES);
meta.result_truncated = truncated;
return record_pending(
paths,
meta,
PendingTerminal::Succeeded { result },
last_stopped_by,
);
}
let mut turn = TurnOptions::default().with_cancel(cancellation.clone());
if let Some(remaining) = remaining {
turn = turn.with_deadline(remaining);
}
let approver = match approver(meta.options.approve, ctx) {
Ok(approver) => approver,
Err(error) => return record_failure(paths, meta, error.to_string(), None),
};
let sink = FileSink {
log: Arc::clone(&events),
ctx: ctx.clone(),
};
let completed_message = message.as_ref().map(|(id, _)| id.clone());
let execution = async {
match message {
Some((_, body)) => run.send_with_options(body, sink, approver, turn).await,
None => {
run.execute_with_approver_and_options(sink, approver, turn)
.await
}
}
};
let report = match remaining {
Some(remaining) => match time::timeout(remaining, execution).await {
Ok(report) => report,
Err(_) => {
cancellation.cancel();
return record_pending(
paths,
meta,
PendingTerminal::Failed {
error: "task deadline elapsed during the turn".to_string(),
},
Some(Bound::Deadline),
);
}
},
None => execution.await,
};
let report = match report {
Ok(report) => report,
Err(error) => return record_failure(paths, meta, error.to_string(), None),
};
meta.usage = meta.usage.plus(report.usage);
meta.updated_ms = now_ms();
save_meta(paths, meta)?;
let stopped_by = report.stopped_by;
match report.outcome {
RunOutcome::Error { message } => {
return if cancel_requested(paths) {
record_pending(paths, meta, PendingTerminal::Cancelled, None)
} else {
record_failure(paths, meta, message, stopped_by)
};
}
RunOutcome::Ok => {
let result = report.final_message.unwrap_or_default();
if let Some(id) = completed_message {
let (reply, result_truncated) = bounded_text(result.clone(), MAX_RESULT_BYTES);
inbox::finish(
paths,
&id,
Some(MessageReply {
result: reply,
result_truncated,
stopped_by,
}),
)?;
}
initial_done = true;
last_result = result;
last_stopped_by = stopped_by;
}
outcome => {
return record_failure(
paths,
meta,
format!("unrecognized run outcome: {outcome:?}"),
stopped_by,
);
}
}
}
}
fn record_pending(
paths: &AgentPaths,
meta: &mut TaskMeta,
completion: PendingTerminal,
stopped_by: Option<Bound>,
) -> Result<(), String> {
if matches!(completion, PendingTerminal::Cancelled) {
meta.result_truncated = false;
meta.stopped_by = None;
} else {
meta.stopped_by = stopped_by;
}
meta.pending_terminal = Some(completion);
meta.updated_ms = now_ms();
save_meta(paths, meta)
}
fn record_failure(
paths: &AgentPaths,
meta: &mut TaskMeta,
message: String,
stopped_by: Option<Bound>,
) -> Result<(), String> {
let (error, _) = bounded_text(message, MAX_RESULT_BYTES);
meta.result_truncated = false;
record_pending(paths, meta, PendingTerminal::Failed { error }, stopped_by)
}
async fn settle(
data: &DataDir,
paths: &AgentPaths,
meta: &mut TaskMeta,
ctx: &DriveContext,
) -> Result<Value, String> {
reconsider_cancel(paths, meta)?;
let cancel_children = !matches!(
meta.pending_terminal,
Some(PendingTerminal::Succeeded { .. })
);
settle_children(data, meta, cancel_children, ctx).await?;
reconsider_cancel(paths, meta)?;
let payload = meta
.terminal_payload()
.expect("a completion was recorded before settling");
let _inbox_lock = inbox::finish_unanswered_durably(paths)?;
write_terminal(paths, &payload)?;
Ok(payload)
}
fn reconsider_cancel(paths: &AgentPaths, meta: &mut TaskMeta) -> Result<(), String> {
if cancel_requested(paths) && !matches!(meta.pending_terminal, Some(PendingTerminal::Cancelled))
{
record_pending(paths, meta, PendingTerminal::Cancelled, None)?;
}
Ok(())
}
async fn settle_children(
data: &DataDir,
meta: &TaskMeta,
cancel_children: bool,
ctx: &DriveContext,
) -> Result<(), String> {
loop {
let mut unfinished = Vec::new();
for child in children_of(data, &meta.id)? {
let Some(paths) = data.agent_dir(&child).filter(AgentPaths::exists) else {
continue;
};
if read_terminal(&paths)?.is_none() {
unfinished.push((child, paths));
}
}
if unfinished.is_empty() {
return Ok(());
}
let cancel = cancel_children || meta.deadline_passed();
let mut remaining = false;
for (child, paths) in unfinished {
if cancel && !cancel_requested(&paths) {
request_cancel(&paths, Some(&meta.id))?;
}
match try_attach(&paths)? {
Some(guard) => {
if Box::pin(drive(data, &child, guard, &ctx.hidden()))
.await?
.is_none()
{
remaining = true;
}
}
None => remaining = true,
}
}
if remaining {
time::sleep(POLL).await;
}
}
}
fn run_parts(meta: &TaskMeta) -> (WorkspaceBuilder, RunSpec) {
let options = &meta.options;
let mut builder = Workspace::builder(Path::new(&meta.workspace))
.with_shell(ShellAccess::from_flag(!options.no_shell));
if let Some(model) = &options.model {
builder = builder.with_model(ModelSelector::Id(model.clone()));
}
if let Some(system_prompt) = options.system_prompt.clone() {
builder = builder.with_system_prompt(system_prompt);
}
let mut spec = RunSpec::new(meta.prompt.clone());
if let Some(effort) = options.effort {
spec = spec.with_effort(effort);
}
if let Some(remaining) = remaining_deadline(meta.deadline_at_ms) {
spec = spec.with_deadline(remaining.max(Duration::from_millis(1)));
}
if let Some(tool_budget) = options.tool_budget {
spec = spec.with_tool_budget(tool_budget);
}
if let Some(token_budget) = options.token_budget {
spec = spec.with_token_budget(token_budget);
}
(builder, spec)
}
fn task_runtime(data: &DataDir, task: &str, meta: &TaskMeta) -> Result<RuntimeBuilder, String> {
let (key, _) =
valid_task_handle(task).ok_or_else(|| format!("malformed task handle {task}"))?;
let mut runtime = Runtime::builder()
.with_store_dir(data.store_dir(key))
.with_command_environment(crate::BASIS_TASK_ID, task)
.with_command_environment(crate::BASIS_DATA_DIR, data.root().to_string_lossy());
if let Some(name) = &meta.options.provider {
runtime = runtime.with_provider(provider::parse(name).map_err(|error| error.to_string())?);
}
if let Some(base_url) = &meta.options.base_url {
runtime = runtime.with_base_url(base_url);
}
if let Some(parent) = &meta.parent {
runtime = runtime.with_command_environment(crate::BASIS_PARENT_TASK_ID, parent);
}
Ok(runtime)
}
fn approver(mode: Approve, ctx: &DriveContext) -> Result<Box<dyn Approver>, Error> {
crate::approve::validate_approval(mode, ctx.can_ask())?;
Ok(match mode {
Approve::Always => Box::new(AllowAll),
Approve::Never => Box::new(DenyAll),
Approve::Prompt => ctx
.approver()
.expect("validate confirmed a host that can ask"),
})
}
pub(crate) fn earlier_deadline(left: Option<u64>, right: Option<u64>) -> Option<u64> {
match (left, right) {
(Some(left), Some(right)) => Some(left.min(right)),
(Some(value), None) | (None, Some(value)) => Some(value),
(None, None) => None,
}
}
fn remaining_deadline(deadline_at: Option<u64>) -> Option<Duration> {
deadline_at.map(|deadline| Duration::from_millis(deadline.saturating_sub(now_ms())))
}
struct FileSink {
log: Arc<Mutex<EventLog>>,
ctx: DriveContext,
}
impl EventSink for FileSink {
fn emit(&mut self, event: Event) -> io::Result<()> {
if let Ok(value) = serde_json::to_value(event) {
self.ctx.show(&value);
if let Ok(mut log) = self.log.lock() {
let _ = log.append(value);
}
}
Ok(())
}
}
#[cfg(test)]
mod tests;