use std::collections::{BTreeMap, HashMap, HashSet, VecDeque};
use std::sync::Arc;
use agentkit_core::{
CancellationHandle, DataRef, Delta, FinishReason, Item, ItemKind, MetadataMap, Modality, Part,
SessionId, TaskId, TextPart, Timestamp, ToolCallId, ToolCallPart, ToolOutput, ToolResultPart,
TurnCancellation, Usage,
};
use agentkit_task_manager::{
PendingLoopUpdates, SimpleTaskManager, TOOL_RESULT_NOT_STARTED_METADATA_KEY, TaskApproval,
TaskLaunchKind, TaskLaunchRequest, TaskManager, TaskResolution, TaskStartContext,
TaskStartOutcome, TurnTaskUpdate,
};
#[cfg(test)]
use agentkit_task_manager::{
TOOL_RESULT_FAILURE_KIND_METADATA_KEY, TOOL_RESULT_FAILURE_KIND_PERMISSION_DENIED,
};
#[cfg(test)]
use agentkit_tools_core::ToolContext;
use agentkit_tools_core::{
AllowAllPermissions, ApprovalDecision, ApprovalRequest, BasicToolExecutor, OwnedToolContext,
PermissionChecker, ToolCatalogEvent, ToolError, ToolExecutionScope, ToolExecutor, ToolRequest,
ToolResources, ToolSource, ToolSpec,
};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use thiserror::Error;
const INTERRUPTED_METADATA_KEY: &str = "agentkit.interrupted";
const INTERRUPT_REASON_METADATA_KEY: &str = "agentkit.interrupt_reason";
const INTERRUPT_STAGE_METADATA_KEY: &str = "agentkit.interrupt_stage";
const USER_CANCELLED_REASON: &str = "user_cancelled";
const DETACHED_NOTIFICATION_TEXT_MAX_CHARS: usize = 512;
const DETACHED_TEXT_PREVIEW_MAX_CHARS: usize = 160;
const DETACHED_CALL_ID_MAX_CHARS: usize = 80;
pub const PROVIDER_FINISH_REASONS_METADATA_KEY: &str = "agentkit.provider_finish_reasons";
pub fn set_provider_finish_reasons<I, S>(metadata: &mut MetadataMap, reasons: I)
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
let mut seen = HashSet::new();
let reasons = reasons
.into_iter()
.map(Into::into)
.filter(|reason: &String| !reason.is_empty() && seen.insert(reason.clone()))
.map(Value::String)
.collect::<Vec<_>>();
metadata.remove(PROVIDER_FINISH_REASONS_METADATA_KEY);
if !reasons.is_empty() {
metadata.insert(
PROVIDER_FINISH_REASONS_METADATA_KEY.into(),
Value::Array(reasons),
);
}
}
fn provider_finish_reasons(metadata: &MetadataMap, fallback: &FinishReason) -> Vec<String> {
metadata
.get(PROVIDER_FINISH_REASONS_METADATA_KEY)
.and_then(Value::as_array)
.map(|values| {
let mut seen = HashSet::new();
values
.iter()
.filter_map(Value::as_str)
.filter(|reason| !reason.is_empty() && seen.insert((*reason).to_owned()))
.map(str::to_owned)
.collect::<Vec<_>>()
})
.filter(|reasons| !reasons.is_empty())
.unwrap_or_else(|| vec![normalized_finish_reason(fallback).into()])
}
fn normalized_finish_reason(reason: &FinishReason) -> &str {
match reason {
FinishReason::Completed => "completed",
FinishReason::ToolCall => "tool_call",
FinishReason::MaxTokens => "max_tokens",
FinishReason::Cancelled => "cancelled",
FinishReason::Blocked => "blocked",
FinishReason::Error => "error",
FinishReason::Other(reason) => reason,
}
}
#[derive(Clone, Copy, Debug, Error, PartialEq, Eq)]
pub enum MessageCaptureError {
#[error("message capture max_messages must be nonzero")]
ZeroMessages,
#[error("message capture max_bytes must be nonzero")]
ZeroBytes,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct MessageCapture {
max_messages: usize,
max_bytes: usize,
}
impl MessageCapture {
pub fn new(max_messages: usize, max_bytes: usize) -> Result<Self, MessageCaptureError> {
if max_messages == 0 {
return Err(MessageCaptureError::ZeroMessages);
}
if max_bytes == 0 {
return Err(MessageCaptureError::ZeroBytes);
}
Ok(Self {
max_messages,
max_bytes,
})
}
pub fn max_messages(self) -> usize {
self.max_messages
}
pub fn max_bytes(self) -> usize {
self.max_bytes
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct TelemetryConfig {
input_messages: Option<MessageCapture>,
output_messages: Option<MessageCapture>,
}
impl TelemetryConfig {
pub fn with_input_messages(mut self, capture: MessageCapture) -> Self {
self.input_messages = Some(capture);
self
}
pub fn with_output_messages(mut self, capture: MessageCapture) -> Self {
self.output_messages = Some(capture);
self
}
pub fn without_input_messages(mut self) -> Self {
self.input_messages = None;
self
}
pub fn without_output_messages(mut self) -> Self {
self.output_messages = None;
self
}
pub fn input_messages(self) -> Option<MessageCapture> {
self.input_messages
}
pub fn output_messages(self) -> Option<MessageCapture> {
self.output_messages
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct SessionConfig {
pub session_id: SessionId,
pub metadata: MetadataMap,
pub cache: Option<PromptCacheRequest>,
}
impl SessionConfig {
pub fn new(session_id: impl Into<SessionId>) -> Self {
Self {
session_id: session_id.into(),
metadata: MetadataMap::new(),
cache: None,
}
}
pub fn with_metadata(mut self, metadata: MetadataMap) -> Self {
self.metadata = metadata;
self
}
pub fn with_cache(mut self, cache: PromptCacheRequest) -> Self {
self.cache = Some(cache);
self
}
pub fn without_cache(mut self) -> Self {
self.cache = None;
self
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum PromptCacheMode {
Disabled,
#[default]
BestEffort,
Required,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum PromptCacheRetention {
Default,
Short,
Extended,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum PromptCacheStrategy {
#[default]
Automatic,
Explicit {
breakpoints: Vec<PromptCacheBreakpoint>,
},
}
impl PromptCacheStrategy {
pub fn automatic() -> Self {
Self::Automatic
}
pub fn explicit(breakpoints: impl IntoIterator<Item = PromptCacheBreakpoint>) -> Self {
Self::Explicit {
breakpoints: breakpoints.into_iter().collect(),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum PromptCacheBreakpoint {
ToolsEnd,
TranscriptItemEnd { index: usize },
TranscriptPartEnd {
item_index: usize,
part_index: usize,
},
}
impl PromptCacheBreakpoint {
pub fn tools_end() -> Self {
Self::ToolsEnd
}
pub fn transcript_item_end(index: usize) -> Self {
Self::TranscriptItemEnd { index }
}
pub fn transcript_part_end(item_index: usize, part_index: usize) -> Self {
Self::TranscriptPartEnd {
item_index,
part_index,
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct PromptCacheRequest {
pub mode: PromptCacheMode,
pub strategy: PromptCacheStrategy,
pub retention: Option<PromptCacheRetention>,
pub key: Option<String>,
}
impl PromptCacheRequest {
pub fn automatic() -> Self {
Self::best_effort(PromptCacheStrategy::automatic())
}
pub fn automatic_required() -> Self {
Self::required(PromptCacheStrategy::automatic())
}
pub fn explicit(breakpoints: impl IntoIterator<Item = PromptCacheBreakpoint>) -> Self {
Self::best_effort(PromptCacheStrategy::explicit(breakpoints))
}
pub fn explicit_required(breakpoints: impl IntoIterator<Item = PromptCacheBreakpoint>) -> Self {
Self::required(PromptCacheStrategy::explicit(breakpoints))
}
pub fn disabled() -> Self {
Self {
mode: PromptCacheMode::Disabled,
strategy: PromptCacheStrategy::Automatic,
retention: None,
key: None,
}
}
pub fn best_effort(strategy: PromptCacheStrategy) -> Self {
Self {
mode: PromptCacheMode::BestEffort,
strategy,
retention: None,
key: None,
}
}
pub fn required(strategy: PromptCacheStrategy) -> Self {
Self {
mode: PromptCacheMode::Required,
strategy,
retention: None,
key: None,
}
}
pub fn with_mode(mut self, mode: PromptCacheMode) -> Self {
self.mode = mode;
self
}
pub fn with_strategy(mut self, strategy: PromptCacheStrategy) -> Self {
self.strategy = strategy;
self
}
pub fn with_retention(mut self, retention: PromptCacheRetention) -> Self {
self.retention = Some(retention);
self
}
pub fn with_key(mut self, key: impl Into<String>) -> Self {
self.key = Some(key.into());
self
}
pub fn without_retention(mut self) -> Self {
self.retention = None;
self
}
pub fn without_key(mut self) -> Self {
self.key = None;
self
}
pub fn is_enabled(&self) -> bool {
!matches!(self.mode, PromptCacheMode::Disabled)
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct TurnRequest {
pub session_id: SessionId,
pub turn_id: agentkit_core::TurnId,
pub transcript: Vec<Item>,
pub available_tools: Vec<ToolSpec>,
pub cache: Option<PromptCacheRequest>,
pub metadata: MetadataMap,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ModelTurnResult {
pub finish_reason: FinishReason,
pub output_items: Vec<Item>,
pub usage: Option<Usage>,
pub metadata: MetadataMap,
#[serde(default)]
pub model: Option<String>,
#[serde(default)]
pub response_id: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum ModelTurnEvent {
Delta(Delta),
ToolCall(ToolCallPart),
Usage(Usage),
Finished(ModelTurnResult),
}
#[async_trait]
pub trait ModelAdapter: Send + Sync {
type Session: ModelSession;
async fn start_session(&self, config: SessionConfig) -> Result<Self::Session, LoopError>;
fn provider_name(&self) -> Option<&str> {
None
}
}
#[async_trait]
pub trait ModelSession: Send {
type Turn: ModelTurn;
async fn begin_turn(
&mut self,
request: TurnRequest,
cancellation: Option<TurnCancellation>,
) -> Result<Self::Turn, LoopError>;
fn model_name(&self) -> Option<&str> {
None
}
fn provider_name(&self) -> Option<&str> {
None
}
}
#[async_trait]
pub trait ModelTurn: Send {
async fn next_event(
&mut self,
cancellation: Option<TurnCancellation>,
) -> Result<Option<ModelTurnEvent>, LoopError>;
}
pub trait LoopObserver: Send + Sync {
fn handle_event(&self, event: ObservedEvent);
}
#[derive(Clone, Debug, PartialEq)]
pub struct ObservedEvent {
pub session_id: Arc<SessionId>,
pub event: AgentEvent,
}
pub trait TranscriptObserver: Send + Sync {
fn on_transcript_event(&self, event: TranscriptEvent<'_>);
}
#[derive(Clone, Debug)]
pub struct TranscriptEvent<'a> {
pub session_id: &'a SessionId,
pub item: &'a Item,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum MutationPoint {
AfterToolResult,
AfterTurnEnded,
}
pub trait EventEmitter: Send + Sync {
fn emit(&self, event: AgentEvent);
}
#[non_exhaustive]
pub struct LoopCtx<'a> {
pub session_id: &'a SessionId,
pub turn_id: Option<&'a agentkit_core::TurnId>,
pub point: MutationPoint,
pub cancellation: Option<TurnCancellation>,
pub emitter: &'a dyn EventEmitter,
}
pub struct TranscriptCursor<'a> {
items: &'a mut Vec<Item>,
pub(crate) dirty: bool,
}
impl<'a> std::ops::Deref for TranscriptCursor<'a> {
type Target = Vec<Item>;
fn deref(&self) -> &Vec<Item> {
self.items
}
}
impl<'a> std::ops::DerefMut for TranscriptCursor<'a> {
fn deref_mut(&mut self) -> &mut Vec<Item> {
self.dirty = true;
self.items
}
}
#[async_trait]
pub trait LoopMutator: Send + Sync {
async fn mutate(
&self,
cursor: &mut TranscriptCursor<'_>,
ctx: LoopCtx<'_>,
) -> Result<(), LoopError> {
let _ = (cursor, ctx);
Ok(())
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum AgentEvent {
RunStarted { session_id: SessionId },
TurnStarted {
session_id: SessionId,
turn_id: agentkit_core::TurnId,
},
InputAccepted {
session_id: SessionId,
items: Vec<Item>,
},
ContentDelta(Delta),
ToolCallRequested(ToolCallPart),
ToolExecutionStarted(ToolCallPart),
ToolExecutionProgress(ToolResultPart),
ToolResultReceived(ToolResultPart),
ApprovalRequired(ApprovalRequest),
ApprovalResolved { approved: bool },
ToolCatalogChanged(ToolCatalogEvent),
MutationStarted {
session_id: SessionId,
turn_id: Option<agentkit_core::TurnId>,
mutator: String,
point: MutationPoint,
},
MutationFinished {
session_id: SessionId,
turn_id: Option<agentkit_core::TurnId>,
mutator: String,
dirty: bool,
metadata: MetadataMap,
},
UsageUpdated(Usage),
Warning { message: String },
RunFailed { message: String },
TurnFinished(TurnResult),
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct PendingApproval {
pub request: ApprovalRequest,
}
impl std::ops::Deref for PendingApproval {
type Target = ApprovalRequest;
fn deref(&self) -> &ApprovalRequest {
&self.request
}
}
impl PendingApproval {
pub fn approve<S: ModelSession>(self, driver: &mut LoopDriver<S>) -> Result<(), LoopError> {
let call_id = self
.request
.call_id
.ok_or_else(|| LoopError::InvalidState("pending approval is missing call id".into()))?;
driver.resolve_approval_for(call_id, ApprovalDecision::Approve)
}
pub fn deny<S: ModelSession>(self, driver: &mut LoopDriver<S>) -> Result<(), LoopError> {
let call_id = self
.request
.call_id
.ok_or_else(|| LoopError::InvalidState("pending approval is missing call id".into()))?;
driver.resolve_approval_for(call_id, ApprovalDecision::Deny { reason: None })
}
pub fn deny_with_reason<S: ModelSession>(
self,
driver: &mut LoopDriver<S>,
reason: impl Into<String>,
) -> Result<(), LoopError> {
let call_id = self
.request
.call_id
.ok_or_else(|| LoopError::InvalidState("pending approval is missing call id".into()))?;
driver.resolve_approval_for(
call_id,
ApprovalDecision::Deny {
reason: Some(reason.into()),
},
)
}
pub fn approve_with_patched_input<S: ModelSession>(
self,
driver: &mut LoopDriver<S>,
input: serde_json::Value,
) -> Result<(), LoopError> {
let call_id = self
.request
.call_id
.ok_or_else(|| LoopError::InvalidState("pending approval is missing call id".into()))?;
driver.resolve_approval_for_with_patched_input(call_id, input)
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct InputRequest {
pub session_id: SessionId,
pub reason: String,
}
impl InputRequest {
pub fn submit<S: ModelSession>(
self,
driver: &mut LoopDriver<S>,
items: Vec<Item>,
) -> Result<(), LoopError> {
driver.submit_input(items)
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct TurnResult {
pub turn_id: agentkit_core::TurnId,
pub finish_reason: FinishReason,
pub items: Vec<Item>,
pub usage: Option<Usage>,
pub metadata: MetadataMap,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum LoopInterrupt {
ApprovalRequest(PendingApproval),
AwaitingInput(InputRequest),
AfterToolResult(ToolRoundInfo),
}
impl LoopInterrupt {
pub fn is_blocking(&self) -> bool {
matches!(self, LoopInterrupt::ApprovalRequest(_))
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ToolRoundInfo {
pub session_id: SessionId,
pub turn_id: agentkit_core::TurnId,
pub transcript_len: usize,
}
impl ToolRoundInfo {
pub fn submit<S: ModelSession>(
self,
driver: &mut LoopDriver<S>,
items: Vec<Item>,
) -> Result<(), LoopError> {
driver.submit_input(items)
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum LoopStep {
Interrupt(LoopInterrupt),
Finished(TurnResult),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct LoopSnapshot {
pub session_id: SessionId,
pub transcript: Vec<Item>,
pub pending_input: Vec<Item>,
}
#[derive(Clone)]
struct PendingApprovalToolCall {
request: ApprovalRequest,
decision: Option<ApprovalDecision>,
surfaced: bool,
presentation_turn_id: agentkit_core::TurnId,
task_id: TaskId,
call: ToolCallPart,
tool_request: ToolRequest,
cancellation: Option<TurnCancellation>,
}
#[derive(Clone, Default)]
struct ActiveToolRound {
presentation_turn_id: agentkit_core::TurnId,
task_turn_id: agentkit_core::TurnId,
pending_calls: VecDeque<(ToolCallPart, ToolRequest)>,
cancellation: Option<TurnCancellation>,
background_pending: bool,
foreground_progressed: bool,
}
#[derive(Default)]
struct DriverLifecycle {
active_turn: Option<agentkit_core::TurnId>,
}
pub struct Agent<M>
where
M: ModelAdapter,
{
model: M,
tool_sources: Vec<Arc<dyn ToolSource>>,
tool_executor: Option<Arc<dyn ToolExecutor>>,
task_manager: Arc<dyn TaskManager>,
permissions: Arc<dyn PermissionChecker>,
resources: Arc<dyn ToolResources>,
cancellation: Option<CancellationHandle>,
mutators: Vec<Arc<dyn LoopMutator>>,
observers: Vec<Arc<dyn LoopObserver>>,
transcript_observers: Vec<Arc<dyn TranscriptObserver>>,
transcript: Vec<Item>,
input: Vec<Item>,
telemetry: TelemetryConfig,
}
impl<M> Agent<M>
where
M: ModelAdapter,
{
pub fn builder() -> AgentBuilder<M> {
AgentBuilder::default()
}
pub async fn start(&self, config: SessionConfig) -> Result<LoopDriver<M::Session>, LoopError> {
let session_id = config.session_id.clone();
let default_cache = config.cache.clone();
let session = self.model.start_session(config).await?;
let provider_name = self.model.provider_name().map(str::to_owned);
let tool_executor = self
.tool_executor
.clone()
.unwrap_or_else(|| Arc::new(BasicToolExecutor::new(self.tool_sources.clone())));
let driver = LoopDriver {
session_id: session_id.clone(),
observed_session_id: Arc::new(session_id.clone()),
provider_name,
telemetry: self.telemetry,
default_cache,
next_turn_cache: None,
session: Some(session),
tool_executor,
task_manager: self.task_manager.clone(),
permissions: self.permissions.clone(),
resources: self.resources.clone(),
cancellation: self.cancellation.clone(),
mutators: self.mutators.clone(),
observers: self.observers.clone(),
transcript_observers: self.transcript_observers.clone(),
transcript: self.transcript.clone(),
pending_input: self.input.clone(),
pending_approvals: BTreeMap::new(),
pending_approval_order: VecDeque::new(),
active_tool_round: None,
pending_round_resume: None,
pending_loop_updates: VecDeque::new(),
next_turn_index: 1,
lifecycle: DriverLifecycle::default(),
background_call_ids: HashSet::new(),
detached_call_ids: HashSet::new(),
interrupted_background_call_ids: HashSet::new(),
tool_cancellations: HashMap::new(),
};
driver.emit(AgentEvent::RunStarted { session_id });
Ok(driver)
}
}
pub struct AgentBuilder<M>
where
M: ModelAdapter,
{
model: Option<M>,
tool_sources: Vec<Arc<dyn ToolSource>>,
tool_executor: Option<Arc<dyn ToolExecutor>>,
task_manager: Option<Arc<dyn TaskManager>>,
permissions: Arc<dyn PermissionChecker>,
resources: Arc<dyn ToolResources>,
cancellation: Option<CancellationHandle>,
mutators: Vec<Arc<dyn LoopMutator>>,
observers: Vec<Arc<dyn LoopObserver>>,
transcript_observers: Vec<Arc<dyn TranscriptObserver>>,
transcript: Vec<Item>,
input: Vec<Item>,
telemetry: TelemetryConfig,
}
impl<M> Default for AgentBuilder<M>
where
M: ModelAdapter,
{
fn default() -> Self {
Self {
model: None,
tool_sources: Vec::new(),
tool_executor: None,
task_manager: None,
permissions: Arc::new(AllowAllPermissions),
resources: Arc::new(()),
cancellation: None,
mutators: Vec::new(),
observers: Vec::new(),
transcript_observers: Vec::new(),
transcript: Vec::new(),
input: Vec::new(),
telemetry: TelemetryConfig::default(),
}
}
}
impl<M> AgentBuilder<M>
where
M: ModelAdapter,
{
pub fn model(mut self, model: M) -> Self {
self.model = Some(model);
self
}
pub fn add_tool_source<S: ToolSource + 'static>(mut self, source: S) -> Self {
self.tool_sources.push(Arc::new(source));
self
}
pub fn tool_executor(mut self, executor: impl ToolExecutor + 'static) -> Self {
self.tool_executor = Some(Arc::new(executor));
self
}
pub fn task_manager(mut self, manager: impl TaskManager + 'static) -> Self {
self.task_manager = Some(Arc::new(manager));
self
}
pub fn permissions(mut self, permissions: impl PermissionChecker + 'static) -> Self {
self.permissions = Arc::new(permissions);
self
}
pub fn resources(mut self, resources: impl ToolResources + 'static) -> Self {
self.resources = Arc::new(resources);
self
}
pub fn cancellation(mut self, handle: CancellationHandle) -> Self {
self.cancellation = Some(handle);
self
}
pub fn mutator<L: LoopMutator + 'static>(mut self, mutator: L) -> Self {
self.mutators.push(Arc::new(mutator));
self
}
pub fn observer<O: LoopObserver + 'static>(mut self, observer: O) -> Self {
self.observers.push(Arc::new(observer));
self
}
pub fn transcript_observer<O: TranscriptObserver + 'static>(mut self, observer: O) -> Self {
self.transcript_observers.push(Arc::new(observer));
self
}
pub fn transcript(mut self, transcript: Vec<Item>) -> Self {
self.transcript = transcript;
self
}
pub fn input(mut self, input: Vec<Item>) -> Self {
self.input = input;
self
}
pub fn telemetry(mut self, telemetry: TelemetryConfig) -> Self {
self.telemetry = telemetry;
self
}
pub fn build(self) -> Result<Agent<M>, LoopError> {
let model = self
.model
.ok_or_else(|| LoopError::InvalidState("model adapter is required".into()))?;
Ok(Agent {
model,
tool_sources: self.tool_sources,
tool_executor: self.tool_executor,
task_manager: self
.task_manager
.unwrap_or_else(|| Arc::new(SimpleTaskManager::new())),
permissions: self.permissions,
resources: self.resources,
cancellation: self.cancellation,
mutators: self.mutators,
observers: self.observers,
transcript_observers: self.transcript_observers,
transcript: self.transcript,
input: self.input,
telemetry: self.telemetry,
})
}
}
pub struct LoopDriver<S>
where
S: ModelSession,
{
session_id: SessionId,
observed_session_id: Arc<SessionId>,
provider_name: Option<String>,
telemetry: TelemetryConfig,
default_cache: Option<PromptCacheRequest>,
next_turn_cache: Option<PromptCacheRequest>,
session: Option<S>,
tool_executor: Arc<dyn ToolExecutor>,
task_manager: Arc<dyn TaskManager>,
permissions: Arc<dyn PermissionChecker>,
resources: Arc<dyn ToolResources>,
cancellation: Option<CancellationHandle>,
mutators: Vec<Arc<dyn LoopMutator>>,
observers: Vec<Arc<dyn LoopObserver>>,
transcript_observers: Vec<Arc<dyn TranscriptObserver>>,
transcript: Vec<Item>,
pending_input: Vec<Item>,
pending_approvals: BTreeMap<ToolCallId, PendingApprovalToolCall>,
pending_approval_order: VecDeque<ToolCallId>,
active_tool_round: Option<ActiveToolRound>,
pending_round_resume: Option<agentkit_core::TurnId>,
pending_loop_updates: VecDeque<TaskResolution>,
next_turn_index: u64,
lifecycle: DriverLifecycle,
background_call_ids: HashSet<ToolCallId>,
detached_call_ids: HashSet<ToolCallId>,
interrupted_background_call_ids: HashSet<ToolCallId>,
tool_cancellations: HashMap<ToolCallId, TurnCancellation>,
}
impl<S> LoopDriver<S>
where
S: ModelSession,
{
fn execute_tool_span(
&self,
request: &ToolRequest,
turn_id: &agentkit_core::TurnId,
launch_kind: &'static str,
) -> tracing::Span {
tracing::info_span!(
"agent.execute_tool",
"otel.name" = %format!("execute_tool {}", request.tool_name),
"gen_ai.operation.name" = "execute_tool",
"gen_ai.tool.name" = %request.tool_name,
"gen_ai.tool.call.id" = %request.call_id,
"gen_ai.conversation.id" = %self.session_id,
"error.type" = tracing::field::Empty,
session.id = %self.session_id,
turn.id = %turn_id,
launch_kind = launch_kind,
)
}
fn start_task_via_manager(
&self,
task_id: Option<TaskId>,
tool_request: ToolRequest,
kind: TaskLaunchKind,
cancellation: Option<TurnCancellation>,
) -> impl std::future::Future<Output = Result<TaskStartOutcome, LoopError>> + Send + 'static
{
let task_manager = self.task_manager.clone();
let tool_executor = self.tool_executor.clone();
let permissions = self.permissions.clone();
let resources = self.resources.clone();
let session_id = self.session_id.clone();
let turn_id = tool_request.turn_id.clone();
let metadata = tool_request.metadata.clone();
async move {
task_manager
.start_task(
TaskLaunchRequest {
task_id,
request: tool_request.clone(),
kind,
},
TaskStartContext {
executor: tool_executor.clone(),
tool_context: {
let execution_scope = ToolExecutionScope {
executor: tool_executor,
session_id: session_id.clone(),
turn_id: turn_id.clone(),
permissions: permissions.clone(),
resources: resources.clone(),
cancellation: cancellation.clone(),
};
OwnedToolContext {
session_id,
turn_id,
metadata,
permissions,
resources,
cancellation,
execution_scope: Some(execution_scope),
approved_request: None,
}
},
},
)
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))
}
}
fn register_tool_cancellation(
&mut self,
call_id: &ToolCallId,
cancellation: Option<TurnCancellation>,
) {
if let Some(cancellation) = cancellation {
self.tool_cancellations
.insert(call_id.clone(), cancellation);
}
}
fn tool_cancellation_for(
&mut self,
call_id: &ToolCallId,
fallback: Option<TurnCancellation>,
) -> Option<TurnCancellation> {
self.tool_cancellations.get(call_id).cloned().or(fallback)
}
fn clear_tool_cancellation(&mut self, call_id: &ToolCallId) {
self.tool_cancellations.remove(call_id);
}
fn has_pending_interrupts(&self) -> bool {
!self.pending_approvals.is_empty()
}
fn start_logical_turn(&mut self) -> agentkit_core::TurnId {
if let Some(turn_id) = &self.lifecycle.active_turn {
return turn_id.clone();
}
let turn_id = agentkit_core::TurnId::new(format!("turn-{}", self.next_turn_index));
self.next_turn_index += 1;
self.start_logical_turn_with(turn_id)
}
fn start_logical_turn_with(&mut self, turn_id: agentkit_core::TurnId) -> agentkit_core::TurnId {
if let Some(active_turn) = &self.lifecycle.active_turn {
return active_turn.clone();
}
self.lifecycle.active_turn = Some(turn_id.clone());
self.emit(AgentEvent::TurnStarted {
session_id: self.session_id.clone(),
turn_id: turn_id.clone(),
});
turn_id
}
fn finish_logical_turn(&mut self, result: &TurnResult) {
if self.pending_round_resume.as_ref() == Some(&result.turn_id) {
self.pending_round_resume = None;
}
if self.lifecycle.active_turn.as_ref() == Some(&result.turn_id) {
self.lifecycle.active_turn = None;
self.emit(AgentEvent::TurnFinished(result.clone()));
}
}
fn emit_tool_catalog_events(&mut self, events: Vec<ToolCatalogEvent>) {
for event in events {
self.emit(AgentEvent::ToolCatalogChanged(event));
}
}
fn enqueue_pending_approval(
&mut self,
presentation_turn_id: &agentkit_core::TurnId,
task: TaskApproval,
cancellation: Option<TurnCancellation>,
) {
let call_id = task.tool_request.call_id.clone();
self.background_call_ids.remove(&call_id);
let cancellation = self.tool_cancellation_for(&call_id, cancellation);
let call = ToolCallPart {
id: call_id.clone(),
name: task.tool_request.tool_name.to_string(),
input: task.tool_request.input.clone(),
metadata: task.tool_request.metadata.clone(),
};
let mut request = task.approval;
request.call_id = Some(call_id.clone());
let pending = PendingApprovalToolCall {
request: request.clone(),
decision: None,
surfaced: false,
presentation_turn_id: presentation_turn_id.clone(),
task_id: task.task_id,
call,
tool_request: task.tool_request,
cancellation,
};
self.pending_approvals.insert(call_id.clone(), pending);
if !self.pending_approval_order.iter().any(|id| id == &call_id) {
self.pending_approval_order.push_back(call_id);
}
self.emit(AgentEvent::ApprovalRequired(request));
}
fn take_next_unsurfaced_approval_interrupt(&mut self) -> Option<LoopStep> {
for call_id in self.pending_approval_order.clone() {
let Some(pending) = self.pending_approvals.get_mut(&call_id) else {
continue;
};
if pending.decision.is_none() && !pending.surfaced {
pending.surfaced = true;
return Some(LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(
PendingApproval {
request: pending.request.clone(),
},
)));
}
}
None
}
fn next_unresolved_approval_interrupt(&self) -> Option<LoopStep> {
self.pending_approval_order.iter().find_map(|call_id| {
self.pending_approvals.get(call_id).and_then(|pending| {
pending.decision.is_none().then(|| {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(PendingApproval {
request: pending.request.clone(),
}))
})
})
})
}
fn take_next_resolved_approval(&mut self) -> Option<PendingApprovalToolCall> {
let call_id = self.pending_approval_order.iter().find_map(|call_id| {
self.pending_approvals
.get(call_id)
.and_then(|pending| pending.decision.as_ref().map(|_| call_id.clone()))
})?;
self.pending_approval_order.retain(|id| id != &call_id);
self.pending_approvals.remove(&call_id)
}
fn queue_resolution_interrupt(
&mut self,
presentation_turn_id: &agentkit_core::TurnId,
resolution: TaskResolution,
cancellation: Option<TurnCancellation>,
) -> Option<LoopStep> {
match resolution {
TaskResolution::Item(item) => {
self.append_tool_result_item(item);
None
}
TaskResolution::Approval(task) => {
self.enqueue_pending_approval(presentation_turn_id, task, cancellation);
self.take_next_unsurfaced_approval_interrupt()
}
}
}
async fn collect_pending_loop_updates(&mut self) -> Result<(), LoopError> {
let PendingLoopUpdates { resolutions } = self
.task_manager
.take_pending_loop_updates()
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?;
self.pending_loop_updates.extend(resolutions);
Ok(())
}
async fn drain_pending_loop_updates(&mut self) -> Result<(bool, Option<LoopStep>), LoopError> {
self.collect_pending_loop_updates().await?;
let mut resolutions = std::mem::take(&mut self.pending_loop_updates);
if !resolutions.is_empty() {
self.start_logical_turn();
}
let mut saw_items = false;
while let Some(resolution) = resolutions.pop_front() {
match resolution {
TaskResolution::Item(item) => {
self.append_tool_result_item(item);
saw_items = true;
}
TaskResolution::Approval(task) => {
let turn_id = self.start_logical_turn();
self.enqueue_pending_approval(&turn_id, task, None);
}
}
}
if let Some(step) = self.finish_cancelled_pending_approval().await? {
return Ok((saw_items, Some(step)));
}
Ok((saw_items, self.take_next_unsurfaced_approval_interrupt()))
}
async fn finish_cancelled_pending_approval(&mut self) -> Result<Option<LoopStep>, LoopError> {
if self.pending_approvals.is_empty() {
return Ok(None);
}
if !self.pending_approvals.values().any(|pending| {
pending
.cancellation
.as_ref()
.is_some_and(TurnCancellation::is_cancelled)
}) {
return Ok(None);
}
self.cancel_pending_approvals().await
}
async fn run_mutators(
&mut self,
point: MutationPoint,
turn_id: Option<&agentkit_core::TurnId>,
cancellation: Option<TurnCancellation>,
) -> Result<(), LoopError> {
if self.mutators.is_empty() {
return Ok(());
}
if cancellation
.as_ref()
.is_some_and(TurnCancellation::is_cancelled)
{
return Err(LoopError::Cancelled);
}
let mutators = self.mutators.clone();
let session_id = self.session_id.clone();
let observed_session_id = Arc::clone(&self.observed_session_id);
let observers = self.observers.clone();
let emitter = DriverEmitter {
session_id: &observed_session_id,
observers: &observers,
};
let mut cursor = TranscriptCursor {
items: &mut self.transcript,
dirty: false,
};
for mutator in &mutators {
if cancellation
.as_ref()
.is_some_and(TurnCancellation::is_cancelled)
{
return Err(LoopError::Cancelled);
}
let ctx = LoopCtx {
session_id: &session_id,
turn_id,
point,
cancellation: cancellation.clone(),
emitter: &emitter,
};
mutator.mutate(&mut cursor, ctx).await?;
}
if cursor.dirty {
validate_transcript_invariants(cursor.items)?;
}
Ok(())
}
async fn continue_active_tool_round(&mut self) -> Result<Option<LoopStep>, LoopError> {
let Some((presentation_turn_id, task_turn_id, cancellation)) =
self.active_tool_round.as_ref().map(|active| {
(
active.presentation_turn_id.clone(),
active.task_turn_id.clone(),
active.cancellation.clone(),
)
})
else {
return Ok(None);
};
loop {
if cancellation
.as_ref()
.is_some_and(TurnCancellation::is_cancelled)
{
self.task_manager
.on_turn_interrupted(&task_turn_id)
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?;
self.active_tool_round = None;
return self
.finish_cancelled(presentation_turn_id, Vec::new())
.map(Some);
}
let next_call = self
.active_tool_round
.as_mut()
.and_then(|active| active.pending_calls.pop_front());
if let Some((call, tool_request)) = next_call {
use tracing::Instrument;
self.register_tool_cancellation(&call.id, cancellation.clone());
let dispatch_span =
self.execute_tool_span(&tool_request, &presentation_turn_id, "plain");
match self
.start_task_via_manager(
None,
tool_request.clone(),
TaskLaunchKind::Plain,
cancellation.clone(),
)
.instrument(dispatch_span.clone())
.await?
{
TaskStartOutcome::Ready(resolution) => {
let resolution = *resolution;
match resolution {
TaskResolution::Item(item) => {
if !tool_result_not_started(&item) {
self.emit(AgentEvent::ToolExecutionStarted(call.clone()));
}
if tool_result_is_error(&item) {
dispatch_span.record("error.type", "tool_error");
}
if let Some(active) = self.active_tool_round.as_mut() {
active.foreground_progressed = true;
}
self.append_tool_result_item(item);
}
TaskResolution::Approval(task) => {
self.enqueue_pending_approval(
&presentation_turn_id,
task,
cancellation.clone(),
);
}
}
continue;
}
TaskStartOutcome::Pending { kind, .. } => {
self.emit(AgentEvent::ToolExecutionStarted(call.clone()));
if kind == agentkit_task_manager::TaskKind::Background {
self.append_detach_placeholder(call.id.clone(), &call.name);
if let Some(active) = self.active_tool_round.as_mut() {
active.background_pending = true;
}
}
continue;
}
}
}
match self
.task_manager
.wait_for_turn(&task_turn_id, cancellation.clone())
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?
{
Some(TurnTaskUpdate::Resolution(resolution)) => {
let resolution = *resolution;
match resolution {
TaskResolution::Item(item) => {
if let Some(active) = self.active_tool_round.as_mut() {
active.foreground_progressed = true;
}
self.append_tool_result_item(item);
}
TaskResolution::Approval(task) => {
self.enqueue_pending_approval(
&presentation_turn_id,
task,
cancellation.clone(),
);
}
}
}
Some(TurnTaskUpdate::Detached(snapshot)) => {
self.append_detach_placeholder(snapshot.call_id, &snapshot.tool_name);
if let Some(active) = self.active_tool_round.as_mut() {
active.background_pending = true;
active.foreground_progressed = true;
}
}
None => {
if cancellation
.as_ref()
.is_some_and(TurnCancellation::is_cancelled)
{
self.task_manager
.on_turn_interrupted(&task_turn_id)
.await
.map_err(|error| {
LoopError::Tool(ToolError::Internal(error.to_string()))
})?;
self.active_tool_round = None;
return self
.finish_cancelled(presentation_turn_id, Vec::new())
.map(Some);
}
let active = self.active_tool_round.take().ok_or_else(|| {
LoopError::InvalidState("missing active tool round".into())
})?;
if let Some(step) = self.take_next_unsurfaced_approval_interrupt() {
return Ok(Some(step));
}
if let Some(step) = self.next_unresolved_approval_interrupt() {
return Ok(Some(step));
}
if active.background_pending && !active.foreground_progressed {
return Ok(None);
}
let info = ToolRoundInfo {
session_id: self.session_id.clone(),
turn_id: presentation_turn_id.clone(),
transcript_len: self.transcript.len(),
};
self.pending_round_resume = Some(presentation_turn_id);
return Ok(Some(LoopStep::Interrupt(LoopInterrupt::AfterToolResult(
info,
))));
}
}
}
}
#[tracing::instrument(
name = "agent.turn",
skip_all,
fields(
otel.name = "invoke_agent",
gen_ai.operation.name = "invoke_agent",
gen_ai.conversation.id = %self.session_id,
gen_ai.provider.name = tracing::field::Empty,
session.id = %self.session_id,
turn.id = %turn_id,
transcript.len = self.transcript.len(),
saw_tool_call = tracing::field::Empty,
finish_reason = tracing::field::Empty,
),
)]
async fn drive_turn(
&mut self,
turn_id: agentkit_core::TurnId,
mutation_point: MutationPoint,
) -> Result<LoopStep, LoopError> {
let cancellation = self
.cancellation
.as_ref()
.map(CancellationHandle::checkpoint);
match self
.run_mutators(mutation_point, Some(&turn_id), cancellation.clone())
.await
{
Ok(()) => {}
Err(LoopError::Cancelled) => {
return self.finish_cancelled(turn_id, interrupted_assistant_items());
}
Err(error) => return Err(error),
}
if !transcript_has_pending_input(&self.transcript) {
let turn_result = TurnResult {
turn_id,
finish_reason: FinishReason::Completed,
items: Vec::new(),
usage: None,
metadata: MetadataMap::new(),
};
self.finish_logical_turn(&turn_result);
return Ok(LoopStep::Finished(turn_result));
}
if cancellation
.as_ref()
.is_some_and(TurnCancellation::is_cancelled)
{
return self.finish_cancelled(turn_id, interrupted_assistant_items());
}
let catalog_events = self.tool_executor.drain_catalog_events();
self.emit_tool_catalog_events(catalog_events);
let request = TurnRequest {
session_id: self.session_id.clone(),
turn_id: turn_id.clone(),
transcript: self.transcript.clone(),
available_tools: self.tool_executor.specs(),
cache: self
.next_turn_cache
.take()
.or_else(|| self.default_cache.clone()),
metadata: MetadataMap::new(),
};
let session = self
.session
.as_mut()
.ok_or_else(|| LoopError::InvalidState("model session is not available".into()))?;
let chat_span = tracing::info_span!(
"chat",
"otel.name" = tracing::field::Empty,
"otel.kind" = "client",
"gen_ai.operation.name" = "chat",
"gen_ai.provider.name" = tracing::field::Empty,
"gen_ai.conversation.id" = %self.session_id,
"gen_ai.request.model" = tracing::field::Empty,
"gen_ai.response.model" = tracing::field::Empty,
"gen_ai.response.id" = tracing::field::Empty,
"gen_ai.response.finish_reasons" = tracing::field::Empty,
"gen_ai.input.messages" = tracing::field::Empty,
"gen_ai.output.messages" = tracing::field::Empty,
"gen_ai.usage.input_tokens" = tracing::field::Empty,
"gen_ai.usage.output_tokens" = tracing::field::Empty,
"gen_ai.usage.cost" = tracing::field::Empty,
);
if let Some(capture) = self.telemetry.input_messages {
record_string_array_attribute(
&chat_span,
"gen_ai.input.messages",
capture_messages(&request.transcript, capture, CaptureOrder::NewestTail),
);
}
let initial_provider_name =
effective_provider_name(session.provider_name(), self.provider_name.as_deref());
if let Some(provider) = &initial_provider_name {
chat_span.record("gen_ai.provider.name", provider.as_str());
tracing::Span::current().record("gen_ai.provider.name", provider.as_str());
}
match session.model_name() {
Some(model) => {
chat_span.record("gen_ai.request.model", model);
chat_span.record("otel.name", format!("chat {model}").as_str());
}
None => {
chat_span.record("otel.name", "chat");
}
}
use tracing::Instrument;
let mut turn = match session
.begin_turn(request, cancellation.clone())
.instrument(chat_span.clone())
.await
{
Ok(turn) => turn,
Err(LoopError::Cancelled) => {
self.task_manager
.on_turn_interrupted(&turn_id)
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?;
return self.finish_cancelled(turn_id, interrupted_assistant_items());
}
Err(error) => return Err(error),
};
let provider_name =
effective_provider_name(session.provider_name(), self.provider_name.as_deref());
if let Some(provider) = &provider_name {
chat_span.record("gen_ai.provider.name", provider.as_str());
tracing::Span::current().record("gen_ai.provider.name", provider.as_str());
}
match session.model_name() {
Some(model) => {
chat_span.record("gen_ai.request.model", model);
chat_span.record("otel.name", format!("chat {model}").as_str());
}
None => {
chat_span.record("otel.name", "chat");
}
}
let mut saw_tool_call = false;
let mut finished_result = None;
let mut latest_usage = None;
while let Some(event) = match turn
.next_event(cancellation.clone())
.instrument(chat_span.clone())
.await
{
Ok(event) => event,
Err(LoopError::Cancelled) => {
self.task_manager
.on_turn_interrupted(&turn_id)
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?;
return self.finish_cancelled(turn_id, interrupted_assistant_items());
}
Err(error) => return Err(error),
} {
if cancellation
.as_ref()
.is_some_and(TurnCancellation::is_cancelled)
{
self.task_manager
.on_turn_interrupted(&turn_id)
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?;
return self.finish_cancelled(turn_id, interrupted_assistant_items());
}
match event {
ModelTurnEvent::Delta(delta) => self.emit(AgentEvent::ContentDelta(delta)),
ModelTurnEvent::Usage(usage) => {
latest_usage = Some(usage.clone());
self.emit(AgentEvent::UsageUpdated(usage));
}
ModelTurnEvent::ToolCall(call) => {
saw_tool_call = true;
self.emit(AgentEvent::ToolCallRequested(call.clone()));
}
ModelTurnEvent::Finished(result) => {
finished_result = Some(result);
break;
}
}
}
let mut result = finished_result.ok_or_else(|| {
LoopError::Provider("model turn ended without a Finished event".into())
})?;
result.usage = merge_usage(result.usage, latest_usage);
if let Some(model) = &result.model {
chat_span.record("gen_ai.response.model", model.as_str());
}
if let Some(id) = &result.response_id {
chat_span.record("gen_ai.response.id", id.as_str());
}
if let Some(tokens) = result
.usage
.as_ref()
.and_then(|usage| usage.tokens.as_ref())
{
record_token_attribute(&chat_span, "gen_ai.usage.input_tokens", tokens.input_tokens);
record_token_attribute(
&chat_span,
"gen_ai.usage.output_tokens",
tokens.output_tokens,
);
}
if let Some(cost) = result.usage.as_ref().and_then(|usage| usage.cost.as_ref()) {
record_f64_attribute(&chat_span, "gen_ai.usage.cost", cost.amount);
}
record_string_array_attribute(
&chat_span,
"gen_ai.response.finish_reasons",
provider_finish_reasons(&result.metadata, &result.finish_reason),
);
if let Some(capture) = self.telemetry.output_messages {
record_string_array_attribute(
&chat_span,
"gen_ai.output.messages",
capture_messages(&result.output_items, capture, CaptureOrder::OldestHead),
);
}
drop(chat_span);
tracing::Span::current().record("saw_tool_call", saw_tool_call);
tracing::Span::current().record(
"finish_reason",
tracing::field::debug(&result.finish_reason),
);
let now = Timestamp::now();
let usage = result.usage.clone();
let finish_reason = result.finish_reason.clone();
let output_items: Vec<Item> = result
.output_items
.drain(..)
.map(|mut item| {
if matches!(item.kind, ItemKind::Assistant) {
if item.usage.is_none() {
item.usage = usage.clone();
}
if item.finish_reason.is_none() {
item.finish_reason = Some(finish_reason.clone());
}
}
if item.created_at.is_none() {
item.created_at = Some(now);
}
item
})
.collect();
self.extend_transcript(output_items.clone());
if saw_tool_call {
let pending_calls = extract_tool_calls(&output_items)
.into_iter()
.map(|call| {
let tool_request = ToolRequest {
call_id: call.id.clone(),
tool_name: agentkit_tools_core::ToolName::new(call.name.clone()),
input: call.input.clone(),
session_id: self.session_id.clone(),
turn_id: turn_id.clone(),
metadata: call.metadata.clone(),
};
(call, tool_request)
})
.collect();
self.active_tool_round = Some(ActiveToolRound {
presentation_turn_id: turn_id.clone(),
task_turn_id: turn_id.clone(),
pending_calls,
cancellation: cancellation.clone(),
background_pending: false,
foreground_progressed: false,
});
if let Some(step) = self.continue_active_tool_round().await? {
return Ok(step);
}
self.finish_logical_turn(&TurnResult {
turn_id,
finish_reason: result.finish_reason,
items: output_items,
usage: result.usage,
metadata: result.metadata,
});
return Ok(LoopStep::Interrupt(LoopInterrupt::AwaitingInput(
InputRequest {
session_id: self.session_id.clone(),
reason: "driver is waiting for input".into(),
},
)));
}
let turn_result = TurnResult {
turn_id,
finish_reason: result.finish_reason,
items: output_items,
usage: result.usage,
metadata: result.metadata,
};
self.finish_logical_turn(&turn_result);
Ok(LoopStep::Finished(turn_result))
}
async fn resume_after_approval(
&mut self,
pending: PendingApprovalToolCall,
) -> Result<LoopStep, LoopError> {
let decision = pending
.decision
.clone()
.ok_or_else(|| LoopError::InvalidState("pending approval has no decision".into()))?;
match decision {
ApprovalDecision::Approve => {
use tracing::Instrument;
self.emit(AgentEvent::ToolExecutionStarted(pending.call.clone()));
let dispatch_span = self.execute_tool_span(
&pending.tool_request,
&pending.presentation_turn_id,
"approved",
);
let cancellation = self
.cancellation
.as_ref()
.map(CancellationHandle::checkpoint);
self.register_tool_cancellation(&pending.call.id, cancellation.clone());
let start = self
.start_task_via_manager(
Some(pending.task_id.clone()),
pending.tool_request.clone(),
TaskLaunchKind::Approved(pending.request.clone()),
cancellation.clone(),
)
.instrument(dispatch_span.clone())
.await;
let outcome = match start {
Ok(outcome) => outcome,
Err(error) => {
self.append_tool_result_item(Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: pending.call.id.clone(),
output: ToolOutput::Text(format!(
"approved task failed to start: {error}"
)),
is_error: true,
metadata: pending.call.metadata.clone(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
});
let turn_id = pending.tool_request.turn_id.clone();
if let Err(cleanup_error) =
self.task_manager.on_turn_interrupted(&turn_id).await
{
tracing::debug!(
%cleanup_error,
%turn_id,
"failed to clean up turn after approved task start error"
);
}
return Err(error);
}
};
match outcome {
TaskStartOutcome::Ready(resolution) => {
let resolution = *resolution;
if let TaskResolution::Item(item) = &resolution
&& tool_result_is_error(item)
{
dispatch_span.record("error.type", "tool_error");
}
if let Some(step) = self.queue_resolution_interrupt(
&pending.presentation_turn_id,
resolution,
cancellation,
) {
return Ok(step);
}
}
TaskStartOutcome::Pending { kind, .. } => {
if kind == agentkit_task_manager::TaskKind::Background {
self.append_detach_placeholder(
pending.call.id.clone(),
&pending.call.name,
);
} else {
self.active_tool_round = Some(ActiveToolRound {
presentation_turn_id: pending.presentation_turn_id.clone(),
task_turn_id: pending.tool_request.turn_id.clone(),
pending_calls: VecDeque::new(),
cancellation: cancellation.clone(),
background_pending: false,
foreground_progressed: false,
});
}
}
}
}
ApprovalDecision::Deny { reason } => {
self.append_tool_result_item(Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: pending.call.id.clone(),
output: ToolOutput::Text(
reason.unwrap_or_else(|| "approval denied".into()),
),
is_error: true,
metadata: pending.call.metadata.clone(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
});
}
}
if let Some(step) = self.continue_active_tool_round().await? {
Ok(step)
} else if let Some(step) = self.take_next_unsurfaced_approval_interrupt() {
Ok(step)
} else if let Some(step) = self.next_unresolved_approval_interrupt() {
Ok(step)
} else {
self.drive_turn(pending.presentation_turn_id, MutationPoint::AfterToolResult)
.await
}
}
fn finish_cancelled(
&mut self,
turn_id: agentkit_core::TurnId,
items: Vec<Item>,
) -> Result<LoopStep, LoopError> {
let pending = self.drain_pending_approval_items();
self.reject_drained_approvals(pending);
self.close_interrupted_tool_calls();
self.extend_transcript(items.clone());
let turn_result = TurnResult {
turn_id,
finish_reason: FinishReason::Cancelled,
items,
usage: None,
metadata: interrupted_metadata("turn"),
};
self.finish_logical_turn(&turn_result);
Ok(LoopStep::Finished(turn_result))
}
pub fn submit_input(&mut self, input: Vec<Item>) -> Result<(), LoopError> {
if self.has_pending_interrupts() {
return Err(LoopError::InvalidState(
"cannot submit input while an interrupt is pending".into(),
));
}
self.emit(AgentEvent::InputAccepted {
session_id: self.session_id.clone(),
items: input.clone(),
});
self.pending_input.extend(input);
Ok(())
}
pub fn set_next_turn_cache(&mut self, cache: PromptCacheRequest) -> Result<(), LoopError> {
if self.has_pending_interrupts() {
return Err(LoopError::InvalidState(
"cannot update next-turn cache while an interrupt is pending".into(),
));
}
self.next_turn_cache = Some(cache);
Ok(())
}
#[cfg(test)]
pub(crate) fn submit_input_with_cache(
&mut self,
input: Vec<Item>,
cache: PromptCacheRequest,
) -> Result<(), LoopError> {
self.set_next_turn_cache(cache)?;
self.submit_input(input)
}
pub fn resolve_approval_for(
&mut self,
call_id: ToolCallId,
decision: ApprovalDecision,
) -> Result<(), LoopError> {
let Some(pending) = self.pending_approvals.get_mut(&call_id) else {
return Err(LoopError::InvalidState(format!(
"no approval request is pending for call {}",
call_id.0
)));
};
pending.decision = Some(decision.clone());
self.emit(AgentEvent::ApprovalResolved {
approved: matches!(decision, ApprovalDecision::Approve),
});
Ok(())
}
pub fn resolve_approval_for_with_patched_input(
&mut self,
call_id: ToolCallId,
input: serde_json::Value,
) -> Result<(), LoopError> {
let Some(pending) = self.pending_approvals.get_mut(&call_id) else {
return Err(LoopError::InvalidState(format!(
"no approval request is pending for call {}",
call_id.0
)));
};
pending.tool_request.input = input;
self.resolve_approval_for(call_id, ApprovalDecision::Approve)
}
pub fn resolve_approval(&mut self, decision: ApprovalDecision) -> Result<(), LoopError> {
let mut unresolved = self
.pending_approval_order
.iter()
.filter(|call_id| {
self.pending_approvals
.get(*call_id)
.is_some_and(|pending| pending.decision.is_none())
})
.cloned();
let Some(call_id) = unresolved.next() else {
return Err(LoopError::InvalidState(
"no approval request is pending".into(),
));
};
if unresolved.next().is_some() {
return Err(LoopError::InvalidState(
"multiple approvals are pending; use resolve_approval_for".into(),
));
}
self.resolve_approval_for(call_id, decision)
}
pub fn cancel_pending_approval_for(&mut self, call_id: ToolCallId) -> Result<(), LoopError> {
let Some(pending) = self.drain_pending_approval_for(&call_id) else {
return Err(LoopError::InvalidState(format!(
"no approval request is pending for call {}",
call_id.0
)));
};
let turn_id = pending.presentation_turn_id.clone();
self.reject_drained_approvals(vec![pending]);
if self.pending_approvals.is_empty() && self.active_tool_round.is_none() {
let _ = self.finish_cancelled(turn_id, Vec::new())?;
}
Ok(())
}
pub async fn cancel_pending_approvals(&mut self) -> Result<Option<LoopStep>, LoopError> {
if self.pending_approvals.is_empty() {
return Ok(None);
}
let Some(turn_id) = self
.pending_approval_order
.iter()
.find_map(|call_id| self.pending_approvals.get(call_id))
.map(|pending| pending.presentation_turn_id.clone())
else {
return Ok(None);
};
let mut seen_turns = HashSet::new();
let mut originating_turns = Vec::new();
for pending in self.pending_approvals.values() {
let originating_turn = pending.tool_request.turn_id.clone();
if seen_turns.insert(originating_turn.clone()) {
originating_turns.push(originating_turn);
}
}
let pending = self.drain_pending_approval_items();
self.active_tool_round = None;
let mut cleanup_error = None;
for originating_turn in originating_turns {
if let Err(error) = self
.task_manager
.on_turn_interrupted(&originating_turn)
.await
&& cleanup_error.is_none()
{
cleanup_error = Some(LoopError::Tool(ToolError::Internal(error.to_string())));
}
}
self.reject_drained_approvals(pending);
if let Some(error) = cleanup_error {
self.close_interrupted_tool_calls();
self.finish_logical_turn(&TurnResult {
turn_id,
finish_reason: FinishReason::Error,
items: Vec::new(),
usage: None,
metadata: MetadataMap::new(),
});
return Err(error);
}
self.finish_cancelled(turn_id, Vec::new()).map(Some)
}
pub fn snapshot(&self) -> LoopSnapshot {
LoopSnapshot {
session_id: self.session_id.clone(),
transcript: self.transcript.clone(),
pending_input: self.pending_input.clone(),
}
}
pub fn wait_for_loop_update(
&self,
) -> impl std::future::Future<Output = Result<(), LoopError>> + Send + 'static {
let has_collected_update = !self.pending_loop_updates.is_empty();
let task_manager = self.task_manager.clone();
async move {
if has_collected_update {
return Ok(());
}
task_manager
.wait_for_loop_update()
.await
.map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))
}
}
pub async fn next(&mut self) -> Result<LoopStep, LoopError> {
if self.lifecycle.active_turn.is_none() {
let continuation_turn = self
.pending_approval_order
.iter()
.find_map(|call_id| self.pending_approvals.get(call_id))
.map(|pending| pending.presentation_turn_id.clone())
.or_else(|| {
self.active_tool_round
.as_ref()
.map(|active| active.presentation_turn_id.clone())
})
.or_else(|| self.pending_round_resume.clone());
if let Some(turn_id) = continuation_turn {
self.start_logical_turn_with(turn_id);
} else if !self.pending_input.is_empty() {
self.start_logical_turn();
}
}
let result = self.next_inner().await;
match &result {
Ok(LoopStep::Finished(turn)) => self.finish_logical_turn(turn),
Err(_) => {
if let Some(turn_id) = self.lifecycle.active_turn.clone() {
self.recover_from_next_error().await;
self.finish_logical_turn(&TurnResult {
turn_id,
finish_reason: FinishReason::Error,
items: Vec::new(),
usage: None,
metadata: MetadataMap::new(),
});
}
}
_ => {}
}
result
}
async fn recover_from_next_error(&mut self) {
let mut seen_turns = HashSet::new();
let mut interrupted_turns = Vec::new();
if let Some(active) = self.active_tool_round.take()
&& seen_turns.insert(active.task_turn_id.clone())
{
interrupted_turns.push(active.task_turn_id);
}
if let Some(turn_id) = self.pending_round_resume.take()
&& seen_turns.insert(turn_id.clone())
{
interrupted_turns.push(turn_id);
}
for pending in self.pending_approvals.values() {
let turn_id = pending.tool_request.turn_id.clone();
if seen_turns.insert(turn_id.clone()) {
interrupted_turns.push(turn_id);
}
}
let pending = self.drain_pending_approval_items();
for turn_id in interrupted_turns {
if let Err(error) = self.task_manager.on_turn_interrupted(&turn_id).await {
tracing::debug!(%error, %turn_id, "failed to clean up turn after loop error");
}
}
self.reject_drained_approvals(pending);
self.close_interrupted_tool_calls();
}
async fn next_inner(&mut self) -> Result<LoopStep, LoopError> {
if let Some(pending) = self.take_next_resolved_approval() {
return self.resume_after_approval(pending).await;
}
if let Some(step) = self.finish_cancelled_pending_approval().await? {
return Ok(step);
}
if let Some(step) = self.take_next_unsurfaced_approval_interrupt() {
return Ok(step);
}
if let Some(step) = self.next_unresolved_approval_interrupt() {
return Ok(step);
}
if let Some(step) = self.continue_active_tool_round().await? {
return Ok(step);
}
if self.pending_round_resume.is_none() && !self.pending_input.is_empty() {
self.collect_pending_loop_updates().await?;
let turn_id = self.start_logical_turn();
let drained: Vec<Item> = std::mem::take(&mut self.pending_input);
self.extend_transcript(drained);
return self
.drive_turn(turn_id, MutationPoint::AfterTurnEnded)
.await;
}
let (had_loop_updates, loop_step) = self.drain_pending_loop_updates().await?;
if let Some(step) = loop_step {
return Ok(step);
}
if let Some(turn_id) = self.pending_round_resume.take() {
let drained: Vec<Item> = std::mem::take(&mut self.pending_input);
self.extend_transcript(drained);
return self
.drive_turn(turn_id, MutationPoint::AfterToolResult)
.await;
}
if self.pending_input.is_empty() && !had_loop_updates {
return Ok(LoopStep::Interrupt(LoopInterrupt::AwaitingInput(
InputRequest {
session_id: self.session_id.clone(),
reason: "driver is waiting for input".into(),
},
)));
}
let turn_id = self.start_logical_turn();
let drained: Vec<Item> = std::mem::take(&mut self.pending_input);
self.extend_transcript(drained);
self.drive_turn(turn_id, MutationPoint::AfterTurnEnded)
.await
}
fn emit(&self, event: AgentEvent) {
fan_out_observed_event(&self.observers, &self.observed_session_id, event);
}
fn append_item(&mut self, mut item: Item) {
if item.created_at.is_none() {
item.created_at = Some(Timestamp::now());
}
for observer in &self.transcript_observers {
observer.on_transcript_event(TranscriptEvent {
session_id: &self.session_id,
item: &item,
});
}
self.transcript.push(item);
}
fn append_detach_placeholder(&mut self, call_id: ToolCallId, tool_name: &str) {
self.background_call_ids.insert(call_id.clone());
if !self.detached_call_ids.insert(call_id.clone()) {
return;
}
let detached_result = ToolResultPart {
call_id: call_id.clone(),
output: ToolOutput::Text(format!(
"Tool {tool_name} is now running in the background. The result will be delivered when it completes."
)),
is_error: false,
metadata: MetadataMap::new(),
};
self.emit(AgentEvent::ToolExecutionProgress(detached_result.clone()));
self.append_item(Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(detached_result)],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
});
}
fn append_tool_result_item(&mut self, item: Item) {
for part in &item.parts {
if let Part::ToolResult(result) = part {
if !self
.interrupted_background_call_ids
.contains(&result.call_id)
{
self.emit(AgentEvent::ToolResultReceived(result.clone()));
}
self.background_call_ids.remove(&result.call_id);
self.clear_tool_cancellation(&result.call_id);
}
}
let item = self.maybe_convert_detached(item);
self.append_item(item);
}
fn drain_pending_approval_for(
&mut self,
call_id: &ToolCallId,
) -> Option<PendingApprovalToolCall> {
let pending = self.pending_approvals.remove(call_id)?;
self.pending_approval_order.retain(|id| id != call_id);
self.clear_tool_cancellation(call_id);
Some(pending)
}
fn drain_pending_approval_items(&mut self) -> Vec<PendingApprovalToolCall> {
let order = std::mem::take(&mut self.pending_approval_order);
let pending = order
.iter()
.filter_map(|call_id| {
let pending = self.pending_approvals.remove(call_id);
self.clear_tool_cancellation(call_id);
pending
})
.collect();
self.pending_approvals.clear();
pending
}
fn reject_drained_approvals(&mut self, pending: Vec<PendingApprovalToolCall>) {
for pending in pending {
self.emit(AgentEvent::ApprovalResolved { approved: false });
self.append_tool_result_item(cancelled_approval_item(pending));
}
}
fn close_interrupted_tool_calls(&mut self) {
for call in unanswered_tool_calls(&self.transcript) {
let call_id = call.id.clone();
let completes_in_background = self.background_call_ids.contains(&call_id);
self.append_tool_result_item(interrupted_tool_result_item(call));
if completes_in_background {
self.detached_call_ids.insert(call_id.clone());
self.interrupted_background_call_ids.insert(call_id);
}
}
}
fn maybe_convert_detached(&mut self, mut item: Item) -> Item {
if !matches!(item.kind, ItemKind::Tool) {
return item;
}
let results: Vec<&ToolResultPart> = item
.parts
.iter()
.filter_map(|p| match p {
Part::ToolResult(r) => Some(r),
_ => None,
})
.collect();
if results.is_empty()
|| !results
.iter()
.all(|r| self.detached_call_ids.contains(&r.call_id))
{
return item;
}
let structured_results = results
.iter()
.map(|result| {
Part::structured(serde_json::to_value(result).unwrap_or_else(
|error| serde_json::json!({ "serialization_error": error.to_string() }),
))
})
.collect::<Vec<_>>();
let failed = results.iter().filter(|result| result.is_error).count();
let with_metadata = results
.iter()
.filter(|result| !result.metadata.is_empty())
.count();
let mut text = format!(
"Background tool results: {} total, {failed} failed, {with_metadata} with metadata. ",
results.len()
);
for (index, result) in results.iter().enumerate() {
self.detached_call_ids.remove(&result.call_id);
self.interrupted_background_call_ids.remove(&result.call_id);
if text.chars().count() >= DETACHED_NOTIFICATION_TEXT_MAX_CHARS {
continue;
}
if index > 0 {
text.push_str("; ");
}
let label = if result.is_error {
"failed"
} else {
"completed"
};
let call_id = truncate_chars(&result.call_id.0, DETACHED_CALL_ID_MAX_CHARS);
let body = render_tool_output_brief(&result.output);
text.push_str(&format!("{call_id} {label}: {body}"));
}
let text = truncate_chars(&text, DETACHED_NOTIFICATION_TEXT_MAX_CHARS);
let mut notification_parts = Vec::with_capacity(1 + structured_results.len());
notification_parts.push(Part::text(text));
notification_parts.extend(structured_results);
item.kind = ItemKind::Notification;
item.parts = notification_parts;
item
}
fn extend_transcript(&mut self, items: impl IntoIterator<Item = Item>) {
let now = Timestamp::now();
for mut item in items {
if item.created_at.is_none() {
item.created_at = Some(now);
}
self.append_item(item);
}
}
}
fn render_tool_output_brief(output: &ToolOutput) -> String {
match output {
ToolOutput::Text(text) => format!(
"text preview: {}",
truncate_chars(text, DETACHED_TEXT_PREVIEW_MAX_CHARS)
),
ToolOutput::Structured(_) => "structured payload".into(),
ToolOutput::Parts(parts) => format!("parts payload ({} parts)", parts.len()),
ToolOutput::Files(files) => format!("files payload ({} files)", files.len()),
}
}
fn truncate_chars(text: &str, max_chars: usize) -> String {
let mut chars = text.chars();
let mut truncated = chars.by_ref().take(max_chars).collect::<String>();
if chars.next().is_some() && max_chars > 0 {
truncated.pop();
truncated.push('…');
}
truncated
}
fn interrupted_metadata(stage: &str) -> MetadataMap {
let mut metadata = MetadataMap::new();
metadata.insert(INTERRUPTED_METADATA_KEY.into(), true.into());
metadata.insert(
INTERRUPT_REASON_METADATA_KEY.into(),
USER_CANCELLED_REASON.into(),
);
metadata.insert(INTERRUPT_STAGE_METADATA_KEY.into(), stage.into());
metadata
}
fn record_token_attribute(span: &tracing::Span, key: &'static str, value: u64) {
match i64::try_from(value) {
Ok(value) => record_i64_attribute(span, key, value),
Err(_) => tracing::warn!(attribute = key, value, "token count exceeds OTEL i64 range"),
}
}
#[cfg(feature = "otel")]
fn record_i64_attribute(span: &tracing::Span, key: &'static str, value: i64) {
use tracing_opentelemetry::OpenTelemetrySpanExt;
span.set_attribute(key, value);
}
#[cfg(not(feature = "otel"))]
fn record_i64_attribute(span: &tracing::Span, key: &'static str, value: i64) {
span.record(key, value);
}
#[cfg(feature = "otel")]
fn record_f64_attribute(span: &tracing::Span, key: &'static str, value: f64) {
use tracing_opentelemetry::OpenTelemetrySpanExt;
span.set_attribute(key, value);
}
#[cfg(not(feature = "otel"))]
fn record_f64_attribute(span: &tracing::Span, key: &'static str, value: f64) {
span.record(key, value);
}
#[cfg(feature = "otel")]
fn otel_string_array(values: Vec<String>) -> opentelemetry::Value {
use opentelemetry::{Array, StringValue, Value as OtelValue};
OtelValue::Array(Array::String(
values.into_iter().map(StringValue::from).collect(),
))
}
#[cfg(feature = "otel")]
fn record_string_array_attribute(span: &tracing::Span, key: &'static str, values: Vec<String>) {
use tracing_opentelemetry::OpenTelemetrySpanExt;
span.set_attribute(key, otel_string_array(values));
}
#[cfg(not(feature = "otel"))]
fn record_string_array_attribute(span: &tracing::Span, key: &'static str, values: Vec<String>) {
span.record(key, tracing::field::debug(&values));
}
#[derive(Clone, Copy)]
enum CaptureOrder {
NewestTail,
OldestHead,
}
fn effective_provider_name(
session_provider: Option<&str>,
adapter_provider: Option<&str>,
) -> Option<String> {
session_provider.or(adapter_provider).map(str::to_owned)
}
fn merge_usage(final_usage: Option<Usage>, streamed_usage: Option<Usage>) -> Option<Usage> {
match (final_usage, streamed_usage) {
(None, streamed) => streamed,
(Some(final_usage), None) => Some(final_usage),
(Some(mut final_usage), Some(streamed)) => {
if final_usage.tokens.is_none() {
final_usage.tokens = streamed.tokens;
}
if final_usage.cost.is_none() {
final_usage.cost = streamed.cost;
}
for (key, value) in streamed.metadata {
final_usage.metadata.entry(key).or_insert(value);
}
Some(final_usage)
}
}
}
fn capture_messages(items: &[Item], capture: MessageCapture, order: CaptureOrder) -> Vec<String> {
let mut captured = Vec::new();
let mut used_bytes = 0usize;
let indices: Box<dyn Iterator<Item = usize>> = match order {
CaptureOrder::NewestTail => Box::new((0..items.len()).rev()),
CaptureOrder::OldestHead => Box::new(0..items.len()),
};
for index in indices.take(capture.max_messages) {
let item = &items[index];
let original_bytes = source_content_bytes(item);
let remaining = capture.max_bytes.saturating_sub(used_bytes);
if original_bytes > remaining {
captured.push(
serde_json::json!({
"type": "truncated",
"original_bytes": original_bytes,
})
.to_string(),
);
break;
}
used_bytes += original_bytes;
captured.push(capture_item_json(item, remaining));
}
if matches!(order, CaptureOrder::NewestTail) {
captured.reverse();
}
captured
}
fn source_content_bytes(item: &Item) -> usize {
item.parts.iter().fold(0, |total, part| {
total.saturating_add(part_source_content_bytes(part))
})
}
fn part_source_content_bytes(part: &Part) -> usize {
match part {
Part::Text(text) => text.text.len(),
Part::Media(media) => media.mime_type.len(),
Part::File(file) => file
.name
.as_deref()
.map_or(0, str::len)
.saturating_add(file.mime_type.as_deref().map_or(0, str::len)),
Part::Structured(_) => 0,
Part::Reasoning(reasoning) => reasoning.summary.as_deref().map_or(0, str::len),
Part::ToolCall(call) => call.id.0.len().saturating_add(call.name.len()),
Part::ToolResult(result) => result
.call_id
.0
.len()
.saturating_add(tool_output_source_content_bytes(&result.output)),
Part::Custom(custom) => custom.kind.len(),
}
}
fn tool_output_source_content_bytes(output: &ToolOutput) -> usize {
match output {
ToolOutput::Text(text) => text.len(),
ToolOutput::Structured(_) | ToolOutput::Parts(_) | ToolOutput::Files(_) => 0,
}
}
struct CaptureBudget {
remaining: usize,
}
impl CaptureBudget {
fn text(&mut self, text: &str) -> (String, bool) {
let end = floor_char_boundary(text, self.remaining.min(text.len()));
self.remaining = self.remaining.saturating_sub(end);
(text[..end].to_owned(), end < text.len())
}
}
fn floor_char_boundary(text: &str, mut index: usize) -> usize {
while index > 0 && !text.is_char_boundary(index) {
index -= 1;
}
index
}
const MAX_CAPTURED_PARTS_PER_ITEM: usize = 256;
fn capture_item_json(item: &Item, max_bytes: usize) -> String {
let mut budget = CaptureBudget {
remaining: max_bytes,
};
let mut parts = item
.parts
.iter()
.take(MAX_CAPTURED_PARTS_PER_ITEM)
.map(|part| sanitized_part(part, &mut budget))
.collect::<Vec<_>>();
if item.parts.len() > parts.len() {
parts.push(serde_json::json!({
"type": "truncated",
"reason": "part_limit",
}));
}
serde_json::json!({
"role": item_kind_name(item.kind),
"parts": parts,
})
.to_string()
}
fn item_kind_name(kind: ItemKind) -> &'static str {
match kind {
ItemKind::System => "system",
ItemKind::Developer => "developer",
ItemKind::User => "user",
ItemKind::Assistant => "assistant",
ItemKind::Tool => "tool",
ItemKind::Context => "context",
ItemKind::Notification => "notification",
}
}
fn modality_name(modality: Modality) -> &'static str {
match modality {
Modality::Audio => "audio",
Modality::Image => "image",
Modality::Video => "video",
Modality::Binary => "binary",
}
}
fn omitted_data_ref(data: &DataRef) -> Value {
let kind = match data {
DataRef::InlineText(_) => "inline_text",
DataRef::InlineBytes(_) => "inline_bytes",
DataRef::Uri(_) => "uri",
DataRef::Handle(_) => "handle",
};
serde_json::json!({ "kind": kind, "omitted": true })
}
fn bounded_field(text: &str, budget: &mut CaptureBudget) -> Value {
let (text, truncated) = budget.text(text);
serde_json::json!({ "value": text, "truncated": truncated })
}
fn sanitized_part(part: &Part, budget: &mut CaptureBudget) -> Value {
match part {
Part::Text(text) => serde_json::json!({
"type": "text",
"text": bounded_field(&text.text, budget),
}),
Part::Media(media) => serde_json::json!({
"type": "media",
"modality": modality_name(media.modality),
"mime_type": bounded_field(&media.mime_type, budget),
"data": omitted_data_ref(&media.data),
}),
Part::File(file) => serde_json::json!({
"type": "file",
"name": file.name.as_deref().map(|name| bounded_field(name, budget)),
"mime_type": file.mime_type.as_deref().map(|mime| bounded_field(mime, budget)),
"data": omitted_data_ref(&file.data),
}),
Part::Structured(_) => serde_json::json!({
"type": "structured",
"truncated": true,
}),
Part::Reasoning(reasoning) => serde_json::json!({
"type": "reasoning",
"summary": reasoning.summary.as_deref().map(|summary| bounded_field(summary, budget)),
"redacted": reasoning.redacted,
"data": reasoning.data.as_ref().map(omitted_data_ref),
}),
Part::ToolCall(call) => serde_json::json!({
"type": "tool_call",
"id": bounded_field(&call.id.0, budget),
"name": bounded_field(&call.name, budget),
"input": { "truncated": true },
}),
Part::ToolResult(result) => serde_json::json!({
"type": "tool_result",
"call_id": bounded_field(&result.call_id.0, budget),
"is_error": result.is_error,
"output": sanitized_tool_output(&result.output, budget),
}),
Part::Custom(custom) => serde_json::json!({
"type": "custom",
"kind": bounded_field(&custom.kind, budget),
"data": custom.data.as_ref().map(omitted_data_ref),
"value": custom.value.as_ref().map(|_| serde_json::json!({ "truncated": true })),
}),
}
}
fn sanitized_tool_output(output: &ToolOutput, budget: &mut CaptureBudget) -> Value {
match output {
ToolOutput::Text(text) => serde_json::json!({
"type": "text",
"text": bounded_field(text, budget),
}),
ToolOutput::Structured(_) => serde_json::json!({
"type": "structured",
"truncated": true,
}),
ToolOutput::Parts(parts) => serde_json::json!({
"type": "parts",
"count": parts.len(),
"truncated": true,
}),
ToolOutput::Files(files) => serde_json::json!({
"type": "files",
"count": files.len(),
"truncated": true,
}),
}
}
#[cfg(test)]
mod telemetry_tests {
use super::*;
#[test]
fn message_capture_is_off_by_default_and_independent() {
let default = TelemetryConfig::default();
assert_eq!(default.input_messages(), None);
assert_eq!(default.output_messages(), None);
let capture = MessageCapture::new(2, 1).unwrap();
let input_only = TelemetryConfig::default().with_input_messages(capture);
assert_eq!(input_only.input_messages().unwrap().max_messages(), 2);
assert_eq!(input_only.input_messages().unwrap().max_bytes(), 1);
assert_eq!(input_only.output_messages(), None);
assert_eq!(
MessageCapture::new(0, 1),
Err(MessageCaptureError::ZeroMessages)
);
assert_eq!(
MessageCapture::new(1, 0),
Err(MessageCaptureError::ZeroBytes)
);
}
#[test]
fn one_source_byte_is_not_rejected_for_json_envelope_overhead() {
let items = vec![Item::text(ItemKind::User, "x")];
let captured = capture_messages(
&items,
MessageCapture::new(1, 1).unwrap(),
CaptureOrder::OldestHead,
);
assert_eq!(captured.len(), 1);
let value: Value = serde_json::from_str(&captured[0]).unwrap();
assert_eq!(value["role"], "user");
assert_eq!(value["parts"][0]["text"]["value"], "x");
assert_eq!(value["parts"][0]["text"]["truncated"], false);
}
#[test]
fn multibyte_source_accounting_preserves_utf8_boundaries() {
let items = vec![Item::text(ItemKind::User, "é")];
let exact = capture_messages(
&items,
MessageCapture::new(1, "é".len()).unwrap(),
CaptureOrder::OldestHead,
);
let value: Value = serde_json::from_str(&exact[0]).unwrap();
assert_eq!(value["parts"][0]["text"]["value"], "é");
assert_eq!(value["parts"][0]["text"]["truncated"], false);
let too_small = capture_messages(
&items,
MessageCapture::new(1, 1).unwrap(),
CaptureOrder::OldestHead,
);
let value: Value = serde_json::from_str(&too_small[0]).unwrap();
assert_eq!(value["type"], "truncated");
assert_eq!(value["original_bytes"], 2);
}
#[test]
fn source_bytes_are_aggregated_without_charging_json_envelopes() {
let items = vec![
Item::text(ItemKind::User, "a"),
Item::text(ItemKind::Assistant, "b"),
Item::text(ItemKind::User, "cd"),
];
let captured = capture_messages(
&items,
MessageCapture::new(3, 3).unwrap(),
CaptureOrder::OldestHead,
);
assert_eq!(captured.len(), 3);
let values = captured
.iter()
.map(|encoded| serde_json::from_str::<Value>(encoded).unwrap())
.collect::<Vec<_>>();
assert_eq!(values[0]["parts"][0]["text"]["value"], "a");
assert_eq!(values[1]["parts"][0]["text"]["value"], "b");
assert_eq!(values[2]["type"], "truncated");
assert_eq!(values[2]["original_bytes"], 2);
}
#[test]
fn tiny_limits_emit_valid_structured_source_byte_truncation() {
let items = vec![Item::text(ItemKind::User, "x".repeat(1_000))];
let captured = capture_messages(
&items,
MessageCapture::new(1, 1).unwrap(),
CaptureOrder::OldestHead,
);
assert_eq!(captured.len(), 1);
let value: Value = serde_json::from_str(&captured[0]).unwrap();
assert_eq!(value["type"], "truncated");
assert_eq!(value["original_bytes"], 1_000);
}
#[test]
fn provider_finish_reason_metadata_round_trips() {
let mut metadata = MetadataMap::new();
set_provider_finish_reasons(&mut metadata, ["end_turn", "", "tool_use", "end_turn"]);
assert_eq!(
provider_finish_reasons(&metadata, &FinishReason::Completed),
vec!["end_turn", "tool_use"]
);
set_provider_finish_reasons(&mut metadata, std::iter::empty::<String>());
assert!(!metadata.contains_key(PROVIDER_FINISH_REASONS_METADATA_KEY));
let fallbacks = [
(FinishReason::Completed, "completed"),
(FinishReason::ToolCall, "tool_call"),
(FinishReason::MaxTokens, "max_tokens"),
(FinishReason::Cancelled, "cancelled"),
(FinishReason::Blocked, "blocked"),
(FinishReason::Error, "error"),
(FinishReason::Other("native".into()), "native"),
];
for (reason, expected) in fallbacks {
assert_eq!(
provider_finish_reasons(&MetadataMap::new(), &reason),
[expected]
);
}
}
#[test]
fn input_is_newest_tail_output_is_head_and_data_refs_are_omitted() {
let items = vec![
Item::text(ItemKind::User, "old"),
Item::text(ItemKind::User, "middle"),
Item::text(ItemKind::User, "new"),
];
let capture = MessageCapture::new(2, 10_000).unwrap();
let input = capture_messages(&items, capture, CaptureOrder::NewestTail);
assert!(input[0].contains("middle"));
assert!(input[1].contains("new"));
let output = capture_messages(&items, capture, CaptureOrder::OldestHead);
assert!(output[0].contains("old"));
assert!(output[1].contains("middle"));
let media = Item::new(
ItemKind::User,
vec![Part::media(
agentkit_core::Modality::Image,
"image/png",
agentkit_core::DataRef::uri("https://secret.invalid/image.png"),
)],
);
let encoded = capture_item_json(&media, 10_000);
assert!(!encoded.contains("secret.invalid"));
assert!(encoded.contains("omitted"));
}
#[test]
fn final_usage_wins_and_streamed_usage_fills_only_missing_fields() {
let mut final_metadata = MetadataMap::new();
final_metadata.insert("shared".into(), serde_json::json!("final"));
let final_usage = Usage {
tokens: Some(agentkit_core::TokenUsage::new(1, 2)),
cost: None,
metadata: final_metadata,
};
let mut streamed_metadata = MetadataMap::new();
streamed_metadata.insert("shared".into(), serde_json::json!("streamed"));
streamed_metadata.insert("stream_only".into(), serde_json::json!(true));
let streamed_usage = Usage {
tokens: Some(agentkit_core::TokenUsage::new(10, 20)),
cost: Some(agentkit_core::CostUsage::new(0.5, "USD")),
metadata: streamed_metadata,
};
let merged = merge_usage(Some(final_usage), Some(streamed_usage)).unwrap();
assert_eq!(merged.tokens.unwrap().input_tokens, 1);
assert_eq!(merged.cost.unwrap().amount, 0.5);
assert_eq!(merged.metadata["shared"], "final");
assert_eq!(merged.metadata["stream_only"], true);
}
#[test]
fn session_provider_precedes_adapter_fallback() {
assert_eq!(
effective_provider_name(Some("session"), Some("adapter")).as_deref(),
Some("session")
);
assert_eq!(
effective_provider_name(None, Some("adapter")).as_deref(),
Some("adapter")
);
}
}
#[cfg(all(test, feature = "otel"))]
mod true_otel_integration_tests {
use std::fs;
use std::process::Command;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
static TEMP_SEQUENCE: AtomicU64 = AtomicU64::new(0);
#[test]
fn actual_layer_exports_driven_loop_spans() {
let root = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.parent()
.unwrap()
.parent()
.unwrap();
let nonce = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
let sequence = TEMP_SEQUENCE.fetch_add(1, Ordering::Relaxed);
let temp = std::env::temp_dir().join(format!(
"agentkit-loop-true-otel-{}-{nonce}-{sequence}",
std::process::id()
));
fs::create_dir(&temp).unwrap();
fs::create_dir(temp.join("src")).unwrap();
let core_path = toml_path(&root.join("crates/agentkit-core"));
let loop_path = toml_path(&root.join("crates/agentkit-loop"));
let async_trait_version = locked_version(root, "async-trait");
let opentelemetry_version = locked_version(root, "opentelemetry");
let serde_json_version = locked_version(root, "serde_json");
let tokio_version = locked_version(root, "tokio");
let tracing_version = locked_version(root, "tracing");
let tracing_otel_version = locked_version(root, "tracing-opentelemetry");
let tracing_subscriber_version = locked_version(root, "tracing-subscriber");
let manifest = format!(
r#"[package]
name = "agentkit-loop-true-otel-test"
version = "0.0.0"
edition = "2024"
[dependencies]
agentkit-core = {{ path = {core_path} }}
agentkit-loop = {{ path = {loop_path}, features = ["otel"] }}
async-trait = "={async_trait_version}"
opentelemetry = {{ version = "={opentelemetry_version}", default-features = false }}
serde_json = "={serde_json_version}"
tokio = {{ version = "={tokio_version}", features = ["rt"] }}
tracing = "={tracing_version}"
tracing-opentelemetry = {{ version = "={tracing_otel_version}", default-features = false }}
tracing-subscriber = "={tracing_subscriber_version}"
"#
);
fs::write(temp.join("Cargo.toml"), manifest).unwrap();
fs::write(temp.join("src/main.rs"), TRUE_OTEL_HARNESS).unwrap();
let output = Command::new(env!("CARGO"))
.args(["run", "--quiet", "--offline"])
.current_dir(&temp)
.env("CARGO_TARGET_DIR", temp.join("target"))
.output()
.unwrap();
let _ = fs::remove_dir_all(&temp);
assert!(
output.status.success(),
"true OTEL harness failed:\nstdout:\n{}\nstderr:\n{}",
String::from_utf8_lossy(&output.stdout),
String::from_utf8_lossy(&output.stderr)
);
}
fn locked_version(root: &std::path::Path, name: &str) -> String {
let lock = fs::read_to_string(root.join("Cargo.lock")).unwrap();
let expected_name = format!("name = {}", serde_json::to_string(name).unwrap());
let versions = lock
.split("[[package]]")
.filter(|package| {
package
.lines()
.any(|line| line.trim() == expected_name.as_str())
})
.filter_map(|package| {
package.lines().find_map(|line| {
line.trim()
.strip_prefix("version = \"")
.and_then(|version| version.strip_suffix('\"'))
.map(str::to_owned)
})
})
.collect::<Vec<_>>();
assert_eq!(versions.len(), 1, "expected one locked version for {name}");
versions.into_iter().next().unwrap()
}
fn toml_path(path: &std::path::Path) -> String {
serde_json::to_string(&path.to_string_lossy()).unwrap()
}
#[test]
fn toml_path_escapes_windows_separators_and_quotes() {
let encoded = toml_path(std::path::Path::new(r#"C:\Users\name\quoted\"dir"#));
assert_eq!(encoded, r#""C:\\Users\\name\\quoted\\\"dir""#);
}
const TRUE_OTEL_HARNESS: &str = r#"
use std::borrow::Cow;
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use std::time::SystemTime;
use agentkit_core::{CostUsage, DataRef, FinishReason, Item, ItemKind, MetadataMap, Modality, Part, TokenUsage, TurnCancellation, Usage};
use agentkit_loop::{Agent, LoopError, LoopStep, MessageCapture, ModelAdapter, ModelSession, ModelTurn, ModelTurnEvent, ModelTurnResult, SessionConfig, TelemetryConfig, TurnRequest, set_provider_finish_reasons};
use async_trait::async_trait;
use opentelemetry::trace::{Span, SpanBuilder, SpanContext, Status, Tracer};
use opentelemetry::{Array, Context, KeyValue, Value};
use tracing_subscriber::layer::SubscriberExt;
#[derive(Clone, Debug)]
struct Exported { name: String, attributes: Vec<KeyValue> }
#[derive(Clone, Default)]
struct MemoryTracer { exported: Arc<Mutex<Vec<Exported>>> }
struct MemorySpan { name: String, attributes: Vec<KeyValue>, exported: Arc<Mutex<Vec<Exported>>>, ended: bool }
impl Tracer for MemoryTracer {
type Span = MemorySpan;
fn build_with_context(&self, builder: SpanBuilder, _: &Context) -> MemorySpan {
MemorySpan { name: builder.name.into_owned(), attributes: builder.attributes.unwrap_or_default(), exported: self.exported.clone(), ended: false }
}
}
impl Span for MemorySpan {
fn add_event_with_timestamp<T>(&mut self, _: T, _: SystemTime, _: Vec<KeyValue>) where T: Into<Cow<'static, str>> {}
fn span_context(&self) -> &SpanContext { &SpanContext::NONE }
fn is_recording(&self) -> bool { !self.ended }
fn set_attribute(&mut self, attribute: KeyValue) { self.attributes.push(attribute); }
fn set_status(&mut self, _: Status) {}
fn update_name<T>(&mut self, name: T) where T: Into<Cow<'static, str>> { self.name = name.into().into_owned(); }
fn add_link(&mut self, _: SpanContext, _: Vec<KeyValue>) {}
fn end_with_timestamp(&mut self, _: SystemTime) {
if !self.ended {
self.ended = true;
self.exported.lock().unwrap().push(Exported { name: self.name.clone(), attributes: self.attributes.clone() });
}
}
}
#[derive(Clone, Copy)]
enum BeginMode { Success, Error, Cancelled }
#[derive(Clone)]
struct ScriptedAdapter { adapter_provider: &'static str, final_usage: bool, overflow: bool, begin_mode: BeginMode }
struct ScriptedSession { selected_provider: Option<&'static str>, model: &'static str, final_usage: bool, overflow: bool, begin_mode: BeginMode }
struct ScriptedTurn { events: VecDeque<ModelTurnEvent> }
#[async_trait]
impl ModelAdapter for ScriptedAdapter {
type Session = ScriptedSession;
async fn start_session(&self, _: SessionConfig) -> Result<Self::Session, LoopError> {
Ok(ScriptedSession { selected_provider: Some("before-begin"), model: "before-model", final_usage: self.final_usage, overflow: self.overflow, begin_mode: self.begin_mode })
}
fn provider_name(&self) -> Option<&str> { Some(self.adapter_provider) }
}
#[async_trait]
impl ModelSession for ScriptedSession {
type Turn = ScriptedTurn;
async fn begin_turn(&mut self, _: TurnRequest, _: Option<TurnCancellation>) -> Result<Self::Turn, LoopError> {
match self.begin_mode {
BeginMode::Error => return Err(LoopError::InvalidState("begin failed".into())),
BeginMode::Cancelled => return Err(LoopError::Cancelled),
BeginMode::Success => {}
}
self.selected_provider = if self.final_usage { Some("session-after-begin") } else { None };
self.model = if self.final_usage { "model-after-begin" } else { "fallback-model-after-begin" };
let mut stream_meta = MetadataMap::new();
stream_meta.insert("stream_only".into(), serde_json::json!(true));
stream_meta.insert("shared".into(), serde_json::json!("stream"));
let streamed = Usage {
tokens: Some(TokenUsage::new(if self.overflow { i64::MAX as u64 + 1 } else { 30 }, 40)),
cost: Some(CostUsage::new(0.75, "USD")),
metadata: stream_meta,
};
let final_usage = if self.final_usage {
let mut metadata = MetadataMap::new();
metadata.insert("shared".into(), serde_json::json!("final"));
Some(Usage { tokens: Some(TokenUsage::new(if self.overflow { i64::MAX as u64 + 1 } else { 1 }, 2)), cost: None, metadata })
} else { None };
let mut result_metadata = MetadataMap::new();
set_provider_finish_reasons(&mut result_metadata, ["native", "native", "done"]);
let outputs = vec![
Item::text(ItemKind::Assistant, "first"),
Item::text(ItemKind::Assistant, "second"),
Item::text(ItemKind::Assistant, "third"),
];
Ok(ScriptedTurn { events: VecDeque::from([
ModelTurnEvent::Usage(streamed),
ModelTurnEvent::Finished(ModelTurnResult { finish_reason: FinishReason::Completed, output_items: outputs, usage: final_usage, metadata: result_metadata, model: Some(self.model.into()), response_id: Some("response-id".into()) }),
]) })
}
fn model_name(&self) -> Option<&str> { Some(self.model) }
fn provider_name(&self) -> Option<&str> { self.selected_provider }
}
#[async_trait]
impl ModelTurn for ScriptedTurn {
async fn next_event(&mut self, _: Option<TurnCancellation>) -> Result<Option<ModelTurnEvent>, LoopError> { Ok(self.events.pop_front()) }
}
fn attr<'a>(span: &'a Exported, key: &str) -> Option<&'a Value> {
span.attributes.iter().rev().find(|a| a.key.as_str() == key).map(|a| &a.value)
}
fn operation(span: &Exported) -> Option<&str> {
match attr(span, "gen_ai.operation.name") { Some(Value::String(value)) => Some(value.as_str()), _ => None }
}
fn json_array(span: &Exported, key: &str) -> Vec<serde_json::Value> {
match attr(span, key) {
Some(Value::Array(Array::String(values))) => values.iter().map(|v| serde_json::from_str(v.as_str()).unwrap()).collect(),
other => panic!("{key} was not Array<String>: {other:?}"),
}
}
fn run_attempt(adapter: ScriptedAdapter) -> (Vec<Exported>, Result<LoopStep, LoopError>) {
let tracer = MemoryTracer::default();
let subscriber = tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer.clone()));
let runtime = tokio::runtime::Builder::new_current_thread().build().unwrap();
let result = tracing::subscriber::with_default(subscriber, || runtime.block_on(async {
let media = Item::new(ItemKind::User, vec![
Part::media(Modality::Image, "image/png", DataRef::inline_bytes([1, 2, 3])),
Part::media(Modality::Audio, "audio/wav", DataRef::uri("https://secret.invalid/audio")),
]);
let agent = Agent::builder().model(adapter).transcript(vec![
Item::text(ItemKind::System, "old"), Item::text(ItemKind::User, "middle"), media,
]).input(vec![Item::text(ItemKind::User, "newest")]).telemetry(
TelemetryConfig::default()
.with_input_messages(MessageCapture::new(3, 100_000).unwrap())
.with_output_messages(MessageCapture::new(2, 100_000).unwrap())
).build().unwrap();
let mut driver = agent.start(SessionConfig::new("otel-test")).await.unwrap();
driver.next().await
}));
(tracer.exported.lock().unwrap().clone(), result)
}
fn run(adapter: ScriptedAdapter) -> (Vec<Exported>, agentkit_loop::TurnResult) {
let (spans, result) = run_attempt(adapter);
let result = match result.unwrap() { LoopStep::Finished(result) => result, other => panic!("unexpected step: {other:?}") };
(spans, result)
}
fn main() {
let (spans, result) = run(ScriptedAdapter { adapter_provider: "adapter", final_usage: true, overflow: true, begin_mode: BeginMode::Success });
let chat = spans.iter().find(|s| operation(s) == Some("chat")).unwrap();
let agent = spans.iter().find(|s| operation(s) == Some("invoke_agent")).unwrap();
assert_eq!(attr(chat, "gen_ai.provider.name"), Some(&Value::String("session-after-begin".into())));
assert_eq!(attr(chat, "gen_ai.request.model"), Some(&Value::String("model-after-begin".into())));
assert!(attr(chat, "gen_ai.usage.input_tokens").is_none());
assert_eq!(attr(chat, "gen_ai.usage.output_tokens"), Some(&Value::I64(2)));
assert_eq!(attr(chat, "gen_ai.usage.cost"), Some(&Value::F64(0.75)));
assert_eq!(attr(agent, "gen_ai.provider.name"), Some(&Value::String("session-after-begin".into())));
assert!(attr(agent, "gen_ai.usage.input_tokens").is_none());
assert!(attr(agent, "gen_ai.usage.cost").is_none());
match attr(chat, "gen_ai.response.finish_reasons") {
Some(Value::Array(Array::String(values))) => assert_eq!(values.iter().map(|v| v.as_str()).collect::<Vec<_>>(), ["native", "done"]),
other => panic!("finish reasons were not Array<String>: {other:?}"),
}
assert_eq!(result.usage.as_ref().unwrap().metadata["shared"], "final");
assert_eq!(result.usage.as_ref().unwrap().metadata["stream_only"], true);
let input = json_array(chat, "gen_ai.input.messages");
assert_eq!(input.len(), 3);
assert!(input[0].to_string().contains("middle"));
assert!(input[1].to_string().contains("omitted"));
assert!(input[2].to_string().contains("newest"));
assert!(!input.iter().any(|message| message.to_string().contains("old")));
let output = json_array(chat, "gen_ai.output.messages");
assert_eq!(output.len(), 2);
assert!(output[0].to_string().contains("first"));
assert!(output[1].to_string().contains("second"));
let encoded = format!("{input:?}{output:?}");
assert!(!encoded.contains("secret.invalid"));
assert!(!encoded.contains("[1,2,3]"));
let (spans, result) = run(ScriptedAdapter { adapter_provider: "adapter-fallback", final_usage: false, overflow: false, begin_mode: BeginMode::Success });
let chat = spans.iter().find(|s| operation(s) == Some("chat")).unwrap();
assert_eq!(attr(chat, "gen_ai.provider.name"), Some(&Value::String("adapter-fallback".into())));
assert_eq!(attr(chat, "gen_ai.request.model"), Some(&Value::String("fallback-model-after-begin".into())));
assert_eq!(attr(chat, "gen_ai.usage.input_tokens"), Some(&Value::I64(30)));
assert_eq!(attr(chat, "gen_ai.usage.cost"), Some(&Value::F64(0.75)));
assert_eq!(result.usage.unwrap().metadata["stream_only"], true);
for begin_mode in [BeginMode::Error, BeginMode::Cancelled] {
let (spans, result) = run_attempt(ScriptedAdapter { adapter_provider: "adapter-before-error", final_usage: false, overflow: false, begin_mode });
match begin_mode {
BeginMode::Error => assert!(matches!(result, Err(LoopError::InvalidState(_)))),
BeginMode::Cancelled => assert!(matches!(result, Ok(LoopStep::Finished(_)))),
BeginMode::Success => unreachable!(),
}
let chat = spans.iter().find(|s| operation(s) == Some("chat")).unwrap();
let agent = spans.iter().find(|s| operation(s) == Some("invoke_agent")).unwrap();
assert_eq!(attr(chat, "gen_ai.provider.name"), Some(&Value::String("before-begin".into())));
assert_eq!(attr(chat, "gen_ai.request.model"), Some(&Value::String("before-model".into())));
assert_eq!(attr(agent, "gen_ai.provider.name"), Some(&Value::String("before-begin".into())));
}
}
"#;
}
fn interrupted_assistant_items() -> Vec<Item> {
vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::Text(TextPart {
text: "Previous assistant response was interrupted by the user before completion."
.into(),
metadata: interrupted_metadata("assistant"),
})],
metadata: interrupted_metadata("assistant"),
usage: None,
finish_reason: None,
created_at: None,
}]
}
fn unanswered_tool_calls(transcript: &[Item]) -> Vec<ToolCallPart> {
let mut open: Vec<ToolCallPart> = Vec::new();
for part in transcript.iter().flat_map(|item| &item.parts) {
match part {
Part::ToolCall(call) => open.push(call.clone()),
Part::ToolResult(result) => open.retain(|call| call.id != result.call_id),
_ => {}
}
}
open
}
fn interrupted_tool_result_item(call: ToolCallPart) -> Item {
Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: call.id,
output: ToolOutput::Text("tool call interrupted before it reported a result".into()),
is_error: true,
metadata: interrupted_metadata("tool"),
})],
metadata: interrupted_metadata("tool"),
usage: None,
finish_reason: None,
created_at: None,
}
}
fn cancelled_approval_item(pending: PendingApprovalToolCall) -> Item {
Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: pending.call.id,
output: ToolOutput::Text("approval cancelled".into()),
is_error: true,
metadata: pending.call.metadata,
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}
}
fn transcript_has_pending_input(transcript: &[Item]) -> bool {
matches!(
transcript.last().map(|item| item.kind),
Some(ItemKind::User | ItemKind::Tool | ItemKind::Notification)
)
}
fn extract_tool_calls(items: &[Item]) -> Vec<ToolCallPart> {
let mut calls = Vec::new();
for item in items {
for part in &item.parts {
if let Part::ToolCall(call) = part {
calls.push(call.clone());
}
}
}
calls
}
fn tool_result_is_error(item: &Item) -> bool {
item.parts
.iter()
.any(|part| matches!(part, Part::ToolResult(result) if result.is_error))
}
fn tool_result_not_started(item: &Item) -> bool {
item.parts.iter().any(|part| {
matches!(
part,
Part::ToolResult(result)
if result
.metadata
.get(TOOL_RESULT_NOT_STARTED_METADATA_KEY)
.and_then(Value::as_bool)
== Some(true)
)
})
}
#[derive(Debug, Error)]
pub enum LoopError {
#[error("invalid driver state: {0}")]
InvalidState(String),
#[error("turn cancelled")]
Cancelled,
#[error("provider error: {0}")]
Provider(String),
#[error("tool error: {0}")]
Tool(#[from] ToolError),
#[error("mutator error: {0}")]
Mutator(String),
#[error("unsupported operation: {0}")]
Unsupported(String),
}
struct DriverEmitter<'a> {
session_id: &'a Arc<SessionId>,
observers: &'a [Arc<dyn LoopObserver>],
}
impl<'a> EventEmitter for DriverEmitter<'a> {
fn emit(&self, event: AgentEvent) {
fan_out_observed_event(self.observers, self.session_id, event);
}
}
fn fan_out_observed_event(
observers: &[Arc<dyn LoopObserver>],
session_id: &Arc<SessionId>,
event: AgentEvent,
) {
if observers.is_empty() {
return;
}
let observed = ObservedEvent {
session_id: Arc::clone(session_id),
event,
};
let last = observers.len() - 1;
for observer in &observers[..last] {
observer.handle_event(observed.clone());
}
observers[last].handle_event(observed);
}
fn validate_transcript_invariants(transcript: &[Item]) -> Result<(), LoopError> {
let mut pending: HashSet<ToolCallId> = HashSet::new();
let mut seen_calls: HashSet<ToolCallId> = HashSet::new();
let mut seen_results: HashSet<ToolCallId> = HashSet::new();
for item in transcript {
for part in &item.parts {
match part {
Part::ToolCall(call) => {
if !seen_calls.insert(call.id.clone()) {
return Err(LoopError::Mutator(format!(
"transcript invariant violation: duplicate tool_use: {}",
call.id.0
)));
}
pending.insert(call.id.clone());
}
Part::ToolResult(result) => {
if !pending.remove(&result.call_id) {
let kind = if seen_results.contains(&result.call_id) {
"duplicate"
} else {
"orphaned"
};
return Err(LoopError::Mutator(format!(
"transcript invariant violation: {kind} tool_result: {}",
result.call_id.0
)));
}
seen_results.insert(result.call_id.clone());
}
_ => {}
}
}
}
if !pending.is_empty() {
let missing: Vec<String> = pending.into_iter().map(|id| id.0).collect();
return Err(LoopError::Mutator(format!(
"transcript invariant violation: tool_use(s) without matching tool_result: {}",
missing.join(", ")
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::collections::VecDeque;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc as StdArc, Mutex as StdMutex};
use agentkit_core::{
CancellationController, ItemKind, Part, TextPart, ToolCallId, ToolCallPart, ToolOutput,
ToolResultPart,
};
use agentkit_task_manager::{
AsyncTaskManager, RoutingDecision, TaskEvent, TaskManager, TaskManagerError,
TaskManagerHandle, TaskRoutingPolicy,
};
use agentkit_tools_core::{
FileSystemPermissionRequest, PermissionCode, PermissionDecision, PermissionDenial, Tool,
ToolAnnotations, ToolCatalogEvent, ToolExecutionOutcome, ToolName, ToolRegistry,
ToolResult, ToolSpec,
};
use serde_json::{Value, json};
use tokio::sync::Notify;
use tokio::time::{Duration, timeout};
use super::*;
struct FakeAdapter;
struct SlowAdapter;
struct RecordingAdapter {
seen_descriptions: StdArc<StdMutex<Vec<Vec<String>>>>,
seen_caches: StdArc<StdMutex<Vec<Option<PromptCacheRequest>>>>,
}
struct MultiToolAdapter;
struct DualApprovalAdapter;
struct FakeSession;
struct SlowSession;
struct RecordingSession {
seen_descriptions: StdArc<StdMutex<Vec<Vec<String>>>>,
seen_caches: StdArc<StdMutex<Vec<Option<PromptCacheRequest>>>>,
}
struct MultiToolSession;
struct DualApprovalSession;
struct FakeTurn {
events: VecDeque<ModelTurnEvent>,
}
struct SlowTurn {
emitted: bool,
}
struct RecordingTurn {
emitted: bool,
}
struct MultiToolTurn {
events: VecDeque<ModelTurnEvent>,
}
struct DualApprovalTurn {
events: VecDeque<ModelTurnEvent>,
}
struct TestTaskManager<T> {
inner: T,
start_error: Option<&'static str>,
approved_start_error: Option<&'static str>,
pending_update_error: Option<(usize, &'static str)>,
pending_update_calls: AtomicUsize,
interrupted: Option<StdArc<StdMutex<Vec<agentkit_core::TurnId>>>>,
interrupt_error: Option<&'static str>,
}
impl<T> TestTaskManager<T> {
fn new(inner: T) -> Self {
Self {
inner,
start_error: None,
approved_start_error: None,
pending_update_error: None,
pending_update_calls: AtomicUsize::new(0),
interrupted: None,
interrupt_error: None,
}
}
fn fail_start(mut self, message: &'static str) -> Self {
self.start_error = Some(message);
self
}
fn fail_approved_start(mut self, message: &'static str) -> Self {
self.approved_start_error = Some(message);
self
}
fn fail_pending_update_on(mut self, call: usize, message: &'static str) -> Self {
self.pending_update_error = Some((call, message));
self
}
fn record_interrupts(
mut self,
interrupted: StdArc<StdMutex<Vec<agentkit_core::TurnId>>>,
) -> Self {
self.interrupted = Some(interrupted);
self
}
fn fail_interrupt(mut self, message: &'static str) -> Self {
self.interrupt_error = Some(message);
self
}
}
#[async_trait]
impl<T: TaskManager> TaskManager for TestTaskManager<T> {
async fn start_task(
&self,
request: TaskLaunchRequest,
ctx: TaskStartContext,
) -> Result<TaskStartOutcome, TaskManagerError> {
if let Some(message) = self.start_error.or_else(|| {
matches!(&request.kind, TaskLaunchKind::Approved(_))
.then_some(self.approved_start_error)
.flatten()
}) {
return Err(TaskManagerError::Internal(message.into()));
}
self.inner.start_task(request, ctx).await
}
async fn wait_for_turn(
&self,
turn_id: &agentkit_core::TurnId,
cancellation: Option<TurnCancellation>,
) -> Result<Option<TurnTaskUpdate>, TaskManagerError> {
self.inner.wait_for_turn(turn_id, cancellation).await
}
async fn take_pending_loop_updates(&self) -> Result<PendingLoopUpdates, TaskManagerError> {
if let Some((call, message)) = self.pending_update_error
&& self.pending_update_calls.fetch_add(1, Ordering::SeqCst) == call
{
return Err(TaskManagerError::Internal(message.into()));
}
self.inner.take_pending_loop_updates().await
}
async fn on_turn_interrupted(
&self,
turn_id: &agentkit_core::TurnId,
) -> Result<(), TaskManagerError> {
if let Some(interrupted) = &self.interrupted {
interrupted.lock().unwrap().push(turn_id.clone());
}
if let Some(message) = self.interrupt_error {
return Err(TaskManagerError::Internal(message.into()));
}
self.inner.on_turn_interrupted(turn_id).await
}
fn handle(&self) -> TaskManagerHandle {
self.inner.handle()
}
}
struct DelayedApprovalExecutor {
entered: StdArc<AtomicBool>,
release: StdArc<Notify>,
approved_entered: Option<StdArc<AtomicBool>>,
approved_release: Option<StdArc<Notify>>,
cancellation: Option<CancellationController>,
spec: ToolSpec,
}
impl DelayedApprovalExecutor {
fn new(entered: StdArc<AtomicBool>, release: StdArc<Notify>) -> Self {
Self {
entered,
release,
approved_entered: None,
approved_release: None,
cancellation: None,
spec: ToolSpec {
name: ToolName::new("echo"),
description: "delayed approval".into(),
input_schema: json!({
"type": "object",
"properties": {
"value": { "type": "string" }
},
"required": ["value"],
"additionalProperties": false
}),
output_schema: None,
annotations: ToolAnnotations::default(),
metadata: MetadataMap::new(),
},
}
}
fn cancelling_on_approval(mut self, controller: CancellationController) -> Self {
self.cancellation = Some(controller);
self
}
fn blocking_after_approval(
mut self,
entered: StdArc<AtomicBool>,
release: StdArc<Notify>,
) -> Self {
self.approved_entered = Some(entered);
self.approved_release = Some(release);
self
}
}
#[async_trait]
impl ToolExecutor for DelayedApprovalExecutor {
fn specs(&self) -> Vec<ToolSpec> {
vec![self.spec.clone()]
}
async fn execute(
&self,
request: ToolRequest,
_ctx: &mut ToolContext<'_>,
) -> ToolExecutionOutcome {
self.entered.store(true, Ordering::SeqCst);
self.release.notified().await;
if let Some(controller) = &self.cancellation {
controller.interrupt();
}
ToolExecutionOutcome::Interrupted(
agentkit_tools_core::ToolInterruption::ApprovalRequired(ApprovalRequest {
task_id: None,
call_id: None,
id: "approval:delayed".into(),
request_kind: "delayed.approval".into(),
reason: agentkit_tools_core::ApprovalReason::PolicyRequiresConfirmation,
summary: "delayed approval".into(),
metadata: request.metadata,
}),
)
}
async fn execute_approved(
&self,
request: ToolRequest,
approved_request: &ApprovalRequest,
ctx: &mut ToolContext<'_>,
) -> ToolExecutionOutcome {
let (Some(entered), Some(release)) = (&self.approved_entered, &self.approved_release)
else {
return self.execute(request, ctx).await;
};
let _ = approved_request;
entered.store(true, Ordering::SeqCst);
release.notified().await;
ToolExecutionOutcome::Completed(ToolResult {
result: ToolResultPart {
call_id: request.call_id,
output: ToolOutput::Text("approved-ok".into()),
is_error: false,
metadata: MetadataMap::new(),
},
duration: None,
metadata: MetadataMap::new(),
})
}
}
#[async_trait]
impl ModelAdapter for FakeAdapter {
type Session = FakeSession;
async fn start_session(&self, _config: SessionConfig) -> Result<Self::Session, LoopError> {
Ok(FakeSession)
}
}
#[async_trait]
impl ModelAdapter for SlowAdapter {
type Session = SlowSession;
async fn start_session(&self, _config: SessionConfig) -> Result<Self::Session, LoopError> {
Ok(SlowSession)
}
}
#[async_trait]
impl ModelAdapter for RecordingAdapter {
type Session = RecordingSession;
async fn start_session(&self, _config: SessionConfig) -> Result<Self::Session, LoopError> {
Ok(RecordingSession {
seen_descriptions: self.seen_descriptions.clone(),
seen_caches: self.seen_caches.clone(),
})
}
}
#[async_trait]
impl ModelAdapter for MultiToolAdapter {
type Session = MultiToolSession;
async fn start_session(&self, _config: SessionConfig) -> Result<Self::Session, LoopError> {
Ok(MultiToolSession)
}
}
#[async_trait]
impl ModelAdapter for DualApprovalAdapter {
type Session = DualApprovalSession;
async fn start_session(&self, _config: SessionConfig) -> Result<Self::Session, LoopError> {
Ok(DualApprovalSession)
}
}
#[async_trait]
impl ModelSession for FakeSession {
type Turn = FakeTurn;
async fn begin_turn(
&mut self,
request: TurnRequest,
_cancellation: Option<TurnCancellation>,
) -> Result<Self::Turn, LoopError> {
let has_tool_result = request.transcript.iter().any(|item| {
item.kind == ItemKind::Tool
&& item
.parts
.iter()
.any(|part| matches!(part, Part::ToolResult(_)))
});
let tool_name = request
.available_tools
.first()
.map(|tool| tool.name.0.clone())
.unwrap_or_else(|| "echo".into());
let events = if has_tool_result {
let result_text = request
.transcript
.iter()
.rev()
.find_map(|item| {
item.parts.iter().find_map(|part| match (item.kind, part) {
(ItemKind::Notification, Part::Text(text)) => Some(text.text.clone()),
(
_,
Part::ToolResult(ToolResultPart {
output: ToolOutput::Text(text),
..
}),
) => Some(text.clone()),
_ => None,
})
})
.unwrap_or_else(|| "missing".into());
VecDeque::from([ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::Completed,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::Text(TextPart {
text: format!("tool said: {result_text}"),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
})])
} else {
VecDeque::from([
ModelTurnEvent::ToolCall(agentkit_core::ToolCallPart {
id: ToolCallId::new("call-1"),
name: tool_name.clone(),
input: json!({ "value": "pong" }),
metadata: MetadataMap::new(),
}),
ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::ToolCall,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::ToolCall(agentkit_core::ToolCallPart {
id: ToolCallId::new("call-1"),
name: tool_name,
input: json!({ "value": "pong" }),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
}),
])
};
Ok(FakeTurn { events })
}
}
#[async_trait]
impl ModelSession for SlowSession {
type Turn = SlowTurn;
async fn begin_turn(
&mut self,
request: TurnRequest,
cancellation: Option<TurnCancellation>,
) -> Result<Self::Turn, LoopError> {
let should_block = request
.transcript
.iter()
.rev()
.find(|item| item.kind == ItemKind::User)
.is_some_and(|item| {
item.parts.iter().any(|part| match part {
Part::Text(text) => text.text == "do the long task",
_ => false,
})
});
if should_block && let Some(cancellation) = cancellation {
cancellation.cancelled().await;
return Err(LoopError::Cancelled);
}
Ok(SlowTurn { emitted: false })
}
}
#[async_trait]
impl ModelSession for RecordingSession {
type Turn = RecordingTurn;
async fn begin_turn(
&mut self,
request: TurnRequest,
_cancellation: Option<TurnCancellation>,
) -> Result<Self::Turn, LoopError> {
let descriptions = request
.available_tools
.iter()
.map(|tool| tool.description.clone())
.collect::<Vec<_>>();
self.seen_descriptions.lock().unwrap().push(descriptions);
self.seen_caches.lock().unwrap().push(request.cache.clone());
Ok(RecordingTurn { emitted: false })
}
}
#[async_trait]
impl ModelSession for MultiToolSession {
type Turn = MultiToolTurn;
async fn begin_turn(
&mut self,
request: TurnRequest,
_cancellation: Option<TurnCancellation>,
) -> Result<Self::Turn, LoopError> {
let has_tool_result = request.transcript.iter().any(|item| {
item.kind == ItemKind::Tool
&& item
.parts
.iter()
.any(|part| matches!(part, Part::ToolResult(_)))
});
let events = if has_tool_result {
VecDeque::from([ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::Completed,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::Text(TextPart {
text: "mixed tools finished".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
})])
} else {
let foreground = agentkit_core::ToolCallPart {
id: ToolCallId::new("call-foreground"),
name: "foreground-wait".into(),
input: json!({}),
metadata: MetadataMap::new(),
};
let background = agentkit_core::ToolCallPart {
id: ToolCallId::new("call-background"),
name: "background-wait".into(),
input: json!({}),
metadata: MetadataMap::new(),
};
VecDeque::from([
ModelTurnEvent::ToolCall(foreground.clone()),
ModelTurnEvent::ToolCall(background.clone()),
ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::ToolCall,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::ToolCall(foreground), Part::ToolCall(background)],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
}),
])
};
Ok(MultiToolTurn { events })
}
}
#[async_trait]
impl ModelSession for DualApprovalSession {
type Turn = DualApprovalTurn;
async fn begin_turn(
&mut self,
request: TurnRequest,
_cancellation: Option<TurnCancellation>,
) -> Result<Self::Turn, LoopError> {
let tool_results = request
.transcript
.iter()
.flat_map(|item| item.parts.iter())
.filter(|part| matches!(part, Part::ToolResult(_)))
.count();
let events = if tool_results >= 2 {
VecDeque::from([ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::Completed,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::Text(TextPart {
text: "both approvals finished".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
})])
} else {
let first = agentkit_core::ToolCallPart {
id: ToolCallId::new("call-1"),
name: "echo".into(),
input: json!({ "value": "first" }),
metadata: MetadataMap::new(),
};
let second = agentkit_core::ToolCallPart {
id: ToolCallId::new("call-2"),
name: "echo".into(),
input: json!({ "value": "second" }),
metadata: MetadataMap::new(),
};
VecDeque::from([
ModelTurnEvent::ToolCall(first.clone()),
ModelTurnEvent::ToolCall(second.clone()),
ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::ToolCall,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::ToolCall(first), Part::ToolCall(second)],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
}),
])
};
Ok(DualApprovalTurn { events })
}
}
#[async_trait]
impl ModelTurn for FakeTurn {
async fn next_event(
&mut self,
_cancellation: Option<TurnCancellation>,
) -> Result<Option<ModelTurnEvent>, LoopError> {
Ok(self.events.pop_front())
}
}
#[async_trait]
impl ModelTurn for SlowTurn {
async fn next_event(
&mut self,
cancellation: Option<TurnCancellation>,
) -> Result<Option<ModelTurnEvent>, LoopError> {
if let Some(cancellation) = cancellation
&& cancellation.is_cancelled()
{
return Err(LoopError::Cancelled);
}
if self.emitted {
Ok(None)
} else {
self.emitted = true;
Ok(Some(ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::Completed,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::Text(TextPart {
text: "done".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
})))
}
}
}
#[async_trait]
impl ModelTurn for RecordingTurn {
async fn next_event(
&mut self,
_cancellation: Option<TurnCancellation>,
) -> Result<Option<ModelTurnEvent>, LoopError> {
if self.emitted {
Ok(None)
} else {
self.emitted = true;
Ok(Some(ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::Completed,
output_items: vec![Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::Text(TextPart {
text: "done".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
usage: None,
metadata: MetadataMap::new(),
})))
}
}
}
#[async_trait]
impl ModelTurn for MultiToolTurn {
async fn next_event(
&mut self,
_cancellation: Option<TurnCancellation>,
) -> Result<Option<ModelTurnEvent>, LoopError> {
Ok(self.events.pop_front())
}
}
#[async_trait]
impl ModelTurn for DualApprovalTurn {
async fn next_event(
&mut self,
_cancellation: Option<TurnCancellation>,
) -> Result<Option<ModelTurnEvent>, LoopError> {
Ok(self.events.pop_front())
}
}
#[derive(Clone)]
struct EchoTool {
spec: ToolSpec,
}
#[derive(Clone)]
struct FailingTool {
spec: ToolSpec,
}
#[derive(Clone)]
struct RunThenDenyTool {
spec: ToolSpec,
}
impl Default for EchoTool {
fn default() -> Self {
Self {
spec: ToolSpec {
name: ToolName::new("echo"),
description: "Echo back a value".into(),
input_schema: json!({
"type": "object",
"properties": {
"value": { "type": "string" }
},
"required": ["value"],
"additionalProperties": false
}),
output_schema: None,
annotations: ToolAnnotations::default(),
metadata: MetadataMap::new(),
},
}
}
}
impl Default for FailingTool {
fn default() -> Self {
Self {
spec: ToolSpec {
name: ToolName::new("failing"),
description: "Always fails after execution starts".into(),
input_schema: json!({
"type": "object",
"properties": {
"value": { "type": "string" }
},
"additionalProperties": true
}),
output_schema: None,
annotations: ToolAnnotations::default(),
metadata: MetadataMap::new(),
},
}
}
}
impl Default for RunThenDenyTool {
fn default() -> Self {
Self {
spec: ToolSpec {
name: ToolName::new("run_then_deny"),
description: "Runs, then returns a permission-denied error".into(),
input_schema: json!({
"type": "object",
"properties": {
"value": { "type": "string" }
},
"additionalProperties": true
}),
output_schema: None,
annotations: ToolAnnotations::default(),
metadata: MetadataMap::new(),
},
}
}
}
#[derive(Clone)]
struct DynamicSpecTool {
spec: ToolSpec,
version: StdArc<AtomicUsize>,
}
impl DynamicSpecTool {
fn new(version: StdArc<AtomicUsize>) -> Self {
Self {
spec: ToolSpec {
name: ToolName::new("dynamic"),
description: "dynamic version 0".into(),
input_schema: json!({
"type": "object",
"properties": {},
"additionalProperties": false
}),
output_schema: None,
annotations: ToolAnnotations::default(),
metadata: MetadataMap::new(),
},
version,
}
}
}
#[async_trait]
impl Tool for EchoTool {
fn spec(&self) -> &ToolSpec {
&self.spec
}
fn proposed_requests(
&self,
request: &agentkit_tools_core::ToolRequest,
) -> Result<
Vec<Box<dyn agentkit_tools_core::PermissionRequest>>,
agentkit_tools_core::ToolError,
> {
Ok(vec![Box::new(FileSystemPermissionRequest::Read {
path: "/tmp/echo".into(),
metadata: request.metadata.clone(),
})])
}
async fn invoke(
&self,
request: agentkit_tools_core::ToolRequest,
_ctx: &mut ToolContext<'_>,
) -> Result<ToolResult, agentkit_tools_core::ToolError> {
let value = request
.input
.get("value")
.and_then(Value::as_str)
.ok_or_else(|| {
agentkit_tools_core::ToolError::InvalidInput("missing value".into())
})?;
Ok(ToolResult {
result: ToolResultPart {
call_id: request.call_id,
output: ToolOutput::Text(value.into()),
is_error: false,
metadata: MetadataMap::new(),
},
duration: None,
metadata: MetadataMap::new(),
})
}
}
#[async_trait]
impl Tool for FailingTool {
fn spec(&self) -> &ToolSpec {
&self.spec
}
async fn invoke(
&self,
_request: agentkit_tools_core::ToolRequest,
_ctx: &mut ToolContext<'_>,
) -> Result<ToolResult, agentkit_tools_core::ToolError> {
Err(agentkit_tools_core::ToolError::ExecutionFailed(
"runtime failed".into(),
))
}
}
#[async_trait]
impl Tool for RunThenDenyTool {
fn spec(&self) -> &ToolSpec {
&self.spec
}
async fn invoke(
&self,
_request: agentkit_tools_core::ToolRequest,
_ctx: &mut ToolContext<'_>,
) -> Result<ToolResult, agentkit_tools_core::ToolError> {
Err(agentkit_tools_core::ToolError::PermissionDenied(
PermissionDenial {
code: PermissionCode::CustomPolicyDenied,
message: "remote 403".into(),
metadata: MetadataMap::new(),
},
))
}
}
#[async_trait]
impl Tool for DynamicSpecTool {
fn spec(&self) -> &ToolSpec {
&self.spec
}
fn current_spec(&self) -> Option<ToolSpec> {
let mut spec = self.spec.clone();
spec.description = format!("dynamic version {}", self.version.load(Ordering::SeqCst));
Some(spec)
}
async fn invoke(
&self,
request: agentkit_tools_core::ToolRequest,
_ctx: &mut ToolContext<'_>,
) -> Result<ToolResult, agentkit_tools_core::ToolError> {
Ok(ToolResult {
result: ToolResultPart {
call_id: request.call_id,
output: ToolOutput::Text("ok".into()),
is_error: false,
metadata: MetadataMap::new(),
},
duration: None,
metadata: MetadataMap::new(),
})
}
}
struct DenyFsReads;
impl PermissionChecker for DenyFsReads {
fn evaluate(
&self,
request: &dyn agentkit_tools_core::PermissionRequest,
) -> PermissionDecision {
if request.kind() == "filesystem.read" {
return PermissionDecision::Deny(PermissionDenial {
code: PermissionCode::PathNotAllowed,
message: "reads denied in test".into(),
metadata: MetadataMap::new(),
});
}
PermissionDecision::Allow
}
}
struct ApproveFsReads;
impl PermissionChecker for ApproveFsReads {
fn evaluate(
&self,
request: &dyn agentkit_tools_core::PermissionRequest,
) -> PermissionDecision {
if request.kind() == "filesystem.read" {
return PermissionDecision::RequireApproval(ApprovalRequest {
task_id: None,
call_id: None,
id: "approval:fs-read".into(),
request_kind: request.kind().into(),
reason: agentkit_tools_core::ApprovalReason::SensitivePath,
summary: request.summary(),
metadata: request.metadata().clone(),
});
}
PermissionDecision::Allow
}
}
struct KeepRecentMutator {
keep: usize,
}
#[async_trait]
impl LoopMutator for KeepRecentMutator {
async fn mutate(
&self,
cursor: &mut TranscriptCursor<'_>,
ctx: LoopCtx<'_>,
) -> Result<(), LoopError> {
if cursor.len() < 2 {
return Ok(());
}
let drop = cursor.len().saturating_sub(self.keep);
ctx.emitter.emit(AgentEvent::MutationStarted {
session_id: ctx.session_id.clone(),
turn_id: ctx.turn_id.cloned(),
mutator: "keep-recent".into(),
point: ctx.point,
});
cursor.drain(..drop);
ctx.emitter.emit(AgentEvent::MutationFinished {
session_id: ctx.session_id.clone(),
turn_id: ctx.turn_id.cloned(),
mutator: "keep-recent".into(),
dirty: true,
metadata: MetadataMap::new(),
});
Ok(())
}
}
struct PointRecordingMutator {
points: StdArc<StdMutex<Vec<MutationPoint>>>,
}
#[async_trait]
impl LoopMutator for PointRecordingMutator {
async fn mutate(
&self,
_cursor: &mut TranscriptCursor<'_>,
ctx: LoopCtx<'_>,
) -> Result<(), LoopError> {
self.points.lock().unwrap().push(ctx.point);
Ok(())
}
}
struct RecordingObserver {
events: StdArc<StdMutex<Vec<AgentEvent>>>,
}
impl LoopObserver for RecordingObserver {
fn handle_event(&self, event: ObservedEvent) {
let event = event.event;
self.events.lock().unwrap().push(event);
}
}
fn turn_lifecycle_events(
events: &[AgentEvent],
) -> Vec<(agentkit_core::TurnId, Option<FinishReason>)> {
events
.iter()
.filter_map(|event| match event {
AgentEvent::TurnStarted { turn_id, .. } => Some((turn_id.clone(), None)),
AgentEvent::TurnFinished(turn) => {
Some((turn.turn_id.clone(), Some(turn.finish_reason.clone())))
}
_ => None,
})
.collect()
}
struct CatalogExecutor {
version: AtomicUsize,
events: StdMutex<Vec<ToolCatalogEvent>>,
}
impl CatalogExecutor {
fn new() -> Self {
Self {
version: AtomicUsize::new(0),
events: StdMutex::new(Vec::new()),
}
}
fn publish_change(&self, version: usize, event: ToolCatalogEvent) {
self.version.store(version, Ordering::SeqCst);
self.events.lock().unwrap().push(event);
}
}
#[async_trait]
impl ToolExecutor for CatalogExecutor {
fn specs(&self) -> Vec<ToolSpec> {
vec![ToolSpec {
name: ToolName::new("dynamic"),
description: format!("dynamic version {}", self.version.load(Ordering::SeqCst)),
input_schema: json!({
"type": "object",
"properties": {},
"additionalProperties": false
}),
output_schema: None,
annotations: ToolAnnotations::default(),
metadata: MetadataMap::new(),
}]
}
fn drain_catalog_events(&self) -> Vec<ToolCatalogEvent> {
std::mem::take(&mut *self.events.lock().unwrap())
}
async fn execute(
&self,
request: ToolRequest,
_ctx: &mut ToolContext<'_>,
) -> ToolExecutionOutcome {
ToolExecutionOutcome::Completed(ToolResult {
result: ToolResultPart {
call_id: request.call_id,
output: ToolOutput::Text("dynamic-ok".into()),
is_error: false,
metadata: MetadataMap::new(),
},
duration: None,
metadata: MetadataMap::new(),
})
}
}
#[derive(Clone)]
struct BlockingTool {
spec: ToolSpec,
entered: StdArc<AtomicBool>,
release: StdArc<Notify>,
output: &'static str,
}
impl BlockingTool {
fn new(
name: &str,
entered: StdArc<AtomicBool>,
release: StdArc<Notify>,
output: &'static str,
) -> Self {
Self {
spec: ToolSpec {
name: ToolName::new(name),
description: format!("blocking tool {name}"),
input_schema: json!({
"type": "object",
"properties": {},
"additionalProperties": false
}),
output_schema: None,
annotations: ToolAnnotations::default(),
metadata: MetadataMap::new(),
},
entered,
release,
output,
}
}
}
#[async_trait]
impl Tool for BlockingTool {
fn spec(&self) -> &ToolSpec {
&self.spec
}
async fn invoke(
&self,
request: agentkit_tools_core::ToolRequest,
_ctx: &mut ToolContext<'_>,
) -> Result<ToolResult, agentkit_tools_core::ToolError> {
self.entered.store(true, Ordering::SeqCst);
self.release.notified().await;
Ok(ToolResult {
result: ToolResultPart {
call_id: request.call_id,
output: ToolOutput::Text(self.output.into()),
is_error: false,
metadata: MetadataMap::new(),
},
duration: None,
metadata: MetadataMap::new(),
})
}
}
struct NameRoutingPolicy {
routes: Vec<(String, RoutingDecision)>,
}
impl NameRoutingPolicy {
fn new(routes: impl IntoIterator<Item = (impl Into<String>, RoutingDecision)>) -> Self {
Self {
routes: routes
.into_iter()
.map(|(name, decision)| (name.into(), decision))
.collect(),
}
}
}
impl TaskRoutingPolicy for NameRoutingPolicy {
fn route(&self, request: &ToolRequest) -> RoutingDecision {
self.routes
.iter()
.find(|(name, _)| name == &request.tool_name.0)
.map(|(_, decision)| *decision)
.unwrap_or(RoutingDecision::Foreground)
}
}
async fn wait_for_task_event(handle: &TaskManagerHandle) -> TaskEvent {
timeout(Duration::from_secs(1), handle.next_event())
.await
.expect("timed out waiting for task event")
.expect("task event stream ended unexpectedly")
}
async fn wait_until_entered(flag: &AtomicBool) {
timeout(Duration::from_secs(1), async {
while !flag.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
})
.await
.expect("task never entered execution");
}
async fn wait_until_completed(handle: &TaskManagerHandle) {
timeout(Duration::from_secs(1), async {
while handle.list_completed().await.is_empty() {
tokio::task::yield_now().await;
}
})
.await
.expect("task never completed");
}
#[tokio::test]
async fn loop_continues_after_completed_tool_call() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-1"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "ping".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let result = run_until_finished(&mut driver).await;
match result {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Completed);
assert_eq!(turn.items.len(), 1);
match &turn.items[0].parts[0] {
Part::Text(text) => assert_eq!(text.text, "tool said: pong"),
other => panic!("unexpected part: {other:?}"),
}
}
other => panic!("unexpected loop step: {other:?}"),
}
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 2);
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[0].1, None);
assert_eq!(lifecycle[1].1, Some(FinishReason::Completed));
}
async fn run_until_finished<S: ModelSession + Send>(driver: &mut LoopDriver<S>) -> LoopStep {
loop {
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::AfterToolResult(_)) => continue,
step => return step,
}
}
}
#[tokio::test]
async fn post_tool_continuation_reports_after_tool_result_mutation_point() {
let points = StdArc::new(StdMutex::new(Vec::<MutationPoint>::new()));
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.mutator(PointRecordingMutator {
points: points.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-mutation-point"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
let _ = run_until_finished(&mut driver).await;
let recorded = points.lock().unwrap().clone();
assert_eq!(
recorded.first(),
Some(&MutationPoint::AfterTurnEnded),
"first drive of a fresh turn must report AfterTurnEnded, got {recorded:?}"
);
assert!(
recorded.contains(&MutationPoint::AfterToolResult),
"post-tool continuation must report AfterToolResult, got {recorded:?}"
);
}
#[tokio::test]
async fn no_work_awaiting_input_emits_no_turn_lifecycle() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(SlowAdapter)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-no-work"))
.await
.unwrap();
for _ in 0..2 {
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))
));
}
assert!(turn_lifecycle_events(&events.lock().unwrap()).is_empty());
}
#[tokio::test]
async fn normal_turn_emits_one_matched_lifecycle_pair() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(SlowAdapter)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-normal-lifecycle"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Finished(TurnResult {
finish_reason: FinishReason::Completed,
..
})
));
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 2);
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[0].1, None);
assert_eq!(lifecycle[1].1, Some(FinishReason::Completed));
}
#[tokio::test]
async fn post_start_error_emits_terminal_error_without_run_failed() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(SlowAdapter)
.mutator(ErrorMutator)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-error-lifecycle"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await,
Err(LoopError::Mutator(message)) if message == "boom"
));
let events = events.lock().unwrap();
let lifecycle = turn_lifecycle_events(&events);
assert_eq!(lifecycle.len(), 2);
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[0].1, None);
assert_eq!(lifecycle[1].1, Some(FinishReason::Error));
assert!(
!events
.iter()
.any(|event| matches!(event, AgentEvent::RunFailed { .. }))
);
}
#[tokio::test]
async fn active_tool_error_repairs_state_and_retry_uses_fresh_lifecycle() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let interrupted = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(ToolRegistry::new().with(EchoTool::default()))
.task_manager(
TestTaskManager::new(SimpleTaskManager::new())
.fail_start("original start failure")
.record_interrupts(interrupted.clone())
.fail_interrupt("cleanup failure"),
)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-active-tool-error"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "first")])
.unwrap();
let error = driver.next().await.unwrap_err();
assert!(error.to_string().contains("original start failure"));
assert!(!error.to_string().contains("cleanup failure"));
assert!(driver.active_tool_round.is_none());
assert!(driver.pending_round_resume.is_none());
assert!(unanswered_tool_calls(&driver.snapshot().transcript).is_empty());
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
assert_eq!(interrupted.lock().unwrap().len(), 1);
driver
.submit_input(vec![Item::text(ItemKind::User, "retry")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Finished(TurnResult {
finish_reason: FinishReason::Completed,
..
})
));
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 4, "{lifecycle:?}");
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[1].1, Some(FinishReason::Error));
assert_eq!(lifecycle[2].0, lifecycle[3].0);
assert_eq!(lifecycle[3].1, Some(FinishReason::Completed));
assert_ne!(lifecycle[0].0, lifecycle[2].0);
}
#[tokio::test]
async fn continuation_error_clears_resume_and_retry_uses_fresh_lifecycle() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let interrupted = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(ToolRegistry::new().with(EchoTool::default()))
.task_manager(
TestTaskManager::new(SimpleTaskManager::new())
.fail_pending_update_on(1, "original continuation failure")
.record_interrupts(interrupted.clone())
.fail_interrupt("cleanup failure"),
)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-continuation-error"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "first")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AfterToolResult(_))
));
let error = driver.next().await.unwrap_err();
assert!(error.to_string().contains("original continuation failure"));
assert!(!error.to_string().contains("cleanup failure"));
assert!(driver.pending_round_resume.is_none());
assert!(driver.active_tool_round.is_none());
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
assert_eq!(interrupted.lock().unwrap().len(), 1);
driver
.submit_input(vec![Item::text(ItemKind::User, "retry")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Finished(TurnResult {
finish_reason: FinishReason::Completed,
..
})
));
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 4, "{lifecycle:?}");
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[1].1, Some(FinishReason::Error));
assert_eq!(lifecycle[2].0, lifecycle[3].0);
assert_eq!(lifecycle[3].1, Some(FinishReason::Completed));
assert_ne!(lifecycle[0].0, lifecycle[2].0);
}
#[test]
fn pending_input_requires_input_bearing_tail_role() {
assert!(!transcript_has_pending_input(&[]));
assert!(!transcript_has_pending_input(&[Item::text(
ItemKind::System,
"system"
)]));
assert!(!transcript_has_pending_input(&[Item::text(
ItemKind::Developer,
"developer"
)]));
assert!(!transcript_has_pending_input(&[Item::text(
ItemKind::Context,
"context"
)]));
assert!(!transcript_has_pending_input(&[Item::text(
ItemKind::Assistant,
"assistant"
)]));
assert!(transcript_has_pending_input(&[Item::text(
ItemKind::User,
"user"
)]));
assert!(transcript_has_pending_input(&[Item::notification(
"background update"
)]));
assert!(transcript_has_pending_input(&[Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: ToolCallId::new("call-test"),
output: ToolOutput::Text("ok".into()),
is_error: false,
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}]));
}
struct DropTrailingUserMutator;
struct ErrorMutator;
#[async_trait]
impl LoopMutator for ErrorMutator {
async fn mutate(
&self,
_cursor: &mut TranscriptCursor<'_>,
_ctx: LoopCtx<'_>,
) -> Result<(), LoopError> {
Err(LoopError::Mutator("boom".into()))
}
}
#[async_trait]
impl LoopMutator for DropTrailingUserMutator {
async fn mutate(
&self,
cursor: &mut TranscriptCursor<'_>,
_ctx: LoopCtx<'_>,
) -> Result<(), LoopError> {
if cursor.last().map(|item| item.kind) == Some(ItemKind::User) {
cursor.pop();
}
Ok(())
}
}
struct RejectAssistantPrefillAdapter {
saw_assistant_tail: StdArc<AtomicBool>,
}
struct RejectAssistantPrefillSession {
saw_assistant_tail: StdArc<AtomicBool>,
}
#[async_trait]
impl ModelAdapter for RejectAssistantPrefillAdapter {
type Session = RejectAssistantPrefillSession;
async fn start_session(&self, _config: SessionConfig) -> Result<Self::Session, LoopError> {
Ok(RejectAssistantPrefillSession {
saw_assistant_tail: self.saw_assistant_tail.clone(),
})
}
}
#[async_trait]
impl ModelSession for RejectAssistantPrefillSession {
type Turn = FakeTurn;
async fn begin_turn(
&mut self,
request: TurnRequest,
_cancellation: Option<TurnCancellation>,
) -> Result<Self::Turn, LoopError> {
if request.transcript.last().map(|item| item.kind) == Some(ItemKind::Assistant) {
self.saw_assistant_tail.store(true, Ordering::SeqCst);
return Err(LoopError::Provider(
"conversation must end with a user message".into(),
));
}
Ok(FakeTurn {
events: VecDeque::from([ModelTurnEvent::Finished(ModelTurnResult {
model: None,
response_id: None,
finish_reason: FinishReason::Completed,
output_items: vec![Item::text(ItemKind::Assistant, "ok")],
usage: None,
metadata: MetadataMap::new(),
})]),
})
}
}
#[tokio::test]
async fn drive_does_not_dispatch_without_valid_trailing_input() {
let saw_assistant_tail = StdArc::new(AtomicBool::new(false));
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(RejectAssistantPrefillAdapter {
saw_assistant_tail: saw_assistant_tail.clone(),
})
.mutator(DropTrailingUserMutator)
.observer(RecordingObserver {
events: events.clone(),
})
.transcript(vec![
Item::text(ItemKind::User, "kickoff"),
Item::text(ItemKind::Assistant, "prior reply"),
])
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-no-valid-input"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "follow up")])
.unwrap();
let outcome = driver.next().await;
assert!(
!saw_assistant_tail.load(Ordering::SeqCst),
"loop dispatched a model turn whose transcript ends in an assistant \
message (outcome: {outcome:?}); with no valid trailing input the turn \
must finish instead of driving"
);
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 2);
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[0].1, None);
assert_eq!(lifecycle[1].1, Some(FinishReason::Completed));
}
#[tokio::test]
async fn loop_uses_injected_permission_checker() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(DenyFsReads)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-2"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "ping".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let result = run_until_finished(&mut driver).await;
match result {
LoopStep::Finished(turn) => match &turn.items[0].parts[0] {
Part::Text(text) => assert!(text.text.contains("tool permission denied")),
other => panic!("unexpected part: {other:?}"),
},
other => panic!("unexpected loop step: {other:?}"),
}
assert!(
events
.lock()
.unwrap()
.iter()
.all(|event| !matches!(event, AgentEvent::ToolExecutionStarted(_))),
"denied tools must not be reported as started"
);
}
#[tokio::test]
async fn failed_tool_execution_still_reports_started() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let tools = ToolRegistry::new().with(FailingTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-failing-start-event"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
match run_until_finished(&mut driver).await {
LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Completed),
other => panic!("unexpected loop step: {other:?}"),
}
let events = events.lock().unwrap();
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolExecutionStarted(call) if call.name == "failing"
)));
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolResultReceived(result) if result.is_error
)));
}
#[tokio::test]
async fn run_then_deny_tool_execution_still_reports_started() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let tools = ToolRegistry::new().with(RunThenDenyTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-run-then-deny-start-event"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
match run_until_finished(&mut driver).await {
LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Completed),
other => panic!("unexpected loop step: {other:?}"),
}
let events = events.lock().unwrap();
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolExecutionStarted(call) if call.name == "run_then_deny"
)));
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.is_error
&& result
.metadata
.get(TOOL_RESULT_FAILURE_KIND_METADATA_KEY)
.and_then(Value::as_str)
== Some(TOOL_RESULT_FAILURE_KIND_PERMISSION_DENIED)
&& result
.metadata
.get(TOOL_RESULT_NOT_STARTED_METADATA_KEY)
.is_none()
)));
}
#[tokio::test]
async fn async_task_manager_background_round_requires_explicit_continue() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"background-wait",
RoutingDecision::Background,
)]));
let handle = task_manager.handle();
let tools = ToolRegistry::new().with(BlockingTool::new(
"background-wait",
entered.clone(),
release.clone(),
"background-done",
));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.task_manager(task_manager)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-background"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "ping".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let first = driver.next().await.unwrap();
match first {
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) => {}
other => panic!("unexpected first loop step: {other:?}"),
}
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 2);
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[1].1, Some(FinishReason::ToolCall));
match wait_for_task_event(&handle).await {
TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "background-wait"),
other => panic!("unexpected task event: {other:?}"),
}
wait_until_entered(entered.as_ref()).await;
release.notify_waiters();
match wait_for_task_event(&handle).await {
TaskEvent::Completed(_, result) => {
assert_eq!(result.output, ToolOutput::Text("background-done".into()))
}
other => panic!("unexpected completion event: {other:?}"),
}
let resumed = driver.next().await.unwrap();
match resumed {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Completed);
match &turn.items[0].parts[0] {
Part::Text(text) => assert_eq!(
text.text,
"tool said: Background tool results: 1 total, 0 failed, 0 with metadata. \
call-1 completed: text preview: background-done"
),
other => panic!("unexpected part after resume: {other:?}"),
}
}
other => panic!("unexpected resumed step: {other:?}"),
}
let events = events.lock().unwrap();
let lifecycle = turn_lifecycle_events(&events);
assert_eq!(lifecycle.len(), 4);
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[2].0, lifecycle[3].0);
assert_ne!(lifecycle[0].0, lifecycle[2].0);
assert_eq!(lifecycle[3].1, Some(FinishReason::Completed));
let terminal_results: Vec<_> = events
.iter()
.filter_map(|event| match event {
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new("call-1") =>
{
Some(result)
}
_ => None,
})
.collect();
assert_eq!(
terminal_results.len(),
1,
"background completion must emit one terminal result event per call: {events:?}"
);
}
#[tokio::test]
async fn detached_parts_notification_preserves_full_output_and_metadata() {
let agent = Agent::builder().model(FakeAdapter).build().unwrap();
let mut driver = agent
.start(SessionConfig::new("session-detached-parts"))
.await
.unwrap();
let call_id = ToolCallId::new("parts-call");
driver.detached_call_ids.insert(call_id.clone());
let parts = vec![
Part::text("part text"),
Part::structured(json!({
"nested": [1, 2, 3]
})),
];
let mut metadata = MetadataMap::new();
metadata.insert("source".into(), json!("background"));
let result = ToolResultPart {
call_id,
output: ToolOutput::Parts(parts.clone()),
is_error: true,
metadata: metadata.clone(),
};
let mut item_metadata = MetadataMap::new();
item_metadata.insert("delivery".into(), json!("deferred"));
let item = Item::new(ItemKind::Tool, vec![Part::ToolResult(result.clone())])
.with_metadata(item_metadata.clone());
let converted = driver.maybe_convert_detached(item);
let (text, structured) = match converted.parts.as_slice() {
[Part::Text(text), Part::Structured(structured)] => (text, structured),
other => panic!("unexpected converted parts: {other:?}"),
};
assert_eq!(converted.kind, ItemKind::Notification);
assert_eq!(converted.metadata, item_metadata);
assert_eq!(structured.value, serde_json::to_value(&result).unwrap());
assert_eq!(
text.text,
"Background tool results: 1 total, 1 failed, 1 with metadata. \
parts-call failed: parts payload (2 parts)"
);
assert!(!text.text.contains("part text"));
assert!(!text.text.contains("background"));
}
#[tokio::test]
async fn detached_files_notification_preserves_full_output() {
let agent = Agent::builder().model(FakeAdapter).build().unwrap();
let mut driver = agent
.start(SessionConfig::new("session-detached-files"))
.await
.unwrap();
let call_id = ToolCallId::new("files-call");
driver.detached_call_ids.insert(call_id.clone());
let files = vec![
agentkit_core::FilePart::named("report.txt", DataRef::inline_text("full file body"))
.with_mime_type("text/plain"),
agentkit_core::FilePart::named(
"remote.json",
DataRef::uri("https://example.test/remote.json"),
),
];
let mut result_metadata = MetadataMap::new();
result_metadata.insert("archive".into(), json!(true));
let result = ToolResultPart::success(call_id, ToolOutput::Files(files.clone()))
.with_metadata(result_metadata);
let mut item_metadata = MetadataMap::new();
item_metadata.insert("delivery".into(), json!("deferred"));
let item = Item::new(ItemKind::Tool, vec![Part::ToolResult(result.clone())])
.with_metadata(item_metadata.clone());
let converted = driver.maybe_convert_detached(item);
let (text, structured) = match converted.parts.as_slice() {
[Part::Text(text), Part::Structured(structured)] => (text, structured),
other => panic!("unexpected converted files: {other:?}"),
};
assert_eq!(converted.kind, ItemKind::Notification);
assert_eq!(converted.metadata, item_metadata);
assert_eq!(structured.value, serde_json::to_value(&result).unwrap());
assert_eq!(
text.text,
"Background tool results: 1 total, 0 failed, 1 with metadata. \
files-call completed: files payload (2 files)"
);
assert!(!text.text.contains("full file body"));
assert!(!text.text.contains("remote.json"));
}
#[test]
fn detached_result_summaries_are_bounded_and_do_not_serialize_structured_payloads() {
let long_text = "é".repeat(DETACHED_TEXT_PREVIEW_MAX_CHARS + 20);
let text_summary = render_tool_output_brief(&ToolOutput::Text(long_text.clone()));
assert_eq!(
text_summary.chars().count(),
"text preview: ".chars().count() + DETACHED_TEXT_PREVIEW_MAX_CHARS
);
assert!(text_summary.ends_with('…'));
assert!(!text_summary.contains(&long_text));
let secret = "structured payload must remain out of notification text";
let structured = ToolOutput::Structured(json!({ "secret": secret }));
assert_eq!(render_tool_output_brief(&structured), "structured payload");
let oversized = "x".repeat(DETACHED_NOTIFICATION_TEXT_MAX_CHARS + 20);
let bounded = truncate_chars(&oversized, DETACHED_NOTIFICATION_TEXT_MAX_CHARS);
assert_eq!(
bounded.chars().count(),
DETACHED_NOTIFICATION_TEXT_MAX_CHARS
);
assert!(bounded.ends_with('…'));
}
#[tokio::test]
async fn detached_tool_placeholder_is_progress_not_terminal_result() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"detaching-wait",
RoutingDecision::ForegroundThenDetachAfter(Duration::from_millis(10)),
)]));
let handle = task_manager.handle();
let tools = ToolRegistry::new().with(BlockingTool::new(
"detaching-wait",
entered.clone(),
release.clone(),
"detached-done",
));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.task_manager(task_manager)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-detached-progress"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::AfterToolResult(_)) => {}
other => panic!("unexpected detach step: {other:?}"),
}
match wait_for_task_event(&handle).await {
TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "detaching-wait"),
other => panic!("unexpected task event: {other:?}"),
}
match wait_for_task_event(&handle).await {
TaskEvent::Detached(snapshot) => assert_eq!(snapshot.tool_name, "detaching-wait"),
other => panic!("unexpected detach event: {other:?}"),
}
wait_until_entered(entered.as_ref()).await;
release.notify_waiters();
match wait_for_task_event(&handle).await {
TaskEvent::Completed(_, result) => {
assert_eq!(result.output, ToolOutput::Text("detached-done".into()))
}
other => panic!("unexpected completion event: {other:?}"),
}
match driver.next().await.unwrap() {
LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Completed),
other => panic!("unexpected resumed step: {other:?}"),
}
let events = events.lock().unwrap();
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolExecutionProgress(result)
if result.call_id == ToolCallId::new("call-1") && !result.is_error
)));
let terminal_results: Vec<_> = events
.iter()
.filter_map(|event| match event {
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new("call-1") =>
{
Some(result)
}
_ => None,
})
.collect();
assert_eq!(
terminal_results.len(),
1,
"detached call must emit one terminal result event: {events:?}"
);
}
#[tokio::test]
async fn cancelled_background_approval_auto_resolves_when_drained() {
let controller = CancellationController::new();
let events = StdArc::new(StdMutex::new(Vec::new()));
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"echo",
RoutingDecision::Background,
)]));
let handle = task_manager.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(DelayedApprovalExecutor::new(
entered.clone(),
release.clone(),
))
.task_manager(task_manager)
.cancellation(controller.handle())
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-cancel-delayed-background-approval"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) => {}
other => panic!("unexpected first step: {other:?}"),
}
match wait_for_task_event(&handle).await {
TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "echo"),
other => panic!("unexpected task event: {other:?}"),
}
wait_until_entered(entered.as_ref()).await;
controller.interrupt();
release.notify_waiters();
wait_until_completed(&handle).await;
match driver.next().await.unwrap() {
LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Cancelled),
other => panic!("cancelled background approval should finish cancelled, got {other:?}"),
}
let events = events.lock().unwrap();
assert!(
events
.iter()
.any(|event| matches!(event, AgentEvent::ApprovalResolved { approved: false }))
);
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new("call-1") && result.is_error
)));
}
#[tokio::test]
async fn approved_foreground_task_waits_for_result_before_model_continuation() {
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let approved_entered = StdArc::new(AtomicBool::new(false));
let approved_release = StdArc::new(Notify::new());
let route_count = StdArc::new(AtomicUsize::new(0));
let routing_count = route_count.clone();
let task_manager = AsyncTaskManager::new().routing(move |_request: &ToolRequest| {
if routing_count.fetch_add(1, Ordering::SeqCst) == 0 {
RoutingDecision::Background
} else {
RoutingDecision::Foreground
}
});
let handle = task_manager.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(
DelayedApprovalExecutor::new(entered.clone(), release.clone())
.blocking_after_approval(approved_entered.clone(), approved_release.clone()),
)
.task_manager(task_manager)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-approved-foreground"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))
));
let task_turn = match wait_for_task_event(&handle).await {
TaskEvent::Started(snapshot) => snapshot.turn_id,
other => panic!("unexpected task event: {other:?}"),
};
wait_until_entered(entered.as_ref()).await;
release.notify_one();
wait_until_completed(&handle).await;
let pending = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending,
other => panic!("unexpected delayed approval step: {other:?}"),
};
let presentation_turn = driver.lifecycle.active_turn.clone().unwrap();
assert_ne!(presentation_turn, task_turn);
pending.approve(&mut driver).unwrap();
let info = {
let next = driver.next();
tokio::pin!(next);
tokio::select! {
() = wait_until_entered(approved_entered.as_ref()) => {}
result = &mut next => {
panic!("model continued before approved foreground result: {result:?}")
}
}
assert!(
timeout(Duration::from_millis(10), &mut next).await.is_err(),
"model continued while approved foreground work was blocked"
);
approved_release.notify_one();
let step = timeout(Duration::from_secs(1), &mut next)
.await
.expect("approved foreground result was not delivered")
.unwrap();
match step {
LoopStep::Interrupt(LoopInterrupt::AfterToolResult(info)) => info,
other => panic!("unexpected approved foreground step: {other:?}"),
}
};
assert_eq!(info.turn_id, presentation_turn);
let turn = match driver.next().await.unwrap() {
LoopStep::Finished(turn) => turn,
other => panic!("model did not continue after approved result: {other:?}"),
};
assert_eq!(turn.finish_reason, FinishReason::Completed);
assert_eq!(turn.turn_id, presentation_turn);
assert!(driver.snapshot().transcript.iter().any(|item| {
item.kind == ItemKind::Notification
&& item.parts.iter().any(
|part| matches!(part, Part::Text(text) if text.text.contains("approved-ok")),
)
}));
}
#[tokio::test]
async fn approved_foreground_then_detach_waits_and_keeps_one_placeholder() {
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let approved_entered = StdArc::new(AtomicBool::new(false));
let approved_release = StdArc::new(Notify::new());
let route_count = StdArc::new(AtomicUsize::new(0));
let routing_count = route_count.clone();
let task_manager = AsyncTaskManager::new().routing(move |_request: &ToolRequest| {
if routing_count.fetch_add(1, Ordering::SeqCst) == 0 {
RoutingDecision::Background
} else {
RoutingDecision::ForegroundThenDetachAfter(Duration::from_millis(10))
}
});
let handle = task_manager.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(
DelayedApprovalExecutor::new(entered.clone(), release.clone())
.blocking_after_approval(approved_entered.clone(), approved_release.clone()),
)
.task_manager(task_manager)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-approved-foreground-detach"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))
));
let _ = wait_for_task_event(&handle).await;
wait_until_entered(entered.as_ref()).await;
release.notify_one();
wait_until_completed(&handle).await;
let pending = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending,
other => panic!("unexpected delayed approval step: {other:?}"),
};
let presentation_turn = driver.lifecycle.active_turn.clone().unwrap();
pending.approve(&mut driver).unwrap();
let step = timeout(Duration::from_secs(1), driver.next())
.await
.expect("approved task did not detach")
.unwrap();
assert!(approved_entered.load(Ordering::SeqCst));
match step {
LoopStep::Interrupt(LoopInterrupt::AfterToolResult(info)) => {
assert_eq!(info.turn_id, presentation_turn);
}
other => panic!("unexpected approved detach step: {other:?}"),
}
let placeholders = driver
.snapshot()
.transcript
.iter()
.filter(|item| item.kind == ItemKind::Tool)
.flat_map(|item| &item.parts)
.filter(|part| {
matches!(
part,
Part::ToolResult(result) if result.call_id == ToolCallId::new("call-1")
)
})
.count();
assert_eq!(placeholders, 1, "detach appended a second tool result");
approved_release.notify_one();
wait_until_completed(&handle).await;
let turn = match driver.next().await.unwrap() {
LoopStep::Finished(turn) => turn,
other => panic!("model did not continue after detached result: {other:?}"),
};
assert_eq!(turn.finish_reason, FinishReason::Completed);
assert_eq!(turn.turn_id, presentation_turn);
let transcript = driver.snapshot().transcript;
assert_eq!(
transcript
.iter()
.filter(|item| item.kind == ItemKind::Tool)
.flat_map(|item| &item.parts)
.filter(|part| {
matches!(
part,
Part::ToolResult(result)
if result.call_id == ToolCallId::new("call-1")
)
})
.count(),
1
);
assert!(transcript.iter().any(|item| {
item.kind == ItemKind::Notification
&& item.parts.iter().any(
|part| matches!(part, Part::Text(text) if text.text.contains("approved-ok")),
)
}));
}
#[tokio::test]
async fn approving_detached_background_call_keeps_one_placeholder() {
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"echo",
RoutingDecision::Background,
)]));
let handle = task_manager.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(DelayedApprovalExecutor::new(
entered.clone(),
release.clone(),
))
.task_manager(task_manager)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-approve-detached-background"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))
));
let _ = wait_for_task_event(&handle).await;
wait_until_entered(entered.as_ref()).await;
release.notify_waiters();
wait_until_completed(&handle).await;
let pending = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending,
other => panic!("unexpected delayed approval step: {other:?}"),
};
pending.approve(&mut driver).unwrap();
release.notify_one();
let _ = driver.next().await.unwrap();
let placeholders = driver
.snapshot()
.transcript
.iter()
.flat_map(|item| &item.parts)
.filter(|part| {
matches!(
part,
Part::ToolResult(result) if result.call_id == ToolCallId::new("call-1")
)
})
.count();
assert_eq!(placeholders, 1, "approval appended a second detach result");
}
#[tokio::test]
async fn failed_background_approval_cleanup_clears_queued_resume() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let inner = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"echo",
RoutingDecision::ForegroundThenDetachAfter(Duration::from_millis(10)),
)]));
let handle = inner.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(DelayedApprovalExecutor::new(
entered.clone(),
release.clone(),
))
.task_manager(TestTaskManager::new(inner).fail_interrupt("cleanup failure"))
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new(
"session-failed-detached-background-approval-cleanup",
))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
let old_turn_id = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::AfterToolResult(info)) => info.turn_id,
other => panic!("unexpected detach step: {other:?}"),
};
assert_eq!(driver.pending_round_resume.as_ref(), Some(&old_turn_id));
driver
.submit_input(vec![Item::text(ItemKind::User, "fresh input")])
.unwrap();
let _ = wait_for_task_event(&handle).await;
wait_until_entered(entered.as_ref()).await;
release.notify_waiters();
wait_until_completed(&handle).await;
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_))
));
let error = driver.cancel_pending_approvals().await.unwrap_err();
assert!(error.to_string().contains("cleanup failure"));
assert!(driver.lifecycle.active_turn.is_none());
assert!(driver.pending_round_resume.is_none());
assert_eq!(driver.snapshot().pending_input.len(), 1);
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
let [started, finished] = &lifecycle[lifecycle.len() - 2..] else {
panic!("missing terminal lifecycle events: {lifecycle:?}");
};
assert_eq!(started.0, finished.0);
assert_eq!(finished.1, Some(FinishReason::Error));
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
let fresh_turn = match driver.next().await.unwrap() {
LoopStep::Finished(turn) => turn,
other => panic!("fresh input did not start a new turn: {other:?}"),
};
assert_ne!(fresh_turn.turn_id, old_turn_id);
}
#[tokio::test]
async fn fresh_input_runs_before_delayed_background_approval() {
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"echo",
RoutingDecision::Background,
)]));
let handle = task_manager.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(DelayedApprovalExecutor::new(
entered.clone(),
release.clone(),
))
.task_manager(task_manager)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new(
"session-input-before-background-approval",
))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))
));
let _ = wait_for_task_event(&handle).await;
wait_until_entered(entered.as_ref()).await;
release.notify_waiters();
wait_until_completed(&handle).await;
driver
.submit_input(vec![Item::text(ItemKind::User, "fresh input")])
.unwrap();
let fresh_turn = match driver.next().await.unwrap() {
LoopStep::Finished(turn) => turn.turn_id,
other => panic!("fresh input was not driven first: {other:?}"),
};
assert!(driver.pending_approvals.is_empty());
assert!(driver.snapshot().pending_input.is_empty());
timeout(Duration::from_millis(100), driver.wait_for_loop_update())
.await
.expect("collected background update did not wake the loop")
.unwrap();
let approval = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(approval)) => approval,
other => panic!("delayed approval was not presented separately: {other:?}"),
};
let approval_turn = driver.lifecycle.active_turn.clone().unwrap();
assert_ne!(fresh_turn, approval_turn);
assert_eq!(
driver
.snapshot()
.transcript
.iter()
.filter(|item| {
item.kind == ItemKind::User
&& item.parts.iter().any(
|part| matches!(part, Part::Text(text) if text.text == "fresh input"),
)
})
.count(),
1,
"fresh input must not be replayed while presenting the approval"
);
approval.deny(&mut driver).unwrap();
}
#[tokio::test]
async fn delayed_background_approval_interrupts_originating_task_turn() {
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let interrupted = StdArc::new(StdMutex::new(Vec::new()));
let inner = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"echo",
RoutingDecision::Background,
)]));
let handle = inner.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(DelayedApprovalExecutor::new(
entered.clone(),
release.clone(),
))
.task_manager(TestTaskManager::new(inner).record_interrupts(interrupted.clone()))
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new(
"session-background-approval-origin-turn",
))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))
));
let _ = wait_for_task_event(&handle).await;
wait_until_entered(entered.as_ref()).await;
release.notify_waiters();
wait_until_completed(&handle).await;
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_))
));
let presentation_turn = driver.lifecycle.active_turn.clone().unwrap();
let task_turn = driver
.pending_approvals
.values()
.next()
.unwrap()
.tool_request
.turn_id
.clone();
assert_ne!(presentation_turn, task_turn);
assert!(matches!(
driver.cancel_pending_approvals().await.unwrap(),
Some(LoopStep::Finished(TurnResult {
finish_reason: FinishReason::Cancelled,
..
}))
));
assert_eq!(interrupted.lock().unwrap().as_slice(), &[task_turn]);
assert!(driver.lifecycle.active_turn.is_none());
}
#[tokio::test]
async fn approved_background_start_error_interrupts_originating_turn() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let interrupted = StdArc::new(StdMutex::new(Vec::new()));
let inner = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"echo",
RoutingDecision::Background,
)]));
let handle = inner.handle();
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(DelayedApprovalExecutor::new(
entered.clone(),
release.clone(),
))
.task_manager(
TestTaskManager::new(inner)
.fail_approved_start("original approved start failure")
.record_interrupts(interrupted.clone())
.fail_interrupt("cleanup failure"),
)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-approved-start-error"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))
));
let _ = wait_for_task_event(&handle).await;
wait_until_entered(entered.as_ref()).await;
release.notify_waiters();
wait_until_completed(&handle).await;
let pending = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending,
other => panic!("unexpected delayed approval step: {other:?}"),
};
let task_turn = driver
.pending_approvals
.values()
.next()
.unwrap()
.tool_request
.turn_id
.clone();
let call_id = pending.request.call_id.clone().expect("approval call id");
assert!(driver.detached_call_ids.contains(&call_id));
pending.approve(&mut driver).unwrap();
let error = driver.next().await.unwrap_err();
assert!(
error
.to_string()
.contains("original approved start failure")
);
assert!(!error.to_string().contains("cleanup failure"));
assert_eq!(interrupted.lock().unwrap().as_slice(), &[task_turn]);
assert!(driver.lifecycle.active_turn.is_none());
assert!(!driver.detached_call_ids.contains(&call_id));
assert!(!driver.background_call_ids.contains(&call_id));
assert!(!driver.tool_cancellations.contains_key(&call_id));
assert!(events.lock().unwrap().iter().any(|event| matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.call_id == call_id && result.is_error
)));
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
}
#[tokio::test]
async fn loop_can_cancel_a_turn_and_continue_after_new_input() {
let controller = CancellationController::new();
let agent = Agent::builder()
.model(SlowAdapter)
.cancellation(controller.handle())
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-cancel"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "do the long task".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let cancelled = tokio::join!(async { driver.next().await }, async {
tokio::task::yield_now().await;
controller.interrupt();
})
.0
.unwrap();
match cancelled {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Cancelled);
assert_eq!(turn.items.len(), 1);
assert_eq!(turn.items[0].kind, ItemKind::Assistant);
assert_eq!(
turn.items[0].metadata.get(INTERRUPTED_METADATA_KEY),
Some(&Value::Bool(true))
);
}
other => panic!("unexpected loop step: {other:?}"),
}
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "try again".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let result = driver.next().await.unwrap();
match result {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Completed);
}
other => panic!("unexpected loop step after retry: {other:?}"),
}
}
#[tokio::test]
async fn loop_interrupt_cancels_foreground_tasks_but_keeps_background_tasks_running() {
let controller = CancellationController::new();
let fg_entered = StdArc::new(AtomicBool::new(false));
let fg_release = StdArc::new(Notify::new());
let bg_entered = StdArc::new(AtomicBool::new(false));
let bg_release = StdArc::new(Notify::new());
let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([
("foreground-wait", RoutingDecision::Foreground),
("background-wait", RoutingDecision::Background),
]));
let handle = task_manager.handle();
let tools = ToolRegistry::new()
.with(BlockingTool::new(
"foreground-wait",
fg_entered.clone(),
fg_release,
"foreground-done",
))
.with(BlockingTool::new(
"background-wait",
bg_entered.clone(),
bg_release.clone(),
"background-done",
));
let agent = Agent::builder()
.model(MultiToolAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.cancellation(controller.handle())
.task_manager(task_manager)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-mixed-cancel"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "run both".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let cancelled = tokio::join!(async { driver.next().await }, async {
let _ = wait_for_task_event(&handle).await;
let _ = wait_for_task_event(&handle).await;
wait_until_entered(fg_entered.as_ref()).await;
wait_until_entered(bg_entered.as_ref()).await;
controller.interrupt();
})
.0
.unwrap();
match cancelled {
LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Cancelled),
other => panic!("unexpected loop step after interrupt: {other:?}"),
}
match wait_for_task_event(&handle).await {
TaskEvent::Cancelled(snapshot) => assert_eq!(snapshot.tool_name, "foreground-wait"),
other => panic!("unexpected post-interrupt event: {other:?}"),
}
let running = handle.list_running().await;
assert_eq!(running.len(), 1);
assert_eq!(running[0].tool_name, "background-wait");
bg_release.notify_waiters();
match wait_for_task_event(&handle).await {
TaskEvent::Completed(snapshot, result) => {
assert_eq!(snapshot.tool_name, "background-wait");
assert_eq!(result.output, ToolOutput::Text("background-done".into()));
}
other => panic!("unexpected background completion event: {other:?}"),
}
}
#[tokio::test]
async fn a_cancelled_turn_answers_the_tool_call_it_abandoned() {
let controller = CancellationController::new();
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
let items = StdArc::new(StdMutex::new(Vec::<Item>::new()));
let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
"wait",
RoutingDecision::Foreground,
)]));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(ToolRegistry::new().with(BlockingTool::new(
"wait",
entered.clone(),
release,
"done",
)))
.permissions(AllowAllPermissions)
.cancellation(controller.handle())
.task_manager(task_manager)
.transcript_observer(RecordingTranscriptObserver {
items: items.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-cancel-mid-call"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "run the tool".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let cancelled = tokio::join!(async { driver.next().await }, async {
wait_until_entered(entered.as_ref()).await;
controller.interrupt();
})
.0
.unwrap();
match cancelled {
LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Cancelled),
other => panic!("unexpected loop step after interrupt: {other:?}"),
}
let transcript = driver.snapshot().transcript;
assert!(
unanswered_tool_calls(&transcript).is_empty(),
"the cancelled turn left a tool call unanswered: {transcript:?}"
);
validate_transcript_invariants(&transcript)
.expect("a cancelled turn must leave a resumable transcript");
let persisted = items.lock().unwrap().clone();
let results: Vec<&ToolResultPart> = persisted
.iter()
.flat_map(|item| &item.parts)
.filter_map(|part| match part {
Part::ToolResult(result) => Some(result),
_ => None,
})
.collect();
assert_eq!(results.len(), 1, "{persisted:?}");
assert_eq!(results[0].call_id, ToolCallId::new("call-1"));
assert!(results[0].is_error);
assert_eq!(
results[0].metadata.get(INTERRUPTED_METADATA_KEY),
Some(&Value::Bool(true))
);
}
#[tokio::test]
async fn regression_cancelled_background_completion_emits_one_terminal_result() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-cancelled-background-event"))
.await
.unwrap();
driver.append_item(Item::new(
ItemKind::Assistant,
vec![Part::ToolCall(ToolCallPart {
id: ToolCallId::new("call-1"),
name: "wait".into(),
input: json!({}),
metadata: MetadataMap::new(),
})],
));
driver.background_call_ids.insert(ToolCallId::new("call-1"));
driver.close_interrupted_tool_calls();
driver.append_tool_result_item(Item::new(
ItemKind::Tool,
vec![Part::ToolResult(ToolResultPart {
call_id: ToolCallId::new("call-1"),
output: ToolOutput::Text("background-done".into()),
is_error: false,
metadata: MetadataMap::new(),
})],
));
let events = events.lock().unwrap();
let terminal_results = events
.iter()
.filter(|event| {
matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new("call-1")
)
})
.count();
assert_eq!(
terminal_results, 1,
"a cancelled background call emitted multiple terminal results: {events:?}"
);
}
#[tokio::test]
async fn regression_cancelled_queued_approval_is_answered_once() {
let controller = CancellationController::new();
let entered = StdArc::new(AtomicBool::new(false));
let release = StdArc::new(Notify::new());
release.notify_one();
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.tool_executor(
DelayedApprovalExecutor::new(entered, release)
.cancelling_on_approval(controller.clone()),
)
.cancellation(controller.handle())
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-cancelled-queued-approval"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
match driver.next().await.unwrap() {
LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Cancelled),
other => panic!("unexpected first cancellation step: {other:?}"),
}
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) => {}
other => panic!("unexpected post-cancellation step: {other:?}"),
}
let events = events.lock().unwrap();
let terminal_results = events
.iter()
.filter(|event| {
matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new("call-1")
)
})
.count();
assert_eq!(
terminal_results, 1,
"a cancelled queued approval was answered more than once: {events:?}"
);
drop(events);
let transcript = driver.snapshot().transcript;
assert!(
!transcript.iter().any(|item| {
item.kind == ItemKind::Notification
&& item.parts.iter().any(|part| {
matches!(part, Part::Text(text) if text.text.contains("Background tool call"))
})
}),
"a queued approval was misreported as a background call: {transcript:?}"
);
}
#[tokio::test]
async fn regression_cancelled_unstarted_call_is_not_tracked_as_detached() {
let agent = Agent::builder().model(FakeAdapter).build().unwrap();
let mut driver = agent
.start(SessionConfig::new("session-cancelled-unstarted-call"))
.await
.unwrap();
driver.append_item(Item::new(
ItemKind::Assistant,
vec![Part::ToolCall(ToolCallPart {
id: ToolCallId::new("call-never-started"),
name: "wait".into(),
input: json!({}),
metadata: MetadataMap::new(),
})],
));
driver
.finish_cancelled(agentkit_core::TurnId::new("turn-cancelled"), Vec::new())
.unwrap();
assert!(
!driver
.detached_call_ids
.contains(&ToolCallId::new("call-never-started")),
"an unstarted call can never deliver a detached result"
);
}
#[tokio::test]
async fn loop_resumes_after_approved_tool_request() {
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(ApproveFsReads)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-approval"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "ping".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let first = driver.next().await.unwrap();
match first {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => {
assert!(pending.request.task_id.is_some());
assert_eq!(pending.request.id.0, "approval:fs-read");
pending.approve(&mut driver).unwrap();
}
other => panic!("unexpected loop step: {other:?}"),
}
let second = driver.next().await.unwrap();
match second {
LoopStep::Finished(turn) => match &turn.items[0].parts[0] {
Part::Text(text) => assert_eq!(text.text, "tool said: pong"),
other => panic!("unexpected part: {other:?}"),
},
other => panic!("unexpected loop step after approval: {other:?}"),
}
}
#[tokio::test]
async fn approval_gated_tool_does_not_start_before_approval() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(ApproveFsReads)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-approval-start-event"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
let pending = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending,
other => panic!("unexpected loop step: {other:?}"),
};
assert!(
events
.lock()
.unwrap()
.iter()
.all(|event| !matches!(event, AgentEvent::ToolExecutionStarted(_))),
"tool start must not be reported before approval"
);
pending.approve(&mut driver).unwrap();
match driver.next().await.unwrap() {
LoopStep::Finished(_) => {}
other => panic!("unexpected loop step after approval: {other:?}"),
}
let started = events
.lock()
.unwrap()
.iter()
.filter(|event| matches!(event, AgentEvent::ToolExecutionStarted(_)))
.count();
assert_eq!(started, 1);
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 2);
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[1].1, Some(FinishReason::Completed));
}
#[tokio::test]
async fn cancelling_pending_approval_resolves_it_and_pairs_tool_result() {
let controller = CancellationController::new();
let events = StdArc::new(StdMutex::new(Vec::new()));
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(ApproveFsReads)
.cancellation(controller.handle())
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-cancel-pending-approval"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_)) => {}
other => panic!("unexpected loop step: {other:?}"),
}
controller.interrupt();
match driver.next().await.unwrap() {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Cancelled);
}
other => panic!("unexpected loop step after cancel: {other:?}"),
}
let events = events.lock().unwrap();
assert!(
events
.iter()
.any(|event| matches!(event, AgentEvent::ApprovalResolved { approved: false })),
"pending approval cancellation should close approval UI state"
);
assert!(
events.iter().any(|event| matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new("call-1") && result.is_error
)),
"pending approval cancellation should pair the assistant tool_use"
);
drop(events);
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
}
#[tokio::test]
async fn cancelling_sole_foreground_approval_for_call_finishes_turn() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(ToolRegistry::new().with(EchoTool::default()))
.permissions(ApproveFsReads)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-cancel-foreground-approval-for"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
let call_id = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => {
pending.request.call_id.expect("approval call id")
}
other => panic!("unexpected loop step: {other:?}"),
};
driver.cancel_pending_approval_for(call_id).unwrap();
assert!(driver.lifecycle.active_turn.is_none());
assert!(driver.pending_approvals.is_empty());
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
let lifecycle = turn_lifecycle_events(&events.lock().unwrap());
assert_eq!(lifecycle.len(), 2, "{lifecycle:?}");
assert_eq!(lifecycle[0].0, lifecycle[1].0);
assert_eq!(lifecycle[1].1, Some(FinishReason::Cancelled));
}
#[tokio::test]
async fn resolved_approval_runs_even_if_cancellation_also_fired() {
let controller = CancellationController::new();
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(ApproveFsReads)
.cancellation(controller.handle())
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-resolved-approval-cancel-race"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
let pending = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending,
other => panic!("unexpected loop step: {other:?}"),
};
controller.interrupt();
pending.approve(&mut driver).unwrap();
match driver.next().await.unwrap() {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Completed);
match &turn.items[0].parts[0] {
Part::Text(text) => assert_eq!(text.text, "tool said: pong"),
other => panic!("unexpected part after approval: {other:?}"),
}
}
other => panic!("unexpected loop step after approved cancel race: {other:?}"),
}
}
#[tokio::test]
async fn loop_resumes_with_patched_input_on_approval() {
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(ApproveFsReads)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-approval-patched"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "ping".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => {
pending
.approve_with_patched_input(&mut driver, json!({ "value": "patched" }))
.unwrap();
}
other => panic!("unexpected loop step: {other:?}"),
}
match driver.next().await.unwrap() {
LoopStep::Finished(turn) => match &turn.items[0].parts[0] {
Part::Text(text) => assert_eq!(text.text, "tool said: patched"),
other => panic!("unexpected part: {other:?}"),
},
other => panic!("unexpected loop step after approval: {other:?}"),
}
}
#[tokio::test]
async fn loop_tracks_multiple_pending_approvals_by_call_id() {
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(DualApprovalAdapter)
.add_tool_source(tools)
.permissions(ApproveFsReads)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-dual-approval"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "run both approvals".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let pending_first = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => {
assert_eq!(
pending.request.call_id.as_ref().map(|id| id.0.as_str()),
Some("call-1")
);
pending
}
other => panic!("unexpected first loop step: {other:?}"),
};
let pending_second = match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => {
assert_eq!(
pending.request.call_id.as_ref().map(|id| id.0.as_str()),
Some("call-2")
);
pending
}
other => panic!("unexpected second loop step: {other:?}"),
};
pending_second.approve(&mut driver).unwrap();
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => {
assert_eq!(
pending.request.call_id.as_ref().map(|id| id.0.as_str()),
Some("call-1")
);
}
other => panic!("unexpected step after approving second request: {other:?}"),
}
pending_first.approve(&mut driver).unwrap();
match driver.next().await.unwrap() {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Completed);
match &turn.items[0].parts[0] {
Part::Text(text) => assert_eq!(text.text, "both approvals finished"),
other => panic!("unexpected final part: {other:?}"),
}
}
other => panic!("unexpected final loop step: {other:?}"),
}
}
#[tokio::test]
async fn failed_pending_approval_cleanup_repairs_and_finishes_error() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(ToolRegistry::new().with(EchoTool::default()))
.permissions(ApproveFsReads)
.task_manager(
TestTaskManager::new(SimpleTaskManager::new())
.fail_interrupt("interrupt cleanup failed"),
)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig::new("session-failed-approval-cleanup"))
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
assert!(matches!(
driver.next().await.unwrap(),
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_))
));
let error = driver.cancel_pending_approvals().await.unwrap_err();
assert!(error.to_string().contains("interrupt cleanup failed"));
assert!(driver.lifecycle.active_turn.is_none());
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
let events = events.lock().unwrap();
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new("call-1") && result.is_error
)));
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::TurnFinished(turn) if turn.finish_reason == FinishReason::Error
)));
}
#[tokio::test]
async fn cancelling_all_pending_approvals_interrupts_every_originating_turn() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let interrupted = StdArc::new(StdMutex::new(Vec::new()));
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(DualApprovalAdapter)
.add_tool_source(tools)
.permissions(ApproveFsReads)
.task_manager(
TestTaskManager::new(SimpleTaskManager::new())
.record_interrupts(interrupted.clone()),
)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-dual-approval-cancel"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "run both approvals".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
for expected_call in ["call-1", "call-2"] {
match driver.next().await.unwrap() {
LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => {
assert_eq!(
pending.request.call_id.as_ref().map(|id| id.0.as_str()),
Some(expected_call)
);
}
other => panic!("unexpected approval step: {other:?}"),
}
}
let first_origin = driver
.pending_approvals
.get(&ToolCallId::new("call-1"))
.unwrap()
.tool_request
.turn_id
.clone();
let second_origin = agentkit_core::TurnId::new("second-originating-turn");
driver
.pending_approvals
.get_mut(&ToolCallId::new("call-2"))
.unwrap()
.tool_request
.turn_id = second_origin.clone();
match driver.cancel_pending_approvals().await.unwrap() {
Some(LoopStep::Finished(turn)) => {
assert_eq!(turn.finish_reason, FinishReason::Cancelled);
}
other => panic!("unexpected cancellation result: {other:?}"),
}
validate_transcript_invariants(&driver.snapshot().transcript).unwrap();
let interrupted = interrupted
.lock()
.unwrap()
.iter()
.cloned()
.collect::<HashSet<_>>();
assert_eq!(interrupted, HashSet::from([first_origin, second_origin]));
let events = events.lock().unwrap();
let cancelled = events
.iter()
.filter(|event| matches!(event, AgentEvent::ApprovalResolved { approved: false }))
.count();
assert_eq!(cancelled, 2);
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::TurnFinished(turn) if turn.finish_reason == FinishReason::Cancelled
)));
for expected_call in ["call-1", "call-2"] {
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolResultReceived(result)
if result.call_id == ToolCallId::new(expected_call) && result.is_error
)));
}
}
#[tokio::test]
async fn loop_compacts_transcript_before_new_turns() {
let events = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.mutator(KeepRecentMutator { keep: 1 })
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-4"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
for text in ["first", "second"] {
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: text.into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let _ = driver.next().await.unwrap();
}
let events = events.lock().unwrap();
assert!(
events
.iter()
.any(|event| matches!(event, AgentEvent::MutationFinished { dirty: true, .. }))
);
}
#[test]
fn transcript_validation_rejects_orphaned_tool_result() {
let transcript = vec![Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: "call-1".into(),
output: ToolOutput::Text("result".into()),
is_error: false,
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}];
let error = validate_transcript_invariants(&transcript).unwrap_err();
assert!(error.to_string().contains("orphaned tool_result"));
}
#[test]
fn transcript_validation_rejects_duplicate_tool_result() {
let transcript = vec![
Item {
id: None,
kind: ItemKind::Assistant,
parts: vec![Part::ToolCall(ToolCallPart {
id: "call-1".into(),
name: "lookup".into(),
input: serde_json::json!({}),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
},
Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: "call-1".into(),
output: ToolOutput::Text("result".into()),
is_error: false,
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
},
Item {
id: None,
kind: ItemKind::Tool,
parts: vec![Part::ToolResult(ToolResultPart {
call_id: "call-1".into(),
output: ToolOutput::Text("again".into()),
is_error: false,
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
},
];
let error = validate_transcript_invariants(&transcript).unwrap_err();
assert!(error.to_string().contains("duplicate tool_result"));
}
#[tokio::test]
async fn loop_refreshes_tool_specs_each_turn() {
let seen_descriptions = StdArc::new(StdMutex::new(Vec::new()));
let version = StdArc::new(AtomicUsize::new(1));
let tools = ToolRegistry::new().with(DynamicSpecTool::new(version.clone()));
let agent = Agent::builder()
.model(RecordingAdapter {
seen_descriptions: seen_descriptions.clone(),
seen_caches: StdArc::new(StdMutex::new(Vec::new())),
})
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-dynamic-tools"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
for text in ["first", "second"] {
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: text.into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let _ = driver.next().await.unwrap();
if text == "first" {
version.store(2, Ordering::SeqCst);
}
}
let seen_descriptions = seen_descriptions.lock().unwrap();
assert_eq!(seen_descriptions.len(), 2);
assert_eq!(seen_descriptions[0], vec!["dynamic version 1".to_string()]);
assert_eq!(seen_descriptions[1], vec!["dynamic version 2".to_string()]);
}
#[tokio::test]
async fn loop_emits_catalog_change_and_uses_updated_specs_next_turn() {
let seen_descriptions = StdArc::new(StdMutex::new(Vec::new()));
let events = StdArc::new(StdMutex::new(Vec::new()));
let executor = StdArc::new(CatalogExecutor::new());
let executor_for_agent: Arc<dyn ToolExecutor> = executor.clone();
let agent = Agent::builder()
.model(RecordingAdapter {
seen_descriptions: seen_descriptions.clone(),
seen_caches: StdArc::new(StdMutex::new(Vec::new())),
})
.tool_executor(executor_for_agent)
.permissions(AllowAllPermissions)
.observer(RecordingObserver {
events: events.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-catalog-events"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "first")])
.unwrap();
let _ = driver.next().await.unwrap();
executor.publish_change(
1,
ToolCatalogEvent {
source: "mcp:mock".into(),
added: vec!["dynamic".into()],
removed: Vec::new(),
changed: Vec::new(),
},
);
driver
.submit_input(vec![Item::text(ItemKind::User, "second")])
.unwrap();
let _ = driver.next().await.unwrap();
let seen_descriptions = seen_descriptions.lock().unwrap();
assert_eq!(seen_descriptions.len(), 2);
assert_eq!(seen_descriptions[0], vec!["dynamic version 0".to_string()]);
assert_eq!(seen_descriptions[1], vec!["dynamic version 1".to_string()]);
let events = events.lock().unwrap();
assert!(events.iter().any(|event| matches!(
event,
AgentEvent::ToolCatalogChanged(ToolCatalogEvent {
source,
added,
removed,
changed,
}) if source == "mcp:mock"
&& added == &vec!["dynamic".to_string()]
&& removed.is_empty()
&& changed.is_empty()
)));
}
#[tokio::test]
async fn loop_passes_session_default_and_next_turn_cache_requests() {
let seen_caches = StdArc::new(StdMutex::new(Vec::new()));
let agent = Agent::builder()
.model(RecordingAdapter {
seen_descriptions: StdArc::new(StdMutex::new(Vec::new())),
seen_caches: seen_caches.clone(),
})
.permissions(AllowAllPermissions)
.build()
.unwrap();
let default_cache = PromptCacheRequest::best_effort(PromptCacheStrategy::Automatic)
.with_retention(PromptCacheRetention::Short);
let override_cache = PromptCacheRequest::required(PromptCacheStrategy::Explicit {
breakpoints: vec![PromptCacheBreakpoint::TranscriptItemEnd { index: 0 }],
});
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("session-cache"),
metadata: MetadataMap::new(),
cache: Some(default_cache.clone()),
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "first".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let _ = driver.next().await.unwrap();
driver
.submit_input_with_cache(
vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "second".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}],
override_cache.clone(),
)
.unwrap();
let _ = driver.next().await.unwrap();
let seen = seen_caches.lock().unwrap();
assert_eq!(seen.len(), 2);
assert_eq!(seen[0], Some(default_cache));
assert_eq!(seen[1], Some(override_cache));
}
#[tokio::test]
async fn loop_yields_after_tool_result_between_rounds() {
let tools = ToolRegistry::new().with(EchoTool::default());
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(tools)
.permissions(AllowAllPermissions)
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("yield-session"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item::text(ItemKind::User, "ping")])
.unwrap();
let step = driver.next().await.unwrap();
let info = match step {
LoopStep::Interrupt(LoopInterrupt::AfterToolResult(info)) => info,
other => panic!("expected AfterToolResult, got {other:?}"),
};
assert_eq!(info.session_id, SessionId::new("yield-session"));
assert_eq!(info.transcript_len, 3);
let interrupt = LoopInterrupt::AfterToolResult(info.clone());
assert!(!interrupt.is_blocking());
driver
.submit_input(vec![Item::text(ItemKind::User, "also: report back")])
.unwrap();
let step = driver.next().await.unwrap();
match step {
LoopStep::Finished(turn) => {
assert_eq!(turn.finish_reason, FinishReason::Completed);
}
other => panic!("expected Finished, got {other:?}"),
}
let snapshot = driver.snapshot();
let has_injected_message = snapshot.transcript.iter().any(|item| {
item.kind == ItemKind::User
&& item.parts.iter().any(|part| match part {
Part::Text(text) => text.text == "also: report back",
_ => false,
})
});
assert!(
has_injected_message,
"injected user message should be in transcript, got: {:?}",
snapshot.transcript
);
}
struct RecordingTranscriptObserver {
items: StdArc<StdMutex<Vec<Item>>>,
}
impl TranscriptObserver for RecordingTranscriptObserver {
fn on_transcript_event(&self, event: TranscriptEvent<'_>) {
self.items.lock().unwrap().push(event.item.clone());
}
}
#[tokio::test]
async fn observers_see_full_tool_round() {
let events = StdArc::new(StdMutex::new(Vec::<AgentEvent>::new()));
let items = StdArc::new(StdMutex::new(Vec::<Item>::new()));
let agent = Agent::builder()
.model(FakeAdapter)
.add_tool_source(ToolRegistry::new().with(EchoTool::default()))
.permissions(AllowAllPermissions)
.observer(RecordingObserver {
events: events.clone(),
})
.transcript_observer(RecordingTranscriptObserver {
items: items.clone(),
})
.build()
.unwrap();
let mut driver = agent
.start(SessionConfig {
session_id: SessionId::new("observer-session"),
metadata: MetadataMap::new(),
cache: None,
})
.await
.unwrap();
driver
.submit_input(vec![Item {
id: None,
kind: ItemKind::User,
parts: vec![Part::Text(TextPart {
text: "ping".into(),
metadata: MetadataMap::new(),
})],
metadata: MetadataMap::new(),
usage: None,
finish_reason: None,
created_at: None,
}])
.unwrap();
let result = run_until_finished(&mut driver).await;
assert!(matches!(result, LoopStep::Finished(_)), "got {result:?}");
let events = events.lock().unwrap().clone();
let tool_call_id = events.iter().find_map(|e| match e {
AgentEvent::ToolCallRequested(c) => Some(c.id.clone()),
_ => None,
});
let tool_results: Vec<_> = events
.iter()
.filter_map(|e| match e {
AgentEvent::ToolResultReceived(r) => Some(r.clone()),
_ => None,
})
.collect();
assert_eq!(tool_results.len(), 1, "events: {events:?}");
assert_eq!(Some(tool_results[0].call_id.clone()), tool_call_id);
assert!(!tool_results[0].is_error);
let items = items.lock().unwrap().clone();
assert_eq!(items.len(), 4, "items: {items:?}");
assert_eq!(items[0].kind, ItemKind::User);
assert_eq!(items[1].kind, ItemKind::Assistant);
assert!(
items[1]
.parts
.iter()
.any(|p| matches!(p, Part::ToolCall(_)))
);
assert_eq!(items[2].kind, ItemKind::Tool);
assert!(
items[2]
.parts
.iter()
.any(|p| matches!(p, Part::ToolResult(_)))
);
assert_eq!(items[3].kind, ItemKind::Assistant);
}
#[test]
fn convenience_cache_builders_construct_expected_defaults() {
let cache = PromptCacheRequest::automatic()
.with_retention(PromptCacheRetention::Short)
.with_key("workspace:demo");
let session = SessionConfig::new("demo").with_cache(cache.clone());
assert_eq!(session.session_id, SessionId::new("demo"));
assert_eq!(session.cache, Some(cache));
let explicit = PromptCacheRequest::explicit([
PromptCacheBreakpoint::tools_end(),
PromptCacheBreakpoint::transcript_item_end(2),
PromptCacheBreakpoint::transcript_part_end(3, 1),
]);
assert_eq!(explicit.mode, PromptCacheMode::BestEffort);
assert_eq!(
explicit.strategy,
PromptCacheStrategy::Explicit {
breakpoints: vec![
PromptCacheBreakpoint::ToolsEnd,
PromptCacheBreakpoint::TranscriptItemEnd { index: 2 },
PromptCacheBreakpoint::TranscriptPartEnd {
item_index: 3,
part_index: 1,
},
],
}
);
}
}