use crate::api::ApiClient;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::cancel::CancelSignal;
use crate::compact::ContextManager;
use crate::config::LoopConfig;
use crate::error::LoopError;
use crate::engine::loop_core::{LoopState, SessionResult, StopReason, ToolCall, TurnResult};
#[cfg(all(test, feature = "hooks"))]
use crate::hooks::Hook;
#[cfg(feature = "hooks")]
use crate::hooks::context::{
CompactTrigger, PostCompactContext, PostToolUseContext, PreCompactContext, PreToolUseContext,
};
#[cfg(all(test, feature = "hooks"))]
use crate::hooks::context::{SessionEndContext as HookSessionEndContext, SessionEndReason};
#[cfg(feature = "hooks")]
use crate::hooks::{HookAction, HookExecutor};
use crate::message::{Message, MessagePart, Role, ToolContent};
use crate::middleware::{ToolDispatchContext, ToolPipeline, ToolPipelineBuilder};
use crate::observer::{
FallbackContext, ModelSwitchedContext, ResponseContext, StreamContext, StreamFailureContext,
TurnEndContext, TurnStartContext,
};
use crate::reflection::{
ExponentialBackoffRecovery, NoopReflector, RecoveryAction, RecoveryStrategy, ReflectionContext,
Reflector,
};
use crate::runtime::LoopRuntime;
use crate::stream::handler::StreamHandler;
use crate::stream::{StreamAccumulator, StreamEvent, StreamStopReason, Usage};
#[cfg(feature = "tool_health")]
use crate::tool::health::ToolHealthRegistry;
use crate::tool::{PermissionCheck, ToolContext, ToolDispatchResult, ToolRegistry, ToolSchema};
mod compact;
mod dispatch;
mod emission;
mod message;
mod stream;
pub struct BareLoop<C: ApiClient> {
client: Arc<C>,
tools: Arc<ToolRegistry>,
config: LoopConfig,
conversation: Vec<Message>,
managers: LoopRuntime,
reflector: Arc<dyn Reflector>,
recovery: Arc<dyn RecoveryStrategy>,
cancelled: Arc<CancelSignal>,
state: LoopState,
budget: SessionResult,
session_start: Option<Instant>,
#[allow(clippy::type_complexity)]
text_streamer: Option<Arc<dyn Fn(&str) + Send + Sync>>,
}
impl<C: ApiClient> BareLoop<C> {
const MAX_RECOVERY_ATTEMPTS: u32 = 5;
pub fn new(client: Arc<C>, tools: ToolRegistry, config: LoopConfig) -> Self {
Self {
client,
tools: Arc::new(tools),
config,
conversation: Vec::new(),
managers: LoopRuntime::new(),
reflector: Arc::new(NoopReflector),
recovery: Arc::new(ExponentialBackoffRecovery::new(3)),
cancelled: Arc::new(CancelSignal::new()),
state: LoopState::Idle,
budget: SessionResult::default(),
session_start: None,
text_streamer: None,
}
}
pub fn new_with_managers(
client: Arc<C>,
tools: ToolRegistry,
config: LoopConfig,
managers: LoopRuntime,
) -> Self {
Self {
client,
tools: Arc::new(tools),
config,
conversation: Vec::new(),
managers,
reflector: Arc::new(NoopReflector),
recovery: Arc::new(ExponentialBackoffRecovery::new(3)),
cancelled: Arc::new(CancelSignal::new()),
state: LoopState::Idle,
budget: SessionResult::default(),
session_start: None,
text_streamer: None,
}
}
pub fn conversation(&self) -> &[Message] {
&self.conversation
}
pub fn config(&self) -> &LoopConfig {
&self.config
}
pub fn tools(&self) -> &ToolRegistry {
&self.tools
}
pub fn is_cancelled(&self) -> bool {
self.cancelled.is_cancelled()
}
pub fn cancel(&self) {
self.cancelled.cancel();
}
pub fn cancel_signal(&self) -> Arc<CancelSignal> {
Arc::clone(&self.cancelled)
}
#[inline]
fn debug_assert_idle(&self) {
debug_assert!(
matches!(self.state, LoopState::Idle),
"BareLoop configuration setters must be called before run() — \
current state is {:?}, expected Idle",
self.state
);
}
pub fn set_reflector(&mut self, reflector: Arc<dyn Reflector>) {
self.debug_assert_idle();
self.reflector = reflector;
}
pub fn set_recovery_strategy(&mut self, strategy: Arc<dyn RecoveryStrategy>) {
self.debug_assert_idle();
self.recovery = strategy;
}
pub fn set_context_manager(&mut self, manager: Arc<ContextManager>) {
self.debug_assert_idle();
self.managers.set_context_manager(manager);
}
pub fn set_stream_handler(&mut self, handler: StreamHandler) {
self.debug_assert_idle();
self.managers.set_stream_handler(handler);
}
#[cfg(feature = "hooks")]
pub fn set_hook_executor(&mut self, executor: Arc<HookExecutor>) {
self.debug_assert_idle();
self.managers.set_hook_executor(executor);
}
#[cfg(feature = "tool_health")]
pub fn set_health_registry(&mut self, registry: Arc<ToolHealthRegistry>) {
self.debug_assert_idle();
self.managers.set_health_registry(registry);
}
pub fn set_pipeline(&mut self, builder: ToolPipelineBuilder) -> Result<(), LoopError> {
self.debug_assert_idle();
let pipeline = builder
.core(Arc::clone(&self.tools))
.build()
.map_err(|e| LoopError::Config(e.to_string()))?;
self.managers.set_pipeline(pipeline);
Ok(())
}
pub fn register_observer(&mut self, observer: Arc<dyn crate::observer::LoopObserver>) {
self.debug_assert_idle();
self.managers.register_observer(observer);
}
pub fn set_text_streamer(&mut self, f: Arc<dyn Fn(&str) + Send + Sync>) {
self.debug_assert_idle();
self.text_streamer = Some(f);
}
pub fn switch_model(&mut self, model: &str) -> ModelSwitch<'_, C> {
ModelSwitch {
loop_: self,
target_model: model.to_string(),
context_window: None,
max_tokens: None,
}
}
fn usage_tokens(usage: Option<&Usage>) -> (u64, u64) {
match usage {
Some(u) => (u64::from(u.input_tokens), u64::from(u.output_tokens)),
None => (0, 0),
}
}
async fn dispatch_and_record(
&mut self,
tool_calls: &[ToolCall],
turn_index: usize,
turn_duration: Duration,
turn_input_tokens: u64,
turn_output_tokens: u64,
budget: &mut SessionResult,
) -> Result<(), LoopError> {
match self.dispatch_tools(tool_calls, turn_index).await {
Ok(results) => {
budget.tool_calls = budget.tool_calls.saturating_add(results.len());
let tool_result_msg = Self::build_tool_result_message(results);
self.conversation.push(tool_result_msg);
self.managers.observers().on_turn_end(&TurnEndContext {
turn: turn_index,
success: true,
error: None,
duration_ms: Self::millis_u64(turn_duration),
input_tokens: turn_input_tokens,
output_tokens: turn_output_tokens,
});
Ok(())
}
Err(e) => {
let err_str = e.to_string();
self.managers.observers().on_turn_end(&TurnEndContext {
turn: turn_index,
success: false,
error: Some(err_str),
duration_ms: Self::millis_u64(turn_duration),
input_tokens: turn_input_tokens,
output_tokens: turn_output_tokens,
});
Err(e)
}
}
}
fn check_cancellation(&mut self) -> Result<(), LoopError> {
if self.is_cancelled() {
self.state = LoopState::Failed {
error: "cancelled".into(),
};
return Err(LoopError::Cancelled);
}
Ok(())
}
async fn do_stream(&mut self) -> Result<(Message, Option<Usage>, StreamStopReason), LoopError> {
match self.stream_turn().await {
Ok((msg, usage, stop)) => {
self.managers.fallback.record_model_success();
let (in_tok, out_tok) = Self::usage_tokens(usage.as_ref());
self.managers.observers().on_stream_success(&StreamContext {
turn: self.budget.total_turns,
model: self.client.model(),
input_tokens: in_tok,
output_tokens: out_tok,
});
Ok((msg, usage, stop))
}
Err(e) => {
let tripped = self.managers.fallback.record_api_failure();
if tripped {
let from = self.client.model();
if let Some(to) = self.managers.fallback.fallback_model() {
tracing::warn!(from = %from, to = %to, "fallback manager tripped");
self.managers
.observers()
.on_fallback(&FallbackContext { from, to });
}
}
self.managers
.observers()
.on_stream_failure(&StreamFailureContext {
turn: self.budget.total_turns,
model: self.client.model(),
error: e.clone(),
});
self.state = LoopState::Failed {
error: e.to_string(),
};
Err(e)
}
}
}
fn accumulate_usage(&mut self, usage: Option<&Usage>) {
if let Some(u) = usage {
self.budget.input_tokens = self
.budget
.input_tokens
.saturating_add(u64::from(u.input_tokens));
self.budget.output_tokens = self
.budget
.output_tokens
.saturating_add(u64::from(u.output_tokens));
}
}
fn finish_turn(&mut self, turn_in: u64, turn_out: u64, duration: Duration) {
self.managers.observers().on_turn_end(&TurnEndContext {
turn: self.budget.total_turns.saturating_sub(1),
success: true,
error: None,
duration_ms: Self::millis_u64(duration),
input_tokens: turn_in,
output_tokens: turn_out,
});
}
fn turn_complete(text: String, turn_in: u64, turn_out: u64, duration: Duration) -> TurnResult {
TurnResult {
text,
tool_calls: Vec::new(),
tool_results: Vec::new(),
input_tokens: turn_in,
output_tokens: turn_out,
duration,
is_complete: true,
stop_reason: StopReason::EndTurn,
}
}
async fn try_compact_context(&mut self) {
if let Err(e) = self.maybe_compact_context(self.budget.total_turns).await {
tracing::warn!(
error = %e,
turn = self.budget.total_turns,
"context compaction failed; continuing with uncompactd history"
);
}
}
}
pub struct ModelSwitch<'a, C: ApiClient> {
loop_: &'a mut BareLoop<C>,
target_model: String,
context_window: Option<u64>,
max_tokens: Option<u32>,
}
impl<C: ApiClient> ModelSwitch<'_, C> {
#[must_use]
pub fn context_window(mut self, tokens: u64) -> Self {
self.context_window = Some(tokens);
self
}
#[must_use]
pub fn max_tokens(mut self, tokens: u32) -> Self {
self.max_tokens = Some(tokens);
self
}
pub fn apply(self) -> Result<(), LoopError> {
let Self {
loop_,
target_model,
context_window,
max_tokens,
} = self;
let trimmed = target_model.trim();
if trimmed.is_empty() {
return Err(LoopError::Config(
"model name must not be empty or whitespace".into(),
));
}
let from = loop_.config.model.clone();
loop_.client.set_model(trimmed);
loop_.config.model = trimmed.to_string();
if let Some(cw) = context_window {
loop_.config.context_window = cw;
}
if let Some(mt) = max_tokens {
loop_.config.max_tokens = mt;
}
loop_.managers.fallback.reset();
loop_
.managers
.fallback
.set_original_model(trimmed.to_string());
loop_
.managers
.observers()
.on_model_switched(&ModelSwitchedContext {
from,
to: trimmed.to_string(),
});
Ok(())
}
}
impl<C: ApiClient> crate::engine::loop_core::Loop for BareLoop<C> {
fn initialize<'a>(
&'a mut self,
config: &'a crate::config::LoopConfig,
) -> Pin<Box<dyn Future<Output = Result<(), LoopError>> + Send + 'a>> {
Box::pin(async move {
config.validate()?;
self.state = LoopState::Processing { turn: 0 };
self.budget = SessionResult::default();
self.session_start = Some(Instant::now());
self.config = config.clone();
self.managers.reset_all();
self.notify_session_start();
Ok(())
})
}
fn process_turn<'a>(
&'a mut self,
input: &'a str,
) -> Pin<Box<dyn Future<Output = Result<TurnResult, LoopError>> + Send + 'a>> {
Box::pin(async move {
if !input.is_empty() {
self.conversation.push(Message::user(input));
}
let turn_start = Instant::now();
self.managers.observers().on_turn_start(&TurnStartContext {
turn: self.budget.total_turns,
query: input.to_string(),
});
self.check_cancellation()?;
let (assistant_msg, usage, _stream_stop) = self.do_stream().await?;
self.accumulate_usage(usage.as_ref());
let text = Self::extract_text(&assistant_msg);
let (turn_in, turn_out) = Self::usage_tokens(usage.as_ref());
let pattern = self.managers.detection.record_response(&text);
self.managers.observers().on_response(&ResponseContext {
turn: self.budget.total_turns,
text: text.clone(),
usage,
});
if let Some(result) = self
.managers
.handle_detected_pattern(&pattern, self.budget.total_turns)
{
return match result {
Ok(_) => {
self.state = LoopState::Completed {
summary: text.clone(),
};
Ok(Self::turn_complete(
text,
turn_in,
turn_out,
turn_start.elapsed(),
))
}
Err(e) => {
self.state = LoopState::Failed {
error: e.to_string(),
};
Err(e)
}
};
}
let tool_calls = Self::extract_tool_calls(&assistant_msg);
self.conversation.push(assistant_msg);
self.budget.total_turns = self.budget.total_turns.saturating_add(1);
if tool_calls.is_empty() {
self.finish_turn(turn_in, turn_out, turn_start.elapsed());
self.state = LoopState::Completed {
summary: text.clone(),
};
return Ok(Self::turn_complete(
text,
turn_in,
turn_out,
turn_start.elapsed(),
));
}
self.state = LoopState::WaitingForTool {
tool: tool_calls
.first()
.map(|tc| tc.tool.clone())
.unwrap_or_default(),
started_at: std::time::SystemTime::now(),
};
let mut budget = std::mem::take(&mut self.budget);
let turn_index = budget.total_turns.saturating_sub(1);
let turn_duration = turn_start.elapsed();
if let Err(e) = self
.dispatch_and_record(
&tool_calls,
turn_index,
turn_duration,
turn_in,
turn_out,
&mut budget,
)
.await
{
self.budget = budget;
self.state = LoopState::Failed {
error: e.to_string(),
};
return Err(e);
}
self.budget = budget;
self.try_compact_context().await;
self.state = LoopState::Processing {
turn: self.budget.total_turns,
};
Ok(TurnResult {
text,
tool_calls,
tool_results: Vec::new(),
input_tokens: turn_in,
output_tokens: turn_out,
duration: turn_start.elapsed(),
is_complete: false,
stop_reason: StopReason::ToolCall,
})
})
}
fn should_continue(&self) -> bool {
if self.is_cancelled() {
return false;
}
self.budget.total_turns < self.config.max_turns
}
fn finalize<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<SessionResult, LoopError>> + Send + 'a>> {
Box::pin(async move {
let duration = self.session_start.map(|s| s.elapsed()).unwrap_or_default();
let success = !matches!(self.state, LoopState::Failed { .. });
self.budget.success = success;
self.budget.session_id = self.config.session_id;
self.budget.total_duration = duration;
if success {
self.budget.final_output = match &self.state {
LoopState::Completed { summary } => Some(summary.clone()),
_ => Some(String::new()),
};
} else {
self.budget.error = Some(match &self.state {
LoopState::Failed { error } => error.clone(),
_ => "session failed".to_string(),
});
}
self.notify_session_end(&self.budget, duration);
Ok(self.budget.clone())
})
}
fn state(&self) -> LoopState {
self.state.clone()
}
fn cancel(&self) {
BareLoop::cancel(self);
}
fn stop_reason(&self) -> Option<LoopError> {
if self.is_cancelled() {
return Some(LoopError::Cancelled);
}
if self.budget.total_turns >= self.config.max_turns {
return Some(LoopError::MaxTurnsExceeded {
max: self.config.max_turns,
});
}
None
}
fn config(&self) -> &LoopConfig {
&self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::api::error::ApiError;
use crate::engine::loop_core::Loop;
use crate::stream::{
DeltaPart, IndexedDelta, MessageDelta, MessageDeltaPayload, MessageMetadata, MessageStart,
PartStart, Usage,
};
use crate::tool::ToolRegistry;
use crate::tool::{Tool, ToolContext, ToolError, ToolOutput, ToolSchema};
use serde_json::{Value, json};
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicUsize, Ordering};
use parking_lot::Mutex;
#[derive(Clone)]
struct MockClient {
responses: Arc<Mutex<Vec<Vec<StreamEvent>>>>,
model_name: Arc<parking_lot::Mutex<String>>,
}
impl MockClient {
fn new(model: &str) -> Self {
Self {
responses: Arc::new(Mutex::new(Vec::new())),
model_name: Arc::new(parking_lot::Mutex::new(model.to_string())),
}
}
fn add_text_response(&self, text: &str) {
let events = vec![
StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_test".into(),
role: "assistant".into(),
model: self.model_name.lock().clone(),
},
}),
StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::text(text)),
}),
StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: text.to_string(),
},
}),
StreamEvent::PartStop,
StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("end_turn".to_string()),
},
usage: Some(Usage::new(10, 20)),
}),
StreamEvent::MessageStop,
];
self.responses.lock().push(events);
}
fn add_events(&self, events: Vec<StreamEvent>) {
self.responses.lock().push(events);
}
fn add_tool_then_text(
&self,
tool_id: &str,
tool_name: &str,
tool_input: Value,
final_text: &str,
) {
let tool_events = vec![
StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_tool".into(),
role: "assistant".into(),
model: self.model_name.lock().clone(),
},
}),
StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call(tool_id, tool_name, tool_input)),
}),
StreamEvent::PartStop,
StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("tool_call".to_string()),
},
usage: Some(Usage::new(50, 10)),
}),
StreamEvent::MessageStop,
];
self.responses.lock().push(tool_events);
let text_events = vec![
StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_final".into(),
role: "assistant".into(),
model: self.model_name.lock().clone(),
},
}),
StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::text(final_text)),
}),
StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: final_text.to_string(),
},
}),
StreamEvent::PartStop,
StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("end_turn".to_string()),
},
usage: Some(Usage::new(30, 15)),
}),
StreamEvent::MessageStop,
];
self.responses.lock().push(text_events);
}
fn add_tool_only_response(&self, tool_id: &str, tool_name: &str, tool_input: Value) {
let tool_events = vec![
StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: format!("msg_{tool_id}"),
role: "assistant".into(),
model: self.model_name.lock().clone(),
},
}),
StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call(tool_id, tool_name, tool_input)),
}),
StreamEvent::PartStop,
StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("tool_call".to_string()),
},
usage: Some(Usage::new(50, 10)),
}),
StreamEvent::MessageStop,
];
self.responses.lock().push(tool_events);
}
#[expect(dead_code)]
fn add_error_response(&self) {
let events = vec![StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_err".into(),
role: "assistant".into(),
model: self.model_name.lock().clone(),
},
})];
self.responses.lock().push(events);
}
}
impl ApiClient for MockClient {
fn model(&self) -> String {
self.model_name.lock().clone()
}
fn set_model(&self, model: &str) -> bool {
if model.trim().is_empty() {
return false;
}
*self.model_name.lock() = model.to_string();
true
}
fn stream_messages(
&self,
_messages: Vec<Message>,
_system: Option<String>,
_tools: Option<Vec<ToolSchema>>,
) -> Pin<Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
let mut guard = self.responses.lock();
if let Some(events) = guard.pop_front() {
let events: Vec<Result<StreamEvent, ApiError>> =
events.into_iter().map(Ok).collect();
Box::pin(futures::stream::iter(events))
} else {
let err = ApiError::api("No more mock responses");
Box::pin(futures::stream::iter(vec![Err(err)]))
}
}
fn create_message(
&self,
_messages: Vec<Message>,
_system: Option<String>,
_tools: Option<Vec<ToolSchema>>,
) -> Pin<Box<dyn Future<Output = Result<Value, ApiError>> + Send + '_>> {
Box::pin(async { Ok(json!({"content": []})) })
}
}
trait PopFront<T> {
fn pop_front(&mut self) -> Option<T>;
}
impl<T> PopFront<T> for Vec<T> {
fn pop_front(&mut self) -> Option<T> {
if self.is_empty() {
None
} else {
Some(self.remove(0))
}
}
}
struct EchoTool;
impl Tool for EchoTool {
fn name(&self) -> &'static str {
"echo"
}
fn description(&self) -> &'static str {
"Echoes back the input"
}
fn schema(&self) -> ToolSchema {
ToolSchema {
tool: "echo".into(),
description: "Echoes back the input".into(),
input_schema: json!({
"type": "object",
"properties": { "message": { "type": "string" } },
"required": ["message"]
}),
}
}
fn call(
&self,
input: Value,
_ctx: &ToolContext,
) -> Pin<Box<dyn Future<Output = Result<ToolOutput, ToolError>> + Send + '_>> {
let msg = input
.get("message")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
Box::pin(async move { Ok(ToolOutput::text(format!("Echo: {msg}"))) })
}
}
struct FailingTool;
impl Tool for FailingTool {
fn name(&self) -> &'static str {
"fail"
}
fn description(&self) -> &'static str {
"Always fails"
}
fn schema(&self) -> ToolSchema {
ToolSchema {
tool: "fail".into(),
description: "Always fails".into(),
input_schema: json!({ "type": "object", "properties": {} }),
}
}
fn call(
&self,
_input: Value,
_ctx: &ToolContext,
) -> Pin<Box<dyn Future<Output = Result<ToolOutput, ToolError>> + Send + '_>> {
Box::pin(async move { Err(ToolError::Execution("Tool intentionally failed".into())) })
}
}
struct CountingObserver {
session_starts: AtomicUsize,
session_ends: AtomicUsize,
turn_starts: AtomicUsize,
turn_ends: AtomicUsize,
tool_pres: AtomicUsize,
tool_posts: AtomicUsize,
}
impl CountingObserver {
fn new() -> Self {
Self {
session_starts: AtomicUsize::new(0),
session_ends: AtomicUsize::new(0),
turn_starts: AtomicUsize::new(0),
turn_ends: AtomicUsize::new(0),
tool_pres: AtomicUsize::new(0),
tool_posts: AtomicUsize::new(0),
}
}
}
impl crate::observer::LoopObserver for CountingObserver {
fn name(&self) -> &'static str {
"counting"
}
fn on_session_start(&self, _ctx: &crate::observer::SessionStartContext) {
self.session_starts.fetch_add(1, Ordering::SeqCst);
}
fn on_session_end(&self, _ctx: &crate::observer::SessionEndContext) {
self.session_ends.fetch_add(1, Ordering::SeqCst);
}
fn on_turn_start(&self, _ctx: &crate::observer::TurnStartContext) {
self.turn_starts.fetch_add(1, Ordering::SeqCst);
}
fn on_turn_end(&self, _ctx: &crate::observer::TurnEndContext) {
self.turn_ends.fetch_add(1, Ordering::SeqCst);
}
fn on_tool_pre(&self, _ctx: &crate::observer::ToolPreContext) {
self.tool_pres.fetch_add(1, Ordering::SeqCst);
}
fn on_tool_post(&self, _ctx: &crate::observer::ToolPostContext) {
self.tool_posts.fetch_add(1, Ordering::SeqCst);
}
}
fn make_config() -> LoopConfig {
LoopConfig {
max_turns: 10,
..Default::default()
}
}
#[tokio::test]
async fn test_bare_loop_single_turn() {
let client = MockClient::new("test-model");
client.add_text_response("Hello! I'm done.");
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let result = agent.run("Hi").await.unwrap();
assert!(result.success);
assert_eq!(result.total_turns, 1);
assert_eq!(result.final_output.as_deref(), Some("Hello! I'm done."));
}
#[tokio::test]
async fn test_bare_loop_with_tool_call() {
let client = MockClient::new("test-model");
client.add_tool_then_text(
"tool_1",
"echo",
json!({"message": "hello"}),
"I echoed your message.",
);
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), registry, config);
let result = agent.run("Echo hello").await.unwrap();
assert!(result.success);
assert_eq!(result.total_turns, 2); assert_eq!(result.tool_calls, 1);
}
#[tokio::test]
async fn test_bare_loop_max_turns_exceeded() {
let client = MockClient::new("test-model");
for i in 0..20 {
client.add_tool_only_response(
&format!("tool_{i}"),
"echo",
json!({"message": format!("msg_{i}")}),
);
}
let mut config = make_config();
config.max_turns = 3;
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let mut agent = BareLoop::new(Arc::new(client), registry, config);
let result = agent.run("Keep going").await;
assert!(result.is_err());
match result.unwrap_err() {
LoopError::MaxTurnsExceeded { max } => assert_eq!(max, 3),
other => panic!("Expected MaxTurnsExceeded, got: {other}"),
}
}
#[tokio::test]
async fn test_bare_loop_cancellation() {
let client = MockClient::new("test-model");
client.add_text_response("Hello!");
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
agent.cancel();
assert!(agent.is_cancelled());
let result = agent.run("Hi").await;
assert!(result.is_err());
match result.unwrap_err() {
LoopError::Cancelled => {}
other => panic!("Expected Cancelled error, got: {other}"),
}
}
#[tokio::test]
async fn test_bare_loop_api_error() {
let client = MockClient::new("test-model");
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let result = agent.run("Hi").await;
assert!(result.is_err());
match result.unwrap_err() {
LoopError::Api(msg) => assert!(msg.contains("No more mock responses")),
other => panic!("Expected Api error, got: {other}"),
}
}
#[tokio::test]
async fn test_tool_not_found_returns_error_result() {
let client = MockClient::new("test-model");
client.add_tool_then_text("tool_1", "nonexistent", json!({}), "I see the tool failed.");
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let result = agent.run("Use nonexistent tool").await.unwrap();
assert!(result.success);
assert_eq!(result.total_turns, 2);
}
#[tokio::test]
async fn test_tool_execution_failure() {
let client = MockClient::new("test-model");
client.add_tool_then_text("tool_1", "fail", json!({}), "The tool failed, moving on.");
let mut registry = ToolRegistry::new();
registry.register(FailingTool);
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), registry, config);
let result = agent.run("Use failing tool").await.unwrap();
assert!(result.success);
assert_eq!(result.total_turns, 2);
}
#[tokio::test]
async fn test_observer_lifecycle_events() {
let client = MockClient::new("test-model");
client.add_text_response("Done!");
let plugin = Arc::new(CountingObserver::new());
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
agent.register_observer(plugin.clone());
let result = agent.run("Hi").await.unwrap();
assert!(result.success);
assert_eq!(plugin.session_starts.load(Ordering::SeqCst), 1);
assert_eq!(plugin.session_ends.load(Ordering::SeqCst), 1);
assert_eq!(plugin.turn_starts.load(Ordering::SeqCst), 1);
assert_eq!(plugin.turn_ends.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_observer_tool_events() {
let client = MockClient::new("test-model");
client.add_tool_then_text("tool_1", "echo", json!({"message": "test"}), "All done!");
let plugin = Arc::new(CountingObserver::new());
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), registry, config);
agent.register_observer(plugin.clone());
let result = agent.run("Echo test").await.unwrap();
assert!(result.success);
assert_eq!(plugin.tool_pres.load(Ordering::SeqCst), 1);
assert_eq!(plugin.tool_posts.load(Ordering::SeqCst), 1);
assert_eq!(plugin.turn_starts.load(Ordering::SeqCst), 2);
assert_eq!(plugin.turn_ends.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_conversation_built_correctly() {
let client = MockClient::new("test-model");
client.add_tool_then_text(
"tool_1",
"echo",
json!({"message": "hello"}),
"Final answer.",
);
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), registry, config);
agent.conversation.push(Message::user("Echo hello"));
let msg = Message::assistant("test");
let tool_calls = BareLoop::<MockClient>::extract_tool_calls(&msg);
assert!(tool_calls.is_empty());
let msg_with_tools = Message::new(
Role::Assistant,
vec![
MessagePart::text("Using tool..."),
MessagePart::tool_call("id1", "echo", json!({"message": "hi"})),
],
);
let tool_calls = BareLoop::<MockClient>::extract_tool_calls(&msg_with_tools);
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].tool, "echo");
}
#[tokio::test]
async fn test_tool_result_message_format() {
let results = vec![super::ToolDispatchResult {
tool_call_id: "tool_123".to_string(),
output: ToolContent::Text("Echo: hello".to_string()),
is_error: false,
duration: Duration::from_millis(100),
resolved_tool_name: String::new(),
}];
let msg = BareLoop::<MockClient>::build_tool_result_message(results);
assert_eq!(msg.role, Role::User);
assert_eq!(msg.parts.len(), 1);
match &msg.parts[0] {
MessagePart::ToolResult {
call_id,
output,
is_error,
} => {
assert_eq!(call_id, "tool_123");
assert!(!is_error.unwrap_or(true));
let text = output.to_string();
assert_eq!(text, "Echo: hello");
}
other => panic!("Expected ToolResult part, got: {other:?}"),
}
}
#[tokio::test]
async fn test_multiple_tool_calls_in_one_turn() {
let client = MockClient::new("test-model");
let tool_events = vec![
StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_multi".into(),
role: "assistant".into(),
model: "test-model".into(),
},
}),
StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call(
"t1",
"echo",
json!({"message": "first"}),
)),
}),
StreamEvent::PartStop,
StreamEvent::PartStart(PartStart {
index: 1,
part: Some(MessagePart::tool_call(
"t2",
"echo",
json!({"message": "second"}),
)),
}),
StreamEvent::PartStop,
StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("tool_call".to_string()),
},
usage: Some(Usage::new(50, 20)),
}),
StreamEvent::MessageStop,
];
client.responses.lock().push(tool_events);
client.add_text_response("Both tools executed.");
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), registry, config);
let result = agent.run("Echo twice").await.unwrap();
assert!(result.success);
assert_eq!(result.total_turns, 2);
assert_eq!(result.tool_calls, 2);
}
#[tokio::test]
async fn test_text_streamer_fires_on_text_delta() {
let client = MockClient::new("test-model");
client.add_text_response("Hello world");
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), make_config());
let received = Arc::new(Mutex::new(Vec::new()));
let buf = Arc::clone(&received);
agent.set_text_streamer(Arc::new(move |delta: &str| {
buf.lock().push(delta.to_string());
}));
let result = agent.run("Hi").await.unwrap();
assert!(result.success);
let received = received.lock();
assert!(!received.is_empty(), "streamer should have fired");
assert!(
received.join("").contains("Hello world"),
"got: {received:?}",
);
}
#[tokio::test]
async fn test_text_streamer_none_works() {
let client = MockClient::new("test-model");
client.add_text_response("No streamer");
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), make_config());
let result = agent.run("Hi").await.unwrap();
assert!(result.success);
}
#[tokio::test]
async fn test_text_streamer_ignores_non_text_deltas() {
let client = MockClient::new("test-model");
let events = vec![
StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg-1".into(),
role: "assistant".into(),
model: "test-model".into(),
},
}),
StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::ToolCall {
id: "call_1".into(),
name: "echo".into(),
input: Value::Null,
}),
}),
StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::InputJson {
partial_json: "{}".into(),
},
}),
StreamEvent::PartStop,
StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("tool_call".into()),
},
usage: None,
}),
StreamEvent::MessageStop,
];
client.add_events(events);
client.add_text_response("Done");
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), make_config());
let received = Arc::new(Mutex::new(String::new()));
let buf = Arc::clone(&received);
agent.set_text_streamer(Arc::new(move |delta: &str| {
buf.lock().push_str(delta);
}));
agent.run("Use tool").await.unwrap();
let received = received.lock();
assert_eq!(&*received, "Done", "only text deltas should fire streamer");
}
#[test]
fn test_accessors() {
let client = MockClient::new("test-model");
let config = make_config();
let session_id = config.session_id;
let agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
assert_eq!(agent.config().session_id, session_id);
assert!(agent.conversation().is_empty());
assert!(!agent.is_cancelled());
}
#[test]
fn test_cancel_signal_shared() {
let client = MockClient::new("test-model");
let config = make_config();
let agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let signal = agent.cancel_signal();
assert!(!signal.is_cancelled());
agent.cancel();
assert!(signal.is_cancelled());
assert!(agent.is_cancelled());
}
#[tokio::test]
async fn test_session_result_fields() {
let client = MockClient::new("test-model");
client.add_text_response("Hello!");
let config = make_config();
let session_id = config.session_id;
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let result = agent.run("Hi").await.unwrap();
assert_eq!(result.session_id, session_id);
assert!(result.total_duration > Duration::ZERO);
assert!(result.input_tokens > 0 || result.output_tokens > 0); }
#[tokio::test]
async fn test_loop_terminates_with_max_turns_1() {
let client = MockClient::new("test-model");
client.add_text_response("One and done.");
let mut config = make_config();
config.max_turns = 1;
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let result = agent.run("Hi").await.unwrap();
assert!(result.success);
assert_eq!(result.total_turns, 1);
}
#[tokio::test]
async fn test_loop_terminates_with_max_turns_0() {
let client = MockClient::new("test-model");
client.add_text_response("Should not be reached.");
let mut config = make_config();
config.max_turns = 0;
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let result = agent.run("Hi").await;
assert!(result.is_err());
match result.unwrap_err() {
LoopError::Config(msg) => assert!(msg.contains("max_turns")),
other => panic!("Expected Config error, got: {other}"),
}
}
#[tokio::test]
async fn test_tool_error_is_soft_not_hard() {
let client = MockClient::new("test-model");
let tool_events = vec![
StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_1".into(),
role: "assistant".into(),
model: "test-model".into(),
},
}),
StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call("t1", "nonexistent", json!({}))),
}),
StreamEvent::PartStop,
StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("tool_call".to_string()),
},
usage: Some(Usage::new(50, 10)),
}),
StreamEvent::MessageStop,
];
client.responses.lock().push(tool_events);
client.add_text_response("Tool wasn't found, but I'll handle it.");
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), config);
let result = agent.run("Use missing tool").await.unwrap();
assert!(result.success);
}
#[tokio::test]
async fn test_loop_detection_hard_stop_propagates_loop_error() {
use crate::detection::{DetectionConfig, DetectionManager};
use crate::runtime::LoopRuntime;
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let client = MockClient::new("test");
for i in 0..10 {
client.add_tool_only_response(&format!("call_{i}"), "echo", json!({ "message": "hi" }));
}
let runtime = LoopRuntime::new().with_detection(
DetectionManager::new_with_config(DetectionConfig {
loop_threshold: 2,
stop_threshold: 2,
..Default::default()
})
.expect("valid detection config"),
);
let mut agent =
BareLoop::new_with_managers(Arc::new(client), registry, make_config(), runtime);
let result = agent.run("test").await;
assert!(
matches!(result, Err(LoopError::LoopDetected { .. })),
"expected Err(LoopError::LoopDetected), got {result:?}"
);
}
#[tokio::test]
async fn test_loop_detection_soft_block_before_stop_threshold() {
use crate::detection::{DetectionConfig, DetectionManager};
use crate::runtime::LoopRuntime;
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let client = MockClient::new("test");
client.add_tool_only_response("c1", "echo", json!({ "message": "hi" }));
client.add_tool_only_response("c2", "echo", json!({ "message": "hi" }));
client.add_text_response("Done");
let runtime = LoopRuntime::new().with_detection(
DetectionManager::new_with_config(DetectionConfig {
loop_threshold: 2,
stop_threshold: 10,
..Default::default()
})
.expect("valid detection config"),
);
let mut agent =
BareLoop::new_with_managers(Arc::new(client), registry, make_config(), runtime);
let result = agent.run("test").await;
assert!(result.is_ok(), "expected Ok, got {result:?}");
}
#[tokio::test]
async fn test_cancelled_before_run_returns_cancelled() {
let client = MockClient::new("test");
client.add_text_response("Hello");
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), make_config());
agent.cancel();
let result = agent.run("test").await;
assert!(
matches!(result, Err(LoopError::Cancelled)),
"expected Err(LoopError::Cancelled), got {result:?}"
);
}
#[tokio::test]
async fn test_default_recovery_on_tool_error_returns_soft_result() {
let mut registry = ToolRegistry::new();
registry.register(FailingTool);
let client = MockClient::new("test");
client.add_tool_then_text("tool_1", "fail", json!({}), "Moving on");
let mut agent = BareLoop::new(Arc::new(client), registry, make_config());
let result = agent.run("Test").await.unwrap();
assert!(result.success);
assert_eq!(result.tool_calls, 1);
}
#[tokio::test]
async fn test_recovery_on_missing_tool_returns_soft_result() {
let client = MockClient::new("test");
client.add_tool_then_text("tool_1", "nonexistent", json!({}), "OK");
let mut agent = BareLoop::new(Arc::new(client), ToolRegistry::new(), make_config());
let result = agent.run("Test").await.unwrap();
assert!(result.success);
assert_eq!(result.tool_calls, 1);
}
#[tokio::test]
async fn test_recovery_noop_reflector_no_retries() {
let mut registry = ToolRegistry::new();
registry.register(FailingTool);
let client = MockClient::new("test");
client.add_tool_then_text("tool_1", "fail", json!({}), "OK");
let mut agent = BareLoop::new(Arc::new(client), registry, make_config());
let result = agent.run("Test").await.unwrap();
assert!(result.success);
assert_eq!(result.tool_calls, 1);
}
#[tokio::test]
async fn test_recovery_respects_cancellation() {
let mut registry = ToolRegistry::new();
registry.register(FailingTool);
let client = MockClient::new("test");
client.add_tool_only_response("tc-1", "fail", json!({}));
let mut agent = BareLoop::new(Arc::new(client), registry, make_config());
agent.cancel();
let result = agent.run("Test").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_set_pipeline_injects_self_tools_registry() {
let client = MockClient::new("test-model");
client.add_tool_then_text("tool_1", "echo", json!({"message": "hello"}), "done");
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let config = make_config();
let mut agent = BareLoop::new(Arc::new(client), registry, config);
let builder = ToolPipeline::builder();
agent.set_pipeline(builder).unwrap();
let result = agent.run("Echo hello").await;
assert!(result.is_ok());
assert!(result.unwrap().success);
}
struct TurnNumberCapture {
turns: Arc<Mutex<Vec<usize>>>,
}
impl TurnNumberCapture {
fn new(shared: Arc<Mutex<Vec<usize>>>) -> Self {
Self { turns: shared }
}
}
impl crate::middleware::ToolMiddleware for TurnNumberCapture {
fn name(&self) -> &'static str {
"turn_capture"
}
fn dispatch<'a>(
&'a self,
ctx: &'a mut ToolDispatchContext,
next: &'a ToolPipeline,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = crate::middleware::ToolDispatchResult> + Send + 'a,
>,
> {
self.turns.lock().push(ctx.turn_number);
next.dispatch(ctx)
}
}
#[tokio::test]
async fn test_turn_number_is_actual_turn_index() {
let client = MockClient::new("test-model");
client.add_tool_only_response("tool_0", "echo", json!({"message": "a"}));
client.add_tool_only_response("tool_1", "echo", json!({"message": "b"}));
client.add_text_response("done");
let mut registry = ToolRegistry::new();
registry.register(EchoTool);
let mut config = make_config();
config.max_turns = 10;
let capture = Arc::new(Mutex::new(Vec::<usize>::new()));
let mut agent = BareLoop::new(Arc::new(client), registry, config);
let builder = ToolPipeline::builder().with(TurnNumberCapture::new(Arc::clone(&capture)));
agent.set_pipeline(builder).unwrap();
let result = agent.run("test").await;
assert!(result.is_ok());
let turns = capture.lock().clone();
assert_eq!(
turns.len(),
2,
"expected tool calls on 2 turns: got {turns:?}"
);
assert_eq!(turns[0], 0, "first tool call should be on turn 0");
assert_eq!(turns[1], 1, "second tool call should be on turn 1");
assert!(
turns.iter().all(|&t| t < 10),
"turn_number must be actual index, not max_turns (10): got {turns:?}"
);
}
#[tokio::test]
async fn switch_model_updates_config_and_client() {
let client = MockClient::new("model-a");
let client_arc = std::sync::Arc::new(client);
let tools = ToolRegistry::new();
let mut config = LoopConfig::default();
config.model = "model-a".to_string();
let mut loop_ = BareLoop::new(client_arc.clone(), tools, config);
loop_.switch_model("model-b").apply().unwrap();
assert_eq!(loop_.config().model, "model-b");
assert_eq!(client_arc.model(), "model-b");
}
#[tokio::test]
async fn switch_model_notifies_observers() {
#[derive(Default)]
struct RecordingObserver {
switches: Mutex<Vec<(String, String)>>,
}
impl crate::observer::LoopObserver for RecordingObserver {
fn name(&self) -> &'static str {
"recording"
}
fn on_model_switched(&self, ctx: &ModelSwitchedContext) {
self.switches
.lock()
.push((ctx.from.clone(), ctx.to.clone()));
}
}
let client = std::sync::Arc::new(MockClient::new("m1"));
let tools = ToolRegistry::new();
let mut loop_ = BareLoop::new(client, tools, LoopConfig::default());
let obs = std::sync::Arc::new(RecordingObserver::default());
let obs_clone = obs.clone();
loop_.register_observer(obs);
loop_.switch_model("m2").apply().unwrap();
loop_.switch_model("m3").apply().unwrap();
let recorded = obs_clone.switches.lock();
assert_eq!(recorded.len(), 2, "should have 2 model-switch events");
assert_eq!(recorded[0], ("default".to_string(), "m2".to_string()));
assert_eq!(recorded[1], ("m2".to_string(), "m3".to_string()));
}
#[tokio::test]
async fn switch_model_unsupported_client() {
struct StaticClient {
model_name: Arc<parking_lot::Mutex<String>>,
}
impl ApiClient for StaticClient {
fn model(&self) -> String {
self.model_name.lock().clone()
}
fn stream_messages(
&self,
_messages: Vec<Message>,
_system: Option<String>,
_tools: Option<Vec<crate::tool::ToolSchema>>,
) -> Pin<
Box<
dyn futures::stream::Stream<Item = Result<StreamEvent, ApiError>>
+ Send
+ 'static,
>,
> {
Box::pin(futures::stream::empty())
}
fn create_message(
&self,
_messages: Vec<Message>,
_system: Option<String>,
_tools: Option<Vec<crate::tool::ToolSchema>>,
) -> Pin<
Box<
dyn std::future::Future<Output = Result<serde_json::Value, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async { Ok(serde_json::Value::Null) })
}
}
let client = std::sync::Arc::new(StaticClient {
model_name: std::sync::Arc::new(parking_lot::Mutex::new("static".to_string())),
});
let tools = ToolRegistry::new();
let mut loop_ = BareLoop::new(client, tools, LoopConfig::default());
loop_.switch_model("new-model").apply().unwrap();
assert_eq!(loop_.config().model, "new-model");
assert_eq!(loop_.client.model(), "static");
}
#[tokio::test]
async fn switch_model_updates_fallback_original() {
let client = std::sync::Arc::new(MockClient::new("primary"));
let tools = ToolRegistry::new();
let mut config = LoopConfig::default();
config.model = "primary".to_string();
let mut loop_ = BareLoop::new(client, tools, config);
assert_eq!(loop_.managers.fallback.original_model(), None);
loop_.switch_model("new-primary").apply().unwrap();
assert_eq!(
loop_.managers.fallback.original_model(),
Some("new-primary".to_string())
);
}
#[tokio::test]
async fn switch_model_rejects_empty() {
let client = std::sync::Arc::new(MockClient::new("model"));
let tools = ToolRegistry::new();
let mut loop_ = BareLoop::new(client, tools, LoopConfig::default());
let result = loop_.switch_model("").apply();
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("empty"));
let result = loop_.switch_model(" ").apply();
assert!(result.is_err());
assert_eq!(loop_.config().model, "default");
}
#[tokio::test]
async fn switch_model_chained() {
let client = std::sync::Arc::new(MockClient::new("a"));
let tools = ToolRegistry::new();
let mut loop_ = BareLoop::new(client, tools, LoopConfig::default());
loop_.switch_model("b").apply().unwrap();
assert_eq!(loop_.config().model, "b");
loop_.switch_model("c").apply().unwrap();
assert_eq!(loop_.config().model, "c");
loop_.switch_model("d").apply().unwrap();
assert_eq!(loop_.config().model, "d");
}
#[tokio::test]
async fn switch_model_updates_context_window() {
let client = std::sync::Arc::new(MockClient::new("big-model"));
let tools = ToolRegistry::new();
let mut loop_ = BareLoop::new(client, tools, LoopConfig::default());
let original_cw = loop_.config().context_window;
assert_ne!(original_cw, 8192);
loop_
.switch_model("small-model")
.context_window(8192)
.apply()
.unwrap();
assert_eq!(loop_.config().model, "small-model");
assert_eq!(loop_.config().context_window, 8192);
}
#[tokio::test]
async fn switch_model_updates_max_tokens() {
let client = std::sync::Arc::new(MockClient::new("m"));
let tools = ToolRegistry::new();
let mut loop_ = BareLoop::new(client, tools, LoopConfig::default());
loop_.switch_model("m2").max_tokens(4096).apply().unwrap();
assert_eq!(loop_.config().model, "m2");
assert_eq!(loop_.config().max_tokens, 4096);
}
#[tokio::test]
async fn switch_model_trims_whitespace() {
let client = std::sync::Arc::new(MockClient::new("m"));
let tools = ToolRegistry::new();
let mut loop_ = BareLoop::new(client, tools, LoopConfig::default());
loop_.switch_model(" gpt-4o ").apply().unwrap();
assert_eq!(loop_.config().model, "gpt-4o");
}
#[tokio::test]
async fn switch_model_resets_fallback_circuit() {
use crate::fallback::FallbackState;
let client = std::sync::Arc::new(MockClient::new("primary"));
let tools = ToolRegistry::new();
let mut config = LoopConfig::default();
config.model = "primary".to_string();
let mut loop_ = BareLoop::new(client, tools, config);
loop_.managers.fallback.set_original_model("primary".into());
loop_.managers.fallback.set_fallback_model("backup");
loop_.managers.fallback.transition_to_fallback();
assert_eq!(loop_.managers.fallback.state(), FallbackState::Fallback);
loop_.switch_model("new-primary").apply().unwrap();
assert_eq!(loop_.managers.fallback.state(), FallbackState::Primary);
assert_eq!(
loop_.managers.fallback.original_model(),
Some("new-primary".to_string())
);
}
#[cfg(feature = "hooks")]
struct ReasonCaptureHook {
reason: Mutex<Option<SessionEndReason>>,
}
#[cfg(feature = "hooks")]
impl ReasonCaptureHook {
fn new() -> Arc<Self> {
Arc::new(Self {
reason: Mutex::new(None),
})
}
fn captured(&self) -> Option<SessionEndReason> {
*self.reason.lock()
}
}
#[cfg(feature = "hooks")]
impl Hook for ReasonCaptureHook {
fn name(&self) -> &'static str {
"ReasonCaptureHook"
}
fn on_session_end(&self, ctx: &HookSessionEndContext) {
*self.reason.lock() = Some(ctx.reason);
}
}
#[cfg(feature = "hooks")]
fn loop_with_reason_hook() -> (BareLoop<MockClient>, Arc<ReasonCaptureHook>) {
let hook = ReasonCaptureHook::new();
let executor = Arc::new(HookExecutor::new().with_hook(hook.clone()));
let config = LoopConfig {
max_turns: 5,
..LoopConfig::default()
};
let mut loop_ = BareLoop::new(
Arc::new(MockClient::new("test")),
ToolRegistry::new(),
config,
);
loop_.set_hook_executor(executor);
(loop_, hook)
}
#[cfg(feature = "hooks")]
#[tokio::test]
async fn session_end_reason_complete() {
let (mut loop_, hook) = loop_with_reason_hook();
loop_.budget.success = true;
loop_.budget.total_turns = 2;
loop_.notify_session_end(&loop_.budget.clone(), Duration::from_millis(100));
assert_eq!(hook.captured(), Some(SessionEndReason::Complete));
}
#[cfg(feature = "hooks")]
#[tokio::test]
async fn session_end_reason_cancelled() {
let (mut loop_, hook) = loop_with_reason_hook();
loop_.budget.success = true;
loop_.budget.total_turns = 2;
loop_.cancelled.cancel();
loop_.notify_session_end(&loop_.budget.clone(), Duration::from_millis(100));
assert_eq!(hook.captured(), Some(SessionEndReason::Cancelled));
}
#[cfg(feature = "hooks")]
#[tokio::test]
async fn session_end_reason_max_turns() {
let (mut loop_, hook) = loop_with_reason_hook();
loop_.budget.success = true;
loop_.budget.total_turns = 5;
loop_.notify_session_end(&loop_.budget.clone(), Duration::from_millis(100));
assert_eq!(hook.captured(), Some(SessionEndReason::MaxTurns));
}
#[cfg(feature = "hooks")]
#[tokio::test]
async fn session_end_reason_error() {
let (mut loop_, hook) = loop_with_reason_hook();
loop_.budget.success = false;
loop_.budget.error = Some("API connection refused".to_string());
loop_.notify_session_end(&loop_.budget.clone(), Duration::from_millis(100));
assert_eq!(hook.captured(), Some(SessionEndReason::Error));
}
#[cfg(feature = "hooks")]
#[tokio::test]
async fn session_end_reason_context_overflow() {
let (mut loop_, hook) = loop_with_reason_hook();
loop_.budget.success = false;
loop_.budget.error = Some("context length exceeded".to_string());
loop_.notify_session_end(&loop_.budget.clone(), Duration::from_millis(100));
assert_eq!(hook.captured(), Some(SessionEndReason::ContextOverflow));
}
}