use kaynine_core::agent_loop::{run_agent_loop, LoopOutcome, LoopParams, NoopHooks, RunHooks};
use kaynine_core::budget::BudgetPolicy;
use kaynine_core::error::{ProviderError, RunFailureReason};
use kaynine_core::event::{EventEnvelope, RealtimeEvent};
use kaynine_core::ids::{BranchId, ModelId, RunId, SessionId, ToolCallId};
use kaynine_core::message::{ContentBlock, FinishReason, Message, ToolResultPayload};
use kaynine_core::provider::{
GenerationOptions, ModelCapabilities, ModelMessage, ProviderEvent, ReasoningLevel,
};
use kaynine_core::testing::{
FakeCredentialProvider, FakeProvider, FakeTokenCounter, ScriptedResponse, ScriptedTool,
};
use kaynine_core::tool::Tool;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
fn caps() -> ModelCapabilities {
ModelCapabilities {
context_tokens: 200_000,
max_output_tokens: 8_192,
supports_tools: true,
supports_images: true,
supports_reasoning: true,
}
}
fn text_events(text: &str) -> Vec<ProviderEvent> {
vec![
ProviderEvent::ResponseStarted,
ProviderEvent::TextDelta {
block: 0,
text: text.into(),
},
ProviderEvent::ResponseCompleted {
finish_reason: FinishReason::Stop,
},
]
}
fn tool_call_events(calls: &[(&str, &str, &str)]) -> Vec<ProviderEvent> {
let mut events = vec![ProviderEvent::ResponseStarted];
for (i, (id, name, args)) in calls.iter().enumerate() {
events.push(ProviderEvent::ToolCallStarted {
block: i as u32,
id: (*id).into(),
name: (*name).into(),
});
events.push(ProviderEvent::ToolCallArgumentsDelta {
block: i as u32,
json: (*args).into(),
});
}
events.push(ProviderEvent::ResponseCompleted {
finish_reason: FinishReason::Stop,
});
events
}
async fn drive(
provider_scripts: Vec<ScriptedResponse>,
tools: Vec<Arc<dyn Tool>>,
history: Vec<Message>,
) -> (
LoopOutcome,
Arc<FakeProvider>,
Vec<EventEnvelope<RealtimeEvent>>,
) {
let provider = Arc::new(FakeProvider::new(caps()));
for script in provider_scripts {
match script {
ScriptedResponse::Events(events) => provider.push_events(events),
ScriptedResponse::EventsThenHang(events) => provider.push_events_then_hang(events),
ScriptedResponse::EventsThenError(events, error) => {
provider.push_events_then_error(events, error)
}
ScriptedResponse::Error(error) => provider.push_error(error),
ScriptedResponse::Hang => provider.push_hang(),
}
}
let (events_tx, mut events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools,
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: "system".into(),
prompt: None,
skills: None,
definition_id: String::new(),
history,
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy {
reserved_output_tokens: 8_192,
extra_safety_margin_tokens: 1_024,
},
events: events_tx,
cancel: CancellationToken::new(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
let mut collected = Vec::new();
while let Ok(env) = events_rx.try_recv() {
collected.push(env);
}
(outcome, provider, collected)
}
#[tokio::test]
async fn single_turn_completes_without_tools() {
let (outcome, provider, events) = drive(
vec![ScriptedResponse::Events(text_events("答案"))],
vec![],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "问".into() }],
}],
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
assert_eq!(messages.len(), 1);
assert!(matches!(&messages[0],
Message::Assistant { blocks, finish_reason: FinishReason::Stop, truncated: false }
if matches!(&blocks[0], ContentBlock::Text { text } if text == "答案")));
}
other => panic!("expected completed, got {other:?}"),
}
assert_eq!(provider.requests().len(), 1);
assert_eq!(provider.requests()[0].system_prompt, "system");
let kinds: Vec<&str> = events
.iter()
.map(|e| match &e.payload {
RealtimeEvent::RunStarted => "run_started",
RealtimeEvent::TurnStarted { .. } => "turn_started",
RealtimeEvent::TextDelta { .. } => "text_delta",
RealtimeEvent::TurnCompleted { .. } => "turn_completed",
RealtimeEvent::RunCompleted => "run_completed",
_ => "other",
})
.collect();
assert_eq!(
kinds,
vec![
"run_started",
"turn_started",
"text_delta",
"turn_completed",
"run_completed"
]
);
for (i, env) in events.iter().enumerate() {
assert_eq!(env.run_seq, Some(i as u64 + 1));
assert_eq!(env.run_id.as_ref().unwrap().as_ref(), "r1");
}
}
#[tokio::test]
async fn multi_turn_tool_flow_appends_results_in_call_order() {
let (outcome, provider, _events) = drive(
vec![
ScriptedResponse::Events(tool_call_events(&[
("c1", "echo", "{\"x\":1}"),
("c2", "echo", "{\"y\":2}"),
])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
vec![Message::User {
blocks: vec![ContentBlock::Text {
text: "跑工具".into(),
}],
}],
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
assert_eq!(messages.len(), 3);
let results = match &messages[1] {
Message::ToolResult { results } => results.clone(),
other => panic!("expected tool result, got {other:?}"),
};
assert_eq!(results[0].call_id, ToolCallId::from("c1"));
assert_eq!(results[1].call_id, ToolCallId::from("c2"));
assert_eq!(results[0].text, "{\"x\":1}");
assert!(!results[0].is_error);
}
other => panic!("expected completed, got {other:?}"),
}
let requests = provider.requests();
assert_eq!(requests.len(), 2);
match &requests[1].messages[..] {
[.., ModelMessage::ToolResults { results }] => {
assert_eq!(results.len(), 2);
assert_eq!(results[0].call_id, ToolCallId::from("c1"));
}
other => panic!("second request should end with tool results, got {other:?}"),
}
}
#[tokio::test]
async fn truncated_text_persists_as_truncated_assistant() {
let (outcome, _provider, _events) = drive(
vec![ScriptedResponse::Events(vec![
ProviderEvent::ResponseStarted,
ProviderEvent::TextDelta {
block: 0,
text: "半截".into(),
},
ProviderEvent::ResponseCompleted {
finish_reason: FinishReason::Length,
},
])],
vec![],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
assert!(matches!(
&messages[0],
Message::Assistant {
finish_reason: FinishReason::Length,
truncated: true,
..
}
));
}
other => panic!("expected completed, got {other:?}"),
}
}
#[tokio::test]
async fn invalid_tool_arguments_fail_run_without_execution() {
let (outcome, provider, _events) = drive(
vec![ScriptedResponse::Events(vec![
ProviderEvent::ResponseStarted,
ProviderEvent::ToolCallStarted {
block: 0,
id: "c1".into(),
name: "echo".into(),
},
ProviderEvent::ToolCallArgumentsDelta {
block: 0,
json: "{\"broken".into(),
},
ProviderEvent::ResponseCompleted {
finish_reason: FinishReason::Stop,
},
])],
vec![Arc::new(ScriptedTool::echo("echo"))],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
)
.await;
match &outcome {
LoopOutcome::Failed { reason, .. } => {
assert!(matches!(
reason,
RunFailureReason::Provider(ProviderError::Protocol(_))
))
}
other => panic!("expected failed, got {other:?}"),
}
assert_eq!(provider.requests().len(), 1);
}
#[tokio::test]
async fn provider_error_before_semantic_output_fails_without_retry() {
let (outcome, provider, _events) = drive(
vec![ScriptedResponse::Error(ProviderError::Network(
"down".into(),
))],
vec![],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
)
.await;
assert!(matches!(
&outcome,
LoopOutcome::Failed {
reason: RunFailureReason::Provider(ProviderError::Network(_)),
..
}
));
assert_eq!(provider.requests().len(), 1);
}
#[tokio::test]
async fn provider_error_after_semantic_output_fails_without_retry() {
let (outcome, provider, _events) = drive(
vec![ScriptedResponse::EventsThenError(
vec![
ProviderEvent::ResponseStarted,
ProviderEvent::TextDelta {
block: 0,
text: "部分".into(),
},
],
ProviderError::Network("mid-stream drop".into()),
)],
vec![],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
)
.await;
assert!(matches!(&outcome, LoopOutcome::Failed { .. }));
assert_eq!(provider.requests().len(), 1);
}
#[tokio::test]
async fn cancel_during_stream_discards_partial_text() {
let provider = Arc::new(FakeProvider::new(caps()));
provider.push_events_then_hang(vec![
ProviderEvent::ResponseStarted,
ProviderEvent::TextDelta {
block: 0,
text: "流式中".into(),
},
]);
let (events_tx, mut events_rx) = mpsc::channel(1024);
let cancel = CancellationToken::new();
let cancel_for_task = cancel.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
cancel_for_task.cancel();
});
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: cancel.clone(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
assert!(matches!(outcome, LoopOutcome::Cancelled { ref messages } if messages.is_empty()));
let mut saw_text_delta = false;
while let Ok(env) = events_rx.try_recv() {
if matches!(env.payload, RealtimeEvent::TextDelta { .. }) {
saw_text_delta = true;
}
}
assert!(saw_text_delta);
}
#[tokio::test]
async fn cancel_during_tool_batch_synthesizes_results_in_call_order() {
let provider = Arc::new(FakeProvider::new(caps()));
provider.push_events(tool_call_events(&[
("c1", "slow", "{}"),
("c2", "echo", "{}"),
]));
let (events_tx, _events_rx) = mpsc::channel(1024);
let cancel = CancellationToken::new();
let cancel_for_task = cancel.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
cancel_for_task.cancel();
});
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![
Arc::new(ScriptedTool::cancel_aware_sleep(
"slow",
Duration::from_secs(30),
)),
Arc::new(ScriptedTool::echo("echo")),
],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: cancel.clone(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
match &outcome {
LoopOutcome::Cancelled { messages } => {
let results = match messages.last() {
Some(Message::ToolResult { results }) => results,
other => panic!("expected tool result message, got {other:?}"),
};
assert_eq!(results.len(), 2);
assert_eq!(results[0].call_id, ToolCallId::from("c1"));
assert_eq!(results[1].call_id, ToolCallId::from("c2"));
assert!(results[0].is_error);
assert!(results[0].text.contains("cancelled"));
assert!(results[1].is_error);
assert!(results[1].text.contains("取消"));
}
other => panic!("expected cancelled, got {other:?}"),
}
}
#[tokio::test]
async fn budget_exceeded_fails_before_provider_call() {
let provider = Arc::new(FakeProvider::new(caps()));
provider.push_events(text_events("不应被调用"));
let (events_tx, _events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(500_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: CancellationToken::new(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
match &outcome {
LoopOutcome::Failed { reason, .. } => {
assert!(matches!(
reason,
RunFailureReason::ContextBudgetExceeded { .. }
))
}
other => panic!("expected failed, got {other:?}"),
}
assert!(provider.requests().is_empty());
}
#[tokio::test]
async fn capability_mismatch_fails_before_any_request() {
let provider = Arc::new(FakeProvider::new(ModelCapabilities {
context_tokens: 200_000,
max_output_tokens: 8_192,
supports_tools: false,
supports_images: true,
supports_reasoning: true,
}));
provider.push_events(text_events("不应被调用"));
let (events_tx, _events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![Arc::new(ScriptedTool::echo("echo"))],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: provider.capabilities,
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: CancellationToken::new(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
match &outcome {
LoopOutcome::Failed { reason, .. } => {
assert!(
matches!(reason, RunFailureReason::ModelCapabilityMismatch { requirement, .. } if requirement == "tools")
)
}
other => panic!("expected failed, got {other:?}"),
}
assert!(provider.requests().is_empty());
}
#[tokio::test]
async fn unknown_tool_yields_error_result_and_run_continues() {
let (outcome, provider, _events) = drive(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "missing_tool", "{}")])),
ScriptedResponse::Events(text_events("已处理错误")),
],
vec![],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
let results = match &messages[1] {
Message::ToolResult { results } => results,
other => panic!("expected tool result message, got {other:?}"),
};
assert_eq!(results[0].call_id, ToolCallId::from("c1"));
assert!(results[0].is_error);
assert!(results[0].text.contains("unknown tool"));
}
other => panic!("expected completed, got {other:?}"),
}
assert_eq!(provider.requests().len(), 2);
}
#[tokio::test]
async fn max_turns_exhaustion_fails_as_internal() {
let provider = Arc::new(FakeProvider::new(caps()));
provider.push_events(tool_call_events(&[("c1", "echo", "{}")]));
provider.push_events(text_events("不应到达"));
let (events_tx, _events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![Arc::new(ScriptedTool::echo("echo"))],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: CancellationToken::new(),
max_turns: Some(1),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
match &outcome {
LoopOutcome::Failed { reason, .. } => {
assert!(matches!(reason, RunFailureReason::Internal))
}
other => panic!("expected failed, got {other:?}"),
}
assert_eq!(provider.requests().len(), 1);
}
#[tokio::test]
async fn dangling_tool_call_in_history_fails_before_provider_call() {
let provider = Arc::new(FakeProvider::new(caps()));
provider.push_events(text_events("不应被调用"));
let (events_tx, _events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![
Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
},
Message::Assistant {
blocks: vec![ContentBlock::ToolCall {
id: ToolCallId::from("c1"),
name: "t".into(),
arguments: serde_json::json!({}),
}],
finish_reason: FinishReason::Stop,
truncated: false,
},
],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: CancellationToken::new(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
assert!(matches!(
&outcome,
LoopOutcome::Failed {
reason: RunFailureReason::Internal,
..
}
));
assert!(provider.requests().is_empty());
}
struct RecordingHooks {
log: Arc<std::sync::Mutex<Vec<String>>>,
}
#[async_trait::async_trait]
impl RunHooks for RecordingHooks {
async fn on_turn_started(&self, turn: u32) -> Result<(), RunFailureReason> {
self.log
.lock()
.unwrap()
.push(format!("turn_started:{turn}"));
Ok(())
}
async fn on_assistant_message(&self, _message: &Message) -> Result<(), RunFailureReason> {
self.log.lock().unwrap().push("assistant".into());
Ok(())
}
async fn on_tool_calls_planned(
&self,
calls: &[(ToolCallId, String, serde_json::Value)],
) -> Result<(), RunFailureReason> {
self.log
.lock()
.unwrap()
.push(format!("planned:{}", calls.len()));
Ok(())
}
async fn on_tool_results(&self, results: &[ToolResultPayload]) -> Result<(), RunFailureReason> {
self.log
.lock()
.unwrap()
.push(format!("results:{}", results.len()));
Ok(())
}
async fn on_turn_completed(
&self,
turn: u32,
_finish_reason: FinishReason,
) -> Result<(), RunFailureReason> {
self.log
.lock()
.unwrap()
.push(format!("turn_completed:{turn}"));
Ok(())
}
}
async fn drive_with_hooks(
provider_scripts: Vec<ScriptedResponse>,
tools: Vec<Arc<dyn Tool>>,
hooks: Arc<dyn RunHooks>,
cancel_grace: Duration,
cancel: CancellationToken,
) -> (LoopOutcome, Arc<FakeProvider>) {
let provider = Arc::new(FakeProvider::new(caps()));
for script in provider_scripts {
match script {
ScriptedResponse::Events(events) => provider.push_events(events),
ScriptedResponse::EventsThenHang(events) => provider.push_events_then_hang(events),
ScriptedResponse::EventsThenError(events, error) => {
provider.push_events_then_error(events, error)
}
ScriptedResponse::Error(error) => provider.push_error(error),
ScriptedResponse::Hang => provider.push_hang(),
}
}
let (events_tx, _events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools,
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel,
max_turns: Some(10),
hooks,
cancel_grace,
policy: Arc::new(AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
(outcome, provider)
}
#[tokio::test]
async fn hooks_called_in_persistence_order() {
let log = Arc::new(std::sync::Mutex::new(Vec::new()));
let hooks = Arc::new(RecordingHooks { log: log.clone() });
let (outcome, provider) = drive_with_hooks(
vec![
ScriptedResponse::Events(tool_call_events(&[
("c1", "echo", "{\"x\":1}"),
("c2", "echo", "{\"y\":2}"),
])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
hooks,
Duration::from_secs(30),
CancellationToken::new(),
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
assert_eq!(provider.requests().len(), 2);
assert_eq!(
*log.lock().unwrap(),
vec![
"turn_started:1",
"assistant",
"turn_completed:1",
"planned:2",
"results:2",
"turn_started:2",
"assistant",
"turn_completed:2",
]
);
}
struct FailingPlannedHooks;
#[async_trait::async_trait]
impl RunHooks for FailingPlannedHooks {
async fn on_tool_calls_planned(
&self,
_calls: &[(ToolCallId, String, serde_json::Value)],
) -> Result<(), RunFailureReason> {
Err(RunFailureReason::Internal)
}
}
#[tokio::test]
async fn hook_failure_on_planned_skips_execution() {
let executed: Arc<std::sync::Mutex<Vec<String>>> = Arc::new(std::sync::Mutex::new(Vec::new()));
let executed_for_tool = executed.clone();
let tool = ScriptedTool::with_behavior("echo", move |_call, _ctx| {
let executed_for_tool = executed_for_tool.clone();
Box::pin(async move {
executed_for_tool.lock().unwrap().push("ran".into());
Ok(kaynine_core::tool::ToolOutput {
is_error: false,
text: "ok".into(),
})
})
});
let (outcome, provider) = drive_with_hooks(
vec![ScriptedResponse::Events(tool_call_events(&[(
"c1", "echo", "{}",
)]))],
vec![Arc::new(tool)],
Arc::new(FailingPlannedHooks),
Duration::from_secs(30),
CancellationToken::new(),
)
.await;
match &outcome {
LoopOutcome::Failed {
reason: RunFailureReason::Internal,
..
} => {}
other => panic!("expected failed internal, got {other:?}"),
}
assert_eq!(provider.requests().len(), 1);
assert!(executed.lock().unwrap().is_empty());
}
struct FailingAssistantHooks;
#[async_trait::async_trait]
impl RunHooks for FailingAssistantHooks {
async fn on_assistant_message(&self, _message: &Message) -> Result<(), RunFailureReason> {
Err(RunFailureReason::Internal)
}
}
#[tokio::test]
async fn hook_failure_on_assistant_message_fails_run() {
let (outcome, _provider) = drive_with_hooks(
vec![ScriptedResponse::Events(text_events("答案"))],
vec![],
Arc::new(FailingAssistantHooks),
Duration::from_secs(30),
CancellationToken::new(),
)
.await;
match &outcome {
LoopOutcome::Failed {
reason: RunFailureReason::Internal,
messages,
} => assert!(messages.is_empty()),
other => panic!("expected failed internal, got {other:?}"),
}
}
#[tokio::test]
async fn uncooperative_tool_hits_cancel_grace_in_doubt() {
let provider = Arc::new(FakeProvider::new(caps()));
provider.push_events(tool_call_events(&[("c1", "stuck", "{}")]));
let (events_tx, _events_rx) = mpsc::channel(1024);
let cancel = CancellationToken::new();
let cancel_for_task = cancel.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
cancel_for_task.cancel();
});
let stuck = ScriptedTool::with_behavior("stuck", move |_call, _ctx| {
Box::pin(async move {
loop {
tokio::time::sleep(Duration::from_secs(1)).await;
}
})
});
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![Arc::new(stuck)],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: cancel.clone(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_millis(100),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let started = std::time::Instant::now();
let outcome = run_agent_loop(params).await;
assert!(started.elapsed() < Duration::from_secs(2));
match &outcome {
LoopOutcome::Cancelled { messages } => {
let results = match messages.last() {
Some(Message::ToolResult { results }) => results,
other => panic!("expected tool result message, got {other:?}"),
};
assert_eq!(results.len(), 1);
assert!(results[0].is_error);
assert!(results[0].text.contains("不得假定可以安全重试"));
}
other => panic!("expected cancelled, got {other:?}"),
}
}
use kaynine_core::policy::{AllowAllPolicy, ApprovalHandler, Policy, SteerSource};
use kaynine_core::testing::{QueueSteerSource, ScriptedApproval, ScriptedPolicy};
#[allow(clippy::too_many_arguments)]
async fn drive_with(
provider_scripts: Vec<ScriptedResponse>,
tools: Vec<Arc<dyn Tool>>,
policy: Arc<dyn Policy>,
approval: Option<Arc<dyn ApprovalHandler>>,
steer: Option<Arc<dyn SteerSource>>,
cancel: CancellationToken,
) -> (
LoopOutcome,
Arc<FakeProvider>,
Vec<EventEnvelope<RealtimeEvent>>,
) {
let provider = Arc::new(FakeProvider::new(caps()));
for script in provider_scripts {
match script {
ScriptedResponse::Events(events) => provider.push_events(events),
ScriptedResponse::EventsThenHang(events) => provider.push_events_then_hang(events),
ScriptedResponse::EventsThenError(events, error) => {
provider.push_events_then_error(events, error)
}
ScriptedResponse::Error(error) => provider.push_error(error),
ScriptedResponse::Hang => provider.push_hang(),
}
}
let (events_tx, mut events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools,
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel,
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy,
approval,
steer,
};
let outcome = run_agent_loop(params).await;
let mut collected = Vec::new();
while let Ok(env) = events_rx.try_recv() {
collected.push(env);
}
(outcome, provider, collected)
}
#[tokio::test]
async fn deny_produces_structured_error_and_run_continues() {
let (outcome, provider, _events) = drive_with(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "echo", "{\"x\":1}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
Arc::new(kaynine_core::policy::DenyAllPolicy {
reason: "机密".into(),
}),
None,
None,
CancellationToken::new(),
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
let results = match &messages[1] {
Message::ToolResult { results } => results,
other => panic!("expected tool result, got {other:?}"),
};
assert!(results[0].is_error);
assert!(results[0].text.contains("权限拒绝"));
assert!(results[0].text.contains("机密"));
}
other => panic!("expected completed, got {other:?}"),
}
assert_eq!(provider.requests().len(), 2);
match &provider.requests()[1].messages[..] {
[.., ModelMessage::ToolResults { results }] => {
assert!(results[0].is_error);
assert!(results[0].text.contains("机密"));
}
other => panic!("second request should end with tool results, got {other:?}"),
}
}
#[tokio::test]
async fn ask_without_handler_fails_closed() {
let (outcome, provider, _events) = drive_with(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "echo", "{}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
Arc::new(ScriptedPolicy::ask()),
None,
None,
CancellationToken::new(),
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
let results = match &messages[1] {
Message::ToolResult { results } => results,
other => panic!("expected tool result, got {other:?}"),
};
assert!(results[0].is_error);
assert!(results[0].text.contains("审批未启用"));
}
other => panic!("expected completed, got {other:?}"),
}
assert_eq!(provider.requests().len(), 2);
}
#[tokio::test]
async fn ask_with_handler_approval_approved_executes_tool() {
let (outcome, _provider, _events) = drive_with(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "echo", "{\"x\":1}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
Arc::new(ScriptedPolicy::ask()),
Some(Arc::new(ScriptedApproval {
decision: kaynine_core::policy::ApprovalDecision::Approved,
wait_cancel_aware: false,
})),
None,
CancellationToken::new(),
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
let results = match &messages[1] {
Message::ToolResult { results } => results,
other => panic!("expected tool result, got {other:?}"),
};
assert!(!results[0].is_error);
assert_eq!(results[0].text, "{\"x\":1}");
}
other => panic!("expected completed, got {other:?}"),
}
}
#[tokio::test]
async fn ask_with_handler_approval_denied_blocks_execution() {
let (outcome, _provider, _events) = drive_with(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "echo", "{}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
Arc::new(ScriptedPolicy::ask()),
Some(Arc::new(ScriptedApproval {
decision: kaynine_core::policy::ApprovalDecision::Denied {
reason: "用户拒绝".into(),
},
wait_cancel_aware: false,
})),
None,
CancellationToken::new(),
)
.await;
match &outcome {
LoopOutcome::Completed { messages } => {
let results = match &messages[1] {
Message::ToolResult { results } => results,
other => panic!("expected tool result, got {other:?}"),
};
assert!(results[0].is_error);
assert!(results[0].text.contains("权限拒绝"));
assert!(results[0].text.contains("用户拒绝"));
}
other => panic!("expected completed, got {other:?}"),
}
}
#[tokio::test]
async fn approval_wait_respects_cancel() {
let executed: Arc<std::sync::Mutex<Vec<String>>> = Arc::new(std::sync::Mutex::new(Vec::new()));
let executed_for_tool = executed.clone();
let tool = ScriptedTool::with_behavior("echo", move |_call, _ctx| {
let executed_for_tool = executed_for_tool.clone();
Box::pin(async move {
executed_for_tool.lock().unwrap().push("ran".into());
Ok(kaynine_core::tool::ToolOutput {
is_error: false,
text: "ok".into(),
})
})
});
let cancel = CancellationToken::new();
let cancel_for_task = cancel.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
cancel_for_task.cancel();
});
let (outcome, _provider, _events) = drive_with(
vec![ScriptedResponse::Events(tool_call_events(&[(
"c1", "echo", "{}",
)]))],
vec![Arc::new(tool)],
Arc::new(ScriptedPolicy::ask()),
Some(Arc::new(ScriptedApproval {
decision: kaynine_core::policy::ApprovalDecision::Approved,
wait_cancel_aware: true,
})),
None,
cancel,
)
.await;
assert!(matches!(outcome, LoopOutcome::Cancelled { .. }));
assert!(executed.lock().unwrap().is_empty());
}
#[tokio::test]
async fn steers_injected_after_toolless_response() {
let (outcome, provider, events) = drive_with(
vec![
ScriptedResponse::Events(text_events("第一答")),
ScriptedResponse::Events(tool_call_events(&[("c1", "echo", "{}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
Arc::new(AllowAllPolicy),
None,
Some(Arc::new(QueueSteerSource::new(vec![
"补充1".into(),
"补充2".into(),
]))),
CancellationToken::new(),
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
let requests = provider.requests();
assert_eq!(requests.len(), 3);
match &requests[1].messages[..] {
[ModelMessage::User { .. }, ModelMessage::Assistant { .. }, ModelMessage::User { blocks: u1 }, ModelMessage::User { blocks: u2 }] =>
{
assert!(matches!(&u1[0], ContentBlock::Text { text } if text == "补充1"));
assert!(matches!(&u2[0], ContentBlock::Text { text } if text == "补充2"));
}
other => panic!("unexpected second request messages: {other:?}"),
}
assert!(events
.iter()
.any(|e| matches!(e.payload, RealtimeEvent::SteerInjected { count: 2 })));
}
#[tokio::test]
async fn steers_injected_after_tool_batch() {
let (outcome, provider, _events) = drive_with(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "echo", "{}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
Arc::new(AllowAllPolicy),
None,
Some(Arc::new(QueueSteerSource::new(vec!["转向".into()]))),
CancellationToken::new(),
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
let requests = provider.requests();
assert_eq!(requests.len(), 2);
match &requests[1].messages[..] {
[.., ModelMessage::ToolResults { .. }, ModelMessage::User { blocks }] => {
assert!(matches!(&blocks[0], ContentBlock::Text { text } if text == "转向"));
}
other => panic!("unexpected second request messages: {other:?}"),
}
}
#[tokio::test]
async fn parallel_safe_tools_run_concurrently() {
let started = std::time::Instant::now();
let (outcome, _provider, _events) = drive(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "a", "{}"), ("c2", "b", "{}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![
Arc::new(ScriptedTool::parallel_sleepy(
"a",
Duration::from_millis(300),
)),
Arc::new(ScriptedTool::parallel_sleepy(
"b",
Duration::from_millis(300),
)),
],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
assert!(started.elapsed() < Duration::from_millis(500));
if let LoopOutcome::Completed { messages } = &outcome {
let results = match &messages[1] {
Message::ToolResult { results } => results,
other => panic!("expected tool result, got {other:?}"),
};
assert_eq!(results[0].call_id, ToolCallId::from("c1"));
assert_eq!(results[1].call_id, ToolCallId::from("c2"));
assert!(!results[0].is_error);
assert!(!results[1].is_error);
}
}
#[tokio::test]
async fn mixed_concurrency_batch_runs_sequentially() {
let started = std::time::Instant::now();
let (outcome, _provider, _events) = drive(
vec![
ScriptedResponse::Events(tool_call_events(&[
("c1", "seq", "{}"),
("c2", "par", "{}"),
])),
ScriptedResponse::Events(text_events("完成")),
],
vec![
Arc::new(ScriptedTool::cancel_aware_sleep(
"seq",
Duration::from_millis(300),
)),
Arc::new(ScriptedTool::parallel_sleepy(
"par",
Duration::from_millis(300),
)),
],
vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
assert!(started.elapsed() >= Duration::from_millis(600));
}
use kaynine_core::error::PromptError;
use kaynine_core::prompt::{PromptComposer, PromptContext, PromptFragment, PromptLayer};
use kaynine_core::testing::FailingLayer;
struct TurnLayer;
#[async_trait::async_trait]
impl PromptLayer for TurnLayer {
fn id(&self) -> &str {
"turn-dynamic"
}
fn priority(&self) -> i32 {
500
}
async fn render(&self, context: &PromptContext) -> Result<Option<PromptFragment>, PromptError> {
Ok(Some(PromptFragment {
content: format!("turn {}", context.turn),
}))
}
}
#[allow(clippy::too_many_arguments)]
async fn drive_with_prompt(
provider_scripts: Vec<ScriptedResponse>,
tools: Vec<Arc<dyn Tool>>,
prompt: Option<Arc<PromptComposer>>,
) -> (LoopOutcome, Arc<FakeProvider>) {
let provider = Arc::new(FakeProvider::new(caps()));
for script in provider_scripts {
match script {
ScriptedResponse::Events(events) => provider.push_events(events),
ScriptedResponse::EventsThenHang(events) => provider.push_events_then_hang(events),
ScriptedResponse::EventsThenError(events, error) => {
provider.push_events_then_error(events, error)
}
ScriptedResponse::Error(error) => provider.push_error(error),
ScriptedResponse::Hang => provider.push_hang(),
}
}
let (events_tx, _events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: Arc::new(FakeTokenCounter::fixed(1_000)),
credentials: Arc::new(FakeCredentialProvider::default()),
tools,
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: "fallback".into(),
prompt,
skills: None,
definition_id: "defs/agent".into(),
history: vec![Message::User {
blocks: vec![ContentBlock::Text { text: "q".into() }],
}],
compaction: None,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget: BudgetPolicy::default(),
events: events_tx,
cancel: CancellationToken::new(),
max_turns: Some(10),
hooks: Arc::new(NoopHooks),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(kaynine_core::policy::AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
(outcome, provider)
}
#[tokio::test]
async fn composer_renders_per_turn_system_prompt() {
let composer = Arc::new(PromptComposer::new(
vec![
Arc::new(TurnLayer),
Arc::new(kaynine_core::testing::StaticLayer::new(
"base",
100,
Some("base"),
)),
],
Vec::new(),
));
let (outcome, provider) = drive_with_prompt(
vec![
ScriptedResponse::Events(tool_call_events(&[("c1", "echo", "{}")])),
ScriptedResponse::Events(text_events("完成")),
],
vec![Arc::new(ScriptedTool::echo("echo"))],
Some(composer),
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
let requests = provider.requests();
assert_eq!(requests.len(), 2);
assert_eq!(requests[0].system_prompt, "base\n\nturn 1");
assert_eq!(requests[1].system_prompt, "base\n\nturn 2");
}
#[tokio::test]
async fn prompt_render_error_fails_run_as_prompt_failure() {
let composer = Arc::new(PromptComposer::new(
vec![Arc::new(FailingLayer)],
Vec::new(),
));
let (outcome, provider) = drive_with_prompt(
vec![ScriptedResponse::Events(text_events("不应到达"))],
vec![],
Some(composer),
)
.await;
match &outcome {
LoopOutcome::Failed { reason, .. } => {
assert!(matches!(reason, RunFailureReason::Prompt(_)))
}
other => panic!("expected failed prompt, got {other:?}"),
}
assert!(provider.requests().is_empty());
}
#[tokio::test]
async fn no_composer_falls_back_to_static_system_prompt() {
let (outcome, provider) = drive_with_prompt(
vec![ScriptedResponse::Events(text_events("答案"))],
vec![],
None,
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
assert_eq!(provider.requests()[0].system_prompt, "fallback");
}
use kaynine_core::compaction::{CompactionConfig, SummaryPayload};
use std::sync::atomic::{AtomicU64, Ordering};
fn history_with_turns(turns: usize, trailing_user: bool) -> Vec<Message> {
let mut history = Vec::new();
for i in 0..turns {
history.push(Message::User {
blocks: vec![ContentBlock::Text {
text: format!("q{i}"),
}],
});
history.push(Message::Assistant {
blocks: vec![ContentBlock::Text {
text: format!("a{i}"),
}],
finish_reason: FinishReason::Stop,
truncated: false,
});
}
if trailing_user {
history.push(Message::User {
blocks: vec![ContentBlock::Text {
text: "当前问题".into(),
}],
});
}
history
}
struct TwoStepCounter {
call: AtomicU64,
first: u64,
rest: u64,
}
impl TwoStepCounter {
fn new(first: u64, rest: u64) -> Arc<Self> {
Arc::new(Self {
call: AtomicU64::new(0),
first,
rest,
})
}
fn counter(self: &Arc<Self>) -> Arc<FakeTokenCounter> {
let inner = self.clone();
Arc::new(FakeTokenCounter {
estimator: Arc::new(move |_| {
let n = inner.call.fetch_add(1, Ordering::SeqCst);
if n == 0 {
inner.first
} else {
inner.rest
}
}),
source: kaynine_core::provider::TokenMeasurementSource::Heuristic,
})
}
}
#[derive(Default)]
struct CheckpointHooks {
checkpoints: std::sync::Mutex<Vec<SummaryPayload>>,
}
#[async_trait::async_trait]
impl RunHooks for CheckpointHooks {
async fn on_summary_checkpoint(
&self,
summary: &SummaryPayload,
) -> Result<(), RunFailureReason> {
self.checkpoints.lock().unwrap().push(summary.clone());
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
async fn drive_compaction(
scripts: Vec<ScriptedResponse>,
counter: Arc<dyn kaynine_core::provider::TokenCounter>,
history: Vec<Message>,
compaction: Option<CompactionConfig>,
budget: BudgetPolicy,
) -> (LoopOutcome, Arc<FakeProvider>, Arc<CheckpointHooks>) {
let provider = Arc::new(FakeProvider::new(caps()));
for script in scripts {
match script {
ScriptedResponse::Events(events) => provider.push_events(events),
ScriptedResponse::EventsThenHang(events) => provider.push_events_then_hang(events),
ScriptedResponse::EventsThenError(events, error) => {
provider.push_events_then_error(events, error)
}
ScriptedResponse::Error(error) => provider.push_error(error),
ScriptedResponse::Hang => provider.push_hang(),
}
}
let hooks = Arc::new(CheckpointHooks::default());
let (events_tx, _events_rx) = mpsc::channel(1024);
let params = LoopParams {
session_id: SessionId::from("s1"),
branch_id: BranchId::from("b1"),
run_id: RunId::from("r1"),
provider: provider.clone(),
token_counter: counter,
credentials: Arc::new(FakeCredentialProvider::default()),
tools: vec![],
model: ModelId::from("test-model"),
reasoning: ReasoningLevel::Off,
generation: GenerationOptions::default(),
capabilities: caps(),
system_prompt: String::new(),
prompt: None,
skills: None,
definition_id: String::new(),
history,
compaction,
compaction_selector: Arc::new(kaynine_core::compaction::CurrentModelSelector),
provider_options: serde_json::Value::Null,
budget,
events: events_tx,
cancel: CancellationToken::new(),
max_turns: Some(10),
hooks: hooks.clone(),
cancel_grace: Duration::from_secs(30),
policy: Arc::new(AllowAllPolicy),
approval: None,
steer: None,
};
let outcome = run_agent_loop(params).await;
(outcome, provider, hooks)
}
fn small_config() -> CompactionConfig {
CompactionConfig {
trigger_ratio: 0.5,
keep_recent_turns: 1,
max_attempts: 3,
}
}
#[tokio::test]
async fn preemptive_trigger_without_compactable_range_proceeds() {
let (outcome, provider, hooks) = drive_compaction(
vec![ScriptedResponse::Events(text_events("答案"))],
Arc::new(FakeTokenCounter::fixed(90_000)),
history_with_turns(1, true),
Some(CompactionConfig {
trigger_ratio: 0.3,
keep_recent_turns: 4,
max_attempts: 3,
}),
BudgetPolicy::default(),
)
.await;
assert!(
matches!(outcome, LoopOutcome::Completed { .. }),
"within-budget run must proceed when nothing is compactable"
);
assert_eq!(provider.requests().len(), 1);
assert!(hooks.checkpoints.lock().unwrap().is_empty());
}
#[tokio::test]
async fn compaction_reclaims_budget_and_completes() {
let counter = TwoStepCounter::new(500_000, 100).counter();
let (outcome, provider, hooks) = drive_compaction(
vec![
ScriptedResponse::Events(text_events("摘要")),
ScriptedResponse::Events(text_events("答案")),
],
counter,
history_with_turns(3, true),
Some(small_config()),
BudgetPolicy::default(),
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
let requests = provider.requests();
assert_eq!(requests.len(), 2);
assert!(requests[0].tools.is_empty(), "summary request is tool-free");
match &requests[1].messages[0] {
ModelMessage::Summary { text } => assert_eq!(text, "摘要"),
other => panic!("second request should start with Summary, got {other:?}"),
}
let checkpoints = hooks.checkpoints.lock().unwrap();
assert_eq!(checkpoints.len(), 1);
assert_eq!(checkpoints[0].covered_message_count, 6);
assert_eq!(checkpoints[0].retain_from, 6);
assert!(!checkpoints[0].source_hash.is_empty());
}
#[tokio::test]
async fn no_strict_decrease_fails() {
let (outcome, provider, hooks) = drive_compaction(
vec![ScriptedResponse::Events(text_events("摘要"))],
Arc::new(FakeTokenCounter::fixed(500_000)),
history_with_turns(3, true),
Some(small_config()),
BudgetPolicy::default(),
)
.await;
match &outcome {
LoopOutcome::Failed {
reason: RunFailureReason::ContextBudgetExceeded { .. },
..
} => {}
other => panic!("expected budget exceeded, got {other:?}"),
}
assert_eq!(provider.requests().len(), 1);
assert_eq!(hooks.checkpoints.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn exceeded_without_compaction_config_fails() {
let (outcome, provider, hooks) = drive_compaction(
vec![],
Arc::new(FakeTokenCounter::fixed(500_000)),
history_with_turns(3, true),
None,
BudgetPolicy::default(),
)
.await;
assert!(matches!(
&outcome,
LoopOutcome::Failed {
reason: RunFailureReason::ContextBudgetExceeded { .. },
..
}
));
assert!(provider.requests().is_empty());
assert!(hooks.checkpoints.lock().unwrap().is_empty());
}
#[tokio::test]
async fn context_overflow_retries_once() {
let (outcome, provider, hooks) = drive_compaction(
vec![
ScriptedResponse::Error(ProviderError::ContextOverflow),
ScriptedResponse::Events(text_events("摘要")),
ScriptedResponse::Events(text_events("答案")),
],
Arc::new(FakeTokenCounter::fixed(500)),
history_with_turns(1, true), Some(small_config()),
BudgetPolicy::default(),
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
let requests = provider.requests();
assert_eq!(requests.len(), 3, "overflow + summary + main");
assert!(matches!(
&requests[2].messages[0],
ModelMessage::Summary { text } if text == "摘要"
));
assert_eq!(hooks.checkpoints.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn second_overflow_fails() {
let (outcome, provider, _hooks) = drive_compaction(
vec![
ScriptedResponse::Error(ProviderError::ContextOverflow),
ScriptedResponse::Error(ProviderError::ContextOverflow),
],
Arc::new(FakeTokenCounter::fixed(500)),
history_with_turns(1, true),
Some(small_config()),
BudgetPolicy::default(),
)
.await;
assert!(matches!(
&outcome,
LoopOutcome::Failed {
reason: RunFailureReason::Provider(ProviderError::ContextOverflow),
..
}
));
assert_eq!(provider.requests().len(), 2);
}
#[tokio::test]
async fn summary_truncation_counts_as_attempt() {
let length_events = vec![
ProviderEvent::ResponseStarted,
ProviderEvent::TextDelta {
block: 0,
text: "截断".into(),
},
ProviderEvent::ResponseCompleted {
finish_reason: FinishReason::Length,
},
];
let (outcome, provider, hooks) = drive_compaction(
vec![
ScriptedResponse::Events(length_events.clone()),
ScriptedResponse::Events(length_events.clone()),
ScriptedResponse::Events(length_events),
],
Arc::new(FakeTokenCounter::fixed(500_000)),
history_with_turns(3, true),
Some(small_config()),
BudgetPolicy::default(),
)
.await;
assert!(matches!(
&outcome,
LoopOutcome::Failed {
reason: RunFailureReason::ContextBudgetExceeded { .. },
..
}
));
assert_eq!(provider.requests().len(), 3, "one stream per attempt");
assert!(hooks.checkpoints.lock().unwrap().is_empty());
}
#[tokio::test]
async fn trigger_ratio_preemptive_compaction() {
let (outcome, provider, hooks) = drive_compaction(
vec![
ScriptedResponse::Events(text_events("摘要")),
ScriptedResponse::Events(text_events("答案")),
],
Arc::new(FakeTokenCounter::fixed(170_000)),
history_with_turns(3, true),
Some(CompactionConfig {
trigger_ratio: 0.8,
keep_recent_turns: 1,
max_attempts: 3,
}),
BudgetPolicy {
reserved_output_tokens: 0,
extra_safety_margin_tokens: 0,
},
)
.await;
assert!(matches!(outcome, LoopOutcome::Completed { .. }));
assert_eq!(provider.requests().len(), 2);
assert_eq!(hooks.checkpoints.lock().unwrap().len(), 1);
}