refactor: 拆分 crate 架构 - models/providers/config/middleware/plugin 迁移到子 crate

- proxycast-core: models, data, logger, errors, backends, connect, session_files,
  orchestrator, session(rate_limit/sticky), config, middleware, plugin
- proxycast-providers: providers, converter, streaming, translator, stream,
  session(SessionManager/signature_store)
- 主 crate 保留 Tauri 业务逻辑,通过 re-export 保持 API 兼容
- 版本号升级到 0.60.0
- 移除已删除的 useFlowEvents hook 导出
- 移除已废弃的 flow_monitor 端到端测试
This commit is contained in:
coso
2026-02-08 03:19:44 +08:00
parent 37054b49f6
commit 19e14c0f4c
246 changed files with 3709 additions and 58200 deletions
-105
View File
@@ -1,105 +0,0 @@
# 图文海报功能实现计划
## 当前进度
### Phase 1: 品牌人设系统扩展 ✅ 完成
**Goal**: 扩展现有人设系统,支持海报设计专用字段(配色、字体、品牌调性)
**Status**: Complete
#### 已完成任务
- [x] 1.1 扩展数据模型 (`project_model.rs`)
- 新增 BrandPersonality, DesignStyle 枚举
- 新增 ColorScheme, Typography, LogoPlacement, ImageStyle, IconStyle 结构体
- 新增 BrandTone, DesignConfig, VisualConfig 结构体
- 新增 BrandPersonaExtension, BrandPersona, BrandPersonaTemplate 结构体
- 新增相关请求类型
- [x] 1.2 扩展数据库 Schema (`schema.rs`)
- 新增 `brand_persona_extensions` 表
- 包含 persona_id, brand_tone_json, design_json, visual_json 字段
- [x] 1.3 创建 BrandPersona DAO (`brand_persona_dao.rs`)
- 实现 create, get, update, delete 方法
- 实现 get_brand_persona 获取完整品牌人设
- 实现 list_templates 获取预设模板
- [x] 1.4 扩展 PersonaService (`persona_service.rs`)
- 新增 get_brand_persona, get_brand_extension 方法
- 新增 save_brand_extension, update_brand_extension 方法
- 新增 delete_brand_extension, list_brand_persona_templates 方法
- [x] 1.5 扩展 Tauri 命令 (`persona_cmd.rs`)
- 新增 get_brand_persona, get_brand_extension 命令
- 新增 save_brand_extension, update_brand_extension 命令
- 新增 delete_brand_extension, list_brand_persona_templates 命令
- 在 runner.rs 中注册新命令
- [x] 1.6 新增前端类型 (`brand-persona.ts`)
- 定义所有品牌人设相关的 TypeScript 类型
- 包含预设配色方案、字体列表、默认值等常量
- [x] 1.7 新增 useBrandPersona Hook (`useBrandPersona.ts`)
- 实现品牌人设的 CRUD 操作
- 支持模板应用功能
- [x] 1.8 新增 BrandPersonaDialog 组件 (`BrandPersonaDialog.tsx`)
- 分步骤创建品牌人设(品牌调性 → 配色方案 → 字体设置 → 预览确认)
- 支持模板快速应用
- 支持预设配色方案选择
- 实时预览效果
#### 验证标准
- [x] 能够创建包含配色方案的品牌人设
- [x] 品牌人设能够正确保存和加载
- [x] 在项目详情页能够管理品牌人设
---
### Phase 2: 素材库扩展
**Goal**: 扩展素材库支持 icon, color, layout 类型
**Status**: Not Started
---
### Phase 3: 海报 Agent 系统
**Goal**: 实现 6 个专用 Agent,支持对话式海报设计
**Status**: Not Started
---
### Phase 4: 工作流系统
**Goal**: 实现 6 步引导工作流
**Status**: Not Started
---
### Phase 5: 多平台导出
**Goal**: 实现多平台尺寸适配和导出
**Status**: Not Started
---
## 新增文件清单
### 后端 (Rust)
- `src-tauri/src/database/dao/brand_persona_dao.rs` - 品牌人设 DAO
### 前端 (TypeScript/React)
- `src/types/brand-persona.ts` - 品牌人设类型定义
- `src/hooks/useBrandPersona.ts` - 品牌人设 Hook
- `src/components/projects/dialogs/BrandPersonaDialog.tsx` - 品牌人设对话框
## 修改文件清单
### 后端 (Rust)
- `src-tauri/src/models/project_model.rs` - 新增品牌人设数据模型
- `src-tauri/src/database/schema.rs` - 新增品牌人设扩展表
- `src-tauri/src/database/dao/mod.rs` - 导出新 DAO
- `src-tauri/src/services/persona_service.rs` - 新增品牌人设服务方法
- `src-tauri/src/commands/persona_cmd.rs` - 新增品牌人设命令
- `src-tauri/src/app/runner.rs` - 注册新命令
### 前端 (TypeScript/React)
- `src/types/index.ts` - 导出新类型
- `src/hooks/index.ts` - 导出新 Hook
- `src/components/projects/dialogs/index.ts` - 导出新组件
-261
View File
@@ -1,261 +0,0 @@
# 流量监控
## 概述
流量监控模块拦截和记录所有 LLM API 请求,提供 Token 统计、历史查询和分析功能。
## 目录结构
```
src-tauri/src/flow_monitor/
├── mod.rs # 模块入口
├── interceptor.rs # 请求拦截器
├── storage.rs # 存储层
├── query.rs # 查询接口
└── stats.rs # 统计分析
```
## 架构
```
┌─────────────────────────────────────────────────────────────────┐
│ HTTP 请求 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Flow Interceptor │
│ ┌─────────────────────────────────────────────────────────────┐│
│ │ 请求拦截 ││
│ │ - 请求 ID 生成 ││
│ │ - 请求体捕获 ││
│ │ - 时间戳记录 ││
│ └─────────────────────────────────────────────────────────────┘│
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Provider 处理 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Flow Interceptor │
│ ┌─────────────────────────────────────────────────────────────┐│
│ │ 响应拦截 ││
│ │ - 响应体捕获 ││
│ │ - Token 计数 ││
│ │ - 延迟计算 ││
│ └─────────────────────────────────────────────────────────────┘│
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Flow Storage │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ SQLite │ │ 内存缓存 │ │ 事件发送 │ │
│ │ 持久化 │ │ (最近 N 条) │ │ (前端通知) │ │
│ └─────────────┘ └─────────────┘ └─────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
```
## 数据模型
### FlowRecord
```rust
pub struct FlowRecord {
pub id: String, // 请求 ID
pub timestamp: i64, // 时间戳
pub provider: String, // Provider 类型
pub model: String, // 模型名称
pub request: FlowRequest, // 请求数据
pub response: Option<FlowResponse>, // 响应数据
pub status: FlowStatus, // 状态
pub latency_ms: Option<u64>, // 延迟 (毫秒)
}
pub struct FlowRequest {
pub messages: Vec<Message>, // 消息列表
pub tools: Option<Vec<Tool>>, // 工具定义
pub stream: bool, // 是否流式
}
pub struct FlowResponse {
pub content: String, // 响应内容
pub tool_calls: Option<Vec<ToolCall>>, // 工具调用
pub usage: TokenUsage, // Token 使用
}
pub struct TokenUsage {
pub prompt_tokens: u32, // 输入 Token
pub completion_tokens: u32, // 输出 Token
pub total_tokens: u32, // 总 Token
}
```
### 数据库表
```sql
CREATE TABLE flow_records (
id TEXT PRIMARY KEY,
timestamp INTEGER NOT NULL,
provider TEXT NOT NULL,
model TEXT NOT NULL,
request_json TEXT NOT NULL,
response_json TEXT,
status TEXT NOT NULL,
latency_ms INTEGER,
prompt_tokens INTEGER,
completion_tokens INTEGER,
total_tokens INTEGER,
created_at INTEGER NOT NULL
);
CREATE INDEX idx_flow_timestamp ON flow_records(timestamp);
CREATE INDEX idx_flow_provider ON flow_records(provider);
CREATE INDEX idx_flow_model ON flow_records(model);
```
## 拦截器实现
```rust
pub struct FlowInterceptor {
storage: Arc<FlowStorage>,
event_sender: mpsc::Sender<FlowEvent>,
}
impl FlowInterceptor {
pub async fn intercept_request(&self, req: &Request) -> String {
let request_id = generate_uuid();
let record = FlowRecord {
id: request_id.clone(),
timestamp: current_timestamp(),
provider: extract_provider(req),
model: extract_model(req),
request: parse_request(req),
response: None,
status: FlowStatus::Pending,
latency_ms: None,
};
self.storage.insert(&record).await;
self.event_sender.send(FlowEvent::RequestStarted(record)).await;
request_id
}
pub async fn intercept_response(
&self,
request_id: &str,
response: &Response,
latency: Duration,
) {
let flow_response = parse_response(response);
self.storage.update_response(
request_id,
&flow_response,
latency.as_millis() as u64,
).await;
self.event_sender.send(FlowEvent::ResponseReceived {
request_id: request_id.to_string(),
response: flow_response,
}).await;
}
}
```
## 查询接口
### 分页查询
```rust
pub struct FlowQuery {
pub provider: Option<String>,
pub model: Option<String>,
pub start_time: Option<i64>,
pub end_time: Option<i64>,
pub status: Option<FlowStatus>,
pub page: u32,
pub page_size: u32,
}
pub async fn query_flows(query: FlowQuery) -> Result<PagedResult<FlowRecord>> {
// 构建 SQL 查询
// 执行分页查询
// 返回结果
}
```
### 统计查询
```rust
pub struct FlowStats {
pub total_requests: u64,
pub total_tokens: u64,
pub avg_latency_ms: f64,
pub by_provider: HashMap<String, ProviderStats>,
pub by_model: HashMap<String, ModelStats>,
}
pub async fn get_stats(time_range: TimeRange) -> Result<FlowStats> {
// 聚合统计
}
```
## 前端事件
### 事件类型
```typescript
interface FlowEvent {
type: 'request_started' | 'response_received' | 'error';
data: FlowRecord;
}
```
### 事件监听
```typescript
// 前端监听
import { listen } from '@tauri-apps/api/event';
listen<FlowEvent>('flow-event', (event) => {
switch (event.payload.type) {
case 'request_started':
addPendingRequest(event.payload.data);
break;
case 'response_received':
updateRequest(event.payload.data);
break;
}
});
```
## Tauri Commands
```rust
#[tauri::command]
async fn get_flow_records(query: FlowQuery) -> Result<PagedResult<FlowRecord>>;
#[tauri::command]
async fn get_flow_stats(time_range: TimeRange) -> Result<FlowStats>;
#[tauri::command]
async fn get_flow_detail(id: String) -> Result<FlowRecord>;
#[tauri::command]
async fn clear_flow_records(before: Option<i64>) -> Result<u64>;
#[tauri::command]
async fn export_flow_records(format: ExportFormat) -> Result<String>;
```
## 相关文档
- [server.md](server.md) - HTTP 服务器
- [database.md](database.md) - 数据库层
- [components.md](components.md) - 前端组件
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.59.0",
"version": "0.60.0",
"type": "module",
"repository": {
"type": "git",
+59 -4
View File
@@ -203,7 +203,6 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50"
[[package]]
name = "aster"
version = "0.11.0"
source = "git+https://github.com/astercloud/aster-rust?tag=v0.11.0#c4a9c40b48bcb0c77375b10ccb1313366fb4e7ba"
dependencies = [
"ahash",
"anyhow",
@@ -6636,7 +6635,7 @@ dependencies = [
[[package]]
name = "proxycast"
version = "0.59.0"
version = "0.60.0"
dependencies = [
"anyhow",
"arboard",
@@ -6672,6 +6671,7 @@ dependencies = [
"proptest",
"proxycast-core",
"proxycast-infra",
"proxycast-providers",
"rand 0.8.5",
"regex",
"reqwest 0.12.28",
@@ -6719,23 +6719,44 @@ dependencies = [
[[package]]
name = "proxycast-core"
version = "0.59.0"
version = "0.60.0"
dependencies = [
"async-trait",
"axum 0.7.9",
"bytes",
"chrono",
"dashmap 5.5.3",
"dirs 5.0.1",
"flate2",
"futures",
"indexmap 2.13.0",
"notify 6.1.1",
"parking_lot",
"proptest",
"rand 0.8.5",
"reqwest 0.12.28",
"rusqlite",
"serde",
"serde_json",
"serde_urlencoded",
"serde_yaml",
"sha2",
"subtle",
"tar",
"tempfile",
"thiserror 1.0.69",
"tokio",
"tower 0.5.3",
"tracing",
"url",
"urlencoding",
"uuid",
"zip",
]
[[package]]
name = "proxycast-infra"
version = "0.59.0"
version = "0.60.0"
dependencies = [
"chrono",
"dashmap 5.5.3",
@@ -6753,6 +6774,40 @@ dependencies = [
"uuid",
]
[[package]]
name = "proxycast-providers"
version = "0.60.0"
dependencies = [
"anyhow",
"async-stream",
"async-trait",
"axum 0.7.9",
"base64 0.22.1",
"bytes",
"chrono",
"dirs 5.0.1",
"flate2",
"futures",
"once_cell",
"open",
"proptest",
"proxycast-core",
"rand 0.8.5",
"regex",
"reqwest 0.12.28",
"serde",
"serde_json",
"serde_urlencoded",
"sha2",
"tempfile",
"thiserror 1.0.69",
"tokio",
"tracing",
"url",
"urlencoding",
"uuid",
]
[[package]]
name = "psl-types"
version = "2.0.11"
+4 -2
View File
@@ -3,7 +3,7 @@ members = ["crates/*"]
resolver = "2"
[workspace.package]
version = "0.59.0"
version = "0.60.0"
edition = "2021"
authors = ["you"]
repository = "https://github.com/aiclientproxy/proxycast"
@@ -13,6 +13,7 @@ homepage = "https://github.com/aiclientproxy/proxycast"
# 项目内 crate 依赖
proxycast-core = { path = "crates/core" }
proxycast-infra = { path = "crates/infra" }
proxycast-providers = { path = "crates/providers" }
voice-core = { path = "crates/voice-core" }
# 序列化
@@ -167,7 +168,7 @@ version = "2.4"
[package]
name = "proxycast"
version = "0.59.0"
version = "0.60.0"
description = "AI API Proxy Desktop App"
authors = ["you"]
edition = "2021"
@@ -185,6 +186,7 @@ tauri-build.workspace = true
# 项目内 crate
proxycast-core.workspace = true
proxycast-infra.workspace = true
proxycast-providers.workspace = true
voice-core.workspace = true
# Tauri
-57
View File
@@ -1,57 +0,0 @@
{
"$schema": "https://schemas.tauri.app/config/2/capability",
"identifier": "default",
"description": "Default capabilities for ProxyCast",
"windows": ["main", "smart-input", "update-notification"],
"permissions": [
"core:default",
"core:webview:default",
"core:webview:allow-webview-close",
"core:webview:allow-webview-position",
"core:webview:allow-webview-size",
"core:webview:allow-set-webview-position",
"core:webview:allow-set-webview-size",
"core:window:default",
"core:window:allow-close",
"core:window:allow-show",
"core:window:allow-hide",
"core:window:allow-set-focus",
"core:window:allow-center",
"core:window:allow-start-dragging",
"shell:allow-open",
"shell:allow-spawn",
"shell:allow-execute",
"shell:allow-kill",
"shell:allow-stdin-write",
"dialog:default",
"global-shortcut:default",
"global-shortcut:allow-is-registered",
"global-shortcut:allow-register",
"global-shortcut:allow-unregister",
{
"identifier": "shell:allow-execute",
"allow": [
{
"name": "binaries/aster-server",
"sidecar": true,
"args": true
},
{
"name": "open",
"cmd": "open",
"args": true
}
]
},
{
"identifier": "shell:allow-spawn",
"allow": [
{
"name": "binaries/aster-server",
"sidecar": true,
"args": true
}
]
}
]
}
+36 -1
View File
@@ -9,10 +9,21 @@ repository.workspace = true
# 序列化
serde.workspace = true
serde_json.workspace = true
serde_urlencoded.workspace = true
# 异步运行时
tokio.workspace = true
async-trait.workspace = true
# 错误处理
thiserror.workspace = true
# 日志
tracing.workspace = true
# HTTP 客户端
reqwest.workspace = true
# 时间和 UUID
chrono.workspace = true
uuid.workspace = true
@@ -22,6 +33,30 @@ indexmap.workspace = true
parking_lot.workspace = true
dirs.workspace = true
sha2.workspace = true
url.workspace = true
urlencoding.workspace = true
bytes.workspace = true
futures.workspace = true
dashmap.workspace = true
notify.workspace = true
rand.workspace = true
# HTTP 服务器(middleware 模块需要)
axum.workspace = true
tower.workspace = true
subtle.workspace = true
# 压缩/归档(plugin installer 需要)
flate2.workspace = true
tar.workspace = true
zip.workspace = true
# YAML 配置
serde_yaml.workspace = true
# 数据库(errors 模块需要 rusqlite::Error)
rusqlite.workspace = true
[dev-dependencies]
proptest.workspace = true
proptest.workspace = true
tempfile.workspace = true
+34
View File
@@ -0,0 +1,34 @@
//! 配置管理模块
//!
//! 提供 YAML 配置文件支持、热重载和配置导入导出功能
//! 同时保持与旧版 JSON 配置的向后兼容性
#![allow(unused_imports)]
mod export;
mod hot_reload;
mod import;
mod path_utils;
mod types;
mod yaml;
pub use export::{ExportBundle, ExportOptions, ExportService, REDACTED_PLACEHOLDER};
pub use hot_reload::{
ConfigChangeEvent as FileChangeEvent, ConfigChangeKind, FileWatcher, HotReloadManager,
ReloadResult,
};
pub use import::{ImportOptions, ImportService, ValidationResult};
pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde};
pub use types::{
generate_secure_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, AsrCredentialEntry,
AsrProviderType, BaiduConfig, Config, ContentCreatorConfig, CredentialEntry,
CredentialPoolConfig, CustomProviderConfig, EndpointProvidersConfig, ExperimentalFeatures,
GeminiApiKeyEntry, InjectionRuleConfig, InjectionSettings, LoggingConfig, ModelInfo,
ModelsConfig, NativeAgentConfig, NavigationConfig, OpenAIAsrConfig, ProviderConfig,
ProviderModelsConfig, ProvidersConfig, QuotaExceededConfig, RemoteManagementConfig,
RetrySettings, RoutingConfig, ScreenshotChatConfig, ServerConfig, TlsConfig, UpdateCheckConfig,
VertexApiKeyEntry, VertexModelAlias, VoiceInputConfig, VoiceInstruction, VoiceOutputConfig,
VoiceOutputMode, VoiceProcessorConfig, WhisperLocalConfig, WhisperModelSize, XunfeiConfig,
DEFAULT_API_KEY,
};
pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService};
File diff suppressed because it is too large Load Diff
@@ -3,7 +3,7 @@
//! 定义 ProxyCast 的配置结构,支持 YAML 和 JSON 序列化/反序列化
//! 保持与旧版 JSON 配置的向后兼容性
use crate::injection::{InjectionMode, InjectionRule};
use crate::models::injection_types::{InjectionMode, InjectionRule};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
@@ -192,34 +192,8 @@ pub struct GeminiApiKeyEntry {
}
/// Vertex AI 模型别名映射
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct VertexModelAlias {
/// 上游模型名称
pub name: String,
/// 客户端可见的别名
pub alias: String,
}
/// Vertex AI 凭证条目
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct VertexApiKeyEntry {
/// 凭证 ID
pub id: String,
/// API Key
pub api_key: String,
/// Base URL
#[serde(default, skip_serializing_if = "Option::is_none")]
pub base_url: Option<String>,
/// 模型别名映射
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub models: Vec<VertexModelAlias>,
/// 单独的代理 URL
#[serde(default, skip_serializing_if = "Option::is_none")]
pub proxy_url: Option<String>,
/// 是否禁用
#[serde(default)]
pub disabled: bool,
}
// Vertex AI 类型从 core crate 重新导出
pub use crate::models::vertex_model::{VertexApiKeyEntry, VertexModelAlias};
#[allow(dead_code)]
fn default_auth_type() -> String {
@@ -871,7 +845,7 @@ impl NativeAgentConfig {
pub fn get_effective_system_prompt(&self) -> Option<String> {
// 优先从文件加载
if let Some(file_path) = &self.system_prompt_file {
let expanded_path = crate::config::expand_tilde(file_path);
let expanded_path = super::path_utils::expand_tilde(file_path);
if let Ok(content) = std::fs::read_to_string(&expanded_path) {
let trimmed = content.trim();
if !trimmed.is_empty() {
@@ -11,7 +11,7 @@
//! ## 使用示例
//!
//! ```rust
//! use proxycast_lib::connect::deep_link::{parse_deep_link, ConnectPayload, DeepLinkError};
//! use proxycast_core::connect::deep_link::{parse_deep_link, ConnectPayload, DeepLinkError};
//!
//! let url = "proxycast://connect?relay=example&key=sk-xxx&name=MyKey";
//! match parse_deep_link(url) {
@@ -80,7 +80,7 @@ impl std::error::Error for DeepLinkError {}
/// # 示例
///
/// ```rust
/// use proxycast_lib::connect::deep_link::parse_deep_link;
/// use proxycast_core::connect::deep_link::parse_deep_link;
///
/// // 完整 URL
/// let result = parse_deep_link("proxycast://connect?relay=example&key=sk-xxx&name=MyKey&ref=abc");
@@ -11,7 +11,7 @@
//! ## 使用示例
//!
//! ```rust,ignore
//! use proxycast_lib::connect::registry::{RelayRegistry, RelayInfo};
//! use proxycast_core::connect::registry::{RelayRegistry, RelayInfo};
//!
//! let registry = RelayRegistry::new(cache_path);
//! registry.load_from_remote().await?;
-3
View File
@@ -1,4 +1 @@
//! 静态数据模块
//!
//! 模型数据现在从 aiclientproxy/models 仓库获取
//! 本地硬编码数据已迁移到独立仓库: https://github.com/aiclientproxy/models
+26 -3
View File
@@ -1,13 +1,36 @@
//! 核心类型模块
//! ProxyCast Core Crate
//!
//! 包含纯数据类型(models)、静态数据(data)、日志配置(logger)
//! 包含纯数据类型、基础模块和无外部业务依赖的独立模块。
//!
//! 本 crate 不包含任何业务逻辑,只提供基础类型定义。
//! ## 模块结构
//! - `models`: 核心数据模型定义
//! - `data`: 静态数据
//! - `logger`: 日志配置
//! - `errors`: 错误类型定义
//! - `backends`: 后端调用层 Trait
//! - `config`: 配置管理(类型、YAML、热重载、导入导出)
//! - `connect`: Deep Link 协议和中转商注册表
//! - `middleware`: HTTP 中间件(认证、限速)
//! - `orchestrator`: 模型选择编排器
//! - `plugin`: 插件系统(加载、管理、UI、安装)
//! - `session`: 会话管理(限速、粘性路由)
//! - `session_files`: 会话文件存储
pub mod data;
pub mod logger;
pub mod models;
// 独立业务模块(无主 crate 依赖)
pub mod backends;
pub mod config;
pub mod connect;
pub mod errors;
pub mod middleware;
pub mod orchestrator;
pub mod plugin;
pub mod session;
pub mod session_files;
// 重新导出常用类型
pub use logger::{LogEntry, LogStore, LogStoreConfig, SharedLogStore};
pub use models::provider_type::ProviderType;
@@ -48,7 +48,7 @@ fn failure_map() -> &'static Mutex<std::collections::HashMap<String, FailureStat
}
#[cfg(test)]
pub(crate) fn clear_auth_failure_state() {
pub fn clear_auth_failure_state() {
let mut map = failure_map().lock().unwrap();
map.clear();
}
@@ -56,7 +56,7 @@ pub(crate) fn clear_auth_failure_state() {
/// 清除特定 client_id 的认证失败状态(用于测试)
/// 只清除指定的条目,不影响其他并行测试
#[cfg(test)]
pub(crate) fn clear_auth_failure_state_for(client_id: &str) {
pub fn clear_auth_failure_state_for(client_id: &str) {
let mut map = failure_map().lock().unwrap();
map.remove(client_id);
}
@@ -0,0 +1,10 @@
//! Middleware 模块
//!
//! 提供 HTTP 请求处理的中间件组件
pub mod management_auth;
#[cfg(test)]
mod tests;
pub use management_auth::ManagementAuthLayer;
@@ -23,25 +23,11 @@ fn arb_secret_key() -> impl Strategy<Value = String> {
"[a-zA-Z0-9_-]{8,32}".prop_map(|s| s)
}
/// 生成随机的无效 secret_key(与有效 key 不同)
fn arb_invalid_secret_key(valid_key: String) -> impl Strategy<Value = String> {
"[a-zA-Z0-9_-]{8,32}".prop_filter_map("must differ from valid key", move |s| {
if s != valid_key {
Some(s)
} else {
None
}
})
}
/// 生成随机的 IP 地址
fn arb_ip_addr() -> impl Strategy<Value = String> {
prop_oneof![
// localhost IPv4
Just("127.0.0.1".to_string()),
// localhost IPv6
Just("::1".to_string()),
// remote IPv4
(1u8..255u8, 0u8..255u8, 0u8..255u8, 1u8..255u8).prop_filter_map(
"not localhost",
|(a, b, c, d)| {
@@ -55,11 +41,6 @@ fn arb_ip_addr() -> impl Strategy<Value = String> {
]
}
/// 生成随机端口
fn arb_port() -> impl Strategy<Value = u16> {
1024u16..65535u16
}
/// Mock service that always returns 200 OK
#[derive(Clone)]
struct MockService;
@@ -88,40 +69,18 @@ impl Service<Request<Body>> for MockService {
/// Helper to create a request with optional Authorization header
fn create_request_with_auth(auth_header: Option<&str>) -> Request<Body> {
let mut builder = Request::builder().uri("/v0/management/status");
if let Some(auth) = auth_header {
builder = builder.header("authorization", auth);
}
builder.body(Body::empty()).unwrap()
}
/// Helper to create a request with X-Management-Key header
fn create_request_with_management_key(key: Option<&str>) -> Request<Body> {
let mut builder = Request::builder().uri("/v0/management/status");
if let Some(k) = key {
builder = builder.header("x-management-key", k);
}
builder.body(Body::empty()).unwrap()
}
/// Helper to create a request with X-Management-Key and X-Forwarded-For headers
fn create_request_with_management_key_and_forwarded(
key: Option<&str>,
forwarded_for: Option<&str>,
) -> Request<Body> {
let mut builder = Request::builder().uri("/v0/management/status");
if let Some(k) = key {
builder = builder.header("x-management-key", k);
}
if let Some(addr) = forwarded_for {
builder = builder.header("x-forwarded-for", addr);
}
builder.body(Body::empty()).unwrap()
}
@@ -137,21 +96,17 @@ fn test_management_auth_rate_limit_after_failures() {
let mut service = layer.layer(MockService);
let rt = tokio::runtime::Runtime::new().unwrap();
// 使用唯一的 IP 地址避免测试间干扰
// 直接使用原子计数器确保唯一性,避免与其他测试冲突
use std::sync::atomic::{AtomicU32, Ordering};
static RATE_LIMIT_TEST_COUNTER: AtomicU32 = AtomicU32::new(1);
let unique_id = RATE_LIMIT_TEST_COUNTER.fetch_add(1, Ordering::SeqCst);
// 使用 TEST-NET-2 (198.51.100.0/24) 范围,确保与其他测试不冲突
let octet3 = ((unique_id >> 8) & 0xFF) as u8;
let octet4 = (unique_id & 0xFF) as u8;
let client_ip = format!("198.51.{}.{}", 100 + (octet3 % 155), octet4.max(1));
let addr: SocketAddr = format!("{client_ip}:12345").parse().unwrap();
// 发送 5 次失败请求,每次都应该返回 401
// 发送 5 次失败请求
for i in 0..5 {
let mut req = create_request_with_management_key(Some("invalid"));
// 安全修复后不再信任 X-Forwarded-For,需要注入 ConnectInfo
req.extensions_mut().insert(ConnectInfo(addr));
let response = rt.block_on(async { service.call(req).await.unwrap() });
assert_eq!(
@@ -162,7 +117,7 @@ fn test_management_auth_rate_limit_after_failures() {
);
}
// 第 6 次请求应该被限速,返回 429
// 第 6 次请求应该被限速
let mut req = create_request_with_management_key(Some("invalid"));
req.extensions_mut().insert(ConnectInfo(addr));
let response = rt.block_on(async { service.call(req).await.unwrap() });
@@ -176,191 +131,96 @@ fn test_management_auth_rate_limit_after_failures() {
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection**
/// *For any* management API request without valid secret_key, the response SHALL be 401 Unauthorized.
/// **Validates: Requirements 9.3**
#[test]
fn prop_management_auth_rejection_missing_key(
secret_key in arb_secret_key()
) {
// 只清除 "unknown" 客户端的状态,避免影响并行测试
clear_auth_failure_state_for("unknown");
// Create config with a valid secret_key
let config = RemoteManagementConfig {
allow_remote: true,
secret_key: Some(secret_key),
disable_control_panel: false,
};
// Create the auth layer and service
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
// Create request WITHOUT any auth header
let req = create_request_with_auth(None);
// Execute the service
let rt = tokio::runtime::Runtime::new().unwrap();
let response = rt.block_on(async {
service.call(req).await.unwrap()
});
// Verify: should return 401 Unauthorized
prop_assert_eq!(
response.status(),
StatusCode::UNAUTHORIZED,
"Request without secret_key should return 401 Unauthorized"
);
let response = rt.block_on(async { service.call(req).await.unwrap() });
prop_assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
/// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection**
/// *For any* management API request with invalid secret_key, the response SHALL be 401 Unauthorized.
/// **Validates: Requirements 9.3**
#[test]
fn prop_management_auth_rejection_invalid_key(
secret_key in arb_secret_key()
) {
// 只清除 "unknown" 客户端的状态,避免影响并行测试
clear_auth_failure_state_for("unknown");
// Create config with a valid secret_key
let config = RemoteManagementConfig {
allow_remote: true,
secret_key: Some(secret_key.clone()),
disable_control_panel: false,
};
// Create the auth layer and service
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
// Create request with WRONG auth header (append "wrong" to make it different)
let wrong_key = format!("{secret_key}wrong");
let req = create_request_with_auth(Some(&format!("Bearer {wrong_key}")));
// Execute the service
let rt = tokio::runtime::Runtime::new().unwrap();
let response = rt.block_on(async {
service.call(req).await.unwrap()
});
// Verify: should return 401 Unauthorized
prop_assert_eq!(
response.status(),
StatusCode::UNAUTHORIZED,
"Request with invalid secret_key should return 401 Unauthorized"
);
let response = rt.block_on(async { service.call(req).await.unwrap() });
prop_assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
/// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection**
/// *For any* management API request with valid secret_key, the response SHALL NOT be 401 Unauthorized.
/// **Validates: Requirements 9.3**
#[test]
fn prop_management_auth_acceptance_valid_key(
secret_key in arb_secret_key()
) {
// 只清除 "unknown" 客户端的状态,避免影响并行测试
clear_auth_failure_state_for("unknown");
// Create config with a valid secret_key
let config = RemoteManagementConfig {
allow_remote: true,
secret_key: Some(secret_key.clone()),
disable_control_panel: false,
};
// Create the auth layer and service
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
// Create request with CORRECT auth header
let req = create_request_with_auth(Some(&format!("Bearer {secret_key}")));
// Execute the service
let rt = tokio::runtime::Runtime::new().unwrap();
let response = rt.block_on(async {
service.call(req).await.unwrap()
});
// Verify: should return 200 OK (passed through to MockService)
prop_assert_eq!(
response.status(),
StatusCode::OK,
"Request with valid secret_key should pass through (200 OK)"
);
let response = rt.block_on(async { service.call(req).await.unwrap() });
prop_assert_eq!(response.status(), StatusCode::OK);
}
/// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection**
/// *For any* management API request with valid X-Management-Key header, the response SHALL NOT be 401 Unauthorized.
/// **Validates: Requirements 9.3**
#[test]
fn prop_management_auth_acceptance_x_management_key(
secret_key in arb_secret_key()
) {
// 只清除 "unknown" 客户端的状态,避免影响并行测试
clear_auth_failure_state_for("unknown");
// Create config with a valid secret_key
let config = RemoteManagementConfig {
allow_remote: true,
secret_key: Some(secret_key.clone()),
disable_control_panel: false,
};
// Create the auth layer and service
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
// Create request with X-Management-Key header
let req = create_request_with_management_key(Some(&secret_key));
// Execute the service
let rt = tokio::runtime::Runtime::new().unwrap();
let response = rt.block_on(async {
service.call(req).await.unwrap()
});
// Verify: should return 200 OK (passed through to MockService)
prop_assert_eq!(
response.status(),
StatusCode::OK,
"Request with valid X-Management-Key should pass through (200 OK)"
);
let response = rt.block_on(async { service.call(req).await.unwrap() });
prop_assert_eq!(response.status(), StatusCode::OK);
}
/// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection**
/// *For any* management API request with invalid X-Management-Key header, the response SHALL be 401 Unauthorized.
/// **Validates: Requirements 9.3**
#[test]
fn prop_management_auth_rejection_invalid_x_management_key(
secret_key in arb_secret_key()
) {
// 只清除 "unknown" 客户端的状态,避免影响并行测试
clear_auth_failure_state_for("unknown");
// Create config with a valid secret_key
let config = RemoteManagementConfig {
allow_remote: true,
secret_key: Some(secret_key.clone()),
disable_control_panel: false,
};
// Create the auth layer and service
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
// Create request with WRONG X-Management-Key header
let wrong_key = format!("{secret_key}wrong");
let req = create_request_with_management_key(Some(&wrong_key));
// Execute the service
let rt = tokio::runtime::Runtime::new().unwrap();
let response = rt.block_on(async {
service.call(req).await.unwrap()
});
// Verify: should return 401 Unauthorized
prop_assert_eq!(
response.status(),
StatusCode::UNAUTHORIZED,
"Request with invalid X-Management-Key should return 401 Unauthorized"
);
let response = rt.block_on(async { service.call(req).await.unwrap() });
prop_assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
}
@@ -375,13 +235,10 @@ mod unit_tests {
secret_key: Some("test-secret-key".to_string()),
disable_control_panel: false,
};
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
let req = create_request_with_auth(None);
let response = service.call(req).await.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
@@ -392,13 +249,10 @@ mod unit_tests {
secret_key: Some("correct-key".to_string()),
disable_control_panel: false,
};
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
let req = create_request_with_auth(Some("Bearer wrong-key"));
let response = service.call(req).await.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
@@ -409,13 +263,10 @@ mod unit_tests {
secret_key: Some("correct-key".to_string()),
disable_control_panel: false,
};
let layer = ManagementAuthLayer::new(config);
let mut service = layer.layer(MockService);
let req = create_request_with_auth(Some("Bearer correct-key"));
let response = service.call(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
}
@@ -1,5 +1,4 @@
//! Anthropic/Claude API 数据模型
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -40,7 +39,7 @@ pub struct ImageSource {
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicMessage {
pub role: String,
pub content: serde_json::Value,
pub content: serde_json::Value, // Can be string or array
}
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -89,6 +88,7 @@ pub struct AnthropicMessagesResponse {
pub usage: AnthropicUsage,
}
// Streaming events
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum AnthropicStreamEvent {
@@ -1,5 +1,3 @@
//! 应用类型定义
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
@@ -1,7 +1,10 @@
//! CodeWhisperer/Kiro API 数据模型
//!
//! 支持标准工具和特殊工具类型(如 web_search)。
//!
//! # 更新日志
//!
//! - 2025-12-27: 添加 CWWebSearchTool 支持,修复 Issue #49
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -50,6 +53,10 @@ pub struct UserInputMessageContext {
}
/// CodeWhisperer 工具项
///
/// 支持两种类型:
/// - 标准工具(带 tool_specification)
/// - 联网搜索工具(仅 type 字段)
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum CWToolItem {
@@ -67,6 +74,9 @@ pub struct CWTool {
}
/// 联网搜索工具
///
/// Codex/Kiro API 支持的特殊工具类型,用于联网搜索。
/// 格式:`{"type": "web_search"}`
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CWWebSearchTool {
#[serde(rename = "type")]
@@ -145,6 +155,7 @@ pub struct CWToolUse {
pub tool_use_id: String,
}
// Response types
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CWStreamEvent {
@@ -72,6 +72,43 @@ impl InjectionRule {
pub fn is_exact(&self) -> bool {
!self.pattern.contains('*')
}
/// 检查模型是否匹配此规则
///
/// 支持的通配符模式:
/// - 精确匹配: `claude-sonnet-4-5`
/// - 前缀匹配: `claude-*`
/// - 后缀匹配: `*-preview`
/// - 包含匹配: `*flash*`
pub fn matches(&self, model: &str) -> bool {
if !self.enabled {
return false;
}
pattern_matches(&self.pattern, model)
}
}
/// 检查模式是否匹配模型名
///
/// 支持的通配符模式:
/// - 精确匹配: `claude-sonnet-4-5`
/// - 前缀匹配: `claude-*`
/// - 后缀匹配: `*-preview`
/// - 包含匹配: `*flash*`
pub fn pattern_matches(pattern: &str, model: &str) -> bool {
if !pattern.contains('*') {
return pattern == model;
}
let parts: Vec<&str> = pattern.split('*').collect();
match parts.as_slice() {
[prefix, ""] => model.starts_with(prefix),
["", suffix] => model.ends_with(suffix),
["", middle, ""] => model.contains(middle),
[prefix, suffix] => model.starts_with(prefix) && model.ends_with(suffix),
_ => false,
}
}
/// 规则排序:精确匹配优先,然后按优先级
@@ -37,6 +37,7 @@ impl KiroFingerprintStore {
.ok_or_else(|| "无法获取应用数据目录".to_string())?
.join("proxycast");
// 确保目录存在
if !app_data_dir.exists() {
fs::create_dir_all(&app_data_dir).map_err(|e| format!("创建应用数据目录失败: {e}"))?;
}
@@ -73,6 +74,8 @@ impl KiroFingerprintStore {
}
/// 获取或创建凭证的指纹绑定
///
/// 如果凭证没有绑定指纹,会基于凭证信息生成一个新的 Machine ID
pub fn get_or_create_binding(
&mut self,
credential_uuid: &str,
@@ -80,6 +83,7 @@ impl KiroFingerprintStore {
client_id: Option<&str>,
) -> Result<&KiroFingerprintBinding, String> {
if !self.bindings.contains_key(credential_uuid) {
// 生成基于凭证的 Machine ID
let machine_id = generate_stable_machine_id(credential_uuid, profile_arn, client_id);
let binding = KiroFingerprintBinding {
@@ -113,6 +117,9 @@ impl KiroFingerprintStore {
}
/// 生成稳定的 Machine ID
///
/// 基于凭证信息生成一个稳定的 UUID 格式 Machine ID。
/// 同一凭证每次生成的 Machine ID 相同,确保账号身份一致。
fn generate_stable_machine_id(
credential_uuid: &str,
profile_arn: Option<&str>,
@@ -120,6 +127,7 @@ fn generate_stable_machine_id(
) -> String {
use sha2::{Digest, Sha256};
// 使用凭证相关信息作为种子
let seed = format!(
"kiro_fingerprint:{}:{}:{}",
credential_uuid,
@@ -131,6 +139,7 @@ fn generate_stable_machine_id(
hasher.update(seed.as_bytes());
let result = hasher.finalize();
// 将哈希结果转换为 UUID 格式
let hex = format!("{result:x}");
format!(
"{}-{}-{}-{}-{}",
@@ -145,10 +154,15 @@ fn generate_stable_machine_id(
/// 切换到本地的结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SwitchToLocalResult {
/// 是否成功
pub success: bool,
/// 结果消息
pub message: String,
/// 是否需要用户操作(如需管理员权限)
pub requires_action: bool,
/// 切换的 Machine ID
pub machine_id: Option<String>,
/// 是否需要重启 Kiro IDE
pub requires_kiro_restart: bool,
}
+42 -3
View File
@@ -1,35 +1,49 @@
//! 机器码相关数据模型
use serde::{Deserialize, Serialize};
/// 机器码信息结构
/// 机器码信息结构 - v0.20.0
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdInfo {
/// 当前机器码
pub current_id: String,
/// 原始机器码(如果有备份)
pub original_id: Option<String>,
/// 操作系统平台
pub platform: String,
/// 是否可以修改
pub can_modify: bool,
/// 是否需要管理员权限
pub requires_admin: bool,
/// 是否存在备份
pub backup_exists: bool,
/// 机器码格式类型
pub format_type: MachineIdFormat,
}
/// 机器码操作结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdResult {
/// 操作是否成功
pub success: bool,
/// 结果消息
pub message: String,
/// 是否需要重启
pub requires_restart: bool,
/// 是否需要管理员权限
pub requires_admin: bool,
/// 新的机器码(如果操作成功)
pub new_machine_id: Option<String>,
}
/// 管理员权限状态
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AdminStatus {
/// 是否具有管理员权限
pub is_admin: bool,
/// 操作系统平台
pub platform: String,
/// 权限提升方法说明
pub elevation_method: Option<String>,
/// 权限检查是否成功
pub check_success: bool,
}
@@ -37,9 +51,12 @@ pub struct AdminStatus {
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum MachineIdFormat {
/// UUID 格式 (xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx)
Uuid,
/// 32位十六进制格式 (xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx)
#[serde(rename = "hex32")]
Hex32,
/// 其他格式
#[serde(rename = "unknown")]
Unknown,
}
@@ -47,20 +64,30 @@ pub enum MachineIdFormat {
/// 机器码备份信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdBackup {
/// 备份的机器码
pub machine_id: String,
/// 备份时间戳
pub timestamp: i64,
/// 操作系统平台
pub platform: String,
/// 机器码格式
pub format: MachineIdFormat,
/// 备份描述
pub description: Option<String>,
}
/// 机器码历史记录
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdHistory {
/// 记录ID
pub id: String,
/// 机器码
pub machine_id: String,
/// 操作时间戳
pub timestamp: String,
/// 操作系统平台
pub platform: String,
/// 备份路径(可选)
pub backup_path: Option<String>,
}
@@ -68,20 +95,30 @@ pub struct MachineIdHistory {
#[derive(Debug, Clone, Serialize, Deserialize)]
#[allow(dead_code)]
pub enum MachineIdOperation {
/// 获取当前机器码
Get,
/// 设置新机器码
Set,
/// 生成随机机器码
Generate,
/// 备份机器码
Backup,
/// 恢复机器码
Restore,
/// 重置为原始机器码
Reset,
}
/// 机器码验证结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdValidation {
/// 是否有效
pub is_valid: bool,
/// 检测到的格式
pub detected_format: MachineIdFormat,
/// 验证错误信息
pub error_message: Option<String>,
/// 格式化后的机器码(如果有效)
pub formatted_id: Option<String>,
}
@@ -90,6 +127,7 @@ impl MachineIdFormat {
pub fn detect(machine_id: &str) -> Self {
let cleaned = machine_id.replace("-", "").replace(" ", "").to_lowercase();
// 检查UUID格式:8-4-4-4-12个十六进制字符
if machine_id.contains("-") && machine_id.len() == 36 {
let parts: Vec<&str> = machine_id.split('-').collect();
if parts.len() == 5
@@ -104,6 +142,7 @@ impl MachineIdFormat {
}
}
// 检查32位十六进制格式
if cleaned.len() == 32 && cleaned.chars().all(|c| c.is_ascii_hexdigit()) {
return MachineIdFormat::Hex32;
}
+143 -2
View File
@@ -1,7 +1,48 @@
//! MCP Server 数据模型
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
/// MCP 服务器配置(类型化)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpServerConfigTyped {
/// 启动命令
pub command: String,
/// 命令参数
#[serde(default)]
pub args: Vec<String>,
/// 环境变量
#[serde(default)]
pub env: HashMap<String, String>,
/// 工作目录
#[serde(skip_serializing_if = "Option::is_none")]
pub cwd: Option<String>,
/// 超时时间(秒)
#[serde(default = "default_timeout")]
pub timeout: u64,
}
fn default_timeout() -> u64 {
30
}
impl Default for McpServerConfigTyped {
fn default() -> Self {
Self {
command: String::new(),
args: Vec::new(),
env: HashMap::new(),
cwd: None,
timeout: 30,
}
}
}
/// 配置验证错误
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConfigValidationError {
pub field: String,
pub message: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpServer {
@@ -37,4 +78,104 @@ impl McpServer {
created_at: Some(chrono::Utc::now().timestamp()),
}
}
/// 解析 server_config 为类型化配置
///
/// 将 JSON Value 解析为 McpServerConfigTyped 结构。
/// 如果解析失败,返回默认配置并尝试提取基本字段。
pub fn parse_config(&self) -> McpServerConfigTyped {
serde_json::from_value(self.server_config.clone()).unwrap_or_else(|_| {
// 尝试手动提取字段
McpServerConfigTyped {
command: self
.server_config
.get("command")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string(),
args: self
.server_config
.get("args")
.and_then(|v| v.as_array())
.map(|arr| {
arr.iter()
.filter_map(|v| v.as_str().map(|s| s.to_string()))
.collect()
})
.unwrap_or_default(),
env: self
.server_config
.get("env")
.and_then(|v| v.as_object())
.map(|obj| {
obj.iter()
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_string())))
.collect()
})
.unwrap_or_default(),
cwd: self
.server_config
.get("cwd")
.and_then(|v| v.as_str())
.map(|s| s.to_string()),
timeout: self
.server_config
.get("timeout")
.and_then(|v| v.as_u64())
.unwrap_or(30),
}
})
}
/// 验证服务器配置
///
/// 检查配置是否有效,返回验证错误列表。
/// 空列表表示配置有效。
pub fn validate_config(&self) -> Vec<ConfigValidationError> {
let mut errors = Vec::new();
let config = self.parse_config();
// 验证 command 不为空
if config.command.trim().is_empty() {
errors.push(ConfigValidationError {
field: "command".to_string(),
message: "启动命令不能为空".to_string(),
});
}
// 验证 name 不为空
if self.name.trim().is_empty() {
errors.push(ConfigValidationError {
field: "name".to_string(),
message: "服务器名称不能为空".to_string(),
});
}
// 验证 name 不包含特殊字符(用于工具名称前缀)
if !self
.name
.chars()
.all(|c| c.is_alphanumeric() || c == '-' || c == '_')
{
errors.push(ConfigValidationError {
field: "name".to_string(),
message: "服务器名称只能包含字母、数字、连字符和下划线".to_string(),
});
}
// 验证 timeout 在合理范围内
if config.timeout == 0 || config.timeout > 300 {
errors.push(ConfigValidationError {
field: "timeout".to_string(),
message: "超时时间必须在 1-300 秒之间".to_string(),
});
}
errors
}
/// 检查配置是否有效
pub fn is_valid(&self) -> bool {
self.validate_config().is_empty()
}
}
+2
View File
@@ -17,6 +17,7 @@ pub mod provider_pool_model;
pub mod provider_type;
pub mod route_model;
pub mod skill_model;
pub mod vertex_model;
#[allow(unused_imports)]
pub use anthropic::*;
@@ -33,3 +34,4 @@ pub use provider_model::Provider;
pub use provider_pool_model::*;
pub use provider_type::ProviderType;
pub use skill_model::{Skill, SkillMetadata, SkillRepo, SkillState, SkillStates};
pub use vertex_model::{VertexApiKeyEntry, VertexModelAlias};
@@ -1,25 +1,38 @@
//! 模型注册表数据结构
//!
//! 借鉴 opencode 的模型管理方式,定义增强的模型元数据结构
use serde::{Deserialize, Serialize};
/// 模型能力
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ModelCapabilities {
/// 是否支持视觉输入
pub vision: bool,
/// 是否支持工具调用
pub tools: bool,
/// 是否支持流式输出
pub streaming: bool,
/// 是否支持 JSON 模式
pub json_mode: bool,
/// 是否支持函数调用
pub function_calling: bool,
/// 是否支持推理/思考
pub reasoning: bool,
}
/// 模型定价
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelPricing {
/// 输入价格(每百万 token)
pub input_per_million: Option<f64>,
/// 输出价格(每百万 token)
pub output_per_million: Option<f64>,
/// 缓存读取价格(每百万 token)
pub cache_read_per_million: Option<f64>,
/// 缓存写入价格(每百万 token)
pub cache_write_per_million: Option<f64>,
/// 货币单位 ("USD" | "CNY")
pub currency: String,
}
@@ -38,9 +51,13 @@ impl Default for ModelPricing {
/// 模型限制
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ModelLimits {
/// 上下文长度
pub context_length: Option<u32>,
/// 最大输出 token 数
pub max_output_tokens: Option<u32>,
/// 每分钟请求数限制
pub requests_per_minute: Option<u32>,
/// 每分钟 token 数限制
pub tokens_per_minute: Option<u32>,
}
@@ -48,11 +65,17 @@ pub struct ModelLimits {
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ModelStatus {
/// 活跃可用
Active,
/// 预览版
Preview,
/// Alpha 测试
Alpha,
/// Beta 测试
Beta,
/// 已弃用
Deprecated,
/// 旧版本
Legacy,
}
@@ -95,8 +118,11 @@ impl std::str::FromStr for ModelStatus {
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ModelTier {
/// 快速响应,适合简单任务
Mini,
/// 均衡性能,适合大多数任务
Pro,
/// 最强能力,适合复杂任务
Max,
}
@@ -133,10 +159,16 @@ impl std::str::FromStr for ModelTier {
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ModelSource {
/// 从内嵌资源加载(构建时打包)
Embedded,
/// 从 models.dev API 获取(已弃用)
ModelsDev,
/// 本地硬编码(国内模型等)
Local,
/// 用户自定义
Custom,
/// 从 Provider API 获取
Api,
}
impl Default for ModelSource {
@@ -152,6 +184,7 @@ impl std::fmt::Display for ModelSource {
Self::ModelsDev => write!(f, "models.dev"),
Self::Local => write!(f, "local"),
Self::Custom => write!(f, "custom"),
Self::Api => write!(f, "api"),
}
}
}
@@ -165,6 +198,7 @@ impl std::str::FromStr for ModelSource {
"models.dev" | "modelsdev" => Ok(Self::ModelsDev),
"local" => Ok(Self::Local),
"custom" => Ok(Self::Custom),
"api" => Ok(Self::Api),
_ => Err(format!("Unknown model source: {s}")),
}
}
@@ -173,25 +207,42 @@ impl std::str::FromStr for ModelSource {
/// 增强的模型元数据
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EnhancedModelMetadata {
/// 模型 ID (如 "claude-sonnet-4-5-20250514")
pub id: String,
/// 显示名称 (如 "Claude Sonnet 4.5")
pub display_name: String,
/// Provider ID (如 "anthropic", "openai", "dashscope")
pub provider_id: String,
/// Provider 显示名称
pub provider_name: String,
/// 模型家族 (如 "sonnet", "gpt-4", "qwen")
pub family: Option<String>,
/// 服务等级
pub tier: ModelTier,
/// 模型能力
pub capabilities: ModelCapabilities,
/// 定价信息
pub pricing: Option<ModelPricing>,
/// 限制信息
pub limits: ModelLimits,
/// 模型状态
pub status: ModelStatus,
/// 发布日期
pub release_date: Option<String>,
/// 是否为最新版本
pub is_latest: bool,
/// 描述
pub description: Option<String>,
/// 数据来源
pub source: ModelSource,
/// 创建时间 (Unix 时间戳)
pub created_at: i64,
/// 最后更新时间 (Unix 时间戳)
pub updated_at: i64,
}
impl EnhancedModelMetadata {
/// 创建新的模型元数据
pub fn new(
id: String,
display_name: String,
@@ -219,51 +270,61 @@ impl EnhancedModelMetadata {
}
}
/// 设置模型家族
pub fn with_family(mut self, family: impl Into<String>) -> Self {
self.family = Some(family.into());
self
}
/// 设置服务等级
pub fn with_tier(mut self, tier: ModelTier) -> Self {
self.tier = tier;
self
}
/// 设置模型能力
pub fn with_capabilities(mut self, capabilities: ModelCapabilities) -> Self {
self.capabilities = capabilities;
self
}
/// 设置定价信息
pub fn with_pricing(mut self, pricing: ModelPricing) -> Self {
self.pricing = Some(pricing);
self
}
/// 设置限制信息
pub fn with_limits(mut self, limits: ModelLimits) -> Self {
self.limits = limits;
self
}
/// 设置模型状态
pub fn with_status(mut self, status: ModelStatus) -> Self {
self.status = status;
self
}
/// 设置发布日期
pub fn with_release_date(mut self, date: impl Into<String>) -> Self {
self.release_date = Some(date.into());
self
}
/// 设置是否为最新版本
pub fn with_is_latest(mut self, is_latest: bool) -> Self {
self.is_latest = is_latest;
self
}
/// 设置描述
pub fn with_description(mut self, description: impl Into<String>) -> Self {
self.description = Some(description.into());
self
}
/// 设置数据来源
pub fn with_source(mut self, source: ModelSource) -> Self {
self.source = source;
self
@@ -273,17 +334,26 @@ impl EnhancedModelMetadata {
/// 用户模型偏好
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UserModelPreference {
/// 模型 ID
pub model_id: String,
/// 是否收藏
pub is_favorite: bool,
/// 是否隐藏
pub is_hidden: bool,
/// 自定义别名
pub custom_alias: Option<String>,
/// 使用次数
pub usage_count: u32,
/// 最后使用时间 (Unix 时间戳)
pub last_used_at: Option<i64>,
/// 创建时间 (Unix 时间戳)
pub created_at: i64,
/// 更新时间 (Unix 时间戳)
pub updated_at: i64,
}
impl UserModelPreference {
/// 创建新的用户偏好
pub fn new(model_id: String) -> Self {
let now = chrono::Utc::now().timestamp();
Self {
@@ -302,45 +372,63 @@ impl UserModelPreference {
/// 模型同步状态
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ModelSyncState {
/// 最后同步时间 (Unix 时间戳)
pub last_sync_at: Option<i64>,
/// 同步的模型数量
pub model_count: u32,
/// 是否正在同步
pub is_syncing: bool,
/// 最后同步错误
pub last_error: Option<String>,
}
// Provider Alias 相关类型
// ============================================================================
// Provider Alias 相关类型(用于 Kiro、Antigravity 等中转服务)
// ============================================================================
/// 单个模型别名映射
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelAlias {
/// 实际模型 ID(如 "claude-sonnet-4-5-20250929")
pub actual: String,
/// 内部 API 名称(如 "CLAUDE_SONNET_4_5_20250929_V1_0")
pub internal_name: Option<String>,
/// 原始 Provider(如 "anthropic")
pub provider: Option<String>,
/// 描述
pub description: Option<String>,
}
/// Provider 的别名配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderAliasConfig {
/// Provider ID(如 "kiro"、"antigravity")
pub provider: String,
/// 描述
pub description: Option<String>,
/// 支持的模型列表
#[serde(default)]
pub models: Vec<String>,
/// 别名映射(模型名 -> 别名配置)
pub aliases: std::collections::HashMap<String, ModelAlias>,
/// 更新时间
pub updated_at: Option<String>,
}
impl ProviderAliasConfig {
/// 检查是否支持指定模型
pub fn supports_model(&self, model: &str) -> bool {
self.models.contains(&model.to_string()) || self.aliases.contains_key(model)
}
/// 获取模型的内部名称
pub fn get_internal_name(&self, model: &str) -> Option<&str> {
self.aliases
.get(model)
.and_then(|a| a.internal_name.as_deref())
}
/// 获取模型的实际 ID
pub fn get_actual_model(&self, model: &str) -> Option<&str> {
self.aliases.get(model).map(|a| a.actual.as_str())
}
@@ -421,6 +509,8 @@ pub struct ModelsDevModalities {
}
impl ModelsDevModel {
/// 转换为 EnhancedModelMetadata
/// 预留:用于从 models.dev API 导入模型数据
#[allow(dead_code)]
pub fn to_enhanced_metadata(
&self,
@@ -429,6 +519,7 @@ impl ModelsDevModel {
) -> EnhancedModelMetadata {
let now = chrono::Utc::now().timestamp();
// 判断是否支持视觉
let supports_vision = self
.modalities
.as_ref()
@@ -436,14 +527,17 @@ impl ModelsDevModel {
.unwrap_or(false)
|| self.attachment;
// 根据模型名称推断服务等级
let tier = infer_model_tier(&self.id, &self.name);
// 解析状态
let status = self
.status
.as_ref()
.and_then(|s| s.parse().ok())
.unwrap_or(ModelStatus::Active);
// 判断是否为最新版本
let is_latest = self.id.contains("latest");
EnhancedModelMetadata {
@@ -456,8 +550,8 @@ impl ModelsDevModel {
capabilities: ModelCapabilities {
vision: supports_vision,
tools: self.tool_call,
streaming: true,
json_mode: true,
streaming: true, // 大多数模型都支持流式
json_mode: true, // 大多数模型都支持 JSON 模式
function_calling: self.tool_call,
reasoning: self.reasoning,
},
@@ -486,28 +580,13 @@ impl ModelsDevModel {
}
/// 根据模型 ID 和名称推断服务等级
/// 用于 to_enhanced_metadata 和测试
#[allow(dead_code)]
fn infer_model_tier(model_id: &str, model_name: &str) -> ModelTier {
let id_lower = model_id.to_lowercase();
let name_lower = model_name.to_lowercase();
let max_patterns = [
"opus",
"gpt-4o",
"gpt-4-turbo",
"gemini-2.5-pro",
"gemini-ultra",
"claude-3-opus",
"qwen-max",
"glm-4-plus",
"deepseek-v3",
];
for pattern in max_patterns {
if id_lower.contains(pattern) || name_lower.contains(pattern) {
return ModelTier::Max;
}
}
// Mini 等级模型(优先检查,因为 gpt-4o-mini 包含 gpt-4o)
let mini_patterns = [
"mini",
"nano",
@@ -525,6 +604,25 @@ fn infer_model_tier(model_id: &str, model_name: &str) -> ModelTier {
}
}
// Max 等级模型
let max_patterns = [
"opus",
"gpt-4o",
"gpt-4-turbo",
"gemini-2.5-pro",
"gemini-ultra",
"claude-3-opus",
"qwen-max",
"glm-4-plus",
"deepseek-v3",
];
for pattern in max_patterns {
if id_lower.contains(pattern) || name_lower.contains(pattern) {
return ModelTier::Max;
}
}
// 默认为 Pro 等级
ModelTier::Pro
}
+51 -1
View File
@@ -1,7 +1,15 @@
//! OpenAI API 数据模型
//!
//! 支持标准 OpenAI 格式以及扩展的工具类型(如 web_search)。
//!
//! # 工具类型支持
//!
//! - `function`: 标准函数调用工具
//! - `web_search`: 联网搜索工具(Claude Code 使用 `web_search_20250305`)
//!
//! # 更新日志
//!
//! - 2025-12-27: 添加 web_search 工具支持,修复 Issue #49
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -50,6 +58,8 @@ pub struct ChatMessage {
pub tool_calls: Option<Vec<ToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
/// 推理内容(DeepSeek R1 等模型的思维链内容)
/// DeepSeek Reasoner 在 Tool Calls 场景下要求此字段
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_content: Option<String>,
}
@@ -74,21 +84,25 @@ impl ChatMessage {
}
/// 提取消息中的图片 URL 列表
/// 返回 (format, base64_data) 元组列表
pub fn get_images(&self) -> Vec<(String, String)> {
match &self.content {
Some(MessageContent::Parts(parts)) => parts
.iter()
.filter_map(|p| {
if let ContentPart::ImageUrl { image_url } = p {
// 解析 data URL: data:image/jpeg;base64,xxxxx
if image_url.url.starts_with("data:") {
let parts: Vec<&str> = image_url.url.splitn(2, ',').collect();
if parts.len() == 2 {
// 提取 media_type: data:image/jpeg;base64 -> image/jpeg
let header = parts[0];
let data = parts[1];
let media_type = header
.strip_prefix("data:")
.and_then(|s| s.split(';').next())
.unwrap_or("image/jpeg");
// 提取格式: image/jpeg -> jpeg
let format =
media_type.split('/').nth(1).unwrap_or("jpeg").to_string();
return Some((format, data.to_string()));
@@ -115,13 +129,21 @@ pub struct FunctionDef {
}
/// 工具定义
///
/// 支持多种工具类型:
/// - `function`: 标准函数调用工具,包含 function 字段
/// - `web_search`: 联网搜索工具,无需额外字段
/// - `web_search_20250305`: Claude Code 的联网搜索工具类型
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum Tool {
/// 标准函数调用工具
#[serde(rename = "function")]
Function { function: FunctionDef },
/// 联网搜索工具(Codex/Kiro 格式)
#[serde(rename = "web_search")]
WebSearch,
/// 联网搜索工具(Claude Code 格式)
#[serde(rename = "web_search_20250305")]
WebSearch20250305,
}
@@ -142,6 +164,7 @@ pub struct ChatCompletionRequest {
pub tools: Option<Vec<Tool>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<serde_json::Value>,
/// 思维链强度:none, low, medium, high
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<String>,
}
@@ -206,24 +229,43 @@ pub struct ChatCompletionChunk {
pub choices: Vec<StreamChoice>,
}
// ============================================================================
// 图像生成 API 数据模型
// ============================================================================
/// OpenAI 图像生成请求
///
/// 兼容 OpenAI Images API,支持通过 Antigravity 生成图像。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationRequest {
/// 图像生成提示词
pub prompt: String,
/// 模型名称 (默认: gemini-3-pro-image-preview)
#[serde(default = "default_image_model")]
pub model: String,
/// 生成图像数量 (默认: 1)
#[serde(default = "default_n")]
pub n: u32,
/// 图像尺寸 (可选,Antigravity 可能忽略)
#[serde(skip_serializing_if = "Option::is_none")]
pub size: Option<String>,
/// 响应格式: "url" 或 "b64_json" (默认: "url")
#[serde(default = "default_response_format")]
pub response_format: String,
/// 图像质量 (可选,Antigravity 可能忽略)
#[serde(skip_serializing_if = "Option::is_none")]
pub quality: Option<String>,
/// 图像风格 (可选,Antigravity 可能忽略)
#[serde(skip_serializing_if = "Option::is_none")]
pub style: Option<String>,
/// 用户标识 (可选)
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
@@ -243,17 +285,25 @@ fn default_response_format() -> String {
/// OpenAI 图像生成响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationResponse {
/// 创建时间戳 (Unix epoch seconds)
pub created: i64,
/// 生成的图像数组
pub data: Vec<ImageData>,
}
/// 单个图像数据
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageData {
/// Base64 编码的图像数据 (当 response_format="b64_json")
#[serde(skip_serializing_if = "Option::is_none")]
pub b64_json: Option<String>,
/// 图像 URL (当 response_format="url",返回 data URL)
#[serde(skip_serializing_if = "Option::is_none")]
pub url: Option<String>,
/// 修订后的提示词 (如果 Antigravity 返回了文本)
#[serde(skip_serializing_if = "Option::is_none")]
pub revised_prompt: Option<String>,
}
@@ -1,5 +1,3 @@
//! Prompt 数据模型
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -10,6 +8,7 @@ pub struct Prompt {
pub content: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
/// Whether this prompt is currently enabled (synced to live file)
#[serde(default)]
pub enabled: bool,
#[serde(rename = "createdAt", skip_serializing_if = "Option::is_none")]
@@ -1,5 +1,3 @@
//! Provider 数据模型
use serde::{Deserialize, Serialize};
use serde_json::Value;
@@ -27,7 +27,7 @@ pub enum CredentialSource {
///
/// 为了向后兼容,PoolProviderType 是 crate::ProviderType 的类型别名。
/// 所有 Provider 类型定义已统一到 lib.rs 中的 ProviderType。
pub type PoolProviderType = crate::ProviderType;
pub type PoolProviderType = super::provider_type::ProviderType;
/// 凭证数据,根据 Provider 类型不同而不同
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -7,38 +7,53 @@ use serde::{Deserialize, Serialize};
/// 单个路由信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteInfo {
/// 路由选择器 (provider 类型或凭证名称)
pub selector: String,
/// Provider 类型
pub provider_type: String,
/// 关联的凭证数量
pub credential_count: usize,
/// 可用的端点列表
pub endpoints: Vec<RouteEndpoint>,
/// 标签 (如 "突破限制", "官方API/三方")
pub tags: Vec<String>,
/// 是否启用
pub enabled: bool,
}
/// 路由端点
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteEndpoint {
/// 端点路径
pub path: String,
pub protocol: String,
/// 协议类型
pub protocol: String, // "openai" 或 "claude"
/// 完整 URL
pub url: String,
}
/// 路由列表响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteListResponse {
/// 服务器基础 URL
pub base_url: String,
/// 默认 Provider
pub default_provider: String,
/// 所有可用路由
pub routes: Vec<RouteInfo>,
}
/// curl 示例
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CurlExample {
/// 描述
pub description: String,
/// curl 命令
pub command: String,
}
impl RouteInfo {
/// 创建新的路由信息
pub fn new(selector: String, provider_type: String) -> Self {
Self {
selector,
@@ -50,6 +65,7 @@ impl RouteInfo {
}
}
/// 添加端点
pub fn add_endpoint(&mut self, base_url: &str, protocol: &str) {
let path = match protocol {
"claude" => format!("/{}/v1/messages", self.selector),
@@ -64,6 +80,7 @@ impl RouteInfo {
});
}
/// 生成 curl 示例
pub fn generate_curl_examples(&self, api_key: &str) -> Vec<CurlExample> {
let mut examples = Vec::new();
@@ -1,5 +1,3 @@
//! Skill 数据模型
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
@@ -77,6 +75,7 @@ impl SkillRepo {
pub fn get_default_skill_repos() -> Vec<SkillRepo> {
vec![
// ProxyCast 官方仓库(排第一位)
SkillRepo {
owner: "proxycast".to_string(),
name: "skills".to_string(),
@@ -109,13 +108,18 @@ pub type SkillStates = HashMap<String, SkillState>;
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
/// Feature: skills-platform-mvp, Property 1: Default Repositories Include ProxyCast Official
/// Validates: Requirements 1.1, 1.2, 1.3
#[test]
fn test_default_repos_include_proxycast_official() {
let repos = get_default_skill_repos();
// 验证列表非空
assert!(!repos.is_empty(), "默认仓库列表不应为空");
// 验证第一个仓库是 ProxyCast 官方仓库
let first_repo = &repos[0];
assert_eq!(
first_repo.owner, "proxycast",
@@ -126,10 +130,34 @@ mod tests {
assert!(first_repo.enabled, "ProxyCast 官方仓库应默认启用");
}
// Property 1: Default Repositories Include ProxyCast Official (Property-Based Test)
// For any call to get_default_skill_repos(), the returned list SHALL contain
// a SkillRepo with owner="proxycast", name="skills", branch="main", and enabled=true,
// and this repo SHALL be the first item in the list.
// Validates: Requirements 1.1, 1.2, 1.3
proptest! {
#[test]
fn prop_default_repos_proxycast_first(_seed in 0u64..1000) {
// 无论调用多少次,结果应该一致
let repos = get_default_skill_repos();
// Property: 列表非空
prop_assert!(!repos.is_empty());
// Property: 第一个仓库是 ProxyCast 官方仓库
let first = &repos[0];
prop_assert_eq!(&first.owner, "proxycast");
prop_assert_eq!(&first.name, "skills");
prop_assert_eq!(&first.branch, "main");
prop_assert!(first.enabled);
}
}
#[test]
fn test_proxycast_repo_exists_in_list() {
let repos = get_default_skill_repos();
// 验证 ProxyCast 仓库存在于列表中
let proxycast_repo = repos
.iter()
.find(|r| r.owner == "proxycast" && r.name == "skills");
@@ -0,0 +1,35 @@
//! Vertex AI 配置模型
//!
//! 定义 Vertex AI 相关的配置类型,供 providers 和 config 模块共享。
use serde::{Deserialize, Serialize};
/// Vertex AI 模型别名映射
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct VertexModelAlias {
/// 上游模型名称
pub name: String,
/// 客户端可见的别名
pub alias: String,
}
/// Vertex AI 凭证条目
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct VertexApiKeyEntry {
/// 凭证 ID
pub id: String,
/// API Key
pub api_key: String,
/// Base URL
#[serde(default, skip_serializing_if = "Option::is_none")]
pub base_url: Option<String>,
/// 模型别名映射
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub models: Vec<VertexModelAlias>,
/// 单独的代理 URL
#[serde(default, skip_serializing_if = "Option::is_none")]
pub proxy_url: Option<String>,
/// 是否禁用
#[serde(default)]
pub disabled: bool,
}
@@ -297,8 +297,8 @@ pub struct PoolStats {
}
/// 全局编排器实例
static GLOBAL_ORCHESTRATOR: once_cell::sync::OnceCell<Arc<ModelOrchestrator>> =
once_cell::sync::OnceCell::new();
static GLOBAL_ORCHESTRATOR: std::sync::OnceLock<Arc<ModelOrchestrator>> =
std::sync::OnceLock::new();
/// 初始化全局编排器
pub fn init_global_orchestrator() -> Arc<ModelOrchestrator> {
+36
View File
@@ -0,0 +1,36 @@
//! 插件系统模块
//!
//! 提供插件扩展功能,支持:
//! - 插件加载和初始化
//! - 请求前/响应后钩子
//! - 插件隔离和错误处理
//! - 插件配置管理
//! - 二进制组件下载和管理
//! - 声明式插件 UI 系统
//! - 插件安装和卸载
pub mod binary_downloader;
pub mod examples;
pub mod installer;
mod loader;
mod manager;
mod types;
pub mod ui_builder;
pub mod ui_trait;
pub mod ui_types;
pub use binary_downloader::BinaryDownloader;
pub use loader::PluginLoader;
pub use manager::PluginManager;
pub use types::{
BinaryComponentStatus, BinaryManifest, HookResult, PlatformBinaries, Plugin, PluginConfig,
PluginContext, PluginError, PluginInfo, PluginManifest, PluginState, PluginStatus, PluginType,
};
pub use ui_trait::{NoUI, PluginUI};
pub use ui_types::{
Action, BoundValue, ChildrenDef, ComponentDef, ComponentType, DataEntry, DataModelUpdate,
SurfaceDefinition, SurfaceUpdate, UIMessage, UserAction,
};
#[cfg(test)]
mod tests;
+16
View File
@@ -0,0 +1,16 @@
//! 会话管理核心模块
//!
//! 提供以下功能:
//! - 增强的限流处理(Duration 解析、指数退避)
//! - 会话粘性管理(会话与账号映射)
//! - 调度模式配置
pub mod rate_limit;
pub mod sticky_config;
pub mod sticky_manager;
pub use rate_limit::{
extract_retry_delay, parse_duration_string, RateLimitReason, RateLimitRecord, RateLimitTracker,
};
pub use sticky_config::{SchedulingMode, StickySessionConfig};
pub use sticky_manager::{AccountInfo, StickySessionManager};
@@ -3,7 +3,7 @@
//! 为每个 Agent 会话提供独立的临时工作目录,
//!
//! ## 目录结构
//! ```
//! ```text
//! ~/.proxycast/sessions/
//! ├── {session-id}/
//! │ ├── .meta.json # 会话元数据
+9 -128
View File
@@ -1,9 +1,13 @@
//! 参数注入类型定义
//!
//! 定义注入规则、注入模式和注入器
//! 基础类型(InjectionMode, InjectionRule)从 proxycast-core 重新导出。
//! 本模块定义注入器(Injector)和注入结果等 infra 层特有类型。
use serde::{Deserialize, Serialize};
// 从 core 重新导出基础类型
pub use proxycast_core::models::injection_types::{InjectionMode, InjectionRule};
/// 允许注入的参数白名单
/// 这些参数是安全的,不会影响请求的核心行为
const ALLOWED_INJECTION_PARAMS: &[&str] = &[
@@ -28,110 +32,6 @@ const BLOCKED_OVERRIDE_PARAMS: &[&str] = &[
"response_format",
];
/// 注入模式
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum InjectionMode {
/// 合并模式:不覆盖已有参数
#[default]
Merge,
/// 覆盖模式:覆盖已有参数
Override,
}
/// 注入规则
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct InjectionRule {
/// 规则 ID
pub id: String,
/// 模型匹配模式(支持通配符)
pub pattern: String,
/// 要注入的参数
pub parameters: serde_json::Value,
/// 注入模式
#[serde(default)]
pub mode: InjectionMode,
/// 优先级(数字越小优先级越高)
#[serde(default = "default_priority")]
pub priority: i32,
/// 是否启用
#[serde(default = "default_enabled")]
pub enabled: bool,
}
fn default_priority() -> i32 {
100
}
fn default_enabled() -> bool {
true
}
impl InjectionRule {
/// 创建新的注入规则
pub fn new(id: &str, pattern: &str, parameters: serde_json::Value) -> Self {
Self {
id: id.to_string(),
pattern: pattern.to_string(),
parameters,
mode: InjectionMode::Merge,
priority: default_priority(),
enabled: true,
}
}
/// 设置注入模式
pub fn with_mode(mut self, mode: InjectionMode) -> Self {
self.mode = mode;
self
}
/// 设置优先级
pub fn with_priority(mut self, priority: i32) -> Self {
self.priority = priority;
self
}
/// 检查模型是否匹配此规则
///
/// 支持的通配符模式:
/// - 精确匹配: `claude-sonnet-4-5`
/// - 前缀匹配: `claude-*`
/// - 后缀匹配: `*-preview`
/// - 包含匹配: `*flash*`
pub fn matches(&self, model: &str) -> bool {
if !self.enabled {
return false;
}
pattern_matches(&self.pattern, model)
}
/// 检查是否为精确匹配规则
pub fn is_exact(&self) -> bool {
!self.pattern.contains('*')
}
}
/// 规则排序:精确匹配优先,然后按优先级
impl Ord for InjectionRule {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
match (self.is_exact(), other.is_exact()) {
(true, false) => return std::cmp::Ordering::Less,
(false, true) => return std::cmp::Ordering::Greater,
_ => {}
}
self.priority.cmp(&other.priority)
}
}
impl PartialOrd for InjectionRule {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Eq for InjectionRule {}
/// 注入结果
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct InjectionResult {
@@ -164,6 +64,10 @@ pub struct InjectionConfig {
pub rules: Vec<InjectionRule>,
}
fn default_enabled() -> bool {
true
}
/// 参数注入器
#[derive(Debug, Clone, Default)]
pub struct Injector {
@@ -273,26 +177,3 @@ impl Injector {
result
}
}
/// 检查模式是否匹配模型名
///
/// 支持的通配符模式:
/// - 精确匹配: `claude-sonnet-4-5`
/// - 前缀匹配: `claude-*`
/// - 后缀匹配: `*-preview`
/// - 包含匹配: `*flash*`
fn pattern_matches(pattern: &str, model: &str) -> bool {
if !pattern.contains('*') {
return pattern == model;
}
let parts: Vec<&str> = pattern.split('*').collect();
match parts.as_slice() {
[prefix, ""] => model.starts_with(prefix),
["", suffix] => model.ends_with(suffix),
["", middle, ""] => model.contains(middle),
[prefix, suffix] => model.starts_with(prefix) && model.ends_with(suffix),
_ => false,
}
}
+55
View File
@@ -0,0 +1,55 @@
[package]
name = "proxycast-providers"
version.workspace = true
edition.workspace = true
authors.workspace = true
repository.workspace = true
[dependencies]
# 项目内 crate
proxycast-core.workspace = true
# 序列化
serde.workspace = true
serde_json.workspace = true
serde_urlencoded.workspace = true
# 异步运行时
tokio.workspace = true
futures.workspace = true
async-stream.workspace = true
async-trait.workspace = true
# 错误处理
anyhow.workspace = true
thiserror.workspace = true
# 日志
tracing.workspace = true
# HTTP 客户端
reqwest.workspace = true
# 时间和 UUID
chrono.workspace = true
uuid.workspace = true
# 工具库
base64.workspace = true
bytes.workspace = true
dirs.workspace = true
flate2.workspace = true
sha2.workspace = true
rand.workspace = true
open.workspace = true
urlencoding.workspace = true
regex.workspace = true
once_cell.workspace = true
url.workspace = true
# HTTP 服务器(streaming 模块需要 axum 类型)
axum.workspace = true
[dev-dependencies]
proptest.workspace = true
tempfile.workspace = true
@@ -1,6 +1,6 @@
//! Anthropic 格式转换为 OpenAI 格式 (支持 Claude Code)
use crate::models::anthropic::*;
use crate::models::openai::*;
use proxycast_core::models::anthropic::*;
use proxycast_core::models::openai::*;
use uuid::Uuid;
/// 将 Anthropic MessagesRequest 转换为 OpenAI ChatCompletionRequest
@@ -1,8 +1,8 @@
//! CodeWhisperer 响应转换为 OpenAI 格式
#![allow(dead_code)]
use crate::models::codewhisperer::*;
use crate::models::openai::*;
use proxycast_core::models::codewhisperer::*;
use proxycast_core::models::openai::*;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
@@ -14,8 +14,8 @@
//! ## 更新日志
//! - 2025-12-28: 修复请求格式,对齐 CLIProxyAPI 实现
use crate::models::openai::*;
use crate::session::{get_thought_signature, SessionManager};
use proxycast_core::models::openai::*;
use serde::{Deserialize, Serialize};
use uuid::Uuid;
@@ -1061,7 +1061,7 @@ pub fn convert_antigravity_to_openai_response(
// 图像生成 API 转换函数
// ============================================================================
use crate::models::openai::{ImageData, ImageGenerationRequest, ImageGenerationResponse};
use proxycast_core::models::openai::{ImageData, ImageGenerationRequest, ImageGenerationResponse};
/// 图像生成模型名称映射
///
@@ -8,8 +8,8 @@
#![allow(dead_code)]
use crate::models::codewhisperer::*;
use crate::models::openai::*;
use proxycast_core::models::codewhisperer::*;
use proxycast_core::models::openai::*;
use std::collections::HashMap;
use uuid::Uuid;
@@ -4,7 +4,7 @@
#![allow(dead_code)]
use crate::models::provider_pool_model::PoolProviderType;
use proxycast_core::models::provider_pool_model::PoolProviderType;
/// 协议类型
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -24,7 +24,7 @@
// 预留功能模块,暂未在主流程中调用
#![allow(dead_code)]
use crate::models::openai::ChatMessage;
use proxycast_core::models::openai::ChatMessage;
/// 模型类型,用于确定推理内容处理策略
#[derive(Debug, Clone, PartialEq)]
@@ -147,7 +147,7 @@ impl ReasoningHandler {
#[cfg(test)]
mod tests {
use super::*;
use crate::models::openai::MessageContent;
use proxycast_core::models::openai::MessageContent;
#[test]
fn test_model_type_detection() {
+18
View File
@@ -0,0 +1,18 @@
//! ProxyCast Providers Crate
//!
//! 包含所有 Provider 实现、协议转换、流式传输等核心业务模块。
//!
//! ## 模块结构
//! - `providers`: Provider 实现(Kiro、Gemini、Claude、OpenAI、Vertex 等)
//! - `converter`: 协议转换(OpenAI ↔ CW、OpenAI ↔ Antigravity 等)
//! - `streaming`: 流式传输管理
//! - `translator`: 请求/响应翻译层
//! - `stream`: 流事件解析和生成
//! - `session`: 会话管理(签名存储、会话 ID 生成)
pub mod converter;
pub mod providers;
pub mod session;
pub mod stream;
pub mod streaming;
pub mod translator;
@@ -2109,11 +2109,11 @@ impl CredentialProvider for AntigravityProvider {
// ============================================================================
use crate::converter::openai_to_antigravity::convert_openai_to_antigravity_with_context;
use crate::models::openai::ChatCompletionRequest;
use crate::providers::ProviderError;
use crate::streaming::traits::{
reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider,
};
use proxycast_core::models::openai::ChatCompletionRequest;
#[async_trait]
impl StreamingProvider for AntigravityProvider {
@@ -1,6 +1,6 @@
//! Claude Custom Provider (自定义 Claude API)
use crate::models::anthropic::AnthropicMessagesRequest;
use crate::models::openai::{ChatCompletionRequest, ContentPart, MessageContent};
use proxycast_core::models::anthropic::AnthropicMessagesRequest;
use proxycast_core::models::openai::{ChatCompletionRequest, ContentPart, MessageContent};
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::error::Error;
@@ -560,7 +560,7 @@ impl StreamingProvider for ClaudeCustomProvider {
.iter()
.filter_map(|tool| {
match tool {
crate::models::openai::Tool::Function { function } => {
proxycast_core::models::openai::Tool::Function { function } => {
Some(serde_json::json!({
"name": function.name,
"description": function.description.clone().unwrap_or_default(),

Some files were not shown because too many files have changed in this diff Show More