2d0d5c1592
新增 7 个示例覆盖全 Phase 公共 API(不依赖 API key 即可运行): - agent_session_demo:AgentBuilder → AgentSession → SessionMemory - custom_tool:BaseTool 注册 + invoke/invoke_all + 权限检查 - prompt_composer:PromptTemplate + PromptComposer + validate_messages - task_agent_demo:JsonPlanParser + Step 状态机 + 错误路径 - conversation_memory_demo:滑动窗口 + 多角色 + 隔离 - knowledge_search_demo:KnowledgeStore + 关键词检索 + 停用词过滤 - streaming_events_demo:submit_stream 事件消费 + 队列耗尽错误路径
130 lines
4.1 KiB
Rust
130 lines
4.1 KiB
Rust
//! agent_session_demo —— Agent 装配 + 会话链路 + SessionMemory 桥接。
|
||
//!
|
||
//! 演示:
|
||
//! 1. 实现 `Agent` trait(定义角色 + system prompt)
|
||
//! 2. 用 `MockProvider` 预设响应(离线可跑)
|
||
//! 3. `AgentBuilder` 装配 `RuntimeBundle`
|
||
//! 4. `AgentSession::submit_turn` 跑多轮对话
|
||
//! 5. `SessionMemory` 读写 + snapshot 输出
|
||
//! 6. 跨 session 数据隔离验证
|
||
//!
|
||
//! 运行:`cargo run --example agent_session_demo`
|
||
|
||
use std::sync::Arc;
|
||
|
||
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||
use agcore::llm::hooks::HookExecutor;
|
||
use agcore::llm::mock::MockProvider;
|
||
use agcore::llm::types::message::{ContentBlock, Message};
|
||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||
use agcore::llm::types::Usage;
|
||
use agcore::tools::ToolRegistry;
|
||
|
||
/// 计算器角色 Agent。
|
||
struct CalculatorAgent;
|
||
|
||
impl Agent for CalculatorAgent {
|
||
fn name(&self) -> &str {
|
||
"calculator"
|
||
}
|
||
fn system_prompt(&self) -> Option<&str> {
|
||
Some("你是一个简洁的计算器助手,每轮回答一句话。")
|
||
}
|
||
}
|
||
|
||
/// 构造预设的纯文本 Assistant 响应。
|
||
fn assistant_text(text: &str) -> MessageResponse {
|
||
MessageResponse {
|
||
id: String::new(),
|
||
model: String::new(),
|
||
message: Message::Assistant {
|
||
content: vec![ContentBlock::Text { text: text.into() }],
|
||
},
|
||
usage: Usage::from_input_output(8, 4),
|
||
stop_reason: StopReason::Stop,
|
||
extra: Default::default(),
|
||
}
|
||
}
|
||
|
||
#[tokio::main]
|
||
async fn main() {
|
||
// 1. MockProvider:预设三轮响应(无须 API key 即可离线运行)
|
||
let provider = Arc::new(MockProvider::new(vec![
|
||
assistant_text("1 + 1 = 2"),
|
||
assistant_text("2 + 2 = 4"),
|
||
assistant_text("会话即将结束。"),
|
||
]));
|
||
|
||
// 2. AgentBuilder 装配 RuntimeBundle(必填:provider / tool_registry / hook_executor)
|
||
let bundle = Arc::new(
|
||
AgentBuilder::new()
|
||
.provider(provider)
|
||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||
.hook_executor(Arc::new(HookExecutor::new()))
|
||
.build()
|
||
.expect("RuntimeBundle 装配失败"),
|
||
);
|
||
|
||
// 3. 创建会话
|
||
let agent: Arc<dyn Agent> = Arc::new(CalculatorAgent);
|
||
let mut session = AgentSession::new(agent, "demo-session", bundle.clone());
|
||
assert_eq!(session.turn_index(), 0);
|
||
|
||
// 4. 提交第一轮
|
||
println!("=== 提交第 1 轮 ===");
|
||
let resp = session.submit_turn("1+1=?").await.expect("submit_turn 失败");
|
||
println!("LLM: {}", resp.text());
|
||
session
|
||
.set_session_data("last_q", "1+1=?")
|
||
.await
|
||
.expect("set_session_data 失败");
|
||
session
|
||
.set_session_data("last_a", resp.text())
|
||
.await
|
||
.expect("set_session_data 失败");
|
||
assert_eq!(session.turn_index(), 1);
|
||
|
||
// 5. 提交第二轮
|
||
println!("\n=== 提交第 2 轮 ===");
|
||
let resp = session.submit_turn("再加一次 2+2=?").await.unwrap();
|
||
println!("LLM: {}", resp.text());
|
||
assert_eq!(session.turn_index(), 2);
|
||
|
||
// 6. 验证 SessionMemory 读取
|
||
println!("\n=== Session Memory 读取 ===");
|
||
println!(
|
||
"last_q = {:?}",
|
||
session.get_session_data("last_q").await.unwrap()
|
||
);
|
||
println!(
|
||
"last_a = {:?}",
|
||
session.get_session_data("last_a").await.unwrap()
|
||
);
|
||
|
||
// 7. Snapshot 格式化输出
|
||
println!("\n=== Session Memory Snapshot ===");
|
||
println!("{}", session.session_memory().snapshot().await.unwrap());
|
||
|
||
// 8. 跨 session 数据隔离验证
|
||
println!("=== 数据隔离验证 ===");
|
||
let other = AgentSession::new(
|
||
Arc::new(CalculatorAgent),
|
||
"other-session",
|
||
bundle,
|
||
);
|
||
assert!(
|
||
other.get_session_data("last_q").await.unwrap().is_none(),
|
||
"新会话不应看到旧 session 的 last_q"
|
||
);
|
||
println!("新会话 last_q = None ✓");
|
||
|
||
// 9. 用量累计验证
|
||
println!("\n=== 用量累计 ===");
|
||
let total = session.usage().total();
|
||
println!(
|
||
"prompt={}, completion={}, total={}",
|
||
total.prompt_tokens, total.completion_tokens, total.total_tokens
|
||
);
|
||
|
||
println!("\n✓ agent_session_demo 完成");
|
||
} |