diff --git a/IMPLEMENTATION_PLAN.md b/IMPLEMENTATION_PLAN.md new file mode 100644 index 000000000..53be1fa85 --- /dev/null +++ b/IMPLEMENTATION_PLAN.md @@ -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 diff --git a/docs/design/settings-redesign.md b/docs/design/settings-redesign.md new file mode 100644 index 000000000..f8f280c9a --- /dev/null +++ b/docs/design/settings-redesign.md @@ -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.Profile); + + return ( + + + + {children} + + + ); +} +``` + +### 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>({ + account: true, + general: true, + agent: true, + system: true, + }); + + const toggleGroup = (key: string) => { + setExpandedGroups(prev => ({ + ...prev, + [key]: !prev[key], + })); + }; + + return ( + + {categoryGroups.map((group) => ( + + toggleGroup(group.key)} + > + {group.title} + + + + {group.items.map((item) => ( + onTabChange(item.key)} + > + + {item.label} + {item.experimental && 实验} + + ))} + + + ))} + + ); +} +``` + +### 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 ( + + + {title} + {extra} + + + + ); +} +``` + +## 六、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) diff --git a/package.json b/package.json index 1aceaab57..7233c9751 100644 --- a/package.json +++ b/package.json @@ -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" } } diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 4804d1177..01c02244f 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -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" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 2e5c0499e..77e7c127e 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -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 diff --git a/src-tauri/crates/agent/Cargo.toml b/src-tauri/crates/agent/Cargo.toml index aa8cdf1dd..ec789faa6 100644 --- a/src-tauri/crates/agent/Cargo.toml +++ b/src-tauri/crates/agent/Cargo.toml @@ -21,6 +21,7 @@ tracing.workspace = true chrono.workspace = true dirs.workspace = true uuid.workspace = true +thiserror.workspace = true [dev-dependencies] tempfile.workspace = true diff --git a/src-tauri/crates/agent/src/aster_state.rs b/src-tauri/crates/agent/src/aster_state.rs index bc554514a..177a69ec3 100644 --- a/src-tauri/crates/agent/src/aster_state.rs +++ b/src-tauri/crates/agent/src/aster_state.rs @@ -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>> 中 + /// - `server_info`: MCP 服务器信息 + pub async fn register_mcp_bridge( + &self, + name: String, + description: String, + client: Arc>>, + server_info: Option, + ) -> 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() diff --git a/src-tauri/crates/agent/src/credential_bridge.rs b/src-tauri/crates/agent/src/credential_bridge.rs index 6c3c33050..7dbaaf365 100644 --- a/src-tauri/crates/agent/src/credential_bridge.rs +++ b/src-tauri/crates/agent/src/credential_bridge.rs @@ -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::() { + 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") + ); + } } diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs index df4561e61..49e816349 100644 --- a/src-tauri/crates/agent/src/lib.rs +++ b/src-tauri/crates/agent/src/lib.rs @@ -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}; diff --git a/src-tauri/crates/agent/src/mcp_bridge.rs b/src-tauri/crates/agent/src/mcp_bridge.rs index 275647b52..6e38b2f5f 100644 --- a/src-tauri/crates/agent/src/mcp_bridge.rs +++ b/src-tauri/crates/agent/src/mcp_bridge.rs @@ -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>, + /// ProxyCast MCP 客户端处理器 + handler: Arc, /// 服务器初始化信息 server_info: Option, + /// 请求超时时间 + timeout: Duration, } impl McpBridgeClient { pub fn new( name: String, service: Arc>, + handler: Arc, server_info: Option, ) -> Self { Self { name, service, + handler, server_info, + timeout: Duration::from_secs(60), // 默认超时 60s } } + + /// 发送请求并处理取消和超时 + async fn send_request( + &self, + request: ClientRequest, + cancel_token: CancellationToken, + ) -> Result { + // 发送请求 + 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::() + .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, + cancel_token: CancellationToken, + ) -> Result { + 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 { + 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, + cancel_token: CancellationToken, + ) -> Result { + 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, + cancel_token: CancellationToken, + ) -> Result { + 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, + cancel_token: CancellationToken, + ) -> Result { + 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 { + 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 { + self.handler.subscribe().await + } + + fn get_info(&self) -> Option<&InitializeResult> { + self.server_info.as_ref() + } +} + diff --git a/src-tauri/crates/agent/src/tools/browser_tool.rs b/src-tauri/crates/agent/src/tools/browser_tool.rs new file mode 100644 index 000000000..1de21baaa --- /dev/null +++ b/src-tauri/crates/agent/src/tools/browser_tool.rs @@ -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 }, +} + +/// Browser Tool 结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct BrowserToolResult { + /// 是否成功 + pub success: bool, + /// 输出内容 + pub output: String, + /// 错误信息 + pub error: Option, +} + +/// Browser Tool 包装器 +/// +/// 提供对 Playwright MCP Server 工具的高级封装 +pub struct BrowserTool { + /// MCP 客户端 + mcp_client: Arc>>>, +} + +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, + ) { + let mut guard = self.mcp_client.lock().await; + *guard = Some(client); + } + + /// 执行浏览器动作 + /// + /// # 参数 + /// - `action`: 浏览器动作 + /// + /// # 返回 + /// 返回工具执行结果 + pub async fn execute( + &self, + action: BrowserAction, + ) -> Result { + 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, + tool_name: &str, + arguments: Value, + ) -> Result { + // 将 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::>() + .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 { + self.execute(BrowserAction::Navigate { + url: url.to_string(), + }) + .await + } + + /// 获取页面快照 + pub async fn snapshot(&self) -> Result { + self.execute(BrowserAction::Snapshot).await + } + + /// 点击元素 + pub async fn click(&self, ref_id: &str) -> Result { + self.execute(BrowserAction::Click { + ref_id: ref_id.to_string(), + }) + .await + } + + /// 输入文本 + pub async fn type_text( + &self, + ref_id: &str, + text: &str, + ) -> Result { + self.execute(BrowserAction::Type { + ref_id: ref_id.to_string(), + text: text.to_string(), + }) + .await + } + + /// 截图 + pub async fn screenshot( + &self, + filename: Option, + ) -> Result { + 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()); + } +} diff --git a/src-tauri/crates/agent/src/tools/mod.rs b/src-tauri/crates/agent/src/tools/mod.rs new file mode 100644 index 000000000..9597cdbcb --- /dev/null +++ b/src-tauri/crates/agent/src/tools/mod.rs @@ -0,0 +1,7 @@ +//! Tools 模块 +//! +//! 提供各种工具的包装器和辅助函数 + +pub mod browser_tool; + +pub use browser_tool::{BrowserAction, BrowserTool, BrowserToolError, BrowserToolResult}; diff --git a/src-tauri/crates/core/src/agent/types.rs b/src-tauri/crates/core/src/agent/types.rs index 7153db92a..5408043a7 100644 --- a/src-tauri/crates/core/src/agent/types.rs +++ b/src-tauri/crates/core/src/agent/types.rs @@ -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 { diff --git a/src-tauri/crates/core/src/database/dao/api_key_provider.rs b/src-tauri/crates/core/src/database/dao/api_key_provider.rs index 8bedf3da0..ec4b042dc 100644 --- a/src-tauri/crates/core/src/database/dao/api_key_provider.rs +++ b/src-tauri/crates/core/src/database/dao/api_key_provider.rs @@ -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; diff --git a/src-tauri/crates/core/src/database/migration_v3.rs b/src-tauri/crates/core/src/database/migration_v3.rs new file mode 100644 index 000000000..b725231ca --- /dev/null +++ b/src-tauri/crates/core/src/database/migration_v3.rs @@ -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, +} + +/// 执行 Playwright MCP Server 迁移 +/// +/// 迁移步骤: +/// 1. 检查是否已迁移 +/// 2. 检查是否已存在 playwright 服务器 +/// 3. 如果不存在,创建默认配置 +/// 4. 标记迁移完成 +pub fn migrate_playwright_mcp_server(conn: &Connection) -> Result { + // 检查是否已经迁移过 + 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 { + // 创建 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(()) +} diff --git a/src-tauri/crates/core/src/database/mod.rs b/src-tauri/crates/core/src/database/mod.rs index 71bcff7c3..a8959d661 100644 --- a/src-tauri/crates/core/src/database/mod.rs +++ b/src-tauri/crates/core/src/database/mod.rs @@ -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 { } } + // 执行 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))) } diff --git a/src-tauri/crates/core/src/models/provider_type.rs b/src-tauri/crates/core/src/models/provider_type.rs index f0c765483..422f0da58 100644 --- a/src-tauri/crates/core/src/models/provider_type.rs +++ b/src-tauri/crates/core/src/models/provider_type.rs @@ -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"); diff --git a/src-tauri/crates/providers/src/providers/openai_custom.rs b/src-tauri/crates/providers/src/providers/openai_custom.rs index 949ebb96b..3af2319f9 100644 --- a/src-tauri/crates/providers/src/providers/openai_custom.rs +++ b/src-tauri/crates/providers/src/providers/openai_custom.rs @@ -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; diff --git a/src-tauri/crates/scheduler/Cargo.toml b/src-tauri/crates/scheduler/Cargo.toml new file mode 100644 index 000000000..2415661fb --- /dev/null +++ b/src-tauri/crates/scheduler/Cargo.toml @@ -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 } diff --git a/src-tauri/crates/scheduler/src/batch.rs b/src-tauri/crates/scheduler/src/batch.rs new file mode 100644 index 000000000..f33a870f4 --- /dev/null +++ b/src-tauri/crates/scheduler/src/batch.rs @@ -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, + + /// 模板变量 + pub variables: HashMap, + + /// 任务元数据 (用于追踪和识别) + #[serde(default)] + pub metadata: HashMap, +} + +/// 单个任务结果 +#[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, + + /// 错误信息 (如果失败) + #[serde(skip_serializing_if = "Option::is_none")] + pub error: Option, + + /// 使用 token 数 + #[serde(default)] + pub usage: TokenUsage, + + /// 开始时间 + pub started_at: chrono::DateTime, + + /// 完成时间 + pub completed_at: Option>, +} + +/// 任务状态 +#[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, + + /// 批量任务选项 + #[serde(default)] + pub options: BatchOptions, + + /// 批量任务状态 + pub status: BatchTaskStatus, + + /// 任务结果 + #[serde(default)] + pub results: Vec, + + /// 创建时间 + pub created_at: chrono::DateTime, + + /// 开始时间 + #[serde(skip_serializing_if = "Option::is_none")] + pub started_at: Option>, + + /// 完成时间 + #[serde(skip_serializing_if = "Option::is_none")] + pub completed_at: Option>, +} + +/// 批量任务状态 +#[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, + 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); + } +} diff --git a/src-tauri/crates/scheduler/src/batch_dao.rs b/src-tauri/crates/scheduler/src/batch_dao.rs new file mode 100644 index 000000000..90c17dd70 --- /dev/null +++ b/src-tauri/crates/scheduler/src/batch_dao.rs @@ -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> { + 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 = row.get(6)?; + let created_at: String = row.get(7)?; + let started_at: Option = row.get(8)?; + let completed_at: Option = 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> { + 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 = row.get(6)?; + let created_at: String = row.get(7)?; + let started_at: Option = row.get(8)?; + let completed_at: Option = 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 { + 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> { + 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> { + 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 { + 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); + } +} diff --git a/src-tauri/crates/scheduler/src/dao.rs b/src-tauri/crates/scheduler/src/dao.rs new file mode 100644 index 000000000..47ba5f2c5 --- /dev/null +++ b/src-tauri/crates/scheduler/src/dao.rs @@ -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, 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, 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 { + let rows = conn.execute("DELETE FROM scheduled_tasks WHERE id = ?", [id])?; + Ok(rows > 0) + } + + /// 获取到期任务 + pub fn get_due_tasks(conn: &Connection, limit: usize) -> Result, 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 { + 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 = 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()); + } +} diff --git a/src-tauri/crates/scheduler/src/executor.rs b/src-tauri/crates/scheduler/src/executor.rs new file mode 100644 index 000000000..084fc8c4b --- /dev/null +++ b/src-tauri/crates/scheduler/src/executor.rs @@ -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; +} + +/// Agent 任务执行器 +/// +/// 通过 CredentialBridge 选择凭证,调用 Aster Agent 执行任务 +pub struct AgentExecutor { + credential_bridge: Arc, +} + +impl AgentExecutor { + /// 创建新的执行器实例 + pub fn new() -> Self { + Self { + credential_bridge: Arc::new(CredentialBridge::new()), + } + } + + /// 使用自定义的 CredentialBridge 创建执行器 + pub fn with_credential_bridge(credential_bridge: Arc) -> 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 { + 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 { + // 从任务参数中提取对话内容 + 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 { + 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 { + 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("不支持的任务类型")); + } +} diff --git a/src-tauri/crates/scheduler/src/lib.rs b/src-tauri/crates/scheduler/src/lib.rs new file mode 100644 index 000000000..7a22818d2 --- /dev/null +++ b/src-tauri/crates/scheduler/src/lib.rs @@ -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}; diff --git a/src-tauri/crates/scheduler/src/scheduler.rs b/src-tauri/crates/scheduler/src/scheduler.rs new file mode 100644 index 000000000..1a65b97d0 --- /dev/null +++ b/src-tauri/crates/scheduler/src/scheduler.rs @@ -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; + + /// 获取任务 + async fn get_task(&self, id: &str) -> Result, String>; + + /// 查询任务列表 + async fn list_tasks(&self, filter: TaskFilter) -> Result, String>; + + /// 更新任务 + async fn update_task(&self, task: ScheduledTask) -> Result<(), String>; + + /// 删除任务 + async fn delete_task(&self, id: &str) -> Result; + + /// 获取到期任务 + async fn get_due_tasks(&self, limit: usize) -> Result, String>; + + /// 标记任务为运行中 + async fn mark_task_running(&self, id: &str) -> Result<(), String>; + + /// 标记任务为完成 + async fn mark_task_completed( + &self, + id: &str, + result: Option, + ) -> 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 { + 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, 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, 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 { + 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, 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, + ) -> 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"); + } +} diff --git a/src-tauri/crates/scheduler/src/template.rs b/src-tauri/crates/scheduler/src/template.rs new file mode 100644 index 000000000..8ee217a9a --- /dev/null +++ b/src-tauri/crates/scheduler/src/template.rs @@ -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, + + /// 模型名称 + pub model: String, + + /// 系统提示词 + pub system_prompt: Option, + + /// 用户消息模板 (支持变量替换,例如 "{{variable_name}}") + pub user_message_template: String, + + /// 温度参数 + pub temperature: Option, + + /// 最大 tokens + pub max_tokens: Option, + + /// 创建时间 + pub created_at: chrono::DateTime, + + /// 更新时间 + pub updated_at: chrono::DateTime, +} + +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 { + 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, "请处理内容: 测试内容, 来自: 测试来源"); + } +} diff --git a/src-tauri/crates/scheduler/src/types.rs b/src-tauri/crates/scheduler/src/types.rs new file mode 100644 index 000000000..617d53f62 --- /dev/null +++ b/src-tauri/crates/scheduler/src/types.rs @@ -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, + /// 任务类型(标识要执行的操作) + 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, + /// 实际执行完成时间(可选) + pub completed_at: Option, + /// 执行结果(可选) + pub result: Option, + /// 错误信息(如果执行失败) + pub error_message: Option, + /// 重试次数 + 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, + ) -> 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) { + 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, + /// Provider 类型 + pub provider_type: Option, + /// 任务类型 + pub task_type: Option, + /// 是否只查询到期的任务 + pub only_due: bool, + /// 限制返回数量 + pub limit: Option, +} + +#[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())); + } +} diff --git a/src-tauri/crates/server/src/handlers/api.rs b/src-tauri/crates/server/src/handlers/api.rs index 85c6c5cf4..6b6d22111 100644 --- a/src-tauri/crates/server/src/handlers/api.rs +++ b/src-tauri/crates/server/src/handlers/api.rs @@ -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, 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 { - 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(); diff --git a/src-tauri/crates/server/src/handlers/api_key_provider_utils.rs b/src-tauri/crates/server/src/handlers/api_key_provider_utils.rs new file mode 100644 index 000000000..45e2e784e --- /dev/null +++ b/src-tauri/crates/server/src/handlers/api_key_provider_utils.rs @@ -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 { + 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::() + .or_else(|_| normalized.replace('_', "-").parse::()); + + 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 { + 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")); + } +} diff --git a/src-tauri/crates/server/src/handlers/batch_api.rs b/src-tauri/crates/server/src/handlers/batch_api.rs new file mode 100644 index 000000000..364d7f3a5 --- /dev/null +++ b/src-tauri/crates/server/src/handlers/batch_api.rs @@ -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, + + /// 批量任务选项 + #[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, +} + +/// 批量任务详情响应 +#[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, + Json(request): Json, +) -> 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, + Path(id): Path, +) -> 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) -> 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, + Path(id): Path, +) -> 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, + Json(template): Json, +) -> 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) -> 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, Path(id): Path) -> 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, Path(id): Path) -> Response { + // TODO: 从数据库删除模板 + state + .logs + .write() + .await + .add("info", &format!("[BATCH] 删除模板: id={}", id)); + + (StatusCode::NO_CONTENT, ()).into_response() +} diff --git a/src-tauri/crates/server/src/handlers/credentials_api.rs b/src-tauri/crates/server/src/handlers/credentials_api.rs index adb7195d1..17d2cc57d 100644 --- a/src-tauri/crates/server/src/handlers/credentials_api.rs +++ b/src-tauri/crates/server/src/handlers/credentials_api.rs @@ -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, 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 { - 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 /// /// 支持多种凭证类型: diff --git a/src-tauri/crates/server/src/handlers/mod.rs b/src-tauri/crates/server/src/handlers/mod.rs index e2e437859..2dca34b28 100644 --- a/src-tauri/crates/server/src/handlers/mod.rs +++ b/src-tauri/crates/server/src/handlers/mod.rs @@ -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 都定义了同名类型) diff --git a/src-tauri/crates/server/src/lib.rs b/src-tauri/crates/server/src/lib.rs index 821603522..d1db14a73 100644 --- a/src-tauri/crates/server/src/lib.rs +++ b/src-tauri/crates/server/src/lib.rs @@ -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); diff --git a/src-tauri/crates/services/src/api_key_provider_service.rs b/src-tauri/crates/services/src/api_key_provider_service.rs index e772ac581..851455083 100644 --- a/src-tauri/crates/services/src/api_key_provider_service.rs +++ b/src-tauri/crates/services/src/api_key_provider_service.rs @@ -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::() + }) + .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 { - 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()); diff --git a/src-tauri/crates/services/src/lib.rs b/src-tauri/crates/services/src/lib.rs index 4078ae6b5..e854c7108 100644 --- a/src-tauri/crates/services/src/lib.rs +++ b/src-tauri/crates/services/src/lib.rs @@ -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; diff --git a/src-tauri/crates/services/src/provider_pool_service.rs b/src-tauri/crates/services/src/provider_pool_service.rs index 74ca3c455..2849ee468 100644 --- a/src-tauri/crates/services/src/provider_pool_service.rs +++ b/src-tauri/crates/services/src/provider_pool_service.rs @@ -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, 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, check_model_name: Option, ) -> Result { - 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, 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 { - 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, 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, source: proxycast_core::models::provider_pool_model::CredentialSource, ) -> Result { - 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 + ); + } } diff --git a/src-tauri/crates/services/src/provider_type_mapping.rs b/src-tauri/crates/services/src/provider_type_mapping.rs new file mode 100644 index 000000000..66cee36ef --- /dev/null +++ b/src-tauri/crates/services/src/provider_type_mapping.rs @@ -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 { + 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 { + 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")); + } +} diff --git a/src-tauri/crates/skills/src/ecommerce_review_reply.rs b/src-tauri/crates/skills/src/ecommerce_review_reply.rs new file mode 100644 index 000000000..e52dbff1f --- /dev/null +++ b/src-tauri/crates/skills/src/ecommerce_review_reply.rs @@ -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, +} + +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, + /// 识别的问题类型 + pub problem_type: Option, + /// 生成的回复 + pub reply: Option, + /// 错误信息 + pub error: Option, +} + +#[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); + } +} diff --git a/src-tauri/crates/skills/src/lib.rs b/src-tauri/crates/skills/src/lib.rs index 641b9dab9..9b136b53f 100644 --- a/src-tauri/crates/skills/src/lib.rs +++ b/src-tauri/crates/skills/src/lib.rs @@ -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, diff --git a/src-tauri/crates/websocket/Cargo.toml b/src-tauri/crates/websocket/Cargo.toml index e9f8f55f3..3400d2a11 100644 --- a/src-tauri/crates/websocket/Cargo.toml +++ b/src-tauri/crates/websocket/Cargo.toml @@ -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 diff --git a/src-tauri/crates/websocket/src/handler.rs b/src-tauri/crates/websocket/src/handler.rs index 9eb20307b..ef9895e76 100644 --- a/src-tauri/crates/websocket/src/handler.rs +++ b/src-tauri/crates/websocket/src/handler.rs @@ -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>, + /// RPC 处理器状态 + pub rpc_state: RpcHandlerState, } impl WsHandlerState { /// 创建新的处理器状态 - pub fn new(config: WsConfig, api_key: String, logs: Arc>) -> Self { + pub fn new( + config: WsConfig, + api_key: String, + logs: Arc>, + db: Option, + scheduler: Option, + ) -> 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; + } + } } } } diff --git a/src-tauri/crates/websocket/src/handlers/mod.rs b/src-tauri/crates/websocket/src/handlers/mod.rs new file mode 100644 index 000000000..cbae84808 --- /dev/null +++ b/src-tauri/crates/websocket/src/handlers/mod.rs @@ -0,0 +1,9 @@ +//! WebSocket RPC 处理器 +//! +//! 实现 RPC 请求处理逻辑 + +pub mod rpc_handler; + +pub use rpc_handler::{ + parse_rpc_request, serialize_rpc_response, RpcHandler, RpcHandlerState, +}; diff --git a/src-tauri/crates/websocket/src/handlers/rpc_handler.rs b/src-tauri/crates/websocket/src/handlers/rpc_handler.rs new file mode 100644 index 000000000..bbe70cd41 --- /dev/null +++ b/src-tauri/crates/websocket/src/handlers/rpc_handler.rs @@ -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>>, + /// Agent 调度器(可选) + pub scheduler: Arc>>, + /// 日志存储 + pub logs: Arc>, +} + +impl RpcHandlerState { + /// 创建新的 RPC 处理器状态 + pub fn new( + db: Option, + scheduler: Option, + logs: Arc>, + ) -> 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, + ) -> Result { + 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, + ) -> Result { + 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, + ) -> Result { + 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 { + // 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, + ) -> Result { + 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 { + // 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, + ) -> Result { + 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 { + serde_json::from_str(msg).map_err(|e| RpcError::parse_error(format!("Invalid JSON: {}", e))) +} + +/// 序列化 RPC 响应 +pub fn serialize_rpc_response(resp: &GatewayRpcResponse) -> Result { + 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")); + } +} diff --git a/src-tauri/crates/websocket/src/lib.rs b/src-tauri/crates/websocket/src/lib.rs index 01cfacb1c..559b34d15 100644 --- a/src-tauri/crates/websocket/src/lib.rs +++ b/src-tauri/crates/websocket/src/lib.rs @@ -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, diff --git a/src-tauri/crates/websocket/src/protocol.rs b/src-tauri/crates/websocket/src/protocol.rs new file mode 100644 index 000000000..d0acc6bf5 --- /dev/null +++ b/src-tauri/crates/websocket/src/protocol.rs @@ -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, +} + +/// 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(skip_serializing_if = "Option::is_none")] + pub error: Option, +} + +/// 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, +} + +impl RpcError { + /// 创建解析错误 + pub fn parse_error(message: impl Into) -> Self { + Self { + code: -32700, + message: message.into(), + data: None, + } + } + + /// 创建无效请求错误 + pub fn invalid_request(message: impl Into) -> Self { + Self { + code: -32600, + message: message.into(), + data: None, + } + } + + /// 创建方法未找到错误 + pub fn method_not_found(method: impl Into) -> Self { + Self { + code: -32601, + message: format!("Method not found: {}", method.into()), + data: None, + } + } + + /// 创建无效参数错误 + pub fn invalid_params(message: impl Into) -> Self { + Self { + code: -32602, + message: message.into(), + data: None, + } + } + + /// 创建内部错误 + pub fn internal_error(message: impl Into) -> Self { + Self { + code: -32603, + message: message.into(), + data: None, + } + } +} + +/// Agent 运行参数 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AgentRunParams { + /// 会话 ID(可选,用于连续对话) + pub session_id: Option, + /// 用户消息 + pub message: String, + /// 模型名称(可选) + pub model: Option, + /// 系统提示词(可选) + pub system_prompt: Option, + /// 温度参数(可选) + pub temperature: Option, + /// 最大 token 数(可选) + pub max_tokens: Option, + /// 是否流式响应 + #[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>, +} + +/// 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, + /// Token 使用量(如果已完成) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage: Option, +} + +/// 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, + /// Token 使用量(如果已完成) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage: Option, +} + +/// 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, +} + +/// 会话详情结果 +#[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, + /// 消息数量 + 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, +} + +/// 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, + /// 下次运行时间 + #[serde(skip_serializing_if = "Option::is_none")] + pub next_run: Option, +} + +/// 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")); + } +} diff --git a/src-tauri/src/app/commands/config.rs b/src-tauri/src/app/commands/config.rs index 5484a99cf..6134ebf84 100644 --- a/src-tauri/src/app/commands/config.rs +++ b/src-tauri/src/app/commands/config.rs @@ -203,6 +203,7 @@ pub async fn update_provider_env_vars( api_host: String, api_key: Option, ) -> 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::().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::() { + 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() + )] + ); + } +} diff --git a/src-tauri/src/app/mod.rs b/src-tauri/src/app/mod.rs index b215766c5..c718879d0 100644 --- a/src-tauri/src/app/mod.rs +++ b/src-tauri/src/app/mod.rs @@ -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::*; diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 2f1681ec2..bcbb2ea2f 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -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, diff --git a/src-tauri/src/app/scheduler_service.rs b/src-tauri/src/app/scheduler_service.rs new file mode 100644 index 000000000..17c21ca7b --- /dev/null +++ b/src-tauri/src/app/scheduler_service.rs @@ -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, + executor: Arc, + 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, + executor: &Arc, + 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 { + self.scheduler.clone() + } + + /// 获取执行器引用 + pub fn executor(&self) -> Arc { + 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()); + } +} diff --git a/src-tauri/src/app/setup.rs b/src-tauri/src/app/setup.rs index 90595180e..933f1400f 100644 --- a/src-tauri/src/app/setup.rs +++ b/src-tauri/src/app/setup.rs @@ -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 { diff --git a/src-tauri/src/commands/ecommerce_review_reply_cmd.rs b/src-tauri/src/commands/ecommerce_review_reply_cmd.rs new file mode 100644 index 000000000..13404842f --- /dev/null +++ b/src-tauri/src/commands/ecommerce_review_reply_cmd.rs @@ -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, + /// AI 模型 + pub model: Option, + /// 执行 ID (可选) + pub execution_id: Option, +} + +/// 执行电商差评回复 +/// +/// 这是一个便捷接口,封装了 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 { + 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"); + } +} diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index c70054af4..9e529c014 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -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; diff --git a/src/App.tsx b/src/App.tsx index 40adb8a82..b2b1e80a1 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -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() { - {/* Settings 页面 */} - - - + {/* Settings 页面 - 使用新版 V2 布局 */} +
+ +
{/* 动态插件页面已移除 */} diff --git a/src/components/provider-pool/api-key/ProviderListItem.tsx b/src/components/provider-pool/api-key/ProviderListItem.tsx index 7a24c887b..1ea2a6cf3 100644 --- a/src/components/provider-pool/api-key/ProviderListItem.tsx +++ b/src/components/provider-pool/api-key/ProviderListItem.tsx @@ -88,6 +88,7 @@ export const ProviderListItem: React.FC = ({ {/* Provider 图标 */} = ({ {/* 图标 */} ` + 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.Account]: true, + [SettingsGroupKey.General]: true, + [SettingsGroupKey.Agent]: true, + [SettingsGroupKey.System]: true, + }); + + const toggleGroup = (key: SettingsGroupKey) => { + setExpandedGroups((prev) => ({ + ...prev, + [key]: !prev[key], + })); + }; + + return ( + + {categoryGroups.map((group) => ( + + toggleGroup(group.key)} + > + {group.title} + + + + {group.items.map((item) => ( + onTabChange(item.key)} + > + + {item.label} + {item.experimental && ( + 实验 + )} + + ))} + + + ))} + + ); +} diff --git a/src/components/settings-v2/_layout/index.tsx b/src/components/settings-v2/_layout/index.tsx new file mode 100644 index 000000000..f1057c4cf --- /dev/null +++ b/src/components/settings-v2/_layout/index.tsx @@ -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 ( + <> + + +

个人资料设置

+

即将推出...

+
+ + ); + + case SettingsTabs.Stats: + return ( + <> + + +

使用统计信息

+

即将推出...

+
+ + ); + + // 通用组 + case SettingsTabs.Appearance: + return ( + <> + + + + ); + + case SettingsTabs.ChatAppearance: + return ( + <> + + +

聊天气泡样式设置

+

即将推出...

+
+ + ); + + case SettingsTabs.Hotkeys: + return ( + <> + + +

快捷键设置

+

即将推出...

+
+ + ); + + // 智能体组 + case SettingsTabs.Providers: + return ( + <> + + + + ); + + case SettingsTabs.Assistant: + return ( + <> + + +

助理配置

+

即将推出...

+
+ + ); + + case SettingsTabs.Skills: + return ( + <> + + + + ); + + case SettingsTabs.Memory: + return ( + <> + + +

记忆管理

+

即将推出...

+
+ + ); + + case SettingsTabs.ImageGen: + return ( + <> + + +

绘画服务配置

+

即将推出...

+
+ + ); + + case SettingsTabs.Voice: + return ( + <> + + +

语音服务配置

+

即将推出...

+
+ + ); + + // 系统组 + case SettingsTabs.Proxy: + return ( + <> + + + + ); + + case SettingsTabs.Storage: + return ( + <> + +
+ + +
+ + ); + + case SettingsTabs.Security: + return ( + <> + +
+ + +
+ + ); + + case SettingsTabs.ExternalTools: + return ( + <> + + + + ); + + case SettingsTabs.Experimental: + return ( + <> + + + + ); + + case SettingsTabs.Developer: + return ( + <> + + + + ); + + case SettingsTabs.About: + return ( + <> + + + + ); + + default: + return ( + +

页面不存在

+
+ ); + } +} + +/** + * 设置页面主组件 + */ +export function SettingsLayoutV2() { + const [activeTab, setActiveTab] = useState( + SettingsTabs.Appearance + ); + + return ( + + + + {renderSettingsContent(activeTab)} + + + ); +} + +export default SettingsLayoutV2; diff --git a/src/components/settings-v2/features/SettingHeader.tsx b/src/components/settings-v2/features/SettingHeader.tsx new file mode 100644 index 000000000..c1caed06d --- /dev/null +++ b/src/components/settings-v2/features/SettingHeader.tsx @@ -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 ( + + + {title} + {extra} + + + + ); +} diff --git a/src/components/settings-v2/hooks/useSettingsCategory.ts b/src/components/settings-v2/hooks/useSettingsCategory.ts new file mode 100644 index 000000000..9c1942292 --- /dev/null +++ b/src/components/settings-v2/hooks/useSettingsCategory.ts @@ -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]); +} diff --git a/src/components/settings-v2/index.ts b/src/components/settings-v2/index.ts new file mode 100644 index 000000000..70bfee796 --- /dev/null +++ b/src/components/settings-v2/index.ts @@ -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'; diff --git a/src/icons/providers/index.tsx b/src/icons/providers/index.tsx index dedc1206a..d5dce005a 100644 --- a/src/icons/providers/index.tsx +++ b/src/icons/providers/index.tsx @@ -226,6 +226,8 @@ const iconComponents: Record>> = { interface ProviderIconProps { /** Provider 类型或 ID */ providerType: string; + /** 回退文本(未命中图标时用于生成首字母) */ + fallbackText?: string; /** 图标大小,支持数字(px)或字符串 */ size?: number | string; /** 额外的 CSS 类名 */ @@ -249,6 +251,7 @@ interface ProviderIconProps { */ export const ProviderIcon: React.FC = ({ providerType, + fallbackText, size = 24, className, showFallback = true, @@ -282,12 +285,27 @@ export const ProviderIcon: React.FC = ({ // 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 ( diff --git a/src/lib/api/ecommerce-review-reply.ts b/src/lib/api/ecommerce-review-reply.ts new file mode 100644 index 000000000..472d1bf9c --- /dev/null +++ b/src/lib/api/ecommerce-review-reply.ts @@ -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 { + return safeInvoke( + "execute_ecommerce_review_reply", + request as unknown as Record + ); + }, +}; + +export default ecommerceReviewReplyApi; diff --git a/src/solutions/ecommerce-review-reply/GuideStep.tsx b/src/solutions/ecommerce-review-reply/GuideStep.tsx new file mode 100644 index 000000000..d84db2759 --- /dev/null +++ b/src/solutions/ecommerce-review-reply/GuideStep.tsx @@ -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 ( + <> + 步骤 1: 选择电商平台 + + + + 选择您要处理差评的电商平台 + + + ); + + case 1: + return ( + <> + 步骤 2: 配置登录凭证 + + + + + + +