Files
agcore/src/engine/switch.rs
T
徐涛 5baa170508 docs: 更新 README feature 表 + 升级指南 + 示例注释 + roadmap 同步
- README 添加 feature 组合表 + 模块级 features 清单 + 升级指南
- 18 个 example 顶部添加 Required features 注释
- roadmap.md 和 roadmap-v0.3.2.md 同步 Phase 26-27 完成状态
- cargo fmt 全量格式化(修复预存格式问题,CI format job 可通过)
2026-07-19 08:18:04 +08:00

219 lines
8.1 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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 sessionRwLock 读锁,返回后释放)
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;
/// 测试用 MockAgentname + 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(_))));
}
}