- README 添加 feature 组合表 + 模块级 features 清单 + 升级指南 - 18 个 example 顶部添加 Required features 注释 - roadmap.md 和 roadmap-v0.3.2.md 同步 Phase 26-27 完成状态 - cargo fmt 全量格式化(修复预存格式问题,CI format job 可通过)
219 lines
8.1 KiB
Rust
219 lines
8.1 KiB
Rust
//! Agent 角色热切换(Phase 18)。
|
||
//!
|
||
//! 提供 `SessionManager::switch_agent()`,运行时替换 session 绑定的 Agent,
|
||
//! 保留 slot 历史 / turn_index / session_memory / cost_so_far。
|
||
//!
|
||
//! **未来扩展方向**(roadmap 备注):
|
||
//! - v0.4 可考虑提供 `switch_agent_with_rollback`,在切换前自动 checkpoint
|
||
//! - 可在 `SessionMeta` 中记录 `previous_agent_name` 支持审计历史
|
||
//! - 可新增 `switch_history` API 暴露切换时间序列
|
||
|
||
use std::sync::Arc;
|
||
|
||
use crate::agent::Agent;
|
||
use crate::engine::error::EngineError;
|
||
use crate::engine::session_manager::{SessionManager, SessionMeta};
|
||
|
||
impl SessionManager {
|
||
/// 热切换指定 session 的 Agent 角色。
|
||
///
|
||
/// # 行为
|
||
///
|
||
/// - **保留上下文**:slot 历史 / turn_index / session_memory / cost_so_far 全部保留
|
||
/// - **更新 SessionMeta**:`agent_name` 替换为新 agent,`created_at`/`parent_id` 保持原始
|
||
/// - **不自动 checkpoint**:与 `auto_checkpoint` 语义一致(仅 `submit_turn`/`finalize_turn` 触发)
|
||
///
|
||
/// # 注意
|
||
///
|
||
/// 切换后**新的 system_prompt 将与已有对话历史共存**。建议在切换后
|
||
/// 发送一条明确的上下文过渡提示(如"你现在以新角色 X 的身份继续对话")
|
||
/// 作为切换后的首条输入,以避免 LLM 误解对话历史。
|
||
///
|
||
/// # 安全提示
|
||
///
|
||
/// `AgentSession.agent` 是 `pub` 字段可直接访问。**绕过 `switch_agent`
|
||
/// 直接修改会导致 SessionMeta 中的 `agent_name` 与内存状态不一致**,
|
||
/// 请始终使用此方法。
|
||
pub async fn switch_agent(
|
||
&self,
|
||
session_id: &str,
|
||
new_agent: Arc<dyn Agent>,
|
||
) -> Result<(), EngineError> {
|
||
// 1. get session(RwLock 读锁,返回后释放)
|
||
let session = self.get(session_id).await?;
|
||
|
||
// 2. lock Mutex,替换 agent,读 name + turn_index
|
||
let (agent_name, turn_index) = {
|
||
let mut guard = session.lock().await;
|
||
guard.agent = new_agent;
|
||
(guard.agent.name().to_string(), guard.turn_index())
|
||
}; // 释放 Mutex
|
||
|
||
// 3. 读取原始 SessionMeta(用于保留 created_at / parent_id)
|
||
let existing_meta = self
|
||
.load_session_meta(session_id)
|
||
.await?
|
||
.ok_or_else(|| EngineError::SessionNotFound(session_id.to_string()))?;
|
||
|
||
// 4. 构造新 meta 并持久化(I/O,无锁)
|
||
let meta = SessionMeta {
|
||
session_id: session_id.to_string(),
|
||
agent_name,
|
||
parent_id: existing_meta.parent_id,
|
||
created_at: existing_meta.created_at,
|
||
turn_count: turn_index,
|
||
};
|
||
self.save_session_meta(&meta).await?;
|
||
|
||
tracing::info!(
|
||
session_id = %session_id,
|
||
agent_name = %meta.agent_name,
|
||
previous_agent = %existing_meta.agent_name,
|
||
"agent switched"
|
||
);
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
use crate::agent::Agent;
|
||
use crate::agent::builder::AgentBuilder;
|
||
use crate::agent::runtime::RuntimeBundle;
|
||
use crate::engine::session_manager::SessionManager;
|
||
use crate::llm::hooks::HookExecutor;
|
||
use crate::llm::mock::MockProvider;
|
||
use crate::llm::types::Usage;
|
||
use crate::llm::types::message::{ContentBlock, Message};
|
||
use crate::llm::types::response_v2::{MessageResponse, StopReason};
|
||
use crate::memory::store::InMemoryStore;
|
||
use crate::tools::ToolRegistry;
|
||
use std::sync::Arc;
|
||
|
||
/// 测试用 MockAgent(name + system_prompt 可控)
|
||
struct MockAgent {
|
||
name: String,
|
||
system_prompt: String,
|
||
}
|
||
|
||
impl MockAgent {
|
||
fn new(name: &str, system_prompt: &str) -> Self {
|
||
Self {
|
||
name: name.to_string(),
|
||
system_prompt: system_prompt.to_string(),
|
||
}
|
||
}
|
||
}
|
||
|
||
impl Agent for MockAgent {
|
||
fn name(&self) -> &str {
|
||
&self.name
|
||
}
|
||
fn system_prompt(&self) -> Option<&str> {
|
||
Some(&self.system_prompt)
|
||
}
|
||
}
|
||
|
||
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(),
|
||
}
|
||
}
|
||
|
||
async fn make_manager_and_bundle() -> (Arc<SessionManager>, Arc<RuntimeBundle>) {
|
||
let store: Arc<dyn crate::memory::store::MemoryStore> = Arc::new(InMemoryStore::new());
|
||
let provider = Arc::new(MockProvider::new(vec![
|
||
assistant_text("response1"),
|
||
assistant_text("response2"),
|
||
assistant_text("response3"),
|
||
]));
|
||
let bundle = Arc::new(
|
||
AgentBuilder::new()
|
||
.provider(provider)
|
||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||
.hook_executor(Arc::new(HookExecutor::new()))
|
||
.session_memory_backend(store.clone())
|
||
.build()
|
||
.expect("RuntimeBundle build"),
|
||
);
|
||
let sm = Arc::new(SessionManager::new(store));
|
||
(sm, bundle)
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn test_switch_agent_basic() {
|
||
let (sm, bundle) = make_manager_and_bundle().await;
|
||
let agent_a: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||
let agent_b: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_b", "I am B"));
|
||
|
||
let sid = sm.create(agent_a, bundle).await.unwrap();
|
||
sm.switch_agent(&sid, agent_b).await.unwrap();
|
||
|
||
let session = sm.get(&sid).await.unwrap();
|
||
let guard = session.lock().await;
|
||
assert_eq!(guard.agent.name(), "agent_b");
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn test_switch_agent_preserves_context() {
|
||
let (sm, bundle) = make_manager_and_bundle().await;
|
||
let agent_a: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||
let agent_b: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_b", "I am B"));
|
||
|
||
let sid = sm.create(agent_a, bundle).await.unwrap();
|
||
|
||
// 在 switch 前写入 session_data 并提交一轮
|
||
{
|
||
let session = sm.get(&sid).await.unwrap();
|
||
let mut guard = session.lock().await;
|
||
guard.set_session_data("key1", "value1").await.unwrap();
|
||
}
|
||
sm.submit_turn(&sid, "hello").await.unwrap();
|
||
|
||
// switch
|
||
sm.switch_agent(&sid, agent_b).await.unwrap();
|
||
|
||
// 验证 turn_index 保留
|
||
let session = sm.get(&sid).await.unwrap();
|
||
let guard = session.lock().await;
|
||
assert_eq!(guard.turn_index(), 1, "turn_index should be preserved");
|
||
// 验证 session_memory 保留
|
||
let val = guard.session_memory().get("key1").await.unwrap();
|
||
assert_eq!(val, Some("value1".to_string()));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn test_switch_agent_updates_session_meta() {
|
||
let (sm, bundle) = make_manager_and_bundle().await;
|
||
let agent_a: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||
let agent_b: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_b", "I am B"));
|
||
|
||
let sid = sm.create(agent_a, bundle).await.unwrap();
|
||
sm.switch_agent(&sid, agent_b).await.unwrap();
|
||
|
||
// 通过 load_session_meta 验证持久化
|
||
let meta = sm.load_session_meta(&sid).await.unwrap();
|
||
assert!(meta.is_some(), "SessionMeta should persist");
|
||
let meta = meta.unwrap();
|
||
assert_eq!(meta.agent_name, "agent_b");
|
||
// created_at / parent_id 保持
|
||
assert_eq!(meta.parent_id, None, "parent_id should remain None");
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn test_switch_agent_session_not_found() {
|
||
let (sm, _bundle) = make_manager_and_bundle().await;
|
||
let agent: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||
let result = sm.switch_agent("nonexistent_session_id", agent).await;
|
||
assert!(matches!(result, Err(EngineError::SessionNotFound(_))));
|
||
}
|
||
}
|