mod common;
use std::sync::Arc;
use async_trait::async_trait;
use serde_json::json;
use common::{FakeProvider, FakeTurn, tool_call};
use hey::agent::{AgentOpts, run_agent};
use hey::config::{CompactionConfig, Config};
use hey::tools::{Registry, Tool, ToolCtx, ToolOutput};
use hey::ui::{Event, TestOutput};
struct FakeTool;
#[async_trait]
impl Tool for FakeTool {
fn name(&self) -> &str {
"bash"
}
fn description(&self) -> &str {
"Run a command (test)"
}
fn schema(&self) -> serde_json::Value {
json!({"type": "object", "properties": {"cmd": {"type": "string"}}})
}
async fn run(&self, _ctx: &ToolCtx, _args: serde_json::Value) -> ToolOutput {
ToolOutput::ok("output-ok")
}
}
fn cfg() -> Config {
let mut c = Config::default();
c.agent.max_turns = Some(10);
c
}
#[tokio::test]
async fn single_turn_text() {
let provider = FakeProvider::new(
"test",
vec![FakeTurn {
text: "hello",
tool_calls: vec![],
}],
);
let tools = Registry::new(); let mut out = TestOutput::new();
let stats = run_agent(
Arc::new(provider),
&tools,
&mut out,
&cfg(),
"hi",
&AgentOpts::default(),
)
.await;
assert_eq!(stats.turns, 1);
assert_eq!(stats.end_reason, hey::agent::EndReason::Complete);
assert_eq!(
out.events,
vec![
Event::TextDelta("hello".into()),
Event::Finished {
reason: "complete".into(),
turns: 1
}
]
);
}
#[tokio::test]
async fn multi_turn_tool_loop() {
let provider = FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![tool_call("c1", "bash", "{\"cmd\":\"ls\"}")],
},
FakeTurn {
text: "done",
tool_calls: vec![],
},
],
);
let mut tools = Registry::new();
tools.register(Arc::new(FakeTool));
let mut out = TestOutput::new();
let stats = run_agent(
Arc::new(provider),
&tools,
&mut out,
&cfg(),
"list files",
&AgentOpts::default(),
)
.await;
assert_eq!(stats.turns, 2);
assert_eq!(stats.tool_calls, 1);
assert_eq!(
out.events[0],
Event::ToolCall {
name: "bash".into(),
args: "{\"cmd\":\"ls\"}".into()
}
);
assert_eq!(
out.events[1],
Event::ToolResult {
name: "bash".into(),
ok: true,
output: "output-ok".into()
}
);
assert_eq!(out.events[2], Event::TextDelta("done".into()));
}
#[tokio::test]
async fn permission_denied_blocks_tool() {
let provider = FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![tool_call("c1", "bash", "{\"cmd\":\"rm -rf /\"}")],
},
FakeTurn {
text: "ok, canceled",
tool_calls: vec![],
},
],
);
let mut tools = Registry::new();
tools.register(Arc::new(FakeTool));
let mut cfg = cfg();
cfg.permission.allow = vec!["bash: echo".into()]; let mut out = TestOutput::new();
let stats = run_agent(
Arc::new(provider),
&tools,
&mut out,
&cfg,
"clean",
&AgentOpts::default(),
)
.await;
assert_eq!(stats.turns, 2);
assert_eq!(
out.events[0],
Event::ToolCall {
name: "bash".into(),
args: "{\"cmd\":\"rm -rf /\"}".into()
}
);
assert_eq!(
out.events[1],
Event::ToolResult {
name: "bash".into(),
ok: false,
output: "blocked by user".into()
}
);
assert!(
out.events
.iter()
.any(|e| e == &Event::TextDelta("ok, canceled".into()))
);
}
#[tokio::test]
async fn max_turns_stops_loop() {
let provider = FakeProvider::new(
"test",
vec![FakeTurn {
text: "",
tool_calls: vec![tool_call("c1", "bash", "{}")],
}],
);
let mut tools = Registry::new();
tools.register(Arc::new(FakeTool));
let mut out = TestOutput::new();
let mut c = cfg();
c.agent.max_turns = Some(3);
let stats = run_agent(
Arc::new(provider),
&tools,
&mut out,
&c,
"loop",
&AgentOpts::default(),
)
.await;
assert_eq!(stats.turns, 3);
assert_eq!(stats.end_reason, hey::agent::EndReason::MaxTurns);
assert!(
out.events
.iter()
.any(|e| matches!(e, Event::Finished { reason, .. } if reason == "max_turns"))
);
}
#[tokio::test]
async fn unknown_tool_is_blocked_without_confirm() {
let provider = FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![tool_call("c1", "nope", "{}")],
},
FakeTurn {
text: "done",
tool_calls: vec![],
},
],
);
let tools = Registry::new(); let mut out = TestOutput::new();
let stats = run_agent(
Arc::new(provider),
&tools,
&mut out,
&cfg(),
"x",
&AgentOpts::default(),
)
.await;
assert!(
!out.events
.iter()
.any(|e| matches!(e, Event::ToolCall { name, .. } if name.starts_with("confirm:")))
);
assert!(matches!(out.events[1], Event::ToolResult { ok: false, .. }));
assert!(
out.events
.iter()
.any(|e| e == &Event::TextDelta("done".into()))
);
assert_eq!(stats.turns, 2);
}
struct CaptureProvider {
reqs: Arc<std::sync::Mutex<Vec<hey::llm::ir::ChatRequest>>>,
}
#[async_trait]
impl hey::llm::Provider for CaptureProvider {
fn model(&self) -> &str {
"test"
}
async fn stream(
&self,
req: &hey::llm::ir::ChatRequest,
on_delta: &mut (dyn FnMut(hey::llm::Delta) + Send),
) -> Result<hey::llm::ir::Completion, hey::llm::LlmError> {
self.reqs.lock().unwrap().push(req.clone());
on_delta(hey::llm::Delta::Text("done".into()));
Ok(hey::llm::ir::Completion {
text: "done".into(),
thinking: String::new(),
tool_calls: vec![],
usage: Default::default(),
})
}
}
static SESSION_LOCK: std::sync::LazyLock<tokio::sync::Mutex<()>> =
std::sync::LazyLock::new(|| tokio::sync::Mutex::new(()));
#[tokio::test]
async fn session_persists_messages() {
let _guard = SESSION_LOCK.lock().await;
let dir = std::env::temp_dir().join(format!("hey-agent-sess-{}", std::process::id()));
hey::sessions::set_dir_override(Some(dir.clone()));
let provider = FakeProvider::new(
"test",
vec![FakeTurn {
text: "final",
tool_calls: vec![],
}],
);
let mut out = TestOutput::new();
let opts = AgentOpts {
session_id: Some("sess-1".into()),
..Default::default()
};
run_agent(
Arc::new(provider),
&Registry::new(),
&mut out,
&cfg(),
"hello session",
&opts,
)
.await;
let msgs = hey::sessions::read_msgs("sess-1");
assert!(
msgs.iter().any(|m| m.plain_text() == "hello session"),
"用户消息写入"
);
assert!(
msgs.iter().any(|m| m.plain_text() == "final"),
"assistant 回复写入"
);
assert!(msgs.iter().all(|m| m.role != hey::llm::ir::Role::System));
let _ = std::fs::remove_dir_all(&dir);
hey::sessions::set_dir_override(None);
}
#[tokio::test]
async fn resume_history_is_included_in_request() {
use hey::llm::ir::{Message, Role};
let history = vec![
Message::text(Role::User, "earlier question"),
Message::text(Role::Assistant, "earlier answer"),
];
let opts = AgentOpts {
resume_msgs: history,
..Default::default()
};
let reqs = Arc::new(std::sync::Mutex::new(Vec::new()));
let provider = CaptureProvider { reqs: reqs.clone() };
let mut out = TestOutput::new();
run_agent(
Arc::new(provider),
&Registry::new(),
&mut out,
&cfg(),
"current prompt",
&opts,
)
.await;
let reqs = reqs.lock().unwrap();
let msgs = &reqs[0].messages;
assert_eq!(msgs.len(), 4, "System + 2 历史 + 本轮");
assert_eq!(msgs[0].role, hey::llm::ir::Role::System);
assert_eq!(msgs[1].plain_text(), "earlier question");
assert_eq!(msgs[2].plain_text(), "earlier answer");
assert_eq!(msgs[3].plain_text(), "current prompt");
}
struct FailingTool;
#[async_trait]
impl Tool for FailingTool {
fn name(&self) -> &str {
"bash"
}
fn description(&self) -> &str {
"failing bash"
}
fn schema(&self) -> serde_json::Value {
serde_json::json!({"type": "object"})
}
async fn run(&self, _ctx: &ToolCtx, _args: serde_json::Value) -> ToolOutput {
ToolOutput::err("bash: no such command (exit code 127)")
}
}
#[tokio::test]
async fn purge_errors_after_removes_failed_inputs() {
let reqs = Arc::new(std::sync::Mutex::new(Vec::new()));
struct LoggingProvider {
inner: FakeProvider,
reqs: Arc<std::sync::Mutex<Vec<hey::llm::ir::ChatRequest>>>,
}
#[async_trait]
impl hey::llm::Provider for LoggingProvider {
fn model(&self) -> &str {
"test"
}
async fn stream(
&self,
req: &hey::llm::ir::ChatRequest,
on_delta: &mut (dyn FnMut(hey::llm::Delta) + Send),
) -> Result<hey::llm::ir::Completion, hey::llm::LlmError> {
self.reqs.lock().unwrap().push(req.clone());
self.inner.stream(req, on_delta).await
}
}
let logging = LoggingProvider {
inner: FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![tool_call("c1", "bash", "{\"cmd\":\"bogus\"}")],
},
FakeTurn {
text: "",
tool_calls: vec![tool_call("c2", "bash", "{\"cmd\":\"still bogus\"}")],
},
FakeTurn {
text: "done",
tool_calls: vec![],
},
],
),
reqs: reqs.clone(),
};
let mut out = TestOutput::new();
let mut c = cfg();
c.compaction = Some(CompactionConfig {
purge_errors_after: 1,
..Default::default()
});
let mut tools = Registry::new();
tools.register(Arc::new(FailingTool));
run_agent(
Arc::new(logging),
&tools,
&mut out,
&c,
"try it",
&AgentOpts::default(),
)
.await;
let reqs = reqs.lock().unwrap();
let last = reqs.last().unwrap();
let texts: Vec<String> = last.messages.iter().map(|m| m.plain_text()).collect();
assert!(
texts
.iter()
.any(|t| t.contains("[pruned: failed call input]")),
"错误输入被剥离: {texts:?}"
);
assert!(
texts.iter().any(|t| t.contains("exit code 127")),
"错误结果保留: {texts:?}"
);
let c1_gone = !last.messages.iter().any(|m| {
m.role == hey::llm::ir::Role::Assistant && m.tool_calls.iter().any(|c| c.id == "c1")
});
assert!(c1_gone, "c1 输入剥离: {texts:?}");
let c2_kept = last.messages.iter().any(|m| {
m.role == hey::llm::ir::Role::Assistant && m.tool_calls.iter().any(|c| c.id == "c2")
});
assert!(c2_kept, "c2 最近保留: {texts:?}");
}
#[tokio::test]
async fn compress_tool_toggle_controls_schema() {
let reqs = Arc::new(std::sync::Mutex::new(Vec::new()));
let capture = CaptureProvider { reqs: reqs.clone() };
let mut out = TestOutput::new();
let mut c = cfg();
c.compaction = Some(CompactionConfig {
compress_tool: false,
..Default::default()
});
run_agent(
Arc::new(capture),
&Registry::new(),
&mut out,
&c,
"x",
&AgentOpts::default(),
)
.await;
{
let reqs = reqs.lock().unwrap();
let names: Vec<String> = reqs[0].tools.iter().map(|t| t.name.clone()).collect();
assert!(
!names.contains(&"compress".to_string()),
"开关关闭时不含 compress: {names:?}"
);
}
let reqs2 = Arc::new(std::sync::Mutex::new(Vec::new()));
let capture2 = CaptureProvider {
reqs: reqs2.clone(),
};
let mut out2 = TestOutput::new();
run_agent(
Arc::new(capture2),
&Registry::new(),
&mut out2,
&cfg(),
"x",
&AgentOpts::default(),
)
.await;
{
let reqs2 = reqs2.lock().unwrap();
let names2: Vec<String> = reqs2[0].tools.iter().map(|t| t.name.clone()).collect();
assert!(
names2.contains(&"compress".to_string()),
"默认含 compress: {names2:?}"
);
}
}
struct SlowTool;
#[async_trait]
impl Tool for SlowTool {
fn name(&self) -> &str {
"slow"
}
fn description(&self) -> &str {
"Slow tool for parallel test"
}
fn schema(&self) -> serde_json::Value {
json!({"type": "object", "properties": {}})
}
async fn run(&self, _ctx: &ToolCtx, _args: serde_json::Value) -> ToolOutput {
tokio::time::sleep(std::time::Duration::from_millis(150)).await;
ToolOutput::ok("slow-done")
}
}
#[tokio::test]
async fn parallel_tool_calls_run_concurrently_and_preserve_order() {
let provider = FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![
tool_call("c1", "slow", "{}"),
tool_call("c2", "slow", "{}"),
tool_call("c3", "slow", "{}"),
],
},
FakeTurn {
text: "all done",
tool_calls: vec![],
},
],
);
let mut tools = Registry::new();
tools.register(Arc::new(SlowTool));
let provider = Arc::new(provider);
let mut out = TestOutput::new();
let t0 = std::time::Instant::now();
let stats = run_agent(
provider.clone(),
&tools,
&mut out,
&cfg(),
"parallel",
&AgentOpts::default(),
)
.await;
let elapsed = t0.elapsed();
assert_eq!(stats.turns, 2);
assert!(
elapsed.as_millis() < 380,
"3 个 150ms 工具应并行完成(实际 {:?})",
elapsed
);
let calls: Vec<String> = out
.events
.iter()
.filter_map(|e| match e {
Event::ToolCall { name, .. } => Some(name.clone()),
_ => None,
})
.collect();
assert_eq!(calls, vec!["slow", "slow", "slow"]);
let results: Vec<String> = out
.events
.iter()
.filter_map(|e| match e {
Event::ToolResult { output, .. } => Some(output.clone()),
_ => None,
})
.collect();
assert_eq!(results, vec!["slow-done", "slow-done", "slow-done"]);
let first_result_at = out
.events
.iter()
.position(|e| matches!(e, Event::ToolResult { .. }))
.expect("有 tool_result 事件");
let last_call_at = out
.events
.iter()
.rposition(|e| matches!(e, Event::ToolCall { .. }))
.expect("有 tool_call 事件");
assert!(
last_call_at < first_result_at,
"所有 tool_call 应先于 tool_result 输出"
);
let mut c = cfg();
c.agent.parallel_tools = Some(false);
let serial_provider = Arc::new(FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![
tool_call("c1", "slow", "{}"),
tool_call("c2", "slow", "{}"),
tool_call("c3", "slow", "{}"),
],
},
FakeTurn {
text: "all done",
tool_calls: vec![],
},
],
));
let mut out2 = TestOutput::new();
let stats2 = run_agent(
serial_provider,
&tools,
&mut out2,
&c,
"serial",
&AgentOpts::default(),
)
.await;
assert_eq!(stats2.turns, 2);
let results2: Vec<String> = out2
.events
.iter()
.filter_map(|e| match e {
Event::ToolResult { output, .. } => Some(output.clone()),
_ => None,
})
.collect();
assert_eq!(results2, vec!["slow-done", "slow-done", "slow-done"]);
}
#[tokio::test]
async fn skills_injected_into_system_prompt() {
use hey::llm::ir::Role;
let reqs = Arc::new(std::sync::Mutex::new(Vec::new()));
let capture = CaptureProvider { reqs: reqs.clone() };
let mut out = TestOutput::new();
let opts = AgentOpts {
skills: vec![hey::skills::Skill {
name: "test-skill".into(),
description: "A test skill".into(),
path: std::path::PathBuf::from("/tmp/fake-skill/SKILL.md"),
}],
..Default::default()
};
run_agent(
Arc::new(capture),
&Registry::new(),
&mut out,
&cfg(),
"hello",
&opts,
)
.await;
let reqs = reqs.lock().unwrap();
let system = reqs[0]
.messages
.iter()
.find(|m| m.role == Role::System)
.expect("应有 System 消息");
assert!(
system.plain_text().contains("test-skill"),
"技能名应注入系统提示"
);
assert!(
system.plain_text().contains("A test skill"),
"技能描述应注入"
);
}
#[tokio::test]
async fn dedup_removes_duplicate_call() {
let provider = FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![tool_call("c1", "bash", "{\"cmd\":\"ls\"}")],
},
FakeTurn {
text: "",
tool_calls: vec![tool_call("c2", "bash", "{\"cmd\":\"ls\"}")],
},
FakeTurn {
text: "done",
tool_calls: vec![],
},
],
);
let mut tools = Registry::new();
tools.register(Arc::new(FakeTool));
let mut out = TestOutput::new();
let mut c = cfg();
c.compaction = Some(CompactionConfig {
dedup: true,
..Default::default()
});
let stats = run_agent(
Arc::new(provider),
&tools,
&mut out,
&c,
"list twice",
&AgentOpts::default(),
)
.await;
assert_eq!(stats.turns, 3);
let results: Vec<&str> = out
.events
.iter()
.filter_map(|e| match e {
Event::ToolResult { output, .. } => Some(output.as_str()),
_ => None,
})
.collect();
assert!(results.len() == 2, "应有 2 个 tool_result, got {results:?}");
assert!(
results[0] == "output-ok" || results[0].contains("pruned"),
"第一次正常输出或 pruned: {results:?}"
);
}
struct MultiTurnCaptureProvider {
reqs: Arc<std::sync::Mutex<Vec<hey::llm::ir::ChatRequest>>>,
turn: std::sync::atomic::AtomicU32,
}
#[async_trait]
impl hey::llm::Provider for MultiTurnCaptureProvider {
fn model(&self) -> &str {
"capture-thinking"
}
async fn stream(
&self,
req: &hey::llm::ir::ChatRequest,
on_delta: &mut (dyn FnMut(hey::llm::Delta) + Send),
) -> Result<hey::llm::ir::Completion, hey::llm::LlmError> {
self.reqs.lock().unwrap().push(req.clone());
let i = self.turn.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if i == 0 {
on_delta(hey::llm::Delta::Thinking("inner reasoning".to_string()));
on_delta(hey::llm::Delta::Text("let me check".to_string()));
Ok(hey::llm::ir::Completion {
text: "let me check".into(),
thinking: "inner reasoning".into(),
tool_calls: vec![tool_call("c1", "bash", "{}")],
usage: Default::default(),
})
} else {
on_delta(hey::llm::Delta::Text("done".to_string()));
Ok(hey::llm::ir::Completion {
text: "done".into(),
thinking: String::new(),
tool_calls: vec![],
usage: Default::default(),
})
}
}
}
#[tokio::test]
async fn thinking_round_trip() {
use hey::llm::ir::ContentBlock;
let reqs = Arc::new(std::sync::Mutex::new(Vec::new()));
let provider = MultiTurnCaptureProvider {
reqs: reqs.clone(),
turn: std::sync::atomic::AtomicU32::new(0),
};
let mut tools = Registry::new();
tools.register(Arc::new(FakeTool));
let mut out = TestOutput::new();
let mut c = cfg();
c.agent.max_turns = Some(3);
run_agent(
Arc::new(provider),
&tools,
&mut out,
&c,
"think",
&AgentOpts::default(),
)
.await;
let reqs = reqs.lock().unwrap();
assert!(reqs.len() >= 2, "至少 2 轮请求");
let second_req = &reqs[1];
let assistant_msgs: Vec<&hey::llm::ir::Message> = second_req
.messages
.iter()
.filter(|m| m.role == hey::llm::ir::Role::Assistant)
.collect();
let has_thinking = assistant_msgs.iter().any(|m| {
m.content
.iter()
.any(|b| matches!(b, ContentBlock::Thinking(_)))
});
assert!(has_thinking, "第二轮请求应回传 Thinking 块");
let thinking_text: String = assistant_msgs
.iter()
.flat_map(|m| &m.content)
.filter_map(|b| match b {
ContentBlock::Thinking(t) => Some(t.as_str()),
_ => None,
})
.collect();
assert!(
thinking_text.contains("inner reasoning"),
"Thinking 内容保留: {thinking_text}"
);
}
#[tokio::test]
async fn context_compaction_triggers() {
use hey::llm::ir::Role;
let reqs = Arc::new(std::sync::Mutex::new(Vec::new()));
let provider = MultiTurnCaptureProvider {
reqs: reqs.clone(),
turn: std::sync::atomic::AtomicU32::new(0),
};
let mut tools = Registry::new();
tools.register(Arc::new(FakeTool));
let mut out = TestOutput::new();
let mut c = cfg();
c.agent.max_turns = Some(3);
c.agent.budget_tokens = Some(10); let history = vec![
hey::llm::ir::Message::text(Role::User, "long history ".repeat(50)),
hey::llm::ir::Message::text(Role::Assistant, "long answer ".repeat(50)),
];
let opts = AgentOpts {
resume_msgs: history,
..Default::default()
};
run_agent(
Arc::new(provider),
&tools,
&mut out,
&c,
"new prompt",
&opts,
)
.await;
let reqs = reqs.lock().unwrap();
assert!(!reqs.is_empty(), "至少发送了请求");
let msg_count = reqs[0].messages.len();
assert!(msg_count < 10, "预算 10 应大幅裁剪: {} 条消息", msg_count);
}
#[tokio::test]
async fn session_no_duplicate_messages_multi_turn() {
let _guard = SESSION_LOCK.lock().await;
let ns = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let dir = std::env::temp_dir().join(format!("hey-dup-{}-{}", std::process::id(), ns));
hey::sessions::set_dir_override(Some(dir.clone()));
let provider = FakeProvider::new(
"test",
vec![
FakeTurn {
text: "",
tool_calls: vec![tool_call("c1", "bash", "{\"cmd\":\"ls\"}")],
},
FakeTurn {
text: "",
tool_calls: vec![tool_call("c2", "bash", "{\"cmd\":\"pwd\"}")],
},
FakeTurn {
text: "done",
tool_calls: vec![],
},
],
);
let mut tools = Registry::new();
tools.register(Arc::new(FakeTool));
let mut out = TestOutput::new();
let opts = AgentOpts {
session_id: Some("dup-1".into()),
..Default::default()
};
run_agent(
Arc::new(provider),
&tools,
&mut out,
&cfg(),
"multi turn",
&opts,
)
.await;
let msgs = hey::sessions::read_msgs("dup-1");
let mut seen = std::collections::HashSet::new();
for m in &msgs {
assert!(
seen.insert(m.id.clone()),
"消息重复持久化: id={} text={}",
m.id,
m.plain_text().chars().take(40).collect::<String>()
);
}
assert_eq!(
msgs.len(),
6,
"期望 6 条消息,实际 {} 条: {:?}",
msgs.len(),
msgs.iter()
.map(|m| m.plain_text().chars().take(15).collect::<String>())
.collect::<Vec<_>>()
);
let _ = std::fs::remove_dir_all(&dir);
hey::sessions::set_dir_override(None);
}