use super::*;
use crate::agent::{Agent, AgentOptions};
use theway_llm_provider::{
AssistantRole, ContentBlock, Message as PiMessage, StopReason, UserContent, UserMessage,
UserRole,
};
#[allow(dead_code)]
fn user_message(text: &str) -> AgentMessage {
AgentMessage::Llm(PiMessage::User(UserMessage {
role: UserRole::User,
content: UserContent::Text(text.into()),
timestamp: 0,
}))
}
fn assistant_message(content: Vec<ContentBlock>) -> AgentMessage {
AgentMessage::Llm(PiMessage::Assistant(theway_llm_provider::AssistantMessage {
role: AssistantRole::Assistant,
content,
api: theway_llm_provider::Api::from("faux"),
provider: theway_llm_provider::Provider::from("faux"),
model: "faux".into(),
response_model: None,
response_id: None,
diagnostics: None,
usage: theway_llm_provider::Usage::default(),
stop_reason: StopReason::Stop,
error_message: None,
timestamp: 0,
}))
}
fn agent() -> Agent {
Agent::new(AgentOptions::default())
}
#[tokio::test]
async fn run_agent_loop_continue_rejects_empty_transcript() {
let agent = agent();
let err = run_agent_loop_continue(agent.inner.clone()).await.unwrap_err();
assert!(err.to_string().contains("No messages to continue from"));
}
#[tokio::test]
async fn drive_loop_returns_ok_when_cancelled() {
let agent = agent();
let cancel = tokio_util::sync::CancellationToken::new();
cancel.cancel();
let result = drive_loop(&agent.inner.clone(), cancel).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn finalize_partial_turn_keeps_only_messages_with_content() {
let agent = agent();
let empty = assistant_message(Vec::new());
agent.state().streaming_message = Some(empty);
let cancel = tokio_util::sync::CancellationToken::new();
finalize_partial_turn(&agent.inner.clone(), &cancel).await;
assert!(agent.state().messages.is_empty());
assert!(agent.state().streaming_message.is_none());
let with_text = assistant_message(vec![ContentBlock::text("partial")]);
agent.state().streaming_message = Some(with_text);
finalize_partial_turn(&agent.inner.clone(), &cancel).await;
assert_eq!(agent.state().messages.len(), 1);
assert!(agent.state().streaming_message.is_none());
let thinking_only = assistant_message(vec![ContentBlock::Thinking(
theway_llm_provider::ThinkingContent {
thinking: "reasoning without any answer yet".into(),
..Default::default()
},
)]);
agent.state().streaming_message = Some(thinking_only);
finalize_partial_turn(&agent.inner.clone(), &cancel).await;
assert_eq!(agent.state().messages.len(), 1);
assert!(agent.state().streaming_message.is_none());
let with_thinking_and_text = assistant_message(vec![
ContentBlock::Thinking(theway_llm_provider::ThinkingContent {
thinking: "reasoning".into(),
..Default::default()
}),
ContentBlock::text("partial"),
]);
agent.state().streaming_message = Some(with_thinking_and_text);
finalize_partial_turn(&agent.inner.clone(), &cancel).await;
assert_eq!(agent.state().messages.len(), 2);
assert!(agent.state().streaming_message.is_none());
}
#[tokio::test]
async fn run_one_blocked_call_returns_error_outcome_without_tool() {
let inner = agent().inner.clone();
let call = PreparedCall::Blocked {
id: "call_1".into(),
name: "blocked".into(),
args: serde_json::json!({}),
result: AgentToolResult {
content: vec![theway_llm_provider::UserContentBlock::text("blocked")],
details: serde_json::Value::Null,
terminate: None,
},
};
let outcome = run_one(inner, call, tokio_util::sync::CancellationToken::new()).await;
assert!(outcome.is_error);
assert_eq!(outcome.id, "call_1");
assert_eq!(outcome.name, "blocked");
}
#[tokio::test]
async fn run_one_unknown_tool_returns_synthesized_error() {
let inner = agent().inner.clone();
let call = PreparedCall::Run {
id: "call_1".into(),
name: "missing".into(),
args: serde_json::json!({}),
tool: None,
};
let outcome = run_one(inner, call, tokio_util::sync::CancellationToken::new()).await;
assert!(outcome.is_error);
assert_eq!(outcome.name, "missing");
match &outcome.result.content[0] {
theway_llm_provider::UserContentBlock::Text(t) => {
assert!(t.text.contains("No tool registered named 'missing'"));
}
_ => panic!("expected text content"),
}
}
fn faux_model() -> theway_llm_provider::Model {
theway_llm_provider::Model {
id: "faux".into(),
name: "Faux".into(),
api: theway_llm_provider::Api::from("faux"),
provider: theway_llm_provider::Provider::from("faux"),
base_url: String::new(),
reasoning: false,
thinking_level_map: None,
input: vec![],
cost: theway_llm_provider::ModelCost::default(),
context_window: 128_000,
max_tokens: 16_384,
headers: None,
compat: None,
}
}
fn assistant_with_stop(
text: &str,
stop: theway_llm_provider::StopReason,
) -> theway_llm_provider::AssistantMessage {
theway_llm_provider::AssistantMessage {
role: theway_llm_provider::AssistantRole::Assistant,
content: vec![ContentBlock::text(text)],
api: theway_llm_provider::Api::from("faux"),
provider: theway_llm_provider::Provider::from("faux"),
model: "faux".into(),
response_model: None,
response_id: None,
diagnostics: None,
usage: theway_llm_provider::Usage::default(),
stop_reason: stop,
error_message: None,
timestamp: 0,
}
}
fn stream_that_returns(text: &'static str, stop: theway_llm_provider::StopReason) -> StreamFn {
Arc::new(move |_, _, _| {
let (stream, mut sender) = theway_llm_provider::AssistantMessageEventStream::new();
tokio::spawn(async move {
let msg = assistant_with_stop(text, stop);
sender.push(theway_llm_provider::AssistantMessageEvent::Start {
partial: msg.clone(),
});
sender.push(theway_llm_provider::AssistantMessageEvent::Done {
reason: match stop {
theway_llm_provider::StopReason::ToolUse => {
theway_llm_provider::DoneReason::ToolUse
}
_ => theway_llm_provider::DoneReason::Stop,
},
message: msg,
});
});
stream
})
}
fn inner_with_model_and_stream(stream: StreamFn) -> Arc<AgentInner> {
let mut state = AgentState::default();
state.model = Some(faux_model());
let agent = Agent::new(AgentOptions {
initial_state: Some(state),
stream_fn: Some(stream),
..Default::default()
});
agent.inner.clone()
}
#[tokio::test]
async fn run_agent_loop_appends_new_messages_and_runs() {
let inner = inner_with_model_and_stream(stream_that_returns(
"ok",
theway_llm_provider::StopReason::Stop,
));
run_agent_loop(
inner.clone(),
vec![user_message("one"), user_message("two")],
)
.await
.unwrap();
let messages = inner.state.lock().messages.clone();
assert_eq!(messages.len(), 3);
assert!(matches!(messages[0], AgentMessage::Llm(PiMessage::User(_))));
assert!(matches!(messages[1], AgentMessage::Llm(PiMessage::User(_))));
assert!(matches!(messages[2], AgentMessage::Llm(PiMessage::Assistant(_))));
}
#[tokio::test]
async fn concurrent_prompts_admit_exactly_one_run() {
let entered = Arc::new(tokio::sync::Notify::new());
let release = Arc::new(tokio::sync::Notify::new());
let stream: StreamFn = Arc::new({
let entered = entered.clone();
let release = release.clone();
move |_, _, _| {
let (stream, mut sender) = theway_llm_provider::AssistantMessageEventStream::new();
let entered = entered.clone();
let release = release.clone();
tokio::spawn(async move {
entered.notify_one();
release.notified().await;
let message = assistant_with_stop("ok", StopReason::Stop);
sender.push(theway_llm_provider::AssistantMessageEvent::Start {
partial: message.clone(),
});
sender.push(theway_llm_provider::AssistantMessageEvent::Done {
reason: theway_llm_provider::DoneReason::Stop,
message,
});
});
stream
}
});
let mut state = AgentState::default();
state.model = Some(faux_model());
let agent = Arc::new(Agent::new(AgentOptions {
initial_state: Some(state),
stream_fn: Some(stream),
..Default::default()
}));
let first = tokio::spawn({
let agent = agent.clone();
async move { agent.prompt(user_message("first")).await }
});
entered.notified().await;
let second = agent.prompt(user_message("second")).await.unwrap_err();
assert!(matches!(second, AgentRunError::AlreadyStreaming));
release.notify_one();
first.await.unwrap().unwrap();
assert!(!agent.is_streaming());
}
#[tokio::test]
async fn drive_loop_errors_when_max_iterations_exceeded() {
let mut inner = inner_with_model_and_stream(stream_that_returns(
"loop",
theway_llm_provider::StopReason::ToolUse,
));
Arc::get_mut(&mut inner).unwrap().max_iterations = Some(1);
let err = drive_loop(&inner, CancellationToken::new()).await.unwrap_err();
assert!(err.to_string().contains("max iterations (1) exceeded"));
assert!(inner
.state
.lock()
.error_message
.as_deref()
.unwrap()
.contains("max iterations"));
}
#[tokio::test]
async fn drive_loop_turn_interrupted_with_queued_steering_continues() {
let call_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let call_count_clone = call_count.clone();
let inner_holder: Arc<std::sync::Mutex<Option<Arc<AgentInner>>>> =
Arc::new(std::sync::Mutex::new(None));
let inner_holder_clone = inner_holder.clone();
let stream: StreamFn = Arc::new(move |_, _, _| {
let nth = call_count_clone.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if nth == 0 {
let holder = inner_holder_clone.clone();
let (stream, sender) = theway_llm_provider::AssistantMessageEventStream::new();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
if let Some(inner) = holder.lock().unwrap().as_ref() {
if let Some(token) = inner.turn_cancel.lock().clone() {
token.cancel();
}
}
drop(sender);
});
stream
} else {
let (stream, mut sender) = theway_llm_provider::AssistantMessageEventStream::new();
tokio::spawn(async move {
let msg = assistant_with_stop("done", theway_llm_provider::StopReason::Stop);
sender.push(theway_llm_provider::AssistantMessageEvent::Start {
partial: msg.clone(),
});
sender.push(theway_llm_provider::AssistantMessageEvent::Done {
reason: theway_llm_provider::DoneReason::Stop,
message: msg,
});
});
stream
}
});
let inner = inner_with_model_and_stream(stream);
*inner_holder.lock().unwrap() = Some(inner.clone());
inner.steering.lock().enqueue(user_message("steer"));
drive_loop(&inner, CancellationToken::new()).await.unwrap();
let messages = inner.state.lock().messages.clone();
assert!(messages.iter().any(|m| matches!(m, AgentMessage::Llm(PiMessage::User(u)) if matches!(&u.content, theway_llm_provider::UserContent::Text(t) if t == "steer"))));
assert!(matches!(
messages.last(),
Some(AgentMessage::Llm(PiMessage::Assistant(_)))
));
}
#[tokio::test]
async fn drive_loop_should_stop_after_turn_hook_stops() {
let mut inner = inner_with_model_and_stream(stream_that_returns(
"ok",
theway_llm_provider::StopReason::Stop,
));
Arc::get_mut(&mut inner).unwrap().options.should_stop_after_turn =
Some(Arc::new(|_ctx| Box::pin(async { true })));
drive_loop(&inner, CancellationToken::new()).await.unwrap();
assert_eq!(inner.state.lock().messages.len(), 1);
}
#[tokio::test]
async fn drive_loop_prepare_next_turn_applies_update() {
let mut inner = inner_with_model_and_stream(stream_that_returns(
"ok",
theway_llm_provider::StopReason::Stop,
));
Arc::get_mut(&mut inner).unwrap().options.prepare_next_turn = Some(Arc::new(|ctx| {
assert_eq!(ctx.message.stop_reason, theway_llm_provider::StopReason::Stop);
Box::pin(async move {
Some(AgentLoopTurnUpdate {
thinking_level: Some(ThinkingLevel::High),
..Default::default()
})
})
}));
drive_loop(&inner, CancellationToken::new()).await.unwrap();
assert_eq!(inner.state.lock().thinking_level, Some(ThinkingLevel::High));
}
struct RunOneTool {
ok: Option<AgentToolResult>,
err: Option<String>,
update: Option<AgentToolResult>,
def: theway_llm_provider::Tool,
}
#[async_trait::async_trait]
impl crate::types::AgentTool for RunOneTool {
fn definition(&self) -> &theway_llm_provider::Tool {
&self.def
}
fn label(&self) -> &str {
"run-one"
}
async fn execute(
&self,
_tool_call_id: &str,
_params: serde_json::Value,
_cancel: CancellationToken,
on_update: Option<AgentToolUpdate>,
) -> Result<AgentToolResult, AgentToolError> {
if let (Some(update), Some(on_update)) = (&self.update, on_update) {
on_update(update.clone());
}
if let Some(err) = &self.err {
return Err(AgentToolError::Message(err.clone()));
}
Ok(self
.ok
.clone()
.unwrap_or_default())
}
}
fn run_one_tool_ok() -> Arc<RunOneTool> {
Arc::new(RunOneTool {
ok: Some(AgentToolResult {
content: vec![UserContentBlock::text("ok")],
details: serde_json::Value::Null,
terminate: None,
}),
err: None,
update: None,
def: theway_llm_provider::Tool {
name: "run_one".into(),
description: String::new(),
parameters: serde_json::Value::Null,
},
})
}
fn run_one_tool_err() -> Arc<RunOneTool> {
Arc::new(RunOneTool {
ok: None,
err: Some("boom".into()),
update: None,
def: theway_llm_provider::Tool {
name: "run_one".into(),
description: String::new(),
parameters: serde_json::Value::Null,
},
})
}
#[tokio::test]
async fn run_one_with_tool_returns_success_and_error_outcomes() {
let ok = run_one(
inner_with_model_and_stream(stream_that_returns(
"unused",
theway_llm_provider::StopReason::Stop,
)),
PreparedCall::Run {
id: "c1".into(),
name: "run_one".into(),
args: serde_json::json!({}),
tool: Some(run_one_tool_ok()),
},
CancellationToken::new(),
)
.await;
assert!(!ok.is_error);
assert_eq!(ok.name, "run_one");
let err = run_one(
inner_with_model_and_stream(stream_that_returns(
"unused",
theway_llm_provider::StopReason::Stop,
)),
PreparedCall::Run {
id: "c2".into(),
name: "run_one".into(),
args: serde_json::json!({}),
tool: Some(run_one_tool_err()),
},
CancellationToken::new(),
)
.await;
assert!(err.is_error);
assert!(matches!(&err.result.content[0], UserContentBlock::Text(t) if t.text == "boom"));
}
#[tokio::test]
async fn run_one_with_tool_streams_update_events() {
let tool = Arc::new(RunOneTool {
ok: Some(AgentToolResult {
content: vec![UserContentBlock::text("done")],
details: serde_json::Value::Null,
terminate: None,
}),
err: None,
update: Some(AgentToolResult {
content: vec![UserContentBlock::text("partial")],
details: serde_json::Value::Null,
terminate: None,
}),
def: theway_llm_provider::Tool {
name: "run_one".into(),
description: String::new(),
parameters: serde_json::Value::Null,
},
});
let inner = inner_with_model_and_stream(stream_that_returns(
"unused",
theway_llm_provider::StopReason::Stop,
));
let mut rx = inner.broadcast_tx.subscribe();
let outcome = run_one(
inner.clone(),
PreparedCall::Run {
id: "c3".into(),
name: "run_one".into(),
args: serde_json::json!({}),
tool: Some(tool),
},
CancellationToken::new(),
)
.await;
assert!(!outcome.is_error);
let mut saw_update = false;
while let Ok(event) = rx.try_recv() {
if matches!(event, LoopEvent::ToolExecutionUpdate { .. }) {
saw_update = true;
}
}
assert!(saw_update, "ToolExecutionUpdate must be emitted for tool streaming");
}
#[tokio::test]
async fn finish_message_keeps_original_when_transform_changes_role() {
let mut agent = agent();
let original = user_message("keep me");
let replacement = assistant_message(vec![ContentBlock::text("changed role")]);
Arc::get_mut(&mut agent.inner)
.unwrap()
.options
.transform_message = Some(Arc::new(move |_message, _cancel| {
let replacement = replacement.clone();
Box::pin(async move { replacement })
}));
let cancel = tokio_util::sync::CancellationToken::new();
let finalized = finish_message(&agent.inner, original.clone(), &cancel).await;
assert!(matches!(
finalized,
AgentMessage::Llm(PiMessage::User(_))
));
let messages = agent.state().messages.clone();
assert_eq!(messages.len(), 1);
assert!(matches!(
messages[0],
AgentMessage::Llm(PiMessage::User(_))
));
}
#[tokio::test]
async fn drive_loop_prepare_next_turn_none_update_is_noop() {
let mut inner = inner_with_model_and_stream(stream_that_returns(
"ok",
theway_llm_provider::StopReason::Stop,
));
Arc::get_mut(&mut inner).unwrap().options.prepare_next_turn =
Some(Arc::new(|_ctx| Box::pin(async { None })));
drive_loop(&inner, CancellationToken::new()).await.unwrap();
assert_eq!(inner.state.lock().thinking_level, None);
}
#[tokio::test]
async fn drive_loop_tool_cancels_outer_token_finishes_cancelled() {
struct CancelOuterTool {
def: theway_llm_provider::Tool,
}
#[async_trait::async_trait]
impl crate::types::AgentTool for CancelOuterTool {
fn definition(&self) -> &theway_llm_provider::Tool {
&self.def
}
fn label(&self) -> &str {
"cancel_outer"
}
async fn execute(
&self,
_tool_call_id: &str,
_params: serde_json::Value,
cancel: CancellationToken,
_on_update: Option<crate::types::AgentToolUpdate>,
) -> Result<crate::types::AgentToolResult, crate::types::AgentToolError> {
cancel.cancel();
Ok(crate::types::AgentToolResult::default())
}
}
let stream: StreamFn = Arc::new(move |_, _, _| {
let (stream, mut sender) = theway_llm_provider::AssistantMessageEventStream::new();
tokio::spawn(async move {
let msg = theway_llm_provider::AssistantMessage {
role: theway_llm_provider::AssistantRole::Assistant,
content: vec![theway_llm_provider::ContentBlock::ToolCall(
theway_llm_provider::ToolCall {
id: "call_1".into(),
name: "cancel_outer".into(),
arguments: serde_json::Map::new(),
thought_signature: None,
},
)],
api: theway_llm_provider::Api::from("faux"),
provider: theway_llm_provider::Provider::from("faux"),
model: "faux".into(),
response_model: None,
response_id: None,
diagnostics: None,
usage: theway_llm_provider::Usage::default(),
stop_reason: theway_llm_provider::StopReason::ToolUse,
error_message: None,
timestamp: 0,
};
sender.push(theway_llm_provider::AssistantMessageEvent::Start {
partial: msg.clone(),
});
sender.push(theway_llm_provider::AssistantMessageEvent::Done {
reason: theway_llm_provider::DoneReason::ToolUse,
message: msg,
});
});
stream
});
let mut state = AgentState::default();
state.model = Some(faux_model());
state.tools = vec![Arc::new(CancelOuterTool {
def: theway_llm_provider::Tool {
name: "cancel_outer".into(),
description: String::new(),
parameters: serde_json::Value::Null,
},
})];
let agent = Agent::new(AgentOptions {
initial_state: Some(state),
stream_fn: Some(stream),
..Default::default()
});
let inner = agent.inner.clone();
drive_loop(&inner, CancellationToken::new()).await.unwrap();
assert!(!inner.state.lock().messages.is_empty());
}
#[tokio::test]
async fn drive_loop_tool_enqueues_steering_and_continues() {
struct SteeringTool {
inner: Arc<std::sync::Mutex<Option<Arc<AgentInner>>>>,
def: theway_llm_provider::Tool,
}
#[async_trait::async_trait]
impl crate::types::AgentTool for SteeringTool {
fn definition(&self) -> &theway_llm_provider::Tool {
&self.def
}
fn label(&self) -> &str {
"steer"
}
async fn execute(
&self,
_tool_call_id: &str,
_params: serde_json::Value,
_cancel: CancellationToken,
_on_update: Option<crate::types::AgentToolUpdate>,
) -> Result<crate::types::AgentToolResult, crate::types::AgentToolError> {
if let Some(inner) = self.inner.lock().unwrap().as_ref() {
inner.steering.lock().enqueue(user_message("steered"));
}
Ok(crate::types::AgentToolResult::default())
}
}
let inner_holder: Arc<std::sync::Mutex<Option<Arc<AgentInner>>>> =
Arc::new(std::sync::Mutex::new(None));
let tool = Arc::new(SteeringTool {
inner: inner_holder.clone(),
def: theway_llm_provider::Tool {
name: "steer".into(),
description: String::new(),
parameters: serde_json::Value::Null,
},
});
let stream_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let stream_calls_clone = stream_calls.clone();
let stream: StreamFn = Arc::new(move |_, _, _| {
let (stream, mut sender) = theway_llm_provider::AssistantMessageEventStream::new();
let nth = stream_calls_clone.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
tokio::spawn(async move {
let stop = if nth == 0 {
theway_llm_provider::StopReason::ToolUse
} else {
theway_llm_provider::StopReason::Stop
};
let msg = theway_llm_provider::AssistantMessage {
role: theway_llm_provider::AssistantRole::Assistant,
content: if nth == 0 {
vec![theway_llm_provider::ContentBlock::ToolCall(
theway_llm_provider::ToolCall {
id: "call_1".into(),
name: "steer".into(),
arguments: serde_json::Map::new(),
thought_signature: None,
},
)]
} else {
vec![theway_llm_provider::ContentBlock::text("ok")]
},
api: theway_llm_provider::Api::from("faux"),
provider: theway_llm_provider::Provider::from("faux"),
model: "faux".into(),
response_model: None,
response_id: None,
diagnostics: None,
usage: theway_llm_provider::Usage::default(),
stop_reason: stop,
error_message: None,
timestamp: 0,
};
sender.push(theway_llm_provider::AssistantMessageEvent::Start {
partial: msg.clone(),
});
sender.push(theway_llm_provider::AssistantMessageEvent::Done {
reason: match stop {
theway_llm_provider::StopReason::ToolUse => {
theway_llm_provider::DoneReason::ToolUse
}
_ => theway_llm_provider::DoneReason::Stop,
},
message: msg,
});
});
stream
});
let mut state = AgentState::default();
state.model = Some(faux_model());
state.tools = vec![tool];
let agent = Agent::new(AgentOptions {
initial_state: Some(state),
stream_fn: Some(stream),
..Default::default()
});
let inner = agent.inner.clone();
*inner_holder.lock().unwrap() = Some(inner.clone());
drive_loop(&inner, CancellationToken::new()).await.unwrap();
let messages = inner.state.lock().messages.clone();
assert!(messages.iter().any(|m| matches!(m, AgentMessage::Llm(PiMessage::User(u))
if matches!(&u.content, theway_llm_provider::UserContent::Text(t) if t == "steered"))));
}
#[tokio::test]
async fn run_one_with_cancelled_token_marks_cancelled() {
let inner = agent().inner.clone();
let cancel = CancellationToken::new();
cancel.cancel();
let call = PreparedCall::Blocked {
id: "call_1".into(),
name: "blocked".into(),
args: serde_json::json!({}),
result: AgentToolResult {
content: vec![theway_llm_provider::UserContentBlock::text("blocked")],
details: serde_json::Value::Null,
terminate: None,
},
};
let outcome = run_one(inner, call, cancel).await;
assert!(outcome.is_error);
}
#[tokio::test]
async fn drive_loop_steering_with_stop_reason_still_continues_once() {
struct SteeringStopTool {
inner: Arc<std::sync::Mutex<Option<Arc<AgentInner>>>>,
def: theway_llm_provider::Tool,
}
#[async_trait::async_trait]
impl crate::types::AgentTool for SteeringStopTool {
fn definition(&self) -> &theway_llm_provider::Tool {
&self.def
}
fn label(&self) -> &str {
"steer"
}
async fn execute(
&self,
_tool_call_id: &str,
_params: serde_json::Value,
_cancel: CancellationToken,
_on_update: Option<crate::types::AgentToolUpdate>,
) -> Result<crate::types::AgentToolResult, crate::types::AgentToolError> {
if let Some(inner) = self.inner.lock().unwrap().as_ref() {
inner.steering.lock().enqueue(user_message("steered"));
}
Ok(crate::types::AgentToolResult::default())
}
}
let inner_holder: Arc<std::sync::Mutex<Option<Arc<AgentInner>>>> =
Arc::new(std::sync::Mutex::new(None));
let tool = Arc::new(SteeringStopTool {
inner: inner_holder.clone(),
def: theway_llm_provider::Tool {
name: "steer".into(),
description: String::new(),
parameters: serde_json::Value::Null,
},
});
let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let calls_clone = calls.clone();
let stream: StreamFn = Arc::new(move |_, _, _| {
let (stream, mut sender) = theway_llm_provider::AssistantMessageEventStream::new();
let nth = calls_clone.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
tokio::spawn(async move {
let msg = if nth == 0 {
theway_llm_provider::AssistantMessage {
role: theway_llm_provider::AssistantRole::Assistant,
content: vec![theway_llm_provider::ContentBlock::ToolCall(
theway_llm_provider::ToolCall {
id: "call_1".into(),
name: "steer".into(),
arguments: serde_json::Map::new(),
thought_signature: None,
},
)],
api: theway_llm_provider::Api::from("faux"),
provider: theway_llm_provider::Provider::from("faux"),
model: "faux".into(),
response_model: None,
response_id: None,
diagnostics: None,
usage: theway_llm_provider::Usage::default(),
stop_reason: theway_llm_provider::StopReason::Stop,
error_message: None,
timestamp: 0,
}
} else {
assistant_with_stop("ok", theway_llm_provider::StopReason::Stop)
};
sender.push(theway_llm_provider::AssistantMessageEvent::Start {
partial: msg.clone(),
});
sender.push(theway_llm_provider::AssistantMessageEvent::Done {
reason: theway_llm_provider::DoneReason::Stop,
message: msg,
});
});
stream
});
let mut state = AgentState::default();
state.model = Some(faux_model());
state.tools = vec![tool];
let agent = Agent::new(AgentOptions {
initial_state: Some(state),
stream_fn: Some(stream),
..Default::default()
});
let inner = agent.inner.clone();
*inner_holder.lock().unwrap() = Some(inner.clone());
drive_loop(&inner, CancellationToken::new()).await.unwrap();
let messages = inner.state.lock().messages.clone();
assert!(messages.iter().any(|m| matches!(m, AgentMessage::Llm(PiMessage::User(u))
if matches!(&u.content, theway_llm_provider::UserContent::Text(t) if t == "steered"))));
}
#[tokio::test(start_paused = true)]
async fn run_one_pump_join_timeout_aborts_pump_when_tool_retains_update() {
static RETAINED_UPDATE: std::sync::Mutex<Option<AgentToolUpdate>> = std::sync::Mutex::new(None);
struct RetainTool {
def: theway_llm_provider::Tool,
}
#[async_trait::async_trait]
impl crate::types::AgentTool for RetainTool {
fn definition(&self) -> &theway_llm_provider::Tool {
&self.def
}
fn label(&self) -> &str {
"retain"
}
async fn execute(
&self,
_tool_call_id: &str,
_params: serde_json::Value,
_cancel: CancellationToken,
on_update: Option<crate::types::AgentToolUpdate>,
) -> Result<crate::types::AgentToolResult, crate::types::AgentToolError> {
*RETAINED_UPDATE.lock().unwrap() = on_update;
Ok(crate::types::AgentToolResult::default())
}
}
let inner = agent().inner.clone();
let tool = Arc::new(RetainTool {
def: theway_llm_provider::Tool {
name: "retain".into(),
description: String::new(),
parameters: serde_json::Value::Null,
},
});
let outcome = run_one(
inner,
PreparedCall::Run {
id: "call_1".into(),
name: "retain".into(),
args: serde_json::json!({}),
tool: Some(tool),
},
CancellationToken::new(),
)
.await;
assert!(!outcome.is_error);
*RETAINED_UPDATE.lock().unwrap() = None;
}