chore: bump version to 0.46.0

- 修复 reasoning_content 字段缺失的编译错误
- 精简 README 为 AI Agent 创作工具平台
- 修复 lint 和格式错误
This commit is contained in:
coso
2026-01-15 02:13:58 +08:00
parent 91fea5dd8c
commit 1f489126d8
50 changed files with 8808 additions and 594 deletions
+18 -199
View File
@@ -2,7 +2,7 @@
# ProxyCast 🚀
**把你的 AI 客户端额度用到任何地方**
**AI Agent 创作工具平台**
[![License: GPL v3](https://img.shields.io/badge/License-GPLv3-blue.svg)](https://www.gnu.org/licenses/gpl-3.0)
[![Tauri](https://img.shields.io/badge/Tauri-2.0-blue.svg)](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` - 从备份恢复
---
## 📸 界面截图
### 仪表盘 - 系统状态与监控
![Dashboard](docs/images/943663ed-b17c-4b32-a74c-c0243ffb3dea.png)
### 凭证池 - 多凭证管理与配额查询
![Provider Pool](docs/images/aee62eb5-3aeb-4454-b14d-24b1d5f9a0fe.png)
### 路由管理 - 智能路由规则和容错策略
![Router](docs/images/067c7d64-e116-4a30-b533-748873166f37.png)
### 配置管理 - 客户端配置切换
![Config](docs/images/25eb018a-5be2-4f82-ba22-e68f39160cac.png)
### 扩展 - MCP/Prompts/Skills 管理
![Extensions](docs/images/ffc70018-aa5f-4738-883d-045614488608.png)
### API Server - 服务控制与 API 测试
![API Server](docs/images/151b4355-821c-4bda-a731-c4367b6b8716.png)
### 设置 - 应用参数和偏好
![Settings](docs/images/c7d8236b-ea6c-4496-ada5-288cd0a01738.png)
- **多 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
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.45.2",
"version": "0.46.0",
"type": "module",
"repository": {
"type": "git",
+351
View File
@@ -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();
+13
View File
@@ -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"
}
}
+3495 -346
View File
File diff suppressed because it is too large Load Diff
+5 -1
View File
@@ -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
+205
View File
@@ -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");
}
}
+293
View File
@@ -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);
}
}
+451
View File
@@ -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
View File
@@ -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};
+6
View File
@@ -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();
}
+20
View File
@@ -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,
})
}
}
+10
View File
@@ -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,
})
}
}
+4
View File
@@ -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
+468
View File
@@ -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"));
}
}
+2
View File
@@ -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};
+14
View File
@@ -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
+4
View File
@@ -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,
+2
View File
@@ -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),
+8 -3
View File
@@ -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);
+13
View File
@@ -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,
+312
View File
@@ -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
View File
@@ -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);
+2 -3
View File
@@ -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。
+120 -12
View File
@@ -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,
});
}
}
+1
View File
@@ -201,6 +201,7 @@ impl AgentDao {
timestamp,
tool_calls,
tool_call_id,
reasoning_content: None,
})
})?;
+4
View File
@@ -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 {
+3
View File
@@ -21,3 +21,6 @@ pub use types::{
NoopProgressCallback, PackageFormat, ProgressCallback,
};
pub use validator::PackageValidator;
#[cfg(test)]
mod tests;
+554
View File
@@ -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());
}
}
+3 -1
View File
@@ -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")
+34
View File
@@ -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
View File
@@ -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 -1
View File
@@ -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
View File
@@ -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;
+35
View File
@@ -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";
}
+49
View File
@@ -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,
};
}
+51 -14
View File
@@ -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
? "开始安装"
: "跳过安装"
+69 -1
View File
@@ -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);
+102
View File
@@ -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">
+148
View File
@@ -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 (终端命令执行)
// ============================================================
+644
View File
@@ -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);
+22
View File
@@ -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";