feat: 实现 Agent 调度器系统

实现完整的 Agent 任务调度系统,支持定时任务、批量处理和后台持续运行。

## 新增功能

### 1. proxycast-scheduler Crate
- 创建独立的调度器 crate
- 实现 SchedulerTrait 接口
- 支持任务 CRUD 操作
- 任务持久化到 SQLite

### 2. 数据模型
- ScheduledTask: 调度任务结构
- TaskStatus: 任务状态枚举 (Pending, Running, Completed, Failed, Cancelled)
- TaskFilter: 任务查询过滤器

### 3. AgentScheduler
- 实现调度器核心逻辑
- 支持任务创建、查询、更新、删除
- 支持到期任务查询
- 自动状态管理

### 4. AgentExecutor
- 实现任务执行器
- 通过 CredentialBridge 自动选择凭证
- 支持多种任务类型:
  - agent_chat: Agent 对话任务
  - batch_process: 批量处理任务
  - scheduled_report: 定时报告任务
- 自动标记凭证健康状态

### 5. SchedulerService (Heartbeat Loop)
- 实现心跳循环服务
- 每 30 秒轮询到期任务
- 自动执行到期任务
- 支持优雅启动和关闭 (CancellationToken)
- 集成到 Tauri setup

### 6. 数据库集成
- 创建 scheduled_tasks 表
- 添加索引 (status, scheduled_at, provider_type)
- 自动初始化表结构

## 技术实现

- 使用 tokio::spawn 启动后台任务
- 使用 tokio::time::interval 实现定时轮询
- 使用 tokio::select! 监听取消信号
- 完整的错误处理和日志记录
- 所有模块包含单元测试

## 文件清单

新增:
- src-tauri/crates/scheduler/Cargo.toml
- src-tauri/crates/scheduler/src/lib.rs
- src-tauri/crates/scheduler/src/types.rs
- src-tauri/crates/scheduler/src/dao.rs
- src-tauri/crates/scheduler/src/scheduler.rs
- src-tauri/crates/scheduler/src/executor.rs
- src-tauri/src/app/scheduler_service.rs

修改:
- src-tauri/Cargo.toml (添加 scheduler 依赖)
- src-tauri/src/app/mod.rs (导出 scheduler_service)
- src-tauri/src/app/setup.rs (集成调度器服务)

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
coso
2026-02-09 21:47:52 +08:00
co-authored by Claude Opus 4.6
parent bf0c1a05fb
commit b14dbe8271
67 changed files with 8612 additions and 689 deletions
+31
View File
@@ -0,0 +1,31 @@
# 批量任务支持实施计划
## Stage 1: 创建 Scheduler Crate
**Goal**: 创建基础调度器模块,定义批量任务数据结构
**Success Criteria**: Crate 编译通过,基础数据结构定义完成
**Tests**: 单元测试验证数据结构序列化
**Status**: Completed
## Stage 2: 实现批量任务执行器
**Goal**: 实现 BatchTaskExecutor,支持并发控制和 Orchestrator Fallback
**Success Criteria**: 执行器能够处理批量任务,支持并发控制
**Tests**: 集成测试验证批量任务执行逻辑
**Status**: In Progress
## Stage 3: 创建 Batch API 端点
**Goal**: 实现批量任务的 REST API
**Success Criteria**: POST /api/batch/tasks 和 GET /api/batch/tasks/:id 可用
**Tests**: API 测试验证创建和查询功能
**Status**: Completed
## Stage 4: 数据库持久化
**Goal**: 实现批量任务和模板的数据库存储
**Success Criteria**: 数据可以持久化到 SQLite
**Tests**: DAO 层单元测试
**Status**: In Progress
## Stage 5: 前端页面实现
**Goal**: 创建批量任务管理界面
**Success Criteria**: 任务列表、创建页面、结果展示完整
**Tests**: 手动测试 UI 交互流程
**Status**: Not Started
+592
View File
@@ -0,0 +1,592 @@
# ProxyCast 设置页面重构设计
> 参考 LobeHub 的设置架构,为 ProxyCast 设计现代化的设置界面
## 一、设计目标
1. **分类清晰**:将设置项按功能分组,便于用户快速定位
2. **侧边导航**:采用左侧菜单 + 右侧内容的布局
3. **可扩展性**:支持动态添加设置模块
4. **一致性**:与 LobeHub 风格保持一致
## 二、设置分类设计
### 分组结构
```
📁 账号 (Account)
├── 👤 个人资料 (Profile)
└── 📊 数据统计 (Stats)
📁 通用 (General)
├── 🎨 外观 (Appearance)
├── 💬 聊天外观 (Chat Appearance)
└── ⌨️ 快捷键 (Hotkeys)
📁 智能体 (Agent)
├── 🧠 AI 服务商 (Providers) → 现有 Provider Pool
├── 🤖 助理服务 (Assistant) → Agent 配置
├── 🔧 技能管理 (Skills/MCP) → 现有 MCP 页面
├── 🧩 记忆设置 (Memory) → 新增
├── 🎨 绘画服务 (Image Gen) → 现有 Image Gen 配置
└── 🎤 语音服务 (Voice/TTS) → 新增
📁 系统 (System)
├── 🌐 网络代理 (Proxy) → 现有 ProxySettings
├── 💾 数据存储 (Storage) → 现有 DirectorySettings
├── 🔒 安全设置 (Security) → 现有 TlsSettings + RemoteManagement
├── 🔌 外部工具 (External Tools) → 现有 ExternalToolsSettings
├── 🧪 实验功能 (Experimental) → 现有 ExperimentalSettings
├── 💻 开发者 (Developer) → 现有 DeveloperSettings
└── ℹ️ 关于 (About) → 现有 AboutSection
```
## 三、路由设计
### 设置页面路由枚举
```typescript
// src/types/settings.ts
export enum SettingsGroupKey {
Account = 'account',
General = 'general',
Agent = 'agent',
System = 'system',
}
export enum SettingsTabs {
// 账号
Profile = 'profile',
Stats = 'stats',
// 通用
Appearance = 'appearance',
ChatAppearance = 'chat-appearance',
Hotkeys = 'hotkeys',
// 智能体
Providers = 'providers',
Assistant = 'assistant',
Skills = 'skills',
Memory = 'memory',
ImageGen = 'image-gen',
Voice = 'voice',
// 系统
Proxy = 'proxy',
Storage = 'storage',
Security = 'security',
ExternalTools = 'external-tools',
Experimental = 'experimental',
Developer = 'developer',
About = 'about',
}
```
## 四、目录结构
```
src/components/settings/
├── _layout/ # 布局层
│ ├── index.tsx # 主布局组件
│ ├── SettingsSidebar.tsx # 设置侧边栏
│ ├── SettingsSidebarBody.tsx # 侧边栏导航菜单
│ └── styles.ts # 布局样式
├── hooks/
│ └── useSettingsCategory.ts # 设置分类定义
├── features/
│ └── SettingHeader.tsx # 设置页头部组件
│
├── account/ # 账号设置组
│ ├── profile/ # 个人资料
│ │ └── index.tsx
│ └── stats/ # 数据统计
│ └── index.tsx
│
├── general/ # 通用设置组
│ ├── appearance/ # 外观设置
│ │ ├── index.tsx
│ │ └── ThemeSelector.tsx
│ ├── chat-appearance/ # 聊天外观
│ │ └── index.tsx
│ └── hotkeys/ # 快捷键
│ └── index.tsx
│
├── agent/ # 智能体设置组
│ ├── providers/ # AI 服务商
│ │ └── index.tsx
│ ├── assistant/ # 助理服务
│ │ └── index.tsx
│ ├── skills/ # 技能/MCP
│ │ └── index.tsx
│ ├── memory/ # 记忆设置
│ │ └── index.tsx
│ ├── image-gen/ # 绘画服务
│ │ └── index.tsx
│ └── voice/ # 语音服务
│ └── index.tsx
│
├── system/ # 系统设置组
│ ├── proxy/ # 网络代理
│ │ └── index.tsx
│ ├── storage/ # 数据存储
│ │ └── index.tsx
│ ├── security/ # 安全设置
│ │ └── index.tsx
│ ├── external-tools/ # 外部工具
│ │ └── index.tsx
│ ├── experimental/ # 实验功能
│ │ └── index.tsx
│ ├── developer/ # 开发者
│ │ └── index.tsx
│ └── about/ # 关于
│ └── index.tsx
│
└── index.tsx # 导出入口
```
## 五、核心组件设计
### 5.1 设置分类 Hook
```typescript
// src/components/settings/hooks/useSettingsCategory.ts
import { useMemo } from 'react';
import { useTranslation } from 'react-i18next';
import {
User,
BarChart3,
Palette,
MessageSquare,
Keyboard,
Brain,
Bot,
Blocks,
BrainCircuit,
Image,
Mic,
Globe,
Database,
Shield,
Wrench,
FlaskConical,
Code,
Info,
LucideIcon,
} from 'lucide-react';
import { SettingsGroupKey, SettingsTabs } from '@/types/settings';
export interface CategoryItem {
key: SettingsTabs;
label: string;
icon: LucideIcon;
experimental?: boolean;
}
export interface CategoryGroup {
key: SettingsGroupKey;
title: string;
items: CategoryItem[];
}
export const useSettingsCategory = (): CategoryGroup[] => {
const { t } = useTranslation('settings');
return useMemo(() => {
const groups: CategoryGroup[] = [];
// 账号组
groups.push({
key: SettingsGroupKey.Account,
title: t('group.account'),
items: [
{ key: SettingsTabs.Profile, label: t('tab.profile'), icon: User },
{ key: SettingsTabs.Stats, label: t('tab.stats'), icon: BarChart3 },
],
});
// 通用组
groups.push({
key: SettingsGroupKey.General,
title: t('group.general'),
items: [
{ key: SettingsTabs.Appearance, label: t('tab.appearance'), icon: Palette },
{ key: SettingsTabs.ChatAppearance, label: t('tab.chatAppearance'), icon: MessageSquare },
{ key: SettingsTabs.Hotkeys, label: t('tab.hotkeys'), icon: Keyboard },
],
});
// 智能体组
groups.push({
key: SettingsGroupKey.Agent,
title: t('group.agent'),
items: [
{ key: SettingsTabs.Providers, label: t('tab.providers'), icon: Brain },
{ key: SettingsTabs.Assistant, label: t('tab.assistant'), icon: Bot },
{ key: SettingsTabs.Skills, label: t('tab.skills'), icon: Blocks },
{ key: SettingsTabs.Memory, label: t('tab.memory'), icon: BrainCircuit },
{ key: SettingsTabs.ImageGen, label: t('tab.imageGen'), icon: Image },
{ key: SettingsTabs.Voice, label: t('tab.voice'), icon: Mic },
],
});
// 系统组
groups.push({
key: SettingsGroupKey.System,
title: t('group.system'),
items: [
{ key: SettingsTabs.Proxy, label: t('tab.proxy'), icon: Globe },
{ key: SettingsTabs.Storage, label: t('tab.storage'), icon: Database },
{ key: SettingsTabs.Security, label: t('tab.security'), icon: Shield },
{ key: SettingsTabs.ExternalTools, label: t('tab.externalTools'), icon: Wrench },
{ key: SettingsTabs.Experimental, label: t('tab.experimental'), icon: FlaskConical, experimental: true },
{ key: SettingsTabs.Developer, label: t('tab.developer'), icon: Code },
{ key: SettingsTabs.About, label: t('tab.about'), icon: Info },
],
});
return groups;
}, [t]);
};
```
### 5.2 设置布局组件
```typescript
// src/components/settings/_layout/index.tsx
import { useState } from 'react';
import styled from 'styled-components';
import { SettingsSidebar } from './SettingsSidebar';
import { SettingsTabs } from '@/types/settings';
const LayoutContainer = styled.div`
display: flex;
height: 100%;
background: hsl(var(--background));
`;
const ContentContainer = styled.main`
flex: 1;
overflow-y: auto;
padding: 24px 32px;
`;
interface SettingsLayoutProps {
children?: React.ReactNode;
}
export function SettingsLayout({ children }: SettingsLayoutProps) {
const [activeTab, setActiveTab] = useState<SettingsTabs>(SettingsTabs.Profile);
return (
<LayoutContainer>
<SettingsSidebar activeTab={activeTab} onTabChange={setActiveTab} />
<ContentContainer>
{children}
</ContentContainer>
</LayoutContainer>
);
}
```
### 5.3 设置侧边栏组件
```typescript
// src/components/settings/_layout/SettingsSidebar.tsx
import styled from 'styled-components';
import { ChevronDown } from 'lucide-react';
import { useState } from 'react';
import { useSettingsCategory, CategoryGroup, CategoryItem } from '../hooks/useSettingsCategory';
import { SettingsTabs } from '@/types/settings';
const SidebarContainer = styled.aside`
width: 240px;
min-width: 240px;
height: 100%;
background: hsl(var(--card));
border-right: 1px solid hsl(var(--border));
overflow-y: auto;
padding: 16px 8px;
`;
const GroupContainer = styled.div`
margin-bottom: 8px;
`;
const GroupHeader = styled.button<{ $expanded: boolean }>`
display: flex;
align-items: center;
justify-content: space-between;
width: 100%;
padding: 8px 12px;
border: none;
background: transparent;
cursor: pointer;
font-size: 12px;
font-weight: 500;
color: hsl(var(--muted-foreground));
text-transform: uppercase;
letter-spacing: 0.5px;
svg {
width: 14px;
height: 14px;
transition: transform 0.2s;
transform: rotate(${({ $expanded }) => $expanded ? '0deg' : '-90deg'});
}
`;
const GroupItems = styled.div<{ $expanded: boolean }>`
display: ${({ $expanded }) => $expanded ? 'flex' : 'none'};
flex-direction: column;
gap: 2px;
padding: 4px 0;
`;
const NavItem = styled.button<{ $active: boolean }>`
display: flex;
align-items: center;
gap: 10px;
width: 100%;
padding: 10px 12px;
border: none;
border-radius: 8px;
background: ${({ $active }) => $active ? 'hsl(var(--accent))' : 'transparent'};
cursor: pointer;
font-size: 14px;
color: ${({ $active }) => $active ? 'hsl(var(--foreground))' : 'hsl(var(--muted-foreground))'};
transition: all 0.15s;
&:hover {
background: hsl(var(--accent));
color: hsl(var(--foreground));
}
svg {
width: 18px;
height: 18px;
}
`;
const ExperimentalBadge = styled.span`
font-size: 10px;
padding: 2px 6px;
background: hsl(var(--destructive) / 0.1);
color: hsl(var(--destructive));
border-radius: 4px;
margin-left: auto;
`;
interface SettingsSidebarProps {
activeTab: SettingsTabs;
onTabChange: (tab: SettingsTabs) => void;
}
export function SettingsSidebar({ activeTab, onTabChange }: SettingsSidebarProps) {
const categoryGroups = useSettingsCategory();
const [expandedGroups, setExpandedGroups] = useState<Record<string, boolean>>({
account: true,
general: true,
agent: true,
system: true,
});
const toggleGroup = (key: string) => {
setExpandedGroups(prev => ({
...prev,
[key]: !prev[key],
}));
};
return (
<SidebarContainer>
{categoryGroups.map((group) => (
<GroupContainer key={group.key}>
<GroupHeader
$expanded={expandedGroups[group.key] ?? true}
onClick={() => toggleGroup(group.key)}
>
{group.title}
<ChevronDown />
</GroupHeader>
<GroupItems $expanded={expandedGroups[group.key] ?? true}>
{group.items.map((item) => (
<NavItem
key={item.key}
$active={activeTab === item.key}
onClick={() => onTabChange(item.key)}
>
<item.icon />
{item.label}
{item.experimental && <ExperimentalBadge>实验</ExperimentalBadge>}
</NavItem>
))}
</GroupItems>
</GroupContainer>
))}
</SidebarContainer>
);
}
```
### 5.4 设置页头组件
```typescript
// src/components/settings/features/SettingHeader.tsx
import styled from 'styled-components';
import { ReactNode } from 'react';
const HeaderContainer = styled.div`
display: flex;
flex-direction: column;
gap: 16px;
margin-bottom: 24px;
`;
const TitleRow = styled.div`
display: flex;
align-items: center;
justify-content: space-between;
`;
const Title = styled.h1`
font-size: 24px;
font-weight: 600;
color: hsl(var(--foreground));
margin: 0;
`;
const Divider = styled.div`
height: 1px;
background: hsl(var(--border));
`;
interface SettingHeaderProps {
title: ReactNode;
extra?: ReactNode;
}
export function SettingHeader({ title, extra }: SettingHeaderProps) {
return (
<HeaderContainer>
<TitleRow>
<Title>{title}</Title>
{extra}
</TitleRow>
<Divider />
</HeaderContainer>
);
}
```
## 六、i18n 配置
```json
// src/i18n/locales/zh-CN/settings.json
{
"group": {
"account": "账号",
"general": "通用",
"agent": "智能体",
"system": "系统"
},
"tab": {
"profile": "个人资料",
"stats": "数据统计",
"appearance": "外观",
"chatAppearance": "聊天外观",
"hotkeys": "快捷键",
"providers": "AI 服务商",
"assistant": "助理服务",
"skills": "技能管理",
"memory": "记忆设置",
"imageGen": "绘画服务",
"voice": "语音服务",
"proxy": "网络代理",
"storage": "数据存储",
"security": "安全设置",
"externalTools": "外部工具",
"experimental": "实验功能",
"developer": "开发者",
"about": "关于"
}
}
```
## 七、迁移计划
### 阶段 1:基础架构(1-2 天)
1. 创建设置类型定义 (`src/types/settings.ts`)
2. 创建设置分类 Hook (`useSettingsCategory.ts`)
3. 创建布局组件 (`_layout/`)
4. 添加 i18n 配置
### 阶段 2:迁移现有组件(2-3 天)
1. 将 `GeneralSettings.tsx` 拆分为 `appearance/` 和 `chat-appearance/`
2. 将 `ProxySettings.tsx` 迁移到 `system/proxy/`
3. 将 `DirectorySettings.tsx` 迁移到 `system/storage/`
4. 将 `TlsSettings.tsx` + `RemoteManagementSettings.tsx` 合并到 `system/security/`
5. 将 `ExternalToolsSettings.tsx` 迁移到 `system/external-tools/`
6. 将其他设置组件按分类迁移
### 阶段 3:新增功能(2-3 天)
1. 添加 `account/profile/` - 用户个人资料
2. 添加 `account/stats/` - 使用统计
3. 添加 `general/hotkeys/` - 快捷键设置
4. 添加 `agent/memory/` - 记忆管理
5. 添加 `agent/voice/` - 语音服务配置
### 阶段 4:集成 & 测试(1 天)
1. 更新 `App.tsx` 路由
2. 更新 `AppSidebar.tsx` 导航
3. 端到端测试
## 八、与 LobeHub 的对应关系
| LobeHub 设置项 | ProxyCast 对应 | 备注 |
|---------------|---------------|------|
| Profile | account/profile | 用户资料 |
| Stats | account/stats | 使用统计 |
| Common (外观) | general/appearance | 主题、语言等 |
| Chat Appearance | general/chat-appearance | 聊天气泡样式 |
| Hotkey | general/hotkeys | 快捷键配置 |
| Provider | agent/providers | AI 服务商配置 |
| Agent | agent/assistant | 助理配置 |
| Skill | agent/skills | MCP/技能管理 |
| Memory | agent/memory | 记忆设置 |
| Image | agent/image-gen | 绘画服务 |
| TTS | agent/voice | 语音服务 |
| Proxy | system/proxy | 网络代理 |
| Storage | system/storage | 数据存储 |
| About | system/about | 关于页面 |
## 九、UI 设计参考
### 颜色方案
- 使用 ProxyCast 现有的 CSS 变量(`hsl(var(--xxx))`)
- 侧边栏背景:`--card`
- 激活项背景:`--accent`
- 分组标题:`--muted-foreground`
### 间距规范
- 侧边栏宽度:240px
- 内容区左右 padding:32px
- 组间距:8px
- 项间距:2px
- 项内 padding:10px 12px
### 动画效果
- 分组展开/收起:0.2s ease
- 悬停效果:0.15s ease
- 页面切换:无动画(保持简洁)
---
**设计完成时间**: 2026-02-09
**预计开发时间**: 5-8 天
**参考项目**: LobeHub (lobehub/lobe-chat)
+2 -2
View File
@@ -111,7 +111,7 @@
"fast-check": "^4.4.0",
"globals": "^15.12.0",
"husky": "^9.1.7",
"jsdom": "^27.3.0",
"jsdom": "^22.1.0",
"postcss": "^8.4.47",
"prettier": "^3.3.3",
"tailwindcss": "^3.4.14",
@@ -119,6 +119,6 @@
"typescript": "^5.6.3",
"vite": "^5.4.21",
"vite-plugin-svgr": "^4.5.0",
"vitest": "^4.0.16"
"vitest": "^3.2.4"
}
}
+28 -56
View File
@@ -255,7 +255,7 @@ dependencies = [
"rand 0.8.5",
"regex",
"reqwest 0.12.28",
"rmcp 0.12.0",
"rmcp",
"schemars 1.2.1",
"scraper",
"serde",
@@ -6563,20 +6563,6 @@ dependencies = [
"unicode-ident",
]
[[package]]
name = "process-wrap"
version = "8.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a3ef4f2f0422f23a82ec9f628ea2acd12871c81a9362b02c43c1aa86acfc3ba1"
dependencies = [
"futures",
"indexmap 2.13.0",
"nix 0.30.1",
"tokio",
"tracing",
"windows 0.61.3",
]
[[package]]
name = "process-wrap"
version = "9.0.1"
@@ -6677,6 +6663,7 @@ dependencies = [
"proxycast-mcp",
"proxycast-processor",
"proxycast-providers",
"proxycast-scheduler",
"proxycast-server",
"proxycast-server-utils",
"proxycast-services",
@@ -6686,7 +6673,7 @@ dependencies = [
"rand 0.8.5",
"regex",
"reqwest 0.12.28",
"rmcp 0.6.4",
"rmcp",
"rusqlite",
"rustls-pemfile 2.2.0",
"scopeguard",
@@ -6740,10 +6727,11 @@ dependencies = [
"proxycast-mcp",
"proxycast-providers",
"proxycast-services",
"rmcp 0.6.4",
"rmcp",
"serde",
"serde_json",
"tempfile",
"thiserror 1.0.69",
"tokio",
"tokio-util",
"tracing",
@@ -6850,7 +6838,7 @@ dependencies = [
"async-trait",
"glob",
"proxycast-core",
"rmcp 0.6.4",
"rmcp",
"serde",
"serde_json",
"thiserror 1.0.69",
@@ -6911,6 +6899,24 @@ dependencies = [
"uuid",
]
[[package]]
name = "proxycast-scheduler"
version = "0.61.0"
dependencies = [
"anyhow",
"async-trait",
"chrono",
"proxycast-agent",
"proxycast-core",
"rusqlite",
"serde",
"serde_json",
"thiserror 1.0.69",
"tokio",
"tracing",
"uuid",
]
[[package]]
name = "proxycast-server"
version = "0.61.0"
@@ -7050,10 +7056,12 @@ name = "proxycast-websocket"
version = "0.61.0"
dependencies = [
"axum 0.7.9",
"chrono",
"dashmap 5.5.3",
"futures",
"parking_lot",
"proptest",
"proxycast-agent",
"proxycast-core",
"serde",
"serde_json",
@@ -7606,29 +7614,6 @@ dependencies = [
"windows-sys 0.52.0",
]
[[package]]
name = "rmcp"
version = "0.6.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "41ab0892f4938752b34ae47cb53910b1b0921e55e77ddb6e44df666cab17939f"
dependencies = [
"base64 0.22.1",
"chrono",
"futures",
"paste",
"pin-project-lite",
"process-wrap 8.2.1",
"rmcp-macros 0.6.4",
"schemars 1.2.1",
"serde",
"serde_json",
"thiserror 2.0.18",
"tokio",
"tokio-stream",
"tokio-util",
"tracing",
]
[[package]]
name = "rmcp"
version = "0.12.0"
@@ -7643,9 +7628,9 @@ dependencies = [
"oauth2",
"pastey",
"pin-project-lite",
"process-wrap 9.0.1",
"process-wrap",
"reqwest 0.12.28",
"rmcp-macros 0.12.0",
"rmcp-macros",
"schemars 1.2.1",
"serde",
"serde_json",
@@ -7658,19 +7643,6 @@ dependencies = [
"url",
]
[[package]]
name = "rmcp-macros"
version = "0.6.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1827cd98dab34cade0513243c6fe0351f0f0b2c9d6825460bcf45b42804bdda0"
dependencies = [
"darling 0.21.3",
"proc-macro2",
"quote",
"serde_json",
"syn 2.0.114",
]
[[package]]
name = "rmcp-macros"
version = "0.12.0"
+4 -1
View File
@@ -25,6 +25,7 @@ proxycast-server = { path = "crates/server" }
proxycast-skills = { path = "crates/skills" }
proxycast-mcp = { path = "crates/mcp" }
proxycast-agent = { path = "crates/agent" }
proxycast-scheduler = { path = "crates/scheduler" }
voice-core = { path = "crates/voice-core" }
# 序列化
@@ -120,7 +121,8 @@ enigo = "0.3"
aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.11.0" }
# MCP (Model Context Protocol)
rmcp = { version = "0.6", features = ["client", "transport-io", "transport-child-process"] }
rmcp = { version = "0.12.0", features = ["client", "transport-io", "transport-child-process"] }
# Tauri
@@ -209,6 +211,7 @@ proxycast-server.workspace = true
proxycast-skills.workspace = true
proxycast-mcp.workspace = true
proxycast-agent.workspace = true
proxycast-scheduler.workspace = true
voice-core.workspace = true
# Tauri
+1
View File
@@ -21,6 +21,7 @@ tracing.workspace = true
chrono.workspace = true
dirs.workspace = true
uuid.workspace = true
thiserror.workspace = true
[dev-dependencies]
tempfile.workspace = true
+42
View File
@@ -466,6 +466,48 @@ impl AsterAgentState {
crate::create_session_config_with_project(db, session_id, project_id)
}
/// 注册 MCP 桥接客户端
///
/// 将 ProxyCast 托管的 MCP 客户端注册到 Aster Agent 的 ExtensionManager,
/// 使 Agent 能够直接调用该 MCP 服务器提供的工具。
///
/// # 参数
/// - `name`: 客户端名称
/// - `description`: 描述
/// - `client`: 实现 McpClientTrait 的客户端,必须包装在 Arc<Mutex<Box<...>>> 中
/// - `server_info`: MCP 服务器信息
pub async fn register_mcp_bridge(
&self,
name: String,
description: String,
client: Arc<tokio::sync::Mutex<Box<dyn aster::agents::mcp_client::McpClientTrait>>>,
server_info: Option<rmcp::model::ServerInfo>,
) -> Result<(), String> {
let agent_guard = self.agent.read().await;
if let Some(agent) = agent_guard.as_ref() {
// 创建 Extension 配置
let config = aster::agents::extension::ExtensionConfig::Builtin {
name: name.clone(),
display_name: Some(name.clone()),
description,
timeout: None,
bundled: Some(false),
available_tools: Vec::new(),
};
// 注册到 ExtensionManager
agent
.extension_manager
.add_client(name, config, client, server_info, None)
.await;
tracing::info!("[AsterAgent] MCP 桥接注册成功");
Ok(())
} else {
Err("Agent 未初始化".to_string())
}
}
/// 检查 Agent 是否已初始化
pub async fn is_initialized(&self) -> bool {
self.agent.read().await.is_some()
@@ -12,6 +12,7 @@
use aster::model::ModelConfig;
use aster::providers::base::Provider;
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
use proxycast_core::database::DbConnection;
use proxycast_core::models::provider_pool_model::{
CredentialData, PoolProviderType, ProviderCredential,
@@ -447,7 +448,15 @@ fn set_provider_env_vars(config: &AsterProviderConfig) {
}
}
"anthropic" => {
// Aster Anthropic Provider 读取 ANTHROPIC_HOST
std::env::set_var("ANTHROPIC_HOST", base_url);
// 兼容历史逻辑,保留旧变量
std::env::set_var("ANTHROPIC_BASE_URL", base_url);
tracing::info!(
"[CredentialBridge] 设置 ANTHROPIC_HOST={}, ANTHROPIC_BASE_URL={}",
base_url,
base_url
);
}
_ => {
// 其他 Provider 使用通用格式
@@ -487,6 +496,10 @@ pub fn map_pool_type_to_aster(pool_type: &PoolProviderType) -> &'static str {
///
/// 支持 60+ API Key Provider,包括 deepseek, moonshot, qwen 等
fn map_provider_type_to_aster(provider_type: &str) -> &'static str {
if let Ok(api_type) = provider_type.parse::<ApiProviderType>() {
return api_type.runtime_spec().aster_provider_name;
}
match provider_type {
// 标准 Provider
"openai" => "openai",
@@ -560,4 +573,26 @@ mod tests {
assert_eq!(host, "https://api.openai.com");
assert_eq!(path, "");
}
#[test]
fn test_set_provider_env_vars_anthropic_sets_host_and_base_url() {
let config = AsterProviderConfig {
provider_name: "anthropic".to_string(),
model_name: "glm-4.7".to_string(),
api_key: Some("test-key".to_string()),
base_url: Some("https://open.bigmodel.cn/api/anthropic".to_string()),
credential_uuid: "test-uuid".to_string(),
};
set_provider_env_vars(&config);
assert_eq!(
std::env::var("ANTHROPIC_HOST").ok().as_deref(),
Some("https://open.bigmodel.cn/api/anthropic")
);
assert_eq!(
std::env::var("ANTHROPIC_BASE_URL").ok().as_deref(),
Some("https://open.bigmodel.cn/api/anthropic")
);
}
}
+2
View File
@@ -11,6 +11,7 @@ pub mod mcp_bridge;
pub mod prompt;
pub mod session_store;
pub mod subagent_scheduler;
pub mod tools;
pub use aster_state::{AsterAgentState, ProviderConfig};
pub use aster_state_support::{
@@ -28,3 +29,4 @@ pub use session_store::{
pub use subagent_scheduler::{
ProxyCastScheduler, ProxyCastSubAgentExecutor, SchedulerEventEmitter, SubAgentProgressEvent,
};
pub use tools::{BrowserAction, BrowserTool, BrowserToolError, BrowserToolResult};
+254 -6
View File
@@ -3,16 +3,28 @@
//! 实现 Aster 的 McpClientTrait,将工具调用转发到
//! ProxyCast 已有的 MCP RunningService,避免重复启动进程。
use rmcp::model::InitializeResult;
use rmcp::service::RunningService;
use rmcp::RoleClient;
use std::sync::Arc;
use aster::agents::mcp_client::{Error as McpError, McpClientTrait};
use aster::session_context::{current_session_id, SESSION_ID_HEADER};
use proxycast_mcp::client::ProxyCastMcpClient;
use rmcp::model::{
CallToolRequest, CallToolRequestParam, CallToolResult, CancelledNotification,
CancelledNotificationMethod, CancelledNotificationParam, ClientRequest, GetPromptRequest,
GetPromptRequestParam, GetPromptResult, InitializeResult, JsonObject, ListPromptsRequest,
ListPromptsResult, ListResourcesRequest, ListResourcesResult, ListToolsRequest,
ListToolsResult, Meta, PaginatedRequestParam, ReadResourceRequest, ReadResourceRequestParam,
ReadResourceResult, ServerNotification, ServerResult,
};
use rmcp::service::{PeerRequestOptions, RunningService, ServiceError};
use rmcp::RoleClient;
use serde_json::Value;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
/// MCP 桥接客户端
///
/// 持有 ProxyCast 的 RunningService 引用,
/// 持有 ProxyCast 的 rmcp RunningService 引用,
/// 将 Aster 的工具调用转发到已有的 MCP 连接。
#[allow(dead_code)]
pub struct McpBridgeClient {
@@ -20,20 +32,256 @@ pub struct McpBridgeClient {
name: String,
/// ProxyCast 的 rmcp RunningService
service: Arc<RunningService<RoleClient, ProxyCastMcpClient>>,
/// ProxyCast MCP 客户端处理器
handler: Arc<ProxyCastMcpClient>,
/// 服务器初始化信息
server_info: Option<InitializeResult>,
/// 请求超时时间
timeout: Duration,
}
impl McpBridgeClient {
pub fn new(
name: String,
service: Arc<RunningService<RoleClient, ProxyCastMcpClient>>,
handler: Arc<ProxyCastMcpClient>,
server_info: Option<InitializeResult>,
) -> Self {
Self {
name,
service,
handler,
server_info,
timeout: Duration::from_secs(60), // 默认超时 60s
}
}
/// 发送请求并处理取消和超时
async fn send_request(
&self,
request: ClientRequest,
cancel_token: CancellationToken,
) -> Result<ServerResult, McpError> {
// 发送请求
let handle = self
.service
.send_cancellable_request(request, PeerRequestOptions::no_options())
.await?;
let request_id = handle.id;
let peer = handle.peer.clone();
// 等待响应,同时处理超时和取消
tokio::select! {
result = handle.rx => {
result.map_err(|_e| ServiceError::TransportClosed)?
}
_ = tokio::time::sleep(self.timeout) => {
// 超时,发送取消通知
let _ = peer.send_notification(
CancelledNotification {
params: CancelledNotificationParam {
request_id,
reason: Some("timed out".to_owned()),
},
method: CancelledNotificationMethod,
extensions: Default::default(),
}
.into(),
).await;
Err(ServiceError::Timeout{timeout: self.timeout})
}
_ = cancel_token.cancelled() => {
// 取消,发送取消通知
let _ = peer.send_notification(
CancelledNotification {
params: CancelledNotificationParam {
request_id,
reason: Some("operation cancelled".to_owned()),
},
method: CancelledNotificationMethod,
extensions: Default::default(),
}
.into(),
).await;
Err(ServiceError::Cancelled { reason: None })
}
}
}
/// 注入 Session ID 到扩展字段
fn inject_session(&self, mut extensions: rmcp::model::Extensions) -> rmcp::model::Extensions {
if let Some(session_id) = current_session_id() {
let mut meta_map = extensions
.get::<Meta>()
.map(|meta| meta.0.clone())
.unwrap_or_default();
// 移除旧的 ID (大小写不敏感)
meta_map.retain(|k, _| !k.eq_ignore_ascii_case(SESSION_ID_HEADER));
// 插入新的 ID
meta_map.insert(SESSION_ID_HEADER.to_string(), Value::String(session_id));
extensions.insert(Meta(meta_map));
}
extensions
}
}
#[async_trait::async_trait]
impl McpClientTrait for McpBridgeClient {
async fn list_resources(
&self,
cursor: Option<String>,
cancel_token: CancellationToken,
) -> Result<ListResourcesResult, McpError> {
let res = self
.send_request(
ClientRequest::ListResourcesRequest(ListResourcesRequest {
params: Some(PaginatedRequestParam { cursor }),
method: Default::default(),
extensions: self.inject_session(Default::default()),
}),
cancel_token,
)
.await?;
match res {
ServerResult::ListResourcesResult(result) => Ok(result),
_ => Err(ServiceError::UnexpectedResponse),
}
}
async fn read_resource(
&self,
uri: &str,
cancel_token: CancellationToken,
) -> Result<ReadResourceResult, McpError> {
let res = self
.send_request(
ClientRequest::ReadResourceRequest(ReadResourceRequest {
params: ReadResourceRequestParam {
uri: uri.to_string(),
},
method: Default::default(),
extensions: self.inject_session(Default::default()),
}),
cancel_token,
)
.await?;
match res {
ServerResult::ReadResourceResult(result) => Ok(result),
_ => Err(ServiceError::UnexpectedResponse),
}
}
async fn list_tools(
&self,
cursor: Option<String>,
cancel_token: CancellationToken,
) -> Result<ListToolsResult, McpError> {
let res = self
.send_request(
ClientRequest::ListToolsRequest(ListToolsRequest {
params: Some(PaginatedRequestParam { cursor }),
method: Default::default(),
extensions: self.inject_session(Default::default()),
}),
cancel_token,
)
.await?;
match res {
ServerResult::ListToolsResult(result) => Ok(result),
_ => Err(ServiceError::UnexpectedResponse),
}
}
async fn call_tool(
&self,
name: &str,
arguments: Option<JsonObject>,
cancel_token: CancellationToken,
) -> Result<CallToolResult, McpError> {
let res = self
.send_request(
ClientRequest::CallToolRequest(CallToolRequest {
params: CallToolRequestParam {
name: name.to_string().into(),
arguments,
},
method: Default::default(),
extensions: self.inject_session(Default::default()),
}),
cancel_token,
)
.await?;
match res {
ServerResult::CallToolResult(result) => Ok(result),
_ => Err(ServiceError::UnexpectedResponse),
}
}
async fn list_prompts(
&self,
cursor: Option<String>,
cancel_token: CancellationToken,
) -> Result<ListPromptsResult, McpError> {
let res = self
.send_request(
ClientRequest::ListPromptsRequest(ListPromptsRequest {
params: Some(PaginatedRequestParam { cursor }),
method: Default::default(),
extensions: self.inject_session(Default::default()),
}),
cancel_token,
)
.await?;
match res {
ServerResult::ListPromptsResult(result) => Ok(result),
_ => Err(ServiceError::UnexpectedResponse),
}
}
async fn get_prompt(
&self,
name: &str,
arguments: Value,
cancel_token: CancellationToken,
) -> Result<GetPromptResult, McpError> {
let arguments = match arguments {
Value::Object(map) => Some(map),
_ => None,
};
let res = self
.send_request(
ClientRequest::GetPromptRequest(GetPromptRequest {
params: GetPromptRequestParam {
name: name.to_string(),
arguments,
},
method: Default::default(),
extensions: self.inject_session(Default::default()),
}),
cancel_token,
)
.await?;
match res {
ServerResult::GetPromptResult(result) => Ok(result),
_ => Err(ServiceError::UnexpectedResponse),
}
}
async fn subscribe(&self) -> mpsc::Receiver<ServerNotification> {
self.handler.subscribe().await
}
fn get_info(&self) -> Option<&InitializeResult> {
self.server_info.as_ref()
}
}
@@ -0,0 +1,253 @@
//! Browser Tool 包装器
//!
//! 将 Playwright MCP Server 的工具映射到 Aster 工具系统
//! 提供浏览器自动化能力,包括导航、快照、点击、输入等操作
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::sync::Arc;
use thiserror::Error;
use tokio::sync::Mutex;
/// Browser Tool 错误类型
#[derive(Debug, Error)]
pub enum BrowserToolError {
#[error("MCP 客户端未初始化")]
ClientNotInitialized,
#[error("工具调用失败: {0}")]
ToolCallFailed(String),
#[error("参数序列化失败: {0}")]
SerializationError(String),
#[error("MCP 错误: {0}")]
McpError(String),
}
/// Browser Tool 动作类型
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum BrowserAction {
/// 导航到 URL
Navigate { url: String },
/// 获取页面快照
Snapshot,
/// 点击元素
Click { ref_id: String },
/// 输入文本
Type { ref_id: String, text: String },
/// 截图
Screenshot { filename: Option<String> },
}
/// Browser Tool 结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BrowserToolResult {
/// 是否成功
pub success: bool,
/// 输出内容
pub output: String,
/// 错误信息
pub error: Option<String>,
}
/// Browser Tool 包装器
///
/// 提供对 Playwright MCP Server 工具的高级封装
pub struct BrowserTool {
/// MCP 客户端
mcp_client: Arc<Mutex<Option<Box<dyn aster::agents::mcp_client::McpClientTrait>>>>,
}
impl BrowserTool {
/// 创建新的 Browser Tool 实例
pub fn new() -> Self {
Self {
mcp_client: Arc::new(Mutex::new(None)),
}
}
/// 设置 MCP 客户端
pub async fn set_mcp_client(
&self,
client: Box<dyn aster::agents::mcp_client::McpClientTrait>,
) {
let mut guard = self.mcp_client.lock().await;
*guard = Some(client);
}
/// 执行浏览器动作
///
/// # 参数
/// - `action`: 浏览器动作
///
/// # 返回
/// 返回工具执行结果
pub async fn execute(
&self,
action: BrowserAction,
) -> Result<BrowserToolResult, BrowserToolError> {
let client_guard = self.mcp_client.lock().await;
let client = client_guard
.as_ref()
.ok_or(BrowserToolError::ClientNotInitialized)?;
// 根据动作类型调用对应的 MCP 工具
match action {
BrowserAction::Navigate { url } => {
self.call_mcp_tool(client, "browser_navigate", serde_json::json!({ "url": url }))
.await
}
BrowserAction::Snapshot => {
self.call_mcp_tool(client, "browser_snapshot", serde_json::json!({}))
.await
}
BrowserAction::Click { ref_id } => {
self.call_mcp_tool(client, "browser_click", serde_json::json!({ "ref": ref_id }))
.await
}
BrowserAction::Type { ref_id, text } => {
self.call_mcp_tool(
client,
"browser_type",
serde_json::json!({ "ref": ref_id, "text": text }),
)
.await
}
BrowserAction::Screenshot { filename } => {
let mut args = serde_json::json!({ "type": "png" });
if let Some(name) = filename {
args["filename"] = Value::String(name);
}
self.call_mcp_tool(client, "browser_take_screenshot", args)
.await
}
}
}
/// 调用 MCP 工具
async fn call_mcp_tool(
&self,
client: &Box<dyn aster::agents::mcp_client::McpClientTrait>,
tool_name: &str,
arguments: Value,
) -> Result<BrowserToolResult, BrowserToolError> {
// 将 Value 转换为 JsonObject
let args = match arguments {
Value::Object(map) => Some(map),
_ => None,
};
// 创建取消令牌
let cancel_token = tokio_util::sync::CancellationToken::new();
// 调用工具
let result = client
.call_tool(tool_name, args, cancel_token)
.await
.map_err(|e| BrowserToolError::McpError(format!("{:?}", e)))?;
// 转换结果
let is_error = result.is_error.unwrap_or(false);
let output = result
.content
.into_iter()
.map(|c| match c.raw {
rmcp::model::RawContent::Text(text) => text.text,
rmcp::model::RawContent::Image(img) => {
format!("[Image: {}]", img.mime_type)
}
_ => "[Unknown content]".to_string(),
})
.collect::<Vec<_>>()
.join("\n");
Ok(BrowserToolResult {
success: !is_error,
output: output.clone(),
error: if is_error { Some(output) } else { None },
})
}
/// 导航到 URL
pub async fn navigate(&self, url: &str) -> Result<BrowserToolResult, BrowserToolError> {
self.execute(BrowserAction::Navigate {
url: url.to_string(),
})
.await
}
/// 获取页面快照
pub async fn snapshot(&self) -> Result<BrowserToolResult, BrowserToolError> {
self.execute(BrowserAction::Snapshot).await
}
/// 点击元素
pub async fn click(&self, ref_id: &str) -> Result<BrowserToolResult, BrowserToolError> {
self.execute(BrowserAction::Click {
ref_id: ref_id.to_string(),
})
.await
}
/// 输入文本
pub async fn type_text(
&self,
ref_id: &str,
text: &str,
) -> Result<BrowserToolResult, BrowserToolError> {
self.execute(BrowserAction::Type {
ref_id: ref_id.to_string(),
text: text.to_string(),
})
.await
}
/// 截图
pub async fn screenshot(
&self,
filename: Option<String>,
) -> Result<BrowserToolResult, BrowserToolError> {
self.execute(BrowserAction::Screenshot { filename }).await
}
}
impl Default for BrowserTool {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_browser_tool_creation() {
let tool = BrowserTool::new();
assert!(tool.mcp_client.lock().await.is_none());
}
#[test]
fn test_browser_action_serialization() {
let action = BrowserAction::Navigate {
url: "https://example.com".to_string(),
};
let json = serde_json::to_string(&action).unwrap();
assert!(json.contains("navigate"));
assert!(json.contains("https://example.com"));
}
#[test]
fn test_browser_tool_result() {
let result = BrowserToolResult {
success: true,
output: "Page loaded".to_string(),
error: None,
};
assert!(result.success);
assert_eq!(result.output, "Page loaded");
assert!(result.error.is_none());
}
}
+7
View File
@@ -0,0 +1,7 @@
//! Tools 模块
//!
//! 提供各种工具的包装器和辅助函数
pub mod browser_tool;
pub use browser_tool::{BrowserAction, BrowserTool, BrowserToolError, BrowserToolResult};
+34 -4
View File
@@ -4,6 +4,7 @@
//! 参考 aster 项目的 Conversation 设计,支持连续对话和工具调用
use serde::{Deserialize, Serialize};
use crate::models::provider_type::is_custom_provider_id;
/// Provider 类型枚举
///
@@ -54,12 +55,13 @@ impl ProviderType {
/// 从 provider 字符串和模型名称推断 provider 类型
///
/// 对于自定义 Provider ID(如 custom-xxx),尝试从模型名称推断协议类型
/// 对于自定义 Provider ID(如 custom-xxx),不做协议推断。
/// 真实协议应由 API Key Provider 的 `type`(即 ApiProviderType)在运行时决定。
pub fn from_provider_and_model(provider: &str, model: &str) -> Self {
// 首先检查是否是自定义 Provider ID(以 custom- 开头)
if provider.starts_with("custom-") {
// 自定义 Provider 使用 Anthropic 兼容协议(Anthropic Compatible)
return Self::AnthropicCompatible;
if is_custom_provider_id(provider) {
// 不对 custom-* 做协议猜测,避免与 DB 中真实 Provider 类型不一致
return Self::OpenAI;
}
// 对于其他 Provider,尝试直接解析
@@ -119,6 +121,34 @@ impl ProviderType {
}
}
#[cfg(test)]
mod tests {
use super::ProviderType;
#[test]
fn test_custom_provider_does_not_force_anthropic_protocol() {
assert_eq!(
ProviderType::from_provider_and_model(
"custom-ba4e7574-dd00-4784-945a-0f383dfa1272",
"claude-sonnet-4-5"
),
ProviderType::OpenAI
);
}
#[test]
fn test_unknown_provider_still_supports_model_based_inference() {
assert_eq!(
ProviderType::from_provider_and_model("unknown-provider", "claude-3-7-sonnet"),
ProviderType::Claude
);
assert_eq!(
ProviderType::from_provider_and_model("unknown-provider", "gemini-2.5-pro"),
ProviderType::Gemini
);
}
}
/// Agent 会话状态
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentSession {
@@ -33,6 +33,118 @@ pub enum ApiProviderType {
Gateway,
}
/// Provider 协议族
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderProtocolFamily {
OpenAiCompatible,
Anthropic,
Gemini,
AzureOpenai,
Vertexai,
AwsBedrock,
Ollama,
Codex,
}
/// Provider 运行时规范
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProviderRuntimeSpec {
pub protocol_family: ProviderProtocolFamily,
pub default_api_host: &'static str,
pub auth_header: &'static str,
pub auth_prefix: Option<&'static str>,
pub extra_headers: &'static [(&'static str, &'static str)],
pub aster_provider_name: &'static str,
}
const NO_EXTRA_HEADERS: [(&str, &str); 0] = [];
const ANTHROPIC_EXTRA_HEADERS: [(&str, &str); 1] = [("anthropic-version", "2023-06-01")];
impl ApiProviderType {
/// 统一的 Provider 运行时规范
pub const fn runtime_spec(&self) -> ProviderRuntimeSpec {
match self {
ApiProviderType::Anthropic | ApiProviderType::AnthropicCompatible => {
ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::Anthropic,
default_api_host: "https://api.anthropic.com",
auth_header: "x-api-key",
auth_prefix: None,
extra_headers: &ANTHROPIC_EXTRA_HEADERS,
aster_provider_name: "anthropic",
}
}
ApiProviderType::Gemini => ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::Gemini,
default_api_host: "https://generativelanguage.googleapis.com",
auth_header: "x-goog-api-key",
auth_prefix: None,
extra_headers: &NO_EXTRA_HEADERS,
aster_provider_name: "google",
},
ApiProviderType::AzureOpenai => ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::AzureOpenai,
default_api_host: "https://api.openai.com",
auth_header: "api-key",
auth_prefix: None,
extra_headers: &NO_EXTRA_HEADERS,
aster_provider_name: "azure",
},
ApiProviderType::Vertexai => ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::Vertexai,
default_api_host: "https://api.openai.com",
auth_header: "Authorization",
auth_prefix: Some("Bearer"),
extra_headers: &NO_EXTRA_HEADERS,
aster_provider_name: "gcpvertexai",
},
ApiProviderType::AwsBedrock => ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::AwsBedrock,
default_api_host: "https://api.openai.com",
auth_header: "Authorization",
auth_prefix: Some("Bearer"),
extra_headers: &NO_EXTRA_HEADERS,
aster_provider_name: "bedrock",
},
ApiProviderType::Ollama => ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::Ollama,
default_api_host: "http://localhost:11434",
auth_header: "Authorization",
auth_prefix: Some("Bearer"),
extra_headers: &NO_EXTRA_HEADERS,
aster_provider_name: "ollama",
},
ApiProviderType::Codex => ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::Codex,
default_api_host: "https://api.openai.com",
auth_header: "Authorization",
auth_prefix: Some("Bearer"),
extra_headers: &NO_EXTRA_HEADERS,
aster_provider_name: "codex",
},
ApiProviderType::Openai
| ApiProviderType::OpenaiResponse
| ApiProviderType::NewApi
| ApiProviderType::Gateway => ProviderRuntimeSpec {
protocol_family: ProviderProtocolFamily::OpenAiCompatible,
default_api_host: "https://api.openai.com",
auth_header: "Authorization",
auth_prefix: Some("Bearer"),
extra_headers: &NO_EXTRA_HEADERS,
aster_provider_name: "openai",
},
}
}
/// 是否属于 Anthropic 协议族
pub const fn is_anthropic_protocol(&self) -> bool {
matches!(
self.runtime_spec().protocol_family,
ProviderProtocolFamily::Anthropic
)
}
}
impl std::fmt::Display for ApiProviderType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
@@ -52,6 +164,179 @@ impl std::fmt::Display for ApiProviderType {
}
}
#[cfg(test)]
mod tests {
use super::{ApiProviderType, ProviderProtocolFamily};
#[test]
fn test_runtime_spec_anthropic_compatible() {
let spec = ApiProviderType::AnthropicCompatible.runtime_spec();
assert_eq!(spec.protocol_family, ProviderProtocolFamily::Anthropic);
assert_eq!(spec.auth_header, "x-api-key");
assert_eq!(spec.auth_prefix, None);
assert_eq!(spec.default_api_host, "https://api.anthropic.com");
assert_eq!(spec.extra_headers[0], ("anthropic-version", "2023-06-01"));
}
#[test]
fn test_runtime_spec_openai_defaults() {
let spec = ApiProviderType::Openai.runtime_spec();
assert_eq!(spec.protocol_family, ProviderProtocolFamily::OpenAiCompatible);
assert_eq!(spec.auth_header, "Authorization");
assert_eq!(spec.auth_prefix, Some("Bearer"));
assert_eq!(spec.default_api_host, "https://api.openai.com");
}
#[test]
fn test_runtime_spec_contract_matrix() {
let cases = [
(
ApiProviderType::Openai,
ProviderProtocolFamily::OpenAiCompatible,
"https://api.openai.com",
"Authorization",
Some("Bearer"),
"openai",
),
(
ApiProviderType::OpenaiResponse,
ProviderProtocolFamily::OpenAiCompatible,
"https://api.openai.com",
"Authorization",
Some("Bearer"),
"openai",
),
(
ApiProviderType::Codex,
ProviderProtocolFamily::Codex,
"https://api.openai.com",
"Authorization",
Some("Bearer"),
"codex",
),
(
ApiProviderType::Anthropic,
ProviderProtocolFamily::Anthropic,
"https://api.anthropic.com",
"x-api-key",
None,
"anthropic",
),
(
ApiProviderType::AnthropicCompatible,
ProviderProtocolFamily::Anthropic,
"https://api.anthropic.com",
"x-api-key",
None,
"anthropic",
),
(
ApiProviderType::Gemini,
ProviderProtocolFamily::Gemini,
"https://generativelanguage.googleapis.com",
"x-goog-api-key",
None,
"google",
),
(
ApiProviderType::AzureOpenai,
ProviderProtocolFamily::AzureOpenai,
"https://api.openai.com",
"api-key",
None,
"azure",
),
(
ApiProviderType::Vertexai,
ProviderProtocolFamily::Vertexai,
"https://api.openai.com",
"Authorization",
Some("Bearer"),
"gcpvertexai",
),
(
ApiProviderType::AwsBedrock,
ProviderProtocolFamily::AwsBedrock,
"https://api.openai.com",
"Authorization",
Some("Bearer"),
"bedrock",
),
(
ApiProviderType::Ollama,
ProviderProtocolFamily::Ollama,
"http://localhost:11434",
"Authorization",
Some("Bearer"),
"ollama",
),
(
ApiProviderType::NewApi,
ProviderProtocolFamily::OpenAiCompatible,
"https://api.openai.com",
"Authorization",
Some("Bearer"),
"openai",
),
(
ApiProviderType::Gateway,
ProviderProtocolFamily::OpenAiCompatible,
"https://api.openai.com",
"Authorization",
Some("Bearer"),
"openai",
),
];
for (
provider_type,
expected_family,
expected_host,
expected_auth_header,
expected_auth_prefix,
expected_aster_provider,
) in cases
{
let spec = provider_type.runtime_spec();
assert_eq!(
spec.protocol_family, expected_family,
"provider_type={provider_type:?}"
);
assert_eq!(
spec.default_api_host, expected_host,
"provider_type={provider_type:?}"
);
assert_eq!(
spec.auth_header, expected_auth_header,
"provider_type={provider_type:?}"
);
assert_eq!(
spec.auth_prefix, expected_auth_prefix,
"provider_type={provider_type:?}"
);
assert_eq!(
spec.aster_provider_name, expected_aster_provider,
"provider_type={provider_type:?}"
);
}
let anthropic_spec = ApiProviderType::Anthropic.runtime_spec();
assert_eq!(
anthropic_spec.extra_headers,
&[("anthropic-version", "2023-06-01")]
);
let anthropic_compatible_spec = ApiProviderType::AnthropicCompatible.runtime_spec();
assert_eq!(
anthropic_compatible_spec.extra_headers,
&[("anthropic-version", "2023-06-01")]
);
let openai_spec = ApiProviderType::Openai.runtime_spec();
assert!(openai_spec.extra_headers.is_empty());
}
}
impl std::str::FromStr for ApiProviderType {
type Err = String;
@@ -0,0 +1,160 @@
//! MCP 服务器默认数据迁移
//!
//! 添加默认的 Playwright MCP Server 配置到数据库
use chrono::Utc;
use rusqlite::{params, Connection};
use serde_json::json;
use uuid::Uuid;
/// 迁移设置键名
const MIGRATION_KEY_PLAYWRIGHT_SERVER: &str = "migrated_playwright_mcp_server_v1";
/// Playwright MCP Server 默认配置
const PLAYWRIGHT_SERVER_NAME: &str = "playwright";
const PLAYWRIGHT_SERVER_DESCRIPTION: &str = "Playwright 浏览器自动化工具";
/// 迁移结果
pub struct MigrationResult {
/// 是否执行了迁移
pub executed: bool,
/// 创建的服务器 ID
pub server_id: Option<String>,
}
/// 执行 Playwright MCP Server 迁移
///
/// 迁移步骤:
/// 1. 检查是否已迁移
/// 2. 检查是否已存在 playwright 服务器
/// 3. 如果不存在,创建默认配置
/// 4. 标记迁移完成
pub fn migrate_playwright_mcp_server(conn: &Connection) -> Result<MigrationResult, String> {
// 检查是否已经迁移过
if is_migration_completed(conn, MIGRATION_KEY_PLAYWRIGHT_SERVER) {
tracing::debug!("[迁移] Playwright MCP Server 已迁移过,跳过");
return Ok(MigrationResult {
executed: false,
server_id: None,
});
}
tracing::info!("[迁移] 开始执行 Playwright MCP Server 迁移");
// 检查是否已存在 playwright 服务器
if server_exists(conn, PLAYWRIGHT_SERVER_NAME) {
tracing::info!(
"[迁移] Playwright MCP Server 已存在,跳过创建并标记迁移完成"
);
mark_migration_completed(conn, MIGRATION_KEY_PLAYWRIGHT_SERVER)?;
return Ok(MigrationResult {
executed: false,
server_id: None,
});
}
// 开始事务
conn.execute("BEGIN TRANSACTION", [])
.map_err(|e| format!("开始事务失败: {e}"))?;
// 执行迁移
let result = execute_playwright_migration(conn);
match result {
Ok(server_id) => {
// 标记迁移完成
mark_migration_completed(conn, MIGRATION_KEY_PLAYWRIGHT_SERVER)?;
// 提交事务
conn.execute("COMMIT", [])
.map_err(|e| format!("提交事务失败: {e}"))?;
tracing::info!("[迁移] Playwright MCP Server 迁移完成: server_id={}", server_id);
Ok(MigrationResult {
executed: true,
server_id: Some(server_id),
})
}
Err(e) => {
// 回滚事务
let _ = conn.execute("ROLLBACK", []);
tracing::error!("[迁移] Playwright MCP Server 迁移失败,已回滚: {}", e);
Err(e)
}
}
}
/// 执行迁移的核心逻辑
fn execute_playwright_migration(conn: &Connection) -> Result<String, String> {
// 创建 Playwright MCP Server 配置
let server_id = Uuid::new_v4().to_string();
let created_at = Utc::now().to_rfc3339();
// 构建服务器配置
let server_config = json!({
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-playwright"],
"env": {
"HEADLESS": "true"
}
});
// 插入服务器配置
conn.execute(
"INSERT INTO mcp_servers (id, name, server_config, description,
enabled_proxycast, enabled_claude, enabled_codex,
enabled_gemini, created_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)",
params![
server_id,
PLAYWRIGHT_SERVER_NAME,
serde_json::to_string(&server_config).unwrap_or_default(),
PLAYWRIGHT_SERVER_DESCRIPTION,
1i32, // enabled_proxycast = true
0i32, // enabled_claude = false
0i32, // enabled_codex = false
0i32, // enabled_gemini = false
created_at,
],
)
.map_err(|e| format!("插入 Playwright MCP Server 失败: {e}"))?;
tracing::info!(
"[迁移] 已创建 Playwright MCP Server: id={}, name={}",
server_id,
PLAYWRIGHT_SERVER_NAME
);
Ok(server_id)
}
/// 检查服务器是否已存在
fn server_exists(conn: &Connection, name: &str) -> bool {
conn.query_row(
"SELECT COUNT(*) FROM mcp_servers WHERE name = ?1",
[name],
|row| row.get::<_, i32>(0),
)
.unwrap_or(0) > 0
}
/// 检查迁移是否已完成
fn is_migration_completed(conn: &Connection, key: &str) -> bool {
conn.query_row(
"SELECT value FROM settings WHERE key = ?1",
[key],
|row| row.get::<_, String>(0),
)
.is_ok()
}
/// 标记迁移已完成
fn mark_migration_completed(conn: &Connection, key: &str) -> Result<(), String> {
conn.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)",
[key, "1"],
)
.map_err(|e| format!("标记迁移完成失败: {e}"))?;
Ok(())
}
+15
View File
@@ -1,6 +1,7 @@
pub mod dao;
pub mod migration;
pub mod migration_v2;
pub mod migration_v3;
pub mod schema;
pub mod system_providers;
@@ -115,5 +116,19 @@ pub fn init_database() -> Result<DbConnection, String> {
}
}
// 执行 Playwright MCP Server 迁移
match migration_v3::migrate_playwright_mcp_server(&conn) {
Ok(result) => {
if result.executed {
if let Some(server_id) = result.server_id {
tracing::info!("[数据库] Playwright MCP Server 迁移完成: server_id={}", server_id);
}
}
}
Err(e) => {
tracing::warn!("[数据库] Playwright MCP Server 迁移失败(非致命): {}", e);
}
}
Ok(Arc::new(Mutex::new(conn)))
}
@@ -4,6 +4,11 @@
use serde::{Deserialize, Serialize};
/// 是否为自定义 Provider ID(`custom-*`)
pub fn is_custom_provider_id(provider_type: &str) -> bool {
provider_type.to_lowercase().starts_with("custom-")
}
/// Provider 类型枚举
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
@@ -91,8 +96,9 @@ impl std::str::FromStr for ProviderType {
"oneapi" | "one-api" | "newapi" | "new-api" => Ok(ProviderType::OpenAI),
"custom" | "custom_openai" => Ok(ProviderType::OpenAI),
// 自定义 Provider(UUID 格式,如 custom-ba4e7574-dd00-4784-945a-0f383dfa1272)
// 这些是用户通过 API Key Provider 添加的自定义服务
s if s.starts_with("custom-") => Ok(ProviderType::OpenAI),
// 注意:这里仅做“通道级”兼容映射(按 OpenAI 兜底),
// 实际协议应在运行时通过 API Key Provider.type 决定。
s if is_custom_provider_id(s) => Ok(ProviderType::OpenAI),
_ => Err(format!("Unknown provider: {s}")),
}
}
@@ -179,6 +185,18 @@ mod tests {
);
}
#[test]
fn test_is_custom_provider_id() {
assert!(is_custom_provider_id(
"custom-ba4e7574-dd00-4784-945a-0f383dfa1272"
));
assert!(is_custom_provider_id(
"CUSTOM-ba4e7574-dd00-4784-945a-0f383dfa1272"
));
assert!(!is_custom_provider_id("custom"));
assert!(!is_custom_provider_id("openai"));
}
#[test]
fn test_provider_type_display() {
assert_eq!(ProviderType::Kiro.to_string(), "kiro");
@@ -42,6 +42,17 @@ impl Default for OpenAICustomProvider {
}
impl OpenAICustomProvider {
fn maybe_log_protocol_mismatch_hint(url: &str, status: StatusCode) {
if (status == StatusCode::UNAUTHORIZED || status == StatusCode::FORBIDDEN)
&& url.contains("/api/anthropic")
{
eprintln!(
"[OPENAI_CUSTOM] 提示: URL '{}' 返回 {},疑似协议不匹配。若上游是 Anthropic 兼容网关,请改用 /v1/messages + x-api-key。",
url, status
);
}
}
pub fn new() -> Self {
Self::default()
}
@@ -222,6 +233,8 @@ impl OpenAICustomProvider {
.send()
.await?;
Self::maybe_log_protocol_mismatch_hint(url, resp.status());
if resp.status() != StatusCode::NOT_FOUND {
return Ok(resp);
}
@@ -258,6 +271,8 @@ impl OpenAICustomProvider {
.send()
.await?;
Self::maybe_log_protocol_mismatch_hint(&url, resp.status());
if resp.status() == StatusCode::NOT_FOUND {
if let Some(fallback_url) = self.build_url_fallback_without_v1("chat/completions") {
if fallback_url != url {
@@ -269,6 +284,7 @@ impl OpenAICustomProvider {
.json(request)
.send()
.await?;
Self::maybe_log_protocol_mismatch_hint(&fallback_url, resp2.status());
return Ok(resp2);
}
}
@@ -297,6 +313,7 @@ impl OpenAICustomProvider {
.header("Authorization", format!("Bearer {api_key}"))
.send()
.await?;
Self::maybe_log_protocol_mismatch_hint(&url, r.status());
if r.status() != StatusCode::NOT_FOUND {
resp = Some(r);
break;
+31
View File
@@ -0,0 +1,31 @@
[package]
name = "proxycast-scheduler"
version = "0.61.0"
edition = "2021"
[dependencies]
# 序列化
serde = { workspace = true, features = ["derive"] }
serde_json = { workspace = true }
# 异步运行时
tokio = { workspace = true }
async-trait = { workspace = true }
# 错误处理
anyhow = { workspace = true }
thiserror = { workspace = true }
# 日志
tracing = { workspace = true }
# 数据库
rusqlite = { workspace = true }
# 时间和 UUID
chrono = { workspace = true }
uuid = { workspace = true, features = ["v4", "serde"] }
# 项目内依赖
proxycast-core = { workspace = true }
proxycast-agent = { workspace = true }
+392
View File
@@ -0,0 +1,392 @@
//! 批量任务定义
//!
//! 定义批量任务相关的数据结构和状态
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use uuid::Uuid;
/// 批量任务选项
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BatchOptions {
/// 并发数量 (默认为 3)
#[serde(default = "default_concurrency")]
pub concurrency: usize,
/// 失败后是否继续 (默认为 true)
#[serde(default = "default_continue_on_error")]
pub continue_on_error: bool,
/// 重试次数 (默认为 0)
#[serde(default)]
pub retry_count: usize,
/// 任务超时时间(秒) (默认为 120)
#[serde(default = "default_timeout")]
pub timeout_seconds: u64,
}
fn default_concurrency() -> usize {
3
}
fn default_continue_on_error() -> bool {
true
}
fn default_timeout() -> u64 {
120
}
impl Default for BatchOptions {
fn default() -> Self {
Self {
concurrency: default_concurrency(),
continue_on_error: default_continue_on_error(),
retry_count: 0,
timeout_seconds: default_timeout(),
}
}
}
/// 单个任务定义
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskDefinition {
/// 任务 ID (可选,如果不提供则自动生成)
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<Uuid>,
/// 模板变量
pub variables: HashMap<String, String>,
/// 任务元数据 (用于追踪和识别)
#[serde(default)]
pub metadata: HashMap<String, String>,
}
/// 单个任务结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskResult {
/// 任务 ID
pub task_id: Uuid,
/// 任务状态
pub status: TaskStatus,
/// 响应内容 (如果成功)
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
/// 错误信息 (如果失败)
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
/// 使用 token 数
#[serde(default)]
pub usage: TokenUsage,
/// 开始时间
pub started_at: chrono::DateTime<chrono::Utc>,
/// 完成时间
pub completed_at: Option<chrono::DateTime<chrono::Utc>>,
}
/// 任务状态
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum TaskStatus {
/// 等待中
Pending,
/// 运行中
Running,
/// 已完成
Completed,
/// 失败
Failed,
/// 已取消
Cancelled,
}
/// Token 使用统计
#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize)]
pub struct TokenUsage {
/// 输入 token 数
#[serde(default)]
pub prompt_tokens: u32,
/// 输出 token 数
#[serde(default)]
pub completion_tokens: u32,
/// 总 token 数
#[serde(default)]
pub total_tokens: u32,
}
impl TokenUsage {
pub fn new(prompt_tokens: u32, completion_tokens: u32) -> Self {
Self {
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
}
}
}
/// 批量任务
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BatchTask {
/// 批量任务 ID
pub id: Uuid,
/// 批量任务名称
pub name: String,
/// 任务模板 ID
pub template_id: Uuid,
/// 任务列表
pub tasks: Vec<TaskDefinition>,
/// 批量任务选项
#[serde(default)]
pub options: BatchOptions,
/// 批量任务状态
pub status: BatchTaskStatus,
/// 任务结果
#[serde(default)]
pub results: Vec<TaskResult>,
/// 创建时间
pub created_at: chrono::DateTime<chrono::Utc>,
/// 开始时间
#[serde(skip_serializing_if = "Option::is_none")]
pub started_at: Option<chrono::DateTime<chrono::Utc>>,
/// 完成时间
#[serde(skip_serializing_if = "Option::is_none")]
pub completed_at: Option<chrono::DateTime<chrono::Utc>>,
}
/// 批量任务状态
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum BatchTaskStatus {
/// 等待中
Pending,
/// 运行中
Running,
/// 已完成
Completed,
/// 部分完成 (部分任务失败)
PartiallyCompleted,
/// 失败 (所有任务失败)
Failed,
/// 已取消
Cancelled,
}
impl BatchTask {
/// 创建新的批量任务
pub fn new(
name: String,
template_id: Uuid,
tasks: Vec<TaskDefinition>,
options: BatchOptions,
) -> Self {
let now = chrono::Utc::now();
Self {
id: Uuid::new_v4(),
name,
template_id,
tasks,
options,
status: BatchTaskStatus::Pending,
results: Vec::new(),
created_at: now,
started_at: None,
completed_at: None,
}
}
/// 获取进度信息
pub fn get_progress(&self) -> (usize, usize, usize) {
// (总数, 成功数, 失败数)
let total = self.tasks.len();
let completed = self
.results
.iter()
.filter(|r| r.status == TaskStatus::Completed)
.count();
let failed = self
.results
.iter()
.filter(|r| r.status == TaskStatus::Failed)
.count();
(total, completed, failed)
}
/// 获取统计信息
pub fn get_statistics(&self) -> BatchTaskStatistics {
let (total, completed, failed) = self.get_progress();
let running = self
.results
.iter()
.filter(|r| r.status == TaskStatus::Running)
.count();
let total_tokens: TokenUsage = self
.results
.iter()
.fold(TokenUsage::default(), |mut acc, r| {
acc.prompt_tokens += r.usage.prompt_tokens;
acc.completion_tokens += r.usage.completion_tokens;
acc.total_tokens += r.usage.total_tokens;
acc
});
BatchTaskStatistics {
total_tasks: total,
completed_tasks: completed,
failed_tasks: failed,
running_tasks: running,
pending_tasks: total - completed - failed - running,
total_tokens,
}
}
}
/// 批量任务统计信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BatchTaskStatistics {
/// 总任务数
pub total_tasks: usize,
/// 已完成任务数
pub completed_tasks: usize,
/// 失败任务数
pub failed_tasks: usize,
/// 运行中任务数
pub running_tasks: usize,
/// 等待中任务数
pub pending_tasks: usize,
/// 总 token 使用量
pub total_tokens: TokenUsage,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_batch_options_default() {
let options = BatchOptions::default();
assert_eq!(options.concurrency, 3);
assert_eq!(options.continue_on_error, true);
assert_eq!(options.retry_count, 0);
assert_eq!(options.timeout_seconds, 120);
}
#[test]
fn test_batch_task_creation() {
let tasks = vec![
TaskDefinition {
id: None,
variables: {
let mut map = HashMap::new();
map.insert("content".to_string(), "测试1".to_string());
map
},
metadata: HashMap::new(),
},
TaskDefinition {
id: None,
variables: {
let mut map = HashMap::new();
map.insert("content".to_string(), "测试2".to_string());
map
},
metadata: HashMap::new(),
},
];
let batch_task = BatchTask::new(
"测试批量任务".to_string(),
Uuid::new_v4(),
tasks,
BatchOptions::default(),
);
assert_eq!(batch_task.name, "测试批量任务");
assert_eq!(batch_task.tasks.len(), 2);
assert_eq!(batch_task.status, BatchTaskStatus::Pending);
}
#[test]
fn test_get_progress() {
let mut batch_task = BatchTask::new(
"测试".to_string(),
Uuid::new_v4(),
vec![
TaskDefinition {
id: Some(Uuid::new_v4()),
variables: HashMap::new(),
metadata: HashMap::new(),
},
TaskDefinition {
id: Some(Uuid::new_v4()),
variables: HashMap::new(),
metadata: HashMap::new(),
},
TaskDefinition {
id: Some(Uuid::new_v4()),
variables: HashMap::new(),
metadata: HashMap::new(),
},
],
BatchOptions::default(),
);
// 添加一些结果
batch_task.results.push(TaskResult {
task_id: batch_task.tasks[0].id.unwrap(),
status: TaskStatus::Completed,
content: Some("完成".to_string()),
error: None,
usage: TokenUsage::default(),
started_at: chrono::Utc::now(),
completed_at: Some(chrono::Utc::now()),
});
batch_task.results.push(TaskResult {
task_id: batch_task.tasks[1].id.unwrap(),
status: TaskStatus::Failed,
content: None,
error: Some("失败".to_string()),
usage: TokenUsage::default(),
started_at: chrono::Utc::now(),
completed_at: Some(chrono::Utc::now()),
});
let (total, completed, failed) = batch_task.get_progress();
assert_eq!(total, 3);
assert_eq!(completed, 1);
assert_eq!(failed, 1);
}
}
+489
View File
@@ -0,0 +1,489 @@
//! 批量任务数据访问对象 (DAO)
//!
//! 提供批量任务和模板的数据库操作
use super::batch::{BatchTask, BatchTaskStatus};
use super::template::TaskTemplate;
use anyhow::{Context, Result};
use proxycast_core::database::DbConnection;
use rusqlite::params;
use std::sync::{Arc, Mutex};
use uuid::Uuid;
/// 批量任务 DAO
pub struct BatchTaskDao;
impl BatchTaskDao {
/// 初始化数据库表
pub fn init_tables(db: &DbConnection) -> Result<()> {
let conn = db.lock().unwrap();
// 创建批量任务表
conn.execute(
"CREATE TABLE IF NOT EXISTS batch_tasks (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
template_id TEXT NOT NULL,
status TEXT NOT NULL,
options_json TEXT NOT NULL,
tasks_json TEXT NOT NULL,
results_json TEXT,
created_at TEXT NOT NULL,
started_at TEXT,
completed_at TEXT
)",
[],
)
.context("创建 batch_tasks 表失败")?;
// 创建模板表
conn.execute(
"CREATE TABLE IF NOT EXISTS batch_templates (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
description TEXT,
model TEXT NOT NULL,
system_prompt TEXT,
user_message_template TEXT NOT NULL,
temperature REAL,
max_tokens INTEGER,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
)",
[],
)
.context("创建 batch_templates 表失败")?;
// 创建索引
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_batch_tasks_status ON batch_tasks(status)",
[],
)?;
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_batch_tasks_created_at ON batch_tasks(created_at DESC)",
[],
)?;
Ok(())
}
/// 保存批量任务
pub fn save(db: &DbConnection, batch_task: &BatchTask) -> Result<()> {
let conn = db.lock().unwrap();
let options_json = serde_json::to_string(&batch_task.options)?;
let tasks_json = serde_json::to_string(&batch_task.tasks)?;
let results_json = if batch_task.results.is_empty() {
None
} else {
Some(serde_json::to_string(&batch_task.results)?)
};
conn.execute(
"INSERT OR REPLACE INTO batch_tasks
(id, name, template_id, status, options_json, tasks_json, results_json,
created_at, started_at, completed_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)",
params![
batch_task.id.to_string(),
batch_task.name,
batch_task.template_id.to_string(),
serde_json::to_string(&batch_task.status)?,
options_json,
tasks_json,
results_json,
batch_task.created_at.to_rfc3339(),
batch_task.started_at.map(|t| t.to_rfc3339()),
batch_task.completed_at.map(|t| t.to_rfc3339()),
],
)
.context("保存批量任务失败")?;
Ok(())
}
/// 根据 ID 查询批量任务
pub fn get_by_id(db: &DbConnection, id: &Uuid) -> Result<Option<BatchTask>> {
let conn = db.lock().unwrap();
let mut stmt = conn.prepare(
"SELECT id, name, template_id, status, options_json, tasks_json, results_json,
created_at, started_at, completed_at
FROM batch_tasks WHERE id = ?1",
)?;
let result = stmt
.query_row(params![id.to_string()], |row| {
let id: String = row.get(0)?;
let name: String = row.get(1)?;
let template_id: String = row.get(2)?;
let status_str: String = row.get(3)?;
let options_json: String = row.get(4)?;
let tasks_json: String = row.get(5)?;
let results_json: Option<String> = row.get(6)?;
let created_at: String = row.get(7)?;
let started_at: Option<String> = row.get(8)?;
let completed_at: Option<String> = row.get(9)?;
Ok((
id,
name,
template_id,
status_str,
options_json,
tasks_json,
results_json,
created_at,
started_at,
completed_at,
))
})
.optional()?;
if let Some((
id,
name,
template_id,
status_str,
options_json,
tasks_json,
results_json,
created_at,
started_at,
completed_at,
)) = result
{
let batch_task = BatchTask {
id: Uuid::parse_str(&id)?,
name,
template_id: Uuid::parse_str(&template_id)?,
status: serde_json::from_str(&status_str)?,
options: serde_json::from_str(&options_json)?,
tasks: serde_json::from_str(&tasks_json)?,
results: results_json
.map(|json| serde_json::from_str(&json))
.transpose()?
.unwrap_or_default(),
created_at: chrono::DateTime::parse_from_rfc3339(&created_at)?.into(),
started_at: started_at
.map(|s| chrono::DateTime::parse_from_rfc3339(&s).map(|dt| dt.into()))
.transpose()?,
completed_at: completed_at
.map(|s| chrono::DateTime::parse_from_rfc3339(&s).map(|dt| dt.into()))
.transpose()?,
};
Ok(Some(batch_task))
} else {
Ok(None)
}
}
/// 查询所有批量任务
pub fn list_all(db: &DbConnection, limit: usize) -> Result<Vec<BatchTask>> {
let conn = db.lock().unwrap();
let mut stmt = conn.prepare(
"SELECT id, name, template_id, status, options_json, tasks_json, results_json,
created_at, started_at, completed_at
FROM batch_tasks
ORDER BY created_at DESC
LIMIT ?1",
)?;
let rows = stmt.query_map(params![limit], |row| {
let id: String = row.get(0)?;
let name: String = row.get(1)?;
let template_id: String = row.get(2)?;
let status_str: String = row.get(3)?;
let options_json: String = row.get(4)?;
let tasks_json: String = row.get(5)?;
let results_json: Option<String> = row.get(6)?;
let created_at: String = row.get(7)?;
let started_at: Option<String> = row.get(8)?;
let completed_at: Option<String> = row.get(9)?;
Ok((
id,
name,
template_id,
status_str,
options_json,
tasks_json,
results_json,
created_at,
started_at,
completed_at,
))
})?;
let mut batch_tasks = Vec::new();
for row in rows {
let (
id,
name,
template_id,
status_str,
options_json,
tasks_json,
results_json,
created_at,
started_at,
completed_at,
) = row?;
let batch_task = BatchTask {
id: Uuid::parse_str(&id)?,
name,
template_id: Uuid::parse_str(&template_id)?,
status: serde_json::from_str(&status_str)?,
options: serde_json::from_str(&options_json)?,
tasks: serde_json::from_str(&tasks_json)?,
results: results_json
.map(|json| serde_json::from_str(&json))
.transpose()?
.unwrap_or_default(),
created_at: chrono::DateTime::parse_from_rfc3339(&created_at)?.into(),
started_at: started_at
.map(|s| chrono::DateTime::parse_from_rfc3339(&s).map(|dt| dt.into()))
.transpose()?,
completed_at: completed_at
.map(|s| chrono::DateTime::parse_from_rfc3339(&s).map(|dt| dt.into()))
.transpose()?,
};
batch_tasks.push(batch_task);
}
Ok(batch_tasks)
}
/// 删除批量任务
pub fn delete(db: &DbConnection, id: &Uuid) -> Result<bool> {
let conn = db.lock().unwrap();
let affected = conn.execute(
"DELETE FROM batch_tasks WHERE id = ?1",
params![id.to_string()],
)?;
Ok(affected > 0)
}
/// 更新批量任务状态
pub fn update_status(
db: &DbConnection,
id: &Uuid,
status: BatchTaskStatus,
) -> Result<()> {
let conn = db.lock().unwrap();
conn.execute(
"UPDATE batch_tasks SET status = ?1 WHERE id = ?2",
params![serde_json::to_string(&status)?, id.to_string()],
)?;
Ok(())
}
}
/// 模板 DAO
pub struct TemplateDao;
impl TemplateDao {
/// 保存模板
pub fn save(db: &DbConnection, template: &TaskTemplate) -> Result<()> {
let conn = db.lock().unwrap();
conn.execute(
"INSERT OR REPLACE INTO batch_templates
(id, name, description, model, system_prompt, user_message_template,
temperature, max_tokens, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)",
params![
template.id.to_string(),
template.name,
template.description,
template.model,
template.system_prompt,
template.user_message_template,
template.temperature,
template.max_tokens,
template.created_at.to_rfc3339(),
template.updated_at.to_rfc3339(),
],
)
.context("保存模板失败")?;
Ok(())
}
/// 根据 ID 查询模板
pub fn get_by_id(db: &DbConnection, id: &Uuid) -> Result<Option<TaskTemplate>> {
let conn = db.lock().unwrap();
let mut stmt = conn.prepare(
"SELECT id, name, description, model, system_prompt, user_message_template,
temperature, max_tokens, created_at, updated_at
FROM batch_templates WHERE id = ?1",
)?;
let result = stmt
.query_row(params![id.to_string()], |row| {
Ok(TaskTemplate {
id: Uuid::parse_str(&row.get::<_, String>(0)?).unwrap(),
name: row.get(1)?,
description: row.get(2)?,
model: row.get(3)?,
system_prompt: row.get(4)?,
user_message_template: row.get(5)?,
temperature: row.get(6)?,
max_tokens: row.get(7)?,
created_at: chrono::DateTime::parse_from_rfc3339(&row.get::<_, String>(8)?)
.unwrap()
.into(),
updated_at: chrono::DateTime::parse_from_rfc3339(&row.get::<_, String>(9)?)
.unwrap()
.into(),
})
})
.optional()?;
Ok(result)
}
/// 查询所有模板
pub fn list_all(db: &DbConnection) -> Result<Vec<TaskTemplate>> {
let conn = db.lock().unwrap();
let mut stmt = conn.prepare(
"SELECT id, name, description, model, system_prompt, user_message_template,
temperature, max_tokens, created_at, updated_at
FROM batch_templates
ORDER BY created_at DESC",
)?;
let rows = stmt.query_map([], |row| {
Ok(TaskTemplate {
id: Uuid::parse_str(&row.get::<_, String>(0)?).unwrap(),
name: row.get(1)?,
description: row.get(2)?,
model: row.get(3)?,
system_prompt: row.get(4)?,
user_message_template: row.get(5)?,
temperature: row.get(6)?,
max_tokens: row.get(7)?,
created_at: chrono::DateTime::parse_from_rfc3339(&row.get::<_, String>(8)?)
.unwrap()
.into(),
updated_at: chrono::DateTime::parse_from_rfc3339(&row.get::<_, String>(9)?)
.unwrap()
.into(),
})
})?;
let mut templates = Vec::new();
for row in rows {
templates.push(row?);
}
Ok(templates)
}
/// 删除模板
pub fn delete(db: &DbConnection, id: &Uuid) -> Result<bool> {
let conn = db.lock().unwrap();
let affected = conn.execute(
"DELETE FROM batch_templates WHERE id = ?1",
params![id.to_string()],
)?;
Ok(affected > 0)
}
}
#[cfg(test)]
mod tests {
use super::*;
use rusqlite::Connection;
use std::collections::HashMap;
fn setup_test_db() -> DbConnection {
let conn = Connection::open_in_memory().unwrap();
let db = Arc::new(Mutex::new(conn));
BatchTaskDao::init_tables(&db).unwrap();
db
}
#[test]
fn test_init_tables() {
let db = setup_test_db();
// 如果能成功创建,说明表初始化成功
assert!(true);
}
#[test]
fn test_save_and_get_template() {
let db = setup_test_db();
let template = TaskTemplate::new(
"测试模板".to_string(),
"gpt-4".to_string(),
"请处理: {{content}}".to_string(),
);
// 保存模板
TemplateDao::save(&db, &template).unwrap();
// 查询模板
let loaded = TemplateDao::get_by_id(&db, &template.id).unwrap();
assert!(loaded.is_some());
let loaded = loaded.unwrap();
assert_eq!(loaded.name, template.name);
assert_eq!(loaded.model, template.model);
assert_eq!(loaded.user_message_template, template.user_message_template);
}
#[test]
fn test_save_and_get_batch_task() {
let db = setup_test_db();
let template = TaskTemplate::new(
"测试模板".to_string(),
"gpt-4".to_string(),
"请处理: {{content}}".to_string(),
);
let tasks = vec![super::super::batch::TaskDefinition {
id: None,
variables: {
let mut map = HashMap::new();
map.insert("content".to_string(), "测试内容".to_string());
map
},
metadata: HashMap::new(),
}];
let batch_task = BatchTask::new(
"测试批量任务".to_string(),
template.id,
tasks,
super::super::batch::BatchOptions::default(),
);
// 保存批量任务
BatchTaskDao::save(&db, &batch_task).unwrap();
// 查询批量任务
let loaded = BatchTaskDao::get_by_id(&db, &batch_task.id).unwrap();
assert!(loaded.is_some());
let loaded = loaded.unwrap();
assert_eq!(loaded.name, batch_task.name);
assert_eq!(loaded.template_id, batch_task.template_id);
assert_eq!(loaded.tasks.len(), 1);
}
}
+406
View File
@@ -0,0 +1,406 @@
//! Agent Scheduler 数据访问层
//!
//! 提供任务的持久化存储功能
use super::types::{ScheduledTask, TaskFilter, TaskStatus};
use rusqlite::{params, Connection};
use tracing::{error, warn};
pub struct SchedulerDao;
impl SchedulerDao {
/// 创建任务表
pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> {
conn.execute(
"CREATE TABLE IF NOT EXISTS scheduled_tasks (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
description TEXT,
task_type TEXT NOT NULL,
params TEXT NOT NULL,
provider_type TEXT NOT NULL,
model TEXT NOT NULL,
status TEXT NOT NULL,
scheduled_at TEXT NOT NULL,
started_at TEXT,
completed_at TEXT,
result TEXT,
error_message TEXT,
retry_count INTEGER NOT NULL DEFAULT 0,
max_retries INTEGER NOT NULL DEFAULT 3,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
)",
[],
)?;
// 创建索引
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_scheduled_tasks_status ON scheduled_tasks(status)",
[],
)?;
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_scheduled_tasks_scheduled_at ON scheduled_tasks(scheduled_at)",
[],
)?;
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_scheduled_tasks_provider_type ON scheduled_tasks(provider_type)",
[],
)?;
Ok(())
}
/// 创建新任务
pub fn create_task(conn: &Connection, task: &ScheduledTask) -> Result<(), rusqlite::Error> {
let params_json = serde_json::to_string(&task.params)
.map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
let result_json = task
.result
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
conn.execute(
"INSERT INTO scheduled_tasks (
id, name, description, task_type, params, provider_type, model,
status, scheduled_at, started_at, completed_at, result, error_message,
retry_count, max_retries, created_at, updated_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17)",
params![
task.id,
task.name,
task.description,
task.task_type,
params_json,
task.provider_type,
task.model,
task.status.to_string(),
task.scheduled_at,
task.started_at,
task.completed_at,
result_json,
task.error_message,
task.retry_count,
task.max_retries,
task.created_at,
task.updated_at,
],
)?;
Ok(())
}
/// 获取任务
pub fn get_task(conn: &Connection, id: &str) -> Result<Option<ScheduledTask>, rusqlite::Error> {
let mut stmt = conn.prepare(
"SELECT id, name, description, task_type, params, provider_type, model,
status, scheduled_at, started_at, completed_at, result, error_message,
retry_count, max_retries, created_at, updated_at
FROM scheduled_tasks WHERE id = ?",
)?;
let mut rows = stmt.query([id])?;
if let Some(row) = rows.next()? {
Ok(Some(Self::row_to_task(row)?))
} else {
Ok(None)
}
}
/// 查询任务列表
pub fn list_tasks(
conn: &Connection,
filter: &TaskFilter,
) -> Result<Vec<ScheduledTask>, rusqlite::Error> {
let mut query = String::from(
"SELECT id, name, description, task_type, params, provider_type, model,
status, scheduled_at, started_at, completed_at, result, error_message,
retry_count, max_retries, created_at, updated_at
FROM scheduled_tasks WHERE 1=1",
);
let mut params = Vec::new();
if let Some(status) = &filter.status {
query.push_str(&format!(" AND status = ?{}", params.len() + 1));
params.push(status.to_string());
}
if let Some(provider_type) = &filter.provider_type {
query.push_str(&format!(" AND provider_type = ?{}", params.len() + 1));
params.push(provider_type.clone());
}
if let Some(task_type) = &filter.task_type {
query.push_str(&format!(" AND task_type = ?{}", params.len() + 1));
params.push(task_type.clone());
}
if filter.only_due {
query.push_str(&format!(" AND status = 'pending' AND scheduled_at <= datetime('now')"));
}
query.push_str(" ORDER BY scheduled_at ASC");
if let Some(limit) = filter.limit {
query.push_str(&format!(" LIMIT {}", limit));
}
let mut stmt = conn.prepare(&query)?;
let param_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p as &dyn rusqlite::ToSql).collect();
let tasks = stmt.query_map(param_refs.as_slice(), |row| Self::row_to_task(row))?;
tasks.collect()
}
/// 更新任务
pub fn update_task(conn: &Connection, task: &ScheduledTask) -> Result<(), rusqlite::Error> {
let params_json = serde_json::to_string(&task.params)
.map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
let result_json = task
.result
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
conn.execute(
"UPDATE scheduled_tasks SET
name = ?1, description = ?2, task_type = ?3, params = ?4,
provider_type = ?5, model = ?6, status = ?7, scheduled_at = ?8,
started_at = ?9, completed_at = ?10, result = ?11, error_message = ?12,
retry_count = ?13, max_retries = ?14, updated_at = ?15
WHERE id = ?16",
params![
task.name,
task.description,
task.task_type,
params_json,
task.provider_type,
task.model,
task.status.to_string(),
task.scheduled_at,
task.started_at,
task.completed_at,
result_json,
task.error_message,
task.retry_count,
task.max_retries,
task.updated_at,
task.id,
],
)?;
Ok(())
}
/// 删除任务
pub fn delete_task(conn: &Connection, id: &str) -> Result<bool, rusqlite::Error> {
let rows = conn.execute("DELETE FROM scheduled_tasks WHERE id = ?", [id])?;
Ok(rows > 0)
}
/// 获取到期任务
pub fn get_due_tasks(conn: &Connection, limit: usize) -> Result<Vec<ScheduledTask>, rusqlite::Error> {
let mut stmt = conn.prepare(
"SELECT id, name, description, task_type, params, provider_type, model,
status, scheduled_at, started_at, completed_at, result, error_message,
retry_count, max_retries, created_at, updated_at
FROM scheduled_tasks
WHERE status = 'pending' AND scheduled_at <= datetime('now')
ORDER BY scheduled_at ASC
LIMIT ?",
)?;
let tasks = stmt.query_map([limit], |row| Self::row_to_task(row))?;
tasks.collect()
}
/// 将数据库行转换为 ScheduledTask
fn row_to_task(row: &rusqlite::Row) -> Result<ScheduledTask, rusqlite::Error> {
let params_json: String = row.get(4)?;
let params: serde_json::Value = serde_json::from_str(&params_json).map_err(|e| {
warn!("Failed to parse params JSON: {}", e);
rusqlite::Error::FromSqlConversionFailure(
4,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
let result_json: Option<String> = row.get(11)?;
let result = result_json
.map(|json| {
serde_json::from_str(&json).map_err(|e| {
warn!("Failed to parse result JSON: {}", e);
rusqlite::Error::FromSqlConversionFailure(
11,
rusqlite::types::Type::Text,
Box::new(e),
)
})
})
.transpose()?;
let status_str: String = row.get(7)?;
let status = match status_str.as_str() {
"pending" => TaskStatus::Pending,
"running" => TaskStatus::Running,
"completed" => TaskStatus::Completed,
"failed" => TaskStatus::Failed,
"cancelled" => TaskStatus::Cancelled,
_ => {
warn!("Unknown task status: {}, defaulting to pending", status_str);
TaskStatus::Pending
}
};
Ok(ScheduledTask {
id: row.get(0)?,
name: row.get(1)?,
description: row.get(2)?,
task_type: row.get(3)?,
params,
provider_type: row.get(5)?,
model: row.get(6)?,
status,
scheduled_at: row.get(8)?,
started_at: row.get(9)?,
completed_at: row.get(10)?,
result,
error_message: row.get(12)?,
retry_count: row.get(13)?,
max_retries: row.get(14)?,
created_at: row.get(15)?,
updated_at: row.get(16)?,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
fn setup_test_db() -> Connection {
let conn = Connection::open_in_memory().unwrap();
SchedulerDao::create_tables(&conn).unwrap();
conn
}
#[test]
fn test_create_and_get_task() {
let conn = setup_test_db();
let task = ScheduledTask::new(
"Test Task".to_string(),
"test_type".to_string(),
serde_json::json!({"key": "value"}),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
SchedulerDao::create_task(&conn, &task).unwrap();
let retrieved = SchedulerDao::get_task(&conn, &task.id).unwrap().unwrap();
assert_eq!(retrieved.id, task.id);
assert_eq!(retrieved.name, task.name);
assert_eq!(retrieved.task_type, task.task_type);
}
#[test]
fn test_list_tasks_with_filter() {
let conn = setup_test_db();
let task1 = ScheduledTask::new(
"Task 1".to_string(),
"type_a".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
let task2 = ScheduledTask::new(
"Task 2".to_string(),
"type_b".to_string(),
serde_json::json!(null),
"anthropic".to_string(),
"claude-3".to_string(),
Utc::now(),
);
SchedulerDao::create_task(&conn, &task1).unwrap();
SchedulerDao::create_task(&conn, &task2).unwrap();
// 查询所有
let all = SchedulerDao::list_tasks(&conn, &TaskFilter::default()).unwrap();
assert_eq!(all.len(), 2);
// 按 task_type 过滤
let filtered = SchedulerDao::list_tasks(
&conn,
&TaskFilter {
task_type: Some("type_a".to_string()),
..Default::default()
},
)
.unwrap();
assert_eq!(filtered.len(), 1);
assert_eq!(filtered[0].task_type, "type_a");
}
#[test]
fn test_update_task() {
let conn = setup_test_db();
let mut task = ScheduledTask::new(
"Test".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
SchedulerDao::create_task(&conn, &task).unwrap();
task.mark_completed(Some(serde_json::json!("done")));
SchedulerDao::update_task(&conn, &task).unwrap();
let updated = SchedulerDao::get_task(&conn, &task.id).unwrap().unwrap();
assert_eq!(updated.status, TaskStatus::Completed);
assert_eq!(updated.result, Some(serde_json::json!("done")));
}
#[test]
fn test_delete_task() {
let conn = setup_test_db();
let task = ScheduledTask::new(
"Test".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
SchedulerDao::create_task(&conn, &task).unwrap();
assert!(SchedulerDao::delete_task(&conn, &task.id).unwrap());
let retrieved = SchedulerDao::get_task(&conn, &task.id).unwrap();
assert!(retrieved.is_none());
}
}
+254
View File
@@ -0,0 +1,254 @@
//! Agent Task Executor
//!
//! 负责执行调度的任务
use super::types::ScheduledTask;
use async_trait::async_trait;
use proxycast_agent::credential_bridge::CredentialBridge;
use proxycast_core::database::DbConnection;
use std::sync::Arc;
/// 任务执行器 Trait
#[async_trait]
pub trait TaskExecutor: Send + Sync {
/// 执行任务
///
/// # 参数
/// - `task`: 要执行的任务
/// - `db`: 数据库连接
///
/// # 返回
/// - 成功返回执行结果(JSON 格式)
/// - 失败返回错误信息
async fn execute(
&self,
task: &ScheduledTask,
db: &DbConnection,
) -> Result<serde_json::Value, String>;
}
/// Agent 任务执行器
///
/// 通过 CredentialBridge 选择凭证,调用 Aster Agent 执行任务
pub struct AgentExecutor {
credential_bridge: Arc<CredentialBridge>,
}
impl AgentExecutor {
/// 创建新的执行器实例
pub fn new() -> Self {
Self {
credential_bridge: Arc::new(CredentialBridge::new()),
}
}
/// 使用自定义的 CredentialBridge 创建执行器
pub fn with_credential_bridge(credential_bridge: Arc<CredentialBridge>) -> Self {
Self { credential_bridge }
}
}
impl Default for AgentExecutor {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl TaskExecutor for AgentExecutor {
async fn execute(
&self,
task: &ScheduledTask,
db: &DbConnection,
) -> Result<serde_json::Value, String> {
tracing::info!(
"[AgentExecutor] 开始执行任务: {} (类型: {}, provider: {}, model: {})",
task.name,
task.task_type,
task.provider_type,
task.model
);
// 1. 从凭证池选择凭证
let aster_config = self
.credential_bridge
.select_and_configure(db, &task.provider_type, &task.model)
.await
.map_err(|e| format!("选择凭证失败: {e}"))?;
tracing::info!(
"[AgentExecutor] 已选择凭证: {} (provider: {}, model: {})",
aster_config.credential_uuid,
aster_config.provider_name,
aster_config.model_name
);
// 2. 根据任务类型执行不同的操作
let result = match task.task_type.as_str() {
"agent_chat" => {
// 执行 Agent 对话任务
self.execute_agent_chat(task, db, &aster_config).await?
}
"batch_process" => {
// 执行批量处理任务
self.execute_batch_process(task, db, &aster_config).await?
}
"scheduled_report" => {
// 执行定时报告任务
self.execute_scheduled_report(task, db, &aster_config)
.await?
}
_ => {
return Err(format!("不支持的任务类型: {}", task.task_type));
}
};
// 3. 标记凭证为健康
if let Err(e) = self
.credential_bridge
.mark_healthy(db, &aster_config.credential_uuid, Some(&task.model))
{
tracing::warn!("[AgentExecutor] 标记凭证健康失败: {}", e);
}
tracing::info!("[AgentExecutor] 任务执行成功: {}", task.name);
Ok(result)
}
}
impl AgentExecutor {
/// 执行 Agent 对话任务
async fn execute_agent_chat(
&self,
task: &ScheduledTask,
_db: &DbConnection,
_aster_config: &proxycast_agent::credential_bridge::AsterProviderConfig,
) -> Result<serde_json::Value, String> {
// 从任务参数中提取对话内容
let prompt = task
.params
.get("prompt")
.and_then(|v| v.as_str())
.ok_or_else(|| "缺少 prompt 参数".to_string())?;
tracing::info!("[AgentExecutor] 执行 Agent 对话: {}", prompt);
// TODO: 实际调用 Aster Agent 执行对话
// 这里需要集成 AsterAgentState 来执行对话
// 暂时返回模拟结果
Ok(serde_json::json!({
"type": "agent_chat",
"prompt": prompt,
"response": "任务已调度执行",
"status": "success"
}))
}
/// 执行批量处理任务
async fn execute_batch_process(
&self,
task: &ScheduledTask,
_db: &DbConnection,
_aster_config: &proxycast_agent::credential_bridge::AsterProviderConfig,
) -> Result<serde_json::Value, String> {
let items = task
.params
.get("items")
.and_then(|v| v.as_array())
.ok_or_else(|| "缺少 items 参数".to_string())?;
tracing::info!("[AgentExecutor] 执行批量处理: {} 项", items.len());
// TODO: 实际执行批量处理逻辑
Ok(serde_json::json!({
"type": "batch_process",
"total": items.len(),
"processed": items.len(),
"status": "success"
}))
}
/// 执行定时报告任务
async fn execute_scheduled_report(
&self,
task: &ScheduledTask,
_db: &DbConnection,
_aster_config: &proxycast_agent::credential_bridge::AsterProviderConfig,
) -> Result<serde_json::Value, String> {
let report_type = task
.params
.get("report_type")
.and_then(|v| v.as_str())
.ok_or_else(|| "缺少 report_type 参数".to_string())?;
tracing::info!("[AgentExecutor] 生成定时报告: {}", report_type);
// TODO: 实际生成报告逻辑
Ok(serde_json::json!({
"type": "scheduled_report",
"report_type": report_type,
"generated_at": chrono::Utc::now().to_rfc3339(),
"status": "success"
}))
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
use rusqlite::Connection;
use std::sync::{Arc, Mutex};
fn setup_test_db() -> DbConnection {
let conn = Connection::open_in_memory().unwrap();
Arc::new(Mutex::new(conn))
}
#[tokio::test]
async fn test_executor_creation() {
let executor = AgentExecutor::new();
assert!(Arc::strong_count(&executor.credential_bridge) >= 1);
}
#[tokio::test]
async fn test_execute_agent_chat_missing_prompt() {
let executor = AgentExecutor::new();
let db = setup_test_db();
let task = ScheduledTask::new(
"Test".to_string(),
"agent_chat".to_string(),
serde_json::json!({}), // 缺少 prompt
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
// 由于缺少凭证池数据,这里会在选择凭证时失败
// 但我们可以测试参数验证逻辑
let result = executor.execute(&task, &db).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_unsupported_task_type() {
let executor = AgentExecutor::new();
let db = setup_test_db();
let task = ScheduledTask::new(
"Test".to_string(),
"unsupported_type".to_string(),
serde_json::json!({}),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
let result = executor.execute(&task, &db).await;
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("不支持的任务类型"));
}
}
+59
View File
@@ -0,0 +1,59 @@
//! ProxyCast Agent Scheduler
//!
//! 提供 Agent 任务调度功能,支持定时任务、重试机制等。
//!
//! ## 功能
//! - 任务创建和管理
//! - 任务持久化到 SQLite
//! - 定时任务调度
//! - 任务状态跟踪
//! - 失败重试机制
//! - 批量任务支持
//!
//! ## 使用示例
//!
//! ```rust,no_run
//! use proxycast_scheduler::{AgentScheduler, ScheduledTask, SchedulerTrait};
//! use chrono::Utc;
//!
//! # async fn example(db: proxycast_core::database::DbConnection) -> Result<(), String> {
//! // 初始化调度器
//! AgentScheduler::init_tables(&db)?;
//! let scheduler = AgentScheduler::new(db);
//!
//! // 创建任务
//! let task = ScheduledTask::new(
//! "测试任务".to_string(),
//! "test_task".to_string(),
//! serde_json::json!({"param": "value"}),
//! "openai".to_string(),
//! "gpt-4".to_string(),
//! Utc::now(),
//! );
//!
//! let task_id = scheduler.create_task(task).await?;
//!
//! // 获取到期任务
//! let due_tasks = scheduler.get_due_tasks(10).await?;
//! # Ok(())
//! # }
//! ```
pub mod batch;
pub mod batch_dao;
pub mod dao;
pub mod executor;
pub mod scheduler;
pub mod template;
pub mod types;
pub use batch::{
BatchOptions, BatchTask, BatchTaskStatistics, BatchTaskStatus, TaskDefinition, TaskResult,
TaskStatus as BatchTaskStatus2, TokenUsage,
};
pub use batch_dao::{BatchTaskDao, TemplateDao};
pub use dao::SchedulerDao;
pub use executor::{AgentExecutor, TaskExecutor};
pub use scheduler::{AgentScheduler, SchedulerTrait};
pub use template::TaskTemplate;
pub use types::{ScheduledTask, TaskFilter, TaskStatus};
+275
View File
@@ -0,0 +1,275 @@
//! Agent Scheduler 核心实现
//!
//! 提供任务调度的核心功能
use super::dao::SchedulerDao;
use super::types::{ScheduledTask, TaskFilter, TaskStatus};
use async_trait::async_trait;
use proxycast_core::database::DbConnection;
use std::sync::Arc;
/// 调度器 Trait
///
/// 定义调度器的核心接口
#[async_trait]
pub trait SchedulerTrait: Send + Sync {
/// 创建新任务
async fn create_task(&self, task: ScheduledTask) -> Result<String, String>;
/// 获取任务
async fn get_task(&self, id: &str) -> Result<Option<ScheduledTask>, String>;
/// 查询任务列表
async fn list_tasks(&self, filter: TaskFilter) -> Result<Vec<ScheduledTask>, String>;
/// 更新任务
async fn update_task(&self, task: ScheduledTask) -> Result<(), String>;
/// 删除任务
async fn delete_task(&self, id: &str) -> Result<bool, String>;
/// 获取到期任务
async fn get_due_tasks(&self, limit: usize) -> Result<Vec<ScheduledTask>, String>;
/// 标记任务为运行中
async fn mark_task_running(&self, id: &str) -> Result<(), String>;
/// 标记任务为完成
async fn mark_task_completed(
&self,
id: &str,
result: Option<serde_json::Value>,
) -> Result<(), String>;
/// 标记任务为失败
async fn mark_task_failed(&self, id: &str, error: String) -> Result<(), String>;
/// 标记任务为取消
async fn mark_task_cancelled(&self, id: &str) -> Result<(), String>;
}
/// Agent Scheduler 实现
pub struct AgentScheduler {
db: DbConnection,
}
impl AgentScheduler {
/// 创建新的调度器实例
pub fn new(db: DbConnection) -> Self {
Self { db }
}
/// 初始化数据库表
pub fn init_tables(db: &DbConnection) -> Result<(), String> {
let conn = proxycast_core::database::lock_db(db)?;
SchedulerDao::create_tables(&conn).map_err(|e| format!("创建调度器表失败: {e}"))
}
}
#[async_trait]
impl SchedulerTrait for AgentScheduler {
async fn create_task(&self, task: ScheduledTask) -> Result<String, String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
let task_id = task.id.clone();
SchedulerDao::create_task(&conn, &task)
.map_err(|e| format!("创建任务失败: {e}"))?;
tracing::info!("[AgentScheduler] 创建任务: {} ({})", task.name, task_id);
Ok(task_id)
}
async fn get_task(&self, id: &str) -> Result<Option<ScheduledTask>, String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
SchedulerDao::get_task(&conn, id).map_err(|e| format!("获取任务失败: {e}"))
}
async fn list_tasks(&self, filter: TaskFilter) -> Result<Vec<ScheduledTask>, String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
SchedulerDao::list_tasks(&conn, &filter).map_err(|e| format!("查询任务列表失败: {e}"))
}
async fn update_task(&self, task: ScheduledTask) -> Result<(), String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
SchedulerDao::update_task(&conn, &task).map_err(|e| format!("更新任务失败: {e}"))
}
async fn delete_task(&self, id: &str) -> Result<bool, String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
let deleted = SchedulerDao::delete_task(&conn, id)
.map_err(|e| format!("删除任务失败: {e}"))?;
if deleted {
tracing::info!("[AgentScheduler] 删除任务: {}", id);
}
Ok(deleted)
}
async fn get_due_tasks(&self, limit: usize) -> Result<Vec<ScheduledTask>, String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
SchedulerDao::get_due_tasks(&conn, limit).map_err(|e| format!("获取到期任务失败: {e}"))
}
async fn mark_task_running(&self, id: &str) -> Result<(), String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
let mut task = SchedulerDao::get_task(&conn, id)
.map_err(|e| format!("获取任务失败: {e}"))?
.ok_or_else(|| format!("任务不存在: {id}"))?;
task.mark_running();
SchedulerDao::update_task(&conn, &task).map_err(|e| format!("更新任务状态失败: {e}"))?;
tracing::info!("[AgentScheduler] 任务开始执行: {}", id);
Ok(())
}
async fn mark_task_completed(
&self,
id: &str,
result: Option<serde_json::Value>,
) -> Result<(), String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
let mut task = SchedulerDao::get_task(&conn, id)
.map_err(|e| format!("获取任务失败: {e}"))?
.ok_or_else(|| format!("任务不存在: {id}"))?;
task.mark_completed(result);
SchedulerDao::update_task(&conn, &task).map_err(|e| format!("更新任务状态失败: {e}"))?;
tracing::info!("[AgentScheduler] 任务执行成功: {}", id);
Ok(())
}
async fn mark_task_failed(&self, id: &str, error: String) -> Result<(), String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
let mut task = SchedulerDao::get_task(&conn, id)
.map_err(|e| format!("获取任务失败: {e}"))?
.ok_or_else(|| format!("任务不存在: {id}"))?;
task.mark_failed(error.clone());
SchedulerDao::update_task(&conn, &task).map_err(|e| format!("更新任务状态失败: {e}"))?;
tracing::warn!("[AgentScheduler] 任务执行失败: {} - {}", id, error);
Ok(())
}
async fn mark_task_cancelled(&self, id: &str) -> Result<(), String> {
let conn = proxycast_core::database::lock_db(&self.db)?;
let mut task = SchedulerDao::get_task(&conn, id)
.map_err(|e| format!("获取任务失败: {e}"))?
.ok_or_else(|| format!("任务不存在: {id}"))?;
task.mark_cancelled();
SchedulerDao::update_task(&conn, &task).map_err(|e| format!("更新任务状态失败: {e}"))?;
tracing::info!("[AgentScheduler] 任务已取消: {}", id);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
use rusqlite::Connection;
use std::sync::{Arc, Mutex};
fn setup_test_scheduler() -> AgentScheduler {
let conn = Connection::open_in_memory().unwrap();
SchedulerDao::create_tables(&conn).unwrap();
let db = Arc::new(Mutex::new(conn));
AgentScheduler::new(db)
}
#[tokio::test]
async fn test_create_and_get_task() {
let scheduler = setup_test_scheduler();
let task = ScheduledTask::new(
"Test Task".to_string(),
"test_type".to_string(),
serde_json::json!({"key": "value"}),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
let task_id = scheduler.create_task(task.clone()).await.unwrap();
let retrieved = scheduler.get_task(&task_id).await.unwrap().unwrap();
assert_eq!(retrieved.id, task_id);
assert_eq!(retrieved.name, task.name);
}
#[tokio::test]
async fn test_mark_task_running() {
let scheduler = setup_test_scheduler();
let task = ScheduledTask::new(
"Test".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
let task_id = scheduler.create_task(task).await.unwrap();
scheduler.mark_task_running(&task_id).await.unwrap();
let updated = scheduler.get_task(&task_id).await.unwrap().unwrap();
assert_eq!(updated.status, TaskStatus::Running);
assert!(updated.started_at.is_some());
}
#[tokio::test]
async fn test_mark_task_completed() {
let scheduler = setup_test_scheduler();
let task = ScheduledTask::new(
"Test".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
let task_id = scheduler.create_task(task).await.unwrap();
scheduler.mark_task_running(&task_id).await.unwrap();
scheduler
.mark_task_completed(&task_id, Some(serde_json::json!("success")))
.await
.unwrap();
let updated = scheduler.get_task(&task_id).await.unwrap().unwrap();
assert_eq!(updated.status, TaskStatus::Completed);
assert_eq!(updated.result, Some(serde_json::json!("success")));
}
#[tokio::test]
async fn test_get_due_tasks() {
let scheduler = setup_test_scheduler();
let past = Utc::now() - chrono::Duration::hours(1);
let future = Utc::now() + chrono::Duration::hours(1);
let past_task = ScheduledTask::new(
"Past".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
past,
);
let future_task = ScheduledTask::new(
"Future".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
future,
);
scheduler.create_task(past_task).await.unwrap();
scheduler.create_task(future_task).await.unwrap();
let due_tasks = scheduler.get_due_tasks(10).await.unwrap();
assert_eq!(due_tasks.len(), 1);
assert_eq!(due_tasks[0].name, "Past");
}
}
+138
View File
@@ -0,0 +1,138 @@
//! 任务模板定义
//!
//! 定义可复用的任务模板
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use uuid::Uuid;
/// 任务模板
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskTemplate {
/// 模板 ID
pub id: Uuid,
/// 模板名称
pub name: String,
/// 模板描述
pub description: Option<String>,
/// 模型名称
pub model: String,
/// 系统提示词
pub system_prompt: Option<String>,
/// 用户消息模板 (支持变量替换,例如 "{{variable_name}}")
pub user_message_template: String,
/// 温度参数
pub temperature: Option<f32>,
/// 最大 tokens
pub max_tokens: Option<u32>,
/// 创建时间
pub created_at: chrono::DateTime<chrono::Utc>,
/// 更新时间
pub updated_at: chrono::DateTime<chrono::Utc>,
}
impl TaskTemplate {
/// 创建新的任务模板
pub fn new(
name: String,
model: String,
user_message_template: String,
) -> Self {
let now = chrono::Utc::now();
Self {
id: Uuid::new_v4(),
name,
description: None,
model,
system_prompt: None,
user_message_template,
temperature: None,
max_tokens: None,
created_at: now,
updated_at: now,
}
}
/// 设置描述
pub fn with_description(mut self, description: String) -> Self {
self.description = Some(description);
self
}
/// 设置系统提示词
pub fn with_system_prompt(mut self, system_prompt: String) -> Self {
self.system_prompt = Some(system_prompt);
self
}
/// 设置温度
pub fn with_temperature(mut self, temperature: f32) -> Self {
self.temperature = Some(temperature);
self
}
/// 设置最大 tokens
pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
self.max_tokens = Some(max_tokens);
self
}
/// 渲染用户消息 (替换变量)
pub fn render_user_message(&self, variables: &HashMap<String, String>) -> String {
let mut message = self.user_message_template.clone();
for (key, value) in variables {
let placeholder = format!("{{{{{}}}}}", key);
message = message.replace(&placeholder, value);
}
message
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_template_creation() {
let template = TaskTemplate::new(
"测试模板".to_string(),
"gpt-4".to_string(),
"请处理: {{content}}".to_string(),
)
.with_description("这是一个测试模板".to_string())
.with_temperature(0.7)
.with_max_tokens(1000);
assert_eq!(template.name, "测试模板");
assert_eq!(template.model, "gpt-4");
assert_eq!(template.temperature, Some(0.7));
assert_eq!(template.max_tokens, Some(1000));
}
#[test]
fn test_render_user_message() {
let template = TaskTemplate::new(
"测试模板".to_string(),
"gpt-4".to_string(),
"请处理内容: {{content}}, 来自: {{source}}".to_string(),
);
let mut variables = HashMap::new();
variables.insert("content".to_string(), "测试内容".to_string());
variables.insert("source".to_string(), "测试来源".to_string());
let rendered = template.render_user_message(&variables);
assert_eq!(rendered, "请处理内容: 测试内容, 来自: 测试来源");
}
}
+300
View File
@@ -0,0 +1,300 @@
//! Agent Scheduler 数据模型
//!
//! 定义调度任务相关的数据结构
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
/// 任务状态
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum TaskStatus {
/// 等待执行
Pending,
/// 正在执行
Running,
/// 执行成功
Completed,
/// 执行失败
Failed,
/// 已取消
Cancelled,
}
impl Default for TaskStatus {
fn default() -> Self {
Self::Pending
}
}
impl std::fmt::Display for TaskStatus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Pending => write!(f, "pending"),
Self::Running => write!(f, "running"),
Self::Completed => write!(f, "completed"),
Self::Failed => write!(f, "failed"),
Self::Cancelled => write!(f, "cancelled"),
}
}
}
/// 调度任务
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ScheduledTask {
/// 任务 UUID
pub id: String,
/// 任务名称
pub name: String,
/// 任务描述
pub description: Option<String>,
/// 任务类型(标识要执行的操作)
pub task_type: String,
/// 任务参数(JSON 格式)
pub params: serde_json::Value,
/// Provider 类型
pub provider_type: String,
/// 模型名称
pub model: String,
/// 任务状态
pub status: TaskStatus,
/// 计划执行时间(RFC3339 格式)
pub scheduled_at: String,
/// 实际执行开始时间(可选)
pub started_at: Option<String>,
/// 实际执行完成时间(可选)
pub completed_at: Option<String>,
/// 执行结果(可选)
pub result: Option<serde_json::Value>,
/// 错误信息(如果执行失败)
pub error_message: Option<String>,
/// 重试次数
pub retry_count: u32,
/// 最大重试次数
pub max_retries: u32,
/// 创建时间
pub created_at: String,
/// 更新时间
pub updated_at: String,
}
impl ScheduledTask {
/// 创建新任务
pub fn new(
name: String,
task_type: String,
params: serde_json::Value,
provider_type: String,
model: String,
scheduled_at: DateTime<Utc>,
) -> Self {
let now = Utc::now();
Self {
id: Uuid::new_v4().to_string(),
name,
description: None,
task_type,
params,
provider_type,
model,
status: TaskStatus::Pending,
scheduled_at: scheduled_at.to_rfc3339(),
started_at: None,
completed_at: None,
result: None,
error_message: None,
retry_count: 0,
max_retries: 3,
created_at: now.to_rfc3339(),
updated_at: now.to_rfc3339(),
}
}
/// 检查任务是否到期
pub fn is_due(&self) -> bool {
if self.status != TaskStatus::Pending {
return false;
}
match DateTime::parse_from_rfc3339(&self.scheduled_at) {
Ok(scheduled_time) => {
let now = Utc::now();
scheduled_time <= now.with_timezone(&chrono::FixedOffset::east_opt(0).unwrap())
}
Err(_) => false,
}
}
/// 检查是否可以重试
pub fn can_retry(&self) -> bool {
self.status == TaskStatus::Failed && self.retry_count < self.max_retries
}
/// 增加重试计数
pub fn increment_retry(&mut self) {
self.retry_count += 1;
self.updated_at = Utc::now().to_rfc3339();
}
/// 标记为运行中
pub fn mark_running(&mut self) {
self.status = TaskStatus::Running;
self.started_at = Some(Utc::now().to_rfc3339());
self.updated_at = Utc::now().to_rfc3339();
}
/// 标记为完成
pub fn mark_completed(&mut self, result: Option<serde_json::Value>) {
self.status = TaskStatus::Completed;
self.completed_at = Some(Utc::now().to_rfc3339());
self.result = result;
self.updated_at = Utc::now().to_rfc3339();
}
/// 标记为失败
pub fn mark_failed(&mut self, error: String) {
self.status = TaskStatus::Failed;
self.completed_at = Some(Utc::now().to_rfc3339());
self.error_message = Some(error);
self.updated_at = Utc::now().to_rfc3339();
}
/// 标记为取消
pub fn mark_cancelled(&mut self) {
self.status = TaskStatus::Cancelled;
self.completed_at = Some(Utc::now().to_rfc3339());
self.updated_at = Utc::now().to_rfc3339();
}
}
/// 任务查询过滤器
#[derive(Debug, Clone, Default)]
pub struct TaskFilter {
/// 任务状态
pub status: Option<TaskStatus>,
/// Provider 类型
pub provider_type: Option<String>,
/// 任务类型
pub task_type: Option<String>,
/// 是否只查询到期的任务
pub only_due: bool,
/// 限制返回数量
pub limit: Option<usize>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_task_status_display() {
assert_eq!(TaskStatus::Pending.to_string(), "pending");
assert_eq!(TaskStatus::Running.to_string(), "running");
assert_eq!(TaskStatus::Completed.to_string(), "completed");
assert_eq!(TaskStatus::Failed.to_string(), "failed");
assert_eq!(TaskStatus::Cancelled.to_string(), "cancelled");
}
#[test]
fn test_scheduled_task_creation() {
let task = ScheduledTask::new(
"Test Task".to_string(),
"test_type".to_string(),
serde_json::json!({"key": "value"}),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
assert_eq!(task.name, "Test Task");
assert_eq!(task.status, TaskStatus::Pending);
assert_eq!(task.retry_count, 0);
assert_eq!(task.max_retries, 3);
}
#[test]
fn test_task_is_due() {
let past = Utc::now() - chrono::Duration::hours(1);
let future = Utc::now() + chrono::Duration::hours(1);
let mut past_task = ScheduledTask::new(
"Past Task".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
past,
);
let future_task = ScheduledTask::new(
"Future Task".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
future,
);
assert!(past_task.is_due());
assert!(!future_task.is_due());
// 非 pending 状态的任务不应被判定为到期
past_task.status = TaskStatus::Running;
assert!(!past_task.is_due());
}
#[test]
fn test_task_can_retry() {
let mut task = ScheduledTask::new(
"Test".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
task.status = TaskStatus::Failed;
assert!(task.can_retry());
task.retry_count = 3;
assert!(!task.can_retry());
task.status = TaskStatus::Completed;
assert!(!task.can_retry());
}
#[test]
fn test_task_status_transitions() {
let mut task = ScheduledTask::new(
"Test".to_string(),
"test".to_string(),
serde_json::json!(null),
"openai".to_string(),
"gpt-4".to_string(),
Utc::now(),
);
task.mark_running();
assert_eq!(task.status, TaskStatus::Running);
assert!(task.started_at.is_some());
task.mark_completed(Some(serde_json::json!("success")));
assert_eq!(task.status, TaskStatus::Completed);
assert!(task.completed_at.is_some());
assert_eq!(task.result, Some(serde_json::json!("success")));
// 重置状态测试失败
task.status = TaskStatus::Pending;
task.started_at = None;
task.completed_at = None;
task.result = None;
task.mark_running();
task.mark_failed("error message".to_string());
assert_eq!(task.status, TaskStatus::Failed);
assert!(task.completed_at.is_some());
assert_eq!(task.error_message, Some("error message".to_string()));
}
}
+119 -405
View File
@@ -23,7 +23,6 @@ use axum::{
Json,
};
use serde_json::json;
use std::collections::HashMap;
use crate::client_detector::ClientType;
use crate::{record_request_telemetry, record_token_usage, AppState};
@@ -40,6 +39,93 @@ use proxycast_server_utils::{
use super::{call_provider_anthropic, call_provider_openai};
async fn select_credential_for_request(
state: &AppState,
selected_provider: &str,
model: &str,
client_type: &ClientType,
explicit_provider_id: Option<&str>,
log_prefix: &str,
include_error_code: bool,
) -> Result<Option<proxycast_core::models::provider_pool_model::ProviderCredential>, Response> {
let db = match &state.db {
Some(db) => db,
None => {
eprintln!("[{log_prefix}] 数据库未初始化!");
return Ok(None);
}
};
if let Some(explicit_provider_id) = explicit_provider_id {
eprintln!(
"[{log_prefix}] 使用 X-Provider-Id 指定的 provider: {explicit_provider_id}"
);
let cred = state
.pool_service
.select_credential_with_client_check(
db,
explicit_provider_id,
Some(model),
Some(client_type),
)
.ok()
.flatten();
if cred.is_none() {
eprintln!(
"[{log_prefix}] X-Provider-Id '{explicit_provider_id}' 没有可用凭证,不进行降级"
);
state.logs.write().await.add(
"error",
&format!(
"[ROUTE] No available credentials for explicitly specified provider '{explicit_provider_id}', refusing to fallback"
),
);
let mut error_body = json!({
"error": {
"type": "provider_unavailable",
"message": format!("No available credentials for provider '{}'", explicit_provider_id)
}
});
if include_error_code {
error_body["error"]["code"] = json!("no_credentials");
}
return Err((StatusCode::SERVICE_UNAVAILABLE, Json(error_body)).into_response());
}
return Ok(cred);
}
let provider_id_hint = selected_provider.to_lowercase();
match state
.pool_service
.select_credential_with_fallback(
db,
&state.api_key_service,
selected_provider,
Some(model),
Some(provider_id_hint.as_str()),
Some(client_type),
)
.await
{
Ok(cred) => {
if cred.is_some() {
eprintln!("[{log_prefix}] 找到凭证: provider={selected_provider}");
} else {
eprintln!("[{log_prefix}] 未找到凭证: provider={selected_provider}");
}
Ok(cred)
}
Err(e) => {
eprintln!("[{log_prefix}] 选择凭证失败: {e}");
Ok(None)
}
}
}
// ============================================================================
// Provider 选择辅助函数
// ============================================================================
@@ -249,231 +335,23 @@ pub async fn chat_completions(
.and_then(|v| v.to_str().ok())
.map(|s| s.to_lowercase());
// 尝试从凭证池中选择凭证(带客户端兼容性检查)
// 如果指定了 X-Provider-Id,优先使用它(不降级)
// 否则使用 selected_provider
// 尝试选择凭证:
// 1) X-Provider-Id 指定时仅走精确匹配(不降级)
// 2) 否则走统一的“池优先 + API Key Provider 智能降级”路径
eprintln!("[CHAT_COMPLETIONS] 开始选择凭证...");
let credential = match &state.db {
Some(db) => {
// 如果指定了 X-Provider-Id,优先使用它(不降级)
if let Some(ref explicit_provider_id) = provider_id_header {
eprintln!(
"[CHAT_COMPLETIONS] 使用 X-Provider-Id 指定的 provider: {explicit_provider_id}"
);
let cred = state
.pool_service
.select_credential_with_client_check(
db,
explicit_provider_id,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
if cred.is_none() {
eprintln!(
"[CHAT_COMPLETIONS] X-Provider-Id '{explicit_provider_id}' 没有可用凭证,不进行降级"
);
state.logs.write().await.add(
"error",
&format!(
"[ROUTE] No available credentials for explicitly specified provider '{explicit_provider_id}', refusing to fallback"
),
);
// 返回错误,不降级
return (
StatusCode::SERVICE_UNAVAILABLE,
Json(json!({
"error": {
"message": format!("No available credentials for provider '{}'", explicit_provider_id),
"type": "provider_unavailable",
"code": "no_credentials"
}
})),
)
.into_response();
}
cred
} else {
// 使用 selected_provider(从 API Server 配置中获取)
eprintln!(
"[CHAT_COMPLETIONS] 尝试从凭证池选择: provider={}, model={}",
selected_provider, request.model
);
let cred = state
.pool_service
.select_credential_with_client_check(
db,
&selected_provider,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
if cred.is_some() {
eprintln!("[CHAT_COMPLETIONS] 找到凭证: provider={selected_provider}");
} else {
eprintln!("[CHAT_COMPLETIONS] 未找到凭证: provider={selected_provider}");
}
cred
}
}
None => {
eprintln!("[CHAT_COMPLETIONS] 数据库未初始化!");
None
}
};
// 如果 Provider Pool 中没有找到凭证,尝试从 API Key Provider 获取(智能降级)
let credential = if credential.is_none() {
eprintln!("[CHAT_COMPLETIONS] Provider Pool 中未找到凭证,尝试 API Key Provider...");
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
let provider_id_lower = selected_provider.to_lowercase();
// 策略 1: 优先按 provider_id 直接查找(支持 deepseek, moonshot 等 60+ Provider)
// 这些 Provider 在 API Key Provider 中有独立配置
let mut found_credential: Option<
proxycast_core::models::provider_pool_model::ProviderCredential,
> = None;
if let Some(db) = &state.db {
// 先尝试按 provider_id 直接查找
eprintln!("[CHAT_COMPLETIONS] 尝试按 provider_id '{provider_id_lower}' 直接查找凭证");
match state
.api_key_service
.get_fallback_credential(
db,
&proxycast_core::models::provider_pool_model::PoolProviderType::OpenAI,
Some(&provider_id_lower),
Some(&client_type),
)
.await
{
Ok(Some(cred)) => {
eprintln!(
"[CHAT_COMPLETIONS] 通过 provider_id '{}' 找到凭证: name={:?}",
provider_id_lower, cred.name
);
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] Using API Key Provider credential by provider_id: {provider_id_lower}"
),
);
found_credential = Some(cred);
}
Ok(None) => {
eprintln!(
"[CHAT_COMPLETIONS] provider_id '{provider_id_lower}' 未找到凭证,尝试按类型查找"
);
}
Err(e) => {
eprintln!("[CHAT_COMPLETIONS] 按 provider_id 查找凭证失败: {e}");
}
}
// 策略 2: 如果按 provider_id 未找到,按类型查找
if found_credential.is_none() {
let api_provider_type = match provider_id_lower.as_str() {
"anthropic" | "claude" => Some(ApiProviderType::Anthropic),
"openai" => Some(ApiProviderType::Openai),
"gemini" => Some(ApiProviderType::Gemini),
// 以下都是 OpenAI 兼容的 Provider,但优先按 provider_id 查找已在上面处理
"deepseek" | "moonshot" | "groq" | "grok" | "mistral" | "perplexity"
| "cohere" | "openrouter" | "silicon" => Some(ApiProviderType::Openai),
_ => None,
};
if let Some(api_type) = api_provider_type {
eprintln!(
"[CHAT_COMPLETIONS] 尝试从 API Key Provider 类型 '{api_type:?}' 获取凭证"
);
match state.api_key_service.get_next_api_key_by_type(db, api_type) {
Ok(Some((_key_id, api_key, provider_info))) => {
eprintln!(
"[CHAT_COMPLETIONS] 从 API Key Provider 获取到凭证: provider={}, api_host={}",
provider_info.name,
provider_info.api_host
);
let base_url = if provider_info.api_host.is_empty() {
None
} else {
Some(provider_info.api_host.clone())
};
let provider_type = match provider_info.provider_type {
ApiProviderType::Anthropic => {
proxycast_core::ProviderType::Anthropic
}
ApiProviderType::Openai | ApiProviderType::OpenaiResponse => {
proxycast_core::ProviderType::OpenAI
}
ApiProviderType::Gemini => {
proxycast_core::ProviderType::GeminiApiKey
}
_ => proxycast_core::ProviderType::OpenAI,
};
let credential_data = match provider_type {
proxycast_core::ProviderType::Anthropic => {
proxycast_core::models::provider_pool_model::CredentialData::AnthropicKey {
api_key: api_key.clone(),
base_url,
}
}
proxycast_core::ProviderType::GeminiApiKey => {
proxycast_core::models::provider_pool_model::CredentialData::GeminiApiKey {
api_key: api_key.clone(),
base_url,
excluded_models: vec![],
}
}
_ => proxycast_core::models::provider_pool_model::CredentialData::OpenAIKey {
api_key: api_key.clone(),
base_url,
},
};
let mut cred =
proxycast_core::models::provider_pool_model::ProviderCredential::new(
provider_type,
credential_data,
);
cred.name = Some(provider_info.name.clone());
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] Using API Key Provider credential: provider={}, type={:?}",
provider_info.name, provider_info.provider_type
),
);
found_credential = Some(cred);
}
Ok(None) => {
eprintln!(
"[CHAT_COMPLETIONS] API Key Provider 类型 '{api_type:?}' 没有可用的 API Key"
);
}
Err(e) => {
eprintln!("[CHAT_COMPLETIONS] 从 API Key Provider 获取凭证失败: {e}");
}
}
}
}
}
found_credential
} else {
credential
let credential = match select_credential_for_request(
&state,
&selected_provider,
&request.model,
&client_type,
provider_id_header.as_deref(),
"CHAT_COMPLETIONS",
true,
)
.await
{
Ok(cred) => cred,
Err(resp) => return resp,
};
// 如果找到凭证池中的凭证,使用它
@@ -983,123 +861,22 @@ pub async fn anthropic_messages(
.and_then(|v| v.to_str().ok())
.map(|s| s.to_lowercase());
// 尝试从凭证池中选择凭证(带客户端兼容性检查)
// 如果指定了 X-Provider-Id,优先使用它(不降级)
// 否则使用 selected_provider
let credential = match &state.db {
Some(db) => {
// 如果指定了 X-Provider-Id,优先使用它(不降级)
if let Some(ref explicit_provider_id) = provider_id_header {
eprintln!(
"[ANTHROPIC_MESSAGES] 使用 X-Provider-Id 指定的 provider: {explicit_provider_id}"
);
let cred = state
.pool_service
.select_credential_with_client_check(
db,
explicit_provider_id,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
if cred.is_none() {
eprintln!(
"[AMP] X-Provider-Id '{explicit_provider_id}' 没有可用凭证,不进行降级"
);
state.logs.write().await.add(
"error",
&format!(
"[ROUTE] No available credentials for explicitly specified provider '{explicit_provider_id}', refusing to fallback"
),
);
// 返回错误,不降级
return (
StatusCode::SERVICE_UNAVAILABLE,
Json(json!({
"error": {
"type": "provider_unavailable",
"message": format!("No available credentials for provider '{}'", explicit_provider_id)
}
})),
)
.into_response();
}
cred
} else {
// 使用 selected_provider(从 API Server 配置中获取)
eprintln!(
"[ANTHROPIC_MESSAGES] 尝试从凭证池选择: provider={}, model={}",
selected_provider, request.model
);
let cred = state
.pool_service
.select_credential_with_client_check(
db,
&selected_provider,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
if cred.is_some() {
eprintln!("[ANTHROPIC_MESSAGES] 找到凭证: provider={selected_provider}");
} else {
eprintln!("[ANTHROPIC_MESSAGES] 未找到凭证: provider={selected_provider}");
}
cred
}
}
None => {
eprintln!("[ANTHROPIC_MESSAGES] 数据库未初始化!");
None
}
};
// 如果 Provider Pool 中没有找到凭证,尝试从 API Key Provider 获取(智能降级)
let credential = if credential.is_none() {
eprintln!("[ANTHROPIC_MESSAGES] Provider Pool 中未找到凭证,尝试 API Key Provider...");
// 策略 1: 优先按 provider_id 直接查找(支持自定义 Provider)
let mut found_credential: Option<
proxycast_core::models::provider_pool_model::ProviderCredential,
> = None;
if let Some(db) = &state.db {
eprintln!("[ANTHROPIC_MESSAGES] 尝试按 provider_id '{selected_provider}' 直接查找凭证");
match state
.api_key_service
.get_fallback_credential(
db,
&proxycast_core::models::provider_pool_model::PoolProviderType::Anthropic,
Some(&selected_provider),
Some(&client_type),
)
.await
{
Ok(Some(cred)) => {
eprintln!(
"[ANTHROPIC_MESSAGES] 通过 provider_id '{}' 找到凭证: name={:?}",
selected_provider, cred.name
);
found_credential = Some(cred);
}
Ok(None) => {
eprintln!("[ANTHROPIC_MESSAGES] provider_id '{selected_provider}' 未找到凭证");
}
Err(e) => {
eprintln!("[ANTHROPIC_MESSAGES] 查找凭证时出错: {e}");
}
}
}
found_credential
} else {
credential
// 尝试选择凭证:
// 1) X-Provider-Id 指定时仅走精确匹配(不降级)
// 2) 否则走统一的“池优先 + API Key Provider 智能降级”路径
let credential = match select_credential_for_request(
&state,
&selected_provider,
&request.model,
&client_type,
provider_id_header.as_deref(),
"ANTHROPIC_MESSAGES",
false,
)
.await
{
Ok(cred) => cred,
Err(resp) => return resp,
};
// 如果找到凭证池中的凭证,使用它
@@ -1651,69 +1428,6 @@ fn build_stream_error_response(
})
}
// ============================================================================
// API Key Provider 辅助函数
// ============================================================================
/// 将 provider_type 映射到 API Key Provider ID
fn map_to_api_key_provider_id(provider_type: &str) -> String {
match provider_type.to_lowercase().as_str() {
"openai" | "gpt" => "openai".to_string(),
"anthropic" | "claude" => "anthropic".to_string(),
"gemini" | "google" => "gemini".to_string(),
"azure" | "azure-openai" | "azure_openai" => "azure-openai".to_string(),
"vertexai" | "vertex" => "vertexai".to_string(),
"bedrock" | "aws-bedrock" | "aws_bedrock" => "aws-bedrock".to_string(),
"ollama" => "ollama".to_string(),
_ => provider_type.to_string(),
}
}
/// 根据 API Provider 类型构建额外的请求头
fn build_api_key_headers(
provider_type: &proxycast_core::database::dao::api_key_provider::ApiProviderType,
api_key: &str,
) -> HashMap<String, String> {
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
let mut headers = HashMap::new();
match provider_type {
ApiProviderType::Anthropic => {
headers.insert("x-api-key".to_string(), api_key.to_string());
headers.insert("anthropic-version".to_string(), "2023-06-01".to_string());
}
ApiProviderType::Gemini => {
headers.insert("x-goog-api-key".to_string(), api_key.to_string());
}
ApiProviderType::AzureOpenai => {
headers.insert("api-key".to_string(), api_key.to_string());
}
_ => {
headers.insert("Authorization".to_string(), format!("Bearer {api_key}"));
}
}
headers
}
/// 获取默认的 API Host
fn get_default_api_host(
provider_type: &proxycast_core::database::dao::api_key_provider::ApiProviderType,
) -> String {
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
match provider_type {
ApiProviderType::Openai | ApiProviderType::OpenaiResponse => {
"https://api.openai.com".to_string()
}
ApiProviderType::Anthropic => "https://api.anthropic.com".to_string(),
ApiProviderType::Gemini => "https://generativelanguage.googleapis.com".to_string(),
ApiProviderType::Ollama => "http://localhost:11434".to_string(),
_ => "https://api.openai.com".to_string(),
}
}
/// 将 OpenAI 格式请求转换为 Anthropic 格式
fn convert_openai_to_anthropic(request: &ChatCompletionRequest) -> serde_json::Value {
let mut messages = Vec::new();
@@ -0,0 +1,125 @@
//! API Key Provider 相关公共工具
//!
//! 统一 provider_id 候选映射和鉴权请求头构建,避免各 handler 规则漂移。
use proxycast_core::database::dao::api_key_provider::{
ApiProviderType, ProviderProtocolFamily,
};
/// 收集 API Key Provider ID 候选列表(按优先级)
///
/// 策略:
/// 1. 优先使用请求方显式传入的 provider_type(支持 custom-* / new-api / gateway)
/// 2. 再尝试别名归一化(如 claude -> anthropic)
/// 3. 最后按协议族回退(如 new-api -> openai)
pub(crate) fn collect_api_key_provider_ids(provider_type: &str) -> Vec<String> {
let mut ids = Vec::new();
let mut push_unique = |candidate: String| {
if !candidate.is_empty() && !ids.iter().any(|id| id == &candidate) {
ids.push(candidate);
}
};
let normalized = provider_type.trim().to_lowercase();
push_unique(normalized.clone());
if normalized.contains('_') {
push_unique(normalized.replace('_', "-"));
}
let parsed = normalized
.parse::<ApiProviderType>()
.or_else(|_| normalized.replace('_', "-").parse::<ApiProviderType>());
if let Ok(api_type) = parsed {
push_unique(api_type.to_string());
let family_fallback = match api_type.runtime_spec().protocol_family {
ProviderProtocolFamily::Anthropic => "anthropic",
ProviderProtocolFamily::Gemini => "gemini",
ProviderProtocolFamily::AzureOpenai => "azure-openai",
ProviderProtocolFamily::Vertexai => "vertexai",
ProviderProtocolFamily::AwsBedrock => "aws-bedrock",
ProviderProtocolFamily::Ollama => "ollama",
ProviderProtocolFamily::OpenAiCompatible | ProviderProtocolFamily::Codex => "openai",
};
push_unique(family_fallback.to_string());
return ids;
}
match normalized.as_str() {
"gpt" => push_unique("openai".to_string()),
"claude" => push_unique("anthropic".to_string()),
"google" => push_unique("gemini".to_string()),
"azure" | "azure_openai" => push_unique("azure-openai".to_string()),
"vertex" => push_unique("vertexai".to_string()),
"bedrock" | "aws_bedrock" => push_unique("aws-bedrock".to_string()),
_ => {}
}
ids
}
/// 根据 API Provider 类型构建额外请求头
pub(crate) fn build_api_key_headers(
provider_type: &ApiProviderType,
api_key: &str,
) -> std::collections::HashMap<String, String> {
let mut headers = std::collections::HashMap::new();
let spec = provider_type.runtime_spec();
let auth_value = match spec.auth_prefix {
Some(prefix) => format!("{prefix} {api_key}"),
None => api_key.to_string(),
};
headers.insert(spec.auth_header.to_string(), auth_value);
for (key, value) in spec.extra_headers {
headers.insert((*key).to_string(), (*value).to_string());
}
headers
}
#[cfg(test)]
mod tests {
use super::{build_api_key_headers, collect_api_key_provider_ids};
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
#[test]
fn test_collect_api_key_provider_ids_keeps_specific_before_family_fallback() {
assert_eq!(
collect_api_key_provider_ids("new-api"),
vec!["new-api".to_string(), "openai".to_string()]
);
assert_eq!(
collect_api_key_provider_ids("gateway"),
vec!["gateway".to_string(), "openai".to_string()]
);
}
#[test]
fn test_collect_api_key_provider_ids_anthropic_compatible_has_anthropic_fallback() {
assert_eq!(
collect_api_key_provider_ids("anthropic-compatible"),
vec!["anthropic-compatible".to_string(), "anthropic".to_string()]
);
}
#[test]
fn test_collect_api_key_provider_ids_custom_keeps_exact_only() {
assert_eq!(
collect_api_key_provider_ids("custom-a32774c6-6fd0-433b-8b81-e95340e08793"),
vec!["custom-a32774c6-6fd0-433b-8b81-e95340e08793".to_string()]
);
}
#[test]
fn test_build_api_key_headers_supports_anthropic_compatible() {
let headers = build_api_key_headers(&ApiProviderType::AnthropicCompatible, "test-key");
assert_eq!(headers.get("x-api-key"), Some(&"test-key".to_string()));
assert_eq!(
headers.get("anthropic-version"),
Some(&"2023-06-01".to_string())
);
assert!(!headers.contains_key("Authorization"));
}
}
@@ -0,0 +1,368 @@
//! 批量任务 API 端点
//!
//! 提供批量任务的创建、查询和管理接口
use axum::{
extract::{Path, State},
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use proxycast_scheduler::{
BatchOptions, BatchTask, BatchTaskDao, BatchTaskStatistics, TaskDefinition, TaskTemplate,
TemplateDao,
};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::AppState;
/// 创建批量任务请求
#[derive(Debug, Deserialize)]
pub struct CreateBatchTaskRequest {
/// 批量任务名称
pub name: String,
/// 任务模板 ID
pub template_id: Uuid,
/// 任务列表
pub tasks: Vec<TaskDefinition>,
/// 批量任务选项
#[serde(default)]
pub options: BatchOptions,
}
/// 创建批量任务响应
#[derive(Debug, Serialize)]
pub struct CreateBatchTaskResponse {
/// 批量任务 ID
pub id: Uuid,
/// 批量任务名称
pub name: String,
/// 任务数量
pub task_count: usize,
/// 创建时间
pub created_at: chrono::DateTime<chrono::Utc>,
}
/// 批量任务详情响应
#[derive(Debug, Serialize)]
pub struct BatchTaskDetailResponse {
/// 批量任务信息
#[serde(flatten)]
pub batch_task: BatchTask,
/// 统计信息
pub statistics: BatchTaskStatistics,
}
/// POST /api/batch/tasks - 创建批量任务
pub async fn create_batch_task(
State(state): State<AppState>,
Json(request): Json<CreateBatchTaskRequest>,
) -> Response {
// 检查数据库是否可用
let db = match &state.db {
Some(db) => db,
None => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": "数据库未初始化",
"type": "database_error"
}
})),
)
.into_response();
}
};
// 验证模板是否存在
let template = match TemplateDao::get_by_id(db, &request.template_id) {
Ok(Some(t)) => t,
Ok(None) => {
return (
StatusCode::NOT_FOUND,
Json(serde_json::json!({
"error": {
"message": format!("模板不存在: {}", request.template_id),
"type": "not_found"
}
})),
)
.into_response();
}
Err(e) => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": format!("查询模板失败: {}", e),
"type": "database_error"
}
})),
)
.into_response();
}
};
// 创建批量任务
let batch_task = BatchTask::new(
request.name.clone(),
request.template_id,
request.tasks,
request.options,
);
let batch_id = batch_task.id;
let task_count = batch_task.tasks.len();
let created_at = batch_task.created_at;
// 保存到数据库
if let Err(e) = BatchTaskDao::save(db, &batch_task) {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": format!("保存批量任务失败: {}", e),
"type": "database_error"
}
})),
)
.into_response();
}
state.logs.write().await.add(
"info",
&format!(
"[BATCH] 创建批量任务: id={}, name={}, task_count={}",
batch_id, request.name, task_count
),
);
// TODO: 启动异步执行任务
// 这里需要集成 BatchTaskExecutor
// 返回响应
(
StatusCode::CREATED,
Json(CreateBatchTaskResponse {
id: batch_id,
name: request.name,
task_count,
created_at,
}),
)
.into_response()
}
"info",
&format!(
"[BATCH] 创建批量任务: id={}, name={}, task_count={}",
batch_id, request.name, task_count
),
);
// TODO: 启动异步执行任务
// 这里需要集成 BatchTaskExecutor
// 返回响应
(
StatusCode::CREATED,
Json(CreateBatchTaskResponse {
id: batch_id,
name: request.name,
task_count,
created_at,
}),
)
.into_response()
}
/// GET /api/batch/tasks/:id - 获取批量任务详情
pub async fn get_batch_task(
State(state): State<AppState>,
Path(id): Path<Uuid>,
) -> Response {
let db = match &state.db {
Some(db) => db,
None => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": "数据库未初始化",
"type": "database_error"
}
})),
)
.into_response();
}
};
match BatchTaskDao::get_by_id(db, &id) {
Ok(Some(batch_task)) => {
let statistics = batch_task.get_statistics();
(
StatusCode::OK,
Json(BatchTaskDetailResponse {
batch_task,
statistics,
}),
)
.into_response()
}
Ok(None) => (
StatusCode::NOT_FOUND,
Json(serde_json::json!({
"error": {
"message": format!("批量任务不存在: {}", id),
"type": "not_found"
}
})),
)
.into_response(),
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": format!("查询批量任务失败: {}", e),
"type": "database_error"
}
})),
)
.into_response(),
}
}
/// GET /api/batch/tasks - 获取批量任务列表
pub async fn list_batch_tasks(State(state): State<AppState>) -> Response {
let db = match &state.db {
Some(db) => db,
None => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": "数据库未初始化",
"type": "database_error"
}
})),
)
.into_response();
}
};
match BatchTaskDao::list_all(db, 100) {
Ok(tasks) => (
StatusCode::OK,
Json(serde_json::json!({
"tasks": tasks
})),
)
.into_response(),
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": format!("查询批量任务列表失败: {}", e),
"type": "database_error"
}
})),
)
.into_response(),
}
}
/// DELETE /api/batch/tasks/:id - 取消批量任务
pub async fn cancel_batch_task(
State(state): State<AppState>,
Path(id): Path<Uuid>,
) -> Response {
// TODO: 取消批量任务
state.logs.write().await.add(
"info",
&format!("[BATCH] 取消批量任务: id={}", id),
);
(
StatusCode::NOT_FOUND,
Json(serde_json::json!({
"error": {
"message": format!("批量任务不存在: {}", id),
"type": "not_found"
}
})),
)
.into_response()
}
/// POST /api/batch/templates - 创建任务模板
pub async fn create_template(
State(state): State<AppState>,
Json(template): Json<TaskTemplate>,
) -> Response {
// TODO: 保存模板到数据库
state.logs.write().await.add(
"info",
&format!("[BATCH] 创建任务模板: id={}, name={}", template.id, template.name),
);
(StatusCode::CREATED, Json(template)).into_response()
}
/// GET /api/batch/templates - 获取模板列表
pub async fn list_templates(State(state): State<AppState>) -> Response {
// TODO: 从数据库查询模板列表
state
.logs
.write()
.await
.add("info", "[BATCH] 查询模板列表");
(
StatusCode::OK,
Json(serde_json::json!({
"templates": []
})),
)
.into_response()
}
/// GET /api/batch/templates/:id - 获取模板详情
pub async fn get_template(State(state): State<AppState>, Path(id): Path<Uuid>) -> Response {
// TODO: 从数据库查询模板
state
.logs
.write()
.await
.add("info", &format!("[BATCH] 查询模板: id={}", id));
(
StatusCode::NOT_FOUND,
Json(serde_json::json!({
"error": {
"message": format!("模板不存在: {}", id),
"type": "not_found"
}
})),
)
.into_response()
}
/// DELETE /api/batch/templates/:id - 删除模板
pub async fn delete_template(State(state): State<AppState>, Path(id): Path<Uuid>) -> Response {
// TODO: 从数据库删除模板
state
.logs
.write()
.await
.add("info", &format!("[BATCH] 删除模板: id={}", id));
(StatusCode::NO_CONTENT, ()).into_response()
}
@@ -17,10 +17,12 @@ use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use crate::AppState;
use proxycast_core::database::dao::api_key_provider::{ApiKeyProviderDao, ApiProviderType};
use proxycast_core::database::dao::api_key_provider::ApiKeyProviderDao;
use proxycast_core::database::dao::provider_pool::ProviderPoolDao;
use proxycast_core::models::provider_pool_model::PoolProviderType;
use super::api_key_provider_utils::{build_api_key_headers, collect_api_key_provider_ids};
/// 选择凭证请求参数
#[derive(Debug, Deserialize)]
pub struct SelectCredentialRequest {
@@ -211,54 +213,58 @@ async fn try_select_api_key_credential(
db: &proxycast_core::database::DbConnection,
request: &SelectCredentialRequest,
) -> Result<Option<CredentialResponse>, CredentialApiError> {
// 将 provider_type 映射到 API Key Provider ID
let provider_id = map_to_api_key_provider_id(&request.provider_type);
let candidate_provider_ids = collect_api_key_provider_ids(&request.provider_type);
// 获取 API Key Provider Service
let api_key_service = &state.api_key_service;
// 尝试获取下一个可用的 API Key
let (key_id, api_key) = match api_key_service.get_next_api_key_entry(db, &provider_id) {
Ok(Some((id, key))) => (id, key),
Ok(None) => return Ok(None),
Err(_) => return Ok(None),
};
for provider_id in candidate_provider_ids {
// 尝试获取下一个可用的 API Key
let (key_id, api_key) = match api_key_service.get_next_api_key_entry(db, &provider_id) {
Ok(Some((id, key))) => (id, key),
Ok(None) => continue,
Err(_) => continue,
};
// 获取 Provider 信息以确定 base_url
let conn = db.lock().map_err(|e| CredentialApiError {
error: "database_lock_error".to_string(),
message: format!("数据库锁定失败: {e}"),
status_code: 500,
})?;
// 获取 Provider 信息以确定 base_url
let provider = {
let conn = db.lock().map_err(|e| CredentialApiError {
error: "database_lock_error".to_string(),
message: format!("数据库锁定失败: {e}"),
status_code: 500,
})?;
let provider = match ApiKeyProviderDao::get_provider_by_id(&conn, &provider_id) {
Ok(Some(p)) => p,
Ok(None) => return Ok(None),
Err(_) => return Ok(None),
};
drop(conn);
match ApiKeyProviderDao::get_provider_by_id(&conn, &provider_id) {
Ok(Some(p)) => p,
Ok(None) => continue,
Err(_) => continue,
}
};
// 构建额外的请求头
let extra_headers = build_api_key_headers(&provider.provider_type, &api_key);
// 构建额外的请求头
let extra_headers = build_api_key_headers(&provider.provider_type, &api_key);
let response = CredentialResponse {
uuid: key_id,
provider_type: request.provider_type.clone(),
credential_type: CredentialType::ApiKey,
access_token: api_key,
base_url: provider.api_host,
expires_at: None, // API Key 通常没有过期时间
name: Some(provider.name),
extra_headers: Some(extra_headers),
};
let response = CredentialResponse {
uuid: key_id,
provider_type: request.provider_type.clone(),
credential_type: CredentialType::ApiKey,
access_token: api_key,
base_url: provider.api_host,
expires_at: None, // API Key 通常没有过期时间
name: Some(provider.name),
extra_headers: Some(extra_headers),
};
tracing::info!(
"[CREDENTIALS_API] API Key 凭证选择成功: {} ({})",
response.name.as_deref().unwrap_or("未命名"),
response.uuid
);
tracing::info!(
"[CREDENTIALS_API] API Key 凭证选择成功: {} ({})",
response.name.as_deref().unwrap_or("未命名"),
response.uuid
);
Ok(Some(response))
return Ok(Some(response));
}
Ok(None)
}
/// 尝试从 OAuth 插件选择凭证(已禁用 - 插件系统已移除)
@@ -284,46 +290,6 @@ fn get_oauth_base_url(provider_type: &PoolProviderType) -> String {
}
}
/// 将 provider_type 映射到 API Key Provider ID
fn map_to_api_key_provider_id(provider_type: &str) -> String {
match provider_type.to_lowercase().as_str() {
"openai" | "gpt" => "openai".to_string(),
"anthropic" | "claude" => "anthropic".to_string(),
"gemini" | "google" => "gemini".to_string(),
"azure" | "azure-openai" => "azure-openai".to_string(),
"vertexai" | "vertex" => "vertexai".to_string(),
"bedrock" | "aws-bedrock" => "aws-bedrock".to_string(),
"ollama" => "ollama".to_string(),
_ => provider_type.to_string(),
}
}
/// 根据 API Provider 类型构建额外的请求头
fn build_api_key_headers(
provider_type: &ApiProviderType,
api_key: &str,
) -> std::collections::HashMap<String, String> {
let mut headers = std::collections::HashMap::new();
match provider_type {
ApiProviderType::Anthropic => {
headers.insert("x-api-key".to_string(), api_key.to_string());
headers.insert("anthropic-version".to_string(), "2023-06-01".to_string());
}
ApiProviderType::Gemini => {
headers.insert("x-goog-api-key".to_string(), api_key.to_string());
}
ApiProviderType::AzureOpenai => {
headers.insert("api-key".to_string(), api_key.to_string());
}
_ => {
headers.insert("Authorization".to_string(), format!("Bearer {api_key}"));
}
}
headers
}
/// GET /v1/credentials/{uuid}/token - 获取指定凭证的 Token
///
/// 支持多种凭证类型:
@@ -3,6 +3,8 @@
//! 将 server 中的各类处理器拆分到独立文件
pub mod api;
pub mod api_key_provider_utils;
pub mod batch_api;
pub mod credentials_api;
pub mod image_handler;
pub mod kiro_credential;
@@ -11,6 +13,7 @@ pub mod provider_calls;
pub mod websocket;
pub use api::*;
pub use batch_api::*;
pub use credentials_api::*;
pub use image_handler::*;
// 避免 SelectCredentialRequest 歧义 glob re-export(credentials_api 和 kiro_credential 都定义了同名类型)
+13
View File
@@ -943,6 +943,17 @@ async fn run_server(
get(handlers::credentials_get_token),
);
// 批量任务 API 路由
let batch_api_routes = Router::new()
.route("/api/batch/tasks", post(handlers::create_batch_task))
.route("/api/batch/tasks", get(handlers::list_batch_tasks))
.route("/api/batch/tasks/:id", get(handlers::get_batch_task))
.route("/api/batch/tasks/:id", axum::routing::delete(handlers::cancel_batch_task))
.route("/api/batch/templates", post(handlers::create_template))
.route("/api/batch/templates", get(handlers::list_templates))
.route("/api/batch/templates/:id", get(handlers::get_template))
.route("/api/batch/templates/:id", axum::routing::delete(handlers::delete_template));
let app = Router::new()
.route("/health", get(health))
.route("/v1/models", get(models))
@@ -985,6 +996,8 @@ async fn run_server(
.merge(kiro_api_routes)
// 凭证 API 路由(用于 aster Agent 集成)
.merge(credentials_api_routes)
// 批量任务 API 路由
.merge(batch_api_routes)
.layer(DefaultBodyLimit::max(body_limit))
.with_state(state);
@@ -7,6 +7,7 @@
use base64::{engine::general_purpose::STANDARD as BASE64, Engine};
use chrono::Utc;
use crate::provider_type_mapping::pool_provider_type_to_api_type;
use proxycast_core::database::dao::api_key_provider::{
ApiKeyEntry, ApiKeyProvider, ApiKeyProviderDao, ApiProviderType, ProviderGroup,
ProviderWithKeys,
@@ -42,6 +43,7 @@ pub struct ConnectionTestResult {
#[cfg(test)]
mod tests {
use super::ApiKeyProviderService;
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
#[test]
fn test_build_codex_responses_request_input_list() {
@@ -62,6 +64,19 @@ data: [DONE]\n";
let content = ApiKeyProviderService::parse_codex_responses_sse_content(body);
assert_eq!(content, "hi!");
}
#[test]
fn test_uses_anthropic_protocol() {
assert!(ApiKeyProviderService::uses_anthropic_protocol(
ApiProviderType::Anthropic
));
assert!(ApiKeyProviderService::uses_anthropic_protocol(
ApiProviderType::AnthropicCompatible
));
assert!(!ApiKeyProviderService::uses_anthropic_protocol(
ApiProviderType::Openai
));
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -214,13 +229,28 @@ impl ApiKeyProviderService {
let start = Instant::now();
// Codex 协议直接走 /responses 端点
let result = if provider.provider_type == ApiProviderType::Codex {
self.test_codex_responses_endpoint(&api_key, &provider.api_host, &test_model, &prompt)
.await
} else {
self.test_openai_chat_once(&api_key, &provider.api_host, &test_model, &prompt)
// 根据 Provider 协议类型选择测试方式
let result = match provider.provider_type {
// Codex 协议直接走 /responses 端点
ApiProviderType::Codex => {
self.test_codex_responses_endpoint(
&api_key,
&provider.api_host,
&test_model,
&prompt,
)
.await
}
// Anthropic / AnthropicCompatible 统一走 /v1/messages
provider_type if Self::uses_anthropic_protocol(provider_type) => {
self.test_anthropic_chat_once(&api_key, &provider.api_host, &test_model, &prompt)
.await
}
// 其余默认 OpenAI 兼容
_ => {
self.test_openai_chat_once(&api_key, &provider.api_host, &test_model, &prompt)
.await
}
};
let latency_ms = start.elapsed().as_millis() as u64;
@@ -242,6 +272,11 @@ impl ApiKeyProviderService {
}
}
#[inline]
fn uses_anthropic_protocol(provider_type: ApiProviderType) -> bool {
provider_type.is_anthropic_protocol()
}
async fn test_openai_chat_once(
&self,
api_key: &str,
@@ -326,6 +361,52 @@ impl ApiKeyProviderService {
Err(format!("API 返回错误: {status} - {body}"))
}
async fn test_anthropic_chat_once(
&self,
api_key: &str,
api_host: &str,
model: &str,
prompt: &str,
) -> Result<(String, String), String> {
use proxycast_providers::providers::claude_custom::ClaudeCustomProvider;
let provider =
ClaudeCustomProvider::with_config(api_key.to_string(), Some(api_host.to_string()));
let request = serde_json::json!({
"model": model,
"max_tokens": 64,
"messages": [{"role": "user", "content": prompt}]
});
let resp = provider
.messages(&request)
.await
.map_err(|e| format!("API 调用失败: {e}"))?;
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
if !status.is_success() {
return Err(format!("API 返回错误: {status} - {body}"));
}
let parsed: serde_json::Value =
serde_json::from_str(&body).map_err(|e| format!("解析响应失败: {e} - {body}"))?;
let content = parsed["content"]
.as_array()
.map(|blocks| {
blocks
.iter()
.filter_map(|block| block["text"].as_str())
.collect::<String>()
})
.unwrap_or_default();
Ok((content, body))
}
fn parse_chat_completions_sse_content(body: &str) -> String {
let mut out = String::new();
@@ -1171,7 +1252,7 @@ impl ApiKeyProviderService {
}
// 策略 2: 通过类型映射查找(降级方案)
if let Some(api_type) = self.map_pool_type_to_api_type(pool_type) {
if let Some(api_type) = pool_provider_type_to_api_type(pool_type) {
eprintln!("[get_fallback_credential] 尝试类型映射: {pool_type:?} -> {api_type:?}");
if let Some(cred) = self.find_by_api_type(db, pool_type, &api_type)? {
eprintln!(
@@ -1188,35 +1269,6 @@ impl ApiKeyProviderService {
Ok(None)
}
/// PoolProviderType → ApiProviderType 映射
///
/// 仅映射有明确对应关系的类型
fn map_pool_type_to_api_type(&self, pool_type: &PoolProviderType) -> Option<ApiProviderType> {
match pool_type {
// API Key 类型 - 直接映射
PoolProviderType::Claude => Some(ApiProviderType::Anthropic),
PoolProviderType::OpenAI => Some(ApiProviderType::Openai),
PoolProviderType::GeminiApiKey => Some(ApiProviderType::Gemini),
PoolProviderType::Vertex => Some(ApiProviderType::Vertexai),
// OAuth 类型 - 可降级到 API Key
PoolProviderType::Gemini => Some(ApiProviderType::Gemini), // Gemini OAuth → Gemini API Key
// API Key Provider 类型 - 直接映射
PoolProviderType::Anthropic => Some(ApiProviderType::Anthropic),
PoolProviderType::AnthropicCompatible => Some(ApiProviderType::AnthropicCompatible),
PoolProviderType::AzureOpenai => Some(ApiProviderType::AzureOpenai),
PoolProviderType::AwsBedrock => Some(ApiProviderType::AwsBedrock),
PoolProviderType::Ollama => Some(ApiProviderType::Ollama),
// OAuth-only,无降级
PoolProviderType::Kiro => None,
PoolProviderType::Codex => None,
PoolProviderType::ClaudeOAuth => None,
PoolProviderType::Antigravity => None,
}
}
/// 通过 ApiProviderType 查找凭证
fn find_by_api_type(
&self,
@@ -1585,8 +1637,8 @@ impl ApiKeyProviderService {
// 根据 Provider 类型选择测试方式
let result = match provider.provider_type {
ApiProviderType::Anthropic => {
// Anthropic 不支持 /models 端点,需要发送测试请求
provider_type if Self::uses_anthropic_protocol(provider_type) => {
// Anthropic / AnthropicCompatible 不支持 /models,统一发送 /messages 测试请求
let test_model = model_name
.or_else(|| provider.custom_models.first().cloned())
.unwrap_or_else(|| "claude-3-haiku-20240307".to_string());
+1
View File
@@ -90,4 +90,5 @@ pub mod kiro_event_service;
// 依赖 providers 的服务
pub mod api_key_provider_service;
pub mod provider_pool_service;
pub mod provider_type_mapping;
pub mod token_cache_service;
@@ -5,6 +5,10 @@
#![allow(dead_code)]
use crate::api_key_provider_service::ApiKeyProviderService;
use crate::provider_type_mapping::{
api_provider_type_to_pool_type, is_custom_provider_id, parse_pool_provider_type,
resolve_pool_provider_type_or_default,
};
use chrono::Utc;
use proxycast_core::database::dao::provider_pool::ProviderPoolDao;
use proxycast_core::database::DbConnection;
@@ -140,7 +144,7 @@ impl ProviderPoolService {
db: &DbConnection,
provider_type: &str,
) -> Result<Vec<CredentialDisplay>, String> {
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
let pt = parse_pool_provider_type(provider_type)?;
let conn = proxycast_core::database::lock_db(db)?;
let mut credentials =
ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?;
@@ -165,7 +169,7 @@ impl ProviderPoolService {
check_health: Option<bool>,
check_model_name: Option<String>,
) -> Result<ProviderCredential, String> {
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
let pt = parse_pool_provider_type(provider_type)?;
let mut cred = ProviderCredential::new(pt, credential);
cred.name = name;
@@ -254,9 +258,17 @@ impl ProviderPoolService {
model: Option<&str>,
client_type: Option<&proxycast_core::models::client_type::ClientType>,
) -> Result<Option<ProviderCredential>, String> {
if is_custom_provider_id(provider_type) {
eprintln!(
"[SELECT_CREDENTIAL] custom provider '{}' 使用智能降级路径",
provider_type
);
return Ok(None);
}
// 对于未知的 provider_type,直接返回 None(不是错误)
// 这样可以让 select_credential_with_fallback 继续尝试智能降级
let pt: PoolProviderType = match provider_type.parse() {
let pt: PoolProviderType = match parse_pool_provider_type(provider_type) {
Ok(pt) => pt,
Err(_) => {
eprintln!(
@@ -427,7 +439,45 @@ impl ProviderPoolService {
eprintln!("[select_credential_with_fallback] Provider Pool 未找到凭证,尝试智能降级");
// Step 2: 智能降级到 API Key Provider
let pt: PoolProviderType = provider_type.parse().unwrap_or(PoolProviderType::OpenAI);
let mut pt = resolve_pool_provider_type_or_default(provider_type);
let mut resolved_provider_id_hint = provider_id_hint;
// 对 custom-* 场景优先查询真实 Provider 类型,避免默认按 OpenAI 协议处理
if is_custom_provider_id(provider_type) {
resolved_provider_id_hint = Some(provider_type);
}
if let Some(custom_provider_id) =
resolved_provider_id_hint.filter(|id| is_custom_provider_id(id))
{
match api_key_service.get_provider(db, custom_provider_id) {
Ok(Some(provider_with_keys)) => {
pt = api_provider_type_to_pool_type(provider_with_keys.provider.provider_type);
eprintln!(
"[select_credential_with_fallback] custom provider '{}' 真实类型 {:?} -> {:?}",
custom_provider_id,
provider_with_keys.provider.provider_type,
pt
);
}
Ok(None) => {
eprintln!(
"[select_credential_with_fallback] custom provider '{}' 不存在,继续使用解析类型 {:?}",
custom_provider_id,
pt
);
}
Err(e) => {
eprintln!(
"[select_credential_with_fallback] 查询 custom provider '{}' 失败: {},继续使用解析类型 {:?}",
custom_provider_id,
e,
pt
);
}
}
}
eprintln!(
"[select_credential_with_fallback] 解析 provider_type '{provider_type}' -> {pt:?}"
);
@@ -435,7 +485,7 @@ impl ProviderPoolService {
// 传入 provider_id_hint 支持 60+ Provider
eprintln!("[select_credential_with_fallback] 调用 get_fallback_credential");
if let Some(cred) = api_key_service
.get_fallback_credential(db, &pt, provider_id_hint, client_type)
.get_fallback_credential(db, &pt, resolved_provider_id_hint, client_type)
.await?
{
eprintln!(
@@ -624,7 +674,7 @@ impl ProviderPoolService {
db: &DbConnection,
provider_type: &str,
) -> Result<usize, String> {
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
let pt = parse_pool_provider_type(provider_type)?;
let conn = proxycast_core::database::lock_db(db)?;
ProviderPoolDao::reset_health_by_type(&conn, &pt).map_err(|e| e.to_string())
}
@@ -960,7 +1010,7 @@ impl ProviderPoolService {
db: &DbConnection,
provider_type: &str,
) -> Result<Vec<HealthCheckResult>, String> {
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
let pt = parse_pool_provider_type(provider_type)?;
let credentials = {
let conn = proxycast_core::database::lock_db(db)?;
ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?
@@ -1798,7 +1848,7 @@ impl ProviderPoolService {
check_model_name: Option<String>,
source: proxycast_core::models::provider_pool_model::CredentialSource,
) -> Result<ProviderCredential, String> {
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
let pt = parse_pool_provider_type(provider_type)?;
let mut cred = ProviderCredential::new_with_source(pt, credential, source);
cred.name = name;
@@ -1990,6 +2040,7 @@ pub struct MigrationResult {
#[cfg(test)]
mod tests {
use super::*;
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
// ==================== Property 3: 不健康凭证排除 ====================
// Feature: antigravity-token-refresh, Property 3: 不健康凭证排除
@@ -2120,4 +2171,24 @@ mod tests {
assert_eq!(deserialized.uuid, info.uuid);
assert_eq!(deserialized.is_healthy, info.is_healthy);
}
#[test]
fn test_api_provider_type_to_pool_type_mapping() {
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::Anthropic),
PoolProviderType::Claude
);
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::AnthropicCompatible),
PoolProviderType::AnthropicCompatible
);
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::Gemini),
PoolProviderType::GeminiApiKey
);
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::Openai),
PoolProviderType::OpenAI
);
}
}
@@ -0,0 +1,133 @@
//! Provider 类型映射与解析工具
//!
//! 统一 services 层中 PoolProviderType 与 ApiProviderType 的映射规则,
//! 避免 `provider_pool_service` 与 `api_key_provider_service` 规则漂移。
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
use proxycast_core::models::provider_pool_model::PoolProviderType;
use proxycast_core::models::provider_type::is_custom_provider_id as core_is_custom_provider_id;
/// 是否为自定义 Provider ID(`custom-*`)
pub(crate) fn is_custom_provider_id(provider_type: &str) -> bool {
core_is_custom_provider_id(provider_type)
}
/// 解析 PoolProviderType
pub(crate) fn parse_pool_provider_type(provider_type: &str) -> Result<PoolProviderType, String> {
provider_type.parse().map_err(|e: String| e)
}
/// 解析 PoolProviderType(失败时回退到 OpenAI)
pub(crate) fn resolve_pool_provider_type_or_default(provider_type: &str) -> PoolProviderType {
provider_type.parse().unwrap_or(PoolProviderType::OpenAI)
}
/// ApiProviderType → PoolProviderType 映射
pub(crate) fn api_provider_type_to_pool_type(api_type: ApiProviderType) -> PoolProviderType {
match api_type {
ApiProviderType::Anthropic => PoolProviderType::Claude,
ApiProviderType::AnthropicCompatible => PoolProviderType::AnthropicCompatible,
ApiProviderType::Gemini => PoolProviderType::GeminiApiKey,
ApiProviderType::Vertexai => PoolProviderType::Vertex,
ApiProviderType::AzureOpenai => PoolProviderType::AzureOpenai,
ApiProviderType::AwsBedrock => PoolProviderType::AwsBedrock,
ApiProviderType::Ollama => PoolProviderType::Ollama,
_ => PoolProviderType::OpenAI,
}
}
/// PoolProviderType → ApiProviderType 映射
pub(crate) fn pool_provider_type_to_api_type(
pool_type: &PoolProviderType,
) -> Option<ApiProviderType> {
match pool_type {
// API Key 类型 - 直接映射
PoolProviderType::Claude => Some(ApiProviderType::Anthropic),
PoolProviderType::OpenAI => Some(ApiProviderType::Openai),
PoolProviderType::GeminiApiKey => Some(ApiProviderType::Gemini),
PoolProviderType::Vertex => Some(ApiProviderType::Vertexai),
// OAuth 类型 - 可降级到 API Key
PoolProviderType::Gemini => Some(ApiProviderType::Gemini), // Gemini OAuth → Gemini API Key
// API Key Provider 类型 - 直接映射
PoolProviderType::Anthropic => Some(ApiProviderType::Anthropic),
PoolProviderType::AnthropicCompatible => Some(ApiProviderType::AnthropicCompatible),
PoolProviderType::AzureOpenai => Some(ApiProviderType::AzureOpenai),
PoolProviderType::AwsBedrock => Some(ApiProviderType::AwsBedrock),
PoolProviderType::Ollama => Some(ApiProviderType::Ollama),
// OAuth-only,无降级
PoolProviderType::Kiro => None,
PoolProviderType::Codex => None,
PoolProviderType::ClaudeOAuth => None,
PoolProviderType::Antigravity => None,
}
}
#[cfg(test)]
mod tests {
use super::{
api_provider_type_to_pool_type, is_custom_provider_id, parse_pool_provider_type,
pool_provider_type_to_api_type, resolve_pool_provider_type_or_default,
};
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
use proxycast_core::models::provider_pool_model::PoolProviderType;
#[test]
fn test_api_provider_type_to_pool_type_mapping() {
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::Anthropic),
PoolProviderType::Claude
);
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::AnthropicCompatible),
PoolProviderType::AnthropicCompatible
);
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::Gemini),
PoolProviderType::GeminiApiKey
);
assert_eq!(
api_provider_type_to_pool_type(ApiProviderType::Openai),
PoolProviderType::OpenAI
);
}
#[test]
fn test_pool_provider_type_to_api_type_mapping() {
assert_eq!(
pool_provider_type_to_api_type(&PoolProviderType::Claude),
Some(ApiProviderType::Anthropic)
);
assert_eq!(
pool_provider_type_to_api_type(&PoolProviderType::AnthropicCompatible),
Some(ApiProviderType::AnthropicCompatible)
);
assert_eq!(
pool_provider_type_to_api_type(&PoolProviderType::Kiro),
None
);
}
#[test]
fn test_pool_provider_type_parser_helpers() {
assert_eq!(
parse_pool_provider_type("openai").unwrap(),
PoolProviderType::OpenAI
);
assert!(parse_pool_provider_type("not-exists").is_err());
assert_eq!(
resolve_pool_provider_type_or_default("not-exists"),
PoolProviderType::OpenAI
);
}
#[test]
fn test_is_custom_provider_id() {
assert!(is_custom_provider_id(
"custom-a32774c6-6fd0-433b-8b81-e95340e08793"
));
assert!(!is_custom_provider_id("openai"));
}
}
@@ -0,0 +1,121 @@
//! 电商差评回复 Skill
//!
//! 实现电商差评自动回复功能,支持淘宝、京东、拼多多等平台
use serde::{Deserialize, Serialize};
/// 电商平台类型
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum EcommercePlatform {
/// 淘宝/天猫
Taobao,
/// 京东
JD,
/// 拼多多
Pinduoduo,
}
impl EcommercePlatform {
/// 获取平台名称
pub fn name(&self) -> &str {
match self {
EcommercePlatform::Taobao => "淘宝",
EcommercePlatform::JD => "京东",
EcommercePlatform::Pinduoduo => "拼多多",
}
}
/// 获取平台的回复风格特点
pub fn reply_style(&self) -> &str {
match self {
EcommercePlatform::Taobao => "亲切、友好、注重客户体验",
EcommercePlatform::JD => "专业、高效、强调服务保障",
EcommercePlatform::Pinduoduo => "热情、实惠、突出性价比",
}
}
}
/// 回复配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReplyConfig {
/// 回复语气: polite(礼貌), sincere(真诚), professional(专业)
pub tone: String,
/// 回复长度: short(简短), medium(中等), long(详细)
pub length: String,
/// 回复模板类型
pub template: Option<String>,
}
impl Default for ReplyConfig {
fn default() -> Self {
Self {
tone: "sincere".to_string(),
length: "medium".to_string(),
template: None,
}
}
}
/// 电商差评回复请求
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EcommerceReviewReplyRequest {
/// 电商平台
pub platform: EcommercePlatform,
/// 差评链接
pub review_url: String,
/// 回复配置
pub config: ReplyConfig,
}
/// 电商差评回复结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EcommerceReviewReplyResult {
/// 是否成功
pub success: bool,
/// 提取的差评内容
pub review_content: Option<String>,
/// 识别的问题类型
pub problem_type: Option<String>,
/// 生成的回复
pub reply: Option<String>,
/// 错误信息
pub error: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_platform_name() {
assert_eq!(EcommercePlatform::Taobao.name(), "淘宝");
assert_eq!(EcommercePlatform::JD.name(), "京东");
assert_eq!(EcommercePlatform::Pinduoduo.name(), "拼多多");
}
#[test]
fn test_reply_style() {
assert!(EcommercePlatform::Taobao.reply_style().contains("亲切"));
assert!(EcommercePlatform::JD.reply_style().contains("专业"));
assert!(EcommercePlatform::Pinduoduo.reply_style().contains("热情"));
}
#[test]
fn test_reply_config_default() {
let config = ReplyConfig::default();
assert_eq!(config.tone, "sincere");
assert_eq!(config.length, "medium");
assert!(config.template.is_none());
}
#[test]
fn test_platform_serialization() {
let platform = EcommercePlatform::Taobao;
let json = serde_json::to_string(&platform).unwrap();
assert_eq!(json, "\"taobao\"");
let deserialized: EcommercePlatform = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized, EcommercePlatform::Taobao);
}
}
+3
View File
@@ -8,6 +8,9 @@ mod llm_provider;
mod proxycast_llm_provider;
mod skill_loader;
// 电商 Skill 模块
pub mod ecommerce_review_reply;
pub use execution_callback::{
events, ExecutionCallback, ExecutionCompletePayload, StepCompletePayload, StepErrorPayload,
StepStartPayload,
+2
View File
@@ -9,6 +9,7 @@ homepage.workspace = true
[dependencies]
# 项目内 crate
proxycast-core.workspace = true
proxycast-agent = { path = "../agent" }
# 序列化
serde.workspace = true
@@ -23,6 +24,7 @@ axum.workspace = true
# 时间和 UUID
uuid.workspace = true
chrono.workspace = true
# 工具库
dashmap.workspace = true
+49 -9
View File
@@ -3,6 +3,7 @@
//! 处理 WebSocket 连接和消息
use super::{
handlers::{parse_rpc_request, serialize_rpc_response, RpcHandler, RpcHandlerState},
WsApiRequest, WsApiResponse, WsConfig, WsConnectionManager, WsEndpoint, WsError, WsMessage,
};
use axum::{
@@ -27,15 +28,25 @@ pub struct WsHandlerState {
pub api_key: String,
/// 日志存储
pub logs: Arc<RwLock<LogStore>>,
/// RPC 处理器状态
pub rpc_state: RpcHandlerState,
}
impl WsHandlerState {
/// 创建新的处理器状态
pub fn new(config: WsConfig, api_key: String, logs: Arc<RwLock<LogStore>>) -> Self {
pub fn new(
config: WsConfig,
api_key: String,
logs: Arc<RwLock<LogStore>>,
db: Option<proxycast_core::database::DbConnection>,
scheduler: Option<proxycast_agent::ProxyCastScheduler>,
) -> Self {
let rpc_state = RpcHandlerState::new(db, scheduler, logs.clone());
Self {
manager: Arc::new(WsConnectionManager::new(config)),
api_key,
logs,
rpc_state,
}
}
}
@@ -152,14 +163,43 @@ async fn handle_socket(socket: WebSocket, state: WsHandlerState, client_info: Op
}
}
}
Err(e) => {
state.manager.on_error();
let error = WsMessage::Error(WsError::invalid_message(format!(
"Failed to parse message: {e}"
)));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender.send(Message::Text(error_text)).await.is_err() {
break;
Err(_) => {
// 尝试解析为 RPC 请求
match parse_rpc_request(&text) {
Ok(rpc_req) => {
let rpc_handler = RpcHandler::new(state.rpc_state.clone());
let rpc_resp = rpc_handler.handle_request(rpc_req).await;
match serialize_rpc_response(&rpc_resp) {
Ok(resp_text) => {
if sender.send(Message::Text(resp_text)).await.is_err() {
break;
}
}
Err(e) => {
state.manager.on_error();
let error = WsMessage::Error(WsError::internal(
None,
format!("Failed to serialize RPC response: {:?}", e),
));
let error_text =
serde_json::to_string(&error).unwrap_or_default();
if sender.send(Message::Text(error_text)).await.is_err() {
break;
}
}
}
}
Err(e) => {
state.manager.on_error();
let error = WsMessage::Error(WsError::invalid_message(format!(
"Failed to parse message as WS or RPC: {}",
e.message
)));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender.send(Message::Text(error_text)).await.is_err() {
break;
}
}
}
}
}
@@ -0,0 +1,9 @@
//! WebSocket RPC 处理器
//!
//! 实现 RPC 请求处理逻辑
pub mod rpc_handler;
pub use rpc_handler::{
parse_rpc_request, serialize_rpc_response, RpcHandler, RpcHandlerState,
};
@@ -0,0 +1,286 @@
//! RPC 处理器
//!
//! 处理 Gateway RPC 请求,集成 Agent 和 Scheduler
use super::super::{protocol::*, WsError};
use std::sync::Arc;
use tokio::sync::RwLock;
/// RPC 处理器状态
#[derive(Clone)]
pub struct RpcHandlerState {
/// 数据库连接
pub db: Arc<RwLock<Option<proxycast_core::database::DbConnection>>>,
/// Agent 调度器(可选)
pub scheduler: Arc<RwLock<Option<proxycast_agent::ProxyCastScheduler>>>,
/// 日志存储
pub logs: Arc<RwLock<proxycast_core::LogStore>>,
}
impl RpcHandlerState {
/// 创建新的 RPC 处理器状态
pub fn new(
db: Option<proxycast_core::database::DbConnection>,
scheduler: Option<proxycast_agent::ProxyCastScheduler>,
logs: Arc<RwLock<proxycast_core::LogStore>>,
) -> Self {
Self {
db: Arc::new(RwLock::new(db)),
scheduler: Arc::new(RwLock::new(scheduler)),
logs,
}
}
}
/// RPC 处理器
pub struct RpcHandler {
state: RpcHandlerState,
}
impl RpcHandler {
/// 创建新的 RPC 处理器
pub fn new(state: RpcHandlerState) -> Self {
Self { state }
}
/// 处理 RPC 请求
pub async fn handle_request(&self, request: GatewayRpcRequest) -> GatewayRpcResponse {
let method = request.method;
let request_id = request.id.clone();
let params = request.params;
// 记录请求
self.state.logs.write().await.add(
"info",
&format!("[RPC] Request: id={} method={:?}", request_id, method),
);
// 路由到具体的处理方法
let result = match method {
RpcMethod::AgentRun => self.handle_agent_run(params).await,
RpcMethod::AgentWait => self.handle_agent_wait(params).await,
RpcMethod::AgentStop => self.handle_agent_stop(params).await,
RpcMethod::SessionsList => self.handle_sessions_list().await,
RpcMethod::SessionsGet => self.handle_sessions_get(params).await,
RpcMethod::CronList => self.handle_cron_list().await,
RpcMethod::CronRun => self.handle_cron_run(params).await,
};
match result {
Ok(data) => GatewayRpcResponse {
jsonrpc: "2.0".to_string(),
id: request_id,
result: Some(data),
error: None,
},
Err(err) => {
self.state.logs.write().await.add(
"error",
&format!("[RPC] Error: id={} error={}", request_id, err.message),
);
GatewayRpcResponse {
jsonrpc: "2.0".to_string(),
id: request_id,
result: None,
error: Some(err),
}
}
}
}
/// 处理 agent.run
async fn handle_agent_run(
&self,
params: Option<serde_json::Value>,
) -> Result<serde_json::Value, RpcError> {
let params: AgentRunParams = params
.and_then(|v| serde_json::from_value(v).ok())
.ok_or_else(|| {
RpcError::invalid_params("Missing or invalid parameters for agent.run")
})?;
// TODO: 实现 Agent 运行逻辑
// 1. 获取或创建会话
// 2. 发送消息到 Agent
// 3. 返回运行 ID
let run_id = uuid::Uuid::new_v4().to_string();
let session_id = params
.session_id
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
let result = AgentRunResult {
run_id: run_id.clone(),
session_id,
completed: false, // 流式模式下未完成
content: None,
usage: None,
};
Ok(serde_json::to_value(result).map_err(|e| RpcError::internal_error(e.to_string()))?)
}
/// 处理 agent.wait
async fn handle_agent_wait(
&self,
params: Option<serde_json::Value>,
) -> Result<serde_json::Value, RpcError> {
let params: AgentWaitParams = params
.and_then(|v| serde_json::from_value(v).ok())
.ok_or_else(|| {
RpcError::invalid_params("Missing or invalid parameters for agent.wait")
})?;
// TODO: 实现 Agent 等待逻辑
// 1. 等待运行完成
// 2. 获取结果
let result = AgentWaitResult {
run_id: params.run_id,
completed: true,
content: Some("Response content".to_string()),
usage: Some(TokenUsage::new(100, 200)),
};
Ok(serde_json::to_value(result).map_err(|e| RpcError::internal_error(e.to_string()))?)
}
/// 处理 agent.stop
async fn handle_agent_stop(
&self,
params: Option<serde_json::Value>,
) -> Result<serde_json::Value, RpcError> {
let params: AgentStopParams = params
.and_then(|v| serde_json::from_value(v).ok())
.ok_or_else(|| {
RpcError::invalid_params("Missing or invalid parameters for agent.stop")
})?;
// TODO: 实现 Agent 停止逻辑
// 1. 取消运行
// 2. 清理资源
let result = AgentStopResult {
run_id: params.run_id,
stopped: true,
};
Ok(serde_json::to_value(result).map_err(|e| RpcError::internal_error(e.to_string()))?)
}
/// 处理 sessions.list
async fn handle_sessions_list(&self) -> Result<serde_json::Value, RpcError> {
// TODO: 从数据库获取会话列表
let sessions = vec![];
let result = SessionsListResult { sessions };
Ok(serde_json::to_value(result).map_err(|e| RpcError::internal_error(e.to_string()))?)
}
/// 处理 sessions.get
async fn handle_sessions_get(
&self,
params: Option<serde_json::Value>,
) -> Result<serde_json::Value, RpcError> {
let params: SessionGetParams = params
.and_then(|v| serde_json::from_value(v).ok())
.ok_or_else(|| {
RpcError::invalid_params("Missing or invalid parameters for sessions.get")
})?;
// TODO: 从数据库获取会话详情
let result = SessionGetResult {
session_id: params.session_id,
model: "claude-sonnet-4-5".to_string(),
system_prompt: None,
message_count: 0,
created_at: chrono::Utc::now().to_rfc3339(),
updated_at: chrono::Utc::now().to_rfc3339(),
};
Ok(serde_json::to_value(result).map_err(|e| RpcError::internal_error(e.to_string()))?)
}
/// 处理 cron.list
async fn handle_cron_list(&self) -> Result<serde_json::Value, RpcError> {
// TODO: 从数据库获取定时任务列表
let tasks = vec![];
let result = CronListResult { tasks };
Ok(serde_json::to_value(result).map_err(|e| RpcError::internal_error(e.to_string()))?)
}
/// 处理 cron.run
async fn handle_cron_run(
&self,
params: Option<serde_json::Value>,
) -> Result<serde_json::Value, RpcError> {
let params: CronRunParams = params
.and_then(|v| serde_json::from_value(v).ok())
.ok_or_else(|| RpcError::invalid_params("Missing or invalid parameters for cron.run"))?;
// TODO: 实现定时任务运行逻辑
// 1. 查找任务
// 2. 提交到调度器
// 3. 返回执行 ID
let execution_id = uuid::Uuid::new_v4().to_string();
let result = CronRunResult {
task_id: params.task_id,
execution_id,
started: true,
};
Ok(serde_json::to_value(result).map_err(|e| RpcError::internal_error(e.to_string()))?)
}
}
/// 从 WsMessage 解析 RPC 请求
pub fn parse_rpc_request(msg: &str) -> Result<GatewayRpcRequest, RpcError> {
serde_json::from_str(msg).map_err(|e| RpcError::parse_error(format!("Invalid JSON: {}", e)))
}
/// 序列化 RPC 响应
pub fn serialize_rpc_response(resp: &GatewayRpcResponse) -> Result<String, WsError> {
serde_json::to_string(resp)
.map_err(|e| WsError::internal(None, format!("Failed to serialize response: {}", e)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_rpc_request() {
let json = r#"{
"jsonrpc": "2.0",
"id": "test-123",
"method": "agent.run",
"params": {
"message": "Hello",
"stream": false
}
}"#;
let request = parse_rpc_request(json).unwrap();
assert_eq!(request.method, RpcMethod::AgentRun);
assert_eq!(request.id, "test-123");
}
#[test]
fn test_serialize_rpc_response() {
let response = GatewayRpcResponse {
jsonrpc: "2.0".to_string(),
id: "test-123".to_string(),
result: Some(serde_json::json!({"success": true})),
error: None,
};
let json = serialize_rpc_response(&response).unwrap();
assert!(json.contains("2.0"));
assert!(json.contains("test-123"));
}
}
+4
View File
@@ -9,11 +9,15 @@
#![allow(dead_code)]
pub mod handler;
pub mod handlers;
pub mod lifecycle;
pub mod processor;
pub mod protocol;
pub mod stream;
pub use handlers::RpcHandler;
pub use processor::MessageProcessor;
pub use protocol::{GatewayRpcRequest, GatewayRpcResponse, RpcError, RpcMethod};
pub use proxycast_core::websocket::types;
pub use proxycast_core::websocket::{
KiroTokenInfo, WsApiRequest, WsApiResponse, WsConfig, WsConnection, WsEndpoint, WsError,
+386
View File
@@ -0,0 +1,386 @@
//! WebSocket RPC 协议定义
//!
//! 定义 JSON-RPC 风格的请求/响应结构,支持 Agent 和 Scheduler 操作
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
/// Gateway RPC 请求
///
/// JSON-RPC 2.0 风格的请求结构
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GatewayRpcRequest {
/// JSON-RPC 版本(固定为 "2.0")
pub jsonrpc: String,
/// 请求 ID(用于关联响应)
pub id: String,
/// 方法名
pub method: RpcMethod,
/// 参数(可选)
#[serde(skip_serializing_if = "Option::is_none")]
pub params: Option<serde_json::Value>,
}
/// RPC 方法名
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RpcMethod {
/// Agent 运行
#[serde(rename = "agent.run")]
AgentRun,
/// Agent 等待完成
#[serde(rename = "agent.wait")]
AgentWait,
/// Agent 停止
#[serde(rename = "agent.stop")]
AgentStop,
/// 列出会话
#[serde(rename = "sessions.list")]
SessionsList,
/// 获取会话详情
#[serde(rename = "sessions.get")]
SessionsGet,
/// 列出定时任务
#[serde(rename = "cron.list")]
CronList,
/// 运行定时任务
#[serde(rename = "cron.run")]
CronRun,
}
/// Gateway RPC 响应
///
/// JSON-RPC 2.0 风格的响应结构
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GatewayRpcResponse {
/// JSON-RPC 版本(固定为 "2.0")
pub jsonrpc: String,
/// 请求 ID(关联请求)
pub id: String,
/// 结果(如果成功)
#[serde(skip_serializing_if = "Option::is_none")]
pub result: Option<serde_json::Value>,
/// 错误(如果失败)
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<RpcError>,
}
/// RPC 错误
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RpcError {
/// 错误码
pub code: i32,
/// 错误消息
pub message: String,
/// 错误数据(可选)
#[serde(skip_serializing_if = "Option::is_none")]
pub data: Option<serde_json::Value>,
}
impl RpcError {
/// 创建解析错误
pub fn parse_error(message: impl Into<String>) -> Self {
Self {
code: -32700,
message: message.into(),
data: None,
}
}
/// 创建无效请求错误
pub fn invalid_request(message: impl Into<String>) -> Self {
Self {
code: -32600,
message: message.into(),
data: None,
}
}
/// 创建方法未找到错误
pub fn method_not_found(method: impl Into<String>) -> Self {
Self {
code: -32601,
message: format!("Method not found: {}", method.into()),
data: None,
}
}
/// 创建无效参数错误
pub fn invalid_params(message: impl Into<String>) -> Self {
Self {
code: -32602,
message: message.into(),
data: None,
}
}
/// 创建内部错误
pub fn internal_error(message: impl Into<String>) -> Self {
Self {
code: -32603,
message: message.into(),
data: None,
}
}
}
/// Agent 运行参数
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentRunParams {
/// 会话 ID(可选,用于连续对话)
pub session_id: Option<String>,
/// 用户消息
pub message: String,
/// 模型名称(可选)
pub model: Option<String>,
/// 系统提示词(可选)
pub system_prompt: Option<String>,
/// 温度参数(可选)
pub temperature: Option<f32>,
/// 最大 token 数(可选)
pub max_tokens: Option<u32>,
/// 是否流式响应
#[serde(default)]
pub stream: bool,
}
/// Agent 等待参数
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentWaitParams {
/// 运行 ID
pub run_id: String,
/// 超时时间(毫秒)
#[serde(default = "default_timeout")]
pub timeout: u64,
}
fn default_timeout() -> u64 {
30000 // 30 秒
}
/// Agent 停止参数
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentStopParams {
/// 运行 ID
pub run_id: String,
}
/// 会话获取参数
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionGetParams {
/// 会话 ID
pub session_id: String,
}
/// Cron 运行参数
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CronRunParams {
/// 任务 ID
pub task_id: String,
/// 任务参数(可选)
pub params: Option<HashMap<String, serde_json::Value>>,
}
/// Agent 运行结果
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AgentRunResult {
/// 运行 ID
pub run_id: String,
/// 会话 ID
pub session_id: String,
/// 是否完成
pub completed: bool,
/// 响应内容(如果已完成)
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
/// Token 使用量(如果已完成)
#[serde(skip_serializing_if = "Option::is_none")]
pub usage: Option<TokenUsage>,
}
/// Agent 等待结果
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AgentWaitResult {
/// 运行 ID
pub run_id: String,
/// 是否完成
pub completed: bool,
/// 响应内容(如果已完成)
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
/// Token 使用量(如果已完成)
#[serde(skip_serializing_if = "Option::is_none")]
pub usage: Option<TokenUsage>,
}
/// Agent 停止结果
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AgentStopResult {
/// 运行 ID
pub run_id: String,
/// 是否成功停止
pub stopped: bool,
}
/// 会话列表结果
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionsListResult {
/// 会话列表
pub sessions: Vec<SessionInfo>,
}
/// 会话详情结果
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionGetResult {
/// 会话 ID
pub session_id: String,
/// 模型
pub model: String,
/// 系统提示词
#[serde(skip_serializing_if = "Option::is_none")]
pub system_prompt: Option<String>,
/// 消息数量
pub message_count: usize,
/// 创建时间
pub created_at: String,
/// 更新时间
pub updated_at: String,
}
/// 会话信息
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionInfo {
/// 会话 ID
pub session_id: String,
/// 模型
pub model: String,
/// 消息数量
pub message_count: usize,
/// 创建时间
pub created_at: String,
/// 更新时间
pub updated_at: String,
}
/// Cron 任务列表结果
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CronListResult {
/// 任务列表
pub tasks: Vec<CronTaskInfo>,
}
/// Cron 任务信息
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CronTaskInfo {
/// 任务 ID
pub task_id: String,
/// 任务名称
pub name: String,
/// Cron 表达式
pub schedule: String,
/// 是否启用
pub enabled: bool,
/// 最后运行时间
#[serde(skip_serializing_if = "Option::is_none")]
pub last_run: Option<String>,
/// 下次运行时间
#[serde(skip_serializing_if = "Option::is_none")]
pub next_run: Option<String>,
}
/// Cron 运行结果
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CronRunResult {
/// 任务 ID
pub task_id: String,
/// 执行 ID
pub execution_id: String,
/// 是否成功启动
pub started: bool,
}
/// Token 使用量
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "camelCase")]
pub struct TokenUsage {
/// 输入 token 数
pub input_tokens: u32,
/// 输出 token 数
pub output_tokens: u32,
}
impl TokenUsage {
/// 创建新的 TokenUsage
pub fn new(input_tokens: u32, output_tokens: u32) -> Self {
Self {
input_tokens,
output_tokens,
}
}
/// 计算总 token 数
pub fn total(&self) -> u32 {
self.input_tokens + self.output_tokens
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_serialize_agent_run_request() {
let request = GatewayRpcRequest {
jsonrpc: "2.0".to_string(),
id: "test-123".to_string(),
method: RpcMethod::AgentRun,
params: Some(serde_json::json!({
"message": "Hello, world!",
"model": "claude-sonnet-4-5",
"stream": false
})),
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains("agent.run"));
}
#[test]
fn test_deserialize_agent_run_request() {
let json = r#"{
"jsonrpc": "2.0",
"id": "test-123",
"method": "agent.run",
"params": {
"message": "Hello, world!",
"stream": false
}
}"#;
let request: GatewayRpcRequest = serde_json::from_str(json).unwrap();
assert_eq!(request.method, RpcMethod::AgentRun);
assert!(request.params.is_some());
}
#[test]
fn test_serialize_error_response() {
let response = GatewayRpcResponse {
jsonrpc: "2.0".to_string(),
id: "test-123".to_string(),
result: None,
error: Some(RpcError::method_not_found("unknown.method")),
};
let json = serde_json::to_string(&response).unwrap();
assert!(json.contains("-32601"));
}
}
+160 -69
View File
@@ -203,6 +203,7 @@ pub async fn update_provider_env_vars(
api_host: String,
api_key: Option<String>,
) -> Result<(), String> {
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
use proxycast_services::live_sync::write_env_to_shell_config;
use serde_json::{json, Value};
use std::fs;
@@ -211,75 +212,15 @@ pub async fn update_provider_env_vars(
// 根据 provider_type 确定要更新的环境变量
// 参考 Claude Code 文档:https://code.claude.com/docs/en/llm-gateway
let env_vars: Vec<(String, String)> = match provider_type.to_lowercase().as_str() {
// Anthropic 兼容类型 - 包括大多数第三方 Provider
// 如 DeepSeek、智谱、MiniMax、OpenRouter、AiHubMix 等
"anthropic" | "new-api" | "gateway" => {
let mut vars = vec![("ANTHROPIC_BASE_URL".to_string(), api_host.clone())];
if let Some(key) = api_key {
vars.push(("ANTHROPIC_AUTH_TOKEN".to_string(), key));
}
vars
}
// OpenAI 兼容类型 - 用于 Codex 等
"openai" | "openai-response" => {
let mut vars = vec![("OPENAI_BASE_URL".to_string(), api_host.clone())];
if let Some(key) = api_key {
vars.push(("OPENAI_API_KEY".to_string(), key));
}
vars
}
// Gemini 类型
"gemini" => {
let mut vars = vec![("GEMINI_API_BASE_URL".to_string(), api_host.clone())];
if let Some(key) = api_key {
vars.push(("GEMINI_API_KEY".to_string(), key));
}
vars
}
// Azure OpenAI 类型
"azure-openai" => {
let mut vars = vec![("AZURE_OPENAI_BASE_URL".to_string(), api_host.clone())];
if let Some(key) = api_key {
vars.push(("AZURE_OPENAI_API_KEY".to_string(), key));
}
vars
}
// Google Vertex AI 类型
"vertexai" => {
let mut vars = vec![("ANTHROPIC_VERTEX_BASE_URL".to_string(), api_host.clone())];
if let Some(key) = api_key {
vars.push(("GOOGLE_APPLICATION_CREDENTIALS".to_string(), key));
}
vars
}
// AWS Bedrock 类型
"aws-bedrock" => {
let vars = vec![
("ANTHROPIC_BEDROCK_BASE_URL".to_string(), api_host.clone()),
("CLAUDE_CODE_USE_BEDROCK".to_string(), "1".to_string()),
];
// Bedrock 通常使用 AWS 凭证,不需要单独的 API Key
vars
}
// Ollama 本地部署
"ollama" => {
let vars = vec![("OLLAMA_BASE_URL".to_string(), api_host.clone())];
vars
}
_ => {
// 未知类型,默认使用 ANTHROPIC_BASE_URL(因为大多数第三方 Provider 都是 Anthropic 兼容的)
logs.write().await.add(
"info",
&format!("Provider 类型 '{provider_type}' 使用默认 ANTHROPIC_BASE_URL"),
);
let mut vars = vec![("ANTHROPIC_BASE_URL".to_string(), api_host.clone())];
if let Some(key) = api_key {
vars.push(("ANTHROPIC_AUTH_TOKEN".to_string(), key));
}
vars
}
};
let parsed_api_provider_type = provider_type.parse::<ApiProviderType>().ok();
let env_vars = build_provider_env_vars(&provider_type, &api_host, api_key.as_deref());
if parsed_api_provider_type.is_none() {
logs.write().await.add(
"info",
&format!("Provider 类型 '{provider_type}' 使用默认 OPENAI_BASE_URL"),
);
}
// 1. 更新 ~/.claude/settings.json
let claude_dir = home.join(".claude");
@@ -343,3 +284,153 @@ pub async fn update_provider_env_vars(
Ok(())
}
fn build_provider_env_vars(
provider_type: &str,
api_host: &str,
api_key: Option<&str>,
) -> Vec<(String, String)> {
use proxycast_core::database::dao::api_key_provider::{
ApiProviderType, ProviderProtocolFamily,
};
let push_if_key = |vars: &mut Vec<(String, String)>, key_name: &str| {
if let Some(value) = api_key {
vars.push((key_name.to_string(), value.to_string()));
}
};
match provider_type.to_lowercase().as_str() {
// Anthropic 兼容类型 - 包括大多数第三方 Provider
// 如 DeepSeek、智谱、MiniMax、OpenRouter、AiHubMix 等
"anthropic" | "anthropic-compatible" => {
let mut vars = vec![
("ANTHROPIC_HOST".to_string(), api_host.to_string()),
("ANTHROPIC_BASE_URL".to_string(), api_host.to_string()),
];
push_if_key(&mut vars, "ANTHROPIC_AUTH_TOKEN");
vars
}
// OpenAI 兼容类型 - 用于 Codex 等
"openai" | "openai-response" => {
let mut vars = vec![("OPENAI_BASE_URL".to_string(), api_host.to_string())];
push_if_key(&mut vars, "OPENAI_API_KEY");
vars
}
// Gemini 类型
"gemini" => {
let mut vars = vec![("GEMINI_API_BASE_URL".to_string(), api_host.to_string())];
push_if_key(&mut vars, "GEMINI_API_KEY");
vars
}
// Azure OpenAI 类型
"azure-openai" => {
let mut vars = vec![("AZURE_OPENAI_BASE_URL".to_string(), api_host.to_string())];
push_if_key(&mut vars, "AZURE_OPENAI_API_KEY");
vars
}
// Google Vertex AI 类型
"vertexai" => {
let mut vars = vec![("ANTHROPIC_VERTEX_BASE_URL".to_string(), api_host.to_string())];
push_if_key(&mut vars, "GOOGLE_APPLICATION_CREDENTIALS");
vars
}
// AWS Bedrock 类型
"aws-bedrock" => {
// Bedrock 通常使用 AWS 凭证,不需要单独的 API Key
vec![
("ANTHROPIC_BEDROCK_BASE_URL".to_string(), api_host.to_string()),
("CLAUDE_CODE_USE_BEDROCK".to_string(), "1".to_string()),
]
}
// Ollama 本地部署
"ollama" => vec![("OLLAMA_BASE_URL".to_string(), api_host.to_string())],
_ => {
// 未知类型尽量按已注册 ApiProviderType 的协议族生成环境变量
if let Ok(api_type) = provider_type.parse::<ApiProviderType>() {
match api_type.runtime_spec().protocol_family {
ProviderProtocolFamily::Anthropic => {
let mut vars = vec![
("ANTHROPIC_HOST".to_string(), api_host.to_string()),
("ANTHROPIC_BASE_URL".to_string(), api_host.to_string()),
];
push_if_key(&mut vars, "ANTHROPIC_AUTH_TOKEN");
vars
}
ProviderProtocolFamily::OpenAiCompatible
| ProviderProtocolFamily::Codex
| ProviderProtocolFamily::Gemini
| ProviderProtocolFamily::AzureOpenai
| ProviderProtocolFamily::Vertexai
| ProviderProtocolFamily::AwsBedrock
| ProviderProtocolFamily::Ollama => {
let mut vars = vec![("OPENAI_BASE_URL".to_string(), api_host.to_string())];
push_if_key(&mut vars, "OPENAI_API_KEY");
vars
}
}
} else {
// 兜底:保守使用 OpenAI 兼容变量,避免误导为 Anthropic
let mut vars = vec![("OPENAI_BASE_URL".to_string(), api_host.to_string())];
push_if_key(&mut vars, "OPENAI_API_KEY");
vars
}
}
}
}
#[cfg(test)]
mod tests {
use super::build_provider_env_vars;
#[test]
fn test_build_provider_env_vars_explicit_anthropic_compatible() {
let vars = build_provider_env_vars(
"anthropic-compatible",
"https://open.bigmodel.cn/api/anthropic",
Some("k1"),
);
assert_eq!(
vars,
vec![
(
"ANTHROPIC_HOST".to_string(),
"https://open.bigmodel.cn/api/anthropic".to_string()
),
(
"ANTHROPIC_BASE_URL".to_string(),
"https://open.bigmodel.cn/api/anthropic".to_string()
),
("ANTHROPIC_AUTH_TOKEN".to_string(), "k1".to_string()),
]
);
}
#[test]
fn test_build_provider_env_vars_openai_family_defaults() {
let vars = build_provider_env_vars("new-api", "https://relay.example.com", Some("k2"));
assert_eq!(
vars,
vec![
(
"OPENAI_BASE_URL".to_string(),
"https://relay.example.com".to_string()
),
("OPENAI_API_KEY".to_string(), "k2".to_string()),
]
);
}
#[test]
fn test_build_provider_env_vars_unknown_provider_fallback() {
let vars = build_provider_env_vars("custom-unknown", "https://proxy.example.com", None);
assert_eq!(
vars,
vec![(
"OPENAI_BASE_URL".to_string(),
"https://proxy.example.com".to_string()
)]
);
}
}
+2
View File
@@ -14,12 +14,14 @@
pub mod bootstrap;
pub mod commands;
pub mod runner;
pub mod scheduler_service;
mod setup;
mod state;
mod types;
mod utils;
pub use runner::run;
pub use scheduler_service::{SchedulerService, SchedulerServiceConfig};
pub use setup::setup_app;
pub use state::*;
pub use types::*;
+2
View File
@@ -760,6 +760,8 @@ pub fn run() {
commands::skill_exec_cmd::execute_skill,
commands::skill_exec_cmd::list_executable_skills,
commands::skill_exec_cmd::get_skill_detail,
// Ecommerce Review Reply commands
commands::ecommerce_review_reply_cmd::execute_ecommerce_review_reply,
// Provider Pool commands
commands::provider_pool_cmd::get_provider_pool_overview,
commands::provider_pool_cmd::get_provider_pool_credentials,
+251
View File
@@ -0,0 +1,251 @@
//! Agent Scheduler 服务
//!
//! 提供后台心跳循环,定期检查并执行到期任务
use proxycast_core::database::DbConnection;
use proxycast_scheduler::{AgentExecutor, AgentScheduler, SchedulerTrait, TaskExecutor};
use std::sync::Arc;
use std::time::Duration;
use tokio::time::interval;
use tokio_util::sync::CancellationToken;
/// 调度器服务配置
#[derive(Debug, Clone)]
pub struct SchedulerServiceConfig {
/// 心跳间隔(秒)
pub heartbeat_interval_secs: u64,
/// 每次轮询的最大任务数
pub max_tasks_per_poll: usize,
/// 是否启用调度器
pub enabled: bool,
}
impl Default for SchedulerServiceConfig {
fn default() -> Self {
Self {
heartbeat_interval_secs: 30,
max_tasks_per_poll: 10,
enabled: true,
}
}
}
/// 调度器服务
///
/// 负责启动和管理后台心跳循环
pub struct SchedulerService {
scheduler: Arc<AgentScheduler>,
executor: Arc<AgentExecutor>,
config: SchedulerServiceConfig,
cancel_token: CancellationToken,
}
impl SchedulerService {
/// 创建新的调度器服务
pub fn new(db: DbConnection, config: SchedulerServiceConfig) -> Self {
let scheduler = Arc::new(AgentScheduler::new(db));
let executor = Arc::new(AgentExecutor::new());
let cancel_token = CancellationToken::new();
Self {
scheduler,
executor,
config,
cancel_token,
}
}
/// 启动心跳循环
///
/// 在后台 tokio 任务中运行,定期轮询到期任务并执行
pub fn start(&self, db: DbConnection) {
if !self.config.enabled {
tracing::info!("[SchedulerService] 调度器已禁用,跳过启动");
return;
}
let scheduler = self.scheduler.clone();
let executor = self.executor.clone();
let config = self.config.clone();
let cancel_token = self.cancel_token.clone();
tokio::spawn(async move {
tracing::info!(
"[SchedulerService] 启动心跳循环,间隔: {} 秒",
config.heartbeat_interval_secs
);
let mut ticker = interval(Duration::from_secs(config.heartbeat_interval_secs));
loop {
tokio::select! {
_ = ticker.tick() => {
if let Err(e) = Self::poll_and_execute(
&scheduler,
&executor,
&db,
config.max_tasks_per_poll,
)
.await
{
tracing::error!("[SchedulerService] 轮询任务失败: {}", e);
}
}
_ = cancel_token.cancelled() => {
tracing::info!("[SchedulerService] 收到取消信号,停止心跳循环");
break;
}
}
}
tracing::info!("[SchedulerService] 心跳循环已停止");
});
}
/// 停止心跳循环
pub fn stop(&self) {
tracing::info!("[SchedulerService] 请求停止心跳循环");
self.cancel_token.cancel();
}
/// 轮询并执行到期任务
async fn poll_and_execute(
scheduler: &Arc<AgentScheduler>,
executor: &Arc<AgentExecutor>,
db: &DbConnection,
max_tasks: usize,
) -> Result<(), String> {
// 1. 获取到期任务
let due_tasks = scheduler.get_due_tasks(max_tasks).await?;
if due_tasks.is_empty() {
tracing::debug!("[SchedulerService] 没有到期任务");
return Ok(());
}
tracing::info!("[SchedulerService] 发现 {} 个到期任务", due_tasks.len());
// 2. 执行每个任务
for task in due_tasks {
let task_id = task.id.clone();
let task_name = task.name.clone();
tracing::info!("[SchedulerService] 开始执行任务: {} ({})", task_name, task_id);
// 标记为运行中
if let Err(e) = scheduler.mark_task_running(&task_id).await {
tracing::error!("[SchedulerService] 标记任务运行失败: {} - {}", task_id, e);
continue;
}
// 执行任务
match executor.execute(&task, db).await {
Ok(result) => {
// 标记为完成
if let Err(e) = scheduler.mark_task_completed(&task_id, Some(result)).await {
tracing::error!(
"[SchedulerService] 标记任务完成失败: {} - {}",
task_id,
e
);
}
}
Err(e) => {
tracing::error!("[SchedulerService] 任务执行失败: {} - {}", task_id, e);
// 标记为失败
if let Err(mark_err) = scheduler.mark_task_failed(&task_id, e).await {
tracing::error!(
"[SchedulerService] 标记任务失败失败: {} - {}",
task_id,
mark_err
);
}
// TODO: 实现重试逻辑
// 如果任务可以重试,重新调度任务
}
}
}
Ok(())
}
/// 获取调度器引用
pub fn scheduler(&self) -> Arc<AgentScheduler> {
self.scheduler.clone()
}
/// 获取执行器引用
pub fn executor(&self) -> Arc<AgentExecutor> {
self.executor.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
use proxycast_scheduler::{ScheduledTask, SchedulerDao};
use rusqlite::Connection;
use std::sync::{Arc, Mutex};
fn setup_test_db() -> DbConnection {
let conn = Connection::open_in_memory().unwrap();
SchedulerDao::create_tables(&conn).unwrap();
Arc::new(Mutex::new(conn))
}
#[test]
fn test_scheduler_service_creation() {
let db = setup_test_db();
let config = SchedulerServiceConfig::default();
let service = SchedulerService::new(db, config);
assert!(Arc::strong_count(&service.scheduler) >= 1);
assert!(Arc::strong_count(&service.executor) >= 1);
}
#[test]
fn test_scheduler_service_config_default() {
let config = SchedulerServiceConfig::default();
assert_eq!(config.heartbeat_interval_secs, 30);
assert_eq!(config.max_tasks_per_poll, 10);
assert!(config.enabled);
}
#[tokio::test]
async fn test_poll_and_execute_no_tasks() {
let db = setup_test_db();
let scheduler = Arc::new(AgentScheduler::new(db.clone()));
let executor = Arc::new(AgentExecutor::new());
let result = SchedulerService::poll_and_execute(&scheduler, &executor, &db, 10).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_poll_and_execute_with_due_task() {
let db = setup_test_db();
let scheduler = Arc::new(AgentScheduler::new(db.clone()));
let executor = Arc::new(AgentExecutor::new());
// 创建一个到期任务
let past = Utc::now() - chrono::Duration::hours(1);
let task = ScheduledTask::new(
"Test Task".to_string(),
"agent_chat".to_string(),
serde_json::json!({"prompt": "test"}),
"openai".to_string(),
"gpt-4".to_string(),
past,
);
scheduler.create_task(task).await.unwrap();
// 轮询并执行(会因为缺少凭证而失败,但不应 panic)
let result = SchedulerService::poll_and_execute(&scheduler, &executor, &db, 10).await;
// 即使执行失败,poll_and_execute 也应该返回 Ok
assert!(result.is_ok());
}
}
+18
View File
@@ -10,10 +10,12 @@ use crate::agent::AsterAgentState;
use crate::database;
use crate::telemetry;
use crate::tray::{TrayIconStatus, TrayManager, TrayStateSnapshot};
use proxycast_scheduler::AgentScheduler;
use proxycast_services::aster_session_store::ProxyCastSessionStore;
use proxycast_services::provider_pool_service::ProviderPoolService;
use proxycast_services::token_cache_service::TokenCacheService;
use super::scheduler_service::{SchedulerService, SchedulerServiceConfig};
use super::types::{AppState, LogState, TrayManagerState};
/// Tauri setup hook
@@ -77,6 +79,22 @@ pub fn setup_app(
.expect("Failed to initialize default skill repos");
}
// 初始化调度器数据库表
if let Err(e) = AgentScheduler::init_tables(&db) {
tracing::error!("[启动] 调度器表初始化失败: {}", e);
} else {
tracing::info!("[启动] 调度器表初始化成功");
}
// 启动调度器服务
let scheduler_config = SchedulerServiceConfig::default();
let scheduler_service = SchedulerService::new(db.clone(), scheduler_config);
scheduler_service.start(db.clone());
tracing::info!("[启动] 调度器服务已启动");
// 将调度器服务注册为 Tauri 状态,以便后续访问
app.manage(Arc::new(scheduler_service));
// 自动启动服务器
let app_handle = app.handle().clone();
tauri::async_runtime::spawn(async move {
@@ -0,0 +1,109 @@
//! 电商差评回复 Tauri 命令
//!
//! 提供电商差评回复的专门接口,封装 Skill 执行逻辑
use serde::{Deserialize, Serialize};
use tauri::State;
use crate::agent::AsterAgentState;
use crate::database::DbConnection;
use crate::commands::skill_exec_cmd::{execute_skill, SkillExecutionResult};
/// 电商差评回复请求
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EcommerceReviewReplyRequest {
/// 电商平台
pub platform: String,
/// 差评链接
pub review_url: String,
/// 回复语气
pub tone: String,
/// 回复长度
pub length: String,
/// 自定义模板 (可选)
pub template: Option<String>,
/// AI 模型
pub model: Option<String>,
/// 执行 ID (可选)
pub execution_id: Option<String>,
}
/// 执行电商差评回复
///
/// 这是一个便捷接口,封装了 execute_skill 的调用
///
/// # Arguments
/// * `app_handle` - Tauri AppHandle
/// * `db` - 数据库连接
/// * `aster_state` - Aster Agent 状态
/// * `request` - 电商差评回复请求
///
/// # Returns
/// * `Ok(SkillExecutionResult)` - 执行结果
/// * `Err(String)` - 错误信息
#[tauri::command]
pub async fn execute_ecommerce_review_reply(
app_handle: tauri::AppHandle,
db: State<'_, DbConnection>,
aster_state: State<'_, AsterAgentState>,
request: EcommerceReviewReplyRequest,
) -> Result<SkillExecutionResult, String> {
tracing::info!(
"[execute_ecommerce_review_reply] 开始执行: platform={}, url={}",
request.platform,
request.review_url
);
// 构建用户输入
let user_input = format!(
"平台: {}\n差评链接: {}\n回复语气: {}\n回复长度: {}{}",
request.platform,
request.review_url,
request.tone,
request.length,
request
.template
.as_ref()
.map(|t| format!("\n自定义模板: {}", t))
.unwrap_or_default()
);
// 调用通用的 execute_skill
execute_skill(
app_handle,
db,
aster_state,
"ecommerce-review-reply".to_string(),
user_input,
Some("anthropic".to_string()), // 优先使用 Anthropic
request.model,
request.execution_id,
None, // session_id
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_request_serialization() {
let request = EcommerceReviewReplyRequest {
platform: "taobao".to_string(),
review_url: "https://example.com/review/123".to_string(),
tone: "sincere".to_string(),
length: "medium".to_string(),
template: None,
model: Some("claude-sonnet-4-5".to_string()),
execution_id: None,
};
let json = serde_json::to_string(&request).unwrap();
let deserialized: EcommerceReviewReplyRequest =
serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.platform, "taobao");
assert_eq!(deserialized.tone, "sincere");
}
}
+1
View File
@@ -9,6 +9,7 @@ pub mod connect_cmd;
pub mod connection_cmd;
pub mod content_cmd;
pub mod context_memory;
pub mod ecommerce_review_reply_cmd;
pub mod external_tools_cmd;
pub mod general_chat_cmd;
+12 -4
View File
@@ -14,6 +14,7 @@ import { withI18nPatch } from "./i18n/withI18nPatch";
import { SplashScreen } from "./components/SplashScreen";
import { AppSidebar } from "./components/AppSidebar";
import { SettingsPage } from "./components/settings";
import { SettingsPageV2 } from "./components/settings-v2";
import { ApiServerPage } from "./components/api-server/ApiServerPage";
import { ProviderPoolPage } from "./components/provider-pool";
import { ToolsPage } from "./components/tools/ToolsPage";
@@ -352,10 +353,17 @@ function AppContent() {
<PluginsPage />
</PageWrapper>
{/* Settings 页面 */}
<PageWrapper $isActive={currentPage === "settings"}>
<SettingsPage />
</PageWrapper>
{/* Settings 页面 - 使用新版 V2 布局 */}
<div
style={{
flex: 1,
minHeight: 0,
display: currentPage === "settings" ? "flex" : "none",
flexDirection: "column",
}}
>
<SettingsPageV2 />
</div>
{/* 动态插件页面已移除 */}
</>
@@ -88,6 +88,7 @@ export const ProviderListItem: React.FC<ProviderListItemProps> = ({
{/* Provider 图标 */}
<ProviderIcon
providerType={provider.id}
fallbackText={provider.name}
size={24}
className="flex-shrink-0"
data-testid="provider-icon"
@@ -170,6 +170,7 @@ export const ProviderSetting: React.FC<ProviderSettingProps> = ({
{/* 图标 */}
<ProviderIcon
providerType={provider.id}
fallbackText={provider.name}
size={40}
className="flex-shrink-0"
data-testid="provider-icon"
@@ -0,0 +1,181 @@
/**
* 设置侧边栏组件
*
* 显示分组的设置导航菜单
* 参考 LobeHub 的 SettingsSidebar 设计
*/
import styled from 'styled-components';
import { ChevronDown } from 'lucide-react';
import { useState } from 'react';
import {
useSettingsCategory,
CategoryGroup,
} from '../hooks/useSettingsCategory';
import { SettingsTabs, SettingsGroupKey } from '@/types/settings';
const SidebarContainer = styled.aside`
width: 240px;
min-width: 240px;
height: 100%;
background: hsl(var(--card));
border-right: 1px solid hsl(var(--border));
overflow-y: auto;
padding: 16px 8px;
&::-webkit-scrollbar {
width: 4px;
}
&::-webkit-scrollbar-track {
background: transparent;
}
&::-webkit-scrollbar-thumb {
background: hsl(var(--border));
border-radius: 2px;
}
`;
const GroupContainer = styled.div`
margin-bottom: 8px;
`;
const GroupHeader = styled.button<{ $expanded: boolean }>`
display: flex;
align-items: center;
justify-content: space-between;
width: 100%;
padding: 8px 12px;
border: none;
background: transparent;
cursor: pointer;
font-size: 12px;
font-weight: 500;
color: hsl(var(--muted-foreground));
text-transform: uppercase;
letter-spacing: 0.5px;
svg {
width: 14px;
height: 14px;
transition: transform 0.2s;
transform: rotate(${({ $expanded }) => ($expanded ? '0deg' : '-90deg')});
}
&:hover {
color: hsl(var(--foreground));
}
`;
const GroupItems = styled.div<{ $expanded: boolean }>`
display: ${({ $expanded }) => ($expanded ? 'flex' : 'none')};
flex-direction: column;
gap: 2px;
padding: 4px 0;
`;
const NavItem = styled.button<{ $active: boolean }>`
display: flex;
align-items: center;
gap: 10px;
width: 100%;
padding: 10px 12px;
border: none;
border-radius: 8px;
background: ${({ $active }) =>
$active ? 'hsl(var(--accent))' : 'transparent'};
cursor: pointer;
font-size: 14px;
color: ${({ $active }) =>
$active ? 'hsl(var(--foreground))' : 'hsl(var(--muted-foreground))'};
transition: all 0.15s;
text-align: left;
&:hover {
background: hsl(var(--accent));
color: hsl(var(--foreground));
}
svg {
width: 18px;
height: 18px;
flex-shrink: 0;
}
`;
const ItemLabel = styled.span`
flex: 1;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
`;
const ExperimentalBadge = styled.span`
font-size: 10px;
padding: 2px 6px;
background: hsl(var(--destructive) / 0.1);
color: hsl(var(--destructive));
border-radius: 4px;
flex-shrink: 0;
`;
interface SettingsSidebarProps {
activeTab: SettingsTabs;
onTabChange: (tab: SettingsTabs) => void;
}
export function SettingsSidebar({
activeTab,
onTabChange,
}: SettingsSidebarProps) {
const categoryGroups = useSettingsCategory();
// 默认展开所有分组
const [expandedGroups, setExpandedGroups] = useState<
Record<SettingsGroupKey, boolean>
>({
[SettingsGroupKey.Account]: true,
[SettingsGroupKey.General]: true,
[SettingsGroupKey.Agent]: true,
[SettingsGroupKey.System]: true,
});
const toggleGroup = (key: SettingsGroupKey) => {
setExpandedGroups((prev) => ({
...prev,
[key]: !prev[key],
}));
};
return (
<SidebarContainer>
{categoryGroups.map((group) => (
<GroupContainer key={group.key}>
<GroupHeader
$expanded={expandedGroups[group.key] ?? true}
onClick={() => toggleGroup(group.key)}
>
{group.title}
<ChevronDown />
</GroupHeader>
<GroupItems $expanded={expandedGroups[group.key] ?? true}>
{group.items.map((item) => (
<NavItem
key={item.key}
$active={activeTab === item.key}
onClick={() => onTabChange(item.key)}
>
<item.icon />
<ItemLabel>{item.label}</ItemLabel>
{item.experimental && (
<ExperimentalBadge>实验</ExperimentalBadge>
)}
</NavItem>
))}
</GroupItems>
</GroupContainer>
))}
</SidebarContainer>
);
}
@@ -0,0 +1,292 @@
/**
* 设置页面主布局组件
*
* 采用左侧边栏 + 右侧内容的布局
* 参考 LobeHub 的设置布局设计
*/
import { useState, ReactNode } from 'react';
import styled from 'styled-components';
import { SettingsSidebar } from './SettingsSidebar';
import { SettingsTabs } from '@/types/settings';
// 外观设置(迁移自原 GeneralSettings)
import { GeneralSettings } from '../../settings/GeneralSettings';
// 网络代理
import { ProxySettings } from '../../settings/ProxySettings';
// 数据存储
import { DirectorySettings } from '../../settings/DirectorySettings';
import { QuotaSettings } from '../../settings/QuotaSettings';
// 安全设置
import { TlsSettings } from '../../settings/TlsSettings';
import { RemoteManagementSettings } from '../../settings/RemoteManagementSettings';
// 外部工具
import { ExternalToolsSettings } from '../../settings/ExternalToolsSettings';
// 实验功能
import { ExperimentalSettings } from '../../settings/ExperimentalSettings';
// 开发者
import { DeveloperSettings } from '../../settings/DeveloperSettings';
// 关于
import { AboutSection } from '../../settings/AboutSection';
// 连接设置
import { ConnectionsSettings } from '../../settings/ConnectionsSettings';
// 扩展设置
import { ExtensionsSettings } from '../../settings/ExtensionsSettings';
import { SettingHeader } from '../features/SettingHeader';
const LayoutContainer = styled.div`
display: flex;
height: 100%;
background: hsl(var(--background));
`;
const ContentContainer = styled.main`
flex: 1;
overflow-y: auto;
padding: 24px 32px;
&::-webkit-scrollbar {
width: 6px;
}
&::-webkit-scrollbar-track {
background: transparent;
}
&::-webkit-scrollbar-thumb {
background: hsl(var(--border));
border-radius: 3px;
}
`;
const ContentWrapper = styled.div`
max-width: 800px;
`;
const PlaceholderPage = styled.div`
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
height: 300px;
color: hsl(var(--muted-foreground));
text-align: center;
p {
margin-top: 8px;
font-size: 14px;
}
`;
/**
* 渲染设置内容
*/
function renderSettingsContent(tab: SettingsTabs): ReactNode {
switch (tab) {
// 账号组
case SettingsTabs.Profile:
return (
<>
<SettingHeader title="个人资料" />
<PlaceholderPage>
<p>个人资料设置</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
case SettingsTabs.Stats:
return (
<>
<SettingHeader title="数据统计" />
<PlaceholderPage>
<p>使用统计信息</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
// 通用组
case SettingsTabs.Appearance:
return (
<>
<SettingHeader title="外观" />
<GeneralSettings />
</>
);
case SettingsTabs.ChatAppearance:
return (
<>
<SettingHeader title="聊天外观" />
<PlaceholderPage>
<p>聊天气泡样式设置</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
case SettingsTabs.Hotkeys:
return (
<>
<SettingHeader title="快捷键" />
<PlaceholderPage>
<p>快捷键设置</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
// 智能体组
case SettingsTabs.Providers:
return (
<>
<SettingHeader title="AI 服务商" />
<ConnectionsSettings />
</>
);
case SettingsTabs.Assistant:
return (
<>
<SettingHeader title="助理服务" />
<PlaceholderPage>
<p>助理配置</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
case SettingsTabs.Skills:
return (
<>
<SettingHeader title="技能管理" />
<ExtensionsSettings />
</>
);
case SettingsTabs.Memory:
return (
<>
<SettingHeader title="记忆设置" />
<PlaceholderPage>
<p>记忆管理</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
case SettingsTabs.ImageGen:
return (
<>
<SettingHeader title="绘画服务" />
<PlaceholderPage>
<p>绘画服务配置</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
case SettingsTabs.Voice:
return (
<>
<SettingHeader title="语音服务" />
<PlaceholderPage>
<p>语音服务配置</p>
<p>即将推出...</p>
</PlaceholderPage>
</>
);
// 系统组
case SettingsTabs.Proxy:
return (
<>
<SettingHeader title="网络代理" />
<ProxySettings />
</>
);
case SettingsTabs.Storage:
return (
<>
<SettingHeader title="数据存储" />
<div className="space-y-4">
<DirectorySettings />
<QuotaSettings />
</div>
</>
);
case SettingsTabs.Security:
return (
<>
<SettingHeader title="安全设置" />
<div className="space-y-6">
<TlsSettings />
<RemoteManagementSettings />
</div>
</>
);
case SettingsTabs.ExternalTools:
return (
<>
<SettingHeader title="外部工具" />
<ExternalToolsSettings />
</>
);
case SettingsTabs.Experimental:
return (
<>
<SettingHeader title="实验功能" />
<ExperimentalSettings />
</>
);
case SettingsTabs.Developer:
return (
<>
<SettingHeader title="开发者" />
<DeveloperSettings />
</>
);
case SettingsTabs.About:
return (
<>
<SettingHeader title="关于" />
<AboutSection />
</>
);
default:
return (
<PlaceholderPage>
<p>页面不存在</p>
</PlaceholderPage>
);
}
}
/**
* 设置页面主组件
*/
export function SettingsLayoutV2() {
const [activeTab, setActiveTab] = useState<SettingsTabs>(
SettingsTabs.Appearance
);
return (
<LayoutContainer>
<SettingsSidebar activeTab={activeTab} onTabChange={setActiveTab} />
<ContentContainer>
<ContentWrapper>{renderSettingsContent(activeTab)}</ContentWrapper>
</ContentContainer>
</LayoutContainer>
);
}
export default SettingsLayoutV2;
@@ -0,0 +1,55 @@
/**
* 设置页头组件
*
* 显示设置页面标题和可选的额外操作
*/
import styled from 'styled-components';
import { ReactNode } from 'react';
const HeaderContainer = styled.div`
display: flex;
flex-direction: column;
gap: 16px;
margin-bottom: 24px;
`;
const TitleRow = styled.div`
display: flex;
align-items: center;
justify-content: space-between;
`;
const Title = styled.h1`
font-size: 24px;
font-weight: 600;
color: hsl(var(--foreground));
margin: 0;
`;
const Divider = styled.div`
height: 1px;
background: hsl(var(--border));
`;
interface SettingHeaderProps {
/** 页面标题 */
title: ReactNode;
/** 额外的操作区域 */
extra?: ReactNode;
}
/**
* 设置页头组件
*/
export function SettingHeader({ title, extra }: SettingHeaderProps) {
return (
<HeaderContainer>
<TitleRow>
<Title>{title}</Title>
{extra}
</TitleRow>
<Divider />
</HeaderContainer>
);
}
@@ -0,0 +1,188 @@
/**
* 设置分类 Hook
*
* 定义设置页面的分组和导航项
* 参考 LobeHub 的 useCategory 设计
*/
import { useMemo } from 'react';
import { useTranslation } from 'react-i18next';
import {
User,
BarChart3,
Palette,
MessageSquare,
Keyboard,
Brain,
Bot,
Blocks,
BrainCircuit,
Image,
Mic,
Globe,
Database,
Shield,
Wrench,
FlaskConical,
Code,
Info,
LucideIcon,
} from 'lucide-react';
import { SettingsGroupKey, SettingsTabs } from '@/types/settings';
/**
* 分类项定义
*/
export interface CategoryItem {
key: SettingsTabs;
label: string;
icon: LucideIcon;
experimental?: boolean;
}
/**
* 分类组定义
*/
export interface CategoryGroup {
key: SettingsGroupKey;
title: string;
items: CategoryItem[];
}
/**
* 设置分类 Hook
*
* 返回按分组组织的设置导航项
*/
export function useSettingsCategory(): CategoryGroup[] {
const { t } = useTranslation();
return useMemo(() => {
const groups: CategoryGroup[] = [];
// 账号组
groups.push({
key: SettingsGroupKey.Account,
title: t('settings.group.account', '账号'),
items: [
{
key: SettingsTabs.Profile,
label: t('settings.tab.profile', '个人资料'),
icon: User,
},
{
key: SettingsTabs.Stats,
label: t('settings.tab.stats', '数据统计'),
icon: BarChart3,
},
],
});
// 通用组
groups.push({
key: SettingsGroupKey.General,
title: t('settings.group.general', '通用'),
items: [
{
key: SettingsTabs.Appearance,
label: t('settings.tab.appearance', '外观'),
icon: Palette,
},
{
key: SettingsTabs.ChatAppearance,
label: t('settings.tab.chatAppearance', '聊天外观'),
icon: MessageSquare,
},
{
key: SettingsTabs.Hotkeys,
label: t('settings.tab.hotkeys', '快捷键'),
icon: Keyboard,
},
],
});
// 智能体组
groups.push({
key: SettingsGroupKey.Agent,
title: t('settings.group.agent', '智能体'),
items: [
{
key: SettingsTabs.Providers,
label: t('settings.tab.providers', 'AI 服务商'),
icon: Brain,
},
{
key: SettingsTabs.Assistant,
label: t('settings.tab.assistant', '助理服务'),
icon: Bot,
},
{
key: SettingsTabs.Skills,
label: t('settings.tab.skills', '技能管理'),
icon: Blocks,
},
{
key: SettingsTabs.Memory,
label: t('settings.tab.memory', '记忆设置'),
icon: BrainCircuit,
},
{
key: SettingsTabs.ImageGen,
label: t('settings.tab.imageGen', '绘画服务'),
icon: Image,
},
{
key: SettingsTabs.Voice,
label: t('settings.tab.voice', '语音服务'),
icon: Mic,
},
],
});
// 系统组
groups.push({
key: SettingsGroupKey.System,
title: t('settings.group.system', '系统'),
items: [
{
key: SettingsTabs.Proxy,
label: t('settings.tab.proxy', '网络代理'),
icon: Globe,
},
{
key: SettingsTabs.Storage,
label: t('settings.tab.storage', '数据存储'),
icon: Database,
},
{
key: SettingsTabs.Security,
label: t('settings.tab.security', '安全设置'),
icon: Shield,
},
{
key: SettingsTabs.ExternalTools,
label: t('settings.tab.externalTools', '外部工具'),
icon: Wrench,
},
{
key: SettingsTabs.Experimental,
label: t('settings.tab.experimental', '实验功能'),
icon: FlaskConical,
experimental: true,
},
{
key: SettingsTabs.Developer,
label: t('settings.tab.developer', '开发者'),
icon: Code,
},
{
key: SettingsTabs.About,
label: t('settings.tab.about', '关于'),
icon: Info,
},
],
});
return groups;
}, [t]);
}
+11
View File
@@ -0,0 +1,11 @@
/**
* 设置页面 V2 导出
*
* 新版设置页面,采用 LobeHub 风格的侧边栏布局
*/
export { SettingsLayoutV2 as SettingsPageV2 } from './_layout';
export { SettingsSidebar } from './_layout/SettingsSidebar';
export { SettingHeader } from './features/SettingHeader';
export { useSettingsCategory } from './hooks/useSettingsCategory';
export type { CategoryItem, CategoryGroup } from './hooks/useSettingsCategory';
+24 -6
View File
@@ -226,6 +226,8 @@ const iconComponents: Record<string, ComponentType<SVGProps<SVGSVGElement>>> = {
interface ProviderIconProps {
/** Provider 类型或 ID */
providerType: string;
/** 回退文本(未命中图标时用于生成首字母) */
fallbackText?: string;
/** 图标大小,支持数字(px)或字符串 */
size?: number | string;
/** 额外的 CSS 类名 */
@@ -249,6 +251,7 @@ interface ProviderIconProps {
*/
export const ProviderIcon: React.FC<ProviderIconProps> = ({
providerType,
fallbackText,
size = 24,
className,
showFallback = true,
@@ -282,12 +285,27 @@ export const ProviderIcon: React.FC<ProviderIconProps> = ({
// Fallback:显示首字母
if (showFallback) {
const initials = providerType
.split(/[-_]/)
.map((word) => word[0])
.join("")
.toUpperCase()
.slice(0, 2);
const source = fallbackText?.trim() || providerType;
const words = source
.split(/[\s-_]+/)
.map((word) => word.trim())
.filter((word) => word.length > 0);
const primaryWord = words[0] || source;
const primaryChars = Array.from(primaryWord);
const firstDigit = primaryChars.find((char) => /\d/.test(char));
const initials =
words.length >= 2
? words
.slice(0, 2)
.map((word) => Array.from(word)[0] || "")
.join("")
.toUpperCase()
: firstDigit && primaryChars.length > 0
? `${primaryChars[0]}${firstDigit}`.toUpperCase()
: primaryChars.slice(0, 2).join("").toUpperCase();
const fallbackFontSize =
typeof size === "number" ? `${Math.max(size * 0.5, 12)}px` : "0.5em";
return (
+69
View File
@@ -0,0 +1,69 @@
/**
* 电商差评回复 API
*
* 封装电商差评回复相关的 Tauri 命令调用
*/
import { safeInvoke } from "@/lib/dev-bridge";
/**
* 电商差评回复请求参数
*/
export interface EcommerceReviewReplyRequest {
/** 电商平台 */
platform: "taobao" | "jd" | "pinduoduo";
/** 差评链接 */
reviewUrl: string;
/** 回复语气 */
tone: "polite" | "sincere" | "professional";
/** 回复长度 */
length: "short" | "medium" | "long";
/** 自定义模板 (可选) */
template?: string;
/** AI 模型 (可选) */
model?: string;
/** 执行 ID (可选) */
executionId?: string;
}
/**
* Skill 执行结果
*/
export interface SkillExecutionResult {
/** 是否成功 */
success: boolean;
/** 最终输出 */
output?: string;
/** 错误信息 */
error?: string;
/** 已完成的步骤结果 */
stepsCompleted: Array<{
stepId: string;
stepName: string;
success: boolean;
output?: string;
error?: string;
}>;
}
/**
* 电商差评回复 API
*/
export const ecommerceReviewReplyApi = {
/**
* 执行电商差评回复
*
* @param request - 请求参数
* @returns 执行结果
*/
async executeReviewReply(
request: EcommerceReviewReplyRequest
): Promise<SkillExecutionResult> {
return safeInvoke(
"execute_ecommerce_review_reply",
request as unknown as Record<string, unknown>
);
},
};
export default ecommerceReviewReplyApi;
@@ -0,0 +1,411 @@
/**
* 电商差评回复 - 配置向导
*
* 5 步配置流程:
* 1. 选择电商平台
* 2. 配置登录凭证
* 3. 配置 AI 模型
* 4. 设置回复规则
* 5. 测试运行
*/
import { useState } from "react";
import styled from "styled-components";
import type { EcommerceConfig } from "./index";
const Container = styled.div`
display: flex;
flex-direction: column;
gap: 24px;
padding: 24px;
background-color: hsl(var(--card));
border-radius: 8px;
border: 1px solid hsl(var(--border));
`;
const StepIndicator = styled.div`
display: flex;
gap: 8px;
margin-bottom: 16px;
`;
const StepDot = styled.div<{ active: boolean; completed: boolean }>`
width: 32px;
height: 32px;
border-radius: 50%;
display: flex;
align-items: center;
justify-content: center;
font-size: 14px;
font-weight: 500;
background-color: ${(props) =>
props.completed
? "hsl(var(--primary))"
: props.active
? "hsl(var(--primary))"
: "hsl(var(--muted))"};
color: ${(props) =>
props.completed || props.active
? "hsl(var(--primary-foreground))"
: "hsl(var(--muted-foreground))"};
`;
const StepTitle = styled.h3`
font-size: 18px;
font-weight: 600;
color: hsl(var(--foreground));
margin-bottom: 16px;
`;
const FormGroup = styled.div`
display: flex;
flex-direction: column;
gap: 8px;
`;
const Label = styled.label`
font-size: 14px;
font-weight: 500;
color: hsl(var(--foreground));
`;
const Select = styled.select`
padding: 8px 12px;
border-radius: 6px;
border: 1px solid hsl(var(--border));
background-color: hsl(var(--background));
color: hsl(var(--foreground));
font-size: 14px;
&:focus {
outline: none;
border-color: hsl(var(--primary));
}
`;
const Input = styled.input`
padding: 8px 12px;
border-radius: 6px;
border: 1px solid hsl(var(--border));
background-color: hsl(var(--background));
color: hsl(var(--foreground));
font-size: 14px;
&:focus {
outline: none;
border-color: hsl(var(--primary));
}
`;
const TextArea = styled.textarea`
padding: 8px 12px;
border-radius: 6px;
border: 1px solid hsl(var(--border));
background-color: hsl(var(--background));
color: hsl(var(--foreground));
font-size: 14px;
min-height: 100px;
resize: vertical;
&:focus {
outline: none;
border-color: hsl(var(--primary));
}
`;
const ButtonGroup = styled.div`
display: flex;
gap: 12px;
justify-content: flex-end;
margin-top: 16px;
`;
const Button = styled.button<{ variant?: "primary" | "secondary" }>`
padding: 8px 16px;
border-radius: 6px;
font-size: 14px;
font-weight: 500;
cursor: pointer;
transition: all 0.2s;
${(props) =>
props.variant === "primary"
? `
background-color: hsl(var(--primary));
color: hsl(var(--primary-foreground));
border: none;
&:hover {
opacity: 0.9;
}
`
: `
background-color: transparent;
color: hsl(var(--foreground));
border: 1px solid hsl(var(--border));
&:hover {
background-color: hsl(var(--muted));
}
`}
&:disabled {
opacity: 0.5;
cursor: not-allowed;
}
`;
const Hint = styled.p`
font-size: 12px;
color: hsl(var(--muted-foreground));
margin-top: 4px;
`;
interface GuideStepProps {
currentStep: number;
onStepChange: (step: number) => void;
onComplete: (config: EcommerceConfig) => void;
}
export function GuideStep({
currentStep,
onStepChange,
onComplete,
}: GuideStepProps) {
const [platform, setPlatform] = useState<"taobao" | "jd" | "pinduoduo">(
"taobao"
);
const [credType, setCredType] = useState<"cookie" | "password">("cookie");
const [credValue, setCredValue] = useState("");
const [primaryModel, setPrimaryModel] = useState("claude-sonnet-4-5");
const [fallbackModel, setFallbackModel] = useState("");
const [tone, setTone] = useState<"polite" | "sincere" | "professional">(
"sincere"
);
const [length, setLength] = useState<"short" | "medium" | "long">("medium");
const [template, setTemplate] = useState("");
const handleNext = () => {
if (currentStep < 4) {
onStepChange(currentStep + 1);
} else {
// 完成配置
const config: EcommerceConfig = {
platform,
credentials: {
type: credType,
value: credValue,
},
aiModel: {
primary: primaryModel,
fallback: fallbackModel || undefined,
},
replyRules: {
tone,
length,
template: template || undefined,
},
};
onComplete(config);
}
};
const handleBack = () => {
if (currentStep > 0) {
onStepChange(currentStep - 1);
}
};
const renderStep = () => {
switch (currentStep) {
case 0:
return (
<>
<StepTitle>步骤 1: 选择电商平台</StepTitle>
<FormGroup>
<Label>电商平台</Label>
<Select
value={platform}
onChange={(e) =>
setPlatform(e.target.value as typeof platform)
}
>
<option value="taobao">淘宝/天猫</option>
<option value="jd">京东</option>
<option value="pinduoduo">拼多多</option>
</Select>
<Hint>选择您要处理差评的电商平台</Hint>
</FormGroup>
</>
);
case 1:
return (
<>
<StepTitle>步骤 2: 配置登录凭证</StepTitle>
<FormGroup>
<Label>凭证类型</Label>
<Select
value={credType}
onChange={(e) => setCredType(e.target.value as typeof credType)}
>
<option value="cookie">Cookie</option>
<option value="password">账号密码</option>
</Select>
</FormGroup>
<FormGroup>
<Label>
{credType === "cookie" ? "Cookie 值" : "账号密码"}
</Label>
<TextArea
value={credValue}
onChange={(e) => setCredValue(e.target.value)}
placeholder={
credType === "cookie"
? "粘贴浏览器 Cookie..."
: "账号:密码"
}
/>
<Hint>
{credType === "cookie"
? "从浏览器开发者工具中复制 Cookie"
: "输入您的账号和密码,用冒号分隔"}
</Hint>
</FormGroup>
</>
);
case 2:
return (
<>
<StepTitle>步骤 3: 配置 AI 模型</StepTitle>
<FormGroup>
<Label>主模型</Label>
<Select
value={primaryModel}
onChange={(e) => setPrimaryModel(e.target.value)}
>
<option value="claude-sonnet-4-5">Claude Sonnet 4.5</option>
<option value="claude-opus-4">Claude Opus 4</option>
<option value="gpt-4">GPT-4</option>
<option value="gpt-4-turbo">GPT-4 Turbo</option>
</Select>
<Hint>用于生成回复的主要 AI 模型</Hint>
</FormGroup>
<FormGroup>
<Label>降级模型 (可选)</Label>
<Select
value={fallbackModel}
onChange={(e) => setFallbackModel(e.target.value)}
>
<option value="">不使用降级模型</option>
<option value="claude-sonnet-3-5">Claude Sonnet 3.5</option>
<option value="gpt-3.5-turbo">GPT-3.5 Turbo</option>
</Select>
<Hint>主模型不可用时使用的备用模型</Hint>
</FormGroup>
</>
);
case 3:
return (
<>
<StepTitle>步骤 4: 设置回复规则</StepTitle>
<FormGroup>
<Label>回复语气</Label>
<Select
value={tone}
onChange={(e) => setTone(e.target.value as typeof tone)}
>
<option value="polite">礼貌</option>
<option value="sincere">真诚</option>
<option value="professional">专业</option>
</Select>
</FormGroup>
<FormGroup>
<Label>回复长度</Label>
<Select
value={length}
onChange={(e) => setLength(e.target.value as typeof length)}
>
<option value="short">简短 (100-150字)</option>
<option value="medium">中等 (200-300字)</option>
<option value="long">详细 (300-500字)</option>
</Select>
</FormGroup>
<FormGroup>
<Label>自定义模板 (可选)</Label>
<TextArea
value={template}
onChange={(e) => setTemplate(e.target.value)}
placeholder="输入自定义回复模板..."
/>
<Hint>留空则使用默认模板</Hint>
</FormGroup>
</>
);
case 4:
return (
<>
<StepTitle>步骤 5: 测试运行</StepTitle>
<FormGroup>
<Label>测试差评链接</Label>
<Input
type="url"
placeholder="粘贴差评链接进行测试..."
/>
<Hint>输入一个差评链接测试配置是否正常工作</Hint>
</FormGroup>
<FormGroup>
<Label>配置摘要</Label>
<div
style={{
padding: "12px",
backgroundColor: "hsl(var(--muted))",
borderRadius: "6px",
fontSize: "14px",
}}
>
<p>平台: {platform === "taobao" ? "淘宝" : platform === "jd" ? "京东" : "拼多多"}</p>
<p>凭证类型: {credType === "cookie" ? "Cookie" : "账号密码"}</p>
<p>主模型: {primaryModel}</p>
<p>回复语气: {tone === "polite" ? "礼貌" : tone === "sincere" ? "真诚" : "专业"}</p>
<p>回复长度: {length === "short" ? "简短" : length === "medium" ? "中等" : "详细"}</p>
</div>
</FormGroup>
</>
);
default:
return null;
}
};
return (
<Container>
<StepIndicator>
{[0, 1, 2, 3, 4].map((step) => (
<StepDot
key={step}
active={step === currentStep}
completed={step < currentStep}
>
{step + 1}
</StepDot>
))}
</StepIndicator>
{renderStep()}
<ButtonGroup>
{currentStep > 0 && (
<Button onClick={handleBack}>上一步</Button>
)}
<Button variant="primary" onClick={handleNext}>
{currentStep === 4 ? "完成配置" : "下一步"}
</Button>
</ButtonGroup>
</Container>
);
}
@@ -0,0 +1,303 @@
/**
* 电商差评回复 - 执行结果
*
* 展示任务执行结果和生成的回复
*/
import styled from "styled-components";
import type { ReviewTask } from "./index";
const Container = styled.div`
display: flex;
flex-direction: column;
gap: 16px;
padding: 24px;
background-color: hsl(var(--card));
border-radius: 8px;
border: 1px solid hsl(var(--border));
height: 100%;
overflow: hidden;
`;
const Header = styled.div`
display: flex;
justify-content: space-between;
align-items: center;
`;
const Title = styled.h3`
font-size: 18px;
font-weight: 600;
color: hsl(var(--foreground));
`;
const Stats = styled.div`
display: flex;
gap: 16px;
padding: 12px;
background-color: hsl(var(--muted));
border-radius: 6px;
`;
const StatItem = styled.div`
display: flex;
flex-direction: column;
align-items: center;
gap: 4px;
`;
const StatValue = styled.div`
font-size: 24px;
font-weight: 600;
color: hsl(var(--foreground));
`;
const StatLabel = styled.div`
font-size: 12px;
color: hsl(var(--muted-foreground));
`;
const ResultList = styled.div`
flex: 1;
display: flex;
flex-direction: column;
gap: 12px;
overflow-y: auto;
`;
const ResultItem = styled.div`
padding: 16px;
background-color: hsl(var(--background));
border-radius: 6px;
border: 1px solid hsl(var(--border));
`;
const ResultHeader = styled.div`
display: flex;
justify-content: space-between;
align-items: center;
margin-bottom: 12px;
`;
const ResultStatus = styled.div<{ status: ReviewTask["status"] }>`
font-size: 12px;
font-weight: 500;
padding: 4px 8px;
border-radius: 4px;
background-color: ${(props) => {
switch (props.status) {
case "pending":
return "hsl(var(--muted))";
case "processing":
return "hsl(var(--primary) / 0.1)";
case "completed":
return "hsl(142, 76%, 36% / 0.1)";
case "failed":
return "hsl(var(--destructive) / 0.1)";
}
}};
color: ${(props) => {
switch (props.status) {
case "pending":
return "hsl(var(--muted-foreground))";
case "processing":
return "hsl(var(--primary))";
case "completed":
return "hsl(142, 76%, 36%)";
case "failed":
return "hsl(var(--destructive))";
}
}};
`;
const ResultTime = styled.div`
font-size: 12px;
color: hsl(var(--muted-foreground));
`;
const ReviewContent = styled.div`
margin-bottom: 12px;
padding: 12px;
background-color: hsl(var(--muted) / 0.5);
border-radius: 4px;
font-size: 14px;
color: hsl(var(--foreground));
`;
const ReplyContent = styled.div`
padding: 12px;
background-color: hsl(var(--primary) / 0.05);
border-radius: 4px;
border-left: 3px solid hsl(var(--primary));
font-size: 14px;
color: hsl(var(--foreground));
line-height: 1.6;
`;
const ErrorContent = styled.div`
padding: 12px;
background-color: hsl(var(--destructive) / 0.05);
border-radius: 4px;
border-left: 3px solid hsl(var(--destructive));
font-size: 14px;
color: hsl(var(--destructive));
`;
const EmptyState = styled.div`
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
padding: 48px 24px;
text-align: center;
color: hsl(var(--muted-foreground));
`;
const CopyButton = styled.button`
padding: 4px 8px;
border-radius: 4px;
font-size: 12px;
background-color: transparent;
color: hsl(var(--primary));
border: 1px solid hsl(var(--primary));
cursor: pointer;
transition: all 0.2s;
&:hover {
background-color: hsl(var(--primary) / 0.1);
}
`;
interface ResultsProps {
tasks: ReviewTask[];
}
export function Results({ tasks }: ResultsProps) {
const completedTasks = tasks.filter((t) => t.status === "completed");
const failedTasks = tasks.filter((t) => t.status === "failed");
const processingTasks = tasks.filter((t) => t.status === "processing");
const handleCopyReply = (reply: string) => {
navigator.clipboard.writeText(reply);
// TODO: 显示复制成功提示
};
const formatTime = (date: Date) => {
return new Intl.DateTimeFormat("zh-CN", {
hour: "2-digit",
minute: "2-digit",
}).format(date);
};
const getStatusText = (status: ReviewTask["status"]) => {
switch (status) {
case "pending":
return "待处理";
case "processing":
return "处理中";
case "completed":
return "已完成";
case "failed":
return "失败";
}
};
return (
<Container>
<Header>
<Title>执行结果</Title>
</Header>
<Stats>
<StatItem>
<StatValue>{tasks.length}</StatValue>
<StatLabel>总任务</StatLabel>
</StatItem>
<StatItem>
<StatValue>{completedTasks.length}</StatValue>
<StatLabel>已完成</StatLabel>
</StatItem>
<StatItem>
<StatValue>{failedTasks.length}</StatValue>
<StatLabel>失败</StatLabel>
</StatItem>
</Stats>
{tasks.length === 0 ? (
<EmptyState>
<p>暂无结果</p>
<p style={{ fontSize: "12px", marginTop: "8px" }}>
添加任务后查看执行结果
</p>
</EmptyState>
) : (
<ResultList>
{[...tasks]
.reverse()
.filter((task) => task.status !== "pending")
.map((task) => (
<ResultItem key={task.id}>
<ResultHeader>
<ResultStatus status={task.status}>
{getStatusText(task.status)}
</ResultStatus>
<ResultTime>{formatTime(task.createdAt)}</ResultTime>
</ResultHeader>
{task.reviewContent && (
<div>
<div
style={{
fontSize: "12px",
color: "hsl(var(--muted-foreground))",
marginBottom: "4px",
}}
>
差评内容:
</div>
<ReviewContent>{task.reviewContent}</ReviewContent>
</div>
)}
{task.reply && (
<div>
<div
style={{
fontSize: "12px",
color: "hsl(var(--muted-foreground))",
marginBottom: "4px",
display: "flex",
justifyContent: "space-between",
alignItems: "center",
}}
>
<span>生成的回复:</span>
<CopyButton onClick={() => handleCopyReply(task.reply!)}>
复制
</CopyButton>
</div>
<ReplyContent>{task.reply}</ReplyContent>
</div>
)}
{task.error && (
<div>
<div
style={{
fontSize: "12px",
color: "hsl(var(--destructive))",
marginBottom: "4px",
}}
>
错误信息:
</div>
<ErrorContent>{task.error}</ErrorContent>
</div>
)}
</ResultItem>
))}
</ResultList>
)}
</Container>
);
}
@@ -0,0 +1,269 @@
/**
* 电商差评回复 - 任务列表
*
* 管理差评回复任务,支持批量添加和执行
*/
import { useState } from "react";
import styled from "styled-components";
import type { ReviewTask, EcommerceConfig } from "./index";
const Container = styled.div`
display: flex;
flex-direction: column;
gap: 16px;
padding: 24px;
background-color: hsl(var(--card));
border-radius: 8px;
border: 1px solid hsl(var(--border));
`;
const Header = styled.div`
display: flex;
justify-content: space-between;
align-items: center;
`;
const Title = styled.h3`
font-size: 18px;
font-weight: 600;
color: hsl(var(--foreground));
`;
const AddTaskForm = styled.div`
display: flex;
gap: 8px;
`;
const Input = styled.input`
flex: 1;
padding: 8px 12px;
border-radius: 6px;
border: 1px solid hsl(var(--border));
background-color: hsl(var(--background));
color: hsl(var(--foreground));
font-size: 14px;
&:focus {
outline: none;
border-color: hsl(var(--primary));
}
`;
const Button = styled.button<{ variant?: "primary" | "secondary" }>`
padding: 8px 16px;
border-radius: 6px;
font-size: 14px;
font-weight: 500;
cursor: pointer;
transition: all 0.2s;
white-space: nowrap;
${(props) =>
props.variant === "primary"
? `
background-color: hsl(var(--primary));
color: hsl(var(--primary-foreground));
border: none;
&:hover {
opacity: 0.9;
}
`
: `
background-color: transparent;
color: hsl(var(--foreground));
border: 1px solid hsl(var(--border));
&:hover {
background-color: hsl(var(--muted));
}
`}
&:disabled {
opacity: 0.5;
cursor: not-allowed;
}
`;
const TaskList = styled.div`
display: flex;
flex-direction: column;
gap: 8px;
max-height: 500px;
overflow-y: auto;
`;
const TaskItem = styled.div`
display: flex;
align-items: center;
gap: 12px;
padding: 12px;
background-color: hsl(var(--background));
border-radius: 6px;
border: 1px solid hsl(var(--border));
`;
const TaskInfo = styled.div`
flex: 1;
min-width: 0;
`;
const TaskUrl = styled.div`
font-size: 14px;
color: hsl(var(--foreground));
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
`;
const TaskStatus = styled.div<{ status: ReviewTask["status"] }>`
font-size: 12px;
color: ${(props) => {
switch (props.status) {
case "pending":
return "hsl(var(--muted-foreground))";
case "processing":
return "hsl(var(--primary))";
case "completed":
return "hsl(142, 76%, 36%)";
case "failed":
return "hsl(var(--destructive))";
}
}};
`;
const TaskActions = styled.div`
display: flex;
gap: 8px;
`;
const EmptyState = styled.div`
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
padding: 48px 24px;
text-align: center;
color: hsl(var(--muted-foreground));
`;
const ConfigInfo = styled.div`
padding: 12px;
background-color: hsl(var(--muted));
border-radius: 6px;
font-size: 12px;
color: hsl(var(--muted-foreground));
`;
interface TasksProps {
tasks: ReviewTask[];
config: EcommerceConfig;
onAddTask: (reviewUrl: string) => void;
onExecuteTask: (taskId: string) => void;
}
export function Tasks({ tasks, config, onAddTask, onExecuteTask }: TasksProps) {
const [newTaskUrl, setNewTaskUrl] = useState("");
const handleAddTask = () => {
if (newTaskUrl.trim()) {
onAddTask(newTaskUrl.trim());
setNewTaskUrl("");
}
};
const getStatusText = (status: ReviewTask["status"]) => {
switch (status) {
case "pending":
return "待处理";
case "processing":
return "处理中...";
case "completed":
return "已完成";
case "failed":
return "失败";
}
};
const getPlatformName = () => {
switch (config.platform) {
case "taobao":
return "淘宝";
case "jd":
return "京东";
case "pinduoduo":
return "拼多多";
}
};
return (
<Container>
<Header>
<Title>差评任务列表</Title>
</Header>
<ConfigInfo>
当前配置: {getPlatformName()} | {config.aiModel.primary} |{" "}
{config.replyRules.tone === "polite"
? "礼貌"
: config.replyRules.tone === "sincere"
? "真诚"
: "专业"}
语气
</ConfigInfo>
<AddTaskForm>
<Input
type="url"
placeholder="粘贴差评链接..."
value={newTaskUrl}
onChange={(e) => setNewTaskUrl(e.target.value)}
onKeyDown={(e) => {
if (e.key === "Enter") {
handleAddTask();
}
}}
/>
<Button variant="primary" onClick={handleAddTask}>
添加任务
</Button>
</AddTaskForm>
{tasks.length === 0 ? (
<EmptyState>
<p>暂无任务</p>
<p style={{ fontSize: "12px", marginTop: "8px" }}>
粘贴差评链接开始处理
</p>
</EmptyState>
) : (
<TaskList>
{tasks.map((task) => (
<TaskItem key={task.id}>
<TaskInfo>
<TaskUrl title={task.reviewUrl}>{task.reviewUrl}</TaskUrl>
<TaskStatus status={task.status}>
{getStatusText(task.status)}
</TaskStatus>
</TaskInfo>
<TaskActions>
{task.status === "pending" && (
<Button
variant="primary"
onClick={() => onExecuteTask(task.id)}
>
执行
</Button>
)}
{task.status === "failed" && (
<Button onClick={() => onExecuteTask(task.id)}>重试</Button>
)}
</TaskActions>
</TaskItem>
))}
</TaskList>
)}
</Container>
);
}
@@ -0,0 +1,197 @@
/**
* 电商差评回复解决方案 - 主页面
*
* 提供电商差评自动回复功能,支持淘宝、京东、拼多多等平台
*/
import { useState } from "react";
import styled from "styled-components";
import { GuideStep } from "./GuideStep";
import { Tasks } from "./Tasks";
import { Results } from "./Results";
import { ecommerceReviewReplyApi } from "@/lib/api/ecommerce-review-reply";
import type { EcommerceReviewReplyRequest } from "@/lib/api/ecommerce-review-reply";
const Container = styled.div`
display: flex;
flex-direction: column;
height: 100%;
width: 100%;
background-color: hsl(var(--background));
padding: 24px;
`;
const Header = styled.div`
margin-bottom: 24px;
`;
const Title = styled.h1`
font-size: 24px;
font-weight: 600;
color: hsl(var(--foreground));
margin-bottom: 8px;
`;
const Description = styled.p`
font-size: 14px;
color: hsl(var(--muted-foreground));
`;
const Content = styled.div`
flex: 1;
display: flex;
gap: 24px;
min-height: 0;
`;
const LeftPanel = styled.div`
flex: 1;
display: flex;
flex-direction: column;
gap: 16px;
min-width: 0;
`;
const RightPanel = styled.div`
width: 400px;
display: flex;
flex-direction: column;
gap: 16px;
`;
export interface EcommerceConfig {
platform: "taobao" | "jd" | "pinduoduo";
credentials: {
type: "cookie" | "password";
value: string;
};
aiModel: {
primary: string;
fallback?: string;
};
replyRules: {
tone: "polite" | "sincere" | "professional";
length: "short" | "medium" | "long";
template?: string;
};
}
export interface ReviewTask {
id: string;
reviewUrl: string;
reviewContent?: string;
status: "pending" | "processing" | "completed" | "failed";
reply?: string;
error?: string;
createdAt: Date;
}
export default function EcommerceReviewReply() {
const [config, setConfig] = useState<EcommerceConfig | null>(null);
const [tasks, setTasks] = useState<ReviewTask[]>([]);
const [currentStep, setCurrentStep] = useState(0);
const handleConfigComplete = (newConfig: EcommerceConfig) => {
setConfig(newConfig);
setCurrentStep(5); // 完成配置,进入任务管理
};
const handleAddTask = (reviewUrl: string) => {
const newTask: ReviewTask = {
id: Date.now().toString(),
reviewUrl,
status: "pending",
createdAt: new Date(),
};
setTasks([...tasks, newTask]);
};
const handleExecuteTask = async (taskId: string) => {
const task = tasks.find((t) => t.id === taskId);
if (!task || !config) return;
// 更新任务状态为处理中
setTasks(
tasks.map((t) =>
t.id === taskId ? { ...t, status: "processing" as const } : t
)
);
try {
// 构建请求参数
const request: EcommerceReviewReplyRequest = {
platform: config.platform,
reviewUrl: task.reviewUrl,
tone: config.replyRules.tone,
length: config.replyRules.length,
template: config.replyRules.template,
model: config.aiModel.primary,
executionId: taskId,
};
// 调用后端 API
const result = await ecommerceReviewReplyApi.executeReviewReply(request);
// 更新任务状态
setTasks(
tasks.map((t) =>
t.id === taskId
? {
...t,
status: result.success ? ("completed" as const) : ("failed" as const),
reply: result.output,
error: result.error,
}
: t
)
);
} catch (error) {
// 处理错误
setTasks(
tasks.map((t) =>
t.id === taskId
? {
...t,
status: "failed" as const,
error: error instanceof Error ? error.message : "执行失败",
}
: t
)
);
}
};
return (
<Container>
<Header>
<Title>电商差评回复助手</Title>
<Description>
智能分析电商差评并生成专业回复,支持淘宝、京东、拼多多等平台
</Description>
</Header>
<Content>
<LeftPanel>
{!config ? (
<GuideStep
currentStep={currentStep}
onStepChange={setCurrentStep}
onComplete={handleConfigComplete}
/>
) : (
<Tasks
tasks={tasks}
config={config}
onAddTask={handleAddTask}
onExecuteTask={handleExecuteTask}
/>
)}
</LeftPanel>
<RightPanel>
<Results tasks={tasks} />
</RightPanel>
</Content>
</Container>
);
}
+93
View File
@@ -0,0 +1,93 @@
/**
* 设置页面类型定义
*
* 定义设置分组和标签页的枚举
*/
/**
* 设置分组 Key
*/
export enum SettingsGroupKey {
Account = 'account',
General = 'general',
Agent = 'agent',
System = 'system',
}
/**
* 设置标签页
*/
export enum SettingsTabs {
// 账号
Profile = 'profile',
Stats = 'stats',
// 通用
Appearance = 'appearance',
ChatAppearance = 'chat-appearance',
Hotkeys = 'hotkeys',
// 智能体
Providers = 'providers',
Assistant = 'assistant',
Skills = 'skills',
Memory = 'memory',
ImageGen = 'image-gen',
Voice = 'voice',
// 系统
Proxy = 'proxy',
Storage = 'storage',
Security = 'security',
ExternalTools = 'external-tools',
Experimental = 'experimental',
Developer = 'developer',
About = 'about',
}
/**
* 分组信息
*/
export interface SettingsGroupInfo {
key: SettingsGroupKey;
labelKey: string; // i18n key
}
/**
* 标签页信息
*/
export interface SettingsTabInfo {
key: SettingsTabs;
labelKey: string; // i18n key
group: SettingsGroupKey;
experimental?: boolean;
}
/**
* 分组到标签页的映射
*/
export const SETTINGS_GROUPS: Record<SettingsGroupKey, SettingsTabs[]> = {
[SettingsGroupKey.Account]: [SettingsTabs.Profile, SettingsTabs.Stats],
[SettingsGroupKey.General]: [
SettingsTabs.Appearance,
SettingsTabs.ChatAppearance,
SettingsTabs.Hotkeys,
],
[SettingsGroupKey.Agent]: [
SettingsTabs.Providers,
SettingsTabs.Assistant,
SettingsTabs.Skills,
SettingsTabs.Memory,
SettingsTabs.ImageGen,
SettingsTabs.Voice,
],
[SettingsGroupKey.System]: [
SettingsTabs.Proxy,
SettingsTabs.Storage,
SettingsTabs.Security,
SettingsTabs.ExternalTools,
SettingsTabs.Experimental,
SettingsTabs.Developer,
SettingsTabs.About,
],
};