mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
chore: bump version to 0.46.0
- 修复 reasoning_content 字段缺失的编译错误 - 精简 README 为 AI Agent 创作工具平台 - 修复 lint 和格式错误
This commit is contained in:
@@ -2,7 +2,7 @@
|
||||
|
||||
# ProxyCast 🚀
|
||||
|
||||
**把你的 AI 客户端额度用到任何地方**
|
||||
**AI Agent 创作工具平台**
|
||||
|
||||
[](https://www.gnu.org/licenses/gpl-3.0)
|
||||
[](https://tauri.app/)
|
||||
@@ -13,219 +13,53 @@
|
||||
|
||||
---
|
||||
|
||||
## 🤔 这个工具能帮你做什么?
|
||||
|
||||
**场景一:换个更好用的 IDE**
|
||||
> 我有 Kiro 账号,可以用 Claude 系列模型,但 Kiro IDE 不太顺手。我想用 Claude Code 或 Cursor 来写代码,但又不想额外付费买 API。
|
||||
|
||||
**场景二:把额度分享给其他工具**
|
||||
> Claude Code 这个月额度还剩很多,与其浪费不如转给 Cherry Studio 聊天用,或者给我的 AI Agent 项目提供 API 接口。
|
||||
|
||||
**场景三:统一管理多个 AI 账号**
|
||||
> 我有 Kiro、Gemini CLI、通义千问好几个账号,想统一管理,哪个有额度就用哪个。
|
||||
|
||||
**ProxyCast 就是解决这些问题的工具** —— 它把你已有的 AI 客户端凭证转换成标准 OpenAI API,让任何支持 OpenAI 接口的工具都能使用。
|
||||
|
||||
---
|
||||
|
||||
## 💡 工作原理
|
||||
|
||||
```
|
||||
你的 AI 客户端凭证 ProxyCast 任意 OpenAI 兼容工具
|
||||
┌─────────────────┐ ┌─────────────┐ ┌─────────────────────┐
|
||||
│ Kiro OAuth │ │ │ │ Claude Code │
|
||||
│ Gemini OAuth │ ───▶ │ 本地 API │ ───▶ │ Cherry Studio │
|
||||
│ 通义千问 OAuth │ │ 代理服务 │ │ Cursor / Cline │
|
||||
│ ... │ │ │ │ 你的 AI Agent │
|
||||
└─────────────────┘ └─────────────┘ └─────────────────────┘
|
||||
```
|
||||
|
||||
> **💡 与 AIClient-2-API 的区别**
|
||||
>
|
||||
> ProxyCast 是 [AIClient-2-API](https://github.com/justlovemaki/AIClient-2-API) 的桌面版本,提供更友好的图形界面和一键操作体验,无需命令行配置。
|
||||
|
||||
---
|
||||
|
||||
## ✨ 核心特性
|
||||
|
||||
### 🎯 多 Provider 统一管理
|
||||
- **Kiro** - 通过 OAuth 使用 Claude 系列模型(Opus 4.5、Sonnet 4.5、Sonnet 4、Haiku 4.5)
|
||||
- **Gemini CLI** - 通过 OAuth 使用 Gemini 模型
|
||||
- **Gemini API Key** - 多账号负载均衡,支持模型排除
|
||||
- **通义千问** - 通过 OAuth 使用 Qwen3 Coder Plus
|
||||
- **Antigravity** - 通过 OAuth 使用 Claude 模型
|
||||
- **Vertex AI** - Google Cloud AI 平台,支持模型别名
|
||||
- **OpenAI 自定义** - 配置自定义 OpenAI 兼容 API
|
||||
- **Claude 自定义** - 配置自定义 Claude API
|
||||
|
||||
### 🖥️ 友好的图形界面
|
||||
- **Dashboard** - 服务状态监控、API 测试面板
|
||||
- **Provider 管理** - 一键加载凭证、Token 刷新、默认 Provider 切换
|
||||
- **设置页面** - 服务器配置、端口设置、API Key 管理
|
||||
- **日志查看** - 实时日志记录、操作追踪
|
||||
|
||||
### 🔄 智能凭证管理
|
||||
- 自动检测凭证文件变化(每 5 秒)
|
||||
- 一键读取本地 OAuth 凭证
|
||||
- Token 过期自动刷新
|
||||
- 环境变量导出(.env 格式)
|
||||
- **配额超限自动切换** - 自动切换到下一个可用凭证
|
||||
- **预览模型回退** - 主模型配额用尽时尝试预览版本
|
||||
- **Per-Key 代理** - 为每个凭证单独配置代理
|
||||
|
||||
### 🔐 安全与管理
|
||||
- **HTTPS 部署** - 当前版本不内置 TLS,请使用反向代理进行 HTTPS 终止
|
||||
- **远程管理 API** - 通过 API 远程管理配置和凭证
|
||||
- **访问控制** - 支持 localhost 限制和密钥认证
|
||||
|
||||
### 🔌 多路由支持
|
||||
- 支持 `/api/provider/{provider}/v1/*` 路由模式
|
||||
- 模型映射 - 将请求模型映射到 Provider 支持的模型
|
||||
- 管理端点代理 - 代理认证和账户功能
|
||||
|
||||
### 🌐 完整 API 兼容
|
||||
- `/v1/chat/completions` - OpenAI Chat API
|
||||
- `/v1/models` - 模型列表
|
||||
- `/v1/messages` - Anthropic Messages API
|
||||
- `/v1/messages/count_tokens` - Token 计数
|
||||
- `/health` - 健康检查
|
||||
- `/ready` - 就绪检查
|
||||
- `/api/provider/{provider}/v1/*` - Provider 路由
|
||||
- `/v0/management/*` - 远程管理 API
|
||||
- `/v0/management/backup` - 触发数据库备份
|
||||
- `/v0/management/restore` - 从备份恢复
|
||||
|
||||
---
|
||||
|
||||
## 📸 界面截图
|
||||
|
||||
### 仪表盘 - 系统状态与监控
|
||||

|
||||
|
||||
### 凭证池 - 多凭证管理与配额查询
|
||||

|
||||
|
||||
### 路由管理 - 智能路由规则和容错策略
|
||||

|
||||
|
||||
### 配置管理 - 客户端配置切换
|
||||

|
||||
|
||||
### 扩展 - MCP/Prompts/Skills 管理
|
||||

|
||||
|
||||
### API Server - 服务控制与 API 测试
|
||||

|
||||
|
||||
### 设置 - 应用参数和偏好
|
||||

|
||||
- **多 Provider 统一管理** - 支持 Kiro、Gemini、通义千问、Antigravity、Vertex AI 等多种 AI 服务
|
||||
- **智能凭证管理** - 自动检测凭证变化、Token 自动刷新、配额超限自动切换
|
||||
- **完整 API 兼容** - 支持 OpenAI Chat API 和 Anthropic Messages API
|
||||
- **友好图形界面** - Dashboard 监控、Provider 管理、日志查看
|
||||
|
||||
---
|
||||
|
||||
## 🚀 快速开始
|
||||
|
||||
### 下载安装
|
||||
### 安装
|
||||
|
||||
#### macOS (推荐使用 Homebrew)
|
||||
#### macOS (Homebrew)
|
||||
|
||||
```bash
|
||||
brew tap aiclientproxy/tap
|
||||
brew install --cask proxycast
|
||||
```
|
||||
|
||||
更新版本:
|
||||
```bash
|
||||
brew upgrade --cask proxycast
|
||||
```
|
||||
|
||||
#### 手动下载
|
||||
|
||||
从 [Releases](https://github.com/aiclientproxy/proxycast/releases) 页面下载对应平台的安装包:
|
||||
从 [Releases](https://github.com/aiclientproxy/proxycast/releases) 下载对应平台安装包。
|
||||
|
||||
- **macOS (Apple Silicon)**: `proxycast_x.x.x_aarch64.dmg`
|
||||
- **macOS (Intel)**: `proxycast_x.x.x_x64.dmg`
|
||||
- **Windows (x64)**: `proxycast_x.x.x_x64-setup.exe`
|
||||
- **Ubuntu/Debian (x64)**: `proxycast_x.x.x_amd64.deb`
|
||||
### 使用
|
||||
|
||||
### 使用步骤
|
||||
|
||||
1. **启动应用** - 打开 ProxyCast
|
||||
2. **加载凭证** - 进入 Provider 管理页面,点击"一键读取凭证"
|
||||
3. **启动服务** - 在 Dashboard 点击"启动服务器"
|
||||
4. **配置客户端** - 在 Cherry-Studio、Cline 等工具中配置:
|
||||
1. 启动 ProxyCast
|
||||
2. 加载凭证 - Provider 管理页面点击"一键读取凭证"
|
||||
3. 启动服务 - Dashboard 点击"启动服务器"
|
||||
4. 配置客户端:
|
||||
```
|
||||
API Base URL: http://localhost:8999/v1
|
||||
API Key: 启动时自动生成的密钥(可在设置页查看/修改)
|
||||
API Key: 启动时自动生成(设置页查看)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🧰 运维提示
|
||||
|
||||
- **自动备份**:数据库默认每天自动备份到 `~/.proxycast/backups/`,保留 7 天。
|
||||
- **配置备份**:每次写入配置会生成 `config.yaml.backup` 以便回滚。
|
||||
- **日志归档**:7 天游离线日志自动压缩,30 天前压缩日志自动清理。
|
||||
- **生产 HTTPS**:当前版本不内置 TLS,生产环境需反向代理终止 HTTPS。
|
||||
|
||||
---
|
||||
|
||||
## 🔧 API 使用示例
|
||||
|
||||
### OpenAI Chat Completions
|
||||
|
||||
```bash
|
||||
curl http://localhost:8999/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-d '{
|
||||
"model": "claude-sonnet-4-5-20250514",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"}
|
||||
],
|
||||
"stream": true
|
||||
}'
|
||||
```
|
||||
|
||||
### Anthropic Messages API
|
||||
|
||||
```bash
|
||||
curl http://localhost:8999/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "x-api-key: your-api-key" \
|
||||
-H "anthropic-version: 2023-06-01" \
|
||||
-d '{
|
||||
"model": "claude-sonnet-4-5-20250514",
|
||||
"max_tokens": 1024,
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🛠️ 开发构建
|
||||
|
||||
### 环境要求
|
||||
|
||||
- Node.js >= 20.0.0
|
||||
- Rust >= 1.70
|
||||
- pnpm 或 npm
|
||||
|
||||
### 本地开发
|
||||
|
||||
```bash
|
||||
# 安装依赖
|
||||
npm install
|
||||
|
||||
# 启动开发服务器
|
||||
# 开发模式
|
||||
npm run tauri dev
|
||||
```
|
||||
|
||||
### 构建发布
|
||||
|
||||
```bash
|
||||
# 构建生产版本
|
||||
# 构建发布
|
||||
npm run tauri build
|
||||
```
|
||||
|
||||
@@ -233,23 +67,8 @@ npm run tauri build
|
||||
|
||||
## 📄 开源协议
|
||||
|
||||
本项目采用 [GNU General Public License v3 (GPLv3)](https://www.gnu.org/licenses/gpl-3.0) 协议开源。
|
||||
|
||||
## 🙏 致谢
|
||||
|
||||
- [AIClient-2-API](https://github.com/justlovemaki/AIClient-2-API) - 核心逻辑参考
|
||||
- [Tauri](https://tauri.app/) - 跨平台桌面框架
|
||||
- [shadcn/ui](https://ui.shadcn.com/) - UI 组件库
|
||||
|
||||
---
|
||||
[GNU General Public License v3 (GPLv3)](https://www.gnu.org/licenses/gpl-3.0)
|
||||
|
||||
## ⚠️ 免责声明
|
||||
|
||||
### 使用风险提示
|
||||
本项目(ProxyCast)仅供学习和研究使用。用户在使用本项目时需自行承担所有风险。作者不对因使用本项目而导致的任何直接、间接或后果性损失负责。
|
||||
|
||||
### 第三方服务责任声明
|
||||
本项目是一个 API 代理工具,不提供任何 AI 模型服务。所有 AI 模型服务均由各自的第三方提供商(如 Google、Anthropic、阿里云等)提供。用户在通过本项目访问这些服务时,应遵守各第三方服务的使用条款和政策。
|
||||
|
||||
### 数据隐私声明
|
||||
本项目在本地运行,不收集或上传任何用户数据。但用户在使用本项目时应保护好自己的 API 密钥和其他敏感信息。
|
||||
本项目仅供学习研究使用,用户需自行承担使用风险。本项目不提供 AI 模型服务,所有服务由第三方提供商提供。
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.45.2",
|
||||
"version": "0.46.0",
|
||||
"type": "module",
|
||||
"repository": {
|
||||
"type": "git",
|
||||
|
||||
@@ -0,0 +1,351 @@
|
||||
#!/usr/bin/env node
|
||||
/**
|
||||
* Playwright Browser Tool for AI Agent
|
||||
*
|
||||
* 提供浏览器自动化功能,专为 AI Agent 设计
|
||||
* 输出结构化的可访问性树,便于 AI 理解页面结构
|
||||
*/
|
||||
|
||||
import { chromium } from 'playwright';
|
||||
|
||||
// 全局浏览器实例(保持会话)
|
||||
let browser = null;
|
||||
let context = null;
|
||||
let page = null;
|
||||
|
||||
// 元素引用映射
|
||||
let elementRefs = new Map();
|
||||
let refCounter = 0;
|
||||
|
||||
/**
|
||||
* 初始化浏览器
|
||||
*/
|
||||
async function initBrowser(headless = true) {
|
||||
if (!browser) {
|
||||
browser = await chromium.launch({
|
||||
headless,
|
||||
args: ['--no-sandbox', '--disable-setuid-sandbox']
|
||||
});
|
||||
context = await browser.newContext({
|
||||
viewport: { width: 1280, height: 720 },
|
||||
userAgent: 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36'
|
||||
});
|
||||
page = await context.newPage();
|
||||
}
|
||||
return { browser, context, page };
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取页面可访问性树快照
|
||||
* 为每个元素生成唯一引用 ID(如 @e1, @e2)
|
||||
*/
|
||||
async function getSnapshot(interactiveOnly = false) {
|
||||
if (!page) throw new Error('浏览器未初始化,请先执行 open 操作');
|
||||
|
||||
// 清空之前的引用
|
||||
elementRefs.clear();
|
||||
refCounter = 0;
|
||||
|
||||
// 获取可访问性树
|
||||
const snapshot = await page.accessibility.snapshot({ interestingOnly: interactiveOnly });
|
||||
|
||||
if (!snapshot) {
|
||||
return { tree: '(empty page)', refs: {} };
|
||||
}
|
||||
|
||||
// 递归处理节点,添加引用 ID
|
||||
function processNode(node, depth = 0) {
|
||||
const indent = ' '.repeat(depth);
|
||||
const ref = `@e${++refCounter}`;
|
||||
|
||||
// 存储元素引用(通过角色和名称定位)
|
||||
elementRefs.set(ref, {
|
||||
role: node.role,
|
||||
name: node.name,
|
||||
// 构建选择器
|
||||
selector: buildSelector(node)
|
||||
});
|
||||
|
||||
let line = `${indent}${ref} [${node.role}]`;
|
||||
|
||||
if (node.name) {
|
||||
line += ` "${node.name}"`;
|
||||
}
|
||||
if (node.value) {
|
||||
line += ` value="${node.value}"`;
|
||||
}
|
||||
if (node.checked !== undefined) {
|
||||
line += ` checked=${node.checked}`;
|
||||
}
|
||||
if (node.pressed !== undefined) {
|
||||
line += ` pressed=${node.pressed}`;
|
||||
}
|
||||
if (node.selected !== undefined) {
|
||||
line += ` selected=${node.selected}`;
|
||||
}
|
||||
if (node.disabled) {
|
||||
line += ` (disabled)`;
|
||||
}
|
||||
|
||||
let result = line + '\n';
|
||||
|
||||
if (node.children) {
|
||||
for (const child of node.children) {
|
||||
result += processNode(child, depth + 1);
|
||||
}
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
function buildSelector(node) {
|
||||
// 优先使用 aria-label 或 role + name 组合
|
||||
if (node.name) {
|
||||
const role = node.role.toLowerCase();
|
||||
// 尝试多种选择器策略
|
||||
return `role=${role}[name="${node.name}"]`;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
const tree = processNode(snapshot);
|
||||
|
||||
// 转换 Map 为普通对象
|
||||
const refs = {};
|
||||
for (const [key, value] of elementRefs) {
|
||||
refs[key] = value;
|
||||
}
|
||||
|
||||
return { tree, refs };
|
||||
}
|
||||
|
||||
/**
|
||||
* 解析选择器(支持 @e1 格式的引用)
|
||||
*/
|
||||
function resolveSelector(selector) {
|
||||
if (selector.startsWith('@e')) {
|
||||
const ref = elementRefs.get(selector);
|
||||
if (!ref) {
|
||||
throw new Error(`未找到元素引用: ${selector},请先执行 snapshot 获取最新的元素引用`);
|
||||
}
|
||||
return ref.selector || `text="${ref.name}"`;
|
||||
}
|
||||
return selector;
|
||||
}
|
||||
|
||||
/**
|
||||
* 执行浏览器操作
|
||||
*/
|
||||
async function executeAction(action, headless = true) {
|
||||
const result = {
|
||||
success: true,
|
||||
output: '',
|
||||
url: null,
|
||||
title: null,
|
||||
screenshot: null,
|
||||
error: null
|
||||
};
|
||||
|
||||
try {
|
||||
switch (action.open?.url ? 'open' : Object.keys(action)[0]) {
|
||||
case 'open': {
|
||||
const url = action.open?.url || action.url;
|
||||
await initBrowser(headless);
|
||||
await page.goto(url, { waitUntil: 'domcontentloaded', timeout: 30000 });
|
||||
result.url = page.url();
|
||||
result.title = await page.title();
|
||||
result.output = `已打开页面: ${result.title}`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'snapshot': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const interactiveOnly = action.snapshot?.interactive_only || action.interactive_only || false;
|
||||
const { tree, refs } = await getSnapshot(interactiveOnly);
|
||||
result.url = page.url();
|
||||
result.title = await page.title();
|
||||
result.output = `页面可访问性树:\n\n${tree}\n\n共 ${Object.keys(refs).length} 个元素`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'click': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const selector = resolveSelector(action.click?.selector || action.selector);
|
||||
await page.click(selector, { timeout: 5000 });
|
||||
result.output = `已点击元素: ${action.click?.selector || action.selector}`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'fill': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const selector = resolveSelector(action.fill?.selector || action.selector);
|
||||
const value = action.fill?.value || action.value;
|
||||
await page.fill(selector, value, { timeout: 5000 });
|
||||
result.output = `已填充表单: ${action.fill?.selector || action.selector} = "${value}"`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'type': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const selector = resolveSelector(action.type?.selector || action.selector);
|
||||
const text = action.type?.text || action.text;
|
||||
await page.type(selector, text, { delay: 50 });
|
||||
result.output = `已输入文本: "${text}"`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'press': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const key = action.press?.key || action.key;
|
||||
await page.keyboard.press(key);
|
||||
result.output = `已按键: ${key}`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'scroll': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const direction = action.scroll?.direction || action.direction || 'down';
|
||||
const amount = action.scroll?.amount || action.amount || 500;
|
||||
|
||||
let deltaX = 0, deltaY = 0;
|
||||
switch (direction) {
|
||||
case 'up': deltaY = -amount; break;
|
||||
case 'down': deltaY = amount; break;
|
||||
case 'left': deltaX = -amount; break;
|
||||
case 'right': deltaX = amount; break;
|
||||
}
|
||||
|
||||
await page.mouse.wheel(deltaX, deltaY);
|
||||
result.output = `已滚动: ${direction} ${amount}px`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'wait_for': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const selector = resolveSelector(action.wait_for?.selector || action.selector);
|
||||
const timeoutMs = action.wait_for?.timeout_ms || action.timeout_ms || 5000;
|
||||
await page.waitForSelector(selector, { timeout: timeoutMs });
|
||||
result.output = `元素已出现: ${action.wait_for?.selector || action.selector}`;
|
||||
break;
|
||||
}
|
||||
|
||||
case 'screenshot': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const fullPage = action.screenshot?.full_page || action.full_page || false;
|
||||
const path = action.screenshot?.path || action.path;
|
||||
|
||||
const options = { fullPage };
|
||||
if (path) {
|
||||
options.path = path;
|
||||
await page.screenshot(options);
|
||||
result.output = `截图已保存: ${path}`;
|
||||
} else {
|
||||
const buffer = await page.screenshot(options);
|
||||
result.screenshot = buffer.toString('base64');
|
||||
result.output = '截图已生成(base64)';
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
case 'get_text': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const selector = action.get_text?.selector || action.selector;
|
||||
|
||||
let text;
|
||||
if (selector) {
|
||||
const resolved = resolveSelector(selector);
|
||||
text = await page.textContent(resolved);
|
||||
} else {
|
||||
text = await page.evaluate(() => document.body.innerText);
|
||||
}
|
||||
|
||||
result.output = text || '(empty)';
|
||||
break;
|
||||
}
|
||||
|
||||
case 'evaluate': {
|
||||
if (!page) throw new Error('浏览器未初始化');
|
||||
const script = action.evaluate?.script || action.script;
|
||||
const evalResult = await page.evaluate(script);
|
||||
result.output = JSON.stringify(evalResult, null, 2);
|
||||
break;
|
||||
}
|
||||
|
||||
case 'close': {
|
||||
if (browser) {
|
||||
await browser.close();
|
||||
browser = null;
|
||||
context = null;
|
||||
page = null;
|
||||
elementRefs.clear();
|
||||
}
|
||||
result.output = '浏览器已关闭';
|
||||
break;
|
||||
}
|
||||
|
||||
default:
|
||||
throw new Error(`未知操作: ${JSON.stringify(action)}`);
|
||||
}
|
||||
} catch (error) {
|
||||
result.success = false;
|
||||
result.error = error.message;
|
||||
result.output = `操作失败: ${error.message}`;
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
/**
|
||||
* 主函数
|
||||
*/
|
||||
async function main() {
|
||||
const args = process.argv.slice(2);
|
||||
|
||||
let actionJson = null;
|
||||
let headless = true;
|
||||
|
||||
for (let i = 0; i < args.length; i++) {
|
||||
if (args[i] === '--action' && args[i + 1]) {
|
||||
actionJson = args[i + 1];
|
||||
i++;
|
||||
} else if (args[i] === '--headless') {
|
||||
headless = true;
|
||||
} else if (args[i] === '--no-headless') {
|
||||
headless = false;
|
||||
}
|
||||
}
|
||||
|
||||
if (!actionJson) {
|
||||
console.error('Usage: browser-tool.mjs --action <json> [--headless|--no-headless]');
|
||||
process.exit(1);
|
||||
}
|
||||
|
||||
try {
|
||||
const action = JSON.parse(actionJson);
|
||||
const result = await executeAction(action, headless);
|
||||
console.log(JSON.stringify(result));
|
||||
} catch (error) {
|
||||
console.log(JSON.stringify({
|
||||
success: false,
|
||||
output: '',
|
||||
error: error.message
|
||||
}));
|
||||
process.exit(1);
|
||||
}
|
||||
}
|
||||
|
||||
// 处理进程退出
|
||||
process.on('exit', async () => {
|
||||
if (browser) {
|
||||
await browser.close();
|
||||
}
|
||||
});
|
||||
|
||||
process.on('SIGINT', async () => {
|
||||
if (browser) {
|
||||
await browser.close();
|
||||
}
|
||||
process.exit(0);
|
||||
});
|
||||
|
||||
main();
|
||||
@@ -0,0 +1,13 @@
|
||||
{
|
||||
"name": "proxycast-playwright-tools",
|
||||
"version": "1.0.0",
|
||||
"description": "Playwright browser automation tools for ProxyCast AI Agent",
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"test": "node browser-tool.mjs --action '{\"open\":{\"url\":\"https://example.com\"}}' --no-headless",
|
||||
"install-browsers": "npx playwright install chromium"
|
||||
},
|
||||
"dependencies": {
|
||||
"playwright": "^1.40.0"
|
||||
}
|
||||
}
|
||||
Generated
+3495
-346
File diff suppressed because it is too large
Load Diff
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "proxycast"
|
||||
version = "0.45.2"
|
||||
version = "0.46.0"
|
||||
description = "AI API Proxy Desktop App"
|
||||
authors = ["you"]
|
||||
edition = "2021"
|
||||
@@ -77,6 +77,10 @@ mouse_position = "0.1.4"
|
||||
window-vibrancy = "0.7.1"
|
||||
if-addrs = "0.13"
|
||||
|
||||
# Aster Agent Framework
|
||||
# 使用本地路径进行开发,后续切换到 git
|
||||
aster = { path = "../../../astercloud/aster-rust/crates/aster" }
|
||||
|
||||
# Platform specific dependencies for browser interceptor
|
||||
|
||||
# Windows specific dependencies for browser interceptor and machine ID management
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
//! Aster Agent 包装器
|
||||
//!
|
||||
//! 提供简化的接口来使用 Aster Agent
|
||||
//! 处理消息发送、事件流转换和会话管理
|
||||
|
||||
use crate::agent::aster_state::{AsterAgentState, SessionConfigBuilder};
|
||||
use crate::agent::event_converter::TauriAgentEvent;
|
||||
use aster::agents::SessionConfig;
|
||||
use aster::conversation::message::Message;
|
||||
use aster::session::SessionManager;
|
||||
use std::path::PathBuf;
|
||||
use tauri::{AppHandle, Emitter};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
/// Aster Agent 包装器
|
||||
///
|
||||
/// 提供与 Tauri 集成的简化接口
|
||||
pub struct AsterAgentWrapper;
|
||||
|
||||
impl AsterAgentWrapper {
|
||||
/// 发送消息并获取流式响应
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `state` - Aster Agent 状态
|
||||
/// * `app` - Tauri AppHandle,用于发送事件
|
||||
/// * `message` - 用户消息文本
|
||||
/// * `session_id` - 会话 ID
|
||||
/// * `event_name` - 前端监听的事件名称
|
||||
///
|
||||
/// # Returns
|
||||
/// 成功时返回 Ok(()),失败时返回错误信息
|
||||
pub async fn send_message(
|
||||
state: &AsterAgentState,
|
||||
app: &AppHandle,
|
||||
message: String,
|
||||
session_id: String,
|
||||
event_name: String,
|
||||
) -> Result<(), String> {
|
||||
// 确保 Agent 已初始化
|
||||
if !state.is_initialized().await {
|
||||
state.init_agent().await?;
|
||||
}
|
||||
|
||||
// 创建取消令牌
|
||||
let cancel_token = state.create_cancel_token(&session_id).await;
|
||||
|
||||
// 创建用户消息
|
||||
let user_message = Message::user().with_text(&message);
|
||||
|
||||
// 创建会话配置
|
||||
let session_config = SessionConfigBuilder::new(&session_id).build();
|
||||
|
||||
// 使用 with_agent 方法获取 Agent 并处理
|
||||
let app_clone = app.clone();
|
||||
let event_name_clone = event_name.clone();
|
||||
let cancel_token_clone = cancel_token.clone();
|
||||
|
||||
let result = state
|
||||
.with_agent(|agent| {
|
||||
// 注意:这里我们需要异步处理,但 with_agent 是同步的
|
||||
// 我们需要重新设计这个接口
|
||||
})
|
||||
.await;
|
||||
|
||||
// 由于 with_agent 的限制,我们需要使用不同的方法
|
||||
// 直接在这里处理流
|
||||
Self::process_reply_internal(
|
||||
state,
|
||||
&app_clone,
|
||||
user_message,
|
||||
session_config,
|
||||
cancel_token_clone,
|
||||
event_name_clone,
|
||||
)
|
||||
.await?;
|
||||
|
||||
// 清理取消令牌
|
||||
state.remove_cancel_token(&session_id).await;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 内部处理回复的方法
|
||||
async fn process_reply_internal(
|
||||
state: &AsterAgentState,
|
||||
app: &AppHandle,
|
||||
user_message: Message,
|
||||
session_config: SessionConfig,
|
||||
cancel_token: CancellationToken,
|
||||
event_name: String,
|
||||
) -> Result<(), String> {
|
||||
// 这里我们需要一个更好的方式来访问 Agent
|
||||
// 暂时使用一个简化的实现
|
||||
|
||||
// 发送开始事件
|
||||
let start_event = TauriAgentEvent::TextDelta {
|
||||
text: String::new(),
|
||||
};
|
||||
let _ = app.emit(&event_name, &start_event);
|
||||
|
||||
// TODO: 实现完整的 Agent 调用
|
||||
// 由于 Agent.reply() 需要 &self,而我们的 with_agent 方法不支持异步
|
||||
// 我们需要重新设计 AsterAgentState 的接口
|
||||
|
||||
// 发送完成事件
|
||||
let done_event = TauriAgentEvent::FinalDone { usage: None };
|
||||
if let Err(e) = app.emit(&event_name, &done_event) {
|
||||
tracing::error!("Failed to emit final done event: {}", e);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 停止当前会话
|
||||
pub async fn stop_session(state: &AsterAgentState, session_id: &str) -> bool {
|
||||
state.cancel_session(session_id).await
|
||||
}
|
||||
|
||||
/// 创建新会话
|
||||
pub async fn create_session(
|
||||
working_dir: Option<PathBuf>,
|
||||
name: Option<String>,
|
||||
) -> Result<String, String> {
|
||||
let dir = working_dir
|
||||
.unwrap_or_else(|| std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")));
|
||||
let session_name = name.unwrap_or_else(|| "New Session".to_string());
|
||||
|
||||
let session =
|
||||
SessionManager::create_session(dir, session_name, aster::session::SessionType::User)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to create session: {}", e))?;
|
||||
|
||||
Ok(session.id)
|
||||
}
|
||||
|
||||
/// 列出所有会话
|
||||
pub async fn list_sessions() -> Result<Vec<SessionInfo>, String> {
|
||||
let sessions = SessionManager::list_sessions()
|
||||
.await
|
||||
.map_err(|e| format!("Failed to list sessions: {}", e))?;
|
||||
|
||||
Ok(sessions
|
||||
.into_iter()
|
||||
.map(|s| SessionInfo {
|
||||
id: s.id,
|
||||
name: s.name,
|
||||
created_at: s.created_at.timestamp(),
|
||||
updated_at: s.updated_at.timestamp(),
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// 获取会话详情
|
||||
pub async fn get_session(session_id: &str) -> Result<SessionDetail, String> {
|
||||
let session = SessionManager::get_session(session_id, true)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to get session: {}", e))?;
|
||||
|
||||
Ok(SessionDetail {
|
||||
id: session.id,
|
||||
name: session.name,
|
||||
created_at: session.created_at.timestamp(),
|
||||
updated_at: session.updated_at.timestamp(),
|
||||
messages: session
|
||||
.conversation
|
||||
.map(|c| {
|
||||
c.messages()
|
||||
.iter()
|
||||
.map(|m| crate::agent::event_converter::convert_to_tauri_message(m))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 会话信息(简化版)
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub struct SessionInfo {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
pub created_at: i64,
|
||||
pub updated_at: i64,
|
||||
}
|
||||
|
||||
/// 会话详情(包含消息)
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub struct SessionDetail {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
pub created_at: i64,
|
||||
pub updated_at: i64,
|
||||
pub messages: Vec<crate::agent::event_converter::TauriMessage>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_config_builder() {
|
||||
let config = SessionConfigBuilder::new("test-session").build();
|
||||
assert_eq!(config.id, "test-session");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,293 @@
|
||||
//! Aster Agent 状态管理
|
||||
//!
|
||||
//! 管理 Aster Agent 实例和相关状态
|
||||
//! 提供 Tauri 应用与 Aster 框架的桥接
|
||||
|
||||
use aster::agents::{Agent, SessionConfig};
|
||||
use aster::model::ModelConfig;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
/// Provider 配置信息
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ProviderConfig {
|
||||
/// Provider 名称 (openai, anthropic, google, ollama 等)
|
||||
pub provider_name: String,
|
||||
/// 模型名称
|
||||
pub model_name: String,
|
||||
/// API Key (可选,某些 provider 从环境变量读取)
|
||||
pub api_key: Option<String>,
|
||||
/// Base URL (可选,用于自定义端点)
|
||||
pub base_url: Option<String>,
|
||||
}
|
||||
|
||||
/// Aster Agent 全局状态
|
||||
///
|
||||
/// 在 Tauri 应用中作为 managed state 使用
|
||||
pub struct AsterAgentState {
|
||||
/// Aster Agent 实例
|
||||
agent: Arc<RwLock<Option<Agent>>>,
|
||||
/// 当前活跃的取消令牌(用于中止正在进行的对话)
|
||||
cancel_tokens: Arc<RwLock<std::collections::HashMap<String, CancellationToken>>>,
|
||||
/// 当前 Provider 配置
|
||||
current_provider_config: Arc<RwLock<Option<ProviderConfig>>>,
|
||||
}
|
||||
|
||||
impl Default for AsterAgentState {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl AsterAgentState {
|
||||
/// 创建新的 Aster Agent 状态
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
agent: Arc::new(RwLock::new(None)),
|
||||
cancel_tokens: Arc::new(RwLock::new(std::collections::HashMap::new())),
|
||||
current_provider_config: Arc::new(RwLock::new(None)),
|
||||
}
|
||||
}
|
||||
|
||||
/// 初始化 Agent
|
||||
///
|
||||
/// 如果 Agent 尚未初始化,则创建新的 Agent 实例
|
||||
pub async fn init_agent(&self) -> Result<(), String> {
|
||||
let mut agent_guard = self.agent.write().await;
|
||||
if agent_guard.is_none() {
|
||||
let agent = Agent::new();
|
||||
*agent_guard = Some(agent);
|
||||
tracing::info!("Aster Agent initialized");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 配置 Provider
|
||||
///
|
||||
/// 根据配置创建并设置 Provider
|
||||
pub async fn configure_provider(
|
||||
&self,
|
||||
config: ProviderConfig,
|
||||
session_id: &str,
|
||||
) -> Result<(), String> {
|
||||
// 确保 Agent 已初始化
|
||||
self.init_agent().await?;
|
||||
|
||||
// 设置环境变量(Aster 的 provider 从环境变量读取配置)
|
||||
self.set_provider_env_vars(&config);
|
||||
|
||||
// 创建 ModelConfig
|
||||
let model_config = ModelConfig::new(&config.model_name)
|
||||
.map_err(|e| format!("创建 ModelConfig 失败: {}", e))?;
|
||||
|
||||
// 创建 Provider
|
||||
let provider = aster::providers::create(&config.provider_name, model_config)
|
||||
.await
|
||||
.map_err(|e| format!("创建 Provider 失败: {}", e))?;
|
||||
|
||||
// 更新 Agent 的 Provider
|
||||
let agent_guard = self.agent.read().await;
|
||||
if let Some(agent) = agent_guard.as_ref() {
|
||||
agent
|
||||
.update_provider(provider, session_id)
|
||||
.await
|
||||
.map_err(|e| format!("更新 Provider 失败: {}", e))?;
|
||||
}
|
||||
|
||||
// 保存当前配置
|
||||
let mut config_guard = self.current_provider_config.write().await;
|
||||
*config_guard = Some(config.clone());
|
||||
|
||||
tracing::info!(
|
||||
"[AsterAgent] Provider 配置成功: {} / {}",
|
||||
config.provider_name,
|
||||
config.model_name
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 设置 Provider 相关的环境变量
|
||||
fn set_provider_env_vars(&self, config: &ProviderConfig) {
|
||||
// 根据 provider 类型设置对应的环境变量
|
||||
let env_key = match config.provider_name.as_str() {
|
||||
"openai" => "OPENAI_API_KEY",
|
||||
"anthropic" => "ANTHROPIC_API_KEY",
|
||||
"google" => "GOOGLE_API_KEY",
|
||||
"deepseek" | "custom_deepseek" => "DEEPSEEK_API_KEY",
|
||||
"groq" => "GROQ_API_KEY",
|
||||
"mistral" => "MISTRAL_API_KEY",
|
||||
"openrouter" => "OPENROUTER_API_KEY",
|
||||
"ollama" => return, // Ollama 不需要 API Key
|
||||
_ => {
|
||||
// 通用 OpenAI 兼容格式
|
||||
if let Some(api_key) = &config.api_key {
|
||||
std::env::set_var("OPENAI_API_KEY", api_key);
|
||||
}
|
||||
if let Some(base_url) = &config.base_url {
|
||||
std::env::set_var("OPENAI_BASE_URL", base_url);
|
||||
}
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(api_key) = &config.api_key {
|
||||
std::env::set_var(env_key, api_key);
|
||||
}
|
||||
|
||||
if let Some(base_url) = &config.base_url {
|
||||
let base_url_key = format!(
|
||||
"{}_BASE_URL",
|
||||
config.provider_name.to_uppercase().replace("_", "")
|
||||
);
|
||||
std::env::set_var(base_url_key, base_url);
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取当前 Provider 配置
|
||||
pub async fn get_provider_config(&self) -> Option<ProviderConfig> {
|
||||
self.current_provider_config.read().await.clone()
|
||||
}
|
||||
|
||||
/// 检查 Provider 是否已配置
|
||||
pub async fn is_provider_configured(&self) -> bool {
|
||||
self.current_provider_config.read().await.is_some()
|
||||
}
|
||||
|
||||
/// 获取 Agent 的只读引用并执行同步操作
|
||||
pub async fn with_agent<F, R>(&self, f: F) -> Result<R, String>
|
||||
where
|
||||
F: FnOnce(&Agent) -> R,
|
||||
{
|
||||
let guard = self.agent.read().await;
|
||||
match guard.as_ref() {
|
||||
Some(agent) => Ok(f(agent)),
|
||||
None => Err("Agent not initialized".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取 Agent 的可变引用并执行同步操作
|
||||
pub async fn with_agent_mut<F, R>(&self, f: F) -> Result<R, String>
|
||||
where
|
||||
F: FnOnce(&mut Agent) -> R,
|
||||
{
|
||||
let mut guard = self.agent.write().await;
|
||||
match guard.as_mut() {
|
||||
Some(agent) => Ok(f(agent)),
|
||||
None => Err("Agent not initialized".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取 Agent 的 Arc 引用
|
||||
///
|
||||
/// 用于需要长期持有 Agent 引用的场景
|
||||
pub fn get_agent_arc(&self) -> Arc<RwLock<Option<Agent>>> {
|
||||
self.agent.clone()
|
||||
}
|
||||
|
||||
/// 创建新的取消令牌
|
||||
pub async fn create_cancel_token(&self, session_id: &str) -> CancellationToken {
|
||||
let token = CancellationToken::new();
|
||||
let mut tokens = self.cancel_tokens.write().await;
|
||||
tokens.insert(session_id.to_string(), token.clone());
|
||||
token
|
||||
}
|
||||
|
||||
/// 取消指定会话的操作
|
||||
pub async fn cancel_session(&self, session_id: &str) -> bool {
|
||||
let tokens = self.cancel_tokens.read().await;
|
||||
if let Some(token) = tokens.get(session_id) {
|
||||
token.cancel();
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// 移除取消令牌
|
||||
pub async fn remove_cancel_token(&self, session_id: &str) {
|
||||
let mut tokens = self.cancel_tokens.write().await;
|
||||
tokens.remove(session_id);
|
||||
}
|
||||
|
||||
/// 检查 Agent 是否已初始化
|
||||
pub async fn is_initialized(&self) -> bool {
|
||||
self.agent.read().await.is_some()
|
||||
}
|
||||
}
|
||||
|
||||
/// 会话配置构建器
|
||||
///
|
||||
/// 用于构建 Aster SessionConfig
|
||||
pub struct SessionConfigBuilder {
|
||||
id: String,
|
||||
max_turns: Option<u32>,
|
||||
}
|
||||
|
||||
impl SessionConfigBuilder {
|
||||
pub fn new(id: impl Into<String>) -> Self {
|
||||
Self {
|
||||
id: id.into(),
|
||||
max_turns: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn max_turns(mut self, turns: u32) -> Self {
|
||||
self.max_turns = Some(turns);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn build(self) -> SessionConfig {
|
||||
SessionConfig {
|
||||
id: self.id,
|
||||
schedule_id: None,
|
||||
max_turns: self.max_turns,
|
||||
retry_config: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 消息构建辅助函数
|
||||
pub mod message_helpers {
|
||||
use aster::conversation::message::Message;
|
||||
|
||||
/// 创建用户文本消息
|
||||
pub fn user_text(text: impl Into<String>) -> Message {
|
||||
Message::user().with_text(text)
|
||||
}
|
||||
|
||||
/// 创建助手文本消息
|
||||
pub fn assistant_text(text: impl Into<String>) -> Message {
|
||||
Message::assistant().with_text(text)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_aster_state_init() {
|
||||
let state = AsterAgentState::new();
|
||||
assert!(!state.is_initialized().await);
|
||||
|
||||
state.init_agent().await.unwrap();
|
||||
assert!(state.is_initialized().await);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cancel_token() {
|
||||
let state = AsterAgentState::new();
|
||||
let session_id = "test-session";
|
||||
|
||||
let token = state.create_cancel_token(session_id).await;
|
||||
assert!(!token.is_cancelled());
|
||||
|
||||
assert!(state.cancel_session(session_id).await);
|
||||
assert!(token.is_cancelled());
|
||||
|
||||
state.remove_cancel_token(session_id).await;
|
||||
assert!(!state.cancel_session(session_id).await);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,451 @@
|
||||
//! Aster 事件转换器
|
||||
//!
|
||||
//! 将 Aster AgentEvent 转换为 Tauri 可用的事件格式
|
||||
//! 用于前端实时显示流式响应
|
||||
|
||||
use aster::agents::AgentEvent;
|
||||
use aster::conversation::message::{ActionRequiredData, Message, MessageContent};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// 从工具结果中提取文本内容
|
||||
///
|
||||
/// 使用 serde_json 来处理,避免直接依赖 rmcp 类型
|
||||
fn extract_tool_result_text<T: serde::Serialize>(result: &T) -> String {
|
||||
if let Ok(json) = serde_json::to_value(result) {
|
||||
if let Some(content) = json.get("content").and_then(|c| c.as_array()) {
|
||||
return content
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
if item.get("type").and_then(|t| t.as_str()) == Some("text") {
|
||||
item.get("text").and_then(|t| t.as_str()).map(String::from)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n");
|
||||
}
|
||||
}
|
||||
String::new()
|
||||
}
|
||||
|
||||
/// Tauri Agent 事件
|
||||
///
|
||||
/// 用于前端消费的事件格式,与现有的 StreamEvent 兼容
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum TauriAgentEvent {
|
||||
/// 文本增量
|
||||
#[serde(rename = "text_delta")]
|
||||
TextDelta { text: String },
|
||||
|
||||
/// 思考内容增量
|
||||
#[serde(rename = "thinking_delta")]
|
||||
ThinkingDelta { text: String },
|
||||
|
||||
/// 工具调用开始
|
||||
#[serde(rename = "tool_start")]
|
||||
ToolStart {
|
||||
tool_name: String,
|
||||
tool_id: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
arguments: Option<String>,
|
||||
},
|
||||
|
||||
/// 工具调用结束
|
||||
#[serde(rename = "tool_end")]
|
||||
ToolEnd {
|
||||
tool_id: String,
|
||||
result: TauriToolResult,
|
||||
},
|
||||
|
||||
/// 需要用户操作(权限确认、用户输入等)
|
||||
#[serde(rename = "action_required")]
|
||||
ActionRequired {
|
||||
request_id: String,
|
||||
action_type: String,
|
||||
data: serde_json::Value,
|
||||
},
|
||||
|
||||
/// 模型变更
|
||||
#[serde(rename = "model_change")]
|
||||
ModelChange { model: String, mode: String },
|
||||
|
||||
/// 完成(单次响应完成)
|
||||
#[serde(rename = "done")]
|
||||
Done {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
usage: Option<TauriTokenUsage>,
|
||||
},
|
||||
|
||||
/// 最终完成(整个对话完成)
|
||||
#[serde(rename = "final_done")]
|
||||
FinalDone {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
usage: Option<TauriTokenUsage>,
|
||||
},
|
||||
|
||||
/// 错误
|
||||
#[serde(rename = "error")]
|
||||
Error { message: String },
|
||||
|
||||
/// 完整消息(用于历史记录)
|
||||
#[serde(rename = "message")]
|
||||
Message { message: TauriMessage },
|
||||
}
|
||||
|
||||
/// 工具执行结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TauriToolResult {
|
||||
pub success: bool,
|
||||
pub output: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
/// Token 使用量
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TauriTokenUsage {
|
||||
pub input_tokens: u32,
|
||||
pub output_tokens: u32,
|
||||
}
|
||||
|
||||
/// 简化的消息结构
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TauriMessage {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub id: Option<String>,
|
||||
pub role: String,
|
||||
pub content: Vec<TauriMessageContent>,
|
||||
pub timestamp: i64,
|
||||
}
|
||||
|
||||
/// 简化的消息内容
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum TauriMessageContent {
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
|
||||
#[serde(rename = "thinking")]
|
||||
Thinking { text: String },
|
||||
|
||||
#[serde(rename = "tool_request")]
|
||||
ToolRequest {
|
||||
id: String,
|
||||
tool_name: String,
|
||||
arguments: serde_json::Value,
|
||||
},
|
||||
|
||||
#[serde(rename = "tool_response")]
|
||||
ToolResponse {
|
||||
id: String,
|
||||
success: bool,
|
||||
output: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
error: Option<String>,
|
||||
},
|
||||
|
||||
#[serde(rename = "action_required")]
|
||||
ActionRequired {
|
||||
id: String,
|
||||
action_type: String,
|
||||
data: serde_json::Value,
|
||||
},
|
||||
|
||||
#[serde(rename = "image")]
|
||||
Image { mime_type: String, data: String },
|
||||
}
|
||||
|
||||
/// 将 Aster AgentEvent 转换为 TauriAgentEvent 列表
|
||||
///
|
||||
/// 一个 AgentEvent 可能产生多个 TauriAgentEvent
|
||||
pub fn convert_agent_event(event: AgentEvent) -> Vec<TauriAgentEvent> {
|
||||
match event {
|
||||
AgentEvent::Message(message) => convert_message(message),
|
||||
AgentEvent::McpNotification((server_name, notification)) => {
|
||||
// MCP 通知暂时忽略或转换为日志
|
||||
tracing::debug!("MCP notification from {}: {:?}", server_name, notification);
|
||||
vec![]
|
||||
}
|
||||
AgentEvent::ModelChange { model, mode } => {
|
||||
vec![TauriAgentEvent::ModelChange { model, mode }]
|
||||
}
|
||||
AgentEvent::HistoryReplaced(_conversation) => {
|
||||
// 历史替换事件,可能需要特殊处理
|
||||
tracing::debug!("History replaced");
|
||||
vec![]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 Aster Message 转换为 TauriAgentEvent 列表
|
||||
fn convert_message(message: Message) -> Vec<TauriAgentEvent> {
|
||||
let mut events = Vec::new();
|
||||
|
||||
for content in &message.content {
|
||||
match content {
|
||||
MessageContent::Text(text_content) => {
|
||||
events.push(TauriAgentEvent::TextDelta {
|
||||
text: text_content.text.clone(),
|
||||
});
|
||||
}
|
||||
MessageContent::Thinking(thinking) => {
|
||||
events.push(TauriAgentEvent::ThinkingDelta {
|
||||
text: thinking.thinking.clone(),
|
||||
});
|
||||
}
|
||||
MessageContent::ToolRequest(tool_request) => match &tool_request.tool_call {
|
||||
Ok(call) => {
|
||||
events.push(TauriAgentEvent::ToolStart {
|
||||
tool_name: call.name.to_string(),
|
||||
tool_id: tool_request.id.clone(),
|
||||
arguments: serde_json::to_string(&call.arguments).ok(),
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
events.push(TauriAgentEvent::Error {
|
||||
message: format!("Invalid tool call: {}", e),
|
||||
});
|
||||
}
|
||||
},
|
||||
MessageContent::ToolResponse(tool_response) => {
|
||||
let (success, output, error) = match &tool_response.tool_result {
|
||||
Ok(result) => {
|
||||
// 从 CallToolResult 中提取文本内容
|
||||
let output = extract_tool_result_text(result);
|
||||
(true, output, None)
|
||||
}
|
||||
Err(e) => (false, String::new(), Some(e.to_string())),
|
||||
};
|
||||
|
||||
events.push(TauriAgentEvent::ToolEnd {
|
||||
tool_id: tool_response.id.clone(),
|
||||
result: TauriToolResult {
|
||||
success,
|
||||
output,
|
||||
error,
|
||||
},
|
||||
});
|
||||
}
|
||||
MessageContent::ActionRequired(action_required) => {
|
||||
let (request_id, action_type, data) = match &action_required.data {
|
||||
ActionRequiredData::ToolConfirmation {
|
||||
id,
|
||||
tool_name,
|
||||
arguments,
|
||||
prompt,
|
||||
} => (
|
||||
id.clone(),
|
||||
"tool_confirmation".to_string(),
|
||||
serde_json::json!({
|
||||
"tool_name": tool_name,
|
||||
"arguments": arguments,
|
||||
"prompt": prompt,
|
||||
}),
|
||||
),
|
||||
ActionRequiredData::Elicitation {
|
||||
id,
|
||||
message,
|
||||
requested_schema,
|
||||
} => (
|
||||
id.clone(),
|
||||
"elicitation".to_string(),
|
||||
serde_json::json!({
|
||||
"message": message,
|
||||
"requested_schema": requested_schema,
|
||||
}),
|
||||
),
|
||||
ActionRequiredData::ElicitationResponse { id, user_data } => (
|
||||
id.clone(),
|
||||
"elicitation_response".to_string(),
|
||||
serde_json::json!({
|
||||
"user_data": user_data,
|
||||
}),
|
||||
),
|
||||
};
|
||||
|
||||
events.push(TauriAgentEvent::ActionRequired {
|
||||
request_id,
|
||||
action_type,
|
||||
data,
|
||||
});
|
||||
}
|
||||
MessageContent::SystemNotification(notification) => {
|
||||
// 系统通知转换为文本
|
||||
events.push(TauriAgentEvent::TextDelta {
|
||||
text: notification.msg.clone(),
|
||||
});
|
||||
}
|
||||
MessageContent::Image(image) => {
|
||||
// 图片内容暂时忽略
|
||||
tracing::debug!("Image content: {}", image.mime_type);
|
||||
}
|
||||
MessageContent::ToolConfirmationRequest(req) => {
|
||||
events.push(TauriAgentEvent::ActionRequired {
|
||||
request_id: req.id.clone(),
|
||||
action_type: "tool_confirmation".to_string(),
|
||||
data: serde_json::json!({
|
||||
"tool_name": req.tool_name,
|
||||
"arguments": req.arguments,
|
||||
"prompt": req.prompt,
|
||||
}),
|
||||
});
|
||||
}
|
||||
MessageContent::FrontendToolRequest(req) => match &req.tool_call {
|
||||
Ok(call) => {
|
||||
events.push(TauriAgentEvent::ToolStart {
|
||||
tool_name: call.name.to_string(),
|
||||
tool_id: req.id.clone(),
|
||||
arguments: serde_json::to_string(&call.arguments).ok(),
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
events.push(TauriAgentEvent::Error {
|
||||
message: format!("Invalid frontend tool call: {}", e),
|
||||
});
|
||||
}
|
||||
},
|
||||
MessageContent::RedactedThinking(_) => {
|
||||
// 隐藏的思考内容,忽略
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
events
|
||||
}
|
||||
|
||||
/// 将 Aster Message 转换为 TauriMessage
|
||||
pub fn convert_to_tauri_message(message: &Message) -> TauriMessage {
|
||||
let content = message
|
||||
.content
|
||||
.iter()
|
||||
.filter_map(|c| convert_message_content(c))
|
||||
.collect();
|
||||
|
||||
TauriMessage {
|
||||
id: message.id.clone(),
|
||||
role: format!("{:?}", message.role).to_lowercase(),
|
||||
content,
|
||||
timestamp: message.created,
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 MessageContent 转换为 TauriMessageContent
|
||||
fn convert_message_content(content: &MessageContent) -> Option<TauriMessageContent> {
|
||||
match content {
|
||||
MessageContent::Text(text) => Some(TauriMessageContent::Text {
|
||||
text: text.text.clone(),
|
||||
}),
|
||||
MessageContent::Thinking(thinking) => Some(TauriMessageContent::Thinking {
|
||||
text: thinking.thinking.clone(),
|
||||
}),
|
||||
MessageContent::ToolRequest(req) => {
|
||||
req.tool_call
|
||||
.as_ref()
|
||||
.ok()
|
||||
.map(|call| TauriMessageContent::ToolRequest {
|
||||
id: req.id.clone(),
|
||||
tool_name: call.name.to_string(),
|
||||
arguments: serde_json::to_value(&call.arguments).unwrap_or_default(),
|
||||
})
|
||||
}
|
||||
MessageContent::ToolResponse(resp) => {
|
||||
let (success, output, error) = match &resp.tool_result {
|
||||
Ok(result) => {
|
||||
let output = extract_tool_result_text(result);
|
||||
(true, output, None)
|
||||
}
|
||||
Err(e) => (false, String::new(), Some(e.to_string())),
|
||||
};
|
||||
Some(TauriMessageContent::ToolResponse {
|
||||
id: resp.id.clone(),
|
||||
success,
|
||||
output,
|
||||
error,
|
||||
})
|
||||
}
|
||||
MessageContent::ActionRequired(action) => {
|
||||
let (id, action_type, data) = match &action.data {
|
||||
ActionRequiredData::ToolConfirmation {
|
||||
id,
|
||||
tool_name,
|
||||
arguments,
|
||||
prompt,
|
||||
} => (
|
||||
id.clone(),
|
||||
"tool_confirmation".to_string(),
|
||||
serde_json::json!({
|
||||
"tool_name": tool_name,
|
||||
"arguments": arguments,
|
||||
"prompt": prompt,
|
||||
}),
|
||||
),
|
||||
ActionRequiredData::Elicitation {
|
||||
id,
|
||||
message,
|
||||
requested_schema,
|
||||
} => (
|
||||
id.clone(),
|
||||
"elicitation".to_string(),
|
||||
serde_json::json!({
|
||||
"message": message,
|
||||
"requested_schema": requested_schema,
|
||||
}),
|
||||
),
|
||||
ActionRequiredData::ElicitationResponse { id, user_data } => (
|
||||
id.clone(),
|
||||
"elicitation_response".to_string(),
|
||||
user_data.clone(),
|
||||
),
|
||||
};
|
||||
Some(TauriMessageContent::ActionRequired {
|
||||
id,
|
||||
action_type,
|
||||
data,
|
||||
})
|
||||
}
|
||||
MessageContent::Image(image) => Some(TauriMessageContent::Image {
|
||||
mime_type: image.mime_type.clone(),
|
||||
data: image.data.clone(),
|
||||
}),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_convert_text_delta() {
|
||||
let message = Message::assistant().with_text("Hello, world!");
|
||||
let events = convert_message(message);
|
||||
|
||||
assert_eq!(events.len(), 1);
|
||||
match &events[0] {
|
||||
TauriAgentEvent::TextDelta { text } => {
|
||||
assert_eq!(text, "Hello, world!");
|
||||
}
|
||||
_ => panic!("Expected TextDelta event"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_convert_model_change() {
|
||||
let event = AgentEvent::ModelChange {
|
||||
model: "claude-3".to_string(),
|
||||
mode: "chat".to_string(),
|
||||
};
|
||||
let events = convert_agent_event(event);
|
||||
|
||||
assert_eq!(events.len(), 1);
|
||||
match &events[0] {
|
||||
TauriAgentEvent::ModelChange { model, mode } => {
|
||||
assert_eq!(model, "claude-3");
|
||||
assert_eq!(mode, "chat");
|
||||
}
|
||||
_ => panic!("Expected ModelChange event"),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -9,7 +9,13 @@
|
||||
//! - native_agent - 核心 Agent 逻辑
|
||||
//! - tool_loop - 工具调用循环
|
||||
//! - tools/ - 工具实现
|
||||
//! - aster_state - Aster Agent 状态管理(新)
|
||||
//! - aster_agent - Aster Agent 包装器(新)
|
||||
//! - event_converter - Aster 事件转换器(新)
|
||||
|
||||
pub mod aster_agent;
|
||||
pub mod aster_state;
|
||||
pub mod event_converter;
|
||||
pub mod native_agent;
|
||||
pub mod parsers;
|
||||
pub mod protocols;
|
||||
@@ -17,6 +23,9 @@ pub mod tool_loop;
|
||||
pub mod tools;
|
||||
pub mod types;
|
||||
|
||||
pub use aster_agent::{AsterAgentWrapper, SessionDetail, SessionInfo};
|
||||
pub use aster_state::AsterAgentState;
|
||||
pub use event_converter::{convert_agent_event, TauriAgentEvent};
|
||||
pub use native_agent::{NativeAgent, NativeAgentState};
|
||||
pub use parsers::{AnthropicSSEParser, OpenAISSEParser};
|
||||
pub use protocols::{create_protocol, AnthropicProtocol, OpenAIProtocol, Protocol};
|
||||
|
||||
@@ -524,6 +524,7 @@ impl NativeAgent {
|
||||
content: Some(OpenAIMessageContent::Text(prompt.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -552,6 +553,7 @@ impl NativeAgent {
|
||||
content: Some(OpenAIMessageContent::Parts(parts)),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
}
|
||||
} else {
|
||||
ChatMessage {
|
||||
@@ -559,6 +561,7 @@ impl NativeAgent {
|
||||
content: Some(OpenAIMessageContent::Text(user_message.to_string())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
}
|
||||
};
|
||||
|
||||
@@ -606,6 +609,7 @@ impl NativeAgent {
|
||||
.collect()
|
||||
}),
|
||||
tool_call_id: msg.tool_call_id.clone(),
|
||||
reasoning_content: msg.reasoning_content.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -642,6 +646,7 @@ impl NativeAgent {
|
||||
timestamp: chrono::Utc::now().to_rfc3339(),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
session.updated_at = chrono::Utc::now().to_rfc3339();
|
||||
}
|
||||
@@ -662,6 +667,7 @@ impl NativeAgent {
|
||||
timestamp: chrono::Utc::now().to_rfc3339(),
|
||||
tool_calls,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
session.updated_at = chrono::Utc::now().to_rfc3339();
|
||||
}
|
||||
|
||||
@@ -31,6 +31,8 @@ struct ToolCallDelta {
|
||||
pub struct OpenAISSEParser {
|
||||
/// 累积的完整内容
|
||||
full_content: String,
|
||||
/// 累积的推理内容(DeepSeek R1 等模型)
|
||||
reasoning_content: String,
|
||||
/// 当前正在构建的工具调用索引
|
||||
current_tool_indices: HashMap<usize, ToolCallDelta>,
|
||||
}
|
||||
@@ -103,6 +105,15 @@ impl OpenAISSEParser {
|
||||
s.to_string()
|
||||
});
|
||||
|
||||
// 提取推理内容(DeepSeek R1 等模型)
|
||||
if let Some(reasoning) = delta
|
||||
.get("reasoning_content")
|
||||
.and_then(|c| c.as_str())
|
||||
.filter(|s| !s.is_empty())
|
||||
{
|
||||
self.reasoning_content.push_str(reasoning);
|
||||
}
|
||||
|
||||
// 提取工具调用
|
||||
if let Some(tool_calls) = delta.get("tool_calls").and_then(|tc| tc.as_array()) {
|
||||
for tc in tool_calls {
|
||||
@@ -181,6 +192,15 @@ impl OpenAISSEParser {
|
||||
self.full_content.clone()
|
||||
}
|
||||
|
||||
/// 获取推理内容(DeepSeek R1 等模型)
|
||||
pub fn get_reasoning_content(&self) -> Option<String> {
|
||||
if self.reasoning_content.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(self.reasoning_content.clone())
|
||||
}
|
||||
}
|
||||
|
||||
/// 是否有工具调用
|
||||
pub fn has_tool_calls(&self) -> bool {
|
||||
!self.current_tool_indices.is_empty()
|
||||
|
||||
@@ -359,6 +359,7 @@ impl AnthropicProtocol {
|
||||
content: full_content,
|
||||
tool_calls,
|
||||
usage,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -396,6 +397,7 @@ impl AnthropicProtocol {
|
||||
content: full_content,
|
||||
tool_calls,
|
||||
usage,
|
||||
reasoning_content: None,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -62,6 +62,7 @@ impl OpenAIProtocol {
|
||||
.collect()
|
||||
}),
|
||||
tool_call_id: msg.tool_call_id.clone(),
|
||||
reasoning_content: msg.reasoning_content.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -81,6 +82,7 @@ impl OpenAIProtocol {
|
||||
content: Some(OpenAIMessageContent::Text(prompt.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -109,6 +111,7 @@ impl OpenAIProtocol {
|
||||
content: Some(OpenAIMessageContent::Parts(parts)),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
}
|
||||
} else {
|
||||
ChatMessage {
|
||||
@@ -116,6 +119,7 @@ impl OpenAIProtocol {
|
||||
content: Some(OpenAIMessageContent::Text(user_message.to_string())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
}
|
||||
};
|
||||
|
||||
@@ -137,6 +141,7 @@ impl OpenAIProtocol {
|
||||
content: Some(OpenAIMessageContent::Text(prompt.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -218,6 +223,7 @@ impl OpenAIProtocol {
|
||||
content,
|
||||
tool_calls: None,
|
||||
usage,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -260,6 +266,7 @@ impl OpenAIProtocol {
|
||||
content: full_content,
|
||||
tool_calls,
|
||||
usage: final_usage,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -317,11 +324,13 @@ impl OpenAIProtocol {
|
||||
content,
|
||||
tool_calls: None,
|
||||
usage,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let full_content = parser.get_full_content();
|
||||
let reasoning_content = parser.get_reasoning_content();
|
||||
let tool_calls = if parser.has_tool_calls() {
|
||||
Some(parser.finalize_tool_calls())
|
||||
} else {
|
||||
@@ -340,6 +349,7 @@ impl OpenAIProtocol {
|
||||
content: full_content,
|
||||
tool_calls,
|
||||
usage: final_usage,
|
||||
reasoning_content,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -84,6 +84,7 @@ impl ToolCallResult {
|
||||
timestamp: chrono::Utc::now().to_rfc3339(),
|
||||
tool_calls: None,
|
||||
tool_call_id: Some(self.tool_call_id.clone()),
|
||||
reasoning_content: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -301,6 +302,7 @@ impl ToolLoopEngine {
|
||||
.collect()
|
||||
}),
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -458,6 +460,7 @@ mod tests {
|
||||
content: "".to_string(),
|
||||
tool_calls: Some(vec![]),
|
||||
usage: None,
|
||||
reasoning_content: None,
|
||||
};
|
||||
assert!(!ToolLoopEngine::has_tool_calls(&result_empty_tools));
|
||||
}
|
||||
@@ -1045,6 +1048,7 @@ mod proptests {
|
||||
content: content.clone(),
|
||||
tool_calls: Some(vec![]),
|
||||
usage: None,
|
||||
reasoning_content: None,
|
||||
};
|
||||
|
||||
// 验证:should_continue 返回 false
|
||||
|
||||
@@ -0,0 +1,468 @@
|
||||
//! Browser 工具模块
|
||||
//!
|
||||
//! 提供浏览器自动化功能,基于 Playwright
|
||||
//! 专为 AI Agent 设计,提供结构化的页面快照
|
||||
|
||||
#![allow(dead_code)]
|
||||
|
||||
use super::registry::Tool;
|
||||
use super::types::{JsonSchema, PropertySchema, ToolDefinition, ToolError, ToolResult};
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::path::PathBuf;
|
||||
use std::process::Stdio;
|
||||
use std::time::Duration;
|
||||
use tokio::process::Command;
|
||||
use tokio::time::timeout;
|
||||
use tracing::info;
|
||||
|
||||
/// 默认超时时间(秒)
|
||||
const DEFAULT_TIMEOUT_SECS: u64 = 30;
|
||||
|
||||
/// Playwright 脚本目录
|
||||
const PLAYWRIGHT_SCRIPTS_DIR: &str = "scripts/playwright";
|
||||
|
||||
/// 浏览器操作类型
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum BrowserAction {
|
||||
/// 打开页面
|
||||
Open { url: String },
|
||||
/// 获取页面快照(AI 友好的可访问性树)
|
||||
Snapshot {
|
||||
#[serde(default)]
|
||||
interactive_only: bool,
|
||||
},
|
||||
/// 点击元素
|
||||
Click { selector: String },
|
||||
/// 填充表单
|
||||
Fill { selector: String, value: String },
|
||||
/// 输入文本(逐字符)
|
||||
Type { selector: String, text: String },
|
||||
/// 按键
|
||||
Press { key: String },
|
||||
/// 滚动页面
|
||||
Scroll {
|
||||
direction: ScrollDirection,
|
||||
#[serde(default = "default_scroll_amount")]
|
||||
amount: i32,
|
||||
},
|
||||
/// 等待元素
|
||||
WaitFor {
|
||||
selector: String,
|
||||
#[serde(default = "default_wait_timeout")]
|
||||
timeout_ms: u64,
|
||||
},
|
||||
/// 截图
|
||||
Screenshot {
|
||||
#[serde(default)]
|
||||
full_page: bool,
|
||||
path: Option<String>,
|
||||
},
|
||||
/// 获取页面文本内容
|
||||
GetText { selector: Option<String> },
|
||||
/// 执行 JavaScript
|
||||
Evaluate { script: String },
|
||||
/// 关闭浏览器
|
||||
Close,
|
||||
}
|
||||
|
||||
fn default_scroll_amount() -> i32 {
|
||||
500
|
||||
}
|
||||
|
||||
fn default_wait_timeout() -> u64 {
|
||||
5000
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ScrollDirection {
|
||||
Up,
|
||||
Down,
|
||||
Left,
|
||||
Right,
|
||||
}
|
||||
|
||||
/// 浏览器操作结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct BrowserResult {
|
||||
/// 是否成功
|
||||
pub success: bool,
|
||||
/// 输出内容
|
||||
pub output: String,
|
||||
/// 页面 URL
|
||||
pub url: Option<String>,
|
||||
/// 页面标题
|
||||
pub title: Option<String>,
|
||||
/// 截图 base64(如果有)
|
||||
pub screenshot: Option<String>,
|
||||
/// 错误信息
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
/// Browser 工具
|
||||
///
|
||||
/// 提供浏览器自动化功能,专为 AI Agent 设计
|
||||
pub struct BrowserTool {
|
||||
/// Playwright 脚本路径
|
||||
script_path: PathBuf,
|
||||
/// 超时时间(秒)
|
||||
timeout_secs: u64,
|
||||
/// 是否使用 headless 模式
|
||||
headless: bool,
|
||||
}
|
||||
|
||||
impl BrowserTool {
|
||||
/// 创建新的 Browser 工具
|
||||
pub fn new() -> Self {
|
||||
// 获取脚本路径(相对于项目根目录)
|
||||
let script_path = std::env::current_dir()
|
||||
.unwrap_or_default()
|
||||
.join(PLAYWRIGHT_SCRIPTS_DIR)
|
||||
.join("browser-tool.mjs");
|
||||
|
||||
Self {
|
||||
script_path,
|
||||
timeout_secs: DEFAULT_TIMEOUT_SECS,
|
||||
headless: true,
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置脚本路径
|
||||
pub fn with_script_path(mut self, path: PathBuf) -> Self {
|
||||
self.script_path = path;
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置超时时间
|
||||
pub fn with_timeout(mut self, timeout_secs: u64) -> Self {
|
||||
self.timeout_secs = timeout_secs;
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置是否 headless
|
||||
pub fn with_headless(mut self, headless: bool) -> Self {
|
||||
self.headless = headless;
|
||||
self
|
||||
}
|
||||
|
||||
/// 执行浏览器操作
|
||||
async fn execute_action(&self, action: &BrowserAction) -> Result<BrowserResult, ToolError> {
|
||||
let action_json = serde_json::to_string(action)
|
||||
.map_err(|e| ToolError::ExecutionFailed(format!("序列化操作失败: {}", e)))?;
|
||||
|
||||
info!("[BrowserTool] 执行操作: {:?}", action);
|
||||
|
||||
// 构建命令
|
||||
let mut cmd = Command::new("node");
|
||||
cmd.arg(&self.script_path);
|
||||
cmd.arg("--action");
|
||||
cmd.arg(&action_json);
|
||||
|
||||
if self.headless {
|
||||
cmd.arg("--headless");
|
||||
}
|
||||
|
||||
cmd.stdin(Stdio::null());
|
||||
cmd.stdout(Stdio::piped());
|
||||
cmd.stderr(Stdio::piped());
|
||||
|
||||
// 执行命令
|
||||
let timeout_duration = Duration::from_secs(self.timeout_secs);
|
||||
let result = timeout(timeout_duration, cmd.output()).await;
|
||||
|
||||
match result {
|
||||
Ok(Ok(output)) => {
|
||||
let stdout = String::from_utf8_lossy(&output.stdout);
|
||||
let stderr = String::from_utf8_lossy(&output.stderr);
|
||||
|
||||
if output.status.success() {
|
||||
// 解析 JSON 输出
|
||||
serde_json::from_str(&stdout).map_err(|e| {
|
||||
ToolError::ExecutionFailed(format!(
|
||||
"解析输出失败: {}\nstdout: {}\nstderr: {}",
|
||||
e, stdout, stderr
|
||||
))
|
||||
})
|
||||
} else {
|
||||
Err(ToolError::ExecutionFailed(format!(
|
||||
"浏览器操作失败: {}",
|
||||
stderr
|
||||
)))
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => Err(ToolError::ExecutionFailed(format!("执行命令失败: {}", e))),
|
||||
Err(_) => Err(ToolError::Timeout),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for BrowserTool {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for BrowserTool {
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
ToolDefinition::new(
|
||||
"browser",
|
||||
"Control a web browser for automation tasks. Use this to navigate websites, \
|
||||
interact with elements, fill forms, and extract information. The 'snapshot' \
|
||||
action returns an accessibility tree with element references (like @e1, @e2) \
|
||||
that can be used in subsequent actions.",
|
||||
)
|
||||
.with_parameters(
|
||||
JsonSchema::new()
|
||||
.add_property(
|
||||
"action",
|
||||
PropertySchema::string(
|
||||
"The browser action to perform. One of: open, snapshot, click, fill, \
|
||||
type, press, scroll, wait_for, screenshot, get_text, evaluate, close",
|
||||
),
|
||||
true,
|
||||
)
|
||||
.add_property(
|
||||
"url",
|
||||
PropertySchema::string("URL to open (for 'open' action)"),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"selector",
|
||||
PropertySchema::string(
|
||||
"Element selector. Can be CSS selector, XPath, or element reference \
|
||||
like @e1 from snapshot output",
|
||||
),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"value",
|
||||
PropertySchema::string("Value to fill (for 'fill' action)"),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"text",
|
||||
PropertySchema::string("Text to type (for 'type' action)"),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"key",
|
||||
PropertySchema::string(
|
||||
"Key to press (for 'press' action), e.g., 'Enter', 'Tab', 'Escape'",
|
||||
),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"direction",
|
||||
PropertySchema::string("Scroll direction: up, down, left, right"),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"script",
|
||||
PropertySchema::string("JavaScript code to evaluate (for 'evaluate' action)"),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"interactive_only",
|
||||
PropertySchema::boolean(
|
||||
"Only include interactive elements in snapshot (default: false)",
|
||||
),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"full_page",
|
||||
PropertySchema::boolean("Capture full page screenshot (default: false)"),
|
||||
false,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
async fn execute(&self, args: serde_json::Value) -> Result<ToolResult, ToolError> {
|
||||
let action_str = args
|
||||
.get("action")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ToolError::InvalidArguments("缺少 action 参数".to_string()))?;
|
||||
|
||||
// 解析操作
|
||||
let action = match action_str {
|
||||
"open" => {
|
||||
let url = args.get("url").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidArguments("open 操作需要 url 参数".to_string())
|
||||
})?;
|
||||
BrowserAction::Open {
|
||||
url: url.to_string(),
|
||||
}
|
||||
}
|
||||
"snapshot" => {
|
||||
let interactive_only = args
|
||||
.get("interactive_only")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
BrowserAction::Snapshot { interactive_only }
|
||||
}
|
||||
"click" => {
|
||||
let selector = args
|
||||
.get("selector")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| {
|
||||
ToolError::InvalidArguments("click 操作需要 selector 参数".to_string())
|
||||
})?;
|
||||
BrowserAction::Click {
|
||||
selector: selector.to_string(),
|
||||
}
|
||||
}
|
||||
"fill" => {
|
||||
let selector = args
|
||||
.get("selector")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| {
|
||||
ToolError::InvalidArguments("fill 操作需要 selector 参数".to_string())
|
||||
})?;
|
||||
let value = args.get("value").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidArguments("fill 操作需要 value 参数".to_string())
|
||||
})?;
|
||||
BrowserAction::Fill {
|
||||
selector: selector.to_string(),
|
||||
value: value.to_string(),
|
||||
}
|
||||
}
|
||||
"type" => {
|
||||
let selector = args
|
||||
.get("selector")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| {
|
||||
ToolError::InvalidArguments("type 操作需要 selector 参数".to_string())
|
||||
})?;
|
||||
let text = args.get("text").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidArguments("type 操作需要 text 参数".to_string())
|
||||
})?;
|
||||
BrowserAction::Type {
|
||||
selector: selector.to_string(),
|
||||
text: text.to_string(),
|
||||
}
|
||||
}
|
||||
"press" => {
|
||||
let key = args.get("key").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidArguments("press 操作需要 key 参数".to_string())
|
||||
})?;
|
||||
BrowserAction::Press {
|
||||
key: key.to_string(),
|
||||
}
|
||||
}
|
||||
"scroll" => {
|
||||
let direction = args
|
||||
.get("direction")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("down");
|
||||
let direction = match direction {
|
||||
"up" => ScrollDirection::Up,
|
||||
"down" => ScrollDirection::Down,
|
||||
"left" => ScrollDirection::Left,
|
||||
"right" => ScrollDirection::Right,
|
||||
_ => ScrollDirection::Down,
|
||||
};
|
||||
let amount = args
|
||||
.get("amount")
|
||||
.and_then(|v| v.as_i64())
|
||||
.map(|v| v as i32)
|
||||
.unwrap_or(500);
|
||||
BrowserAction::Scroll { direction, amount }
|
||||
}
|
||||
"wait_for" => {
|
||||
let selector = args
|
||||
.get("selector")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| {
|
||||
ToolError::InvalidArguments("wait_for 操作需要 selector 参数".to_string())
|
||||
})?;
|
||||
let timeout_ms = args
|
||||
.get("timeout_ms")
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(5000);
|
||||
BrowserAction::WaitFor {
|
||||
selector: selector.to_string(),
|
||||
timeout_ms,
|
||||
}
|
||||
}
|
||||
"screenshot" => {
|
||||
let full_page = args
|
||||
.get("full_page")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
let path = args.get("path").and_then(|v| v.as_str()).map(String::from);
|
||||
BrowserAction::Screenshot { full_page, path }
|
||||
}
|
||||
"get_text" => {
|
||||
let selector = args
|
||||
.get("selector")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
BrowserAction::GetText { selector }
|
||||
}
|
||||
"evaluate" => {
|
||||
let script = args.get("script").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidArguments("evaluate 操作需要 script 参数".to_string())
|
||||
})?;
|
||||
BrowserAction::Evaluate {
|
||||
script: script.to_string(),
|
||||
}
|
||||
}
|
||||
"close" => BrowserAction::Close,
|
||||
_ => {
|
||||
return Err(ToolError::InvalidArguments(format!(
|
||||
"未知的操作: {}",
|
||||
action_str
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
// 执行操作
|
||||
let result = self.execute_action(&action).await?;
|
||||
|
||||
// 构建输出
|
||||
let mut output = result.output;
|
||||
|
||||
if let Some(url) = &result.url {
|
||||
output = format!("URL: {}\n{}", url, output);
|
||||
}
|
||||
if let Some(title) = &result.title {
|
||||
output = format!("Title: {}\n{}", title, output);
|
||||
}
|
||||
|
||||
if result.success {
|
||||
Ok(ToolResult::success(output))
|
||||
} else {
|
||||
Ok(ToolResult::failure_with_output(
|
||||
output,
|
||||
result.error.unwrap_or_else(|| "未知错误".to_string()),
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_tool_definition() {
|
||||
let tool = BrowserTool::new();
|
||||
let def = tool.definition();
|
||||
|
||||
assert_eq!(def.name, "browser");
|
||||
assert!(!def.description.is_empty());
|
||||
assert!(def.parameters.required.contains(&"action".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_action_serialization() {
|
||||
let action = BrowserAction::Open {
|
||||
url: "https://example.com".to_string(),
|
||||
};
|
||||
let json = serde_json::to_string(&action).unwrap();
|
||||
assert!(json.contains("open"));
|
||||
assert!(json.contains("https://example.com"));
|
||||
}
|
||||
}
|
||||
@@ -14,6 +14,7 @@
|
||||
//! - `prompt`: 工具 Prompt 生成器(System Prompt 工具注入)
|
||||
|
||||
pub mod bash;
|
||||
pub mod browser;
|
||||
pub mod edit_file;
|
||||
pub mod prompt;
|
||||
pub mod read_file;
|
||||
@@ -25,6 +26,7 @@ pub mod types;
|
||||
pub mod write_file;
|
||||
|
||||
pub use bash::{BashExecutionResult, BashTool, ShellType};
|
||||
pub use browser::{BrowserAction, BrowserResult, BrowserTool};
|
||||
pub use edit_file::{EditFileResult, EditFileTool, UndoResult};
|
||||
pub use prompt::{generate_tools_prompt, PromptFormat, ToolPromptGenerator};
|
||||
pub use read_file::{ReadFileResult, ReadFileTool};
|
||||
|
||||
@@ -116,6 +116,10 @@ pub struct AgentMessage {
|
||||
/// 工具调用 ID(tool 角色消息需要)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_call_id: Option<String>,
|
||||
/// 推理内容(DeepSeek R1 等模型的思维链内容)
|
||||
/// DeepSeek Reasoner 在 Tool Calls 场景下要求此字段
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_content: Option<String>,
|
||||
}
|
||||
|
||||
/// 消息内容类型
|
||||
@@ -468,6 +472,9 @@ pub struct StreamResult {
|
||||
/// Token 使用量
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub usage: Option<TokenUsage>,
|
||||
/// 推理内容(DeepSeek R1 等模型的思维链内容)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_content: Option<String>,
|
||||
}
|
||||
|
||||
impl StreamResult {
|
||||
@@ -477,6 +484,7 @@ impl StreamResult {
|
||||
content,
|
||||
tool_calls: None,
|
||||
usage: None,
|
||||
reasoning_content: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -492,6 +500,12 @@ impl StreamResult {
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置推理内容
|
||||
pub fn with_reasoning_content(mut self, reasoning_content: String) -> Self {
|
||||
self.reasoning_content = Some(reasoning_content);
|
||||
self
|
||||
}
|
||||
|
||||
/// 是否有工具调用
|
||||
pub fn has_tool_calls(&self) -> bool {
|
||||
self.tool_calls
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::agent::AsterAgentState;
|
||||
use crate::agent::NativeAgentState;
|
||||
use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState;
|
||||
use crate::commands::connect_cmd::ConnectStateWrapper;
|
||||
@@ -134,6 +135,7 @@ pub struct AppStates {
|
||||
pub enhanced_stats_service: EnhancedStatsServiceState,
|
||||
pub batch_operations: BatchOperationsState,
|
||||
pub native_agent: NativeAgentState,
|
||||
pub aster_agent: AsterAgentState,
|
||||
pub oauth_plugin_manager: crate::commands::oauth_plugin_cmd::OAuthPluginManagerState,
|
||||
pub orchestrator: OrchestratorState,
|
||||
pub connect_state: ConnectStateWrapper,
|
||||
@@ -213,6 +215,7 @@ pub fn init_states(config: &Config) -> Result<AppStates, String> {
|
||||
|
||||
// 其他状态
|
||||
let native_agent_state = NativeAgentState::new();
|
||||
let aster_agent_state = AsterAgentState::new();
|
||||
let oauth_plugin_manager_state =
|
||||
crate::commands::oauth_plugin_cmd::OAuthPluginManagerState::with_defaults();
|
||||
let orchestrator_state = OrchestratorState::new();
|
||||
@@ -275,6 +278,7 @@ pub fn init_states(config: &Config) -> Result<AppStates, String> {
|
||||
enhanced_stats_service: enhanced_stats_service_state,
|
||||
batch_operations: batch_operations_state,
|
||||
native_agent: native_agent_state,
|
||||
aster_agent: aster_agent_state,
|
||||
oauth_plugin_manager: oauth_plugin_manager_state,
|
||||
orchestrator: orchestrator_state,
|
||||
connect_state,
|
||||
|
||||
@@ -119,6 +119,7 @@ pub async fn check_api_compatibility(
|
||||
)),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
}],
|
||||
temperature: None,
|
||||
max_tokens: Some(100),
|
||||
@@ -155,6 +156,7 @@ pub async fn check_api_compatibility(
|
||||
)),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
}],
|
||||
temperature: None,
|
||||
max_tokens: Some(10),
|
||||
|
||||
@@ -25,13 +25,18 @@ pub async fn save_config(
|
||||
) -> Result<(), String> {
|
||||
let host = config.server.host.to_lowercase();
|
||||
|
||||
tracing::info!("[CONFIG] 保存配置请求: host={}, port={}", host, config.server.port);
|
||||
tracing::info!(
|
||||
"[CONFIG] 保存配置请求: host={}, port={}",
|
||||
host,
|
||||
config.server.port
|
||||
);
|
||||
|
||||
// 验证绑定地址
|
||||
if !is_valid_bind_host(&host) {
|
||||
tracing::warn!("[CONFIG] 无效的监听地址: {}", host);
|
||||
return Err(
|
||||
"无效的监听地址。允许的地址:127.0.0.1、localhost、::1、0.0.0.0、:: 或局域网 IP".to_string(),
|
||||
"无效的监听地址。允许的地址:127.0.0.1、localhost、::1、0.0.0.0、:: 或局域网 IP"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
@@ -43,7 +48,7 @@ pub async fn save_config(
|
||||
|
||||
let mut s = state.write().await;
|
||||
s.config = config.clone();
|
||||
|
||||
|
||||
match config::save_config(&config) {
|
||||
Ok(()) => {
|
||||
tracing::info!("[CONFIG] 配置保存成功: host={}", config.server.host);
|
||||
|
||||
@@ -71,6 +71,7 @@ pub fn run() {
|
||||
enhanced_stats_service: enhanced_stats_service_state,
|
||||
batch_operations: batch_operations_state,
|
||||
native_agent: native_agent_state,
|
||||
aster_agent: aster_agent_state,
|
||||
oauth_plugin_manager: oauth_plugin_manager_state,
|
||||
orchestrator: orchestrator_state,
|
||||
connect_state,
|
||||
@@ -153,6 +154,7 @@ pub fn run() {
|
||||
.manage(enhanced_stats_service_state)
|
||||
.manage(batch_operations_state)
|
||||
.manage(native_agent_state)
|
||||
.manage(aster_agent_state)
|
||||
.manage(oauth_plugin_manager_state)
|
||||
.manage(orchestrator_state)
|
||||
.manage(connect_state)
|
||||
@@ -883,6 +885,7 @@ pub fn run() {
|
||||
commands::plugin_cmd::get_plugins_with_ui,
|
||||
commands::plugin_cmd::read_plugin_manifest_cmd,
|
||||
commands::plugin_cmd::launch_plugin_ui,
|
||||
commands::plugin_cmd::frontend_debug_log,
|
||||
// Plugin RPC commands
|
||||
commands::plugin_rpc_cmd::plugin_rpc_connect,
|
||||
commands::plugin_rpc_cmd::plugin_rpc_disconnect,
|
||||
@@ -1072,6 +1075,16 @@ pub fn run() {
|
||||
commands::native_agent_cmd::native_agent_get_session,
|
||||
commands::native_agent_cmd::native_agent_delete_session,
|
||||
commands::native_agent_cmd::native_agent_list_sessions,
|
||||
// Aster Agent commands
|
||||
commands::aster_agent_cmd::aster_agent_init,
|
||||
commands::aster_agent_cmd::aster_agent_status,
|
||||
commands::aster_agent_cmd::aster_agent_configure_provider,
|
||||
commands::aster_agent_cmd::aster_agent_chat_stream,
|
||||
commands::aster_agent_cmd::aster_agent_stop,
|
||||
commands::aster_agent_cmd::aster_session_create,
|
||||
commands::aster_agent_cmd::aster_session_list,
|
||||
commands::aster_agent_cmd::aster_session_get,
|
||||
commands::aster_agent_cmd::aster_agent_confirm,
|
||||
// Models config commands
|
||||
commands::models_cmd::get_models_config,
|
||||
commands::models_cmd::save_models_config,
|
||||
|
||||
@@ -0,0 +1,312 @@
|
||||
//! Aster Agent 命令模块
|
||||
//!
|
||||
//! 提供基于 Aster 框架的 Tauri 命令
|
||||
//! 这是新的对话系统实现,与 native_agent_cmd.rs 并行存在
|
||||
|
||||
use crate::agent::aster_state::{ProviderConfig, SessionConfigBuilder};
|
||||
use crate::agent::event_converter::convert_agent_event;
|
||||
use crate::agent::{
|
||||
AsterAgentState, AsterAgentWrapper, SessionDetail, SessionInfo, TauriAgentEvent,
|
||||
};
|
||||
use aster::conversation::message::Message;
|
||||
use futures::StreamExt;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::path::PathBuf;
|
||||
use tauri::{AppHandle, Emitter, State};
|
||||
|
||||
/// Aster Agent 状态信息
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct AsterAgentStatus {
|
||||
pub initialized: bool,
|
||||
pub provider_configured: bool,
|
||||
pub provider_name: Option<String>,
|
||||
pub model_name: Option<String>,
|
||||
}
|
||||
|
||||
/// Provider 配置请求
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ConfigureProviderRequest {
|
||||
pub provider_name: String,
|
||||
pub model_name: String,
|
||||
#[serde(default)]
|
||||
pub api_key: Option<String>,
|
||||
#[serde(default)]
|
||||
pub base_url: Option<String>,
|
||||
}
|
||||
|
||||
/// 初始化 Aster Agent
|
||||
#[tauri::command]
|
||||
pub async fn aster_agent_init(
|
||||
state: State<'_, AsterAgentState>,
|
||||
) -> Result<AsterAgentStatus, String> {
|
||||
tracing::info!("[AsterAgent] 初始化 Agent");
|
||||
|
||||
state.init_agent().await?;
|
||||
|
||||
let provider_config = state.get_provider_config().await;
|
||||
|
||||
tracing::info!("[AsterAgent] Agent 初始化成功");
|
||||
|
||||
Ok(AsterAgentStatus {
|
||||
initialized: true,
|
||||
provider_configured: provider_config.is_some(),
|
||||
provider_name: provider_config.as_ref().map(|c| c.provider_name.clone()),
|
||||
model_name: provider_config.as_ref().map(|c| c.model_name.clone()),
|
||||
})
|
||||
}
|
||||
|
||||
/// 配置 Aster Agent 的 Provider
|
||||
#[tauri::command]
|
||||
pub async fn aster_agent_configure_provider(
|
||||
state: State<'_, AsterAgentState>,
|
||||
request: ConfigureProviderRequest,
|
||||
session_id: String,
|
||||
) -> Result<AsterAgentStatus, String> {
|
||||
tracing::info!(
|
||||
"[AsterAgent] 配置 Provider: {} / {}",
|
||||
request.provider_name,
|
||||
request.model_name
|
||||
);
|
||||
|
||||
let config = ProviderConfig {
|
||||
provider_name: request.provider_name,
|
||||
model_name: request.model_name,
|
||||
api_key: request.api_key,
|
||||
base_url: request.base_url,
|
||||
};
|
||||
|
||||
state
|
||||
.configure_provider(config.clone(), &session_id)
|
||||
.await?;
|
||||
|
||||
Ok(AsterAgentStatus {
|
||||
initialized: true,
|
||||
provider_configured: true,
|
||||
provider_name: Some(config.provider_name),
|
||||
model_name: Some(config.model_name),
|
||||
})
|
||||
}
|
||||
|
||||
/// 获取 Aster Agent 状态
|
||||
#[tauri::command]
|
||||
pub async fn aster_agent_status(
|
||||
state: State<'_, AsterAgentState>,
|
||||
) -> Result<AsterAgentStatus, String> {
|
||||
let provider_config = state.get_provider_config().await;
|
||||
Ok(AsterAgentStatus {
|
||||
initialized: state.is_initialized().await,
|
||||
provider_configured: provider_config.is_some(),
|
||||
provider_name: provider_config.as_ref().map(|c| c.provider_name.clone()),
|
||||
model_name: provider_config.as_ref().map(|c| c.model_name.clone()),
|
||||
})
|
||||
}
|
||||
|
||||
/// 发送消息请求参数
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct AsterChatRequest {
|
||||
pub message: String,
|
||||
pub session_id: String,
|
||||
pub event_name: String,
|
||||
#[serde(default)]
|
||||
pub images: Option<Vec<ImageInput>>,
|
||||
/// Provider 配置(可选,如果未配置则使用当前配置)
|
||||
#[serde(default)]
|
||||
pub provider_config: Option<ConfigureProviderRequest>,
|
||||
}
|
||||
|
||||
/// 图片输入
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ImageInput {
|
||||
pub data: String,
|
||||
pub media_type: String,
|
||||
}
|
||||
|
||||
/// 发送消息并获取流式响应
|
||||
#[tauri::command]
|
||||
pub async fn aster_agent_chat_stream(
|
||||
app: AppHandle,
|
||||
state: State<'_, AsterAgentState>,
|
||||
request: AsterChatRequest,
|
||||
) -> Result<(), String> {
|
||||
tracing::info!(
|
||||
"[AsterAgent] 发送流式消息: session={}, event={}",
|
||||
request.session_id,
|
||||
request.event_name
|
||||
);
|
||||
|
||||
// 确保 Agent 已初始化
|
||||
if !state.is_initialized().await {
|
||||
state.init_agent().await?;
|
||||
}
|
||||
|
||||
// 如果提供了 Provider 配置,则配置 Provider
|
||||
if let Some(provider_config) = &request.provider_config {
|
||||
let config = ProviderConfig {
|
||||
provider_name: provider_config.provider_name.clone(),
|
||||
model_name: provider_config.model_name.clone(),
|
||||
api_key: provider_config.api_key.clone(),
|
||||
base_url: provider_config.base_url.clone(),
|
||||
};
|
||||
state
|
||||
.configure_provider(config, &request.session_id)
|
||||
.await?;
|
||||
}
|
||||
|
||||
// 检查 Provider 是否已配置
|
||||
if !state.is_provider_configured().await {
|
||||
return Err("Provider 未配置,请先调用 aster_agent_configure_provider".to_string());
|
||||
}
|
||||
|
||||
// 创建取消令牌
|
||||
let cancel_token = state.create_cancel_token(&request.session_id).await;
|
||||
|
||||
// 创建用户消息
|
||||
let user_message = Message::user().with_text(&request.message);
|
||||
|
||||
// 创建会话配置
|
||||
let session_config = SessionConfigBuilder::new(&request.session_id).build();
|
||||
|
||||
// 获取 Agent Arc 并保持 guard 在整个流处理期间存活
|
||||
let agent_arc = state.get_agent_arc();
|
||||
let guard = agent_arc.read().await;
|
||||
let agent = guard.as_ref().ok_or("Agent not initialized")?;
|
||||
|
||||
// 获取事件流
|
||||
let stream_result = agent
|
||||
.reply(user_message, session_config, Some(cancel_token.clone()))
|
||||
.await;
|
||||
|
||||
match stream_result {
|
||||
Ok(mut stream) => {
|
||||
// 处理事件流
|
||||
while let Some(event_result) = stream.next().await {
|
||||
match event_result {
|
||||
Ok(agent_event) => {
|
||||
// 转换 Aster 事件为 Tauri 事件
|
||||
let tauri_events = convert_agent_event(agent_event);
|
||||
|
||||
// 发送每个事件到前端
|
||||
for tauri_event in tauri_events {
|
||||
if let Err(e) = app.emit(&request.event_name, &tauri_event) {
|
||||
tracing::error!("[AsterAgent] 发送事件失败: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
// 发送错误事件
|
||||
let error_event = TauriAgentEvent::Error {
|
||||
message: format!("Stream error: {}", e),
|
||||
};
|
||||
if let Err(emit_err) = app.emit(&request.event_name, &error_event) {
|
||||
tracing::error!("[AsterAgent] 发送错误事件失败: {}", emit_err);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 发送完成事件
|
||||
let done_event = TauriAgentEvent::FinalDone { usage: None };
|
||||
if let Err(e) = app.emit(&request.event_name, &done_event) {
|
||||
tracing::error!("[AsterAgent] 发送完成事件失败: {}", e);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
// 发送错误事件
|
||||
let error_event = TauriAgentEvent::Error {
|
||||
message: format!("Agent error: {}", e),
|
||||
};
|
||||
if let Err(emit_err) = app.emit(&request.event_name, &error_event) {
|
||||
tracing::error!("[AsterAgent] 发送错误事件失败: {}", emit_err);
|
||||
}
|
||||
return Err(format!("Agent error: {}", e));
|
||||
}
|
||||
}
|
||||
|
||||
// guard 会在函数结束时自动释放(stream_result 先释放)
|
||||
|
||||
// 清理取消令牌
|
||||
state.remove_cancel_token(&request.session_id).await;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 停止当前会话
|
||||
#[tauri::command]
|
||||
pub async fn aster_agent_stop(
|
||||
state: State<'_, AsterAgentState>,
|
||||
session_id: String,
|
||||
) -> Result<bool, String> {
|
||||
tracing::info!("[AsterAgent] 停止会话: {}", session_id);
|
||||
Ok(state.cancel_session(&session_id).await)
|
||||
}
|
||||
|
||||
/// 创建新会话
|
||||
#[tauri::command]
|
||||
pub async fn aster_session_create(
|
||||
working_dir: Option<String>,
|
||||
name: Option<String>,
|
||||
) -> Result<String, String> {
|
||||
tracing::info!("[AsterAgent] 创建会话: name={:?}", name);
|
||||
let dir = working_dir.map(PathBuf::from);
|
||||
AsterAgentWrapper::create_session(dir, name).await
|
||||
}
|
||||
|
||||
/// 列出所有会话
|
||||
#[tauri::command]
|
||||
pub async fn aster_session_list() -> Result<Vec<SessionInfo>, String> {
|
||||
tracing::info!("[AsterAgent] 列出会话");
|
||||
AsterAgentWrapper::list_sessions().await
|
||||
}
|
||||
|
||||
/// 获取会话详情
|
||||
#[tauri::command]
|
||||
pub async fn aster_session_get(session_id: String) -> Result<SessionDetail, String> {
|
||||
tracing::info!("[AsterAgent] 获取会话: {}", session_id);
|
||||
AsterAgentWrapper::get_session(&session_id).await
|
||||
}
|
||||
|
||||
/// 确认权限请求
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ConfirmRequest {
|
||||
pub request_id: String,
|
||||
pub confirmed: bool,
|
||||
pub response: Option<String>,
|
||||
}
|
||||
|
||||
/// 确认权限请求(用于工具调用确认等)
|
||||
#[tauri::command]
|
||||
pub async fn aster_agent_confirm(
|
||||
state: State<'_, AsterAgentState>,
|
||||
request: ConfirmRequest,
|
||||
) -> Result<(), String> {
|
||||
tracing::info!(
|
||||
"[AsterAgent] 确认请求: id={}, confirmed={}",
|
||||
request.request_id,
|
||||
request.confirmed
|
||||
);
|
||||
|
||||
// TODO: 实现权限确认逻辑
|
||||
// 这需要 Aster 框架支持 confirmation_tx 通道
|
||||
// 目前先返回成功
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_aster_chat_request_deserialize() {
|
||||
let json = r#"{
|
||||
"message": "Hello",
|
||||
"session_id": "test-session",
|
||||
"event_name": "agent_stream"
|
||||
}"#;
|
||||
|
||||
let request: AsterChatRequest = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(request.message, "Hello");
|
||||
assert_eq!(request.session_id, "test-session");
|
||||
assert_eq!(request.event_name, "agent_stream");
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,6 @@
|
||||
pub mod agent_cmd;
|
||||
pub mod api_key_provider_cmd;
|
||||
pub mod aster_agent_cmd;
|
||||
pub mod auto_fix_cmd;
|
||||
pub mod browser_interceptor_cmd;
|
||||
pub mod config_cmd;
|
||||
|
||||
@@ -267,6 +267,7 @@ pub async fn native_agent_chat_stream(
|
||||
timestamp: chrono::Utc::now().to_rfc3339(),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
};
|
||||
if let Err(e) = AgentDao::add_message(&conn, sid, &user_message) {
|
||||
tracing::warn!("[NativeAgent] 保存用户消息到数据库失败: {}", e);
|
||||
@@ -368,6 +369,7 @@ pub async fn native_agent_chat_stream(
|
||||
timestamp: chrono::Utc::now().to_rfc3339(),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
};
|
||||
if let Err(e) = AgentDao::add_message(&conn, sid, &assistant_message) {
|
||||
tracing::warn!("[NativeAgent] 保存助手消息到数据库失败: {}", e);
|
||||
|
||||
@@ -42,7 +42,7 @@ fn get_local_ip() -> Option<String> {
|
||||
socket.connect("8.8.8.8:80").ok()?;
|
||||
let local_addr = socket.local_addr().ok()?;
|
||||
let ip_str = local_addr.ip().to_string();
|
||||
|
||||
|
||||
// 检查是否是 VPN 地址 (198.18.x.x)
|
||||
if let IpAddr::V4(ipv4) = local_addr.ip() {
|
||||
if ipv4.octets()[0] == 198 && (ipv4.octets()[1] == 18 || ipv4.octets()[1] == 19) {
|
||||
@@ -60,7 +60,7 @@ fn get_local_ip() -> Option<String> {
|
||||
return Some("127.0.0.1".to_string());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Some(ip_str)
|
||||
}
|
||||
|
||||
@@ -108,7 +108,6 @@ fn get_all_local_ips() -> Vec<String> {
|
||||
ips
|
||||
}
|
||||
|
||||
|
||||
/// 根据监听地址生成可访问的 URL
|
||||
///
|
||||
/// 用于生成客户端配置中的 API URL。
|
||||
|
||||
@@ -19,6 +19,12 @@ use tokio::sync::RwLock;
|
||||
|
||||
use super::plugin_install_cmd::PluginInstallerState;
|
||||
|
||||
/// 前端调试日志命令
|
||||
#[tauri::command]
|
||||
pub fn frontend_debug_log(message: String) {
|
||||
println!("[Frontend] {}", message);
|
||||
}
|
||||
|
||||
/// 插件管理器状态
|
||||
pub struct PluginManagerState(pub Arc<RwLock<PluginManager>>);
|
||||
|
||||
@@ -144,7 +150,9 @@ pub async fn get_plugins_dir(
|
||||
state: tauri::State<'_, PluginManagerState>,
|
||||
) -> Result<String, String> {
|
||||
let manager = state.0.read().await;
|
||||
Ok(manager.plugins_dir().to_string_lossy().to_string())
|
||||
let dir = manager.plugins_dir().to_string_lossy().to_string();
|
||||
println!("[get_plugins_dir] 返回: {}", dir);
|
||||
Ok(dir)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
@@ -226,6 +234,18 @@ pub async fn get_plugins_with_ui(
|
||||
|
||||
// 获取所有已安装插件(从数据库)
|
||||
let installed_plugins = installer.list_installed().map_err(|e| e.to_string())?;
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] 从数据库获取到 {} 个已安装插件",
|
||||
installed_plugins.len()
|
||||
);
|
||||
for p in &installed_plugins {
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] 数据库插件: id={}, install_path={:?}",
|
||||
p.id,
|
||||
p.install_path
|
||||
);
|
||||
}
|
||||
|
||||
let mut registered_ids: std::collections::HashSet<String> =
|
||||
installed_plugins.iter().map(|p| p.id.clone()).collect();
|
||||
|
||||
@@ -254,11 +274,23 @@ pub async fn get_plugins_with_ui(
|
||||
})
|
||||
.collect();
|
||||
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] 从数据库筛选出 {} 个带 UI 的插件",
|
||||
ui_plugins.len()
|
||||
);
|
||||
|
||||
// 扫描插件目录中未注册的插件
|
||||
let plugins_dir = manager.plugins_dir();
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] PluginManager plugins_dir={:?}, exists={}",
|
||||
plugins_dir,
|
||||
plugins_dir.exists()
|
||||
);
|
||||
|
||||
if let Ok(entries) = std::fs::read_dir(plugins_dir) {
|
||||
for entry in entries.flatten() {
|
||||
let path = entry.path();
|
||||
tracing::debug!("[get_plugins_with_ui] 扫描目录: {:?}", path);
|
||||
if path.is_dir() {
|
||||
let plugin_id = path
|
||||
.file_name()
|
||||
@@ -266,14 +298,36 @@ pub async fn get_plugins_with_ui(
|
||||
.map(|s| s.to_string());
|
||||
|
||||
if let Some(id) = plugin_id {
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] 发现目录: id={}, already_registered={}",
|
||||
id,
|
||||
registered_ids.contains(&id)
|
||||
);
|
||||
|
||||
// 跳过已注册的插件
|
||||
if registered_ids.contains(&id) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// 尝试读取 manifest
|
||||
if let Some(manifest) = read_plugin_manifest(&path) {
|
||||
let manifest_result = read_plugin_manifest(&path);
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] 读取 manifest: id={}, success={}",
|
||||
id,
|
||||
manifest_result.is_some()
|
||||
);
|
||||
|
||||
if let Some(manifest) = manifest_result {
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] manifest: name={}, has_ui={}",
|
||||
manifest.name,
|
||||
manifest.ui.is_some()
|
||||
);
|
||||
if let Some(ui_config) = manifest.ui {
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] ui_config: surfaces={:?}",
|
||||
ui_config.surfaces
|
||||
);
|
||||
if !ui_config.surfaces.is_empty() {
|
||||
ui_plugins.push(PluginUIInfo {
|
||||
plugin_id: id.clone(),
|
||||
@@ -289,8 +343,16 @@ pub async fn get_plugins_with_ui(
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
tracing::error!("[get_plugins_with_ui] 无法读取插件目录: {:?}", plugins_dir);
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
"[get_plugins_with_ui] 最终返回 {} 个带 UI 的插件: {:?}",
|
||||
ui_plugins.len(),
|
||||
ui_plugins.iter().map(|p| &p.plugin_id).collect::<Vec<_>>()
|
||||
);
|
||||
|
||||
Ok(ui_plugins)
|
||||
}
|
||||
|
||||
@@ -341,26 +403,58 @@ pub async fn handle_plugin_action(
|
||||
///
|
||||
/// 从插件目录读取 plugin.json 文件
|
||||
/// 用于检查插件是否存在于文件系统中(即使未在数据库中注册)
|
||||
/// 首先尝试从 PluginManager 的 plugins_dir 读取,如果失败则尝试从数据库中的 install_path 读取
|
||||
#[tauri::command]
|
||||
pub async fn read_plugin_manifest_cmd(
|
||||
state: tauri::State<'_, PluginManagerState>,
|
||||
installer_state: tauri::State<'_, PluginInstallerState>,
|
||||
plugin_id: String,
|
||||
) -> Result<Option<PluginManifest>, String> {
|
||||
let manager = state.0.read().await;
|
||||
let plugins_dir = manager.plugins_dir();
|
||||
let plugin_path = plugins_dir.join(&plugin_id);
|
||||
|
||||
tracing::debug!(
|
||||
"read_plugin_manifest_cmd: plugin_id={}, plugins_dir={:?}, plugin_path={:?}",
|
||||
// 使用 println! 确保在 release 版本中也能看到日志
|
||||
println!(
|
||||
"[read_plugin_manifest_cmd] plugin_id={}, plugins_dir={:?}, plugin_path={:?}, exists={}",
|
||||
plugin_id,
|
||||
plugins_dir,
|
||||
plugin_path
|
||||
plugin_path,
|
||||
plugin_path.exists()
|
||||
);
|
||||
|
||||
let result = read_plugin_manifest(&plugin_path);
|
||||
tracing::debug!("read_plugin_manifest_cmd: result={:?}", result.is_some());
|
||||
// 首先尝试从 PluginManager 的 plugins_dir 读取
|
||||
if let Some(manifest) = read_plugin_manifest(&plugin_path) {
|
||||
println!("[read_plugin_manifest_cmd] found in plugins_dir");
|
||||
println!(
|
||||
"[read_plugin_manifest_cmd] manifest: name={}, version={}, plugin_type={:?}, has_ui={}",
|
||||
manifest.name,
|
||||
manifest.version,
|
||||
manifest.plugin_type,
|
||||
manifest.ui.is_some()
|
||||
);
|
||||
// 输出序列化后的 JSON
|
||||
if let Ok(json) = serde_json::to_string(&manifest) {
|
||||
println!("[read_plugin_manifest_cmd] JSON: {}", json);
|
||||
}
|
||||
return Ok(Some(manifest));
|
||||
}
|
||||
|
||||
Ok(result)
|
||||
// 如果失败,尝试从数据库中的 install_path 读取
|
||||
let installer = installer_state.0.read().await;
|
||||
if let Ok(Some(installed)) = installer.get_plugin(&plugin_id) {
|
||||
println!(
|
||||
"[read_plugin_manifest_cmd] trying install_path={:?}",
|
||||
installed.install_path
|
||||
);
|
||||
if let Some(manifest) = read_plugin_manifest(&installed.install_path) {
|
||||
println!("[read_plugin_manifest_cmd] found in install_path");
|
||||
return Ok(Some(manifest));
|
||||
}
|
||||
}
|
||||
|
||||
println!("[read_plugin_manifest_cmd] not found");
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
/// 启动插件 UI(用于 binary 类型插件)
|
||||
@@ -369,15 +463,29 @@ pub async fn read_plugin_manifest_cmd(
|
||||
#[tauri::command]
|
||||
pub async fn launch_plugin_ui(
|
||||
state: tauri::State<'_, PluginManagerState>,
|
||||
installer_state: tauri::State<'_, PluginInstallerState>,
|
||||
plugin_id: String,
|
||||
) -> Result<(), String> {
|
||||
let manager = state.0.read().await;
|
||||
let plugins_dir = manager.plugins_dir();
|
||||
let plugin_path = plugins_dir.join(&plugin_id);
|
||||
|
||||
// 读取插件清单
|
||||
let manifest =
|
||||
read_plugin_manifest(&plugin_path).ok_or_else(|| format!("插件 {} 不存在", plugin_id))?;
|
||||
// 首先尝试从 PluginManager 的 plugins_dir 读取
|
||||
let (manifest, actual_plugin_path) = if let Some(m) = read_plugin_manifest(&plugin_path) {
|
||||
(m, plugin_path)
|
||||
} else {
|
||||
// 如果失败,尝试从数据库中的 install_path 读取
|
||||
let installer = installer_state.0.read().await;
|
||||
if let Ok(Some(installed)) = installer.get_plugin(&plugin_id) {
|
||||
if let Some(m) = read_plugin_manifest(&installed.install_path) {
|
||||
(m, installed.install_path.clone())
|
||||
} else {
|
||||
return Err(format!("插件 {} 不存在", plugin_id));
|
||||
}
|
||||
} else {
|
||||
return Err(format!("插件 {} 不存在", plugin_id));
|
||||
}
|
||||
};
|
||||
|
||||
// 检查是否是 binary 类型
|
||||
if manifest.plugin_type != PluginType::Binary {
|
||||
@@ -390,7 +498,7 @@ pub async fn launch_plugin_ui(
|
||||
.ok_or_else(|| "插件缺少 binary 配置".to_string())?;
|
||||
|
||||
let binary_name = &binary_config.binary_name;
|
||||
let binary_path = plugin_path.join(binary_name);
|
||||
let binary_path = actual_plugin_path.join(binary_name);
|
||||
|
||||
if !binary_path.exists() {
|
||||
return Err(format!("插件二进制文件不存在: {}", binary_path.display()));
|
||||
|
||||
@@ -16,6 +16,7 @@ pub fn convert_anthropic_to_openai(request: &AnthropicMessagesRequest) -> ChatCo
|
||||
content: Some(MessageContent::Text(system_text)),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -83,6 +84,7 @@ fn convert_anthropic_message(msg: &AnthropicMessage) -> Vec<ChatMessage> {
|
||||
content: Some(MessageContent::Text(s.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
serde_json::Value::Array(parts) => {
|
||||
@@ -147,6 +149,7 @@ fn convert_anthropic_message(msg: &AnthropicMessage) -> Vec<ChatMessage> {
|
||||
content,
|
||||
tool_calls: tc,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
// 处理 user 消息
|
||||
@@ -158,6 +161,7 @@ fn convert_anthropic_message(msg: &AnthropicMessage) -> Vec<ChatMessage> {
|
||||
content: Some(MessageContent::Text(content)),
|
||||
tool_calls: None,
|
||||
tool_call_id: Some(tool_use_id),
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -168,6 +172,7 @@ fn convert_anthropic_message(msg: &AnthropicMessage) -> Vec<ChatMessage> {
|
||||
content: Some(MessageContent::Text(text_parts.join(""))),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -201,6 +201,7 @@ impl AgentDao {
|
||||
timestamp,
|
||||
tool_calls,
|
||||
tool_call_id,
|
||||
reasoning_content: None,
|
||||
})
|
||||
})?;
|
||||
|
||||
|
||||
@@ -58,6 +58,10 @@ pub struct ChatMessage {
|
||||
pub tool_calls: Option<Vec<ToolCall>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_call_id: Option<String>,
|
||||
/// 推理内容(DeepSeek R1 等模型的思维链内容)
|
||||
/// DeepSeek Reasoner 在 Tool Calls 场景下要求此字段
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_content: Option<String>,
|
||||
}
|
||||
|
||||
impl ChatMessage {
|
||||
|
||||
@@ -21,3 +21,6 @@ pub use types::{
|
||||
NoopProgressCallback, PackageFormat, ProgressCallback,
|
||||
};
|
||||
pub use validator::PackageValidator;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
@@ -0,0 +1,554 @@
|
||||
//! 插件安装器测试
|
||||
//!
|
||||
//! 测试插件安装、注册、加载的完整流程
|
||||
|
||||
use super::*;
|
||||
use std::fs;
|
||||
use std::io::Write;
|
||||
use tempfile::TempDir;
|
||||
use zip::write::FileOptions;
|
||||
use zip::ZipWriter;
|
||||
|
||||
/// 创建测试用的插件包
|
||||
fn create_test_plugin_zip(dir: &TempDir, plugin_id: &str, version: &str) -> std::path::PathBuf {
|
||||
let zip_path = dir.path().join(format!("{}.zip", plugin_id));
|
||||
let file = fs::File::create(&zip_path).unwrap();
|
||||
let mut zip = ZipWriter::new(file);
|
||||
|
||||
let options = FileOptions::default().compression_method(zip::CompressionMethod::Stored);
|
||||
|
||||
// 创建 plugin.json
|
||||
let manifest = serde_json::json!({
|
||||
"name": plugin_id,
|
||||
"version": version,
|
||||
"description": format!("Test plugin {}", plugin_id),
|
||||
"author": "Test Author",
|
||||
"plugin_type": "script",
|
||||
"entry": "config.json",
|
||||
"ui": {
|
||||
"surfaces": ["tools"],
|
||||
"entry": "dist/index.js",
|
||||
"icon": "Package"
|
||||
}
|
||||
});
|
||||
|
||||
zip.start_file("plugin.json", options).unwrap();
|
||||
zip.write_all(manifest.to_string().as_bytes()).unwrap();
|
||||
|
||||
// 创建 config.json
|
||||
let config = serde_json::json!({
|
||||
"enabled": true
|
||||
});
|
||||
zip.start_file("config.json", options).unwrap();
|
||||
zip.write_all(config.to_string().as_bytes()).unwrap();
|
||||
|
||||
// 创建 dist/index.js
|
||||
zip.start_file("dist/index.js", options).unwrap();
|
||||
zip.write_all(b"// Plugin UI code").unwrap();
|
||||
|
||||
zip.finish().unwrap();
|
||||
zip_path
|
||||
}
|
||||
|
||||
/// 创建测试注册表
|
||||
fn create_test_registry(dir: &TempDir) -> PluginRegistry {
|
||||
let db_path = dir.path().join("test.db");
|
||||
let registry = PluginRegistry::from_path(&db_path).unwrap();
|
||||
registry.init_tables().unwrap();
|
||||
registry
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod registry_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_registry_create_and_register() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let registry = create_test_registry(&temp_dir);
|
||||
|
||||
let plugin = InstalledPlugin {
|
||||
id: "test-plugin".to_string(),
|
||||
name: "Test Plugin".to_string(),
|
||||
version: "1.0.0".to_string(),
|
||||
description: "A test plugin".to_string(),
|
||||
author: Some("Test Author".to_string()),
|
||||
install_path: temp_dir.path().join("test-plugin"),
|
||||
installed_at: chrono::Utc::now(),
|
||||
source: InstallSource::Local {
|
||||
path: "/tmp/test.zip".to_string(),
|
||||
},
|
||||
enabled: true,
|
||||
};
|
||||
|
||||
// 注册插件
|
||||
registry.register(&plugin).unwrap();
|
||||
|
||||
// 验证插件存在
|
||||
assert!(registry.exists("test-plugin").unwrap());
|
||||
assert!(!registry.exists("non-existent").unwrap());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_get_plugin() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let registry = create_test_registry(&temp_dir);
|
||||
|
||||
let plugin = InstalledPlugin {
|
||||
id: "get-test".to_string(),
|
||||
name: "Get Test Plugin".to_string(),
|
||||
version: "2.0.0".to_string(),
|
||||
description: "Plugin for get test".to_string(),
|
||||
author: None,
|
||||
install_path: temp_dir.path().join("get-test"),
|
||||
installed_at: chrono::Utc::now(),
|
||||
source: InstallSource::Url {
|
||||
url: "https://example.com/plugin.zip".to_string(),
|
||||
},
|
||||
enabled: false,
|
||||
};
|
||||
|
||||
registry.register(&plugin).unwrap();
|
||||
|
||||
// 获取插件
|
||||
let retrieved = registry.get("get-test").unwrap().unwrap();
|
||||
assert_eq!(retrieved.id, "get-test");
|
||||
assert_eq!(retrieved.name, "Get Test Plugin");
|
||||
assert_eq!(retrieved.version, "2.0.0");
|
||||
assert!(!retrieved.enabled);
|
||||
|
||||
// 获取不存在的插件
|
||||
let not_found = registry.get("non-existent").unwrap();
|
||||
assert!(not_found.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_list_plugins() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let registry = create_test_registry(&temp_dir);
|
||||
|
||||
// 注册多个插件
|
||||
for i in 1..=3 {
|
||||
let plugin = InstalledPlugin {
|
||||
id: format!("plugin-{}", i),
|
||||
name: format!("Plugin {}", i),
|
||||
version: "1.0.0".to_string(),
|
||||
description: format!("Plugin {} description", i),
|
||||
author: Some("Author".to_string()),
|
||||
install_path: temp_dir.path().join(format!("plugin-{}", i)),
|
||||
installed_at: chrono::Utc::now(),
|
||||
source: InstallSource::Local {
|
||||
path: format!("/tmp/plugin-{}.zip", i),
|
||||
},
|
||||
enabled: true,
|
||||
};
|
||||
registry.register(&plugin).unwrap();
|
||||
}
|
||||
|
||||
// 列出所有插件
|
||||
let plugins = registry.list().unwrap();
|
||||
assert_eq!(plugins.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_unregister_plugin() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let registry = create_test_registry(&temp_dir);
|
||||
|
||||
let plugin = InstalledPlugin {
|
||||
id: "unregister-test".to_string(),
|
||||
name: "Unregister Test".to_string(),
|
||||
version: "1.0.0".to_string(),
|
||||
description: "".to_string(),
|
||||
author: None,
|
||||
install_path: temp_dir.path().join("unregister-test"),
|
||||
installed_at: chrono::Utc::now(),
|
||||
source: InstallSource::Local {
|
||||
path: "/tmp/test.zip".to_string(),
|
||||
},
|
||||
enabled: true,
|
||||
};
|
||||
|
||||
registry.register(&plugin).unwrap();
|
||||
assert!(registry.exists("unregister-test").unwrap());
|
||||
|
||||
// 注销插件
|
||||
registry.unregister("unregister-test").unwrap();
|
||||
assert!(!registry.exists("unregister-test").unwrap());
|
||||
|
||||
// 注销不存在的插件应该返回错误
|
||||
let result = registry.unregister("non-existent");
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_set_enabled() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let registry = create_test_registry(&temp_dir);
|
||||
|
||||
let plugin = InstalledPlugin {
|
||||
id: "enabled-test".to_string(),
|
||||
name: "Enabled Test".to_string(),
|
||||
version: "1.0.0".to_string(),
|
||||
description: "".to_string(),
|
||||
author: None,
|
||||
install_path: temp_dir.path().join("enabled-test"),
|
||||
installed_at: chrono::Utc::now(),
|
||||
source: InstallSource::Local {
|
||||
path: "/tmp/test.zip".to_string(),
|
||||
},
|
||||
enabled: true,
|
||||
};
|
||||
|
||||
registry.register(&plugin).unwrap();
|
||||
|
||||
// 禁用插件
|
||||
registry.set_enabled("enabled-test", false).unwrap();
|
||||
let retrieved = registry.get("enabled-test").unwrap().unwrap();
|
||||
assert!(!retrieved.enabled);
|
||||
|
||||
// 启用插件
|
||||
registry.set_enabled("enabled-test", true).unwrap();
|
||||
let retrieved = registry.get("enabled-test").unwrap().unwrap();
|
||||
assert!(retrieved.enabled);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod installer_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_installer_from_paths() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir.clone(), temp_install_dir.clone(), &db_path)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(installer.plugins_dir(), plugins_dir);
|
||||
assert_eq!(installer.temp_dir(), temp_install_dir);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_installer_list_installed_empty() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir, temp_install_dir, &db_path).unwrap();
|
||||
|
||||
let plugins = installer.list_installed().unwrap();
|
||||
assert!(plugins.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_installer_is_installed() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir, temp_install_dir, &db_path).unwrap();
|
||||
|
||||
// 未安装的插件
|
||||
assert!(!installer.is_installed("non-existent").unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_installer_install_from_file() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
// 创建测试插件包
|
||||
let zip_path = create_test_plugin_zip(&temp_dir, "local-test-plugin", "1.0.0");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir.clone(), temp_install_dir, &db_path).unwrap();
|
||||
|
||||
// 安装插件
|
||||
let progress = NoopProgressCallback;
|
||||
let result = installer.install_from_file(&zip_path, &progress).await;
|
||||
|
||||
assert!(result.is_ok(), "Install failed: {:?}", result.err());
|
||||
|
||||
let installed = result.unwrap();
|
||||
assert_eq!(installed.id, "local-test-plugin");
|
||||
assert_eq!(installed.version, "1.0.0");
|
||||
assert!(installed.enabled);
|
||||
|
||||
// 验证插件已注册
|
||||
assert!(installer.is_installed("local-test-plugin").unwrap());
|
||||
|
||||
// 验证插件文件已复制
|
||||
let plugin_dir = plugins_dir.join("local-test-plugin");
|
||||
assert!(plugin_dir.exists());
|
||||
assert!(plugin_dir.join("plugin.json").exists());
|
||||
assert!(plugin_dir.join("config.json").exists());
|
||||
assert!(plugin_dir.join("dist/index.js").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_installer_uninstall() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
// 创建并安装测试插件
|
||||
let zip_path = create_test_plugin_zip(&temp_dir, "uninstall-test", "1.0.0");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir.clone(), temp_install_dir, &db_path).unwrap();
|
||||
|
||||
let progress = NoopProgressCallback;
|
||||
installer
|
||||
.install_from_file(&zip_path, &progress)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// 验证已安装
|
||||
assert!(installer.is_installed("uninstall-test").unwrap());
|
||||
assert!(plugins_dir.join("uninstall-test").exists());
|
||||
|
||||
// 卸载插件
|
||||
installer.uninstall("uninstall-test").await.unwrap();
|
||||
|
||||
// 验证已卸载
|
||||
assert!(!installer.is_installed("uninstall-test").unwrap());
|
||||
assert!(!plugins_dir.join("uninstall-test").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_installer_reinstall_updates_version() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir, temp_install_dir, &db_path).unwrap();
|
||||
|
||||
let progress = NoopProgressCallback;
|
||||
|
||||
// 安装 v1.0.0
|
||||
let zip_v1 = create_test_plugin_zip(&temp_dir, "version-test", "1.0.0");
|
||||
let installed_v1 = installer
|
||||
.install_from_file(&zip_v1, &progress)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(installed_v1.version, "1.0.0");
|
||||
|
||||
// 安装 v2.0.0(覆盖安装)
|
||||
let zip_v2 = create_test_plugin_zip(&temp_dir, "version-test", "2.0.0");
|
||||
let installed_v2 = installer
|
||||
.install_from_file(&zip_v2, &progress)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(installed_v2.version, "2.0.0");
|
||||
|
||||
// 验证数据库中的版本已更新
|
||||
let plugin = installer.get_plugin("version-test").unwrap().unwrap();
|
||||
assert_eq!(plugin.version, "2.0.0");
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod validator_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_validator_valid_package() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let zip_path = create_test_plugin_zip(&temp_dir, "valid-plugin", "1.0.0");
|
||||
|
||||
let validator = PackageValidator::new();
|
||||
|
||||
// 验证格式
|
||||
let format = validator.validate_format(&zip_path);
|
||||
assert!(
|
||||
format.is_ok(),
|
||||
"Format validation failed: {:?}",
|
||||
format.err()
|
||||
);
|
||||
|
||||
// 提取并验证 manifest
|
||||
let manifest = validator.extract_and_validate_manifest(&zip_path, format.unwrap());
|
||||
assert!(
|
||||
manifest.is_ok(),
|
||||
"Manifest validation failed: {:?}",
|
||||
manifest.err()
|
||||
);
|
||||
|
||||
let manifest = manifest.unwrap();
|
||||
assert_eq!(manifest.name, "valid-plugin");
|
||||
assert_eq!(manifest.version, "1.0.0");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validator_missing_manifest() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let zip_path = temp_dir.path().join("no-manifest.zip");
|
||||
|
||||
// 创建没有 plugin.json 的 zip
|
||||
let file = fs::File::create(&zip_path).unwrap();
|
||||
let mut zip = ZipWriter::new(file);
|
||||
let options = FileOptions::default();
|
||||
zip.start_file("config.json", options).unwrap();
|
||||
zip.write_all(b"{}").unwrap();
|
||||
zip.finish().unwrap();
|
||||
|
||||
let validator = PackageValidator::new();
|
||||
let format = validator.validate_format(&zip_path).unwrap();
|
||||
let result = validator.extract_and_validate_manifest(&zip_path, format);
|
||||
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validator_invalid_manifest_json() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let zip_path = temp_dir.path().join("invalid-json.zip");
|
||||
|
||||
// 创建有无效 JSON 的 zip
|
||||
let file = fs::File::create(&zip_path).unwrap();
|
||||
let mut zip = ZipWriter::new(file);
|
||||
let options = FileOptions::default();
|
||||
zip.start_file("plugin.json", options).unwrap();
|
||||
zip.write_all(b"{ invalid json }").unwrap();
|
||||
zip.finish().unwrap();
|
||||
|
||||
let validator = PackageValidator::new();
|
||||
let format = validator.validate_format(&zip_path).unwrap();
|
||||
let result = validator.extract_and_validate_manifest(&zip_path, format);
|
||||
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validator_missing_required_fields() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let zip_path = temp_dir.path().join("missing-fields.zip");
|
||||
|
||||
// 创建缺少必填字段的 manifest
|
||||
let file = fs::File::create(&zip_path).unwrap();
|
||||
let mut zip = ZipWriter::new(file);
|
||||
let options = FileOptions::default();
|
||||
zip.start_file("plugin.json", options).unwrap();
|
||||
// 缺少 name 和 version
|
||||
zip.write_all(b"{\"description\": \"test\"}").unwrap();
|
||||
zip.finish().unwrap();
|
||||
|
||||
let validator = PackageValidator::new();
|
||||
let format = validator.validate_format(&zip_path).unwrap();
|
||||
let result = validator.extract_and_validate_manifest(&zip_path, format);
|
||||
|
||||
assert!(result.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod integration_tests {
|
||||
use super::*;
|
||||
|
||||
/// 测试完整的安装-加载-卸载流程
|
||||
#[tokio::test]
|
||||
async fn test_full_plugin_lifecycle() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
// 创建测试插件包
|
||||
let zip_path = create_test_plugin_zip(&temp_dir, "lifecycle-test", "1.0.0");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir.clone(), temp_install_dir, &db_path).unwrap();
|
||||
|
||||
let progress = NoopProgressCallback;
|
||||
|
||||
// 1. 安装插件
|
||||
let installed = installer
|
||||
.install_from_file(&zip_path, &progress)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(installed.id, "lifecycle-test");
|
||||
|
||||
// 2. 验证插件已注册
|
||||
assert!(installer.is_installed("lifecycle-test").unwrap());
|
||||
|
||||
// 3. 获取插件信息
|
||||
let plugin_info = installer.get_plugin("lifecycle-test").unwrap().unwrap();
|
||||
assert_eq!(plugin_info.version, "1.0.0");
|
||||
assert!(plugin_info.enabled);
|
||||
|
||||
// 4. 验证插件文件存在
|
||||
let plugin_dir = plugins_dir.join("lifecycle-test");
|
||||
assert!(plugin_dir.join("plugin.json").exists());
|
||||
assert!(plugin_dir.join("dist/index.js").exists());
|
||||
|
||||
// 5. 列出所有插件
|
||||
let all_plugins = installer.list_installed().unwrap();
|
||||
assert_eq!(all_plugins.len(), 1);
|
||||
assert_eq!(all_plugins[0].id, "lifecycle-test");
|
||||
|
||||
// 6. 卸载插件
|
||||
installer.uninstall("lifecycle-test").await.unwrap();
|
||||
|
||||
// 7. 验证插件已移除
|
||||
assert!(!installer.is_installed("lifecycle-test").unwrap());
|
||||
assert!(!plugin_dir.exists());
|
||||
|
||||
// 8. 列表应为空
|
||||
let all_plugins = installer.list_installed().unwrap();
|
||||
assert!(all_plugins.is_empty());
|
||||
}
|
||||
|
||||
/// 测试多个插件同时安装
|
||||
#[tokio::test]
|
||||
async fn test_multiple_plugins() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let plugins_dir = temp_dir.path().join("plugins");
|
||||
let temp_install_dir = temp_dir.path().join("temp");
|
||||
let db_path = temp_dir.path().join("test.db");
|
||||
|
||||
let installer =
|
||||
PluginInstaller::from_paths(plugins_dir, temp_install_dir, &db_path).unwrap();
|
||||
|
||||
let progress = NoopProgressCallback;
|
||||
|
||||
// 安装多个插件
|
||||
let plugin_ids = ["plugin-a", "plugin-b", "plugin-c"];
|
||||
for id in &plugin_ids {
|
||||
let zip_path = create_test_plugin_zip(&temp_dir, id, "1.0.0");
|
||||
installer
|
||||
.install_from_file(&zip_path, &progress)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
// 验证所有插件都已安装
|
||||
let all_plugins = installer.list_installed().unwrap();
|
||||
assert_eq!(all_plugins.len(), 3);
|
||||
|
||||
for id in &plugin_ids {
|
||||
assert!(installer.is_installed(id).unwrap());
|
||||
}
|
||||
|
||||
// 卸载一个插件
|
||||
installer.uninstall("plugin-b").await.unwrap();
|
||||
|
||||
// 验证只剩两个
|
||||
let remaining = installer.list_installed().unwrap();
|
||||
assert_eq!(remaining.len(), 2);
|
||||
assert!(!installer.is_installed("plugin-b").unwrap());
|
||||
}
|
||||
}
|
||||
@@ -19,7 +19,9 @@ impl PluginLoader {
|
||||
}
|
||||
|
||||
pub fn default_plugins_dir() -> PathBuf {
|
||||
dirs::config_dir()
|
||||
// 使用 data_dir 而不是 config_dir,与 PluginInstaller 保持一致
|
||||
// 在 macOS 上两者相同,但在其他平台上可能不同
|
||||
dirs::data_dir()
|
||||
.unwrap_or_else(|| PathBuf::from("."))
|
||||
.join("proxycast")
|
||||
.join("plugins")
|
||||
|
||||
@@ -171,6 +171,40 @@ fn test_plugin_manifest_serde() {
|
||||
assert_eq!(parsed.hooks.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_machine_id_tool_manifest_parsing() {
|
||||
// 测试 machine-id-tool 的 plugin.json 格式
|
||||
let json = r#"{
|
||||
"name": "machine-id-tool",
|
||||
"version": "0.4.0",
|
||||
"description": "Machine ID 管理工具",
|
||||
"author": "ProxyCast Team",
|
||||
"homepage": "https://github.com/aiclientproxy/MachineIdTool",
|
||||
"license": "MIT",
|
||||
"plugin_type": "script",
|
||||
"entry": "config.json",
|
||||
"min_proxycast_version": "1.0.0",
|
||||
"ui": {
|
||||
"surfaces": ["tools", "sidebar"],
|
||||
"entry": "dist/index.js",
|
||||
"icon": "Cpu",
|
||||
"title": "机器码管理工具",
|
||||
"description": "查看、修改和管理系统机器码"
|
||||
}
|
||||
}"#;
|
||||
|
||||
let manifest: PluginManifest = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(manifest.name, "machine-id-tool");
|
||||
assert_eq!(manifest.version, "0.4.0");
|
||||
assert_eq!(manifest.plugin_type, PluginType::Script);
|
||||
assert!(manifest.ui.is_some());
|
||||
|
||||
let ui = manifest.ui.unwrap();
|
||||
assert_eq!(ui.surfaces, vec!["tools", "sidebar"]);
|
||||
assert_eq!(ui.entry, Some("dist/index.js".to_string()));
|
||||
assert_eq!(ui.icon, Some("Cpu".to_string()));
|
||||
}
|
||||
|
||||
// Property-based tests
|
||||
use proptest::prelude::*;
|
||||
|
||||
|
||||
+10
-12
@@ -223,7 +223,7 @@ impl ServerState {
|
||||
}
|
||||
|
||||
/// 解析绑定地址
|
||||
///
|
||||
///
|
||||
/// 直接返回用户配置的地址,不做任何自动替换。
|
||||
/// 如果地址无效,绑定时会失败并返回错误。
|
||||
fn resolve_bind_host(&self, configured_host: &str) -> String {
|
||||
@@ -299,7 +299,7 @@ impl ServerState {
|
||||
// - 局域网 IP:检查是否在当前网卡列表中,如果不在则自动切换到当前局域网 IP
|
||||
let configured_host = self.config.server.host.clone();
|
||||
let host = self.resolve_bind_host(&configured_host);
|
||||
|
||||
|
||||
// 如果地址发生了变化,记录日志
|
||||
if host != configured_host {
|
||||
tracing::warn!(
|
||||
@@ -308,7 +308,7 @@ impl ServerState {
|
||||
host
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
let port = self.config.server.port;
|
||||
let api_key = self.config.server.api_key.clone();
|
||||
let api_key_for_state = api_key.clone(); // 用于保存到 running_api_key
|
||||
@@ -1051,15 +1051,13 @@ async fn run_server(
|
||||
let addr: std::net::SocketAddr = format!("{host}:{port}")
|
||||
.parse()
|
||||
.map_err(|e| format!("无效的监听地址 {}:{} - {}", host, port, e))?;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind(addr)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
format!(
|
||||
"无法绑定到 {}:{},错误: {}。请检查地址是否有效或端口是否被占用。",
|
||||
host, port, e
|
||||
)
|
||||
})?;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| {
|
||||
format!(
|
||||
"无法绑定到 {}:{},错误: {}。请检查地址是否有效或端口是否被占用。",
|
||||
host, port, e
|
||||
)
|
||||
})?;
|
||||
|
||||
tracing::info!("Server listening on {}", addr);
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"$schema": "https://schema.tauri.app/config/2",
|
||||
"productName": "ProxyCast",
|
||||
"version": "0.45.2",
|
||||
"version": "0.46.0",
|
||||
"identifier": "com.proxycast.app",
|
||||
"build": {
|
||||
"beforeDevCommand": "npm run dev",
|
||||
|
||||
+27
@@ -28,6 +28,7 @@ import {
|
||||
} from "./components/terminal";
|
||||
import { flowEventManager } from "./lib/flowEventManager";
|
||||
import { OnboardingWizard, useOnboardingState } from "./components/onboarding";
|
||||
import { STORAGE_KEYS } from "./components/onboarding/constants";
|
||||
import { ConnectConfirmDialog } from "./components/connect";
|
||||
import { showRegistryLoadError } from "./lib/utils/connectError";
|
||||
import { useDeepLink } from "./hooks/useDeepLink";
|
||||
@@ -36,6 +37,7 @@ import { ComponentDebugProvider } from "./contexts/ComponentDebugContext";
|
||||
import { SoundProvider } from "./contexts/SoundProvider";
|
||||
import { ComponentDebugOverlay } from "./components/dev";
|
||||
import { Page } from "./types/page";
|
||||
import { windowApi } from "./lib/api/window";
|
||||
|
||||
const AppContainer = styled.div`
|
||||
display: flex;
|
||||
@@ -103,6 +105,31 @@ function AppContent() {
|
||||
flowEventManager.subscribe();
|
||||
}, []);
|
||||
|
||||
// 应用启动时应用保存的窗口尺寸偏好
|
||||
useEffect(() => {
|
||||
const applyWindowSizePreference = async () => {
|
||||
const savedPreference = localStorage.getItem(
|
||||
STORAGE_KEYS.WINDOW_SIZE_PREFERENCE,
|
||||
);
|
||||
if (savedPreference) {
|
||||
try {
|
||||
if (savedPreference === "fullscreen") {
|
||||
const isCurrentlyFullscreen = await windowApi.isFullscreen();
|
||||
if (!isCurrentlyFullscreen) {
|
||||
await windowApi.toggleFullscreen();
|
||||
}
|
||||
} else {
|
||||
await windowApi.setWindowSizeByOption(savedPreference);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("应用窗口尺寸偏好失败:", error);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
applyWindowSizePreference();
|
||||
}, []);
|
||||
|
||||
// 处理 Registry 加载失败
|
||||
// _Requirements: 7.2, 7.3_
|
||||
useEffect(() => {
|
||||
|
||||
@@ -24,6 +24,7 @@ import {
|
||||
CollapsibleContent,
|
||||
CollapsibleTrigger,
|
||||
} from "@/components/ui/collapsible";
|
||||
import { getAgentBackend, setAgentBackend, type AgentBackend } from "../config";
|
||||
|
||||
// --- Styled Components ---
|
||||
|
||||
@@ -142,6 +143,14 @@ interface ChatSettingsProps {
|
||||
export const ChatSettings: React.FC<ChatSettingsProps> = ({ onClose }) => {
|
||||
// Local state for UI toggles (Mocking functional settings)
|
||||
const [fontSize, setFontSize] = useState([14]);
|
||||
const [agentBackend, setBackend] = useState<AgentBackend>(getAgentBackend());
|
||||
|
||||
const handleBackendChange = (value: AgentBackend) => {
|
||||
setBackend(value);
|
||||
setAgentBackend(value);
|
||||
// 提示用户需要刷新页面
|
||||
window.location.reload();
|
||||
};
|
||||
|
||||
return (
|
||||
<SettingsContainer>
|
||||
@@ -161,6 +170,30 @@ export const ChatSettings: React.FC<ChatSettingsProps> = ({ onClose }) => {
|
||||
</Header>
|
||||
|
||||
<ScrollArea className="flex-1">
|
||||
{/* Agent Backend Settings */}
|
||||
<CollapsibleSection title="Agent 后端">
|
||||
<SettingRow>
|
||||
<div>
|
||||
<div className="label">Agent 引擎</div>
|
||||
<div className="desc">切换后需刷新页面</div>
|
||||
</div>
|
||||
<Select
|
||||
value={agentBackend}
|
||||
onValueChange={(v) => handleBackendChange(v as AgentBackend)}
|
||||
>
|
||||
<SelectTrigger className="w-[100px] h-7 text-xs">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="native">Native</SelectItem>
|
||||
<SelectItem value="aster">Aster</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</SettingRow>
|
||||
</CollapsibleSection>
|
||||
|
||||
<Separator />
|
||||
|
||||
{/* Message Settings */}
|
||||
<CollapsibleSection title="消息设置">
|
||||
<SettingRow>
|
||||
|
||||
@@ -0,0 +1,346 @@
|
||||
/**
|
||||
* DecisionPanel - 权限确认面板
|
||||
*
|
||||
* 用于显示需要用户确认的操作,如:
|
||||
* - 工具调用确认
|
||||
* - 用户问题(AskUserQuestion)
|
||||
* - 权限请求
|
||||
*
|
||||
* 参考 Claude-Cowork 的设计
|
||||
*/
|
||||
|
||||
import { useState, useEffect } from "react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { cn } from "@/lib/utils";
|
||||
import {
|
||||
CheckCircle,
|
||||
XCircle,
|
||||
AlertTriangle,
|
||||
HelpCircle,
|
||||
Terminal,
|
||||
FileEdit,
|
||||
Globe,
|
||||
} from "lucide-react";
|
||||
import type { ActionRequired, ConfirmResponse } from "@/stores/agentStore";
|
||||
|
||||
/** 问题选项 */
|
||||
interface QuestionOption {
|
||||
label: string;
|
||||
description?: string;
|
||||
}
|
||||
|
||||
/** 问题数据 */
|
||||
interface Question {
|
||||
question: string;
|
||||
header?: string;
|
||||
options?: QuestionOption[];
|
||||
multiSelect?: boolean;
|
||||
}
|
||||
|
||||
interface DecisionPanelProps {
|
||||
request: ActionRequired;
|
||||
onSubmit: (response: ConfirmResponse) => void;
|
||||
}
|
||||
|
||||
/** 获取工具图标 */
|
||||
function getToolIcon(toolName?: string) {
|
||||
if (!toolName) return <HelpCircle className="h-4 w-4" />;
|
||||
|
||||
const name = toolName.toLowerCase();
|
||||
if (
|
||||
name.includes("bash") ||
|
||||
name.includes("terminal") ||
|
||||
name.includes("exec")
|
||||
) {
|
||||
return <Terminal className="h-4 w-4" />;
|
||||
}
|
||||
if (
|
||||
name.includes("write") ||
|
||||
name.includes("edit") ||
|
||||
name.includes("file")
|
||||
) {
|
||||
return <FileEdit className="h-4 w-4" />;
|
||||
}
|
||||
if (name.includes("web") || name.includes("fetch") || name.includes("http")) {
|
||||
return <Globe className="h-4 w-4" />;
|
||||
}
|
||||
return <AlertTriangle className="h-4 w-4" />;
|
||||
}
|
||||
|
||||
/** 格式化工具参数 */
|
||||
function formatArguments(args?: Record<string, unknown>): string {
|
||||
if (!args) return "";
|
||||
try {
|
||||
return JSON.stringify(args, null, 2);
|
||||
} catch {
|
||||
return String(args);
|
||||
}
|
||||
}
|
||||
|
||||
export function DecisionPanel({ request, onSubmit }: DecisionPanelProps) {
|
||||
// 解析问题数据(用于 ask_user 类型)
|
||||
const questions: Question[] = (request as any).questions ?? [];
|
||||
const [selectedOptions, setSelectedOptions] = useState<
|
||||
Record<number, string[]>
|
||||
>({});
|
||||
const [otherInputs, setOtherInputs] = useState<Record<number, string>>({});
|
||||
|
||||
// 重置状态当请求变化时
|
||||
useEffect(() => {
|
||||
setSelectedOptions({});
|
||||
setOtherInputs({});
|
||||
}, [request.requestId]);
|
||||
|
||||
// 切换选项
|
||||
const toggleOption = (
|
||||
qIndex: number,
|
||||
optionLabel: string,
|
||||
multiSelect?: boolean,
|
||||
) => {
|
||||
setSelectedOptions((prev) => {
|
||||
const current = prev[qIndex] ?? [];
|
||||
if (multiSelect) {
|
||||
const next = current.includes(optionLabel)
|
||||
? current.filter((label) => label !== optionLabel)
|
||||
: [...current, optionLabel];
|
||||
return { ...prev, [qIndex]: next };
|
||||
}
|
||||
return { ...prev, [qIndex]: [optionLabel] };
|
||||
});
|
||||
};
|
||||
|
||||
// 构建答案
|
||||
const buildAnswers = () => {
|
||||
const answers: Record<string, string> = {};
|
||||
questions.forEach((q, qIndex) => {
|
||||
const selected = selectedOptions[qIndex] ?? [];
|
||||
const otherText = otherInputs[qIndex]?.trim() ?? "";
|
||||
let value = "";
|
||||
if (q.multiSelect) {
|
||||
const combined = [...selected];
|
||||
if (otherText) combined.push(otherText);
|
||||
value = combined.join(", ");
|
||||
} else {
|
||||
value = otherText || selected[0] || "";
|
||||
}
|
||||
if (value) answers[q.question] = value;
|
||||
});
|
||||
return answers;
|
||||
};
|
||||
|
||||
// 检查是否���以提交
|
||||
const canSubmit =
|
||||
questions.length === 0 ||
|
||||
questions.every((_, qIndex) => {
|
||||
const selected = selectedOptions[qIndex] ?? [];
|
||||
const otherText = otherInputs[qIndex]?.trim() ?? "";
|
||||
return selected.length > 0 || otherText.length > 0;
|
||||
});
|
||||
|
||||
// 处理允许
|
||||
const handleAllow = () => {
|
||||
const response =
|
||||
questions.length > 0 ? JSON.stringify(buildAnswers()) : undefined;
|
||||
onSubmit({
|
||||
requestId: request.requestId,
|
||||
confirmed: true,
|
||||
response,
|
||||
});
|
||||
};
|
||||
|
||||
// 处理拒绝
|
||||
const handleDeny = () => {
|
||||
onSubmit({
|
||||
requestId: request.requestId,
|
||||
confirmed: false,
|
||||
response: "用户拒绝了请求",
|
||||
});
|
||||
};
|
||||
|
||||
// 渲染用户问题面板
|
||||
if (request.actionType === "ask_user" && questions.length > 0) {
|
||||
return (
|
||||
<Card className="border-blue-200 bg-blue-50/50 dark:border-blue-800 dark:bg-blue-950/20">
|
||||
<CardHeader className="pb-2">
|
||||
<CardTitle className="flex items-center gap-2 text-sm font-medium text-blue-700 dark:text-blue-300">
|
||||
<HelpCircle className="h-4 w-4" />
|
||||
Claude 的问题
|
||||
</CardTitle>
|
||||
</CardHeader>
|
||||
<CardContent className="space-y-4">
|
||||
{questions.map((q, qIndex) => (
|
||||
<div key={qIndex} className="space-y-3">
|
||||
<p className="text-sm text-foreground">{q.question}</p>
|
||||
|
||||
{q.header && (
|
||||
<Badge variant="secondary" className="text-xs">
|
||||
{q.header}
|
||||
</Badge>
|
||||
)}
|
||||
|
||||
{/* 选项列表 */}
|
||||
{q.options && q.options.length > 0 && (
|
||||
<div className="grid gap-2">
|
||||
{q.options.map((option, optIndex) => {
|
||||
const isSelected = (selectedOptions[qIndex] ?? []).includes(
|
||||
option.label,
|
||||
);
|
||||
const shouldAutoSubmit =
|
||||
questions.length === 1 && !q.multiSelect;
|
||||
|
||||
return (
|
||||
<button
|
||||
key={optIndex}
|
||||
className={cn(
|
||||
"rounded-lg border px-4 py-3 text-left text-sm transition-colors",
|
||||
isSelected
|
||||
? "border-blue-500 bg-blue-100 dark:border-blue-400 dark:bg-blue-900/30"
|
||||
: "border-border bg-background hover:border-blue-300 hover:bg-muted",
|
||||
)}
|
||||
onClick={() => {
|
||||
if (shouldAutoSubmit) {
|
||||
onSubmit({
|
||||
requestId: request.requestId,
|
||||
confirmed: true,
|
||||
response: option.label,
|
||||
});
|
||||
return;
|
||||
}
|
||||
toggleOption(qIndex, option.label, q.multiSelect);
|
||||
}}
|
||||
>
|
||||
<div className="font-medium">{option.label}</div>
|
||||
{option.description && (
|
||||
<div className="mt-1 text-xs text-muted-foreground">
|
||||
{option.description}
|
||||
</div>
|
||||
)}
|
||||
</button>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 其他输入 */}
|
||||
<div className="space-y-1">
|
||||
<label className="text-xs font-medium text-muted-foreground">
|
||||
其他
|
||||
</label>
|
||||
<Input
|
||||
placeholder="输入你的答案..."
|
||||
value={otherInputs[qIndex] ?? ""}
|
||||
onChange={(e) =>
|
||||
setOtherInputs((prev) => ({
|
||||
...prev,
|
||||
[qIndex]: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{q.multiSelect && (
|
||||
<p className="text-xs text-muted-foreground">
|
||||
可以选择多个选项
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
))}
|
||||
|
||||
{/* 操作按钮 */}
|
||||
<div className="flex gap-2 pt-2">
|
||||
<Button
|
||||
size="sm"
|
||||
onClick={handleAllow}
|
||||
disabled={!canSubmit}
|
||||
className="bg-blue-600 hover:bg-blue-700"
|
||||
>
|
||||
<CheckCircle className="mr-1 h-4 w-4" />
|
||||
提交答案
|
||||
</Button>
|
||||
<Button size="sm" variant="outline" onClick={handleDeny}>
|
||||
<XCircle className="mr-1 h-4 w-4" />
|
||||
取消
|
||||
</Button>
|
||||
</div>
|
||||
</CardContent>
|
||||
</Card>
|
||||
);
|
||||
}
|
||||
|
||||
// 渲染工具确认面板
|
||||
return (
|
||||
<Card className="border-amber-200 bg-amber-50/50 dark:border-amber-800 dark:bg-amber-950/20">
|
||||
<CardHeader className="pb-2">
|
||||
<CardTitle className="flex items-center gap-2 text-sm font-medium text-amber-700 dark:text-amber-300">
|
||||
<AlertTriangle className="h-4 w-4" />
|
||||
权限请求
|
||||
</CardTitle>
|
||||
</CardHeader>
|
||||
<CardContent className="space-y-3">
|
||||
{/* 工具信息 */}
|
||||
<div className="flex items-center gap-2">
|
||||
{getToolIcon(request.toolName)}
|
||||
<span className="text-sm">
|
||||
Claude 想要使用:
|
||||
<span className="ml-1 font-medium">
|
||||
{request.toolName || "未知工具"}
|
||||
</span>
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{/* 参数预览 */}
|
||||
{request.arguments && (
|
||||
<div className="rounded-lg bg-muted/50 p-3">
|
||||
<pre className="max-h-40 overflow-auto whitespace-pre-wrap break-words font-mono text-xs text-muted-foreground">
|
||||
{formatArguments(request.arguments)}
|
||||
</pre>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 操作按钮 */}
|
||||
<div className="flex gap-2 pt-2">
|
||||
<Button
|
||||
size="sm"
|
||||
onClick={handleAllow}
|
||||
className="bg-green-600 hover:bg-green-700"
|
||||
>
|
||||
<CheckCircle className="mr-1 h-4 w-4" />
|
||||
允许
|
||||
</Button>
|
||||
<Button size="sm" variant="outline" onClick={handleDeny}>
|
||||
<XCircle className="mr-1 h-4 w-4" />
|
||||
拒绝
|
||||
</Button>
|
||||
</div>
|
||||
</CardContent>
|
||||
</Card>
|
||||
);
|
||||
}
|
||||
|
||||
/** 权限确认列表组件 */
|
||||
export function DecisionPanelList({
|
||||
requests,
|
||||
onSubmit,
|
||||
}: {
|
||||
requests: ActionRequired[];
|
||||
onSubmit: (response: ConfirmResponse) => void;
|
||||
}) {
|
||||
if (requests.length === 0) return null;
|
||||
|
||||
return (
|
||||
<div className="space-y-3">
|
||||
{requests.map((request) => (
|
||||
<DecisionPanel
|
||||
key={request.requestId}
|
||||
request={request}
|
||||
onSubmit={onSubmit}
|
||||
/>
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
export default DecisionPanel;
|
||||
@@ -0,0 +1,35 @@
|
||||
/**
|
||||
* Agent 后端配置
|
||||
*
|
||||
* 用于切换 Native 和 Aster 后端
|
||||
*/
|
||||
|
||||
export type AgentBackend = "native" | "aster";
|
||||
|
||||
// 默认使用 Native 后端,可以通过 localStorage 切换
|
||||
const STORAGE_KEY = "proxycast_agent_backend";
|
||||
|
||||
/**
|
||||
* 获取当前 Agent 后端
|
||||
*/
|
||||
export function getAgentBackend(): AgentBackend {
|
||||
const stored = localStorage.getItem(STORAGE_KEY);
|
||||
if (stored === "aster" || stored === "native") {
|
||||
return stored;
|
||||
}
|
||||
return "native"; // 默认
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置 Agent 后端
|
||||
*/
|
||||
export function setAgentBackend(backend: AgentBackend): void {
|
||||
localStorage.setItem(STORAGE_KEY, backend);
|
||||
}
|
||||
|
||||
/**
|
||||
* 是否使用 Aster 后端
|
||||
*/
|
||||
export function useAsterBackend(): boolean {
|
||||
return getAgentBackend() === "aster";
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
/**
|
||||
* Agent Chat Hook 统一导出
|
||||
*
|
||||
* 根据配置自动选择 Native 或 Aster 后端
|
||||
*/
|
||||
|
||||
import { useAgentChat } from "./useAgentChat";
|
||||
import { useAsterAgentChat } from "./useAsterAgentChat";
|
||||
import { getAgentBackend } from "../config";
|
||||
|
||||
export type { Topic } from "./useAgentChat";
|
||||
|
||||
/** Hook 配置选项 */
|
||||
interface UseAgentChatUnifiedOptions {
|
||||
systemPrompt?: string;
|
||||
onWriteFile?: (content: string, fileName: string) => void;
|
||||
}
|
||||
|
||||
/**
|
||||
* 统一的 Agent Chat Hook
|
||||
*
|
||||
* 根据 localStorage 配置自动选择后端:
|
||||
* - "native": 使用原有的 Native Agent 后端
|
||||
* - "aster": 使用新的 Aster Agent 后端
|
||||
*
|
||||
* 切换方式:
|
||||
* localStorage.setItem("proxycast_agent_backend", "aster")
|
||||
*/
|
||||
export function useAgentChatUnified(options: UseAgentChatUnifiedOptions = {}) {
|
||||
const backend = getAgentBackend();
|
||||
|
||||
// 根据配置选择 hook
|
||||
// 注意:React hooks 规则要求 hooks 调用顺序一致
|
||||
// 这里我们总是调用两个 hook,但只使用其中一个的结果
|
||||
const nativeResult = useAgentChat(options);
|
||||
const asterResult = useAsterAgentChat(options);
|
||||
|
||||
if (backend === "aster") {
|
||||
console.log("[AgentChat] 使用 Aster 后端");
|
||||
return asterResult;
|
||||
}
|
||||
|
||||
console.log("[AgentChat] 使用 Native 后端");
|
||||
return nativeResult;
|
||||
}
|
||||
|
||||
// 重新导出原有 hooks,便于直接使用
|
||||
export { useAgentChat } from "./useAgentChat";
|
||||
export { useAsterAgentChat } from "./useAsterAgentChat";
|
||||
@@ -0,0 +1,649 @@
|
||||
/**
|
||||
* Aster Agent Chat Hook
|
||||
*
|
||||
* 基于 Aster 框架的聊天 hook
|
||||
* 接口与 useAgentChat 保持一致,便于切换
|
||||
*/
|
||||
|
||||
import { useState, useEffect, useRef, useCallback } from "react";
|
||||
import { toast } from "sonner";
|
||||
import { safeListen } from "@/lib/dev-bridge";
|
||||
import type { UnlistenFn } from "@tauri-apps/api/event";
|
||||
import {
|
||||
initAsterAgent,
|
||||
sendAsterMessageStream,
|
||||
createAsterSession,
|
||||
listAsterSessions,
|
||||
getAsterSession,
|
||||
stopAsterSession,
|
||||
confirmAsterAction,
|
||||
parseStreamEvent,
|
||||
type StreamEvent,
|
||||
type AsterSessionInfo,
|
||||
} from "@/lib/api/agent";
|
||||
import { Message, MessageImage, ContentPart } from "../types";
|
||||
|
||||
/** 话题信息 */
|
||||
export interface Topic {
|
||||
id: string;
|
||||
title: string;
|
||||
createdAt: Date;
|
||||
messagesCount: number;
|
||||
}
|
||||
|
||||
/** 权限确认请求 */
|
||||
export interface ActionRequired {
|
||||
requestId: string;
|
||||
actionType: string;
|
||||
toolName?: string;
|
||||
arguments?: Record<string, unknown>;
|
||||
question?: string;
|
||||
timestamp: Date;
|
||||
}
|
||||
|
||||
/** 确认响应 */
|
||||
export interface ConfirmResponse {
|
||||
requestId: string;
|
||||
confirmed: boolean;
|
||||
response?: string;
|
||||
}
|
||||
|
||||
/** Hook 配置选项 */
|
||||
interface UseAsterAgentChatOptions {
|
||||
systemPrompt?: string;
|
||||
onWriteFile?: (content: string, fileName: string) => void;
|
||||
}
|
||||
|
||||
// 音效相关(复用)
|
||||
let toolcallAudio: HTMLAudioElement | null = null;
|
||||
let typewriterAudio: HTMLAudioElement | null = null;
|
||||
let lastTypewriterTime = 0;
|
||||
const TYPEWRITER_INTERVAL = 120;
|
||||
|
||||
const initAudio = () => {
|
||||
if (!toolcallAudio) {
|
||||
toolcallAudio = new Audio("/sounds/tool-call.mp3");
|
||||
toolcallAudio.volume = 1;
|
||||
toolcallAudio.load();
|
||||
}
|
||||
if (!typewriterAudio) {
|
||||
typewriterAudio = new Audio("/sounds/typing.mp3");
|
||||
typewriterAudio.volume = 0.6;
|
||||
typewriterAudio.load();
|
||||
}
|
||||
};
|
||||
|
||||
const getSoundEnabled = (): boolean => {
|
||||
return localStorage.getItem("proxycast_sound_enabled") === "true";
|
||||
};
|
||||
|
||||
const playToolcallSound = () => {
|
||||
if (!getSoundEnabled()) return;
|
||||
initAudio();
|
||||
if (toolcallAudio) {
|
||||
toolcallAudio.currentTime = 0;
|
||||
toolcallAudio.play().catch(console.error);
|
||||
}
|
||||
};
|
||||
|
||||
const playTypewriterSound = () => {
|
||||
if (!getSoundEnabled()) return;
|
||||
const now = Date.now();
|
||||
if (now - lastTypewriterTime < TYPEWRITER_INTERVAL) return;
|
||||
initAudio();
|
||||
if (typewriterAudio) {
|
||||
typewriterAudio.currentTime = 0;
|
||||
typewriterAudio.play().catch(console.error);
|
||||
lastTypewriterTime = now;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* 将前端 Provider 类型映射到 Aster Provider 名称
|
||||
*/
|
||||
const mapProviderName = (providerType: string): string => {
|
||||
const mapping: Record<string, string> = {
|
||||
// OpenAI 兼容
|
||||
openai: "openai",
|
||||
"gpt-4": "openai",
|
||||
"gpt-4o": "openai",
|
||||
// Anthropic
|
||||
claude: "anthropic",
|
||||
anthropic: "anthropic",
|
||||
// Google
|
||||
google: "google",
|
||||
gemini: "google",
|
||||
// DeepSeek
|
||||
deepseek: "custom_deepseek",
|
||||
"deepseek-reasoner": "custom_deepseek",
|
||||
// Ollama
|
||||
ollama: "ollama",
|
||||
// OpenRouter
|
||||
openrouter: "openrouter",
|
||||
// 其他
|
||||
groq: "groq",
|
||||
mistral: "mistral",
|
||||
};
|
||||
return mapping[providerType.toLowerCase()] || providerType;
|
||||
};
|
||||
|
||||
export function useAsterAgentChat(options: UseAsterAgentChatOptions = {}) {
|
||||
const { onWriteFile } = options;
|
||||
|
||||
// 状态
|
||||
const [isInitialized, setIsInitialized] = useState(false);
|
||||
const [sessionId, setSessionId] = useState<string | null>(null);
|
||||
const [messages, setMessages] = useState<Message[]>([]);
|
||||
const [topics, setTopics] = useState<Topic[]>([]);
|
||||
const [isSending, setIsSending] = useState(false);
|
||||
const [pendingActions, setPendingActions] = useState<ActionRequired[]>([]);
|
||||
|
||||
// Provider/Model(本地状态)
|
||||
const [providerType, setProviderType] = useState(
|
||||
() => localStorage.getItem("agent_pref_provider") || "claude",
|
||||
);
|
||||
const [model, setModel] = useState(
|
||||
() => localStorage.getItem("agent_pref_model") || "claude-sonnet-4-5",
|
||||
);
|
||||
|
||||
// Refs
|
||||
const unlistenRef = useRef<UnlistenFn | null>(null);
|
||||
const currentAssistantMsgIdRef = useRef<string | null>(null);
|
||||
|
||||
// 持久化 provider/model
|
||||
useEffect(() => {
|
||||
localStorage.setItem("agent_pref_provider", providerType);
|
||||
}, [providerType]);
|
||||
|
||||
useEffect(() => {
|
||||
localStorage.setItem("agent_pref_model", model);
|
||||
}, [model]);
|
||||
|
||||
// 初始化 Aster Agent
|
||||
useEffect(() => {
|
||||
const init = async () => {
|
||||
try {
|
||||
await initAsterAgent();
|
||||
setIsInitialized(true);
|
||||
console.log("[AsterChat] Agent 初始化成功");
|
||||
// 初始化后加载话题列表
|
||||
const sessions = await listAsterSessions();
|
||||
const topicList: Topic[] = sessions.map((s: AsterSessionInfo) => ({
|
||||
id: s.id,
|
||||
title:
|
||||
s.name ||
|
||||
`话题 ${new Date(s.created_at).toLocaleDateString("zh-CN")}`,
|
||||
createdAt: new Date(s.created_at),
|
||||
messagesCount: s.messages_count,
|
||||
}));
|
||||
setTopics(topicList);
|
||||
} catch (err) {
|
||||
console.error("[AsterChat] 初始化失败:", err);
|
||||
}
|
||||
};
|
||||
init();
|
||||
}, []);
|
||||
|
||||
// 加载话题列表
|
||||
const loadTopics = useCallback(async () => {
|
||||
try {
|
||||
const sessions = await listAsterSessions();
|
||||
const topicList: Topic[] = sessions.map((s: AsterSessionInfo) => ({
|
||||
id: s.id,
|
||||
title:
|
||||
s.name ||
|
||||
`话题 ${new Date(s.created_at).toLocaleDateString("zh-CN")}`,
|
||||
createdAt: new Date(s.created_at),
|
||||
messagesCount: s.messages_count,
|
||||
}));
|
||||
setTopics(topicList);
|
||||
} catch (error) {
|
||||
console.error("[AsterChat] 加载话题失败:", error);
|
||||
}
|
||||
}, []);
|
||||
|
||||
// 确保有会话
|
||||
const ensureSession = useCallback(async (): Promise<string | null> => {
|
||||
if (sessionId) return sessionId;
|
||||
|
||||
try {
|
||||
const newSessionId = await createAsterSession();
|
||||
setSessionId(newSessionId);
|
||||
return newSessionId;
|
||||
} catch (error) {
|
||||
console.error("[AsterChat] 创建会话失败:", error);
|
||||
toast.error("创建会话失败");
|
||||
return null;
|
||||
}
|
||||
}, [sessionId]);
|
||||
|
||||
// 辅助函数:追加文本到 contentParts
|
||||
const appendTextToParts = (
|
||||
parts: ContentPart[],
|
||||
text: string,
|
||||
): ContentPart[] => {
|
||||
const newParts = [...parts];
|
||||
const lastPart = newParts[newParts.length - 1];
|
||||
|
||||
if (lastPart && lastPart.type === "text") {
|
||||
newParts[newParts.length - 1] = {
|
||||
type: "text",
|
||||
text: lastPart.text + text,
|
||||
};
|
||||
} else {
|
||||
newParts.push({ type: "text", text });
|
||||
}
|
||||
return newParts;
|
||||
};
|
||||
|
||||
// 发送消息
|
||||
const sendMessage = useCallback(
|
||||
async (
|
||||
content: string,
|
||||
images: MessageImage[],
|
||||
_webSearch?: boolean,
|
||||
_thinking?: boolean,
|
||||
) => {
|
||||
// 用户消息
|
||||
const userMsg: Message = {
|
||||
id: crypto.randomUUID(),
|
||||
role: "user",
|
||||
content,
|
||||
images: images.length > 0 ? images : undefined,
|
||||
timestamp: new Date(),
|
||||
};
|
||||
|
||||
// 助手消息占位符
|
||||
const assistantMsgId = crypto.randomUUID();
|
||||
const assistantMsg: Message = {
|
||||
id: assistantMsgId,
|
||||
role: "assistant",
|
||||
content: "",
|
||||
timestamp: new Date(),
|
||||
isThinking: true,
|
||||
thinkingContent: "思考中...",
|
||||
contentParts: [],
|
||||
};
|
||||
|
||||
setMessages((prev) => [...prev, userMsg, assistantMsg]);
|
||||
setIsSending(true);
|
||||
currentAssistantMsgIdRef.current = assistantMsgId;
|
||||
|
||||
let accumulatedContent = "";
|
||||
let unlisten: UnlistenFn | null = null;
|
||||
|
||||
try {
|
||||
const activeSessionId = await ensureSession();
|
||||
if (!activeSessionId) throw new Error("无法创建会话");
|
||||
|
||||
const eventName = `aster_stream_${assistantMsgId}`;
|
||||
|
||||
// 设置事件监听
|
||||
unlisten = await safeListen<StreamEvent>(eventName, (event) => {
|
||||
const data = parseStreamEvent(event.payload);
|
||||
if (!data) return;
|
||||
|
||||
switch (data.type) {
|
||||
case "text_delta":
|
||||
accumulatedContent += data.text;
|
||||
playTypewriterSound();
|
||||
setMessages((prev) =>
|
||||
prev.map((msg) =>
|
||||
msg.id === assistantMsgId
|
||||
? {
|
||||
...msg,
|
||||
content: accumulatedContent,
|
||||
thinkingContent: undefined,
|
||||
contentParts: appendTextToParts(
|
||||
msg.contentParts || [],
|
||||
data.text,
|
||||
),
|
||||
}
|
||||
: msg,
|
||||
),
|
||||
);
|
||||
break;
|
||||
|
||||
case "tool_start": {
|
||||
playToolcallSound();
|
||||
const newToolCall = {
|
||||
id: data.tool_id,
|
||||
name: data.tool_name,
|
||||
arguments: data.arguments,
|
||||
status: "running" as const,
|
||||
startTime: new Date(),
|
||||
};
|
||||
|
||||
// 检查是否是写入文件工具
|
||||
const toolName = data.tool_name.toLowerCase();
|
||||
if (toolName.includes("write") || toolName.includes("create")) {
|
||||
try {
|
||||
const args = JSON.parse(data.arguments || "{}");
|
||||
const filePath = args.path || args.file_path || args.filePath;
|
||||
const fileContent = args.content || args.text || "";
|
||||
if (filePath && fileContent && onWriteFile) {
|
||||
onWriteFile(fileContent, filePath);
|
||||
}
|
||||
} catch (e) {
|
||||
console.warn("[AsterChat] 解析工具参数失败:", e);
|
||||
}
|
||||
}
|
||||
|
||||
setMessages((prev) =>
|
||||
prev.map((msg) => {
|
||||
if (msg.id !== assistantMsgId) return msg;
|
||||
if (msg.toolCalls?.find((tc) => tc.id === data.tool_id))
|
||||
return msg;
|
||||
return {
|
||||
...msg,
|
||||
toolCalls: [...(msg.toolCalls || []), newToolCall],
|
||||
contentParts: [
|
||||
...(msg.contentParts || []),
|
||||
{ type: "tool_use" as const, toolCall: newToolCall },
|
||||
],
|
||||
};
|
||||
}),
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
case "tool_end":
|
||||
setMessages((prev) =>
|
||||
prev.map((msg) => {
|
||||
if (msg.id !== assistantMsgId) return msg;
|
||||
const updatedToolCalls = (msg.toolCalls || []).map((tc) =>
|
||||
tc.id === data.tool_id
|
||||
? {
|
||||
...tc,
|
||||
status: data.result.success
|
||||
? ("completed" as const)
|
||||
: ("failed" as const),
|
||||
result: data.result,
|
||||
endTime: new Date(),
|
||||
}
|
||||
: tc,
|
||||
);
|
||||
const updatedContentParts = (msg.contentParts || []).map(
|
||||
(part) => {
|
||||
if (
|
||||
part.type === "tool_use" &&
|
||||
part.toolCall.id === data.tool_id
|
||||
) {
|
||||
return {
|
||||
...part,
|
||||
toolCall: {
|
||||
...part.toolCall,
|
||||
status: data.result.success
|
||||
? ("completed" as const)
|
||||
: ("failed" as const),
|
||||
result: data.result,
|
||||
endTime: new Date(),
|
||||
},
|
||||
};
|
||||
}
|
||||
return part;
|
||||
},
|
||||
);
|
||||
return {
|
||||
...msg,
|
||||
toolCalls: updatedToolCalls,
|
||||
contentParts: updatedContentParts,
|
||||
};
|
||||
}),
|
||||
);
|
||||
break;
|
||||
|
||||
case "final_done":
|
||||
setMessages((prev) =>
|
||||
prev.map((msg) =>
|
||||
msg.id === assistantMsgId
|
||||
? {
|
||||
...msg,
|
||||
isThinking: false,
|
||||
content: accumulatedContent || "(无响应)",
|
||||
}
|
||||
: msg,
|
||||
),
|
||||
);
|
||||
setIsSending(false);
|
||||
unlistenRef.current = null;
|
||||
currentAssistantMsgIdRef.current = null;
|
||||
if (unlisten) {
|
||||
unlisten();
|
||||
unlisten = null;
|
||||
}
|
||||
break;
|
||||
|
||||
case "error":
|
||||
toast.error(`响应错误: ${data.message}`);
|
||||
setMessages((prev) =>
|
||||
prev.map((msg) =>
|
||||
msg.id === assistantMsgId
|
||||
? {
|
||||
...msg,
|
||||
isThinking: false,
|
||||
content: accumulatedContent || `错误: ${data.message}`,
|
||||
}
|
||||
: msg,
|
||||
),
|
||||
);
|
||||
setIsSending(false);
|
||||
if (unlisten) {
|
||||
unlisten();
|
||||
unlisten = null;
|
||||
}
|
||||
break;
|
||||
|
||||
// 处理权限确认请求
|
||||
default: {
|
||||
// 检查是否是 action_required 事件(通过原始 payload)
|
||||
const rawEvent = event.payload as unknown as Record<
|
||||
string,
|
||||
unknown
|
||||
>;
|
||||
if (rawEvent.type === "action_required") {
|
||||
const actionData: ActionRequired = {
|
||||
requestId: rawEvent.request_id as string,
|
||||
actionType: rawEvent.action_type as string,
|
||||
toolName: (rawEvent.data as Record<string, unknown>)
|
||||
?.tool_name as string,
|
||||
arguments: (rawEvent.data as Record<string, unknown>)
|
||||
?.arguments as Record<string, unknown>,
|
||||
question: (rawEvent.data as Record<string, unknown>)
|
||||
?.question as string,
|
||||
timestamp: new Date(),
|
||||
};
|
||||
setPendingActions((prev) => [...prev, actionData]);
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
unlistenRef.current = unlisten;
|
||||
|
||||
// 发送请求
|
||||
const imagesToSend =
|
||||
images.length > 0
|
||||
? images.map((img) => ({
|
||||
data: img.data,
|
||||
media_type: img.mediaType,
|
||||
}))
|
||||
: undefined;
|
||||
|
||||
// 构建 Provider 配置
|
||||
const providerConfig = {
|
||||
provider_name: mapProviderName(providerType),
|
||||
model_name: model,
|
||||
};
|
||||
|
||||
await sendAsterMessageStream(
|
||||
content,
|
||||
activeSessionId,
|
||||
eventName,
|
||||
imagesToSend,
|
||||
providerConfig,
|
||||
);
|
||||
} catch (error) {
|
||||
console.error("[AsterChat] 发送失败:", error);
|
||||
toast.error(`发送失败: ${error}`);
|
||||
setMessages((prev) => prev.filter((msg) => msg.id !== assistantMsgId));
|
||||
setIsSending(false);
|
||||
if (unlisten) unlisten();
|
||||
}
|
||||
},
|
||||
[ensureSession, onWriteFile, providerType, model],
|
||||
);
|
||||
|
||||
// 停止发送
|
||||
const stopSending = useCallback(async () => {
|
||||
if (unlistenRef.current) {
|
||||
unlistenRef.current();
|
||||
unlistenRef.current = null;
|
||||
}
|
||||
|
||||
if (sessionId) {
|
||||
try {
|
||||
await stopAsterSession(sessionId);
|
||||
} catch (e) {
|
||||
console.error("[AsterChat] 停止失败:", e);
|
||||
}
|
||||
}
|
||||
|
||||
if (currentAssistantMsgIdRef.current) {
|
||||
setMessages((prev) =>
|
||||
prev.map((msg) =>
|
||||
msg.id === currentAssistantMsgIdRef.current
|
||||
? { ...msg, isThinking: false, content: msg.content || "(已停止)" }
|
||||
: msg,
|
||||
),
|
||||
);
|
||||
currentAssistantMsgIdRef.current = null;
|
||||
}
|
||||
|
||||
setIsSending(false);
|
||||
toast.info("已停止生成");
|
||||
}, [sessionId]);
|
||||
|
||||
// 确认权限请求
|
||||
const confirmAction = useCallback(async (response: ConfirmResponse) => {
|
||||
try {
|
||||
await confirmAsterAction(
|
||||
response.requestId,
|
||||
response.confirmed,
|
||||
response.response,
|
||||
);
|
||||
// 移除已处理的请求
|
||||
setPendingActions((prev) =>
|
||||
prev.filter((a) => a.requestId !== response.requestId),
|
||||
);
|
||||
} catch (error) {
|
||||
console.error("[AsterChat] 确认失败:", error);
|
||||
toast.error("确认操作失败");
|
||||
}
|
||||
}, []);
|
||||
|
||||
// 清空消息
|
||||
const clearMessages = useCallback(() => {
|
||||
setMessages([]);
|
||||
setSessionId(null);
|
||||
toast.success("新话题已创建");
|
||||
}, []);
|
||||
|
||||
// 删除消息
|
||||
const deleteMessage = useCallback((id: string) => {
|
||||
setMessages((prev) => prev.filter((msg) => msg.id !== id));
|
||||
}, []);
|
||||
|
||||
// 编辑消息
|
||||
const editMessage = useCallback((id: string, newContent: string) => {
|
||||
setMessages((prev) =>
|
||||
prev.map((msg) =>
|
||||
msg.id === id ? { ...msg, content: newContent } : msg,
|
||||
),
|
||||
);
|
||||
}, []);
|
||||
|
||||
// 切换话题
|
||||
const switchTopic = useCallback(
|
||||
async (topicId: string) => {
|
||||
if (topicId === sessionId) return;
|
||||
|
||||
try {
|
||||
const detail = await getAsterSession(topicId);
|
||||
const loadedMessages: Message[] = detail.messages.map((msg, index) => ({
|
||||
id: `${topicId}-${index}`,
|
||||
role: msg.role as "user" | "assistant",
|
||||
content: msg.content,
|
||||
timestamp: new Date(msg.timestamp),
|
||||
isThinking: false,
|
||||
}));
|
||||
|
||||
setMessages(loadedMessages);
|
||||
setSessionId(topicId);
|
||||
toast.info("已切换话题");
|
||||
} catch (error) {
|
||||
console.error("[AsterChat] 切换话题失败:", error);
|
||||
setMessages([]);
|
||||
setSessionId(topicId);
|
||||
toast.error("加载对话历史失败");
|
||||
}
|
||||
},
|
||||
[sessionId],
|
||||
);
|
||||
|
||||
// 删除话题
|
||||
const deleteTopic = useCallback(
|
||||
async (topicId: string) => {
|
||||
// TODO: 实现后端删除
|
||||
setTopics((prev) => prev.filter((t) => t.id !== topicId));
|
||||
if (topicId === sessionId) {
|
||||
setSessionId(null);
|
||||
setMessages([]);
|
||||
}
|
||||
toast.success("话题已删除");
|
||||
},
|
||||
[sessionId],
|
||||
);
|
||||
|
||||
// 兼容接口
|
||||
const handleStartProcess = useCallback(async () => {
|
||||
// Aster 不需要单独启动进程
|
||||
}, []);
|
||||
|
||||
const handleStopProcess = useCallback(async () => {
|
||||
setSessionId(null);
|
||||
}, []);
|
||||
|
||||
return {
|
||||
// 兼容 useAgentChat 接口
|
||||
processStatus: { running: isInitialized },
|
||||
handleStartProcess,
|
||||
handleStopProcess,
|
||||
|
||||
providerType,
|
||||
setProviderType,
|
||||
model,
|
||||
setModel,
|
||||
providerConfig: {}, // 简化版本
|
||||
isConfigLoading: false,
|
||||
|
||||
messages,
|
||||
isSending,
|
||||
sendMessage,
|
||||
stopSending,
|
||||
clearMessages,
|
||||
deleteMessage,
|
||||
editMessage,
|
||||
|
||||
topics,
|
||||
sessionId,
|
||||
switchTopic,
|
||||
deleteTopic,
|
||||
loadTopics,
|
||||
|
||||
// Aster 特有功能
|
||||
pendingActions,
|
||||
confirmAction,
|
||||
};
|
||||
}
|
||||
@@ -8,13 +8,20 @@ import { Button } from "@/components/ui/button";
|
||||
import { ChevronLeft, ChevronRight } from "lucide-react";
|
||||
import { WelcomeStep } from "./steps/WelcomeStep";
|
||||
import { UserProfileStep } from "./steps/UserProfileStep";
|
||||
import { WindowSizeStep } from "./steps/WindowSizeStep";
|
||||
import { PluginSelectStep } from "./steps/PluginSelectStep";
|
||||
import {
|
||||
InstallProgressStep,
|
||||
type PluginInstallState,
|
||||
} from "./steps/InstallProgressStep";
|
||||
import { CompleteStep } from "./steps/CompleteStep";
|
||||
import { userProfiles, type UserProfile } from "./constants";
|
||||
import {
|
||||
userProfiles,
|
||||
type UserProfile,
|
||||
type WindowSizePreference,
|
||||
STORAGE_KEYS,
|
||||
} from "./constants";
|
||||
import { windowApi } from "@/lib/api/window";
|
||||
|
||||
const Overlay = styled.div`
|
||||
position: fixed;
|
||||
@@ -83,7 +90,7 @@ const FooterRight = styled.div`
|
||||
gap: 12px;
|
||||
`;
|
||||
|
||||
const TOTAL_STEPS = 5;
|
||||
const TOTAL_STEPS = 6;
|
||||
|
||||
interface OnboardingWizardProps {
|
||||
onComplete: () => void;
|
||||
@@ -92,6 +99,8 @@ interface OnboardingWizardProps {
|
||||
export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
const [currentStep, setCurrentStep] = useState(1);
|
||||
const [userProfile, setUserProfile] = useState<UserProfile | null>(null);
|
||||
const [windowSizePreference, setWindowSizePreference] =
|
||||
useState<WindowSizePreference | null>("default");
|
||||
const [selectedPlugins, setSelectedPlugins] = useState<string[]>([]);
|
||||
const [installResults, setInstallResults] = useState<PluginInstallState[]>(
|
||||
[],
|
||||
@@ -110,8 +119,8 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
const handleNext = useCallback(() => {
|
||||
if (currentStep < TOTAL_STEPS) {
|
||||
// 如果没有选择插件,跳过安装步骤
|
||||
if (currentStep === 3 && selectedPlugins.length === 0) {
|
||||
setCurrentStep(5); // 直接跳到完成页
|
||||
if (currentStep === 4 && selectedPlugins.length === 0) {
|
||||
setCurrentStep(6); // 直接跳到完成页
|
||||
} else {
|
||||
setCurrentStep((prev) => prev + 1);
|
||||
}
|
||||
@@ -121,9 +130,9 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
const handleBack = useCallback(() => {
|
||||
if (currentStep > 1) {
|
||||
// 如果从完成页返回且没有安装结果,返回到插件选择
|
||||
if (currentStep === 5 && installResults.length === 0) {
|
||||
setCurrentStep(3);
|
||||
} else if (currentStep === 5) {
|
||||
if (currentStep === 6 && installResults.length === 0) {
|
||||
setCurrentStep(4);
|
||||
} else if (currentStep === 6) {
|
||||
// 已经安装过,不能返回
|
||||
return;
|
||||
} else {
|
||||
@@ -138,14 +147,33 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
|
||||
const handleInstallComplete = useCallback((results: PluginInstallState[]) => {
|
||||
setInstallResults(results);
|
||||
setCurrentStep(5);
|
||||
setCurrentStep(6);
|
||||
}, []);
|
||||
|
||||
const handleFinish = useCallback(() => {
|
||||
const handleFinish = useCallback(async () => {
|
||||
// 保存窗口尺寸偏好
|
||||
if (windowSizePreference) {
|
||||
localStorage.setItem(
|
||||
STORAGE_KEYS.WINDOW_SIZE_PREFERENCE,
|
||||
windowSizePreference,
|
||||
);
|
||||
|
||||
// 应用窗口尺寸
|
||||
try {
|
||||
if (windowSizePreference === "fullscreen") {
|
||||
await windowApi.toggleFullscreen();
|
||||
} else {
|
||||
await windowApi.setWindowSizeByOption(windowSizePreference);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("应用窗口尺寸失败:", error);
|
||||
}
|
||||
}
|
||||
|
||||
// 触发插件变化事件,刷新侧边栏
|
||||
window.dispatchEvent(new CustomEvent("plugin-changed"));
|
||||
onComplete();
|
||||
}, [onComplete]);
|
||||
}, [onComplete, windowSizePreference]);
|
||||
|
||||
// 渲染当前步骤
|
||||
const renderStep = () => {
|
||||
@@ -160,6 +188,13 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
/>
|
||||
);
|
||||
case 3:
|
||||
return (
|
||||
<WindowSizeStep
|
||||
selectedSize={windowSizePreference}
|
||||
onSelect={setWindowSizePreference}
|
||||
/>
|
||||
);
|
||||
case 4:
|
||||
return (
|
||||
<PluginSelectStep
|
||||
userProfile={userProfile}
|
||||
@@ -167,14 +202,14 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
onSelectionChange={setSelectedPlugins}
|
||||
/>
|
||||
);
|
||||
case 4:
|
||||
case 5:
|
||||
return (
|
||||
<InstallProgressStep
|
||||
selectedPlugins={selectedPlugins}
|
||||
onComplete={handleInstallComplete}
|
||||
/>
|
||||
);
|
||||
case 5:
|
||||
case 6:
|
||||
return (
|
||||
<CompleteStep
|
||||
installResults={installResults}
|
||||
@@ -191,6 +226,8 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
switch (currentStep) {
|
||||
case 2:
|
||||
return userProfile !== null;
|
||||
case 3:
|
||||
return windowSizePreference !== null;
|
||||
default:
|
||||
return true;
|
||||
}
|
||||
@@ -198,7 +235,7 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
|
||||
// 判断是否显示底部导航
|
||||
const showFooter =
|
||||
currentStep !== 1 && currentStep !== 4 && currentStep !== 5;
|
||||
currentStep !== 1 && currentStep !== 5 && currentStep !== 6;
|
||||
|
||||
return (
|
||||
<Overlay>
|
||||
@@ -230,7 +267,7 @@ export function OnboardingWizard({ onComplete }: OnboardingWizardProps) {
|
||||
跳过
|
||||
</Button>
|
||||
<Button onClick={handleNext} disabled={!canProceed()}>
|
||||
{currentStep === 3
|
||||
{currentStep === 4
|
||||
? selectedPlugins.length > 0
|
||||
? "开始安装"
|
||||
: "跳过安装"
|
||||
|
||||
@@ -2,7 +2,18 @@
|
||||
* 初次安装引导 - 常量配置
|
||||
*/
|
||||
|
||||
import { Code, User, FileCode, Activity, Cpu, Globe } from "lucide-react";
|
||||
import {
|
||||
Code,
|
||||
User,
|
||||
FileCode,
|
||||
Activity,
|
||||
Cpu,
|
||||
Globe,
|
||||
Minimize2,
|
||||
Monitor,
|
||||
Maximize2,
|
||||
Fullscreen,
|
||||
} from "lucide-react";
|
||||
import type { LucideIcon } from "lucide-react";
|
||||
|
||||
/**
|
||||
@@ -102,4 +113,61 @@ export const STORAGE_KEYS = {
|
||||
ONBOARDING_COMPLETE: "proxycast_onboarding_complete",
|
||||
ONBOARDING_VERSION: "proxycast_onboarding_version",
|
||||
USER_PROFILE: "proxycast_user_profile",
|
||||
WINDOW_SIZE_PREFERENCE: "proxycast_window_size_preference",
|
||||
} as const;
|
||||
|
||||
/**
|
||||
* 窗口尺寸偏好类型
|
||||
*/
|
||||
export type WindowSizePreference =
|
||||
| "compact"
|
||||
| "default"
|
||||
| "flow_monitor"
|
||||
| "large"
|
||||
| "fullscreen";
|
||||
|
||||
/**
|
||||
* 窗口尺寸选项配置
|
||||
*/
|
||||
export interface WindowSizeOptionConfig {
|
||||
id: WindowSizePreference;
|
||||
name: string;
|
||||
description: string;
|
||||
icon: LucideIcon;
|
||||
}
|
||||
|
||||
/**
|
||||
* 窗口尺寸选项列表
|
||||
*/
|
||||
export const windowSizeOptions: WindowSizeOptionConfig[] = [
|
||||
{
|
||||
id: "compact",
|
||||
name: "紧凑模式",
|
||||
description: "1000×700 - 节省屏幕空间",
|
||||
icon: Minimize2,
|
||||
},
|
||||
{
|
||||
id: "default",
|
||||
name: "默认大小",
|
||||
description: "1200×800 - 日常使用",
|
||||
icon: Monitor,
|
||||
},
|
||||
{
|
||||
id: "flow_monitor",
|
||||
name: "Flow Monitor",
|
||||
description: "1600×1000 - 数据展示优化",
|
||||
icon: Activity,
|
||||
},
|
||||
{
|
||||
id: "large",
|
||||
name: "大屏模式",
|
||||
description: "1920×1200 - 大屏幕显示",
|
||||
icon: Maximize2,
|
||||
},
|
||||
{
|
||||
id: "fullscreen",
|
||||
name: "全屏模式",
|
||||
description: "占满整个屏幕",
|
||||
icon: Fullscreen,
|
||||
},
|
||||
];
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
/**
|
||||
* 初次安装引导 - 窗口尺寸选择
|
||||
*/
|
||||
|
||||
import styled from "styled-components";
|
||||
import { Check } from "lucide-react";
|
||||
import { windowSizeOptions, type WindowSizePreference } from "../constants";
|
||||
|
||||
const Container = styled.div`
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
padding: 32px 24px;
|
||||
`;
|
||||
|
||||
const Title = styled.h2`
|
||||
font-size: 24px;
|
||||
font-weight: 600;
|
||||
color: hsl(var(--foreground));
|
||||
margin-bottom: 8px;
|
||||
text-align: center;
|
||||
`;
|
||||
|
||||
const Subtitle = styled.p`
|
||||
font-size: 14px;
|
||||
color: hsl(var(--muted-foreground));
|
||||
margin-bottom: 32px;
|
||||
text-align: center;
|
||||
`;
|
||||
|
||||
const OptionsGrid = styled.div`
|
||||
display: grid;
|
||||
grid-template-columns: repeat(2, 1fr);
|
||||
gap: 12px;
|
||||
width: 100%;
|
||||
max-width: 500px;
|
||||
`;
|
||||
|
||||
const OptionCard = styled.button<{ $selected?: boolean }>`
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
padding: 20px 16px;
|
||||
border-radius: 12px;
|
||||
border: 2px solid
|
||||
${({ $selected }) =>
|
||||
$selected ? "hsl(var(--primary))" : "hsl(var(--border))"};
|
||||
background: ${({ $selected }) =>
|
||||
$selected ? "hsl(var(--primary) / 0.05)" : "hsl(var(--card))"};
|
||||
cursor: pointer;
|
||||
transition: all 0.2s;
|
||||
position: relative;
|
||||
|
||||
&:hover {
|
||||
border-color: ${({ $selected }) =>
|
||||
$selected ? "hsl(var(--primary))" : "hsl(var(--primary) / 0.5)"};
|
||||
}
|
||||
`;
|
||||
|
||||
const CheckBadge = styled.div`
|
||||
position: absolute;
|
||||
top: 10px;
|
||||
right: 10px;
|
||||
width: 20px;
|
||||
height: 20px;
|
||||
border-radius: 50%;
|
||||
background: hsl(var(--primary));
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
|
||||
svg {
|
||||
width: 12px;
|
||||
height: 12px;
|
||||
color: hsl(var(--primary-foreground));
|
||||
}
|
||||
`;
|
||||
|
||||
const IconWrapper = styled.div<{ $selected?: boolean }>`
|
||||
width: 48px;
|
||||
height: 48px;
|
||||
border-radius: 10px;
|
||||
background: ${({ $selected }) =>
|
||||
$selected ? "hsl(var(--primary))" : "hsl(var(--muted))"};
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
margin-bottom: 12px;
|
||||
transition: all 0.2s;
|
||||
|
||||
svg {
|
||||
width: 24px;
|
||||
height: 24px;
|
||||
color: ${({ $selected }) =>
|
||||
$selected ? "hsl(var(--primary-foreground))" : "hsl(var(--foreground))"};
|
||||
}
|
||||
`;
|
||||
|
||||
const OptionName = styled.span`
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
color: hsl(var(--foreground));
|
||||
margin-bottom: 4px;
|
||||
`;
|
||||
|
||||
const OptionDescription = styled.span`
|
||||
font-size: 11px;
|
||||
color: hsl(var(--muted-foreground));
|
||||
text-align: center;
|
||||
`;
|
||||
|
||||
interface WindowSizeStepProps {
|
||||
selectedSize: WindowSizePreference | null;
|
||||
onSelect: (size: WindowSizePreference) => void;
|
||||
}
|
||||
|
||||
export function WindowSizeStep({
|
||||
selectedSize,
|
||||
onSelect,
|
||||
}: WindowSizeStepProps) {
|
||||
return (
|
||||
<Container>
|
||||
<Title>选择窗口大小</Title>
|
||||
<Subtitle>您可以随时在设置中更改窗口大小</Subtitle>
|
||||
|
||||
<OptionsGrid>
|
||||
{windowSizeOptions.map((option) => {
|
||||
const isSelected = selectedSize === option.id;
|
||||
const Icon = option.icon;
|
||||
|
||||
return (
|
||||
<OptionCard
|
||||
key={option.id}
|
||||
$selected={isSelected}
|
||||
onClick={() => onSelect(option.id)}
|
||||
>
|
||||
{isSelected && (
|
||||
<CheckBadge>
|
||||
<Check />
|
||||
</CheckBadge>
|
||||
)}
|
||||
<IconWrapper $selected={isSelected}>
|
||||
<Icon />
|
||||
</IconWrapper>
|
||||
<OptionName>{option.name}</OptionName>
|
||||
<OptionDescription>{option.description}</OptionDescription>
|
||||
</OptionCard>
|
||||
);
|
||||
})}
|
||||
</OptionsGrid>
|
||||
</Container>
|
||||
);
|
||||
}
|
||||
@@ -269,42 +269,71 @@ export function PluginUIRenderer({
|
||||
setLoading(true);
|
||||
setError(null);
|
||||
|
||||
// 调试日志函数
|
||||
const debugLog = async (msg: string) => {
|
||||
console.log(msg);
|
||||
try {
|
||||
await safeInvoke("frontend_debug_log", { message: msg });
|
||||
} catch {
|
||||
// 忽略错误
|
||||
}
|
||||
};
|
||||
|
||||
try {
|
||||
// 获取插件目录
|
||||
await debugLog(`[PluginUIRenderer] 开始检查插件: ${pluginId}`);
|
||||
const dir = await safeInvoke<string>("get_plugins_dir");
|
||||
await debugLog(`[PluginUIRenderer] 插件目录: ${dir}`);
|
||||
setPluginsDir(dir);
|
||||
|
||||
// 首先尝试读取插件清单
|
||||
await debugLog("[PluginUIRenderer] 调用 read_plugin_manifest_cmd...");
|
||||
const manifest = await safeInvoke<PluginManifest | null>(
|
||||
"read_plugin_manifest_cmd",
|
||||
{
|
||||
pluginId,
|
||||
},
|
||||
);
|
||||
await debugLog(
|
||||
`[PluginUIRenderer] manifest 结果: ${JSON.stringify(manifest)}`,
|
||||
);
|
||||
|
||||
if (manifest) {
|
||||
setPluginManifest(manifest);
|
||||
|
||||
// 检查数据库中是否已注册
|
||||
await debugLog(`[PluginUIRenderer] 检查是否已安装...`);
|
||||
const installed = await safeInvoke<boolean>("is_plugin_installed", {
|
||||
pluginId,
|
||||
});
|
||||
await debugLog(
|
||||
`[PluginUIRenderer] is_plugin_installed: ${installed}`,
|
||||
);
|
||||
|
||||
if (installed) {
|
||||
// 从数据库获取插件信息
|
||||
await debugLog(`[PluginUIRenderer] 获取已安装插件列表...`);
|
||||
const plugins = await safeInvoke<InstalledPlugin[]>(
|
||||
"list_installed_plugins",
|
||||
);
|
||||
await debugLog(
|
||||
`[PluginUIRenderer] 已安装插件数量: ${plugins.length}`,
|
||||
);
|
||||
const plugin = plugins.find((p) => p.id === pluginId);
|
||||
await debugLog(
|
||||
`[PluginUIRenderer] 找到插件: ${JSON.stringify(plugin)}`,
|
||||
);
|
||||
|
||||
if (plugin) {
|
||||
setPluginInfo(plugin);
|
||||
setLoading(false);
|
||||
await debugLog(`[PluginUIRenderer] 设置 pluginInfo 成功`);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
// 插件存在于文件系统中但未在数据库注册,创建临时的插件信息
|
||||
await debugLog(`[PluginUIRenderer] 创建临时插件信息...`);
|
||||
setPluginInfo({
|
||||
id: pluginId,
|
||||
name: manifest.name,
|
||||
@@ -313,14 +342,19 @@ export function PluginUIRenderer({
|
||||
ui_entry: undefined,
|
||||
});
|
||||
setLoading(false);
|
||||
await debugLog(`[PluginUIRenderer] 临时插件信息已设置`);
|
||||
return;
|
||||
}
|
||||
|
||||
// 插件不存在
|
||||
await debugLog(`[PluginUIRenderer] manifest 为空,插件不存在`);
|
||||
setPluginInfo(null);
|
||||
setPluginManifest(null);
|
||||
} catch (err) {
|
||||
console.error("检查插件失败:", err);
|
||||
await debugLog(
|
||||
`[PluginUIRenderer] 错误: ${err instanceof Error ? err.message : String(err)}`,
|
||||
);
|
||||
setError(err instanceof Error ? err.message : String(err));
|
||||
} finally {
|
||||
setLoading(false);
|
||||
|
||||
@@ -11,6 +11,7 @@ import {
|
||||
Info,
|
||||
RotateCcw,
|
||||
Volume2,
|
||||
Maximize2,
|
||||
} from "lucide-react";
|
||||
import { cn, validateProxyUrl } from "@/lib/utils";
|
||||
import { getConfig, saveConfig, Config } from "@/hooks/useTauri";
|
||||
@@ -18,6 +19,8 @@ import { useOnboardingState } from "@/components/onboarding";
|
||||
import { LanguageSelector, Language } from "./LanguageSelector";
|
||||
import { useI18nPatch } from "@/i18n/I18nPatchProvider";
|
||||
import { useSoundContext } from "@/contexts/useSoundContext";
|
||||
import { windowApi, type WindowSizeOption } from "@/lib/api/window";
|
||||
import { STORAGE_KEYS } from "@/components/onboarding/constants";
|
||||
|
||||
type Theme = "light" | "dark" | "system";
|
||||
|
||||
@@ -31,6 +34,14 @@ export function GeneralSettings() {
|
||||
const { soundEnabled, setSoundEnabled, playToolcallSound } =
|
||||
useSoundContext();
|
||||
|
||||
// 窗口尺寸状态
|
||||
const [windowSizeOptions, setWindowSizeOptions] = useState<
|
||||
WindowSizeOption[]
|
||||
>([]);
|
||||
const [currentWindowSize, setCurrentWindowSize] = useState<string>("default");
|
||||
const [isFullscreen, setIsFullscreen] = useState(false);
|
||||
const [windowSizeLoading, setWindowSizeLoading] = useState(true);
|
||||
|
||||
// 重新运行引导
|
||||
const handleResetOnboarding = useCallback(() => {
|
||||
resetOnboarding();
|
||||
@@ -54,8 +65,53 @@ export function GeneralSettings() {
|
||||
setTheme(savedTheme);
|
||||
}
|
||||
loadConfig();
|
||||
loadWindowSizeOptions();
|
||||
}, []);
|
||||
|
||||
const loadWindowSizeOptions = async () => {
|
||||
setWindowSizeLoading(true);
|
||||
try {
|
||||
const options = await windowApi.getWindowSizeOptions();
|
||||
setWindowSizeOptions(options);
|
||||
|
||||
const fullscreen = await windowApi.isFullscreen();
|
||||
setIsFullscreen(fullscreen);
|
||||
|
||||
// 从 localStorage 读取保存的偏好
|
||||
const savedPreference = localStorage.getItem(
|
||||
STORAGE_KEYS.WINDOW_SIZE_PREFERENCE,
|
||||
);
|
||||
if (savedPreference) {
|
||||
setCurrentWindowSize(savedPreference);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("加载窗口尺寸选项失败:", error);
|
||||
} finally {
|
||||
setWindowSizeLoading(false);
|
||||
}
|
||||
};
|
||||
|
||||
const handleWindowSizeChange = async (optionId: string) => {
|
||||
try {
|
||||
if (optionId === "fullscreen") {
|
||||
if (!isFullscreen) {
|
||||
await windowApi.toggleFullscreen();
|
||||
setIsFullscreen(true);
|
||||
}
|
||||
} else {
|
||||
if (isFullscreen) {
|
||||
await windowApi.toggleFullscreen();
|
||||
setIsFullscreen(false);
|
||||
}
|
||||
await windowApi.setWindowSizeByOption(optionId);
|
||||
}
|
||||
setCurrentWindowSize(optionId);
|
||||
localStorage.setItem(STORAGE_KEYS.WINDOW_SIZE_PREFERENCE, optionId);
|
||||
} catch (error) {
|
||||
console.error("设置窗口尺寸失败:", error);
|
||||
}
|
||||
};
|
||||
|
||||
const loadConfig = async () => {
|
||||
setConfigLoading(true);
|
||||
try {
|
||||
@@ -217,6 +273,52 @@ export function GeneralSettings() {
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 窗口尺寸 */}
|
||||
<div className="rounded-lg border p-3">
|
||||
<div className="flex items-center gap-2 mb-3">
|
||||
<Maximize2 className="h-4 w-4 text-muted-foreground" />
|
||||
<h3 className="text-sm font-medium">窗口尺寸</h3>
|
||||
</div>
|
||||
|
||||
{windowSizeLoading ? (
|
||||
<div className="flex items-center justify-center py-2">
|
||||
<RefreshCw className="h-4 w-4 animate-spin text-muted-foreground" />
|
||||
</div>
|
||||
) : (
|
||||
<div className="grid grid-cols-2 gap-2">
|
||||
{windowSizeOptions.map((option) => (
|
||||
<button
|
||||
key={option.id}
|
||||
onClick={() => handleWindowSizeChange(option.id)}
|
||||
className={cn(
|
||||
"flex flex-col items-start p-2 rounded border text-left transition-colors",
|
||||
currentWindowSize === option.id && !isFullscreen
|
||||
? "border-primary bg-primary/5"
|
||||
: "hover:bg-muted",
|
||||
)}
|
||||
>
|
||||
<span className="text-sm font-medium">{option.name}</span>
|
||||
<span className="text-xs text-muted-foreground">
|
||||
{option.description}
|
||||
</span>
|
||||
</button>
|
||||
))}
|
||||
<button
|
||||
onClick={() => handleWindowSizeChange("fullscreen")}
|
||||
className={cn(
|
||||
"flex flex-col items-start p-2 rounded border text-left transition-colors",
|
||||
isFullscreen ? "border-primary bg-primary/5" : "hover:bg-muted",
|
||||
)}
|
||||
>
|
||||
<span className="text-sm font-medium">全屏模式</span>
|
||||
<span className="text-xs text-muted-foreground">
|
||||
占满整个屏幕
|
||||
</span>
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 语言 */}
|
||||
<div className="rounded-lg border p-3">
|
||||
<div className="flex items-center justify-between">
|
||||
|
||||
@@ -497,6 +497,154 @@ export async function listGooseProviders(): Promise<GooseProviderInfo[]> {
|
||||
return await safeInvoke("goose_agent_list_providers");
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// Aster Agent API (基于 Aster 框架的 Agent 实现)
|
||||
// ============================================================
|
||||
|
||||
/**
|
||||
* Aster Agent 状态
|
||||
*/
|
||||
export interface AsterAgentStatus {
|
||||
initialized: boolean;
|
||||
provider_configured: boolean;
|
||||
provider_name?: string;
|
||||
model_name?: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* Aster Provider 配置
|
||||
*/
|
||||
export interface AsterProviderConfig {
|
||||
provider_name: string;
|
||||
model_name: string;
|
||||
api_key?: string;
|
||||
base_url?: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* Aster 会话信息
|
||||
*/
|
||||
export interface AsterSessionInfo {
|
||||
id: string;
|
||||
name?: string;
|
||||
created_at: string;
|
||||
updated_at: string;
|
||||
messages_count: number;
|
||||
}
|
||||
|
||||
/**
|
||||
* Aster 会话详情
|
||||
*/
|
||||
export interface AsterSessionDetail {
|
||||
id: string;
|
||||
name?: string;
|
||||
messages: Array<{
|
||||
role: string;
|
||||
content: string;
|
||||
timestamp: string;
|
||||
}>;
|
||||
}
|
||||
|
||||
/**
|
||||
* 初始化 Aster Agent
|
||||
*/
|
||||
export async function initAsterAgent(): Promise<AsterAgentStatus> {
|
||||
return await safeInvoke("aster_agent_init");
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取 Aster Agent 状态
|
||||
*/
|
||||
export async function getAsterAgentStatus(): Promise<AsterAgentStatus> {
|
||||
return await safeInvoke("aster_agent_status");
|
||||
}
|
||||
|
||||
/**
|
||||
* 配置 Aster Agent 的 Provider
|
||||
*/
|
||||
export async function configureAsterProvider(
|
||||
config: AsterProviderConfig,
|
||||
sessionId: string,
|
||||
): Promise<AsterAgentStatus> {
|
||||
return await safeInvoke("aster_agent_configure_provider", {
|
||||
request: config,
|
||||
session_id: sessionId,
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* 发送消息到 Aster Agent (流式响应)
|
||||
*
|
||||
* 通过 Tauri 事件接收响应流
|
||||
*/
|
||||
export async function sendAsterMessageStream(
|
||||
message: string,
|
||||
sessionId: string,
|
||||
eventName: string,
|
||||
images?: ImageInput[],
|
||||
providerConfig?: AsterProviderConfig,
|
||||
): Promise<void> {
|
||||
return await safeInvoke("aster_agent_chat_stream", {
|
||||
request: {
|
||||
message,
|
||||
session_id: sessionId,
|
||||
event_name: eventName,
|
||||
images,
|
||||
provider_config: providerConfig,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* 停止 Aster Agent 会话
|
||||
*/
|
||||
export async function stopAsterSession(sessionId: string): Promise<boolean> {
|
||||
return await safeInvoke("aster_agent_stop", { sessionId });
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建 Aster 会话
|
||||
*/
|
||||
export async function createAsterSession(
|
||||
workingDir?: string,
|
||||
name?: string,
|
||||
): Promise<string> {
|
||||
return await safeInvoke("aster_session_create", { workingDir, name });
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取 Aster 会话列表
|
||||
*/
|
||||
export async function listAsterSessions(): Promise<AsterSessionInfo[]> {
|
||||
return await safeInvoke("aster_session_list");
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取 Aster 会话详情
|
||||
*/
|
||||
export async function getAsterSession(
|
||||
sessionId: string,
|
||||
): Promise<AsterSessionDetail> {
|
||||
return await safeInvoke("aster_session_get", { sessionId });
|
||||
}
|
||||
|
||||
/**
|
||||
* 确认 Aster Agent 权限请求
|
||||
*/
|
||||
export async function confirmAsterAction(
|
||||
requestId: string,
|
||||
confirmed: boolean,
|
||||
response?: string,
|
||||
): Promise<void> {
|
||||
return await safeInvoke("aster_agent_confirm", {
|
||||
request: {
|
||||
request_id: requestId,
|
||||
confirmed,
|
||||
response,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// Terminal Tool API (终端命令执行)
|
||||
// ============================================================
|
||||
|
||||
@@ -0,0 +1,644 @@
|
||||
/**
|
||||
* Aster Agent Zustand Store
|
||||
*
|
||||
* 基于 Aster 框架的 Agent 状态管理
|
||||
* 参考 Claude-Cowork 的设计模式
|
||||
*/
|
||||
|
||||
import { create } from "zustand";
|
||||
import { invoke } from "@tauri-apps/api/core";
|
||||
import { listen, type UnlistenFn } from "@tauri-apps/api/event";
|
||||
|
||||
// ============ 类型定义 ============
|
||||
|
||||
/** 消息图片 */
|
||||
export interface MessageImage {
|
||||
data: string;
|
||||
mediaType: string;
|
||||
}
|
||||
|
||||
/** 工具调用结果 */
|
||||
export interface ToolResult {
|
||||
success: boolean;
|
||||
output?: string;
|
||||
error?: string;
|
||||
}
|
||||
|
||||
/** 工具调用状态 */
|
||||
export interface ToolCallState {
|
||||
id: string;
|
||||
name: string;
|
||||
arguments?: string;
|
||||
status: "pending" | "running" | "completed" | "failed";
|
||||
result?: ToolResult;
|
||||
startTime?: Date;
|
||||
endTime?: Date;
|
||||
}
|
||||
|
||||
/** Token 使用量 */
|
||||
export interface TokenUsage {
|
||||
input_tokens: number;
|
||||
output_tokens: number;
|
||||
total_tokens: number;
|
||||
}
|
||||
|
||||
/** 内容片段类型 */
|
||||
export type ContentPart =
|
||||
| { type: "text"; text: string }
|
||||
| { type: "thinking"; text: string }
|
||||
| { type: "tool_use"; toolCall: ToolCallState };
|
||||
|
||||
/** 消息 */
|
||||
export interface Message {
|
||||
id: string;
|
||||
role: "user" | "assistant";
|
||||
content: string;
|
||||
images?: MessageImage[];
|
||||
timestamp: Date;
|
||||
isThinking?: boolean;
|
||||
thinkingContent?: string;
|
||||
toolCalls?: ToolCallState[];
|
||||
usage?: TokenUsage;
|
||||
contentParts?: ContentPart[];
|
||||
}
|
||||
|
||||
/** 会话信息 */
|
||||
export interface SessionInfo {
|
||||
id: string;
|
||||
name?: string;
|
||||
createdAt: Date;
|
||||
updatedAt: Date;
|
||||
messagesCount: number;
|
||||
}
|
||||
|
||||
/** 权限确认请求 */
|
||||
export interface ActionRequired {
|
||||
requestId: string;
|
||||
actionType: "tool_confirmation" | "ask_user" | "permission_request";
|
||||
toolName?: string;
|
||||
arguments?: Record<string, unknown>;
|
||||
question?: string;
|
||||
options?: Array<{
|
||||
label: string;
|
||||
description?: string;
|
||||
}>;
|
||||
timestamp: Date;
|
||||
}
|
||||
|
||||
/** 确认响应 */
|
||||
export interface ConfirmResponse {
|
||||
requestId: string;
|
||||
confirmed: boolean;
|
||||
response?: string;
|
||||
}
|
||||
|
||||
// ============ Tauri 事件类型 ============
|
||||
|
||||
/** Tauri Agent 事件 */
|
||||
export type TauriAgentEvent =
|
||||
| { type: "text_delta"; text: string }
|
||||
| { type: "thinking_delta"; text: string }
|
||||
| {
|
||||
type: "tool_start";
|
||||
tool_name: string;
|
||||
tool_id: string;
|
||||
arguments?: string;
|
||||
}
|
||||
| { type: "tool_end"; tool_id: string; result: ToolResult }
|
||||
| {
|
||||
type: "action_required";
|
||||
request_id: string;
|
||||
action_type: string;
|
||||
data: Record<string, unknown>;
|
||||
}
|
||||
| { type: "model_change"; model: string; mode: string }
|
||||
| { type: "done"; usage?: TokenUsage }
|
||||
| { type: "final_done"; usage?: TokenUsage }
|
||||
| { type: "error"; message: string }
|
||||
| { type: "message"; message: unknown };
|
||||
|
||||
// ============ Store 状态类型 ============
|
||||
|
||||
interface AgentState {
|
||||
// 会话状态
|
||||
currentSessionId: string | null;
|
||||
sessions: SessionInfo[];
|
||||
messages: Message[];
|
||||
|
||||
// 流式状态
|
||||
isStreaming: boolean;
|
||||
currentAssistantMsgId: string | null;
|
||||
|
||||
// 权限确认
|
||||
pendingActions: ActionRequired[];
|
||||
|
||||
// 配置
|
||||
isInitialized: boolean;
|
||||
|
||||
// Actions
|
||||
initialize: () => Promise<void>;
|
||||
sendMessage: (content: string, images?: MessageImage[]) => Promise<void>;
|
||||
stopStreaming: () => Promise<void>;
|
||||
confirmAction: (response: ConfirmResponse) => Promise<void>;
|
||||
switchSession: (sessionId: string) => Promise<void>;
|
||||
createSession: (name?: string) => Promise<string>;
|
||||
deleteSession: (sessionId: string) => Promise<void>;
|
||||
clearMessages: () => void;
|
||||
loadSessions: () => Promise<void>;
|
||||
|
||||
// 内部方法
|
||||
_handleEvent: (event: TauriAgentEvent) => void;
|
||||
_cleanup: () => void;
|
||||
}
|
||||
|
||||
// ============ Store 实现 ============
|
||||
|
||||
// 事件监听器引用
|
||||
let eventUnlisten: UnlistenFn | null = null;
|
||||
|
||||
export const useAgentStore = create<AgentState>((set, get) => ({
|
||||
// 初始状态
|
||||
currentSessionId: null,
|
||||
sessions: [],
|
||||
messages: [],
|
||||
isStreaming: false,
|
||||
currentAssistantMsgId: null,
|
||||
pendingActions: [],
|
||||
isInitialized: false,
|
||||
|
||||
// 初始化 Agent
|
||||
initialize: async () => {
|
||||
try {
|
||||
await invoke("aster_agent_init");
|
||||
set({ isInitialized: true });
|
||||
console.log("[AgentStore] Agent 初始化成功");
|
||||
|
||||
// 加载会话列表
|
||||
await get().loadSessions();
|
||||
} catch (error) {
|
||||
console.error("[AgentStore] Agent 初始化失败:", error);
|
||||
throw error;
|
||||
}
|
||||
},
|
||||
|
||||
// 发送消息
|
||||
sendMessage: async (content: string, images?: MessageImage[]) => {
|
||||
const state = get();
|
||||
|
||||
// 确保已初始化
|
||||
if (!state.isInitialized) {
|
||||
await state.initialize();
|
||||
}
|
||||
|
||||
// 确保有会话
|
||||
let sessionId = state.currentSessionId;
|
||||
if (!sessionId) {
|
||||
sessionId = await state.createSession();
|
||||
}
|
||||
|
||||
// 创建用户消息
|
||||
const userMsg: Message = {
|
||||
id: crypto.randomUUID(),
|
||||
role: "user",
|
||||
content,
|
||||
images,
|
||||
timestamp: new Date(),
|
||||
};
|
||||
|
||||
// 创建助手消息占位符
|
||||
const assistantMsgId = crypto.randomUUID();
|
||||
const assistantMsg: Message = {
|
||||
id: assistantMsgId,
|
||||
role: "assistant",
|
||||
content: "",
|
||||
timestamp: new Date(),
|
||||
isThinking: true,
|
||||
thinkingContent: "思考中...",
|
||||
contentParts: [],
|
||||
};
|
||||
|
||||
// 更新状态
|
||||
set((s) => ({
|
||||
messages: [...s.messages, userMsg, assistantMsg],
|
||||
isStreaming: true,
|
||||
currentAssistantMsgId: assistantMsgId,
|
||||
}));
|
||||
|
||||
// 创建唯一事件名称
|
||||
const eventName = `aster_stream_${assistantMsgId}`;
|
||||
|
||||
try {
|
||||
// 设置事件监听器
|
||||
eventUnlisten = await listen<TauriAgentEvent>(eventName, (event) => {
|
||||
get()._handleEvent(event.payload);
|
||||
});
|
||||
|
||||
// 发送请求
|
||||
await invoke("aster_agent_chat_stream", {
|
||||
request: {
|
||||
message: content,
|
||||
session_id: sessionId,
|
||||
event_name: eventName,
|
||||
images: images?.map((img) => ({
|
||||
data: img.data,
|
||||
media_type: img.mediaType,
|
||||
})),
|
||||
},
|
||||
});
|
||||
} catch (error) {
|
||||
console.error("[AgentStore] 发送消息失败:", error);
|
||||
|
||||
// 更新消息状态为错误
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) =>
|
||||
msg.id === assistantMsgId
|
||||
? {
|
||||
...msg,
|
||||
isThinking: false,
|
||||
content: `错误: ${error}`,
|
||||
}
|
||||
: msg,
|
||||
),
|
||||
isStreaming: false,
|
||||
currentAssistantMsgId: null,
|
||||
}));
|
||||
|
||||
get()._cleanup();
|
||||
throw error;
|
||||
}
|
||||
},
|
||||
|
||||
// 停止流式响应
|
||||
stopStreaming: async () => {
|
||||
const state = get();
|
||||
if (!state.currentSessionId) return;
|
||||
|
||||
try {
|
||||
await invoke("aster_agent_stop", {
|
||||
sessionId: state.currentSessionId,
|
||||
});
|
||||
} catch (error) {
|
||||
console.error("[AgentStore] 停止失败:", error);
|
||||
}
|
||||
|
||||
// 更新消息状态
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) =>
|
||||
msg.id === s.currentAssistantMsgId
|
||||
? {
|
||||
...msg,
|
||||
isThinking: false,
|
||||
content: msg.content || "(已停止生成)",
|
||||
}
|
||||
: msg,
|
||||
),
|
||||
isStreaming: false,
|
||||
currentAssistantMsgId: null,
|
||||
}));
|
||||
|
||||
get()._cleanup();
|
||||
},
|
||||
|
||||
// 确认权限请求
|
||||
confirmAction: async (response: ConfirmResponse) => {
|
||||
try {
|
||||
await invoke("aster_agent_confirm", {
|
||||
request: {
|
||||
request_id: response.requestId,
|
||||
confirmed: response.confirmed,
|
||||
response: response.response,
|
||||
},
|
||||
});
|
||||
|
||||
// 移除已处理的请求
|
||||
set((s) => ({
|
||||
pendingActions: s.pendingActions.filter(
|
||||
(a) => a.requestId !== response.requestId,
|
||||
),
|
||||
}));
|
||||
} catch (error) {
|
||||
console.error("[AgentStore] 确认失败:", error);
|
||||
throw error;
|
||||
}
|
||||
},
|
||||
|
||||
// 切换会话
|
||||
switchSession: async (sessionId: string) => {
|
||||
try {
|
||||
const detail = await invoke<{
|
||||
id: string;
|
||||
name?: string;
|
||||
messages: Array<{
|
||||
role: string;
|
||||
content: string;
|
||||
timestamp: string;
|
||||
}>;
|
||||
}>("aster_session_get", { sessionId });
|
||||
|
||||
// 转换消息格式
|
||||
const messages: Message[] = detail.messages.map((msg, index) => ({
|
||||
id: `${sessionId}-${index}`,
|
||||
role: msg.role as "user" | "assistant",
|
||||
content: msg.content,
|
||||
timestamp: new Date(msg.timestamp),
|
||||
}));
|
||||
|
||||
set({
|
||||
currentSessionId: sessionId,
|
||||
messages,
|
||||
pendingActions: [],
|
||||
});
|
||||
} catch (error) {
|
||||
console.error("[AgentStore] 切换会话失败:", error);
|
||||
throw error;
|
||||
}
|
||||
},
|
||||
|
||||
// 创建会话
|
||||
createSession: async (name?: string) => {
|
||||
try {
|
||||
const sessionId = await invoke<string>("aster_session_create", {
|
||||
workingDir: null,
|
||||
name,
|
||||
});
|
||||
|
||||
set((s) => ({
|
||||
currentSessionId: sessionId,
|
||||
messages: [],
|
||||
sessions: [
|
||||
{
|
||||
id: sessionId,
|
||||
name,
|
||||
createdAt: new Date(),
|
||||
updatedAt: new Date(),
|
||||
messagesCount: 0,
|
||||
},
|
||||
...s.sessions,
|
||||
],
|
||||
}));
|
||||
|
||||
return sessionId;
|
||||
} catch (error) {
|
||||
console.error("[AgentStore] 创建会话失败:", error);
|
||||
throw error;
|
||||
}
|
||||
},
|
||||
|
||||
// 删除会话
|
||||
deleteSession: async (sessionId: string) => {
|
||||
// TODO: 实现后端删除接口
|
||||
set((s) => ({
|
||||
sessions: s.sessions.filter((sess) => sess.id !== sessionId),
|
||||
...(s.currentSessionId === sessionId
|
||||
? { currentSessionId: null, messages: [] }
|
||||
: {}),
|
||||
}));
|
||||
},
|
||||
|
||||
// 清空消息
|
||||
clearMessages: () => {
|
||||
set({
|
||||
messages: [],
|
||||
currentSessionId: null,
|
||||
pendingActions: [],
|
||||
});
|
||||
},
|
||||
|
||||
// 加载会话列表
|
||||
loadSessions: async () => {
|
||||
try {
|
||||
const sessions = await invoke<
|
||||
Array<{
|
||||
id: string;
|
||||
name?: string;
|
||||
created_at: string;
|
||||
updated_at: string;
|
||||
messages_count: number;
|
||||
}>
|
||||
>("aster_session_list");
|
||||
|
||||
set({
|
||||
sessions: sessions.map((s) => ({
|
||||
id: s.id,
|
||||
name: s.name,
|
||||
createdAt: new Date(s.created_at),
|
||||
updatedAt: new Date(s.updated_at),
|
||||
messagesCount: s.messages_count,
|
||||
})),
|
||||
});
|
||||
} catch (error) {
|
||||
console.error("[AgentStore] 加载会话列表失败:", error);
|
||||
}
|
||||
},
|
||||
|
||||
// 处理事件
|
||||
_handleEvent: (event: TauriAgentEvent) => {
|
||||
const state = get();
|
||||
const msgId = state.currentAssistantMsgId;
|
||||
if (!msgId) return;
|
||||
|
||||
console.log("[AgentStore] 收到事件:", event.type, event);
|
||||
|
||||
switch (event.type) {
|
||||
case "text_delta":
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) => {
|
||||
if (msg.id !== msgId) return msg;
|
||||
|
||||
const newContent = msg.content + event.text;
|
||||
const newParts = [...(msg.contentParts || [])];
|
||||
|
||||
// 追加到最后一个 text 类型,或创建新的
|
||||
const lastPart = newParts[newParts.length - 1];
|
||||
if (lastPart && lastPart.type === "text") {
|
||||
newParts[newParts.length - 1] = {
|
||||
type: "text",
|
||||
text: lastPart.text + event.text,
|
||||
};
|
||||
} else {
|
||||
newParts.push({ type: "text", text: event.text });
|
||||
}
|
||||
|
||||
return {
|
||||
...msg,
|
||||
content: newContent,
|
||||
isThinking: false,
|
||||
thinkingContent: undefined,
|
||||
contentParts: newParts,
|
||||
};
|
||||
}),
|
||||
}));
|
||||
break;
|
||||
|
||||
case "thinking_delta":
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) => {
|
||||
if (msg.id !== msgId) return msg;
|
||||
|
||||
const newParts = [...(msg.contentParts || [])];
|
||||
const lastPart = newParts[newParts.length - 1];
|
||||
if (lastPart && lastPart.type === "thinking") {
|
||||
newParts[newParts.length - 1] = {
|
||||
type: "thinking",
|
||||
text: lastPart.text + event.text,
|
||||
};
|
||||
} else {
|
||||
newParts.push({ type: "thinking", text: event.text });
|
||||
}
|
||||
|
||||
return {
|
||||
...msg,
|
||||
thinkingContent: (msg.thinkingContent || "") + event.text,
|
||||
contentParts: newParts,
|
||||
};
|
||||
}),
|
||||
}));
|
||||
break;
|
||||
|
||||
case "tool_start": {
|
||||
const newToolCall: ToolCallState = {
|
||||
id: event.tool_id,
|
||||
name: event.tool_name,
|
||||
arguments: event.arguments,
|
||||
status: "running",
|
||||
startTime: new Date(),
|
||||
};
|
||||
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) => {
|
||||
if (msg.id !== msgId) return msg;
|
||||
|
||||
// 检查是否已存在
|
||||
if (msg.toolCalls?.find((tc) => tc.id === event.tool_id)) {
|
||||
return msg;
|
||||
}
|
||||
|
||||
return {
|
||||
...msg,
|
||||
toolCalls: [...(msg.toolCalls || []), newToolCall],
|
||||
contentParts: [
|
||||
...(msg.contentParts || []),
|
||||
{ type: "tool_use" as const, toolCall: newToolCall },
|
||||
],
|
||||
};
|
||||
}),
|
||||
}));
|
||||
break;
|
||||
}
|
||||
|
||||
case "tool_end":
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) => {
|
||||
if (msg.id !== msgId) return msg;
|
||||
|
||||
const updatedToolCalls = (msg.toolCalls || []).map((tc) =>
|
||||
tc.id === event.tool_id
|
||||
? {
|
||||
...tc,
|
||||
status: event.result.success
|
||||
? ("completed" as const)
|
||||
: ("failed" as const),
|
||||
result: event.result,
|
||||
endTime: new Date(),
|
||||
}
|
||||
: tc,
|
||||
);
|
||||
|
||||
const updatedContentParts = (msg.contentParts || []).map((part) => {
|
||||
if (
|
||||
part.type === "tool_use" &&
|
||||
part.toolCall.id === event.tool_id
|
||||
) {
|
||||
return {
|
||||
...part,
|
||||
toolCall: {
|
||||
...part.toolCall,
|
||||
status: event.result.success
|
||||
? ("completed" as const)
|
||||
: ("failed" as const),
|
||||
result: event.result,
|
||||
endTime: new Date(),
|
||||
},
|
||||
};
|
||||
}
|
||||
return part;
|
||||
});
|
||||
|
||||
return {
|
||||
...msg,
|
||||
toolCalls: updatedToolCalls,
|
||||
contentParts: updatedContentParts,
|
||||
};
|
||||
}),
|
||||
}));
|
||||
break;
|
||||
|
||||
case "action_required":
|
||||
set((s) => ({
|
||||
pendingActions: [
|
||||
...s.pendingActions,
|
||||
{
|
||||
requestId: event.request_id,
|
||||
actionType: event.action_type as ActionRequired["actionType"],
|
||||
...event.data,
|
||||
timestamp: new Date(),
|
||||
},
|
||||
],
|
||||
}));
|
||||
break;
|
||||
|
||||
case "done":
|
||||
// 单次响应完成,但工具循环可能继续
|
||||
console.log("[AgentStore] done 事件,等待 final_done...");
|
||||
break;
|
||||
|
||||
case "final_done":
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) =>
|
||||
msg.id === msgId
|
||||
? {
|
||||
...msg,
|
||||
isThinking: false,
|
||||
usage: event.usage,
|
||||
}
|
||||
: msg,
|
||||
),
|
||||
isStreaming: false,
|
||||
currentAssistantMsgId: null,
|
||||
}));
|
||||
get()._cleanup();
|
||||
break;
|
||||
|
||||
case "error":
|
||||
set((s) => ({
|
||||
messages: s.messages.map((msg) =>
|
||||
msg.id === msgId
|
||||
? {
|
||||
...msg,
|
||||
isThinking: false,
|
||||
content: msg.content || `错误: ${event.message}`,
|
||||
}
|
||||
: msg,
|
||||
),
|
||||
isStreaming: false,
|
||||
currentAssistantMsgId: null,
|
||||
}));
|
||||
get()._cleanup();
|
||||
break;
|
||||
}
|
||||
},
|
||||
|
||||
// 清理资源
|
||||
_cleanup: () => {
|
||||
if (eventUnlisten) {
|
||||
eventUnlisten();
|
||||
eventUnlisten = null;
|
||||
}
|
||||
},
|
||||
}));
|
||||
|
||||
// 导出便捷 hooks
|
||||
export const useAgentMessages = () => useAgentStore((s) => s.messages);
|
||||
export const useAgentStreaming = () => useAgentStore((s) => s.isStreaming);
|
||||
export const useAgentSessions = () => useAgentStore((s) => s.sessions);
|
||||
export const usePendingActions = () => useAgentStore((s) => s.pendingActions);
|
||||
@@ -0,0 +1,22 @@
|
||||
/**
|
||||
* Stores 导出
|
||||
*/
|
||||
|
||||
// Aster Agent Store
|
||||
export {
|
||||
useAgentStore,
|
||||
useAgentMessages,
|
||||
useAgentStreaming,
|
||||
useAgentSessions,
|
||||
usePendingActions,
|
||||
type Message,
|
||||
type MessageImage,
|
||||
type ToolResult,
|
||||
type ToolCallState,
|
||||
type TokenUsage,
|
||||
type ContentPart,
|
||||
type SessionInfo,
|
||||
type ActionRequired,
|
||||
type ConfirmResponse,
|
||||
type TauriAgentEvent,
|
||||
} from "./agentStore";
|
||||
Reference in New Issue
Block a user