Compare commits
32
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2af92cd554 | ||
|
|
635942248b | ||
|
|
fe51961202 | ||
|
|
212cfcc916 | ||
|
|
88d00ac927 | ||
|
|
358e971094 | ||
|
|
85b92ae9d4 | ||
|
|
57b2fbaaed | ||
|
|
2c8e31919d | ||
|
|
c6651c9b75 | ||
|
|
e636e16820 | ||
|
|
1c89d23ba2 | ||
|
|
5b4343a051 | ||
|
|
6e1182e64c | ||
|
|
b8f4fe0fe3 | ||
|
|
7574f9c24c | ||
|
|
c82af60f81 | ||
|
|
c8a91f6eaf | ||
|
|
13edacd775 | ||
|
|
821cea8e60 | ||
|
|
517ef7db32 | ||
|
|
4cf5918b9c | ||
|
|
9da9b83167 | ||
|
|
b1875192fd | ||
|
|
d3067e2f53 | ||
|
|
5648b1d217 | ||
|
|
98dfe6c1ed | ||
|
|
9e476e79bb | ||
|
|
3bd135ec98 | ||
|
|
76f3235ed7 | ||
|
|
6315f2d008 | ||
|
|
fba78f5f33 |
@@ -198,6 +198,18 @@ pub use vector_store::VectorStore;
|
||||
5. **风险评估** - 潜在风险、缓解措施
|
||||
6. **验收标准** - 可验证的完成条件
|
||||
|
||||
### 进度同步规范 (docs/roadmap.md)
|
||||
|
||||
完成一项实施后,必须检查 `docs/roadmap.md` 是否存在对应内容;若存在,必须同步标记为完成:
|
||||
|
||||
- **Step / Phase 状态行**:对应 Step 加 ✅ 标记;Phase 章节末尾「状态」行从 ⏳ 改为 ✅ Phase X 全部交付物已完成
|
||||
- **里程碑表**:更新对应里程碑状态从 ⏳ 改为 ✅ + 完成日期
|
||||
- **依赖关系图(Mermaid)**:节点 `class` 从 `pending` / `core` 改为 `done`,必要时更新节点摘要
|
||||
- **文末「已完成 / 进行中阶段」列表**:追加一行 `- ✅ Phase X — 一句话要点`
|
||||
- **顶部「当前状态」**:补充新完成 Phase,更新「下一步」指向
|
||||
|
||||
参考案例:2026-07-05 完成 Phase 7 SqliteStore 时同步更新 6 处(顶部状态 / Phase 章节 / 依赖图 / M3 / 下一步行动 / 已完成列表)。
|
||||
|
||||
---
|
||||
|
||||
## 项目特定规则
|
||||
|
||||
@@ -2,6 +2,67 @@
|
||||
|
||||
本项目所有重要变更均记录于此文件。格式参考 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/)。
|
||||
|
||||
## [0.2.0-rc.1] - 2026-07-05
|
||||
|
||||
v0.2.0 候选发布。Phase 5-7 三大 P0 全部交付完成,API 稳定性扫尾,新增 2 个面向新用户的集成示例。
|
||||
|
||||
### Added
|
||||
|
||||
**Phase 5 — 热身准备**
|
||||
- `ProviderConfig::from_env(prefix)`:从 `{prefix}_BASE_URL` / `{prefix}_API_KEY` / `{prefix}_MODEL` / `{prefix}_TIMEOUT_SECS` / `{prefix}_MAX_RETRIES` 环境变量构造配置
|
||||
- `ProviderConfig::timeout_secs` / `max_retries` 字段(默认 30 / 3)
|
||||
- `OllamaProvider`:本地推理 Provider(OpenAI-compatible,`http://localhost:11434/v1` 默认端点)
|
||||
- `ProviderType::Ollama` 变体 + `FromStr` 解析
|
||||
|
||||
**Phase 6 — ToolDef IR 正式化**
|
||||
- `ToolDef` 结构体(name / description / parameters)替代已废弃的 `OpenaiToolDefinition`
|
||||
- `MessageRequest.tools` 切换为 `Vec<ToolDef>`
|
||||
- `OpenaiToolDefinition` 降级为 `#[doc(hidden)]`,仅供 OpenAI 适配层内部消费
|
||||
|
||||
**Phase 7 — SqliteStore 持久化**
|
||||
- `SqliteStore`:`MemoryStore` 的 SQLite 后端实现,基于 `rusqlite 0.32` bundled
|
||||
- WAL 模式 + `synchronous=NORMAL` + `busy_timeout=5s` 兼顾崩溃安全与吞吐
|
||||
- `Arc<Mutex<Connection>>` + `spawn_blocking` 卸载 IO;10×10 并发写入无 race
|
||||
- `MemoryStore::save / get / delete / list` CRUD + prefix / since / offset+limit 过滤
|
||||
- 进程重启数据不丢的 round-trip 验证
|
||||
|
||||
**Phase 8 — MVP 集成出口**
|
||||
- `examples/quick_start`:30 行最小可运行示例(MockProvider + EchoTool + submit_turn)
|
||||
- `examples/end_to_end`:3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证
|
||||
|
||||
### Changed
|
||||
|
||||
- **API 稳定性护栏**:14 个公开枚举标记 `#[non_exhaustive]`,覆盖 P0 核心 IR(`Message` / `ContentBlock` / `ContentBlockType` / `StreamEvent` / `HookEvent`)、P0 Error(`AgentError` / `LlmError` / `ToolError` / `MemoryError` / `PromptError`)、P1 其他(`MemoryStrategy` / `StepStatus` / `ToolChoice` / `ResponseFormat`)
|
||||
- **`StepStatus::Completed` 字段类型**:从废弃的 `ChatResponse` 切换为 IR 层 `MessageResponse`(同时清理 `task_agent_demo.rs` 的 `ChatResponse` / `OpenaiChatMessage` / `FinishReason` 三处废弃类型引用)
|
||||
|
||||
### Non-exhaustive 清单
|
||||
|
||||
为防止未来新增变体时下游 exhaustive match 静默失效,14 个枚举追加 `#[non_exhaustive]`:
|
||||
|
||||
| 优先级 | 枚举 |
|
||||
|--------|------|
|
||||
| P0 核心 IR | `Message`, `ContentBlock`, `ContentBlockType`, `StreamEvent`, `HookEvent` |
|
||||
| P0 Error | `AgentError`, `LlmError`, `ToolError`, `MemoryError`, `PromptError` |
|
||||
| P1 其他 | `MemoryStrategy`, `StepStatus`, `ToolChoice`, `ResponseFormat` |
|
||||
|
||||
明确不加:内部 wire-format(`OpenaiChatMessage` 等)/ 语义已收敛(`Role` / `ServiceTier` / `Modality` / `ImageDetail` / `AudioFormat` / `StopSequence`)/ 使用面窄(`Permission` / `McpTransport` 等)。
|
||||
|
||||
### Deprecated
|
||||
|
||||
(继承自 0.1.0,无新增)`ChatResponse` / `ToolDefinition` 保持 `#[deprecated]` 标记。
|
||||
|
||||
### Fixed
|
||||
|
||||
- 修复 `StepStatus::Completed(ChatResponse)` 字段类型与 IR 体系不一致问题(已完成迁移)
|
||||
|
||||
### Migration Guide (v0.1 → v0.2.0-rc.1)
|
||||
|
||||
1. **枚举 match**:14 个 `#[non_exhaustive]` 枚举在 crate 外必须使用 `_ =>` 通配分支
|
||||
2. **`StepStatus::Completed`**:字段类型从 `ChatResponse` 切换为 `MessageResponse`,需做字段映射(参考 `docs/15-phase8-mvp-integration.md` §3.1.2)
|
||||
3. **`ToolDefinition` → `ToolDef`**:Phase 6 已彻底替换 `#[deprecated]` 别名,需全局重命名
|
||||
|
||||
---
|
||||
|
||||
## [0.1.0] - 2026-07-04
|
||||
|
||||
首个公开版本。涵盖 Phase 0-4c 的全部核心能力、Provider IR 重构、LlmCycle 简化,以及面向用户的 7 个离线示例。
|
||||
|
||||
+5
-2
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "agcore"
|
||||
version = "0.1.0"
|
||||
version = "0.2.0-rc.1"
|
||||
edition = "2024"
|
||||
|
||||
[dependencies]
|
||||
@@ -19,8 +19,11 @@ futures-core = "0.3"
|
||||
bytes = "1"
|
||||
async-stream = "0.3"
|
||||
tokio-util = { version = "0.7", features = ["rt"] }
|
||||
time = { version = "0.3", features = ["serde"] }
|
||||
time = { version = "0.3", features = ["serde", "parsing", "formatting", "macros"] }
|
||||
rusqlite = { version = "0.32", features = ["bundled"] }
|
||||
|
||||
[dev-dependencies]
|
||||
dotenvy = "0.15.7"
|
||||
wiremock = "0.6"
|
||||
temp-env = "0.3"
|
||||
tempfile = "3"
|
||||
|
||||
@@ -26,7 +26,7 @@ AG Core 不是 Agent 产品,而是 Agent 的**底层依赖库**:上层应用
|
||||
|
||||
```toml
|
||||
[dependencies]
|
||||
agcore = "0.1"
|
||||
agcore = "0.2"
|
||||
tokio = { version = "1", features = ["macros", "rt-multi-thread"] }
|
||||
```
|
||||
|
||||
@@ -110,10 +110,12 @@ let provider = create_provider(
|
||||
).expect("创建 Provider 失败");
|
||||
```
|
||||
|
||||
更多端到端示例见 [`examples/`](./examples/) 目录(共 7 个,全部可 `cargo run --example <name>`):
|
||||
更多端到端示例见 [`examples/`](./examples/) 目录(共 10 个,全部可 `cargo run --example <name>`):
|
||||
|
||||
| 示例 | 说明 |
|
||||
|------|------|
|
||||
| `quick_start` | **30 行最小示例**:MockProvider + EchoTool + submit_turn,新用户 5 分钟上手 |
|
||||
| `end_to_end` | **完整集成示例**:3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证 |
|
||||
| `agent_session_demo` | Agent + 会话 + SessionMemory 完整链路(MockProvider 离线) |
|
||||
| `custom_tool` | 自定义工具注册、单次 / 并行调用、权限检查 |
|
||||
| `prompt_composer` | 提示词模板与组合器(纯离线) |
|
||||
@@ -121,6 +123,7 @@ let provider = create_provider(
|
||||
| `conversation_memory_demo` | 对话记忆滑动窗口与隔离 |
|
||||
| `knowledge_search_demo` | 知识页面关键词检索 |
|
||||
| `streaming_events_demo` | LLM 流式响应事件消费(含错误路径) |
|
||||
| `simple_visit` | 真实 LLM 调用(OpenAI / Anthropic,设置 `OPENAI_*` / `ANTHROPIC_*` 环境变量) |
|
||||
|
||||
## 核心模块
|
||||
|
||||
|
||||
@@ -0,0 +1,522 @@
|
||||
# Phase 5:热身准备 — 实施方案
|
||||
|
||||
## 1. 背景与目标
|
||||
|
||||
Phase 5 是 v0.2.0 发布周期的**热身准备阶段**,包含三个互不依赖的 Step,为后续 Phase 6-12 的端到端集成提供基础设施。
|
||||
|
||||
**核心目标**:
|
||||
- 为 Phase 8(端到端示例)提供零 API key 的运行路径(Ollama)
|
||||
- 为公共枚举的向后兼容性加上编译期护栏(`#[non_exhaustive]`)
|
||||
- 为 Provider 构造提供统一的超时与重试配置入口(`ProviderConfig` 扩展)
|
||||
|
||||
三个 Step 之间**无依赖关系**,但出于实现效率考虑,按 **5.2 → 5.3 → 5.1** 顺序执行。理由:5.2 先新增 `Ollama` 枚举变体,5.3 再加 `#[non_exhaustive]`,避免枚举标记后添加变体需要在外部 crate 加 `_ =>` 兜底分支的困扰。
|
||||
|
||||
## 2. 需求分析
|
||||
|
||||
### Step 5.2 — Ollama Provider
|
||||
|
||||
| 维度 | 内容 |
|
||||
|------|------|
|
||||
| **需求** | 新增 `OllamaProvider`,newtype 包装 `GenericOpenaiProvider`,默认连接本地 Ollama 实例 |
|
||||
| **优先级** | P0 — 为 Phase 8 端到端示例提供无需 API key 的运行路径 |
|
||||
| **预期交付物** | `src/llm/provider/ollama.rs` 新建文件;`ProviderType` 新增 `Ollama` 变体 |
|
||||
| **代码量** | ~55 行 |
|
||||
|
||||
### Step 5.3 — `#[non_exhaustive]` 前置标记
|
||||
|
||||
| 维度 | 内容 |
|
||||
|------|------|
|
||||
| **需求** | 为 4 个公共枚举添加 `#[non_exhaustive]` 属性,避免后续新增变体时破坏下游 match |
|
||||
| **优先级** | P1 — 编译期兼容性保障 |
|
||||
| **预期交付物** | 修改 4 个枚举定义,各加一行属性 |
|
||||
| **代码量** | ~4 行 |
|
||||
|
||||
### Step 5.1 — ProviderConfig 扩展
|
||||
|
||||
| 维度 | 内容 |
|
||||
|------|------|
|
||||
| **需求** | `ProviderConfig` 新增 `timeout_secs` 和 `max_retries` 字段;实现 `Default`、`from_env()` 构造;timeout 传导到各 Provider HTTP Client |
|
||||
| **优先级** | P0 — 与 Roadmap 一致,Phase 8(MVP 出口)依赖 from_env |
|
||||
| **预期交付物** | `ProviderConfig` 扩展;`create_provider()` 超时注入;`from_env()` + 单元测试 |
|
||||
| **代码量** | ~60 行 + 测试 |
|
||||
|
||||
## 3. 方案设计
|
||||
|
||||
### 3.1 Step 5.2 — Ollama Provider(先执行)
|
||||
|
||||
#### 改动文件清单
|
||||
|
||||
| 文件 | 操作 | 说明 |
|
||||
|------|------|------|
|
||||
| `src/llm/provider/ollama.rs` | **新建** | OllamaProvider newtype 包装 |
|
||||
| `src/llm/provider.rs` | 修改 | `ProviderType` 新增 `Ollama` 变体;`FromStr` 加解析;`create_provider()` 加分支 |
|
||||
| `src/llm/provider/mod.rs` 或其他模块注册文件 | 修改(如需要) | 注册 `pub mod ollama` |
|
||||
|
||||
#### 关键代码
|
||||
|
||||
**`src/llm/provider/ollama.rs`**(新建):
|
||||
|
||||
```rust
|
||||
//! Ollama Provider —— OpenAI-compatible 协议的 newtype 包装,零 API key。
|
||||
//!
|
||||
//! 默认 base_url = `http://localhost:11434/v1`,空 api_key 也可工作。
|
||||
//! 实现方式同 DeepSeekProvider / QwenProvider,共享 GenericOpenaiProvider 的 HTTP/SSE/转换逻辑。
|
||||
|
||||
use reqwest::Client;
|
||||
use std::pin::Pin;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use futures_core::Stream;
|
||||
|
||||
use super::openai::GenericOpenaiProvider;
|
||||
use super::{LlmProvider, ProviderCapabilities};
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
||||
|
||||
pub struct OllamaProvider(pub GenericOpenaiProvider);
|
||||
|
||||
impl OllamaProvider {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"http://localhost:11434/v1".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
Self(GenericOpenaiProvider::new_with_name(
|
||||
url,
|
||||
api_key,
|
||||
model,
|
||||
"ollama",
|
||||
))
|
||||
}
|
||||
|
||||
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
|
||||
/// 与 `OpenaiChatProvider::with_client` 和 `DeepSeekProvider::with_client` 一致。
|
||||
pub fn with_client(self, client: Client) -> Self {
|
||||
Self(self.0.with_client(client))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for OllamaProvider {
|
||||
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
self.0.chat(request).await
|
||||
}
|
||||
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
self.0.chat_stream(request).await
|
||||
}
|
||||
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
let mut caps = self.0.capabilities();
|
||||
caps.provider_name = "ollama";
|
||||
caps
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**`src/llm/provider.rs`** 的修改:
|
||||
|
||||
```rust
|
||||
// ProviderType 新增变体
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ProviderType {
|
||||
OpenaiChat,
|
||||
OpenaiResponse,
|
||||
Anthropic,
|
||||
DeepSeek,
|
||||
Qwen,
|
||||
/// Ollama(本地),默认 base_url = `http://localhost:11434/v1`。
|
||||
Ollama,
|
||||
}
|
||||
|
||||
// FromStr 加解析
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
match s.to_lowercase().as_str() {
|
||||
// ... 已有条目 ...
|
||||
"ollama" => Ok(ProviderType::Ollama),
|
||||
_ => Err(format!("未知的 Provider 类型: {s}")),
|
||||
}
|
||||
}
|
||||
|
||||
// create_provider() 加分支
|
||||
// Step 5.2 阶段仅展示基本构造。Step 5.1(ProviderConfig 扩展)
|
||||
// 执行到此分支时,将同步补充 with_client 链式调用注入 timeout:
|
||||
//
|
||||
// let client = Client::builder()
|
||||
// .timeout(Duration::from_secs(config.timeout_secs))
|
||||
// .build()?;
|
||||
// Ok(Box::new(
|
||||
// ollama::OllamaProvider::new(config.base_url, config.api_key, config.model)
|
||||
// .with_client(client),
|
||||
// ))
|
||||
ProviderType::Ollama => Ok(Box::new(ollama::OllamaProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
```
|
||||
|
||||
#### 集成方式
|
||||
|
||||
OllamaProvider 的 newtype 包装模式与 `DeepSeekProvider`、`QwenProvider` 完全一致,`LlmProvider` trait 委托给 `self.0`。`capabilities().provider_name` 返回 `"ollama"`。
|
||||
|
||||
### 3.2 Step 5.3 — `#[non_exhaustive]` 前置标记
|
||||
|
||||
#### 改动文件清单
|
||||
|
||||
| 文件 | 行号 | 枚举 | 操作 |
|
||||
|------|------|------|------|
|
||||
| `src/llm/provider.rs` | ~21 | `ProviderType` | 加 `#[non_exhaustive]` |
|
||||
| `src/llm/types/response_v2.rs` | ~22 | `StopReason` | 加 `#[non_exhaustive]` |
|
||||
| `src/llm/types/shared.rs` | ~16 | `FinishReason` | 加 `#[non_exhaustive]` |
|
||||
| `src/memory/store.rs` | ~35 | `EvictionPolicy` | 加 `#[non_exhaustive]` |
|
||||
|
||||
**排除清单**:`SlotMode`。
|
||||
|
||||
**决策理由**:`SlotMode` 枚举在 Phase 10(`src/llm/context.rs`)中才实际定义,Phase 5 尚不存在此类型。`#[non_exhaustive]` 无法标注不存在的枚举,因此排除标注。Roadmap(v0.2.0 §Phase 5 Step 5.3)列出的 `SlotMode`(预置) 推迟到 Phase 10 实现时一并添加。
|
||||
|
||||
#### 关键代码
|
||||
|
||||
每个枚举在 `derive` 上方或下方加一行属性:
|
||||
|
||||
```rust
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[non_exhaustive]
|
||||
pub enum ProviderType {
|
||||
// ...
|
||||
}
|
||||
```
|
||||
|
||||
#### 影响分析
|
||||
|
||||
- `#[non_exhaustive]` 是纯编译期属性,不影响运行时行为
|
||||
- 同一 crate 内的 exhaustive match 不受影响(同 crate 可穷举)
|
||||
- 下游 crate 的 match 必须加 `_ =>` 兜底分支,这是期望行为——确保未来新增变体时不会 silent break
|
||||
- **单向门**:此步骤一旦通过 `v0.2.0` 发布到公共 API 后,**不可回退**。回退意味着移除 `#[non_exhaustive]`,可能破坏已添加 `_ =>` 的下游代码。因此必须在发布前完成并确认所有枚举变体正确
|
||||
|
||||
### 3.3 Step 5.1 — ProviderConfig 扩展(最后执行)
|
||||
|
||||
#### 改动文件清单
|
||||
|
||||
| 文件 | 操作 | 说明 |
|
||||
|------|------|------|
|
||||
| `src/llm/provider.rs` | 修改 | `ProviderConfig` 加字段;加 `impl Default`;加 `from_env()`;`create_provider` 注入 timeout |
|
||||
| `src/llm/provider/openai.rs` | 修改 | `GenericOpenaiProvider` 新增 `timeout_secs` 字段;`new_with_name` 接受 timeout 参数;`map_reqwest_error` 参数化 |
|
||||
| `src/llm/provider/anthropic.rs` | 修改 | 新增 `timeout_secs` 字段;`new()` 接受 timeout 参数;`map_reqwest_error` 参数化 |
|
||||
| `src/llm/provider/anthropic.rs` | 修改 | 新增 `with_timeout()` 方法(返回 `Result<Self, LlmError>`) |
|
||||
| `src/llm/provider/openai_compat.rs` | 修改 | `DeepSeekProvider` 和 `QwenProvider` 新增公开 `with_client()` 方法 |
|
||||
| `src/llm/provider/ollama.rs` | 修改 | `OllamaProvider` 新增公开 `with_client()` 方法 |
|
||||
| `Cargo.toml` | 修改 | 加 `temp_env` dev-dependency |
|
||||
| 测试文件(`provider.rs` 内联或独立) | 新增 | `from_env` 单元测试 + timeout 传导集成测试 |
|
||||
|
||||
#### 数据结构
|
||||
|
||||
```rust
|
||||
/// Provider 构造参数 —— 通用 base_url + api_key + model + timeout/retry 配置。
|
||||
pub struct ProviderConfig {
|
||||
pub base_url: String,
|
||||
pub api_key: String,
|
||||
pub model: String,
|
||||
/// 请求超时秒数(默认 30)。应用于 Provider 的 HTTP Client 级别。
|
||||
pub timeout_secs: u64,
|
||||
/// 最大重试次数(默认 3)。当前此字段仅由 `from_env()` 采集,
|
||||
/// 实际重试逻辑由 `CycleConfig.retry.max_retries` 控制。
|
||||
/// 未来可合并到统一的 retry 配置。
|
||||
pub max_retries: u32,
|
||||
}
|
||||
|
||||
impl Default for ProviderConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
base_url: String::new(),
|
||||
api_key: String::new(),
|
||||
model: String::new(),
|
||||
timeout_secs: 30,
|
||||
max_retries: 3,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ProviderConfig {
|
||||
/// 从环境变量构造 ProviderConfig。
|
||||
///
|
||||
/// 必填变量:
|
||||
/// - `{prefix}_BASE_URL`
|
||||
/// - `{prefix}_API_KEY`
|
||||
/// - `{prefix}_MODEL`
|
||||
///
|
||||
/// 可选变量(有默认值):
|
||||
/// - `{prefix}_TIMEOUT_SECS`(默认 30)
|
||||
/// - `{prefix}_MAX_RETRIES`(默认 3)
|
||||
pub fn from_env(prefix: &str) -> Result<Self, String> {
|
||||
let base_url = std::env::var(format!("{prefix}_BASE_URL"))
|
||||
.map_err(|_| format!("{prefix}_BASE_URL 环境变量未设置"))?;
|
||||
let api_key = std::env::var(format!("{prefix}_API_KEY"))
|
||||
.map_err(|_| format!("{prefix}_API_KEY 环境变量未设置"))?;
|
||||
let model = std::env::var(format!("{prefix}_MODEL"))
|
||||
.map_err(|_| format!("{prefix}_MODEL 环境变量未设置"))?;
|
||||
let timeout_secs = match std::env::var(format!("{prefix}_TIMEOUT_SECS")) {
|
||||
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||
tracing::warn!("{prefix}_TIMEOUT_SECS='{v}' 解析失败,使用默认值 30");
|
||||
30
|
||||
}),
|
||||
Err(_) => 30,
|
||||
};
|
||||
let max_retries = match std::env::var(format!("{prefix}_MAX_RETRIES")) {
|
||||
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||
tracing::warn!("{prefix}_MAX_RETRIES='{v}' 解析失败,使用默认值 3");
|
||||
3
|
||||
}),
|
||||
Err(_) => 3,
|
||||
};
|
||||
|
||||
// ponytail: max_retries 当前仅采集,不传入 Provider。
|
||||
// 实际重试由 CycleConfig.retry.max_retries 控制。
|
||||
// 此 warn 在应用启动时通常只触发一次,多次调用 from_env 时
|
||||
// 重复输出的风险低。如有噪声,可改用 std::sync::Once 控制。
|
||||
if max_retries != 3 {
|
||||
tracing::warn!(
|
||||
"ProviderConfig.max_retries={} 已采集但当前未生效;\
|
||||
重试次数由 CycleConfig.retry.max_retries 控制",
|
||||
max_retries,
|
||||
);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
base_url,
|
||||
api_key,
|
||||
model,
|
||||
timeout_secs,
|
||||
max_retries,
|
||||
})
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
#### Timeout 传导模式
|
||||
|
||||
在 `create_provider()` 中,对基于 `GenericOpenaiProvider` 的 Provider(OpenAI / DeepSeek / Qwen / Ollama),通过同一模式注入 timeout:构造带 timeout 的 `Client` 后调用 `with_client(client)`。
|
||||
|
||||
所有 OpenAI-compatible 分支新增的 `with_client()` 公开方法:
|
||||
|
||||
| Provider | 方法 | 位置 |
|
||||
|----------|------|------|
|
||||
| `OpenaiChatProvider` | 已有 `with_client(Client) -> Self` | `openai.rs` |
|
||||
| `DeepSeekProvider` | 新增 `with_client(Client) -> Self` | `openai_compat.rs` |
|
||||
| `QwenProvider` | 新增 `with_client(Client) -> Self` | `openai_compat.rs` |
|
||||
| `OllamaProvider` | 新增 `with_client(Client) -> Self` | `ollama.rs`(新建文件) |
|
||||
|
||||
**关于 `new_with_client` 的说明**:`DeepSeekProvider` 和 `QwenProvider` 当前已有测试用的 `new_with_client(base_url, api_key, model, client)` 方法(通过 `inner.http_client = client` 直接写字段)。新增 `with_client` 后,`new_with_client` 应重构为 `Self::new(base_url, api_key, model).with_client(client)` 代理,统一走公开 API 路径。
|
||||
|
||||
代码示例(以 DeepSeek 为例,OpenAI/Qwen/Ollama 模式完全一致):
|
||||
|
||||
```rust
|
||||
ProviderType::DeepSeek => {
|
||||
let client = Client::builder()
|
||||
.timeout(Duration::from_secs(config.timeout_secs))
|
||||
.build()
|
||||
.map_err(|e| LlmError::Other(format!("创建 HTTP 客户端失败: {e}")))?;
|
||||
Ok(Box::new(
|
||||
openai_compat::DeepSeekProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
)
|
||||
.with_client(client),
|
||||
))
|
||||
}
|
||||
```
|
||||
|
||||
Anthropic 由于需要保留 `default_headers`,使用独立的 `with_timeout` 模式:
|
||||
|
||||
AnthropicProvider 新增 `with_timeout` 方法:
|
||||
|
||||
```rust
|
||||
impl AnthropicProvider {
|
||||
/// 替换默认 HTTP Client 的超时配置。
|
||||
///
|
||||
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后原有通过 `with_client`
|
||||
/// 注入的 Client 将被替换。headers 逻辑与 `new()` 中的构造保持一致。
|
||||
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
|
||||
// ponytail: 重建 http_client 时保留已有默认 headers(x-api-key / anthropic-version)。
|
||||
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
|
||||
let key_header = HeaderValue::from_str(&self.api_key)
|
||||
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
|
||||
let version_header = HeaderValue::from_static("2023-06-01");
|
||||
|
||||
self.http_client = Client::builder()
|
||||
.timeout(Duration::from_secs(secs))
|
||||
.default_headers({
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-api-key", key_header);
|
||||
headers.insert("anthropic-version", version_header);
|
||||
headers
|
||||
})
|
||||
.build()
|
||||
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
|
||||
Ok(self)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
#### `map_reqwest_error` 中的硬编码超时修复
|
||||
|
||||
`openai.rs` 和 `anthropic.rs` 中的 `map_reqwest_error` 辅助函数当前在超时错误中返回硬编码的 `Duration::from_secs(120)`:
|
||||
|
||||
```rust
|
||||
// 现状 —— 硬编码 120s,与可配置 timeout 脱节
|
||||
LlmError::Timeout { duration: Duration::from_secs(120) }
|
||||
```
|
||||
|
||||
**修复方式**:采用**方案 A**——在 Provider struct 中存储 `timeout_secs` 字段,`map_reqwest_error` 读取该字段的值而非硬编码 120s。
|
||||
|
||||
```rust
|
||||
// 修复后 —— 参数化,从 Provider 存储的 timeout_secs 读取
|
||||
// GenericOpenaiProvider 新增 timeout_secs 字段:
|
||||
pub struct GenericOpenaiProvider {
|
||||
http_client: Client,
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
provider_name: &'static str,
|
||||
extra_headers: Vec<(String, String)>,
|
||||
timeout_secs: u64, // ← 新增,由 new_with_name 的参数传入
|
||||
}
|
||||
|
||||
// map_reqwest_error 使用 self.timeout_secs 而非硬编码 120:
|
||||
LlmError::Timeout { duration: Duration::from_secs(self.timeout_secs) }
|
||||
```
|
||||
|
||||
**方案 B(从 reqwest::Client 提取 timeout)已被否决**:`reqwest::Client` 不提供 timeout getter,无法从已构造的 client 中反向读取超时配置。
|
||||
|
||||
如果漏掉此修复,用户设置 `AG_LLM_TIMEOUT_SECS=60` 后超时,错误消息仍显示 "LLM 请求超时(120s)",与实际配置不符。
|
||||
|
||||
---
|
||||
|
||||
#### max_retries 说明
|
||||
|
||||
`ProviderConfig.max_retries` 当前仅由 `from_env()` 采集存储,**实际重试操作由 `CycleConfig.retry.max_retries` 控制**。两者之间的关系通过文档注释声明:
|
||||
|
||||
```rust
|
||||
/// 最大重试次数(默认 3)。当前此字段仅由 `from_env()` 采集,
|
||||
/// 实际重试逻辑由 `CycleConfig.retry.max_retries` 控制。
|
||||
/// 未来 Phase 6+ 可统一合并此字段到 CycleConfig。
|
||||
```
|
||||
|
||||
#### 测试设计
|
||||
|
||||
使用 `temp_env` 在单元测试中隔离环境变量:
|
||||
|
||||
```rust
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn provider_config_from_env_requires_all_vars() {
|
||||
// 未设置任何变量时应返回 Err
|
||||
let result = ProviderConfig::from_env("TEST_PROVIDER");
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_from_env_uses_defaults() {
|
||||
temp_env::with_vars([
|
||||
("TEST_PROVIDER_BASE_URL", Some("http://localhost:11434/v1")),
|
||||
("TEST_PROVIDER_API_KEY", Some("")),
|
||||
("TEST_PROVIDER_MODEL", Some("llama3")),
|
||||
], || {
|
||||
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||
assert_eq!(config.timeout_secs, 30);
|
||||
assert_eq!(config.max_retries, 3);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_from_env_reads_custom_timeout() {
|
||||
temp_env::with_vars([
|
||||
("TEST_PROVIDER_BASE_URL", Some("http://x")),
|
||||
("TEST_PROVIDER_API_KEY", Some("k")),
|
||||
("TEST_PROVIDER_MODEL", Some("m")),
|
||||
("TEST_PROVIDER_TIMEOUT_SECS", Some("60")),
|
||||
("TEST_PROVIDER_MAX_RETRIES", Some("5")),
|
||||
], || {
|
||||
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||
assert_eq!(config.timeout_secs, 60);
|
||||
assert_eq!(config.max_retries, 5);
|
||||
});
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## 4. 实现计划
|
||||
|
||||
### Step 5.2 — Ollama Provider(~55 行)
|
||||
|
||||
| 步骤 | 操作 | 验证 |
|
||||
|------|------|------|
|
||||
| 1 | 创建 `src/llm/provider/ollama.rs`,实现 `OllamaProvider` newtype | 编译通过 |
|
||||
| 2 | 在 `provider.rs` 注册 `pub mod ollama` | 编译通过 |
|
||||
| 3 | `ProviderType` 新增 `Ollama` 变体 | 编译通过 |
|
||||
| 4 | `FromStr` 加 `"ollama"` 解析 | 编译通过 |
|
||||
| 5 | `create_provider()` 加 `Ollama =>` 分支 | 编译通过 |
|
||||
| 6 | 运行 `cargo build` | 无错误 |
|
||||
|
||||
### Step 5.3 — `#[non_exhaustive]` 前置标记(~4 行)
|
||||
|
||||
| 步骤 | 操作 | 验证 |
|
||||
|------|------|------|
|
||||
| 1 | `ProviderType`(`provider.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||
| 2 | `StopReason`(`response_v2.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||
| 3 | `FinishReason`(`shared.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||
| 4 | `EvictionPolicy`(`memory/store.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||
| 5 | 运行 `cargo build --all-targets` | 无 warning |
|
||||
|
||||
### Step 5.1 — ProviderConfig 扩展(~60 行 + 测试)
|
||||
|
||||
| 步骤 | 操作 | 验证 |
|
||||
|------|------|------|
|
||||
| 1 | `ProviderConfig` 加 `timeout_secs` / `max_retries` 字段 | 编译通过 |
|
||||
| 2 | 实现 `impl Default for ProviderConfig` | 编译通过 |
|
||||
| 3 | 实现 `ProviderConfig::from_env()` | 编译通过 |
|
||||
| 4 | `GenericOpenaiProvider` 和 `AnthropicProvider` 新增 `timeout_secs` 字段,`new_with_name`/`new()` 接受 timeout 参数 | 编译通过 |
|
||||
| 5 | `map_reqwest_error` 在各 Provider 中改为从 `self.timeout_secs` 读取,移除硬编码 120s | 编译通过 |
|
||||
| 6 | `create_provider()` 中各分支注入 timeout(OpenAI-compatible 用 `Client::builder().timeout()` + `with_client`;Anthropic 用 `with_timeout()`) | 编译通过 |
|
||||
| 7 | `DeepSeekProvider`/`QwenProvider` 的 `new_with_client` 重构为 `Self::new(...).with_client(client)` 代理 | 测试通过 |
|
||||
| 8 | `Cargo.toml` 添加 `temp_env` dev-dependency | `cargo build` 通过 |
|
||||
| 8 | 添加 `from_env` 单元测试 + timeout 传导集成测试 | `cargo test` 通过 |
|
||||
| 9 | 完整验证 | 见第 6 节 |
|
||||
|
||||
## 5. 风险评估
|
||||
|
||||
| 风险 | 影响 | 概率 | 缓解措施 |
|
||||
|------|------|------|----------|
|
||||
| `create_provider()` 中 `Client::builder().build()` 返回 `Result`,当前代码使用 `.expect()`,改为 `map_err` 转为 `LlmError` 后需确保所有分支正确转换 | 编译期强制处理,遗漏分支直接报错 | 低 | `create_provider` 返回 `Result<Box<dyn LlmProvider>, LlmError>`,`map_err` 天然适配。新增的 timeout 注入路径逐一检查 |
|
||||
| `AnthropicProvider` 的 `default_headers` 在 `with_timeout` 中重建时与 `new()` 中的 headers 不一致 | Anthropic 认证失败 | 低 | `with_timeout` 方法复制 `new()` 中的 headers 构造逻辑。通过已有测试验证认证通过 |
|
||||
| Ollama 实际运行时行为差异:版本兼容性、API 路径、模型名等 | 运行时才能发现 | 中 | Phase 5 仅做类型级验证(`cargo build`),Phase 8 端到端测试时通过 Ollama mock 或真实实例验证 |
|
||||
| `max_retries` 存储了却未实际使用,造成困惑 | 开发者误以为已生效 | 中 | 通过文档注释明确声明 `max_retries` 当前仅采集,实际重试由 `CycleConfig.retry.max_retries` 控制 |
|
||||
| `temp_env` 测试在多线程并发测试中互相污染环境变量 | 偶发测试失败 | 中(Rust 默认单线程测试用 `--test-threads=1` 可避免) | 将 `from_env` 测试控制在同一测试文件,避免并行执行。必要时在 CI 中确保 `--test-threads=1` |
|
||||
|
||||
## 6. 验收标准
|
||||
|
||||
以下条件**全部满足**方可认为 Phase 5 完成:
|
||||
|
||||
- [ ] `cargo build --all-targets` 通过,无错误
|
||||
- [ ] `cargo test --all-targets` 通过,新增测试覆盖 `from_env` 的必填/选填/默认值场景
|
||||
- [ ] `cargo clippy --all-targets -- -D warnings` 通过,无任何 warning
|
||||
- [ ] `cargo doc --no-deps -D warnings` 通过,所有公共 API 有文档注释(`///`)
|
||||
- [ ] 新增文件:1(`ollama.rs`)
|
||||
- [ ] 修改文件:9(`provider.rs`、`openai.rs`、`anthropic.rs`、`openai_compat.rs`、`response_v2.rs`、`shared.rs`、`store.rs`、`Cargo.toml`、测试文件)
|
||||
- [ ] 净代码增量:~160 行
|
||||
- [ ] `ProviderType` 新增 `Ollama` 变体,`"ollama"` 字符串可解析
|
||||
- [ ] 4 个公共枚举带有 `#[non_exhaustive]` 属性
|
||||
- [ ] `ProviderConfig` 可从环境变量构造(`from_env()`),含默认值
|
||||
- [ ] timeout 值已传导到 `create_provider()` 中各 Provider 的 HTTP Client 配置
|
||||
- [ ] timeout 传导验证通过至少一个端到端 wiremock 集成测试(模拟 HTTP 服务在超时后返回 408,验证 Provider 返回 `LlmError::Timeout`)
|
||||
- [ ] `DeepSeekProvider`、`QwenProvider`、`OllamaProvider` 均有公开 `with_client()` 方法,可在 `create_provider` 中注入 timeout Client
|
||||
- [ ] `map_reqwest_error` 中不再硬编码 `Duration::from_secs(120)`,改为参数化读取
|
||||
@@ -0,0 +1,236 @@
|
||||
# Phase 6 — ToolDefinition IR 正式化实施方案
|
||||
|
||||
## 背景与目标
|
||||
|
||||
在 agcore v0.2 路线图中,Phase 6 旨在引入 `ToolDef` 新类型,替换已标记 `#[deprecated(since = "0.1.0")]` 的 `ToolDefinition`(即 `OpenaiToolDefinition` 类型别名),消除 OpenAI wire format 对核心类型系统的泄漏,建立 Provider 无关的工具定义中间表示(IR)。
|
||||
|
||||
预期成果:
|
||||
- 核心类型系统不再直接依赖 `OpenaiToolDefinition`
|
||||
- 所有 Provider 适配层从统一的 `ToolDef` IR 出发,各自转换为对应 wire format
|
||||
- 消除 `#[allow(deprecated)]` 抑制点,恢复 clippy 零警告状态
|
||||
|
||||
## 当前状态分析
|
||||
|
||||
当前代码库中工具定义相关的关键状态如下:
|
||||
|
||||
1. **`OpenaiToolDefinition` 结构体**定义于 `llm/types/tool.rs`,包含 4 个字段:
|
||||
- `name: String`
|
||||
- `description: Option<String>`
|
||||
- `parameters: Value`
|
||||
- `strict: Option<bool>`
|
||||
|
||||
当前存在多处 `#[allow(deprecated)]` 抑制点,分布在 `llm/cycle.rs`、`tools/registry.rs`、`tools/mcp.rs`、`agent/agent.rs` 等文件中。
|
||||
|
||||
2. **`ToolDefinition` 类型别名**定义于 `llm/types/mod.rs:105`,标记为 `#[deprecated]`:
|
||||
```rust
|
||||
#[deprecated(since = "0.1.0", note = "use OpenaiToolDefinition directly")]
|
||||
pub type ToolDefinition = OpenaiToolDefinition;
|
||||
```
|
||||
|
||||
3. **`MessageRequest.tools` 字段**类型为 `Vec<OpenaiToolDefinition>`(直接引用原始类型,而非别名)。
|
||||
|
||||
4. **引用该类型的 4 个源文件**:
|
||||
- `llm/cycle.rs`:4 个方法参数使用 `Vec<ToolDefinition>`
|
||||
- `tools/registry.rs`:`definitions() -> Vec<ToolDefinition>` 返回类型 + struct literal 构造(含 `strict: None`)
|
||||
- `tools/mcp.rs`:`list_tools() -> Vec<ToolDefinition>` 返回类型 + struct literal 构造(含 `strict: None`)
|
||||
- `agent/agent.rs`:`fn tool_definitions() -> Vec<ToolDefinition>` trait 默认实现
|
||||
|
||||
5. **Provider 适配层**:`openai.rs` 和 `anthropic.rs` 从 `MessageRequest.tools` 读取数据并转换为各自的 wire format。`openai_compat.rs` 和 `ollama.rs` 委托给 `GenericOpenaiProvider`,无需直接改动。
|
||||
|
||||
## 需求推演
|
||||
|
||||
### 决策 1:`strict` 字段的处理
|
||||
|
||||
当前所有构造路径均硬编码 `strict: None`(`registry.rs`、`mcp.rs`),`BaseTool` trait 无 `strict` 方法,用户 API 无法设置该值。
|
||||
|
||||
**结论:移除。** `strict` 是 OpenAI 的 Structured Outputs 专属字段,不属于 Provider 无关的 IR。未来如需支持,走 `MessageRequest.extra` 逃生舱,在各 Provider 适配层自行消费。
|
||||
|
||||
### 决策 2:`description` 保持 `Option<String>`
|
||||
|
||||
Anthropic 要求 `description` 为必填(`String`),但 MCP 等来源可能缺失该字段。保持 `Option`,由 Anthropic 适配层以 `unwrap_or_default()` 兜底。
|
||||
|
||||
### 决策 3:`parameters` 保持 `Value`
|
||||
|
||||
所有 Provider 的 wire format 均接受 JSON Schema 格式的 Value。当前不做 typed 方案,保留 `Value`。
|
||||
|
||||
### 决策 4:采用直接切断而非阶段性 deprecation
|
||||
|
||||
v0.1 已标记 `#[deprecated]`,用户已有预期。pre-1.0 阶段的 breaking change 是合理的。`OpenaiToolDefinition` 保留但降级为 `#[doc(hidden)]`。
|
||||
|
||||
## 方案设计
|
||||
|
||||
### ToolDef 结构体
|
||||
|
||||
位置:`src/llm/types/tool.rs`(与 `OpenaiToolDefinition` 同文件,不建独立文件/模块)。
|
||||
|
||||
```rust
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
pub struct ToolDef {
|
||||
pub name: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
#[serde(default)]
|
||||
pub parameters: Value,
|
||||
}
|
||||
```
|
||||
|
||||
实现双向 `From` 转换:
|
||||
|
||||
```rust
|
||||
impl From<ToolDef> for OpenaiToolDefinition {
|
||||
fn from(t: ToolDef) -> Self {
|
||||
Self {
|
||||
name: t.name,
|
||||
description: t.description,
|
||||
parameters: t.parameters,
|
||||
strict: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<OpenaiToolDefinition> for ToolDef {
|
||||
fn from(t: OpenaiToolDefinition) -> Self {
|
||||
Self {
|
||||
name: t.name,
|
||||
description: t.description,
|
||||
parameters: t.parameters,
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
serde 属性与 `OpenaiToolDefinition` 原有属性一致,保证 JSON 序列化兼容。
|
||||
|
||||
明确不做:
|
||||
- builder 模式(Rust struct literal + `..Default::default()` 已足够)
|
||||
- `#[non_exhaustive]`(IR 类型自有完整控制权,不需要)
|
||||
- 独立文件(7 行 struct 无需独立模块)
|
||||
|
||||
### 4 单元切割计划
|
||||
|
||||
每步设计为可编译的安全 checkpoint。
|
||||
|
||||
#### 单元 6.1 — 新增 ToolDef + From 实现
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 涉及文件 | `llm/types/tool.rs` |
|
||||
| 变更内容 | 新增 `ToolDef` struct(约 7 行)、2 个 `From` impl(约 12 行) |
|
||||
| 验证标准 | `cargo build` 编译通过(旧代码照常编译,零影响) |
|
||||
| 检查点 | 新类型存在但未被消费,安全 checkpoint |
|
||||
|
||||
#### 单元 6.2 — 别名切换 + 构造同步修复
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 涉及文件 | `llm/types/mod.rs`、`llm/types/request_v2.rs`、`llm/types/request.rs`、`tools/registry.rs`、`tools/mcp.rs` |
|
||||
| 变更内容 | 切换别名 `pub type ToolDefinition = ToolDef`;`MessageRequest.tools` 改为 `Vec<ToolDef>`;registry/mcp 构造去掉 `strict: None` |
|
||||
| 验证标准 | `cargo build` 编译通过 |
|
||||
| 风险提示 | `cycle.rs` 方法参数使用别名,自动生效无需修改;`agent/agent.rs` trait 默认实现使用别名,自动适配;暂不移除 `#[allow(deprecated)]` |
|
||||
|
||||
#### 单元 6.3 — Provider 适配
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 涉及文件 | `llm/provider/openai.rs` |
|
||||
| 变更内容 | `convert_request()` 中 `tool_defs.into_iter().map(|t| OpenaiTool::Function { function: t })` → 改为 `.map(|t| OpenaiTool::Function { function: t.into() })`;`t` 类型从 `OpenaiToolDefinition` 变为 `ToolDef`,需 `Into` 转换 |
|
||||
| 验证标准 | `cargo test --all-targets` 全部通过 |
|
||||
| 不修改的文件 | `anthropic.rs`(同名字段访问自动适配)、`openai_compat.rs`/`ollama.rs`(委托给 `GenericOpenaiProvider`) |
|
||||
|
||||
#### 单元 6.4 — 清理
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 涉及文件 | `llm/cycle.rs`、`tools/registry.rs`、`tools/mcp.rs`、`agent/agent.rs`、`llm/types/mod.rs`、`llm/types/tool.rs` |
|
||||
| 变更内容 | 移除所有与 `ToolDefinition` 相关的 `#[allow(deprecated)]`;`llm/types/mod.rs` 移除旧 `#[deprecated]` 别名(仅保留 `pub use tool::ToolDef`);`OpenaiToolDefinition` 降级为 `#[doc(hidden)]`。精确列表由 `cargo clippy -D warnings` 检出——clippy 会标记所有不再需要的 `#[allow]` |
|
||||
| 验证标准 | `cargo clippy --all-targets -- -D warnings` 零警告;`cargo test --all-targets` 全部通过 |
|
||||
| 无需改动 | 测试代码(无一直接引用 `ToolDefinition`) |
|
||||
|
||||
### Provider 适配策略
|
||||
|
||||
Provider 适配层改动最小化,仅在序列化入口处加一层 `From` 转换:
|
||||
|
||||
| Provider | 适配方式 | 改动 |
|
||||
|----------|---------|------|
|
||||
| OpenAI(`GenericOpenaiProvider`) | `MessageRequest.tools: Vec<ToolDef>` → lambda 内改为 `.map(\|t\| OpenaiTool::Function { function: t.into() })`,将 `ToolDef` 通过 `Into` 转为 `OpenaiToolDefinition` | `convert_request` lambda 内 +`.into()` |
|
||||
| Anthropic | `t.name` / `t.description` / `t.parameters` 字段名不变,直接访问 | 零改动 |
|
||||
| OpenAI Compat(DeepSeek、Qwen) | 委托给 `GenericOpenaiProvider` | 零改动 |
|
||||
| Ollama | 委托给 `GenericOpenaiProvider` | 零改动 |
|
||||
|
||||
### 变更清单汇总
|
||||
|
||||
| 文件 | 改动类型 | 估计行数 |
|
||||
|------|---------|---------|
|
||||
| `llm/types/tool.rs` | +`ToolDef` + 2x `From` | +19 |
|
||||
| `llm/types/mod.rs` | 改别名 + re-export | ~3 |
|
||||
| `llm/types/request_v2.rs` | 改 `tools` 字段 + import | ~2 |
|
||||
| `llm/cycle.rs` | 移 `#[allow(deprecated)]` | -1 |
|
||||
| `tools/registry.rs` | 构造去掉 `strict` + 移 `allow` | ~4 |
|
||||
| `tools/mcp.rs` | 构造去掉 `strict` + 移 `allow` | ~4 |
|
||||
| `agent/agent.rs` | 移 `#[allow(deprecated)]` | -1 |
|
||||
| `llm/provider/openai.rs` | `convert_request` 加 `.map(Into::into)` | +2 |
|
||||
| 新增 roundtrip 测试 | `MessageRequest` 序列化 roundtrip 验证 | +15 |
|
||||
| **合计** | | **约 47 行(+ 约 15 行测试)** |
|
||||
|
||||
### 测试策略
|
||||
|
||||
新增一条 `MessageRequest` 序列化 roundtrip 测试,覆盖 `ToolDef` 的 JSON 序列化/反序列化兼容性。该测试验证 `ToolDef` 的 serde 属性与 `OpenaiToolDefinition` 一致,确保 wire format 兼容。
|
||||
|
||||
## 实施步骤
|
||||
|
||||
1. 新建分支 `phase-6-tooldef-ir`
|
||||
2. 按单元 6.1 → 6.2 → 6.3 → 6.4 顺序执行,每步提交一个 commit
|
||||
3. 每步执行对应的验证标准
|
||||
4. 全量通过后创建 PR
|
||||
|
||||
```
|
||||
git checkout -b phase-6-tooldef-ir
|
||||
# 执行单元 6.1 → commit
|
||||
# 执行单元 6.2 → commit
|
||||
# 执行单元 6.3 → commit
|
||||
# 执行单元 6.4 → commit
|
||||
# 全量验证
|
||||
```
|
||||
|
||||
### 用户迁移指引
|
||||
|
||||
Phase 6 涉及公共 API 类型替换,下游用户升级到 v0.2 时需注意:
|
||||
|
||||
| 旧用法 | 新用法 |
|
||||
|--------|--------|
|
||||
| `use agcore::llm::types::ToolDefinition` | `use agcore::llm::types::ToolDef`(别名已移除) |
|
||||
| `use agcore::llm::types::OpenaiToolDefinition` | `use agcore::llm::types::ToolDef`(`OpenaiToolDefinition` 已降级为 `#[doc(hidden)]`) |
|
||||
| 直接构造 `ToolDefinition { strict: None, .. }` | 构造 `ToolDef { .. }`(去掉 `strict` 字段) |
|
||||
|
||||
`OpenaiToolDefinition` 仍保留但标记 `#[doc(hidden)]`,极端情况仍需使用时可通过全路径访问。
|
||||
|
||||
## 验证标准
|
||||
|
||||
| 阶段 | 验证命令 |
|
||||
|------|---------|
|
||||
| 单元 6.1 | `cargo build` 编译通过 |
|
||||
| 单元 6.2 | `cargo build` 编译通过(新旧代码全量编译) |
|
||||
| 单元 6.3 | `cargo test --all-targets` 全部通过 |
|
||||
| 单元 6.4 | `cargo clippy --all-targets -- -D warnings` 零警告;`cargo test --all-targets` 全部通过 |
|
||||
| 最终 | `cargo build --all-targets` + `cargo test --all-targets` + `cargo clippy --all-targets -- -D warnings` 全绿 |
|
||||
|
||||
## 风险与缓解
|
||||
|
||||
| 风险 | 说明 | 缓解措施 |
|
||||
|------|------|---------|
|
||||
| Struct literal 断层 | Step 6.2 切别名与构造修复若不同步,registry/mcp 中使用 `OpenaiToolDefinition` struct literal 的构造代码会编译失败 | 别名切换与构造修复合并在同一单元,原子化提交 |
|
||||
| 遗漏 `#[allow(deprecated)]` | 部分抑制点因 grep 遗漏而未在 6.4 移除 | clippy `-D warnings` 可检出;6.4 前做一次全库 grep 确认无遗漏 |
|
||||
| JSON 兼容性 | `ToolDef` serde 属性与 `OpenaiToolDefinition` 不一致导致 wire format 变化 | `ToolDef` serde 属性与 `OpenaiToolDefinition` 保持一致;roundtrip 测试验证 |
|
||||
| Provider 适配遗漏 | 部分 Provider 分支未经测试覆盖 | `cargo test --all-targets` 包含 Provider 测试 |
|
||||
|
||||
## 否决记录
|
||||
|
||||
| 否决方案 | 原因 |
|
||||
|---------|------|
|
||||
| 保留 `strict` 字段 | OpenAI 专属字段,当前所有构造路径传 `None`。不属于 Provider 无关的 IR。未来支持走 `MessageRequest.extra` 逃生舱 |
|
||||
| 逐步 deprecation 过渡 | pre-1.0 阶段 breaking change 合理,v0.1 已标记 deprecation,用户已有预期 |
|
||||
| `ToolDef` 建独立文件 | 约 7 行的 struct 不需要独立文件,与 `OpenaiToolDefinition` 共享 `types/tool.rs` 即可 |
|
||||
| `ToolDef` 放在 `tools/` 模块 | 会创造 `llm` → `tools` 的逆向依赖,破坏模块分层 |
|
||||
| 添加 builder 模式 | Rust struct literal + `..Default::default()` 已足够覆盖使用场景 |
|
||||
| 添加 `#[non_exhaustive]` | IR 类型自有完整控制权,不需要对外隐藏字段 |
|
||||
| 同 Phase 净化 `parameters` 类型化 | 属于独立工作,留给 v0.3+ 阶段处理 |
|
||||
@@ -0,0 +1,526 @@
|
||||
# Phase 7 — SqliteStore 持久化实现方案
|
||||
|
||||
- **文档编号**:14
|
||||
- **标题**:Phase 7 — SqliteStore 持久化实现方案
|
||||
- **日期**:2026-07-05
|
||||
- **状态**:已定稿
|
||||
- **涉及模块**:memory/store
|
||||
- **关联文档**:roadmap.md, 6-memory-system.md
|
||||
|
||||
---
|
||||
|
||||
## 背景与目标
|
||||
|
||||
Phase 7 的核心任务是完成 MemoryStore trait 的 SQLite 后端实现,使 Agent 进程重启后记忆数据不丢失。这是 v0.2.0 从"内存玩具"走向"可用工具"的关键门槛,也是后续 Phase 8(MVP 出口)和 Phase 10(ContextSlot)的前置依赖。
|
||||
|
||||
**成功标准**:
|
||||
- SqliteStore 完整实现 MemoryStore trait(4 个方法:save/get/delete/list)
|
||||
- 进程关闭后重新打开同一数据库文件,数据完整可读
|
||||
- 与现有 InMemoryStore 通过 MemoryStore trait 可互换,消费者零改动
|
||||
- 所有现有测试保持通过,clippy 0 警告
|
||||
|
||||
### Scope & Non-goals
|
||||
|
||||
| 范围 | 内容 |
|
||||
|------|------|
|
||||
| 包含 | 单表 CRUD + prefix/since/offset/limit 查询 + WAL 并发 + Mutex 串行化 + 错误映射 |
|
||||
| 不包含(Phase 7) | 淘汰策略(EvictionPolicy,仅 InMemoryStore 持有,需 v0.3 纳入 SqliteStore) |
|
||||
| 不包含(Phase 7) | Schema 迁移框架(PRAGMA user_version 足矣,不引入 refinery/sea-query) |
|
||||
| 不包含(Phase 7) | 批量写入 / 事务 API(N+1 clear 延迟可接受,优化后置) |
|
||||
| 不包含(Phase 7) | 跨进程共享同一数据库文件(Mutex 为单进程设计) |
|
||||
|
||||
---
|
||||
|
||||
## 当前状态分析
|
||||
|
||||
### 现有实现
|
||||
- MemoryStore trait 已在 v0.1 Phase 3 就绪,定义 4 个异步方法
|
||||
- InMemoryStore 实现稳定运行,使用 `Mutex<HashMap>` 作为后端
|
||||
- 全量测试 191 个通过,clippy 0 警告
|
||||
- 项目当前无 SQLite 或其他数据库依赖
|
||||
|
||||
### 现有消费者
|
||||
通过 `crate::memory::store::MemoryStore` 路径引用的模块:
|
||||
|
||||
| 模块 | 文件 | 使用方式 |
|
||||
|------|------|----------|
|
||||
| Agent Builder | `agent/builder.rs` | RuntimeBundle 中引用 MemoryStore |
|
||||
| Session Memory | `agent/session_memory.rs` | SessionMemory 实现 |
|
||||
| Agent Runtime | `agent/runtime.rs` | 类型标注 |
|
||||
| Agent Session | `agent/session.rs` | 默认 InMemoryStore 兜底 |
|
||||
| Conversation | `memory/conversation.rs` | ConversationMemory 测试 |
|
||||
| Knowledge | `memory/knowledge.rs` | KnowledgeStore 测试 |
|
||||
| Retriever | `memory/retriever.rs` | MemoryRetriever 测试 |
|
||||
|
||||
所有消费者均通过 `MemoryStore` trait 访问,不依赖具体实现类型,因此新增 SqliteStore 不会产生编译或运行时影响。
|
||||
|
||||
### 目录结构现状
|
||||
|
||||
```
|
||||
src/memory/
|
||||
├── mod.rs
|
||||
├── store.rs ← 包含 MemoryStore trait + InMemoryStore + EvictionPolicy
|
||||
├── conversation.rs
|
||||
├── knowledge.rs
|
||||
├── retriever.rs
|
||||
└── vector.rs
|
||||
```
|
||||
|
||||
`store.rs` 目前是一个单体文件,同时承载 trait 定义和 InMemoryStore 实现。
|
||||
|
||||
---
|
||||
|
||||
## 调研发现
|
||||
|
||||
### MemoryStore trait 定义
|
||||
|
||||
```rust
|
||||
#[async_trait]
|
||||
pub trait MemoryStore: Send + Sync {
|
||||
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError>;
|
||||
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError>;
|
||||
async fn delete(&self, id: &str) -> Result<(), MemoryError>;
|
||||
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError>;
|
||||
}
|
||||
```
|
||||
|
||||
### 关键类型
|
||||
|
||||
| 类型 | 定义 |
|
||||
|------|------|
|
||||
| `MemoryItem` | `{ id: String, content: String, metadata: Value, created_at: OffsetDateTime }` |
|
||||
| `MemoryFilter` | `{ prefix: Option<String>, since: Option<OffsetDateTime>, offset: Option<usize>, limit: Option<usize> }` |
|
||||
| `MemoryError` | 变体:`NotFound` / `Storage` / `Serialization` / `InvalidInput` / `RetrievalError` |
|
||||
| `EvictionPolicy` | `None` / `Ttl { ttl_secs }` / `Capacity { max_items }` |
|
||||
| `EvictionConfig` | `{ policy, check_interval }` |
|
||||
|
||||
### 并发模型参考
|
||||
|
||||
InMemoryStore 当前使用 `Mutex<HashMap>` 实现 `Send + Sync`。SqliteStore 将遵循相同模式,使用 `Arc<Mutex<Connection>>` + `spawn_blocking` 满足异步 trait 约束。
|
||||
|
||||
### Schema 设计考虑
|
||||
|
||||
- `created_at` 使用 TEXT(ISO 8601) 存储——`.to_string()` 零转换,字典序与时间序一致(前提:所有时间戳归一化到 UTC;`OffsetDateTime::to_string()` 在 UTC 下输出 `"2026-07-05T12:00:00Z"` 格式,字典序与时间序严格对应)
|
||||
- 初始 schema 即创建 `created_at` 索引,避免后续大数据量全表排序
|
||||
- Schema 版本管理通过 `PRAGMA user_version` 实现,零外部依赖,后续加字段只需追加 `if version < N { ALTER TABLE }`
|
||||
|
||||
---
|
||||
|
||||
## 可选方案
|
||||
|
||||
### A. rusqlite + Mutex\<Connection\>(推荐)
|
||||
|
||||
| 维度 | 评估 |
|
||||
|------|------|
|
||||
| 新增依赖 | 1 个(rusqlite 0.32 + bundled features) |
|
||||
| 实现量 | ~200 行 |
|
||||
| SQL 支持 | 原生支持 prefix LIKE 过滤 + ORDER BY 排序 |
|
||||
| 性能 | 有索引时查询 O(log n),写入串行化 |
|
||||
| 并发 | WAL 模式 + Mutex 串行化写入,适合单进程 Agent |
|
||||
| 事务支持 | 完整 ACID |
|
||||
| 崩溃安全 | WAL 模式,崩溃恢复有保障 |
|
||||
|
||||
**适用场景**:单进程 Agent 本地持久化、嵌入式场景、需要关系查询能力的通用存储。
|
||||
|
||||
**外部依赖评估**:
|
||||
- rusqlite 0.32 — 最新稳定版(2025-12 发布),维护活跃(月均 2+ 次提交),Apache-2.0 许可证
|
||||
- `bundled` feature 编译 SQLite 源码(Public Domain)进二进制,无系统级 SQLite 依赖,零外部 C 库安装步骤
|
||||
- 供应链风险:bundled 模式依赖 crate 发布节奏同步 SQLite 安全更新;SQLite 安全公告频率极低(年均 <3 例),此风险可接受
|
||||
|
||||
### B. JSONL 文件
|
||||
|
||||
| 维度 | 评估 |
|
||||
|------|------|
|
||||
| 新增依赖 | 0 |
|
||||
| 实现量 | ~150 行 |
|
||||
| get() 复杂度 | O(n) 全量扫描 |
|
||||
| delete() 复杂度 | O(n) 全量重写 |
|
||||
| 并发 | 需文件锁(flock) |
|
||||
| 事务支持 | 无 |
|
||||
| 崩溃安全 | 无保障,写入中断可能丢失或损坏数据 |
|
||||
|
||||
**否决理由**:核心的 get() 查询场景不可接受 O(n) 性能;在 Agent 运行时频繁读写记忆的场景下,全量扫描的成本会随着数据积累线性增长,不符合可用性要求。
|
||||
|
||||
### C. sled 嵌入式 KV
|
||||
|
||||
| 维度 | 评估 |
|
||||
|------|------|
|
||||
| 新增依赖 | 1 个(纯 Rust) |
|
||||
| 实现量 | ~150 行 |
|
||||
| prefix scan | 原生支持 |
|
||||
| 排序 | 需手动实现 |
|
||||
| 关系模型 | 不如 SQL 匹配当前查询模式 |
|
||||
| 社区成熟度 | 较新,API 仍在演进 |
|
||||
|
||||
**否决理由**:当前查询模式(prefix 过滤 + 按 created_at 排序)在关系模型中用一条 SQL 即可表达,引入 KV 存储反而需要手动处理排序逻辑。非必要不引入新存储范式。
|
||||
|
||||
---
|
||||
|
||||
## 推荐方案
|
||||
|
||||
### 总体方向:方案 A(rusqlite + Mutex\<Connection\>)
|
||||
|
||||
选择理由:
|
||||
|
||||
1. **最少依赖,最高匹配**:1 个新增依赖即可完整支持 MemoryFilter 的所有查询维度(prefix LIKE、created_at 范围、offset/limit)
|
||||
2. **生产就绪**:rusqlite 是 SQLite 的 Rust 绑定事实标准,bundled 模式免去系统 SQLite 依赖
|
||||
3. **Schema 演进简单**:PRAGMA user_version + 逐版本迁移,零外部迁移工具依赖
|
||||
4. **与 InMemoryStore 语义一致**:Mutex 串行化 + spawn_blocking 适配 async trait,与现有并发模型同构
|
||||
|
||||
### 关键设计决策
|
||||
|
||||
| 决策 | 选择 | 理由 |
|
||||
|------|------|------|
|
||||
| Schema 版本管理 | PRAGMA user_version | 零外部依赖,~20 行,后续 ALTER TABLE 即可 |
|
||||
| created_at 存储格式 | TEXT(ISO 8601) + UTC 归一化 | 零转换代码;UTC 下输出 `"2026-07-05T12:00:00Z"`,字典序与时间序严格一致 |
|
||||
| Upsert SQL 策略 | `INSERT ... ON CONFLICT(id) DO UPDATE SET ...` | 保留调用方传入的 `created_at`,避免被 `DEFAULT` 覆盖 |
|
||||
| 性能索引 | 初始 schema 加 created_at 索引 | 避免大数据量全表排序 |
|
||||
| 配置参数 | 仅 `path`、`busy_timeout=5s` | 其余内置默认值;5s 超时避免 `SQLITE_BUSY` 快速失败 |
|
||||
| Mutex 中毒恢复 | `.lock().unwrap_or_else(\|e\| e.into_inner())` | 不 panic,恢复执行 |
|
||||
| 批量操作 | 不加 | N+1 clear ~250ms(N=50),可接受,优化后置 |
|
||||
| spawn_blocking 取消安全性 | 短事务模式(auto-commit) | 每个操作独立事务,取消时后台 task 自然完成/panic,不 Cross 操作持有 Mutex |
|
||||
|
||||
**我们放弃了什么**(集中 Trade-off 记录):
|
||||
- **写入串行化**:`Mutex<Connection>` 确保 SQLite 写入安全,代价是同一时刻只能有一个写入者。Agent 场景下写入频率低(每次 LLM 调用触发 1-2 次),串行化不构成瓶颈
|
||||
- **单进程锁**:无法跨进程共享同一数据库文件。多进程场景需要网络后端(PostgreSQL/Redis)
|
||||
- **无横向扩展**:单文件 SQLite 无分片能力。需扩展时切换到分布式后端
|
||||
|
||||
### Schema 定义(初始版本)
|
||||
|
||||
```sql
|
||||
CREATE TABLE IF NOT EXISTS memory_items (
|
||||
id TEXT PRIMARY KEY,
|
||||
content TEXT NOT NULL,
|
||||
metadata TEXT NOT NULL DEFAULT '{}',
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_items_created_at
|
||||
ON memory_items(created_at);
|
||||
```
|
||||
|
||||
### 数据完整性防御
|
||||
|
||||
| 异常场景 | 防御措施 | 错误映射 |
|
||||
|---------|---------|---------|
|
||||
| 数据库文件损坏 | `migrate()` 中执行 `PRAGMA quick_check`;失败时 `open()` 返回 `MemoryError::Storage` | `Storage` |
|
||||
| created_at 解析失败 | `get()`/`list()` 中 `OffsetDateTime::parse` 失败不 panic,返回 `MemoryError::Serialization` | `Serialization` |
|
||||
| content / metadata 为 NULL | `get()` 中检测 SQLite 返回值,NULL 时返回 `MemoryError::Storage` | `Storage` |
|
||||
| 约束冲突(PRIMARY KEY / NOT NULL) | 映射为 `MemoryError::InvalidInput` | `InvalidInput` |
|
||||
| 序列化/反序列化失败 | `serde_json::to_string`/`from_str` 错误映射为 `MemoryError::Serialization` | `Serialization` |
|
||||
|
||||
### 性能预算(目标延迟,单条操作)
|
||||
|
||||
| 操作 | 目标延迟 | 说明 |
|
||||
|------|---------|------|
|
||||
| `save(1KB item)` | < 5ms | 含 serde_json 序列化 + spawn_blocking + SQLite INSERT |
|
||||
| `get(1KB item)` | < 3ms | 含 SQLite SELECT + 反序列化 |
|
||||
| `list(prefix 匹配 100 行)` | < 20ms | 含索引 B-tree 遍历 + ORDER BY + LIMIT |
|
||||
| 并发 10 writer | p99 < 50ms | Mutex 串行化排队,每 writer 等待 9×5ms 内 |
|
||||
|
||||
实施后通过 Step C 测试验证以上预算。未达标时不阻塞发布,但记录为可观测告警阈值。
|
||||
|
||||
---
|
||||
|
||||
## 实施建议
|
||||
|
||||
### 实施计划
|
||||
|
||||
#### Step A — 目录重构(纯搬移,零行为变化)
|
||||
|
||||
目标:将单体 `store.rs` 拆分为模块目录架构,为新增 SqliteStore 做准备。
|
||||
|
||||
```
|
||||
src/memory/
|
||||
├── store.rs ← 模块根:MemoryStore trait + EvictionPolicy/EvictionConfig
|
||||
│ + pub mod in_memory;
|
||||
│ + pub mod sqlite_store;
|
||||
│ + pub use in_memory::InMemoryStore;
|
||||
├── store/
|
||||
│ ├── in_memory.rs ← InMemoryStore 提取至此(struct + impl + 6 个内联测试)
|
||||
│ └── sqlite_store.rs ← 新增 SqliteStore
|
||||
```
|
||||
|
||||
模式参考:`llm/provider.rs` → `llm/provider/{openai,anthropic,ollama}.rs`
|
||||
|
||||
**外部消费者的导入路径不变**(`crate::memory::store::MemoryStore`),零改动风险。
|
||||
|
||||
重构步骤:
|
||||
1. 创建 `src/memory/store/` 目录
|
||||
2. 创建 `src/memory/store/in_memory.rs`,从原 `store.rs` 提取 InMemoryStore 全部代码(struct + impl + Default + 6 个测试)
|
||||
3. 修改 `src/memory/store.rs`:保留 MemoryStore trait + EvictionPolicy/EvictionConfig,加 `pub mod in_memory;` + `pub use in_memory::InMemoryStore;`
|
||||
4. 验证:`cargo test --all-targets` 全绿,测试数量不变(191 pass)
|
||||
|
||||
#### Step B — SqliteStore 实现
|
||||
|
||||
1. `Cargo.toml` 添加 `rusqlite = { version = "0.32", features = ["bundled"] }`
|
||||
2. 创建 `src/memory/store/sqlite_store.rs`,实现:
|
||||
- `SqliteStore` 结构体:`{ conn: Arc<Mutex<Connection>> }`
|
||||
- `SqliteStore::open(path)` 构造函数,支持 `":memory:"`
|
||||
- `lock_conn()` 辅助方法(Mutex 中毒恢复)
|
||||
- `migrate()` Schema 初始化 + 版本管理
|
||||
- `MemoryStore` trait 的 4 个方法
|
||||
- `From<rusqlite::Error> for MemoryError`
|
||||
3. 修改 `src/memory/store.rs`:加 `pub mod sqlite_store;` + `pub use sqlite_store::SqliteStore;`
|
||||
4. 修改 `src/memory.rs`:加 `pub use store::SqliteStore;`
|
||||
5. 编写测试覆盖:
|
||||
- CRUD 基本操作
|
||||
- Upsert(同 id 重复 save 覆盖)
|
||||
- prefix 过滤
|
||||
- 由于/until 时间范围过滤
|
||||
- 并发 10 个 writer × 10 次操作
|
||||
- 持久化恢复(write → drop → reopen → read)
|
||||
6. 验证:`cargo test --all-targets` 全绿 + `cargo clippy --all-targets -- -D warnings` 0 警告
|
||||
|
||||
#### Step C — 验证确认
|
||||
|
||||
1. 确认现有 memory 模块内测试全部通过
|
||||
2. 确认 agent/llm/tools/prompt 模块不受影响
|
||||
3. 确认 SqliteStore 与 InMemoryStore 通过 MemoryStore trait 可互换
|
||||
4. 确认 clippy 无新增警告
|
||||
|
||||
### Commit 安排
|
||||
|
||||
| 顺序 | 类型 | Scope | 描述 |
|
||||
|------|------|-------|------|
|
||||
| 1 | refactor | memory | 将 store.rs 拆分为模块目录,仅结构搬移 |
|
||||
| 2 | feat | memory | 实现 SqliteStore 持久化 |
|
||||
|
||||
### 风险与缓解
|
||||
|
||||
| 风险 | 严重度 | 缓解措施 |
|
||||
|------|--------|----------|
|
||||
| Mutex 中毒导致后续操作全部失败 | 中 | `lock_conn()` 使用 `.lock().unwrap_or_else(\|e\| e.into_inner())` 恢复模式,不 panic |
|
||||
| spawn_blocking 取消后连接状态不一致 | 中 | 每个操作使用短事务(auto-commit),不跨操作持有 Mutex;取消时遗留 task 自然完成或 panic,Mutex 通过 `.into_inner()` 恢复 |
|
||||
| WAL 文件无限增长 | 低 | 内置 auto-checkpoint 阈值 + 启动时执行 `PRAGMA wal_checkpoint(TRUNCATE)` |
|
||||
| list 无索引导致全表扫描 | 中(大数据量) | 初始 schema 即创建 `idx_memory_items_created_at` 索引 |
|
||||
| 父目录不存在导致 open 失败 | 低 | `open()` 内部调用 `fs::create_dir_all()` 确保目录存在 |
|
||||
| clear() N+1 删除性能 | 低 | 不走 trait 接口的逐条删除,可后续优化为直接 `DELETE FROM memory_items` |
|
||||
| 数据库文件损坏 | 低 | `migrate()` 中执行 `PRAGMA quick_check`;失败时返回 `MemoryError::Storage`,调用方可切换 InMemoryStore |
|
||||
|
||||
### 可观测性(实施时落实)
|
||||
|
||||
- 所有 MemoryStore 方法通过 `tracing::instrument` 记录延迟和结果(`info!` 正常完成,`warn!` 超过性能预算阈值,`error!` 操作失败)
|
||||
- `list()` 返回行数通过 `tracing::debug` 记录(调优参考)
|
||||
- WAL 文件大小在 `migrate()` 后检查一次,超过 100MB 时记录 `warn!`
|
||||
- 操作计数(读写次数、错误率)暂不暴露为独立 metrics,v0.3 按需添加
|
||||
|
||||
---
|
||||
|
||||
## 已知假设
|
||||
|
||||
| 假设 | 验证状态 | Fallback |
|
||||
|------|---------|----------|
|
||||
| 单进程独享 SQLite 文件,无跨进程竞争 | ✅ 设计前提(Mutex 为单进程设计) | 多进程场景使用网络后端(PostgreSQL/Redis,v0.3+) |
|
||||
| ISO 8601 TEXT 字典序等价于时间序 | ✅ 条件成立(需 UTC 归一化) | 若时区异常,切换 INTEGER(unix_timestamp) 存储后重建索引 |
|
||||
| SqliteStore 初始化失败可 fallback 到 InMemoryStore | ✅ 调用方自行控制 | `open()` 返回 `MemoryError`,消费者 catch 后改用 `InMemoryStore::new()` |
|
||||
| N+1 clear ~250ms(N=50) 可接受 | 🟡 未实测(基于 N×5ms 推算) | 若成为瓶颈,SqliteStore 内部加 `delete_by_prefix()` 方法(不走 trait 接口) |
|
||||
| rusqlite bundled SQLite 版本足够新 | ✅ 0.32 版内置 SQLite 3.46 | 如需特定版本,切换 `bundled` 为指定版本或使用系统 SQLite |
|
||||
| busy_timeout=5s 覆盖所有竞争场景 | 🟡 未实测(WAL 下写写冲突概率低) | 若观测到 `SQLITE_BUSY`,增大超时或在重试逻辑中处理 |
|
||||
| spawn_blocking 线程池不会被耗尽 | ✅ 默认 512 线程,Agent 场景占用 ≤10 | 若观测到阻塞任务排队,启动时 `tokio::task::spawn_blocking` 已有兜底排队机制 |
|
||||
|
||||
---
|
||||
|
||||
## 参考来源
|
||||
|
||||
- [rusqlite crate](https://crates.io/crates/rusqlite) — 官方文档
|
||||
- [SQLite PRAGMA user_version](https://www.sqlite.org/pragma.html#pragma_user_version) — Schema 版本管理机制
|
||||
- [SQLite WAL mode](https://www.sqlite.org/wal.html) — 并发读写性能优化
|
||||
- `docs/6-memory-system.md` — Phase 3 MemoryStore trait 原始设计
|
||||
- `docs/roadmap.md` — 项目里程碑规划(Phase 7/8/10 依赖关系)
|
||||
- `src/llm/provider.rs` → `src/llm/provider/` — 目录重构模式参考
|
||||
|
||||
---
|
||||
|
||||
## 实施计划
|
||||
|
||||
### 任务总览
|
||||
|
||||
3 个阶段、8 个任务单元、2 个 Commit。
|
||||
|
||||
### 阶段一:目录重构
|
||||
|
||||
#### Task A1 — 创建 store/ 目录并提取 InMemoryStore
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 创建 `src/memory/store/` 目录,新建 `store/in_memory.rs`,从 `store.rs` 完整提取 InMemoryStore 结构体、impl MemoryStore、impl Default、6 个内联测试 |
|
||||
| 涉及文件 | `src/memory/store.rs` → 分割到 `src/memory/store/in_memory.rs`(新增) |
|
||||
| 前置依赖 | 无 |
|
||||
| 预估工作量 | S(< 1h) |
|
||||
| 风险等级 | 低 — 纯搬移,编译器可验证 |
|
||||
| 验收条件 | `cargo build` 通过(此时 store.rs 尚未修改,store/in_memory.rs 应被 crate 忽略) |
|
||||
|
||||
注意:需要先在 store.rs 顶部添加 `pub mod in_memory;` 声明,否则子模块不会被编译。或者可以先创建目录和文件,等 Task A2 再统一加声明路径。
|
||||
|
||||
实际做法:先创建文件但不声明,A2 统一声明。这样 A1 和 A2 之间可以有一个无编译的中间状态。
|
||||
|
||||
#### Task A2 — 修改 store.rs 模块根
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 修改 `store.rs` 为纯模块根:保留 `MemoryStore` trait、`EvictionPolicy`、`EvictionConfig`;添加 `pub mod in_memory;` + `pub use in_memory::InMemoryStore;`;删除已提取到 in_memory.rs 中的代码 |
|
||||
| 涉及文件 | `src/memory/store.rs`(修改) |
|
||||
| 前置依赖 | Task A1(文件已存在) |
|
||||
| 预估工作量 | S(< 1h) |
|
||||
| 风险等级 | 低 — 保留部分不变,提取部分在子模块中 |
|
||||
| 验收条件 | `cargo test --all-targets` 全绿,测试数量不变(191 pass),clippy 0 warning |
|
||||
|
||||
#### Task A3 — 验证阶段一
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 运行全量测试链确认目录重构零行为变化 |
|
||||
| 涉及文件 | 全量 |
|
||||
| 前置依赖 | Task A2 |
|
||||
| 预估工作量 | XS(验证) |
|
||||
| 风险等级 | 低 |
|
||||
| 验收条件 | `cargo test --all-targets` 191 pass、`cargo clippy --all-targets -- -D warnings` 0 警告、`cargo build` 通过 |
|
||||
|
||||
### 阶段二:SqliteStore 实现
|
||||
|
||||
#### Task B1 — 添加 rusqlite 及 dev-dependencies
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 在 `Cargo.toml` 中添加依赖:`[dependencies]` 加 `rusqlite = { version = "0.32", features = ["bundled"] }`,`[dev-dependencies]` 加 `tempfile = "3"`(用于测试隔离);运行 `cargo build` 确认编译通过,`cargo test --no-run` 验证 dev-dependencies 可用 |
|
||||
| 涉及文件 | `Cargo.toml`(修改)、`Cargo.lock`(自动更新) |
|
||||
| 前置依赖 | 无(可与阶段一并行) |
|
||||
| 预估工作量 | XS(< 15min) |
|
||||
| 风险等级 | 低 — 标准依赖添加 |
|
||||
| 验收条件 | `cargo build` 成功,`cargo test --no-run` 成功,Cargo.lock 中生成 rusqlite 和 tempfile 条目 |
|
||||
|
||||
#### Task B2 — 实现 SqliteStore 核心
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 创建 `src/memory/store/sqlite_store.rs`,实现以下 8 个子模块: |
|
||||
| | (1)`SqliteStore` 结构体 `{ conn: Arc<Mutex<Connection>> }` |
|
||||
| | (2)`SqliteStore::open(path)` — 支持 `":memory:"`,内部调用 `fs::create_dir_all` 确保父目录存在 |
|
||||
| | (3)`lock_conn()` — 内部辅助方法,`.lock().unwrap_or_else(\|e\| e.into_inner())` 处理 Mutex 中毒 |
|
||||
| | (4)`migrate()` — 按版本递增执行迁移:`PRAGMA user_version` 检查(初始版本号=1)→ 建表 `memory_items` + 索引 `idx_memory_items_created_at` + 设置 WAL 模式 + `busy_timeout=5s` + `PRAGMA synchronous = NORMAL` + `PRAGMA wal_autocheckpoint=1000` + `PRAGMA quick_check`(检测数据库损坏)+ `PRAGMA wal_checkpoint(TRUNCATE)` |
|
||||
| | (5)`MemoryStore` trait 的 4 个方法实现(save/get/delete/list),全部使用 spawn_blocking 包裹;**生命周期注意**:`spawn_blocking` 闭包前先 `.conn.clone()` 取 `Arc<Connection>`,参数调 `.clone()` 取 owned 值,再传入 `spawn_blocking(move \|{ ... })` |
|
||||
| | — `save`: `INSERT INTO memory_items (id, content, metadata, created_at) VALUES (?1, ?2, ?3, ?4) ON CONFLICT(id) DO UPDATE SET content=excluded.content, metadata=excluded.metadata, created_at=excluded.created_at`(全字段覆盖 upsert,与 InMemoryStore 行为一致) |
|
||||
| | — `get`: `SELECT content, metadata, created_at FROM memory_items WHERE id = ?1` → Ok(None) 当无结果 |
|
||||
| | — `delete`: `DELETE FROM memory_items WHERE id = ?1`(幂等,不返回 NotFound) |
|
||||
| | — `list`: 根据 filter 字段(prefix/since/offset/limit)组合动态构造 WHERE 子句 + `ORDER BY created_at ASC` + `LIMIT ? OFFSET ?`,全部使用参数化查询防注入 |
|
||||
| | (6)序列化转换层:`time::OffsetDateTime` 存为 TEXT(ISO 8601),通过 `.to_string()` 绑定 `String` 参数;`serde_json::Value` 存为 TEXT,通过 `serde_json::to_string()` 绑定 `String` 参数;读取时通过 `OffsetDateTime::parse` 和 `serde_json::from_str` 反序列化。在 SqliteStore 内部实现 `to_sql_params()` / `from_sql_row()` 辅助方法集中处理 |
|
||||
| | (7)错误映射:不在 blanket impl From 中处理所有 rusqlite Error,而是在每个方法内部按数据完整性防御表的映射规则逐类处理: |
|
||||
| | — 数据库文件损坏 / IO 错误 → `MemoryError::Storage` |
|
||||
| | — created_at 解析失败 → `MemoryError::Serialization`(不 panic) |
|
||||
| | — content/metadata 为 NULL → `MemoryError::Storage` |
|
||||
| | — serde_json 序列化/反序列化失败 → `MemoryError::Serialization` |
|
||||
| | — 约束冲突 → `MemoryError::InvalidInput` |
|
||||
| | (8)可观测性:在 4 个 trait 方法和 `open()` 上添加 `#[tracing::instrument(skip(self))]`;正常完成记录 `trace!`,超过性能预算阈值记录 `warn!`,操作失败记录 `error!` |
|
||||
| 涉及文件 | `src/memory/store/sqlite_store.rs`(新增) |
|
||||
| 前置依赖 | Task B1(rusqlite + tempfile 依赖)、Task A1(store/ 目录存在) |
|
||||
| 预估工作量 | M(1-4h) |
|
||||
| 风险等级 | 高 — 3 个技术点需留意:`time::OffsetDateTime` 无 `rusqlite::ToSql` 实现,需显式处理 String 绑定;`spawn_blocking` + `&self` 生命周期需 clone 后才能传闭包;错误映射需精细匹配 `rusqlite::Error` 嵌套变体(`SqliteFailure` 内含 `ErrorCode`) |
|
||||
| 验收条件 | 单元测试通过(见 Task B4)、`cargo build` 通过 |
|
||||
|
||||
#### Task B3 — 注册模块并重导出
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 在 `store.rs` 添加 `pub mod sqlite_store;` + `pub use sqlite_store::SqliteStore;`;在 `memory.rs` 添加 `pub use store::SqliteStore;` |
|
||||
| 涉及文件 | `src/memory/store.rs`(修改)、`src/memory.rs`(修改) |
|
||||
| 前置依赖 | Task B2(sqlite_store.rs 文件存在) |
|
||||
| 预估工作量 | XS(< 15min) |
|
||||
| 风险等级 | 低 |
|
||||
| 验收条件 | `cargo build` 通过,`SqliteStore` 可从 `agcore::memory::SqliteStore` 路径访问 |
|
||||
|
||||
#### Task B4 — 编写测试
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 在 `sqlite_store.rs` 中编写 `#[cfg(test)] mod tests`,覆盖: |
|
||||
| | (1)CRUD 基本操作(save → get → list → delete → get None) |
|
||||
| | (2)Upsert 语义(同 id 重复 save 覆盖内容,created_at 保持调用方传入值) |
|
||||
| | (3)prefix 过滤(MemoryFilter.prefix) |
|
||||
| | (4)时间范围过滤(MemoryFilter.since) |
|
||||
| | (5)offset/limit 分页 |
|
||||
| | (6)并发 10 个 writer × 10 次写入,验证无数据丢失 |
|
||||
| | (7)持久化恢复(write → drop store → reopen 同一文件 → read) |
|
||||
| | (8)错误路径:`open("/nonexistent_dir/ag.db")` 返回 Storage 错误 |
|
||||
| | 辅助函数:`make_item(id)` 创建 MemoryItem,使用 `tempfile::TempDir`(或自定义 tmp 路径)隔离测试数据库文件 |
|
||||
| 涉及文件 | `src/memory/store/sqlite_store.rs`(修改追加 test mod) |
|
||||
| 前置依赖 | Task B2(实现完成) |
|
||||
| 预估工作量 | M(1-4h) |
|
||||
| 风险等级 | 中 — 并发测试的时序控制、持久化恢复测试的 TempDir 管理 |
|
||||
| 验收条件 | 全量测试通过,新增测试 ≥ 8 个 |
|
||||
|
||||
#### Task B5 — 验证阶段二
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | 运行全量测试链确认 SqliteStore 实现正确,不影响已有模块 |
|
||||
| 涉及文件 | 全量 |
|
||||
| 前置依赖 | Task B3、Task B4 |
|
||||
| 预估工作量 | S(< 1h) |
|
||||
| 风险等级 | 低 |
|
||||
| 验收条件 | `cargo test --all-targets` 全绿(191 + 新增测试)、`cargo clippy --all-targets -- -D warnings` 0 警告、`cargo build` 通过 |
|
||||
|
||||
### 阶段三:验证收尾
|
||||
|
||||
#### Task C1 — 跨模块兼容性验证
|
||||
|
||||
| 项目 | 内容 |
|
||||
|------|------|
|
||||
| 任务描述 | (1)确认 `ConversationMemory` / `KnowledgeStore` / `MemoryRetriever` / `AgentSession` / `SessionMemory` 以 `Arc<dyn MemoryStore>` 接受 SqliteStore 时编译通过且测试全绿 |
|
||||
| | (2)确认 SqliteStore 与 InMemoryStore 可互换——修改一个现有测试将后端从 InMemoryStore 换为 SqliteStore(使用 `":memory:"`),测试全绿 |
|
||||
| | (3)验证性能预算:在测试环境下测量 save(1KB)/get(1KB)/list(100行) 的单次延迟,确认 < 5ms / < 3ms / < 20ms |
|
||||
| 涉及文件 | 测试文件(memory/ 模块内各 test mod) |
|
||||
| 前置依赖 | Task B5 |
|
||||
| 预估工作量 | S(< 1h) |
|
||||
| 风险等级 | 低 |
|
||||
| 验收条件 | 全量测试通过 + 互换测试通过 + 性能预算大致满足(未达标不阻塞发布) |
|
||||
|
||||
### 依赖关系图
|
||||
|
||||
```
|
||||
阶段一(目录重构) 阶段二(SqliteStore 实现)
|
||||
┌──────────────┐ ┌──────────────┐
|
||||
│ Task A1 │ │ Task B1 │ ← 无依赖,可与 A 并行
|
||||
│ 创建目录+提取 │ │ Cargo.toml │
|
||||
└──────┬───────┘ └──────┬───────┘
|
||||
↓ ↓
|
||||
┌──────────────┐ ┌──────────────┐
|
||||
│ Task A2 │ │ Task B2 │
|
||||
│ 修改store.rs │← A1 ───→│ 核心实现 │← B1 + A1
|
||||
└──────┬───────┘ └──────┬───────┘
|
||||
↓ ↓
|
||||
┌──────────────┐ ┌──────────────┐ ┌──────────────┐
|
||||
│ Task A3 │ │ Task B3 │← B2 ──→│ Task B4 │
|
||||
│ 验证阶段一 │ │ 注册+重导出 │ │ 编写测试 │
|
||||
└──────────────┘ └──────┬───────┘ └──────┬───────┘
|
||||
↓ ↓
|
||||
┌──────────────┐────────────────┘
|
||||
│ Task B5 │← B3 + B4
|
||||
│ 验证阶段二 │
|
||||
└──────┬───────┘
|
||||
↓
|
||||
阶段三(验证收尾)
|
||||
┌──────────────┐
|
||||
│ Task C1 │
|
||||
│ 跨模块兼容性 │
|
||||
└──────────────┘
|
||||
```
|
||||
|
||||
### Commit 安排
|
||||
|
||||
| 顺序 | Commit 类型 | Scope | 描述 | 包含 Task |
|
||||
|------|------------|-------|------|-----------|
|
||||
| 1 | refactor | memory | 将 store.rs 拆分为模块目录,仅结构搬移 | A1 → A2 → A3 |
|
||||
| 2 | feat | memory | 实现 SqliteStore 持久化(含错误映射、测试、WAL 模式) | B1 → B2 → B3 → B4 → B5 → C1 |
|
||||
|
||||
注意:Task B1 与阶段一无依赖,可以在 Commit 1 合并进行或在 Commit 2 开头。建议在 Commit 2 开头,因为 Cargo.toml 变更属于功能变更而非重构。
|
||||
|
||||
### 验证全链
|
||||
|
||||
实施完毕后整体认证链路:
|
||||
|
||||
1. `cargo test --all-targets` — 全量测试通过
|
||||
2. `cargo clippy --all-targets -- -D warnings` — 0 警告
|
||||
3. `cargo build --release` — release 构建通过
|
||||
4. 确认 `cargo doc --no-deps` 无 warning(新增公开类型文档注释)
|
||||
5. 确认测试数量:191 + (8 个 new sqlite_store tests) = 199+ pass
|
||||
@@ -0,0 +1,543 @@
|
||||
# Phase 8 — MVP 集成出口实现方案
|
||||
|
||||
- **文档编号**:15
|
||||
- **标题**:Phase 8 — MVP 集成出口实现方案
|
||||
- **日期**:2026-07-05
|
||||
- **状态**:已定稿
|
||||
- **涉及模块**:全局(llm/types、agent、tools、memory、prompt、examples)
|
||||
- **关联文档**:roadmap.md(§Phase 8)、14-phase7-sqlite-store.md
|
||||
|
||||
---
|
||||
|
||||
## 1. 背景与目标
|
||||
|
||||
Phase 5-7 已交付 P0 功能闭环:ProviderConfig `from_env()`(Phase 5)、ToolDef IR 正式化(Phase 6)、SqliteStore 持久化(Phase 7)。当前 200 个测试全绿、clippy 0 警告,但缺乏一个"可被人依赖"的集成出口。
|
||||
|
||||
Phase 8 的目标是完成 API 稳定性扫尾 + Quick Start 示例 + 端到端示例,产出 **v0.2.0-rc.1** 标签。三个 Step 分别对应三类用户群体:
|
||||
|
||||
| Step | 受众 | 交付物 |
|
||||
|------|------|--------|
|
||||
| **8.1** | 存量升级者(v0.1 → v0.2) | API 稳定性扫尾 + CHANGELOG |
|
||||
| **8.2** | 新用户评估者("30 秒决定要不要用") | Quick Start 示例 |
|
||||
| **8.3** | 技术决策者("这框架能跑真实场景吗") | 端到端集成示例 |
|
||||
|
||||
---
|
||||
|
||||
## 2. 当前状态
|
||||
|
||||
| 度量 | 数值 |
|
||||
|------|------|
|
||||
| `cargo test --all-targets` | ✅ 200 passed / 0 failed |
|
||||
| `cargo clippy --all-targets -- -D warnings` | ✅ 0 警告 |
|
||||
| 已存在 `#[non_exhaustive]` 枚举 | 4 个(StopReason / FinishReason / EvictionPolicy / ProviderType) |
|
||||
| 已存在 `#[deprecated]` 项 | 3 个(ChatResponse / ToolDefinition / task_agent_demo 中旧类型使用) |
|
||||
| 已有示例 | 8 个 |
|
||||
| `StepStatus::Completed` 使用类型 | `ChatResponse`(已 `#[deprecated]`) |
|
||||
|
||||
### 2.1 关键技术债
|
||||
|
||||
```
|
||||
// agent/task.rs —— StepStatus 当前使用已废弃类型
|
||||
#[allow(deprecated)]
|
||||
pub enum StepStatus {
|
||||
Completed(ChatResponse), // ← ChatResponse 已在 0.1.0 标记 #[deprecated]
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
`task_agent_demo.rs` 中同时使用了 `ChatResponse` / `OpenaiChatMessage` / `FinishReason` 三个废弃类型,入口处有 `#![allow(deprecated)]`。
|
||||
|
||||
---
|
||||
|
||||
## 3. 实施方案
|
||||
|
||||
### 3.1 Step 8.1 — API 稳定性扫尾
|
||||
|
||||
拆为 4 个增量 commit:
|
||||
|
||||
| Commit | 内容 | 涉及文件 |
|
||||
|--------|------|---------|
|
||||
| **commit 1** | 14 个公开枚举追加 `#[non_exhaustive]` | 各枚举定义文件(详见 §3.1.1) |
|
||||
| **commit 2** | `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` + `task_agent_demo.rs` 清理全部 3 个废弃类型(`ChatResponse` / `OpenaiChatMessage` / `FinishReason`),移除 `#![allow(deprecated)]` | `src/agent/task.rs`、`examples/task_agent_demo.rs` |
|
||||
| **commit 3** | CHANGELOG v0.2 条目 + Cargo.toml version → `0.2.0-rc.1` + README 更新 | `CHANGELOG.md`、`Cargo.toml`、`README.md` |
|
||||
| **commit 4** | 验证:`cargo test + clippy + cargo doc` 零告警 | 无代码改动 |
|
||||
|
||||
#### 3.1.1 `#[non_exhaustive]` 追加清单(14 个枚举)
|
||||
|
||||
按优先级分级:
|
||||
|
||||
| 优先级 | 枚举 | 模块路径 | 理由 |
|
||||
|--------|------|---------|------|
|
||||
| **P0 核心** | `Message` | `llm/types/message.rs` | 核心 IR 类型,未来可能新增变体(MultiModal 扩展) |
|
||||
| | `ContentBlock` | `llm/types/message.rs` | 同上 |
|
||||
| | `ContentBlockType` | `llm/types/message.rs` | 同上 |
|
||||
| | `StreamEvent` | `llm/types/response_v2.rs` | 流式事件集,Provider 扩展可能新增事件 |
|
||||
| | `HookEvent` | `llm/hooks.rs` | 生命周期钩子,框架扩展需要新增事件点 |
|
||||
| **P0 Error** | `AgentError` | `agent/error.rs` | 顶层错误,下游 match 需保护 |
|
||||
| | `LlmError` | `llm/error.rs` | LLM 调用错误 |
|
||||
| | `ToolError` | `tools/error.rs` | 工具系统错误 |
|
||||
| | `MemoryError` | `memory/error.rs` | 记忆系统错误 |
|
||||
| | `PromptError` | `prompt/error.rs` | 提示词工程错误 |
|
||||
| **P1 其他** | `MemoryStrategy` | `memory/conversation.rs` | 对话策略,未来可扩展(如 Summarize) |
|
||||
| | `StepStatus` | `agent/task.rs` | 步骤状态机,可扩展(如 Cancelled) |
|
||||
| | `ToolChoice` | `llm/types/request.rs` | Provider 工具选择策略 |
|
||||
| | `ResponseFormat` | `llm/types/shared.rs` | 响应格式枚举 |
|
||||
|
||||
**明确不加的**:
|
||||
|
||||
| 类别 | 枚举 | 原因 |
|
||||
|------|------|------|
|
||||
| 内部 wire-format | `OpenaiChatMessage` / `OpenaiTool` / `OpenaiToolCall` / `ContentField` / `OpenaiContentPart` / `LegacyStreamEvent` | 内部转换层,不构成公共 API 契约 |
|
||||
| 语义稳定 | `Role` / `ServiceTier` / `Modality` / `ImageDetail` / `AudioFormat` / `StopSequence` | 语义已收敛,协议层无新增变体预期 |
|
||||
| 使用面窄 | `TemplateValue` / `Permission` / `McpTransport` / `ContentBlockBuilder` / `ExtraError` | 内部实现细节或使用频率极低,下游不直接 match |
|
||||
|
||||
> **`#[non_exhaustive]` 的不可逆性**:一旦 v0.2.0-rc.1 发布,以下游代码可能依赖 `_ =>` 通配分支。在 v0.3+ 中移除 `#[non_exhaustive]` 将构成 semver breaking change(新增变体不再触发编译警告,下游 match 可能遗漏新变体),因此当前追加的标记应视为永久 API 契约。
|
||||
|
||||
#### 3.1.2 StepStatus 迁移细节
|
||||
|
||||
```rust
|
||||
// 变更前
|
||||
#[allow(deprecated)]
|
||||
pub enum StepStatus {
|
||||
Completed(ChatResponse), // ChatResponse 已 #[deprecated]
|
||||
...
|
||||
}
|
||||
|
||||
// 变更后
|
||||
#[non_exhaustive]
|
||||
pub enum StepStatus {
|
||||
Completed(MessageResponse),
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
**字段映射差异**:`MessageResponse` 不是 `ChatResponse` 的简单改名——两者结构不同,迁移需要做字段适配:
|
||||
|
||||
| ChatResponse 字段 | 类型 | MessageResponse 字段 | 类型 | 映射方式 |
|
||||
|-------------------|------|---------------------|------|---------|
|
||||
| `message` | `OpenaiChatMessage` | `message` | `Message` | 类型替换:`OpenaiChatMessage::assistant_text(t)` → `Message::Assistant { content: vec![ContentBlock::Text { text: t.into() }] }` |
|
||||
| `usage` | `Usage` | `usage` | `Usage` | ✅ 同类型,直接迁移 |
|
||||
| `stop_reason` | `Option<FinishReason>` | `stop_reason` | `StopReason` | 类型替换:`Some(FinishReason::Stop)` → `StopReason::Stop`;无 Option 包裹 |
|
||||
| — | — | `id` | `String` | 新增必填字段,使用空字符串 `""` 占位 |
|
||||
| — | — | `model` | `String` | 新增必填字段,使用 `"mock"` 或空字符串占位 |
|
||||
| — | — | `extra` | `HashMap<String, Value>` | 新增字段,使用 `HashMap::new()` 占位 |
|
||||
|
||||
**迁移示例**(`task_agent_demo.rs` 的构造代码):
|
||||
|
||||
```rust
|
||||
// 旧代码(3 个废弃类型)
|
||||
StepStatus::Completed(ChatResponse {
|
||||
message: OpenaiChatMessage::assistant_text("天气:晴,22°C"),
|
||||
usage: Usage::from_input_output(10, 5),
|
||||
stop_reason: Some(FinishReason::Stop),
|
||||
})
|
||||
|
||||
// 新代码(纯 MessageResponse)
|
||||
StepStatus::Completed(MessageResponse {
|
||||
id: String::new(),
|
||||
model: "mock".into(),
|
||||
message: Message::assistant("天气:晴,22°C"),
|
||||
usage: Usage::from_input_output(10, 5),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: HashMap::new(),
|
||||
})
|
||||
```
|
||||
|
||||
涉及文件:
|
||||
- `src/agent/task.rs`:枚举定义 + `#[allow(deprecated)]` 移除 + `#[non_exhaustive]` 追加
|
||||
- `examples/task_agent_demo.rs`:`ChatResponse{...}` → `MessageResponse{...}` 构造替换,同时替换 `OpenaiChatMessage` / `FinishReason` 引用,移除 `#![allow(deprecated)]`
|
||||
|
||||
### 3.2 Step 8.2 — Quick Start 示例
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 文件 | `examples/quick_start.rs` |
|
||||
| 规模 | ~36 行 |
|
||||
| Provider | `MockProvider`(FIFO 单响应队列) |
|
||||
| 工具 | `EchoTool`(回传 `"收到: {input}"`,完整 JSON Schema 参数声明) |
|
||||
| 执行 | `submit_turn("你好")` → 验证输出包含 `"收到"` |
|
||||
| 验证 | `cargo run --example quick_start` exit 0 |
|
||||
|
||||
设计要点:
|
||||
- 展示四层抽象:Agent trait / BaseTool 自定义 / AgentBuilder 装配 / AgentSession 执行
|
||||
- 无外部依赖、无 API key、零配置
|
||||
|
||||
### 3.3 Step 8.3 — 端到端示例
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 文件 | `examples/end_to_end.rs` |
|
||||
| 规模 | ~160 行(**最小可行边界**:3 工具 + 3 轮 + 持久化验证,防止实施中进一步膨胀) |
|
||||
| Provider | 自动检测 `AG_LLM_*` → `from_env()`,fallback 到 `MockProvider` |
|
||||
| 工具组合 | EchoTool(回显)+ CalcTool(四则运算,本地执行)+ NoteTool(笔记,通过 MemoryStore trait 操作 SessionMemory) |
|
||||
| 持久化 | `tempfile::TempDir` + `SqliteStore`,drop 后重建连接验证数据不丢 |
|
||||
| 对话 | 3 轮:计算 → 记笔记 → 回忆 |
|
||||
| 验证 | `cargo run --example end_to_end` exit 0(无需任何外部配置) |
|
||||
|
||||
**真实 Provider 切换**:示例在文件顶部注释中说明 "设置 `AG_LLM_BASE_URL` / `AG_LLM_API_KEY` / `AG_LLM_MODEL` 环境变量即可使用真实 LLM Provider(支持 OpenAI / Ollama 等);未设置时自动降级为 MockProvider,零配置可运行。"
|
||||
|
||||
**`from_env()` 部分环境变量策略**:`from_env()` 要求完整的三件套(`{prefix}_BASE_URL` + `{prefix}_API_KEY` + `{prefix}_MODEL`)。当环境变量部分设置时,示例**整体降级到 MockProvider**——不在"半配置"状态下尝试部分初始化。日志输出形如 `"AG_LLM_* 环境变量不完整(检测到: {found_vars}),回退到 MockProvider"`。
|
||||
|
||||
架构亮点:
|
||||
|
||||
```
|
||||
┌─────────────────────────┐
|
||||
│ AgentSession │
|
||||
│ (submit_turn × 3) │
|
||||
└────┬──────┬──────┬──────┘
|
||||
│ │ │
|
||||
┌────┘ │ └──────┐
|
||||
▼ ▼ ▼
|
||||
┌──────────┐ ┌────────┐ ┌──────────┐
|
||||
│ EchoTool │ │CalcTool│ │ NoteTool │
|
||||
│ (回显) │ │(四则) │ │ (记忆) │
|
||||
└──────────┘ └────────┘ └────┬─────┘
|
||||
│
|
||||
┌──────▼──────┐
|
||||
│ SessionMemory│
|
||||
│ (MemoryStore)│
|
||||
└──────┬──────┘
|
||||
│
|
||||
┌──────▼──────┐
|
||||
│ SqliteStore │
|
||||
│ (temp dir) │
|
||||
└─────────────┘
|
||||
```
|
||||
|
||||
NoteTool 展示 `MemoryStore` trait 解耦能力:不绑定 SqliteStore,上层 `AgentSession` 通过 `SessionMemory` 操作,底层可互换。
|
||||
|
||||
---
|
||||
|
||||
## 4. 否决项记录
|
||||
|
||||
| 否决方案 | 否决原因 |
|
||||
|---------|---------|
|
||||
| `#[non_exhaustive]` 仅加 5 个核心类型 | 全面覆盖 Error enums 为零运行时成本,对下游更友好。Error 枚举是下游 match 最密集的地方,漏标会在 v0.3 引入 breakage |
|
||||
| StepStatus::Completed 留到 v0.3 再修 | rc.1 前清理 deprecated 类型污染最划算——越晚 migration cost 越高,且当前仅 1 个示例 + 1 个测试引用 |
|
||||
| Quick Start 纯文本路线(不展示自定义工具) | 含 EchoTool 展示核心差异化,仅多 5 行代码但传递了"可以自定义工具"的关键信息 |
|
||||
| 端到端仅 Echo + Calc(无 NoteTool) | NoteTool 展示 MemoryStore trait 解耦能力是架构亮点,跳过后新用户无法理解 memory 如何集成到 Agent 流程 |
|
||||
| 持久化仅注释说明不实际运行(方案 Y) | 进程内实操验证(create → drop → reopen → assert)比注释更有说服力,增加约 15 行代码 |
|
||||
|
||||
---
|
||||
|
||||
## 5. 关键假设
|
||||
|
||||
1. **MockProvider FIFO 队列满足 auto-tool-loop 消费顺序**:MockProvider 的 `pop()` 按预设顺序弹出。当 LLM 返回多个 tool call 时队列消费顺序与预设一致,无需额外同步
|
||||
2. **StepStatus 切换需做字段适配**:`ChatResponse`(3 字段) 到 `MessageResponse`(6 字段) 存在字段类型差异(`message` 类型不同、`stop_reason` 类型 + Option 有无不同、`id`/`model`/`extra` 为新增必填字段),消费者需按字段映射表提供占位值。但消费者仅 1 个(`task_agent_demo.rs`)+ 1 个内联测试,手动适配工作量极小。`StepStatus` 的 `is_terminal()` / `is_pending()` 行为不受影响
|
||||
3. **所有 10 个示例零外部配置 exit 0**:已有 8 个示例已验证,新增 2 个(quick_start + end_to_end)均使用 MockProvider fallback,无需 API key
|
||||
4. **`#[non_exhaustive]` × 14 不触发额外 clippy warning**:当前无代码对以上枚举做 exhaustive match(不含 `_`),追加 `#[non_exhaustive]` 是纯安全标记
|
||||
|
||||
---
|
||||
|
||||
## 6. 实施顺序与验证标准
|
||||
|
||||
### 6.1 提交顺序
|
||||
|
||||
```
|
||||
Step 8.1 (4 commits)
|
||||
→ commit 1: #[non_exhaustive] × 14
|
||||
→ commit 2: StepStatus 修复(Completed(ChatResponse) → Completed(MessageResponse))
|
||||
→ commit 3: CHANGELOG v0.2 + Cargo.toml version 0.2.0-rc.1 + README 更新
|
||||
→ commit 4: 验证(test / clippy / doc 零告警)
|
||||
|
||||
Step 8.2
|
||||
→ commit 5: examples/quick_start.rs(~36 行)
|
||||
|
||||
Step 8.3
|
||||
→ commit 6: examples/end_to_end.rs(~160 行)
|
||||
|
||||
最终验证
|
||||
→ cargo test --all-targets
|
||||
→ cargo clippy --all-targets -- -D warnings
|
||||
→ cargo doc --no-deps
|
||||
→ git tag v0.2.0-rc.1
|
||||
```
|
||||
|
||||
### 6.2 验收标准
|
||||
|
||||
| 指标 | 要求 |
|
||||
|------|------|
|
||||
| `cargo test --all-targets` | 全绿 |
|
||||
| `cargo clippy --all-targets -- -D warnings` | 0 警告 |
|
||||
| `cargo doc --no-deps` | 0 warning |
|
||||
| 所有 10 个示例 | `cargo run --example <name>` exit 0 |
|
||||
| Cargo.toml version | `0.2.0-rc.1` |
|
||||
| CHANGELOG | v0.2 条目完整(Added / Changed / Deprecated / Fixed / Removed 各节) |
|
||||
| README | 示例列表 + 版本号更新 |
|
||||
| git tag | `v0.2.0-rc.1` |
|
||||
|
||||
---
|
||||
|
||||
## 7. 参考来源
|
||||
|
||||
- **roadmap.md** — Phase 8 原始定义(Step 8.1/8.2/8.3)、依赖关系(Phase 5/6/7 → Phase 8)
|
||||
- **`src/agent/task.rs`** — `StepStatus` 当前实现,`Completed(ChatResponse)` 类型
|
||||
- **`src/llm/types/message.rs`** — `Message` / `ContentBlock` / `ContentBlockType` 枚举定义
|
||||
- **`src/llm/types/response_v2.rs`** — `StreamEvent` / `StopReason` 枚举定义(StopReason 已有 `#[non_exhaustive]`)
|
||||
- **`src/llm/types/shared.rs`** — `ResponseFormat` / `Role` / `FinishReason` 等枚举(FinishReason 已有 `#[non_exhaustive]`)
|
||||
- **`src/llm/types/request.rs`** — `ToolChoice` 枚举定义
|
||||
- **`src/llm/hooks.rs`** — `HookEvent` 枚举定义
|
||||
- **`src/llm/error.rs`** — `LlmError` 枚举定义
|
||||
- **`src/agent/error.rs`** — `AgentError` 枚举定义
|
||||
- **`src/tools/error.rs`** — `ToolError` 枚举定义
|
||||
- **`src/memory/error.rs`** — `MemoryError` 枚举定义
|
||||
- **`src/memory/conversation.rs`** — `MemoryStrategy` 枚举定义
|
||||
- **`src/prompt/error.rs`** — `PromptError` 枚举定义
|
||||
- **`examples/task_agent_demo.rs`** — 当前使用 `#[allow(deprecated)]` + `ChatResponse` 的示例
|
||||
|
||||
---
|
||||
|
||||
## 8. 实施计划
|
||||
|
||||
### 8.1 实施步骤
|
||||
|
||||
#### Step 8.1 — API 稳定性扫尾
|
||||
|
||||
拆为 4 个增量 commit,依次提交。
|
||||
|
||||
##### commit 1: #[non_exhaustive] × 14
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 涉及文件 | 14 个枚举定义所在文件(见下方清单) |
|
||||
| 前置依赖 | 无 |
|
||||
| 预估工作量 | S(<1h) |
|
||||
| 风险等级 | 低 |
|
||||
|
||||
在每个目标枚举定义处的 `pub enum` 之前加一行 `#[non_exhaustive]`,纯文本属性追加,无逻辑变更。
|
||||
|
||||
| 目标枚举 | 文件路径 | 行号附近 |
|
||||
|---------|---------|---------|
|
||||
| `Message` | `src/llm/types/message.rs` | `pub enum Message` (L22) |
|
||||
| `ContentBlock` | `src/llm/types/message.rs` | `pub enum ContentBlock` (L99) |
|
||||
| `ContentBlockType` | `src/llm/types/message.rs` | `pub enum ContentBlockType` (L134) |
|
||||
| `StreamEvent` | `src/llm/types/response_v2.rs` | `pub enum StreamEvent` (L167) |
|
||||
| `HookEvent` | `src/llm/hooks.rs` | `pub enum HookEvent` (L9) |
|
||||
| `AgentError` | `src/agent/error.rs` | `pub enum AgentError` (L20) |
|
||||
| `LlmError` | `src/llm/error.rs` | `pub enum LlmError` (L10) |
|
||||
| `ToolError` | `src/tools/error.rs` | `pub enum ToolError` (L6) |
|
||||
| `MemoryError` | `src/memory/error.rs` | `pub enum MemoryError` (L8) |
|
||||
| `PromptError` | `src/prompt/error.rs` | `pub enum PromptError` (L3) |
|
||||
| `MemoryStrategy` | `src/memory/conversation.rs` | `pub enum MemoryStrategy` (L14) |
|
||||
| `StepStatus` | `src/agent/task.rs` | `pub enum StepStatus` (L59) |
|
||||
| `ToolChoice` | `src/llm/types/request.rs` | `pub enum ToolChoice` (L14) |
|
||||
| `ResponseFormat` | `src/llm/types/shared.rs` | `pub enum ResponseFormat` (L70) |
|
||||
|
||||
> **注意**:`StepStatus` 在 commit 2 中会同时被修改(variant 类型替换 + 移除 `#[allow(deprecated)]`)。commit 1 仅追加 `#[non_exhaustive]` 属性,commit 2 再处理变体变更和清理。
|
||||
|
||||
**验收条件**:`cargo build --all-targets` 通过
|
||||
|
||||
##### commit 2: StepStatus 修复 + 废弃类型清理
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 涉及文件 | `src/agent/task.rs`,`examples/task_agent_demo.rs` |
|
||||
| 前置依赖 | commit 1(StepStatus 先标记 `#[non_exhaustive]`,此处改 variant 时一并保留,无实际冲突) |
|
||||
| 预估工作量 | S(<1h,约 20 行改动) |
|
||||
| 风险等级 | 低 |
|
||||
|
||||
两步操作:
|
||||
|
||||
1. **`src/agent/task.rs`**(L59-L71):
|
||||
- `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)`
|
||||
- 移除 `#[allow(deprecated)]`(第 13、59 行两处)
|
||||
|
||||
2. **`examples/task_agent_demo.rs`**:
|
||||
- 替换 3 个废弃类型:`ChatResponse` → `MessageResponse`,`OpenaiChatMessage::assistant_text(t)` → `Message::assistant(t)`,`FinishReason::Stop` → `StopReason::Stop`
|
||||
- 补充 `id: String::new()`,`model: "mock".into()`,`extra: HashMap::new()` 占位字段
|
||||
- 移除 `#![allow(deprecated)]`(第 26 行)
|
||||
- 移除 `use` 中的 `ChatResponse`、`OpenaiChatMessage`、`FinishReason`
|
||||
- 添加 `use std::collections::HashMap`,`use agcore::llm::types::{Message, MessageResponse, StopReason}`(注意:`Message::assistant_text(t)` 不存在,需使用 `Message::assistant(t)`)
|
||||
|
||||
字段映射参见 §3.1.2 的字段映射表和迁移示例。
|
||||
|
||||
**验收条件**:`cargo build --all-targets` 通过,零 deprecated warning
|
||||
|
||||
##### commit 3: CHANGELOG + 版本号 + README
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 涉及文件 | `CHANGELOG.md`,`Cargo.toml`,`README.md` |
|
||||
| 前置依赖 | commit 1+2(CHANGELOG 需记录实际变更) |
|
||||
| 预估工作量 | S(<1h) |
|
||||
| 风险等级 | 低 |
|
||||
|
||||
1. **`CHANGELOG.md`**:新增 `[0.2.0-rc.1]` 条目,包含:
|
||||
- **Added**:SqliteStore 持久化 / OllamaProvider / ProviderConfig::from_env / ToolDef IR / Quick Start 和 end_to_end 示例
|
||||
- **Changed**:MessageRequest.tools 切换 ToolDef / StepStatus::Completed 类型替换
|
||||
- **Deprecated**:ChatResponse / with_system_prompt() / with_client()
|
||||
- **Non-exhaustive**:14 个枚举标记清单
|
||||
|
||||
2. **`Cargo.toml`**:第 3 行 `version = "0.1.0"` → `version = "0.2.0-rc.1"`
|
||||
|
||||
3. **`README.md`**:更新示例列表从 7 个改为 10 个(含新增 2 个),版本号同步
|
||||
|
||||
**验收条件**:人工 review CHANGELOG + `git diff` 确认版本号
|
||||
|
||||
##### commit 4: 验证
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 涉及文件 | 无代码改动 |
|
||||
| 前置依赖 | commit 3 |
|
||||
| 预估工作量 | S(<1h,主要等待编译) |
|
||||
| 风险等级 | 低 |
|
||||
|
||||
运行三条命令:
|
||||
|
||||
```bash
|
||||
cargo test --all-targets
|
||||
cargo clippy --all-targets -- -D warnings
|
||||
cargo doc --no-deps 2>&1 | grep "^warning:" && echo "WARNINGS FOUND" || echo "0 warnings"
|
||||
```
|
||||
|
||||
**验收条件**:前两条 0 错误,第三条输出 `0 warnings`
|
||||
|
||||
#### Step 8.2 — Quick Start 示例
|
||||
|
||||
##### commit 5: examples/quick_start.rs
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 涉及文件 | `examples/quick_start.rs` |
|
||||
| 前置依赖 | 无(可从 Phase 7 独立创建) |
|
||||
| 预估工作量 | S(<1h) |
|
||||
| 风险等级 | 低 |
|
||||
|
||||
新文件 `examples/quick_start.rs`,~36 行,结构如下:
|
||||
|
||||
```
|
||||
1- 6 use 块(agcore 类型 + Arrow/std 类型)
|
||||
7- 8 struct Greeter + impl Agent(name / system_prompt)
|
||||
9-14 struct EchoTool + #[async_trait] impl BaseTool(完整 JSON Schema 带 text 参数)
|
||||
15-20 fn mock_response() -> MessageResponse 辅助函数(构造纯文本响应)
|
||||
21-33 #[tokio::main] async fn main():
|
||||
- ToolRegistry::new() + register EchoTool
|
||||
- MockProvider 预设 1 条 mock_response
|
||||
- AgentBuilder::new() + provider + tool_registry + hook_executor → build
|
||||
- AgentSession::new + submit_turn("你好")
|
||||
- println!("{}", response.text())
|
||||
```
|
||||
|
||||
**设计约束**:
|
||||
- EchoTool 的 `parameters()` 返回完整 JSON Schema:`{"type":"object","properties":{"text":{"type":"string"}},"required":["text"]}`
|
||||
- 无外部依赖、无 API key、零配置
|
||||
- 展示四层抽象:Agent trait / BaseTool 自定义 / AgentBuilder 装配 / AgentSession 执行
|
||||
|
||||
**验收条件**:`cargo run --example quick_start` exit 0,输出包含 `"收到"`
|
||||
|
||||
#### Step 8.3 — 端到端示例
|
||||
|
||||
##### commit 6: examples/end_to_end.rs
|
||||
|
||||
| 属性 | 值 |
|
||||
|------|-----|
|
||||
| 涉及文件 | `examples/end_to_end.rs` |
|
||||
| 前置依赖 | commit 5(示例编写模式已建立);SqliteStore(Phase 7 已完成) |
|
||||
| 预估工作量 | M(1-4h) |
|
||||
| 风险等级 | 中 |
|
||||
|
||||
新文件 `examples/end_to_end.rs`,~160 行,最小可行边界(3 工具 + 3 轮 + 持久化验证)。
|
||||
|
||||
**Provider 初始化策略**:
|
||||
|
||||
```
|
||||
if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
|
||||
// 使用真实 Provider(AG_LLM_MODEL 非必填,from_env 内部会处理默认值)
|
||||
let provider: Arc<dyn LlmProvider> = Arc::from(create_provider(
|
||||
ProviderType::OpenaiChat, ProviderConfig::from_env("AG_LLM").unwrap()
|
||||
)?);
|
||||
} else {
|
||||
// MockProvider fallback,预设 4 条响应序列
|
||||
let found = ["AG_LLM_BASE_URL", "AG_LLM_API_KEY"].iter()
|
||||
.filter(|k| env::var(k).is_ok()).collect::<Vec<_>>();
|
||||
eprintln!("AG_LLM_* 环境变量不完整(检测到: {:?}),回退到 MockProvider", found);
|
||||
}
|
||||
```
|
||||
|
||||
**工具定义**:
|
||||
|
||||
| 工具 | 功能 | 关键技术点 |
|
||||
|------|------|-----------|
|
||||
| `EchoTool` | 回显输入 | 基础工具注册模式 |
|
||||
| `CalcTool` | 本地执行四则运算 | 手动解析算术表达式(ponytail:基础 +-*/ 运算无需引入 `rhai` 依赖) |
|
||||
| `NoteTool` | 通过 MemoryStore trait 读写笔记 | 直接持有 `Arc<dyn MemoryStore>`,key 前缀 `"note:"`;save 用 `MemoryStore::save(MemoryItem { id: "note:{key}", content, .. })`,query 用 `MemoryStore::list(MemoryFilter { prefix: Some("note:"), .. })` |
|
||||
|
||||
**持久化验证**:
|
||||
|
||||
```rust
|
||||
let dir = tempfile::TempDir::new()?;
|
||||
let db_path = dir.path().join("agcore.db");
|
||||
let backend = Arc::new(SqliteStore::open(&db_path)?);
|
||||
// ... 构建 RuntimeBundle + AgentSession,写入数据 ...
|
||||
drop(bundle); // 释放所有对 backend 的 Arc 引用
|
||||
drop(session);
|
||||
// 此时 backend 无活跃引用,SQLite 连接自动关闭
|
||||
let backend2 = Arc::new(SqliteStore::open(&db_path)?); // 重建连接
|
||||
// assert 数据仍在
|
||||
```
|
||||
|
||||
**输出示范**:
|
||||
|
||||
```
|
||||
=== agcore 端到端演示 ===
|
||||
🔄 Provider: MockProvider (离线回退模式)
|
||||
💾 SqliteStore: /tmp/agcore_XXXXX/agcore.db
|
||||
🔧 注册工具: echo, calc, note
|
||||
|
||||
第 1 轮 用户: 帮我算 25 * 4
|
||||
→ 调用 calc(...) → 100
|
||||
→ 回答: 25 * 4 = 100
|
||||
|
||||
第 2 轮 用户: 记下来:结果是 100
|
||||
→ 调用 note(save, ...)
|
||||
→ 回答: 已记录
|
||||
|
||||
第 3 轮 用户: 我刚才算了什么?
|
||||
→ 调用 note(query)
|
||||
→ 回答: 您刚才的计算结果是 100
|
||||
|
||||
📊 用量: prompt=XX, completion=XX
|
||||
|
||||
=== 持久化验证 ===
|
||||
✓ 跨连接数据存活验证通过
|
||||
|
||||
✓ 端到端演示完成
|
||||
```
|
||||
|
||||
**设计约束**:
|
||||
- 文件顶部注释说明 `AG_LLM_*` 环境变量切换真实 Provider
|
||||
- 零外部配置可运行(Mock fallback)
|
||||
- 最小可行边界:3 工具 + 3 轮 + 持久化验证,不膨胀
|
||||
|
||||
**验收条件**:`cargo run --example end_to_end` exit 0(零外部配置)
|
||||
|
||||
### 8.2 并行机会
|
||||
|
||||
commit 1 和 commit 5 可以并行执行(零文件重叠)。commit 5 也可与 commit 2 并行。commit 6 实质上也仅依赖「代码库状态稳定」而非某个具体 commit。
|
||||
|
||||
| 并行组 | commit A | commit B | 前提 |
|
||||
|--------|---------|---------|------|
|
||||
| 1 | commit 1(#[non_exhaustive]) | commit 5(Quick Start) | 零文件重叠 |
|
||||
| 2 | commit 2(StepStatus 修复) | commit 5(Quick Start) | 零文件重叠 |
|
||||
| 3 | commit 5(Quick Start) | commit 6(端到端) | 零文件重叠,但存在知识依赖——commit 6 需参考 commit 5 的 `MessageResponse` 构造、`MockProvider` 用法、`AgentBuilder` 装配模式。推荐 commit 5 先行或实施前同步这些模式 |
|
||||
|
||||
### 8.3 风险与应对
|
||||
|
||||
| 风险 | 影响 | 可能性 | 应对 |
|
||||
|------|------|--------|------|
|
||||
| MockProvider 响应序列与 tool-loop 消费顺序不匹配 | commit 6 端到端示例不通过 | 中 | 按 §5 假设 1:设计响应队列时确保每条 Mock 响应的 `stop_reason` 与 ToolUse/Stop 匹配。出现不匹配时改用完整 `MessageResponse` 构造显式控制 |
|
||||
| NoteTool 与 AgentSession 的数据传递路径需要扩展现有 API | commit 6 需要修改 `session.rs` | 低 | ponytail 方案:NoteTool 直接持有 `Arc<dyn MemoryStore>` 引用,在 execute 时直接操作 `MemoryStore::save/get`,绕过 AgentSession 的 session_memory 封装 |
|
||||
| `#[non_exhaustive]` 在某个 enum 上导致 crate 内 match 编译失败 | commit 1 不通过 | 低 | 实施前先运行 `rg "match.*(Message|ContentBlock|ContentBlockType|StreamEvent|HookEvent|AgentError|LlmError|ToolError|MemoryError|PromptError|MemoryStrategy|StepStatus|ToolChoice|ResponseFormat)" src/ --include="*.rs"` 快速扫描 exhaustive match。若某 enum 编译失败,回退该 enum 上的 `#[non_exhaustive]` 属性,标注原因 |
|
||||
|
||||
### 8.4 测试策略
|
||||
|
||||
| commit | 测试 | 方式 |
|
||||
|--------|------|------|
|
||||
| commit 1 | 编译测试 | `cargo build --all-targets` |
|
||||
| commit 2 | 编译 + 单测 + 无 deprecated warning | `cargo build --all-targets && cargo test` |
|
||||
| commit 3 | 人工 review | `git diff` |
|
||||
| commit 4 | 全量自动化 | `cargo test + clippy + doc` |
|
||||
| commit 5 | 示例运行 | `cargo run --example quick_start` |
|
||||
| commit 6 | 示例运行 | `cargo run --example end_to_end` |
|
||||
| 最终 | 全量回归 | 全部三项 + 所有 10 个示例 |
|
||||
@@ -0,0 +1,821 @@
|
||||
# Phase 9 — 流式体验增强实施方案
|
||||
|
||||
- **文档编号**:16
|
||||
- **标题**:Phase 9 — 流式体验增强实施方案
|
||||
- **日期**:2026-07-05
|
||||
- **状态**:已定稿
|
||||
- **涉及模块**:llm/cycle、llm/types/response_v2、agent/session
|
||||
- **关联文档**:roadmap.md(§Phase 9)、15-phase8-mvp-integration.md
|
||||
|
||||
---
|
||||
|
||||
## 1. 背景与目标
|
||||
|
||||
agcore 已发布 v0.2.0-rc.1,Phase 0-8 全部完成。当前 Agent 会话只有非流式 API(`submit_turn`),开发者无法看到实时 token 输出和工具执行过程。Phase 9 的目标是为 `AgentSession` 新增流式方法 `submit_turn_stream`,让开发者能实时看到 LLM token 生成和工具执行状态。
|
||||
|
||||
### 1.1 现有能力
|
||||
|
||||
| 能力 | 方法 | 流式 | 自动工具循环 | 状态 |
|
||||
|------|------|------|-------------|------|
|
||||
| LLM 流式请求 | `LlmCycle::submit_stream` | ✅ | ❌ | 已就绪 |
|
||||
| LLM 工具循环 | `LlmCycle::submit_with_tools` | ❌ | ✅ | 已就绪 |
|
||||
| Agent 会话 | `AgentSession::submit_turn` | ❌ | ✅ | 已就绪 |
|
||||
| 流事件枚举 | `StreamEvent`(11 变体) | — | — | 缺工具执行事件 |
|
||||
| Mock 流 | `MockProvider::chat_stream` | ✅ | — | 可模拟流事件序列 |
|
||||
|
||||
### 1.2 核心矛盾
|
||||
|
||||
流式能力和工具循环能力分别存在于两个方法中,从未被组合。`submit_stream` 只管将 LLM 流事件原样转发,不理解工具调用;`submit_with_tools` 自动执行工具循环但全程阻塞。Phase 9 就是要组合它们:**在工具循环中,每一轮 LLM 调用都是流式的,并在工具执行前后插入语义事件**。
|
||||
|
||||
---
|
||||
|
||||
## 2. 需求分析
|
||||
|
||||
### 2.1 功能需求
|
||||
|
||||
1. **`AgentSession::submit_turn_stream(user_input)`** — 返回 `StreamEvent` 流,开发者通过 `while let Some(event) = stream.next().await` 逐事件消费
|
||||
2. **流式工具循环** — 多轮工具调用过程中流不卡死,每轮工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
3. **`finalize_turn(response)`** — 流消费完成后同步 session 状态(cost 累计 + `OnTurnEnd` hook 触发)
|
||||
4. **新增 `StreamEvent` 变体** — `ToolExecutionStarted` + `ToolExecutionCompleted`,携带工具名称、调用 ID、参数/结果摘要
|
||||
|
||||
### 2.2 非功能需求
|
||||
|
||||
- **零影响**:现有 `submit_turn` 和 `submit_with_tools` 行为不变,存量测试 0 回归
|
||||
- **异步流**:消费者通过 `futures_util::StreamExt::next()` 逐事件消费
|
||||
- **错误事件化**:错误通过 `StreamEvent::Error` 事件表达,不通过 `Result` 通道终止流
|
||||
- **最少代码**:复用现有 `submit_with_tools` 的工具循环逻辑模式和 `submit_stream` 的流管道模式
|
||||
|
||||
### 2.3 不做事项
|
||||
|
||||
| 事项 | 理由 |
|
||||
|------|------|
|
||||
| 新增示例(Phase 9.2 再加) | 缩窄 Phase 9 范围至核心能力 |
|
||||
| `OnTurnEnd` 自动触发 | Rust 所有权约束:流是延迟求值,`&mut self` 无法进入闭包;由消费者收到 `MessageComplete` 后手动调用 `finalize_turn` |
|
||||
| 修复 cost 统计 | 中间轮 cost 丢失是已知限制,与 `submit_turn` 行为一致 |
|
||||
| 跨 turn 消息历史保留 | Phase 10 `ContextSlot` 的职责 |
|
||||
| 并行 tool 调用的事件细化 | 当前工具调用是顺序 `for` 循环,并行化留待后续优化 |
|
||||
| `run_tool_loop` 内消息压缩 | `run_tool_loop` 不接收 `compact_config` 参数,不执行上下文压缩。长工具循环中消息增长可能导致 context window 溢出,这是流式实现的已知限制。后续可通过传递 `compact_config` 给 `run_tool_loop` 支持 |
|
||||
| LLM 请求自动 retry | 流式版本不在 `run_tool_loop` 内部实现 retry(详见 §3.6 说明)。调用方可自行包装 `RetryProvider` 或在 `LlmProvider` 实现层处理 |
|
||||
|
||||
---
|
||||
|
||||
## 3. 方案设计
|
||||
|
||||
### 3.1 架构总览
|
||||
|
||||
```
|
||||
┌──────────────────────────────────────────────────────────────┐
|
||||
│ AgentSession │
|
||||
│ ┌──────────────────────────────────────────────────────┐ │
|
||||
│ │ submit_turn_stream() │ │
|
||||
│ │ ├─ OnTurnStart hook(同步触发,返回流之前) │ │
|
||||
│ │ ├─ 组装 LlmCycle(system_prompt / compact_config) │ │
|
||||
│ │ ├─ 调用 submit_with_tools_stream() │ │
|
||||
│ │ ├─ turn_index += 1 │ │
|
||||
│ │ └─ 返回流 │ │
|
||||
│ │ │ │
|
||||
│ │ finalize_turn(response) │ │
|
||||
│ │ ├─ cost_so_far.add(&response.usage) │ │
|
||||
│ │ └─ OnTurnEnd hook(turn_index - 1) │ │
|
||||
│ └──────────────────────────────────────────────────────┘ │
|
||||
│ submit_with_tools_stream(prompt, Arc<ToolRegistry>)
|
||||
▼
|
||||
┌──────────────────────────────────────────────────────────────┐
|
||||
│ LlmCycle (tokio::spawn task — run_tool_loop 状态机) │
|
||||
│ │
|
||||
│ max_turns = max_tool_turns.unwrap_or(10) │
|
||||
│ for round in 1..=max_turns { │
|
||||
│ ① build_request(messages, tools) │
|
||||
│ ② provider.chat_stream(request).await │
|
||||
│ 匹配 Err → tx.send(Error{..}) + return(不 panic) │
|
||||
│ ③ 消费 LLM 流,所有事件 → mpsc unbounded tx(全量转发) │
|
||||
│ ④ partial.finalize() → MessageResponse │
|
||||
│ ⑤ if has_tool_use: │
|
||||
│ ├─ tx → ToolExecutionStarted { tool_name, id, args } │
|
||||
│ ├─ registry.invoke_all(calls, timeout).await │
|
||||
│ ├─ for result: tx → ToolExecutionCompleted { ... } │
|
||||
│ ├─ push tool results → messages │
|
||||
│ └─ continue(新一轮) │
|
||||
│ else: break(最终轮,已发出 MessageComplete) │
|
||||
│ } │
|
||||
│ │
|
||||
│ 产出事件序列(通过 mpsc::unbounded_channel): │
|
||||
│ MessageStart → ... → ToolCallEnd → ToolExecutionStarted → │
|
||||
│ ToolExecutionCompleted → MessageStart → TextDelta → ... → │
|
||||
│ CostUpdate → MessageComplete │
|
||||
└──────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
> **`max_tool_turns` 语义**:与非流式 `submit_with_tools` 一致——`None` 退化为 `10`(`unwrap_or(10)`)。默认值 `Some(10)` 已提供安全上限;如需增大限制,手动设置为 `Some(N)`。⚠️ 生产环境建议始终设有限值防止无限循环。
|
||||
|
||||
### 3.2 事件序列约定
|
||||
|
||||
**纯文本流**(无 tool_use):
|
||||
|
||||
```
|
||||
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd → CostUpdate → MessageComplete
|
||||
```
|
||||
|
||||
**单轮工具调用**:
|
||||
|
||||
```
|
||||
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd
|
||||
→ ContentBlockStart → ToolCallArgumentsDelta* → ToolCallEnd
|
||||
→ CostUpdate → MessageComplete { stop_reason: ToolUse }
|
||||
→ ToolExecutionStarted → [工具执行] → ToolExecutionCompleted
|
||||
→ ContentBlockStart → TextDelta* → ContentBlockEnd
|
||||
→ CostUpdate → MessageComplete { stop_reason: Stop }
|
||||
```
|
||||
|
||||
**多轮工具调用**:
|
||||
|
||||
```
|
||||
... → ToolExecutionCompleted(第 1 轮)
|
||||
→ ToolCallArgumentsDelta* → ToolCallEnd
|
||||
→ ToolExecutionStarted → ToolExecutionCompleted(第 2 轮)
|
||||
→ ... → CostUpdate → MessageComplete(最终轮)
|
||||
```
|
||||
|
||||
**工具不可恢复错误**:
|
||||
|
||||
```
|
||||
... → ToolCallEnd → ToolExecutionStarted
|
||||
→ Error { "tool 'search' 不可恢复错误: ..." } → MessageComplete
|
||||
```
|
||||
|
||||
> **`MessageComplete.full_response` 内容范围**:每轮 LLM 调用独立产生一个 `MessageComplete`,其中 `full_response` 仅包含**该轮 LLM 的单个响应**(不累积前面工具轮次的结果)。中间轮(`stop_reason: ToolUse`)的 `full_response` 通常只包含 `ToolUse` block,无文本。最终轮(`stop_reason: Stop`)的 `full_response` 包含 LLM 的最终输出。消费者如需追踪完整对话历史,应自行累加所有轮次的 `Message`。
|
||||
|
||||
### 3.3 StreamEvent 新增变体
|
||||
|
||||
在 `src/llm/types/response_v2.rs` 的 `StreamEvent` 枚举中追加两个变体:
|
||||
|
||||
```rust
|
||||
/// 工具开始执行 —— 在 ToolCallEnd 之后、registry.invoke 之前发出。
|
||||
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
|
||||
ToolExecutionStarted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 工具参数(JSON 字符串形式)
|
||||
arguments: String,
|
||||
},
|
||||
|
||||
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
|
||||
ToolExecutionCompleted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 结果摘要(前 200 字符)
|
||||
result_summary: String,
|
||||
/// 是否出错
|
||||
is_error: bool,
|
||||
},
|
||||
```
|
||||
|
||||
在 `PartialMessageResponse::apply_to` 中追加:
|
||||
|
||||
```rust
|
||||
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
|
||||
```
|
||||
|
||||
这两个是**元事件**,不参与内容块累积,`apply_to` 直接返回 `true`。
|
||||
|
||||
### 3.4 新增方法签名
|
||||
|
||||
**`LlmCycle` 层**(`src/llm/cycle.rs`):
|
||||
|
||||
```rust
|
||||
/// 提交消息并自动处理工具调用循环,流式产出所有事件。
|
||||
///
|
||||
/// 与 `submit_with_tools` 的区别:
|
||||
/// - LLM 响应是流式的(全程 `chat_stream` 而非 `chat`)
|
||||
/// - 工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
/// - 错误以 `StreamEvent::Error` 形式出现在流中,而非终止 `Result`
|
||||
/// - 消费方需手动 `push_message()` 同步消息历史
|
||||
///
|
||||
/// **运行时要求**:内部使用 `tokio::spawn`,需要 tokio 多线程运行时。
|
||||
pub async fn submit_with_tools_stream(
|
||||
&mut self,
|
||||
prompt: String,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, LlmError>
|
||||
```
|
||||
|
||||
```rust
|
||||
/// 运行工具循环的核心异步状态机。
|
||||
///
|
||||
/// 接收 owned 字段,通过 mpsc::unbounded_channel 产出事件序列。
|
||||
/// 由 `submit_with_tools_stream` 在 tokio::spawn 中调用。
|
||||
///
|
||||
/// **运行时要求**:此函数内部使用 `tokio::spawn`,要求调用方运行在
|
||||
/// tokio 多线程运行时中(`#[tokio::main]` 或 `#[tokio::test(flavor = "multi_thread")]`)。
|
||||
/// 不在 WASM 目标下可用。
|
||||
async fn run_tool_loop(
|
||||
messages: Vec<Message>,
|
||||
provider: Arc<dyn LlmProvider>,
|
||||
config: CycleConfig,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
tools: Vec<ToolDef>,
|
||||
tx: mpsc::UnboundedSender<StreamEvent>,
|
||||
hook_executor: Option<Arc<HookExecutor>>,
|
||||
)
|
||||
```
|
||||
|
||||
**`AgentSession` 层**(`src/agent/session.rs`):
|
||||
|
||||
```rust
|
||||
/// 提交一轮对话(流式,含自动 tool 循环),返回 `StreamEvent` 流。
|
||||
///
|
||||
/// 与 `submit_turn` 的区别:
|
||||
/// - 以流事件序列而非 `MessageResponse` 返回
|
||||
/// - 工具执行期间插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
/// - 消费方在收到 `MessageComplete` 后需手动调用 `finalize_turn` 同步状态
|
||||
///
|
||||
/// **运行时要求**:内部委托 `submit_with_tools_stream`,需要 tokio 多线程运行时。
|
||||
pub async fn submit_turn_stream(
|
||||
&mut self,
|
||||
user_input: impl Into<String>,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, AgentError>
|
||||
|
||||
/// 完成一轮 turn:累计 cost + 触发 OnTurnEnd hook。
|
||||
///
|
||||
/// 由消费者在收到 `MessageComplete.full_response` 后调用。
|
||||
pub async fn finalize_turn(&mut self, response: &MessageResponse)
|
||||
```
|
||||
|
||||
### 3.5 消费者使用模式
|
||||
|
||||
```rust
|
||||
use futures_util::StreamExt;
|
||||
|
||||
let mut stream = session.submit_turn_stream("计算 1+2").await?;
|
||||
|
||||
let mut final_response = None;
|
||||
while let Some(event) = stream.next().await {
|
||||
match &event {
|
||||
StreamEvent::TextDelta { text } => print!("{}", text),
|
||||
StreamEvent::ToolExecutionStarted { tool_name, arguments, .. } => {
|
||||
println!("\n🔧 [{}({})]", tool_name, arguments);
|
||||
}
|
||||
StreamEvent::ToolExecutionCompleted { result_summary, .. } => {
|
||||
println!(" → {}", result_summary);
|
||||
}
|
||||
StreamEvent::MessageComplete { full_response } => {
|
||||
final_response = Some(full_response.clone());
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
std::io::stdout().flush().ok();
|
||||
|
||||
if let Some(response) = final_response {
|
||||
session.finalize_turn(&response).await;
|
||||
}
|
||||
```
|
||||
|
||||
> **⚠️ 消费者注意**:`finalize_turn` 是开发者责任 —— 遗漏调用会导致 `cost_so_far` 不累计、`OnTurnEnd` hook 不触发。session 状态仍然可用,后续 `submit_turn` 也能正常执行,但 cost 信息不完整。`finalize_turn` 无自动补偿机制,建议使用 `Drop` guard 或在 `while` 循环的 `finally` 块中确保调用。
|
||||
|
||||
### 3.6 run_tool_loop 核心逻辑
|
||||
|
||||
`run_tool_loop` 是此方案的核心状态机(约 90 行),其伪代码逻辑如下:
|
||||
|
||||
```
|
||||
1. 接收 owned 字段:messages, provider, config, tool_registry, tools, tx, hook_executor
|
||||
2. max_turns = config.max_tool_turns.unwrap_or(10)
|
||||
// None → 10(退化为默认值),Some(n) → n
|
||||
// 与非流式 submit_with_tools 行为一致
|
||||
3. 工具循环(for round in 1..=max_turns):
|
||||
a. build_request(messages, tools)
|
||||
// 空 tool_registry 时 tools 为空列表,流退化为纯文本流(可安全运行)
|
||||
b. PreRequest hook(如果有 hook_executor)
|
||||
c. 发起流式 LLM 调用:
|
||||
let stream = match provider.chat_stream(request).await {
|
||||
Ok(s) => s,
|
||||
Err(e) => {
|
||||
// 第一层错误:chat_stream 自身失败(网络/认证/限流)
|
||||
// 这里不做 retry:retry 逻辑留给上层循环的 submit_request 模式,
|
||||
// 流式场景中 retry 需重新建立 mpsc 通道,复杂度与收益不匹配
|
||||
tx.send(StreamEvent::Error { message: e.to_string() }).ok();
|
||||
return; // 直接结束 task
|
||||
}
|
||||
};
|
||||
d. 消费 LLM 流:
|
||||
- PartialMessageResponse::new()
|
||||
- while let Some(result) = stream.next().await
|
||||
- match result:
|
||||
Ok(event) → apply_to + tx.send(event)
|
||||
Err(e) → tx.send(Error { message }) + break
|
||||
// 第二层错误:stream 内部事件错误(如 chunk 解析失败)
|
||||
e. partial.finalize()? → response
|
||||
f. push response.message → messages
|
||||
g. 检查 has_tool_calls_in_response(&response)
|
||||
h. 如果没有 tool_use: break(最终轮,流已自然结束)
|
||||
i. 如果有 tool_use:
|
||||
- extract_tool_calls_from_response(&response)
|
||||
- tx.send(ToolExecutionStarted { tool_name, tool_call_id, arguments })
|
||||
- registry.invoke_all(calls, tool_timeout).await
|
||||
- for result in results:
|
||||
tx.send(ToolExecutionCompleted { tool_name, tool_call_id, result_summary, is_error })
|
||||
- push tool results → messages
|
||||
- continue(新一轮 LLM 流)
|
||||
4. 流结束(tokio::spawn 自然退出)
|
||||
```
|
||||
|
||||
> **关于 LLM 请求 retry**:非流式 `submit_with_tools` 内部通过 `submit_request` 的 retry 循环处理临时错误。流式版本 `run_tool_loop` **不在内部实现 retry**。原因:(1)retry 需要重新建立 mpsc 通道和事件流上下文,复杂度与收益不匹配;(2)`unbounded_channel` 已发出的事件无法撤回。如果需要 retry 语义,调用方应在上层做 fallback 策略,或在 `llm provider` 实现层完成 retry(如 `RetryProvider` 包装器)。
|
||||
|
||||
**错误处理**:
|
||||
|
||||
| 场景 | 行为 |
|
||||
|------|------|
|
||||
| LLM 请求失败(`chat_stream` 返回 `Err`) | `tx.send(Error { message })` + `return` 结束 task。**不做 retry**(见上方说明) |
|
||||
| LLM 流内事件错误(stream Item 的 `Err`) | `tx.send(Error { message })` + `break` 结束当轮流,终止循环 |
|
||||
| 可恢复工具错误(`is_recoverable() == true`) | 作为 tool result 回传 LLM,流继续,不出 Error 事件 |
|
||||
| 不可恢复工具错误(`is_recoverable() == false`) | `tx.send(Error { message })` + 终止循环 |
|
||||
| 工具超时(`tokio::time::timeout`) | 视为不可恢复,`tx.send(Error)` + 终止 |
|
||||
| 最大工具循环轮次超限 | `tx.send(Error { "达到最大工具循环轮次" })` + 终止 |
|
||||
| spawn task 内部 panic | 由于 `JoinHandle` 不保存(detached),panic 由 tokio 运行时静默捕获;消费者看到 stream 直接结束(返回 `None`),无 `Error` 事件。建议在 `run_tool_loop` 内部避免 `unwrap()`,所有可失败路径通过 `Result` + `?` 传播 |
|
||||
|
||||
### 3.7 修改文件清单
|
||||
|
||||
| # | 文件 | 改动 | 估算行数 |
|
||||
|---|------|------|---------|
|
||||
| 1 | `llm/types/response_v2.rs` | +2 `StreamEvent` 变体 +2 `apply_to` arm | ~20 |
|
||||
| 2 | `llm/cycle.rs` | +`submit_with_tools_stream` 方法 + `run_tool_loop` 模块函数 | ~140 |
|
||||
| 3 | `llm/cycle.rs` | `CycleConfig` 加 `#[derive(Clone)]` | ~1 |
|
||||
| 4 | `agent/session.rs` | +`submit_turn_stream` + `finalize_turn` | ~70 |
|
||||
| — | **测试**(内联) | 4 个场景测试(纯度本、单轮、多轮、超限) | ~150 |
|
||||
| | **合计** | | **~380** |
|
||||
|
||||
> 注:`RetryConfig` 已标注 `#[derive(Debug, Clone)]`,无需额外修改。
|
||||
|
||||
---
|
||||
|
||||
## 4. 实现计划
|
||||
|
||||
按 5 个 Step 增量实施,每步可独立编译和测试。
|
||||
|
||||
### Step 1 — 基础设施准备
|
||||
|
||||
**目标**:数据层就绪,为流事件新增变体和配置 Clone 奠基。
|
||||
|
||||
**改动**:
|
||||
|
||||
- `llm/types/response_v2.rs`:
|
||||
- `StreamEvent` 枚举追加 `ToolExecutionStarted` / `ToolExecutionCompleted` 变体
|
||||
- `PartialMessageResponse::apply_to` 追加两个新变体的 arm(均返回 `true`)
|
||||
- `llm/cycle.rs`:
|
||||
- `CycleConfig` 加 `#[derive(Clone)]`(所有字段为基础类型 + `RetryConfig`)
|
||||
|
||||
**验证**:`cargo build` 通过
|
||||
|
||||
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
|
||||
|
||||
**目标**:实现流式工具循环的核心状态机,这是整个 Phase 9 的技术关键。
|
||||
|
||||
**改动**:
|
||||
|
||||
- `llm/cycle.rs`:
|
||||
- 新增 `run_tool_loop()` 模块函数(约 90 行),基于 `mpsc::unbounded_channel` 通信
|
||||
- 新增 `submit_with_tools_stream()` 公开方法,入口参数为 `prompt` + `Arc<ToolRegistry>`
|
||||
- 内部 `tokio::spawn` 启动 `run_tool_loop`,返回 `rx` 端作为 `dyn Stream`
|
||||
|
||||
**验证**:`cargo build` 通过
|
||||
|
||||
### Step 3 — 单元测试(LlmCycle 层)
|
||||
|
||||
**目标**:验证 `submit_with_tools_stream` 在 8 个核心场景下的行为和事件序列正确性(含 §8 Step 3 扩展的工具错误路径)。
|
||||
|
||||
**新增**(`llm/cycle.rs` 内联测试 `#[cfg(test)]`):
|
||||
|
||||
| 场景 | Mock 响应序列 | 验证点 |
|
||||
|------|---------------|--------|
|
||||
| 1 — 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` |
|
||||
| 2 — 单轮工具调用 | 2 个响应:tool_use → text | 包含 `ToolExecutionStarted` + `ToolExecutionCompleted`;最终 `stop_reason` 为 `Stop` |
|
||||
| 3 — 多轮工具调用 | 4 个响应:3 × tool_use → 1 × text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度正确 |
|
||||
| 4 — 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `Error` 事件;消息历史停在第 2 轮 |
|
||||
|
||||
**验证**:`cargo test` 全部通过
|
||||
|
||||
### Step 4 — `AgentSession` 层包装
|
||||
|
||||
**目标**:为 `AgentSession` 新增流式会话接口,保持与 `submit_turn` 一致的行为语义。
|
||||
|
||||
**改动**:
|
||||
|
||||
- `agent/session.rs`:
|
||||
- `submit_turn_stream(user_input)` — 触发 `OnTurnStart` hook → 组装 `LlmCycle` → 调用 `submit_with_tools_stream` → `turn_index += 1` → 返回流
|
||||
- `finalize_turn(response)` — `cost_so_far.add(&response.usage)` → 触发 `OnTurnEnd` hook
|
||||
|
||||
**验证**:`cargo build` 通过
|
||||
|
||||
### Step 5 — 集成测试 + 扫尾
|
||||
|
||||
**目标**:端到端验证 `submit_turn_stream` + `finalize_turn` 的完整链路,确保零回归。
|
||||
|
||||
**新增**(`agent/session.rs` 内联测试):
|
||||
|
||||
- **场景**:`submit_turn_stream` 跑通 mock provider → 消费流(验证各事件到达) → `finalize_turn` 后 cost 更新正确
|
||||
- **场景**:verify `OnTurnStart` hook 在 `submit_turn_stream` 返回流之前已触发
|
||||
|
||||
**验证**:
|
||||
|
||||
```bash
|
||||
cargo test --all-targets # 全绿,存量测试 0 回归
|
||||
cargo clippy --all-targets -- -D warnings # 0 警告
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 5. 运行细节
|
||||
|
||||
### 5.1 `run_tool_loop` 的 spawn 生命周期
|
||||
|
||||
#### 执行模型:立即执行 vs 惰性流
|
||||
|
||||
`submit_with_tools_stream` 采用 **立即执行** 模型(`tokio::spawn` + `mpsc`),这与 `submit_stream` 的 **惰性执行**(`async_stream::stream!` 宏,消费者首次 `next()` 时才触发 LLM 调用)不同。
|
||||
|
||||
**选择理由**:工具循环是 **不确定轮次的** —— 每个工具执行的结果可能影响后续 LLM 调用。惰性流无法表达这种"边消费边控制"的语义。通过 `tokio::spawn` 将工具循环移到独立 task 中运行,使得:
|
||||
- 消费者可以随时开始消费(不丢失事件)
|
||||
- 工具循环在后台独立运行,不受消费者消费节奏影响
|
||||
- `mpsc::unbounded_channel` 作为事件缓冲区,解耦生产者与消费者
|
||||
|
||||
**对消费者的影响**:`submit_with_tools_stream().await?` 返回时,工具循环可能已经开始执行(事件已开始写入 channel)。消费者应尽快开始 `while let Some(event) = stream.next().await`,避免 channel 缓冲过多事件。如果在返回流后长时间不消费,事件会堆积在 mpsc buffer 中(内存开销,无阻塞风险 —— 见 §6 风险表)。
|
||||
|
||||
#### 生命周期
|
||||
|
||||
```
|
||||
submit_with_tools_stream()
|
||||
│
|
||||
├─ mpsc::unbounded_channel() → (tx, rx)
|
||||
├─ messages.push(user_text(prompt))
|
||||
├─ compact check
|
||||
├─ tokio::spawn(run_tool_loop(messages, provider, config, ..., tx))
|
||||
└─ return Box::pin(rx) as dyn Stream
|
||||
|
||||
[用户消费 stream]
|
||||
└─ while let Some(event) = rx.recv().await { yield event }
|
||||
|
||||
[用户 drop rx / 结束循环]
|
||||
└─ rx 被 drop → tx.send() 返回 Err
|
||||
→ run_tool_loop 检测到 tx.closed()
|
||||
→ break → task 自然终止
|
||||
```
|
||||
|
||||
#### JoinHandle 与 panic 处理
|
||||
|
||||
`run_tool_loop` 的 `JoinHandle` 在 spawn 后**不保存**(detached pattern)。panic 由 tokio 运行时捕获并通过 `tracing::error` 记录:
|
||||
|
||||
```rust
|
||||
// submit_with_tools_stream 内部
|
||||
tokio::spawn(async move {
|
||||
run_tool_loop(..., tx).await;
|
||||
});
|
||||
```
|
||||
|
||||
如果 `run_tool_loop` 内部发生 panic(如 `unwrap()`),tokio 的 `spawn` 会静默吞掉 panic 并终止 task。消费者此时看到 stream 直接返回 `None`,不会收到 `StreamEvent::Error`。实际编码中应避免 `unwrap()`,所有 `Result` 使用 `?` 或 `match` 处理。
|
||||
|
||||
Rx 侧实现 `Stream` trait:使用 `tokio_stream::wrappers::UnboundedReceiverStream` 包装 `mpsc::UnboundedReceiver`,因为 `mpsc::UnboundedReceiver` 本身不实现 `Stream`(`tokio-stream = "0.1"` 已在 `Cargo.toml` 中存在)。
|
||||
|
||||
### 5.2 消息历史同步
|
||||
|
||||
`submit_with_tools_stream` 内部由 `run_tool_loop` 管理 `messages` 的拷贝,不会写入 `self.messages`。消费方在收到 `MessageComplete` 后需手动:
|
||||
|
||||
```rust
|
||||
let response = full_response.clone();
|
||||
cycle.push_message(response.message.clone());
|
||||
```
|
||||
|
||||
在 `AgentSession::submit_turn_stream` 中,由于流是延迟求值且 `&mut self` 无法进入 spawn 闭包,消息历史同步交由消费方在 `finalize_turn` 前自行决定。当前方案中 `submit_turn_stream` **不自动同步消息历史**,这与 `submit_stream` 的已有行为一致(ponytail: Phase 2 FIX-E 注释)。
|
||||
|
||||
---
|
||||
|
||||
## 6. 风险评估
|
||||
|
||||
| 风险 | 影响 | 缓解措施 |
|
||||
|------|------|---------|
|
||||
| `&mut self` 约束导致流内无法访问 session 状态 | 中 | 复用 `submit_stream` 已有模式:方法体内读取 `self` 后构建 owned 数据,spawn 闭包不捕获 `&mut self` |
|
||||
| spawn task 生命周期管理 | 低 | 用户 drop rx → `tx.send` 返回 `Err` → `run_tool_loop` 自然终止 |
|
||||
| spawn task panic 静默丢失 | 中 | `run_tool_loop` 内部使用 `match`/`?` 避免 `unwrap()`;`JoinHandle` 不做 `await`(detached),panic 由 tokio 运行时记录日志。消费者看到 stream 提前结束(收到 `None`)但无 Error 事件 |
|
||||
| 中间轮 cost 不累加到 `cost_so_far` | 低 | 与现有 `submit_turn` 行为一致(仅最终轮计入),标记为已知限制,不在此 Phase 修复 |
|
||||
| 工具循环中 hook 可用性 | 低 | `PreRequest`/`PostRequest` hook 通过 `hook_executor.clone()` 进入 spawn task;hook 在 `run_tool_loop` 循环内触发 |
|
||||
| `run_tool_loop` 不支持消息压缩 | 中 | 长工具循环中消息不断增长,可能超出 context window。当前不传递 `compact_config`,后续可扩展 `run_tool_loop` 签名增添此参数 |
|
||||
| `unbounded_channel` 在消费慢于生产时内存增长 | 低 | LLM 流式输出天然有节流(token 生成速度远慢于 CPU 处理速度),消费者通常快于生产者。后续如需背压可切换为 `mpsc::channel(N)` + backpressure |
|
||||
| 流式版本不做 LLM retry | 低 | 非流式 `submit_with_tools` 通过 `submit_request` 的 retry 循环处理临时错误。流式版本中 retry 需重建 mpsc 通道,复杂度不匹配。调用方可使用 `RetryProvider` 包装器或在 Provider 层实现 retry |
|
||||
| 执行模式与 `submit_stream` 不一致(立即 vs 惰性) | 低 | `submit_stream` 的惰性语义不适配需要后台执行的工具循环。消费者应在 `submit_turn_stream` 返回后尽快消费流事件 |
|
||||
| `tokio::spawn` 要求 tokio 多线程运行时 | 低 | agcore 已依赖 tokio,涉及 IO 的 API 均使用 async。`#[tokio::test]` 单线程运行时不支持 `spawn`,测试中将 `run_tool_loop` 提取为可独立调用的函数,测试不走 spawn 直接调用 |
|
||||
| `CycleConfig` 加 `Clone` 影响现有代码 | 无 | 纯配置 struct,所有字段是基础类型或已 `Clone` 的 `RetryConfig` |
|
||||
|
||||
---
|
||||
|
||||
## 7. 验收标准
|
||||
|
||||
| # | 验收项 | 验证方式 |
|
||||
|---|--------|---------|
|
||||
| 1 | `cargo build --all-targets` 通过 | ✅ 编译器无错误 |
|
||||
| 2 | `submit_with_tools_stream` 纯文本流事件序列正确 | 单元测试验证:事件类型、顺序与 `submit_stream` 一致 |
|
||||
| 3 | `submit_with_tools_stream` 单轮工具调用事件序列正确 | 单元测试验证:含 `ToolExecutionStarted` / `ToolExecutionCompleted` |
|
||||
| 4 | `submit_with_tools_stream` 多轮工具调用事件序列正确 | 单元测试验证:多对 `ToolExecutionStarted`/`ToolExecutionCompleted` |
|
||||
| 5 | `submit_with_tools_stream` 最大轮次超限产生 Error 事件 | 单元测试验证:流中出现 `StreamEvent::Error` |
|
||||
| 6 | `submit_turn_stream` + `finalize_turn` 端到端链路 | 集成测试验证:cost 更新 + hook 触发 |
|
||||
| 7 | `cargo test --all-targets` 全绿,存量测试 0 回归 | ✅ 无回归 |
|
||||
| 8 | `cargo clippy --all-targets -- -D warnings` 0 警告 | ✅ 无警告 |
|
||||
| 9 | 现有 `submit_turn` / `submit_with_tools` / `submit_stream` 行为零影响 | ✅ 存量测试通过 |
|
||||
|
||||
---
|
||||
|
||||
---
|
||||
|
||||
## 8. 实施计划
|
||||
|
||||
按 5 个 Step 分阶段实施,每步产出独立 commit,可验证后退。
|
||||
|
||||
### 依赖关系
|
||||
|
||||
```mermaid
|
||||
graph LR
|
||||
S1["Step 1: 基础设施"]:::s1
|
||||
S2["Step 2: 核心状态机"]:::s2
|
||||
S3["Step 3: LlmCycle 单元测试"]:::s3
|
||||
S4["Step 4: AgentSession 包装"]:::s4
|
||||
S5["Step 5: 集成测试 + 扫尾"]:::s5
|
||||
|
||||
S1 --> S2
|
||||
S1 --> S4
|
||||
S2 --> S3
|
||||
S2 --> S4
|
||||
S3 --> S5
|
||||
S4 --> S5
|
||||
|
||||
classDef s1 fill:#e2e8f0,stroke:#94a3b8
|
||||
classDef s2 fill:#fbbf24,stroke:#d97706
|
||||
classDef s3 fill:#93c5fd,stroke:#2563eb
|
||||
classDef s4 fill:#93c5fd,stroke:#2563eb
|
||||
classDef s5 fill:#4ade80,stroke:#16a34a
|
||||
```
|
||||
|
||||
| Step | 依赖 | 并行机会 |
|
||||
|------|------|---------|
|
||||
| S1 | 无 | — |
|
||||
| S2 | S1 | 可与 S4 并行 |
|
||||
| S3 | S2 | 阻塞,需 S2 完成 |
|
||||
| S4 | S1, S2 | 功能依赖 S2(调用 `submit_with_tools_stream`);文件级无重叠但需先编译过 S2 |
|
||||
| S5 | S3 + S4 | 需 S3 和 S4 都完成 |
|
||||
|
||||
### Step 1 — 基础设施准备
|
||||
|
||||
**工作量**:S(< 1h)
|
||||
**风险**:低(纯新增,不影响现有代码逻辑)
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 1.1 | `StreamEvent` 枚举追加 `ToolExecutionStarted` 变体 | `llm/types/response_v2.rs` | 无 | 低 |
|
||||
| 1.2 | `StreamEvent` 枚举追加 `ToolExecutionCompleted` 变体 | `llm/types/response_v2.rs` | 1.1 | 低 |
|
||||
| 1.3 | `PartialMessageResponse::apply_to` 追加两个元事件 arm(均返回 `true`) | `llm/types/response_v2.rs` | 1.2 | 低 |
|
||||
| 1.4 | `CycleConfig` 加 `#[derive(Clone)]` | `llm/cycle.rs` | 无 | 低 |
|
||||
|
||||
**验收条件**:
|
||||
- `cargo build` 通过,编译器无 warning
|
||||
- 新增的 `StreamEvent` 变体可通过 `serde` roundtrip 序列化/反序列化
|
||||
- `CycleConfig` 可正常 clone
|
||||
|
||||
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
|
||||
|
||||
**工作量**:M(1-4h)
|
||||
**风险**:中(核心实现,需正确设计 spawn + mpsc 生命周期)
|
||||
**前置依赖**:S1
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 2.1 | 实现 `run_tool_loop()` 模块函数:消息循环构建请求 → `chat_stream` → 消费流 → 检测 tool_use → 工具执行 → 新一轮 | `llm/cycle.rs` | S1 | 中 |
|
||||
| 2.2 | 实现 `submit_with_tools_stream()` 公开方法:提取字段 → spawn `run_tool_loop` → 返回 `UnboundedReceiverStream` | `llm/cycle.rs` | 2.1 | 中 |
|
||||
| 2.3 | 新增导入:`tokio::sync::mpsc`、`tokio_stream::wrappers::UnboundedReceiverStream` | `llm/cycle.rs` | 2.2 | 低 |
|
||||
|
||||
**关键实现细节**:
|
||||
|
||||
```rust
|
||||
// run_tool_loop 函数签名
|
||||
async fn run_tool_loop(
|
||||
mut messages: Vec<Message>,
|
||||
provider: Arc<dyn LlmProvider>,
|
||||
config: CycleConfig,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
tools: Vec<ToolDef>,
|
||||
tx: mpsc::UnboundedSender<StreamEvent>,
|
||||
hook_executor: Option<Arc<HookExecutor>>,
|
||||
) {
|
||||
let max_turns = config.max_tool_turns.unwrap_or(10);
|
||||
let tool_timeout = config.tool_timeout_secs;
|
||||
let max_bytes = config.max_tool_result_bytes;
|
||||
|
||||
let mut round = 0u32;
|
||||
loop {
|
||||
round += 1;
|
||||
if round > max_turns {
|
||||
// §3.6 错误表:最大轮次超限 → Error 事件 + 终止
|
||||
let _ = tx.send(StreamEvent::Error { message: "达到最大工具循环轮次".to_string() });
|
||||
break;
|
||||
}
|
||||
// ① 构建请求
|
||||
let request = MessageRequest {
|
||||
model: config.model.clone(),
|
||||
messages: messages.clone(),
|
||||
tools: tools.clone(),
|
||||
tool_choice: ToolChoice::Auto,
|
||||
max_tokens: config.max_tokens,
|
||||
temperature: config.temperature,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
// ② PreRequest hook
|
||||
// ...
|
||||
|
||||
// ③ chat_stream
|
||||
let stream = match provider.chat_stream(request).await {
|
||||
Ok(s) => s,
|
||||
Err(e) => {
|
||||
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// ④ 消费流
|
||||
let mut partial = PartialMessageResponse::new();
|
||||
let mut stream = stream;
|
||||
while let Some(result) = stream.next().await {
|
||||
match result {
|
||||
Ok(event) => {
|
||||
partial.apply_to(&event);
|
||||
if tx.send(event).is_err() { return; }
|
||||
}
|
||||
Err(e) => {
|
||||
// ponytail: 流内事件错误后 partial 处于损坏状态,
|
||||
// 不能继续执行 finalize/finalize —— 直接 return 结束 task
|
||||
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ⑤ finalize
|
||||
let response = match partial.finalize() {
|
||||
Ok(r) => r,
|
||||
Err(e) => { let _ = tx.send(StreamEvent::Error { .. }); return; }
|
||||
};
|
||||
messages.push(response.message.clone());
|
||||
|
||||
// ⑥ 检测 tool_use
|
||||
if !has_tool_calls_in_response(&response) {
|
||||
break; // 最终轮
|
||||
}
|
||||
|
||||
// ⑦ 执行工具
|
||||
let tool_calls = extract_tool_calls_from_response(&response);
|
||||
let calls: Vec<_> = tool_calls.into_iter()
|
||||
.map(|(id, name, args)| {
|
||||
let value = serde_json::from_str(&args).unwrap_or(Value::Null);
|
||||
(id, name, value)
|
||||
}).collect();
|
||||
|
||||
for (tool_call_id, tool_name, args_value) in &calls {
|
||||
let args_json = serde_json::to_string(&args_value).unwrap_or_default();
|
||||
if tx.send(StreamEvent::ToolExecutionStarted {
|
||||
tool_name: tool_name.clone(),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
arguments: args_json,
|
||||
}).is_err() { return; }
|
||||
}
|
||||
|
||||
let results = tool_registry.invoke_all(calls, tool_timeout).await;
|
||||
|
||||
for result in &results {
|
||||
let summary = match &result.output {
|
||||
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
|
||||
Err(e) => e.to_string(),
|
||||
};
|
||||
// ponytail: 复用现有 truncate_tool_result 函数(cycle.rs 末尾),
|
||||
// 确保多字节 UTF-8 字符不被截断破坏。上限 200 字符。
|
||||
let truncated = truncate_tool_result(&summary, 200);
|
||||
if tx.send(StreamEvent::ToolExecutionCompleted {
|
||||
tool_name: result.tool_name.clone(),
|
||||
tool_call_id: result.tool_call_id.clone(),
|
||||
result_summary: truncated,
|
||||
is_error: result.output.is_err(),
|
||||
}).is_err() { return; }
|
||||
}
|
||||
|
||||
for result in results {
|
||||
let is_error = result.output.is_err();
|
||||
let content = match &result.output {
|
||||
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
|
||||
Err(e) if e.is_recoverable() => format!("错误: {}", e),
|
||||
Err(e) => {
|
||||
let _ = tx.send(StreamEvent::Error { .. });
|
||||
return;
|
||||
}
|
||||
};
|
||||
messages.push(Message::tool_result(result.tool_call_id, content, is_error));
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**验收条件**:
|
||||
- `cargo build` 通过
|
||||
- 新增方法签名与方案设计一致
|
||||
- 未修改现有 `submit_with_tools`/`submit_stream` 的行为
|
||||
|
||||
### Step 3 — 单元测试(LlmCycle 层)
|
||||
|
||||
**工作量**:M(1-4h)
|
||||
**风险**:低(与现有测试模式一致,使用已有 MockProvider)
|
||||
**前置依赖**:S2
|
||||
|
||||
测试策略:直接使用公开的 `crate::llm::mock::MockProvider`(已完整实现 `chat_stream` + 预设响应队列),避免改造 `cycle.rs` 测试模块内的内联 Stub。测试中调用 `submit_with_tools_stream` 时通过 `#[tokio::test(flavor = "multi_thread")]` 满足 spawn 运行时要求,或在单元级将 `run_tool_loop` 作为独立函数直接测试(不走 spawn)。
|
||||
|
||||
| # | 测试场景 | Mock 响应序列 | 验证点 | 覆盖路径 |
|
||||
|---|---------|---------------|--------|---------|
|
||||
| 3.1 | 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` | 正常路径:单轮 LLM → 文本返回 |
|
||||
| 3.2 | 单轮工具调用 | 2 个响应:tool_use + text | 包含一对 `ToolExecutionStarted`/`ToolExecutionCompleted`;最终 `stop_reason` 为 `Stop` | 正常路径:LLM → 工具 → LLM |
|
||||
| 3.3 | 多轮工具调用 | 4 个响应:3×tool_use + 1×text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度为 8(user + 3×(assistant+tool) + final assistant) | 正常路径:LLM → 工具 → LLM → 工具 → LLM |
|
||||
| 3.4 | 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `StreamEvent::Error`;消息历史停在第 2 轮 | 边界条件:超出上限 |
|
||||
| 3.5 | `chat_stream` 返回 Err | Mock `chat_stream` 返回 `Err(LlmError::Other(...))` | 流中第一个事件为 `StreamEvent::Error`;随后流结束 | 异常路径:LLM 不可用 |
|
||||
| 3.6 | 空 tool_registry | 1 个 text 响应,registry 中无工具 | 流退化为纯文本流,事件序列与 3.1 一致 | 退化场景:无工具可用 |
|
||||
| 3.7 | 不可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(不可恢复) | 流中出现 `StreamEvent::Error`;消息历史中不含该工具结果(循环终止前未 push) | 异常路径:工具执行失败 |
|
||||
| 3.8 | 可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(可恢复) | 工具结果作为 `ToolResult { is_error: true }` 回传 LLM;流正常结束,无 `Error` 事件 | 异常路径:工具出错但可恢复 |
|
||||
| 3.9 | 工具超时 | 2 个响应:tool_use → text,`tool_timeout_secs: 1`,模拟工具耗时 10 秒 | 流中出现 `StreamEvent::Error`;循环终止前未 push 工具结果 | 异常路径:工具执行超时 |
|
||||
|
||||
**验收条件**:
|
||||
- `cargo test` 新增 8 个测试全部通过
|
||||
- `cargo test` 存量测试 0 回归
|
||||
|
||||
### Step 4 — `AgentSession` 层包装
|
||||
|
||||
**工作量**:S(< 1h)
|
||||
**风险**:低(薄包装层,逻辑简单)
|
||||
**前置依赖**:S1
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 4.1 | 实现 `submit_turn_stream()`:触发 `OnTurnStart` → 组装 `LlmCycle` → 调用 `submit_with_tools_stream` → `turn_index += 1` → 返回流 | `agent/session.rs` | S1 | 低 |
|
||||
| 4.2 | 实现 `finalize_turn()`:`cost_so_far.add()` → 触发 `OnTurnEnd` hook | `agent/session.rs` | S1 | 低 |
|
||||
|
||||
**验收条件**:
|
||||
- `cargo build` 通过
|
||||
- 新增方法签名与方案设计一致
|
||||
- 与 `submit_turn` 的 system_prompt / compact_config / bundle 使用方式一致
|
||||
|
||||
### Step 5 — 集成测试 + 扫尾
|
||||
|
||||
**工作量**:S(< 1h)
|
||||
**风险**:低(基于现有测试框架)
|
||||
**前置依赖**:S3 + S4
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 5.1 | `submit_turn_stream` 端到端测试:跑通 mock provider → 消费流验证各事件到达 → `finalize_turn` 后 cost 更新正确 | `agent/session.rs`(内联测试) | S4 | 低 |
|
||||
| 5.2 | Hook 触发验证:`OnTurnStart` 在 `submit_turn_stream` 返回流之前触发;`finalize_turn` 调用后 `OnTurnEnd` 正确触发 | `agent/session.rs`(内联测试) | S4 | 低 |
|
||||
| 5.3 | `cargo test --all-targets` 全绿验证 | 全仓 | S5.1+S5.2 | 低 |
|
||||
| 5.4 | `cargo clippy --all-targets -- -D warnings` 0 警告 | 全仓 | S5.3 | 低 |
|
||||
| 5.5 | `cargo build --all-targets` 发布模式验证 | 全仓 | S5.4 | 低 |
|
||||
|
||||
**验收条件**:
|
||||
- 全量测试通过,存量 0 回归
|
||||
- clippy 0 警告
|
||||
- 发布模式零 warning
|
||||
|
||||
### 实施总览
|
||||
|
||||
| | Step 1 | Step 2 | Step 3 | Step 4 | Step 5 | **合计** |
|
||||
|--|--------|--------|--------|--------|--------|---------|
|
||||
| **工作量** | S | M | M | S | S | **M-L** |
|
||||
| **文件数** | 2 | 1 | 1(内联) | 1 | 1(内联) | **~4** |
|
||||
| **代码行** | ~20 | ~140 | ~150 含测试 | ~70 | ~80 含测试 | **~380** |
|
||||
| **风险** | 低 | 中 | 低 | 低 | 低 | 中 |
|
||||
| **并行** | — | 阻塞(S4 依赖 S2) | 阻塞 | 阻塞(依赖 S2) | 阻塞 | — |
|
||||
|
||||
---
|
||||
|
||||
## 附录 A:新增 StreamEvent 变体的 apply_to 语义
|
||||
|
||||
```rust
|
||||
// 在 PartialMessageResponse::apply_to 中追加:
|
||||
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => {
|
||||
// 元事件:不参与内容块累积,不修改 partial response 状态
|
||||
true
|
||||
}
|
||||
```
|
||||
|
||||
## 附录 B:CycleConfig 的 Clone 推导
|
||||
|
||||
```rust
|
||||
/// LLM 调用周期配置。
|
||||
#[derive(Debug, Clone)] // ← 追加 Clone
|
||||
pub struct CycleConfig {
|
||||
pub model: String,
|
||||
pub max_tokens: Option<u32>,
|
||||
pub temperature: Option<f32>,
|
||||
pub max_turns: Option<u32>,
|
||||
pub retry: RetryConfig, // 已 #[derive(Clone)]
|
||||
pub max_tool_turns: Option<u32>,
|
||||
pub tool_timeout_secs: u64,
|
||||
pub max_tool_result_bytes: usize,
|
||||
}
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,344 @@
|
||||
# 笔记:opencode 子代理调度、分发与合并及工作流推进
|
||||
|
||||
> 基于 `/Users/midnite/Samples/opencode` 源码调研,2026-07-04
|
||||
|
||||
---
|
||||
|
||||
## 一、整体架构
|
||||
|
||||
```
|
||||
LLM(主 Agent)
|
||||
│
|
||||
├── 调用 Task tool(tool call)
|
||||
│ ↓
|
||||
│ TaskTool.execute() ← packages/opencode/src/tool/task.ts
|
||||
│ │
|
||||
│ ├── agent.get() ← 查找 Agent 定义(agent.ts)
|
||||
│ ├── deriveSubagentPermission() ← 权限合并(subagent-permissions.ts)
|
||||
│ ├── sessions.create() ← 创建子 session
|
||||
│ │
|
||||
│ ├── [前台] background.wait() + background.waitForPromotion() race
|
||||
│ │ ↓ 完成
|
||||
│ │ renderOutput() → XML <task> 标签返回
|
||||
│ │
|
||||
│ └── [后台] background.start() → notify() 异步注入结果
|
||||
│
|
||||
└── 会话循环(runLoop) ← prompt.ts
|
||||
│
|
||||
├── 检测 subtask type part → handleSubtask()
|
||||
├── 检测 compaction → compaction.process()
|
||||
└── 正常流程 → LLM.stream() → processor.handleEvent()
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 二、子代理调度(Dispatch)
|
||||
|
||||
### 2.1 三种触发入口
|
||||
|
||||
| 入口 | 触发方式 | 调用链路 |
|
||||
|------|---------|---------|
|
||||
| A — LLM 自主 | LLM 调用 `task` tool | 系统提示词中注入了 Task tool 描述 + `describeTask()` 输出子代理列表 → LLM 决策 |
|
||||
| B — `subtask` part | 消息中有 `type: "subtask"` 的 part | `handleSubtask()` 直接执行 TaskTool,不走 LLM |
|
||||
| C — `agent` part | 消息中有 `type: "agent"` 的 part | 转为"调用 task tool 带 subagent: XXX"的提示词,引导 LLM |
|
||||
|
||||
### 2.2 TaskTool.execute() 完整流程(task.ts)
|
||||
|
||||
```
|
||||
execute(params, ctx):
|
||||
1. background 开关检查(需 experimental flag)
|
||||
2. ctx.ask() 权限询问
|
||||
3. agent.get(subagent_type) 查找子代理定义
|
||||
4. task_id 存在 → sessions.get(task_id) 恢复已有子 session
|
||||
task_id 不存在 → sessions.create() 创建新子 session
|
||||
5. deriveSubagentSessionPermission() 合并权限
|
||||
6. 添加默认 deny 规则(todowrite / task)
|
||||
7. 确定 model(继承或子代理自定义)
|
||||
8. 执行 runTask() → ops.resolvePromptParts() + ops.prompt()
|
||||
9. 结果格式化为 XML ← renderOutput()
|
||||
```
|
||||
|
||||
### 2.3 关键:子 session 创建(task.ts lines 121-158)
|
||||
|
||||
```typescript
|
||||
// 权限继承
|
||||
const childPermission = deriveSubagentSessionPermission({
|
||||
parentSessionPermission: parent.permission ?? [],
|
||||
subagent: next,
|
||||
})
|
||||
|
||||
// 默认 deny 规则
|
||||
const childToolDenies = [
|
||||
// 子代理自己的 permission 没允许 todowrite → 默认 deny
|
||||
...(next.permission.some(r => r.permission === "todowrite") ? []
|
||||
: [{ permission: "todowrite", pattern: "*", action: "deny" }]),
|
||||
// 子代理自己的 permission 没允许 task → 默认 deny(防嵌套)
|
||||
...(next.permission.some(r => r.permission === "task") ? []
|
||||
: [{ permission: "task", pattern: "*", action: "deny" }]),
|
||||
// 主 agent 专有工具也不给子代理
|
||||
...(cfg.experimental?.primary_tools?.map(p => ({ permission: p, ... })) ?? []),
|
||||
]
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 三、通信格式:Tool Call / Tool Result
|
||||
|
||||
### 3.1 父→子:Task tool 参数
|
||||
|
||||
```
|
||||
{
|
||||
subagent_type: "explore" | "general" | ...,
|
||||
description: "简短描述(3-5词)",
|
||||
prompt: "子代理的完整任务描述",
|
||||
task_id?: "恢复已有子 session 时使用",
|
||||
command?: "触发该调用的 CLI 命令(可选)",
|
||||
background?: true // 后台模式(需 experimental flag)
|
||||
}
|
||||
```
|
||||
|
||||
### 3.2 子→父:XML 包装的纯文本(renderOutput)
|
||||
|
||||
```xml
|
||||
<task id="ses_xxxxx" state="completed">
|
||||
<summary>任务简述</summary>
|
||||
<task_result>
|
||||
子 agent 输出的完整文本内容...
|
||||
</task_result>
|
||||
</task>
|
||||
```
|
||||
|
||||
错误时:
|
||||
|
||||
```xml
|
||||
<task id="ses_xxxxx" state="error">
|
||||
<summary>任务失败</summary>
|
||||
<task_error>
|
||||
Error: 具体错误信息...
|
||||
</task_error>
|
||||
</task>
|
||||
```
|
||||
|
||||
### 3.3 传递给 LLM 的方式
|
||||
|
||||
**前台模式**:
|
||||
```
|
||||
TaskTool.execute() 返回 { output: "<task>...</task>" }
|
||||
↓
|
||||
AI SDK 将其转为 tool result,存入数据库 tool part
|
||||
↓
|
||||
下一轮 LLM 调用时,tool result 作为消息历史的一部分传入
|
||||
↓
|
||||
LLM 看到 XML,自行解析使用
|
||||
```
|
||||
|
||||
**后台模式**:
|
||||
```
|
||||
TaskTool.execute() 立即返回 <task state="running">...
|
||||
↓
|
||||
子 agent 完成后 → background.wait() 触发 → inject()
|
||||
↓
|
||||
向父 session 注入合成 text part(synthetic: true)
|
||||
携带 <task state="completed">... 结果
|
||||
↓
|
||||
父 LLM 在下一轮循环中看到该消息
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 四、分发与合并(Distribution & Merge)
|
||||
|
||||
### 4.1 并行分发
|
||||
|
||||
- **无专用分发层**。依赖 LLM 在单条消息中发出多个 tool call
|
||||
- `task.txt` 引导 LLM:*"Launch multiple agents concurrently whenever possible"*
|
||||
- 底层通过 Effect.ts 的 `Effect.forkIn(scope, { startImmediately: true })` 实现同一消息内多 tool call 并发
|
||||
- **子 agent 之间完全隔离**,无直接通信
|
||||
|
||||
### 4.2 结果合并
|
||||
|
||||
**无专用合并逻辑。** 合并完全通过 LLM 的上下文理解完成:
|
||||
|
||||
- 前台:tool result 自然进入消息历史,LLM 下一轮读取
|
||||
- CLI 命令:额外注入 "Summarize the task tool output above and continue with your task." 引导 LLM 总结
|
||||
- LLM 自主调用:无额外引导,LLM 自行决定如何使用
|
||||
|
||||
### 4.3 前台/后台切换机制(task.ts lines 303-333)
|
||||
|
||||
```typescript
|
||||
// 前台执行
|
||||
return yield* Effect.raceFirst(
|
||||
background.wait({ id: nextSession.id }), // 等完成
|
||||
background.waitForPromotion(nextSession.id), // 等 promote 到后台
|
||||
)
|
||||
```
|
||||
|
||||
当用户将前台任务 promote 到后台时,`waitForPromotion` 先返回(标记 `metadata.background = true`),TaskTool 转而返回后台模式的输出。
|
||||
|
||||
### 4.4 后台作业引擎(core/background-job.ts)
|
||||
|
||||
纯内存、非持久化注册表。使用 Effect.ts 的 `SynchronizedRef` 做并发控制。
|
||||
|
||||
| 操作 | 行为 |
|
||||
|------|------|
|
||||
| `start()` | 创建 job,fork run effect,返回 info |
|
||||
| `extend()` | 追加顺序执行的 run(通过 `Deferred` 链式等待前一个完成) |
|
||||
| `wait()` | `Deferred.await(done)`,可选 timeout |
|
||||
| `waitForPromotion()` | 等待 `promoted` Deferred 或检测 `background` 标记 |
|
||||
| `promote()` | 标记 `background = true`,触发 `onPromote` callback |
|
||||
| `cancel()` | 设置 `cancelled`,close scope(中断所有子 fork) |
|
||||
|
||||
---
|
||||
|
||||
## 五、工作流推进(Workflow Progression)
|
||||
|
||||
### 5.1 核心循环(prompt.ts → runLoop)
|
||||
|
||||
```
|
||||
runLoop(sessionID):
|
||||
while true:
|
||||
1. MessageV2.filterCompactedEffect() 获取消息
|
||||
2. MessageV2.latest() 取最近 user/assistant/tasks
|
||||
3. 检查 finish 状态
|
||||
- 不是 tool-calls 且有 finish → break(退出循环)
|
||||
4. 取 tasks(subtask / compaction 队列)
|
||||
- subtask → handleSubtask() → continue
|
||||
- compaction → compaction.process() → continue/break
|
||||
5. 检查 overflow → 自动创建 compaction task → continue
|
||||
6. 构建 assistant message
|
||||
7. SessionProcessor.create() 创建 handle
|
||||
8. SessionTools.resolve() 解析所有工具
|
||||
9. 构建 system prompt(环境信息 + skills + MCP + instructions)
|
||||
10. handle.process() — 启动 LLM stream
|
||||
11. 检查 result:
|
||||
- "compact" → 返回给外层触发 compaction
|
||||
- "stop" → break
|
||||
- "continue" → 继续循环
|
||||
```
|
||||
|
||||
### 5.2 SessionProcessor 事件处理(processor.ts)
|
||||
|
||||
| Stream 事件 | 处理逻辑 |
|
||||
|------------|---------|
|
||||
| `reasoning-start/delta/end` | 创建 reasoning part → 增量追加 → 最终持久化 |
|
||||
| `tool-input-start/delta/end` | 创建/更新 tool part(pending 状态) |
|
||||
| `tool-call` | 标记 running → 设置 input → **doom loop 检测** |
|
||||
| `tool-result` | `completeToolCall()` → 持久化结果 + 附件 |
|
||||
| `tool-error` | `failToolCall()` → 标记错误 |
|
||||
| `provider-error` | 抛出异常 → 触发重试 |
|
||||
| `text-start/delta/end` | 流式文本 → `updatePartDelta()` **增量持久化** |
|
||||
| `step-start` | 创建快照(snapshot) |
|
||||
| `step-finish` | 生成 patch diff → 更新 usage/tokens → **overflow 检测** → 触发 summary |
|
||||
| `finish` | stream 结束 |
|
||||
|
||||
### 5.3 Doom Loop 检测(processor.ts lines 351-377)
|
||||
|
||||
连续 3 次**完全相同的 tool call**(相同名称 + 相同输入)触发权限询问:
|
||||
|
||||
```typescript
|
||||
const recentParts = parts.slice(-DOOM_LOOP_THRESHOLD) // DOOM_LOOP_THRESHOLD = 3
|
||||
if (recentParts.length === DOOM_LOOP_THRESHOLD &&
|
||||
recentParts.every(part =>
|
||||
part.type === "tool" &&
|
||||
part.tool === value.name &&
|
||||
part.state.status !== "pending" &&
|
||||
JSON.stringify(part.state.input) === JSON.stringify(input)
|
||||
)) {
|
||||
yield* permission.ask({ permission: "doom_loop", ... })
|
||||
}
|
||||
```
|
||||
|
||||
### 5.4 Compaction 工作流
|
||||
|
||||
两种触发方式:
|
||||
|
||||
| 触发条件 | 行为 |
|
||||
|---------|------|
|
||||
| step-finish 检测到 `isOverflow()` + `auto: true` | 创建 compaction task → 下一轮循环执行 → 压缩后 continue |
|
||||
| step-finish 检测到 `isOverflow()` + `auto: false` | 标记 `assistantMessage.error` → idle 等待用户干预 |
|
||||
|
||||
Compaction 使用专门的 `compaction` agent(hidden, mode=primary, `*=deny`)执行。
|
||||
压缩后的消息标记 `compacted: true`,后续通过 `MessageV2.filterCompactedEffect()` 过滤。
|
||||
|
||||
### 5.5 重试机制(processor.ts lines 658-672)
|
||||
|
||||
```typescript
|
||||
Effect.retry(
|
||||
SessionRetry.policy({
|
||||
provider: input.model.providerID,
|
||||
parse, // 错误解析(区分可重试/不可重试)
|
||||
set: (info) => status.set(sessionID, { type: "retry", ... }),
|
||||
}),
|
||||
)
|
||||
```
|
||||
|
||||
遇 provider 错误自动重试,LLM stream 完成后 `Effect.ensuring(cleanup)` 保证资源释放。
|
||||
|
||||
---
|
||||
|
||||
## 六、六种内置 Agent
|
||||
|
||||
| 名称 | Mode | Hidden | 用途 | 核心权限特征 |
|
||||
|------|------|--------|------|-------------|
|
||||
| `build` | primary | 否 | 默认 agent,全部工具 | question/plan_enter=allow |
|
||||
| `plan` | primary | 否 | 计划模式,禁用编辑 | edit=deny(除 plans), task(general)=deny |
|
||||
| `general` | subagent | 否 | 通用子代理 | todowrite=deny(默认禁止改 todo) |
|
||||
| `explore` | subagent | 否 | 只读代码探索 | `*=deny`,仅 read/grep/glob/bash/webfetch/websearch |
|
||||
| `compaction` | primary | 是 | 会话压缩(自动) | `*=deny` |
|
||||
| `title` | primary | 是 | 生成会话标题 | `*=deny`(step=1 时异步 fork) |
|
||||
| `summary` | primary | 是 | 生成消息摘要 | `*=deny`(每个 step-finish 时异步 fork) |
|
||||
|
||||
用户可通过 `config.agent` 自定义 agent(支持 `mode: "all"`),也可通过 `agent.generate` 让 LLM 辅助生成。
|
||||
|
||||
---
|
||||
|
||||
## 七、权限模型总结
|
||||
|
||||
```
|
||||
父 session permission
|
||||
│
|
||||
├── 仅继承 deny 规则 + external_directory 规则 ← subagent-permissions.ts
|
||||
│ (父 agent 的 allow 规则不传播到子代理)
|
||||
│
|
||||
├── 子代理自身 permission(来自 agent 定义)
|
||||
│
|
||||
├── 默认 deny:
|
||||
│ - todowrite(除非子代理明确允许)
|
||||
│ - task(除非子代理明确允许,默认防嵌套)
|
||||
│
|
||||
└── 主 agent 专有工具 deny(来自 config.experimental.primary_tools)
|
||||
```
|
||||
|
||||
子代理的 session 权限 = **父 deny + 父 external_directory + 自身 permission - 默认 deny - primary_tools deny**。
|
||||
|
||||
---
|
||||
|
||||
## 八、关键设计决策
|
||||
|
||||
| 决策 | 意图 | 效果/局限 |
|
||||
|------|------|----------|
|
||||
| 结果以 XML 纯文本嵌入上下文 | 简单、LLM 可直接理解 | LLM 自行解析 XML;大结果可能被截断 |
|
||||
| 无专用 merge 逻辑 | 简洁,不引入额外抽象 | 依赖 LLM 的理解能力处理返回结果 |
|
||||
| 默认禁止子代理嵌套 task | 防止无限递归 | 限制了多级分解场景 |
|
||||
| 同一消息多 tool call 并发 | 利用 LLM 并行能力 | 子 agent 隔离,无法协作 |
|
||||
| Effect.ts 贯穿全程 | 类型安全、结构化并发 | 学习曲线陡峭 |
|
||||
| session 作为隔离边界 | 天然权限/消息隔离 | 每个子 session 独立数据库记录,开销较大 |
|
||||
| 后台引擎纯内存 | 有意识取舍(注释说明) | 进程重启丢失状态 |
|
||||
|
||||
---
|
||||
|
||||
## 九、参考源码路径
|
||||
|
||||
| 文件 | 角色 |
|
||||
|------|------|
|
||||
| `packages/opencode/src/tool/task.ts` | Task tool 核心实现(调度入口) |
|
||||
| `packages/opencode/src/tool/task.txt` | Task tool 的 LLM 使用说明 |
|
||||
| `packages/opencode/src/agent/agent.ts` | Agent 定义注册中心 |
|
||||
| `packages/opencode/src/agent/subagent-permissions.ts` | 子代理权限推导 |
|
||||
| `packages/opencode/src/tool/registry.ts` | 工具注册 + `describeTask()` 列出可用子代理 |
|
||||
| `packages/opencode/src/session/prompt.ts` | 会话循环 + `handleSubtask()` + 提示词构建 |
|
||||
| `packages/opencode/src/session/processor.ts` | LLM stream 事件处理器 |
|
||||
| `packages/opencode/src/session/tools.ts` | Tool ↔ AI SDK 桥接 |
|
||||
| `packages/opencode/src/session/system.ts` | 系统提示词生成(含 Task tool 说明) |
|
||||
| `packages/opencode/src/background/job.ts` | 后台作业包装层 |
|
||||
| `packages/core/src/background-job.ts` | 后台作业核心引擎(内存注册表) |
|
||||
+375
-63
@@ -1,13 +1,13 @@
|
||||
# AG Core Roadmap
|
||||
|
||||
> 定稿日期:2026-05-11
|
||||
> 最后更新:2026-07-04(v0.1 发布完成)
|
||||
> 最后更新:2026-07-07
|
||||
|
||||
## 愿景
|
||||
|
||||
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
|
||||
|
||||
**当前状态**:Phase 0-4c 全部完成;Provider IR 重构(统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen Provider)已完成;LlmCycle 简化(IR 消息类型切换 + 桥接层移除)已完成;v0.1 发布就绪(**182 个测试通过、0 clippy 警告、7 个离线示例可运行**)。
|
||||
**当前状态**:v0.1.0 已发布(2026-07-04)。Phase 0-10 全部完成,v0.2.0-rc.1 已打标签。Provider IR 重构 + LlmCycle 简化 + 11 个离线示例(含 `quick_start` 30 行最小示例、`end_to_end` 完整集成示例、`context_slot_demo` 分支对话示例)+ SqliteStore 持久化 + 14 个公开枚举 `#[non_exhaustive]` 护栏 + `StepStatus` IR 迁移 + `submit_turn_stream` 流式体验 + ContextSlot 多上下文分区管理已交付。下一步进入 Phase 11(测试与检索补强)。
|
||||
|
||||
---
|
||||
|
||||
@@ -240,95 +240,400 @@ graph BT
|
||||
|
||||
---
|
||||
|
||||
## 扩展计划(v0.2+)
|
||||
## v0.2.0 — 生产就绪(Production-Ready Core)
|
||||
|
||||
> 以下功能在已完成的 phase 中已实现基础能力或在 Phase 4 阶段明确了边界,后续可按维度增量扩展。
|
||||
> 设计参考:见 `docs/note-agent-harness-references.md`(OpenClaw / Hermes / OpenHuman / OpenHarness 横向对比)。
|
||||
> OpenCode 借鉴:见 `docs/note-opencode-agent-switching.md`(Agent 切换 + System Prompt 拼接机制)。
|
||||
**目标**:解决 Rust Agent 工具箱从"能跑"到"能被人依赖"的鸿沟。持久化、配置层、上下文管理三大块补齐后,开发者可在 30 分钟内写出生产可用的 Agent 服务。
|
||||
|
||||
### 已有扩展项(沿用)
|
||||
**总体规模**:8 个增量 Phase(Phase 5-12),17 个可验证 Step。
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| Prompt Optimizer | `prompt` | 提示词自动优化 | P3 | 待实现 |
|
||||
| 流式接口优化 | `llm/stream` | 流式响应解析与事件化 | P0 | ✅ 已完成基础实现 |
|
||||
### 功能清单
|
||||
|
||||
### v0.2+ 新增扩展项
|
||||
#### P0 — 必须交付
|
||||
|
||||
> 以下为基于 Phase 4 设计讨论确定的 v0.2+ 候选扩展方向,按维度分组。
|
||||
> 标注为"v0.2 待评估"表示在 Phase 4 完成后再决定是否启动。
|
||||
| # | 功能 | 模块 | 方案要点 |
|
||||
|---|------|------|---------|
|
||||
| 1 | SqliteStore | `memory` | `rusqlite` + `bundled` feature,`MemoryStore` 的 SQLite 实现,进程重启数据不丢 |
|
||||
| 2 | ProviderConfig 扩展 + `from_env()` | `llm` | 补全 `timeout_secs` / `max_retries` 字段;`AG_LLM_*` 环境变量辅助函数 |
|
||||
| 3 | ToolDefinition IR 正式化 | `tools` | 移除 deprecated OpenAI wire 格式,替换为自定义 `ToolDef` 结构体 |
|
||||
| 4 | API 稳定性管理 | `*` | 公开枚举加 `#[non_exhaustive]`;CHANGELOG 记录 Breaking Changes;废弃 API 用 `#[deprecated]` 标记 |
|
||||
| 5 | Quick Start + 端到端示例 | `examples/` | 30 行 `main.rs` 快速开始;一个"SQLite 持久化 + Provider + 工具调用 + 多轮对话"的可运行示例(`cargo run --example`) |
|
||||
|
||||
#### Multi-Agent / 协同
|
||||
#### P1 — 重要但不阻塞
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| Multi-Agent 协同(Swarm) | `agent` | 子 Agent 委派、并行子任务、结果聚合 | P2 | v0.2 待评估 |
|
||||
| # | 功能 | 模块 | 方案要点 |
|
||||
|---|------|------|---------|
|
||||
| 6 | Ollama Provider | `llm/provider` | OpenAI Compat,本地 LLM 支持,实现量极小 |
|
||||
| 7 | VectorRetriever trait | `memory` | 语义检索 trait 抽象(`index` / `search`),不绑定后端实现 |
|
||||
| 8 | 流式 `submit_turn_stream` | `agent` | `AgentSession` 新增 `submit_turn_stream()`,返回 `Stream<Item = StreamEvent>` |
|
||||
| 9 | 测试补强 | `*` | wiremock Provider roundtrip 测试;多线程并发写入 MemoryStore 测试 |
|
||||
|
||||
#### 技能(Skills)
|
||||
#### P2 — 有时间再做
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| Markdown 技能按需加载 | `agent` / `prompt` | 兼容 `SKILL.md` 格式(Hermes / OpenHarness 风格),按 prompt 上下文动态加载 | P2 | v0.2 待评估 |
|
||||
| # | 功能 | 模块 | 备注 |
|
||||
|---|------|------|------|
|
||||
| 10 | MCP StreamableHttp | `tools` | 当前仅预留枚举变体 |
|
||||
| 11 | Gemini Provider | `llm/provider` | 协议差异大,实现成本较高 |
|
||||
| 12 | 文件系统 MemoryStore 后端 | `memory` | JSON/JSONL 轻量持久化 |
|
||||
|
||||
#### 记忆(Memory)
|
||||
### ContextSlot 上下文管理
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| 多通道检索(hybrid) | `memory/retriever` | 在 TextOverlap 之上叠加向量检索通道 | P2 | v0.2 待评估 |
|
||||
| KnowledgeGraph 深度记忆 | `memory` | 实体-关系图、`note-knowledge-graph-design.md` 已记录设计 | P3 | v0.2 待评估 |
|
||||
| TokenJuice 智能压缩 | `memory` / `llm/compact` | 借鉴 OpenHuman TokenJuice,对工具结果做语义压缩而非字节截断 | P3 | v0.2 待评估 |
|
||||
**模块归属**:`src/llm/context.rs`(与 `compact.rs` 同级)
|
||||
|
||||
#### 交互层(TUI / Gateway)
|
||||
**核心概念**:`ContextSlot` 是一段带策略配置的消息列表,以 `slot_id` 为 namespace 独立持久化到 `MemoryStore`。支持三种模式、三种来源和派生关联(记录 `parent_id`)。
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| TUI / 多平台 Gateway | 应用层 | OpenClaw / Hermes 风格的消息平台桥接(Feishu / Telegram / Discord 等) | P3 | v0.2+ 应用层 |
|
||||
**核心类型**:
|
||||
|
||||
#### 训练基础设施
|
||||
```rust
|
||||
pub struct ContextSlot { id, session_id, config, messages, store }
|
||||
pub struct SlotConfig { mode: SlotMode, source: SlotSource, budget, compact }
|
||||
pub enum SlotMode {
|
||||
Full, // 完整对话历史
|
||||
Focused(FocusedConfig), // 聚焦:保持 LLM 注意力
|
||||
Readonly, // 只读参考上下文
|
||||
}
|
||||
pub struct FocusedConfig { keep_system, recent_turns, inject_summary }
|
||||
pub enum SlotSource {
|
||||
New, // 全新空槽,独立持久化
|
||||
Derived { parent_id, strategy: DeriveStrategy }, // 从父 slot 派生
|
||||
Static(Vec<Message>), // 预置消息,不持久化
|
||||
}
|
||||
pub enum DeriveStrategy { Full, Focused(FocusedConfig) }
|
||||
pub struct ContextBudget { system, history, tools, tool_results, reserve }
|
||||
```
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| RL 轨迹导出 | `agent` | ShareGPT 格式轨迹、Atropos 集成(Hermes 风格) | P3 | v0.3+ 探索 |
|
||||
**持久化 Key 命名**:
|
||||
- `slot_msg:{session_id}:{slot_id}:{index}` → 消息内容
|
||||
- `slot_meta:{session_id}:{slot_id}` → `SlotMeta`(含 `parent_id`)
|
||||
- `slot_rel:{session_id}:{child_id}:parent` → `"{parent_id}"`
|
||||
|
||||
#### 安全治理
|
||||
**`AgentSession` 扩展**:
|
||||
- `create_slot(id, config)` — 创建新 slot
|
||||
- `switch_slot(id)` — 切换当前 slot
|
||||
- `list_slots()` — 列出所有 slot
|
||||
- `derive_slot(id, parent_id, strategy)` — 从父 slot 派生
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| Human-in-the-loop 审批 | `agent` / `tools/permission` | 高危工具执行前的异步审批回调(OpenHarness `permission_prompt` 模式) | P2 | v0.2 待评估 |
|
||||
**与 `ConversationMemory` 的关系**:保留不废除。`ConversationMemory` 继续服务传统对话场景。
|
||||
|
||||
#### 流式 / 实时
|
||||
**v0.2 不做**:
|
||||
- ❌ `slot.fork()` / `merge()` — 分支方法推迟到 v0.3+
|
||||
- ❌ `inject_summary` 自动生成 — v0.2 仅消费端(从 `SessionMemory` 读取),生成在 v0.3+
|
||||
- ❌ 血缘关系图遍历 — 只存 `parent_id`,不做查询
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| 流式 `submit_turn` | `agent/session` | Phase 4 v1 只暴露非流式 `submit_turn()`;v0.2 包装 `LlmCycle::submit_stream` 暴露流式入口 | P2 | v0.2 待评估 |
|
||||
**依赖**:Phase 0(MemoryStore trait)、Phase 3(MemoryStore 持久化)
|
||||
**优先级**:P1
|
||||
|
||||
#### Agent 切换 / Prompt 动态(OpenCode 借鉴)
|
||||
---
|
||||
|
||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
||||
|-------|---------|------|--------|------|
|
||||
| Agent 身份切换(角色轮换) | `agent` | 借鉴 OpenCode Tab 键切换 build/plan:同一 `AgentSession` 持有可热替换的 `Agent` 引用,切换时不重置消息历史,在末尾追加 `synthetic: true` 的状态变更消息。详见 `docs/note-opencode-agent-switching.md` §4 | P2 | v0.2 待评估 |
|
||||
| System Prompt 多层动态拼接 | `agent/session` | 借鉴 OpenCode `request.ts:58-66`:拆分 `base_prompt + agent_prompt + env_context` 三层,`AgentSession::submit_turn` 每轮重算(不缓存),便于按 agent 类型动态切换 | P2 | v0.2 待评估 |
|
||||
| **多 Context 切换** | `agent` | **Phase 4c 的 SessionMemory 数据结构已预留信息桥接通道,v0.2+ 在其上包装 `ContextManager` 实现完整的多 context 切换:创建/销毁/切换 context、通过 SessionMemory 桥接关键信息。详见 `docs/note-context-switch-design.md`** | P2 | v0.2 待评估 |
|
||||
### v0.2.0 实施计划 — 8 个增量 Phase
|
||||
|
||||
> **编号说明**:Phase 5-12 接续 v0.1 的 Phase 0-4c,按开发顺序排列。
|
||||
|
||||
#### Phase 5: 热身准备(Warmup)
|
||||
|
||||
**目标**:快速交付三个互不依赖的独立改动,建立交付节奏。
|
||||
|
||||
| Step | 内容 | 文件范围 | 验证标准 |
|
||||
|------|------|---------|---------|
|
||||
| **5.1** ✅ | `ProviderConfig` 扩展:补 `timeout_secs`(def=30) + `max_retries`(def=3);新增 `ProviderConfig::from_env(prefix)` | `llm/provider.rs` + 各 Provider `new()` 构造函数 | `cargo test` + `from_env()` 单元测试 |
|
||||
| **5.2** ✅ | `OllamaProvider`:基于 `GenericOpenaiProvider` 包装,改 base_url 为 `http://localhost:11434`;`ProviderType` 新增 `Ollama` | `llm/provider/provider.rs` + `llm/provider/ollama.rs`(新增) | `cargo build` — 纯类型级验证 |
|
||||
| **5.3** ✅ | 公开枚举 `#[non_exhaustive]` 前置标记:`ProviderType` / `StopReason` / `FinishReason` / `EvictionPolicy` / `SlotMode`(预置) | 各枚举定义处 | 编译通过 + `cargo clippy` 0 警告 |
|
||||
|
||||
**实际新增**(2026-07-05 commit `98dfe6c`):
|
||||
- 新增文件 1 个(`llm/provider/ollama.rs`,72 行)
|
||||
- 修改文件 2 个(`llm/provider.rs` 加 `from_env` + `Default` + 4 个字段;`memory/store.rs` EvictionPolicy 加 `#[non_exhaustive]`)
|
||||
- `ProviderType::Ollama` 变体 + `FromStr` 解析("ollama" → Ollama)
|
||||
- `OllamaProvider::new(base_url, api_key, model, timeout_secs)` + `with_client()` 构造函数
|
||||
- `ProviderConfig::from_env(prefix)` 解析 `{prefix}_API_KEY` / `{prefix}_BASE_URL` / `{prefix}_MODEL` 环境变量
|
||||
- 全量测试 182 → 190(+8,phase 5 新增 from_env 与 Ollama 相关单测)
|
||||
- clippy 0 警告
|
||||
|
||||
**依赖**:无(三个 Step 互不冲突)
|
||||
**优先级**:P0(5.1)+ P1(5.2)+ P0 前置(5.3)
|
||||
**为何独立成 Phase**:三个改动零文件重叠,可以并行推进。它们是后续所有 Phase 的"门把手"——先做完热身再进入核心工作。
|
||||
**状态**:✅ Phase 5 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### Phase 6: ToolDefinition IR 正式化
|
||||
|
||||
**目标**:引入 `ToolDef` 新类型,替换已标记 `#[deprecated]` 的 `ToolDefinition`(`OpenaiToolDefinition` 别名)。
|
||||
|
||||
**这是 v0.2 技术风险最高的 Phase**,影响 4 个模块约 8 个文件。通过 5 个 Step 逐文件切割确保每步可编译。
|
||||
|
||||
| Step | 内容 | 验证标准 |
|
||||
|------|------|---------|
|
||||
| **6.1** ✅ | `types/tool.rs` 新增 `ToolDef` 结构体 + `From<ToolDef> for OpenaiToolDefinition` + 反向 `From` | 单元测试 roundtrip |
|
||||
| **6.2** ✅ | `types/mod.rs` 切别名 `pub type ToolDefinition = ToolDef`;`MessageRequest.tools` 改 `Vec<ToolDef>` | `cargo build` 编译断点 |
|
||||
| **6.3** ✅ | `cycle.rs` 4 个方法签名 + `registry.rs` `definitions()` 签名更新 | `cargo build` |
|
||||
| **6.4** ✅ | Provider 适配层(openai.rs / anthropic.rs / openai_compat.rs):`build_request()` 内做 `ToolDef → wire-format` 转换 | `cargo test` 每个 provider 测试 |
|
||||
| **6.5** ✅ | 所有测试/示例中 `ToolDefinition` → `ToolDef` 修复;移除旧 `#[deprecated]` alias | `cargo test --all-targets` 全绿 |
|
||||
|
||||
**边界切割技巧**:
|
||||
- Step 6.1 → 6.2 之间是安全 checkpoint:新类型存在但旧代码照常编译
|
||||
- Provider 层不改序列化逻辑,只加一层 `From` 转换
|
||||
- 当前代码中 `ToolDefinition` 已是 `#[deprecated(since = "0.1.0")]`,用户已有迁移预期
|
||||
|
||||
**依赖**:无(仅与 Phase 5.3 有枚举兼容关系)
|
||||
**优先级**:P0
|
||||
|
||||
**实际新增**(2026-07-05 commit `4cf5918` / `9da9b83` / `b187519`,详见 `docs/13-phase6-tooldef-ir.md`):
|
||||
- 修改文件 8 个:`llm/types/tool.rs`、`llm/types/mod.rs`、`llm/types/request_v2.rs`、`llm/cycle.rs`、`llm/provider/openai.rs`、`tools/registry.rs`、`tools/mcp.rs`、`agent/agent.rs`
|
||||
- `ToolDef` IR(name / description / parameters,无 `strict`)新增于 `types/tool.rs`,配套双向 `From` 转换
|
||||
- `OpenaiToolDefinition` 降级为 `#[doc(hidden)]`,仅供 OpenAI 适配层内部消费
|
||||
- `MessageRequest.tools` 切换为 `Vec<ToolDef>`
|
||||
- `ToolDefinition` 别名最终完全移除(直接使用 `ToolDef`)
|
||||
- 4 处 `#[allow(deprecated)]` 抑制点全部清理(cycle/registry/mcp/agent);残留 `#[allow(deprecated)]` 均与 `ChatResponse` / `with_system_prompt` 等其他弃用项无关
|
||||
- 新增 roundtrip 测试 `message_request_with_tools_roundtrip`(断言 `strict` 字段不泄漏到序列化输出)
|
||||
- Anthropic 适配层字段名一致零改动;openai_compat/ollama 委托 `GenericOpenaiProvider` 零改动
|
||||
- 全量测试 190 → 191(+1,Phase 6 新增 roundtrip);clippy 0 警告
|
||||
|
||||
**状态**:✅ Phase 6 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### Phase 7: SqliteStore 持久化
|
||||
|
||||
**目标**:实现 `MemoryStore` 的 SQLite 后端,进程重启数据不丢。
|
||||
|
||||
**与 Phase 6 无耦合,可重叠开发。**
|
||||
|
||||
| Step | 内容 | 文件 | 验证标准 |
|
||||
|------|------|-----|---------|
|
||||
| **7.1** ✅ | 新增 `memory/store/sqlite.rs`:`Mutex<Connection>` + `spawn_blocking`,实现 `save/get/delete/list` + prefix 过滤 | `memory/store/sqlite.rs` + `Cargo.toml`(add `rusqlite`) | 单元测试 CRUD + prefix 查询 |
|
||||
| **7.2** ✅ | WAL 模式 + 并发安全 + 集成测试(`tokio::spawn` 10 个并发 task) | `sqlite.rs` 扩展 | 并发写入 100 轮无 race |
|
||||
|
||||
**设计决策**:
|
||||
- 用 `Mutex<Connection>` 而非连接池(ponytail:一个连接够用就不加 r2d2)
|
||||
- WAL 模式:`PRAGMA journal_mode=WAL` 解决读写锁
|
||||
|
||||
**依赖**:`MemoryStore` trait(v0.1 Phase 3 已就绪)
|
||||
**优先级**:P0
|
||||
|
||||
**实际新增**(2026-07-05 commit `13edacd` / `c8a91f6` / `c82af60`,详见 `docs/14-phase7-sqlite-store.md`):
|
||||
- 方案文档:`docs/14-phase7-sqlite-store.md`(526 行,Phase 7 设计推演与权衡记录)
|
||||
- 结构重组:`src/memory/store.rs` 单体文件 → `src/memory/store/{mod.rs(in_memory.rs, sqlite_store.rs)}` 模块目录;外部导入路径 `crate::memory::store::MemoryStore` 不变
|
||||
- 新增文件 2 个:`src/memory/store/sqlite_store.rs`(545 行 SqliteStore 实现 + 9 个内联测试)、`src/memory/store/in_memory.rs`(266 行,结构搬移)
|
||||
- 核心实现要点:
|
||||
- `Arc<Mutex<Connection>>` 串行化所有 IO;`spawn_blocking` 卸载到阻塞线程池
|
||||
- WAL 模式 + `synchronous=NORMAL` + `busy_timeout=5s` + `wal_autocheckpoint=1000`
|
||||
- `PRAGMA user_version` schema 版本管理(`INITIAL_USER_VERSION = 1`)
|
||||
- `created_at` 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
|
||||
- 错误精细映射:`SqliteFailure` / `InvalidQuery` → `InvalidInput`;`FromSqlConversionFailure` → `Serialization`;其他 → `Storage`
|
||||
- 9 个内联测试覆盖:CRUD、upsert、prefix / since / offset+limit 过滤、10 写者 × 10 次并发写入、持久化 round-trip(重启连接不丢数据)、`InMemoryStore ↔ SqliteStore` trait-box 互换兼容性
|
||||
- 依赖:`rusqlite = { version = "0.32", features = ["bundled"] }`;`time` 增补 `parsing` / `formatting` / `macros` features;`dev-dependencies` 新增 `tempfile = "3"`
|
||||
- 全量测试 191 → 200(+9,Phase 7 新增 SqliteStore 单测);clippy 0 警告
|
||||
|
||||
**状态**:✅ Phase 7 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### Phase 8: MVP 集成出口(v0.2.0-rc.1 候选)
|
||||
|
||||
**目标**:P0 五项全部交付。开发者 clone 仓库后 10 分钟跑起持久化 Agent。
|
||||
|
||||
| Step | 内容 | 验证标准 |
|
||||
|------|------|---------|
|
||||
| **8.1** ✅ | API 稳定性扫尾:`#[non_exhaustive]` × 14 公开枚举 + `StepStatus::Completed` 切 `MessageResponse` + CHANGELOG v0.2.0-rc.1 + Cargo.toml 0.2.0-rc.1 | `cargo doc --no-deps` 0 warning + 零 deprecated warning |
|
||||
| **8.2** ✅ | Quick Start 示例(57 行 `main.rs`):MockProvider + EchoTool + submit_turn 真实工具调用 | `cargo run --example quick_start` exit 0 |
|
||||
| **8.3** ✅ | 端到端示例:SqliteStore + AG_LLM_* from_env 自动检测 + 3 工具 + 3 轮对话 + 持久化跨连接验证 | `cargo run --example end_to_end`(Mock fallback,无需 API key)|
|
||||
|
||||
**Phase 8 全部完成**。**已打 `v0.2.0-rc.1` 标签**。
|
||||
|
||||
**实际新增**(2026-07-05,7 commits):
|
||||
- `feat(core)` —— 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他)
|
||||
- `refactor(agent)` —— `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` + `task_agent_demo.rs` 清理 3 处废弃类型
|
||||
- `docs` —— CHANGELOG v0.2.0-rc.1 条目 + Cargo.toml version 0.1.0 → 0.2.0-rc.1 + README 示例列表 7 → 10
|
||||
- `test(core)` —— 验证 commit 1-3 零回归(test 200 passed + clippy 0 警告 + doc 0 warning)
|
||||
- `feat(examples)` —— `quick_start.rs`(60 行)+ `end_to_end.rs`(246 行)
|
||||
- `docs(roadmap)` —— 标记 Phase 8 全部完成 + M4 里程碑 ✅
|
||||
- `fix(examples)` —— 实施后 PM/SA/Code Reviewer 三方审查发现 6 项问题(🔴 CalcTool 除零 panic + 🟡 drop 注释准确性 + 🟡 EchoTool 错误处理 + 💭 断言一致性 + 💭 工具两端语义统一 + 💭 trailing newline),全部修复
|
||||
|
||||
**依赖**:Phase 5(ProviderConfig from_env)+ Phase 6(ToolDef)+ Phase 7(SqliteStore)
|
||||
**优先级**:P0
|
||||
**状态**:✅ Phase 8 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### Phase 9: 流式体验增强
|
||||
|
||||
**目标**:Agent 会话支持流式输出,开发者看到实时 token。
|
||||
|
||||
| Step | 内容 | 文件 | 验证标准 |
|
||||
|------|------|-----|---------|
|
||||
| **9.1** ✅ | `AgentSession::submit_turn_stream(user_input) -> impl Stream<Item=StreamEvent>` | `agent/session.rs` | 单元测试验证流事件序列:`TextDelta → ... → MessageComplete` |
|
||||
|
||||
**注意**:tool 自动循环时流中插入 `ToolExecutionStarted` 事件,用户端 UI 显示"正在调用工具..."。
|
||||
|
||||
**依赖**:Phase 6(ToolDef)+ `LlmProvider.chat_stream`(v0.1 已有)
|
||||
**优先级**:P1
|
||||
|
||||
**实际新增**(2026-07-06 commit `212cfcc`,详见 `docs/16-phase9-streaming-experience.md`):
|
||||
- 方案文档:`docs/16-phase9-streaming-experience.md`(821 行,含状态机设计推演与边界情况)
|
||||
- 修改文件 3 个:`src/agent/session.rs`(+208,含 `submit_turn_stream` / `finalize_turn`)、`src/llm/cycle.rs`(+784,含 `submit_with_tools_stream` / `run_tool_loop` spawn + mpsc 状态机)、`src/llm/types/response_v2.rs`(+21,含 `StreamEvent::ToolExecutionStarted`/`Completed` 变体 + `apply_to` 元事件)
|
||||
- 关键设计:`CycleConfig` 加 `Clone` derive 以支持 spawn 跨 task;`finalize_turn` 手动同步状态(`submit_turn_stream` 返回流前不落库,避免半成品被 hook 误读)
|
||||
- 测试:新增 9 个单元测试 + 2 个集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock provider 流消费 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211(+11,0 失败)
|
||||
- clippy 0 警告
|
||||
- 无新增外部依赖
|
||||
|
||||
**状态**:✅ Phase 9 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### Phase 10: ContextSlot 上下文管理
|
||||
|
||||
**目标**:支持多上下文分区管理,Agent 可在不同 slot 之间切换。
|
||||
|
||||
| Step | 内容 | 验证标准 |
|
||||
|------|------|---------|
|
||||
| **10.1** ✅ | `src/agent/context.rs`:`ContextSlot` + `SlotConfig` / `SlotMode` / `FocusedConfig` / `SlotSource` / `DeriveStrategy` / `ContextBudget` / `SlotMeta` 核心类型 | `cargo build --all-targets` |
|
||||
| **10.2** ✅ | ContextSlot 持久化:基于 `MemoryStore` trait(不绑定 SqliteStore)实现 save/load/list/delete + slot 命名空间 key 策略 + `load_messages()` Focused 读时过滤 + `append_messages()` Readonly 阻断 + colon 注入防护 | 单元测试:持久化 roundtrip / session 隔离 / Focused 边界 / delete 保护 / 派生 / load_messages() |
|
||||
| **10.3** ✅ | `AgentSession` 扩展:`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot` + `new()` 自动创建 `"default"` slot + `submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回 + 新示例 `context_slot_demo` | 集成测试 + `cargo run --example context_slot_demo` exit 0 |
|
||||
|
||||
**如何保证简单场景无感**:`AgentSession::new()` 内部检查,自动创建 `"default"` slot → `submit_turn` 默认写到 default slot。
|
||||
|
||||
**实际新增**(2026-07-07 commit `6359422`,详见 `docs/17-phase10-contextslot.md`):
|
||||
- 方案文档:`docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复)
|
||||
- 新增文件 3 个:`src/agent/context.rs`(~430 行 ContextSlot 核心类型 + 持久化方法 + 22 个测试)、`src/agent/context.rs` 中的 `ContextSlot::filter_focused` 静态方法(被 `load_messages` 和 `derive_slot` 复用,消除代码重复)、`examples/context_slot_demo.rs`(~160 行分支对话示例:法律咨询 → 派生两个方向 → 切换 → 隔离验证 → 删除保护)
|
||||
- 修改文件 3 个:`src/agent.rs`(+5 行 module 声明 + re-export)、`src/agent/error.rs`(+56 行:3 个新变体 `SlotReadonly`/`SlotNotFound`/`SlotAlreadyExists` + 4 个测试)、`src/agent/session.rs`(+825/-197 行:slots 字段 + 6 个管理方法 + submit_turn/finalize_turn 改造 + 17 个测试)
|
||||
- 关键设计:
|
||||
- **模块归属**:`agent/context.rs`(零新依赖方向,遵循 `agent → memory` 已有依赖)
|
||||
- **持久化**:JSON blob 批次存储,每 slot 3-4 条 `MemoryItem`(`slot_data` / `slot_meta` / `slot_config` / `slot_rel`)
|
||||
- **submit_turn 签名不变**:方案 A(内部 `current_slot_id` 状态),向后兼容
|
||||
- **Focused 模式读时过滤**:`load_messages() -> Vec<Message>`,避免 Rust 借用检查问题
|
||||
- **增量追加写回**:`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式数据不丢失
|
||||
- **delete_slot 双重保护**:禁止删 `"default"` + 至少保留一个 slot
|
||||
- **colon 注入防护**:`assert_no_colon` 在 key 构造时 panic
|
||||
- **错误传播**:`serde_json` / `MemoryStore` 所有错误用 `?` 传播,无静默吞掉
|
||||
- 验证:211 → 254 测试(+43 新测试),clippy 0 警告,doc 0 warning,10 + 1 示例全部 exit 0
|
||||
- finalize_turn 签名变更(破坏性):新增 `new_messages_from_cycle: Vec<Message>` 参数,返回从 `()` 改为 `Result<(), AgentError>`——影响 Phase 9 的 `submit_turn_stream_triggers_turn_hooks` 和 `submit_turn_stream_end_to_end` 2 个测试,已适配
|
||||
|
||||
**依赖**:Phase 5(`#[non_exhaustive]` 预置 SlotMode 等枚举)、Phase 7(SqliteStore 推荐持久化后端;`MemoryStore` trait 即可)
|
||||
**优先级**:P1
|
||||
**状态**:✅ Phase 10 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### Phase 11: 测试与检索补强
|
||||
|
||||
**目标**:补全测试覆盖 + 语义检索抽象。
|
||||
|
||||
| Step | 内容 | 验证标准 |
|
||||
|------|------|---------|
|
||||
| **11.1** | `VectorRetriever` trait:`index(id, embeddings)` + `search(query, k)` | 编译 + mock 测试 |
|
||||
| **11.2** | wiremock Provider roundtrip 测试:模拟 OpenAI/Anthropic HTTP 端点 | `cargo test` 新增 10+ roundtrip 测试 |
|
||||
| **11.3** | 并发测试补强:InMemoryStore + SqliteStore 多线程写入验证 | 跑 100 轮无 race |
|
||||
|
||||
**依赖**:无(可随时做)
|
||||
**优先级**:P1
|
||||
|
||||
---
|
||||
|
||||
#### Phase 12: P2 锦上添花(可选)
|
||||
|
||||
**目标**:时间允许时按优先级交付。
|
||||
|
||||
| 优先级 | 功能 | 实现量估计 | 备注 |
|
||||
|--------|------|-----------|------|
|
||||
| **12.1** | 文件系统 MemoryStore(JSON/JSONL) | ~80 行 | 最简单,适合练手 |
|
||||
| **12.2** | MCP StreamableHttp 传输 | ~150 行 | 协议还在演进 |
|
||||
| **12.3** | Gemini Provider | ~300 行 | 协议差异大,建议推迟到 v0.3 |
|
||||
|
||||
**依赖**:无(独立交付)
|
||||
|
||||
---
|
||||
|
||||
### v0.2.0 Phase 依赖关系图
|
||||
|
||||
```mermaid
|
||||
graph BT
|
||||
P5["<b>Phase 5: 热身准备</b><br/>ProviderConfig::from_env<br/>Ollama Provider<br/>#[non_exhaustive] 标记"]:::done
|
||||
P6["<b>Phase 6: ToolDef IR</b><br/>Provider 无关工具定义"]:::done
|
||||
P7["<b>Phase 7: SqliteStore</b><br/>rusqlite + WAL<br/>9 个内联测试<br/>持久化 round-trip"]:::done
|
||||
P8["<b>Phase 8: MVP 出口</b><br/>rc.1 标签<br/>14 枚举 #[non_exhaustive]<br/>StepStatus IR 迁移<br/>quick_start + end_to_end"]:::done
|
||||
P9["<b>Phase 9: 流式体验增强</b><br/>submit_turn_stream<br/>submit_with_tools_stream<br/>9 单元测试 + 2 集成测试"]:::done
|
||||
P10["<b>Phase 10: ContextSlot</b><br/>ContextSlot 类型<br/>JSON blob 持久化<br/>AgentSession 集成<br/>43 个新测试"]:::done
|
||||
P11["Phase 11<br/>测试与检索"]:::p1
|
||||
P12["Phase 12<br/>P2 锦上添花"]:::p2
|
||||
|
||||
P8 --> P5
|
||||
P8 --> P6
|
||||
P8 --> P7
|
||||
|
||||
P9 --> P6
|
||||
|
||||
P10 --> P7
|
||||
P10 --> P8
|
||||
|
||||
P11 -.-> P7
|
||||
|
||||
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
|
||||
classDef warmup fill:#e2e8f0,stroke:#94a3b8
|
||||
classDef core fill:#fbbf24,stroke:#d97706
|
||||
classDef mvp fill:#4ade80,stroke:#16a34a
|
||||
classDef p1 fill:#93c5fd,stroke:#2563eb
|
||||
classDef p2 fill:#c4b5fd,stroke:#7c3aed
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 关键里程碑
|
||||
|
||||
| 里程碑 | Phase 完成条件 | 可验证指标 | 状态 |
|
||||
|--------|---------------|-----------|------|
|
||||
| **M1** | Phase 5 | 热身三项完成:`from_env()` 可用 / Ollama 类型存在 / `#[non_exhaustive]` 就位 | ✅ 2026-07-05 |
|
||||
| **M2** | Phase 6 | `ToolDef` 全量切换,`cargo test --all-targets` 全绿 | ✅ 2026-07-05 |
|
||||
| **M3** | Phase 7 | SqliteStore CRUD + 并发测试通过,进程重启数据不丢 | ✅ 2026-07-05 |
|
||||
| **M4** | **Phase 8 (rc.1)** | P0 五项全部交付,`cargo run --example quick_start` 跑通 | ✅ 2026-07-05 |
|
||||
| **M5** | Phase 9 | `submit_turn_stream` 流式事件序列验证通过 | ✅ 2026-07-06 |
|
||||
| **M6** | Phase 10 | ContextSlot 创建/切换/派生集成测试通过 | ✅ 2026-07-07 |
|
||||
| **M7** | Phase 11 | wiremock + 并发测试补强,测试总量 200+ | ⏳ |
|
||||
| **M8** | Phase 12(可选) | P2 功能按需交付 | ⏳ |
|
||||
|
||||
---
|
||||
|
||||
## v0.3+ 展望
|
||||
|
||||
### 已规划的功能
|
||||
|
||||
| 功能 | 说明 | 预计版本 |
|
||||
|------|------|---------|
|
||||
| ContextSlot 分支(fork/merge) | 在决策点 fork 出子上下文,分支独立演进,可合并/丢弃 | v0.3 |
|
||||
| 摘要自动生成 | Hook 驱动,`OnTurnEnd` 自动将对话摘要写入 `SessionMemory`,`inject_summary` 消费端已在 v0.2 就绪 | v0.3 |
|
||||
| 知识图谱 | 实体-关系图,`docs/note-knowledge-graph-design.md` 已记录设计 | v0.3+ |
|
||||
| Multi-Agent 协同(Swarm) | 子 Agent 委派、并行子任务、结果聚合 | v0.4+ |
|
||||
| 精确 tokenizer 计数 | 绑定具体模型的 tokenizer 计数,替代当前的字符估算 | v0.3+ |
|
||||
| 血缘关系图遍历 | 以 `parent_id` 为基础,提供 slot 血缘链查询 | v0.3+ |
|
||||
| Markdown 技能按需加载 | 兼容 `SKILL.md` 格式,按 prompt 上下文动态加载 | v0.3+ |
|
||||
| TokenJuice 语义压缩 | 对工具结果做语义压缩而非字节截断 | v0.3+ |
|
||||
| Human-in-the-loop 审批 | 高危工具执行前的异步审批回调 | v0.3+ |
|
||||
| RL 轨迹导出 | ShareGPT 格式轨迹、Atropos 集成 | v0.4+ |
|
||||
|
||||
### 明确不做(agcore 范围外)
|
||||
|
||||
| 功能 | 原因 |
|
||||
|------|------|
|
||||
| TUI / 多平台 Gateway | 应用层职责(Feishu / Telegram / Discord 桥接) |
|
||||
| 配置自动加载(config/figment) | 配置来源策略应由上游应用决定,agcore 不定义配置格式 |
|
||||
| 提示词自动优化 | 属于智能层,不应内建于 core 库 |
|
||||
|
||||
---
|
||||
|
||||
## 风险与建议
|
||||
|
||||
1. **Phase 0 已完成**:LLM 调用周期基础设施已全部实现,可以支撑后续模块开发
|
||||
2. **并行可能性**:Phase 0 和 Phase 1 可并行开展(无相互依赖),可加速早期交付
|
||||
3. **MCP 协议复杂性**:MCP 涉及协议握手、session 管理、长期连接,建议预留充足时间调研协议细节
|
||||
4. **Scope 蔓延风险**:当前 specs 只有 1 份文档,建议每个模块上线前都产出对应 spec,避免边实现边设计
|
||||
5. **Phase 4 抽象化边界**:AG Core 定位为"支持库"而非"Agent 产品",Phase 4(4a/4b/4c)需严格控制范围——只暴露 trait + 最小 reference impl,业务循环(多轮 turn 编排、对话记忆自动回写、Task 拆解策略)留给上层应用。`SessionMemory`(Phase 4c)提供信息桥接通道但不实现 context 切换逻辑。多 context 切换管理延后至 v0.2+。详细设计决策见 `docs/7-agent-runtime.md`
|
||||
6. **参考项目语言差异**:OpenClaw / Hermes / OpenHarness 均为 Python/TypeScript 实现,OpenHuman 虽是 Rust + Tauri 但定位是桌面应用。借鉴时**只取架构模式**,不照搬具体实现(如 Pydantic 工具校验、SQLite Memory Tree、Node+Python 双进程等)
|
||||
1. **持久化依赖**:`rusqlite` + `bundled` 零外部依赖编译,但 SQLite 不适配所有场景(分布式/高并发写)。`MemoryStore` trait 的抽象层允许下游自行实现 Redis / PostgreSQL 后端
|
||||
2. **ContextSlot 心智负担**:`ContextSlot` 引入了一等抽象的复杂度。建议通过 `AgentBuilder` 默认创建 `"default"` slot,让简单场景无感使用
|
||||
3. **向量检索生态**:`VectorRetriever` trait-only 不绑定实现,需社区贡献或用户自行适配 pgvector / qdrant / lancedb
|
||||
4. **Scope 蔓延**:agcore 定位为"支持库"而非"Agent 产品",始终以 trait + reference impl 为边界,业务循环留给上层
|
||||
5. **API 稳定性**:v0.2 引入 `#[non_exhaustive]` 和 `#[deprecated]` 机制,但不承诺 SemVer 稳定——仍在快速迭代期
|
||||
|
||||
---
|
||||
|
||||
## 下一步行动
|
||||
|
||||
1. **Phase 4c 已完成**:Phase 4a + 4b + 4c 已交付(116 测试通过,0 clippy 警告)。可启动 v0.2+ 扩展评估(如多 Context 切换、Multi-Agent 协同等)
|
||||
2. **Context 切换备忘**:`docs/note-context-switch-design.md` 记录了多 context 切换方案讨论,作为 v0.2+ 扩展项的输入
|
||||
3. **参考项目调研沉淀**:已完成 OpenClaw / Hermes / OpenHuman / OpenHarness 横向调研,结果沉淀至 `docs/note-agent-harness-references.md`,作为 v0.2+ 扩展项的输入
|
||||
4. **Phase 3 备用设计就绪**:`docs/note-knowledge-graph-design.md` 记录了 KnowledgeGraph、高级评分、RecallBased 淘汰等设计,v0.2+ 记忆扩展可直接参考
|
||||
1. **Phase 11 启动**:测试与检索补强(`VectorRetriever` trait + wiremock Provider roundtrip + 并发写入验证),P1 功能
|
||||
2. **示例先行**:每完成一个 Phase 立即更新对应示例,验证通过后再合入
|
||||
3. **里程碑追踪**:以 Phase 10(ContextSlot,已完成)为最新节点,逐 Phase 验收
|
||||
4. **v0.2.0 正式版**:Phase 8-11 全部完成后,去掉 rc 后缀打 `v0.2.0` 正式版
|
||||
|
||||
**已完成 / 进行中阶段**:
|
||||
- ✅ Phase 0 Foundation — 全部交付物已完成
|
||||
@@ -338,9 +643,16 @@ graph BT
|
||||
- ✅ Phase 4a Core Glue — 全部交付物已完成
|
||||
- ✅ Phase 4b Task Execution — 全部交付物已完成
|
||||
- ✅ Phase 4c Session Memory — 全部交付物已完成
|
||||
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen 适配(方案:`docs/10-llm-provider-refinement.md`、`docs/10a-phase0-types-and-trait.md`、`docs/10b-phase1-provider-adaptation.md`)
|
||||
- ✅ LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除(方案:`docs/10c-phase2-llm-cycle-simplify.md`)
|
||||
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、7 个离线示例、README + 错误消息友好化、Roadmap 同步、CHANGELOG 初始化(计划:`docs/11-v0.1-release-plan.md`)
|
||||
- ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy)
|
||||
- ✅ Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
|
||||
- ✅ Phase 7 SqliteStore — `rusqlite 0.32` + WAL 模式 + `Arc<Mutex<Connection>>` + `spawn_blocking`;`memory/store.rs` → `store/{in_memory,sqlite_store}.rs` 模块化;9 个内联测试覆盖 CRUD/upsert/过滤/10×10 并发/持久化 round-trip;`InMemoryStore ↔ SqliteStore` trait-box 互换兼容
|
||||
- ✅ **Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 60 行 + `end_to_end` 246 行),10 个离线示例全部 exit 0;**v0.2.0-rc.1 标签已打**;实施后三方审查发现 6 项问题(1 🔴 + 2 🟡 + 3 💭)已全部修复
|
||||
- ✅ **Phase 9 流式体验增强** — `AgentSession::submit_turn_stream` 流式事件序列 + `LlmCycle::submit_with_tools_stream` spawn + mpsc 状态机 + `StreamEvent::ToolExecutionStarted`/`Completed` 新变体 + 9 单元测试 + 2 集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock 验证 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211;`CycleConfig` 加 `Clone` derive;方案文档 `docs/16-phase9-streaming-experience.md`(821 行)
|
||||
- ✅ **Phase 10 ContextSlot 上下文管理** — `src/agent/context.rs` 新增 `ContextSlot` 核心类型(Full / Focused / Readonly 三种模式,New / Derived / Static 三种来源)+ JSON blob 批次持久化(每 slot 3-4 条 MemoryItem,`slot_config` key 自恢复支持旧版本兼容);`AgentSession` 扩展 slots 字段 + 5 个管理方法(`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot`,自动创建 `"default"` slot,`delete_slot` 双重保护禁止删 default/最后一个);`submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回(`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式"读时过滤"语义不丢失数据);`finalize_turn` 签名变更(新增 `new_messages_from_cycle: Vec<Message>` 参数,返回 `Result<(), AgentError>`);`agent/error.rs` 新增 3 个 Slot 错误变体(`SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`);`examples/context_slot_demo.rs` 新增分支对话示例(法律咨询入口 → 两个派生方向 → 切换 → 隔离验证 → 删除保护);方案文档 `docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复);全量 211 → 254(+43 新测试),clippy 0 警告,doc 0 warning,11 个离线示例全部 exit 0
|
||||
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen/Ollama 适配
|
||||
- ✅ LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除
|
||||
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、8 个离线示例(含 `simple_visit`)、README + 错误消息友好化、CHANGELOG 初始化
|
||||
- 📋 **v0.2 规划细化完成** — 8 个增量 Phase(Phase 5-12),17 个可验证 Step,覆盖 P0-P2 全部 12 项功能 + ContextSlot
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -15,9 +15,9 @@ use std::sync::Arc;
|
||||
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::types::Usage;
|
||||
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。
|
||||
@@ -72,7 +72,10 @@ async fn main() {
|
||||
|
||||
// 4. 提交第一轮
|
||||
println!("=== 提交第 1 轮 ===");
|
||||
let resp = session.submit_turn("1+1=?").await.expect("submit_turn 失败");
|
||||
let resp = session
|
||||
.submit_turn("1+1=?")
|
||||
.await
|
||||
.expect("submit_turn 失败");
|
||||
println!("LLM: {}", resp.text());
|
||||
session
|
||||
.set_session_data("last_q", "1+1=?")
|
||||
@@ -107,11 +110,7 @@ async fn main() {
|
||||
|
||||
// 8. 跨 session 数据隔离验证
|
||||
println!("=== 数据隔离验证 ===");
|
||||
let other = AgentSession::new(
|
||||
Arc::new(CalculatorAgent),
|
||||
"other-session",
|
||||
bundle,
|
||||
);
|
||||
let other = AgentSession::new(Arc::new(CalculatorAgent), "other-session", bundle);
|
||||
assert!(
|
||||
other.get_session_data("last_q").await.unwrap().is_none(),
|
||||
"新会话不应看到旧 session 的 last_q"
|
||||
@@ -127,4 +126,4 @@ async fn main() {
|
||||
);
|
||||
|
||||
println!("\n✓ agent_session_demo 完成");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
//! context_slot_demo —— 多上下文槽位管理示例。
|
||||
//!
|
||||
//! 场景:法律咨询入口 → 派生两个独立探索方向 → 切换 → 隔离验证 → 删除。
|
||||
//!
|
||||
//! 展示:
|
||||
//! - 默认 slot 自动创建
|
||||
//! - 多 slot 间的消息隔离
|
||||
//! - 派生 slot 从父 slot 复制消息
|
||||
//! - 删除非 default slot 后自动回退到 default
|
||||
//!
|
||||
//! 运行:`cargo run --example context_slot_demo`(离线,零配置)
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::provider::LlmProvider;
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use agcore::llm::types::Usage;
|
||||
use agcore::tools::ToolRegistry;
|
||||
|
||||
struct LegalAdvisor;
|
||||
|
||||
impl Agent for LegalAdvisor {
|
||||
fn name(&self) -> &str {
|
||||
"legal-advisor"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("你是法律顾问。请用一句话回答用户问题。")
|
||||
}
|
||||
}
|
||||
|
||||
/// 构造一个简单的 Assistant 响应(用于 MockProvider)。
|
||||
fn assistant_resp(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: "mock".into(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: Usage::from_input_output(5, 5),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
// 1. 构造 session(自动包含 default slot)
|
||||
let provider: Arc<dyn LlmProvider> = Arc::new(MockProvider::new(vec![
|
||||
assistant_resp("您好,我可以帮您处理法律问题。"),
|
||||
assistant_resp("管辖权问题:建议选择合同签订地法院。"),
|
||||
assistant_resp("条款修改:建议将上限调整为 80 万。"),
|
||||
assistant_resp("已回到主对话。"),
|
||||
]));
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
let mut session = AgentSession::new(Arc::new(LegalAdvisor), "legal-001", bundle);
|
||||
|
||||
println!("=== 1. 默认 slot 自动创建 ===");
|
||||
assert_eq!(session.current_slot_id(), "default");
|
||||
let slots: Vec<_> = session.list_slots().collect();
|
||||
println!("初始 slots: {slots:?}");
|
||||
assert_eq!(slots.len(), 1);
|
||||
assert!(slots.contains(&&"default".to_string()));
|
||||
|
||||
println!("\n=== 2. 在 default slot 中提交一轮 ===");
|
||||
let r1 = session.submit_turn("我需要法律援助").await.unwrap();
|
||||
println!("default slot response: {}", r1.text());
|
||||
|
||||
println!("\n=== 3. 派生两个独立探索方向的 slot ===");
|
||||
session
|
||||
.derive_slot("option_jurisdiction", "default", agcore::agent::DeriveStrategy::Full)
|
||||
.await
|
||||
.unwrap();
|
||||
session
|
||||
.derive_slot("option_amendment", "default", agcore::agent::DeriveStrategy::Full)
|
||||
.await
|
||||
.unwrap();
|
||||
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||
println!("派生后 slots: {slots:?}");
|
||||
assert_eq!(slots.len(), 3);
|
||||
|
||||
println!("\n=== 4. 切到 option_jurisdiction 并提交 ===");
|
||||
session.switch_slot("option_jurisdiction").await.unwrap();
|
||||
assert_eq!(session.current_slot_id(), "option_jurisdiction");
|
||||
let r2 = session.submit_turn("如果用户质疑管辖权?").await.unwrap();
|
||||
println!("option_jurisdiction response: {}", r2.text());
|
||||
|
||||
println!("\n=== 5. 切到 option_amendment 并提交 ===");
|
||||
session.switch_slot("option_amendment").await.unwrap();
|
||||
let r3 = session.submit_turn("用户要求提高赔偿上限?").await.unwrap();
|
||||
println!("option_amendment response: {}", r3.text());
|
||||
|
||||
println!("\n=== 6. 切回 default,验证消息隔离 ===");
|
||||
session.switch_slot("default").await.unwrap();
|
||||
let r4 = session.submit_turn("汇总一下我们的讨论").await.unwrap();
|
||||
println!("default response: {}", r4.text());
|
||||
// 验证 default slot 不包含 option_jurisdiction 的"管辖权"问题
|
||||
let (_, default_slot) = session.slots().find(|(id, _)| *id == "default").unwrap();
|
||||
let default_has_jurisdiction = default_slot
|
||||
.messages
|
||||
.iter()
|
||||
.any(|m| message_contains(m, "管辖权"));
|
||||
assert!(
|
||||
!default_has_jurisdiction,
|
||||
"default slot 不应包含 option_jurisdiction 的消息"
|
||||
);
|
||||
|
||||
println!("\n=== 7. 删除 option_amendment,验证回退到 default ===");
|
||||
session.delete_slot("option_amendment").await.unwrap();
|
||||
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||
println!("删除后 slots: {slots:?}");
|
||||
assert!(!slots.contains(&"option_amendment".to_string()));
|
||||
assert_eq!(slots.len(), 2);
|
||||
|
||||
println!("\n=== 8. 切到 option_jurisdiction 并删除,验证 current 回退 ===");
|
||||
session.switch_slot("option_jurisdiction").await.unwrap();
|
||||
session.delete_slot("option_jurisdiction").await.unwrap();
|
||||
assert_eq!(session.current_slot_id(), "default");
|
||||
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||
println!("删除后 slots: {slots:?}");
|
||||
assert_eq!(slots.len(), 1);
|
||||
assert_eq!(slots[0], "default");
|
||||
|
||||
println!("\n=== 9. 验证 delete_slot 保护逻辑 ===");
|
||||
let err = session.delete_slot("default").await.unwrap_err();
|
||||
println!("删除 default 返回错误: {err}");
|
||||
assert!(matches!(err, agcore::agent::AgentError::Config(_)));
|
||||
|
||||
println!("\n✓ context_slot_demo 完成");
|
||||
}
|
||||
|
||||
/// 检查 Message 是否包含指定文本(提取第一个 Text block)。
|
||||
fn message_contains(msg: &Message, needle: &str) -> bool {
|
||||
use agcore::llm::types::message::ContentBlock;
|
||||
let blocks = match msg {
|
||||
Message::System { content }
|
||||
| Message::User { content }
|
||||
| Message::Assistant { content } => content,
|
||||
Message::UserImage { .. } => return false,
|
||||
Message::ToolResult { content, .. } => content,
|
||||
_ => return false,
|
||||
};
|
||||
for block in blocks {
|
||||
if let ContentBlock::Text { text } = block
|
||||
&& text.contains(needle)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
@@ -29,6 +29,7 @@ fn message_text(msg: &Message) -> &str {
|
||||
.next()
|
||||
.unwrap_or(""),
|
||||
Message::UserImage { .. } => "[image]",
|
||||
_ => "",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -80,11 +81,8 @@ async fn main() {
|
||||
// 3. 多角色混合 + clear
|
||||
println!("\n=== 多角色写入 + clear ===");
|
||||
let store3 = Arc::new(InMemoryStore::new());
|
||||
let mut memory3 = ConversationMemory::new(
|
||||
store3,
|
||||
"session-3",
|
||||
ConversationMemoryConfig::default(),
|
||||
);
|
||||
let mut memory3 =
|
||||
ConversationMemory::new(store3, "session-3", ConversationMemoryConfig::default());
|
||||
memory3
|
||||
.add_message(Message::user_text("你好"))
|
||||
.await
|
||||
@@ -98,7 +96,9 @@ async fn main() {
|
||||
.await
|
||||
.unwrap();
|
||||
memory3
|
||||
.add_message(Message::assistant("我无法查询实时天气,但你可以查看天气应用。"))
|
||||
.add_message(Message::assistant(
|
||||
"我无法查询实时天气,但你可以查看天气应用。",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
println!(
|
||||
@@ -119,16 +119,8 @@ async fn main() {
|
||||
// 4. Session 隔离
|
||||
println!("\n=== Session 隔离(共用 InMemoryStore)===");
|
||||
let store4 = Arc::new(InMemoryStore::new());
|
||||
let mut a = ConversationMemory::new(
|
||||
store4.clone(),
|
||||
"s-a",
|
||||
ConversationMemoryConfig::default(),
|
||||
);
|
||||
let mut b = ConversationMemory::new(
|
||||
store4.clone(),
|
||||
"s-b",
|
||||
ConversationMemoryConfig::default(),
|
||||
);
|
||||
let mut a = ConversationMemory::new(store4.clone(), "s-a", ConversationMemoryConfig::default());
|
||||
let mut b = ConversationMemory::new(store4.clone(), "s-b", ConversationMemoryConfig::default());
|
||||
a.add_message(Message::user_text("A 的消息")).await.unwrap();
|
||||
b.add_message(Message::user_text("B 的消息")).await.unwrap();
|
||||
println!(
|
||||
@@ -140,4 +132,4 @@ async fn main() {
|
||||
assert_eq!(b.len(), 1);
|
||||
|
||||
println!("\n✓ conversation_memory_demo 完成");
|
||||
}
|
||||
}
|
||||
|
||||
+11
-16
@@ -17,7 +17,7 @@ use agcore::tools::{
|
||||
ToolRegistry,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{json, Value};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
/// 天气查询工具 —— 模拟根据城市返回天气数据。
|
||||
struct WeatherTool;
|
||||
@@ -42,11 +42,7 @@ impl BaseTool for WeatherTool {
|
||||
fn required_permissions(&self) -> Vec<Permission> {
|
||||
vec![Permission::Network]
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
args: Value,
|
||||
_ctx: &ToolContext<'_>,
|
||||
) -> Result<Value, ToolError> {
|
||||
async fn execute(&self, args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||
let city = args["city"].as_str().unwrap_or("未知");
|
||||
// 模拟查询:根据城市名给出不同温度
|
||||
let (temperature, condition) = match city {
|
||||
@@ -84,11 +80,7 @@ impl BaseTool for DeleteFileTool {
|
||||
fn required_permissions(&self) -> Vec<Permission> {
|
||||
vec![Permission::Delete]
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_args: Value,
|
||||
_ctx: &ToolContext<'_>,
|
||||
) -> Result<Value, ToolError> {
|
||||
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||
Ok(json!({"deleted": true}))
|
||||
}
|
||||
}
|
||||
@@ -138,9 +130,8 @@ async fn main() {
|
||||
|
||||
// 5. 权限检查:默认 PermissionConfig 黑名单含 Delete
|
||||
println!("\n=== 权限检查(默认 PermissionConfig,denied = [Delete, Shell])===");
|
||||
let mut registry_with_checker = ToolRegistry::new().with_permission_checker(PermissionChecker::new(
|
||||
PermissionConfig::default(),
|
||||
));
|
||||
let mut registry_with_checker = ToolRegistry::new()
|
||||
.with_permission_checker(PermissionChecker::new(PermissionConfig::default()));
|
||||
registry_with_checker
|
||||
.register(Arc::new(WeatherTool) as ToolRef)
|
||||
.unwrap();
|
||||
@@ -155,7 +146,11 @@ async fn main() {
|
||||
.unwrap();
|
||||
println!(
|
||||
"get_weather 权限检查: {}",
|
||||
if r.output.is_ok() { "通过 ✓" } else { "阻断 ✗" }
|
||||
if r.output.is_ok() {
|
||||
"通过 ✓"
|
||||
} else {
|
||||
"阻断 ✗"
|
||||
}
|
||||
);
|
||||
|
||||
// delete_file 声明 Delete → 在 denied 列表 → 阻断
|
||||
@@ -166,4 +161,4 @@ async fn main() {
|
||||
println!("delete_file 权限检查: 阻断 ✗ ({err})");
|
||||
|
||||
println!("\n✓ custom_tool 完成");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,247 @@
|
||||
//! end_to_end —— 3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证。
|
||||
//!
|
||||
//! 运行:`cargo run --example end_to_end`(离线,零配置)
|
||||
//!
|
||||
//! ## 真实 LLM Provider 切换
|
||||
//!
|
||||
//! 设置环境变量即可使用真实 LLM Provider:
|
||||
//! - `AG_LLM_BASE_URL` —— API 端点(如 `https://api.openai.com/v1`)
|
||||
//! - `AG_LLM_API_KEY` —— API key
|
||||
//! - `AG_LLM_MODEL` —— 模型名(如 `gpt-4o-mini`)
|
||||
//! - `AG_LLM_PROVIDER`(可选)—— Provider 类型,默认 OpenaiChat(OpenAI / DeepSeek / Qwen / Ollama)
|
||||
//!
|
||||
//! 未设置上述变量时自动降级为 MockProvider,零配置可运行。
|
||||
|
||||
use std::env;
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::provider::{create_provider, LlmProvider, ProviderConfig, ProviderType};
|
||||
use agcore::llm::types::{Usage, message::{ContentBlock, Message}, response_v2::{MessageResponse, StopReason}};
|
||||
use agcore::memory::store::{MemoryStore, SqliteStore};
|
||||
use agcore::memory::types::{MemoryFilter, MemoryItem};
|
||||
use agcore::tools::{BaseTool, ToolContext, ToolError, ToolRegistry};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{Value, json};
|
||||
use tempfile::TempDir;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
// === Agent ===
|
||||
|
||||
struct AssistantAgent;
|
||||
impl Agent for AssistantAgent {
|
||||
fn name(&self) -> &str { "end-to-end assistant" }
|
||||
fn system_prompt(&self) -> Option<&str> { Some("简洁助手,必要时调用工具完成任务。") }
|
||||
}
|
||||
|
||||
// === Tools ===
|
||||
|
||||
struct EchoTool;
|
||||
#[async_trait]
|
||||
impl BaseTool for EchoTool {
|
||||
fn name(&self) -> &str { "echo" }
|
||||
fn description(&self) -> &str { "回显输入文本" }
|
||||
fn parameters(&self) -> Value {
|
||||
json!({"type":"object","properties":{"text":{"type":"string"}},"required":["text"]})
|
||||
}
|
||||
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||
let text = args.get("text").and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ToolError::InvalidArguments("text".into(), "需要 string 类型的 text 参数".into()))?;
|
||||
Ok(json!({"echoed": format!("收到: {text}")}))
|
||||
}
|
||||
}
|
||||
|
||||
/// 四则运算:'a op b' 格式(ponytail: 基础 +-*/ 不引入 rhai 依赖)。
|
||||
struct CalcTool;
|
||||
#[async_trait]
|
||||
impl BaseTool for CalcTool {
|
||||
fn name(&self) -> &str { "calc" }
|
||||
fn description(&self) -> &str { "四则运算:'a op b' 格式,op ∈ {+, -, *, /}" }
|
||||
fn parameters(&self) -> Value {
|
||||
json!({"type":"object","properties":{"expr":{"type":"string"}},"required":["expr"]})
|
||||
}
|
||||
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||
let expr = args["expr"].as_str().unwrap_or("");
|
||||
let parts: Vec<&str> = expr.split_whitespace().collect();
|
||||
if parts.len() != 3 {
|
||||
return Err(ToolError::InvalidArguments("expr".into(), "需要 'a op b' 三段式".into()));
|
||||
}
|
||||
let a: i64 = parts[0].parse().map_err(|_| ToolError::InvalidArguments("expr".into(), format!("无法解析 '{}'", parts[0])))?;
|
||||
let b: i64 = parts[2].parse().map_err(|_| ToolError::InvalidArguments("expr".into(), format!("无法解析 '{}'", parts[2])))?;
|
||||
let result = match parts[1] {
|
||||
"+" => a + b,
|
||||
"-" => a - b,
|
||||
"*" => a * b,
|
||||
"/" => a.checked_div(b).ok_or_else(|| {
|
||||
ToolError::InvalidArguments("expr".into(), "除数不能为 0".into())
|
||||
})?,
|
||||
op => return Err(ToolError::InvalidArguments("expr".into(), format!("不支持的运算符: {op}"))),
|
||||
};
|
||||
Ok(json!({"result": result}))
|
||||
}
|
||||
}
|
||||
|
||||
/// 通过 MemoryStore trait 读写笔记:直接持有 Arc<dyn MemoryStore>,
|
||||
/// 绕开 AgentSession 封装(NoteTool 在 tool.execute 中直接操作 store)。
|
||||
/// 关键前缀 "note:" 用于 list 过滤。
|
||||
struct NoteTool { store: Arc<dyn MemoryStore> }
|
||||
impl NoteTool { const PREFIX: &'static str = "note:"; }
|
||||
|
||||
#[async_trait]
|
||||
impl BaseTool for NoteTool {
|
||||
fn name(&self) -> &str { "note" }
|
||||
fn description(&self) -> &str { "笔记 save/query: save(key, content) / query()" }
|
||||
fn parameters(&self) -> Value {
|
||||
json!({
|
||||
"type":"object",
|
||||
"properties":{
|
||||
"action":{"type":"string","enum":["save","query"]},
|
||||
"key":{"type":"string"},
|
||||
"content":{"type":"string"}
|
||||
},
|
||||
"required":["action"]
|
||||
})
|
||||
}
|
||||
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||
let action = args["action"].as_str().unwrap_or("");
|
||||
match action {
|
||||
"save" => {
|
||||
let key = args["key"].as_str().unwrap_or("");
|
||||
let content = args["content"].as_str().unwrap_or("");
|
||||
let item = MemoryItem {
|
||||
id: format!("{}{}", Self::PREFIX, key),
|
||||
content: content.to_string(),
|
||||
metadata: json!({}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
};
|
||||
self.store.save(item).await
|
||||
.map_err(|e| ToolError::ExecutionFailed("note".into(), e.to_string()))?;
|
||||
Ok(json!({"saved": key}))
|
||||
}
|
||||
"query" => {
|
||||
let filter = MemoryFilter { prefix: Some(Self::PREFIX.into()), ..Default::default() };
|
||||
let items = self.store.list(&filter).await
|
||||
.map_err(|e| ToolError::ExecutionFailed("note".into(), e.to_string()))?;
|
||||
let notes: Vec<String> = items.into_iter().map(|i| i.content).collect();
|
||||
Ok(json!({"notes": notes}))
|
||||
}
|
||||
_ => Err(ToolError::InvalidArguments("action".into(), format!("未知 action: {action}"))),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// === Mock response helper ===
|
||||
|
||||
fn resp(content: Vec<ContentBlock>, stop: StopReason, u: (u32, u32)) -> MessageResponse {
|
||||
MessageResponse { id: String::new(), model: "mock".into(),
|
||||
message: Message::Assistant { content },
|
||||
usage: Usage::from_input_output(u.0, u.1),
|
||||
stop_reason: stop, extra: Default::default() }
|
||||
}
|
||||
|
||||
fn mock_responses() -> Vec<MessageResponse> {
|
||||
vec![
|
||||
// 第 1 轮:calc(25 * 4) → tool_result(100) → 文本回答
|
||||
resp(vec![ContentBlock::ToolUse { id: "t1".into(), name: "calc".into(),
|
||||
input: json!({"expr": "25 * 4"}) }], StopReason::ToolUse, (5, 8)),
|
||||
resp(vec![ContentBlock::Text { text: "25 * 4 = 100".into() }], StopReason::Stop, (8, 12)),
|
||||
// 第 2 轮:note(save, last_calc, "100") → tool_result(saved) → 文本回答
|
||||
resp(vec![ContentBlock::ToolUse { id: "t2".into(), name: "note".into(),
|
||||
input: json!({"action": "save", "key": "last_calc", "content": "100"}) }],
|
||||
StopReason::ToolUse, (10, 14)),
|
||||
resp(vec![ContentBlock::Text { text: "已记录:last_calc = 100".into() }], StopReason::Stop, (12, 16)),
|
||||
// 第 3 轮:note(query) → tool_result([100]) → 文本回答
|
||||
resp(vec![ContentBlock::ToolUse { id: "t3".into(), name: "note".into(),
|
||||
input: json!({"action": "query"}) }], StopReason::ToolUse, (8, 8)),
|
||||
resp(vec![ContentBlock::Text { text: "您刚才的计算结果是 100".into() }], StopReason::Stop, (10, 14)),
|
||||
// 后续冗余响应(防止队列耗尽报错)
|
||||
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
|
||||
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
|
||||
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
|
||||
]
|
||||
}
|
||||
|
||||
// === Provider selection ===
|
||||
|
||||
fn select_provider() -> Arc<dyn LlmProvider> {
|
||||
if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
|
||||
let cfg = ProviderConfig::from_env("AG_LLM").expect("AG_LLM_* 环境变量解析失败");
|
||||
let provider_type = env::var("AG_LLM_PROVIDER").ok()
|
||||
.and_then(|s| s.parse::<ProviderType>().ok())
|
||||
.unwrap_or(ProviderType::OpenaiChat);
|
||||
Arc::from(create_provider(provider_type, cfg).expect("Provider 创建失败"))
|
||||
} else {
|
||||
let found: Vec<&str> = ["AG_LLM_BASE_URL", "AG_LLM_API_KEY", "AG_LLM_MODEL"]
|
||||
.iter().filter(|k| env::var(k).is_ok()).copied().collect();
|
||||
eprintln!("AG_LLM_* 环境变量不完整(检测到: {:?}),回退到 MockProvider", found);
|
||||
Arc::new(MockProvider::new(mock_responses()))
|
||||
}
|
||||
}
|
||||
|
||||
// === Main ===
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
println!("=== agcore 端到端演示 ===");
|
||||
let dir = TempDir::new().expect("TempDir 创建失败");
|
||||
let db_path = dir.path().join("agcore.db");
|
||||
let backend: Arc<dyn MemoryStore> =
|
||||
Arc::new(SqliteStore::open(&db_path).expect("SqliteStore 打开失败"));
|
||||
println!("💾 SqliteStore: {}", db_path.display());
|
||||
let provider_label = if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
|
||||
"真实 LLM Provider"
|
||||
} else {
|
||||
"MockProvider (离线回退模式)"
|
||||
};
|
||||
println!("🔄 Provider: {provider_label}");
|
||||
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(Arc::new(EchoTool)).unwrap();
|
||||
registry.register(Arc::new(CalcTool)).unwrap();
|
||||
registry.register(Arc::new(NoteTool { store: backend.clone() })).unwrap();
|
||||
println!("🔧 注册工具: {:?}", registry.list_tools());
|
||||
|
||||
let bundle = Arc::new(AgentBuilder::new()
|
||||
.provider(select_provider())
|
||||
.tool_registry(Arc::new(registry))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.build().expect("RuntimeBundle 装配失败"));
|
||||
|
||||
let mut session = AgentSession::new(Arc::new(AssistantAgent), "e2e-1", bundle.clone());
|
||||
|
||||
println!("\n第 1 轮 用户: 帮我算 25 * 4");
|
||||
let r1 = session.submit_turn("帮我算 25 * 4").await.expect("turn 1 失败");
|
||||
println!(" → 回答: {}", r1.text());
|
||||
|
||||
println!("\n第 2 轮 用户: 记下来:结果是 100");
|
||||
let r2 = session.submit_turn("记下来:结果是 100").await.expect("turn 2 失败");
|
||||
println!(" → 回答: {}", r2.text());
|
||||
|
||||
println!("\n第 3 轮 用户: 我刚才算了什么?");
|
||||
let r3 = session.submit_turn("我刚才算了什么?").await.expect("turn 3 失败");
|
||||
println!(" → 回答: {}", r3.text());
|
||||
|
||||
let total = session.usage().total();
|
||||
println!("\n📊 用量: prompt={}, completion={}, total={}",
|
||||
total.prompt_tokens, total.completion_tokens, total.total_tokens);
|
||||
|
||||
println!("\n=== 持久化验证 ===");
|
||||
// 显式释放所有对 backend 的 Arc 引用,确保 SqliteStore Connection 真正关闭。
|
||||
// 释放顺序:session → bundle(间接持有 NoteTool → backend clone)→ backend 局部变量。
|
||||
drop(session); // session.bundle Arc 计数 -1
|
||||
drop(bundle); // bundle Arc 计数归零 → registry → NoteTool → backend clone Arc 计数 2→1
|
||||
drop(backend); // backend 局部变量 Arc 计数 1→0 → SqliteStore::drop → Connection 自动 close
|
||||
let backend2: Arc<dyn MemoryStore> =
|
||||
Arc::new(SqliteStore::open(&db_path).expect("重开 SqliteStore 失败"));
|
||||
let filter = MemoryFilter { prefix: Some("note:".into()), ..Default::default() };
|
||||
let items = backend2.list(&filter).await.expect("list 失败");
|
||||
println!("✓ 跨连接数据存活: 找到 {} 条 note", items.len());
|
||||
assert!(!items.is_empty(), "持久化验证失败:重开后无数据");
|
||||
for i in &items {
|
||||
println!(" - {} = {}", i.id, i.content);
|
||||
}
|
||||
|
||||
println!("\n✓ 端到端演示完成");
|
||||
}
|
||||
@@ -38,14 +38,26 @@ async fn main() {
|
||||
let ks = KnowledgeStore::new(store);
|
||||
|
||||
let pages = vec![
|
||||
make_page("rust-1", "Rust 入门", "Rust 是一门系统级编程语言,注重安全性与并发。"),
|
||||
make_page("python-1", "Python 简介", "Python 是一门动态类型的高级编程语言。"),
|
||||
make_page(
|
||||
"rust-1",
|
||||
"Rust 入门",
|
||||
"Rust 是一门系统级编程语言,注重安全性与并发。",
|
||||
),
|
||||
make_page(
|
||||
"python-1",
|
||||
"Python 简介",
|
||||
"Python 是一门动态类型的高级编程语言。",
|
||||
),
|
||||
make_page(
|
||||
"langgraph-1",
|
||||
"LangGraph 框架",
|
||||
"LangGraph 是 LangChain 的状态图扩展,用于构建多步 Agent。",
|
||||
),
|
||||
make_page("rust-async", "Rust 异步编程", "Rust 异步基于 tokio 与 futures 抽象。"),
|
||||
make_page(
|
||||
"rust-async",
|
||||
"Rust 异步编程",
|
||||
"Rust 异步基于 tokio 与 futures 抽象。",
|
||||
),
|
||||
];
|
||||
for p in &pages {
|
||||
ks.add_page(p.clone()).await.expect("保存页面失败");
|
||||
@@ -62,14 +74,8 @@ async fn main() {
|
||||
let result = retriever.retrieve("Rust 异步").await.unwrap();
|
||||
println!("query: {}", result.query);
|
||||
for item in &result.items {
|
||||
println!(
|
||||
" 命中: {} (score={:.3})",
|
||||
item.page.title, item.score
|
||||
);
|
||||
assert!(
|
||||
(0.0..=1.0).contains(&item.score),
|
||||
"score 应在 [0, 1] 区间"
|
||||
);
|
||||
println!(" 命中: {} (score={:.3})", item.page.title, item.score);
|
||||
assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
|
||||
}
|
||||
assert!(!result.items.is_empty(), "应至少命中一个页面");
|
||||
|
||||
@@ -85,14 +91,8 @@ async fn main() {
|
||||
min_score: 0.5,
|
||||
};
|
||||
let retriever2 = MemoryRetriever::new(ks2, cfg);
|
||||
let result = retriever2
|
||||
.retrieve("完全不相关的火锅配方")
|
||||
.await
|
||||
.unwrap();
|
||||
println!(
|
||||
"无关 query → items.len = {} (期望 0)",
|
||||
result.items.len()
|
||||
);
|
||||
let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
|
||||
println!("无关 query → items.len = {} (期望 0)", result.items.len());
|
||||
assert!(result.items.is_empty());
|
||||
|
||||
// 4. max_results 截断
|
||||
@@ -146,4 +146,4 @@ async fn main() {
|
||||
assert!(only_stop.items.is_empty(), "纯停用词 query 必须返回空结果");
|
||||
|
||||
println!("\n✓ knowledge_search_demo 完成");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,7 +11,7 @@
|
||||
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::prompt::{
|
||||
validate_messages, PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext,
|
||||
PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext, validate_messages,
|
||||
};
|
||||
|
||||
fn message_text(msg: &Message) -> String {
|
||||
@@ -27,16 +27,16 @@ fn message_text(msg: &Message) -> String {
|
||||
})
|
||||
.collect(),
|
||||
Message::UserImage { .. } => "[image]".into(),
|
||||
_ => String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn main() {
|
||||
// 1. PromptTemplate::compile + render —— 直接构造模板
|
||||
println!("=== PromptTemplate::compile + render ===");
|
||||
let tpl = PromptTemplate::compile(
|
||||
"今日 {{location}} 天气:{{condition}},温度 {{temperature}}",
|
||||
)
|
||||
.expect("编译失败");
|
||||
let tpl =
|
||||
PromptTemplate::compile("今日 {{location}} 天气:{{condition}},温度 {{temperature}}")
|
||||
.expect("编译失败");
|
||||
let mut ctx = TemplateContext::new();
|
||||
ctx.insert("location", "北京");
|
||||
ctx.insert("condition", "晴");
|
||||
@@ -58,7 +58,10 @@ fn main() {
|
||||
.register("weather", "今日 {{location}}:{{condition}}")
|
||||
.expect("注册失败");
|
||||
registry
|
||||
.register("greet", "你好 {{name}}!{{#if formal}} 见到您很荣幸。{{/if}}")
|
||||
.register(
|
||||
"greet",
|
||||
"你好 {{name}}!{{#if formal}} 见到您很荣幸。{{/if}}",
|
||||
)
|
||||
.expect("注册失败");
|
||||
|
||||
let mut ctx = TemplateContext::new();
|
||||
@@ -88,6 +91,7 @@ fn main() {
|
||||
Message::User { .. } | Message::UserImage { .. } => "user",
|
||||
Message::Assistant { .. } => "assistant",
|
||||
Message::ToolResult { .. } => "tool",
|
||||
_ => "unknown",
|
||||
};
|
||||
println!("[{i}] {role}: {}", message_text(m));
|
||||
}
|
||||
@@ -105,4 +109,4 @@ fn main() {
|
||||
}
|
||||
|
||||
println!("\n✓ prompt_composer 完成");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
//! quick_start —— 30 行最小可运行示例,展示 Agent / BaseTool / Builder / Session 四层抽象。
|
||||
//!
|
||||
//! 运行:`cargo run --example quick_start`(离线,零配置)
|
||||
|
||||
use std::sync::Arc;
|
||||
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::provider::LlmProvider;
|
||||
use agcore::llm::types::{Usage, message::{ContentBlock, Message}, response_v2::{MessageResponse, StopReason}};
|
||||
use agcore::tools::{BaseTool, ToolContext, ToolError, ToolRegistry};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
struct Greeter;
|
||||
impl Agent for Greeter {
|
||||
fn name(&self) -> &str { "greeter" }
|
||||
fn system_prompt(&self) -> Option<&str> { Some("中文助手,先调用 echo 工具,再总结。") }
|
||||
}
|
||||
|
||||
struct EchoTool;
|
||||
#[async_trait]
|
||||
impl BaseTool for EchoTool {
|
||||
fn name(&self) -> &str { "echo" }
|
||||
fn description(&self) -> &str { "回显文本" }
|
||||
fn parameters(&self) -> Value {
|
||||
json!({"type":"object","properties":{"text":{"type":"string"}},"required":["text"]})
|
||||
}
|
||||
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||
let text = args.get("text").and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ToolError::InvalidArguments("text".into(), "需要 string 类型的 text 参数".into()))?;
|
||||
Ok(json!({"echoed": format!("收到: {text}")}))
|
||||
}
|
||||
}
|
||||
|
||||
fn resp(content: Vec<ContentBlock>, stop: StopReason, u: (u32, u32)) -> MessageResponse {
|
||||
MessageResponse { id: String::new(), model: "mock".into(), message: Message::Assistant { content },
|
||||
usage: Usage::from_input_output(u.0, u.1), stop_reason: stop, extra: Default::default() }
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(Arc::new(EchoTool)).unwrap();
|
||||
let provider: Arc<dyn LlmProvider> = Arc::new(MockProvider::new(vec![
|
||||
resp(vec![ContentBlock::ToolUse { id: "c1".into(), name: "echo".into(),
|
||||
input: json!({"text": "你好"}) }], StopReason::ToolUse, (5, 8)),
|
||||
resp(vec![ContentBlock::Text { text: "EchoTool 已收到您的消息并完成回传。".into() }],
|
||||
StopReason::Stop, (8, 16)),
|
||||
]));
|
||||
let bundle = Arc::new(AgentBuilder::new()
|
||||
.provider(provider).tool_registry(Arc::new(registry))
|
||||
.hook_executor(Arc::new(HookExecutor::new())).build().unwrap());
|
||||
let mut session = AgentSession::new(Arc::new(Greeter), "qs", bundle);
|
||||
let resp = session.submit_turn("你好").await.unwrap();
|
||||
let text = resp.text();
|
||||
println!("LLM: {text}");
|
||||
assert!(text.contains("收到"), "响应应包含'收到'字样: {text}");
|
||||
println!("\n✓ quick_start 完成");
|
||||
}
|
||||
@@ -3,7 +3,7 @@ use std::env;
|
||||
use agcore::init_tracing;
|
||||
use agcore::llm::{
|
||||
cycle::{CycleConfig, LlmCycle},
|
||||
provider::{create_provider, ProviderConfig, ProviderType},
|
||||
provider::{ProviderConfig, ProviderType, create_provider},
|
||||
types::{message::ContentBlock, message::Message, response_v2::MessageResponse},
|
||||
};
|
||||
|
||||
@@ -51,10 +51,11 @@ async fn main() {
|
||||
base_url,
|
||||
api_key,
|
||||
model: model.clone(),
|
||||
timeout_secs: 30,
|
||||
max_retries: 3,
|
||||
};
|
||||
|
||||
let provider = create_provider(provider_type, config)
|
||||
.expect("创建 Provider 失败");
|
||||
let provider = create_provider(provider_type, config).expect("创建 Provider 失败");
|
||||
|
||||
let cycle_config = CycleConfig {
|
||||
model,
|
||||
@@ -63,9 +64,9 @@ async fn main() {
|
||||
..CycleConfig::default()
|
||||
};
|
||||
|
||||
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![
|
||||
Message::system("你是一个简洁的助手,对于任何问题都是用一句话回答。"),
|
||||
]);
|
||||
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![Message::system(
|
||||
"你是一个简洁的助手,对于任何问题都是用一句话回答。",
|
||||
)]);
|
||||
|
||||
println!("发送请求...");
|
||||
|
||||
|
||||
@@ -17,9 +17,9 @@ use std::sync::Arc;
|
||||
use agcore::llm::cycle::{CycleConfig, LlmCycle};
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::provider::LlmProvider;
|
||||
use agcore::llm::types::Usage;
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
|
||||
use agcore::llm::types::Usage;
|
||||
use futures_util::StreamExt;
|
||||
|
||||
/// 构造预设的纯文本响应。
|
||||
@@ -99,9 +99,7 @@ async fn main() {
|
||||
// 上层 Agent 通过 `match` 或 `?` 处理 `AgentError::Llm(_)`。
|
||||
println!("\n=== 阶段 2:错误路径(队列耗尽)===");
|
||||
let mut cycle = LlmCycle::new_with_arc(dyn_provider, CycleConfig::default());
|
||||
let result = cycle
|
||||
.submit_stream("第二次提问".to_string(), vec![])
|
||||
.await;
|
||||
let result = cycle.submit_stream("第二次提问".to_string(), vec![]).await;
|
||||
match result {
|
||||
Ok(_) => panic!("阶段 2 必须失败(队列耗尽)"),
|
||||
Err(e) => {
|
||||
@@ -114,4 +112,4 @@ async fn main() {
|
||||
}
|
||||
|
||||
println!("\n✓ streaming_events_demo 完成");
|
||||
}
|
||||
}
|
||||
|
||||
+25
-24
@@ -8,27 +8,13 @@
|
||||
//! 5. 错误路径:非法 JSON / 空 steps / 缺字段 → `AgentError::PlanParse`
|
||||
//!
|
||||
//! 运行:`cargo run --example task_agent_demo`
|
||||
//!
|
||||
//! ## 已知技术债(v0.2 迁移指南)
|
||||
//!
|
||||
//! 本示例使用 `#[deprecated]` 标记的旧 wire-format 类型:
|
||||
//! - `ChatResponse`、`OpenaiChatMessage`、`FinishReason` —— `OpenaiChatProvider::chat_inner()`
|
||||
//! 内部转换层仍在使用(参见 `docs/10a-phase0-types-and-trait.md` §2.5.1),
|
||||
//! 故结构体定义保留。
|
||||
//! - `StepStatus::Completed(ChatResponse)` —— 因为 `Step` 的"已完成"变体需携带
|
||||
//! provider 响应,目前沿用旧的 `ChatResponse`。
|
||||
//!
|
||||
//! **触发迁移的条件**:v0.2 引入 IR 层的 `StepResult` / 切换为 `MessageResponse`。
|
||||
//! **迁移路径**:将本文件 `ChatResponse`/`OpenaiChatMessage`/`FinishReason` 替换为
|
||||
//! `MessageResponse`/`Message`/`StopReason`,移除顶部 `#![allow(deprecated)]`。
|
||||
//! 上层应用代码(`TaskAgent` 消费者)也可同步迁移。
|
||||
|
||||
#![allow(deprecated)]
|
||||
use std::collections::HashMap;
|
||||
|
||||
use agcore::agent::{AgentError, JsonPlanParser, PlanParser, Step, StepStatus};
|
||||
use agcore::llm::types::openai_message::OpenaiChatMessage;
|
||||
use agcore::llm::types::shared::FinishReason;
|
||||
use agcore::llm::types::{ChatResponse, Usage};
|
||||
use agcore::llm::types::message::Message;
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use agcore::llm::types::Usage;
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
@@ -67,21 +53,36 @@ async fn main() {
|
||||
assert!(step.status.is_pending());
|
||||
|
||||
step.status = StepStatus::Running;
|
||||
println!("Running: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
|
||||
println!(
|
||||
"Running: pending={}, terminal={}",
|
||||
step.status.is_pending(),
|
||||
step.status.is_terminal()
|
||||
);
|
||||
|
||||
step.status = StepStatus::Completed(ChatResponse {
|
||||
message: OpenaiChatMessage::assistant_text("天气:晴,22°C"),
|
||||
step.status = StepStatus::Completed(MessageResponse {
|
||||
id: String::new(),
|
||||
model: "mock".into(),
|
||||
message: Message::assistant("天气:晴,22°C"),
|
||||
usage: Usage::from_input_output(5, 10),
|
||||
stop_reason: Some(FinishReason::Stop),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: HashMap::new(),
|
||||
});
|
||||
println!("Completed: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
|
||||
println!(
|
||||
"Completed: pending={}, terminal={}",
|
||||
step.status.is_pending(),
|
||||
step.status.is_terminal()
|
||||
);
|
||||
assert!(step.status.is_terminal());
|
||||
|
||||
// 3. 失败路径
|
||||
println!("\n=== Step 状态机:失败路径 ===");
|
||||
let mut fail_step = Step::new(0, "调用天气 API");
|
||||
fail_step.status = StepStatus::Failed(AgentError::Other("API 不可用".into()));
|
||||
println!("Failed: pending={}, terminal={}", fail_step.status.is_pending(), fail_step.status.is_terminal());
|
||||
println!(
|
||||
"Failed: pending={}, terminal={}",
|
||||
fail_step.status.is_pending(),
|
||||
fail_step.status.is_terminal()
|
||||
);
|
||||
assert!(fail_step.status.is_terminal());
|
||||
|
||||
// 4. 跳过路径
|
||||
|
||||
+6
-1
@@ -11,6 +11,7 @@
|
||||
|
||||
pub mod agent;
|
||||
pub mod builder;
|
||||
pub mod context;
|
||||
pub mod error;
|
||||
pub mod runtime;
|
||||
pub mod session;
|
||||
@@ -20,9 +21,13 @@ pub mod task;
|
||||
// 重导出公共 API(按使用频度排序)
|
||||
pub use agent::Agent;
|
||||
pub use builder::AgentBuilder;
|
||||
pub use context::{
|
||||
ContextBudget, ContextSlot, DeriveStrategy, FocusedConfig, SlotConfig, SlotMeta, SlotMode,
|
||||
SlotSource,
|
||||
};
|
||||
pub use error::AgentError;
|
||||
pub use runtime::{AgentConfig, RuntimeBundle};
|
||||
pub use session::AgentSession;
|
||||
pub use session_memory::SessionMemory;
|
||||
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
|
||||
pub use task::JsonPlanParser;
|
||||
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
|
||||
|
||||
+2
-4
@@ -7,14 +7,12 @@
|
||||
//! - **不绑定业务循环**:`submit_turn` 在 `AgentSession` 上,不在 trait 上
|
||||
|
||||
use crate::agent::runtime::RuntimeBundle;
|
||||
#[allow(deprecated)]
|
||||
use crate::llm::types::ToolDefinition;
|
||||
use crate::llm::types::tool::ToolDef;
|
||||
|
||||
/// Agent 角色抽象。
|
||||
///
|
||||
/// 实现此 trait 即可接入 Agent Runtime。典型实现是 struct 持有静态配置(name、system prompt 模板),
|
||||
/// 也可以是基于配置动态生成的轻量实现。
|
||||
#[allow(deprecated)]
|
||||
pub trait Agent: Send + Sync {
|
||||
/// 角色名(用于日志、调试、UI 展示)。
|
||||
fn name(&self) -> &str;
|
||||
@@ -26,7 +24,7 @@ pub trait Agent: Send + Sync {
|
||||
///
|
||||
/// **默认实现**:从 `bundle.tool_registry` 取全部工具(最常用模式)。
|
||||
/// **子 trait / 具体实现可覆盖**:做白名单、过滤、按状态动态调整等。
|
||||
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDefinition> {
|
||||
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDef> {
|
||||
bundle.tool_registry.definitions()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -92,15 +92,17 @@ impl AgentBuilder {
|
||||
/// `AgentError::Config(...)`,提示调用 `.provider(...)` / `.tool_registry(...)` /
|
||||
/// `.hook_executor(...)` 补齐。不 panic。
|
||||
pub fn build(self) -> Result<RuntimeBundle, AgentError> {
|
||||
let provider = self
|
||||
.provider
|
||||
.ok_or_else(|| AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into()))?;
|
||||
let provider = self.provider.ok_or_else(|| {
|
||||
AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into())
|
||||
})?;
|
||||
let tool_registry = self
|
||||
.tool_registry
|
||||
.ok_or_else(|| AgentError::Config("缺少 tool_registry,请先调用 .tool_registry(...)(即使是空 ToolRegistry 也需要传入)".into()))?;
|
||||
let hook_executor = self
|
||||
.hook_executor
|
||||
.ok_or_else(|| AgentError::Config("缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into()))?;
|
||||
let hook_executor = self.hook_executor.ok_or_else(|| {
|
||||
AgentError::Config(
|
||||
"缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into(),
|
||||
)
|
||||
})?;
|
||||
|
||||
let config = self.config.unwrap_or_default();
|
||||
|
||||
|
||||
@@ -0,0 +1,900 @@
|
||||
//! ContextSlot —— 多上下文槽位管理。
|
||||
//!
|
||||
//! 设计要点(参见 `docs/17-phase10-contextslot.md`):
|
||||
//!
|
||||
//! - **多上下文分区**:单个 session 内可创建/切换/派生多个独立消息上下文
|
||||
//! - **三种模式**:Full(完整历史)/ Focused(读时过滤)/ Readonly(禁止写入)
|
||||
//! - **三种来源**:New(全新)/ Derived(派生)/ Static(静态)
|
||||
//! - **基于 MemoryStore trait 持久化**:JSON blob 批次存储,每 slot 3-4 条 MemoryItem 记录
|
||||
//! - **零新依赖方向**:放在 `agent/` 下利用已有的 `agent → memory` 依赖
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use time::OffsetDateTime;
|
||||
|
||||
use crate::agent::error::AgentError;
|
||||
use crate::llm::types::message::Message;
|
||||
use crate::memory::store::MemoryStore;
|
||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||
|
||||
/// 上下文槽 —— 一段带策略配置的消息列表。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ContextSlot {
|
||||
/// 当前 slot 的唯一标识(同一个 session_id 内唯一)。
|
||||
pub id: String,
|
||||
/// 所属 session。
|
||||
pub session_id: String,
|
||||
/// 槽配置。
|
||||
pub config: SlotConfig,
|
||||
/// 消息列表(全量,Focused/Readonly 在读取时做策略过滤)。
|
||||
pub messages: Vec<Message>,
|
||||
/// 槽元数据。
|
||||
pub meta: SlotMeta,
|
||||
}
|
||||
|
||||
/// 槽配置。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SlotConfig {
|
||||
/// 槽模式(Full / Focused / Readonly)。
|
||||
pub mode: SlotMode,
|
||||
/// 槽来源(New / Derived / Static)。
|
||||
pub source: SlotSource,
|
||||
/// 上下文预算(v0.2 纯数据结构,无消费逻辑)。
|
||||
pub budget: ContextBudget,
|
||||
/// 是否启用自动压缩(v0.2 保留字段,LlmCycle 内部自行判断)。
|
||||
pub compact: bool,
|
||||
}
|
||||
|
||||
impl Default for SlotConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
mode: SlotMode::Full,
|
||||
source: SlotSource::New,
|
||||
budget: ContextBudget::default(),
|
||||
compact: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 槽模式。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[non_exhaustive]
|
||||
pub enum SlotMode {
|
||||
/// 完整对话历史(全部消息)。
|
||||
Full,
|
||||
/// 聚焦模式 —— 读取时按策略过滤,保持 LLM 注意力。
|
||||
Focused(FocusedConfig),
|
||||
/// 只读参考上下文 —— 禁止写入。
|
||||
Readonly,
|
||||
}
|
||||
|
||||
/// 聚焦模式配置。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FocusedConfig {
|
||||
/// 是否保留 system prompt。
|
||||
pub keep_system: bool,
|
||||
/// 保留的最近消息条数(以消息条数而非对话轮次为单位,因为一轮对话可能包含多条 tool 消息)。
|
||||
pub recent_messages: usize,
|
||||
/// 摘要覆盖(v0.2 仅消费端:手动设置则注入,不自动生成)。
|
||||
/// v0.3 将支持 Hook 驱动的自动摘要生成。
|
||||
pub summary_override: Option<String>,
|
||||
}
|
||||
|
||||
/// 槽来源。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[non_exhaustive]
|
||||
pub enum SlotSource {
|
||||
/// 全新空槽。
|
||||
New,
|
||||
/// 从父 slot 派生(记录 parent_id)。
|
||||
Derived {
|
||||
parent_id: String,
|
||||
strategy: DeriveStrategy,
|
||||
},
|
||||
/// 预置静态消息(不持久化,随 session 生命周期存在)。
|
||||
Static(Vec<Message>),
|
||||
}
|
||||
|
||||
/// 派生策略。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub enum DeriveStrategy {
|
||||
/// 完整复制父 slot 的消息。
|
||||
Full,
|
||||
/// 按聚焦策略复制父 slot 的消息。
|
||||
Focused(FocusedConfig),
|
||||
}
|
||||
|
||||
/// 上下文预算(v0.2 纯数据结构,无消费逻辑)。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ContextBudget {
|
||||
/// system prompt 预算。
|
||||
pub system: u32,
|
||||
/// 对话历史预算。
|
||||
pub history: u32,
|
||||
/// 工具定义预算。
|
||||
pub tools: u32,
|
||||
/// 工具结果预算。
|
||||
pub tool_results: u32,
|
||||
/// 预留 buffer。
|
||||
pub reserve: u32,
|
||||
}
|
||||
|
||||
impl Default for ContextBudget {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
system: 8_000,
|
||||
history: 80_000,
|
||||
tools: 10_000,
|
||||
tool_results: 20_000,
|
||||
reserve: 10_000,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ContextBudget {
|
||||
/// 自动分配:按上下文窗口的固定比例分配预算。
|
||||
/// v0.2 只做占位实现,v0.3 将根据实际 provider 的 context_window 计算。
|
||||
pub fn auto() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// 槽元数据。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SlotMeta {
|
||||
/// 父 slot id(仅 Derived 来源有值)。
|
||||
pub parent_id: Option<String>,
|
||||
/// 消息总数。
|
||||
pub message_count: usize,
|
||||
/// 总 token 估算值(由 add_messages 时累计,v0.2 为近似值)。
|
||||
pub total_tokens: u32,
|
||||
/// 创建时间(Unix 时间戳,秒)。
|
||||
pub created_at: u64,
|
||||
}
|
||||
|
||||
impl SlotMeta {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
parent_id: None,
|
||||
message_count: 0,
|
||||
total_tokens: 0,
|
||||
created_at: std::time::SystemTime::now()
|
||||
.duration_since(std::time::SystemTime::UNIX_EPOCH)
|
||||
.map(|d| d.as_secs())
|
||||
.unwrap_or(0),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for SlotMeta {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl ContextSlot {
|
||||
/// 持久化 key 前缀。
|
||||
const KEY_DATA: &'static str = "slot_data";
|
||||
const KEY_META: &'static str = "slot_meta";
|
||||
const KEY_CONFIG: &'static str = "slot_config";
|
||||
const KEY_REL: &'static str = "slot_rel";
|
||||
|
||||
/// 校验 id 不含冒号(避免破坏 key 格式与 list prefix 过滤)。
|
||||
/// 失败时 panic —— 这是开发者错误而非用户错误。
|
||||
fn assert_no_colon(id: &str, field: &str) {
|
||||
if id.contains(':') {
|
||||
panic!(
|
||||
"{field} '{id}' contains ':' which would break key format. \
|
||||
Use only letters, digits, hyphens and underscores."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn data_key(session_id: &str, slot_id: &str) -> String {
|
||||
Self::assert_no_colon(session_id, "session_id");
|
||||
Self::assert_no_colon(slot_id, "slot_id");
|
||||
format!("{}:{}:{}", Self::KEY_DATA, session_id, slot_id)
|
||||
}
|
||||
pub(crate) fn meta_key(session_id: &str, slot_id: &str) -> String {
|
||||
Self::assert_no_colon(session_id, "session_id");
|
||||
Self::assert_no_colon(slot_id, "slot_id");
|
||||
format!("{}:{}:{}", Self::KEY_META, session_id, slot_id)
|
||||
}
|
||||
pub(crate) fn config_key(session_id: &str, slot_id: &str) -> String {
|
||||
Self::assert_no_colon(session_id, "session_id");
|
||||
Self::assert_no_colon(slot_id, "slot_id");
|
||||
format!("{}:{}:{}", Self::KEY_CONFIG, session_id, slot_id)
|
||||
}
|
||||
pub(crate) fn rel_key(session_id: &str, child_id: &str) -> String {
|
||||
Self::assert_no_colon(session_id, "session_id");
|
||||
Self::assert_no_colon(child_id, "child_id");
|
||||
format!("{}:{}:{}", Self::KEY_REL, session_id, child_id)
|
||||
}
|
||||
|
||||
/// 构造 MemoryItem 的辅助函数。
|
||||
fn make_item(key: String, content: String) -> MemoryItem {
|
||||
MemoryItem {
|
||||
id: key,
|
||||
content,
|
||||
metadata: serde_json::json!({}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建一个新的空 ContextSlot(不持久化,仅内存构造)。
|
||||
pub fn new(
|
||||
session_id: impl Into<String>,
|
||||
slot_id: impl Into<String>,
|
||||
config: SlotConfig,
|
||||
) -> Self {
|
||||
Self {
|
||||
id: slot_id.into(),
|
||||
session_id: session_id.into(),
|
||||
config,
|
||||
messages: Vec::new(),
|
||||
meta: SlotMeta::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 保存 slot 数据到存储后端(全量写入,含 config)。
|
||||
pub async fn save(&self, store: &dyn MemoryStore) -> Result<(), AgentError> {
|
||||
let data = serde_json::to_string(&self.messages)
|
||||
.map_err(|e| AgentError::Other(e.to_string()))?;
|
||||
let meta = serde_json::to_string(&self.meta)
|
||||
.map_err(|e| AgentError::Other(e.to_string()))?;
|
||||
let config = serde_json::to_string(&self.config)
|
||||
.map_err(|e| AgentError::Other(e.to_string()))?;
|
||||
|
||||
store
|
||||
.save(Self::make_item(
|
||||
Self::data_key(&self.session_id, &self.id),
|
||||
data,
|
||||
))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
store
|
||||
.save(Self::make_item(
|
||||
Self::meta_key(&self.session_id, &self.id),
|
||||
meta,
|
||||
))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
store
|
||||
.save(Self::make_item(
|
||||
Self::config_key(&self.session_id, &self.id),
|
||||
config,
|
||||
))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
|
||||
// 派生关系
|
||||
if let SlotSource::Derived { parent_id, .. } = &self.config.source {
|
||||
store
|
||||
.save(Self::make_item(
|
||||
Self::rel_key(&self.session_id, &self.id),
|
||||
parent_id.clone(),
|
||||
))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 从存储加载 slot,config 从 `slot_config` key 自行恢复。
|
||||
/// 若 config 记录不存在(旧版本升级场景),使用 `SlotConfig::default()`。
|
||||
pub async fn load(
|
||||
id: &str,
|
||||
session_id: &str,
|
||||
store: &dyn MemoryStore,
|
||||
) -> Result<Option<Self>, AgentError> {
|
||||
let meta_item = store
|
||||
.get(&Self::meta_key(session_id, id))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
let data_item = store
|
||||
.get(&Self::data_key(session_id, id))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
let config_item = store
|
||||
.get(&Self::config_key(session_id, id))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
|
||||
match (meta_item, data_item) {
|
||||
(Some(m), Some(d)) => {
|
||||
let meta: SlotMeta = serde_json::from_str(&m.content)
|
||||
.map_err(|e| AgentError::Other(e.to_string()))?;
|
||||
let messages: Vec<Message> = serde_json::from_str(&d.content)
|
||||
.map_err(|e| AgentError::Other(e.to_string()))?;
|
||||
// config 从存储恢复;不存在则使用 default(兼容旧版本)
|
||||
let config = match config_item {
|
||||
Some(c) => serde_json::from_str(&c.content)
|
||||
.map_err(|e| AgentError::Other(e.to_string()))?,
|
||||
None => SlotConfig::default(),
|
||||
};
|
||||
Ok(Some(Self {
|
||||
id: id.to_string(),
|
||||
session_id: session_id.to_string(),
|
||||
config,
|
||||
messages,
|
||||
meta,
|
||||
}))
|
||||
}
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// 列出某 session 下的所有 slot 元数据。
|
||||
pub async fn list(
|
||||
session_id: &str,
|
||||
store: &dyn MemoryStore,
|
||||
) -> Result<Vec<SlotMeta>, AgentError> {
|
||||
let prefix_str = format!("{}:{}:", Self::KEY_META, session_id);
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some(prefix_str),
|
||||
..Default::default()
|
||||
};
|
||||
let items = store.list(&filter).await.map_err(AgentError::Memory)?;
|
||||
let mut metas = Vec::new();
|
||||
for item in items {
|
||||
if let Ok(meta) = serde_json::from_str::<SlotMeta>(&item.content) {
|
||||
metas.push(meta);
|
||||
}
|
||||
}
|
||||
Ok(metas)
|
||||
}
|
||||
|
||||
/// 删除 slot 的所有存储记录(slot_data + slot_meta + slot_config + slot_rel)。
|
||||
pub async fn delete(
|
||||
id: &str,
|
||||
session_id: &str,
|
||||
store: &dyn MemoryStore,
|
||||
) -> Result<(), AgentError> {
|
||||
store
|
||||
.delete(&Self::data_key(session_id, id))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
store
|
||||
.delete(&Self::meta_key(session_id, id))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
store
|
||||
.delete(&Self::config_key(session_id, id))
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
// slot_rel 是 best-effort(仅 Derived 来源的 slot 才有此 key)
|
||||
let _ = store.delete(&Self::rel_key(session_id, id)).await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 追加消息(Readonly 模式下返回 `SlotReadonly` 错误)。
|
||||
/// Full / Focused 模式下允许追加。
|
||||
pub fn append_messages(&mut self, new_messages: Vec<Message>) -> Result<(), AgentError> {
|
||||
if matches!(self.config.mode, SlotMode::Readonly) {
|
||||
return Err(AgentError::SlotReadonly(
|
||||
"Readonly slot does not allow writes".into(),
|
||||
));
|
||||
}
|
||||
let count = new_messages.len();
|
||||
self.messages.extend(new_messages);
|
||||
self.meta.message_count += count;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 按 FocusedConfig 过滤消息(静态辅助函数,被 `load_messages` 和 `derive_slot` 复用)。
|
||||
///
|
||||
/// 过滤逻辑:
|
||||
/// 1. 保留第一条 system 消息(如果 `keep_system=true`)
|
||||
/// 2. 取最近 `recent_messages` 条非 system 消息(如果 `recent_messages > 0`)
|
||||
/// 3. 追加摘要消息(如果 `summary_override` 存在)
|
||||
pub fn filter_focused(messages: &[Message], cfg: &FocusedConfig) -> Vec<Message> {
|
||||
let mut result = Vec::new();
|
||||
// 保留 system prompt
|
||||
if cfg.keep_system
|
||||
&& let Some(msg) = messages
|
||||
.iter()
|
||||
.find(|m| matches!(m, Message::System { .. }))
|
||||
{
|
||||
result.push(msg.clone());
|
||||
}
|
||||
// 处理 recent_messages=0 边界:上面已处理 system,下面仅取最近 N 条
|
||||
if cfg.recent_messages > 0 {
|
||||
let recent: Vec<&Message> = messages
|
||||
.iter()
|
||||
.filter(|m| !matches!(m, Message::System { .. }))
|
||||
.collect();
|
||||
let start = recent.len().saturating_sub(cfg.recent_messages);
|
||||
for msg in recent.iter().skip(start) {
|
||||
result.push((*msg).clone());
|
||||
}
|
||||
}
|
||||
// 注入摘要
|
||||
if let Some(summary) = &cfg.summary_override {
|
||||
result.push(Message::system(format!("[上下文摘要] {}", summary)));
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
/// 返回消息列表。Focused 模式下按策略过滤(裁剪到最近 recent_messages 条)。
|
||||
pub fn load_messages(&self) -> Vec<Message> {
|
||||
match &self.config.mode {
|
||||
SlotMode::Focused(cfg) => Self::filter_focused(&self.messages, cfg),
|
||||
_ => self.messages.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::memory::store::InMemoryStore;
|
||||
|
||||
fn make_store() -> std::sync::Arc<dyn MemoryStore> {
|
||||
std::sync::Arc::new(InMemoryStore::new())
|
||||
}
|
||||
|
||||
fn make_slot(id: &str, session: &str) -> ContextSlot {
|
||||
ContextSlot::new(session, id, SlotConfig::default())
|
||||
}
|
||||
|
||||
/// 提取 `Message` 的第一个 Text block 的内容(用于测试断言)。
|
||||
/// 返回 None 表示该消息不含纯文本 block。
|
||||
fn extract_text(msg: &Message) -> &str {
|
||||
use crate::llm::types::message::ContentBlock;
|
||||
let blocks = match msg {
|
||||
Message::System { content }
|
||||
| Message::User { content }
|
||||
| Message::Assistant { content } => content,
|
||||
Message::UserImage { .. } => return "",
|
||||
Message::ToolResult { content, .. } => content,
|
||||
};
|
||||
for block in blocks {
|
||||
if let ContentBlock::Text { text } = block {
|
||||
return text;
|
||||
}
|
||||
}
|
||||
""
|
||||
}
|
||||
|
||||
// ===== 持久化 =====
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_save_load_roundtrip() {
|
||||
let store = make_store();
|
||||
let mut slot = make_slot("default", "s1");
|
||||
slot.append_messages(vec![Message::user_text("hi")]).unwrap();
|
||||
slot.append_messages(vec![Message::assistant("hello")]).unwrap();
|
||||
|
||||
slot.save(&*store).await.unwrap();
|
||||
let loaded = ContextSlot::load("default", "s1", &*store).await.unwrap();
|
||||
let loaded = loaded.expect("slot should exist after save");
|
||||
assert_eq!(loaded.id, "default");
|
||||
assert_eq!(loaded.session_id, "s1");
|
||||
assert_eq!(loaded.messages.len(), 2);
|
||||
assert_eq!(loaded.meta.message_count, 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_session_isolation() {
|
||||
let store = make_store();
|
||||
let mut a = make_slot("main", "sA");
|
||||
a.append_messages(vec![Message::user_text("only in A")])
|
||||
.unwrap();
|
||||
a.save(&*store).await.unwrap();
|
||||
|
||||
let mut b = make_slot("main", "sB");
|
||||
b.append_messages(vec![Message::user_text("only in B")])
|
||||
.unwrap();
|
||||
b.save(&*store).await.unwrap();
|
||||
|
||||
let loaded_a = ContextSlot::load("main", "sA", &*store).await.unwrap().unwrap();
|
||||
let loaded_b = ContextSlot::load("main", "sB", &*store).await.unwrap().unwrap();
|
||||
assert_eq!(extract_text(&loaded_a.messages[0]), "only in A");
|
||||
assert_eq!(extract_text(&loaded_b.messages[0]), "only in B");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_derived_parent_id_recorded() {
|
||||
let store = make_store();
|
||||
let slot = ContextSlot::new(
|
||||
"s1",
|
||||
"child",
|
||||
SlotConfig {
|
||||
mode: SlotMode::Full,
|
||||
source: SlotSource::Derived {
|
||||
parent_id: "default".to_string(),
|
||||
strategy: DeriveStrategy::Full,
|
||||
},
|
||||
budget: ContextBudget::default(),
|
||||
compact: true,
|
||||
},
|
||||
);
|
||||
slot.save(&*store).await.unwrap();
|
||||
|
||||
// rel_key 直接读
|
||||
let rel = store
|
||||
.get(&ContextSlot::rel_key("s1", "child"))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(rel.content, "default");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_readonly_rejects_write() {
|
||||
let mut slot = make_slot("ro", "s1");
|
||||
slot.config.mode = SlotMode::Readonly;
|
||||
let result = slot.append_messages(vec![Message::user_text("nope")]);
|
||||
assert!(matches!(result, Err(AgentError::SlotReadonly(_))));
|
||||
assert!(slot.messages.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_delete_then_load_none() {
|
||||
let store = make_store();
|
||||
let mut slot = make_slot("to_delete", "s1");
|
||||
slot.append_messages(vec![Message::user_text("hi")]).unwrap();
|
||||
slot.save(&*store).await.unwrap();
|
||||
|
||||
ContextSlot::delete("to_delete", "s1", &*store).await.unwrap();
|
||||
let loaded = ContextSlot::load("to_delete", "s1", &*store).await.unwrap();
|
||||
assert!(loaded.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_list_multiple() {
|
||||
let store = make_store();
|
||||
for id in ["alpha", "beta", "gamma"] {
|
||||
let mut s = make_slot(id, "sX");
|
||||
s.append_messages(vec![Message::user_text(id)]).unwrap();
|
||||
s.save(&*store).await.unwrap();
|
||||
}
|
||||
// 不同 session 不该列出
|
||||
let mut s2 = make_slot("alpha", "sY");
|
||||
s2.append_messages(vec![Message::user_text("y")]).unwrap();
|
||||
s2.save(&*store).await.unwrap();
|
||||
|
||||
let metas = ContextSlot::list("sX", &*store).await.unwrap();
|
||||
assert_eq!(metas.len(), 3);
|
||||
let metas_y = ContextSlot::list("sY", &*store).await.unwrap();
|
||||
assert_eq!(metas_y.len(), 1);
|
||||
}
|
||||
|
||||
// ===== Focused 模式 =====
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_focused_recent_messages() {
|
||||
let mut slot = make_slot("f", "s1");
|
||||
slot.append_messages(vec![Message::system("sys")]).unwrap();
|
||||
for i in 0..5 {
|
||||
slot.append_messages(vec![Message::user_text(format!("u{i}"))])
|
||||
.unwrap();
|
||||
slot.append_messages(vec![Message::assistant(format!("a{i}"))])
|
||||
.unwrap();
|
||||
}
|
||||
slot.config.mode = SlotMode::Focused(FocusedConfig {
|
||||
keep_system: true,
|
||||
recent_messages: 3,
|
||||
summary_override: None,
|
||||
});
|
||||
|
||||
let loaded = slot.load_messages();
|
||||
// system + 最近 3 条 (assistant 4, user 4, assistant 5 实际是按 vec 顺序取最近 3 条非 system)
|
||||
let has_sys = loaded.iter().any(|m| matches!(m, Message::System { .. }));
|
||||
assert!(has_sys, "system 提示应保留");
|
||||
// 最近 3 条非 system 应该是 a4, u4, a5 (按 messages 存储顺序的最后 3 条)
|
||||
assert_eq!(loaded.len(), 1 + 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_focused_summary_override() {
|
||||
let mut slot = make_slot("f", "s1");
|
||||
slot.append_messages(vec![Message::user_text("u")]).unwrap();
|
||||
slot.append_messages(vec![Message::assistant("a")]).unwrap();
|
||||
slot.config.mode = SlotMode::Focused(FocusedConfig {
|
||||
keep_system: false,
|
||||
recent_messages: 100,
|
||||
summary_override: Some("讨论了 X".to_string()),
|
||||
});
|
||||
|
||||
let loaded = slot.load_messages();
|
||||
// 2 条原始 + 1 条摘要 system = 3
|
||||
assert_eq!(loaded.len(), 3);
|
||||
// 最后一条是摘要
|
||||
if let Message::System { content } = &loaded[2] {
|
||||
let text = format!("{:?}", content);
|
||||
assert!(text.contains("上下文摘要"));
|
||||
assert!(text.contains("讨论了 X"));
|
||||
} else {
|
||||
panic!("最后一条应为 system 摘要");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_focused_zero_messages() {
|
||||
let mut slot = make_slot("f", "s1");
|
||||
slot.append_messages(vec![Message::system("sys")]).unwrap();
|
||||
slot.append_messages(vec![Message::user_text("u")]).unwrap();
|
||||
slot.config.mode = SlotMode::Focused(FocusedConfig {
|
||||
keep_system: true,
|
||||
recent_messages: 0,
|
||||
summary_override: None,
|
||||
});
|
||||
|
||||
let loaded = slot.load_messages();
|
||||
// recent_messages=0 但 keep_system=true 应只含 system
|
||||
assert_eq!(loaded.len(), 1);
|
||||
assert!(matches!(loaded[0], Message::System { .. }));
|
||||
}
|
||||
|
||||
// ===== 边界 =====
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_empty_messages_roundtrip() {
|
||||
let store = make_store();
|
||||
let slot = make_slot("empty", "s1");
|
||||
slot.save(&*store).await.unwrap();
|
||||
let loaded = ContextSlot::load("empty", "s1", &*store).await.unwrap().unwrap();
|
||||
assert!(loaded.messages.is_empty());
|
||||
assert_eq!(loaded.meta.message_count, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_save_on_readonly_side_effect() {
|
||||
let store = make_store();
|
||||
let mut slot = make_slot("ro", "s1");
|
||||
slot.config.mode = SlotMode::Readonly;
|
||||
// save 本身允许(只禁止 append)
|
||||
slot.save(&*store).await.unwrap();
|
||||
let loaded = ContextSlot::load("ro", "s1", &*store).await.unwrap();
|
||||
assert!(loaded.is_some());
|
||||
}
|
||||
|
||||
// ===== 派生 (derive_slot 行为) =====
|
||||
|
||||
#[tokio::test]
|
||||
async fn derive_full_copies_parent_messages() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
for i in 0..3 {
|
||||
parent.append_messages(vec![Message::user_text(format!("u{i}"))])
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
// 模拟 derive_slot 内部 Full 策略
|
||||
let child_messages = parent.messages.clone();
|
||||
let child = ContextSlot::new(
|
||||
"s1",
|
||||
"c",
|
||||
SlotConfig {
|
||||
mode: SlotMode::Full,
|
||||
source: SlotSource::Derived {
|
||||
parent_id: "p".to_string(),
|
||||
strategy: DeriveStrategy::Full,
|
||||
},
|
||||
budget: ContextBudget::default(),
|
||||
compact: true,
|
||||
},
|
||||
);
|
||||
let mut child = child;
|
||||
child.messages = child_messages;
|
||||
let store = make_store();
|
||||
child.save(&*store).await.unwrap();
|
||||
|
||||
let loaded = ContextSlot::load("c", "s1", &*store).await.unwrap().unwrap();
|
||||
assert_eq!(loaded.messages.len(), 3);
|
||||
assert!(matches!(loaded.config.source, SlotSource::Derived { .. }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn derive_focused_filters_parent_messages() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::system("sys")]).unwrap();
|
||||
for i in 0..5 {
|
||||
parent.append_messages(vec![Message::user_text(format!("u{i}"))])
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
// 模拟 derive_slot 内部 Focused 策略:按 FocusedConfig 过滤
|
||||
let cfg = FocusedConfig {
|
||||
keep_system: true,
|
||||
recent_messages: 2,
|
||||
summary_override: None,
|
||||
};
|
||||
// 应用 load_messages 同样的过滤
|
||||
let mut filtered = Vec::new();
|
||||
if cfg.keep_system
|
||||
&& let Some(m) = parent
|
||||
.messages
|
||||
.iter()
|
||||
.find(|m| matches!(m, Message::System { .. }))
|
||||
{
|
||||
filtered.push(m.clone());
|
||||
}
|
||||
if cfg.recent_messages > 0 {
|
||||
let recent: Vec<&Message> = parent
|
||||
.messages
|
||||
.iter()
|
||||
.filter(|m| !matches!(m, Message::System { .. }))
|
||||
.collect();
|
||||
let start = recent.len().saturating_sub(cfg.recent_messages);
|
||||
for m in recent.iter().skip(start) {
|
||||
filtered.push((*m).clone());
|
||||
}
|
||||
}
|
||||
|
||||
assert_eq!(filtered.len(), 1 + 2); // system + 2 条
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn derived_slot_loadable_independently() {
|
||||
let store = make_store();
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::user_text("u")]).unwrap();
|
||||
parent.save(&*store).await.unwrap();
|
||||
|
||||
// 派生 child
|
||||
let mut child = ContextSlot::new(
|
||||
"s1",
|
||||
"c",
|
||||
SlotConfig {
|
||||
mode: SlotMode::Full,
|
||||
source: SlotSource::Derived {
|
||||
parent_id: "p".to_string(),
|
||||
strategy: DeriveStrategy::Full,
|
||||
},
|
||||
budget: ContextBudget::default(),
|
||||
compact: true,
|
||||
},
|
||||
);
|
||||
child.append_messages(vec![Message::user_text("derived msg")])
|
||||
.unwrap();
|
||||
child.save(&*store).await.unwrap();
|
||||
|
||||
// child 可独立加载
|
||||
let loaded = ContextSlot::load("c", "s1", &*store).await.unwrap().unwrap();
|
||||
assert_eq!(loaded.messages.len(), 1);
|
||||
assert_eq!(extract_text(&loaded.messages[0]), "derived msg");
|
||||
}
|
||||
|
||||
// ===== delete 保护 (AgentSession 层,但 ContextSlot.delete 不保护;逻辑测试在 session.rs) =====
|
||||
|
||||
#[tokio::test]
|
||||
async fn slot_delete_cleans_all_records() {
|
||||
let store = make_store();
|
||||
let mut slot = ContextSlot::new(
|
||||
"s1",
|
||||
"x",
|
||||
SlotConfig {
|
||||
mode: SlotMode::Full,
|
||||
source: SlotSource::Derived {
|
||||
parent_id: "p".to_string(),
|
||||
strategy: DeriveStrategy::Full,
|
||||
},
|
||||
budget: ContextBudget::default(),
|
||||
compact: true,
|
||||
},
|
||||
);
|
||||
slot.append_messages(vec![Message::user_text("u")]).unwrap();
|
||||
slot.save(&*store).await.unwrap();
|
||||
|
||||
// 确认所有记录存在
|
||||
assert!(store.get(&ContextSlot::data_key("s1", "x")).await.unwrap().is_some());
|
||||
assert!(store.get(&ContextSlot::meta_key("s1", "x")).await.unwrap().is_some());
|
||||
assert!(store.get(&ContextSlot::config_key("s1", "x")).await.unwrap().is_some());
|
||||
assert!(store.get(&ContextSlot::rel_key("s1", "x")).await.unwrap().is_some());
|
||||
|
||||
ContextSlot::delete("x", "s1", &*store).await.unwrap();
|
||||
|
||||
// data/meta/config 已删
|
||||
assert!(store.get(&ContextSlot::data_key("s1", "x")).await.unwrap().is_none());
|
||||
assert!(store.get(&ContextSlot::meta_key("s1", "x")).await.unwrap().is_none());
|
||||
assert!(store.get(&ContextSlot::config_key("s1", "x")).await.unwrap().is_none());
|
||||
}
|
||||
|
||||
// ===== 基础类型测试 =====
|
||||
|
||||
#[test]
|
||||
fn slot_meta_new_sets_zero_message_count() {
|
||||
let m = SlotMeta::new();
|
||||
assert_eq!(m.message_count, 0);
|
||||
assert_eq!(m.total_tokens, 0);
|
||||
assert!(m.parent_id.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn context_budget_default_sum_128k() {
|
||||
let b = ContextBudget::default();
|
||||
assert_eq!(b.system + b.history + b.tools + b.tool_results + b.reserve, 128_000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn slot_config_default_is_full_new() {
|
||||
let c = SlotConfig::default();
|
||||
assert!(matches!(c.mode, SlotMode::Full));
|
||||
assert!(matches!(c.source, SlotSource::New));
|
||||
assert!(c.compact);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn focused_config_serializes_roundtrip() {
|
||||
let cfg = FocusedConfig {
|
||||
keep_system: true,
|
||||
recent_messages: 5,
|
||||
summary_override: Some("sum".into()),
|
||||
};
|
||||
let json = serde_json::to_string(&cfg).unwrap();
|
||||
let back: FocusedConfig = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(back.recent_messages, 5);
|
||||
assert_eq!(back.summary_override.as_deref(), Some("sum"));
|
||||
}
|
||||
|
||||
// ====== filter_focused 静态方法(被 load_messages 和 derive_slot 复用) ======
|
||||
|
||||
#[test]
|
||||
fn filter_focused_keeps_system_and_recent() {
|
||||
let mut messages = vec![Message::system("sys")];
|
||||
for i in 0..5 {
|
||||
messages.push(Message::user_text(format!("u{i}")));
|
||||
messages.push(Message::assistant(format!("a{i}")));
|
||||
}
|
||||
let cfg = FocusedConfig {
|
||||
keep_system: true,
|
||||
recent_messages: 3,
|
||||
summary_override: None,
|
||||
};
|
||||
let filtered = ContextSlot::filter_focused(&messages, &cfg);
|
||||
// system + 3 条最近的非 system 消息
|
||||
assert_eq!(filtered.len(), 1 + 3);
|
||||
assert!(matches!(filtered[0], Message::System { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn filter_focused_injects_summary() {
|
||||
let messages = vec![
|
||||
Message::user_text("u"),
|
||||
Message::assistant("a"),
|
||||
];
|
||||
let cfg = FocusedConfig {
|
||||
keep_system: false,
|
||||
recent_messages: 100,
|
||||
summary_override: Some("讨论了 X".into()),
|
||||
};
|
||||
let filtered = ContextSlot::filter_focused(&messages, &cfg);
|
||||
// 2 条原始 + 1 条摘要 system
|
||||
assert_eq!(filtered.len(), 3);
|
||||
if let Message::System { content } = &filtered[2] {
|
||||
let text = format!("{:?}", content);
|
||||
assert!(text.contains("上下文摘要"));
|
||||
} else {
|
||||
panic!("最后一条应为 system 摘要");
|
||||
}
|
||||
}
|
||||
|
||||
// ====== Colon 校验(key 格式保护) ======
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "session_id 's:1' contains ':'")]
|
||||
fn key_constructor_rejects_colon_in_session_id() {
|
||||
// 通过 make_slot 间接调用 slot.save 时会触发 data_key -> assert_no_colon
|
||||
let store = make_store();
|
||||
let slot = ContextSlot::new("s:1", "default", SlotConfig::default());
|
||||
let _ = tokio_test_runtime(slot.save(&*store));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "slot_id 'a:b' contains ':'")]
|
||||
fn key_constructor_rejects_colon_in_slot_id() {
|
||||
let store = make_store();
|
||||
let slot = ContextSlot::new("s1", "a:b", SlotConfig::default());
|
||||
let _ = tokio_test_runtime(slot.save(&*store));
|
||||
}
|
||||
|
||||
/// 在同步测试中运行 future 的辅助函数。
|
||||
fn tokio_test_runtime<F: std::future::Future>(f: F) -> F::Output {
|
||||
tokio::runtime::Builder::new_current_thread()
|
||||
.enable_all()
|
||||
.build()
|
||||
.unwrap()
|
||||
.block_on(f)
|
||||
}
|
||||
}
|
||||
+54
-3
@@ -18,6 +18,7 @@ use crate::tools::error::ToolError;
|
||||
/// **不实现 `Clone`**:透传内层 `LlmError` / `MemoryError`,两者均未派生 `Clone`(保留
|
||||
/// 完整错误信息,传递所有权)。如需在多 session 间共享错误状态,用 `Arc<AgentError>` 包装。
|
||||
#[derive(Debug, Error)]
|
||||
#[non_exhaustive]
|
||||
pub enum AgentError {
|
||||
/// LLM 调用错误(透传 Phase 0)。
|
||||
#[error("LLM 错误: {0}")]
|
||||
@@ -35,6 +36,18 @@ pub enum AgentError {
|
||||
#[error("Plan 解析错误: {0}")]
|
||||
PlanParse(String),
|
||||
|
||||
/// Readonly slot 不允许写入(Phase 10 新增)。
|
||||
#[error("Readonly slot 不允许写入: {0}")]
|
||||
SlotReadonly(String),
|
||||
|
||||
/// Slot 不存在(Phase 10 新增)。
|
||||
#[error("Slot '{0}' 不存在")]
|
||||
SlotNotFound(String),
|
||||
|
||||
/// Slot 已存在(Phase 10 新增)。
|
||||
#[error("Slot '{0}' 已存在")]
|
||||
SlotAlreadyExists(String),
|
||||
|
||||
/// 钩子阻断操作(Agent 层特有)。
|
||||
#[error("钩子阻断: {0}")]
|
||||
HookBlocked(String),
|
||||
@@ -59,6 +72,7 @@ impl AgentError {
|
||||
/// - `Tool`:由内层 `is_recoverable()` 决定
|
||||
/// - `HookBlocked` / `LimitExceeded`:不可恢复(需人工介入或终止循环)
|
||||
/// - `Config` / `Other`:不可恢复
|
||||
/// - `SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`:不可恢复(结构性错误)
|
||||
pub fn is_recoverable(&self) -> bool {
|
||||
match self {
|
||||
Self::Llm(e) => matches!(
|
||||
@@ -68,9 +82,13 @@ impl AgentError {
|
||||
Self::Tool(e) => e.is_recoverable(),
|
||||
Self::Memory(e) => e.is_recoverable(),
|
||||
Self::PlanParse(_) => false,
|
||||
Self::HookBlocked(_) | Self::LimitExceeded(_) | Self::Config(_) | Self::Other(_) => {
|
||||
false
|
||||
}
|
||||
Self::SlotReadonly(_)
|
||||
| Self::SlotNotFound(_)
|
||||
| Self::SlotAlreadyExists(_)
|
||||
| Self::HookBlocked(_)
|
||||
| Self::LimitExceeded(_)
|
||||
| Self::Config(_)
|
||||
| Self::Other(_) => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -180,4 +198,37 @@ mod tests {
|
||||
let err = caller().unwrap_err();
|
||||
assert!(matches!(err, AgentError::Memory(_)));
|
||||
}
|
||||
|
||||
// ====== Phase 10: Slot 错误变体测试 ======
|
||||
|
||||
#[test]
|
||||
fn slot_readonly_not_recoverable() {
|
||||
assert!(!AgentError::SlotReadonly("readonly".into()).is_recoverable());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn slot_not_found_not_recoverable() {
|
||||
assert!(!AgentError::SlotNotFound("missing".into()).is_recoverable());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn slot_already_exists_not_recoverable() {
|
||||
assert!(!AgentError::SlotAlreadyExists("dup".into()).is_recoverable());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn slot_error_messages() {
|
||||
assert_eq!(
|
||||
format!("{}", AgentError::SlotReadonly("readonly".into())),
|
||||
"Readonly slot 不允许写入: readonly"
|
||||
);
|
||||
assert_eq!(
|
||||
format!("{}", AgentError::SlotNotFound("foo".into())),
|
||||
"Slot 'foo' 不存在"
|
||||
);
|
||||
assert_eq!(
|
||||
format!("{}", AgentError::SlotAlreadyExists("bar".into())),
|
||||
"Slot 'bar' 已存在"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,8 +16,8 @@ use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::llm::compact::CompactConfig;
|
||||
use crate::llm::provider::LlmProvider;
|
||||
use crate::llm::hooks::HookExecutor;
|
||||
use crate::llm::provider::LlmProvider;
|
||||
use crate::memory::retriever::MemoryRetriever;
|
||||
use crate::memory::store::MemoryStore;
|
||||
use crate::tools::ToolRegistry;
|
||||
|
||||
+746
-102
@@ -1,29 +1,41 @@
|
||||
//! AgentSession —— 智能体"会话"实例。
|
||||
//!
|
||||
//! 设计要点(参见 `docs/7-agent-runtime.md` §3.2.3):
|
||||
//! 设计要点(参见 `docs/7-agent-runtime.md` §3.2.3 与 `docs/17-phase10-contextslot.md`):
|
||||
//!
|
||||
//! - **会话 = 角色 + 状态**:绑定 `session_id` / `agent` / `bundle`,累计 `turn_index` 和 `cost_so_far`
|
||||
//! - **多上下文分区**(Phase 10):通过 `ContextSlot` 管理多个独立的消息上下文
|
||||
//! - **最小 reference impl**:`submit_turn` 演示"组装 LlmCycle → submit_with_tools → 累计 cost"的标准流程
|
||||
//! - **不做业务循环**:多轮策略、错误重试、记忆回写由上层应用或具体 `TaskAgent` 决定
|
||||
//! - **不持有 ConversationMemory**:上层可独立 new 一个 `ConversationMemory`,在合适的时机调 `add_message`
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
|
||||
use futures_core::Stream;
|
||||
|
||||
use crate::agent::agent::Agent;
|
||||
use crate::agent::context::{
|
||||
ContextSlot, DeriveStrategy, SlotConfig, SlotMode, SlotSource,
|
||||
};
|
||||
use crate::agent::error::AgentError;
|
||||
use crate::agent::runtime::RuntimeBundle;
|
||||
use crate::agent::session_memory::SessionMemory;
|
||||
use crate::llm::cycle::{CostTracker, CycleConfig, LlmCycle};
|
||||
use crate::llm::hooks::{HookContext, HookEvent};
|
||||
use crate::llm::stream::StreamEvent;
|
||||
use crate::llm::types::message::Message;
|
||||
use crate::llm::types::response_v2::MessageResponse;
|
||||
use crate::memory::store::InMemoryStore;
|
||||
use crate::memory::store::{InMemoryStore, MemoryStore};
|
||||
|
||||
/// Agent 会话实例。
|
||||
///
|
||||
/// 同一 `Agent` 可被多个 `AgentSession` 复用(不同 session_id 互不干扰)。
|
||||
/// `submit_turn` 一次只跑一轮 LLM 调用(含自动 tool 循环)。
|
||||
///
|
||||
/// **Phase 10 新增**:通过 `slots: HashMap<String, ContextSlot>` 管理多个独立的对话上下文。
|
||||
/// `submit_turn` 默认写入当前活跃 slot(`current_slot_id`)。
|
||||
///
|
||||
/// **不实现 `Clone`**:session 持有累计 `turn_index` / `cost_so_far` / `session_memory`,
|
||||
/// 共享这些状态需要显式 sync 语义;如果上层需要并发访问,自己用 `Arc<Mutex<_>>` 包装。
|
||||
pub struct AgentSession {
|
||||
@@ -36,6 +48,10 @@ pub struct AgentSession {
|
||||
cost_so_far: CostTracker,
|
||||
/// 会话级记忆(Phase 4c 替换内联 HashMap)。
|
||||
pub session_memory: SessionMemory,
|
||||
/// Phase 10 新增:所有 slot(id → ContextSlot)。
|
||||
slots: HashMap<String, ContextSlot>,
|
||||
/// Phase 10 新增:当前活跃 slot 的 id。
|
||||
current_slot_id: String,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for AgentSession {
|
||||
@@ -46,6 +62,8 @@ impl std::fmt::Debug for AgentSession {
|
||||
.field("turn_index", &self.turn_index)
|
||||
.field("cost_so_far", &self.cost_so_far.total())
|
||||
.field("session_memory", &"<SessionMemory>")
|
||||
.field("slots", &self.slots.keys().collect::<Vec<_>>())
|
||||
.field("current_slot_id", &self.current_slot_id)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
@@ -54,6 +72,14 @@ impl AgentSession {
|
||||
/// 创建一个新的会话实例。
|
||||
///
|
||||
/// `agent` 与 `bundle` 共同决定 `submit_turn` 行为:system_prompt / 工具集 / LLM 后端均来自它们。
|
||||
///
|
||||
/// Phase 10 新增:自动创建 `"default"` slot,确保简单场景无感使用。
|
||||
///
|
||||
/// **注意**:`new()` 是同步函数,无法执行异步的 `ContextSlot::load()`。
|
||||
/// 因此始终创建空的 default slot。若需要从存储恢复 session 历史,
|
||||
/// 可在创建后调用 `switch_slot("default")` 尝试从存储加载。
|
||||
/// (v0.2 简化:`switch_slot` 在 HashMap 中已有 key 时不会重载——如需恢复,
|
||||
/// 请在清空 `slots` 后调用 `switch_slot`,或等待 v0.3 的懒加载支持。)
|
||||
pub fn new(
|
||||
agent: Arc<dyn Agent>,
|
||||
session_id: impl Into<String>,
|
||||
@@ -65,6 +91,16 @@ impl AgentSession {
|
||||
.clone()
|
||||
.unwrap_or_else(|| Arc::new(InMemoryStore::new()));
|
||||
let session_memory = SessionMemory::new(backend, &session_id_str);
|
||||
|
||||
// 自动创建 "default" slot
|
||||
let default_slot = ContextSlot::new(
|
||||
&session_id_str,
|
||||
"default",
|
||||
SlotConfig::default(),
|
||||
);
|
||||
let mut slots = HashMap::new();
|
||||
slots.insert("default".to_string(), default_slot);
|
||||
|
||||
Self {
|
||||
session_id: session_id_str,
|
||||
agent,
|
||||
@@ -72,6 +108,8 @@ impl AgentSession {
|
||||
turn_index: 0,
|
||||
cost_so_far: CostTracker::default(),
|
||||
session_memory,
|
||||
slots,
|
||||
current_slot_id: "default".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -96,7 +134,9 @@ impl AgentSession {
|
||||
key: impl Into<String>,
|
||||
value: impl Into<String>,
|
||||
) -> Result<(), AgentError> {
|
||||
self.session_memory.set(&key.into(), &value.into()).await
|
||||
self.session_memory
|
||||
.set(&key.into(), &value.into())
|
||||
.await
|
||||
}
|
||||
|
||||
/// 读取一条会话级数据。
|
||||
@@ -104,20 +144,158 @@ impl AgentSession {
|
||||
self.session_memory.get(key).await
|
||||
}
|
||||
|
||||
/// Phase 10: 当前 slot id。
|
||||
pub fn current_slot_id(&self) -> &str {
|
||||
&self.current_slot_id
|
||||
}
|
||||
|
||||
/// Phase 10: 列出所有 slot 的不可变引用(按 id 顺序)。
|
||||
pub fn slots(&self) -> impl Iterator<Item = (&String, &ContextSlot)> {
|
||||
self.slots.iter()
|
||||
}
|
||||
|
||||
/// Phase 10: 解析存储后端。
|
||||
/// fallback 链:`session_memory_backend` → `memory_store` → `InMemoryStore`。
|
||||
fn resolve_store(&self) -> Arc<dyn MemoryStore> {
|
||||
self.bundle
|
||||
.session_memory_backend
|
||||
.clone()
|
||||
.or_else(|| self.bundle.memory_store.clone())
|
||||
.unwrap_or_else(|| Arc::new(InMemoryStore::new()))
|
||||
}
|
||||
|
||||
// ====== Phase 10: Slot 管理方法 ======
|
||||
|
||||
/// 创建新 slot(config 可选,不传则使用默认值)。
|
||||
pub async fn create_slot(
|
||||
&mut self,
|
||||
id: impl Into<String>,
|
||||
config: Option<SlotConfig>,
|
||||
) -> Result<(), AgentError> {
|
||||
let id = id.into();
|
||||
if self.slots.contains_key(&id) {
|
||||
return Err(AgentError::SlotAlreadyExists(id));
|
||||
}
|
||||
let slot = ContextSlot::new(
|
||||
&self.session_id,
|
||||
&id,
|
||||
config.unwrap_or_default(),
|
||||
);
|
||||
slot.save(&*self.resolve_store()).await?;
|
||||
self.slots.insert(id, slot);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 切换到指定 slot。
|
||||
/// - 如果 slot 已在内存中,直接切换 current_slot_id
|
||||
/// - 如果不在内存中,尝试从存储加载(config 自动从 slot_config key 恢复)
|
||||
/// - 存储中也不存在则返回 `SlotNotFound`
|
||||
pub async fn switch_slot(&mut self, id: &str) -> Result<(), AgentError> {
|
||||
if !self.slots.contains_key(id) {
|
||||
let store = self.resolve_store();
|
||||
match ContextSlot::load(id, &self.session_id, &*store).await? {
|
||||
Some(slot) => {
|
||||
self.slots.insert(id.to_string(), slot);
|
||||
}
|
||||
None => return Err(AgentError::SlotNotFound(id.to_string())),
|
||||
}
|
||||
}
|
||||
self.current_slot_id = id.to_string();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 列出所有 slot id。
|
||||
pub fn list_slots(&self) -> impl Iterator<Item = &String> {
|
||||
self.slots.keys()
|
||||
}
|
||||
|
||||
/// 从父 slot 派生新 slot(继承父 slot 的全量或聚焦消息)。
|
||||
pub async fn derive_slot(
|
||||
&mut self,
|
||||
id: impl Into<String>,
|
||||
parent_id: &str,
|
||||
strategy: DeriveStrategy,
|
||||
) -> Result<(), AgentError> {
|
||||
let slot_id = id.into();
|
||||
if self.slots.contains_key(&slot_id) {
|
||||
return Err(AgentError::SlotAlreadyExists(slot_id));
|
||||
}
|
||||
let parent = self
|
||||
.slots
|
||||
.get(parent_id)
|
||||
.ok_or_else(|| AgentError::SlotNotFound(parent_id.to_string()))?;
|
||||
|
||||
let parent_messages = parent.messages.clone();
|
||||
let (messages, focused_cfg) = match &strategy {
|
||||
DeriveStrategy::Full => (parent_messages, None),
|
||||
DeriveStrategy::Focused(cfg) => {
|
||||
let filtered = ContextSlot::filter_focused(&parent_messages, cfg);
|
||||
(filtered, Some(cfg.clone()))
|
||||
}
|
||||
};
|
||||
|
||||
let mode = match focused_cfg {
|
||||
Some(cfg) => SlotMode::Focused(cfg),
|
||||
None => SlotMode::Full,
|
||||
};
|
||||
|
||||
let slot = ContextSlot::new(
|
||||
&self.session_id,
|
||||
&slot_id,
|
||||
SlotConfig {
|
||||
mode,
|
||||
source: SlotSource::Derived {
|
||||
parent_id: parent_id.to_string(),
|
||||
strategy,
|
||||
},
|
||||
budget: Default::default(),
|
||||
compact: true,
|
||||
},
|
||||
);
|
||||
let mut slot = slot;
|
||||
slot.messages = messages;
|
||||
slot.save(&*self.resolve_store()).await?;
|
||||
self.slots.insert(slot_id, slot);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 删除一个 slot。
|
||||
/// - 禁止删除 `"default"` slot
|
||||
/// - 至少保留一个 slot
|
||||
/// - 已删除后再 load 返回 None
|
||||
pub async fn delete_slot(&mut self, id: &str) -> Result<(), AgentError> {
|
||||
if id == "default" {
|
||||
return Err(AgentError::Config("Cannot delete the 'default' slot".into()));
|
||||
}
|
||||
if self.slots.len() <= 1 {
|
||||
return Err(AgentError::Config("Cannot delete the last slot".into()));
|
||||
}
|
||||
ContextSlot::delete(id, &self.session_id, &*self.resolve_store()).await?;
|
||||
self.slots.remove(id);
|
||||
if self.current_slot_id == id {
|
||||
self.current_slot_id = "default".to_string();
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ====== 原始方法 ======
|
||||
|
||||
/// 提交一轮对话(含自动 tool 循环),返回 LLM 响应。
|
||||
///
|
||||
/// 流程:
|
||||
/// 1. 触发 `OnTurnStart` hook
|
||||
/// 2. 组装 `LlmCycle`(注入 system_prompt / hook_executor / compact_config / 消息历史)
|
||||
/// 3. `submit_with_tools` 跑单轮对话
|
||||
/// 4. 累计 `cost_so_far`
|
||||
/// 5. 触发 `OnTurnEnd` hook
|
||||
/// 6. `turn_index += 1`
|
||||
/// Phase 10 改造:
|
||||
/// - 从当前 slot 加载历史(Focused 模式读时过滤)
|
||||
/// - 提交完成后**增量追加**本轮新增消息到当前 slot(不覆盖,确保 Focused 语义不丢数据)
|
||||
///
|
||||
/// **不做**:
|
||||
/// - 不持有 `ConversationMemory`(由上层独立 task 决定何时回写)
|
||||
/// - 不做 Plan 拆解(Phase 4b 才加 `TaskAgent`)
|
||||
/// - 不做 session_data 持久化(Phase 4c 替换为 `SessionMemory`)
|
||||
/// 流程:
|
||||
/// 1. 检查当前 slot 不是 Readonly
|
||||
/// 2. 触发 `OnTurnStart` hook
|
||||
/// 3. 加载当前 slot 的历史消息
|
||||
/// 4. 组装 `LlmCycle`(注入 system_prompt / compact_config / 历史)
|
||||
/// 5. `submit_with_tools` 跑单轮对话
|
||||
/// 6. 累计 `cost_so_far`
|
||||
/// 7. **增量追加**本轮新增消息到当前 slot + 保存到 store
|
||||
/// 8. 触发 `OnTurnEnd` hook
|
||||
/// 9. `turn_index += 1`
|
||||
pub async fn submit_turn(
|
||||
&mut self,
|
||||
user_input: impl Into<String>,
|
||||
@@ -125,50 +303,196 @@ impl AgentSession {
|
||||
let turn_index = self.turn_index;
|
||||
let hook_executor = Arc::clone(&self.bundle.hook_executor);
|
||||
|
||||
// 0. Readonly slot 拒绝 submit_turn
|
||||
{
|
||||
let slot = self
|
||||
.slots
|
||||
.get(&self.current_slot_id)
|
||||
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?;
|
||||
if matches!(slot.config.mode, SlotMode::Readonly) {
|
||||
return Err(AgentError::SlotReadonly(format!(
|
||||
"Cannot submit turn on Readonly slot '{}'",
|
||||
self.current_slot_id
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
// 1. 触发 OnTurnStart hook
|
||||
let start_ctx =
|
||||
HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
|
||||
let start_ctx = HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
|
||||
hook_executor
|
||||
.execute(HookEvent::OnTurnStart, &start_ctx)
|
||||
.await;
|
||||
|
||||
// 2. 组装 LlmCycle —— 共享 bundle 中的 provider 句柄
|
||||
// 工具列表从 agent.tool_definitions(bundle) 派生(默认 = bundle 全量);
|
||||
// submit_with_tools 内部从 registry 自行取 definitions,此处仅消费以触发
|
||||
// 子 trait 覆盖(白名单/过滤)的副作用。
|
||||
// 2. 从当前 slot 加载历史消息
|
||||
let history = self
|
||||
.slots
|
||||
.get(&self.current_slot_id)
|
||||
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?
|
||||
.load_messages();
|
||||
|
||||
// 3. 组装 LlmCycle
|
||||
let _ = self.agent.tool_definitions(&self.bundle);
|
||||
let mut cycle = LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default())
|
||||
.with_messages(Vec::new());
|
||||
// Phase 2 切换 system_prompt 字段为 Message::System(FIX-D)。
|
||||
// 若 agent 自带 system prompt,预置到 messages 列表头部。
|
||||
let mut initial_messages: Vec<Message> = Vec::new();
|
||||
let mut cycle =
|
||||
LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default());
|
||||
let mut messages_with_prompt = history;
|
||||
if let Some(prompt) = self.agent.system_prompt() {
|
||||
initial_messages.push(Message::system(prompt));
|
||||
}
|
||||
if !initial_messages.is_empty() {
|
||||
cycle = cycle.with_messages(initial_messages);
|
||||
messages_with_prompt.insert(0, Message::system(prompt));
|
||||
}
|
||||
let input_len = messages_with_prompt.len();
|
||||
cycle = cycle.with_messages(messages_with_prompt);
|
||||
if let Some(cfg) = self.bundle.config.compact_config.clone() {
|
||||
cycle = cycle.with_compact_config(cfg);
|
||||
}
|
||||
|
||||
// 3. 提交(HookExecutor 不在这里传——内部 hook 由 LlmCycle 在 PreRequest/PostRequest 触发)
|
||||
// 4. 提交
|
||||
let response = cycle
|
||||
.submit_with_tools(user_input.into(), &self.bundle.tool_registry)
|
||||
.await?;
|
||||
|
||||
// 4. 累计 cost
|
||||
// 5. 累计 cost
|
||||
self.cost_so_far.add(&response.usage);
|
||||
|
||||
// 5. 触发 OnTurnEnd hook
|
||||
// 6. 只将本轮新增消息追加到当前 slot(保留全量历史,确保 Focused 模式的"读时过滤"语义不丢失数据)
|
||||
// cycle.messages() 包含 [system_prompt?, history..., user_input, tool_calls..., final_response]
|
||||
// 新增消息 = cycle.messages()[input_len..](跳过 initial_messages,即跳过已被持久化的内容)
|
||||
let new_messages: Vec<Message> = cycle
|
||||
.messages()
|
||||
.iter()
|
||||
.skip(input_len)
|
||||
.cloned()
|
||||
.collect();
|
||||
let store = self.resolve_store();
|
||||
if let Some(slot) = self.slots.get_mut(&self.current_slot_id) {
|
||||
slot.append_messages(new_messages)?;
|
||||
slot.save(&*store).await?;
|
||||
}
|
||||
|
||||
// 7. 触发 OnTurnEnd hook
|
||||
let end_ctx = HookContext::new(HookEvent::OnTurnEnd).with_turn_index(turn_index);
|
||||
hook_executor.execute(HookEvent::OnTurnEnd, &end_ctx).await;
|
||||
|
||||
// 6. turn_index 递增
|
||||
// 8. turn_index 递增
|
||||
self.turn_index += 1;
|
||||
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
/// 提交一轮对话(流式版本,含自动 tool 循环),返回 `StreamEvent` 流。
|
||||
///
|
||||
/// Phase 10 改造:
|
||||
/// - 从当前 slot 加载历史(Focused 模式读时过滤)
|
||||
/// - finalize_turn 需要传入本轮新增消息列表
|
||||
///
|
||||
/// **运行时要求**:内部委托 `submit_with_tools_stream`,需要 tokio 多线程运行时。
|
||||
pub async fn submit_turn_stream(
|
||||
&mut self,
|
||||
user_input: impl Into<String>,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, AgentError> {
|
||||
let turn_index = self.turn_index;
|
||||
let hook_executor = Arc::clone(&self.bundle.hook_executor);
|
||||
|
||||
// 0. Readonly 检查
|
||||
{
|
||||
let slot = self
|
||||
.slots
|
||||
.get(&self.current_slot_id)
|
||||
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?;
|
||||
if matches!(slot.config.mode, SlotMode::Readonly) {
|
||||
return Err(AgentError::SlotReadonly(format!(
|
||||
"Cannot submit turn stream on Readonly slot '{}'",
|
||||
self.current_slot_id
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
// 1. 触发 OnTurnStart hook
|
||||
let start_ctx = HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
|
||||
hook_executor
|
||||
.execute(HookEvent::OnTurnStart, &start_ctx)
|
||||
.await;
|
||||
|
||||
// 2. 触发子 trait 覆盖(白名单/过滤)的副作用
|
||||
let _ = self.agent.tool_definitions(&self.bundle);
|
||||
|
||||
// 3. 从当前 slot 加载历史
|
||||
let history = self
|
||||
.slots
|
||||
.get(&self.current_slot_id)
|
||||
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?
|
||||
.load_messages();
|
||||
|
||||
// 4. 组装 LlmCycle
|
||||
let mut cycle =
|
||||
LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default());
|
||||
let mut messages_with_prompt = history;
|
||||
if let Some(prompt) = self.agent.system_prompt() {
|
||||
messages_with_prompt.insert(0, Message::system(prompt));
|
||||
}
|
||||
cycle = cycle.with_messages(messages_with_prompt);
|
||||
if let Some(cfg) = self.bundle.config.compact_config.clone() {
|
||||
cycle = cycle.with_compact_config(cfg);
|
||||
}
|
||||
|
||||
// 5. 调用流式工具循环
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream(
|
||||
user_input.into(),
|
||||
Arc::clone(&self.bundle.tool_registry),
|
||||
)
|
||||
.await?;
|
||||
|
||||
// 6. turn_index 递增 —— 配合 finalize_turn 用 (turn_index - 1) 传递正确的 OnTurnEnd 序号
|
||||
self.turn_index += 1;
|
||||
|
||||
// 注:hook_executor 不显式 drop,生命周期由 Arc 自动管理
|
||||
let _ = hook_executor;
|
||||
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
/// 完成一轮 turn:累计 cost + 触发 OnTurnEnd hook + 增量追加消息到当前 slot。
|
||||
///
|
||||
/// Phase 10 改造:
|
||||
/// - 新增 `new_messages_from_cycle` 参数:流式场景下,本轮新增的消息列表
|
||||
/// (由消费者在流消费完毕后从 `cycle.messages()[input_len..]` 获取并传入)
|
||||
/// - 仅**增量追加**到当前 slot(不覆盖已有消息),与 submit_turn 行为一致
|
||||
/// - 返回类型从 `()` 改为 `Result<(), AgentError>`,错误传播更清晰
|
||||
///
|
||||
/// 由消费者在收到 `MessageComplete.full_response` 后调用。
|
||||
pub async fn finalize_turn(
|
||||
&mut self,
|
||||
response: &MessageResponse,
|
||||
new_messages_from_cycle: Vec<Message>,
|
||||
) -> Result<(), AgentError> {
|
||||
self.cost_so_far.add(&response.usage);
|
||||
|
||||
// 防御性检查:current_slot_id 必须在 slots 中(与 submit_turn 行为一致)。
|
||||
// 正常流程:submit_turn_stream 已注册 slot,finalize_turn 不应触发此分支。
|
||||
if !self.slots.contains_key(&self.current_slot_id) {
|
||||
return Err(AgentError::SlotNotFound(self.current_slot_id.clone()));
|
||||
}
|
||||
// 增量追加到当前 slot(仅 Full/Focused 模式允许,Readonly 阻断)
|
||||
let store = self.resolve_store();
|
||||
if let Some(slot) = self.slots.get_mut(&self.current_slot_id) {
|
||||
if matches!(slot.config.mode, SlotMode::Readonly) {
|
||||
return Err(AgentError::SlotReadonly(format!(
|
||||
"Cannot finalize turn on Readonly slot '{}'",
|
||||
self.current_slot_id
|
||||
)));
|
||||
}
|
||||
slot.append_messages(new_messages_from_cycle)?;
|
||||
slot.save(&*store).await?;
|
||||
}
|
||||
|
||||
// 防御性 saturating_sub 防止误用 panic。
|
||||
let end_ctx = HookContext::new(HookEvent::OnTurnEnd)
|
||||
.with_turn_index(self.turn_index.saturating_sub(1));
|
||||
self.bundle
|
||||
.hook_executor
|
||||
.execute(HookEvent::OnTurnEnd, &end_ctx)
|
||||
.await;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -223,9 +547,7 @@ mod tests {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text {
|
||||
text: text.into(),
|
||||
}],
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: crate::llm::types::Usage::from_input_output(10, 5),
|
||||
stop_reason: StopReason::Stop,
|
||||
@@ -233,65 +555,7 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
/// 烟雾测试 1:AgentSession::submit_turn 跑通 mock provider。
|
||||
#[tokio::test]
|
||||
async fn submit_turn_runs_with_mock_provider() {
|
||||
let provider = Arc::new(MockProvider::new(vec![assistant_text("hello back")]));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: Some("you are a test agent".into()),
|
||||
});
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let mut session = AgentSession::new(agent, "s1", bundle);
|
||||
assert_eq!(session.turn_index(), 0);
|
||||
|
||||
let response = session.submit_turn("hi").await.unwrap();
|
||||
assert_eq!(response.text(), "hello back");
|
||||
assert_eq!(session.turn_index(), 1);
|
||||
assert_eq!(session.usage().total().prompt_tokens, 10);
|
||||
assert_eq!(session.usage().total().completion_tokens, 5);
|
||||
}
|
||||
|
||||
/// 烟雾测试 2:session_data 读写。
|
||||
#[tokio::test]
|
||||
async fn session_data_set_get() {
|
||||
let provider = Arc::new(MockProvider::new(vec![]));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: None,
|
||||
});
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
let mut session = AgentSession::new(agent, "s2", bundle);
|
||||
|
||||
assert!(session.get_session_data("k").await.unwrap().is_none());
|
||||
session.set_session_data("k", "v").await.unwrap();
|
||||
assert_eq!(session.get_session_data("k").await.unwrap(), Some("v".into()));
|
||||
// 覆盖写
|
||||
session.set_session_data("k", "v2").await.unwrap();
|
||||
assert_eq!(
|
||||
session.get_session_data("k").await.unwrap(),
|
||||
Some("v2".into())
|
||||
);
|
||||
}
|
||||
|
||||
/// 烟雾测试 3:submit_turn 触发 OnTurnStart / OnTurnEnd hook。
|
||||
#[tokio::test]
|
||||
async fn submit_turn_triggers_turn_hooks() {
|
||||
fn build_session(provider_responses: Vec<MessageResponse>) -> (AgentSession, Arc<CountHook>, Arc<CountHook>) {
|
||||
let mut hook_executor = HookExecutor::new();
|
||||
let start_count = Arc::new(CountHook(AtomicU32::new(0)));
|
||||
let end_count = Arc::new(CountHook(AtomicU32::new(0)));
|
||||
@@ -304,13 +568,10 @@ mod tests {
|
||||
Box::new(CountHookAdapter(end_count.clone())),
|
||||
);
|
||||
|
||||
let provider = Arc::new(MockProvider::new(vec![
|
||||
assistant_text("ok"),
|
||||
assistant_text("ok 2"),
|
||||
]));
|
||||
let provider = Arc::new(MockProvider::new(provider_responses));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: None,
|
||||
prompt: Some("you are a test agent".into()),
|
||||
});
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
@@ -320,7 +581,53 @@ mod tests {
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
let mut session = AgentSession::new(agent, "s3", bundle);
|
||||
|
||||
let session = AgentSession::new(agent, "test-session", bundle);
|
||||
(session, start_count, end_count)
|
||||
}
|
||||
|
||||
/// 烟雾测试 1:AgentSession::submit_turn 跑通 mock provider(向后兼容)。
|
||||
#[tokio::test]
|
||||
async fn submit_turn_runs_with_mock_provider() {
|
||||
let (mut session, start_count, end_count) = build_session(vec![assistant_text("hello back")]);
|
||||
assert_eq!(session.turn_index(), 0);
|
||||
|
||||
let response = session.submit_turn("hi").await.unwrap();
|
||||
assert_eq!(extract_text(&response.message), "hello back");
|
||||
assert_eq!(session.turn_index(), 1);
|
||||
assert_eq!(session.usage().total().prompt_tokens, 10);
|
||||
assert_eq!(session.usage().total().completion_tokens, 5);
|
||||
|
||||
// hook 触发
|
||||
assert_eq!(start_count.0.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(end_count.0.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
/// 烟雾测试 2:session_data 读写。
|
||||
#[tokio::test]
|
||||
async fn session_data_set_get() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
assert!(session.get_session_data("k").await.unwrap().is_none());
|
||||
session.set_session_data("k", "v").await.unwrap();
|
||||
assert_eq!(
|
||||
session.get_session_data("k").await.unwrap(),
|
||||
Some("v".into())
|
||||
);
|
||||
// 覆盖写
|
||||
session.set_session_data("k", "v2").await.unwrap();
|
||||
assert_eq!(
|
||||
session.get_session_data("k").await.unwrap(),
|
||||
Some("v2".into())
|
||||
);
|
||||
}
|
||||
|
||||
/// 烟雾测试 3:submit_turn 触发 OnTurnStart / OnTurnEnd hook。
|
||||
#[tokio::test]
|
||||
async fn submit_turn_triggers_turn_hooks() {
|
||||
let (mut session, start_count, end_count) = build_session(vec![
|
||||
assistant_text("ok"),
|
||||
assistant_text("ok 2"),
|
||||
]);
|
||||
|
||||
session.submit_turn("hi").await.unwrap();
|
||||
assert_eq!(start_count.0.load(Ordering::SeqCst), 1);
|
||||
@@ -330,4 +637,341 @@ mod tests {
|
||||
assert_eq!(start_count.0.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(end_count.0.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
}
|
||||
|
||||
/// 提取 Message 的第一个 Text block(测试辅助)。
|
||||
fn extract_text(msg: &Message) -> &str {
|
||||
use crate::llm::types::message::ContentBlock;
|
||||
let blocks = match msg {
|
||||
Message::System { content }
|
||||
| Message::User { content }
|
||||
| Message::Assistant { content } => content,
|
||||
Message::UserImage { .. } => return "",
|
||||
Message::ToolResult { content, .. } => content,
|
||||
};
|
||||
for block in blocks {
|
||||
if let ContentBlock::Text { text } = block {
|
||||
return text;
|
||||
}
|
||||
}
|
||||
""
|
||||
}
|
||||
|
||||
// ====== Phase 10 新增测试 ======
|
||||
|
||||
/// Phase 10: 默认 slot 自动创建。
|
||||
#[tokio::test]
|
||||
async fn default_slot_auto_created() {
|
||||
let (session, _, _) = build_session(vec![]);
|
||||
assert_eq!(session.current_slot_id(), "default");
|
||||
let slots: Vec<_> = session.list_slots().collect();
|
||||
assert_eq!(slots.len(), 1);
|
||||
assert!(slots.contains(&&"default".to_string()));
|
||||
}
|
||||
|
||||
/// Phase 10: submit_turn 写入当前 slot。
|
||||
#[tokio::test]
|
||||
async fn submit_turn_writes_to_current_slot() {
|
||||
let (mut session, _, _) = build_session(vec![assistant_text("resp")]);
|
||||
session.submit_turn("user input").await.unwrap();
|
||||
|
||||
// 检查 default slot 内存中的消息
|
||||
let slot = session.slots.get("default").expect("default slot exists");
|
||||
// submit_turn 增量追加的是 cycle.messages()[input_len..] 部分,
|
||||
// 即 [user_input, tool_results?, final_response](不含 system_prompt,system 由 agent 提供)
|
||||
assert!(slot.messages.len() >= 2, "应至少包含 user 和 assistant");
|
||||
// 验证 user 输入和 assistant 响应都已写入
|
||||
let has_user = slot.messages.iter().any(|m| extract_text(m) == "user input");
|
||||
let has_resp = slot.messages.iter().any(|m| extract_text(m) == "resp");
|
||||
assert!(has_user && has_resp, "slot 应包含 user input 和 assistant response");
|
||||
}
|
||||
|
||||
/// Phase 10: create_slot 创建新 slot。
|
||||
#[tokio::test]
|
||||
async fn create_slot_basic() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
session.create_slot("scratch", None).await.unwrap();
|
||||
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||
assert!(slots.contains(&"default".to_string()));
|
||||
assert!(slots.contains(&"scratch".to_string()));
|
||||
assert_eq!(slots.len(), 2);
|
||||
}
|
||||
|
||||
/// Phase 10: create_slot 拒绝重复 id。
|
||||
#[tokio::test]
|
||||
async fn create_slot_rejects_duplicate() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
session.create_slot("dup", None).await.unwrap();
|
||||
let err = session.create_slot("dup", None).await.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotAlreadyExists(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: switch_slot 切换并保留各自消息。
|
||||
#[tokio::test]
|
||||
async fn switch_slot_isolates_messages() {
|
||||
let (mut session, _, _) = build_session(vec![
|
||||
assistant_text("resp a"),
|
||||
assistant_text("resp b"),
|
||||
assistant_text("resp c"),
|
||||
]);
|
||||
|
||||
// 1. 在 default 中提交一次
|
||||
session.submit_turn("msg in default").await.unwrap();
|
||||
|
||||
// 2. 创建 slot_a
|
||||
session.create_slot("slot_a", None).await.unwrap();
|
||||
session.switch_slot("slot_a").await.unwrap();
|
||||
assert_eq!(session.current_slot_id(), "slot_a");
|
||||
session.submit_turn("msg in slot_a").await.unwrap();
|
||||
|
||||
// 3. 检查 slot_a 的消息数
|
||||
let slot_a = session.slots.get("slot_a").unwrap();
|
||||
let slot_a_count = slot_a.messages.len();
|
||||
assert!(slot_a_count >= 2, "slot_a 至少 2 条消息,实际 {}", slot_a_count);
|
||||
|
||||
// 4. 切回 default,验证 default 不包含 slot_a 的消息
|
||||
session.switch_slot("default").await.unwrap();
|
||||
let slot_default = session.slots.get("default").unwrap();
|
||||
let default_count = slot_default.messages.len();
|
||||
assert!(default_count >= 2);
|
||||
// 验证 default 中没有 "msg in slot_a"
|
||||
let default_has_a = slot_default
|
||||
.messages
|
||||
.iter()
|
||||
.any(|m| extract_text(m) == "msg in slot_a");
|
||||
assert!(!default_has_a, "default 不应包含 slot_a 的消息");
|
||||
// 验证 slot_a 中没有 "msg in default"
|
||||
let slot_a = session.slots.get("slot_a").unwrap();
|
||||
let a_has_default = slot_a
|
||||
.messages
|
||||
.iter()
|
||||
.any(|m| extract_text(m) == "msg in default");
|
||||
assert!(!a_has_default, "slot_a 不应包含 default 的消息");
|
||||
}
|
||||
|
||||
/// Phase 10: Readonly slot 拒绝写入。
|
||||
#[tokio::test]
|
||||
async fn readonly_slot_rejects_submit_turn() {
|
||||
let (mut session, _, _) = build_session(vec![assistant_text("resp")]);
|
||||
session
|
||||
.create_slot(
|
||||
"ro",
|
||||
Some(SlotConfig {
|
||||
mode: SlotMode::Readonly,
|
||||
source: SlotSource::New,
|
||||
budget: Default::default(),
|
||||
compact: true,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
session.switch_slot("ro").await.unwrap();
|
||||
let err = session.submit_turn("blocked").await.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotReadonly(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: delete_slot 删除非 default。
|
||||
#[tokio::test]
|
||||
async fn delete_slot_removes_non_default() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
session.create_slot("to_delete", None).await.unwrap();
|
||||
session.delete_slot("to_delete").await.unwrap();
|
||||
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||
assert!(!slots.contains(&"to_delete".to_string()));
|
||||
assert_eq!(slots.len(), 1);
|
||||
}
|
||||
|
||||
/// Phase 10: delete_slot 禁止删 default。
|
||||
#[tokio::test]
|
||||
async fn delete_slot_rejects_default() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
let err = session.delete_slot("default").await.unwrap_err();
|
||||
assert!(matches!(err, AgentError::Config(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: delete_slot 禁止删最后一个 slot。
|
||||
#[tokio::test]
|
||||
async fn delete_slot_rejects_last() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
// 只有 default 一个 slot
|
||||
let err = session.delete_slot("default").await.unwrap_err();
|
||||
assert!(matches!(err, AgentError::Config(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: delete_slot 后 current 回退到 default。
|
||||
#[tokio::test]
|
||||
async fn delete_slot_falls_back_to_default() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
session.create_slot("temp", None).await.unwrap();
|
||||
session.switch_slot("temp").await.unwrap();
|
||||
assert_eq!(session.current_slot_id(), "temp");
|
||||
session.delete_slot("temp").await.unwrap();
|
||||
assert_eq!(session.current_slot_id(), "default");
|
||||
}
|
||||
|
||||
/// Phase 10: derive_slot Full 策略复制父 slot 消息。
|
||||
#[tokio::test]
|
||||
async fn derive_slot_full_copies_parent() {
|
||||
let (mut session, _, _) = build_session(vec![assistant_text("resp")]);
|
||||
session.submit_turn("parent msg").await.unwrap();
|
||||
session
|
||||
.derive_slot("child", "default", DeriveStrategy::Full)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let child = session.slots.get("child").unwrap();
|
||||
assert!(matches!(child.config.source, SlotSource::Derived { .. }));
|
||||
// child 应有 parent 的消息拷贝
|
||||
let has_parent = child
|
||||
.messages
|
||||
.iter()
|
||||
.any(|m| extract_text(m) == "parent msg");
|
||||
assert!(has_parent);
|
||||
}
|
||||
|
||||
/// Phase 10: derive_slot 拒绝重复 id。
|
||||
#[tokio::test]
|
||||
async fn derive_slot_rejects_duplicate() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
session.create_slot("child", None).await.unwrap();
|
||||
let err = session
|
||||
.derive_slot("child", "default", DeriveStrategy::Full)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotAlreadyExists(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: derive_slot 父 slot 不存在返回 SlotNotFound。
|
||||
#[tokio::test]
|
||||
async fn derive_slot_parent_not_found() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
let err = session
|
||||
.derive_slot("child", "nonexistent", DeriveStrategy::Full)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotNotFound(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: switch_slot 加载不存在的 slot 返回 SlotNotFound。
|
||||
#[tokio::test]
|
||||
async fn switch_slot_not_found() {
|
||||
let (mut session, _, _) = build_session(vec![]);
|
||||
let err = session.switch_slot("missing").await.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotNotFound(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: slot 数据持久化到 storage,switch 时可恢复。
|
||||
/// 使用 session_memory_backend 配置可验证持久化。
|
||||
#[tokio::test]
|
||||
async fn slot_persistence_roundtrip() {
|
||||
// 创建一个共享的 InMemoryStore 作为后端
|
||||
let backend = Arc::new(InMemoryStore::new());
|
||||
|
||||
let provider = Arc::new(MockProvider::new(vec![assistant_text("resp")]));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: None,
|
||||
});
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.session_memory_backend(backend.clone())
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let mut session = AgentSession::new(agent, "persist-session", bundle);
|
||||
session.create_slot("persist_test", None).await.unwrap();
|
||||
session.switch_slot("persist_test").await.unwrap();
|
||||
session.submit_turn("hi").await.unwrap();
|
||||
|
||||
// 验证 data/meta/config 三个 key 都已写入共享 backend
|
||||
let stored_data = backend
|
||||
.get(&ContextSlot::data_key("persist-session", "persist_test"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(stored_data.is_some(), "data 应已持久化");
|
||||
|
||||
let stored_meta = backend
|
||||
.get(&ContextSlot::meta_key("persist-session", "persist_test"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(stored_meta.is_some(), "meta 应已持久化");
|
||||
|
||||
let stored_config = backend
|
||||
.get(&ContextSlot::config_key("persist-session", "persist_test"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(stored_config.is_some(), "config 应已持久化");
|
||||
}
|
||||
|
||||
/// Phase 10: 当只有 `memory_store`(无 `session_memory_backend`)时,resolve_store
|
||||
/// 应 fallback 到 `memory_store`。
|
||||
#[tokio::test]
|
||||
async fn resolve_store_falls_back_to_memory_store() {
|
||||
let backend = Arc::new(InMemoryStore::new());
|
||||
|
||||
let provider = Arc::new(MockProvider::new(vec![assistant_text("ok")]));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: None,
|
||||
});
|
||||
// 注意:这里只设置 memory_store,不设置 session_memory_backend
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.memory_store(backend.clone())
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let mut session = AgentSession::new(agent, "fb-session", bundle);
|
||||
session.create_slot("fb_slot", None).await.unwrap();
|
||||
|
||||
// 验证 backend 中已存在 fb_slot 的数据
|
||||
let stored = backend
|
||||
.get(&ContextSlot::data_key("fb-session", "fb_slot"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(stored.is_some(), "memory_store fallback 应生效");
|
||||
}
|
||||
|
||||
/// Phase 10: finalize_turn 在 current_slot 不存在时返回 SlotNotFound(与 submit_turn 一致)。
|
||||
#[tokio::test]
|
||||
async fn finalize_turn_slot_not_found() {
|
||||
let (mut session, _, _) = build_session(vec![assistant_text("ok")]);
|
||||
// 强制 current_slot_id 指向不存在的 slot(模拟异常状态)
|
||||
session.current_slot_id = "ghost".to_string();
|
||||
let response = assistant_text("ok");
|
||||
let err = session.finalize_turn(&response, vec![]).await.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotNotFound(_)));
|
||||
}
|
||||
|
||||
/// Phase 10: finalize_turn 在 Readonly slot 上返回 SlotReadonly。
|
||||
#[tokio::test]
|
||||
async fn finalize_turn_readonly_rejects() {
|
||||
let (mut session, _, _) = build_session(vec![assistant_text("ok")]);
|
||||
session
|
||||
.create_slot(
|
||||
"ro",
|
||||
Some(SlotConfig {
|
||||
mode: SlotMode::Readonly,
|
||||
source: SlotSource::New,
|
||||
budget: Default::default(),
|
||||
compact: true,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
session.switch_slot("ro").await.unwrap();
|
||||
let response = assistant_text("ok");
|
||||
let err = session
|
||||
.finalize_turn(&response, vec![Message::user_text("x")])
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotReadonly(_)));
|
||||
}
|
||||
}
|
||||
@@ -78,11 +78,7 @@ impl SessionMemory {
|
||||
prefix: Some(format!("{}:", self.namespace)),
|
||||
..Default::default()
|
||||
};
|
||||
let items = self
|
||||
.store
|
||||
.list(&filter)
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
|
||||
|
||||
let mut lines = Vec::with_capacity(items.len() + 2);
|
||||
lines.push("<session-context>".to_string());
|
||||
@@ -113,11 +109,7 @@ impl SessionMemory {
|
||||
prefix: Some(format!("{}:", self.namespace)),
|
||||
..Default::default()
|
||||
};
|
||||
let items = self
|
||||
.store
|
||||
.list(&filter)
|
||||
.await
|
||||
.map_err(AgentError::Memory)?;
|
||||
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
|
||||
|
||||
for item in items {
|
||||
self.store
|
||||
@@ -181,4 +173,4 @@ mod tests {
|
||||
assert!(mem_a.get("key").await.unwrap().is_none());
|
||||
assert_eq!(mem_b.get("key").await.unwrap(), Some("val_b".into()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+5
-11
@@ -10,8 +10,7 @@
|
||||
//! - 重试由上层新建 `Plan` 实现,`TaskAgent` 不做自动重试
|
||||
|
||||
use crate::agent::error::AgentError;
|
||||
#[allow(deprecated)]
|
||||
use crate::llm::types::ChatResponse;
|
||||
use crate::llm::types::response_v2::MessageResponse;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
@@ -56,14 +55,14 @@ impl Step {
|
||||
/// 均未派生 `Clone`(保留原始错误信息,传递所有权而非克隆)。如需复制 `Plan`,
|
||||
/// 只能 clone 处于 `Pending` / `Running` / `Completed` / `Skipped` 状态的步骤。
|
||||
#[derive(Debug)]
|
||||
#[allow(deprecated)]
|
||||
#[non_exhaustive]
|
||||
pub enum StepStatus {
|
||||
/// 初始状态 —— 等待执行。
|
||||
Pending,
|
||||
/// 正在执行(`TaskAgent::execute_plan` 进入)。
|
||||
Running,
|
||||
/// 已完成(含 LLM 响应)。
|
||||
Completed(ChatResponse),
|
||||
Completed(MessageResponse),
|
||||
/// 失败(含错误)。
|
||||
Failed(AgentError),
|
||||
/// 跳过(上层主动跳过)。
|
||||
@@ -130,9 +129,7 @@ impl PlanParser for JsonPlanParser {
|
||||
.collect::<Result<Vec<_>, AgentError>>()?;
|
||||
|
||||
if steps.is_empty() {
|
||||
return Err(AgentError::PlanParse(
|
||||
"Plan 至少需要一个步骤".into(),
|
||||
));
|
||||
return Err(AgentError::PlanParse("Plan 至少需要一个步骤".into()));
|
||||
}
|
||||
|
||||
Ok(Plan {
|
||||
@@ -203,10 +200,7 @@ mod tests {
|
||||
let plan = Plan {
|
||||
id: "p1".into(),
|
||||
goal: "test goal".into(),
|
||||
steps: vec![
|
||||
Step::new(0, "first"),
|
||||
Step::new(1, "second"),
|
||||
],
|
||||
steps: vec![Step::new(0, "first"), Step::new(1, "second")],
|
||||
};
|
||||
assert_eq!(plan.steps.len(), 2);
|
||||
assert_eq!(plan.steps[0].index, 0);
|
||||
|
||||
+32
-16
@@ -73,10 +73,7 @@ impl CompactState {
|
||||
|
||||
/// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。
|
||||
pub fn estimate_message_tokens(messages: &[Message]) -> u32 {
|
||||
messages
|
||||
.iter()
|
||||
.map(estimate_single_message_tokens)
|
||||
.sum()
|
||||
messages.iter().map(estimate_single_message_tokens).sum()
|
||||
}
|
||||
|
||||
fn estimate_single_message_tokens(msg: &Message) -> u32 {
|
||||
@@ -99,9 +96,7 @@ fn estimate_block_tokens(block: &ContentBlock) -> u32 {
|
||||
match block {
|
||||
ContentBlock::Text { text } => estimate_text_tokens(text),
|
||||
ContentBlock::Thinking { text, .. } => estimate_text_tokens(text),
|
||||
ContentBlock::ToolUse { input, .. } => {
|
||||
estimate_text_tokens(&input.to_string())
|
||||
}
|
||||
ContentBlock::ToolUse { input, .. } => estimate_text_tokens(&input.to_string()),
|
||||
ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
|
||||
// ponytail: Image / Audio / File / Extension 在 IR 中固定估算。
|
||||
// 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。
|
||||
@@ -148,14 +143,25 @@ pub fn microcompact(messages: &mut [Message], keep_recent: usize) -> u32 {
|
||||
|
||||
// 第一遍:计算可释放 token(仅非错误 ToolResult)
|
||||
for msg in &messages[..prune_start] {
|
||||
if matches!(msg, Message::ToolResult { is_error: false, .. }) {
|
||||
if matches!(
|
||||
msg,
|
||||
Message::ToolResult {
|
||||
is_error: false,
|
||||
..
|
||||
}
|
||||
) {
|
||||
freed_tokens += estimate_single_message_tokens(msg);
|
||||
}
|
||||
}
|
||||
|
||||
// 第二遍:替换内容(仅非错误 ToolResult)
|
||||
for msg in &mut messages[..prune_start] {
|
||||
if let Message::ToolResult { content, is_error: false, .. } = msg {
|
||||
if let Message::ToolResult {
|
||||
content,
|
||||
is_error: false,
|
||||
..
|
||||
} = msg
|
||||
{
|
||||
*content = vec![ContentBlock::Text {
|
||||
text: "[pruned]".to_string(),
|
||||
}];
|
||||
@@ -177,13 +183,15 @@ mod tests {
|
||||
fn estimate_message_tokens_handles_all_variants() {
|
||||
let messages = vec![
|
||||
Message::System {
|
||||
content: vec![ContentBlock::Text {
|
||||
text: "sys".into(),
|
||||
}],
|
||||
content: vec![ContentBlock::Text { text: "sys".into() }],
|
||||
},
|
||||
Message::user_text("hi"),
|
||||
Message::assistant("ans"),
|
||||
Message::user_image("b64", "image/png", crate::llm::types::shared::ImageDetail::Auto),
|
||||
Message::user_image(
|
||||
"b64",
|
||||
"image/png",
|
||||
crate::llm::types::shared::ImageDetail::Auto,
|
||||
),
|
||||
Message::tool_result("call_1", "tool res", false),
|
||||
];
|
||||
let tokens = estimate_message_tokens(&messages);
|
||||
@@ -205,7 +213,10 @@ mod tests {
|
||||
assert!(freed > 0);
|
||||
assert_eq!(messages.len(), before_len); // 只改内容,不删消息
|
||||
// 索引 1 是被压缩的 ToolResult
|
||||
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
|
||||
if let Message::ToolResult {
|
||||
content, is_error, ..
|
||||
} = &messages[1]
|
||||
{
|
||||
assert_eq!(content.len(), 1);
|
||||
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
|
||||
assert!(!is_error);
|
||||
@@ -228,9 +239,14 @@ mod tests {
|
||||
assert_eq!(freed, 0); // 错误 ToolResult 不计入
|
||||
assert_eq!(messages.len(), before_len);
|
||||
// 错误信息保留完整
|
||||
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
|
||||
if let Message::ToolResult {
|
||||
content, is_error, ..
|
||||
} = &messages[1]
|
||||
{
|
||||
assert!(is_error);
|
||||
assert!(matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down")));
|
||||
assert!(
|
||||
matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down"))
|
||||
);
|
||||
} else {
|
||||
panic!("expected ToolResult at index 1");
|
||||
}
|
||||
|
||||
+29
-30
@@ -8,11 +8,9 @@
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
use crate::llm::types::openai_message::{
|
||||
ContentField, OpenaiChatMessage, OpenaiContentPart,
|
||||
};
|
||||
use crate::llm::types::OpenaiToolCall;
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage, OpenaiContentPart};
|
||||
|
||||
/// `OpenaiChatMessage` → IR `Message`。
|
||||
///
|
||||
@@ -24,11 +22,10 @@ use crate::llm::types::OpenaiToolCall;
|
||||
/// - `Function`(已废弃)→ `Message::ToolResult`(`name` 作为 `tool_call_id` 兜底)
|
||||
pub fn from_openai(msg: &OpenaiChatMessage) -> Message {
|
||||
match msg {
|
||||
OpenaiChatMessage::Developer { content, .. } | OpenaiChatMessage::System { content, .. } => {
|
||||
Message::System {
|
||||
content: content_to_blocks(content),
|
||||
}
|
||||
}
|
||||
OpenaiChatMessage::Developer { content, .. }
|
||||
| OpenaiChatMessage::System { content, .. } => Message::System {
|
||||
content: content_to_blocks(content),
|
||||
},
|
||||
OpenaiChatMessage::User { content, .. } => Message::User {
|
||||
content: content_to_blocks(content),
|
||||
},
|
||||
@@ -86,7 +83,11 @@ pub fn to_openai(msg: &Message) -> OpenaiChatMessage {
|
||||
content: blocks_to_content(content),
|
||||
name: None,
|
||||
},
|
||||
Message::UserImage { data, mime_type, detail } => {
|
||||
Message::UserImage {
|
||||
data,
|
||||
mime_type,
|
||||
detail,
|
||||
} => {
|
||||
// ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。
|
||||
let mime = mime_type.clone();
|
||||
let is_url = data.starts_with("http://") || data.starts_with("https://");
|
||||
@@ -167,26 +168,25 @@ pub fn content_to_blocks(field: &ContentField) -> Vec<ContentBlock> {
|
||||
ContentField::Array(parts) => parts
|
||||
.iter()
|
||||
.filter_map(|p| match p {
|
||||
OpenaiContentPart::Text { text } => {
|
||||
Some(ContentBlock::Text { text: text.clone() })
|
||||
}
|
||||
OpenaiContentPart::Refusal { refusal } => {
|
||||
Some(ContentBlock::Text { text: refusal.clone() })
|
||||
}
|
||||
OpenaiContentPart::Text { text } => Some(ContentBlock::Text { text: text.clone() }),
|
||||
OpenaiContentPart::Refusal { refusal } => Some(ContentBlock::Text {
|
||||
text: refusal.clone(),
|
||||
}),
|
||||
OpenaiContentPart::Image { image_url, .. } => {
|
||||
// ponytail: 简化处理 —— URL 直接通过,data URI 拆出
|
||||
// data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。
|
||||
let url = &image_url.url;
|
||||
if let Some(rest) = url.strip_prefix("data:")
|
||||
&& let Some((mime, b64)) = rest.split_once(";base64,") {
|
||||
return Some(ContentBlock::Image {
|
||||
source: crate::llm::types::message::ImageSource {
|
||||
data: b64.to_string(),
|
||||
mime_type: mime.to_string(),
|
||||
is_url: false,
|
||||
},
|
||||
});
|
||||
}
|
||||
&& let Some((mime, b64)) = rest.split_once(";base64,")
|
||||
{
|
||||
return Some(ContentBlock::Image {
|
||||
source: crate::llm::types::message::ImageSource {
|
||||
data: b64.to_string(),
|
||||
mime_type: mime.to_string(),
|
||||
is_url: false,
|
||||
},
|
||||
});
|
||||
}
|
||||
Some(ContentBlock::Image {
|
||||
source: crate::llm::types::message::ImageSource {
|
||||
data: url.clone(),
|
||||
@@ -263,7 +263,9 @@ mod tests {
|
||||
match ir {
|
||||
Message::System { content } => {
|
||||
assert_eq!(content.len(), 1);
|
||||
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful"));
|
||||
assert!(
|
||||
matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful")
|
||||
);
|
||||
}
|
||||
_ => panic!("expected System variant"),
|
||||
}
|
||||
@@ -385,10 +387,7 @@ mod tests {
|
||||
assert_eq!(parts.len(), 1);
|
||||
match &parts[0] {
|
||||
OpenaiContentPart::Image { image_url, .. } => {
|
||||
assert_eq!(
|
||||
image_url.url,
|
||||
"data:image/png;base64,BASE64DATA"
|
||||
);
|
||||
assert_eq!(image_url.url, "data:image/png;base64,BASE64DATA");
|
||||
}
|
||||
_ => panic!("expected Image part"),
|
||||
}
|
||||
|
||||
+828
-43
File diff suppressed because it is too large
Load Diff
+11
-4
@@ -8,9 +8,12 @@ use std::time::Duration;
|
||||
///
|
||||
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查 API key、减少上下文)。
|
||||
#[derive(thiserror::Error, Debug)]
|
||||
#[non_exhaustive]
|
||||
pub enum LlmError {
|
||||
/// API 认证失败(API key 无效、过期或权限不足)。
|
||||
#[error("LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确")]
|
||||
#[error(
|
||||
"LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确"
|
||||
)]
|
||||
Authentication(String),
|
||||
|
||||
/// 请求被限流,可选地附带重试等待时间。可重试。
|
||||
@@ -18,7 +21,9 @@ pub enum LlmError {
|
||||
RateLimit { retry_after: Option<Duration> },
|
||||
|
||||
/// HTTP 请求失败(网络错误或非 2xx 状态码),包含状态码与响应体。
|
||||
#[error("LLM 请求失败(HTTP {status}): {body}。请检查 Provider 端点地址(base_url)和网络连通性")]
|
||||
#[error(
|
||||
"LLM 请求失败(HTTP {status}): {body}。请检查 Provider 端点地址(base_url)和网络连通性"
|
||||
)]
|
||||
Request { status: u16, body: String },
|
||||
|
||||
/// 请求超时。可重试。
|
||||
@@ -30,10 +35,12 @@ pub enum LlmError {
|
||||
Stream(String),
|
||||
|
||||
/// 上下文长度超出模型窗口限制。
|
||||
#[error("LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compaction(llm::compact)")]
|
||||
#[error(
|
||||
"LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compaction(llm::compact)"
|
||||
)]
|
||||
ContextLength { actual: u32, limit: u32 },
|
||||
|
||||
/// 其他未分类的 LLM 调用失败。
|
||||
#[error("LLM 调用失败: {0}")]
|
||||
Other(String),
|
||||
}
|
||||
}
|
||||
|
||||
+2
-3
@@ -7,6 +7,7 @@ use crate::llm::types::request_v2::MessageRequest;
|
||||
|
||||
/// 生命周期钩子事件点。
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
#[non_exhaustive]
|
||||
pub enum HookEvent {
|
||||
/// LLM 请求发起之前(可阻断)。
|
||||
PreRequest,
|
||||
@@ -130,9 +131,7 @@ impl Default for HookExecutor {
|
||||
impl HookExecutor {
|
||||
/// 创建一个空的执行器。
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
hooks: Vec::new(),
|
||||
}
|
||||
Self { hooks: Vec::new() }
|
||||
}
|
||||
|
||||
/// 注册一个钩子到指定事件点。
|
||||
|
||||
+11
-9
@@ -97,8 +97,7 @@ impl LlmProvider for MockProvider {
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
_request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
let response = self.pop()?;
|
||||
// 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。
|
||||
let id = response.id.clone();
|
||||
@@ -206,10 +205,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn chat_returns_queued_response() {
|
||||
let provider = MockProvider::new(vec![text_response("hello")]);
|
||||
let resp = provider
|
||||
.chat(MessageRequest::default())
|
||||
.await
|
||||
.unwrap();
|
||||
let resp = provider.chat(MessageRequest::default()).await.unwrap();
|
||||
assert_eq!(resp.text(), "hello");
|
||||
assert_eq!(provider.remaining(), 0);
|
||||
}
|
||||
@@ -231,7 +227,10 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn chat_stream_emits_text_delta_sequence() {
|
||||
let provider = MockProvider::new(vec![text_response("hi")]);
|
||||
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
|
||||
let mut stream = provider
|
||||
.chat_stream(MessageRequest::default())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seen_start = false;
|
||||
let mut seen_block_start = false;
|
||||
@@ -283,7 +282,10 @@ mod tests {
|
||||
extra: Default::default(),
|
||||
};
|
||||
let provider = MockProvider::new(vec![response]);
|
||||
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
|
||||
let mut stream = provider
|
||||
.chat_stream(MessageRequest::default())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut saw_tool_args = false;
|
||||
let mut saw_tool_end = false;
|
||||
@@ -302,4 +304,4 @@ mod tests {
|
||||
assert!(saw_tool_args);
|
||||
assert!(saw_tool_end);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+456
-21
@@ -1,11 +1,14 @@
|
||||
pub mod anthropic;
|
||||
pub mod ollama;
|
||||
pub mod openai;
|
||||
pub mod openai_compat;
|
||||
pub mod registry;
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_core::Stream;
|
||||
use reqwest::Client;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::llm::error::LlmError;
|
||||
@@ -18,6 +21,7 @@ use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
||||
/// 当前协议数量(5 种以内)完全可控,enum 的编译期安全检查优于运行时的 `HashMap::get()`。
|
||||
/// 未来如果扩展到 15+ 种以上,再改为注册表模式。
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
#[non_exhaustive]
|
||||
pub enum ProviderType {
|
||||
/// OpenAI Chat Completions API(兼容 DeepSeek / Qwen 等 `/chat/completions` 端点)。
|
||||
OpenaiChat,
|
||||
@@ -29,6 +33,8 @@ pub enum ProviderType {
|
||||
DeepSeek,
|
||||
/// Qwen / 阿里云百炼(OpenAI-compatible `/chat/completions`)。
|
||||
Qwen,
|
||||
/// Ollama 本地推理(OpenAI-compatible `/chat/completions`,默认 `http://localhost:11434/v1`)。
|
||||
Ollama,
|
||||
}
|
||||
|
||||
impl std::str::FromStr for ProviderType {
|
||||
@@ -41,47 +47,211 @@ impl std::str::FromStr for ProviderType {
|
||||
"anthropic" | "claude" => Ok(ProviderType::Anthropic),
|
||||
"deepseek" => Ok(ProviderType::DeepSeek),
|
||||
"qwen" | "dashscope" | "tongyi" => Ok(ProviderType::Qwen),
|
||||
"ollama" => Ok(ProviderType::Ollama),
|
||||
_ => Err(format!("未知的 Provider 类型: {s}")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Provider 构造参数 —— 通用 base_url + api_key + model。
|
||||
/// Provider 构造参数 —— 通用 base_url + api_key + model + timeout / retry 配置。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ProviderConfig {
|
||||
/// API base URL(如 `https://api.openai.com/v1`)。为空时由 Provider 选择默认值。
|
||||
pub base_url: String,
|
||||
/// API key。Ollama 等本地 Provider 可为空。
|
||||
pub api_key: String,
|
||||
/// 模型名(如 `gpt-4o` / `claude-sonnet-4-20250514`)。
|
||||
pub model: String,
|
||||
/// 请求超时秒数(默认 30)。应用于 Provider 的 HTTP Client 级别。
|
||||
pub timeout_secs: u64,
|
||||
/// 最大重试次数(默认 3)。
|
||||
///
|
||||
/// 当前此字段仅由 `from_env()` 采集,**实际重试逻辑由 `CycleConfig.retry.max_retries` 控制**。
|
||||
/// 此处保留字段以与 Roadmap §Phase 5 Step 5.1 对齐;未来 Phase 6+ 可统一合并到 `CycleConfig`。
|
||||
pub max_retries: u32,
|
||||
}
|
||||
|
||||
impl Default for ProviderConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
base_url: String::new(),
|
||||
api_key: String::new(),
|
||||
model: String::new(),
|
||||
timeout_secs: 30,
|
||||
max_retries: 3,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ProviderConfig {
|
||||
/// 从环境变量构造 `ProviderConfig`。
|
||||
///
|
||||
/// 必填变量:
|
||||
/// - `{prefix}_BASE_URL`
|
||||
/// - `{prefix}_API_KEY`
|
||||
/// - `{prefix}_MODEL`
|
||||
///
|
||||
/// 可选变量(有默认值):
|
||||
/// - `{prefix}_TIMEOUT_SECS`(默认 30,解析失败回退 30 并 warn)
|
||||
/// - `{prefix}_MAX_RETRIES`(默认 3,解析失败回退 3 并 warn)
|
||||
pub fn from_env(prefix: &str) -> Result<Self, String> {
|
||||
let base_url = std::env::var(format!("{prefix}_BASE_URL"))
|
||||
.map_err(|_| format!("{prefix}_BASE_URL 环境变量未设置"))?;
|
||||
let api_key = std::env::var(format!("{prefix}_API_KEY"))
|
||||
.map_err(|_| format!("{prefix}_API_KEY 环境变量未设置"))?;
|
||||
let model = std::env::var(format!("{prefix}_MODEL"))
|
||||
.map_err(|_| format!("{prefix}_MODEL 环境变量未设置"))?;
|
||||
let timeout_secs = match std::env::var(format!("{prefix}_TIMEOUT_SECS")) {
|
||||
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||
tracing::warn!("{prefix}_TIMEOUT_SECS='{v}' 解析失败,使用默认值 30");
|
||||
30
|
||||
}),
|
||||
Err(_) => 30,
|
||||
};
|
||||
let max_retries = match std::env::var(format!("{prefix}_MAX_RETRIES")) {
|
||||
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||
tracing::warn!("{prefix}_MAX_RETRIES='{v}' 解析失败,使用默认值 3");
|
||||
3
|
||||
}),
|
||||
Err(_) => 3,
|
||||
};
|
||||
|
||||
// ponytail: max_retries 当前仅采集,不传入 Provider。
|
||||
// 实际重试由 CycleConfig.retry.max_retries 控制。
|
||||
if max_retries != 3 {
|
||||
tracing::warn!(
|
||||
"ProviderConfig.max_retries={} 已采集但当前未生效;\
|
||||
重试次数由 CycleConfig.retry.max_retries 控制",
|
||||
max_retries,
|
||||
);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
base_url,
|
||||
api_key,
|
||||
model,
|
||||
timeout_secs,
|
||||
max_retries,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 构造带 timeout 的 `reqwest::Client`(OpenAI-compatible 共享)。
|
||||
fn build_client_with_timeout(timeout_secs: u64) -> Result<Client, LlmError> {
|
||||
Client::builder()
|
||||
.timeout(Duration::from_secs(timeout_secs))
|
||||
.build()
|
||||
.map_err(|e| LlmError::Other(format!("创建 HTTP 客户端失败: {e}")))
|
||||
}
|
||||
|
||||
/// 构造带 Anthropic 默认 headers + timeout 的 `reqwest::Client`。
|
||||
///
|
||||
/// Anthropic 由于需要保留 `x-api-key` / `anthropic-version` 默认 headers,
|
||||
/// 与 OpenAI-compatible 共享的 `build_client_with_timeout` 不同。
|
||||
fn build_anthropic_client(api_key: &str, timeout_secs: u64) -> Result<Client, LlmError> {
|
||||
use reqwest::header::{HeaderMap, HeaderValue};
|
||||
|
||||
let key_header = HeaderValue::from_str(api_key)
|
||||
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
|
||||
let version_header = HeaderValue::from_static("2023-06-01");
|
||||
|
||||
Client::builder()
|
||||
.timeout(Duration::from_secs(timeout_secs))
|
||||
.default_headers({
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-api-key", key_header);
|
||||
headers.insert("anthropic-version", version_header);
|
||||
headers
|
||||
})
|
||||
.build()
|
||||
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))
|
||||
}
|
||||
|
||||
/// Provider 工厂 —— exhaustive match 在编译期保证新 Provider 被注册。
|
||||
///
|
||||
/// `config.timeout_secs` 注入到 Provider 的 HTTP Client 超时配置。
|
||||
/// 每个分支通过 `from_parts` (pub(crate)) 一次性构造,无冗余 client 创建。
|
||||
pub fn create_provider(
|
||||
provider_type: ProviderType,
|
||||
config: ProviderConfig,
|
||||
) -> Result<Box<dyn LlmProvider>, LlmError> {
|
||||
match provider_type {
|
||||
ProviderType::OpenaiChat => Ok(Box::new(openai::OpenaiChatProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::OpenaiChat => {
|
||||
let client = build_client_with_timeout(config.timeout_secs)?;
|
||||
Ok(Box::new(openai::OpenaiChatProvider(
|
||||
openai::GenericOpenaiProvider::from_parts(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
"openai",
|
||||
client,
|
||||
Vec::new(),
|
||||
config.timeout_secs,
|
||||
),
|
||||
)))
|
||||
}
|
||||
ProviderType::OpenaiResponse => Err(LlmError::Other(
|
||||
"OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(),
|
||||
)),
|
||||
ProviderType::Anthropic => Ok(Box::new(anthropic::AnthropicProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::DeepSeek => Ok(Box::new(openai_compat::DeepSeekProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::Qwen => Ok(Box::new(openai_compat::QwenProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::Anthropic => {
|
||||
let client = build_anthropic_client(&config.api_key, config.timeout_secs)?;
|
||||
Ok(Box::new(anthropic::AnthropicProvider::from_parts(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
client,
|
||||
config.timeout_secs,
|
||||
)))
|
||||
}
|
||||
ProviderType::DeepSeek => {
|
||||
let client = build_client_with_timeout(config.timeout_secs)?;
|
||||
Ok(Box::new(openai_compat::DeepSeekProvider(
|
||||
openai::GenericOpenaiProvider::from_parts(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
"deepseek",
|
||||
client,
|
||||
Vec::new(),
|
||||
config.timeout_secs,
|
||||
),
|
||||
)))
|
||||
}
|
||||
ProviderType::Qwen => {
|
||||
let client = build_client_with_timeout(config.timeout_secs)?;
|
||||
Ok(Box::new(openai_compat::QwenProvider(
|
||||
openai::GenericOpenaiProvider::from_parts(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
"qwen",
|
||||
client,
|
||||
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
||||
config.timeout_secs,
|
||||
),
|
||||
)))
|
||||
}
|
||||
ProviderType::Ollama => {
|
||||
let client = build_client_with_timeout(config.timeout_secs)?;
|
||||
// ponytail: Ollama 默认 base_url 由 OllamaProvider 构造处理 —— 但 from_parts 不走
|
||||
// OllamaProvider::new 的默认 URL 回退。这里保留 base_url(可能为空 → http://localhost:11434/v1)。
|
||||
let base_url = if config.base_url.is_empty() {
|
||||
"http://localhost:11434/v1".to_string()
|
||||
} else {
|
||||
config.base_url
|
||||
};
|
||||
Ok(Box::new(ollama::OllamaProvider(
|
||||
openai::GenericOpenaiProvider::from_parts(
|
||||
base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
"ollama",
|
||||
client,
|
||||
Vec::new(),
|
||||
config.timeout_secs,
|
||||
),
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -142,3 +312,268 @@ pub trait LlmProvider: Send + Sync {
|
||||
/// 返回 Provider 静态能力描述。
|
||||
fn capabilities(&self) -> ProviderCapabilities;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::time::Duration;
|
||||
|
||||
#[test]
|
||||
fn provider_config_default_values() {
|
||||
let config = ProviderConfig::default();
|
||||
assert_eq!(config.timeout_secs, 30);
|
||||
assert_eq!(config.max_retries, 3);
|
||||
assert_eq!(config.base_url, "");
|
||||
assert_eq!(config.api_key, "");
|
||||
assert_eq!(config.model, "");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_from_env_requires_all_three() {
|
||||
// 使用 temp_env 移除所有相关变量,避免外部环境意外设置导致测试 flaky
|
||||
temp_env::with_vars(
|
||||
[
|
||||
("TEST_PROVIDER_MISSING_BASE_URL", None::<&str>),
|
||||
("TEST_PROVIDER_MISSING_API_KEY", None::<&str>),
|
||||
("TEST_PROVIDER_MISSING_MODEL", None::<&str>),
|
||||
("TEST_PROVIDER_MISSING_TIMEOUT_SECS", None::<&str>),
|
||||
("TEST_PROVIDER_MISSING_MAX_RETRIES", None::<&str>),
|
||||
],
|
||||
|| {
|
||||
let result = ProviderConfig::from_env("TEST_PROVIDER_MISSING");
|
||||
assert!(result.is_err());
|
||||
let msg = result.unwrap_err();
|
||||
assert!(
|
||||
msg.contains("TEST_PROVIDER_MISSING_BASE_URL"),
|
||||
"error should mention missing var, got: {msg}"
|
||||
);
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_from_env_uses_defaults_when_only_required_set() {
|
||||
temp_env::with_vars(
|
||||
[
|
||||
("TEST_PROVIDER_BASE_URL", Some("http://localhost:11434/v1")),
|
||||
("TEST_PROVIDER_API_KEY", Some("")),
|
||||
("TEST_PROVIDER_MODEL", Some("llama3")),
|
||||
],
|
||||
|| {
|
||||
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||
assert_eq!(config.base_url, "http://localhost:11434/v1");
|
||||
assert_eq!(config.api_key, "");
|
||||
assert_eq!(config.model, "llama3");
|
||||
assert_eq!(config.timeout_secs, 30);
|
||||
assert_eq!(config.max_retries, 3);
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_from_env_reads_custom_values() {
|
||||
temp_env::with_vars(
|
||||
[
|
||||
("TEST_PROVIDER_BASE_URL", Some("http://x")),
|
||||
("TEST_PROVIDER_API_KEY", Some("k")),
|
||||
("TEST_PROVIDER_MODEL", Some("m")),
|
||||
("TEST_PROVIDER_TIMEOUT_SECS", Some("60")),
|
||||
("TEST_PROVIDER_MAX_RETRIES", Some("5")),
|
||||
],
|
||||
|| {
|
||||
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||
assert_eq!(config.timeout_secs, 60);
|
||||
assert_eq!(config.max_retries, 5);
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_from_env_falls_back_on_invalid_numbers() {
|
||||
temp_env::with_vars(
|
||||
[
|
||||
("TEST_PROVIDER_BASE_URL", Some("http://x")),
|
||||
("TEST_PROVIDER_API_KEY", Some("k")),
|
||||
("TEST_PROVIDER_MODEL", Some("m")),
|
||||
("TEST_PROVIDER_TIMEOUT_SECS", Some("not-a-number")),
|
||||
("TEST_PROVIDER_MAX_RETRIES", Some("also-bad")),
|
||||
],
|
||||
|| {
|
||||
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||
// 解析失败回退默认值
|
||||
assert_eq!(config.timeout_secs, 30);
|
||||
assert_eq!(config.max_retries, 3);
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
/// Timeout 传导集成测试:构造 `ProviderConfig` timeout=1s,
|
||||
/// `create_provider` 注入 1s 超时 client,请求一个故意延迟 3s 的 mock server,
|
||||
/// 验证返回 `LlmError::Timeout { duration: 1s }`。
|
||||
#[tokio::test]
|
||||
async fn create_provider_injects_timeout_into_openai_chat() {
|
||||
use serde_json::json;
|
||||
use wiremock::matchers::{method, path};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||
|
||||
let server = MockServer::start().await;
|
||||
// 故意延迟 3s 触发超时
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_delay(Duration::from_secs(3))
|
||||
.set_body_json(json!({
|
||||
"id": "x",
|
||||
"object": "chat.completion",
|
||||
"created": 0,
|
||||
"model": "gpt-4o",
|
||||
"choices": [],
|
||||
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = create_provider(
|
||||
ProviderType::OpenaiChat,
|
||||
ProviderConfig {
|
||||
base_url: server.uri(),
|
||||
api_key: "sk-test".into(),
|
||||
model: "gpt-4o".into(),
|
||||
timeout_secs: 1,
|
||||
max_retries: 3,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let err = provider
|
||||
.chat(crate::llm::types::request_v2::MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![],
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.unwrap_err();
|
||||
match err {
|
||||
LlmError::Timeout { duration } => {
|
||||
assert_eq!(duration, Duration::from_secs(1));
|
||||
}
|
||||
other => panic!("expected Timeout, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Timeout 传导验证:`create_provider` 生成的 DeepSeek Provider 也带 1s 超时,
|
||||
/// 错误消息中的 duration 与 timeout_secs 一致(而非硬编码 120s)。
|
||||
#[tokio::test]
|
||||
async fn create_provider_injects_timeout_into_deepseek() {
|
||||
use serde_json::json;
|
||||
use wiremock::matchers::{method, path};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_delay(Duration::from_secs(3))
|
||||
.set_body_json(json!({
|
||||
"id": "x",
|
||||
"object": "chat.completion",
|
||||
"created": 0,
|
||||
"model": "deepseek-chat",
|
||||
"choices": [],
|
||||
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = create_provider(
|
||||
ProviderType::DeepSeek,
|
||||
ProviderConfig {
|
||||
base_url: server.uri(),
|
||||
api_key: "sk-test".into(),
|
||||
model: "deepseek-chat".into(),
|
||||
timeout_secs: 1,
|
||||
max_retries: 3,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let err = provider
|
||||
.chat(crate::llm::types::request_v2::MessageRequest {
|
||||
model: "deepseek-chat".into(),
|
||||
messages: vec![],
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.unwrap_err();
|
||||
match err {
|
||||
LlmError::Timeout { duration } => {
|
||||
assert_eq!(duration, Duration::from_secs(1));
|
||||
}
|
||||
other => panic!("expected Timeout, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Timeout 传导验证:`create_provider` 生成的 Anthropic Provider 通过 `with_timeout`
|
||||
/// 注入 1s 超时。
|
||||
///
|
||||
/// 与 OpenAI-compatible 路径不同,Anthropic 走 `AnthropicProvider::with_timeout()`
|
||||
/// 重建底层 client(保留 default_headers),独立于 OpenAI-compatible 的 `build_client_with_timeout`。
|
||||
/// 单独覆盖此路径以验证 `with_timeout` 不会因服务端延迟而返回硬编码 120s 的超时错误。
|
||||
#[tokio::test]
|
||||
async fn create_provider_injects_timeout_into_anthropic() {
|
||||
use serde_json::json;
|
||||
use wiremock::matchers::{method, path};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||
|
||||
let server = MockServer::start().await;
|
||||
// Anthropic Messages API 端点:`POST /v1/messages`
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/messages"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_delay(Duration::from_secs(3))
|
||||
.set_body_json(json!({
|
||||
"id": "msg_timeout_test",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"model": "claude-sonnet-4-20250514",
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1}
|
||||
})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = create_provider(
|
||||
ProviderType::Anthropic,
|
||||
ProviderConfig {
|
||||
base_url: server.uri(),
|
||||
api_key: "sk-ant-test".into(),
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
timeout_secs: 1,
|
||||
max_retries: 3,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let err = provider
|
||||
.chat(crate::llm::types::request_v2::MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![],
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.unwrap_err();
|
||||
match err {
|
||||
LlmError::Timeout { duration } => {
|
||||
assert_eq!(duration, Duration::from_secs(1));
|
||||
}
|
||||
other => panic!("expected Timeout, got {other:?}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+132
-55
@@ -12,10 +12,10 @@ use async_trait::async_trait;
|
||||
use bytes::Bytes;
|
||||
use futures_core::Stream;
|
||||
use futures_util::StreamExt;
|
||||
use reqwest::header::{HeaderMap, HeaderValue};
|
||||
use reqwest::Client;
|
||||
use reqwest::header::{HeaderMap, HeaderValue};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{json, Value};
|
||||
use serde_json::{Value, json};
|
||||
use tracing::{debug, error, info, warn};
|
||||
|
||||
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
|
||||
@@ -39,16 +39,20 @@ pub struct AnthropicProvider {
|
||||
#[allow(dead_code)]
|
||||
api_key: String,
|
||||
model: String,
|
||||
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
|
||||
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter,
|
||||
/// 因此单独存储以便错误消息与配置保持一致。
|
||||
timeout_secs: u64,
|
||||
}
|
||||
|
||||
impl AnthropicProvider {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
let key_header = HeaderValue::from_str(&api_key)
|
||||
.expect("Anthropic API key 包含无效的 HTTP 头部字符");
|
||||
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||
let key_header =
|
||||
HeaderValue::from_str(&api_key).expect("Anthropic API key 包含无效的 HTTP 头部字符");
|
||||
let version_header = HeaderValue::from_static("2023-06-01");
|
||||
|
||||
let http_client = Client::builder()
|
||||
.timeout(Duration::from_secs(120))
|
||||
.timeout(Duration::from_secs(timeout_secs))
|
||||
.default_headers({
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-api-key", key_header);
|
||||
@@ -67,14 +71,80 @@ impl AnthropicProvider {
|
||||
},
|
||||
api_key,
|
||||
model,
|
||||
timeout_secs,
|
||||
}
|
||||
}
|
||||
|
||||
/// ⚠️ 替换 HTTP Client,**丢弃** `new()` 中设置的默认 headers(`x-api-key` / `anthropic-version`)。
|
||||
///
|
||||
/// 调用此方法后,所有 Anthropic API 请求将以**无认证头**发送出去,预期会 401/403 失败。
|
||||
/// 推荐改用 [`Self::with_timeout`],它会重建 client 并保留默认 headers。
|
||||
///
|
||||
/// 此方法仍保留以兼容调用方自定义 client 但不需要默认 headers 的极端场景。
|
||||
#[deprecated(
|
||||
since = "0.2.0",
|
||||
note = "此方法会丢弃默认 headers(x-api-key / anthropic-version),改为使用 `with_timeout` 或带 headers 的 `Client::builder()`"
|
||||
)]
|
||||
pub fn with_client(mut self, client: Client) -> Self {
|
||||
self.http_client = client;
|
||||
self
|
||||
}
|
||||
|
||||
/// 替换 HTTP Client 的超时配置(重建底层 client,保留默认 headers)。
|
||||
///
|
||||
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后通过 `with_client` 注入的 Client
|
||||
/// 将被替换。headers 构造逻辑与 `new()` 中的保持一致(`x-api-key` / `anthropic-version`)。
|
||||
///
|
||||
/// ponytail: 同值调用短路。当 `secs == self.timeout_secs` 时跳过 client 重建,
|
||||
/// 避免 `create_provider` 路径 `new(timeout).with_timeout(timeout)` 的双重构造。
|
||||
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
|
||||
if secs == self.timeout_secs {
|
||||
return Ok(self);
|
||||
}
|
||||
// ponytail: 重建 http_client 时保留已有默认 headers(x-api-key / anthropic-version)。
|
||||
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
|
||||
let key_header = HeaderValue::from_str(&self.api_key)
|
||||
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
|
||||
let version_header = HeaderValue::from_static("2023-06-01");
|
||||
|
||||
self.http_client = Client::builder()
|
||||
.timeout(Duration::from_secs(secs))
|
||||
.default_headers({
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-api-key", key_header);
|
||||
headers.insert("anthropic-version", version_header);
|
||||
headers
|
||||
})
|
||||
.build()
|
||||
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
|
||||
self.timeout_secs = secs;
|
||||
Ok(self)
|
||||
}
|
||||
|
||||
/// 一次性构造 —— `create_provider` 路径专用,避免 `new(...)` + `with_timeout(...)` 的双重 client 构造。
|
||||
///
|
||||
/// 调用方负责预先构造好符合 Anthropic 协议要求的 `http_client`(带正确的 `x-api-key` /
|
||||
/// `anthropic-version` 默认 headers + 指定 timeout)。
|
||||
pub(crate) fn from_parts(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
http_client: Client,
|
||||
timeout_secs: u64,
|
||||
) -> Self {
|
||||
Self {
|
||||
http_client,
|
||||
base_url: if base_url.is_empty() {
|
||||
"https://api.anthropic.com".to_string()
|
||||
} else {
|
||||
base_url
|
||||
},
|
||||
api_key,
|
||||
model,
|
||||
timeout_secs,
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
|
||||
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
|
||||
}
|
||||
@@ -104,12 +174,14 @@ impl AnthropicProvider {
|
||||
Message::User { content } => {
|
||||
api_messages.push(AnthropicMessage::user(content));
|
||||
}
|
||||
Message::UserImage { data, mime_type, detail } => {
|
||||
Message::UserImage {
|
||||
data,
|
||||
mime_type,
|
||||
detail,
|
||||
} => {
|
||||
// Anthropic image format: {type: "image", source: {type: "base64", media_type, data}}
|
||||
let source = if data.starts_with("http://") || data.starts_with("https://") {
|
||||
AnthropicImageSource::Url {
|
||||
url: data.clone(),
|
||||
}
|
||||
AnthropicImageSource::Url { url: data.clone() }
|
||||
} else {
|
||||
AnthropicImageSource::Base64 {
|
||||
media_type: mime_type.clone(),
|
||||
@@ -194,7 +266,7 @@ impl AnthropicProvider {
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(Self::map_reqwest_error)?;
|
||||
.map_err(|e| self.map_reqwest_error(e))?;
|
||||
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
@@ -215,8 +287,7 @@ impl AnthropicProvider {
|
||||
async fn chat_stream_inner(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
let mut body = self.build_request_body(request)?;
|
||||
body.stream = Some(true);
|
||||
|
||||
@@ -230,16 +301,16 @@ impl AnthropicProvider {
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(Self::map_reqwest_error)?;
|
||||
.map_err(|e| self.map_reqwest_error(e))?;
|
||||
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
return Err(Self::handle_error_response(response).await);
|
||||
}
|
||||
|
||||
let byte_stream = response.bytes_stream().map(|r| {
|
||||
r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}")))
|
||||
});
|
||||
let byte_stream = response
|
||||
.bytes_stream()
|
||||
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}"))));
|
||||
|
||||
let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> =
|
||||
Box::pin(byte_stream);
|
||||
@@ -247,10 +318,10 @@ impl AnthropicProvider {
|
||||
Ok(Box::pin(AnthropicSseStream::new(byte_stream)))
|
||||
}
|
||||
|
||||
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
|
||||
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
|
||||
if e.is_timeout() {
|
||||
LlmError::Timeout {
|
||||
duration: Duration::from_secs(120),
|
||||
duration: Duration::from_secs(self.timeout_secs),
|
||||
}
|
||||
} else if e.is_connect() {
|
||||
LlmError::Other(format!("连接失败: {e}"))
|
||||
@@ -291,13 +362,12 @@ impl AnthropicProvider {
|
||||
blocks.push(ContentBlock::Text { text });
|
||||
}
|
||||
AnthropicContentBlockResp::ToolUse { id, name, input } => {
|
||||
blocks.push(ContentBlock::ToolUse {
|
||||
id,
|
||||
name,
|
||||
input,
|
||||
});
|
||||
blocks.push(ContentBlock::ToolUse { id, name, input });
|
||||
}
|
||||
AnthropicContentBlockResp::Thinking { thinking, signature } => {
|
||||
AnthropicContentBlockResp::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
} => {
|
||||
blocks.push(ContentBlock::Thinking {
|
||||
text: thinking,
|
||||
signature,
|
||||
@@ -336,8 +406,7 @@ impl LlmProvider for AnthropicProvider {
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
self.chat_stream_inner(request).await
|
||||
}
|
||||
|
||||
@@ -418,7 +487,9 @@ impl AnthropicMessage {
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum AnthropicContentPart {
|
||||
Text { text: String },
|
||||
Text {
|
||||
text: String,
|
||||
},
|
||||
Image {
|
||||
source: AnthropicImageSource,
|
||||
},
|
||||
@@ -437,13 +508,8 @@ enum AnthropicContentPart {
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum AnthropicImageSource {
|
||||
Base64 {
|
||||
media_type: String,
|
||||
data: String,
|
||||
},
|
||||
Url {
|
||||
url: String,
|
||||
},
|
||||
Base64 { media_type: String, data: String },
|
||||
Url { url: String },
|
||||
}
|
||||
|
||||
fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> {
|
||||
@@ -523,9 +589,18 @@ struct AnthropicUsage {
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum AnthropicContentBlockResp {
|
||||
Text { text: String },
|
||||
ToolUse { id: String, name: String, input: Value },
|
||||
Thinking { thinking: String, signature: Option<String> },
|
||||
Text {
|
||||
text: String,
|
||||
},
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: Value,
|
||||
},
|
||||
Thinking {
|
||||
thinking: String,
|
||||
signature: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -662,18 +737,13 @@ impl AnthropicSseStream {
|
||||
// 先把所有字段提前,避免 match 中 part-move
|
||||
let block_type = match &content_block {
|
||||
AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text,
|
||||
AnthropicContentBlockStart::ToolUse { id, name } => {
|
||||
ContentBlockType::ToolUse {
|
||||
id: id.clone(),
|
||||
name: name.clone(),
|
||||
}
|
||||
}
|
||||
AnthropicContentBlockStart::ToolUse { id, name } => ContentBlockType::ToolUse {
|
||||
id: id.clone(),
|
||||
name: name.clone(),
|
||||
},
|
||||
AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking,
|
||||
};
|
||||
events.push(StreamEvent::ContentBlockStart {
|
||||
index,
|
||||
block_type,
|
||||
});
|
||||
events.push(StreamEvent::ContentBlockStart { index, block_type });
|
||||
let builder = match content_block {
|
||||
AnthropicContentBlockStart::Text { text } => {
|
||||
crate::llm::types::response_v2::ContentBlockBuilder::Text(text)
|
||||
@@ -742,7 +812,9 @@ impl AnthropicSseStream {
|
||||
completion_tokens_details: None,
|
||||
prompt_tokens_details: None,
|
||||
};
|
||||
events.push(StreamEvent::CostUpdate { usage: partial_usage });
|
||||
events.push(StreamEvent::CostUpdate {
|
||||
usage: partial_usage,
|
||||
});
|
||||
}
|
||||
}
|
||||
AnthropicSseEvent::MessageStop => {
|
||||
@@ -753,7 +825,9 @@ impl AnthropicSseStream {
|
||||
self.saw_terminal = true;
|
||||
match self.partial.clone().finalize() {
|
||||
Ok(full) => {
|
||||
events.push(StreamEvent::MessageComplete { full_response: full });
|
||||
events.push(StreamEvent::MessageComplete {
|
||||
full_response: full,
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
events.push(StreamEvent::Error {
|
||||
@@ -780,10 +854,7 @@ fn _unused_marker() {}
|
||||
impl Stream for AnthropicSseStream {
|
||||
type Item = Result<StreamEvent, LlmError>;
|
||||
|
||||
fn poll_next(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Option<Self::Item>> {
|
||||
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
loop {
|
||||
if let Some(data) = self.next_event_line() {
|
||||
let mut events = self.handle_event_json(&data);
|
||||
@@ -835,7 +906,12 @@ mod tests {
|
||||
|
||||
fn make_provider(base_url: String) -> AnthropicProvider {
|
||||
// 跳过默认 header 注入:测试用自定义 base_url 直接 mock
|
||||
AnthropicProvider::new(base_url, "sk-ant-test".into(), "claude-sonnet-4-20250514".into())
|
||||
AnthropicProvider::new(
|
||||
base_url,
|
||||
"sk-ant-test".into(),
|
||||
"claude-sonnet-4-20250514".into(),
|
||||
30,
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -979,6 +1055,7 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
|
||||
"http://x".into(),
|
||||
"k".into(),
|
||||
"claude-sonnet-4-20250514".into(),
|
||||
30,
|
||||
)
|
||||
.capabilities();
|
||||
assert_eq!(caps.provider_name, "anthropic");
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
//! Ollama Provider —— OpenAI-compatible 协议的 newtype 包装,零 API key。
|
||||
//!
|
||||
//! 默认 base_url = `http://localhost:11434/v1`,空 api_key 也可工作。
|
||||
//! 实现方式同 `DeepSeekProvider` / `QwenProvider`,共享 `GenericOpenaiProvider`
|
||||
//! 的 HTTP/SSE/转换逻辑,仅配置不同。
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use futures_core::Stream;
|
||||
use reqwest::Client;
|
||||
|
||||
use super::openai::GenericOpenaiProvider;
|
||||
use super::{LlmProvider, ProviderCapabilities};
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
||||
|
||||
/// Ollama 本地 Provider —— OpenAI-compatible 协议的 newtype 包装。
|
||||
///
|
||||
/// Ollama 在 `localhost:11434` 暴露与 OpenAI 兼容的 `/v1/chat/completions`
|
||||
/// 接口,因此完全复用 `GenericOpenaiProvider` 的实现。允许空 `api_key`。
|
||||
pub struct OllamaProvider(pub GenericOpenaiProvider);
|
||||
|
||||
impl OllamaProvider {
|
||||
/// 构造 Ollama Provider。
|
||||
///
|
||||
/// - `base_url` 为空时使用默认 `http://localhost:11434/v1`
|
||||
/// - `api_key` 可为空字符串(Ollama 不校验)
|
||||
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"http://localhost:11434/v1".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
Self(GenericOpenaiProvider::new_with_name(
|
||||
url,
|
||||
api_key,
|
||||
model,
|
||||
"ollama",
|
||||
timeout_secs,
|
||||
))
|
||||
}
|
||||
|
||||
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
|
||||
///
|
||||
/// 与 `OpenaiChatProvider::with_client`、`DeepSeekProvider::with_client`、
|
||||
/// `QwenProvider::with_client` 签名一致。
|
||||
pub fn with_client(self, client: Client) -> Self {
|
||||
Self(self.0.with_client(client))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for OllamaProvider {
|
||||
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
self.0.chat(request).await
|
||||
}
|
||||
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
self.0.chat_stream(request).await
|
||||
}
|
||||
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
let mut caps = self.0.capabilities();
|
||||
caps.provider_name = "ollama";
|
||||
caps
|
||||
}
|
||||
}
|
||||
+93
-56
@@ -52,31 +52,63 @@ pub struct GenericOpenaiProvider {
|
||||
model: String,
|
||||
provider_name: &'static str,
|
||||
extra_headers: Vec<(String, String)>,
|
||||
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
|
||||
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter,
|
||||
/// 因此单独存储以便错误消息与配置保持一致。
|
||||
timeout_secs: u64,
|
||||
}
|
||||
|
||||
impl GenericOpenaiProvider {
|
||||
/// 基础构造器。
|
||||
pub fn new_with_name(
|
||||
/// 一次性构造 —— `create_provider` 路径专用,避免 `new_with_name` + `with_client` 的双重 client 构造。
|
||||
///
|
||||
/// 调用方负责预先构造好带正确 timeout 的 `http_client`。`extra_headers` 与 `timeout_secs`
|
||||
/// 一并设置字段,避免后续修改。
|
||||
pub(crate) fn from_parts(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
provider_name: &'static str,
|
||||
http_client: Client,
|
||||
extra_headers: Vec<(String, String)>,
|
||||
timeout_secs: u64,
|
||||
) -> Self {
|
||||
let http_client = Client::builder()
|
||||
.timeout(Duration::from_secs(120))
|
||||
.build()
|
||||
.expect("创建 HTTP 客户端失败");
|
||||
|
||||
Self {
|
||||
http_client,
|
||||
base_url,
|
||||
api_key,
|
||||
model,
|
||||
provider_name,
|
||||
extra_headers: Vec::new(),
|
||||
extra_headers,
|
||||
timeout_secs,
|
||||
}
|
||||
}
|
||||
|
||||
/// 基础构造器。
|
||||
///
|
||||
/// `timeout_secs` 应用于 `reqwest::Client` 的请求超时配置。
|
||||
/// 应由 `ProviderConfig::timeout_secs` 传入(调用方如不知道,可传 30)。
|
||||
pub fn new_with_name(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
provider_name: &'static str,
|
||||
timeout_secs: u64,
|
||||
) -> Self {
|
||||
let http_client = Client::builder()
|
||||
.timeout(Duration::from_secs(timeout_secs))
|
||||
.build()
|
||||
.expect("创建 HTTP 客户端失败");
|
||||
Self::from_parts(
|
||||
base_url,
|
||||
api_key,
|
||||
model,
|
||||
provider_name,
|
||||
http_client,
|
||||
Vec::new(),
|
||||
timeout_secs,
|
||||
)
|
||||
}
|
||||
|
||||
/// 带额外请求头的构造器(如 Qwen 需要 SSE 启用头)。
|
||||
pub fn new_with_name_and_headers(
|
||||
base_url: String,
|
||||
@@ -84,10 +116,21 @@ impl GenericOpenaiProvider {
|
||||
model: String,
|
||||
provider_name: &'static str,
|
||||
extra_headers: Vec<(String, String)>,
|
||||
timeout_secs: u64,
|
||||
) -> Self {
|
||||
let mut base = Self::new_with_name(base_url, api_key, model, provider_name);
|
||||
base.extra_headers = extra_headers;
|
||||
base
|
||||
let http_client = Client::builder()
|
||||
.timeout(Duration::from_secs(timeout_secs))
|
||||
.build()
|
||||
.expect("创建 HTTP 客户端失败");
|
||||
Self::from_parts(
|
||||
base_url,
|
||||
api_key,
|
||||
model,
|
||||
provider_name,
|
||||
http_client,
|
||||
extra_headers,
|
||||
timeout_secs,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn with_client(mut self, client: Client) -> Self {
|
||||
@@ -119,10 +162,10 @@ impl GenericOpenaiProvider {
|
||||
Ok(builder.json(body))
|
||||
}
|
||||
|
||||
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
|
||||
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
|
||||
if e.is_timeout() {
|
||||
LlmError::Timeout {
|
||||
duration: Duration::from_secs(120),
|
||||
duration: Duration::from_secs(self.timeout_secs),
|
||||
}
|
||||
} else if e.is_connect() {
|
||||
LlmError::Other(format!("连接失败: {}", e))
|
||||
@@ -135,9 +178,7 @@ impl GenericOpenaiProvider {
|
||||
///
|
||||
/// ponytail: Qwen 等部分 OpenAI-compatible 提供方可能返回非标准 error body
|
||||
/// (无法解析为 JSON),此处直接用 status code + 原始 body 兜底。
|
||||
async fn handle_error_response(
|
||||
response: reqwest::Response,
|
||||
) -> LlmError {
|
||||
async fn handle_error_response(response: reqwest::Response) -> LlmError {
|
||||
let status = response.status().as_u16();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
|
||||
@@ -148,10 +189,7 @@ impl GenericOpenaiProvider {
|
||||
// 与 OpenAI 完全一致;DeepSeek/Qwen 通常遵循。
|
||||
LlmError::RateLimit { retry_after: None }
|
||||
}
|
||||
_ if status >= 500 => LlmError::Request {
|
||||
status,
|
||||
body,
|
||||
},
|
||||
_ if status >= 500 => LlmError::Request { status, body },
|
||||
_ if status == 400 && body.contains("context_length_exceeded") => {
|
||||
LlmError::ContextLength {
|
||||
actual: 0,
|
||||
@@ -188,7 +226,7 @@ impl GenericOpenaiProvider {
|
||||
Some(
|
||||
tool_defs
|
||||
.into_iter()
|
||||
.map(|t| OpenaiTool::Function { function: t })
|
||||
.map(|t| OpenaiTool::Function { function: t.into() })
|
||||
.collect(),
|
||||
)
|
||||
};
|
||||
@@ -200,7 +238,9 @@ impl GenericOpenaiProvider {
|
||||
stop_sequences[0].clone(),
|
||||
))
|
||||
} else {
|
||||
Some(crate::llm::types::shared::StopSequence::Multiple(stop_sequences))
|
||||
Some(crate::llm::types::shared::StopSequence::Multiple(
|
||||
stop_sequences,
|
||||
))
|
||||
};
|
||||
|
||||
let frequency_penalty = request.get_extra_opt("frequency_penalty");
|
||||
@@ -245,9 +285,7 @@ impl GenericOpenaiProvider {
|
||||
let stop_reason = match choice.finish_reason {
|
||||
Some(FinishReason::Stop) => StopReason::Stop,
|
||||
Some(FinishReason::Length) => StopReason::Length,
|
||||
Some(FinishReason::ToolCalls) | Some(FinishReason::FunctionCall) => {
|
||||
StopReason::ToolUse
|
||||
}
|
||||
Some(FinishReason::ToolCalls) | Some(FinishReason::FunctionCall) => StopReason::ToolUse,
|
||||
Some(FinishReason::ContentFilter) => StopReason::ContentFilter,
|
||||
Some(FinishReason::Other) | None => StopReason::Stop,
|
||||
};
|
||||
@@ -263,7 +301,10 @@ impl GenericOpenaiProvider {
|
||||
}
|
||||
|
||||
/// 非流式 `chat()` 入口。
|
||||
pub async fn chat_blocking(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
pub async fn chat_blocking(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<MessageResponse, LlmError> {
|
||||
let req = self.convert_request(request)?;
|
||||
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
|
||||
|
||||
@@ -280,7 +321,7 @@ impl GenericOpenaiProvider {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!(error = %e, "请求失败");
|
||||
Self::map_reqwest_error(e)
|
||||
self.map_reqwest_error(e)
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
@@ -303,8 +344,7 @@ impl GenericOpenaiProvider {
|
||||
pub async fn chat_stream_inner(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
let mut req = self.convert_request(request)?;
|
||||
req.stream = Some(true);
|
||||
req.stream_options = Some(StreamOptions {
|
||||
@@ -322,7 +362,7 @@ impl GenericOpenaiProvider {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!(error = %e, "流式请求失败");
|
||||
Self::map_reqwest_error(e)
|
||||
self.map_reqwest_error(e)
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
@@ -331,9 +371,9 @@ impl GenericOpenaiProvider {
|
||||
}
|
||||
|
||||
let byte_stream: std::pin::Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> = {
|
||||
let s = response.bytes_stream().map(|r| {
|
||||
r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e)))
|
||||
});
|
||||
let s = response
|
||||
.bytes_stream()
|
||||
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e))));
|
||||
Box::pin(s)
|
||||
};
|
||||
|
||||
@@ -368,8 +408,7 @@ impl LlmProvider for GenericOpenaiProvider {
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
self.chat_stream_inner(request).await
|
||||
}
|
||||
|
||||
@@ -389,12 +428,13 @@ impl LlmProvider for GenericOpenaiProvider {
|
||||
pub struct OpenaiChatProvider(pub GenericOpenaiProvider);
|
||||
|
||||
impl OpenaiChatProvider {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||
Self(GenericOpenaiProvider::new_with_name(
|
||||
base_url,
|
||||
api_key,
|
||||
model,
|
||||
"openai",
|
||||
timeout_secs,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -412,8 +452,7 @@ impl LlmProvider for OpenaiChatProvider {
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
self.0.chat_stream(request).await
|
||||
}
|
||||
|
||||
@@ -436,7 +475,10 @@ enum BlockState {
|
||||
/// 正在累积 text block。
|
||||
InText { block_index: u32 },
|
||||
/// 正在累积 tool_call block。
|
||||
InTool { block_index: u32, tool_call_index: u32 },
|
||||
InTool {
|
||||
block_index: u32,
|
||||
tool_call_index: u32,
|
||||
},
|
||||
/// 正在累积 refusal block。
|
||||
InRefusal { block_index: u32 },
|
||||
}
|
||||
@@ -456,9 +498,7 @@ pub struct ChunkToEventStream {
|
||||
}
|
||||
|
||||
impl ChunkToEventStream {
|
||||
fn new(
|
||||
chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>,
|
||||
) -> Self {
|
||||
fn new(chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>) -> Self {
|
||||
Self {
|
||||
chunks,
|
||||
buffer: String::new(),
|
||||
@@ -500,12 +540,7 @@ impl ChunkToEventStream {
|
||||
let mut events = Vec::new();
|
||||
|
||||
// 元信息:MessageStart(仅在第一次见到 role=assistant 时)。
|
||||
if self.partial.id.is_none()
|
||||
&& chunk
|
||||
.choices
|
||||
.iter()
|
||||
.any(|c| c.delta.role.is_some())
|
||||
{
|
||||
if self.partial.id.is_none() && chunk.choices.iter().any(|c| c.delta.role.is_some()) {
|
||||
events.push(StreamEvent::MessageStart {
|
||||
id: chunk.id.clone(),
|
||||
model: chunk.model.clone(),
|
||||
@@ -650,9 +685,7 @@ impl ChunkToEventStream {
|
||||
self.partial.stop_reason = Some(match fr {
|
||||
FinishReason::Stop => StopReason::Stop,
|
||||
FinishReason::Length => StopReason::Length,
|
||||
FinishReason::ToolCalls | FinishReason::FunctionCall => {
|
||||
StopReason::ToolUse
|
||||
}
|
||||
FinishReason::ToolCalls | FinishReason::FunctionCall => StopReason::ToolUse,
|
||||
FinishReason::ContentFilter => StopReason::ContentFilter,
|
||||
FinishReason::Other => StopReason::Other,
|
||||
});
|
||||
@@ -692,7 +725,9 @@ impl ChunkToEventStream {
|
||||
|
||||
match self.partial.clone().finalize() {
|
||||
Ok(full) => {
|
||||
events.push(StreamEvent::MessageComplete { full_response: full });
|
||||
events.push(StreamEvent::MessageComplete {
|
||||
full_response: full,
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
events.push(StreamEvent::Error {
|
||||
@@ -707,10 +742,7 @@ impl ChunkToEventStream {
|
||||
impl Stream for ChunkToEventStream {
|
||||
type Item = Result<StreamEvent, LlmError>;
|
||||
|
||||
fn poll_next(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Option<Self::Item>> {
|
||||
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
loop {
|
||||
// 先尝试从 buffer 取一行处理
|
||||
if let Some(line) = self.next_line() {
|
||||
@@ -814,6 +846,7 @@ mod tests {
|
||||
"sk-test".into(),
|
||||
"gpt-4o".into(),
|
||||
"openai",
|
||||
30,
|
||||
);
|
||||
let response = provider
|
||||
.chat_blocking(MessageRequest {
|
||||
@@ -842,6 +875,7 @@ mod tests {
|
||||
"sk-test".into(),
|
||||
"gpt-4o".into(),
|
||||
"openai",
|
||||
30,
|
||||
);
|
||||
let err = provider
|
||||
.chat_blocking(MessageRequest {
|
||||
@@ -868,6 +902,7 @@ mod tests {
|
||||
"sk-test".into(),
|
||||
"gpt-4o".into(),
|
||||
"openai",
|
||||
30,
|
||||
);
|
||||
let err = provider
|
||||
.chat_blocking(MessageRequest {
|
||||
@@ -906,6 +941,7 @@ data: [DONE]\n\n";
|
||||
"sk-test".into(),
|
||||
"gpt-4o".into(),
|
||||
"openai",
|
||||
30,
|
||||
);
|
||||
let mut stream = provider
|
||||
.chat_stream_inner(MessageRequest {
|
||||
@@ -974,6 +1010,7 @@ data: [DONE]\n\n";
|
||||
"k".into(),
|
||||
"gpt-4o".into(),
|
||||
"openai",
|
||||
30,
|
||||
);
|
||||
let ir = provider.convert_response(resp).unwrap();
|
||||
assert_eq!(ir.stop_reason, StopReason::ToolUse);
|
||||
|
||||
@@ -15,12 +15,12 @@ use std::pin::Pin;
|
||||
use async_trait::async_trait;
|
||||
use futures_core::Stream;
|
||||
|
||||
use super::openai::GenericOpenaiProvider;
|
||||
use super::ProviderCapabilities;
|
||||
use super::openai::GenericOpenaiProvider;
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::provider::LlmProvider;
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
||||
use crate::llm::provider::LlmProvider;
|
||||
|
||||
// =============================================================================
|
||||
// DeepSeek
|
||||
@@ -29,7 +29,7 @@ use crate::llm::provider::LlmProvider;
|
||||
pub struct DeepSeekProvider(pub GenericOpenaiProvider);
|
||||
|
||||
impl DeepSeekProvider {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://api.deepseek.com".to_string()
|
||||
} else {
|
||||
@@ -40,26 +40,27 @@ impl DeepSeekProvider {
|
||||
api_key,
|
||||
model,
|
||||
"deepseek",
|
||||
timeout_secs,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
impl DeepSeekProvider {
|
||||
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
|
||||
pub fn with_client(self, client: reqwest::Client) -> Self {
|
||||
Self(self.0.with_client(client))
|
||||
}
|
||||
|
||||
/// 测试中(带 mock_client)使用的构造器。
|
||||
///
|
||||
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
|
||||
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
|
||||
/// 无超时),不影响行为。
|
||||
pub fn new_with_client(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
client: reqwest::Client,
|
||||
) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://api.deepseek.com".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
let mut inner = GenericOpenaiProvider::new_with_name(url, api_key, model, "deepseek");
|
||||
inner.http_client = client;
|
||||
Self(inner)
|
||||
Self::new(base_url, api_key, model, 30).with_client(client)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -72,8 +73,7 @@ impl LlmProvider for DeepSeekProvider {
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
self.0.chat_stream(request).await
|
||||
}
|
||||
|
||||
@@ -91,7 +91,7 @@ impl LlmProvider for DeepSeekProvider {
|
||||
pub struct QwenProvider(pub GenericOpenaiProvider);
|
||||
|
||||
impl QwenProvider {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
|
||||
} else {
|
||||
@@ -104,31 +104,28 @@ impl QwenProvider {
|
||||
model,
|
||||
"qwen",
|
||||
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
||||
timeout_secs,
|
||||
);
|
||||
Self(inner)
|
||||
}
|
||||
|
||||
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
|
||||
pub fn with_client(self, client: reqwest::Client) -> Self {
|
||||
Self(self.0.with_client(client))
|
||||
}
|
||||
|
||||
/// 测试构造器。
|
||||
///
|
||||
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
|
||||
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
|
||||
/// 无超时),不影响行为。
|
||||
pub fn new_with_client(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
client: reqwest::Client,
|
||||
) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
let mut inner = GenericOpenaiProvider::new_with_name_and_headers(
|
||||
url,
|
||||
api_key,
|
||||
model,
|
||||
"qwen",
|
||||
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
||||
);
|
||||
inner.http_client = client;
|
||||
Self(inner)
|
||||
Self::new(base_url, api_key, model, 30).with_client(client)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -141,8 +138,7 @@ impl LlmProvider for QwenProvider {
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||
self.0.chat_stream(request).await
|
||||
}
|
||||
|
||||
@@ -156,8 +152,8 @@ impl LlmProvider for QwenProvider {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::message::Message as IrMessage;
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use serde_json::json;
|
||||
use wiremock::matchers::{method, path};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||
@@ -182,11 +178,8 @@ mod tests {
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = DeepSeekProvider::new(
|
||||
server.uri(),
|
||||
"sk-test".into(),
|
||||
"deepseek-chat".into(),
|
||||
);
|
||||
let provider =
|
||||
DeepSeekProvider::new(server.uri(), "sk-test".into(), "deepseek-chat".into(), 30);
|
||||
let response = provider
|
||||
.chat(MessageRequest {
|
||||
model: "deepseek-chat".into(),
|
||||
@@ -219,7 +212,7 @@ mod tests {
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into());
|
||||
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into(), 30);
|
||||
let response = provider
|
||||
.chat(MessageRequest {
|
||||
model: "qwen-plus".into(),
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::provider::{create_provider, LlmProvider, ProviderConfig, ProviderType};
|
||||
use crate::llm::provider::{LlmProvider, ProviderConfig, ProviderType, create_provider};
|
||||
|
||||
/// Provider 注册表 —— 管理多个 LLM Provider 实例。
|
||||
///
|
||||
@@ -61,8 +61,6 @@ impl ProviderRegistry {
|
||||
|
||||
/// 获取默认 Provider。
|
||||
pub fn get_default(&self) -> Option<&dyn LlmProvider> {
|
||||
self.default_name
|
||||
.as_ref()
|
||||
.and_then(|name| self.get(name))
|
||||
self.default_name.as_ref().and_then(|name| self.get(name))
|
||||
}
|
||||
}
|
||||
|
||||
+7
-8
@@ -14,8 +14,8 @@ use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use futures_core::stream::Stream;
|
||||
use futures_util::future::poll_fn;
|
||||
use futures_util::FutureExt;
|
||||
use futures_util::future::poll_fn;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::llm::error::LlmError;
|
||||
@@ -95,9 +95,7 @@ impl Stream for ChunkToLegacyEventStream {
|
||||
}
|
||||
|
||||
if let Some(usage) = &chunk.usage {
|
||||
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate {
|
||||
usage: *usage,
|
||||
}));
|
||||
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate { usage: *usage }));
|
||||
}
|
||||
|
||||
Poll::Ready(None)
|
||||
@@ -143,9 +141,7 @@ fn empty_message_response() -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![],
|
||||
},
|
||||
message: Message::Assistant { content: vec![] },
|
||||
usage: Usage::default(),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: HashMap::new(),
|
||||
@@ -172,7 +168,10 @@ fn map_legacy_to_ir(legacy: LegacyStreamEvent) -> StreamEvent {
|
||||
LegacyStreamEvent::AssistantTextDelta { text } => StreamEvent::TextDelta { text },
|
||||
LegacyStreamEvent::ToolExecutionStarted { input, .. } => {
|
||||
let arguments = serde_json::to_string(&input).unwrap_or_default();
|
||||
StreamEvent::ToolCallArgumentsDelta { index: 0, arguments }
|
||||
StreamEvent::ToolCallArgumentsDelta {
|
||||
index: 0,
|
||||
arguments,
|
||||
}
|
||||
}
|
||||
LegacyStreamEvent::ToolExecutionCompleted { .. } => {
|
||||
// 旧 ToolExecutionCompleted 不在 IR 流协议中——工具执行是消费方职责。
|
||||
|
||||
@@ -20,15 +20,12 @@ use crate::llm::types::shared::ImageDetail;
|
||||
/// 消费方 match 可直接区分文本和图片输入。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[non_exhaustive]
|
||||
pub enum Message {
|
||||
/// 系统提示(User & Assistant 之外的引导指令)。
|
||||
System {
|
||||
content: Vec<ContentBlock>,
|
||||
},
|
||||
System { content: Vec<ContentBlock> },
|
||||
/// 用户输入。
|
||||
User {
|
||||
content: Vec<ContentBlock>,
|
||||
},
|
||||
User { content: Vec<ContentBlock> },
|
||||
/// 用户的图片输入(快捷构造,免去构造 ContentBlock 的 boilerplate)。
|
||||
UserImage {
|
||||
data: String,
|
||||
@@ -36,9 +33,7 @@ pub enum Message {
|
||||
detail: ImageDetail,
|
||||
},
|
||||
/// Assistant 回复内容块(可能包含 text、thinking、tool_use 等多种 block 的混合)。
|
||||
Assistant {
|
||||
content: Vec<ContentBlock>,
|
||||
},
|
||||
Assistant { content: Vec<ContentBlock> },
|
||||
/// 工具调用结果。
|
||||
ToolResult {
|
||||
tool_call_id: String,
|
||||
@@ -103,6 +98,7 @@ impl Message {
|
||||
/// block 的逃生舱。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[non_exhaustive]
|
||||
pub enum ContentBlock {
|
||||
/// 纯文本。
|
||||
Text { text: String },
|
||||
@@ -130,10 +126,7 @@ pub enum ContentBlock {
|
||||
signature: Option<String>,
|
||||
},
|
||||
/// 逃生舱:Provider 特定 block 透传(OpenAI Response 内置工具等)。
|
||||
Extension {
|
||||
kind: String,
|
||||
data: Value,
|
||||
},
|
||||
Extension { kind: String, data: Value },
|
||||
}
|
||||
|
||||
/// 内容块类型标签 —— 用于 `StreamEvent::ContentBlockStart.block_type`。
|
||||
@@ -141,6 +134,7 @@ pub enum ContentBlock {
|
||||
/// 用途:在流式场景中,Provider 先下发 block 类型,再下发 block 内容增量。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[non_exhaustive]
|
||||
pub enum ContentBlockType {
|
||||
/// 文本块。
|
||||
Text,
|
||||
@@ -349,9 +343,7 @@ mod tests {
|
||||
fn message_roundtrip_each_variant() {
|
||||
let msgs = vec![
|
||||
Message::System {
|
||||
content: vec![ContentBlock::Text {
|
||||
text: "sys".into(),
|
||||
}],
|
||||
content: vec![ContentBlock::Text { text: "sys".into() }],
|
||||
},
|
||||
Message::User {
|
||||
content: vec![ContentBlock::Text {
|
||||
@@ -376,9 +368,7 @@ mod tests {
|
||||
},
|
||||
Message::ToolResult {
|
||||
tool_call_id: "call_1".into(),
|
||||
content: vec![ContentBlock::Text {
|
||||
text: "ok".into(),
|
||||
}],
|
||||
content: vec![ContentBlock::Text { text: "ok".into() }],
|
||||
is_error: true,
|
||||
},
|
||||
];
|
||||
|
||||
@@ -26,7 +26,7 @@ pub use shared::{
|
||||
AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier,
|
||||
StopSequence,
|
||||
};
|
||||
pub use tool::{FunctionCall, OpenaiToolCall, OpenaiToolDefinition};
|
||||
pub use tool::{FunctionCall, OpenaiToolCall, ToolDef};
|
||||
pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage};
|
||||
|
||||
// Re-export IR 内容块 / 消息类型供 `types::ContentBlock` 等历史路径消费。
|
||||
@@ -96,7 +96,3 @@ impl From<ChatResponse> for OpenaiChatChunk {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 工具定义别名(无新类型冲突,保留)。
|
||||
#[deprecated(since = "0.1.0", note = "ToolDefinition 仍直接对应 OpenAI wire-format;未来 v0.2 引入 IR 工具类型后会再次更新")]
|
||||
pub type ToolDefinition = OpenaiToolDefinition;
|
||||
|
||||
@@ -11,18 +11,21 @@ pub struct StreamOptions {
|
||||
pub include_obfuscation: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
#[derive(Default)]
|
||||
#[derive(Debug, Clone, Default)]
|
||||
#[non_exhaustive]
|
||||
pub enum ToolChoice {
|
||||
#[default]
|
||||
None,
|
||||
Auto,
|
||||
Required,
|
||||
Named { name: String },
|
||||
AllowedTools { tool_names: Vec<String> },
|
||||
Named {
|
||||
name: String,
|
||||
},
|
||||
AllowedTools {
|
||||
tool_names: Vec<String>,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
impl Serialize for ToolChoice {
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
|
||||
+48
-16
@@ -10,14 +10,14 @@ use thiserror::Error;
|
||||
|
||||
use crate::llm::types::message::Message;
|
||||
use crate::llm::types::request::ToolChoice;
|
||||
use crate::llm::types::tool::OpenaiToolDefinition;
|
||||
use crate::llm::types::tool::ToolDef;
|
||||
|
||||
/// Provider 无关的请求类型。
|
||||
///
|
||||
/// 设计要点:
|
||||
/// - `system` 字段不存在;system 提示由调用方通过 `Message::System` 在 `messages` 中表达。
|
||||
/// - `tools` / `tool_choice` 直接复用现有 `OpenaiToolDefinition` / `ToolChoice`
|
||||
/// (10a §251 决策:先复用旧类型,Phase 2 切换为新 `ToolDefinition` 后再调整)。
|
||||
/// - `tools` 使用 Provider 无关的 `ToolDef` IR;各 Provider 适配层在 `convert_request`
|
||||
/// 中转换为对应 wire format。`tool_choice` 复用现有 `ToolChoice`。
|
||||
/// - `extra` 作为逃生舱:Provider 特定字段(`web_search_options`、`previous_response_id` 等)
|
||||
/// 通过 `extra.set_extra / get_extra` 传递,避免持续膨胀本结构体。
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
@@ -26,8 +26,8 @@ pub struct MessageRequest {
|
||||
pub model: String,
|
||||
/// 消息列表(包含 system / user / assistant / tool_result 等所有变体)。
|
||||
pub messages: Vec<Message>,
|
||||
/// 工具定义列表。
|
||||
pub tools: Vec<OpenaiToolDefinition>,
|
||||
/// 工具定义列表(Provider 无关 IR)。
|
||||
pub tools: Vec<ToolDef>,
|
||||
/// 工具选择策略。
|
||||
pub tool_choice: ToolChoice,
|
||||
/// 最大输出 token 数。
|
||||
@@ -127,14 +127,9 @@ mod tests {
|
||||
#[test]
|
||||
fn extra_set_and_get_roundtrip() {
|
||||
let mut req = MessageRequest::default();
|
||||
req.set_extra(
|
||||
"previous_response_id",
|
||||
"resp_abc123",
|
||||
);
|
||||
req.set_extra("previous_response_id", "resp_abc123");
|
||||
|
||||
let v: Option<String> = req
|
||||
.get_extra("previous_response_id")
|
||||
.expect("get_extra ok");
|
||||
let v: Option<String> = req.get_extra("previous_response_id").expect("get_extra ok");
|
||||
assert_eq!(v.as_deref(), Some("resp_abc123"));
|
||||
|
||||
let missing: Option<String> = req.get_extra("missing").expect("missing ok");
|
||||
@@ -174,10 +169,7 @@ mod tests {
|
||||
}
|
||||
|
||||
let opts: Options = req.get_extra_as().expect("get_extra_as ok");
|
||||
assert_eq!(
|
||||
opts.web_search_options.search_context_size,
|
||||
"high"
|
||||
);
|
||||
assert_eq!(opts.web_search_options.search_context_size, "high");
|
||||
assert_eq!(opts.user.as_deref(), Some("u_123"));
|
||||
}
|
||||
|
||||
@@ -206,4 +198,44 @@ mod tests {
|
||||
assert_eq!(decoded.stream, req.stream);
|
||||
assert_eq!(decoded.extra.get("trace_id"), Some(&json!("t-1")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn message_request_with_tools_roundtrip() {
|
||||
// 验证 ToolDef 的 serde 属性与 OpenaiToolDefinition 一致:
|
||||
// 同名字段(name/description/parameters)序列化结果应一致。
|
||||
let params = json!({
|
||||
"type": "object",
|
||||
"properties": {"x": {"type": "number"}},
|
||||
"required": ["x"],
|
||||
});
|
||||
let tool = super::ToolDef {
|
||||
name: "add".to_string(),
|
||||
description: Some("add two numbers".to_string()),
|
||||
parameters: params.clone(),
|
||||
};
|
||||
let req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
tools: vec![tool],
|
||||
tool_choice: ToolChoice::Auto,
|
||||
max_tokens: None,
|
||||
temperature: None,
|
||||
top_p: None,
|
||||
stop_sequences: vec![],
|
||||
stream: false,
|
||||
thinking: None,
|
||||
extra: HashMap::new(),
|
||||
};
|
||||
|
||||
let json = serde_json::to_string(&req).expect("serialize");
|
||||
// 验证反序列化能还原所有字段(包括嵌套 parameters)
|
||||
let decoded: MessageRequest = serde_json::from_str(&json).expect("deserialize");
|
||||
assert_eq!(decoded.tools.len(), 1);
|
||||
assert_eq!(decoded.tools[0].name, "add");
|
||||
assert_eq!(decoded.tools[0].description.as_deref(), Some("add two numbers"));
|
||||
assert_eq!(decoded.tools[0].parameters, params);
|
||||
|
||||
// 验证序列化 JSON 不含 ToolDef 没有的字段(如 strict),保持 wire-format 兼容
|
||||
assert!(!json.contains("strict"), "ToolDef 序列化不应包含 strict 字段");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,8 +2,8 @@ use crate::llm::types::openai_message::OpenaiChatMessage;
|
||||
use crate::llm::types::shared::{FinishReason, ServiceTier};
|
||||
use crate::llm::types::tool::OpenaiToolCall;
|
||||
use crate::llm::types::usage::Usage;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use crate::llm::types::{ContentField, OpenaiContentPart};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TokenLogprob {
|
||||
@@ -135,11 +135,7 @@ impl From<OpenaiChatMessage> for Delta {
|
||||
text.push_str(&t);
|
||||
}
|
||||
}
|
||||
if text.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(text)
|
||||
}
|
||||
if text.is_empty() { None } else { Some(text) }
|
||||
}
|
||||
},
|
||||
refusal: None,
|
||||
|
||||
@@ -19,6 +19,7 @@ use crate::llm::types::usage::{CompletionTokensDetails, PromptTokensDetails, Usa
|
||||
/// Phase 2 完成时统一收敛。
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[non_exhaustive]
|
||||
pub enum StopReason {
|
||||
/// 自然停止。
|
||||
Stop,
|
||||
@@ -164,11 +165,15 @@ pub enum ContentBlockBuilder {
|
||||
/// `thinking_signature`,最终通过 `finalize()` 回填到 `full_response` 的 `Thinking` block 中。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[non_exhaustive]
|
||||
pub enum StreamEvent {
|
||||
/// 消息开始(元信息)。
|
||||
MessageStart { id: String, model: String },
|
||||
/// 内容块开始(告知块类型,携带 id/name for ToolUse)。
|
||||
ContentBlockStart { index: u32, block_type: ContentBlockType },
|
||||
ContentBlockStart {
|
||||
index: u32,
|
||||
block_type: ContentBlockType,
|
||||
},
|
||||
/// 内容块结束标记。
|
||||
ContentBlockEnd { index: u32 },
|
||||
/// 文本增量。
|
||||
@@ -187,6 +192,24 @@ pub enum StreamEvent {
|
||||
MessageComplete { full_response: MessageResponse },
|
||||
/// 错误事件。
|
||||
Error { message: String },
|
||||
/// 工具开始执行 —— 在 `ToolCallEnd` 之后、`registry.invoke_all` 之前发出。
|
||||
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
|
||||
ToolExecutionStarted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 工具参数(JSON 字符串形式),用于 UI 展示
|
||||
arguments: String,
|
||||
},
|
||||
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
|
||||
ToolExecutionCompleted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 结果摘要(由 `CycleConfig.max_tool_result_bytes` 截断,默认 65536 字节/字符边界安全),
|
||||
/// 用于 UI 反馈。完整结果已在内部 `messages` 中作为 `ToolResult` 回传给 LLM。
|
||||
result_summary: String,
|
||||
/// 是否执行出错
|
||||
is_error: bool,
|
||||
},
|
||||
}
|
||||
|
||||
/// 流式响应累积状态。
|
||||
@@ -320,9 +343,8 @@ impl PartialMessageResponse {
|
||||
true
|
||||
}
|
||||
StreamEvent::ToolCallArgumentsDelta { index, arguments } => {
|
||||
if let Some(ContentBlockBuilder::ToolUse {
|
||||
arguments: buf, ..
|
||||
}) = self.blocks.get_mut(index)
|
||||
if let Some(ContentBlockBuilder::ToolUse { arguments: buf, .. }) =
|
||||
self.blocks.get_mut(index)
|
||||
{
|
||||
buf.push_str(arguments);
|
||||
}
|
||||
@@ -348,6 +370,9 @@ impl PartialMessageResponse {
|
||||
self.is_errored = true;
|
||||
false
|
||||
}
|
||||
// 元事件:不参与内容块累积,不修改 partial 状态
|
||||
//(Phase 9 —— 工具执行透明化,由 run_tool_loop 在工具前后插入)
|
||||
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -360,9 +385,7 @@ impl PartialMessageResponse {
|
||||
let mut content_blocks = Vec::with_capacity(self.blocks.len());
|
||||
for (idx, builder) in self.blocks {
|
||||
let block = Self::builder_to_block(idx, builder, self.thinking_signature.as_deref())
|
||||
.map_err(|e| LlmError::Other(format!(
|
||||
"partial 块 #{idx} finalize 失败: {e}"
|
||||
)))?;
|
||||
.map_err(|e| LlmError::Other(format!("partial 块 #{idx} finalize 失败: {e}")))?;
|
||||
content_blocks.push(block);
|
||||
}
|
||||
|
||||
@@ -743,10 +766,7 @@ mod tests {
|
||||
Message::Assistant { content } => {
|
||||
assert_eq!(content.len(), 2);
|
||||
match (&content[0], &content[1]) {
|
||||
(
|
||||
ContentBlock::Text { text: t1 },
|
||||
ContentBlock::Text { text: t2 },
|
||||
) => {
|
||||
(ContentBlock::Text { text: t1 }, ContentBlock::Text { text: t2 }) => {
|
||||
assert_eq!(t1, "first");
|
||||
assert_eq!(t2, "second");
|
||||
}
|
||||
@@ -772,9 +792,7 @@ mod tests {
|
||||
index: 0,
|
||||
block_type: ContentBlockType::Text,
|
||||
},
|
||||
StreamEvent::TextDelta {
|
||||
text: "x".into(),
|
||||
},
|
||||
StreamEvent::TextDelta { text: "x".into() },
|
||||
StreamEvent::ContentBlockEnd { index: 0 },
|
||||
StreamEvent::MessageComplete {
|
||||
full_response: empty_response(),
|
||||
@@ -832,15 +850,9 @@ mod tests {
|
||||
block_type: ContentBlockType::Text,
|
||||
},
|
||||
StreamEvent::ContentBlockEnd { index: 0 },
|
||||
StreamEvent::TextDelta {
|
||||
text: "t".into(),
|
||||
},
|
||||
StreamEvent::ThinkingDelta {
|
||||
text: "p".into(),
|
||||
},
|
||||
StreamEvent::RefusalDelta {
|
||||
text: "r".into(),
|
||||
},
|
||||
StreamEvent::TextDelta { text: "t".into() },
|
||||
StreamEvent::ThinkingDelta { text: "p".into() },
|
||||
StreamEvent::RefusalDelta { text: "r".into() },
|
||||
StreamEvent::ToolCallArgumentsDelta {
|
||||
index: 1,
|
||||
arguments: "{\"x\":1}".into(),
|
||||
|
||||
@@ -13,6 +13,7 @@ pub enum Role {
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[non_exhaustive]
|
||||
pub enum FinishReason {
|
||||
Stop,
|
||||
Length,
|
||||
@@ -67,6 +68,7 @@ pub enum StopSequence {
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case", tag = "type")]
|
||||
#[non_exhaustive]
|
||||
pub enum ResponseFormat {
|
||||
Text,
|
||||
JsonObject,
|
||||
|
||||
@@ -1,6 +1,25 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
/// Provider 无关的工具定义 IR(v0.2 引入,替换 `ToolDefinition` 别名)。
|
||||
///
|
||||
/// 字段最小化:仅承载跨 Provider 公共的概念(name、description、parameters)。
|
||||
/// OpenAI 专属 `strict` 字段不在此表达,由 OpenAI 适配层通过
|
||||
/// `MessageRequest.extra` 逃生舱在 `convert_request` 内补充。
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
pub struct ToolDef {
|
||||
pub name: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
#[serde(default)]
|
||||
pub parameters: Value,
|
||||
}
|
||||
|
||||
/// 旧 OpenAI wire-format 工具定义(v0.2 降级为 `#[doc(hidden)]`)。
|
||||
///
|
||||
/// 由 `ToolDef` 替代;保留仅供 OpenAI 适配层消费 `ToolDef → OpenaiToolDefinition`
|
||||
/// 转换与外部反序列化兼容路径使用,不作为公共 API。
|
||||
#[doc(hidden)]
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
pub struct OpenaiToolDefinition {
|
||||
pub name: String,
|
||||
@@ -12,6 +31,27 @@ pub struct OpenaiToolDefinition {
|
||||
pub strict: Option<bool>,
|
||||
}
|
||||
|
||||
impl From<ToolDef> for OpenaiToolDefinition {
|
||||
fn from(t: ToolDef) -> Self {
|
||||
Self {
|
||||
name: t.name,
|
||||
description: t.description,
|
||||
parameters: t.parameters,
|
||||
strict: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<OpenaiToolDefinition> for ToolDef {
|
||||
fn from(t: OpenaiToolDefinition) -> Self {
|
||||
Self {
|
||||
name: t.name,
|
||||
description: t.description,
|
||||
parameters: t.parameters,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct FunctionCall {
|
||||
pub name: String,
|
||||
|
||||
+3
-3
@@ -12,11 +12,11 @@ pub use conversation::{ConversationMemory, ConversationMemoryConfig};
|
||||
pub use error::MemoryError;
|
||||
pub use knowledge::KnowledgeStore;
|
||||
pub use retriever::MemoryRetriever;
|
||||
pub use store::{InMemoryStore, MemoryStore};
|
||||
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
|
||||
|
||||
// 低频类型(配置/高级使用)
|
||||
pub use conversation::MemoryStrategy;
|
||||
pub use knowledge::{PageIndexEntry, KNOWLEDGE_PREFIX};
|
||||
pub use retriever::{RetrieverConfig, RetrievalResult, ScoredItem};
|
||||
pub use knowledge::{KNOWLEDGE_PREFIX, PageIndexEntry};
|
||||
pub use retriever::{RetrievalResult, RetrieverConfig, ScoredItem};
|
||||
pub use store::{EvictionConfig, EvictionPolicy};
|
||||
pub use types::{KnowledgePage, MemoryFilter, MemoryItem};
|
||||
|
||||
+21
-14
@@ -12,6 +12,7 @@ use crate::memory::types::MemoryItem;
|
||||
|
||||
/// 对话消息管理策略。
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
#[non_exhaustive]
|
||||
pub enum MemoryStrategy {
|
||||
/// 滑动窗口:达到上限时删除最旧消息。
|
||||
SlidingWindow,
|
||||
@@ -160,7 +161,12 @@ impl ConversationMemory {
|
||||
}
|
||||
|
||||
fn make_message_id(&self, index: usize, now: &OffsetDateTime) -> String {
|
||||
format!("{}{:010}_{}", self.session_prefix(), index, now.unix_timestamp_nanos())
|
||||
format!(
|
||||
"{}{:010}_{}",
|
||||
self.session_prefix(),
|
||||
index,
|
||||
now.unix_timestamp_nanos()
|
||||
)
|
||||
}
|
||||
|
||||
async fn maybe_evict_and_compact(&mut self) {
|
||||
@@ -175,15 +181,16 @@ impl ConversationMemory {
|
||||
}
|
||||
|
||||
if let Some(ref compact_config) = self.config.compact_config
|
||||
&& should_compact(&self.messages, compact_config, &self.compact_state) {
|
||||
let keep_recent = compact_config.keep_recent;
|
||||
let freed = microcompact(&mut self.messages, keep_recent);
|
||||
if freed > 0 {
|
||||
self.compact_state.record_success();
|
||||
} else {
|
||||
let _ = self.compact_state.record_failure();
|
||||
}
|
||||
&& should_compact(&self.messages, compact_config, &self.compact_state)
|
||||
{
|
||||
let keep_recent = compact_config.keep_recent;
|
||||
let freed = microcompact(&mut self.messages, keep_recent);
|
||||
if freed > 0 {
|
||||
self.compact_state.record_success();
|
||||
} else {
|
||||
let _ = self.compact_state.record_failure();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -196,7 +203,8 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn add_and_get_history() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let mut conv = ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
|
||||
let mut conv =
|
||||
ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
|
||||
conv.add_message(Message::user_text("hello")).await.unwrap();
|
||||
conv.add_message(Message::user_text("world")).await.unwrap();
|
||||
assert_eq!(conv.len(), 2);
|
||||
@@ -211,9 +219,7 @@ mod tests {
|
||||
conv.add_message(Message::tool_result("call_1", "ok", false))
|
||||
.await
|
||||
.unwrap();
|
||||
conv.add_message(Message::assistant("done"))
|
||||
.await
|
||||
.unwrap();
|
||||
conv.add_message(Message::assistant("done")).await.unwrap();
|
||||
|
||||
let original = conv.get_history().to_vec();
|
||||
assert_eq!(original.len(), 2);
|
||||
@@ -263,7 +269,8 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn clear_empties_messages() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let mut conv = ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
|
||||
let mut conv =
|
||||
ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
|
||||
conv.add_message(Message::user_text("hello")).await.unwrap();
|
||||
assert!(!conv.is_empty());
|
||||
conv.clear().await.unwrap();
|
||||
|
||||
+2
-1
@@ -6,6 +6,7 @@ use thiserror::Error;
|
||||
///
|
||||
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查环境变量、重试)。
|
||||
#[derive(Debug, Error)]
|
||||
#[non_exhaustive]
|
||||
pub enum MemoryError {
|
||||
/// 按 ID 未找到指定记忆条目。可重试——通常是 namespace 拼写错误或条目已被淘汰。
|
||||
#[error("未找到记忆条目 '{0}',请检查 ID 或 namespace 是否正确")]
|
||||
@@ -33,4 +34,4 @@ impl MemoryError {
|
||||
pub fn is_recoverable(&self) -> bool {
|
||||
matches!(self, Self::NotFound(_) | Self::RetrievalError(_))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -57,8 +57,8 @@ impl KnowledgeStore {
|
||||
}
|
||||
let now = OffsetDateTime::now_utc();
|
||||
let id = format!("{KNOWLEDGE_PREFIX}{}", page.id);
|
||||
let content = serde_json::to_string(&page)
|
||||
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
|
||||
let content =
|
||||
serde_json::to_string(&page).map_err(|e| MemoryError::Serialization(e.to_string()))?;
|
||||
let item = MemoryItem {
|
||||
id,
|
||||
content,
|
||||
@@ -128,7 +128,10 @@ impl KnowledgeStore {
|
||||
.filter(|entry| {
|
||||
entry.title.to_lowercase().contains(&needle)
|
||||
|| entry.summary.to_lowercase().contains(&needle)
|
||||
|| entry.tags.iter().any(|t| t.to_lowercase().contains(&needle))
|
||||
|| entry
|
||||
.tags
|
||||
.iter()
|
||||
.any(|t| t.to_lowercase().contains(&needle))
|
||||
})
|
||||
.map(|entry| entry.id.clone())
|
||||
.collect()
|
||||
|
||||
+15
-8
@@ -97,7 +97,11 @@ impl MemoryRetriever {
|
||||
|
||||
// 4. 过滤 → 排序 → 截取
|
||||
items.retain(|i| i.score >= self.config.min_score);
|
||||
items.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
|
||||
items.sort_by(|a, b| {
|
||||
b.score
|
||||
.partial_cmp(&a.score)
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
items.truncate(self.config.max_results);
|
||||
|
||||
Ok(RetrievalResult {
|
||||
@@ -159,12 +163,12 @@ fn char_bigrams(s: &str) -> Vec<String> {
|
||||
|
||||
fn default_stop_words() -> HashSet<String> {
|
||||
[
|
||||
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has",
|
||||
"had", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall",
|
||||
"can", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what",
|
||||
"which", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor",
|
||||
"so", "if", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at",
|
||||
"by", "of", "as", "into", "through", "during", "before", "after", "above", "below",
|
||||
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had",
|
||||
"do", "does", "did", "will", "would", "should", "could", "may", "might", "shall", "can",
|
||||
"this", "that", "these", "those", "it", "its", "they", "them", "their", "what", "which",
|
||||
"who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor", "so", "if",
|
||||
"then", "else", "with", "without", "for", "to", "from", "in", "on", "at", "by", "of", "as",
|
||||
"into", "through", "during", "before", "after", "above", "below",
|
||||
]
|
||||
.iter()
|
||||
.map(|s| s.to_string())
|
||||
@@ -236,7 +240,10 @@ mod tests {
|
||||
min_score: 0.99,
|
||||
};
|
||||
let retriever = MemoryRetriever::new(ks, config);
|
||||
let result = retriever.retrieve("totally unrelated content").await.unwrap();
|
||||
let result = retriever
|
||||
.retrieve("totally unrelated content")
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(result.items.is_empty());
|
||||
}
|
||||
|
||||
|
||||
+7
-259
@@ -1,14 +1,16 @@
|
||||
//! MemoryStore 抽象接口与默认实现。
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
use crate::memory::error::MemoryError;
|
||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||
|
||||
pub mod in_memory;
|
||||
pub mod sqlite_store;
|
||||
|
||||
pub use in_memory::InMemoryStore;
|
||||
pub use sqlite_store::SqliteStore;
|
||||
|
||||
/// 底层记忆存储抽象接口。
|
||||
///
|
||||
/// 下游可实现此 trait 以对接持久化后端(JSON 文件、SQLite、Redis 等)。
|
||||
@@ -32,6 +34,7 @@ pub trait MemoryStore: Send + Sync {
|
||||
|
||||
/// 淘汰策略。
|
||||
#[derive(Debug, Clone)]
|
||||
#[non_exhaustive]
|
||||
pub enum EvictionPolicy {
|
||||
/// 不淘汰(默认)。
|
||||
None,
|
||||
@@ -57,258 +60,3 @@ impl Default for EvictionConfig {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
|
||||
pub struct InMemoryStore {
|
||||
items: Mutex<HashMap<String, MemoryItem>>,
|
||||
eviction: EvictionConfig,
|
||||
/// 自上次淘汰检查以来的写入次数。
|
||||
writes_since_check: Mutex<usize>,
|
||||
}
|
||||
|
||||
impl InMemoryStore {
|
||||
/// 创建一个无淘汰策略的 InMemoryStore。
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
items: Mutex::new(HashMap::new()),
|
||||
eviction: EvictionConfig::default(),
|
||||
writes_since_check: Mutex::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建一个带淘汰配置的 InMemoryStore。
|
||||
pub fn with_eviction(eviction: EvictionConfig) -> Self {
|
||||
Self {
|
||||
items: Mutex::new(HashMap::new()),
|
||||
eviction,
|
||||
writes_since_check: Mutex::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
fn maybe_evict(&self) {
|
||||
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
|
||||
let should_check = {
|
||||
let mut counter = self.writes_since_check.lock().unwrap();
|
||||
*counter += 1;
|
||||
if *counter >= self.eviction.check_interval {
|
||||
*counter = 0;
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
};
|
||||
if !should_check {
|
||||
return;
|
||||
}
|
||||
|
||||
let policy = self.eviction.policy.clone();
|
||||
match policy {
|
||||
EvictionPolicy::None => {}
|
||||
EvictionPolicy::Ttl { ttl_secs } => {
|
||||
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
|
||||
let mut items = self.items.lock().unwrap();
|
||||
items.retain(|_, v| v.created_at > cutoff);
|
||||
}
|
||||
EvictionPolicy::Capacity { max_items } => {
|
||||
let mut items = self.items.lock().unwrap();
|
||||
if items.len() > max_items {
|
||||
let mut vec: Vec<_> = items.drain().collect();
|
||||
// O(n) 部分排序:保留 created_at 最大的 max_items 个
|
||||
vec.select_nth_unstable_by(max_items, |a, b| {
|
||||
b.1.created_at.cmp(&a.1.created_at)
|
||||
});
|
||||
vec.truncate(max_items);
|
||||
*items = vec.into_iter().collect();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for InMemoryStore {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MemoryStore for InMemoryStore {
|
||||
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
|
||||
{
|
||||
let mut items = self.items.lock().unwrap();
|
||||
items.insert(item.id.clone(), item);
|
||||
}
|
||||
self.maybe_evict();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
|
||||
let items = self.items.lock().unwrap();
|
||||
Ok(items.get(id).cloned())
|
||||
}
|
||||
|
||||
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
|
||||
let mut items = self.items.lock().unwrap();
|
||||
items.remove(id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
|
||||
let items = self.items.lock().unwrap();
|
||||
let mut result: Vec<MemoryItem> = items
|
||||
.values()
|
||||
.filter(|v| match &filter.prefix {
|
||||
Some(p) => v.id.starts_with(p),
|
||||
None => true,
|
||||
})
|
||||
.filter(|v| match filter.since {
|
||||
Some(t) => v.created_at > t,
|
||||
None => true,
|
||||
})
|
||||
.cloned()
|
||||
.collect();
|
||||
// 按 created_at 升序排列(最旧在前)
|
||||
result.sort_by_key(|v| v.created_at);
|
||||
// 应用 offset
|
||||
if let Some(offset) = filter.offset {
|
||||
if offset < result.len() {
|
||||
result.drain(..offset);
|
||||
} else {
|
||||
result.clear();
|
||||
}
|
||||
}
|
||||
// 应用 limit
|
||||
if let Some(limit) = filter.limit {
|
||||
result.truncate(limit);
|
||||
}
|
||||
Ok(result)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
fn make_item(id: &str) -> MemoryItem {
|
||||
MemoryItem {
|
||||
id: id.to_string(),
|
||||
content: format!("content-{id}"),
|
||||
metadata: serde_json::json!({}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn save_get_delete_list() {
|
||||
let store = InMemoryStore::new();
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
store.save(make_item("b")).await.unwrap();
|
||||
|
||||
let got = store.get("a").await.unwrap();
|
||||
assert!(got.is_some());
|
||||
assert_eq!(got.unwrap().id, "a");
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
|
||||
store.delete("a").await.unwrap();
|
||||
assert!(store.get("a").await.unwrap().is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn save_is_upsert() {
|
||||
let store = InMemoryStore::new();
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
let mut item = make_item("a");
|
||||
item.content = "updated".to_string();
|
||||
store.save(item).await.unwrap();
|
||||
let got = store.get("a").await.unwrap().unwrap();
|
||||
assert_eq!(got.content, "updated");
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_with_prefix_and_limit() {
|
||||
let store = InMemoryStore::new();
|
||||
store.save(make_item("foo_a")).await.unwrap();
|
||||
store.save(make_item("foo_b")).await.unwrap();
|
||||
store.save(make_item("bar_a")).await.unwrap();
|
||||
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some("foo_".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let list = store.list(&filter).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some("foo_".to_string()),
|
||||
limit: Some(1),
|
||||
..Default::default()
|
||||
};
|
||||
let list = store.list(&filter).await.unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn capacity_eviction() {
|
||||
// 强制每次写入都检查
|
||||
let eviction = EvictionConfig {
|
||||
policy: EvictionPolicy::Capacity { max_items: 2 },
|
||||
check_interval: 1,
|
||||
};
|
||||
let store = InMemoryStore::with_eviction(eviction);
|
||||
// 第一条和第二条共存
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
store.save(make_item("b")).await.unwrap();
|
||||
// 第三条写入触发淘汰:a 或 b 之一被淘汰
|
||||
store.save(make_item("c")).await.unwrap();
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
// 留下的应该是 b 和 c(最新的两个)
|
||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||
assert!(ids.contains(&"b"));
|
||||
assert!(ids.contains(&"c"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ttl_eviction() {
|
||||
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
|
||||
// 改用小 TTL + 睡眠:先 save a,sleep,save b 时 a 已过期被淘汰。
|
||||
let eviction = EvictionConfig {
|
||||
policy: EvictionPolicy::Ttl { ttl_secs: 1 },
|
||||
check_interval: 1,
|
||||
};
|
||||
let store = InMemoryStore::with_eviction(eviction);
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
// 等待超过 1 秒
|
||||
std::thread::sleep(std::time::Duration::from_millis(1100));
|
||||
// 触发淘汰:a 已超过 ttl_secs=1,应被淘汰
|
||||
store.save(make_item("b")).await.unwrap();
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
// 由于 ttl_secs=1,且 b 刚写入,可能刚好处于临界值。
|
||||
// 我们只断言 list 不包含 "a" 即可。
|
||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||
assert!(
|
||||
!ids.contains(&"a"),
|
||||
"expected 'a' to be evicted, but found in {ids:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn none_policy_no_eviction() {
|
||||
let eviction = EvictionConfig {
|
||||
policy: EvictionPolicy::None,
|
||||
check_interval: 1,
|
||||
};
|
||||
let store = InMemoryStore::with_eviction(eviction);
|
||||
for i in 0..100 {
|
||||
store.save(make_item(&format!("item_{i}"))).await.unwrap();
|
||||
}
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 100);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,266 @@
|
||||
//! 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
use crate::memory::error::MemoryError;
|
||||
use crate::memory::store::{EvictionConfig, EvictionPolicy, MemoryStore};
|
||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||
|
||||
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
|
||||
pub struct InMemoryStore {
|
||||
items: Mutex<HashMap<String, MemoryItem>>,
|
||||
eviction: EvictionConfig,
|
||||
/// 自上次淘汰检查以来的写入次数。
|
||||
writes_since_check: Mutex<usize>,
|
||||
}
|
||||
|
||||
impl InMemoryStore {
|
||||
/// 创建一个无淘汰策略的 InMemoryStore。
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
items: Mutex::new(HashMap::new()),
|
||||
eviction: EvictionConfig::default(),
|
||||
writes_since_check: Mutex::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建一个带淘汰配置的 InMemoryStore。
|
||||
pub fn with_eviction(eviction: EvictionConfig) -> Self {
|
||||
Self {
|
||||
items: Mutex::new(HashMap::new()),
|
||||
eviction,
|
||||
writes_since_check: Mutex::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
fn maybe_evict(&self) {
|
||||
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
|
||||
let should_check = {
|
||||
let mut counter = self.writes_since_check.lock().unwrap();
|
||||
*counter += 1;
|
||||
if *counter >= self.eviction.check_interval {
|
||||
*counter = 0;
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
};
|
||||
if !should_check {
|
||||
return;
|
||||
}
|
||||
|
||||
let policy = self.eviction.policy.clone();
|
||||
match policy {
|
||||
EvictionPolicy::None => {}
|
||||
EvictionPolicy::Ttl { ttl_secs } => {
|
||||
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
|
||||
let mut items = self.items.lock().unwrap();
|
||||
items.retain(|_, v| v.created_at > cutoff);
|
||||
}
|
||||
EvictionPolicy::Capacity { max_items } => {
|
||||
let mut items = self.items.lock().unwrap();
|
||||
if items.len() > max_items {
|
||||
let mut vec: Vec<_> = items.drain().collect();
|
||||
// O(n) 部分排序:保留 created_at 最大的 max_items 个
|
||||
vec.select_nth_unstable_by(max_items, |a, b| {
|
||||
b.1.created_at.cmp(&a.1.created_at)
|
||||
});
|
||||
vec.truncate(max_items);
|
||||
*items = vec.into_iter().collect();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for InMemoryStore {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MemoryStore for InMemoryStore {
|
||||
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
|
||||
{
|
||||
let mut items = self.items.lock().unwrap();
|
||||
items.insert(item.id.clone(), item);
|
||||
}
|
||||
self.maybe_evict();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
|
||||
let items = self.items.lock().unwrap();
|
||||
Ok(items.get(id).cloned())
|
||||
}
|
||||
|
||||
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
|
||||
let mut items = self.items.lock().unwrap();
|
||||
items.remove(id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
|
||||
let items = self.items.lock().unwrap();
|
||||
let mut result: Vec<MemoryItem> = items
|
||||
.values()
|
||||
.filter(|v| match &filter.prefix {
|
||||
Some(p) => v.id.starts_with(p),
|
||||
None => true,
|
||||
})
|
||||
.filter(|v| match filter.since {
|
||||
Some(t) => v.created_at > t,
|
||||
None => true,
|
||||
})
|
||||
.cloned()
|
||||
.collect();
|
||||
// 按 created_at 升序排列(最旧在前)
|
||||
result.sort_by_key(|v| v.created_at);
|
||||
// 应用 offset
|
||||
if let Some(offset) = filter.offset {
|
||||
if offset < result.len() {
|
||||
result.drain(..offset);
|
||||
} else {
|
||||
result.clear();
|
||||
}
|
||||
}
|
||||
// 应用 limit
|
||||
if let Some(limit) = filter.limit {
|
||||
result.truncate(limit);
|
||||
}
|
||||
Ok(result)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
fn make_item(id: &str) -> MemoryItem {
|
||||
MemoryItem {
|
||||
id: id.to_string(),
|
||||
content: format!("content-{id}"),
|
||||
metadata: serde_json::json!({}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn save_get_delete_list() {
|
||||
let store = InMemoryStore::new();
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
store.save(make_item("b")).await.unwrap();
|
||||
|
||||
let got = store.get("a").await.unwrap();
|
||||
assert!(got.is_some());
|
||||
assert_eq!(got.unwrap().id, "a");
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
|
||||
store.delete("a").await.unwrap();
|
||||
assert!(store.get("a").await.unwrap().is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn save_is_upsert() {
|
||||
let store = InMemoryStore::new();
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
let mut item = make_item("a");
|
||||
item.content = "updated".to_string();
|
||||
store.save(item).await.unwrap();
|
||||
let got = store.get("a").await.unwrap().unwrap();
|
||||
assert_eq!(got.content, "updated");
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_with_prefix_and_limit() {
|
||||
let store = InMemoryStore::new();
|
||||
store.save(make_item("foo_a")).await.unwrap();
|
||||
store.save(make_item("foo_b")).await.unwrap();
|
||||
store.save(make_item("bar_a")).await.unwrap();
|
||||
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some("foo_".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let list = store.list(&filter).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some("foo_".to_string()),
|
||||
limit: Some(1),
|
||||
..Default::default()
|
||||
};
|
||||
let list = store.list(&filter).await.unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn capacity_eviction() {
|
||||
// 强制每次写入都检查
|
||||
let eviction = EvictionConfig {
|
||||
policy: EvictionPolicy::Capacity { max_items: 2 },
|
||||
check_interval: 1,
|
||||
};
|
||||
let store = InMemoryStore::with_eviction(eviction);
|
||||
// 第一条和第二条共存
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
store.save(make_item("b")).await.unwrap();
|
||||
// 第三条写入触发淘汰:a 或 b 之一被淘汰
|
||||
store.save(make_item("c")).await.unwrap();
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
// 留下的应该是 b 和 c(最新的两个)
|
||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||
assert!(ids.contains(&"b"));
|
||||
assert!(ids.contains(&"c"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ttl_eviction() {
|
||||
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
|
||||
// 改用小 TTL + 睡眠:先 save a,sleep,save b 时 a 已过期被淘汰。
|
||||
let eviction = EvictionConfig {
|
||||
policy: EvictionPolicy::Ttl { ttl_secs: 1 },
|
||||
check_interval: 1,
|
||||
};
|
||||
let store = InMemoryStore::with_eviction(eviction);
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
// 等待超过 1 秒
|
||||
std::thread::sleep(std::time::Duration::from_millis(1100));
|
||||
// 触发淘汰:a 已超过 ttl_secs=1,应被淘汰
|
||||
store.save(make_item("b")).await.unwrap();
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
// 由于 ttl_secs=1,且 b 刚写入,可能刚好处于临界值。
|
||||
// 我们只断言 list 不包含 "a" 即可。
|
||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||
assert!(
|
||||
!ids.contains(&"a"),
|
||||
"expected 'a' to be evicted, but found in {ids:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn none_policy_no_eviction() {
|
||||
let eviction = EvictionConfig {
|
||||
policy: EvictionPolicy::None,
|
||||
check_interval: 1,
|
||||
};
|
||||
let store = InMemoryStore::with_eviction(eviction);
|
||||
for i in 0..100 {
|
||||
store.save(make_item(&format!("item_{i}"))).await.unwrap();
|
||||
}
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 100);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,545 @@
|
||||
//! SqliteStore —— 基于 rusqlite 的持久化 MemoryStore 实现。
|
||||
//!
|
||||
//! 单进程独享、写入串行化(WAL + Mutex),适合本地 Agent 长期持久化场景。
|
||||
|
||||
use std::path::Path;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use rusqlite::{params, params_from_iter, Connection, ErrorCode};
|
||||
use time::format_description::well_known::Rfc3339;
|
||||
use time::OffsetDateTime;
|
||||
use tracing::{debug, error, instrument, warn};
|
||||
|
||||
use crate::memory::error::MemoryError;
|
||||
use crate::memory::store::MemoryStore;
|
||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||
|
||||
const INITIAL_USER_VERSION: i64 = 1;
|
||||
const BUSY_TIMEOUT_MS: i64 = 5000;
|
||||
const WAL_AUTOCHECKPOINT_PAGES: i64 = 1000;
|
||||
|
||||
/// SQLite 持久化后端的 MemoryStore 实现。
|
||||
///
|
||||
/// 设计要点:
|
||||
/// - 单进程独享:`Arc<Mutex<Connection>>` 串行化所有 IO
|
||||
/// - WAL 模式 + `synchronous=NORMAL` 兼顾崩溃安全与吞吐
|
||||
/// - `created_at` 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
|
||||
/// - 所有 IO 通过 `tokio::task::spawn_blocking` 卸载到阻塞线程池
|
||||
pub struct SqliteStore {
|
||||
conn: Arc<Mutex<Connection>>,
|
||||
}
|
||||
|
||||
impl SqliteStore {
|
||||
/// 打开或创建一个 SQLite 数据库。
|
||||
///
|
||||
/// - `path = ":memory:"` 使用内存数据库(测试场景)
|
||||
/// - 其他路径:自动创建父目录;文件已存在则附加打开
|
||||
/// - 启动时执行 `migrate()`,失败立即返回错误
|
||||
#[instrument(skip(path), fields(path = %path.as_ref().display()))]
|
||||
pub fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
|
||||
let path_ref = path.as_ref();
|
||||
let path_str = path_ref.to_string_lossy();
|
||||
|
||||
let conn = if path_str == ":memory:" {
|
||||
Connection::open_in_memory()
|
||||
} else {
|
||||
if let Some(parent) = path_ref.parent()
|
||||
&& !parent.as_os_str().is_empty()
|
||||
{
|
||||
std::fs::create_dir_all(parent).map_err(|e| {
|
||||
MemoryError::Storage(format!(
|
||||
"创建数据库父目录失败 ({}): {}",
|
||||
parent.display(),
|
||||
e
|
||||
))
|
||||
})?;
|
||||
}
|
||||
Connection::open(path_ref)
|
||||
}
|
||||
.map_err(|e| map_sqlite_error(e, "打开数据库"))?;
|
||||
|
||||
migrate(&conn)?;
|
||||
Ok(Self {
|
||||
conn: Arc::new(Mutex::new(conn)),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MemoryStore for SqliteStore {
|
||||
#[instrument(skip(self, item), fields(id = %item.id))]
|
||||
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
|
||||
let conn = Arc::clone(&self.conn);
|
||||
let created_at_str = item
|
||||
.created_at
|
||||
.to_offset(time::UtcOffset::UTC)
|
||||
.format(&Rfc3339)
|
||||
.map_err(|e| MemoryError::Serialization(format!("format created_at: {e}")))?;
|
||||
let metadata_str = serde_json::to_string(&item.metadata)
|
||||
.map_err(|e| MemoryError::Serialization(format!("serialize metadata: {e}")))?;
|
||||
let id = item.id;
|
||||
let content = item.content;
|
||||
|
||||
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
|
||||
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||
conn.execute(
|
||||
"INSERT INTO memory_items (id, content, metadata, created_at) \
|
||||
VALUES (?1, ?2, ?3, ?4) \
|
||||
ON CONFLICT(id) DO UPDATE SET \
|
||||
content=excluded.content, \
|
||||
metadata=excluded.metadata, \
|
||||
created_at=excluded.created_at",
|
||||
params![id, content, metadata_str, created_at_str],
|
||||
)
|
||||
.map_err(|e| map_sqlite_error(e, "保存记忆"))?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||
}
|
||||
|
||||
#[instrument(skip(self, id))]
|
||||
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
|
||||
let conn = Arc::clone(&self.conn);
|
||||
let id_owned = id.to_string();
|
||||
|
||||
tokio::task::spawn_blocking(move || -> Result<Option<MemoryItem>, MemoryError> {
|
||||
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let mut stmt = conn
|
||||
.prepare("SELECT id, content, metadata, created_at FROM memory_items WHERE id = ?1")
|
||||
.map_err(|e| map_sqlite_error(e, "prepare get"))?;
|
||||
let mut rows = stmt
|
||||
.query_map(params![id_owned], row_to_item)
|
||||
.map_err(|e| map_sqlite_error(e, "query get"))?;
|
||||
match rows.next() {
|
||||
None => Ok(None),
|
||||
Some(row) => row
|
||||
.map(Some)
|
||||
.map_err(|e| map_sqlite_error(e, "decode row")),
|
||||
}
|
||||
})
|
||||
.await
|
||||
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||
}
|
||||
|
||||
#[instrument(skip(self, id))]
|
||||
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
|
||||
let conn = Arc::clone(&self.conn);
|
||||
let id_owned = id.to_string();
|
||||
|
||||
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
|
||||
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||
conn.execute(
|
||||
"DELETE FROM memory_items WHERE id = ?1",
|
||||
params![id_owned],
|
||||
)
|
||||
.map_err(|e| map_sqlite_error(e, "delete"))?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||
}
|
||||
|
||||
#[instrument(skip(self, filter))]
|
||||
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
|
||||
let mut sql = String::from(
|
||||
"SELECT id, content, metadata, created_at FROM memory_items WHERE 1=1",
|
||||
);
|
||||
let mut param_values: Vec<String> = Vec::new();
|
||||
let mut ph_idx = 0usize;
|
||||
|
||||
if filter.prefix.is_some() {
|
||||
ph_idx += 1;
|
||||
sql.push_str(&format!(" AND id LIKE ?{ph_idx} || '%'"));
|
||||
}
|
||||
if filter.since.is_some() {
|
||||
ph_idx += 1;
|
||||
sql.push_str(&format!(" AND created_at > ?{ph_idx}"));
|
||||
}
|
||||
// ORDER BY created_at ASC(按时间升序,最旧在前)
|
||||
sql.push_str(" ORDER BY created_at ASC");
|
||||
|
||||
let limit_sql: String = match (filter.limit, filter.offset) {
|
||||
(Some(_), Some(_)) => {
|
||||
ph_idx += 1;
|
||||
let limit_p = ph_idx;
|
||||
ph_idx += 1;
|
||||
let offset_p = ph_idx;
|
||||
format!(" LIMIT ?{limit_p} OFFSET ?{offset_p}")
|
||||
}
|
||||
(Some(_), None) => {
|
||||
ph_idx += 1;
|
||||
let limit_p = ph_idx;
|
||||
format!(" LIMIT ?{limit_p}")
|
||||
}
|
||||
(None, Some(_)) => {
|
||||
// SQLite 中 LIMIT -1 表示无限制
|
||||
ph_idx += 1;
|
||||
let offset_p = ph_idx;
|
||||
format!(" LIMIT -1 OFFSET ?{offset_p}")
|
||||
}
|
||||
(None, None) => String::new(),
|
||||
};
|
||||
sql.push_str(&limit_sql);
|
||||
|
||||
if let Some(p) = &filter.prefix {
|
||||
param_values.push(p.clone());
|
||||
}
|
||||
if let Some(t) = filter.since {
|
||||
let s = t
|
||||
.to_offset(time::UtcOffset::UTC)
|
||||
.format(&Rfc3339)
|
||||
.map_err(|e| MemoryError::Serialization(format!("format since: {e}")))?;
|
||||
param_values.push(s);
|
||||
}
|
||||
if let Some(l) = filter.limit {
|
||||
param_values.push(l.to_string());
|
||||
}
|
||||
if let Some(o) = filter.offset {
|
||||
param_values.push(o.to_string());
|
||||
}
|
||||
|
||||
let conn = Arc::clone(&self.conn);
|
||||
let sql_owned = sql;
|
||||
let param_values_owned = param_values;
|
||||
|
||||
tokio::task::spawn_blocking(move || -> Result<Vec<MemoryItem>, MemoryError> {
|
||||
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let mut stmt = conn
|
||||
.prepare(&sql_owned)
|
||||
.map_err(|e| map_sqlite_error(e, "list prepare"))?;
|
||||
let params_iter: Vec<&dyn rusqlite::ToSql> = param_values_owned
|
||||
.iter()
|
||||
.map(|s| s as &dyn rusqlite::ToSql)
|
||||
.collect();
|
||||
let rows = stmt
|
||||
.query_map(params_from_iter(params_iter), row_to_item)
|
||||
.map_err(|e| map_sqlite_error(e, "list query"))?;
|
||||
let mut result = Vec::new();
|
||||
for row in rows {
|
||||
result.push(row.map_err(|e| map_sqlite_error(e, "list row"))?);
|
||||
}
|
||||
debug!(count = result.len(), "SqliteStore::list 完成");
|
||||
Ok(result)
|
||||
})
|
||||
.await
|
||||
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||
}
|
||||
}
|
||||
|
||||
fn row_to_item(row: &rusqlite::Row<'_>) -> Result<MemoryItem, rusqlite::Error> {
|
||||
let id: String = row.get(0)?;
|
||||
let content: String = row.get(1)?;
|
||||
let metadata_str: String = row.get(2)?;
|
||||
let created_at_str: String = row.get(3)?;
|
||||
|
||||
let metadata: serde_json::Value = serde_json::from_str(&metadata_str).map_err(|e| {
|
||||
rusqlite::Error::FromSqlConversionFailure(2, rusqlite::types::Type::Text, Box::new(e))
|
||||
})?;
|
||||
let created_at = OffsetDateTime::parse(&created_at_str, &Rfc3339).map_err(|e| {
|
||||
rusqlite::Error::FromSqlConversionFailure(3, rusqlite::types::Type::Text, Box::new(e))
|
||||
})?;
|
||||
|
||||
Ok(MemoryItem {
|
||||
id,
|
||||
content,
|
||||
metadata,
|
||||
created_at: created_at.to_offset(time::UtcOffset::UTC),
|
||||
})
|
||||
}
|
||||
|
||||
fn migrate(conn: &Connection) -> Result<(), MemoryError> {
|
||||
conn.pragma_update(None, "journal_mode", "WAL")
|
||||
.map_err(|e| map_sqlite_error(e, "PRAGMA journal_mode"))?;
|
||||
conn.pragma_update(None, "synchronous", "NORMAL")
|
||||
.map_err(|e| map_sqlite_error(e, "PRAGMA synchronous"))?;
|
||||
conn.execute_batch(&format!("PRAGMA busy_timeout = {BUSY_TIMEOUT_MS};"))
|
||||
.map_err(|e| map_sqlite_error(e, "PRAGMA busy_timeout"))?;
|
||||
conn.execute_batch(&format!(
|
||||
"PRAGMA wal_autocheckpoint = {WAL_AUTOCHECKPOINT_PAGES};"
|
||||
))
|
||||
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_autocheckpoint"))?;
|
||||
|
||||
let version: i64 = conn
|
||||
.query_row("PRAGMA user_version", [], |row| row.get(0))
|
||||
.map_err(|e| map_sqlite_error(e, "PRAGMA user_version"))?;
|
||||
|
||||
if version < INITIAL_USER_VERSION {
|
||||
conn.execute_batch(
|
||||
"CREATE TABLE IF NOT EXISTS memory_items (
|
||||
id TEXT PRIMARY KEY,
|
||||
content TEXT NOT NULL,
|
||||
metadata TEXT NOT NULL DEFAULT '{}',
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_items_created_at
|
||||
ON memory_items(created_at);
|
||||
PRAGMA user_version = 1;",
|
||||
)
|
||||
.map_err(|e| map_sqlite_error(e, "create schema v1"))?;
|
||||
}
|
||||
|
||||
let check_result: String = conn
|
||||
.query_row("PRAGMA quick_check", [], |row| row.get(0))
|
||||
.map_err(|e| map_sqlite_error(e, "PRAGMA quick_check"))?;
|
||||
if check_result != "ok" {
|
||||
error!(result = %check_result, "数据库文件 quick_check 失败");
|
||||
return Err(MemoryError::Storage(format!(
|
||||
"数据库文件损坏: {check_result}"
|
||||
)));
|
||||
}
|
||||
|
||||
conn.execute_batch("PRAGMA wal_checkpoint(TRUNCATE);")
|
||||
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_checkpoint"))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn map_sqlite_error(e: rusqlite::Error, ctx: &str) -> MemoryError {
|
||||
match &e {
|
||||
rusqlite::Error::SqliteFailure(err, _) => match err.code {
|
||||
ErrorCode::ConstraintViolation => MemoryError::InvalidInput(format!("{ctx}: {e}")),
|
||||
ErrorCode::DatabaseBusy | ErrorCode::DatabaseLocked => {
|
||||
warn!("SQLite 忙: {e}");
|
||||
MemoryError::Storage(format!("{ctx}: {e}"))
|
||||
}
|
||||
_ => MemoryError::Storage(format!("{ctx}: {e}")),
|
||||
},
|
||||
rusqlite::Error::InvalidQuery
|
||||
| rusqlite::Error::InvalidParameterName(_)
|
||||
| rusqlite::Error::InvalidColumnIndex(_)
|
||||
| rusqlite::Error::InvalidColumnName(_) => {
|
||||
MemoryError::InvalidInput(format!("{ctx}: {e}"))
|
||||
}
|
||||
rusqlite::Error::FromSqlConversionFailure(_, _, _)
|
||||
| rusqlite::Error::ToSqlConversionFailure(_) => {
|
||||
MemoryError::Serialization(format!("{ctx}: {e}"))
|
||||
}
|
||||
_ => MemoryError::Storage(format!("{ctx}: {e}")),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::memory::store::InMemoryStore;
|
||||
use std::sync::Arc;
|
||||
use tempfile::TempDir;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
fn make_item(id: &str) -> MemoryItem {
|
||||
MemoryItem {
|
||||
id: id.to_string(),
|
||||
content: format!("content-{id}"),
|
||||
metadata: serde_json::json!({"id_key": id}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
}
|
||||
}
|
||||
|
||||
fn make_item_at(id: &str, when: OffsetDateTime) -> MemoryItem {
|
||||
MemoryItem {
|
||||
id: id.to_string(),
|
||||
content: format!("content-{id}"),
|
||||
metadata: serde_json::json!({}),
|
||||
created_at: when,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn crud_basic() {
|
||||
let store = SqliteStore::open(":memory:").unwrap();
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
store.save(make_item("b")).await.unwrap();
|
||||
|
||||
let got_a = store.get("a").await.unwrap();
|
||||
assert!(got_a.is_some());
|
||||
assert_eq!(got_a.unwrap().id, "a");
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
|
||||
store.delete("a").await.unwrap();
|
||||
assert!(store.get("a").await.unwrap().is_none());
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
assert_eq!(list[0].id, "b");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn save_is_upsert() {
|
||||
let store = SqliteStore::open(":memory:").unwrap();
|
||||
store.save(make_item("a")).await.unwrap();
|
||||
let mut item = make_item("a");
|
||||
item.content = "updated".to_string();
|
||||
item.metadata = serde_json::json!({"rev": 2});
|
||||
let original_created_at = item.created_at;
|
||||
store.save(item).await.unwrap();
|
||||
|
||||
let got = store.get("a").await.unwrap().unwrap();
|
||||
assert_eq!(got.content, "updated");
|
||||
assert_eq!(got.metadata["rev"], serde_json::json!(2));
|
||||
// created_at 保持调用方传入值
|
||||
assert_eq!(got.created_at, original_created_at);
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_with_prefix() {
|
||||
let store = SqliteStore::open(":memory:").unwrap();
|
||||
store.save(make_item("foo_a")).await.unwrap();
|
||||
store.save(make_item("foo_b")).await.unwrap();
|
||||
store.save(make_item("bar_a")).await.unwrap();
|
||||
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some("foo_".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let list = store.list(&filter).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||
assert!(ids.contains(&"foo_a"));
|
||||
assert!(ids.contains(&"foo_b"));
|
||||
assert!(!ids.contains(&"bar_a"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_with_since_filter() {
|
||||
let store = SqliteStore::open(":memory:").unwrap();
|
||||
let t0 = OffsetDateTime::now_utc();
|
||||
store
|
||||
.save(make_item_at("early", t0 - time::Duration::seconds(60)))
|
||||
.await
|
||||
.unwrap();
|
||||
store.save(make_item_at("middle", t0)).await.unwrap();
|
||||
store
|
||||
.save(make_item_at("late", t0 + time::Duration::seconds(60)))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let filter = MemoryFilter {
|
||||
since: Some(t0 - time::Duration::seconds(1)),
|
||||
..Default::default()
|
||||
};
|
||||
let list = store.list(&filter).await.unwrap();
|
||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||
assert_eq!(list.len(), 2);
|
||||
assert!(ids.contains(&"middle"));
|
||||
assert!(ids.contains(&"late"));
|
||||
assert!(!ids.contains(&"early"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_with_offset_and_limit() {
|
||||
let store = SqliteStore::open(":memory:").unwrap();
|
||||
// 写入 5 条时间递增的记录
|
||||
let base = OffsetDateTime::now_utc() - time::Duration::seconds(5);
|
||||
for i in 0..5 {
|
||||
let mut item = make_item(&format!("item_{i}"));
|
||||
item.created_at = base + time::Duration::seconds(i);
|
||||
store.save(item).await.unwrap();
|
||||
}
|
||||
|
||||
// offset=1, limit=2 -> item_1, item_2
|
||||
let filter = MemoryFilter {
|
||||
offset: Some(1),
|
||||
limit: Some(2),
|
||||
..Default::default()
|
||||
};
|
||||
let list = store.list(&filter).await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
assert_eq!(list[0].id, "item_1");
|
||||
assert_eq!(list[1].id, "item_2");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_writers_no_data_loss() {
|
||||
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for w in 0..10 {
|
||||
let s = Arc::clone(&store);
|
||||
handles.push(tokio::spawn(async move {
|
||||
for i in 0..10 {
|
||||
let id = format!("w{w}_i{i}");
|
||||
s.save(make_item(&id)).await.unwrap();
|
||||
}
|
||||
}));
|
||||
}
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 100);
|
||||
// 验证所有 id 唯一
|
||||
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
|
||||
ids.sort();
|
||||
ids.dedup();
|
||||
assert_eq!(ids.len(), 100);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn persistence_round_trip() {
|
||||
let dir = TempDir::new().unwrap();
|
||||
let path = dir.path().join("memory.db");
|
||||
|
||||
// 阶段 1:写入 3 条
|
||||
{
|
||||
let store = SqliteStore::open(&path).unwrap();
|
||||
store.save(make_item("alpha")).await.unwrap();
|
||||
store.save(make_item("beta")).await.unwrap();
|
||||
store.save(make_item("gamma")).await.unwrap();
|
||||
}
|
||||
|
||||
// 阶段 2:重新打开,验证数据完整
|
||||
{
|
||||
let store = SqliteStore::open(&path).unwrap();
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 3);
|
||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||
assert!(ids.contains(&"alpha"));
|
||||
assert!(ids.contains(&"beta"));
|
||||
assert!(ids.contains(&"gamma"));
|
||||
|
||||
// 单条读回
|
||||
let got = store.get("beta").await.unwrap().unwrap();
|
||||
assert_eq!(got.content, "content-beta");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn open_invalid_path_returns_error() {
|
||||
// 路径指向已存在的目录而非文件,open 应失败
|
||||
let dir = TempDir::new().unwrap();
|
||||
match SqliteStore::open(dir.path()) {
|
||||
Err(MemoryError::Storage(_)) => {}
|
||||
Err(other) => panic!("expected Storage error, got {other:?}"),
|
||||
Ok(_) => panic!("expected error when opening a directory as database"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn trait_object_compatibility() {
|
||||
// ponytail: 回归验证 SqliteStore 可作为 Arc<dyn MemoryStore> 与 InMemoryStore 互换
|
||||
// 所有现有消费者(Conversation / Knowledge / Retriever / SessionMemory)均通过 trait object 引用,
|
||||
// 此测试确保 trait 接口契约在 SqliteStore 上同样成立。
|
||||
let sqlite: Arc<dyn MemoryStore> =
|
||||
Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||
let in_mem: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
|
||||
let stores: Vec<Arc<dyn MemoryStore>> = vec![Arc::clone(&sqlite), Arc::clone(&in_mem)];
|
||||
for store in &stores {
|
||||
store.save(make_item("x")).await.unwrap();
|
||||
let got = store.get("x").await.unwrap();
|
||||
assert_eq!(got.unwrap().id, "x");
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
store.delete("x").await.unwrap();
|
||||
assert!(store.get("x").await.unwrap().is_none());
|
||||
}
|
||||
}
|
||||
}
|
||||
+2
-2
@@ -1,7 +1,7 @@
|
||||
pub mod composer;
|
||||
pub mod error;
|
||||
pub mod template;
|
||||
pub mod composer;
|
||||
|
||||
pub use composer::{PromptComposer, validate_messages};
|
||||
pub use error::PromptError;
|
||||
pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue};
|
||||
pub use composer::{validate_messages, PromptComposer};
|
||||
|
||||
+17
-15
@@ -48,7 +48,11 @@ impl PromptComposer {
|
||||
|
||||
/// 添加一条 Tool 消息(工具执行结果回传)。
|
||||
pub fn tool(mut self, tool_call_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
self.push_message(Message::tool_result(tool_call_id.into(), content.into(), false));
|
||||
self.push_message(Message::tool_result(
|
||||
tool_call_id.into(),
|
||||
content.into(),
|
||||
false,
|
||||
));
|
||||
self
|
||||
}
|
||||
|
||||
@@ -133,11 +137,7 @@ impl PromptComposer {
|
||||
}
|
||||
|
||||
/// 添加一条含指定 ContentBlock 的 Tool 消息。
|
||||
pub fn tool_content(
|
||||
mut self,
|
||||
tool_call_id: impl Into<String>,
|
||||
block: ContentBlock,
|
||||
) -> Self {
|
||||
pub fn tool_content(mut self, tool_call_id: impl Into<String>, block: ContentBlock) -> Self {
|
||||
self.push_message(Message::ToolResult {
|
||||
tool_call_id: tool_call_id.into(),
|
||||
content: vec![block],
|
||||
@@ -187,9 +187,7 @@ impl PromptComposer {
|
||||
/// 验证消息序列是否符合 LLM API 要求(Tool 消息必须紧跟含 tool_calls 的 Assistant)。
|
||||
pub fn validate_messages(messages: &[Message]) -> Result<(), PromptError> {
|
||||
if messages.is_empty() {
|
||||
return Err(PromptError::InvalidSequence(
|
||||
"消息列表不能为空".to_string(),
|
||||
));
|
||||
return Err(PromptError::InvalidSequence("消息列表不能为空".to_string()));
|
||||
}
|
||||
|
||||
let mut last_tool_call_ids: Vec<String> = Vec::new();
|
||||
@@ -297,7 +295,8 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn test_template_if() {
|
||||
let tpl = PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
|
||||
let tpl =
|
||||
PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
|
||||
let mut ctx = TemplateContext::new();
|
||||
ctx.insert("name", "Bob");
|
||||
|
||||
@@ -312,11 +311,14 @@ mod tests {
|
||||
fn test_template_each() {
|
||||
let tpl = PromptTemplate::compile("Items: {{#each items}}{{item}}, {{/each}}").unwrap();
|
||||
let mut ctx = TemplateContext::new();
|
||||
ctx.insert("items", TemplateValue::Array(vec![
|
||||
TemplateValue::String("a".to_string()),
|
||||
TemplateValue::String("b".to_string()),
|
||||
TemplateValue::String("c".to_string()),
|
||||
]));
|
||||
ctx.insert(
|
||||
"items",
|
||||
TemplateValue::Array(vec![
|
||||
TemplateValue::String("a".to_string()),
|
||||
TemplateValue::String("b".to_string()),
|
||||
TemplateValue::String("c".to_string()),
|
||||
]),
|
||||
);
|
||||
|
||||
let result = tpl.render(&ctx).unwrap();
|
||||
assert_eq!(result, "Items: a, b, c, ");
|
||||
|
||||
+7
-2
@@ -1,6 +1,7 @@
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Error, Debug)]
|
||||
#[non_exhaustive]
|
||||
pub enum PromptError {
|
||||
#[error("模板解析错误: {0}。请检查模板语法({{var}} / {{#if}} / {{#each}})")]
|
||||
Parse(String),
|
||||
@@ -8,7 +9,9 @@ pub enum PromptError {
|
||||
#[error("渲染错误: 变量 '{0}' 未找到。请在 TemplateContext 中插入该变量")]
|
||||
VariableNotFound(String),
|
||||
|
||||
#[error("渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板")]
|
||||
#[error(
|
||||
"渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板"
|
||||
)]
|
||||
PartialNotFound(String),
|
||||
|
||||
#[error("渲染错误: '{0}' 不是数组,无法遍历。请确认传入的是数组或先判空")]
|
||||
@@ -20,7 +23,9 @@ pub enum PromptError {
|
||||
#[error("渲染错误: {0}")]
|
||||
Render(String),
|
||||
|
||||
#[error("消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)")]
|
||||
#[error(
|
||||
"消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)"
|
||||
)]
|
||||
InvalidSequence(String),
|
||||
|
||||
#[error("文件读取错误: {0}。请检查模板文件路径与权限")]
|
||||
|
||||
+21
-34
@@ -1,6 +1,6 @@
|
||||
use serde_json::Value;
|
||||
use std::collections::HashMap;
|
||||
use std::fmt;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::prompt::error::PromptError;
|
||||
|
||||
@@ -140,7 +140,9 @@ fn json_to_template_value(v: &Value) -> Result<TemplateValue, PromptError> {
|
||||
#[derive(Debug, Clone)]
|
||||
enum Fragment {
|
||||
Literal(String),
|
||||
Variable { name: String },
|
||||
Variable {
|
||||
name: String,
|
||||
},
|
||||
If {
|
||||
condition: String,
|
||||
body: Vec<Fragment>,
|
||||
@@ -223,8 +225,7 @@ fn compile_fragments(template: &str) -> Result<Vec<Fragment>, PromptError> {
|
||||
|
||||
let tag = tag_content.trim();
|
||||
if let Some(rest) = tag.strip_prefix("#if ") {
|
||||
let (body, else_body, new_i) =
|
||||
parse_block(template, i, "if")?;
|
||||
let (body, else_body, new_i) = parse_block(template, i, "if")?;
|
||||
let condition = rest.trim().to_string();
|
||||
fragments.push(Fragment::If {
|
||||
condition,
|
||||
@@ -331,10 +332,7 @@ fn parse_block(
|
||||
Err(PromptError::Parse(format!("未闭合的 {{#{}}} 块", kind)))
|
||||
}
|
||||
|
||||
fn parse_each_block(
|
||||
template: &str,
|
||||
start: usize,
|
||||
) -> Result<(Vec<Fragment>, usize), PromptError> {
|
||||
fn parse_each_block(template: &str, start: usize) -> Result<(Vec<Fragment>, usize), PromptError> {
|
||||
let bytes = template.as_bytes();
|
||||
let len = bytes.len();
|
||||
let mut depth = 1u32;
|
||||
@@ -368,9 +366,7 @@ fn parse_each_block(
|
||||
}
|
||||
}
|
||||
|
||||
Err(PromptError::Parse(
|
||||
"未闭合的 {{#each}} 块".to_string(),
|
||||
))
|
||||
Err(PromptError::Parse("未闭合的 {{#each}} 块".to_string()))
|
||||
}
|
||||
|
||||
fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), PromptError> {
|
||||
@@ -395,9 +391,7 @@ fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), Prom
|
||||
}
|
||||
}
|
||||
|
||||
Err(PromptError::Parse(
|
||||
"未闭合的 {{#raw}} 块".to_string(),
|
||||
))
|
||||
Err(PromptError::Parse("未闭合的 {{#raw}} 块".to_string()))
|
||||
}
|
||||
|
||||
// ===== Renderer =====
|
||||
@@ -418,33 +412,28 @@ fn render_fragments(
|
||||
Fragment::Literal(text) => {
|
||||
output.push_str(text);
|
||||
}
|
||||
Fragment::Variable { name } => {
|
||||
match ctx.get(name) {
|
||||
Some(val) => {
|
||||
output.push_str(&format!("{}", val));
|
||||
}
|
||||
None => {
|
||||
return Err(PromptError::VariableNotFound(name.clone()));
|
||||
}
|
||||
Fragment::Variable { name } => match ctx.get(name) {
|
||||
Some(val) => {
|
||||
output.push_str(&format!("{}", val));
|
||||
}
|
||||
}
|
||||
None => {
|
||||
return Err(PromptError::VariableNotFound(name.clone()));
|
||||
}
|
||||
},
|
||||
Fragment::If {
|
||||
condition,
|
||||
body,
|
||||
else_body,
|
||||
} => {
|
||||
let truthy = ctx
|
||||
.get(condition)
|
||||
.map(|v| v.is_truthy())
|
||||
.unwrap_or(false);
|
||||
let truthy = ctx.get(condition).map(|v| v.is_truthy()).unwrap_or(false);
|
||||
let target = if truthy { body } else { else_body };
|
||||
render_fragments(target, ctx, partials, output, depth + 1)?;
|
||||
}
|
||||
Fragment::Each { variable, body } => {
|
||||
let arr = match ctx.get(variable) {
|
||||
Some(val) => val.as_array().ok_or_else(|| {
|
||||
PromptError::NotAnArray(variable.clone())
|
||||
})?,
|
||||
Some(val) => val
|
||||
.as_array()
|
||||
.ok_or_else(|| PromptError::NotAnArray(variable.clone()))?,
|
||||
None => {
|
||||
return Err(PromptError::VariableNotFound(variable.clone()));
|
||||
}
|
||||
@@ -504,10 +493,8 @@ impl PromptTemplateRegistry {
|
||||
|
||||
/// 延迟编译注册:只存储原始字符串,首次渲染时编译。
|
||||
pub fn register_lazy(&mut self, name: &str, template: &str) {
|
||||
self.templates.insert(
|
||||
name.to_string(),
|
||||
StoredTemplate::Raw(template.to_string()),
|
||||
);
|
||||
self.templates
|
||||
.insert(name.to_string(), StoredTemplate::Raw(template.to_string()));
|
||||
}
|
||||
|
||||
/// 从文件读取并编译注册。
|
||||
|
||||
+10
-3
@@ -4,9 +4,12 @@ use std::sync::Arc;
|
||||
|
||||
/// 工具调用过程中可能发生的所有错误。
|
||||
#[derive(thiserror::Error, Debug, Clone)]
|
||||
#[non_exhaustive]
|
||||
pub enum ToolError {
|
||||
/// 工具未注册。不可恢复——需调用方先 `registry.register(...)`。
|
||||
#[error("工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写")]
|
||||
#[error(
|
||||
"工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写"
|
||||
)]
|
||||
NotFound(String),
|
||||
|
||||
/// 工具执行失败(可恢复——文本回传 LLM 由其决定重试或放弃)。
|
||||
@@ -14,11 +17,15 @@ pub enum ToolError {
|
||||
ExecutionFailed(String, String),
|
||||
|
||||
/// 工具参数无效(可恢复——文本回传 LLM)。
|
||||
#[error("工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema")]
|
||||
#[error(
|
||||
"工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema"
|
||||
)]
|
||||
InvalidArguments(String, String),
|
||||
|
||||
/// 权限被拒绝(不可恢复——终止循环)。
|
||||
#[error("权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过")]
|
||||
#[error(
|
||||
"权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过"
|
||||
)]
|
||||
PermissionDenied(String, String),
|
||||
|
||||
/// MCP 协议错误(不可恢复)。
|
||||
|
||||
+17
-37
@@ -9,19 +9,18 @@
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::process::Stdio;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::time::Duration;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{json, Value};
|
||||
use serde_json::{Value, json};
|
||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
||||
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
|
||||
use tokio::sync::{oneshot, Mutex};
|
||||
use tokio::sync::{Mutex, oneshot};
|
||||
|
||||
#[allow(deprecated)]
|
||||
use crate::llm::types::ToolDefinition;
|
||||
use crate::llm::types::tool::ToolDef;
|
||||
use crate::tools::base::{BaseTool, ToolContext, ToolRef};
|
||||
use crate::tools::error::ToolError;
|
||||
|
||||
@@ -136,7 +135,6 @@ impl std::fmt::Debug for McpClient {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(deprecated)]
|
||||
impl McpClient {
|
||||
/// 创建一个 MCP 客户端。
|
||||
pub fn new(server_name: impl Into<String>, transport: McpTransport) -> Self {
|
||||
@@ -226,9 +224,7 @@ impl McpClient {
|
||||
"version": env!("CARGO_PKG_VERSION")
|
||||
}
|
||||
});
|
||||
let _response = self
|
||||
.send_request("initialize", Some(init_params))
|
||||
.await?;
|
||||
let _response = self.send_request("initialize", Some(init_params)).await?;
|
||||
|
||||
// 发送 initialized 通知(无 id)
|
||||
self.send_notification("notifications/initialized", Some(json!({})))
|
||||
@@ -239,7 +235,7 @@ impl McpClient {
|
||||
}
|
||||
|
||||
/// 列出服务器支持的工具(调用 `tools/list`)。
|
||||
pub async fn list_tools(&mut self) -> Result<Vec<ToolDefinition>, ToolError> {
|
||||
pub async fn list_tools(&mut self) -> Result<Vec<ToolDef>, ToolError> {
|
||||
if !self.is_initialized() {
|
||||
return Err(ToolError::McpNotInitialized(self.server_name.clone()));
|
||||
}
|
||||
@@ -274,11 +270,10 @@ impl McpClient {
|
||||
description: description.clone(),
|
||||
input_schema: input_schema.clone(),
|
||||
});
|
||||
defs.push(ToolDefinition {
|
||||
defs.push(ToolDef {
|
||||
name,
|
||||
description,
|
||||
parameters: input_schema,
|
||||
strict: None,
|
||||
});
|
||||
}
|
||||
Ok(defs)
|
||||
@@ -337,11 +332,7 @@ impl McpClient {
|
||||
if let Some(state) = self.process.take() {
|
||||
let mut state = state.lock().await;
|
||||
// 优雅等待 5 秒
|
||||
let graceful = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
state.child.wait(),
|
||||
)
|
||||
.await;
|
||||
let graceful = tokio::time::timeout(Duration::from_secs(5), state.child.wait()).await;
|
||||
if graceful.is_err() {
|
||||
// 超时则强杀
|
||||
let _ = state.child.kill().await;
|
||||
@@ -372,11 +363,7 @@ impl McpClient {
|
||||
tools
|
||||
}
|
||||
|
||||
async fn send_request(
|
||||
&self,
|
||||
method: &str,
|
||||
params: Option<Value>,
|
||||
) -> Result<Value, ToolError> {
|
||||
async fn send_request(&self, method: &str, params: Option<Value>) -> Result<Value, ToolError> {
|
||||
let state_arc = self
|
||||
.process
|
||||
.as_ref()
|
||||
@@ -412,9 +399,11 @@ impl McpClient {
|
||||
.write_all(b"\n")
|
||||
.await
|
||||
.map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?;
|
||||
state.stdin.flush().await.map_err(|e| {
|
||||
ToolError::McpError(format!("flush stdin 失败: {e}"))
|
||||
})?;
|
||||
state
|
||||
.stdin
|
||||
.flush()
|
||||
.await
|
||||
.map_err(|e| ToolError::McpError(format!("flush stdin 失败: {e}")))?;
|
||||
}
|
||||
|
||||
// 等待响应(带超时)
|
||||
@@ -471,10 +460,7 @@ impl McpClient {
|
||||
}
|
||||
|
||||
/// 持续读取 stdout,将响应分发到对应的 oneshot sender。
|
||||
async fn read_loop(
|
||||
mut reader: BufReader<ChildStdout>,
|
||||
state: Arc<Mutex<ChildProcessState>>,
|
||||
) {
|
||||
async fn read_loop(mut reader: BufReader<ChildStdout>, state: Arc<Mutex<ChildProcessState>>) {
|
||||
let mut line = String::new();
|
||||
loop {
|
||||
line.clear();
|
||||
@@ -542,7 +528,6 @@ enum McpClientHandle {
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
#[allow(deprecated)]
|
||||
impl BaseTool for McpToolAdapter {
|
||||
fn name(&self) -> &str {
|
||||
&self.name
|
||||
@@ -556,11 +541,7 @@ impl BaseTool for McpToolAdapter {
|
||||
self.parameters.clone()
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
&self,
|
||||
_args: Value,
|
||||
_ctx: &ToolContext<'_>,
|
||||
) -> Result<Value, ToolError> {
|
||||
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||
// 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。
|
||||
// 实际生产中应持有 Arc<McpClient> 并通过 mcp.call_tool() 执行。
|
||||
// 这里返回错误,提示需要通过其他方式调用 MCP 工具。
|
||||
@@ -617,8 +598,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn test_jsonrpc_response_parse_error() {
|
||||
let s =
|
||||
r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
|
||||
let s = r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
|
||||
let resp: JsonRpcResponse = serde_json::from_str(s).unwrap();
|
||||
assert_eq!(resp.id, 1);
|
||||
assert!(resp.result.is_none());
|
||||
|
||||
+21
-18
@@ -148,9 +148,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_default_config_denies_delete() {
|
||||
let checker = PermissionChecker::new(PermissionConfig::default());
|
||||
assert!(checker
|
||||
.check("rm_file", &p(Permission::Delete))
|
||||
.is_err());
|
||||
assert!(checker.check("rm_file", &p(Permission::Delete)).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -246,12 +244,16 @@ mod tests {
|
||||
allow_unspecified: false,
|
||||
};
|
||||
let checker = PermissionChecker::new(cfg);
|
||||
assert!(checker
|
||||
.check("t", &[Permission::Custom("db:read".into())])
|
||||
.is_ok());
|
||||
assert!(checker
|
||||
.check("t", &[Permission::Custom("db:write".into())])
|
||||
.is_err());
|
||||
assert!(
|
||||
checker
|
||||
.check("t", &[Permission::Custom("db:read".into())])
|
||||
.is_ok()
|
||||
);
|
||||
assert!(
|
||||
checker
|
||||
.check("t", &[Permission::Custom("db:write".into())])
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -262,12 +264,11 @@ mod tests {
|
||||
allow_unspecified: false,
|
||||
};
|
||||
let checker = PermissionChecker::new(cfg);
|
||||
assert!(checker
|
||||
.check(
|
||||
"t",
|
||||
&[Permission::Read, Permission::Network]
|
||||
)
|
||||
.is_ok());
|
||||
assert!(
|
||||
checker
|
||||
.check("t", &[Permission::Read, Permission::Network])
|
||||
.is_ok()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -279,8 +280,10 @@ mod tests {
|
||||
};
|
||||
let checker = PermissionChecker::new(cfg);
|
||||
// 任一权限不在白名单则拒绝
|
||||
assert!(checker
|
||||
.check("t", &[Permission::Read, Permission::Write])
|
||||
.is_err());
|
||||
assert!(
|
||||
checker
|
||||
.check("t", &[Permission::Read, Permission::Write])
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+10
-10
@@ -7,8 +7,7 @@ use std::time::Duration;
|
||||
use futures::future::join_all;
|
||||
use serde_json::Value;
|
||||
|
||||
#[allow(deprecated)]
|
||||
use crate::llm::types::ToolDefinition;
|
||||
use crate::llm::types::tool::ToolDef;
|
||||
use crate::tools::base::{ToolContext, ToolRef};
|
||||
use crate::tools::error::ToolError;
|
||||
use crate::tools::permission::PermissionChecker;
|
||||
@@ -71,7 +70,6 @@ impl std::fmt::Debug for ToolRegistry {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(deprecated)]
|
||||
impl ToolRegistry {
|
||||
/// 创建一个新的工具注册表。
|
||||
pub fn new() -> Self {
|
||||
@@ -127,16 +125,15 @@ impl ToolRegistry {
|
||||
self.inner.tools.keys().cloned().collect()
|
||||
}
|
||||
|
||||
/// 获取所有工具的 `ToolDefinition` 列表(用于传递给 LLM)。
|
||||
pub fn definitions(&self) -> Vec<ToolDefinition> {
|
||||
/// 获取所有工具的 `ToolDef` 列表(用于传递给 LLM)。
|
||||
pub fn definitions(&self) -> Vec<ToolDef> {
|
||||
self.inner
|
||||
.tools
|
||||
.values()
|
||||
.map(|tool| ToolDefinition {
|
||||
.map(|tool| ToolDef {
|
||||
name: tool.name().to_string(),
|
||||
description: Some(tool.description().to_string()),
|
||||
parameters: tool.parameters(),
|
||||
strict: None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -348,7 +345,10 @@ mod tests {
|
||||
async fn test_invoke_success() {
|
||||
let mut reg = ToolRegistry::new();
|
||||
reg.register(Arc::new(AddTool { base: 100 })).unwrap();
|
||||
let result = reg.invoke("call_1", "add", json!({ "n": 5 })).await.unwrap();
|
||||
let result = reg
|
||||
.invoke("call_1", "add", json!({ "n": 5 }))
|
||||
.await
|
||||
.unwrap();
|
||||
let value = result.output.unwrap();
|
||||
assert_eq!(value["result"], 105);
|
||||
assert_eq!(result.tool_call_id, "call_1");
|
||||
@@ -372,8 +372,8 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_invoke_with_permission_denied() {
|
||||
let mut reg = ToolRegistry::new()
|
||||
.with_permission_checker(PermissionChecker::new(Default::default()));
|
||||
let mut reg =
|
||||
ToolRegistry::new().with_permission_checker(PermissionChecker::new(Default::default()));
|
||||
reg.register(Arc::new(ShellTool)).unwrap();
|
||||
let result = reg.invoke("call_z", "shell", json!({})).await;
|
||||
assert!(matches!(result, Err(ToolError::PermissionDenied(_, _))));
|
||||
|
||||
Reference in New Issue
Block a user