mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
Merge branch 'main' into main
This commit is contained in:
+20
-17
@@ -1,24 +1,27 @@
|
||||
{
|
||||
"permissions": {
|
||||
"allow": [
|
||||
"Bash(npm run format:*)",
|
||||
"Bash(npm run lint)",
|
||||
"Bash(npx tsc:*)",
|
||||
"Bash(cargo test:*)",
|
||||
"Bash(cargo build:*)",
|
||||
"Bash(npm run check:*)",
|
||||
"Bash(npm run:*)",
|
||||
"Bash(tree:*)",
|
||||
"Bash(find:*)",
|
||||
"Bash(cargo check:*)",
|
||||
"Bash(rm:*)",
|
||||
"Bash(cargo clippy:*)",
|
||||
"Bash(cargo fmt:*)",
|
||||
"Bash(lsof:*)",
|
||||
"Bash(xargs kill:*)",
|
||||
"Bash(cargo run:*)"
|
||||
"Bash",
|
||||
"Read(*)",
|
||||
"Write(*)",
|
||||
"Edit(*)",
|
||||
"MultiEdit(*)",
|
||||
"Glob(*)",
|
||||
"Grep(*)",
|
||||
"Task(*)",
|
||||
"TaskOutput(*)",
|
||||
"LSP(*)",
|
||||
"NotebookEdit(*)",
|
||||
"TodoWrite(*)",
|
||||
"AskUserQuestion(*)",
|
||||
"EnterPlanMode(*)",
|
||||
"ExitPlanMode(*)",
|
||||
"KillShell(*)",
|
||||
"WebFetch(domain:*)",
|
||||
"Skill(*)",
|
||||
"SlashCommand(*)"
|
||||
],
|
||||
"deny": [],
|
||||
"ask": []
|
||||
"defaultMode": "bypassPermissions"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -47,3 +47,17 @@ jobs:
|
||||
- name: Check Rust build
|
||||
working-directory: src-tauri
|
||||
run: cargo check --all-targets
|
||||
|
||||
- name: Lint frontend
|
||||
run: npm run lint
|
||||
|
||||
- name: Test frontend
|
||||
run: npm test
|
||||
|
||||
- name: Clippy
|
||||
working-directory: src-tauri
|
||||
run: cargo clippy --all-targets
|
||||
|
||||
- name: Test Rust
|
||||
working-directory: src-tauri
|
||||
run: cargo test
|
||||
|
||||
@@ -97,8 +97,8 @@ jobs:
|
||||
```
|
||||
|
||||
### 默认配置
|
||||
- **端口**: 3001
|
||||
- **API Key**: proxycast-key
|
||||
- **端口**: 8999
|
||||
- **API Key**: 首次启动自动生成,可在设置页查看/修改
|
||||
releaseDraft: false
|
||||
prerelease: false
|
||||
args: --target ${{ matrix.target }}
|
||||
|
||||
@@ -74,7 +74,7 @@
|
||||
- **Per-Key 代理** - 为每个凭证单独配置代理
|
||||
|
||||
### 🔐 安全与管理
|
||||
- **TLS/HTTPS 支持** - 可选启用 HTTPS 加密通信
|
||||
- **HTTPS 部署** - 当前版本不内置 TLS,请使用反向代理进行 HTTPS 终止
|
||||
- **远程管理 API** - 通过 API 远程管理配置和凭证
|
||||
- **访问控制** - 支持 localhost 限制和密钥认证
|
||||
|
||||
@@ -88,8 +88,12 @@
|
||||
- `/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` - 从备份恢复
|
||||
|
||||
---
|
||||
|
||||
@@ -135,20 +139,29 @@
|
||||
3. **启动服务** - 在 Dashboard 点击"启动服务器"
|
||||
4. **配置客户端** - 在 Cherry-Studio、Cline 等工具中配置:
|
||||
```
|
||||
API Base URL: http://localhost:3001/v1
|
||||
API Key: proxycast-key
|
||||
API Base URL: http://localhost:8999/v1
|
||||
API Key: 启动时自动生成的密钥(可在设置页查看/修改)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🧰 运维提示
|
||||
|
||||
- **自动备份**:数据库默认每天自动备份到 `~/.proxycast/backups/`,保留 7 天。
|
||||
- **配置备份**:每次写入配置会生成 `config.yaml.backup` 以便回滚。
|
||||
- **日志归档**:7 天游离线日志自动压缩,30 天前压缩日志自动清理。
|
||||
- **生产 HTTPS**:当前版本不内置 TLS,生产环境需反向代理终止 HTTPS。
|
||||
|
||||
---
|
||||
|
||||
## 🔧 API 使用示例
|
||||
|
||||
### OpenAI Chat Completions
|
||||
|
||||
```bash
|
||||
curl http://localhost:3001/v1/chat/completions \
|
||||
curl http://localhost:8999/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer proxycast-key" \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-d '{
|
||||
"model": "claude-sonnet-4-5-20250514",
|
||||
"messages": [
|
||||
@@ -161,9 +174,9 @@ curl http://localhost:3001/v1/chat/completions \
|
||||
### Anthropic Messages API
|
||||
|
||||
```bash
|
||||
curl http://localhost:3001/v1/messages \
|
||||
curl http://localhost:8999/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "x-api-key: proxycast-key" \
|
||||
-H "x-api-key: your-api-key" \
|
||||
-H "anthropic-version: 2023-06-01" \
|
||||
-d '{
|
||||
"model": "claude-sonnet-4-5-20250514",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+4
-4
@@ -85,8 +85,8 @@ src/
|
||||
### 路由模式
|
||||
|
||||
```
|
||||
http://localhost:3000/{provider}/v1/chat/completions
|
||||
http://localhost:3000/{provider}/v1/messages
|
||||
http://localhost:8999/{provider}/v1/chat/completions
|
||||
http://localhost:8999/{provider}/v1/messages
|
||||
```
|
||||
|
||||
### 支持的端点
|
||||
@@ -104,8 +104,8 @@ http://localhost:3000/{provider}/v1/messages
|
||||
{
|
||||
"server": {
|
||||
"host": "127.0.0.1",
|
||||
"port": 3000,
|
||||
"apiKey": "proxycast-key"
|
||||
"port": 8999,
|
||||
"apiKey": "your-api-key"
|
||||
},
|
||||
"providers": {
|
||||
"kiro": {
|
||||
|
||||
@@ -46,7 +46,7 @@ ProxyCast 会自动检测本地的 AI 客户端凭证文件。
|
||||
|
||||
1. 在仪表盘点击 **启动服务**
|
||||
2. 服务状态变为"运行中"
|
||||
3. 记下 API 地址(默认 `http://127.0.0.1:9090`)
|
||||
3. 记下 API 地址(默认 `http://127.0.0.1:8999`)
|
||||
|
||||
## 步骤 4: 测试 API
|
||||
|
||||
@@ -61,7 +61,7 @@ ProxyCast 会自动检测本地的 AI 客户端凭证文件。
|
||||
**OpenAI 格式:**
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
curl http://127.0.0.1:8999/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-d '{
|
||||
@@ -73,7 +73,7 @@ curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
**Claude 格式:**
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/messages \
|
||||
curl http://127.0.0.1:8999/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "x-api-key: your-api-key" \
|
||||
-H "anthropic-version: 2023-06-01" \
|
||||
@@ -90,7 +90,7 @@ curl http://127.0.0.1:9090/v1/messages \
|
||||
|
||||
在 Cursor 设置中配置 OpenAI API:
|
||||
|
||||
- API Base URL: `http://127.0.0.1:9090/v1`
|
||||
- API Base URL: `http://127.0.0.1:8999/v1`
|
||||
- API Key: 你在 ProxyCast 中设置的 API Key
|
||||
|
||||
### Continue 配置
|
||||
@@ -103,7 +103,7 @@ curl http://127.0.0.1:9090/v1/messages \
|
||||
"title": "ProxyCast Claude",
|
||||
"provider": "openai",
|
||||
"model": "claude-sonnet-4-20250514",
|
||||
"apiBase": "http://127.0.0.1:9090/v1",
|
||||
"apiBase": "http://127.0.0.1:8999/v1",
|
||||
"apiKey": "your-api-key"
|
||||
}]
|
||||
}
|
||||
|
||||
@@ -29,7 +29,7 @@ navigation:
|
||||
|
||||
服务运行时显示:
|
||||
|
||||
- **API 地址**: 本地 API 端点(如 `http://127.0.0.1:9090`)
|
||||
- **API 地址**: 本地 API 端点(如 `http://127.0.0.1:8999`)
|
||||
- **API Key**: 当前配置的访问密钥
|
||||
- **复制按钮**: 一键复制 API 地址或 Key
|
||||
|
||||
|
||||
@@ -101,7 +101,7 @@ coding:
|
||||
### 通过 API 调用
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
curl http://127.0.0.1:8999/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-d '{
|
||||
|
||||
@@ -16,7 +16,7 @@ navigation:
|
||||
server:
|
||||
host: "127.0.0.1"
|
||||
port: 8999
|
||||
api_key: "proxy_cast"
|
||||
api_key: "your-api-key"
|
||||
|
||||
# TLS/HTTPS 配置
|
||||
tls:
|
||||
@@ -24,6 +24,8 @@ server:
|
||||
cert_path: "/path/to/cert.pem"
|
||||
key_path: "/path/to/key.pem"
|
||||
|
||||
# 注意:当前版本暂不支持 TLS。启用后服务将无法启动,请使用反向代理做 TLS 终止。
|
||||
|
||||
# 全局代理 URL(支持 socks5/http/https)
|
||||
proxy_url: "socks5://127.0.0.1:1080"
|
||||
|
||||
@@ -278,7 +280,7 @@ injection:
|
||||
server:
|
||||
host: "127.0.0.1"
|
||||
port: 8999
|
||||
api_key: "proxy_cast"
|
||||
api_key: "your-api-key"
|
||||
tls:
|
||||
enable: false
|
||||
cert_path: ""
|
||||
|
||||
@@ -56,7 +56,7 @@ resilience:
|
||||
|
||||
server:
|
||||
host: "127.0.0.1"
|
||||
port: 9090
|
||||
port: 8999
|
||||
```
|
||||
|
||||
### 敏感信息处理
|
||||
@@ -107,15 +107,15 @@ server:
|
||||
|
||||
```bash
|
||||
# ProxyCast API Configuration
|
||||
PROXYCAST_API_BASE=http://127.0.0.1:9090/v1
|
||||
PROXYCAST_API_BASE=http://127.0.0.1:8999/v1
|
||||
PROXYCAST_API_KEY=your-api-key
|
||||
|
||||
# OpenAI Compatible
|
||||
OPENAI_API_BASE=http://127.0.0.1:9090/v1
|
||||
OPENAI_API_BASE=http://127.0.0.1:8999/v1
|
||||
OPENAI_API_KEY=your-api-key
|
||||
|
||||
# Claude Compatible
|
||||
ANTHROPIC_API_BASE=http://127.0.0.1:9090
|
||||
ANTHROPIC_API_BASE=http://127.0.0.1:8999
|
||||
ANTHROPIC_API_KEY=your-api-key
|
||||
```
|
||||
|
||||
@@ -142,6 +142,25 @@ ProxyCast 会自动备份配置:
|
||||
3. 选择要恢复的版本
|
||||
4. 点击 **恢复**
|
||||
|
||||
## 完整备份与恢复(生产建议)
|
||||
|
||||
仅导出配置无法覆盖数据库与凭证文件。生产环境建议定期备份以下路径:
|
||||
|
||||
- 配置文件:macOS `~/Library/Application Support/proxycast/config.yaml`;Linux `~/.config/proxycast/config.yaml`;Windows `%APPDATA%\\proxycast\\config.yaml`
|
||||
- 凭证副本目录:macOS `~/Library/Application Support/proxycast/credentials/`;Linux `~/.local/share/proxycast/credentials/`;Windows `%APPDATA%\\proxycast\\credentials\\`
|
||||
- 数据库与日志:`~/.proxycast/`(含 `proxycast.db`、`logs/`、`request_logs/`、`auth/`)
|
||||
|
||||
```bash
|
||||
# 示例:备份数据库与日志目录
|
||||
cp -a ~/.proxycast ~/.proxycast.backup-$(date +%Y%m%d%H%M%S)
|
||||
```
|
||||
|
||||
恢复时将备份内容替换回原路径,并确保应用已退出。
|
||||
|
||||
## 旧版本迁移说明
|
||||
|
||||
如果检测到旧版 `~/.proxycast/config.json`,当前版本会阻止启动并提示手动迁移。请先导出旧配置或重新导入 YAML 配置,再启动应用。
|
||||
|
||||
## 配置同步
|
||||
|
||||
### 跨设备同步
|
||||
|
||||
@@ -16,7 +16,7 @@ API Server 是 ProxyCast 的核心组件,提供 OpenAI/Claude 兼容的 API
|
||||
| 选项 | 默认值 | 说明 |
|
||||
|------|--------|------|
|
||||
| 主机地址 | `127.0.0.1` | 监听地址 |
|
||||
| 端口 | `9090` | 监听端口 |
|
||||
| 端口 | `8999` | 监听端口 |
|
||||
| API Key | 自动生成 | 访问密钥 |
|
||||
|
||||
### 配置步骤
|
||||
@@ -31,10 +31,10 @@ API Server 是 ProxyCast 的核心组件,提供 OpenAI/Claude 兼容的 API
|
||||
| 地址 | 说明 |
|
||||
|------|------|
|
||||
| `127.0.0.1` | 仅本机访问 |
|
||||
| `0.0.0.0` | 允许局域网访问 |
|
||||
| `localhost` | 仅本机访问 |
|
||||
|
||||
::alert{type="warning"}
|
||||
使用 `0.0.0.0` 时请确保配置了 API Key 认证。
|
||||
当前版本仅支持本地监听(127.0.0.1/localhost/::1),不支持对外开放。
|
||||
::
|
||||
|
||||
## API 端点
|
||||
@@ -99,7 +99,7 @@ API Server 是 ProxyCast 的核心组件,提供 OpenAI/Claude 兼容的 API
|
||||
|
||||
**OpenAI 格式:**
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
curl http://127.0.0.1:8999/v1/chat/completions \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"model": "gpt-4", "messages": [...]}'
|
||||
@@ -107,7 +107,7 @@ curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
|
||||
**Claude 格式:**
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/messages \
|
||||
curl http://127.0.0.1:8999/v1/messages \
|
||||
-H "x-api-key: your-api-key" \
|
||||
-H "anthropic-version: 2023-06-01" \
|
||||
-H "Content-Type: application/json" \
|
||||
|
||||
@@ -93,7 +93,7 @@ MCP 定义了 AI 模型与外部系统交互的标准方式:
|
||||
通过 API 调用 MCP 工具:
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
curl http://127.0.0.1:8999/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-d '{
|
||||
|
||||
@@ -11,7 +11,8 @@ navigation:
|
||||
|
||||
## 概述
|
||||
|
||||
Codex Provider 允许你使用 OpenAI Codex 的 OAuth 凭证访问 GPT 模型,无需 API Key。
|
||||
Codex Provider 允许你使用 OpenAI Codex 的 OAuth 凭证访问 GPT 模型。
|
||||
同时,为了兼容 Codex CLI 的「API Key 登录」用户(`~/.codex/auth.json` 只有 `api_key`),ProxyCast 也支持读取 `api_key` 并作为 Bearer Token 使用(无需刷新)。
|
||||
|
||||
## 支持的模型
|
||||
|
||||
@@ -54,6 +55,8 @@ credential_pool:
|
||||
|
||||
### Token 文件格式
|
||||
|
||||
#### OAuth 模式(推荐用于 Codex OAuth)
|
||||
|
||||
```json
|
||||
{
|
||||
"access_token": "eyJ...",
|
||||
@@ -63,6 +66,19 @@ credential_pool:
|
||||
}
|
||||
```
|
||||
|
||||
#### API Key 模式(兼容 Codex CLI)
|
||||
|
||||
```json
|
||||
{
|
||||
"api_key": "sk-xxx",
|
||||
"api_base_url": "https://api.openai.com"
|
||||
}
|
||||
```
|
||||
|
||||
说明:
|
||||
- `api_key` / `apiKey`:必填
|
||||
- `api_base_url` / `apiBaseUrl`:可选;可填写 `https://api.openai.com` 或带 `/v1` 的地址(例如网关、反代、Azure 兼容地址)
|
||||
|
||||
## Token 刷新
|
||||
|
||||
ProxyCast 会自动在 Token 过期前刷新:
|
||||
|
||||
@@ -50,7 +50,7 @@ ProxyCast 提供 OpenAI 和 Claude 兼容的 API 端点。
|
||||
使用 `Authorization` 头:
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
curl http://127.0.0.1:8999/v1/chat/completions \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '...'
|
||||
@@ -61,7 +61,7 @@ curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
使用 `x-api-key` 头:
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/messages \
|
||||
curl http://127.0.0.1:8999/v1/messages \
|
||||
-H "x-api-key: your-api-key" \
|
||||
-H "anthropic-version: 2023-06-01" \
|
||||
-H "Content-Type: application/json" \
|
||||
@@ -70,7 +70,7 @@ curl http://127.0.0.1:9090/v1/messages \
|
||||
|
||||
## 基础 URL
|
||||
|
||||
默认地址:`http://127.0.0.1:9090`
|
||||
默认地址:`http://127.0.0.1:8999`
|
||||
|
||||
可在设置中修改主机和端口。
|
||||
|
||||
|
||||
@@ -92,7 +92,7 @@ Authorization: Bearer your-api-key
|
||||
设置 `stream: true` 启用流式响应:
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
curl http://127.0.0.1:8999/v1/chat/completions \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
@@ -202,7 +202,7 @@ Authorization: Bearer your-api-key
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(
|
||||
base_url="http://127.0.0.1:9090/v1",
|
||||
base_url="http://127.0.0.1:8999/v1",
|
||||
api_key="your-api-key"
|
||||
)
|
||||
|
||||
@@ -220,7 +220,7 @@ print(response.choices[0].message.content)
|
||||
import OpenAI from 'openai';
|
||||
|
||||
const client = new OpenAI({
|
||||
baseURL: 'http://127.0.0.1:9090/v1',
|
||||
baseURL: 'http://127.0.0.1:8999/v1',
|
||||
apiKey: 'your-api-key'
|
||||
});
|
||||
|
||||
|
||||
@@ -93,7 +93,7 @@ anthropic-version: 2023-06-01
|
||||
设置 `stream: true` 启用流式响应:
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/messages \
|
||||
curl http://127.0.0.1:8999/v1/messages \
|
||||
-H "x-api-key: your-api-key" \
|
||||
-H "anthropic-version: 2023-06-01" \
|
||||
-H "Content-Type: application/json" \
|
||||
@@ -200,7 +200,7 @@ anthropic-version: 2023-06-01
|
||||
import anthropic
|
||||
|
||||
client = anthropic.Anthropic(
|
||||
base_url="http://127.0.0.1:9090",
|
||||
base_url="http://127.0.0.1:8999",
|
||||
api_key="your-api-key"
|
||||
)
|
||||
|
||||
@@ -219,7 +219,7 @@ print(message.content[0].text)
|
||||
import Anthropic from '@anthropic-ai/sdk';
|
||||
|
||||
const client = new Anthropic({
|
||||
baseURL: 'http://127.0.0.1:9090',
|
||||
baseURL: 'http://127.0.0.1:8999',
|
||||
apiKey: 'your-api-key'
|
||||
});
|
||||
|
||||
|
||||
@@ -26,6 +26,10 @@ Authorization: Bearer your-secret-key
|
||||
| `secret_key` | 管理密钥,为空时禁用所有管理端点(返回 404) |
|
||||
| `allow_remote` | 是否允许远程访问,为 false 时仅允许 localhost |
|
||||
|
||||
::alert{type="warning"}
|
||||
当前版本未启用 TLS,仅支持本地访问,`allow_remote` 必须保持为 `false`。
|
||||
::
|
||||
|
||||
## /v0/management/status
|
||||
|
||||
获取服务器状态信息。
|
||||
|
||||
@@ -30,10 +30,10 @@ navigation:
|
||||
```bash
|
||||
# 查找占用端口的进程
|
||||
# macOS/Linux
|
||||
lsof -i :9090
|
||||
lsof -i :8999
|
||||
|
||||
# Windows
|
||||
netstat -ano | findstr :9090
|
||||
netstat -ano | findstr :8999
|
||||
```
|
||||
|
||||
或在设置中更改端口号。
|
||||
|
||||
@@ -152,7 +152,7 @@ export NO_PROXY=localhost,127.0.0.1
|
||||
| 端口 | 用途 |
|
||||
|------|------|
|
||||
| 443 | HTTPS 请求 |
|
||||
| 9090 | ProxyCast API(默认) |
|
||||
| 8999 | ProxyCast API(默认) |
|
||||
|
||||
### macOS 防火墙
|
||||
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
---
|
||||
title: 上线运维
|
||||
description: 生产就绪的最小运维清单
|
||||
navigation:
|
||||
icon: i-heroicons-wrench-screwdriver
|
||||
---
|
||||
|
||||
# 上线运行与运维(生产就绪最小版)
|
||||
|
||||
本页面用于“马上上线且长期稳定运行”的最小运维闭环,避免上线后因配置、备份或回滚缺失导致不可恢复的问题。
|
||||
|
||||
## 上线前检查
|
||||
|
||||
- 确认服务仅本地监听:`server.host = 127.0.0.1`(当前版本仅支持本地监听)
|
||||
- 设置强 API Key:不要使用默认值 `proxy_cast`
|
||||
- 确认日志保留策略:`logging.retention_days` 合理(建议 >= 7 天)
|
||||
- 确认凭证与配置已正确导入,并完成一次启动 + 健康检查
|
||||
|
||||
## 运行健康检查
|
||||
|
||||
- HTTP 健康检查:`GET /health`
|
||||
- 关键字段应包含 `status=healthy` 与 `version`
|
||||
- 建议在上线后做一次 API 冒烟请求(如 `/v1/models`)
|
||||
|
||||
## 备份与恢复(必须)
|
||||
|
||||
当前版本需要手动备份以下路径:
|
||||
|
||||
- 配置文件(macOS: `~/Library/Application Support/proxycast/config.yaml`,Linux: `~/.config/proxycast/config.yaml`,Windows: `%APPDATA%\\proxycast\\config.yaml`)
|
||||
- 凭证池副本目录(导入的凭证文件):macOS `~/Library/Application Support/proxycast/credentials/`,Linux `~/.local/share/proxycast/credentials/`,Windows `%APPDATA%\\proxycast\\credentials\\`
|
||||
- OAuth 凭证目录:`~/.proxycast/auth/`
|
||||
- 数据库文件:`~/.proxycast/proxycast.db`
|
||||
- 日志目录:`~/.proxycast/logs/`、`~/.proxycast/request_logs/`
|
||||
|
||||
恢复步骤(顺序建议):
|
||||
|
||||
1. 停止应用
|
||||
2. 恢复 `config.yaml`
|
||||
3. 恢复 `credentials` 目录、`auth` 目录与数据库 `proxycast.db`
|
||||
4. 如需保留历史日志,恢复 `logs/` 与 `request_logs/`
|
||||
5. 启动应用并验证 `/health` 与关键功能
|
||||
|
||||
## 回滚策略
|
||||
|
||||
- 如果升级失败,恢复备份的 `config.yaml` 与 `proxycast.db`
|
||||
- 使用上一版本安装包覆盖安装
|
||||
- 完成健康检查与冒烟测试
|
||||
|
||||
## 发布质量门槛(最小)
|
||||
|
||||
- `cd src-tauri && cargo test`
|
||||
- `cd src-tauri && cargo clippy`
|
||||
- `npm test`
|
||||
- `npm run lint`
|
||||
- `npm run build`
|
||||
@@ -81,12 +81,12 @@ ProxyCast 会自动检测本地的 AI 客户端凭证文件:
|
||||
|
||||
### 3. 启动服务
|
||||
|
||||
点击仪表盘的「启动服务」按钮,API Server 默认运行在 `http://127.0.0.1:9090`。
|
||||
点击仪表盘的「启动服务」按钮,API Server 默认运行在 `http://127.0.0.1:8999`。
|
||||
|
||||
### 4. 测试 API
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:9090/v1/chat/completions \
|
||||
curl http://127.0.0.1:8999/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer your-api-key" \
|
||||
-d '{
|
||||
|
||||
+87
@@ -0,0 +1,87 @@
|
||||
# 生产运维与上线就绪清单
|
||||
|
||||
本文档用于生产环境部署、运行、备份与恢复的最小操作规范,避免上线后缺少可执行流程。
|
||||
|
||||
## 部署前检查
|
||||
|
||||
- 确认 API Key 已更换(禁止使用默认值 `proxy_cast`)。
|
||||
- 确认监听地址:
|
||||
- 本机使用 `127.0.0.1`/`localhost`。
|
||||
- 当前版本仅支持本地监听,不支持对外服务。
|
||||
- 若需要 HTTPS,请使用反向代理终止 TLS;当前服务端未启用内置 TLS。
|
||||
- 确认磁盘权限可写:`~/.proxycast/`、`~/.proxycast/request_logs/`、应用数据目录(macOS: `~/Library/Application Support/proxycast/`,Linux: `~/.local/share/proxycast/`,Windows: `%APPDATA%\\proxycast\\`)。
|
||||
|
||||
## 配置路径与加载顺序
|
||||
|
||||
- YAML 配置(优先):
|
||||
- macOS: `~/Library/Application Support/proxycast/config.yaml`
|
||||
- Linux: `~/.config/proxycast/config.yaml`
|
||||
- Windows: `%APPDATA%\\proxycast\\config.yaml`
|
||||
- JSON 配置(兼容):macOS `~/Library/Application Support/proxycast/config.json`,Linux `~/.config/proxycast/config.json`,Windows `%APPDATA%\\proxycast\\config.json`
|
||||
- 旧版遗留路径:`~/.proxycast/config.json`(检测到会提示手动迁移)
|
||||
- 两者都不存在时使用默认配置。
|
||||
- 首次启动会自动生成强随机 API Key 并写入配置。
|
||||
|
||||
## 数据与日志位置
|
||||
|
||||
- SQLite 数据库:`~/.proxycast/proxycast.db`
|
||||
- 凭证池副本目录:macOS `~/Library/Application Support/proxycast/credentials/`,Linux `~/.local/share/proxycast/credentials/`,Windows `%APPDATA%\\proxycast\\credentials\\`
|
||||
- OAuth/Token 目录(默认):`~/.proxycast/auth/`
|
||||
- 日志目录:`~/.proxycast/logs/`
|
||||
- 请求日志目录:`~/.proxycast/request_logs/`
|
||||
- 数据库备份目录:`~/.proxycast/backups/`
|
||||
|
||||
## 备份与恢复
|
||||
|
||||
### 备份
|
||||
|
||||
1. 可使用管理端点触发备份(需配置管理密钥):
|
||||
- `POST /v0/management/backup`
|
||||
2. 或手动备份(建议停服后执行):
|
||||
- 复制以下路径:
|
||||
- 配置文件(macOS: `~/Library/Application Support/proxycast/config.yaml`,Linux: `~/.config/proxycast/config.yaml`,Windows: `%APPDATA%\\proxycast\\config.yaml`)
|
||||
- 配置备份文件:`config.yaml.backup`
|
||||
- 凭证池副本目录(macOS: `~/Library/Application Support/proxycast/credentials/`,Linux: `~/.local/share/proxycast/credentials/`,Windows: `%APPDATA%\\proxycast\\credentials\\`)
|
||||
- `~/.proxycast/proxycast.db`
|
||||
- `~/.proxycast/auth/`(如需要保留 OAuth/Token)
|
||||
- `~/.proxycast/logs/`、`~/.proxycast/request_logs/`(如需保留日志)
|
||||
3. 将备份文件存入受控存储(加密磁盘或安全存储)。
|
||||
|
||||
### 自动备份
|
||||
|
||||
- 服务运行期间每 24 小时自动创建数据库备份到 `~/.proxycast/backups/`。
|
||||
- 备份默认保留 7 天,过期文件会被清理。
|
||||
|
||||
### 恢复
|
||||
|
||||
1. 停止 ProxyCast 服务。
|
||||
2. 使用管理端点恢复(建议停服后执行,执行时会锁定数据库并短暂阻塞请求):
|
||||
- `POST /v0/management/restore`,请求体:`{"backup_path": "/path/to/proxycast_YYYYMMDD_HHMMSS.db"}`
|
||||
3. 或手动恢复上述文件到原路径。
|
||||
4. 启动服务并检查 `/health` 与 `/ready`。
|
||||
|
||||
## 升级与回滚
|
||||
|
||||
- 升级前执行备份流程。
|
||||
- 升级后若出现异常:
|
||||
- 恢复备份文件。
|
||||
- 回滚到上一个稳定版本的安装包。
|
||||
|
||||
## 运行与排障
|
||||
|
||||
- 健康检查:`GET /health`
|
||||
- 就绪检查:`GET /ready`
|
||||
- 常见问题排查:
|
||||
- 端口占用:修改配置端口或释放占用端口。
|
||||
- 配置解析失败:检查 YAML/JSON 语法,确认缩进正确。
|
||||
- 数据库初始化失败:检查 `~/.proxycast/` 权限与磁盘空间。
|
||||
|
||||
## 安全基线
|
||||
|
||||
- 禁止默认 API key。
|
||||
- 当前版本未实现内置 TLS,远程管理必须保持关闭且仅本地访问。
|
||||
|
||||
## 管理 API 基线
|
||||
|
||||
- 管理 API 启用后会对失败认证进行短期限制,避免暴力尝试。
|
||||
- 建议仅在内网使用,并配合独立强密钥。
|
||||
Generated
+20
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"version": "0.14.2",
|
||||
"version": "0.15.0",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "proxycast",
|
||||
"version": "0.14.2",
|
||||
"version": "0.15.0",
|
||||
"dependencies": {
|
||||
"@radix-ui/react-dialog": "^1.1.2",
|
||||
"@radix-ui/react-dropdown-menu": "^2.1.2",
|
||||
@@ -97,6 +97,7 @@
|
||||
"integrity": "sha512-e7jT4DxYvIDLk1ZHmU/m/mB19rex9sv0c2ftBtjSBv+kVM/902eh0fINUzD7UwLLNR+jU585GxUJ8/EBfAM5fw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"@babel/code-frame": "^7.27.1",
|
||||
"@babel/generator": "^7.28.5",
|
||||
@@ -2780,6 +2781,7 @@
|
||||
"integrity": "sha512-LPM2G3Syo1GLzXLGJAKdqoU35XvrWzGJ21/7sgZTUpbkBaOasTj8tjwn6w+hCkqaa1TfJ/w67rJSwYItlJ2mYw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"undici-types": "~6.21.0"
|
||||
}
|
||||
@@ -2797,6 +2799,7 @@
|
||||
"integrity": "sha512-cisd7gxkzjBKU2GgdYrTdtQx1SORymWyaAFhaxQPK9bYO9ot3Y5OikQRvY0VYQtvwjeQnizCINJAenh/V7MK2w==",
|
||||
"devOptional": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"@types/prop-types": "*",
|
||||
"csstype": "^3.2.2"
|
||||
@@ -2808,6 +2811,7 @@
|
||||
"integrity": "sha512-MEe3UeoENYVFXzoXEWsvcpg6ZvlrFNlOQ7EOsvhI3CfAXwzPfO8Qwuxd40nepsYKqyyVQnTdEfv68q91yLcKrQ==",
|
||||
"devOptional": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"peerDependencies": {
|
||||
"@types/react": "^18.0.0"
|
||||
}
|
||||
@@ -2847,6 +2851,7 @@
|
||||
"integrity": "sha512-N9lBGA9o9aqb1hVMc9hzySbhKibHmB+N3IpoShyV6HyQYRGIhlrO5rQgttypi+yEeKsKI4idxC8Jw6gXKD4THA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"@typescript-eslint/scope-manager": "8.49.0",
|
||||
"@typescript-eslint/types": "8.49.0",
|
||||
@@ -3169,6 +3174,7 @@
|
||||
"integrity": "sha512-NZyJarBfL7nWwIq+FDL6Zp/yHEhePMNnnJ0y3qfieCrmNvYct8uvtiV41UvlSe6apAfk0fY1FbWx+NwfmpvtTg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"bin": {
|
||||
"acorn": "bin/acorn"
|
||||
},
|
||||
@@ -3387,6 +3393,7 @@
|
||||
}
|
||||
],
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"baseline-browser-mapping": "^2.9.0",
|
||||
"caniuse-lite": "^1.0.30001759",
|
||||
@@ -3734,6 +3741,7 @@
|
||||
"integrity": "sha512-LEyamqS7W5HB3ujJyvi0HQK/dtVINZvd5mAAp9eT5S/ujByGjiZLCzPcHVzuXbpJDJF/cxwHlfceVUDZ2lnSTw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"@eslint-community/eslint-utils": "^4.8.0",
|
||||
"@eslint-community/regexpp": "^4.12.1",
|
||||
@@ -4368,6 +4376,7 @@
|
||||
"integrity": "sha512-/imKNG4EbWNrVjoNC/1H5/9GFy+tqjGBHCaSsN+P2RnPqjsLmv6UD3Ej+Kj8nBWaRAwyk7kK5ZUc+OEatnTR3A==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"bin": {
|
||||
"jiti": "bin/jiti.js"
|
||||
}
|
||||
@@ -4846,6 +4855,7 @@
|
||||
}
|
||||
],
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"nanoid": "^3.3.11",
|
||||
"picocolors": "^1.1.1",
|
||||
@@ -5068,6 +5078,7 @@
|
||||
"resolved": "https://registry.npmjs.org/react/-/react-18.3.1.tgz",
|
||||
"integrity": "sha512-wS+hAgJShR0KhEvPJArfuPVN1+Hz1t0Y6n5jLrGQbkb4urgPE/0Rve+1kMB1v/oWgHgm4WIcV+i7F2pTVj+2iQ==",
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"loose-envify": "^1.1.0"
|
||||
},
|
||||
@@ -5080,6 +5091,7 @@
|
||||
"resolved": "https://registry.npmjs.org/react-dom/-/react-dom-18.3.1.tgz",
|
||||
"integrity": "sha512-5m4nQKp+rZRb09LNH59GM4BxTh9251/ylbKIbpe7TpGxfJ+9kv6BLkLBXIjjspbgbnIBNqlI23tRnTWT0snUIw==",
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"loose-envify": "^1.1.0",
|
||||
"scheduler": "^0.23.2"
|
||||
@@ -5572,6 +5584,7 @@
|
||||
"integrity": "sha512-5gTmgEY/sqK6gFXLIsQNH19lWb4ebPDLA4SdLP7dsWkIXHWlG66oPuVvXSGFPppYZz8ZDZq0dYYrbHfBCVUb1Q==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"engines": {
|
||||
"node": ">=12"
|
||||
},
|
||||
@@ -5647,6 +5660,7 @@
|
||||
"integrity": "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw==",
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"peer": true,
|
||||
"bin": {
|
||||
"tsc": "bin/tsc",
|
||||
"tsserver": "bin/tsserver"
|
||||
@@ -5759,6 +5773,7 @@
|
||||
"integrity": "sha512-o5a9xKjbtuhY6Bi5S3+HvbRERmouabWbyUcpXXUA1u+GNUKoROi9byOJ8M0nHbHYHkYICiMlqxkg1KkYmm25Sw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"esbuild": "^0.21.3",
|
||||
"postcss": "^8.4.43",
|
||||
@@ -5819,6 +5834,7 @@
|
||||
"integrity": "sha512-E4t7DJ9pESL6E3I8nFjPa4xGUd3PmiWDLsDztS2qXSJWfHtbQnwAWylaBvSNY48I3vr8PTqIZlyK8TE3V3CA4Q==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"@vitest/expect": "4.0.16",
|
||||
"@vitest/mocker": "4.0.16",
|
||||
@@ -6375,6 +6391,7 @@
|
||||
"integrity": "sha512-5gTmgEY/sqK6gFXLIsQNH19lWb4ebPDLA4SdLP7dsWkIXHWlG66oPuVvXSGFPppYZz8ZDZq0dYYrbHfBCVUb1Q==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"engines": {
|
||||
"node": ">=12"
|
||||
},
|
||||
@@ -6388,6 +6405,7 @@
|
||||
"integrity": "sha512-dZwN5L1VlUBewiP6H9s2+B3e3Jg96D0vzN+Ry73sOefebhYr9f94wwkMNN/9ouoU8pV1BqA1d1zGk8928cx0rg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"peer": true,
|
||||
"dependencies": {
|
||||
"esbuild": "^0.27.0",
|
||||
"fdir": "^6.5.0",
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.15.0",
|
||||
"version": "0.17.1",
|
||||
"type": "module",
|
||||
"repository": {
|
||||
"type": "git",
|
||||
|
||||
Generated
+16
-1
@@ -1289,6 +1289,16 @@ dependencies = [
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "fs2"
|
||||
version = "0.4.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9564fc758e15025b46aa6643b1b77d047d1a56a1aea6e01002ac0c7026876213"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"winapi",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "fs_extra"
|
||||
version = "1.3.0"
|
||||
@@ -3367,7 +3377,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast"
|
||||
version = "0.15.1"
|
||||
version = "0.17.1"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"async-stream",
|
||||
@@ -3375,9 +3385,12 @@ dependencies = [
|
||||
"axum",
|
||||
"axum-server",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"chrono",
|
||||
"dashmap",
|
||||
"dirs 5.0.1",
|
||||
"flate2",
|
||||
"fs2",
|
||||
"futures",
|
||||
"indexmap 2.12.1",
|
||||
"md5",
|
||||
@@ -3396,6 +3409,7 @@ dependencies = [
|
||||
"serde_urlencoded",
|
||||
"serde_yaml",
|
||||
"sha2",
|
||||
"subtle",
|
||||
"tauri",
|
||||
"tauri-build",
|
||||
"tauri-plugin-autostart",
|
||||
@@ -3405,6 +3419,7 @@ dependencies = [
|
||||
"thiserror 1.0.69",
|
||||
"tiktoken-rs",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
"tower 0.4.13",
|
||||
"tower-http 0.5.2",
|
||||
"tracing",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "proxycast"
|
||||
version = "0.15.1"
|
||||
version = "0.17.1"
|
||||
description = "AI API Proxy Desktop App"
|
||||
authors = ["you"]
|
||||
edition = "2021"
|
||||
@@ -38,7 +38,10 @@ async-stream = "0.3"
|
||||
regex = "1"
|
||||
md5 = "0.7"
|
||||
urlencoding = "2"
|
||||
rusqlite = { version = "0.31", features = ["bundled"] }
|
||||
subtle = "2.5"
|
||||
flate2 = "1"
|
||||
fs2 = "0.4"
|
||||
rusqlite = { version = "0.31", features = ["bundled", "backup"] }
|
||||
serde_yaml = "0.9"
|
||||
indexmap = { version = "2", features = ["serde"] }
|
||||
zip = "0.6"
|
||||
@@ -50,12 +53,14 @@ tiktoken-rs = "0.6"
|
||||
async-trait = "0.1"
|
||||
thiserror = "1"
|
||||
base64 = "0.22"
|
||||
bytes = "1"
|
||||
rand = "0.8"
|
||||
sha2 = "0.10"
|
||||
serde_urlencoded = "0.7"
|
||||
open = "5"
|
||||
url = "2"
|
||||
once_cell = "1"
|
||||
tokio-util = "0.7"
|
||||
|
||||
[dev-dependencies]
|
||||
proptest = "1"
|
||||
|
||||
@@ -1,3 +1,10 @@
|
||||
fn main() {
|
||||
// tauri::generate_context! 在编译期会校验 `frontendDist` 路径是否存在。
|
||||
// 开发/CI 场景下可能只跑 `cargo check/test` 而未先构建前端,从而导致宏 panic。
|
||||
// 这里提前创建配置中的 `../dist` 目录,避免无关的编译阻塞。
|
||||
if let Ok(manifest_dir) = std::env::var("CARGO_MANIFEST_DIR") {
|
||||
let dist_dir = std::path::PathBuf::from(manifest_dir).join("../dist");
|
||||
let _ = std::fs::create_dir_all(dist_dir);
|
||||
}
|
||||
tauri_build::build()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
# Seeds for failure cases proptest has generated in the past. It is
|
||||
# automatically read and these particular cases re-run before any
|
||||
# novel cases are generated.
|
||||
#
|
||||
# It is recommended to check this file in to source control so that
|
||||
# everyone who runs the test benefits from these saved cases.
|
||||
cc 93bc59c922237bf2c4ffc251c66c90f6bdcf7433a6a188b325fe65f38a8a6fd4 # shrinks to flow_count = 1
|
||||
cc 5f7ef8a17a79bad4599803df567d02ed7c7a29bae560e3e90dfb8b2f920cadf2 # shrinks to flow_count = 5
|
||||
@@ -0,0 +1,7 @@
|
||||
# Seeds for failure cases proptest has generated in the past. It is
|
||||
# automatically read and these particular cases re-run before any
|
||||
# novel cases are generated.
|
||||
#
|
||||
# It is recommended to check this file in to source control so that
|
||||
# everyone who runs the test benefits from these saved cases.
|
||||
cc 4ef289c82d7068ccd05f549e93999e0d88c53e9d4ad4ebbcac1b33d00993d259 # shrinks to prefix = "ot"
|
||||
@@ -0,0 +1,7 @@
|
||||
# Seeds for failure cases proptest has generated in the past. It is
|
||||
# automatically read and these particular cases re-run before any
|
||||
# novel cases are generated.
|
||||
#
|
||||
# It is recommended to check this file in to source control so that
|
||||
# everyone who runs the test benefits from these saved cases.
|
||||
cc a025afdf39438a61f1e2fdf53f6ee71ea4d4add523c76c9d060acfce58b901d6 # shrinks to initial_window = 30, new_window = 10, request_count = 14
|
||||
@@ -0,0 +1,8 @@
|
||||
# Seeds for failure cases proptest has generated in the past. It is
|
||||
# automatically read and these particular cases re-run before any
|
||||
# novel cases are generated.
|
||||
#
|
||||
# It is recommended to check this file in to source control so that
|
||||
# everyone who runs the test benefits from these saved cases.
|
||||
cc be956d5aa14123ea1af5d3e310121b3dec52581c16da860aa2622f1689edd520 # shrinks to tool_call = ("call_00aa00aa", "aa_", "{\"value\":\"aAaA_a\"}")
|
||||
cc f35fbd7673ecfb016ab0ee416718b5561162ae8ac1b04fbffe0fcebd8467a341 # shrinks to tool_call = ("call_a000a0a0", "__a", "{\"value\":\"aaAaaA\"}")
|
||||
@@ -38,7 +38,7 @@ pub fn get_config_status(app_type: String) -> Result<ConfigStatus, String> {
|
||||
AppType::Claude => config_dir.join("settings.json"),
|
||||
AppType::Codex => config_dir.join("auth.json"),
|
||||
AppType::Gemini => config_dir.join(".env"),
|
||||
AppType::ProxyCast => config_dir.join("config.json"),
|
||||
AppType::ProxyCast => config_dir.join("config.yaml"),
|
||||
};
|
||||
|
||||
let has_env = match app {
|
||||
@@ -50,11 +50,20 @@ pub fn get_config_status(app_type: String) -> Result<ConfigStatus, String> {
|
||||
}
|
||||
AppType::Codex => config_dir.join("auth.json").exists(),
|
||||
AppType::Gemini => config_dir.join(".env").exists(),
|
||||
AppType::ProxyCast => config_dir.join("config.json").exists(),
|
||||
AppType::ProxyCast => {
|
||||
config_dir.join("config.yaml").exists() || config_dir.join("config.json").exists()
|
||||
}
|
||||
};
|
||||
|
||||
let exists = match app {
|
||||
AppType::ProxyCast => {
|
||||
config_dir.join("config.yaml").exists() || config_dir.join("config.json").exists()
|
||||
}
|
||||
_ => main_config.exists(),
|
||||
};
|
||||
|
||||
Ok(ConfigStatus {
|
||||
exists: main_config.exists(),
|
||||
exists,
|
||||
path: config_dir.to_string_lossy().to_string(),
|
||||
has_env,
|
||||
})
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -8,6 +8,7 @@ use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
/// 注入配置状态
|
||||
#[allow(dead_code)]
|
||||
pub struct InjectionConfigState(pub Arc<RwLock<InjectionSettings>>);
|
||||
|
||||
/// 注入配置响应
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
pub mod config_cmd;
|
||||
pub mod flow_monitor_cmd;
|
||||
pub mod injection_cmd;
|
||||
pub mod mcp_cmd;
|
||||
pub mod oauth_cmd;
|
||||
@@ -14,3 +15,4 @@ pub mod telemetry_cmd;
|
||||
pub mod tray_cmd;
|
||||
pub mod usage_cmd;
|
||||
pub mod websocket_cmd;
|
||||
pub mod window_cmd;
|
||||
|
||||
@@ -206,11 +206,12 @@ pub async fn refresh_oauth_token(
|
||||
};
|
||||
|
||||
match result {
|
||||
Ok(token) => {
|
||||
Ok(_token) => {
|
||||
logs.write()
|
||||
.await
|
||||
.add("info", &format!("[{display_name}] Token 刷新成功"));
|
||||
Ok(token)
|
||||
// P0 安全修复:不返回明文 token
|
||||
Ok("Token 刷新成功".to_string())
|
||||
}
|
||||
Err(e) => {
|
||||
logs.write()
|
||||
@@ -234,31 +235,32 @@ pub async fn get_oauth_env_variables(
|
||||
match provider_type {
|
||||
OAuthProvider::Kiro => {
|
||||
let creds = &s.kiro_provider.credentials;
|
||||
// P0 安全修复:不返回明文敏感凭证
|
||||
if let Some(token) = &creds.access_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_ACCESS_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(token) = &creds.refresh_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_REFRESH_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(id) = &creds.client_id {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_CLIENT_ID".to_string(),
|
||||
value: id.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(id),
|
||||
});
|
||||
}
|
||||
if let Some(secret) = &creds.client_secret {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_CLIENT_SECRET".to_string(),
|
||||
value: secret.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(secret),
|
||||
});
|
||||
}
|
||||
@@ -286,17 +288,18 @@ pub async fn get_oauth_env_variables(
|
||||
}
|
||||
OAuthProvider::Gemini => {
|
||||
let creds = &s.gemini_provider.credentials;
|
||||
// P0 安全修复:不返回明文敏感凭证
|
||||
if let Some(token) = &creds.access_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "GEMINI_ACCESS_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(token) = &creds.refresh_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "GEMINI_REFRESH_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
@@ -311,17 +314,18 @@ pub async fn get_oauth_env_variables(
|
||||
}
|
||||
OAuthProvider::Qwen => {
|
||||
let creds = &s.qwen_provider.credentials;
|
||||
// P0 安全修复:不返回明文敏感凭证
|
||||
if let Some(token) = &creds.access_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "QWEN_ACCESS_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(token) = &creds.refresh_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "QWEN_REFRESH_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(),
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -741,6 +741,7 @@ pub fn add_codex_oauth_credential(
|
||||
db: State<'_, DbConnection>,
|
||||
pool_service: State<'_, ProviderPoolServiceState>,
|
||||
creds_file_path: String,
|
||||
api_base_url: Option<String>,
|
||||
name: Option<String>,
|
||||
) -> Result<ProviderCredential, String> {
|
||||
// 复制并重命名文件到应用存储目录
|
||||
@@ -751,6 +752,7 @@ pub fn add_codex_oauth_credential(
|
||||
"codex",
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path: stored_file_path,
|
||||
api_base_url,
|
||||
},
|
||||
name,
|
||||
Some(true),
|
||||
@@ -854,6 +856,8 @@ pub fn get_pool_credential_oauth_status(
|
||||
}
|
||||
|
||||
/// 调试 Kiro 凭证加载(从默认路径)
|
||||
/// P0 安全修复:仅在 debug 构建中可用
|
||||
#[cfg(debug_assertions)]
|
||||
#[tauri::command]
|
||||
pub async fn debug_kiro_credentials() -> Result<String, String> {
|
||||
use crate::providers::kiro::KiroProvider;
|
||||
@@ -883,31 +887,15 @@ pub async fn debug_kiro_credentials() -> Result<String, String> {
|
||||
provider.credentials.client_id_hash.is_some()
|
||||
));
|
||||
|
||||
if let Some(hash) = &provider.credentials.client_id_hash {
|
||||
result.push_str(&format!("🔗 clientIdHash: {}\n", hash));
|
||||
}
|
||||
|
||||
// P0 安全修复:不再输出敏感信息(clientIdHash、token 前缀等)
|
||||
let detected_method = provider.detect_auth_method();
|
||||
result.push_str(&format!("🎯 检测到的认证方式: {}\n", detected_method));
|
||||
|
||||
let refresh_url = provider.get_refresh_url();
|
||||
result.push_str(&format!("🌐 刷新端点: {}\n", refresh_url));
|
||||
|
||||
if let Some(client_id) = &provider.credentials.client_id {
|
||||
result.push_str(&format!(
|
||||
"🆔 client_id 前缀: {}...\n",
|
||||
&client_id[..std::cmp::min(20, client_id.len())]
|
||||
));
|
||||
}
|
||||
|
||||
result.push_str("\n🚀 尝试刷新 token...\n");
|
||||
match provider.refresh_token().await {
|
||||
Ok(token) => {
|
||||
result.push_str(&format!("✅ Token 刷新成功! Token 长度: {}\n", token.len()));
|
||||
result.push_str(&format!(
|
||||
"🎫 Token 前缀: {}...\n",
|
||||
&token[..std::cmp::min(50, token.len())]
|
||||
));
|
||||
// 不再输出 token 前缀
|
||||
}
|
||||
Err(e) => {
|
||||
result.push_str(&format!("❌ Token 刷新失败: {}\n", e));
|
||||
@@ -922,7 +910,16 @@ pub async fn debug_kiro_credentials() -> Result<String, String> {
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// P0 安全修复:release 构建中禁用 debug 命令
|
||||
#[cfg(not(debug_assertions))]
|
||||
#[tauri::command]
|
||||
pub async fn debug_kiro_credentials() -> Result<String, String> {
|
||||
Err("此调试命令仅在开发构建中可用".to_string())
|
||||
}
|
||||
|
||||
/// 测试用户上传的凭证文件
|
||||
/// P0 安全修复:仅在 debug 构建中可用,且不输出敏感信息
|
||||
#[cfg(debug_assertions)]
|
||||
#[tauri::command]
|
||||
pub async fn test_user_credentials() -> Result<String, String> {
|
||||
use crate::providers::kiro::KiroProvider;
|
||||
@@ -937,7 +934,8 @@ pub async fn test_user_credentials() -> Result<String, String> {
|
||||
"Library/Application Support/proxycast/credentials/kiro_d8da9d58_1765757992_kiro.json",
|
||||
);
|
||||
|
||||
result.push_str(&format!("📂 用户凭证路径: {}\n", user_creds_path.display()));
|
||||
// P0 安全修复:不输出完整路径,仅显示文件是否存在
|
||||
result.push_str("📂 检查用户凭证文件...\n");
|
||||
|
||||
// 检查文件是否存在
|
||||
if !user_creds_path.exists() {
|
||||
@@ -959,88 +957,27 @@ pub async fn test_user_credentials() -> Result<String, String> {
|
||||
Ok(json) => {
|
||||
result.push_str("✅ JSON 格式有效\n");
|
||||
|
||||
// 检查关键字段
|
||||
// 检查关键字段(仅显示是否存在,不显示值)
|
||||
let has_access_token =
|
||||
json.get("accessToken").and_then(|v| v.as_str()).is_some();
|
||||
let has_refresh_token =
|
||||
json.get("refreshToken").and_then(|v| v.as_str()).is_some();
|
||||
let auth_method = json.get("authMethod").and_then(|v| v.as_str());
|
||||
let client_id_hash = json.get("clientIdHash").and_then(|v| v.as_str());
|
||||
let has_client_id_hash =
|
||||
json.get("clientIdHash").and_then(|v| v.as_str()).is_some();
|
||||
let region = json.get("region").and_then(|v| v.as_str());
|
||||
|
||||
result.push_str(&format!("🔑 有 accessToken: {}\n", has_access_token));
|
||||
result.push_str(&format!("🔄 有 refreshToken: {}\n", has_refresh_token));
|
||||
result.push_str(&format!("📄 authMethod: {:?}\n", auth_method));
|
||||
result.push_str(&format!("🏷️ clientIdHash: {:?}\n", client_id_hash));
|
||||
// P0 安全修复:不输出 clientIdHash 值
|
||||
result.push_str(&format!("🏷️ 有 clientIdHash: {}\n", has_client_id_hash));
|
||||
result.push_str(&format!("🌍 region: {:?}\n", region));
|
||||
|
||||
if let Some(hash) = client_id_hash {
|
||||
// 检查 clientIdHash 对应的文件
|
||||
let hash_file_path = dirs::home_dir()
|
||||
.unwrap()
|
||||
.join(".aws/sso/cache")
|
||||
.join(format!("{}.json", hash));
|
||||
|
||||
result.push_str(&format!(
|
||||
"\n🔗 检查 clientIdHash 文件: {}\n",
|
||||
hash_file_path.display()
|
||||
));
|
||||
|
||||
if hash_file_path.exists() {
|
||||
result.push_str("✅ clientIdHash 文件存在\n");
|
||||
|
||||
match std::fs::read_to_string(&hash_file_path) {
|
||||
Ok(hash_content) => {
|
||||
match serde_json::from_str::<serde_json::Value>(&hash_content) {
|
||||
Ok(hash_json) => {
|
||||
let has_client_id = hash_json
|
||||
.get("clientId")
|
||||
.and_then(|v| v.as_str())
|
||||
.is_some();
|
||||
let has_client_secret = hash_json
|
||||
.get("clientSecret")
|
||||
.and_then(|v| v.as_str())
|
||||
.is_some();
|
||||
|
||||
result.push_str(&format!(
|
||||
"🆔 hash 文件有 clientId: {}\n",
|
||||
has_client_id
|
||||
));
|
||||
result.push_str(&format!(
|
||||
"🔒 hash 文件有 clientSecret: {}\n",
|
||||
has_client_secret
|
||||
));
|
||||
|
||||
if has_client_id && has_client_secret {
|
||||
result.push_str("✅ IdC 认证配置完整!\n");
|
||||
} else {
|
||||
result.push_str(
|
||||
"⚠️ IdC 认证配置不完整,将使用 social 认证\n",
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
result.push_str(&format!(
|
||||
"❌ 无法解析 hash 文件 JSON: {}\n",
|
||||
e
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
result.push_str(&format!("❌ 无法读取 hash 文件: {}\n", e));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
result.push_str("❌ clientIdHash 文件不存在\n");
|
||||
}
|
||||
}
|
||||
|
||||
// 现在使用我们的 KiroProvider 来测试加载
|
||||
// 使用 KiroProvider 测试加载
|
||||
result.push_str("\n🔧 使用 KiroProvider 测试加载...\n");
|
||||
|
||||
let mut provider = KiroProvider::new();
|
||||
// 设置凭证路径到用户文件
|
||||
provider.creds_path = Some(user_creds_path.clone());
|
||||
|
||||
match provider
|
||||
@@ -1065,9 +1002,6 @@ pub async fn test_user_credentials() -> Result<String, String> {
|
||||
let detected_method = provider.detect_auth_method();
|
||||
result.push_str(&format!("🎯 检测到的认证方式: {}\n", detected_method));
|
||||
|
||||
let refresh_url = provider.get_refresh_url();
|
||||
result.push_str(&format!("🌐 刷新端点: {}\n", refresh_url));
|
||||
|
||||
result.push_str("\n🚀 尝试刷新 token...\n");
|
||||
match provider.refresh_token().await {
|
||||
Ok(token) => {
|
||||
@@ -1075,10 +1009,7 @@ pub async fn test_user_credentials() -> Result<String, String> {
|
||||
"✅ Token 刷新成功! Token 长度: {}\n",
|
||||
token.len()
|
||||
));
|
||||
result.push_str(&format!(
|
||||
"🎫 Token 前缀: {}...\n",
|
||||
&token[..std::cmp::min(50, token.len())]
|
||||
));
|
||||
// P0 安全修复:不输出 token 前缀
|
||||
}
|
||||
Err(e) => {
|
||||
result.push_str(&format!("❌ Token 刷新失败: {}\n", e));
|
||||
@@ -1103,6 +1034,13 @@ pub async fn test_user_credentials() -> Result<String, String> {
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// P0 安全修复:release 构建中禁用 test_user_credentials 命令
|
||||
#[cfg(not(debug_assertions))]
|
||||
#[tauri::command]
|
||||
pub async fn test_user_credentials() -> Result<String, String> {
|
||||
Err("此调试命令仅在开发构建中可用".to_string())
|
||||
}
|
||||
|
||||
/// 迁移 Private 配置到凭证池
|
||||
///
|
||||
/// 从 providers 配置中读取单个凭证配置,迁移到凭证池中并标记为 Private 来源
|
||||
@@ -1299,6 +1237,7 @@ pub async fn get_codex_auth_url_and_wait(
|
||||
"codex",
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path: result.creds_file_path,
|
||||
api_base_url: None,
|
||||
},
|
||||
name,
|
||||
Some(true),
|
||||
@@ -1339,6 +1278,7 @@ pub async fn start_codex_oauth_login(
|
||||
"codex",
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path: result.creds_file_path,
|
||||
api_base_url: None,
|
||||
},
|
||||
name,
|
||||
Some(true),
|
||||
|
||||
@@ -161,6 +161,7 @@ pub async fn clear_switch_log(
|
||||
}
|
||||
|
||||
/// 添加切换日志条目(内部使用)
|
||||
#[allow(dead_code)]
|
||||
pub async fn add_switch_log_entry(
|
||||
state: &ResilienceConfigState,
|
||||
from_provider: &str,
|
||||
|
||||
@@ -67,7 +67,8 @@ pub async fn get_route_curl_examples(
|
||||
// 查找匹配的路由
|
||||
let route = routes.iter().find(|r| r.selector == selector);
|
||||
|
||||
let api_key = &config.server.api_key;
|
||||
// P0 安全修复:curl 示例使用占位符,不暴露真实 API Key
|
||||
let api_key = "${PROXYCAST_API_KEY}";
|
||||
|
||||
match route {
|
||||
Some(r) => Ok(r.generate_curl_examples(api_key)),
|
||||
|
||||
@@ -45,6 +45,7 @@ impl From<WsConnection> for WsConnectionInfo {
|
||||
}
|
||||
|
||||
/// WebSocket 状态封装(用于 Tauri State)
|
||||
#[allow(dead_code)]
|
||||
pub struct WsServiceState {
|
||||
pub enabled: Arc<RwLock<bool>>,
|
||||
pub stats: Arc<RwLock<WsStatsSnapshot>>,
|
||||
@@ -73,6 +74,7 @@ impl Default for WsServiceState {
|
||||
}
|
||||
|
||||
/// 获取 WebSocket 服务状态
|
||||
#[allow(dead_code)]
|
||||
#[tauri::command]
|
||||
pub async fn get_websocket_status(
|
||||
state: tauri::State<'_, WsServiceState>,
|
||||
@@ -90,6 +92,7 @@ pub async fn get_websocket_status(
|
||||
}
|
||||
|
||||
/// 获取 WebSocket 连接列表
|
||||
#[allow(dead_code)]
|
||||
#[tauri::command]
|
||||
pub async fn get_websocket_connections(
|
||||
state: tauri::State<'_, WsServiceState>,
|
||||
@@ -99,6 +102,7 @@ pub async fn get_websocket_connections(
|
||||
}
|
||||
|
||||
/// 启用/禁用 WebSocket 服务
|
||||
#[allow(dead_code)]
|
||||
#[tauri::command]
|
||||
pub async fn set_websocket_enabled(
|
||||
state: tauri::State<'_, WsServiceState>,
|
||||
|
||||
@@ -0,0 +1,396 @@
|
||||
//! 窗口控制命令
|
||||
//!
|
||||
//! 提供窗口大小调整、位置控制等功能
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tauri::{AppHandle, Manager, PhysicalSize, Window};
|
||||
|
||||
/// 窗口大小预设
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct WindowSize {
|
||||
pub width: u32,
|
||||
pub height: u32,
|
||||
}
|
||||
|
||||
/// 预定义的窗口大小
|
||||
impl WindowSize {
|
||||
/// 默认窗口大小
|
||||
pub fn default() -> Self {
|
||||
Self {
|
||||
width: 1200,
|
||||
height: 800,
|
||||
}
|
||||
}
|
||||
|
||||
/// Flow Monitor 优化大小(更宽更高,适合数据展示)
|
||||
pub fn flow_monitor() -> Self {
|
||||
Self {
|
||||
width: 1600,
|
||||
height: 1000,
|
||||
}
|
||||
}
|
||||
|
||||
/// 紧凑模式
|
||||
pub fn compact() -> Self {
|
||||
Self {
|
||||
width: 1000,
|
||||
height: 700,
|
||||
}
|
||||
}
|
||||
|
||||
/// 大屏模式
|
||||
pub fn large() -> Self {
|
||||
Self {
|
||||
width: 1920,
|
||||
height: 1200,
|
||||
}
|
||||
}
|
||||
|
||||
/// 超大屏模式
|
||||
pub fn extra_large() -> Self {
|
||||
Self {
|
||||
width: 2560,
|
||||
height: 1440,
|
||||
}
|
||||
}
|
||||
|
||||
/// 4K 模式
|
||||
pub fn ultra_wide() -> Self {
|
||||
Self {
|
||||
width: 3440,
|
||||
height: 1440,
|
||||
}
|
||||
}
|
||||
|
||||
/// 4K 标准模式
|
||||
pub fn four_k() -> Self {
|
||||
Self {
|
||||
width: 3840,
|
||||
height: 2160,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 窗口大小选项
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct WindowSizeOption {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
pub description: String,
|
||||
pub size: WindowSize,
|
||||
}
|
||||
|
||||
impl WindowSizeOption {
|
||||
/// 获取所有可用的窗口大小选项
|
||||
pub fn all_options() -> Vec<Self> {
|
||||
vec![
|
||||
Self {
|
||||
id: "compact".to_string(),
|
||||
name: "紧凑模式".to_string(),
|
||||
description: "1000×700 - 节省屏幕空间".to_string(),
|
||||
size: WindowSize::compact(),
|
||||
},
|
||||
Self {
|
||||
id: "default".to_string(),
|
||||
name: "默认大小".to_string(),
|
||||
description: "1200×800 - 日常使用".to_string(),
|
||||
size: WindowSize::default(),
|
||||
},
|
||||
Self {
|
||||
id: "flow_monitor".to_string(),
|
||||
name: "Flow Monitor".to_string(),
|
||||
description: "1600×1000 - 数据展示优化".to_string(),
|
||||
size: WindowSize::flow_monitor(),
|
||||
},
|
||||
Self {
|
||||
id: "large".to_string(),
|
||||
name: "大屏模式".to_string(),
|
||||
description: "1920×1200 - 大屏幕显示".to_string(),
|
||||
size: WindowSize::large(),
|
||||
},
|
||||
Self {
|
||||
id: "extra_large".to_string(),
|
||||
name: "超大屏模式".to_string(),
|
||||
description: "2560×1440 - 超大屏幕".to_string(),
|
||||
size: WindowSize::extra_large(),
|
||||
},
|
||||
Self {
|
||||
id: "ultra_wide".to_string(),
|
||||
name: "超宽屏模式".to_string(),
|
||||
description: "3440×1440 - 超宽屏显示".to_string(),
|
||||
size: WindowSize::ultra_wide(),
|
||||
},
|
||||
Self {
|
||||
id: "four_k".to_string(),
|
||||
name: "4K 模式".to_string(),
|
||||
description: "3840×2160 - 4K 显示器".to_string(),
|
||||
size: WindowSize::four_k(),
|
||||
},
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取所有可用的窗口大小选项
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Vec<WindowSizeOption>` - 所有可用的窗口大小选项
|
||||
#[tauri::command]
|
||||
pub async fn get_window_size_options() -> Vec<WindowSizeOption> {
|
||||
WindowSizeOption::all_options()
|
||||
}
|
||||
|
||||
/// 设置窗口为指定的预设大小
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
/// * `option_id` - 窗口大小选项 ID
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(WindowSize)` - 成功时返回之前的窗口大小
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn set_window_size_by_option(
|
||||
app: AppHandle,
|
||||
option_id: String,
|
||||
) -> Result<WindowSize, String> {
|
||||
// 获取当前大小
|
||||
let current_size = get_window_size(app.clone()).await?;
|
||||
|
||||
// 查找对应的窗口大小选项
|
||||
let options = WindowSizeOption::all_options();
|
||||
let option = options
|
||||
.iter()
|
||||
.find(|opt| opt.id == option_id)
|
||||
.ok_or_else(|| format!("未找到窗口大小选项: {}", option_id))?;
|
||||
|
||||
// 设置新的窗口大小
|
||||
set_window_size(app.clone(), option.size.clone()).await?;
|
||||
|
||||
// 居中窗口
|
||||
center_window(app).await?;
|
||||
|
||||
Ok(current_size)
|
||||
}
|
||||
|
||||
/// 切换全屏模式
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(bool)` - 成功时返回是否进入了全屏模式
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn toggle_fullscreen(app: AppHandle) -> Result<bool, String> {
|
||||
let window = app.get_webview_window("main").ok_or("无法获取主窗口")?;
|
||||
|
||||
let is_fullscreen = window
|
||||
.is_fullscreen()
|
||||
.map_err(|e| format!("获取全屏状态失败: {}", e))?;
|
||||
|
||||
window
|
||||
.set_fullscreen(!is_fullscreen)
|
||||
.map_err(|e| format!("切换全屏模式失败: {}", e))?;
|
||||
|
||||
Ok(!is_fullscreen)
|
||||
}
|
||||
|
||||
/// 检查是否处于全屏模式
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(bool)` - 成功时返回是否处于全屏模式
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn is_fullscreen(app: AppHandle) -> Result<bool, String> {
|
||||
let window = app.get_webview_window("main").ok_or("无法获取主窗口")?;
|
||||
|
||||
window
|
||||
.is_fullscreen()
|
||||
.map_err(|e| format!("获取全屏状态失败: {}", e))
|
||||
}
|
||||
|
||||
/// 获取当前窗口大小
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(WindowSize)` - 成功时返回当前窗口大小
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn get_window_size(app: AppHandle) -> Result<WindowSize, String> {
|
||||
let window = app.get_webview_window("main").ok_or("无法获取主窗口")?;
|
||||
|
||||
let size = window
|
||||
.inner_size()
|
||||
.map_err(|e| format!("获取窗口大小失败: {}", e))?;
|
||||
|
||||
Ok(WindowSize {
|
||||
width: size.width,
|
||||
height: size.height,
|
||||
})
|
||||
}
|
||||
|
||||
/// 设置窗口大小
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
/// * `size` - 新的窗口大小
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(())` - 成功
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn set_window_size(app: AppHandle, size: WindowSize) -> Result<(), String> {
|
||||
let window = app.get_webview_window("main").ok_or("无法获取主窗口")?;
|
||||
|
||||
let physical_size = PhysicalSize::new(size.width, size.height);
|
||||
|
||||
window
|
||||
.set_size(physical_size)
|
||||
.map_err(|e| format!("设置窗口大小失败: {}", e))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 切换到 Flow Monitor 优化大小
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(WindowSize)` - 成功时返回之前的窗口大小
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn resize_for_flow_monitor(app: AppHandle) -> Result<WindowSize, String> {
|
||||
// 先获取当前大小,用于恢复
|
||||
let current_size = get_window_size(app.clone()).await?;
|
||||
|
||||
// 设置为 Flow Monitor 优化大小
|
||||
let flow_monitor_size = WindowSize::flow_monitor();
|
||||
set_window_size(app, flow_monitor_size).await?;
|
||||
|
||||
Ok(current_size)
|
||||
}
|
||||
|
||||
/// 恢复窗口到指定大小
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
/// * `size` - 要恢复的窗口大小
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(())` - 成功
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn restore_window_size(app: AppHandle, size: WindowSize) -> Result<(), String> {
|
||||
set_window_size(app, size).await
|
||||
}
|
||||
|
||||
/// 切换窗口大小(在默认大小和 Flow Monitor 大小之间切换)
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(bool)` - 成功时返回是否切换到了 Flow Monitor 大小
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn toggle_window_size(app: AppHandle) -> Result<bool, String> {
|
||||
let current_size = get_window_size(app.clone()).await?;
|
||||
let flow_monitor_size = WindowSize::flow_monitor();
|
||||
let default_size = WindowSize::default();
|
||||
|
||||
// 判断当前是否接近 Flow Monitor 大小(允许一些误差)
|
||||
let is_flow_monitor_size = (current_size.width as i32 - flow_monitor_size.width as i32).abs()
|
||||
< 50
|
||||
&& (current_size.height as i32 - flow_monitor_size.height as i32).abs() < 50;
|
||||
|
||||
if is_flow_monitor_size {
|
||||
// 当前是 Flow Monitor 大小,切换到默认大小
|
||||
set_window_size(app, default_size).await?;
|
||||
Ok(false)
|
||||
} else {
|
||||
// 当前不是 Flow Monitor 大小,切换到 Flow Monitor 大小
|
||||
set_window_size(app, flow_monitor_size).await?;
|
||||
Ok(true)
|
||||
}
|
||||
}
|
||||
|
||||
/// 居中窗口
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `app` - Tauri AppHandle
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(())` - 成功
|
||||
/// * `Err(String)` - 失败时返回错误消息
|
||||
#[tauri::command]
|
||||
pub async fn center_window(app: AppHandle) -> Result<(), String> {
|
||||
let window = app.get_webview_window("main").ok_or("无法获取主窗口")?;
|
||||
|
||||
window
|
||||
.center()
|
||||
.map_err(|e| format!("居中窗口失败: {}", e))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_window_size_presets() {
|
||||
let default = WindowSize::default();
|
||||
assert_eq!(default.width, 1200);
|
||||
assert_eq!(default.height, 800);
|
||||
|
||||
let flow_monitor = WindowSize::flow_monitor();
|
||||
assert_eq!(flow_monitor.width, 1600);
|
||||
assert_eq!(flow_monitor.height, 1000);
|
||||
|
||||
let compact = WindowSize::compact();
|
||||
assert_eq!(compact.width, 1000);
|
||||
assert_eq!(compact.height, 700);
|
||||
|
||||
let large = WindowSize::large();
|
||||
assert_eq!(large.width, 1920);
|
||||
assert_eq!(large.height, 1200);
|
||||
|
||||
let extra_large = WindowSize::extra_large();
|
||||
assert_eq!(extra_large.width, 2560);
|
||||
assert_eq!(extra_large.height, 1440);
|
||||
|
||||
let ultra_wide = WindowSize::ultra_wide();
|
||||
assert_eq!(ultra_wide.width, 3440);
|
||||
assert_eq!(ultra_wide.height, 1440);
|
||||
|
||||
let four_k = WindowSize::four_k();
|
||||
assert_eq!(four_k.width, 3840);
|
||||
assert_eq!(four_k.height, 2160);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_window_size_options() {
|
||||
let options = WindowSizeOption::all_options();
|
||||
assert_eq!(options.len(), 7);
|
||||
|
||||
// 验证每个选项都有有效的 ID 和名称
|
||||
for option in &options {
|
||||
assert!(!option.id.is_empty());
|
||||
assert!(!option.name.is_empty());
|
||||
assert!(!option.description.is_empty());
|
||||
assert!(option.size.width > 0);
|
||||
assert!(option.size.height > 0);
|
||||
}
|
||||
|
||||
// 验证特定选项
|
||||
let default_option = options.iter().find(|opt| opt.id == "default").unwrap();
|
||||
assert_eq!(default_option.size.width, 1200);
|
||||
assert_eq!(default_option.size.height, 800);
|
||||
}
|
||||
}
|
||||
@@ -34,6 +34,7 @@ impl Default for ExportOptions {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
impl ExportOptions {
|
||||
/// 创建仅配置导出选项
|
||||
pub fn config_only() -> Self {
|
||||
@@ -95,6 +96,7 @@ pub struct ExportBundle {
|
||||
pub redacted: bool,
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
impl ExportBundle {
|
||||
/// 当前导出格式版本
|
||||
pub const CURRENT_VERSION: &'static str = "1.0";
|
||||
@@ -139,6 +141,7 @@ impl ExportBundle {
|
||||
|
||||
/// 导出错误类型
|
||||
#[derive(Debug, Clone)]
|
||||
#[allow(dead_code)]
|
||||
pub enum ExportError {
|
||||
/// 配置错误
|
||||
ConfigError(String),
|
||||
@@ -180,6 +183,7 @@ pub const REDACTED_PLACEHOLDER: &str = "***REDACTED***";
|
||||
/// 提供配置和凭证的统一导出功能
|
||||
pub struct ExportService;
|
||||
|
||||
#[allow(dead_code)]
|
||||
impl ExportService {
|
||||
/// 导出配置为 YAML 字符串
|
||||
///
|
||||
|
||||
@@ -17,6 +17,7 @@ use tokio::sync::mpsc;
|
||||
|
||||
/// 热重载错误类型
|
||||
#[derive(Debug, Clone)]
|
||||
#[allow(dead_code)]
|
||||
pub enum HotReloadError {
|
||||
/// 文件监控错误
|
||||
WatchError(String),
|
||||
@@ -46,6 +47,7 @@ impl std::error::Error for HotReloadError {}
|
||||
|
||||
/// 热重载结果
|
||||
#[derive(Debug, Clone)]
|
||||
#[allow(dead_code)]
|
||||
pub enum ReloadResult {
|
||||
/// 重载成功
|
||||
Success {
|
||||
@@ -370,6 +372,8 @@ impl HotReloadManager {
|
||||
|
||||
/// 验证配置
|
||||
fn validate_config(&self, config: &Config) -> Result<(), HotReloadError> {
|
||||
let is_localhost = is_localhost_host(&config.server.host);
|
||||
|
||||
// 验证端口范围
|
||||
if config.server.port == 0 {
|
||||
return Err(HotReloadError::ValidationError(
|
||||
@@ -377,6 +381,12 @@ impl HotReloadManager {
|
||||
));
|
||||
}
|
||||
|
||||
if !is_localhost {
|
||||
return Err(HotReloadError::ValidationError(
|
||||
"当前版本仅支持本地监听,请使用 127.0.0.1/localhost/::1".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
// 验证重试配置
|
||||
if config.retry.max_retries > 100 {
|
||||
return Err(HotReloadError::ValidationError(
|
||||
@@ -397,6 +407,32 @@ impl HotReloadManager {
|
||||
));
|
||||
}
|
||||
|
||||
if config.server.api_key.trim().is_empty() {
|
||||
return Err(HotReloadError::ValidationError(
|
||||
"API Key 不能为空".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
if (!is_localhost || config.remote_management.allow_remote)
|
||||
&& crate::config::is_default_api_key(&config.server.api_key)
|
||||
{
|
||||
return Err(HotReloadError::ValidationError(
|
||||
"非本地访问场景下禁止使用默认 API Key,请设置强口令".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
if config.server.tls.enable {
|
||||
return Err(HotReloadError::ValidationError(
|
||||
"当前版本暂不支持 TLS,请关闭 TLS 配置".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
if config.remote_management.allow_remote {
|
||||
return Err(HotReloadError::ValidationError(
|
||||
"当前版本未启用 TLS,禁止开启远程管理".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -437,6 +473,15 @@ impl HotReloadManager {
|
||||
}
|
||||
}
|
||||
|
||||
fn is_localhost_host(host: &str) -> bool {
|
||||
if host == "localhost" {
|
||||
return true;
|
||||
}
|
||||
host.parse::<std::net::IpAddr>()
|
||||
.map(|addr| addr.is_loopback())
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// 热重载状态
|
||||
#[derive(Debug, Clone, serde::Serialize)]
|
||||
pub struct HotReloadStatus {
|
||||
|
||||
@@ -313,7 +313,10 @@ impl ImportService {
|
||||
|
||||
// 如果是脱敏数据,清理凭证池中的占位符
|
||||
if bundle.redacted {
|
||||
Self::clean_redacted_credentials(&mut config);
|
||||
let server_key_cleared = Self::clean_redacted_credentials(&mut config);
|
||||
if server_key_cleared {
|
||||
warnings.push("检测到脱敏的服务器 API Key,已清空,需要手动设置".to_string());
|
||||
}
|
||||
}
|
||||
|
||||
Ok(ImportResult::success_with_warnings(config, warnings))
|
||||
@@ -469,7 +472,9 @@ impl ImportService {
|
||||
/// 清理脱敏的凭证数据
|
||||
///
|
||||
/// 移除凭证池中使用占位符的条目
|
||||
fn clean_redacted_credentials(config: &mut Config) {
|
||||
fn clean_redacted_credentials(config: &mut Config) -> bool {
|
||||
let mut server_key_cleared = false;
|
||||
|
||||
// 清理 OpenAI 凭证池中的脱敏条目
|
||||
config
|
||||
.credential_pool
|
||||
@@ -490,10 +495,13 @@ impl ImportService {
|
||||
config.providers.claude.api_key = None;
|
||||
}
|
||||
|
||||
// 清理服务器 API 密钥(如果是脱敏的,恢复默认值)
|
||||
// 清理服务器 API 密钥(如果是脱敏的,清空并提示手动设置)
|
||||
if config.server.api_key == REDACTED_PLACEHOLDER {
|
||||
config.server.api_key = "proxy_cast".to_string();
|
||||
config.server.api_key = String::new();
|
||||
server_key_cleared = true;
|
||||
}
|
||||
|
||||
server_key_cleared
|
||||
}
|
||||
|
||||
/// 从文件导入配置
|
||||
@@ -638,7 +646,7 @@ server:
|
||||
let current = Config::default();
|
||||
let yaml = r#"
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
host: 127.0.0.1
|
||||
port: 9000
|
||||
api_key: new_key
|
||||
"#;
|
||||
@@ -646,7 +654,7 @@ server:
|
||||
let result = ImportService::import_yaml(yaml, ¤t, &options).expect("导入应成功");
|
||||
|
||||
assert!(result.success);
|
||||
assert_eq!(result.config.server.host, "0.0.0.0");
|
||||
assert_eq!(result.config.server.host, "127.0.0.1");
|
||||
assert_eq!(result.config.server.port, 9000);
|
||||
assert_eq!(result.config.server.api_key, "new_key");
|
||||
}
|
||||
@@ -664,7 +672,7 @@ server:
|
||||
|
||||
let yaml = r#"
|
||||
server:
|
||||
host: 0.0.0.0
|
||||
host: 127.0.0.1
|
||||
port: 9000
|
||||
api_key: new_key
|
||||
credential_pool:
|
||||
@@ -677,7 +685,7 @@ credential_pool:
|
||||
|
||||
assert!(result.success);
|
||||
// 服务器配置应被更新
|
||||
assert_eq!(result.config.server.host, "0.0.0.0");
|
||||
assert_eq!(result.config.server.host, "127.0.0.1");
|
||||
// 凭证池应合并
|
||||
assert_eq!(result.config.credential_pool.openai.len(), 2);
|
||||
}
|
||||
@@ -757,10 +765,11 @@ credential_pool:
|
||||
proxy_url: None,
|
||||
});
|
||||
|
||||
ImportService::clean_redacted_credentials(&mut config);
|
||||
let server_key_cleared = ImportService::clean_redacted_credentials(&mut config);
|
||||
|
||||
// 服务器 API 密钥应恢复默认值
|
||||
assert_eq!(config.server.api_key, "proxy_cast");
|
||||
// 服务器 API 密钥应被清空并提示手动设置
|
||||
assert!(server_key_cleared);
|
||||
assert_eq!(config.server.api_key, "");
|
||||
// Provider API 密钥应被清除
|
||||
assert!(config.providers.openai.api_key.is_none());
|
||||
// 凭证池中脱敏的条目应被移除
|
||||
|
||||
@@ -15,8 +15,8 @@ use tempfile::NamedTempFile;
|
||||
fn arb_host() -> impl Strategy<Value = String> {
|
||||
prop_oneof![
|
||||
Just("127.0.0.1".to_string()),
|
||||
Just("0.0.0.0".to_string()),
|
||||
Just("localhost".to_string()),
|
||||
Just("::1".to_string()),
|
||||
"[0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}".prop_map(|s| s),
|
||||
]
|
||||
}
|
||||
@@ -2120,40 +2120,36 @@ proptest! {
|
||||
/// **Validates: Requirements 5.5**
|
||||
#[test]
|
||||
fn prop_export_import_redacted_loses_secrets(config in arb_config_with_secrets()) {
|
||||
// 导出为 YAML(脱敏)
|
||||
let yaml = ExportService::export_yaml(&config, true)
|
||||
// 导出为脱敏 bundle
|
||||
let options = ExportOptions::redacted();
|
||||
let bundle = ExportService::export(&config, &options, "1.0.0")
|
||||
.expect("导出应成功");
|
||||
|
||||
// 导入 YAML
|
||||
// 导入 bundle(脱敏数据会触发清理)
|
||||
let empty_config = Config::default();
|
||||
let options = ImportOptions::replace();
|
||||
let result = ImportService::import_yaml(&yaml, &empty_config, &options)
|
||||
let import_options = ImportOptions::replace();
|
||||
let result = ImportService::import(
|
||||
&bundle,
|
||||
&empty_config,
|
||||
&import_options,
|
||||
&config.auth_dir,
|
||||
)
|
||||
.expect("导入应成功");
|
||||
|
||||
// 清理脱敏数据
|
||||
let mut imported = result.config;
|
||||
ImportService::import(
|
||||
&ExportBundle::new("1.0.0"),
|
||||
&imported,
|
||||
&ImportOptions::merge(),
|
||||
&config.auth_dir,
|
||||
).ok(); // 忽略结果,只是为了触发清理
|
||||
|
||||
// 验证脱敏后的配置不包含原始敏感信息
|
||||
// 服务器 API 密钥应为脱敏占位符或默认值
|
||||
prop_assert!(
|
||||
imported.server.api_key == REDACTED_PLACEHOLDER ||
|
||||
imported.server.api_key == "proxy_cast",
|
||||
"脱敏后服务器 API 密钥应为占位符或默认值: {}",
|
||||
imported.server.api_key
|
||||
// 服务器 API 密钥应被清空
|
||||
prop_assert_eq!(
|
||||
result.config.server.api_key,
|
||||
"",
|
||||
"脱敏后服务器 API 密钥应被清空"
|
||||
);
|
||||
|
||||
// 如果原始配置有 OpenAI API 密钥,导入后应为脱敏占位符
|
||||
if config.providers.openai.api_key.is_some() {
|
||||
prop_assert_eq!(
|
||||
imported.providers.openai.api_key,
|
||||
Some(REDACTED_PLACEHOLDER.to_string()),
|
||||
"脱敏后 OpenAI API 密钥应为占位符"
|
||||
result.config.providers.openai.api_key,
|
||||
None,
|
||||
"脱敏后 OpenAI API 密钥应被清空"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -334,8 +334,28 @@ fn default_port() -> u16 {
|
||||
8999
|
||||
}
|
||||
|
||||
pub const DEFAULT_API_KEY: &str = "proxy_cast";
|
||||
|
||||
fn default_api_key() -> String {
|
||||
"proxy_cast".to_string()
|
||||
DEFAULT_API_KEY.to_string()
|
||||
}
|
||||
|
||||
/// 生成安全 API Key(32 字节随机)
|
||||
pub fn generate_secure_api_key() -> String {
|
||||
use rand::distributions::Alphanumeric;
|
||||
use rand::Rng;
|
||||
|
||||
let token: String = rand::thread_rng()
|
||||
.sample_iter(&Alphanumeric)
|
||||
.take(32)
|
||||
.map(char::from)
|
||||
.collect();
|
||||
format!("pc_{token}")
|
||||
}
|
||||
|
||||
/// 是否为默认 API Key
|
||||
pub fn is_default_api_key(api_key: &str) -> bool {
|
||||
api_key == DEFAULT_API_KEY
|
||||
}
|
||||
|
||||
impl Default for ServerConfig {
|
||||
|
||||
@@ -104,6 +104,10 @@ impl ConfigManager {
|
||||
std::fs::create_dir_all(parent).map_err(|e| ConfigError::WriteError(e.to_string()))?;
|
||||
}
|
||||
|
||||
if path.exists() {
|
||||
let backup_path = path.with_extension("yaml.backup");
|
||||
let _ = std::fs::copy(path, backup_path);
|
||||
}
|
||||
let yaml = Self::to_yaml(&self.config)?;
|
||||
std::fs::write(path, yaml).map_err(|e| ConfigError::WriteError(e.to_string()))
|
||||
}
|
||||
@@ -656,30 +660,65 @@ fn json_config_path() -> std::path::PathBuf {
|
||||
/// 加载配置(向后兼容)
|
||||
///
|
||||
/// 优先加载 YAML 配置,如果不存在则尝试加载 JSON 配置
|
||||
/// 首次启动时自动生成强随机 API Key 并保存配置
|
||||
pub fn load_config() -> Result<Config, Box<dyn std::error::Error>> {
|
||||
use super::types::{generate_secure_api_key, is_default_api_key};
|
||||
|
||||
let yaml_path = ConfigManager::default_config_path();
|
||||
let json_path = json_config_path();
|
||||
|
||||
// 优先尝试 YAML 配置
|
||||
if yaml_path.exists() {
|
||||
let content = std::fs::read_to_string(&yaml_path)?;
|
||||
let config = serde_yaml::from_str(&content)?;
|
||||
let mut config: Config = serde_yaml::from_str(&content)?;
|
||||
// 如果配置中使用默认 API Key,生成强随机 Key 并保存
|
||||
if is_default_api_key(&config.server.api_key) {
|
||||
let new_key = generate_secure_api_key();
|
||||
tracing::warn!("[CONFIG] 检测到默认 API Key,已自动生成强随机 Key");
|
||||
config.server.api_key = new_key;
|
||||
// 保存更新后的配置
|
||||
if let Err(e) = save_config_yaml(&config) {
|
||||
tracing::error!("[CONFIG] 保存配置失败: {}", e);
|
||||
}
|
||||
}
|
||||
return Ok(config);
|
||||
}
|
||||
|
||||
// 回退到 JSON 配置
|
||||
if json_path.exists() {
|
||||
let content = std::fs::read_to_string(&json_path)?;
|
||||
let config = serde_json::from_str(&content)?;
|
||||
let mut config: Config = serde_json::from_str(&content)?;
|
||||
// 如果配置中使用默认 API Key,生成强随机 Key 并保存
|
||||
if is_default_api_key(&config.server.api_key) {
|
||||
let new_key = generate_secure_api_key();
|
||||
tracing::warn!("[CONFIG] 检测到默认 API Key,已自动生成强随机 Key");
|
||||
config.server.api_key = new_key;
|
||||
// 保存更新后的配置(迁移到 YAML)
|
||||
if let Err(e) = save_config_yaml(&config) {
|
||||
tracing::error!("[CONFIG] 保存配置失败: {}", e);
|
||||
}
|
||||
}
|
||||
return Ok(config);
|
||||
}
|
||||
|
||||
// 都不存在,返回默认配置
|
||||
Ok(Config::default())
|
||||
// 都不存在,创建默认配置并生成强随机 API Key
|
||||
let mut config = Config::default();
|
||||
let new_key = generate_secure_api_key();
|
||||
tracing::info!("[CONFIG] 首次启动,已生成强随机 API Key");
|
||||
config.server.api_key = new_key;
|
||||
// 保存初始配置
|
||||
if let Err(e) = save_config_yaml(&config) {
|
||||
tracing::error!("[CONFIG] 保存初始配置失败: {}", e);
|
||||
}
|
||||
Ok(config)
|
||||
}
|
||||
|
||||
/// 保存配置(向后兼容,使用 JSON 格式)
|
||||
/// 保存配置(同时写入 YAML 与 JSON,兼容旧版)
|
||||
pub fn save_config(config: &Config) -> Result<(), Box<dyn std::error::Error>> {
|
||||
// 主配置优先写入 YAML
|
||||
save_config_yaml(config)?;
|
||||
|
||||
// 兼容旧版 JSON 配置
|
||||
let path = json_config_path();
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)?;
|
||||
@@ -695,6 +734,10 @@ pub fn save_config_yaml(config: &Config) -> Result<(), Box<dyn std::error::Error
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)?;
|
||||
}
|
||||
if path.exists() {
|
||||
let backup_path = path.with_extension("yaml.backup");
|
||||
let _ = std::fs::copy(&path, &backup_path);
|
||||
}
|
||||
let content = serde_yaml::to_string(config)?;
|
||||
std::fs::write(&path, content)?;
|
||||
Ok(())
|
||||
@@ -708,7 +751,7 @@ mod unit_tests {
|
||||
fn test_parse_yaml_minimal() {
|
||||
let yaml = r#"
|
||||
server:
|
||||
host: "0.0.0.0"
|
||||
host: "127.0.0.1"
|
||||
port: 9000
|
||||
api_key: "test-key"
|
||||
providers:
|
||||
@@ -716,7 +759,7 @@ providers:
|
||||
enabled: true
|
||||
"#;
|
||||
let config = ConfigManager::parse_yaml(yaml).unwrap();
|
||||
assert_eq!(config.server.host, "0.0.0.0");
|
||||
assert_eq!(config.server.host, "127.0.0.1");
|
||||
assert_eq!(config.server.port, 9000);
|
||||
assert_eq!(config.server.api_key, "test-key");
|
||||
assert!(config.providers.kiro.enabled);
|
||||
|
||||
@@ -326,8 +326,10 @@ pub fn convert_openai_to_codewhisperer(
|
||||
CWTool {
|
||||
tool_specification: ToolSpecification {
|
||||
name: t.function.name.clone(),
|
||||
// P1 安全修复:使用字符边界安全的截断,防止 UTF-8 panic
|
||||
description: if desc.len() > 500 {
|
||||
format!("{}...", &desc[..497])
|
||||
let truncated: String = desc.chars().take(497).collect();
|
||||
format!("{}...", truncated)
|
||||
} else {
|
||||
desc
|
||||
},
|
||||
|
||||
@@ -1,11 +1,8 @@
|
||||
use rusqlite::Connection;
|
||||
use serde_json::Value;
|
||||
|
||||
/// 从旧的 JSON 配置迁移数据到 SQLite
|
||||
#[allow(dead_code)]
|
||||
pub fn migrate_from_json(
|
||||
conn: &Connection,
|
||||
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
pub fn migrate_from_json(conn: &Connection) -> Result<(), String> {
|
||||
// 检查是否已经迁移过
|
||||
let migrated: bool = conn
|
||||
.query_row(
|
||||
@@ -20,27 +17,30 @@ pub fn migrate_from_json(
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// 读取旧配置文件
|
||||
let home = dirs::home_dir().ok_or("Cannot find home directory")?;
|
||||
// 读取旧配置文件(历史路径)
|
||||
let home = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?;
|
||||
let config_path = home.join(".proxycast").join("config.json");
|
||||
|
||||
if config_path.exists() {
|
||||
let content = std::fs::read_to_string(&config_path)?;
|
||||
let _config: Value = serde_json::from_str(&content)?;
|
||||
// 备份旧配置,避免误覆盖
|
||||
let backup_path = config_path.with_file_name("config.json.backup");
|
||||
if !backup_path.exists() {
|
||||
std::fs::copy(&config_path, &backup_path)
|
||||
.map_err(|e| format!("备份旧配置失败: {}", e))?;
|
||||
}
|
||||
|
||||
// TODO: 解析旧配置并插入到数据库
|
||||
// 这里需要根据实际的旧配置格式来实现
|
||||
|
||||
// 备份旧配置
|
||||
let backup_path = home.join(".proxycast").join("config.json.backup");
|
||||
std::fs::copy(&config_path, &backup_path)?;
|
||||
return Err(
|
||||
"检测到旧版 config.json(~/.proxycast/config.json),当前版本尚未支持自动迁移。请手动导出/重建配置后再启动。"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
// 标记迁移完成
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_from_json', 'true')",
|
||||
[],
|
||||
)?;
|
||||
)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -9,20 +9,22 @@ use std::sync::{Arc, Mutex};
|
||||
pub type DbConnection = Arc<Mutex<Connection>>;
|
||||
|
||||
/// 获取数据库文件路径
|
||||
pub fn get_db_path() -> PathBuf {
|
||||
let home = dirs::home_dir().expect("Cannot find home directory");
|
||||
pub fn get_db_path() -> Result<PathBuf, String> {
|
||||
let home = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?;
|
||||
let db_dir = home.join(".proxycast");
|
||||
std::fs::create_dir_all(&db_dir).expect("Cannot create .proxycast directory");
|
||||
db_dir.join("proxycast.db")
|
||||
std::fs::create_dir_all(&db_dir)
|
||||
.map_err(|e| format!("无法创建数据库目录 {:?}: {}", db_dir, e))?;
|
||||
Ok(db_dir.join("proxycast.db"))
|
||||
}
|
||||
|
||||
/// 初始化数据库连接
|
||||
pub fn init_database() -> Result<DbConnection, rusqlite::Error> {
|
||||
let db_path = get_db_path();
|
||||
let conn = Connection::open(&db_path)?;
|
||||
pub fn init_database() -> Result<DbConnection, String> {
|
||||
let db_path = get_db_path()?;
|
||||
let conn = Connection::open(&db_path).map_err(|e| e.to_string())?;
|
||||
|
||||
// 创建表结构
|
||||
schema::create_tables(&conn)?;
|
||||
schema::create_tables(&conn).map_err(|e| e.to_string())?;
|
||||
migration::migrate_from_json(&conn)?;
|
||||
|
||||
Ok(Arc::new(Mutex::new(conn)))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,570 @@
|
||||
//! 批量操作服务
|
||||
//!
|
||||
//! 该模块实现 Flow 批量操作功能,支持对多个 Flow 进行批量收藏、
|
||||
//! 添加标签、导出、删除等操作。
|
||||
//!
|
||||
//! **Validates: Requirements 11.2-11.6**
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
use thiserror::Error;
|
||||
|
||||
use super::exporter::{ExportFormat, ExportOptions, FlowExporter};
|
||||
use super::models::LLMFlow;
|
||||
use super::monitor::FlowMonitor;
|
||||
use super::session::SessionManager;
|
||||
|
||||
/// 批量操作错误
|
||||
#[derive(Debug, Error)]
|
||||
pub enum BatchOpsError {
|
||||
#[error("Flow 不存在: {0}")]
|
||||
FlowNotFound(String),
|
||||
#[error("会话不存在: {0}")]
|
||||
SessionNotFound(String),
|
||||
#[error("导出错误: {0}")]
|
||||
ExportError(String),
|
||||
#[error("操作失败: {0}")]
|
||||
OperationFailed(String),
|
||||
}
|
||||
|
||||
pub type Result<T> = std::result::Result<T, BatchOpsError>;
|
||||
|
||||
/// 批量操作类型
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum BatchOperation {
|
||||
Star,
|
||||
Unstar,
|
||||
AddTags { tags: Vec<String> },
|
||||
RemoveTags { tags: Vec<String> },
|
||||
Export { format: ExportFormat },
|
||||
Delete,
|
||||
AddToSession { session_id: String },
|
||||
}
|
||||
|
||||
/// 批量操作结果
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
pub struct BatchResult {
|
||||
pub total: usize,
|
||||
pub success: usize,
|
||||
pub failed: usize,
|
||||
pub errors: Vec<(String, String)>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub export_data: Option<String>,
|
||||
}
|
||||
|
||||
impl BatchResult {
|
||||
pub fn new(total: usize) -> Self {
|
||||
Self {
|
||||
total,
|
||||
success: 0,
|
||||
failed: 0,
|
||||
errors: Vec::new(),
|
||||
export_data: None,
|
||||
}
|
||||
}
|
||||
pub fn record_success(&mut self) {
|
||||
self.success += 1;
|
||||
}
|
||||
pub fn record_failure(&mut self, flow_id: impl Into<String>, error: impl Into<String>) {
|
||||
self.failed += 1;
|
||||
self.errors.push((flow_id.into(), error.into()));
|
||||
}
|
||||
pub fn is_all_success(&self) -> bool {
|
||||
self.failed == 0
|
||||
}
|
||||
pub fn is_all_failed(&self) -> bool {
|
||||
self.success == 0 && self.total > 0
|
||||
}
|
||||
pub fn is_partial_success(&self) -> bool {
|
||||
self.success > 0 && self.failed > 0
|
||||
}
|
||||
}
|
||||
|
||||
/// 批量操作服务
|
||||
pub struct BatchOperations {
|
||||
flow_monitor: Arc<FlowMonitor>,
|
||||
session_manager: Option<Arc<SessionManager>>,
|
||||
}
|
||||
|
||||
impl BatchOperations {
|
||||
pub fn new(
|
||||
flow_monitor: Arc<FlowMonitor>,
|
||||
session_manager: Option<Arc<SessionManager>>,
|
||||
) -> Self {
|
||||
Self {
|
||||
flow_monitor,
|
||||
session_manager,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(&self, flow_ids: &[String], operation: BatchOperation) -> BatchResult {
|
||||
self.execute_with_progress(flow_ids, operation, |_, _| {})
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn execute_with_progress<F>(
|
||||
&self,
|
||||
flow_ids: &[String],
|
||||
operation: BatchOperation,
|
||||
progress: F,
|
||||
) -> BatchResult
|
||||
where
|
||||
F: Fn(usize, usize) + Send + Sync,
|
||||
{
|
||||
let mut result = BatchResult::new(flow_ids.len());
|
||||
match operation {
|
||||
BatchOperation::Star => {
|
||||
self.batch_star(flow_ids, true, &mut result, &progress)
|
||||
.await
|
||||
}
|
||||
BatchOperation::Unstar => {
|
||||
self.batch_star(flow_ids, false, &mut result, &progress)
|
||||
.await
|
||||
}
|
||||
BatchOperation::AddTags { tags } => {
|
||||
self.batch_add_tags(flow_ids, &tags, &mut result, &progress)
|
||||
.await
|
||||
}
|
||||
BatchOperation::RemoveTags { tags } => {
|
||||
self.batch_remove_tags(flow_ids, &tags, &mut result, &progress)
|
||||
.await
|
||||
}
|
||||
BatchOperation::Export { format } => {
|
||||
self.batch_export(flow_ids, format, &mut result, &progress)
|
||||
.await
|
||||
}
|
||||
BatchOperation::Delete => self.batch_delete(flow_ids, &mut result, &progress).await,
|
||||
BatchOperation::AddToSession { session_id } => {
|
||||
self.batch_add_to_session(flow_ids, &session_id, &mut result, &progress)
|
||||
.await
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
async fn batch_star<F>(
|
||||
&self,
|
||||
flow_ids: &[String],
|
||||
starred: bool,
|
||||
result: &mut BatchResult,
|
||||
progress: &F,
|
||||
) where
|
||||
F: Fn(usize, usize),
|
||||
{
|
||||
let total = flow_ids.len();
|
||||
for (i, flow_id) in flow_ids.iter().enumerate() {
|
||||
progress(i + 1, total);
|
||||
let memory_store = self.flow_monitor.memory_store();
|
||||
let store = memory_store.read().await;
|
||||
let current_starred = store
|
||||
.get(flow_id)
|
||||
.and_then(|f| f.read().ok().map(|flow| flow.annotations.starred));
|
||||
drop(store);
|
||||
match current_starred {
|
||||
Some(current) if current != starred => {
|
||||
if self.flow_monitor.toggle_starred(flow_id).await {
|
||||
result.record_success();
|
||||
} else {
|
||||
result.record_failure(flow_id, "更新收藏状态失败");
|
||||
}
|
||||
}
|
||||
Some(_) => {
|
||||
result.record_success();
|
||||
}
|
||||
None => {
|
||||
result.record_failure(flow_id, "Flow 不存在");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn batch_add_tags<F>(
|
||||
&self,
|
||||
flow_ids: &[String],
|
||||
tags: &[String],
|
||||
result: &mut BatchResult,
|
||||
progress: &F,
|
||||
) where
|
||||
F: Fn(usize, usize),
|
||||
{
|
||||
let total = flow_ids.len();
|
||||
for (i, flow_id) in flow_ids.iter().enumerate() {
|
||||
progress(i + 1, total);
|
||||
let memory_store = self.flow_monitor.memory_store();
|
||||
let store = memory_store.read().await;
|
||||
let exists = store.get(flow_id).is_some();
|
||||
drop(store);
|
||||
if !exists {
|
||||
result.record_failure(flow_id, "Flow 不存在");
|
||||
continue;
|
||||
}
|
||||
let mut all_success = true;
|
||||
for tag in tags {
|
||||
if !self.flow_monitor.add_tag(flow_id, tag.clone()).await {
|
||||
all_success = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if all_success {
|
||||
result.record_success();
|
||||
} else {
|
||||
result.record_failure(flow_id, "添加标签失败");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn batch_remove_tags<F>(
|
||||
&self,
|
||||
flow_ids: &[String],
|
||||
tags: &[String],
|
||||
result: &mut BatchResult,
|
||||
progress: &F,
|
||||
) where
|
||||
F: Fn(usize, usize),
|
||||
{
|
||||
let total = flow_ids.len();
|
||||
for (i, flow_id) in flow_ids.iter().enumerate() {
|
||||
progress(i + 1, total);
|
||||
let memory_store = self.flow_monitor.memory_store();
|
||||
let store = memory_store.read().await;
|
||||
let exists = store.get(flow_id).is_some();
|
||||
drop(store);
|
||||
if !exists {
|
||||
result.record_failure(flow_id, "Flow 不存在");
|
||||
continue;
|
||||
}
|
||||
for tag in tags {
|
||||
let _ = self.flow_monitor.remove_tag(flow_id, tag).await;
|
||||
}
|
||||
result.record_success();
|
||||
}
|
||||
}
|
||||
|
||||
async fn batch_export<F>(
|
||||
&self,
|
||||
flow_ids: &[String],
|
||||
format: ExportFormat,
|
||||
result: &mut BatchResult,
|
||||
progress: &F,
|
||||
) where
|
||||
F: Fn(usize, usize),
|
||||
{
|
||||
let total = flow_ids.len();
|
||||
let mut flows: Vec<LLMFlow> = Vec::with_capacity(total);
|
||||
for (i, flow_id) in flow_ids.iter().enumerate() {
|
||||
progress(i + 1, total);
|
||||
let memory_store = self.flow_monitor.memory_store();
|
||||
let store = memory_store.read().await;
|
||||
if let Some(flow_lock) = store.get(flow_id) {
|
||||
if let Ok(flow) = flow_lock.read() {
|
||||
flows.push(flow.clone());
|
||||
result.record_success();
|
||||
} else {
|
||||
result.record_failure(flow_id, "无法读取 Flow");
|
||||
}
|
||||
} else {
|
||||
result.record_failure(flow_id, "Flow 不存在");
|
||||
}
|
||||
}
|
||||
if !flows.is_empty() {
|
||||
let options = ExportOptions {
|
||||
format,
|
||||
..Default::default()
|
||||
};
|
||||
let exporter = FlowExporter::new(options);
|
||||
let export_result = exporter.export(&flows);
|
||||
result.export_data = Some(export_result.to_string_pretty());
|
||||
}
|
||||
}
|
||||
|
||||
async fn batch_delete<F>(&self, flow_ids: &[String], result: &mut BatchResult, progress: &F)
|
||||
where
|
||||
F: Fn(usize, usize),
|
||||
{
|
||||
let total = flow_ids.len();
|
||||
for (i, flow_id) in flow_ids.iter().enumerate() {
|
||||
progress(i + 1, total);
|
||||
let memory_store = self.flow_monitor.memory_store();
|
||||
let mut store = memory_store.write().await;
|
||||
if store.remove(flow_id) {
|
||||
result.record_success();
|
||||
} else {
|
||||
result.record_failure(flow_id, "Flow 不存在或删除失败");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn batch_add_to_session<F>(
|
||||
&self,
|
||||
flow_ids: &[String],
|
||||
session_id: &str,
|
||||
result: &mut BatchResult,
|
||||
progress: &F,
|
||||
) where
|
||||
F: Fn(usize, usize),
|
||||
{
|
||||
let total = flow_ids.len();
|
||||
let session_manager = match &self.session_manager {
|
||||
Some(sm) => sm,
|
||||
None => {
|
||||
for flow_id in flow_ids {
|
||||
result.record_failure(flow_id, "会话管理器不可用");
|
||||
}
|
||||
return;
|
||||
}
|
||||
};
|
||||
match session_manager.get_session(session_id) {
|
||||
Ok(Some(_)) => {}
|
||||
Ok(None) => {
|
||||
for flow_id in flow_ids {
|
||||
result.record_failure(flow_id, format!("会话不存在: {}", session_id));
|
||||
}
|
||||
return;
|
||||
}
|
||||
Err(e) => {
|
||||
for flow_id in flow_ids {
|
||||
result.record_failure(flow_id, format!("查询会话失败: {}", e));
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
for (i, flow_id) in flow_ids.iter().enumerate() {
|
||||
progress(i + 1, total);
|
||||
let memory_store = self.flow_monitor.memory_store();
|
||||
let store = memory_store.read().await;
|
||||
let exists = store.get(flow_id).is_some();
|
||||
drop(store);
|
||||
if !exists {
|
||||
result.record_failure(flow_id, "Flow 不存在");
|
||||
continue;
|
||||
}
|
||||
match session_manager.add_flow(session_id, flow_id) {
|
||||
Ok(_) => {
|
||||
result.record_success();
|
||||
}
|
||||
Err(e) => {
|
||||
result.record_failure(flow_id, format!("添加到会话失败: {}", e));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 属性测试
|
||||
// ============================================================================
|
||||
|
||||
#[cfg(test)]
|
||||
mod property_tests {
|
||||
use super::*;
|
||||
use crate::flow_monitor::models::{FlowMetadata, FlowType, LLMRequest};
|
||||
use crate::flow_monitor::monitor::FlowMonitorConfig;
|
||||
use proptest::prelude::*;
|
||||
|
||||
fn create_test_flow_monitor() -> Arc<FlowMonitor> {
|
||||
let config = FlowMonitorConfig::default();
|
||||
Arc::new(FlowMonitor::new(config, None))
|
||||
}
|
||||
|
||||
async fn create_test_flow(monitor: &FlowMonitor, flow_id: &str) -> String {
|
||||
let request = LLMRequest {
|
||||
method: "POST".to_string(),
|
||||
path: "/v1/chat/completions".to_string(),
|
||||
model: "gpt-4".to_string(),
|
||||
..Default::default()
|
||||
};
|
||||
let metadata = FlowMetadata::default();
|
||||
let mut flow = crate::flow_monitor::models::LLMFlow::new(
|
||||
flow_id.to_string(),
|
||||
FlowType::ChatCompletions,
|
||||
request,
|
||||
metadata,
|
||||
);
|
||||
flow.state = crate::flow_monitor::models::FlowState::Completed;
|
||||
let store = monitor.memory_store();
|
||||
store.write().await.add(flow);
|
||||
flow_id.to_string()
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// **Feature: flow-monitor-enhancement, Property 20: 批量操作正确性**
|
||||
/// **Validates: Requirements 11.2-11.6**
|
||||
///
|
||||
/// *对于任意* Flow 集合和批量操作,操作后所有 Flow 应该被正确更新。
|
||||
#[test]
|
||||
fn prop_batch_star_correctness(flow_count in 1usize..10usize) {
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
rt.block_on(async {
|
||||
let monitor = create_test_flow_monitor();
|
||||
let batch_ops = BatchOperations::new(monitor.clone(), None);
|
||||
|
||||
// 创建测试 Flow
|
||||
let mut flow_ids = Vec::new();
|
||||
for i in 0..flow_count {
|
||||
let id = create_test_flow(&monitor, &format!("flow-{}", i)).await;
|
||||
flow_ids.push(id);
|
||||
}
|
||||
|
||||
// 执行批量收藏
|
||||
let result = batch_ops.execute(&flow_ids, BatchOperation::Star).await;
|
||||
|
||||
// 验证结果
|
||||
prop_assert_eq!(result.total, flow_count);
|
||||
prop_assert_eq!(result.success, flow_count);
|
||||
prop_assert_eq!(result.failed, 0);
|
||||
|
||||
// 验证所有 Flow 都被收藏
|
||||
let store = monitor.memory_store();
|
||||
let s = store.read().await;
|
||||
for flow_id in &flow_ids {
|
||||
if let Some(flow_lock) = s.get(flow_id) {
|
||||
let flow = flow_lock.read().unwrap();
|
||||
prop_assert!(flow.annotations.starred, "Flow {} 应该被收藏", flow_id);
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
})?;
|
||||
}
|
||||
|
||||
/// **Feature: flow-monitor-enhancement, Property 21: 批量操作原子性**
|
||||
/// **Validates: Requirements 11.2-11.6**
|
||||
///
|
||||
/// *对于任意* 批量操作,如果部分失败,成功的部分应该被正确应用,失败的部分应该被正确报告。
|
||||
#[test]
|
||||
fn prop_batch_operation_atomicity(
|
||||
valid_flow_count in 1usize..8usize,
|
||||
invalid_flow_count in 1usize..5usize,
|
||||
) {
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
rt.block_on(async {
|
||||
let monitor = create_test_flow_monitor();
|
||||
let batch_ops = BatchOperations::new(monitor.clone(), None);
|
||||
|
||||
// 创建有效的 Flow
|
||||
let mut valid_flow_ids = Vec::new();
|
||||
for i in 0..valid_flow_count {
|
||||
let id = create_test_flow(&monitor, &format!("valid-flow-{}", i)).await;
|
||||
valid_flow_ids.push(id);
|
||||
}
|
||||
|
||||
// 创建无效的 Flow ID(不存在的)
|
||||
let mut invalid_flow_ids = Vec::new();
|
||||
for i in 0..invalid_flow_count {
|
||||
invalid_flow_ids.push(format!("invalid-flow-{}", i));
|
||||
}
|
||||
|
||||
// 混合有效和无效的 Flow ID
|
||||
let mut all_flow_ids = valid_flow_ids.clone();
|
||||
all_flow_ids.extend(invalid_flow_ids.clone());
|
||||
|
||||
// 执行批量收藏操作
|
||||
let result = batch_ops.execute(&all_flow_ids, BatchOperation::Star).await;
|
||||
|
||||
// 验证结果统计
|
||||
prop_assert_eq!(result.total, valid_flow_count + invalid_flow_count);
|
||||
prop_assert_eq!(result.success, valid_flow_count);
|
||||
prop_assert_eq!(result.failed, invalid_flow_count);
|
||||
prop_assert_eq!(result.errors.len(), invalid_flow_count);
|
||||
|
||||
// 验证成功的 Flow 被正确更新
|
||||
let store = monitor.memory_store();
|
||||
let s = store.read().await;
|
||||
for flow_id in &valid_flow_ids {
|
||||
if let Some(flow_lock) = s.get(flow_id) {
|
||||
let flow = flow_lock.read().unwrap();
|
||||
prop_assert!(flow.annotations.starred, "有效的 Flow {} 应该被收藏", flow_id);
|
||||
}
|
||||
}
|
||||
|
||||
// 验证失败的 Flow ID 被正确报告
|
||||
for invalid_id in &invalid_flow_ids {
|
||||
let found_error = result.errors.iter().any(|(id, _)| id == invalid_id);
|
||||
prop_assert!(found_error, "无效的 Flow ID {} 应该在错误列表中", invalid_id);
|
||||
}
|
||||
|
||||
// 验证部分成功状态
|
||||
prop_assert!(result.is_partial_success(), "应该是部分成功状态");
|
||||
prop_assert!(!result.is_all_success(), "不应该是全部成功");
|
||||
prop_assert!(!result.is_all_failed(), "不应该是全部失败");
|
||||
|
||||
Ok(())
|
||||
})?;
|
||||
}
|
||||
|
||||
/// **Feature: flow-monitor-enhancement, Property 21b: 批量标签操作原子性**
|
||||
/// **Validates: Requirements 11.2-11.6**
|
||||
///
|
||||
/// *对于任意* 批量标签操作,部分失败时应该正确处理成功和失败的情况。
|
||||
#[test]
|
||||
fn prop_batch_tag_operation_atomicity(
|
||||
valid_flow_count in 1usize..6usize,
|
||||
invalid_flow_count in 1usize..4usize,
|
||||
tag_count in 1usize..4usize,
|
||||
) {
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
rt.block_on(async {
|
||||
let monitor = create_test_flow_monitor();
|
||||
let batch_ops = BatchOperations::new(monitor.clone(), None);
|
||||
|
||||
// 创建有效的 Flow
|
||||
let mut valid_flow_ids = Vec::new();
|
||||
for i in 0..valid_flow_count {
|
||||
let id = create_test_flow(&monitor, &format!("valid-flow-{}", i)).await;
|
||||
valid_flow_ids.push(id);
|
||||
}
|
||||
|
||||
// 创建无效的 Flow ID
|
||||
let mut invalid_flow_ids = Vec::new();
|
||||
for i in 0..invalid_flow_count {
|
||||
invalid_flow_ids.push(format!("invalid-flow-{}", i));
|
||||
}
|
||||
|
||||
// 创建标签列表
|
||||
let tags: Vec<String> = (0..tag_count).map(|i| format!("tag-{}", i)).collect();
|
||||
|
||||
// 混合有效和无效的 Flow ID
|
||||
let mut all_flow_ids = valid_flow_ids.clone();
|
||||
all_flow_ids.extend(invalid_flow_ids.clone());
|
||||
|
||||
// 执行批量添加标签操作
|
||||
let result = batch_ops.execute(
|
||||
&all_flow_ids,
|
||||
BatchOperation::AddTags { tags: tags.clone() }
|
||||
).await;
|
||||
|
||||
// 验证结果统计
|
||||
prop_assert_eq!(result.total, valid_flow_count + invalid_flow_count);
|
||||
prop_assert_eq!(result.success, valid_flow_count);
|
||||
prop_assert_eq!(result.failed, invalid_flow_count);
|
||||
|
||||
// 验证成功的 Flow 被正确添加标签
|
||||
let store = monitor.memory_store();
|
||||
let s = store.read().await;
|
||||
for flow_id in &valid_flow_ids {
|
||||
if let Some(flow_lock) = s.get(flow_id) {
|
||||
let flow = flow_lock.read().unwrap();
|
||||
for tag in &tags {
|
||||
prop_assert!(
|
||||
flow.annotations.tags.contains(tag),
|
||||
"有效的 Flow {} 应该包含标签 {}",
|
||||
flow_id,
|
||||
tag
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 验证失败的 Flow ID 被正确报告
|
||||
for invalid_id in &invalid_flow_ids {
|
||||
let found_error = result.errors.iter().any(|(id, _)| id == invalid_id);
|
||||
prop_assert!(found_error, "无效的 Flow ID {} 应该在错误列表中", invalid_id);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
})?;
|
||||
}
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,148 @@
|
||||
//! LLM Flow Monitor 模块
|
||||
//!
|
||||
//! 该模块提供完整的 LLM API 流量监控功能,参考 mitmproxy 的 Flow 模型设计。
|
||||
//! 用于捕获、存储、分析和回放 AI Agent 与大模型之间的完整交互数据。
|
||||
//!
|
||||
//! # 主要组件
|
||||
//!
|
||||
//! - `models`: 核心数据模型,包括 LLMFlow、LLMRequest、LLMResponse 等
|
||||
//! - `stream_rebuilder`: SSE 流式响应重建器
|
||||
//! - `memory_store`: 内存存储,支持 LRU 驱逐策略
|
||||
//! - `file_store`: 文件存储,支持 JSONL 格式和 SQLite 索引
|
||||
//! - `query_service`: 查询服务,支持多维度过滤、排序、分页和全文搜索
|
||||
//! - `exporter`: 导出服务,支持 HAR、JSON、JSONL、Markdown、CSV 格式
|
||||
//! - `monitor`: 核心监控服务
|
||||
//! - `filter_parser`: 高级过滤表达式解析器,支持类似 mitmproxy 的语法
|
||||
|
||||
pub mod batch_ops;
|
||||
pub mod bookmark;
|
||||
pub mod code_exporter;
|
||||
pub mod diff;
|
||||
pub mod enhanced_stats;
|
||||
pub mod exporter;
|
||||
pub mod file_store;
|
||||
pub mod filter_parser;
|
||||
pub mod interceptor;
|
||||
pub mod memory_store;
|
||||
pub mod models;
|
||||
pub mod monitor;
|
||||
pub mod query_service;
|
||||
pub mod quick_filter;
|
||||
pub mod replayer;
|
||||
pub mod session;
|
||||
pub mod stream_rebuilder;
|
||||
|
||||
// 重新导出核心类型
|
||||
pub use models::{
|
||||
ClientInfo,
|
||||
ContentPart,
|
||||
FlowAnnotations,
|
||||
// 错误
|
||||
FlowError,
|
||||
FlowErrorType,
|
||||
// 元数据
|
||||
FlowMetadata,
|
||||
FlowState,
|
||||
FlowTimestamps,
|
||||
FlowType,
|
||||
// 核心 Flow 结构
|
||||
LLMFlow,
|
||||
// 请求相关
|
||||
LLMRequest,
|
||||
// 响应相关
|
||||
LLMResponse,
|
||||
Message,
|
||||
MessageContent,
|
||||
MessageRole,
|
||||
RequestParameters,
|
||||
RoutingInfo,
|
||||
StopReason,
|
||||
StreamChunk,
|
||||
StreamInfo,
|
||||
ThinkingContent,
|
||||
TokenUsage,
|
||||
ToolCall,
|
||||
ToolCallDelta,
|
||||
ToolDefinition,
|
||||
ToolResult,
|
||||
};
|
||||
|
||||
// 重新导出流重建器
|
||||
pub use stream_rebuilder::{StreamFormat, StreamRebuilder, StreamRebuilderError};
|
||||
|
||||
// 重新导出内存存储
|
||||
pub use memory_store::{FlowFilter, FlowMemoryStore, LatencyRange, TimeRange, TokenRange};
|
||||
|
||||
// 重新导出文件存储
|
||||
pub use file_store::{
|
||||
CleanupResult, FileStoreError, FlowFileStore, FlowIndexRecord, FtsSearchResult, RotationConfig,
|
||||
};
|
||||
|
||||
// 重新导出查询服务
|
||||
pub use query_service::{
|
||||
FlowQueryResult, FlowQueryService, FlowSearchResult, FlowSortBy, FlowStats, ModelStats,
|
||||
ProviderStats, QueryWithExpressionError, StateStats,
|
||||
};
|
||||
|
||||
// 重新导出导出服务
|
||||
pub use exporter::{
|
||||
default_redaction_rules, ExportFormat, ExportOptions, ExportResult, FlowExporter, HarArchive,
|
||||
HarEntry, HarLlmExtension, HarLog, RedactionRule, Redactor,
|
||||
};
|
||||
|
||||
// 重新导出监控服务
|
||||
pub use monitor::{
|
||||
FlowEvent, FlowMonitor, FlowMonitorConfig, FlowSummary, FlowUpdate, RequestRateTracker,
|
||||
ThresholdCheckResult, ThresholdConfig,
|
||||
};
|
||||
|
||||
// 重新导出过滤表达式解析器
|
||||
pub use filter_parser::{
|
||||
get_filter_help, Comparison, ComparisonOp, FilterExpr, FilterParseError, FilterParser,
|
||||
FilterToken, FILTER_HELP,
|
||||
};
|
||||
|
||||
// 重新导出拦截器
|
||||
pub use interceptor::{
|
||||
FlowInterceptor, InterceptAction, InterceptConfig, InterceptEvent, InterceptState,
|
||||
InterceptType, InterceptedFlow, InterceptorError, ModifiedData, TimeoutAction,
|
||||
};
|
||||
|
||||
// 重新导出重放器
|
||||
pub use replayer::{
|
||||
BatchReplayResult, FlowReplayer, ReplayConfig, ReplayResult, ReplayerError, RequestModification,
|
||||
};
|
||||
|
||||
// 重新导出差异对比器
|
||||
pub use diff::{
|
||||
DiffConfig, DiffItem, DiffType, FlowDiff, FlowDiffResult, MessageDiffItem, TokenDiff,
|
||||
};
|
||||
|
||||
// 重新导出会话管理器
|
||||
pub use session::{
|
||||
AutoSessionConfig, FlowSession, SessionError, SessionExportResult, SessionManager,
|
||||
};
|
||||
|
||||
// 重新导出快速过滤器管理器
|
||||
pub use quick_filter::{
|
||||
QuickFilter, QuickFilterError, QuickFilterExport, QuickFilterManager, QuickFilterUpdate,
|
||||
PRESET_FILTERS,
|
||||
};
|
||||
|
||||
// 重新导出代码导出器
|
||||
pub use code_exporter::{CodeExporter, CodeFormat};
|
||||
|
||||
// 重新导出书签管理器
|
||||
pub use bookmark::{BookmarkError, BookmarkExport, BookmarkManager, FlowBookmark};
|
||||
|
||||
// 重新导出增强统计服务
|
||||
pub use enhanced_stats::{
|
||||
Distribution, EnhancedStats, EnhancedStatsService, ReportFormat, StatsTimeRange,
|
||||
TimeSeriesPoint, TrendData,
|
||||
};
|
||||
|
||||
// 重新导出批量操作服务
|
||||
pub use batch_ops::{BatchOperation, BatchOperations, BatchOpsError, BatchResult};
|
||||
|
||||
// 重新导出 ProviderType(从 lib.rs)
|
||||
pub use crate::ProviderType;
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -4,6 +4,30 @@
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// 允许注入的参数白名单
|
||||
/// 这些参数是安全的,不会影响请求的核心行为
|
||||
const ALLOWED_INJECTION_PARAMS: &[&str] = &[
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
"top_p",
|
||||
"top_k",
|
||||
"frequency_penalty",
|
||||
"presence_penalty",
|
||||
"stop",
|
||||
"seed",
|
||||
"n",
|
||||
];
|
||||
|
||||
/// 禁止注入的参数黑名单(即使在白名单中也不允许 Override 模式)
|
||||
const BLOCKED_OVERRIDE_PARAMS: &[&str] = &[
|
||||
"model",
|
||||
"messages",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"stream",
|
||||
"response_format",
|
||||
];
|
||||
|
||||
/// 注入模式
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
@@ -213,6 +237,20 @@ impl Injector {
|
||||
let mut rule_applied = false;
|
||||
|
||||
for (key, value) in params {
|
||||
// 安全修复:检查参数是否在白名单中
|
||||
if !ALLOWED_INJECTION_PARAMS.contains(&key.as_str()) {
|
||||
tracing::warn!("[INJECTION] 参数 {} 不在白名单中,跳过注入", key);
|
||||
continue;
|
||||
}
|
||||
|
||||
// 安全修复:Override 模式下检查黑名单
|
||||
if rule.mode == InjectionMode::Override
|
||||
&& BLOCKED_OVERRIDE_PARAMS.contains(&key.as_str())
|
||||
{
|
||||
tracing::warn!("[INJECTION] 参数 {} 禁止使用 Override 模式", key);
|
||||
continue;
|
||||
}
|
||||
|
||||
let should_inject = match rule.mode {
|
||||
InjectionMode::Merge => !obj.contains_key(key),
|
||||
InjectionMode::Override => true,
|
||||
|
||||
+323
-11
@@ -3,6 +3,7 @@ mod config;
|
||||
mod converter;
|
||||
pub mod credential;
|
||||
mod database;
|
||||
pub mod flow_monitor;
|
||||
pub mod injection;
|
||||
mod logger;
|
||||
pub mod middleware;
|
||||
@@ -16,6 +17,7 @@ pub mod router;
|
||||
mod server;
|
||||
mod server_utils;
|
||||
mod services;
|
||||
pub mod streaming;
|
||||
pub mod telemetry;
|
||||
pub mod tray;
|
||||
pub mod websocket;
|
||||
@@ -25,11 +27,21 @@ use std::sync::Arc;
|
||||
use tauri::{Manager, Runtime};
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use commands::flow_monitor_cmd::{
|
||||
BatchOperationsState, BookmarkManagerState, EnhancedStatsServiceState, FlowInterceptorState,
|
||||
FlowMonitorState, FlowQueryServiceState, FlowReplayerState, QuickFilterManagerState,
|
||||
SessionManagerState,
|
||||
};
|
||||
use commands::plugin_cmd::PluginManagerState;
|
||||
use commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState};
|
||||
use commands::resilience_cmd::ResilienceConfigState;
|
||||
use commands::router_cmd::RouterConfigState;
|
||||
use commands::skill_cmd::SkillServiceState;
|
||||
use flow_monitor::{
|
||||
BatchOperations, BookmarkManager, EnhancedStatsService, FlowFileStore, FlowInterceptor,
|
||||
FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowReplayer, InterceptConfig,
|
||||
QuickFilterManager, SessionManager,
|
||||
};
|
||||
use services::provider_pool_service::ProviderPoolService;
|
||||
use services::skill_service::SkillService;
|
||||
use services::token_cache_service::TokenCacheService;
|
||||
@@ -188,6 +200,10 @@ mod tests {
|
||||
pub type AppState = Arc<RwLock<server::ServerState>>;
|
||||
pub type LogState = Arc<RwLock<logger::LogStore>>;
|
||||
|
||||
fn generate_api_key() -> String {
|
||||
config::generate_secure_api_key()
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
async fn start_server(
|
||||
state: tauri::State<'_, AppState>,
|
||||
@@ -246,6 +262,20 @@ async fn save_config(
|
||||
state: tauri::State<'_, AppState>,
|
||||
config: config::Config,
|
||||
) -> Result<(), String> {
|
||||
// P0 安全修复:禁止危险的网络配置
|
||||
let host = config.server.host.to_lowercase();
|
||||
if host == "0.0.0.0" || host == "::" {
|
||||
return Err(
|
||||
"安全限制:不允许监听所有网络接口 (0.0.0.0 或 ::)。请使用 127.0.0.1 或 localhost"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
// 禁止开启远程管理
|
||||
if config.remote_management.allow_remote {
|
||||
return Err("安全限制:不允许开启远程管理功能".to_string());
|
||||
}
|
||||
|
||||
let mut s = state.write().await;
|
||||
s.config = config.clone();
|
||||
config::save_config(&config).map_err(|e| e.to_string())
|
||||
@@ -365,31 +395,32 @@ async fn get_env_variables(state: tauri::State<'_, AppState>) -> Result<Vec<EnvV
|
||||
let creds = &s.kiro_provider.credentials;
|
||||
let mut vars = Vec::new();
|
||||
|
||||
// P0 安全修复:不再返回明文敏感凭证,仅返回 masked 版本
|
||||
if let Some(token) = &creds.access_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_ACCESS_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(token) = &creds.refresh_token {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_REFRESH_TOKEN".to_string(),
|
||||
value: token.clone(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(token),
|
||||
});
|
||||
}
|
||||
if let Some(id) = &creds.client_id {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_CLIENT_ID".to_string(),
|
||||
value: id.clone(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(id),
|
||||
});
|
||||
}
|
||||
if let Some(secret) = &creds.client_secret {
|
||||
vars.push(EnvVariable {
|
||||
key: "KIRO_CLIENT_SECRET".to_string(),
|
||||
value: secret.clone(),
|
||||
value: String::new(), // 不返回明文
|
||||
masked: mask_token(secret),
|
||||
});
|
||||
}
|
||||
@@ -1363,12 +1394,57 @@ async fn test_api(
|
||||
|
||||
#[cfg_attr(mobile, tauri::mobile_entry_point)]
|
||||
pub fn run() {
|
||||
let config = config::load_config().unwrap_or_default();
|
||||
let state: AppState = Arc::new(RwLock::new(server::ServerState::new(config)));
|
||||
let logs: LogState = Arc::new(RwLock::new(logger::LogStore::new()));
|
||||
let mut config = match config::load_config() {
|
||||
Ok(cfg) => cfg,
|
||||
Err(err) => {
|
||||
tracing::error!("配置加载失败,已中止启动: {}", err);
|
||||
eprintln!("配置加载失败,已中止启动: {}", err);
|
||||
return;
|
||||
}
|
||||
};
|
||||
if config.server.api_key == config::DEFAULT_API_KEY {
|
||||
let new_key = generate_api_key();
|
||||
config.server.api_key = new_key.clone();
|
||||
if let Err(err) = config::save_config(&config) {
|
||||
tracing::error!("自动生成 API key 失败,无法保存配置,已中止启动: {}", err);
|
||||
eprintln!("自动生成 API key 失败,无法保存配置,已中止启动: {}", err);
|
||||
return;
|
||||
}
|
||||
tracing::info!("检测到默认 API key,已自动生成并保存新密钥");
|
||||
eprintln!("检测到默认 API key,已自动生成并保存新密钥");
|
||||
}
|
||||
if !is_loopback_host(&config.server.host) {
|
||||
tracing::error!("当前版本仅支持本地监听,请使用 127.0.0.1/localhost/::1。");
|
||||
eprintln!("当前版本仅支持本地监听,请使用 127.0.0.1/localhost/::1。");
|
||||
return;
|
||||
}
|
||||
if config.server.api_key == config::DEFAULT_API_KEY {
|
||||
tracing::error!("检测到使用默认 API key,已中止启动。请配置强密钥。");
|
||||
eprintln!("检测到使用默认 API key,已中止启动。请配置强密钥。");
|
||||
return;
|
||||
}
|
||||
if config.server.tls.enable {
|
||||
tracing::error!("检测到 TLS 配置已启用,但当前版本尚未支持 TLS,已中止启动。");
|
||||
eprintln!("检测到 TLS 配置已启用,但当前版本尚未支持 TLS,已中止启动。");
|
||||
return;
|
||||
}
|
||||
if config.remote_management.allow_remote {
|
||||
tracing::error!("检测到远程管理已开启,但当前版本未启用 TLS,已中止启动。");
|
||||
eprintln!("检测到远程管理已开启,但当前版本未启用 TLS,已中止启动。");
|
||||
return;
|
||||
}
|
||||
let state: AppState = Arc::new(RwLock::new(server::ServerState::new(config.clone())));
|
||||
let logs: LogState = Arc::new(RwLock::new(logger::LogStore::with_config(&config.logging)));
|
||||
|
||||
// Initialize database for Switch functionality
|
||||
let db = database::init_database().expect("Failed to initialize database");
|
||||
let db = match database::init_database() {
|
||||
Ok(conn) => conn,
|
||||
Err(err) => {
|
||||
tracing::error!("数据库初始化失败,已中止启动: {}", err);
|
||||
eprintln!("数据库初始化失败,已中止启动: {}", err);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// Initialize SkillService
|
||||
let skill_service = SkillService::new().expect("Failed to initialize SkillService");
|
||||
@@ -1405,8 +1481,14 @@ pub fn run() {
|
||||
let shared_tokens = Arc::new(parking_lot::RwLock::new(
|
||||
telemetry::TokenTracker::with_defaults(),
|
||||
));
|
||||
let log_rotation = telemetry::LogRotationConfig {
|
||||
max_memory_logs: 10000,
|
||||
retention_days: config.logging.retention_days,
|
||||
max_file_size: 10 * 1024 * 1024,
|
||||
enable_file_logging: config.logging.enabled,
|
||||
};
|
||||
let shared_logger = Arc::new(
|
||||
telemetry::RequestLogger::with_defaults().expect("Failed to create RequestLogger"),
|
||||
telemetry::RequestLogger::new(log_rotation).expect("Failed to create RequestLogger"),
|
||||
);
|
||||
|
||||
// Initialize TelemetryState with shared instances
|
||||
@@ -1417,6 +1499,91 @@ pub fn run() {
|
||||
)
|
||||
.expect("Failed to create TelemetryState");
|
||||
|
||||
// Initialize FlowMonitor and FlowQueryService
|
||||
let flow_monitor_config = FlowMonitorConfig::default();
|
||||
let flow_file_store = {
|
||||
// 获取应用数据目录
|
||||
let data_dir = dirs::data_dir()
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("."))
|
||||
.join("proxycast")
|
||||
.join("flows");
|
||||
|
||||
// 创建目录(如果不存在)
|
||||
if let Err(e) = std::fs::create_dir_all(&data_dir) {
|
||||
tracing::warn!("无法创建 Flow 存储目录: {}", e);
|
||||
}
|
||||
|
||||
let rotation_config = flow_monitor::RotationConfig::default();
|
||||
match FlowFileStore::new(data_dir, rotation_config) {
|
||||
Ok(store) => Some(Arc::new(store)),
|
||||
Err(e) => {
|
||||
tracing::warn!("无法初始化 Flow 文件存储: {}", e);
|
||||
None
|
||||
}
|
||||
}
|
||||
};
|
||||
let flow_monitor = Arc::new(FlowMonitor::new(
|
||||
flow_monitor_config,
|
||||
flow_file_store.clone(),
|
||||
));
|
||||
let flow_monitor_state = FlowMonitorState(flow_monitor.clone());
|
||||
|
||||
// 初始化 Flow 拦截器
|
||||
let flow_interceptor = Arc::new(FlowInterceptor::new(InterceptConfig::default()));
|
||||
let flow_interceptor_state = FlowInterceptorState(flow_interceptor.clone());
|
||||
|
||||
// 初始化 Flow 重放器
|
||||
let flow_replayer = Arc::new(FlowReplayer::new(
|
||||
flow_monitor.clone(),
|
||||
provider_pool_service_state.0.clone(),
|
||||
db.clone(),
|
||||
));
|
||||
let flow_replayer_state = FlowReplayerState(flow_replayer);
|
||||
|
||||
// 初始化会话管理器
|
||||
let db_path = database::get_db_path();
|
||||
let session_manager =
|
||||
Arc::new(SessionManager::new(db_path.clone()).expect("Failed to create SessionManager"));
|
||||
let session_manager_state = SessionManagerState(session_manager);
|
||||
|
||||
// 初始化快速过滤器管理器
|
||||
let quick_filter_manager = Arc::new(
|
||||
QuickFilterManager::new(db_path.clone()).expect("Failed to create QuickFilterManager"),
|
||||
);
|
||||
let quick_filter_manager_state = QuickFilterManagerState(quick_filter_manager);
|
||||
|
||||
// 初始化书签管理器
|
||||
let bookmark_manager =
|
||||
Arc::new(BookmarkManager::new(db_path).expect("Failed to create BookmarkManager"));
|
||||
let bookmark_manager_state = BookmarkManagerState(bookmark_manager);
|
||||
|
||||
// 初始化增强统计服务
|
||||
let enhanced_stats_service = Arc::new(EnhancedStatsService::new(flow_monitor.memory_store()));
|
||||
let enhanced_stats_service_state = EnhancedStatsServiceState(enhanced_stats_service);
|
||||
|
||||
// 初始化批量操作服务
|
||||
let batch_operations = Arc::new(BatchOperations::new(
|
||||
flow_monitor.clone(),
|
||||
Some(session_manager_state.0.clone()),
|
||||
));
|
||||
let batch_operations_state = BatchOperationsState(batch_operations);
|
||||
|
||||
// FlowQueryService 需要 file_store,如果没有则创建一个临时的
|
||||
let flow_query_service_state = if let Some(file_store) = flow_file_store {
|
||||
let query_service = FlowQueryService::new(flow_monitor.memory_store(), file_store);
|
||||
FlowQueryServiceState(Arc::new(query_service))
|
||||
} else {
|
||||
// 如果没有文件存储,创建一个临时的内存存储
|
||||
let temp_dir = std::env::temp_dir().join("proxycast_flows");
|
||||
let _ = std::fs::create_dir_all(&temp_dir);
|
||||
let rotation_config = flow_monitor::RotationConfig::default();
|
||||
let temp_store = FlowFileStore::new(temp_dir, rotation_config)
|
||||
.expect("Failed to create temp FlowFileStore");
|
||||
let query_service =
|
||||
FlowQueryService::new(flow_monitor.memory_store(), Arc::new(temp_store));
|
||||
FlowQueryServiceState(Arc::new(query_service))
|
||||
};
|
||||
|
||||
// Initialize default skill repos
|
||||
{
|
||||
let conn = db.lock().expect("Failed to lock database");
|
||||
@@ -1433,6 +1600,8 @@ pub fn run() {
|
||||
let shared_stats_clone = shared_stats.clone();
|
||||
let shared_tokens_clone = shared_tokens.clone();
|
||||
let shared_logger_clone = shared_logger.clone();
|
||||
let flow_monitor_clone = flow_monitor.clone();
|
||||
let flow_interceptor_clone = flow_interceptor.clone();
|
||||
|
||||
tauri::Builder::default()
|
||||
.plugin(tauri_plugin_shell::init())
|
||||
@@ -1452,6 +1621,15 @@ pub fn run() {
|
||||
.manage(resilience_config_state)
|
||||
.manage(telemetry_state)
|
||||
.manage(plugin_manager_state)
|
||||
.manage(flow_monitor_state)
|
||||
.manage(flow_query_service_state)
|
||||
.manage(flow_interceptor_state)
|
||||
.manage(flow_replayer_state)
|
||||
.manage(session_manager_state)
|
||||
.manage(quick_filter_manager_state)
|
||||
.manage(bookmark_manager_state)
|
||||
.manage(enhanced_stats_service_state)
|
||||
.manage(batch_operations_state)
|
||||
.setup(move |app| {
|
||||
// 初始化托盘管理器
|
||||
// Requirements 1.4: 应用启动时显示停止状态图标
|
||||
@@ -1480,6 +1658,7 @@ pub fn run() {
|
||||
let shared_stats = shared_stats_clone.clone();
|
||||
let shared_tokens = shared_tokens_clone.clone();
|
||||
let shared_logger = shared_logger_clone.clone();
|
||||
let shared_flow_monitor = flow_monitor_clone.clone();
|
||||
let app_handle = app.handle().clone();
|
||||
tauri::async_runtime::spawn(async move {
|
||||
// 先加载凭证
|
||||
@@ -1493,7 +1672,7 @@ pub fn run() {
|
||||
logs.write().await.add("info", "[启动] Kiro 凭证已加载");
|
||||
}
|
||||
}
|
||||
// 启动服务器(使用共享的遥测实例)
|
||||
// 启动服务器(使用共享的遥测实例和 Flow Monitor)
|
||||
let server_started;
|
||||
let server_address;
|
||||
{
|
||||
@@ -1502,7 +1681,7 @@ pub fn run() {
|
||||
.await
|
||||
.add("info", "[启动] 正在自动启动服务器...");
|
||||
match s
|
||||
.start_with_telemetry(
|
||||
.start_with_telemetry_and_flow_monitor(
|
||||
logs.clone(),
|
||||
pool_service,
|
||||
token_cache,
|
||||
@@ -1510,6 +1689,8 @@ pub fn run() {
|
||||
Some(shared_stats),
|
||||
Some(shared_tokens),
|
||||
Some(shared_logger),
|
||||
Some(shared_flow_monitor),
|
||||
Some(flow_interceptor_clone),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -1779,7 +1960,138 @@ pub fn run() {
|
||||
commands::plugin_cmd::reload_plugins,
|
||||
commands::plugin_cmd::unload_plugin,
|
||||
commands::plugin_cmd::get_plugins_dir,
|
||||
// Flow Monitor commands
|
||||
commands::flow_monitor_cmd::query_flows,
|
||||
commands::flow_monitor_cmd::get_flow_detail,
|
||||
commands::flow_monitor_cmd::search_flows,
|
||||
commands::flow_monitor_cmd::get_flow_stats,
|
||||
commands::flow_monitor_cmd::export_flows,
|
||||
commands::flow_monitor_cmd::update_flow_annotations,
|
||||
commands::flow_monitor_cmd::toggle_flow_starred,
|
||||
commands::flow_monitor_cmd::add_flow_comment,
|
||||
commands::flow_monitor_cmd::add_flow_tag,
|
||||
commands::flow_monitor_cmd::remove_flow_tag,
|
||||
commands::flow_monitor_cmd::set_flow_marker,
|
||||
commands::flow_monitor_cmd::cleanup_flows,
|
||||
commands::flow_monitor_cmd::get_recent_flows,
|
||||
commands::flow_monitor_cmd::get_flow_monitor_status,
|
||||
commands::flow_monitor_cmd::get_flow_monitor_debug_info,
|
||||
commands::flow_monitor_cmd::create_test_flows,
|
||||
commands::flow_monitor_cmd::enable_flow_monitor,
|
||||
commands::flow_monitor_cmd::disable_flow_monitor,
|
||||
commands::flow_monitor_cmd::subscribe_flow_events,
|
||||
commands::flow_monitor_cmd::get_all_flow_tags,
|
||||
// Flow Monitor filter expression commands
|
||||
commands::flow_monitor_cmd::parse_filter,
|
||||
commands::flow_monitor_cmd::validate_filter,
|
||||
commands::flow_monitor_cmd::get_filter_help_items,
|
||||
commands::flow_monitor_cmd::get_filter_help_text,
|
||||
commands::flow_monitor_cmd::query_flows_with_expression,
|
||||
// Flow Interceptor commands
|
||||
commands::flow_monitor_cmd::intercept_config_get,
|
||||
commands::flow_monitor_cmd::intercept_config_set,
|
||||
commands::flow_monitor_cmd::intercept_continue,
|
||||
commands::flow_monitor_cmd::intercept_cancel,
|
||||
commands::flow_monitor_cmd::intercept_get_flow,
|
||||
commands::flow_monitor_cmd::intercept_list_flows,
|
||||
commands::flow_monitor_cmd::intercept_count,
|
||||
commands::flow_monitor_cmd::intercept_is_enabled,
|
||||
commands::flow_monitor_cmd::intercept_enable,
|
||||
commands::flow_monitor_cmd::intercept_disable,
|
||||
commands::flow_monitor_cmd::intercept_set_editing,
|
||||
commands::flow_monitor_cmd::subscribe_intercept_events,
|
||||
// Flow Monitor realtime enhancement commands
|
||||
commands::flow_monitor_cmd::get_threshold_config,
|
||||
commands::flow_monitor_cmd::update_threshold_config,
|
||||
commands::flow_monitor_cmd::get_request_rate,
|
||||
commands::flow_monitor_cmd::set_rate_window,
|
||||
// Flow Replayer commands
|
||||
commands::flow_monitor_cmd::replay_flow,
|
||||
commands::flow_monitor_cmd::replay_flows_batch,
|
||||
// Flow Diff commands
|
||||
commands::flow_monitor_cmd::diff_flows,
|
||||
// Session Management commands
|
||||
commands::flow_monitor_cmd::create_session,
|
||||
commands::flow_monitor_cmd::get_session,
|
||||
commands::flow_monitor_cmd::list_sessions,
|
||||
commands::flow_monitor_cmd::add_flow_to_session,
|
||||
commands::flow_monitor_cmd::remove_flow_from_session,
|
||||
commands::flow_monitor_cmd::update_session,
|
||||
commands::flow_monitor_cmd::archive_session,
|
||||
commands::flow_monitor_cmd::unarchive_session,
|
||||
commands::flow_monitor_cmd::delete_session,
|
||||
commands::flow_monitor_cmd::export_session,
|
||||
commands::flow_monitor_cmd::get_session_flow_count,
|
||||
commands::flow_monitor_cmd::is_flow_in_session,
|
||||
commands::flow_monitor_cmd::get_sessions_for_flow,
|
||||
commands::flow_monitor_cmd::get_auto_session_config,
|
||||
commands::flow_monitor_cmd::set_auto_session_config,
|
||||
commands::flow_monitor_cmd::register_active_session,
|
||||
// Quick Filter commands
|
||||
commands::flow_monitor_cmd::save_quick_filter,
|
||||
commands::flow_monitor_cmd::get_quick_filter,
|
||||
commands::flow_monitor_cmd::update_quick_filter,
|
||||
commands::flow_monitor_cmd::delete_quick_filter,
|
||||
commands::flow_monitor_cmd::list_quick_filters,
|
||||
commands::flow_monitor_cmd::list_quick_filters_by_group,
|
||||
commands::flow_monitor_cmd::list_quick_filter_groups,
|
||||
commands::flow_monitor_cmd::export_quick_filters,
|
||||
commands::flow_monitor_cmd::import_quick_filters,
|
||||
commands::flow_monitor_cmd::find_quick_filter_by_name,
|
||||
// Code Export commands
|
||||
commands::flow_monitor_cmd::export_flow_as_code,
|
||||
commands::flow_monitor_cmd::export_flows_as_code,
|
||||
commands::flow_monitor_cmd::get_code_export_formats,
|
||||
// Bookmark Management commands
|
||||
commands::flow_monitor_cmd::add_bookmark,
|
||||
commands::flow_monitor_cmd::get_bookmark,
|
||||
commands::flow_monitor_cmd::get_bookmark_by_flow_id,
|
||||
commands::flow_monitor_cmd::remove_bookmark,
|
||||
commands::flow_monitor_cmd::remove_bookmark_by_flow_id,
|
||||
commands::flow_monitor_cmd::update_bookmark,
|
||||
commands::flow_monitor_cmd::list_bookmarks,
|
||||
commands::flow_monitor_cmd::list_bookmark_groups,
|
||||
commands::flow_monitor_cmd::is_flow_bookmarked,
|
||||
commands::flow_monitor_cmd::get_bookmark_count,
|
||||
commands::flow_monitor_cmd::export_bookmarks,
|
||||
commands::flow_monitor_cmd::import_bookmarks,
|
||||
commands::flow_monitor_cmd::toggle_bookmark,
|
||||
// Enhanced Stats commands
|
||||
commands::flow_monitor_cmd::get_enhanced_stats,
|
||||
commands::flow_monitor_cmd::get_request_trend,
|
||||
commands::flow_monitor_cmd::get_token_distribution,
|
||||
commands::flow_monitor_cmd::get_latency_histogram,
|
||||
commands::flow_monitor_cmd::export_stats_report,
|
||||
// Batch Operations commands
|
||||
commands::flow_monitor_cmd::batch_star_flows,
|
||||
commands::flow_monitor_cmd::batch_unstar_flows,
|
||||
commands::flow_monitor_cmd::batch_add_tags,
|
||||
commands::flow_monitor_cmd::batch_remove_tags,
|
||||
commands::flow_monitor_cmd::batch_export_flows,
|
||||
commands::flow_monitor_cmd::batch_delete_flows,
|
||||
commands::flow_monitor_cmd::batch_add_to_session,
|
||||
// Window control commands
|
||||
commands::window_cmd::get_window_size,
|
||||
commands::window_cmd::set_window_size,
|
||||
commands::window_cmd::resize_for_flow_monitor,
|
||||
commands::window_cmd::restore_window_size,
|
||||
commands::window_cmd::toggle_window_size,
|
||||
commands::window_cmd::center_window,
|
||||
commands::window_cmd::get_window_size_options,
|
||||
commands::window_cmd::set_window_size_by_option,
|
||||
commands::window_cmd::toggle_fullscreen,
|
||||
commands::window_cmd::is_fullscreen,
|
||||
])
|
||||
.run(tauri::generate_context!())
|
||||
.expect("error while running tauri application");
|
||||
}
|
||||
|
||||
fn is_loopback_host(host: &str) -> bool {
|
||||
if host == "localhost" {
|
||||
return true;
|
||||
}
|
||||
match host.parse::<std::net::IpAddr>() {
|
||||
Ok(addr) => addr.is_loopback(),
|
||||
Err(_) => false,
|
||||
}
|
||||
}
|
||||
|
||||
+276
-10
@@ -1,13 +1,35 @@
|
||||
//! 日志管理模块
|
||||
use chrono::{Local, Utc};
|
||||
use chrono::{Duration, Local, Utc};
|
||||
use flate2::write::GzEncoder;
|
||||
use flate2::Compression;
|
||||
use regex::Regex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::VecDeque;
|
||||
use std::fs::{self, OpenOptions};
|
||||
use std::io::Write;
|
||||
use std::io::{Read, Write};
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct LogStoreConfig {
|
||||
pub max_logs: usize,
|
||||
pub retention_days: u32,
|
||||
pub max_file_size: u64,
|
||||
pub enable_file_logging: bool,
|
||||
}
|
||||
|
||||
impl Default for LogStoreConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_logs: 1000,
|
||||
retention_days: 7,
|
||||
max_file_size: 10 * 1024 * 1024,
|
||||
enable_file_logging: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct LogEntry {
|
||||
pub timestamp: String,
|
||||
@@ -18,6 +40,7 @@ pub struct LogEntry {
|
||||
pub struct LogStore {
|
||||
logs: VecDeque<LogEntry>,
|
||||
max_logs: usize,
|
||||
config: LogStoreConfig,
|
||||
log_file_path: Option<PathBuf>,
|
||||
}
|
||||
|
||||
@@ -34,9 +57,12 @@ impl Default for LogStore {
|
||||
|
||||
let log_file = log_dir.join("proxycast.log");
|
||||
|
||||
let config = LogStoreConfig::default();
|
||||
|
||||
Self {
|
||||
logs: VecDeque::new(),
|
||||
max_logs: 1000,
|
||||
max_logs: config.max_logs,
|
||||
config,
|
||||
log_file_path: Some(log_file),
|
||||
}
|
||||
}
|
||||
@@ -47,23 +73,36 @@ impl LogStore {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
pub fn with_config(logging: &crate::config::LoggingConfig) -> Self {
|
||||
let mut store = Self::default();
|
||||
store.config.retention_days = logging.retention_days;
|
||||
store.config.enable_file_logging = logging.enabled;
|
||||
store.max_logs = store.config.max_logs;
|
||||
store
|
||||
}
|
||||
|
||||
pub fn add(&mut self, level: &str, message: &str) {
|
||||
let sanitized = sanitize_log_message(message);
|
||||
let now = Utc::now();
|
||||
let entry = LogEntry {
|
||||
timestamp: now.to_rfc3339(),
|
||||
level: level.to_string(),
|
||||
message: message.to_string(),
|
||||
message: sanitized.clone(),
|
||||
};
|
||||
|
||||
self.logs.push_back(entry.clone());
|
||||
|
||||
// 写入日志文件
|
||||
if let Some(ref path) = self.log_file_path {
|
||||
let local_time = Local::now().format("%Y-%m-%d %H:%M:%S%.3f");
|
||||
let log_line = format!("{} [{}] {}\n", local_time, level.to_uppercase(), message);
|
||||
if self.config.enable_file_logging {
|
||||
if let Some(ref path) = self.log_file_path {
|
||||
self.rotate_log_file_if_needed(path);
|
||||
let local_time = Local::now().format("%Y-%m-%d %H:%M:%S%.3f");
|
||||
let log_line = format!("{} [{}] {}\n", local_time, level.to_uppercase(), sanitized);
|
||||
|
||||
if let Ok(mut file) = OpenOptions::new().create(true).append(true).open(path) {
|
||||
let _ = file.write_all(log_line.as_bytes());
|
||||
if let Ok(mut file) = OpenOptions::new().create(true).append(true).open(path) {
|
||||
let _ = file.write_all(log_line.as_bytes());
|
||||
}
|
||||
self.prune_old_logs(path);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,6 +117,7 @@ impl LogStore {
|
||||
if let Some(ref log_path) = self.log_file_path {
|
||||
let log_dir = log_path.parent().unwrap_or(std::path::Path::new("."));
|
||||
let raw_file = log_dir.join(format!("raw_response_{request_id}.txt"));
|
||||
let sanitized = sanitize_log_message(body);
|
||||
|
||||
if let Ok(mut file) = OpenOptions::new()
|
||||
.create(true)
|
||||
@@ -85,7 +125,7 @@ impl LogStore {
|
||||
.write(true)
|
||||
.open(&raw_file)
|
||||
{
|
||||
let _ = file.write_all(body.as_bytes());
|
||||
let _ = file.write_all(sanitized.as_bytes());
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -103,7 +143,233 @@ impl LogStore {
|
||||
.as_ref()
|
||||
.map(|p| p.to_string_lossy().to_string())
|
||||
}
|
||||
|
||||
fn rotate_log_file_if_needed(&self, path: &PathBuf) {
|
||||
let Ok(metadata) = fs::metadata(path) else {
|
||||
return;
|
||||
};
|
||||
|
||||
if metadata.len() <= self.config.max_file_size {
|
||||
return;
|
||||
}
|
||||
|
||||
let suffix = Local::now().format("%Y%m%d-%H%M%S");
|
||||
let rotated = path.with_file_name(format!(
|
||||
"{}.{}",
|
||||
path.file_name().unwrap_or_default().to_string_lossy(),
|
||||
suffix
|
||||
));
|
||||
|
||||
let _ = fs::rename(path, &rotated);
|
||||
self.prune_old_logs(path);
|
||||
}
|
||||
|
||||
fn prune_old_logs(&self, path: &PathBuf) {
|
||||
let Some(dir) = path.parent() else {
|
||||
return;
|
||||
};
|
||||
self.archive_old_logs(path);
|
||||
let Ok(entries) = fs::read_dir(dir) else {
|
||||
return;
|
||||
};
|
||||
let cutoff = Utc::now() - Duration::days(self.config.retention_days as i64);
|
||||
let prefix = format!(
|
||||
"{}.",
|
||||
path.file_name().unwrap_or_default().to_string_lossy()
|
||||
);
|
||||
|
||||
for entry in entries.flatten() {
|
||||
let file_name = entry.file_name();
|
||||
let file_name = file_name.to_string_lossy();
|
||||
if !file_name.starts_with(&prefix) {
|
||||
continue;
|
||||
}
|
||||
let Ok(metadata) = entry.metadata() else {
|
||||
continue;
|
||||
};
|
||||
let Ok(modified) = metadata.modified() else {
|
||||
continue;
|
||||
};
|
||||
let modified = chrono::DateTime::<Utc>::from(modified);
|
||||
if modified < cutoff {
|
||||
let _ = fs::remove_file(entry.path());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn archive_old_logs(&self, path: &PathBuf) {
|
||||
let Some(dir) = path.parent() else {
|
||||
return;
|
||||
};
|
||||
let Ok(entries) = fs::read_dir(dir) else {
|
||||
return;
|
||||
};
|
||||
let archive_cutoff = Utc::now() - Duration::days(7);
|
||||
let delete_cutoff = Utc::now() - Duration::days(30);
|
||||
let prefix = format!(
|
||||
"{}.",
|
||||
path.file_name().unwrap_or_default().to_string_lossy()
|
||||
);
|
||||
|
||||
for entry in entries.flatten() {
|
||||
let file_name = entry.file_name();
|
||||
let file_name = file_name.to_string_lossy();
|
||||
if !file_name.starts_with(&prefix) {
|
||||
continue;
|
||||
}
|
||||
let path = entry.path();
|
||||
let Ok(metadata) = entry.metadata() else {
|
||||
continue;
|
||||
};
|
||||
let Ok(modified) = metadata.modified() else {
|
||||
continue;
|
||||
};
|
||||
let modified = chrono::DateTime::<Utc>::from(modified);
|
||||
|
||||
if file_name.ends_with(".gz") {
|
||||
if modified < delete_cutoff {
|
||||
let _ = fs::remove_file(path);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
if modified >= archive_cutoff {
|
||||
continue;
|
||||
}
|
||||
|
||||
let mut input = Vec::new();
|
||||
if let Ok(mut file) = fs::File::open(&path) {
|
||||
if file.read_to_end(&mut input).is_err() {
|
||||
continue;
|
||||
}
|
||||
} else {
|
||||
continue;
|
||||
}
|
||||
|
||||
let gz_path = path.with_extension(format!(
|
||||
"{}.gz",
|
||||
path.extension().unwrap_or_default().to_string_lossy()
|
||||
));
|
||||
if let Ok(gz_file) = fs::File::create(&gz_path) {
|
||||
let mut encoder = GzEncoder::new(gz_file, Compression::default());
|
||||
if encoder.write_all(&input).is_ok() && encoder.finish().is_ok() {
|
||||
let _ = fs::remove_file(&path);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub type SharedLogStore = Arc<RwLock<LogStore>>;
|
||||
|
||||
/// P2 安全修复:扩展日志脱敏规则,覆盖更多敏感字段
|
||||
pub fn sanitize_log_message(message: &str) -> String {
|
||||
let patterns = [
|
||||
// Bearer token
|
||||
(r"Bearer\s+[A-Za-z0-9._-]+", "Bearer ***"),
|
||||
// API key 各种格式
|
||||
(
|
||||
r#"api[_-]?key["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#,
|
||||
"api_key: ***",
|
||||
),
|
||||
// 通用 token
|
||||
(r#"token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#, "token: ***"),
|
||||
// P2 新增:access_token
|
||||
(
|
||||
r#"access[_-]?token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#,
|
||||
"access_token: ***",
|
||||
),
|
||||
// P2 新增:refresh_token
|
||||
(
|
||||
r#"refresh[_-]?token["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#,
|
||||
"refresh_token: ***",
|
||||
),
|
||||
// P2 新增:client_secret
|
||||
(
|
||||
r#"client[_-]?secret["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#,
|
||||
"client_secret: ***",
|
||||
),
|
||||
// P2 新增:authorization header
|
||||
(
|
||||
r#"[Aa]uthorization["']?\s*[:=]\s*["']?[A-Za-z0-9._\s-]+"#,
|
||||
"authorization: ***",
|
||||
),
|
||||
// P2 新增:password
|
||||
(r#"password["']?\s*[:=]\s*["']?[^\s"',}]+"#, "password: ***"),
|
||||
// P2 新增:secret
|
||||
(
|
||||
r#"secret["']?\s*[:=]\s*["']?[A-Za-z0-9._-]+"#,
|
||||
"secret: ***",
|
||||
),
|
||||
];
|
||||
|
||||
let mut sanitized = message.to_string();
|
||||
for (pattern, replacement) in patterns {
|
||||
if let Ok(re) = Regex::new(pattern) {
|
||||
sanitized = re.replace_all(&sanitized, replacement).to_string();
|
||||
}
|
||||
}
|
||||
sanitized
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::sanitize_log_message;
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_bearer_token() {
|
||||
let input = "Authorization: Bearer abcDEF123._-XYZ";
|
||||
let output = sanitize_log_message(input);
|
||||
// 验证敏感 token 被脱敏
|
||||
assert!(!output.contains("abcDEF123"));
|
||||
assert!(output.contains("***"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_api_key() {
|
||||
let input = r#"request api_key="sk-test_123.456-ABC" end"#;
|
||||
let output = sanitize_log_message(input);
|
||||
assert!(output.contains("api_key: ***"));
|
||||
assert!(!output.contains("sk-test_123"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_access_token() {
|
||||
let input = "access_token=atk_12345";
|
||||
let output = sanitize_log_message(input);
|
||||
assert!(output.contains("access_token: ***"));
|
||||
assert!(!output.contains("atk_12345"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_refresh_token() {
|
||||
let input = "refresh_token: rtk_ABCDE-123";
|
||||
let output = sanitize_log_message(input);
|
||||
assert!(output.contains("refresh_token: ***"));
|
||||
assert!(!output.contains("rtk_ABCDE"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_client_secret() {
|
||||
let input = "client_secret = \"cs_SeCreT-999\"";
|
||||
let output = sanitize_log_message(input);
|
||||
assert!(output.contains("client_secret: ***"));
|
||||
assert!(!output.contains("cs_SeCreT"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_password() {
|
||||
let input = r#"{"password":"p@ssW0rd!"}"#;
|
||||
let output = sanitize_log_message(input);
|
||||
assert!(output.contains("password: ***"));
|
||||
assert!(!output.contains("p@ssW0rd!"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_plain_text_unchanged() {
|
||||
let input = "这是一段普通日志,不包含任何敏感字段。";
|
||||
let output = sanitize_log_message(input);
|
||||
assert_eq!(output, input);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,10 +20,39 @@ use futures::future::BoxFuture;
|
||||
use std::{
|
||||
net::{IpAddr, SocketAddr},
|
||||
sync::Arc,
|
||||
sync::Mutex,
|
||||
task::{Context, Poll},
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
use subtle::ConstantTimeEq;
|
||||
use tower::{Layer, Service};
|
||||
|
||||
const MAX_AUTH_FAILURES: u32 = 5;
|
||||
const FAILURE_WINDOW_SECS: u64 = 60;
|
||||
const BLOCK_SECS: u64 = 300;
|
||||
// 安全修复:限制 failure_map 最大条目数,防止内存 DoS
|
||||
const MAX_FAILURE_ENTRIES: usize = 10000;
|
||||
const ENTRY_EXPIRE_SECS: u64 = 3600;
|
||||
|
||||
struct FailureState {
|
||||
count: u32,
|
||||
window_start: Instant,
|
||||
blocked_until: Option<Instant>,
|
||||
last_access: Instant,
|
||||
}
|
||||
|
||||
fn failure_map() -> &'static Mutex<std::collections::HashMap<String, FailureState>> {
|
||||
static FAILURES: std::sync::OnceLock<Mutex<std::collections::HashMap<String, FailureState>>> =
|
||||
std::sync::OnceLock::new();
|
||||
FAILURES.get_or_init(|| Mutex::new(std::collections::HashMap::new()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn clear_auth_failure_state() {
|
||||
let mut map = failure_map().lock().unwrap();
|
||||
map.clear();
|
||||
}
|
||||
|
||||
/// Management API 认证层
|
||||
///
|
||||
/// 用于包装需要认证的管理端点
|
||||
@@ -98,6 +127,76 @@ impl<S> ManagementAuthService<S> {
|
||||
.get::<axum::extract::ConnectInfo<SocketAddr>>()
|
||||
.map(|ci| ci.0)
|
||||
}
|
||||
|
||||
fn get_client_id(req: &Request<Body>) -> String {
|
||||
// 安全修复:只使用真实的连接地址,不信任 X-Forwarded-For
|
||||
// X-Forwarded-For 可被伪造,用于绕过限速或导致 failure_map 无界增长
|
||||
if let Some(addr) = Self::get_client_addr(req) {
|
||||
return addr.ip().to_string();
|
||||
}
|
||||
"unknown".to_string()
|
||||
}
|
||||
|
||||
fn check_rate_limit(client_id: &str) -> bool {
|
||||
let now = Instant::now();
|
||||
let mut map = failure_map().lock().unwrap();
|
||||
if let Some(state) = map.get_mut(client_id) {
|
||||
state.last_access = now;
|
||||
if let Some(blocked_until) = state.blocked_until {
|
||||
if blocked_until > now {
|
||||
return false;
|
||||
}
|
||||
state.blocked_until = None;
|
||||
state.count = 0;
|
||||
state.window_start = now;
|
||||
}
|
||||
if now.duration_since(state.window_start).as_secs() > FAILURE_WINDOW_SECS {
|
||||
state.count = 0;
|
||||
state.window_start = now;
|
||||
}
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
fn record_failure(client_id: &str) {
|
||||
let now = Instant::now();
|
||||
let mut map = failure_map().lock().unwrap();
|
||||
|
||||
// 安全修复:容量保护,超过上限时清理长时间未访问的条目
|
||||
if map.len() > MAX_FAILURE_ENTRIES {
|
||||
map.retain(|_, state| {
|
||||
now.duration_since(state.last_access).as_secs() <= ENTRY_EXPIRE_SECS
|
||||
});
|
||||
}
|
||||
|
||||
let entry = map.entry(client_id.to_string()).or_insert(FailureState {
|
||||
count: 0,
|
||||
window_start: now,
|
||||
blocked_until: None,
|
||||
last_access: now,
|
||||
});
|
||||
entry.last_access = now;
|
||||
|
||||
if now.duration_since(entry.window_start).as_secs() > FAILURE_WINDOW_SECS {
|
||||
entry.count = 0;
|
||||
entry.window_start = now;
|
||||
entry.blocked_until = None;
|
||||
}
|
||||
|
||||
entry.count += 1;
|
||||
if entry.count >= MAX_AUTH_FAILURES {
|
||||
entry.blocked_until = Some(now + Duration::from_secs(BLOCK_SECS));
|
||||
}
|
||||
}
|
||||
|
||||
fn record_success(client_id: &str) {
|
||||
let mut map = failure_map().lock().unwrap();
|
||||
map.remove(client_id);
|
||||
}
|
||||
|
||||
fn secret_key_matches(provided: &str, expected: &str) -> bool {
|
||||
provided.as_bytes().ct_eq(expected.as_bytes()).into()
|
||||
}
|
||||
}
|
||||
|
||||
impl<S> Service<Request<Body>> for ManagementAuthService<S>
|
||||
@@ -118,6 +217,14 @@ where
|
||||
let mut inner = self.inner.clone();
|
||||
|
||||
Box::pin(async move {
|
||||
let client_id = Self::get_client_id(&req);
|
||||
if !Self::check_rate_limit(&client_id) {
|
||||
return Ok(create_error_response(
|
||||
StatusCode::TOO_MANY_REQUESTS,
|
||||
"Too many failed authentication attempts",
|
||||
));
|
||||
}
|
||||
|
||||
// 1. 检查 secret_key 是否为空(禁用管理 API)
|
||||
let secret_key = match &config.secret_key {
|
||||
Some(key) if !key.is_empty() => key.clone(),
|
||||
@@ -148,9 +255,10 @@ where
|
||||
// 3. 验证 secret_key
|
||||
let provided_key = Self::extract_secret_key(&req);
|
||||
match provided_key {
|
||||
Some(key) if key == secret_key => {
|
||||
Some(key) if Self::secret_key_matches(&key, &secret_key) => {
|
||||
// 认证成功,继续处理请求
|
||||
tracing::debug!("[MANAGEMENT_AUTH] Auth successful from {:?}", client_addr);
|
||||
Self::record_success(&client_id);
|
||||
inner.call(req).await
|
||||
}
|
||||
Some(_) => {
|
||||
@@ -158,6 +266,7 @@ where
|
||||
"[MANAGEMENT_AUTH] Invalid secret_key from {:?}",
|
||||
client_addr
|
||||
);
|
||||
Self::record_failure(&client_id);
|
||||
Ok(create_error_response(
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"Invalid secret key",
|
||||
@@ -168,6 +277,7 @@ where
|
||||
"[MANAGEMENT_AUTH] Missing secret_key from {:?}",
|
||||
client_addr
|
||||
);
|
||||
Self::record_failure(&client_id);
|
||||
Ok(create_error_response(
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"Missing secret key",
|
||||
|
||||
@@ -3,9 +3,12 @@
|
||||
//! 使用 proptest 进行属性测试
|
||||
|
||||
use crate::config::RemoteManagementConfig;
|
||||
use crate::middleware::management_auth::{ManagementAuthLayer, ManagementAuthService};
|
||||
use crate::middleware::management_auth::{
|
||||
clear_auth_failure_state, ManagementAuthLayer, ManagementAuthService,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
extract::ConnectInfo,
|
||||
http::{Request, Response, StatusCode},
|
||||
};
|
||||
use proptest::prelude::*;
|
||||
@@ -102,6 +105,54 @@ fn create_request_with_management_key(key: Option<&str>) -> Request<Body> {
|
||||
builder.body(Body::empty()).unwrap()
|
||||
}
|
||||
|
||||
/// Helper to create a request with X-Management-Key and X-Forwarded-For headers
|
||||
fn create_request_with_management_key_and_forwarded(
|
||||
key: Option<&str>,
|
||||
forwarded_for: Option<&str>,
|
||||
) -> Request<Body> {
|
||||
let mut builder = Request::builder().uri("/v0/management/status");
|
||||
|
||||
if let Some(k) = key {
|
||||
builder = builder.header("x-management-key", k);
|
||||
}
|
||||
|
||||
if let Some(addr) = forwarded_for {
|
||||
builder = builder.header("x-forwarded-for", addr);
|
||||
}
|
||||
|
||||
builder.body(Body::empty()).unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_management_auth_rate_limit_after_failures() {
|
||||
clear_auth_failure_state();
|
||||
let config = RemoteManagementConfig {
|
||||
allow_remote: true,
|
||||
secret_key: Some("valid_key".to_string()),
|
||||
disable_control_panel: false,
|
||||
};
|
||||
let layer = ManagementAuthLayer::new(config);
|
||||
let mut service = layer.layer(MockService);
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
|
||||
// 使用唯一的 IP 地址避免测试间干扰
|
||||
let client_ip = format!("203.0.113.{}", std::process::id() % 256);
|
||||
let addr: SocketAddr = format!("{}:12345", client_ip).parse().unwrap();
|
||||
|
||||
for _ in 0..5 {
|
||||
let mut req = create_request_with_management_key(Some("invalid"));
|
||||
// 安全修复后不再信任 X-Forwarded-For,需要注入 ConnectInfo
|
||||
req.extensions_mut().insert(ConnectInfo(addr));
|
||||
let response = rt.block_on(async { service.call(req).await.unwrap() });
|
||||
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
let mut req = create_request_with_management_key(Some("invalid"));
|
||||
req.extensions_mut().insert(ConnectInfo(addr));
|
||||
let response = rt.block_on(async { service.call(req).await.unwrap() });
|
||||
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
@@ -112,6 +163,7 @@ proptest! {
|
||||
fn prop_management_auth_rejection_missing_key(
|
||||
secret_key in arb_secret_key()
|
||||
) {
|
||||
clear_auth_failure_state();
|
||||
// Create config with a valid secret_key
|
||||
let config = RemoteManagementConfig {
|
||||
allow_remote: true,
|
||||
@@ -147,6 +199,7 @@ proptest! {
|
||||
fn prop_management_auth_rejection_invalid_key(
|
||||
secret_key in arb_secret_key()
|
||||
) {
|
||||
clear_auth_failure_state();
|
||||
// Create config with a valid secret_key
|
||||
let config = RemoteManagementConfig {
|
||||
allow_remote: true,
|
||||
@@ -183,6 +236,7 @@ proptest! {
|
||||
fn prop_management_auth_acceptance_valid_key(
|
||||
secret_key in arb_secret_key()
|
||||
) {
|
||||
clear_auth_failure_state();
|
||||
// Create config with a valid secret_key
|
||||
let config = RemoteManagementConfig {
|
||||
allow_remote: true,
|
||||
@@ -218,6 +272,7 @@ proptest! {
|
||||
fn prop_management_auth_acceptance_x_management_key(
|
||||
secret_key in arb_secret_key()
|
||||
) {
|
||||
clear_auth_failure_state();
|
||||
// Create config with a valid secret_key
|
||||
let config = RemoteManagementConfig {
|
||||
allow_remote: true,
|
||||
@@ -253,6 +308,7 @@ proptest! {
|
||||
fn prop_management_auth_rejection_invalid_x_management_key(
|
||||
secret_key in arb_secret_key()
|
||||
) {
|
||||
clear_auth_failure_state();
|
||||
// Create config with a valid secret_key
|
||||
let config = RemoteManagementConfig {
|
||||
allow_remote: true,
|
||||
|
||||
@@ -72,7 +72,12 @@ pub enum CredentialData {
|
||||
excluded_models: Vec<String>,
|
||||
},
|
||||
/// Codex OAuth 凭证(OpenAI Codex)
|
||||
CodexOAuth { creds_file_path: String },
|
||||
CodexOAuth {
|
||||
creds_file_path: String,
|
||||
/// API Base URL(可选,默认使用凭证文件中的配置)
|
||||
#[serde(default)]
|
||||
api_base_url: Option<String>,
|
||||
},
|
||||
/// Claude OAuth 凭证(Anthropic OAuth)
|
||||
ClaudeOAuth { creds_file_path: String },
|
||||
/// iFlow OAuth 凭证
|
||||
@@ -113,7 +118,9 @@ impl CredentialData {
|
||||
CredentialData::GeminiApiKey { api_key, .. } => {
|
||||
format!("Gemini API Key: {}", mask_key(api_key))
|
||||
}
|
||||
CredentialData::CodexOAuth { creds_file_path } => {
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path, ..
|
||||
} => {
|
||||
format!("Codex OAuth: {}", mask_path(creds_file_path))
|
||||
}
|
||||
CredentialData::ClaudeOAuth { creds_file_path } => {
|
||||
@@ -550,7 +557,9 @@ pub fn get_oauth_creds_path(cred: &CredentialData) -> Option<String> {
|
||||
CredentialData::AntigravityOAuth {
|
||||
creds_file_path, ..
|
||||
} => Some(creds_file_path.clone()),
|
||||
CredentialData::CodexOAuth { creds_file_path } => Some(creds_file_path.clone()),
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path, ..
|
||||
} => Some(creds_file_path.clone()),
|
||||
CredentialData::ClaudeOAuth { creds_file_path } => Some(creds_file_path.clone()),
|
||||
CredentialData::IFlowOAuth { creds_file_path } => Some(creds_file_path.clone()),
|
||||
CredentialData::IFlowCookie { creds_file_path } => Some(creds_file_path.clone()),
|
||||
|
||||
@@ -126,6 +126,18 @@ impl ProcessError {
|
||||
ProcessError::Cancelled => "cancelled",
|
||||
}
|
||||
}
|
||||
|
||||
/// 记录带上下文的错误日志
|
||||
pub fn log_with_context(&self, request_id: &str, provider: &str, model: &str) {
|
||||
tracing::error!(
|
||||
request_id = %request_id,
|
||||
provider = %provider,
|
||||
model = %model,
|
||||
error_type = %self.error_type(),
|
||||
error_message = %self.to_string(),
|
||||
"Request processing failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
@@ -58,6 +58,8 @@ pub struct RequestProcessor {
|
||||
pub tokens: Arc<ParkingLotRwLock<TokenTracker>>,
|
||||
/// 凭证池服务
|
||||
pub pool_service: Arc<ProviderPoolService>,
|
||||
/// 热重载协调锁(避免配置更新期间请求读取不一致的配置)
|
||||
pub reload_lock: Arc<RwLock<()>>,
|
||||
}
|
||||
|
||||
impl RequestProcessor {
|
||||
@@ -85,6 +87,7 @@ impl RequestProcessor {
|
||||
stats,
|
||||
tokens,
|
||||
pool_service,
|
||||
reload_lock: Arc::new(RwLock::new(())),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -102,6 +105,7 @@ impl RequestProcessor {
|
||||
stats: Arc::new(ParkingLotRwLock::new(StatsAggregator::with_defaults())),
|
||||
tokens: Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults())),
|
||||
pool_service,
|
||||
reload_lock: Arc::new(RwLock::new(())),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -126,6 +130,7 @@ impl RequestProcessor {
|
||||
stats,
|
||||
tokens,
|
||||
pool_service,
|
||||
reload_lock: Arc::new(RwLock::new(())),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
use super::traits::{PipelineStep, StepError};
|
||||
use crate::processor::RequestContext;
|
||||
use async_trait::async_trait;
|
||||
use subtle::ConstantTimeEq;
|
||||
|
||||
/// 认证步骤
|
||||
///
|
||||
@@ -34,7 +35,7 @@ impl AuthStep {
|
||||
/// 验证 API Key
|
||||
pub fn verify(&self, provided_key: Option<&str>) -> Result<(), StepError> {
|
||||
match provided_key {
|
||||
Some(key) if key == self.expected_key => Ok(()),
|
||||
Some(key) if key.as_bytes().ct_eq(self.expected_key.as_bytes()).into() => Ok(()),
|
||||
Some(_) => Err(StepError::Auth("Invalid API key".to_string())),
|
||||
None => Err(StepError::Auth("No API key provided".to_string())),
|
||||
}
|
||||
|
||||
@@ -1567,3 +1567,159 @@ impl CredentialProvider for AntigravityProvider {
|
||||
"antigravity"
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// StreamingProvider Trait 实现
|
||||
// ============================================================================
|
||||
|
||||
use crate::models::openai::ChatCompletionRequest;
|
||||
use crate::providers::ProviderError;
|
||||
use crate::streaming::traits::{
|
||||
reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider,
|
||||
};
|
||||
|
||||
#[async_trait]
|
||||
impl StreamingProvider for AntigravityProvider {
|
||||
/// 发起流式 API 调用
|
||||
///
|
||||
/// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。
|
||||
/// Antigravity 使用 Gemini 流式格式。
|
||||
///
|
||||
/// # 需求覆盖
|
||||
/// - 需求 1.4: AntigravityProvider 流式支持
|
||||
async fn call_api_stream(
|
||||
&self,
|
||||
request: &ChatCompletionRequest,
|
||||
) -> Result<StreamResponse, ProviderError> {
|
||||
let token = self
|
||||
.credentials
|
||||
.access_token
|
||||
.as_ref()
|
||||
.ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?;
|
||||
|
||||
let project_id = self.project_id.clone().unwrap_or_else(generate_project_id);
|
||||
let actual_model = alias_to_model_name(&request.model);
|
||||
|
||||
// 构建 Antigravity 请求体
|
||||
// 将 OpenAI 格式转换为 Gemini/Antigravity 格式
|
||||
let mut contents = Vec::new();
|
||||
let mut system_instruction = None;
|
||||
|
||||
for msg in &request.messages {
|
||||
let role = &msg.role;
|
||||
let text = match &msg.content {
|
||||
Some(crate::models::openai::MessageContent::Text(t)) => t.clone(),
|
||||
Some(crate::models::openai::MessageContent::Parts(parts)) => parts
|
||||
.iter()
|
||||
.filter_map(|p| {
|
||||
if let crate::models::openai::ContentPart::Text { text } = p {
|
||||
Some(text.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(""),
|
||||
None => String::new(),
|
||||
};
|
||||
|
||||
if role == "system" {
|
||||
system_instruction = Some(serde_json::json!({
|
||||
"parts": [{ "text": text }]
|
||||
}));
|
||||
} else {
|
||||
let gemini_role = if role == "assistant" { "model" } else { "user" };
|
||||
contents.push(serde_json::json!({
|
||||
"role": gemini_role,
|
||||
"parts": [{ "text": text }]
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
let mut request_body = serde_json::json!({
|
||||
"contents": contents
|
||||
});
|
||||
|
||||
if let Some(sys) = system_instruction {
|
||||
request_body["systemInstruction"] = sys;
|
||||
}
|
||||
|
||||
// 构建 Antigravity 请求
|
||||
let mut payload = request_body.clone();
|
||||
payload["model"] = serde_json::json!(actual_model);
|
||||
payload["userAgent"] = serde_json::json!("antigravity");
|
||||
payload["project"] = serde_json::json!(project_id);
|
||||
payload["requestId"] = serde_json::json!(generate_request_id());
|
||||
|
||||
if payload.get("request").is_none() {
|
||||
payload["request"] = serde_json::json!({});
|
||||
}
|
||||
payload["request"]["sessionId"] = serde_json::json!(generate_session_id());
|
||||
|
||||
// 尝试多个 base URL
|
||||
let mut last_error: Option<ProviderError> = None;
|
||||
|
||||
for base_url in &self.base_urls {
|
||||
let url = format!(
|
||||
"{}/{ANTIGRAVITY_API_VERSION}:streamGenerateContent",
|
||||
base_url
|
||||
);
|
||||
|
||||
tracing::info!(
|
||||
"[ANTIGRAVITY_STREAM] 发起流式请求: url={} model={}",
|
||||
url,
|
||||
actual_model
|
||||
);
|
||||
|
||||
let result = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {token}"))
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "text/event-stream")
|
||||
.header("User-Agent", "antigravity/1.11.5 windows/amd64")
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await;
|
||||
|
||||
match result {
|
||||
Ok(resp) => {
|
||||
let status = resp.status();
|
||||
if status.is_success() {
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] 流式响应开始: status={}", status);
|
||||
return Ok(reqwest_stream_to_stream_response(resp));
|
||||
} else {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::warn!(
|
||||
"[ANTIGRAVITY_STREAM] 请求失败 ({}): {} - {}",
|
||||
base_url,
|
||||
status,
|
||||
body
|
||||
);
|
||||
last_error = Some(ProviderError::from_http_status(status.as_u16(), &body));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("[ANTIGRAVITY_STREAM] 连接失败 ({}): {}", base_url, e);
|
||||
last_error = Some(ProviderError::from_reqwest_error(&e));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Err(last_error.unwrap_or_else(|| {
|
||||
ProviderError::NetworkError("All Antigravity base URLs failed".to_string())
|
||||
}))
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
self.credentials.access_token.is_some() && self.credentials.enable != Some(false)
|
||||
}
|
||||
|
||||
fn provider_name(&self) -> &'static str {
|
||||
"AntigravityProvider"
|
||||
}
|
||||
|
||||
fn stream_format(&self) -> StreamFormat {
|
||||
StreamFormat::GeminiStream
|
||||
}
|
||||
}
|
||||
|
||||
@@ -319,3 +319,127 @@ impl ClaudeCustomProvider {
|
||||
Ok(data)
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// StreamingProvider Trait 实现
|
||||
// ============================================================================
|
||||
|
||||
use crate::providers::ProviderError;
|
||||
use crate::streaming::traits::{
|
||||
reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
|
||||
#[async_trait]
|
||||
impl StreamingProvider for ClaudeCustomProvider {
|
||||
/// 发起流式 API 调用
|
||||
///
|
||||
/// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。
|
||||
/// Claude 使用 Anthropic SSE 格式。
|
||||
///
|
||||
/// # 需求覆盖
|
||||
/// - 需求 1.2: ClaudeCustomProvider 流式支持
|
||||
async fn call_api_stream(
|
||||
&self,
|
||||
request: &ChatCompletionRequest,
|
||||
) -> Result<StreamResponse, ProviderError> {
|
||||
let api_key = self.config.api_key.as_ref().ok_or_else(|| {
|
||||
ProviderError::ConfigurationError("Claude API key not configured".to_string())
|
||||
})?;
|
||||
|
||||
// 转换 OpenAI 请求为 Anthropic 格式
|
||||
let mut anthropic_messages = Vec::new();
|
||||
let mut system_content = None;
|
||||
|
||||
for msg in &request.messages {
|
||||
let role = &msg.role;
|
||||
|
||||
// 提取消息内容
|
||||
let content = match &msg.content {
|
||||
Some(MessageContent::Text(text)) => text.clone(),
|
||||
Some(MessageContent::Parts(parts)) => parts
|
||||
.iter()
|
||||
.filter_map(|p| {
|
||||
if let ContentPart::Text { text } = p {
|
||||
Some(text.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(""),
|
||||
None => String::new(),
|
||||
};
|
||||
|
||||
if role == "system" {
|
||||
system_content = Some(content);
|
||||
} else {
|
||||
let anthropic_role = if role == "assistant" {
|
||||
"assistant"
|
||||
} else {
|
||||
"user"
|
||||
};
|
||||
anthropic_messages.push(serde_json::json!({
|
||||
"role": anthropic_role,
|
||||
"content": content
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
let mut anthropic_body = serde_json::json!({
|
||||
"model": request.model,
|
||||
"max_tokens": request.max_tokens.unwrap_or(4096),
|
||||
"messages": anthropic_messages,
|
||||
"stream": true
|
||||
});
|
||||
|
||||
if let Some(sys) = system_content {
|
||||
anthropic_body["system"] = serde_json::json!(sys);
|
||||
}
|
||||
|
||||
let url = self.build_url("messages");
|
||||
|
||||
tracing::info!(
|
||||
"[CLAUDE_STREAM] 发起流式请求: url={} model={}",
|
||||
url,
|
||||
request.model
|
||||
);
|
||||
|
||||
let resp = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("x-api-key", api_key)
|
||||
.header("anthropic-version", "2023-06-01")
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "text/event-stream")
|
||||
.json(&anthropic_body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| ProviderError::from_reqwest_error(&e))?;
|
||||
|
||||
// 检查响应状态
|
||||
let status = resp.status();
|
||||
if !status.is_success() {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::error!("[CLAUDE_STREAM] 请求失败: {} - {}", status, body);
|
||||
return Err(ProviderError::from_http_status(status.as_u16(), &body));
|
||||
}
|
||||
|
||||
tracing::info!("[CLAUDE_STREAM] 流式响应开始: status={}", status);
|
||||
|
||||
// 将 reqwest 响应转换为 StreamResponse
|
||||
Ok(reqwest_stream_to_stream_response(resp))
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
self.is_configured()
|
||||
}
|
||||
|
||||
fn provider_name(&self) -> &'static str {
|
||||
"ClaudeCustomProvider"
|
||||
}
|
||||
|
||||
fn stream_format(&self) -> StreamFormat {
|
||||
StreamFormat::AnthropicSse
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@ const OPENAI_TOKEN_URL: &str = "https://auth.openai.com/oauth/token";
|
||||
const OPENAI_CLIENT_ID: &str = "app_EMoamEEZ73f0CkXaXp7hrann";
|
||||
const DEFAULT_CALLBACK_PORT: u16 = 1455;
|
||||
const CODEX_API_BASE_URL: &str = "https://chatgpt.com/backend-api/codex";
|
||||
const DEFAULT_API_BASE_URL: &str = "https://api.openai.com";
|
||||
|
||||
/// Codex OAuth credentials storage
|
||||
///
|
||||
@@ -26,6 +27,10 @@ const CODEX_API_BASE_URL: &str = "https://chatgpt.com/backend-api/codex";
|
||||
/// Supports multiple field name formats:
|
||||
/// - snake_case: `refresh_token`, `access_token`, `id_token`, `account_id`, `last_refresh`
|
||||
/// - camelCase: `refreshToken`, `accessToken`, `idToken`, `accountId`, `lastRefresh`
|
||||
///
|
||||
/// 同时兼容 Codex CLI 的 API Key 登录格式:
|
||||
/// - `api_key` / `apiKey`
|
||||
/// - `api_base_url` / `apiBaseUrl`
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CodexCredentials {
|
||||
/// JWT ID token containing user claims
|
||||
@@ -45,6 +50,18 @@ pub struct CodexCredentials {
|
||||
alias = "refreshToken"
|
||||
)]
|
||||
pub refresh_token: Option<String>,
|
||||
/// API Key(Codex CLI 支持通过 API Key 登录)
|
||||
/// 支持字段名: api_key, apiKey, OPENAI_API_KEY
|
||||
#[serde(
|
||||
default,
|
||||
skip_serializing_if = "Option::is_none",
|
||||
alias = "apiKey",
|
||||
alias = "OPENAI_API_KEY"
|
||||
)]
|
||||
pub api_key: Option<String>,
|
||||
/// API Base URL(可选)
|
||||
#[serde(default, skip_serializing_if = "Option::is_none", alias = "apiBaseUrl")]
|
||||
pub api_base_url: Option<String>,
|
||||
/// OpenAI account identifier
|
||||
#[serde(default, skip_serializing_if = "Option::is_none", alias = "accountId")]
|
||||
pub account_id: Option<String>,
|
||||
@@ -66,8 +83,7 @@ pub struct CodexCredentials {
|
||||
#[serde(
|
||||
default,
|
||||
skip_serializing_if = "Option::is_none",
|
||||
rename = "expired",
|
||||
alias = "expires_at",
|
||||
alias = "expired",
|
||||
alias = "expiresAt"
|
||||
)]
|
||||
pub expires_at: Option<String>,
|
||||
@@ -83,6 +99,8 @@ impl Default for CodexCredentials {
|
||||
id_token: None,
|
||||
access_token: None,
|
||||
refresh_token: None,
|
||||
api_key: None,
|
||||
api_base_url: None,
|
||||
account_id: None,
|
||||
last_refresh: None,
|
||||
email: None,
|
||||
@@ -428,6 +446,38 @@ impl CodexProvider {
|
||||
CODEX_API_BASE_URL
|
||||
}
|
||||
|
||||
/// 获取已配置的 API Key(trim 后的非空值)
|
||||
fn get_api_key(&self) -> Option<&str> {
|
||||
self.credentials
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(|s| s.trim())
|
||||
.filter(|s| !s.is_empty())
|
||||
}
|
||||
|
||||
pub(crate) fn build_responses_url(base_url: &str) -> String {
|
||||
let base = base_url.trim_end_matches('/');
|
||||
|
||||
// 规则说明:
|
||||
// - 如果 base_url 以 /v1 结尾:直接拼 /responses
|
||||
// - 如果 base_url 只有域名(path 为空或 /):拼 /v1/responses(OpenAI 标准)
|
||||
// - 如果 base_url 已包含路径前缀(如 https://yunyi.cfd/codex):认为前缀已包含路由信息,拼 /responses
|
||||
if base.ends_with("/v1") {
|
||||
return format!("{}/responses", base);
|
||||
}
|
||||
|
||||
if let Ok(parsed) = url::Url::parse(base) {
|
||||
let path = parsed.path().trim_end_matches('/');
|
||||
if path.is_empty() || path == "/" {
|
||||
return format!("{}/v1/responses", base);
|
||||
}
|
||||
return format!("{}/responses", base);
|
||||
}
|
||||
|
||||
// 兜底:保持旧行为
|
||||
format!("{}/v1/responses", base)
|
||||
}
|
||||
|
||||
/// Load credentials from the default path
|
||||
pub async fn load_credentials(&mut self) -> Result<(), Box<dyn Error + Send + Sync>> {
|
||||
let path = Self::default_creds_path();
|
||||
@@ -457,9 +507,14 @@ impl CodexProvider {
|
||||
})?;
|
||||
|
||||
// 检查关键字段
|
||||
if creds.refresh_token.is_none() {
|
||||
let has_api_key = creds
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(|s| !s.trim().is_empty())
|
||||
.unwrap_or(false);
|
||||
if creds.refresh_token.is_none() && !has_api_key {
|
||||
tracing::warn!(
|
||||
"[CODEX] 凭证文件缺少 refresh_token 字段。支持的字段名: refresh_token, refreshToken"
|
||||
"[CODEX] 凭证文件缺少 refresh_token/api_key 字段。支持的字段名: refresh_token, refreshToken, api_key, apiKey"
|
||||
);
|
||||
// 打印文件中的顶级字段名,帮助调试
|
||||
if let Ok(json_value) = serde_json::from_str::<serde_json::Value>(&content) {
|
||||
@@ -471,9 +526,10 @@ impl CodexProvider {
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
"[CODEX] 凭证加载成功: has_access={}, has_refresh={}, email={:?}, path={:?}",
|
||||
"[CODEX] 凭证加载成功: has_access={}, has_refresh={}, has_api_key={}, email={:?}, path={:?}",
|
||||
creds.access_token.is_some(),
|
||||
creds.refresh_token.is_some(),
|
||||
has_api_key,
|
||||
creds.email,
|
||||
path
|
||||
);
|
||||
@@ -505,6 +561,11 @@ impl CodexProvider {
|
||||
|
||||
/// Check if the access token is expired
|
||||
pub fn is_token_expired(&self) -> bool {
|
||||
// API Key 模式:不涉及过期概念
|
||||
if self.get_api_key().is_some() {
|
||||
return false;
|
||||
}
|
||||
|
||||
if let Some(expires_str) = &self.credentials.expires_at {
|
||||
if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expires_str) {
|
||||
let now = chrono::Utc::now();
|
||||
@@ -518,6 +579,9 @@ impl CodexProvider {
|
||||
|
||||
/// Check if credentials are valid (has access token and not expired)
|
||||
pub fn is_valid(&self) -> bool {
|
||||
if self.get_api_key().is_some() {
|
||||
return true;
|
||||
}
|
||||
self.credentials.access_token.is_some() && !self.is_token_expired()
|
||||
}
|
||||
|
||||
@@ -608,6 +672,8 @@ impl CodexProvider {
|
||||
id_token,
|
||||
access_token: Some(access_token),
|
||||
refresh_token,
|
||||
api_key: None,
|
||||
api_base_url: None,
|
||||
account_id,
|
||||
last_refresh: Some(chrono::Utc::now().to_rfc3339()),
|
||||
email,
|
||||
@@ -626,13 +692,55 @@ impl CodexProvider {
|
||||
}
|
||||
|
||||
/// Refresh the access token using the refresh token
|
||||
///
|
||||
/// Supports three authentication modes (in priority order):
|
||||
/// 1. **API Key Mode**: Returns the API key directly (no refresh needed)
|
||||
/// 2. **OAuth Mode**: Refreshes the access token using the refresh token
|
||||
/// 3. **Access Token Mode**: Returns the existing access token (may be expired)
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(String)` - The access token or API key
|
||||
/// * `Err` - If no credentials are available
|
||||
///
|
||||
/// # Examples
|
||||
/// ```ignore
|
||||
/// // API Key mode
|
||||
/// provider.credentials.api_key = Some("sk-test".to_string());
|
||||
/// let token = provider.refresh_token().await?; // Returns "sk-test"
|
||||
///
|
||||
/// // OAuth mode
|
||||
/// provider.credentials.refresh_token = Some("refresh_token".to_string());
|
||||
/// let token = provider.refresh_token().await?; // Refreshes and returns new access_token
|
||||
///
|
||||
/// // Access Token mode (fallback)
|
||||
/// provider.credentials.access_token = Some("access_token".to_string());
|
||||
/// let token = provider.refresh_token().await?; // Returns "access_token" (with warning)
|
||||
/// ```
|
||||
pub async fn refresh_token(&mut self) -> Result<String, Box<dyn Error + Send + Sync>> {
|
||||
let refresh_token = self.credentials.refresh_token.as_ref().ok_or_else(|| {
|
||||
create_config_error(
|
||||
"没有可用的 refresh_token。请确保凭证文件包含 refresh_token 或 refreshToken 字段,\
|
||||
或使用 OAuth 登录功能重新获取凭证",
|
||||
)
|
||||
})?;
|
||||
// 1. API Key 模式无需刷新(优先级最高)
|
||||
if let Some(api_key) = self.get_api_key() {
|
||||
return Ok(api_key.to_string());
|
||||
}
|
||||
|
||||
// 2. 无 refresh_token 时的降级处理
|
||||
if self.credentials.refresh_token.is_none() {
|
||||
// 2a. 有 access_token:返回(可能过期,由上层处理)
|
||||
if let Some(ref access_token) = self.credentials.access_token {
|
||||
tracing::warn!("[CODEX] 没有 refresh_token,返回现有 access_token(可能已过期)");
|
||||
return Ok(access_token.clone());
|
||||
}
|
||||
|
||||
// 2b. 无任何凭证:清晰的错误指导
|
||||
return Err(create_config_error(
|
||||
"没有可用的认证凭证。请配置以下任一方式:\n\
|
||||
1. API Key 模式:在凭证文件中添加 api_key/apiKey 字段\n\
|
||||
2. OAuth 模式:使用 OAuth 登录获取 refresh_token\n\
|
||||
3. Access Token 模式:在凭证文件中添加 access_token/accessToken 字段",
|
||||
));
|
||||
}
|
||||
|
||||
// 3. OAuth 刷新流程(标准流程)
|
||||
let refresh_token = self.credentials.refresh_token.as_ref().unwrap();
|
||||
|
||||
tracing::info!("[CODEX] 正在刷新 access token");
|
||||
|
||||
@@ -768,6 +876,11 @@ impl CodexProvider {
|
||||
|
||||
/// Check if token needs refresh (expiring within the specified duration)
|
||||
pub fn needs_refresh(&self, lead_time: chrono::Duration) -> bool {
|
||||
// API Key 模式无需刷新
|
||||
if self.get_api_key().is_some() {
|
||||
return false;
|
||||
}
|
||||
|
||||
if self.credentials.access_token.is_none() {
|
||||
return true;
|
||||
}
|
||||
@@ -788,6 +901,11 @@ impl CodexProvider {
|
||||
/// This is the recommended method to call before making API requests.
|
||||
/// It will automatically refresh the token if it's expired or about to expire.
|
||||
pub async fn ensure_valid_token(&mut self) -> Result<String, Box<dyn Error + Send + Sync>> {
|
||||
// 兼容 Codex CLI 的 API Key 登录:auth.json 只有 api_key,没有 refresh_token
|
||||
if let Some(api_key) = self.get_api_key() {
|
||||
return Ok(api_key.to_string());
|
||||
}
|
||||
|
||||
// Refresh if token expires within 5 minutes
|
||||
let lead_time = chrono::Duration::minutes(5);
|
||||
|
||||
@@ -811,6 +929,11 @@ impl CodexProvider {
|
||||
|
||||
/// Get the access token, refreshing if necessary
|
||||
pub async fn get_access_token(&mut self) -> Result<String, Box<dyn Error + Send + Sync>> {
|
||||
// API Key 模式直接返回
|
||||
if let Some(api_key) = self.get_api_key() {
|
||||
return Ok(api_key.to_string());
|
||||
}
|
||||
|
||||
if self.is_token_expired() {
|
||||
self.refresh_token().await?;
|
||||
}
|
||||
@@ -919,42 +1042,99 @@ impl CodexProvider {
|
||||
&self,
|
||||
request: &serde_json::Value,
|
||||
) -> Result<reqwest::Response, Box<dyn Error + Send + Sync>> {
|
||||
let token = self
|
||||
.credentials
|
||||
.access_token
|
||||
.as_ref()
|
||||
.ok_or("No access token available")?;
|
||||
enum AuthMode {
|
||||
ApiKey,
|
||||
OAuth,
|
||||
}
|
||||
|
||||
let (token, mode) = match self.get_api_key() {
|
||||
Some(api_key) => (api_key, AuthMode::ApiKey),
|
||||
None => (
|
||||
self.credentials
|
||||
.access_token
|
||||
.as_deref()
|
||||
.ok_or("No access token or api_key available")?,
|
||||
AuthMode::OAuth,
|
||||
),
|
||||
};
|
||||
|
||||
// Build the Codex API URL
|
||||
let url = format!("{}/responses", CODEX_API_BASE_URL);
|
||||
let url = match mode {
|
||||
AuthMode::ApiKey => {
|
||||
let has_custom_base_url = self
|
||||
.credentials
|
||||
.api_base_url
|
||||
.as_deref()
|
||||
.map(|s| s.trim())
|
||||
.filter(|s| !s.is_empty())
|
||||
.is_some();
|
||||
|
||||
let base_url = self
|
||||
.credentials
|
||||
.api_base_url
|
||||
.as_deref()
|
||||
.map(|s| s.trim())
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or(DEFAULT_API_BASE_URL);
|
||||
|
||||
// Warn if API key doesn't look like OpenAI format but no custom base URL is set
|
||||
if !has_custom_base_url && !token.starts_with("sk-") {
|
||||
tracing::warn!(
|
||||
"[CODEX] API key does not appear to be an OpenAI key (doesn't start with 'sk-'), \
|
||||
but no api_base_url is configured. Requests will be sent to {}. \
|
||||
If you're using a third-party API provider, please add 'api_base_url' to ~/.codex/auth.json",
|
||||
DEFAULT_API_BASE_URL
|
||||
);
|
||||
}
|
||||
|
||||
Self::build_responses_url(base_url)
|
||||
}
|
||||
AuthMode::OAuth => format!("{}/responses", CODEX_API_BASE_URL),
|
||||
};
|
||||
|
||||
// Transform OpenAI chat completion request to Codex format
|
||||
let codex_request = transform_to_codex_format(request)?;
|
||||
|
||||
tracing::debug!("[CODEX] Calling API: {}", url);
|
||||
|
||||
let resp = self
|
||||
let mut req = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {}", token))
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "text/event-stream")
|
||||
.header("Version", "0.21.0")
|
||||
.header("Openai-Beta", "responses=experimental")
|
||||
.header(
|
||||
"User-Agent",
|
||||
"codex_cli_rs/0.50.0 (Mac OS 26.0.1; arm64) Apple_Terminal/464",
|
||||
)
|
||||
.header("Originator", "codex_cli_rs")
|
||||
.header("Session_id", uuid::Uuid::new_v4().to_string())
|
||||
// Add account ID header if available
|
||||
.header(
|
||||
"Chatgpt-Account-Id",
|
||||
self.credentials.account_id.as_deref().unwrap_or(""),
|
||||
)
|
||||
.json(&codex_request)
|
||||
.send()
|
||||
.await?;
|
||||
.json(&codex_request);
|
||||
|
||||
// 部分三方 Codex 代理(如 Yunyi)会依赖 Codex CLI 的特征 headers;
|
||||
// 仅在 OAuth 模式或显式配置了自定义 base_url 时附加,避免影响 OpenAI 官方 Key 模式。
|
||||
let should_add_codex_cli_headers = matches!(mode, AuthMode::OAuth)
|
||||
|| (matches!(mode, AuthMode::ApiKey)
|
||||
&& self
|
||||
.credentials
|
||||
.api_base_url
|
||||
.as_deref()
|
||||
.map(|s| !s.trim().is_empty())
|
||||
.unwrap_or(false));
|
||||
|
||||
if should_add_codex_cli_headers {
|
||||
req = req
|
||||
.header("Version", "0.21.0")
|
||||
.header(
|
||||
"User-Agent",
|
||||
"codex_cli_rs/0.50.0 (Mac OS 26.0.1; arm64) Apple_Terminal/464",
|
||||
)
|
||||
.header("Originator", "codex_cli_rs")
|
||||
.header("Session_id", uuid::Uuid::new_v4().to_string())
|
||||
.header("Conversation_id", uuid::Uuid::new_v4().to_string())
|
||||
// Add account ID header if available
|
||||
.header(
|
||||
"Chatgpt-Account-Id",
|
||||
self.credentials.account_id.as_deref().unwrap_or(""),
|
||||
);
|
||||
}
|
||||
|
||||
let resp = req.send().await?;
|
||||
|
||||
Ok(resp)
|
||||
}
|
||||
@@ -1161,6 +1341,7 @@ mod tests {
|
||||
let creds = CodexCredentials::default();
|
||||
assert!(creds.access_token.is_none());
|
||||
assert!(creds.refresh_token.is_none());
|
||||
assert!(creds.api_key.is_none());
|
||||
assert_eq!(creds.r#type, "codex");
|
||||
}
|
||||
|
||||
@@ -1230,6 +1411,32 @@ mod tests {
|
||||
assert_eq!(creds.expires_at, Some("2024-12-31T23:59:59Z".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_codex_credentials_api_key_fields() {
|
||||
let json = r#"{
|
||||
"api_key": "sk-test",
|
||||
"api_base_url": "https://api.openai.com/v1"
|
||||
}"#;
|
||||
|
||||
let creds: CodexCredentials = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(creds.api_key, Some("sk-test".to_string()));
|
||||
assert_eq!(
|
||||
creds.api_base_url,
|
||||
Some("https://api.openai.com/v1".to_string())
|
||||
);
|
||||
|
||||
let json2 = r#"{
|
||||
"apiKey": "sk-test-2",
|
||||
"apiBaseUrl": "https://example.com/v1"
|
||||
}"#;
|
||||
let creds2: CodexCredentials = serde_json::from_str(json2).unwrap();
|
||||
assert_eq!(creds2.api_key, Some("sk-test-2".to_string()));
|
||||
assert_eq!(
|
||||
creds2.api_base_url,
|
||||
Some("https://example.com/v1".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_codex_credentials_expires_at_alias() {
|
||||
// 测试 expires_at 字段的多种别名
|
||||
@@ -1262,6 +1469,35 @@ mod tests {
|
||||
assert!(provider.credentials.access_token.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_responses_url() {
|
||||
assert_eq!(
|
||||
CodexProvider::build_responses_url("https://api.openai.com"),
|
||||
"https://api.openai.com/v1/responses"
|
||||
);
|
||||
assert_eq!(
|
||||
CodexProvider::build_responses_url("https://api.openai.com/v1"),
|
||||
"https://api.openai.com/v1/responses"
|
||||
);
|
||||
assert_eq!(
|
||||
CodexProvider::build_responses_url("https://example.com/v1/"),
|
||||
"https://example.com/v1/responses"
|
||||
);
|
||||
assert_eq!(
|
||||
CodexProvider::build_responses_url("https://yunyi.cfd/codex"),
|
||||
"https://yunyi.cfd/codex/responses"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_ensure_valid_token_prefers_api_key() {
|
||||
let mut provider = CodexProvider::new();
|
||||
provider.credentials.api_key = Some("sk-test".to_string());
|
||||
|
||||
let token = provider.ensure_valid_token().await.unwrap();
|
||||
assert_eq!(token, "sk-test");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_generate_auth_url() {
|
||||
let provider = CodexProvider::new();
|
||||
@@ -1340,7 +1576,12 @@ mod tests {
|
||||
// No expiry - should be considered expired
|
||||
assert!(provider.is_token_expired());
|
||||
|
||||
// API Key 模式 - 不应视为过期
|
||||
provider.credentials.api_key = Some("sk-test".to_string());
|
||||
assert!(!provider.is_token_expired());
|
||||
|
||||
// Expired token
|
||||
provider.credentials.api_key = None;
|
||||
provider.credentials.expires_at = Some("2020-01-01T00:00:00Z".to_string());
|
||||
assert!(provider.is_token_expired());
|
||||
|
||||
@@ -1438,6 +1679,65 @@ mod tests {
|
||||
assert_eq!(result["max_output_tokens"], 1000);
|
||||
assert_eq!(result["top_p"], 0.9);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_token_with_only_access_token() {
|
||||
// 场景:只有 access_token(无 refresh_token 和 api_key)
|
||||
let mut provider = CodexProvider::new();
|
||||
provider.credentials.access_token = Some("test_access_token".to_string());
|
||||
provider.credentials.refresh_token = None;
|
||||
provider.credentials.api_key = None;
|
||||
|
||||
let result = provider.refresh_token().await;
|
||||
assert!(result.is_ok());
|
||||
assert_eq!(result.unwrap(), "test_access_token");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_token_with_no_credentials() {
|
||||
// 场景:无任何凭证(api_key、refresh_token、access_token 均为 None)
|
||||
let mut provider = CodexProvider::new();
|
||||
provider.credentials.api_key = None;
|
||||
provider.credentials.refresh_token = None;
|
||||
provider.credentials.access_token = None;
|
||||
|
||||
let result = provider.refresh_token().await;
|
||||
assert!(result.is_err());
|
||||
let error_msg = result.unwrap_err().to_string();
|
||||
assert!(error_msg.contains("没有可用的认证凭证"));
|
||||
assert!(error_msg.contains("API Key 模式"));
|
||||
assert!(error_msg.contains("OAuth 模式"));
|
||||
assert!(error_msg.contains("Access Token 模式"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_api_key_priority_over_refresh_token() {
|
||||
// 场景:同时有 api_key 和 refresh_token
|
||||
let mut provider = CodexProvider::new();
|
||||
provider.credentials.api_key = Some("sk-test-api-key".to_string());
|
||||
provider.credentials.refresh_token = Some("test_refresh_token".to_string());
|
||||
provider.credentials.access_token = Some("test_access_token".to_string());
|
||||
|
||||
let result = provider.refresh_token().await;
|
||||
assert!(result.is_ok());
|
||||
// 应该返回 API Key(优先级最高)
|
||||
assert_eq!(result.unwrap(), "sk-test-api-key");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_token_with_expired_access_token() {
|
||||
// 场景:只有 access_token(已过期)
|
||||
let mut provider = CodexProvider::new();
|
||||
provider.credentials.access_token = Some("expired_access_token".to_string());
|
||||
provider.credentials.expires_at = Some("2020-01-01T00:00:00Z".to_string());
|
||||
provider.credentials.refresh_token = None;
|
||||
provider.credentials.api_key = None;
|
||||
|
||||
let result = provider.refresh_token().await;
|
||||
assert!(result.is_ok());
|
||||
// 应该返回 access_token(即使已过期,由上层处理)
|
||||
assert_eq!(result.unwrap(), "expired_access_token");
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
@@ -1714,6 +2014,8 @@ pub async fn start_codex_oauth_server_and_get_url() -> Result<
|
||||
id_token,
|
||||
access_token: Some(access_token.to_string()),
|
||||
refresh_token,
|
||||
api_key: None,
|
||||
api_base_url: None,
|
||||
account_id,
|
||||
last_refresh: Some(now.to_rfc3339()),
|
||||
email: email.clone(),
|
||||
|
||||
+148
-63
@@ -322,27 +322,8 @@ impl KiroProvider {
|
||||
tracing::info!("[KIRO] 没有 clientIdHash 字段");
|
||||
}
|
||||
|
||||
// 读取目录中其他 JSON 文件
|
||||
if tokio::fs::try_exists(dir).await.unwrap_or(false) {
|
||||
let mut entries = tokio::fs::read_dir(dir).await?;
|
||||
while let Some(entry) = entries.next_entry().await? {
|
||||
let file_path = entry.path();
|
||||
if file_path.extension().map(|e| e == "json").unwrap_or(false) && file_path != path
|
||||
{
|
||||
if let Ok(content) = tokio::fs::read_to_string(&file_path).await {
|
||||
if let Ok(creds) = serde_json::from_str::<KiroCredentials>(&content) {
|
||||
tracing::info!(
|
||||
"[KIRO] Extra file {:?}: has_client_id={}, has_client_secret={}",
|
||||
file_path.file_name(),
|
||||
creds.client_id.is_some(),
|
||||
creds.client_secret.is_some()
|
||||
);
|
||||
merge_credentials(&mut merged, &creds);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// 安全修复:不再遍历目录中其他 JSON 文件,避免串凭证/串账号风险
|
||||
// 只信任主凭证文件和 clientIdHash 指向的文件
|
||||
|
||||
tracing::info!(
|
||||
"[KIRO] Final merged: has_access={}, has_refresh={}, has_client_id={}, has_client_secret={}, auth_method={:?}",
|
||||
@@ -625,11 +606,8 @@ impl KiroProvider {
|
||||
token_len < 100 || refresh_token.ends_with("...") || refresh_token.contains("...");
|
||||
|
||||
if is_truncated {
|
||||
tracing::error!(
|
||||
"[KIRO] 检测到 refreshToken 被截断!长度: {}, 内容: {}...",
|
||||
token_len,
|
||||
&refresh_token[..std::cmp::min(30, token_len)]
|
||||
);
|
||||
// 安全修复:不打印 token 内容,只打印长度
|
||||
tracing::error!("[KIRO] 检测到 refreshToken 被截断!长度: {}", token_len);
|
||||
return Err(format!(
|
||||
"refreshToken 已被截断(长度: {} 字符)。\n\n⚠️ 这通常是 Kiro IDE 为了防止凭证被第三方工具使用而故意截断的。\n\n💡 解决方案:\n1. 使用 Kir-Manager 工具获取完整的凭证\n2. 或者使用其他方式获取未截断的凭证文件\n3. 正常的 refreshToken 长度应该在 500+ 字符",
|
||||
token_len
|
||||
@@ -957,45 +935,49 @@ impl KiroProvider {
|
||||
let cw_request = convert_openai_to_codewhisperer(request, profile_arn.clone());
|
||||
let url = self.get_base_url();
|
||||
|
||||
// Debug: 记录转换后的请求
|
||||
if let Ok(json_str) = serde_json::to_string_pretty(&cw_request) {
|
||||
// 保存到文件用于调试
|
||||
let uuid_prefix = uuid::Uuid::new_v4()
|
||||
.to_string()
|
||||
.split('-')
|
||||
.next()
|
||||
.unwrap_or("unknown")
|
||||
.to_string();
|
||||
let debug_path = dirs::home_dir()
|
||||
.unwrap_or_default()
|
||||
.join(".proxycast")
|
||||
.join("logs")
|
||||
.join(format!("cw_request_{uuid_prefix}.json"));
|
||||
let _ = tokio::fs::write(&debug_path, &json_str).await;
|
||||
tracing::debug!("[CW_REQ] Request saved to {:?}", debug_path);
|
||||
|
||||
// 记录历史消息数量和 tool_results 情况
|
||||
let history_len = cw_request
|
||||
.conversation_state
|
||||
.history
|
||||
.as_ref()
|
||||
.map(|h| h.len())
|
||||
.unwrap_or(0);
|
||||
let current_has_tools = cw_request
|
||||
.conversation_state
|
||||
.current_message
|
||||
.user_input_message
|
||||
.user_input_message_context
|
||||
.as_ref()
|
||||
.map(|ctx| ctx.tool_results.as_ref().map(|tr| tr.len()).unwrap_or(0))
|
||||
.unwrap_or(0);
|
||||
tracing::info!(
|
||||
"[CW_REQ] history={} current_tool_results={}",
|
||||
history_len,
|
||||
current_has_tools
|
||||
);
|
||||
// 安全修复:仅在 PROXYCAST_DEBUG=1 时写入请求调试文件,避免泄露敏感信息
|
||||
let debug_enabled = std::env::var("PROXYCAST_DEBUG")
|
||||
.map(|v| v == "1")
|
||||
.unwrap_or(false);
|
||||
if debug_enabled {
|
||||
if let Ok(json_str) = serde_json::to_string_pretty(&cw_request) {
|
||||
let uuid_prefix = uuid::Uuid::new_v4()
|
||||
.to_string()
|
||||
.split('-')
|
||||
.next()
|
||||
.unwrap_or("unknown")
|
||||
.to_string();
|
||||
let debug_path = dirs::home_dir()
|
||||
.unwrap_or_default()
|
||||
.join(".proxycast")
|
||||
.join("logs")
|
||||
.join(format!("cw_request_{uuid_prefix}.json"));
|
||||
let _ = tokio::fs::write(&debug_path, &json_str).await;
|
||||
tracing::debug!("[CW_REQ] Request saved to {:?}", debug_path);
|
||||
}
|
||||
}
|
||||
|
||||
// 记录历史消息数量和 tool_results 情况(不落盘)
|
||||
let history_len = cw_request
|
||||
.conversation_state
|
||||
.history
|
||||
.as_ref()
|
||||
.map(|h| h.len())
|
||||
.unwrap_or(0);
|
||||
let current_has_tools = cw_request
|
||||
.conversation_state
|
||||
.current_message
|
||||
.user_input_message
|
||||
.user_input_message_context
|
||||
.as_ref()
|
||||
.map(|ctx| ctx.tool_results.as_ref().map(|tr| tr.len()).unwrap_or(0))
|
||||
.unwrap_or(0);
|
||||
tracing::info!(
|
||||
"[CW_REQ] history={} current_tool_results={}",
|
||||
history_len,
|
||||
current_has_tools
|
||||
);
|
||||
|
||||
// 生成基于凭证的唯一 Machine ID(关键改进:每个账号独立指纹)
|
||||
let machine_id = generate_machine_id_from_credentials(
|
||||
profile_arn.as_deref(),
|
||||
@@ -1113,3 +1095,106 @@ impl CredentialProvider for KiroProvider {
|
||||
"kiro"
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// StreamingProvider Trait 实现
|
||||
// ============================================================================
|
||||
|
||||
use crate::providers::ProviderError;
|
||||
use crate::streaming::traits::{
|
||||
reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider,
|
||||
};
|
||||
|
||||
#[async_trait]
|
||||
impl StreamingProvider for KiroProvider {
|
||||
/// 发起流式 API 调用
|
||||
///
|
||||
/// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。
|
||||
/// Kiro/CodeWhisperer 使用 AWS Event Stream 格式。
|
||||
///
|
||||
/// # 需求覆盖
|
||||
/// - 需求 1.1: KiroProvider 流式支持
|
||||
async fn call_api_stream(
|
||||
&self,
|
||||
request: &ChatCompletionRequest,
|
||||
) -> Result<StreamResponse, ProviderError> {
|
||||
let token = self
|
||||
.credentials
|
||||
.access_token
|
||||
.as_ref()
|
||||
.ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?;
|
||||
|
||||
let profile_arn = if self.credentials.auth_method.as_deref() == Some("social") {
|
||||
self.credentials.profile_arn.clone()
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let cw_request = convert_openai_to_codewhisperer(request, profile_arn.clone());
|
||||
let url = self.get_base_url();
|
||||
|
||||
// 生成基于凭证的唯一 Machine ID
|
||||
let machine_id = generate_machine_id_from_credentials(
|
||||
profile_arn.as_deref(),
|
||||
self.credentials.client_id.as_deref(),
|
||||
);
|
||||
let kiro_version = get_kiro_version();
|
||||
let (os_name, node_version) = get_system_runtime_info();
|
||||
|
||||
tracing::debug!(
|
||||
"[KIRO_STREAM] 发起流式请求: url={} machine_id={}...",
|
||||
url,
|
||||
&machine_id[..16]
|
||||
);
|
||||
|
||||
let resp = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {token}"))
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "application/vnd.amazon.eventstream")
|
||||
.header("amz-sdk-invocation-id", uuid::Uuid::new_v4().to_string())
|
||||
.header("amz-sdk-request", "attempt=1; max=1")
|
||||
.header("x-amzn-kiro-agent-mode", "vibe")
|
||||
.header(
|
||||
"x-amz-user-agent",
|
||||
format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}"),
|
||||
)
|
||||
.header(
|
||||
"user-agent",
|
||||
format!(
|
||||
"aws-sdk-js/1.0.0 ua/2.1 os/{os_name} lang/js md/nodejs#{node_version} api/codewhispererruntime#1.0.0 m/E KiroIDE-{kiro_version}-{machine_id}"
|
||||
),
|
||||
)
|
||||
.header("Connection", "close")
|
||||
.json(&cw_request)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| ProviderError::from_reqwest_error(&e))?;
|
||||
|
||||
// 检查响应状态
|
||||
let status = resp.status();
|
||||
if !status.is_success() {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::error!("[KIRO_STREAM] 请求失败: {} - {}", status, body);
|
||||
return Err(ProviderError::from_http_status(status.as_u16(), &body));
|
||||
}
|
||||
|
||||
tracing::info!("[KIRO_STREAM] 流式响应开始: status={}", status);
|
||||
|
||||
// 将 reqwest 响应转换为 StreamResponse
|
||||
Ok(reqwest_stream_to_stream_response(resp))
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn provider_name(&self) -> &'static str {
|
||||
"KiroProvider"
|
||||
}
|
||||
|
||||
fn stream_format(&self) -> StreamFormat {
|
||||
StreamFormat::AwsEventStream
|
||||
}
|
||||
}
|
||||
|
||||
@@ -143,3 +143,80 @@ impl OpenAICustomProvider {
|
||||
Ok(data)
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// StreamingProvider Trait 实现
|
||||
// ============================================================================
|
||||
|
||||
use crate::providers::ProviderError;
|
||||
use crate::streaming::traits::{
|
||||
reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
|
||||
#[async_trait]
|
||||
impl StreamingProvider for OpenAICustomProvider {
|
||||
/// 发起流式 API 调用
|
||||
///
|
||||
/// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。
|
||||
/// OpenAI 使用 OpenAI SSE 格式。
|
||||
///
|
||||
/// # 需求覆盖
|
||||
/// - 需求 1.3: OpenAICustomProvider 流式支持
|
||||
async fn call_api_stream(
|
||||
&self,
|
||||
request: &ChatCompletionRequest,
|
||||
) -> Result<StreamResponse, ProviderError> {
|
||||
let api_key = self.config.api_key.as_ref().ok_or_else(|| {
|
||||
ProviderError::ConfigurationError("OpenAI API key not configured".to_string())
|
||||
})?;
|
||||
|
||||
// 确保请求启用流式
|
||||
let mut stream_request = request.clone();
|
||||
stream_request.stream = true;
|
||||
|
||||
let url = self.build_url("chat/completions");
|
||||
|
||||
tracing::info!(
|
||||
"[OPENAI_STREAM] 发起流式请求: url={} model={}",
|
||||
url,
|
||||
request.model
|
||||
);
|
||||
|
||||
let resp = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {api_key}"))
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "text/event-stream")
|
||||
.json(&stream_request)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| ProviderError::from_reqwest_error(&e))?;
|
||||
|
||||
// 检查响应状态
|
||||
let status = resp.status();
|
||||
if !status.is_success() {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::error!("[OPENAI_STREAM] 请求失败: {} - {}", status, body);
|
||||
return Err(ProviderError::from_http_status(status.as_u16(), &body));
|
||||
}
|
||||
|
||||
tracing::info!("[OPENAI_STREAM] 流式响应开始: status={}", status);
|
||||
|
||||
// 将 reqwest 响应转换为 StreamResponse
|
||||
Ok(reqwest_stream_to_stream_response(resp))
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
self.is_configured()
|
||||
}
|
||||
|
||||
fn provider_name(&self) -> &'static str {
|
||||
"OpenAICustomProvider"
|
||||
}
|
||||
|
||||
fn stream_format(&self) -> StreamFormat {
|
||||
StreamFormat::OpenAiSse
|
||||
}
|
||||
}
|
||||
|
||||
@@ -138,8 +138,10 @@ impl QwenProvider {
|
||||
if let Some(expire_str) = &self.credentials.expire {
|
||||
if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) {
|
||||
let now = chrono::Utc::now();
|
||||
// 安全修复:显式转换为 Utc 时区再比较
|
||||
let expires_utc = expires.with_timezone(&chrono::Utc);
|
||||
// Token 有效期需要超过 30 秒
|
||||
return expires > now + chrono::Duration::seconds(30);
|
||||
return expires_utc > now + chrono::Duration::seconds(30);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -149,7 +151,8 @@ impl QwenProvider {
|
||||
return expiry > now + 30_000;
|
||||
}
|
||||
|
||||
true
|
||||
// 安全修复:没有过期时间时采用保守策略,认为 token 无效
|
||||
false
|
||||
}
|
||||
|
||||
pub fn get_base_url(&self) -> String {
|
||||
|
||||
@@ -20,6 +20,14 @@ fn arb_time_offset_secs() -> impl Strategy<Value = i64> {
|
||||
-3600i64..7200i64
|
||||
}
|
||||
|
||||
/// 生成不会与 lead_time 边界冲突的时间偏移
|
||||
/// 避免 time_offset_secs 恰好等于 lead_time_mins * 60 的情况
|
||||
fn arb_time_offset_avoiding_boundary(lead_time_mins: i64) -> impl Strategy<Value = i64> {
|
||||
let boundary = lead_time_mins * 60;
|
||||
// 生成不等于边界值的时间偏移
|
||||
(-3600i64..7200i64).prop_filter("避免边界值", move |&offset| offset != boundary)
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
@@ -36,11 +44,14 @@ proptest! {
|
||||
#[test]
|
||||
fn test_codex_token_refresh_timing(
|
||||
lead_time_mins in arb_lead_time_mins(),
|
||||
time_offset_secs in arb_time_offset_secs(),
|
||||
time_offset_secs in -3600i64..7200i64,
|
||||
) {
|
||||
let lead_time = Duration::minutes(lead_time_mins);
|
||||
let lead_time_secs = lead_time_mins * 60;
|
||||
|
||||
// 跳过边界条件,因为时间精度问题可能导致不确定行为
|
||||
prop_assume!(time_offset_secs != lead_time_secs);
|
||||
|
||||
let mut provider = CodexProvider::new();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
|
||||
@@ -72,11 +83,14 @@ proptest! {
|
||||
#[test]
|
||||
fn test_iflow_token_refresh_timing(
|
||||
lead_time_mins in arb_lead_time_mins(),
|
||||
time_offset_secs in arb_time_offset_secs(),
|
||||
time_offset_secs in -3600i64..7200i64,
|
||||
) {
|
||||
let lead_time = Duration::minutes(lead_time_mins);
|
||||
let lead_time_secs = lead_time_mins * 60;
|
||||
|
||||
// 跳过边界条件,因为时间精度问题可能导致不确定行为
|
||||
prop_assume!(time_offset_secs != lead_time_secs);
|
||||
|
||||
let mut provider = IFlowProvider::new();
|
||||
provider.credentials.auth_type = "oauth".to_string();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
|
||||
@@ -5,11 +5,20 @@
|
||||
use crate::proxy::{ProxyClientFactory, ProxyError, ProxyProtocol};
|
||||
use proptest::prelude::*;
|
||||
|
||||
/// 生成有效的主机名(必须以字母开头,避免纯数字被误认为 IP)
|
||||
fn arb_hostname() -> impl Strategy<Value = String> {
|
||||
(
|
||||
"[a-z]", // 首字母必须是字母
|
||||
"[a-z0-9]{0,19}", // 后续字符可以是字母或数字
|
||||
)
|
||||
.prop_map(|(first, rest)| format!("{}{}", first, rest))
|
||||
}
|
||||
|
||||
/// 生成有效的 socks5 代理 URL
|
||||
fn arb_socks5_url() -> impl Strategy<Value = String> {
|
||||
(
|
||||
"[a-z0-9]{1,20}", // host
|
||||
1024u16..65535u16, // port
|
||||
"[a-z][a-z0-9]{0,19}", // host: 必须以字母开头
|
||||
1024u16..65535u16, // port
|
||||
)
|
||||
.prop_map(|(host, port)| format!("socks5://{}:{}", host, port))
|
||||
}
|
||||
@@ -17,8 +26,8 @@ fn arb_socks5_url() -> impl Strategy<Value = String> {
|
||||
/// 生成有效的 http 代理 URL
|
||||
fn arb_http_url() -> impl Strategy<Value = String> {
|
||||
(
|
||||
"[a-z0-9]{1,20}", // host
|
||||
1024u16..65535u16, // port
|
||||
"[a-z][a-z0-9]{0,19}", // host: 必须以字母开头
|
||||
1024u16..65535u16, // port
|
||||
)
|
||||
.prop_map(|(host, port)| format!("http://{}:{}", host, port))
|
||||
}
|
||||
@@ -26,8 +35,8 @@ fn arb_http_url() -> impl Strategy<Value = String> {
|
||||
/// 生成有效的 https 代理 URL
|
||||
fn arb_https_url() -> impl Strategy<Value = String> {
|
||||
(
|
||||
"[a-z0-9]{1,20}", // host
|
||||
1024u16..65535u16, // port
|
||||
"[a-z][a-z0-9]{0,19}", // host: 必须以字母开头
|
||||
1024u16..65535u16, // port
|
||||
)
|
||||
.prop_map(|(host, port)| format!("https://{}:{}", host, port))
|
||||
}
|
||||
|
||||
@@ -101,22 +101,16 @@ impl ProviderRouter {
|
||||
// /{selector}/v1/messages
|
||||
[selector, "v1", "messages"] => {
|
||||
let registry = self.registry.read().await;
|
||||
let route = registry
|
||||
.find_by_selector(selector)
|
||||
.cloned()
|
||||
.unwrap_or_else(|| {
|
||||
// 创建一个临时的选择器路由
|
||||
RegisteredRoute {
|
||||
path_pattern: format!("/{}/v1/messages", selector),
|
||||
route_type: RouteType::CredentialSelector,
|
||||
provider_type: None,
|
||||
credential_uuid: None,
|
||||
credential_name: Some(selector.to_string()),
|
||||
protocols: vec!["claude".to_string()],
|
||||
enabled: true,
|
||||
priority: 50,
|
||||
}
|
||||
});
|
||||
let route = registry.find_by_selector(selector).cloned();
|
||||
|
||||
// 安全修复:未注册的 selector 不创建临时路由,直接返回 None
|
||||
let route = match route {
|
||||
Some(r) => r,
|
||||
None => {
|
||||
tracing::warn!("[ROUTER] 未注册的 selector: {},拒绝请求", selector);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
Some(RouteMatch {
|
||||
route,
|
||||
@@ -128,19 +122,16 @@ impl ProviderRouter {
|
||||
// /{selector}/v1/chat/completions
|
||||
[selector, "v1", "chat", "completions"] => {
|
||||
let registry = self.registry.read().await;
|
||||
let route = registry
|
||||
.find_by_selector(selector)
|
||||
.cloned()
|
||||
.unwrap_or_else(|| RegisteredRoute {
|
||||
path_pattern: format!("/{}/v1/chat/completions", selector),
|
||||
route_type: RouteType::CredentialSelector,
|
||||
provider_type: None,
|
||||
credential_uuid: None,
|
||||
credential_name: Some(selector.to_string()),
|
||||
protocols: vec!["openai".to_string()],
|
||||
enabled: true,
|
||||
priority: 50,
|
||||
});
|
||||
let route = registry.find_by_selector(selector).cloned();
|
||||
|
||||
// 安全修复:未注册的 selector 不创建临时路由,直接返回 None
|
||||
let route = match route {
|
||||
Some(r) => r,
|
||||
None => {
|
||||
tracing::warn!("[ROUTER] 未注册的 selector: {},拒绝请求", selector);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
Some(RouteMatch {
|
||||
route,
|
||||
@@ -231,6 +222,11 @@ mod tests {
|
||||
let registry = Arc::new(RwLock::new(RouteRegistry::new()));
|
||||
let router = ProviderRouter::new(registry);
|
||||
|
||||
// 安全修复后,未注册的 selector 会返回 None,需要先注册
|
||||
router
|
||||
.register_credential("kiro", "uuid-selector-test", Some("my-kiro"))
|
||||
.await;
|
||||
|
||||
let match1 = router.resolve("/my-kiro/v1/messages").await.unwrap();
|
||||
assert_eq!(match1.protocol, "claude");
|
||||
assert_eq!(match1.selector, Some("my-kiro".to_string()));
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -383,6 +383,7 @@ pub async fn management_add_credential(
|
||||
if let Some(token_file) = request.token_file {
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path: token_file,
|
||||
api_base_url: request.base_url.clone(),
|
||||
}
|
||||
} else {
|
||||
return (
|
||||
|
||||
@@ -1,6 +1,18 @@
|
||||
//! Provider 调用处理器
|
||||
//!
|
||||
//! 根据凭证类型调用不同的 Provider API
|
||||
//!
|
||||
//! # 流式传输支持
|
||||
//!
|
||||
//! 本模块支持真正的端到端流式传输,通过以下组件实现:
|
||||
//! - `StreamManager`: 管理流式请求的生命周期
|
||||
//! - `StreamingProvider`: Provider 的流式 API 接口
|
||||
//! - `FlowMonitor`: 实时捕获流式响应
|
||||
//!
|
||||
//! # 需求覆盖
|
||||
//!
|
||||
//! - 需求 4.2: 调用 process_chunk 更新流重建器
|
||||
//! - 需求 5.1: 在收到 chunk 后立即转发给客户端
|
||||
|
||||
use axum::{
|
||||
body::Body,
|
||||
@@ -13,6 +25,7 @@ use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
|
||||
use crate::converter::openai_to_antigravity::{
|
||||
convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context,
|
||||
};
|
||||
use crate::flow_monitor::stream_rebuilder::StreamFormat;
|
||||
use crate::models::anthropic::AnthropicMessagesRequest;
|
||||
use crate::models::openai::ChatCompletionRequest;
|
||||
use crate::models::provider_pool_model::{CredentialData, ProviderCredential};
|
||||
@@ -24,12 +37,38 @@ use crate::server_utils::{
|
||||
build_anthropic_response, build_anthropic_stream_response, parse_cw_response, safe_truncate,
|
||||
CWParsedResponse,
|
||||
};
|
||||
use crate::streaming::{
|
||||
StreamConfig, StreamContext, StreamError, StreamFormat as StreamingFormat, StreamManager,
|
||||
StreamResponse,
|
||||
};
|
||||
|
||||
/// 根据凭证调用 Provider (Anthropic 格式)
|
||||
///
|
||||
/// # 参数
|
||||
/// - `state`: 应用状态
|
||||
/// - `credential`: 凭证信息
|
||||
/// - `request`: Anthropic 格式请求
|
||||
/// - `flow_id`: Flow ID(可选,用于流式响应处理)
|
||||
pub async fn call_provider_anthropic(
|
||||
state: &AppState,
|
||||
credential: &ProviderCredential,
|
||||
request: &AnthropicMessagesRequest,
|
||||
flow_id: Option<&str>,
|
||||
) -> Response {
|
||||
// 如果是流式请求且有 flow_id,设置流式状态
|
||||
if request.stream {
|
||||
if let Some(fid) = flow_id {
|
||||
// 根据凭证类型确定流格式
|
||||
let format = match &credential.credential {
|
||||
CredentialData::KiroOAuth { .. } => StreamFormat::OpenAI,
|
||||
CredentialData::ClaudeKey { .. } => StreamFormat::Anthropic,
|
||||
CredentialData::AntigravityOAuth { .. } => StreamFormat::Gemini,
|
||||
_ => StreamFormat::Unknown,
|
||||
};
|
||||
state.flow_monitor.set_streaming(fid, format).await;
|
||||
}
|
||||
}
|
||||
|
||||
match &credential.credential {
|
||||
CredentialData::KiroOAuth { creds_file_path } => {
|
||||
// 使用 TokenCacheService 获取有效 token
|
||||
@@ -642,11 +681,19 @@ pub async fn call_provider_anthropic(
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 根据凭证调用 Provider (OpenAI 格式)
|
||||
///
|
||||
/// # 参数
|
||||
/// - `state`: 应用状态
|
||||
/// - `credential`: 凭证信息
|
||||
/// - `request`: OpenAI 格式请求
|
||||
/// - `flow_id`: Flow ID(可选,用于流式响应处理)
|
||||
pub async fn call_provider_openai(
|
||||
state: &AppState,
|
||||
credential: &ProviderCredential,
|
||||
request: &ChatCompletionRequest,
|
||||
flow_id: Option<&str>,
|
||||
) -> Response {
|
||||
let _start_time = std::time::Instant::now();
|
||||
match &credential.credential {
|
||||
@@ -921,3 +968,508 @@ pub async fn call_provider_openai(
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 流式传输支持
|
||||
// ============================================================================
|
||||
|
||||
/// 获取凭证对应的流式格式
|
||||
///
|
||||
/// 根据凭证类型返回对应的流式响应格式。
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential`: 凭证信息
|
||||
///
|
||||
/// # 返回
|
||||
/// 流式格式枚举
|
||||
pub fn get_stream_format_for_credential(credential: &ProviderCredential) -> StreamingFormat {
|
||||
match &credential.credential {
|
||||
CredentialData::KiroOAuth { .. } => StreamingFormat::AwsEventStream,
|
||||
CredentialData::ClaudeKey { .. } => StreamingFormat::AnthropicSse,
|
||||
CredentialData::OpenAIKey { .. } => StreamingFormat::OpenAiSse,
|
||||
// TODO: 任务 6 完成后,将这些改为 GeminiStream
|
||||
CredentialData::AntigravityOAuth { .. } => StreamingFormat::OpenAiSse,
|
||||
CredentialData::GeminiOAuth { .. } => StreamingFormat::OpenAiSse,
|
||||
CredentialData::GeminiApiKey { .. } => StreamingFormat::OpenAiSse,
|
||||
CredentialData::VertexKey { .. } => StreamingFormat::OpenAiSse,
|
||||
_ => StreamingFormat::OpenAiSse,
|
||||
}
|
||||
}
|
||||
|
||||
/// 处理流式响应
|
||||
///
|
||||
/// 使用 StreamManager 处理流式响应,集成 Flow Monitor。
|
||||
///
|
||||
/// # 参数
|
||||
/// - `state`: 应用状态
|
||||
/// - `flow_id`: Flow ID(用于 Flow Monitor 集成)
|
||||
/// - `source_stream`: 源字节流
|
||||
/// - `source_format`: 源流格式
|
||||
/// - `target_format`: 目标流格式
|
||||
/// - `model`: 模型名称
|
||||
///
|
||||
/// # 返回
|
||||
/// SSE 格式的 HTTP 响应
|
||||
///
|
||||
/// # 需求覆盖
|
||||
/// - 需求 4.2: 调用 process_chunk 更新流重建器
|
||||
/// - 需求 5.1: 在收到 chunk 后立即转发给客户端
|
||||
pub async fn handle_streaming_response(
|
||||
state: &AppState,
|
||||
flow_id: Option<&str>,
|
||||
source_stream: StreamResponse,
|
||||
source_format: StreamingFormat,
|
||||
target_format: StreamingFormat,
|
||||
model: &str,
|
||||
) -> Response {
|
||||
// 创建流式管理器
|
||||
let manager = StreamManager::with_default_config();
|
||||
|
||||
// 创建流式上下文
|
||||
let context = StreamContext::new(
|
||||
flow_id.map(|s| s.to_string()),
|
||||
source_format,
|
||||
target_format,
|
||||
model,
|
||||
);
|
||||
|
||||
// 获取 flow_id 的克隆用于回调
|
||||
let flow_id_for_callback = flow_id.map(|s| s.to_string());
|
||||
let flow_monitor = state.flow_monitor.clone();
|
||||
|
||||
// 创建带回调的流式处理
|
||||
let managed_stream = if let Some(fid) = flow_id_for_callback {
|
||||
// 使用带回调的流式处理,集成 Flow Monitor
|
||||
let on_chunk = move |event: &str, _metrics: &crate::streaming::StreamMetrics| {
|
||||
// 解析 SSE 事件并调用 process_chunk
|
||||
// SSE 格式: "event: xxx\ndata: {...}\n\n"
|
||||
let lines: Vec<&str> = event.lines().collect();
|
||||
let mut event_type: Option<&str> = None;
|
||||
let mut data: Option<&str> = None;
|
||||
|
||||
for line in lines {
|
||||
if line.starts_with("event: ") {
|
||||
event_type = Some(&line[7..]);
|
||||
} else if line.starts_with("data: ") {
|
||||
data = Some(&line[6..]);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(d) = data {
|
||||
// 使用 tokio::spawn 异步调用 process_chunk
|
||||
let flow_monitor_clone = flow_monitor.clone();
|
||||
let fid_clone = fid.clone();
|
||||
let event_type_owned = event_type.map(|s| s.to_string());
|
||||
let data_owned = d.to_string();
|
||||
|
||||
tokio::spawn(async move {
|
||||
flow_monitor_clone
|
||||
.process_chunk(&fid_clone, event_type_owned.as_deref(), &data_owned)
|
||||
.await;
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
let stream = manager.handle_stream_with_callback(context, source_stream, on_chunk);
|
||||
|
||||
// 转换为 Body 流
|
||||
let body_stream = stream.map(|result| -> Result<axum::body::Bytes, std::io::Error> {
|
||||
match result {
|
||||
Ok(event) => Ok(axum::body::Bytes::from(event)),
|
||||
Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())),
|
||||
}
|
||||
});
|
||||
|
||||
Body::from_stream(body_stream)
|
||||
} else {
|
||||
// 没有 flow_id,使用普通流式处理
|
||||
let stream = manager.handle_stream(context, source_stream);
|
||||
|
||||
let body_stream = stream.map(|result| -> Result<axum::body::Bytes, std::io::Error> {
|
||||
match result {
|
||||
Ok(event) => Ok(axum::body::Bytes::from(event)),
|
||||
Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())),
|
||||
}
|
||||
});
|
||||
|
||||
Body::from_stream(body_stream)
|
||||
};
|
||||
|
||||
// 构建 SSE 响应
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "text/event-stream")
|
||||
.header(header::CACHE_CONTROL, "no-cache")
|
||||
.header(header::CONNECTION, "keep-alive")
|
||||
.header("X-Accel-Buffering", "no")
|
||||
.body(managed_stream)
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(
|
||||
serde_json::json!({"error": {"message": "Failed to build streaming response"}}),
|
||||
),
|
||||
)
|
||||
.into_response()
|
||||
})
|
||||
}
|
||||
|
||||
/// 处理流式响应(带超时)
|
||||
///
|
||||
/// 与 `handle_streaming_response` 类似,但添加了超时保护。
|
||||
///
|
||||
/// # 参数
|
||||
/// - `state`: 应用状态
|
||||
/// - `flow_id`: Flow ID
|
||||
/// - `source_stream`: 源字节流
|
||||
/// - `source_format`: 源流格式
|
||||
/// - `target_format`: 目标流格式
|
||||
/// - `model`: 模型名称
|
||||
/// - `timeout_ms`: 超时时间(毫秒)
|
||||
///
|
||||
/// # 返回
|
||||
/// SSE 格式的 HTTP 响应
|
||||
///
|
||||
/// # 需求覆盖
|
||||
/// - 需求 6.2: 超时错误处理
|
||||
/// - 需求 6.5: 可配置的流式响应超时
|
||||
pub async fn handle_streaming_response_with_timeout(
|
||||
state: &AppState,
|
||||
flow_id: Option<&str>,
|
||||
source_stream: StreamResponse,
|
||||
source_format: StreamingFormat,
|
||||
target_format: StreamingFormat,
|
||||
model: &str,
|
||||
timeout_ms: u64,
|
||||
) -> Response {
|
||||
use futures::stream::BoxStream;
|
||||
|
||||
// 创建带超时配置的流式管理器
|
||||
let config = StreamConfig::new()
|
||||
.with_timeout_ms(timeout_ms)
|
||||
.with_chunk_timeout_ms(30_000); // 30 秒 chunk 超时
|
||||
|
||||
let manager = StreamManager::new(config.clone());
|
||||
|
||||
// 创建流式上下文
|
||||
let context = StreamContext::new(
|
||||
flow_id.map(|s| s.to_string()),
|
||||
source_format,
|
||||
target_format,
|
||||
model,
|
||||
);
|
||||
|
||||
// 获取 flow_id 的克隆用于回调
|
||||
let flow_id_for_callback = flow_id.map(|s| s.to_string());
|
||||
let flow_monitor = state.flow_monitor.clone();
|
||||
|
||||
// 创建带超时的流式处理,使用 BoxStream 统一类型
|
||||
let timeout_stream: BoxStream<'static, Result<String, crate::streaming::StreamError>> =
|
||||
if let Some(fid) = flow_id_for_callback {
|
||||
let on_chunk = move |event: &str, _metrics: &crate::streaming::StreamMetrics| {
|
||||
let lines: Vec<&str> = event.lines().collect();
|
||||
let mut event_type: Option<&str> = None;
|
||||
let mut data: Option<&str> = None;
|
||||
|
||||
for line in lines {
|
||||
if line.starts_with("event: ") {
|
||||
event_type = Some(&line[7..]);
|
||||
} else if line.starts_with("data: ") {
|
||||
data = Some(&line[6..]);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(d) = data {
|
||||
let flow_monitor_clone = flow_monitor.clone();
|
||||
let fid_clone = fid.clone();
|
||||
let event_type_owned = event_type.map(|s| s.to_string());
|
||||
let data_owned = d.to_string();
|
||||
|
||||
tokio::spawn(async move {
|
||||
flow_monitor_clone
|
||||
.process_chunk(&fid_clone, event_type_owned.as_deref(), &data_owned)
|
||||
.await;
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
let stream = manager.handle_stream_with_callback(context, source_stream, on_chunk);
|
||||
Box::pin(crate::streaming::with_timeout(stream, &config))
|
||||
} else {
|
||||
let stream = manager.handle_stream(context, source_stream);
|
||||
Box::pin(crate::streaming::with_timeout(stream, &config))
|
||||
};
|
||||
|
||||
// 转换为 Body 流
|
||||
let body_stream = timeout_stream.map(|result| -> Result<axum::body::Bytes, std::io::Error> {
|
||||
match result {
|
||||
Ok(event) => Ok(axum::body::Bytes::from(event)),
|
||||
Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())),
|
||||
}
|
||||
});
|
||||
|
||||
// 构建 SSE 响应
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "text/event-stream")
|
||||
.header(header::CACHE_CONTROL, "no-cache")
|
||||
.header(header::CONNECTION, "keep-alive")
|
||||
.header("X-Accel-Buffering", "no")
|
||||
.body(Body::from_stream(body_stream))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(
|
||||
serde_json::json!({"error": {"message": "Failed to build streaming response"}}),
|
||||
),
|
||||
)
|
||||
.into_response()
|
||||
})
|
||||
}
|
||||
|
||||
/// 将 reqwest 响应转换为 StreamResponse
|
||||
///
|
||||
/// 用于将 Provider 的 HTTP 响应转换为统一的流式响应类型。
|
||||
///
|
||||
/// # 参数
|
||||
/// - `response`: reqwest HTTP 响应
|
||||
///
|
||||
/// # 返回
|
||||
/// 统一的流式响应类型
|
||||
pub fn response_to_stream(response: reqwest::Response) -> StreamResponse {
|
||||
crate::streaming::reqwest_stream_to_stream_response(response)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 客户端断开检测
|
||||
// ============================================================================
|
||||
|
||||
/// 带客户端断开检测的流式响应处理
|
||||
///
|
||||
/// 在流式传输过程中检测客户端是否断开连接,并在断开时:
|
||||
/// 1. 停止处理上游数据
|
||||
/// 2. 标记 Flow 为取消状态
|
||||
/// 3. 清理资源
|
||||
///
|
||||
/// # 参数
|
||||
/// - `state`: 应用状态
|
||||
/// - `flow_id`: Flow ID
|
||||
/// - `source_stream`: 源字节流
|
||||
/// - `source_format`: 源流格式
|
||||
/// - `target_format`: 目标流格式
|
||||
/// - `model`: 模型名称
|
||||
/// - `cancel_token`: 取消令牌(用于取消上游请求)
|
||||
///
|
||||
/// # 返回
|
||||
/// SSE 格式的 HTTP 响应
|
||||
///
|
||||
/// # 需求覆盖
|
||||
/// - 需求 5.4: 客户端断开时取消上游请求
|
||||
pub async fn handle_streaming_with_disconnect_detection(
|
||||
state: &AppState,
|
||||
flow_id: Option<&str>,
|
||||
source_stream: StreamResponse,
|
||||
source_format: StreamingFormat,
|
||||
target_format: StreamingFormat,
|
||||
model: &str,
|
||||
cancel_token: Option<tokio_util::sync::CancellationToken>,
|
||||
) -> Response {
|
||||
use futures::StreamExt;
|
||||
|
||||
// 创建流式管理器
|
||||
let manager = StreamManager::with_default_config();
|
||||
|
||||
// 创建流式上下文
|
||||
let context = StreamContext::new(
|
||||
flow_id.map(|s| s.to_string()),
|
||||
source_format,
|
||||
target_format,
|
||||
model,
|
||||
);
|
||||
|
||||
// 获取 flow_id 的克隆
|
||||
let flow_id_for_callback = flow_id.map(|s| s.to_string());
|
||||
let flow_id_for_cancel = flow_id.map(|s| s.to_string());
|
||||
let flow_monitor = state.flow_monitor.clone();
|
||||
let flow_monitor_for_cancel = state.flow_monitor.clone();
|
||||
|
||||
// 创建带回调的流式处理
|
||||
// 使用 BoxStream 统一类型
|
||||
let managed_stream: futures::stream::BoxStream<
|
||||
'static,
|
||||
Result<String, crate::streaming::StreamError>,
|
||||
> = if let Some(fid) = flow_id_for_callback {
|
||||
let on_chunk = move |event: &str, _metrics: &crate::streaming::StreamMetrics| {
|
||||
let lines: Vec<&str> = event.lines().collect();
|
||||
let mut event_type: Option<&str> = None;
|
||||
let mut data: Option<&str> = None;
|
||||
|
||||
for line in lines {
|
||||
if line.starts_with("event: ") {
|
||||
event_type = Some(&line[7..]);
|
||||
} else if line.starts_with("data: ") {
|
||||
data = Some(&line[6..]);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(d) = data {
|
||||
let flow_monitor_clone = flow_monitor.clone();
|
||||
let fid_clone = fid.clone();
|
||||
let event_type_owned = event_type.map(|s| s.to_string());
|
||||
let data_owned = d.to_string();
|
||||
|
||||
tokio::spawn(async move {
|
||||
flow_monitor_clone
|
||||
.process_chunk(&fid_clone, event_type_owned.as_deref(), &data_owned)
|
||||
.await;
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
Box::pin(manager.handle_stream_with_callback(context, source_stream, on_chunk))
|
||||
} else {
|
||||
// 没有 flow_id,使用普通流式处理
|
||||
Box::pin(manager.handle_stream(context, source_stream))
|
||||
};
|
||||
|
||||
// 如果有取消令牌,创建一个可取消的流
|
||||
let body_stream = if let Some(token) = cancel_token {
|
||||
// 创建一个可取消的流
|
||||
let cancellable_stream = CancellableStream::new(managed_stream, token.clone());
|
||||
|
||||
// 当流被取消时,标记 Flow 为取消状态
|
||||
let cancel_handler = {
|
||||
let token = token.clone();
|
||||
let flow_id = flow_id_for_cancel.clone();
|
||||
async move {
|
||||
token.cancelled().await;
|
||||
if let Some(fid) = flow_id {
|
||||
flow_monitor_for_cancel.cancel_flow(&fid).await;
|
||||
tracing::info!("[STREAM] 客户端断开,已取消 Flow: {}", fid);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// 在后台运行取消处理器
|
||||
tokio::spawn(cancel_handler);
|
||||
|
||||
// 转换为 Body 流
|
||||
let stream =
|
||||
cancellable_stream.map(|result| -> Result<axum::body::Bytes, std::io::Error> {
|
||||
match result {
|
||||
Ok(event) => Ok(axum::body::Bytes::from(event)),
|
||||
Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())),
|
||||
}
|
||||
});
|
||||
|
||||
Body::from_stream(stream)
|
||||
} else {
|
||||
// 没有取消令牌,使用普通流
|
||||
let stream = managed_stream.map(|result| -> Result<axum::body::Bytes, std::io::Error> {
|
||||
match result {
|
||||
Ok(event) => Ok(axum::body::Bytes::from(event)),
|
||||
Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())),
|
||||
}
|
||||
});
|
||||
|
||||
Body::from_stream(stream)
|
||||
};
|
||||
|
||||
// 构建 SSE 响应
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "text/event-stream")
|
||||
.header(header::CACHE_CONTROL, "no-cache")
|
||||
.header(header::CONNECTION, "keep-alive")
|
||||
.header("X-Accel-Buffering", "no")
|
||||
.body(body_stream)
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(
|
||||
serde_json::json!({"error": {"message": "Failed to build streaming response"}}),
|
||||
),
|
||||
)
|
||||
.into_response()
|
||||
})
|
||||
}
|
||||
|
||||
/// 可取消的流包装器
|
||||
///
|
||||
/// 包装一个流,使其可以通过取消令牌取消。
|
||||
/// 当取消令牌被触发时,流将返回 ClientDisconnected 错误。
|
||||
pub struct CancellableStream<S> {
|
||||
inner: S,
|
||||
cancel_token: tokio_util::sync::CancellationToken,
|
||||
cancelled: bool,
|
||||
}
|
||||
|
||||
impl<S> CancellableStream<S> {
|
||||
/// 创建新的可取消流
|
||||
pub fn new(inner: S, cancel_token: tokio_util::sync::CancellationToken) -> Self {
|
||||
Self {
|
||||
inner,
|
||||
cancel_token,
|
||||
cancelled: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<S> futures::Stream for CancellableStream<S>
|
||||
where
|
||||
S: futures::Stream<Item = Result<String, StreamError>> + Unpin,
|
||||
{
|
||||
type Item = Result<String, StreamError>;
|
||||
|
||||
fn poll_next(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Option<Self::Item>> {
|
||||
use std::task::Poll;
|
||||
|
||||
// 检查是否已取消
|
||||
if self.cancelled {
|
||||
return Poll::Ready(None);
|
||||
}
|
||||
|
||||
// 检查取消令牌
|
||||
if self.cancel_token.is_cancelled() {
|
||||
self.cancelled = true;
|
||||
return Poll::Ready(Some(Err(StreamError::ClientDisconnected)));
|
||||
}
|
||||
|
||||
// 轮询内部流
|
||||
std::pin::Pin::new(&mut self.inner).poll_next(cx)
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建取消令牌
|
||||
///
|
||||
/// 创建一个可用于取消流式请求的令牌。
|
||||
///
|
||||
/// # 返回
|
||||
/// 取消令牌
|
||||
pub fn create_cancel_token() -> tokio_util::sync::CancellationToken {
|
||||
tokio_util::sync::CancellationToken::new()
|
||||
}
|
||||
|
||||
/// 检测客户端断开并触发取消
|
||||
///
|
||||
/// 监控客户端连接状态,当检测到断开时触发取消令牌。
|
||||
///
|
||||
/// # 参数
|
||||
/// - `cancel_token`: 取消令牌
|
||||
///
|
||||
/// # 注意
|
||||
/// 此函数应该在单独的任务中运行,与流式响应并行。
|
||||
/// 实际的断开检测依赖于 axum 的连接管理。
|
||||
pub async fn monitor_client_disconnect(cancel_token: tokio_util::sync::CancellationToken) {
|
||||
// 在实际应用中,这里会监控客户端连接状态
|
||||
// 当检测到断开时,调用 cancel_token.cancel()
|
||||
//
|
||||
// 由于 axum 的 SSE 响应会自动处理客户端断开,
|
||||
// 这个函数主要用于需要主动检测断开的场景
|
||||
|
||||
// 等待取消令牌被触发(由其他地方触发)
|
||||
cancel_token.cancelled().await;
|
||||
}
|
||||
|
||||
@@ -6,12 +6,15 @@ use axum::{
|
||||
body::Body,
|
||||
extract::{
|
||||
ws::{Message as WsMessage, WebSocket, WebSocketUpgrade},
|
||||
State,
|
||||
Query, State,
|
||||
},
|
||||
http::HeaderMap,
|
||||
response::IntoResponse,
|
||||
};
|
||||
use futures::{SinkExt, StreamExt as FuturesStreamExt};
|
||||
use serde::Deserialize;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
|
||||
use crate::converter::openai_to_antigravity::{
|
||||
@@ -27,40 +30,57 @@ use crate::providers::{
|
||||
use crate::server::AppState;
|
||||
use crate::server_utils::parse_cw_response;
|
||||
use crate::websocket::{
|
||||
WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsMessage as WsProtoMessage,
|
||||
WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsFlowEvent, WsMessage as WsProtoMessage,
|
||||
};
|
||||
|
||||
/// WebSocket 查询参数
|
||||
#[derive(Debug, Deserialize, Default)]
|
||||
pub struct WsQueryParams {
|
||||
/// API 密钥(通过 URL 参数传递)
|
||||
pub api_key: Option<String>,
|
||||
/// Token(通过 URL 参数传递,与 api_key 等效)
|
||||
pub token: Option<String>,
|
||||
}
|
||||
|
||||
/// WebSocket 升级处理器
|
||||
pub async fn ws_upgrade_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<AppState>,
|
||||
Query(params): Query<WsQueryParams>,
|
||||
headers: HeaderMap,
|
||||
) -> impl IntoResponse {
|
||||
// 验证 API 密钥
|
||||
// 验证 API 密钥:优先从 header 获取,其次从 URL 参数获取
|
||||
let auth = headers
|
||||
.get("authorization")
|
||||
.or_else(|| headers.get("x-api-key"))
|
||||
.and_then(|v| v.to_str().ok());
|
||||
|
||||
let key = match auth {
|
||||
Some(s) if s.starts_with("Bearer ") => &s[7..],
|
||||
Some(s) => s,
|
||||
Some(s) if s.starts_with("Bearer ") => Some(&s[7..]),
|
||||
Some(s) => Some(s),
|
||||
None => {
|
||||
return axum::http::Response::builder()
|
||||
.status(401)
|
||||
.body(Body::from("No API key provided"))
|
||||
.unwrap()
|
||||
.into_response();
|
||||
// 尝试从 URL 参数获取
|
||||
params.api_key.as_deref().or(params.token.as_deref())
|
||||
}
|
||||
};
|
||||
|
||||
if key != state.api_key {
|
||||
return axum::http::Response::builder()
|
||||
.status(401)
|
||||
.body(Body::from("Invalid API key"))
|
||||
.unwrap()
|
||||
.into_response();
|
||||
}
|
||||
// 如果没有提供任何认证信息,允许连接(用于内部 Flow Monitor)
|
||||
// 但会在日志中记录
|
||||
let authenticated = match key {
|
||||
Some(k) if k == state.api_key => true,
|
||||
Some(_) => {
|
||||
return axum::http::Response::builder()
|
||||
.status(401)
|
||||
.body(Body::from("Invalid API key"))
|
||||
.unwrap()
|
||||
.into_response();
|
||||
}
|
||||
None => {
|
||||
// 允许无认证连接(仅用于本地 Flow Monitor UI)
|
||||
tracing::debug!("[WS] Allowing unauthenticated connection for Flow Monitor");
|
||||
false
|
||||
}
|
||||
};
|
||||
|
||||
// 获取客户端信息
|
||||
let client_info = headers
|
||||
@@ -68,11 +88,16 @@ pub async fn ws_upgrade_handler(
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
ws.on_upgrade(move |socket| handle_websocket(socket, state, client_info))
|
||||
ws.on_upgrade(move |socket| handle_websocket(socket, state, client_info, authenticated))
|
||||
}
|
||||
|
||||
/// 处理 WebSocket 连接
|
||||
pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: Option<String>) {
|
||||
pub async fn handle_websocket(
|
||||
socket: WebSocket,
|
||||
state: AppState,
|
||||
client_info: Option<String>,
|
||||
authenticated: bool,
|
||||
) {
|
||||
let conn_id = uuid::Uuid::new_v4().to_string();
|
||||
|
||||
// 注册连接
|
||||
@@ -90,13 +115,73 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[WS] New connection: {} (client: {:?})",
|
||||
"[WS] New connection: {} (client: {:?}, authenticated: {})",
|
||||
&conn_id[..8],
|
||||
client_info
|
||||
client_info,
|
||||
authenticated
|
||||
),
|
||||
);
|
||||
|
||||
let (mut sender, mut receiver) = socket.split();
|
||||
let (sender, mut receiver) = socket.split();
|
||||
let sender = Arc::new(Mutex::new(sender));
|
||||
|
||||
// Flow 事件订阅状态
|
||||
let flow_subscribed = Arc::new(std::sync::atomic::AtomicBool::new(false));
|
||||
|
||||
// 启动 Flow 事件转发任务
|
||||
let flow_sender = sender.clone();
|
||||
let flow_subscribed_clone = flow_subscribed.clone();
|
||||
let flow_monitor = state.flow_monitor.clone();
|
||||
let conn_id_clone = conn_id.clone();
|
||||
let _logs_clone = state.logs.clone();
|
||||
|
||||
let flow_task = tokio::spawn(async move {
|
||||
let mut flow_receiver = flow_monitor.subscribe();
|
||||
|
||||
loop {
|
||||
match flow_receiver.recv().await {
|
||||
Ok(event) => {
|
||||
// 只有在订阅状态下才转发事件
|
||||
if !flow_subscribed_clone.load(std::sync::atomic::Ordering::Relaxed) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// 转换为 WebSocket 消息
|
||||
let ws_event: WsFlowEvent = event.into();
|
||||
let ws_msg = WsProtoMessage::FlowEvent(ws_event);
|
||||
|
||||
if let Ok(msg_text) = serde_json::to_string(&ws_msg) {
|
||||
let mut sender_guard = flow_sender.lock().await;
|
||||
if sender_guard
|
||||
.send(WsMessage::Text(msg_text.into()))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
tracing::debug!(
|
||||
"[WS] Flow event send failed for connection {}",
|
||||
&conn_id_clone[..8]
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
|
||||
tracing::warn!(
|
||||
"[WS] Flow event receiver lagged by {} messages for connection {}",
|
||||
n,
|
||||
&conn_id_clone[..8]
|
||||
);
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Closed) => {
|
||||
tracing::debug!(
|
||||
"[WS] Flow event channel closed for connection {}",
|
||||
&conn_id_clone[..8]
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// 消息处理循环
|
||||
while let Some(msg) = receiver.next().await {
|
||||
@@ -107,7 +192,8 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O
|
||||
|
||||
match serde_json::from_str::<WsProtoMessage>(&text) {
|
||||
Ok(ws_msg) => {
|
||||
let response = handle_ws_message(&state, &conn_id, ws_msg).await;
|
||||
let response =
|
||||
handle_ws_message(&state, &conn_id, ws_msg, &flow_subscribed).await;
|
||||
if let Some(resp) = response {
|
||||
let resp_text = serde_json::to_string(&resp).unwrap_or_default();
|
||||
if sender.send(WsMessage::Text(resp_text)).await.is_err() {
|
||||
@@ -139,7 +225,8 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O
|
||||
}
|
||||
}
|
||||
Ok(WsMessage::Ping(data)) => {
|
||||
if sender.send(WsMessage::Pong(data)).await.is_err() {
|
||||
let mut sender_guard = sender.lock().await;
|
||||
if sender_guard.send(WsMessage::Pong(data)).await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
@@ -159,6 +246,9 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O
|
||||
}
|
||||
}
|
||||
|
||||
// 取消 Flow 事件转发任务
|
||||
flow_task.abort();
|
||||
|
||||
// 清理连接
|
||||
state.ws_manager.unregister(&conn_id);
|
||||
state.logs.write().await.add(
|
||||
@@ -172,10 +262,56 @@ async fn handle_ws_message(
|
||||
state: &AppState,
|
||||
conn_id: &str,
|
||||
msg: WsProtoMessage,
|
||||
flow_subscribed: &Arc<std::sync::atomic::AtomicBool>,
|
||||
) -> Option<WsProtoMessage> {
|
||||
match msg {
|
||||
WsProtoMessage::Ping { timestamp } => Some(WsProtoMessage::Pong { timestamp }),
|
||||
WsProtoMessage::Pong { .. } => None,
|
||||
WsProtoMessage::SubscribeFlowEvents => {
|
||||
// 订阅 Flow 事件
|
||||
flow_subscribed.store(true, std::sync::atomic::Ordering::Relaxed);
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[WS] Connection {} subscribed to flow events",
|
||||
&conn_id[..8]
|
||||
),
|
||||
);
|
||||
// 返回确认消息
|
||||
Some(WsProtoMessage::Response(WsApiResponse {
|
||||
request_id: "subscribe_flow_events".to_string(),
|
||||
payload: serde_json::json!({
|
||||
"status": "subscribed",
|
||||
"message": "Successfully subscribed to flow events"
|
||||
}),
|
||||
}))
|
||||
}
|
||||
WsProtoMessage::UnsubscribeFlowEvents => {
|
||||
// 取消订阅 Flow 事件
|
||||
flow_subscribed.store(false, std::sync::atomic::Ordering::Relaxed);
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
"[WS] Connection {} unsubscribed from flow events",
|
||||
&conn_id[..8]
|
||||
),
|
||||
);
|
||||
// 返回确认消息
|
||||
Some(WsProtoMessage::Response(WsApiResponse {
|
||||
request_id: "unsubscribe_flow_events".to_string(),
|
||||
payload: serde_json::json!({
|
||||
"status": "unsubscribed",
|
||||
"message": "Successfully unsubscribed from flow events"
|
||||
}),
|
||||
}))
|
||||
}
|
||||
WsProtoMessage::FlowEvent(_) => {
|
||||
// 客户端不应该发送 FlowEvent 消息
|
||||
Some(WsProtoMessage::Error(WsError::invalid_request(
|
||||
None,
|
||||
"FlowEvent messages are server-to-client only",
|
||||
)))
|
||||
}
|
||||
WsProtoMessage::Request(request) => {
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
|
||||
@@ -7,6 +7,7 @@ use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
|
||||
use crate::credential::CredentialSyncService;
|
||||
use crate::database::dao::provider_pool::ProviderPoolDao;
|
||||
use crate::database::DbConnection;
|
||||
use crate::flow_monitor::{FlowInterceptor, FlowMonitor, FlowMonitorConfig};
|
||||
use crate::injection::Injector;
|
||||
use crate::logger::LogStore;
|
||||
use crate::models::anthropic::*;
|
||||
@@ -225,6 +226,37 @@ impl ServerState {
|
||||
shared_stats: Option<Arc<parking_lot::RwLock<crate::telemetry::StatsAggregator>>>,
|
||||
shared_tokens: Option<Arc<parking_lot::RwLock<crate::telemetry::TokenTracker>>>,
|
||||
shared_logger: Option<Arc<crate::telemetry::RequestLogger>>,
|
||||
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
self.start_with_telemetry_and_flow_monitor(
|
||||
logs,
|
||||
pool_service,
|
||||
token_cache,
|
||||
db,
|
||||
shared_stats,
|
||||
shared_tokens,
|
||||
shared_logger,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// 启动服务器(使用共享的遥测实例和 Flow Monitor)
|
||||
///
|
||||
/// 这允许服务器与 TelemetryState 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger,
|
||||
/// 以及与 FlowMonitorState 共享同一个 FlowMonitor,
|
||||
/// 使得请求处理过程中记录的统计数据和 Flow 数据能够在前端监控页面中显示。
|
||||
pub async fn start_with_telemetry_and_flow_monitor(
|
||||
&mut self,
|
||||
logs: Arc<RwLock<LogStore>>,
|
||||
pool_service: Arc<ProviderPoolService>,
|
||||
token_cache: Arc<TokenCacheService>,
|
||||
db: Option<DbConnection>,
|
||||
shared_stats: Option<Arc<parking_lot::RwLock<crate::telemetry::StatsAggregator>>>,
|
||||
shared_tokens: Option<Arc<parking_lot::RwLock<crate::telemetry::TokenTracker>>>,
|
||||
shared_logger: Option<Arc<crate::telemetry::RequestLogger>>,
|
||||
shared_flow_monitor: Option<Arc<FlowMonitor>>,
|
||||
shared_flow_interceptor: Option<Arc<FlowInterceptor>>,
|
||||
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
if self.running {
|
||||
return Ok(());
|
||||
@@ -275,6 +307,8 @@ impl ServerState {
|
||||
shared_stats,
|
||||
shared_tokens,
|
||||
shared_logger,
|
||||
shared_flow_monitor,
|
||||
shared_flow_interceptor,
|
||||
Some(config),
|
||||
Some(config_path),
|
||||
)
|
||||
@@ -343,6 +377,10 @@ pub struct AppState {
|
||||
pub request_logger: Option<Arc<crate::telemetry::RequestLogger>>,
|
||||
/// Amp CLI 路由器
|
||||
pub amp_router: Arc<crate::router::AmpRouter>,
|
||||
/// Flow 监控服务
|
||||
pub flow_monitor: Arc<FlowMonitor>,
|
||||
/// Flow 拦截器
|
||||
pub flow_interceptor: Arc<FlowInterceptor>,
|
||||
}
|
||||
|
||||
/// 启动配置文件监控
|
||||
@@ -620,6 +658,8 @@ async fn run_server(
|
||||
shared_stats: Option<Arc<parking_lot::RwLock<crate::telemetry::StatsAggregator>>>,
|
||||
shared_tokens: Option<Arc<parking_lot::RwLock<crate::telemetry::TokenTracker>>>,
|
||||
shared_logger: Option<Arc<crate::telemetry::RequestLogger>>,
|
||||
shared_flow_monitor: Option<Arc<FlowMonitor>>,
|
||||
shared_flow_interceptor: Option<Arc<FlowInterceptor>>,
|
||||
config: Option<Config>,
|
||||
config_path: Option<PathBuf>,
|
||||
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
@@ -673,6 +713,14 @@ async fn run_server(
|
||||
.unwrap_or_default(),
|
||||
));
|
||||
|
||||
// 使用共享的 Flow 监控服务,如果没有则创建新的
|
||||
let flow_monitor = shared_flow_monitor
|
||||
.unwrap_or_else(|| Arc::new(FlowMonitor::new(FlowMonitorConfig::default(), None)));
|
||||
|
||||
// 使用共享的 Flow 拦截器,如果没有则创建新的
|
||||
let flow_interceptor =
|
||||
shared_flow_interceptor.unwrap_or_else(|| Arc::new(FlowInterceptor::default()));
|
||||
|
||||
let state = AppState {
|
||||
api_key: api_key.to_string(),
|
||||
base_url,
|
||||
@@ -693,6 +741,8 @@ async fn run_server(
|
||||
hot_reload_manager: hot_reload_manager.clone(),
|
||||
request_logger: shared_logger,
|
||||
amp_router,
|
||||
flow_monitor,
|
||||
flow_interceptor,
|
||||
};
|
||||
|
||||
// 启动配置文件监控
|
||||
@@ -1152,7 +1202,8 @@ async fn anthropic_messages_with_selector(
|
||||
);
|
||||
|
||||
// 根据凭证类型调用相应的 Provider
|
||||
handlers::call_provider_anthropic(&state, &cred, &request).await
|
||||
// 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求
|
||||
handlers::call_provider_anthropic(&state, &cred, &request, None).await
|
||||
}
|
||||
None => {
|
||||
// 回退到默认 Kiro provider
|
||||
@@ -1224,7 +1275,8 @@ async fn chat_completions_with_selector(
|
||||
),
|
||||
);
|
||||
|
||||
handlers::call_provider_openai(&state, &cred, &request).await
|
||||
// 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求
|
||||
handlers::call_provider_openai(&state, &cred, &request, None).await
|
||||
}
|
||||
None => {
|
||||
state.logs.write().await.add(
|
||||
@@ -1320,7 +1372,8 @@ async fn amp_chat_completions(
|
||||
&cred.uuid[..8]
|
||||
),
|
||||
);
|
||||
handlers::call_provider_openai(&state, &cred, &request).await
|
||||
// 注意:这里没有 Flow 捕获,因为是通过 AMP CLI 路由的请求
|
||||
handlers::call_provider_openai(&state, &cred, &request, None).await
|
||||
}
|
||||
None => {
|
||||
state.logs.write().await.add(
|
||||
@@ -1415,7 +1468,8 @@ async fn amp_messages(
|
||||
&cred.uuid[..8]
|
||||
),
|
||||
);
|
||||
handlers::call_provider_anthropic(&state, &cred, &request).await
|
||||
// 注意:这里没有 Flow 捕获,因为是通过 AMP CLI 路由的请求
|
||||
handlers::call_provider_anthropic(&state, &cred, &request, None).await
|
||||
}
|
||||
None => {
|
||||
state.logs.write().await.add(
|
||||
|
||||
@@ -0,0 +1,143 @@
|
||||
//! 备份服务
|
||||
//!
|
||||
//! 提供数据库与配置备份的基础能力
|
||||
|
||||
use crate::database::{get_db_path, DbConnection};
|
||||
use chrono::{DateTime, Duration, Utc};
|
||||
use rusqlite::DatabaseName;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct BackupService {
|
||||
backup_dir: PathBuf,
|
||||
retention_days: u32,
|
||||
}
|
||||
|
||||
impl BackupService {
|
||||
pub fn new(backup_dir: PathBuf, retention_days: u32) -> Result<Self, String> {
|
||||
std::fs::create_dir_all(&backup_dir)
|
||||
.map_err(|e| format!("无法创建备份目录 {:?}: {}", backup_dir, e))?;
|
||||
Ok(Self {
|
||||
backup_dir,
|
||||
retention_days,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn with_defaults() -> Result<Self, String> {
|
||||
let home = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?;
|
||||
let backup_dir = home.join(".proxycast").join("backups");
|
||||
Self::new(backup_dir, 7)
|
||||
}
|
||||
|
||||
pub fn backup_database(&self) -> Result<PathBuf, String> {
|
||||
let db_path = get_db_path()?;
|
||||
let timestamp = Utc::now().format("%Y%m%d_%H%M%S");
|
||||
let backup_path = self.backup_dir.join(format!("proxycast_{}.db", timestamp));
|
||||
|
||||
std::fs::copy(&db_path, &backup_path).map_err(|e| format!("备份失败: {}", e))?;
|
||||
|
||||
self.cleanup_old_backups()?;
|
||||
Ok(backup_path)
|
||||
}
|
||||
|
||||
pub fn backup_database_with_connection(&self, db: &DbConnection) -> Result<PathBuf, String> {
|
||||
let timestamp = Utc::now().format("%Y%m%d_%H%M%S");
|
||||
let backup_path = self.backup_dir.join(format!("proxycast_{}.db", timestamp));
|
||||
let conn = db.lock().map_err(|_| "数据库锁已被占用".to_string())?;
|
||||
let progress: Option<fn(rusqlite::backup::Progress)> = None;
|
||||
conn.backup(DatabaseName::Main, &backup_path, progress)
|
||||
.map_err(|e| format!("备份失败: {}", e))?;
|
||||
|
||||
self.cleanup_old_backups()?;
|
||||
Ok(backup_path)
|
||||
}
|
||||
|
||||
pub fn restore_database(&self, backup_path: &Path) -> Result<(), String> {
|
||||
// P1 安全修复:验证备份路径在白名单目录内
|
||||
let canonical_backup = backup_path
|
||||
.canonicalize()
|
||||
.map_err(|e| format!("无法解析备份路径: {}", e))?;
|
||||
let canonical_backup_dir = self
|
||||
.backup_dir
|
||||
.canonicalize()
|
||||
.map_err(|e| format!("无法解析备份目录: {}", e))?;
|
||||
|
||||
if !canonical_backup.starts_with(&canonical_backup_dir) {
|
||||
return Err("安全限制:只能从备份目录恢复数据库".to_string());
|
||||
}
|
||||
|
||||
if !backup_path.exists() {
|
||||
return Err("备份文件不存在".to_string());
|
||||
}
|
||||
let db_path = get_db_path()?;
|
||||
std::fs::copy(backup_path, db_path).map_err(|e| format!("恢复失败: {}", e))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn restore_database_with_connection(
|
||||
&self,
|
||||
db: &DbConnection,
|
||||
backup_path: &Path,
|
||||
) -> Result<(), String> {
|
||||
// P1 安全修复:验证备份路径在白名单目录内
|
||||
let canonical_backup = backup_path
|
||||
.canonicalize()
|
||||
.map_err(|e| format!("无法解析备份路径: {}", e))?;
|
||||
let canonical_backup_dir = self
|
||||
.backup_dir
|
||||
.canonicalize()
|
||||
.map_err(|e| format!("无法解析备份目录: {}", e))?;
|
||||
|
||||
if !canonical_backup.starts_with(&canonical_backup_dir) {
|
||||
return Err("安全限制:只能从备份目录恢复数据库".to_string());
|
||||
}
|
||||
|
||||
if !backup_path.exists() {
|
||||
return Err("备份文件不存在".to_string());
|
||||
}
|
||||
let mut conn = db.lock().map_err(|_| "数据库锁已被占用".to_string())?;
|
||||
let progress: Option<fn(rusqlite::backup::Progress)> = None;
|
||||
conn.restore(DatabaseName::Main, backup_path, progress)
|
||||
.map_err(|e| format!("恢复失败: {}", e))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn list_backups(&self) -> Result<Vec<PathBuf>, String> {
|
||||
let mut backups = Vec::new();
|
||||
let entries =
|
||||
std::fs::read_dir(&self.backup_dir).map_err(|e| format!("无法读取备份目录: {}", e))?;
|
||||
for entry in entries.flatten() {
|
||||
let path = entry.path();
|
||||
if path.extension().map(|e| e == "db").unwrap_or(false) {
|
||||
backups.push(path);
|
||||
}
|
||||
}
|
||||
backups.sort();
|
||||
Ok(backups)
|
||||
}
|
||||
|
||||
pub fn cleanup_old_backups(&self) -> Result<(), String> {
|
||||
let entries =
|
||||
std::fs::read_dir(&self.backup_dir).map_err(|e| format!("无法读取备份目录: {}", e))?;
|
||||
let cutoff = Utc::now() - Duration::days(self.retention_days as i64);
|
||||
|
||||
for entry in entries.flatten() {
|
||||
let path = entry.path();
|
||||
let Ok(metadata) = entry.metadata() else {
|
||||
continue;
|
||||
};
|
||||
let Ok(modified) = metadata.modified() else {
|
||||
continue;
|
||||
};
|
||||
let modified = DateTime::<Utc>::from(modified);
|
||||
if modified < cutoff {
|
||||
let _ = std::fs::remove_file(path);
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn backup_dir(&self) -> &PathBuf {
|
||||
&self.backup_dir
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,23 @@ use crate::models::{AppType, McpServer};
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::path::PathBuf;
|
||||
|
||||
/// P0 安全修复:校验 TOML 键名是否合法(仅允许字母、数字、下划线和连字符)
|
||||
fn is_valid_toml_key(key: &str) -> bool {
|
||||
!key.is_empty()
|
||||
&& key
|
||||
.chars()
|
||||
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
|
||||
}
|
||||
|
||||
/// P0 安全修复:转义 TOML 字符串值中的特殊字符
|
||||
fn escape_toml_string(s: &str) -> String {
|
||||
s.replace('\\', "\\\\")
|
||||
.replace('"', "\\\"")
|
||||
.replace('\n', "\\n")
|
||||
.replace('\r', "\\r")
|
||||
.replace('\t', "\\t")
|
||||
}
|
||||
|
||||
/// Get the MCP config file path for an app type
|
||||
#[allow(dead_code)]
|
||||
pub fn get_mcp_config_path(app_type: &AppType) -> Option<PathBuf> {
|
||||
@@ -136,20 +153,30 @@ fn sync_mcp_to_codex(
|
||||
|
||||
// Add new MCP server sections - use name as key
|
||||
for server in servers {
|
||||
// P0 安全修复:校验 server.name 防止 TOML 注入
|
||||
if !is_valid_toml_key(&server.name) {
|
||||
tracing::warn!(
|
||||
"[MCP Sync] 跳过无效的服务器名称: {} (仅允许字母、数字、下划线和连字符)",
|
||||
server.name
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
new_lines.push(String::new());
|
||||
new_lines.push(format!("[mcp_servers.{}]", server.name));
|
||||
|
||||
if let Some(config) = server.server_config.as_object() {
|
||||
// Convert JSON config to TOML format
|
||||
if let Some(command) = config.get("command").and_then(|v| v.as_str()) {
|
||||
new_lines.push(format!("command = \"{command}\""));
|
||||
// P0 安全修复:转义 TOML 字符串值
|
||||
new_lines.push(format!("command = \"{}\"", escape_toml_string(command)));
|
||||
}
|
||||
|
||||
if let Some(args) = config.get("args").and_then(|v| v.as_array()) {
|
||||
let args_str: Vec<String> = args
|
||||
.iter()
|
||||
.filter_map(|a| a.as_str())
|
||||
.map(|s| format!("\"{s}\""))
|
||||
.map(|s| format!("\"{}\"", escape_toml_string(s)))
|
||||
.collect();
|
||||
new_lines.push(format!("args = [{}]", args_str.join(", ")));
|
||||
}
|
||||
@@ -157,8 +184,13 @@ fn sync_mcp_to_codex(
|
||||
if let Some(env) = config.get("env").and_then(|v| v.as_object()) {
|
||||
new_lines.push("[mcp_servers.".to_string() + &server.name + ".env]");
|
||||
for (key, value) in env {
|
||||
// P0 安全修复:校验 env key 并转义值
|
||||
if !is_valid_toml_key(key) {
|
||||
tracing::warn!("[MCP Sync] 跳过无效的环境变量名: {}", key);
|
||||
continue;
|
||||
}
|
||||
if let Some(val) = value.as_str() {
|
||||
new_lines.push(format!("{key} = \"{val}\""));
|
||||
new_lines.push(format!("{} = \"{}\"", key, escape_toml_string(val)));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -532,3 +564,72 @@ pub fn import_mcp_from_app(
|
||||
AppType::ProxyCast => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{escape_toml_string, is_valid_toml_key};
|
||||
|
||||
#[test]
|
||||
fn test_valid_toml_key_accepts_alphanumeric() {
|
||||
assert!(is_valid_toml_key("abc"));
|
||||
assert!(is_valid_toml_key("ABC123"));
|
||||
assert!(is_valid_toml_key("test_server"));
|
||||
assert!(is_valid_toml_key("my-server"));
|
||||
assert!(is_valid_toml_key("server_1-test"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_valid_toml_key_rejects_invalid() {
|
||||
// 含 ] 的注入尝试
|
||||
assert!(!is_valid_toml_key("bad]"));
|
||||
assert!(!is_valid_toml_key("bad]\n[evil]"));
|
||||
// 含换行
|
||||
assert!(!is_valid_toml_key("bad\nkey"));
|
||||
// 含空格
|
||||
assert!(!is_valid_toml_key("bad key"));
|
||||
// 含点号
|
||||
assert!(!is_valid_toml_key("bad.key"));
|
||||
// 空字符串
|
||||
assert!(!is_valid_toml_key(""));
|
||||
// 含特殊字符
|
||||
assert!(!is_valid_toml_key("bad=key"));
|
||||
assert!(!is_valid_toml_key("bad[key"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_escape_toml_string_backslash() {
|
||||
assert_eq!(escape_toml_string(r"path\to\file"), r"path\\to\\file");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_escape_toml_string_quote() {
|
||||
assert_eq!(escape_toml_string(r#"say "hello""#), r#"say \"hello\""#);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_escape_toml_string_newline() {
|
||||
assert_eq!(escape_toml_string("line1\nline2"), r"line1\nline2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_escape_toml_string_carriage_return() {
|
||||
assert_eq!(escape_toml_string("line1\rline2"), r"line1\rline2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_escape_toml_string_tab() {
|
||||
assert_eq!(escape_toml_string("col1\tcol2"), r"col1\tcol2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_escape_toml_string_combined() {
|
||||
let input = "path\\to\\file\nwith \"quotes\"\tand\rtabs";
|
||||
let output = escape_toml_string(input);
|
||||
assert!(!output.contains('\n'));
|
||||
assert!(!output.contains('\r'));
|
||||
assert!(!output.contains('\t'));
|
||||
assert!(output.contains("\\n"));
|
||||
assert!(output.contains("\\r"));
|
||||
assert!(output.contains("\\t"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
pub mod backup_service;
|
||||
pub mod live_sync;
|
||||
pub mod mcp_service;
|
||||
pub mod mcp_sync;
|
||||
|
||||
@@ -487,8 +487,12 @@ impl ProviderPoolService {
|
||||
self.check_gemini_api_key_health(api_key, base_url.as_deref(), model)
|
||||
.await
|
||||
}
|
||||
CredentialData::CodexOAuth { creds_file_path } => {
|
||||
self.check_codex_health(creds_file_path, model).await
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path,
|
||||
api_base_url,
|
||||
} => {
|
||||
self.check_codex_health(creds_file_path, api_base_url.as_deref(), model)
|
||||
.await
|
||||
}
|
||||
CredentialData::ClaudeOAuth { creds_file_path } => {
|
||||
self.check_claude_oauth_health(creds_file_path, model).await
|
||||
@@ -924,7 +928,13 @@ impl ProviderPoolService {
|
||||
}
|
||||
|
||||
// Codex 健康检查
|
||||
async fn check_codex_health(&self, creds_path: &str, model: &str) -> Result<(), String> {
|
||||
// 支持 Yunyi 等代理使用 responses API 格式
|
||||
async fn check_codex_health(
|
||||
&self,
|
||||
creds_path: &str,
|
||||
override_base_url: Option<&str>,
|
||||
model: &str,
|
||||
) -> Result<(), String> {
|
||||
use crate::providers::codex::CodexProvider;
|
||||
|
||||
let mut provider = CodexProvider::new();
|
||||
@@ -933,33 +943,85 @@ impl ProviderPoolService {
|
||||
.await
|
||||
.map_err(|e| format!("加载 Codex 凭证失败: {}", e))?;
|
||||
|
||||
let token = provider
|
||||
.ensure_valid_token()
|
||||
.await
|
||||
.map_err(|e| format!("获取 Codex Token 失败: {}", e))?;
|
||||
let token = provider.ensure_valid_token().await.map_err(|e| {
|
||||
format!(
|
||||
"获取 Codex Token 失败: 配置错误,请检查凭证设置。详情:{}",
|
||||
e
|
||||
)
|
||||
})?;
|
||||
|
||||
// 使用 OpenAI 兼容 API 进行健康检查
|
||||
let url = "https://api.openai.com/v1/chat/completions";
|
||||
let request_body = serde_json::json!({
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "Say OK"}],
|
||||
"max_tokens": 10
|
||||
});
|
||||
// 优先使用 override_base_url(来自 CredentialData),其次使用凭证文件中的配置
|
||||
let base_url = override_base_url
|
||||
.map(|s| s.trim())
|
||||
.filter(|s| !s.is_empty())
|
||||
.or_else(|| {
|
||||
provider
|
||||
.credentials
|
||||
.api_base_url
|
||||
.as_deref()
|
||||
.map(|s| s.trim())
|
||||
.filter(|s| !s.is_empty())
|
||||
});
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(url)
|
||||
.header("Authorization", format!("Bearer {}", token))
|
||||
.json(&request_body)
|
||||
.timeout(self.health_check_timeout)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("请求失败: {}", e))?;
|
||||
match base_url {
|
||||
Some(base) => {
|
||||
// 使用自定义 base_url (如 Yunyi),与 CodexProvider 的 URL/headers 行为保持一致
|
||||
let url = CodexProvider::build_responses_url(base);
|
||||
|
||||
if response.status().is_success() {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!("HTTP {}", response.status()))
|
||||
// Codex/Yunyi 使用 responses API 格式;云驿等代理要求 stream 必须为 true
|
||||
let request_body = serde_json::json!({
|
||||
"model": model,
|
||||
"input": [{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "Say OK"}]
|
||||
}],
|
||||
"max_output_tokens": 10,
|
||||
"stream": true
|
||||
});
|
||||
|
||||
tracing::debug!(
|
||||
"[HEALTH_CHECK] Codex responses API URL: {}, model: {}",
|
||||
url,
|
||||
model
|
||||
);
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(&url)
|
||||
.bearer_auth(&token)
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "text/event-stream")
|
||||
.header("Openai-Beta", "responses=experimental")
|
||||
.header("Originator", "codex_cli_rs")
|
||||
.header("Session_id", uuid::Uuid::new_v4().to_string())
|
||||
.header("Conversation_id", uuid::Uuid::new_v4().to_string())
|
||||
.header(
|
||||
"User-Agent",
|
||||
"codex_cli_rs/0.77.0 (ProxyCast health check; Mac OS; arm64)",
|
||||
)
|
||||
.json(&request_body)
|
||||
.timeout(self.health_check_timeout)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("请求失败: {}", e))?;
|
||||
|
||||
if response.status().is_success() {
|
||||
Ok(())
|
||||
} else {
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
Err(format!(
|
||||
"HTTP {} - {}",
|
||||
status,
|
||||
body.chars().take(200).collect::<String>()
|
||||
))
|
||||
}
|
||||
}
|
||||
None => {
|
||||
// 没有自定义 base_url,使用 OpenAI 官方 chat/completions API
|
||||
self.check_openai_health(&token, None, model).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1165,12 +1227,20 @@ impl ProviderPoolService {
|
||||
let creds: serde_json::Value =
|
||||
serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {}", e))?;
|
||||
|
||||
let has_access_token = creds
|
||||
let has_api_key = creds
|
||||
.get("apiKey")
|
||||
.or_else(|| creds.get("api_key"))
|
||||
.map(|v| v.as_str().is_some())
|
||||
.unwrap_or(false);
|
||||
|
||||
let has_oauth_access_token = creds
|
||||
.get("accessToken")
|
||||
.or_else(|| creds.get("access_token"))
|
||||
.map(|v| v.as_str().is_some())
|
||||
.unwrap_or(false);
|
||||
|
||||
let has_access_token = has_oauth_access_token || has_api_key;
|
||||
|
||||
let has_refresh_token = creds
|
||||
.get("refreshToken")
|
||||
.or_else(|| creds.get("refresh_token"))
|
||||
@@ -1201,6 +1271,20 @@ impl ProviderPoolService {
|
||||
(has_access_token, None)
|
||||
}
|
||||
}
|
||||
"codex" => {
|
||||
// Codex: 兼容 OAuth token 或 Codex CLI 的 API Key 登录
|
||||
if has_api_key {
|
||||
(true, None)
|
||||
} else {
|
||||
let expires_at = creds
|
||||
.get("expiresAt")
|
||||
.or_else(|| creds.get("expires_at"))
|
||||
.or_else(|| creds.get("expired"))
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
(has_oauth_access_token, expires_at)
|
||||
}
|
||||
}
|
||||
_ => (has_access_token, None),
|
||||
};
|
||||
|
||||
|
||||
@@ -282,9 +282,9 @@ impl TokenCacheService {
|
||||
last_refresh_error: None,
|
||||
})
|
||||
}
|
||||
CredentialData::CodexOAuth { creds_file_path } => {
|
||||
self.refresh_codex(creds_file_path).await
|
||||
}
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path, ..
|
||||
} => self.refresh_codex(creds_file_path).await,
|
||||
CredentialData::ClaudeOAuth { creds_file_path } => {
|
||||
self.refresh_claude_oauth(creds_file_path).await
|
||||
}
|
||||
@@ -764,7 +764,9 @@ impl TokenCacheService {
|
||||
refresh_error_count: 0,
|
||||
last_refresh_error: None,
|
||||
}),
|
||||
CredentialData::CodexOAuth { creds_file_path } => {
|
||||
CredentialData::CodexOAuth {
|
||||
creds_file_path, ..
|
||||
} => {
|
||||
let content = tokio::fs::read_to_string(creds_file_path)
|
||||
.await
|
||||
.map_err(|e| format!("读取 Codex 凭证文件失败: {}", e))?;
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,281 @@
|
||||
//! 流式传输错误类型
|
||||
//!
|
||||
//! 定义流式传输过程中可能发生的各种错误类型。
|
||||
//!
|
||||
//! # 需求覆盖
|
||||
//!
|
||||
//! - 需求 6.1: 网络错误处理
|
||||
//! - 需求 6.2: 超时错误处理
|
||||
//! - 需求 6.3: Provider 错误转发
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
|
||||
/// 流式传输错误类型
|
||||
///
|
||||
/// 涵盖流式传输过程中可能发生的所有错误情况。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(tag = "type", content = "details")]
|
||||
pub enum StreamError {
|
||||
/// 网络错误
|
||||
///
|
||||
/// 当网络连接失败、DNS 解析失败或连接被重置时发生。
|
||||
/// 对应需求 6.1
|
||||
Network(String),
|
||||
|
||||
/// 超时错误
|
||||
///
|
||||
/// 当流式响应超过配置的超时时间时发生。
|
||||
/// 对应需求 6.2
|
||||
Timeout,
|
||||
|
||||
/// 解析错误
|
||||
///
|
||||
/// 当无法解析流式数据(如无效的 AWS Event Stream 或 SSE 格式)时发生。
|
||||
ParseError(String),
|
||||
|
||||
/// Provider 错误
|
||||
///
|
||||
/// 当上游 Provider 返回错误响应时发生。
|
||||
/// 对应需求 6.3
|
||||
ProviderError {
|
||||
/// HTTP 状态码
|
||||
status: u16,
|
||||
/// 错误消息
|
||||
message: String,
|
||||
},
|
||||
|
||||
/// 客户端断开连接
|
||||
///
|
||||
/// 当客户端在流式传输过程中断开连接时发生。
|
||||
ClientDisconnected,
|
||||
|
||||
/// 缓冲区溢出
|
||||
///
|
||||
/// 当流式数据超过配置的缓冲区大小时发生。
|
||||
BufferOverflow,
|
||||
|
||||
/// 内部错误
|
||||
///
|
||||
/// 其他内部错误。
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
impl fmt::Display for StreamError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
StreamError::Network(msg) => write!(f, "网络错误: {}", msg),
|
||||
StreamError::Timeout => write!(f, "流式响应超时"),
|
||||
StreamError::ParseError(msg) => write!(f, "解析错误: {}", msg),
|
||||
StreamError::ProviderError { status, message } => {
|
||||
write!(f, "Provider 错误 ({}): {}", status, message)
|
||||
}
|
||||
StreamError::ClientDisconnected => write!(f, "客户端已断开连接"),
|
||||
StreamError::BufferOverflow => write!(f, "缓冲区溢出"),
|
||||
StreamError::Internal(msg) => write!(f, "内部错误: {}", msg),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for StreamError {}
|
||||
|
||||
// ============================================================================
|
||||
// From trait 实现 - 用于错误转换
|
||||
// ============================================================================
|
||||
|
||||
impl From<std::io::Error> for StreamError {
|
||||
fn from(err: std::io::Error) -> Self {
|
||||
StreamError::Network(err.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl From<reqwest::Error> for StreamError {
|
||||
fn from(err: reqwest::Error) -> Self {
|
||||
if err.is_timeout() {
|
||||
StreamError::Timeout
|
||||
} else if err.is_connect() {
|
||||
StreamError::Network(format!("连接失败: {}", err))
|
||||
} else if err.is_request() {
|
||||
StreamError::Network(format!("请求错误: {}", err))
|
||||
} else {
|
||||
StreamError::Network(err.to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<serde_json::Error> for StreamError {
|
||||
fn from(err: serde_json::Error) -> Self {
|
||||
StreamError::ParseError(err.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for StreamError {
|
||||
fn from(msg: String) -> Self {
|
||||
StreamError::Internal(msg)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for StreamError {
|
||||
fn from(msg: &str) -> Self {
|
||||
StreamError::Internal(msg.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 辅助方法
|
||||
// ============================================================================
|
||||
|
||||
impl StreamError {
|
||||
/// 创建网络错误
|
||||
pub fn network(msg: impl Into<String>) -> Self {
|
||||
StreamError::Network(msg.into())
|
||||
}
|
||||
|
||||
/// 创建解析错误
|
||||
pub fn parse_error(msg: impl Into<String>) -> Self {
|
||||
StreamError::ParseError(msg.into())
|
||||
}
|
||||
|
||||
/// 创建 Provider 错误
|
||||
pub fn provider_error(status: u16, message: impl Into<String>) -> Self {
|
||||
StreamError::ProviderError {
|
||||
status,
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建内部错误
|
||||
pub fn internal(msg: impl Into<String>) -> Self {
|
||||
StreamError::Internal(msg.into())
|
||||
}
|
||||
|
||||
/// 判断错误是否可重试
|
||||
///
|
||||
/// 网络错误、超时和某些 Provider 错误(如 429、5xx)可以重试。
|
||||
pub fn is_retryable(&self) -> bool {
|
||||
match self {
|
||||
StreamError::Network(_) => true,
|
||||
StreamError::Timeout => true,
|
||||
StreamError::ProviderError { status, .. } => *status == 429 || *status >= 500,
|
||||
StreamError::ParseError(_) => false,
|
||||
StreamError::ClientDisconnected => false,
|
||||
StreamError::BufferOverflow => false,
|
||||
StreamError::Internal(_) => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// 判断是否为客户端错误
|
||||
pub fn is_client_error(&self) -> bool {
|
||||
matches!(self, StreamError::ClientDisconnected)
|
||||
}
|
||||
|
||||
/// 获取 HTTP 状态码(如果适用)
|
||||
pub fn status_code(&self) -> Option<u16> {
|
||||
match self {
|
||||
StreamError::ProviderError { status, .. } => Some(*status),
|
||||
StreamError::Timeout => Some(504), // Gateway Timeout
|
||||
StreamError::Network(_) => Some(502), // Bad Gateway
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 转换为 SSE 错误事件格式
|
||||
pub fn to_sse_error(&self) -> String {
|
||||
let error_json = serde_json::json!({
|
||||
"error": {
|
||||
"type": self.error_type_string(),
|
||||
"message": self.to_string(),
|
||||
}
|
||||
});
|
||||
format!("event: error\ndata: {}\n\n", error_json)
|
||||
}
|
||||
|
||||
/// 获取错误类型字符串
|
||||
fn error_type_string(&self) -> &'static str {
|
||||
match self {
|
||||
StreamError::Network(_) => "network_error",
|
||||
StreamError::Timeout => "timeout",
|
||||
StreamError::ParseError(_) => "parse_error",
|
||||
StreamError::ProviderError { .. } => "provider_error",
|
||||
StreamError::ClientDisconnected => "client_disconnected",
|
||||
StreamError::BufferOverflow => "buffer_overflow",
|
||||
StreamError::Internal(_) => "internal_error",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 测试模块
|
||||
// ============================================================================
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_stream_error_display() {
|
||||
let err = StreamError::Network("connection refused".to_string());
|
||||
assert_eq!(err.to_string(), "网络错误: connection refused");
|
||||
|
||||
let err = StreamError::Timeout;
|
||||
assert_eq!(err.to_string(), "流式响应超时");
|
||||
|
||||
let err = StreamError::provider_error(429, "rate limited");
|
||||
assert_eq!(err.to_string(), "Provider 错误 (429): rate limited");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_stream_error_is_retryable() {
|
||||
assert!(StreamError::Network("test".to_string()).is_retryable());
|
||||
assert!(StreamError::Timeout.is_retryable());
|
||||
assert!(StreamError::provider_error(429, "rate limited").is_retryable());
|
||||
assert!(StreamError::provider_error(500, "server error").is_retryable());
|
||||
assert!(!StreamError::provider_error(400, "bad request").is_retryable());
|
||||
assert!(!StreamError::ParseError("invalid json".to_string()).is_retryable());
|
||||
assert!(!StreamError::ClientDisconnected.is_retryable());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_stream_error_status_code() {
|
||||
assert_eq!(StreamError::Timeout.status_code(), Some(504));
|
||||
assert_eq!(
|
||||
StreamError::Network("test".to_string()).status_code(),
|
||||
Some(502)
|
||||
);
|
||||
assert_eq!(
|
||||
StreamError::provider_error(429, "test").status_code(),
|
||||
Some(429)
|
||||
);
|
||||
assert_eq!(StreamError::ClientDisconnected.status_code(), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_stream_error_from_io_error() {
|
||||
let io_err = std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "refused");
|
||||
let stream_err: StreamError = io_err.into();
|
||||
assert!(matches!(stream_err, StreamError::Network(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_stream_error_from_serde_json_error() {
|
||||
let json_err = serde_json::from_str::<serde_json::Value>("invalid").unwrap_err();
|
||||
let stream_err: StreamError = json_err.into();
|
||||
assert!(matches!(stream_err, StreamError::ParseError(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_stream_error_serialization() {
|
||||
let err = StreamError::provider_error(500, "internal server error");
|
||||
let json = serde_json::to_string(&err).unwrap();
|
||||
let deserialized: StreamError = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(err, deserialized);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_stream_error_to_sse_error() {
|
||||
let err = StreamError::Timeout;
|
||||
let sse = err.to_sse_error();
|
||||
assert!(sse.starts_with("event: error\n"));
|
||||
assert!(sse.contains("timeout"));
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user