mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
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:
@@ -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
|
||||
@@ -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
@@ -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"
|
||||
}
|
||||
}
|
||||
|
||||
Generated
+28
-56
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -21,6 +21,7 @@ tracing.workspace = true
|
||||
chrono.workspace = true
|
||||
dirs.workspace = true
|
||||
uuid.workspace = true
|
||||
thiserror.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
tempfile.workspace = true
|
||||
|
||||
@@ -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")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
//! Tools 模块
|
||||
//!
|
||||
//! 提供各种工具的包装器和辅助函数
|
||||
|
||||
pub mod browser_tool;
|
||||
|
||||
pub use browser_tool::{BrowserAction, BrowserTool, BrowserToolError, BrowserToolResult};
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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 }
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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(¶ms_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());
|
||||
}
|
||||
}
|
||||
@@ -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("不支持的任务类型"));
|
||||
}
|
||||
}
|
||||
@@ -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};
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -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, "请处理内容: 测试内容, 来自: 测试来源");
|
||||
}
|
||||
}
|
||||
@@ -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()));
|
||||
}
|
||||
}
|
||||
@@ -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 都定义了同名类型)
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
)]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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]);
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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 (
|
||||
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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,
|
||||
],
|
||||
};
|
||||
Reference in New Issue
Block a user