Compare commits

..

10 Commits

Author SHA1 Message Date
YoVinchen d11a12a746 Merge branch 'main' into feature/error-request-logging 2025-12-16 18:09:19 +08:00
YoVinchen e76ee66d13 chore(clippy): fix uninlined format args 2025-12-16 18:02:08 +08:00
YoVinchen a540bf92ca fix(speedtest): skip client build for invalid inputs 2025-12-16 17:50:41 +08:00
YoVinchen 24391fc431 Merge remote-tracking branch 'origin/main' into feature/error-request-logging
# Conflicts:
#	src-tauri/src/proxy/handlers.rs
#	src-tauri/src/proxy/mod.rs
#	src-tauri/src/proxy/provider_router.rs
#	src-tauri/src/services/proxy.rs
#	src/components/providers/ProviderActions.tsx
2025-12-16 17:40:49 +08:00
YoVinchen 2e6ba77187 feat(proxy): add settings button to proxy panel
Add configuration buttons in both running and stopped states to
provide easy access to proxy settings dialog.
2025-12-15 00:16:12 +08:00
YoVinchen 53ccd5f70d style: apply code formatting
- Remove trailing whitespace in misc.rs
- Add trailing comma in App.tsx
- Format multi-line className in ProviderCard.tsx
2025-12-14 16:16:56 +08:00
YoVinchen 2af8dd2dac style: fix clippy warnings and typescript errors
- Add allow(dead_code) for CircuitBreaker::get_state (reserved for future)
- Fix all uninlined format string warnings (27 instances)
- Use inline format syntax for better readability
- Fix unused import and parameter warnings in ProviderActions.tsx
- Achieve zero warnings in both Rust and TypeScript
2025-12-14 16:05:38 +08:00
YoVinchen 4cf4654863 feat(proxy): implement error capture and logging in all handlers
- Capture and log all failed requests in handle_messages (Claude)
- Capture and log all failed requests in handle_gemini (Gemini)
- Capture and log all failed requests in handle_responses (Codex)
- Capture and log all failed requests in handle_chat_completions (Codex)
- Record error status codes, messages, and latency for all failures
- Generate unique session_id for each request
- Support both streaming and non-streaming error scenarios
2025-12-14 16:03:28 +08:00
YoVinchen 8b202ea988 feat(proxy): enhance error logging with context support
- Add log_error_with_context() method for detailed error recording
- Support streaming flag, session_id, and provider_type fields
- Remove dead_code warning from log_error() method
- Enable comprehensive error request tracking in database
2025-12-14 16:03:02 +08:00
YoVinchen 14bc8a00e5 feat(proxy): add error mapper for HTTP status code mapping
- Add error_mapper.rs module to map ProxyError to HTTP status codes
- Implement map_proxy_error_to_status() for error classification
- Implement get_error_message() for user-friendly error messages
- Support all error types: upstream, timeout, connection, provider failures
- Include comprehensive unit tests for all mappings
2025-12-14 16:02:04 +08:00
183 changed files with 3421 additions and 13664 deletions
+1 -1
View File
@@ -1 +1 @@
22.12.0 v22.4.1
-152
View File
@@ -5,158 +5,6 @@ All notable changes to CC Switch will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [3.9.0-3] - 2025-12-29
### Beta Release
Third beta release with important bug fixes for Windows compatibility, UI improvements, and new features.
### Added
- **Universal Provider** - Support for universal provider configurations (#348)
- **Provider Search Filter** - Quick filter to find providers by name (#435)
- **Keyboard Shortcut** - Open settings with Command+comma / Ctrl+comma (#436)
- **Xiaomi MiMo Icon** - Added MiMo icon and Claude provider configuration (#470)
- **Usage Model Extraction** - Extract model info from usage statistics (#455)
- **Skip First-Run Confirmation** - Option to skip Claude Code first-run confirmation dialog
- **Exit Animations** - Added exit animation to FullScreenPanel dialogs
- **Fade Transitions** - Smooth fade transitions for app/view/panel switching
### Fixed
#### Windows
- Wrap npx/npm commands with `cmd /c` for MCP export
- Prevent terminal windows from appearing during version check
#### macOS
- Use .app bundle path for autostart to prevent terminal window popup
#### UI
- Resolve Dialog/Modal not opening on first click (#492)
- Improve dark mode text contrast for form labels
- Reduce header spacing and fix layout shift on view switch
- Prevent header layout shift when switching views
#### Database & Schema
- Add missing base columns migration for proxy_config
- Add backward compatibility check for proxy_config seed insert
#### Other
- Use local timezone and robust DST handling in usage stats (#500)
- Remove deprecated `sync_enabled_to_codex` call
- Gracefully handle invalid Codex config.toml during MCP sync
- Add missing translations for reasoning model and OpenRouter compat mode
### Improved
- **macOS Tray** - Use macOS tray template icon
- **Header Alignment** - Remove macOS titlebar tint, align custom header
- **Shadow Removal** - Cleaner UI by removing shadow styles
- **Code Inspector** - Added code-inspector-plugin for development
- **i18n** - Complete internationalization for usage panel and settings
- **Sponsor Logos** - Made sponsor logos clickable
### Stats
- 35 commits since v3.9.0-2
- 5 files changed in test/lint fixes
---
## [3.9.0-1] - 2025-12-18
### Beta Release
This beta release introduces the **Local API Proxy** feature, along with Skills multi-app support, UI improvements, and numerous bug fixes.
### Major Features
#### Local Proxy Server
- **Local HTTP Proxy** - High-performance proxy server built on Axum framework
- **Multi-app Support** - Unified proxy for Claude Code, Codex, and Gemini CLI API requests
- **Per-app Takeover** - Independent control over which apps route through the proxy
- **Live Config Takeover** - Automatically backs up and redirects CLI configurations to local proxy
#### Auto Failover
- **Circuit Breaker** - Automatically detects provider failures and triggers protection
- **Smart Failover** - Automatically switches to backup provider when current one is unavailable
- **Health Tracking** - Real-time monitoring of provider availability
- **Independent Failover Queues** - Each app maintains its own failover queue
#### Monitoring
- **Request Logging** - Detailed logging of all proxy requests
- **Usage Statistics** - Token consumption, latency, success rate metrics
- **Real-time Status** - Frontend displays proxy status and statistics
#### Skills Multi-App Support
- **Multi-app Support** - Skills now support both Claude and Codex (#365)
- **Multi-app Migration** - Existing Skills auto-migrate to multi-app structure (#378)
- **Installation Path Fix** - Use directory basename for skill installation path (#358)
### Added
- **Provider Icon Colors** - Customize provider icon colors (#385)
- **Deeplink Usage Config** - Import usage query config via deeplink (#400)
- **Error Request Logging** - Detailed logging for proxy requests (#401)
- **Closable Toast** - Added close button to switch notification toast (#350)
- **Icon Color Component** - ProviderIcon component supports color prop (#384)
### Fixed
#### Proxy Related
- Takeover Codex base_url via model_provider
- Harden crash recovery with fallback detection
- Sync UI when active provider differs from current setting
- Resolve circuit breaker race condition and error classification
- Stabilize live takeover and provider editing
- Reset health badges when proxy stops
- Retry failover for all HTTP errors including 4xx
- Fix HalfOpen counter underflow and config field inconsistencies
- Resolve circuit breaker state persistence and HalfOpen deadlock
- Auto-recover live config after abnormal exit
- Update live backup when hot-switching provider in proxy mode
- Wait for server shutdown before exiting app
- Disable auto-start on app launch by resetting enabled flag on stop
- Sync live config tokens to database before takeover
- Resolve 404 error and auto-setup proxy targets
#### MCP Related
- Skip sync when target CLI app is not installed
- Improve upsert and import robustness
- Use browser-compatible platform detection for MCP presets
#### UI Related
- Restore fade transition for Skills button
- Add close button to all success toasts
- Prevent card jitter when health badge appears
- Update SettingsPage tab styles (#342)
#### Other
- Fix Azure website link (#407)
- Add fallback to provider config for usage credentials (#360)
- Fix Windows black screen on startup (use system titlebar)
- Add fallback for crypto.randomUUID() on older WebViews
- Use correct npm package for Codex CLI version check
- Security fixes for JavaScript executor and usage script (#151)
### Improved
- **Proxy Active Theme** - Apply emerald theme when proxy takeover is active
- **Card Animation** - Improved provider card hover animation
- **Remove Restart Prompt** - No longer prompts restart when switching providers
### Technical
- Implement per-app takeover mode
- Proxy module contains 20+ Rust files with complete layered architecture
- Add 5 new database tables for proxy functionality
- Modularize handlers.rs to reduce code duplication
- Remove is_proxy_target in favor of failover_queue
### Stats
- 55 commits since v3.8.2
- 164 files changed
- +22,164 / -570 lines
---
## [3.8.0] - 2025-11-28 ## [3.8.0] - 2025-11-28
### Major Updates ### Major Updates
+9 -8
View File
@@ -2,7 +2,7 @@
# All-in-One Assistant for Claude Code, Codex & Gemini CLI # All-in-One Assistant for Claude Code, Codex & Gemini CLI
[![Version](https://img.shields.io/badge/version-3.8.3-blue.svg)](https://github.com/farion1231/cc-switch/releases) [![Version](https://img.shields.io/badge/version-3.8.2-blue.svg)](https://github.com/farion1231/cc-switch/releases)
[![Platform](https://img.shields.io/badge/platform-Windows%20%7C%20macOS%20%7C%20Linux-lightgrey.svg)](https://github.com/farion1231/cc-switch/releases) [![Platform](https://img.shields.io/badge/platform-Windows%20%7C%20macOS%20%7C%20Linux-lightgrey.svg)](https://github.com/farion1231/cc-switch/releases)
[![Built with Tauri](https://img.shields.io/badge/built%20with-Tauri%202-orange.svg)](https://tauri.app/) [![Built with Tauri](https://img.shields.io/badge/built%20with-Tauri%202-orange.svg)](https://tauri.app/)
[![Downloads](https://img.shields.io/endpoint?url=https://api.pinstudios.net/api/badges/downloads/farion1231/cc-switch/total)](https://github.com/farion1231/cc-switch/releases/latest) [![Downloads](https://img.shields.io/endpoint?url=https://api.pinstudios.net/api/badges/downloads/farion1231/cc-switch/total)](https://github.com/farion1231/cc-switch/releases/latest)
@@ -15,7 +15,7 @@ English | [中文](README_ZH.md) | [日本語](README_JA.md) | [Changelog](CHANG
## ❤️Sponsor ## ❤️Sponsor
[![Zhipu GLM](assets/partners/banners/glm-en.jpg)](https://z.ai/subscribe?ic=8JVLJQFSKB) ![Zhipu GLM](assets/partners/banners/glm-en.jpg)
This project is sponsored by Z.ai, supporting us with their GLM CODING PLAN.GLM CODING PLAN is a subscription service designed for AI coding, starting at just $3/month. It provides access to their flagship GLM-4.6 model across 10+ popular AI coding tools (Claude Code, Cline, Roo Code, etc.), offering developers top-tier, fast, and stable coding experiences.Get 10% OFF the GLM CODING PLAN with [this link](https://z.ai/subscribe?ic=8JVLJQFSKB)! This project is sponsored by Z.ai, supporting us with their GLM CODING PLAN.GLM CODING PLAN is a subscription service designed for AI coding, starting at just $3/month. It provides access to their flagship GLM-4.6 model across 10+ popular AI coding tools (Claude Code, Cline, Roo Code, etc.), offering developers top-tier, fast, and stable coding experiences.Get 10% OFF the GLM CODING PLAN with [this link](https://z.ai/subscribe?ic=8JVLJQFSKB)!
@@ -23,18 +23,19 @@ This project is sponsored by Z.ai, supporting us with their GLM CODING PLAN.GLM
<table> <table>
<tr> <tr>
<td width="180"><a href="https://www.packyapi.com/register?aff=cc-switch"><img src="assets/partners/logos/packycode.png" alt="PackyCode" width="150"></a></td> <td width="180"><img src="assets/partners/logos/packycode.png" alt="PackyCode" width="150"></td>
<td>Thanks to PackyCode for sponsoring this project! PackyCode is a reliable and efficient API relay service provider, offering relay services for Claude Code, Codex, Gemini, and more. PackyCode provides special discounts for our software users: register using <a href="https://www.packyapi.com/register?aff=cc-switch">this link</a> and enter the "cc-switch" promo code during recharge to get 10% off.</td> <td>Thanks to PackyCode for sponsoring this project! PackyCode is a reliable and efficient API relay service provider, offering relay services for Claude Code, Codex, Gemini, and more. PackyCode provides special discounts for our software users: register using <a href="https://www.packyapi.com/register?aff=cc-switch">this link</a> and enter the "cc-switch" promo code during recharge to get 10% off.</td>
</tr> </tr>
<tr> <tr>
<td width="180"><a href="https://aigocode.com/invite/CC-SWITCH"><img src="assets/partners/logos/aigocode.png" alt="AIGoCode" width="150"></a></td> <td width="180"><img src="assets/partners/logos/sds-en.png" alt="ShanDianShuo" width="150"></td>
<td>Thanks to AIGoCode for sponsoring this project! AIGoCode is an all-in-one platform that integrates Claude Code, Codex, and the latest Gemini models, providing you with stable, efficient, and highly cost-effective AI coding services. The platform offers flexible subscription plans, zero risk of account suspension, direct access with no VPN required, and lightning-fast responses. AIGoCode has prepared a special benefit for CC Switch users: if you register via <a href="https://aigocode.com/invite/CC-SWITCH">this link</a>, you'll receive an extra 10% bonus credit on your first top-up!</td> <td>Thanks to ShanDianShuo for sponsoring this project! ShanDianShuo is a local-first AI voice input: Millisecond latency, data stays on device, 4x faster than typing, AI-powered correction, Privacy-first, completely free. Doubles your coding efficiency with Claude Code! <a href="https://www.shandianshuo.cn">Free download</a> for Mac/Win</td>
</tr> </tr>
<tr> <tr>
<td width="180"><a href="https://www.dmxapi.cn/register?aff=bUHu"><img src="assets/partners/logos/dmx-en.jpg" alt="DMXAPI" width="150"></a></td> <td width="180"><img src="assets/partners/logos/aigocode.png" alt="AIGoCode" width="150"></td>
<td>Thanks to DMXAPI for sponsoring this project! DMXAPI provides global large model API services to 200+ enterprise users. One API key for all global models. Features include: instant invoicing, unlimited concurrency, starting from $0.15, 24/7 technical support. GPT/Claude/Gemini all at 32% off, domestic models 20-50% off, Claude Code exclusive models at 66% off! <a href="https://www.dmxapi.cn/register?aff=bUHu">Register here</a></td> <td>Thanks to AIGoCode for sponsoring this project! AIGoCode is an all-in-one platform that integrates Claude Code, Codex, and the latest Gemini models, providing you with stable, efficient, and highly cost-effective AI coding services. The platform offers flexible subscription plans, zero risk of account suspension, direct access with no VPN required, and lightning-fast responses.AIGoCode has prepared a special benefit for CC Switch users: if you register via <a href="https://aigocode.com/invite/CC-SWITCH">this link</a>, youll receive an extra 10% bonus credit on your first top-up!
</td>
</tr> </tr>
</table> </table>
@@ -47,7 +48,7 @@ This project is sponsored by Z.ai, supporting us with their GLM CODING PLAN.GLM
## Features ## Features
### Current Version: v3.8.3 | [Full Changelog](CHANGELOG.md) | [Release Notes](docs/release-note-v3.8.0-en.md) ### Current Version: v3.8.2 | [Full Changelog](CHANGELOG.md) | [Release Notes](docs/release-note-v3.8.0-en.md)
**v3.8.0 Major Update (2025-11-28)** **v3.8.0 Major Update (2025-11-28)**
+10 -8
View File
@@ -2,7 +2,7 @@
# Claude Code / Codex / Gemini CLI オールインワン・アシスタント # Claude Code / Codex / Gemini CLI オールインワン・アシスタント
[![Version](https://img.shields.io/badge/version-3.8.3-blue.svg)](https://github.com/farion1231/cc-switch/releases) [![Version](https://img.shields.io/badge/version-3.8.2-blue.svg)](https://github.com/farion1231/cc-switch/releases)
[![Platform](https://img.shields.io/badge/platform-Windows%20%7C%20macOS%20%7C%20Linux-lightgrey.svg)](https://github.com/farion1231/cc-switch/releases) [![Platform](https://img.shields.io/badge/platform-Windows%20%7C%20macOS%20%7C%20Linux-lightgrey.svg)](https://github.com/farion1231/cc-switch/releases)
[![Built with Tauri](https://img.shields.io/badge/built%20with-Tauri%202-orange.svg)](https://tauri.app/) [![Built with Tauri](https://img.shields.io/badge/built%20with-Tauri%202-orange.svg)](https://tauri.app/)
[![Downloads](https://img.shields.io/endpoint?url=https://api.pinstudios.net/api/badges/downloads/farion1231/cc-switch/total)](https://github.com/farion1231/cc-switch/releases/latest) [![Downloads](https://img.shields.io/endpoint?url=https://api.pinstudios.net/api/badges/downloads/farion1231/cc-switch/total)](https://github.com/farion1231/cc-switch/releases/latest)
@@ -15,7 +15,7 @@
## ❤️スポンサー ## ❤️スポンサー
[![Zhipu GLM](assets/partners/banners/glm-en.jpg)](https://z.ai/subscribe?ic=8JVLJQFSKB) ![Zhipu GLM](assets/partners/banners/glm-en.jpg)
本プロジェクトは Z.ai の GLM CODING PLAN による支援を受けています。GLM CODING PLAN は AI コーディング向けのサブスクリプションで、月額わずか 3 ドルから。Claude Code、Cline、Roo Code など 10 以上の人気 AI コーディングツールでフラッグシップモデル GLM-4.6 を利用でき、速く安定した開発体験を提供します。[このリンク](https://z.ai/subscribe?ic=8JVLJQFSKB) から申し込むと 10% オフになります! 本プロジェクトは Z.ai の GLM CODING PLAN による支援を受けています。GLM CODING PLAN は AI コーディング向けのサブスクリプションで、月額わずか 3 ドルから。Claude Code、Cline、Roo Code など 10 以上の人気 AI コーディングツールでフラッグシップモデル GLM-4.6 を利用でき、速く安定した開発体験を提供します。[このリンク](https://z.ai/subscribe?ic=8JVLJQFSKB) から申し込むと 10% オフになります!
@@ -23,18 +23,20 @@
<table> <table>
<tr> <tr>
<td width="180"><a href="https://www.packyapi.com/register?aff=cc-switch"><img src="assets/partners/logos/packycode.png" alt="PackyCode" width="150"></a></td> <td width="180"><img src="assets/partners/logos/packycode.png" alt="PackyCode" width="150"></td>
<td>PackyCode のご支援に感謝します!PackyCode は Claude Code、Codex、Gemini などのリレーサービスを提供する信頼性の高い API 中継プラットフォームです。本ソフト利用者向けに特別割引があります:<a href="https://www.packyapi.com/register?aff=cc-switch">このリンク</a>で登録し、チャージ時に「cc-switch」クーポンを入力すると 10% オフになります。</td> <td>PackyCode のご支援に感謝します!PackyCode は Claude Code、Codex、Gemini などのリレーサービスを提供する信頼性の高い API 中継プラットフォームです。本ソフト利用者向けに特別割引があります:<a href="https://www.packyapi.com/register?aff=cc-switch">このリンク</a>で登録し、チャージ時に「cc-switch」クーポンを入力すると 10% オフになります。</td>
</tr> </tr>
<tr> <tr>
<td width="180"><a href="https://aigocode.com/invite/CC-SWITCH"><img src="assets/partners/logos/aigocode.png" alt="AIGoCode" width="150"></a></td> <td width="180"><img src="assets/partners/logos/sds-en.png" alt="ShanDianShuo" width="150"></td>
<td>本プロジェクトは AIGoCode のスポンサー提供でお届けしています。AIGoCode は、Claude Code・Codex・最新の Gemini モデルを統合したオールインワンのAIコーディングプラットフォームで、安定性・高速性・コストパフォーマンスに優れた開発サービスを提供します。柔軟なサブスクリプションプランを備え、レスポンスも非常に高速です。さらに、CC Switch ユーザー向けの特典として、<a href="https://aigocode.com/invite/CC-SWITCH">このリンク</a>から登録すると、初回チャージ時に10%分のボーナスクレジットが付与されます!</td> <td>ShanDianShuo のご支援に感謝します!ShanDianShuo はローカルファーストの音声入力ツールで、ミリ秒遅延・データは端末から外に出ず・キーボード入力の 4 倍の速度・AI 自動補正・プライバシー優先で完全無料。Claude Code と組み合わせればコーディング効率が倍増します。<a href="https://www.shandianshuo.cn">Mac/Win 版を無料ダウンロード</a></td>
</tr> </tr>
<tr> <tr>
<td width="180"><a href="https://www.dmxapi.cn/register?aff=bUHu"><img src="assets/partners/logos/dmx-en.jpg" alt="DMXAPI" width="150"></a></td> <td width="180"><img src="assets/partners/logos/aigocode.png" alt="AIGoCode" width="150"></td>
<td>DMXAPI のご支援に感謝します!DMXAPI は 200 社以上の企業ユーザーにグローバル大規模モデル API サービスを提供しています。1 つの API キーで全世界のモデルにアクセス可能。即時請求書発行、同時接続数無制限、最低 $0.15 から、24 時間年中無休のテクニカルサポート。GPT/Claude/Gemini が全て 32% オフ、国内モデルは 20〜50% オフ、Claude Code 専用モデルは 66% オフ実施中!<a href="https://www.dmxapi.cn/register?aff=bUHu">登録はこちら</a></td> <td>本プロジェクトは AIGoCode のスポンサー提供でお届けしています。AIGoCode は、Claude Code・Codex・最新の Gemini モデルを統合したオールインワンのAIコーディングプラットフォームで、安定性・高速性・コストパフォーマンスに優れた開発サービスを提供します。柔軟なサブスクリプションプランを備え、レスポンスも非常に高速です。さらに、CC Switch ユーザー向けの特典として、<a href="https://aigocode.com/invite/CC-SWITCH">このリンク</a>から登録すると、初回チャージ時に10%分のボーナスクレジットが付与されます!
</td>
</tr> </tr>
</table> </table>
@@ -47,7 +49,7 @@
## 特長 ## 特長
### 現在のバージョン:v3.8.3 | [完全な更新履歴](CHANGELOG.md) | [リリースノート](docs/release-note-v3.8.0-en.md) ### 現在のバージョン:v3.8.2 | [完全な更新履歴](CHANGELOG.md) | [リリースノート](docs/release-note-v3.8.0-en.md)
**v3.8.0 メジャーアップデート (2025-11-28)** **v3.8.0 メジャーアップデート (2025-11-28)**
+10 -10
View File
@@ -2,7 +2,7 @@
# Claude Code / Codex / Gemini CLI 全方位辅助工具 # Claude Code / Codex / Gemini CLI 全方位辅助工具
[![Version](https://img.shields.io/badge/version-3.8.3-blue.svg)](https://github.com/farion1231/cc-switch/releases) [![Version](https://img.shields.io/badge/version-3.8.2-blue.svg)](https://github.com/farion1231/cc-switch/releases)
[![Platform](https://img.shields.io/badge/platform-Windows%20%7C%20macOS%20%7C%20Linux-lightgrey.svg)](https://github.com/farion1231/cc-switch/releases) [![Platform](https://img.shields.io/badge/platform-Windows%20%7C%20macOS%20%7C%20Linux-lightgrey.svg)](https://github.com/farion1231/cc-switch/releases)
[![Built with Tauri](https://img.shields.io/badge/built%20with-Tauri%202-orange.svg)](https://tauri.app/) [![Built with Tauri](https://img.shields.io/badge/built%20with-Tauri%202-orange.svg)](https://tauri.app/)
[![Downloads](https://img.shields.io/endpoint?url=https://api.pinstudios.net/api/badges/downloads/farion1231/cc-switch/total)](https://github.com/farion1231/cc-switch/releases/latest) [![Downloads](https://img.shields.io/endpoint?url=https://api.pinstudios.net/api/badges/downloads/farion1231/cc-switch/total)](https://github.com/farion1231/cc-switch/releases/latest)
@@ -15,7 +15,7 @@
## ❤️赞助商 ## ❤️赞助商
[![智谱 GLM](assets/partners/banners/glm-zh.jpg)](https://www.bigmodel.cn/claude-code?ic=RRVJPB5SII) ![智谱 GLM](assets/partners/banners/glm-zh.jpg)
感谢智谱AI的 GLM CODING PLAN 赞助了本项目!GLM CODING PLAN 是专为AI编码打造的订阅套餐,每月最低仅需20元,即可在十余款主流AI编码工具如 Claude Code、Cline 中畅享智谱旗舰模型 GLM-4.6,为开发者提供顶尖、高速、稳定的编码体验。CC Switch 已经预设了智谱GLM,只需要填写 key 即可一键导入编程工具。智谱AI为本软件的用户提供了特别优惠,使用[此链接](https://www.bigmodel.cn/claude-code?ic=RRVJPB5SII)购买可以享受九折优惠。 感谢智谱AI的 GLM CODING PLAN 赞助了本项目!GLM CODING PLAN 是专为AI编码打造的订阅套餐,每月最低仅需20元,即可在十余款主流AI编码工具如 Claude Code、Cline 中畅享智谱旗舰模型 GLM-4.6,为开发者提供顶尖、高速、稳定的编码体验。CC Switch 已经预设了智谱GLM,只需要填写 key 即可一键导入编程工具。智谱AI为本软件的用户提供了特别优惠,使用[此链接](https://www.bigmodel.cn/claude-code?ic=RRVJPB5SII)购买可以享受九折优惠。
@@ -23,20 +23,20 @@
<table> <table>
<tr> <tr>
<td width="180"><a href="https://www.packyapi.com/register?aff=cc-switch"><img src="assets/partners/logos/packycode.png" alt="PackyCode" width="150"></a></td> <td width="180"><img src="assets/partners/logos/packycode.png" alt="PackyCode" width="150"></td>
<td>感谢 PackyCode 赞助了本项目!PackyCode 是一家稳定、高效的API中转服务商,提供 Claude Code、Codex、Gemini 等多种中转服务。PackyCode 为本软件的用户提供了特别优惠,使用<a href="https://www.packyapi.com/register?aff=cc-switch">此链接</a>注册并在充值时填写"cc-switch"优惠码,可以享受9折优惠。</td> <td>感谢 PackyCode 赞助了本项目!PackyCode 是一家稳定、高效的API中转服务商,提供 Claude Code、Codex、Gemini 等多种中转服务。PackyCode 为本软件的用户提供了特别优惠,使用<a href="https://www.packyapi.com/register?aff=cc-switch">此链接</a>注册并在充值时填写"cc-switch"优惠码,可以享受9折优惠。</td>
</tr> </tr>
<tr> <tr>
<td width="180"><a href="https://aigocode.com/invite/CC-SWITCH"><img src="assets/partners/logos/aigocode.png" alt="AIGoCode" width="150"></a></td> <td width="180"><img src="assets/partners/logos/sds-zh.png" alt="ShanDianShuo" width="150"></td>
<td>感谢闪电说赞助了本项目!闪电说是本地优先的 AI 语音输入法:毫秒级响应,数据不离设备;打字速度提升 4 倍,AI 智能纠错;绝对隐私安全,完全免费,配合 Claude Code 写代码效率翻倍!支持 Mac/Win 双平台,<a href="https://www.shandianshuo.cn">免费下载</a></td>
</tr>
<tr>
<td width="180"><img src="assets/partners/logos/aigocode.png" alt="AIGoCode" width="150"></td>
<td>感谢 AIGoCode 赞助了本项目!AIGoCode 是一个集成了 Claude Code、Codex 以及 Gemini 最新模型的一站式平台,为你提供稳定、高效且高性价比的AI编程服务。本站提供灵活的订阅计划,零封号风险,国内直连,无需魔法,极速响应。AIGoCode 为 CC Switch 的用户提供了特别福利,通过<a href="https://aigocode.com/invite/CC-SWITCH">此链接</a>注册的用户首次充值可以获得额外10%奖励额度!</td> <td>感谢 AIGoCode 赞助了本项目!AIGoCode 是一个集成了 Claude Code、Codex 以及 Gemini 最新模型的一站式平台,为你提供稳定、高效且高性价比的AI编程服务。本站提供灵活的订阅计划,零封号风险,国内直连,无需魔法,极速响应。AIGoCode 为 CC Switch 的用户提供了特别福利,通过<a href="https://aigocode.com/invite/CC-SWITCH">此链接</a>注册的用户首次充值可以获得额外10%奖励额度!</td>
</tr> </tr>
<tr>
<td width="180"><a href="https://www.dmxapi.cn/register?aff=bUHu"><img src="assets/partners/logos/dmx-zh.jpeg" alt="DMXAPI" width="150"></a></td>
<td>感谢 DMXAPI(大模型API)赞助了本项目! DMXAPI,一个Key用全球大模型。
为200多家企业用户提供全球大模型API服务。· 充值即开票 ·当天开票 ·并发不限制 ·1元起充 · 7x24 在线技术辅导,GPT/Claude/Gemini全部6.8折,国内模型5~8折,Claude Code 专属模型3.4折进行中!<a href="https://www.dmxapi.cn/register?aff=bUHu">点击这里注册</a></td>
</tr>
</table> </table>
## 界面预览 ## 界面预览
@@ -47,7 +47,7 @@
## 功能特性 ## 功能特性
### 当前版本:v3.8.3 | [完整更新日志](CHANGELOG.md) ### 当前版本:v3.8.2 | [完整更新日志](CHANGELOG.md)
**v3.8.0 重大更新(2025-11-28** **v3.8.0 重大更新(2025-11-28**
Binary file not shown.

Before

Width:  |  Height:  |  Size: 264 KiB

After

Width:  |  Height:  |  Size: 102 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 299 KiB

After

Width:  |  Height:  |  Size: 110 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 41 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 22 KiB

+1 -1
View File
@@ -4,7 +4,7 @@
"rsc": false, "rsc": false,
"tsx": true, "tsx": true,
"tailwind": { "tailwind": {
"config": "tailwind.config.cjs", "config": "tailwind.config.js",
"css": "src/index.css", "css": "src/index.css",
"baseColor": "neutral", "baseColor": "neutral",
"cssVariables": true, "cssVariables": true,
-165
View File
@@ -1,165 +0,0 @@
# CC Switch 代理功能使用指南
## 功能介绍
CC Switch 的代理功能是一个本地 HTTP 代理服务器,可以统一管理 Claude Code、Codex 和 Gemini CLI 的 API 请求。主要特性包括:
- **统一代理入口** - 所有 CLI 应用的请求通过本地代理转发
- **自动故障转移** - 当前供应商故障时自动切换到备用供应商
- **按应用控制** - 可独立控制每个应用是否启用代理
- **配置保护** - 自动备份原始配置,停止代理时安全恢复
## 快速开始
### 1. 启动代理
在 CC Switch 主界面,点击右上角的 **Proxy** 按钮,可以看到代理控制面板。
点击 **启动代理** 按钮启动本地代理服务器。代理默认监听 `127.0.0.1:15721`
### 2. 启用应用接管
代理启动后,你可以选择让哪些应用的请求通过代理:
- **Claude** - 接管 Claude Code 的 API 请求
- **Codex** - 接管 Codex CLI 的 API 请求
- **Gemini** - 接管 Gemini CLI 的 API 请求
点击对应应用的开关即可启用/禁用接管。
> **注意**:启用接管后,CC Switch 会自动修改对应应用的配置文件,将 API 端点指向本地代理。原始配置会被安全备份。
### 3. 正常使用 CLI
启用接管后,你可以正常使用各个 CLI 工具。所有请求都会经过 CC Switch 代理转发到配置的供应商。
### 4. 停止代理
当你不再需要代理时,点击 **停止代理** 按钮。CC Switch 会:
1. 安全关闭代理服务器
2. 自动恢复所有应用的原始配置
3. 清除代理状态
## 自动故障转移
### 工作原理
代理功能内置了智能故障转移机制:
1. **健康监控** - 实时监控每个供应商的响应状态
2. **熔断器** - 连续失败 5 次后触发熔断,暂停使用该供应商
3. **自动切换** - 熔断后自动切换到列表中的下一个供应商
4. **自动恢复** - 30 秒后尝试恢复熔断的供应商
### 配置故障转移
要使用故障转移功能,你需要:
1. 在对应应用下添加多个供应商(至少 2 个)
2. 启动代理并启用接管
3. 当主供应商故障时,代理会自动切换到备用供应商
### 健康状态指示
在供应商卡片上可以看到健康状态指示:
- **绿色** - 供应商正常
- **红色** - 供应商故障/熔断中
- **灰色** - 未使用代理或未检测
## 按应用接管
v3.9.0 新增了按应用分粒度控制功能:
- 你可以只接管 Claude,而让 Codex 使用原始配置
- 每个应用的接管状态独立管理
- 启用/禁用不会影响其他应用
### 接管状态检测
CC Switch 通过检测配置备份来判断接管状态:
- 存在备份 = 已接管
- 无备份 = 未接管
这确保了即使 CC Switch 异常退出,重新启动后也能正确识别状态。
## 代理配置
在代理面板中,你可以配置以下参数:
| 参数 | 默认值 | 说明 |
|------|--------|------|
| 监听地址 | 127.0.0.1 | 代理服务器绑定地址 |
| 监听端口 | 15721 | 代理服务器端口 |
| 最大重试 | 3 | 请求失败时的最大重试次数 |
| 请求超时 | 120 秒 | 单个请求的超时时间 |
| 启用日志 | 是 | 是否记录请求日志 |
## 常见问题
### Q: 代理启动失败,提示端口被占用?
A: 默认端口 15721 可能被其他程序占用。你可以:
- 关闭占用该端口的程序
- 在代理配置中修改端口号
### Q: 启用接管后 CLI 无法使用?
A: 请检查:
1. 代理服务器是否正常运行(查看代理面板状态)
2. 供应商配置是否正确(API Key 等)
3. 网络连接是否正常
### Q: 如何恢复原始配置?
A: 点击 **停止代理** 按钮,CC Switch 会自动恢复所有应用的原始配置。
如果 CC Switch 异常退出,重新启动后会检测到之前的备份,你可以:
- 点击停止代理来恢复配置
- 或继续使用代理功能
### Q: 故障转移没有生效?
A: 请确保:
1. 配置了至少 2 个供应商
2. 代理已启动且接管已启用
3. 故障转移只在代理模式下工作
### Q: 代理会影响性能吗?
A: 本地代理的延迟开销非常小(通常 < 1ms)。但如果启用了请求日志,在高频请求场景下可能会有少量性能影响。
## 技术细节
### 配置文件位置
启用接管后,CC Switch 会修改以下配置文件:
| 应用 | 配置文件 | 修改内容 |
|------|----------|----------|
| Claude | `~/.claude/settings.json` | `apiBaseUrl` 指向代理 |
| Codex | `~/.codex/config.toml` | `[api] baseUrl` 指向代理 |
| Gemini | `~/.gemini/.env` | `GEMINI_BASE_URL` 指向代理 |
原始配置备份在 CC Switch 数据库中,停止代理时自动恢复。
### 代理模式
代理服务器运行在接管模式下,会:
1. 接收来自 CLI 的 HTTPS 请求
2. 根据当前供应商配置转发到真实 API 端点
3. 返回响应给 CLI
4. 记录请求日志和健康状态
### 数据库表
代理功能使用以下数据库表:
- `proxy_config` - 代理配置
- `provider_health` - 供应商健康状态
- `proxy_request_logs` - 请求日志
- `circuit_breaker_config` - 熔断器配置
- `proxy_live_backup` - Live 配置备份
+3 -5
View File
@@ -1,8 +1,7 @@
{ {
"name": "cc-switch", "name": "cc-switch",
"version": "3.9.0-3", "version": "3.8.2",
"description": "All-in-One Assistant for Claude Code, Codex & Gemini CLI", "description": "All-in-One Assistant for Claude Code, Codex & Gemini CLI",
"type": "module",
"scripts": { "scripts": {
"dev": "pnpm tauri dev", "dev": "pnpm tauri dev",
"build": "pnpm tauri build", "build": "pnpm tauri build",
@@ -28,7 +27,6 @@
"@types/react-dom": "^18.2.0", "@types/react-dom": "^18.2.0",
"@vitejs/plugin-react": "^4.2.0", "@vitejs/plugin-react": "^4.2.0",
"autoprefixer": "^10.4.20", "autoprefixer": "^10.4.20",
"code-inspector-plugin": "^1.3.3",
"cross-fetch": "^4.1.0", "cross-fetch": "^4.1.0",
"jsdom": "^25.0.0", "jsdom": "^25.0.0",
"msw": "^2.11.6", "msw": "^2.11.6",
@@ -36,7 +34,7 @@
"prettier": "^3.6.2", "prettier": "^3.6.2",
"tailwindcss": "^3.4.17", "tailwindcss": "^3.4.17",
"typescript": "^5.3.0", "typescript": "^5.3.0",
"vite": "^7.3.0", "vite": "^5.0.0",
"vitest": "^2.0.5" "vitest": "^2.0.5"
}, },
"dependencies": { "dependencies": {
@@ -87,4 +85,4 @@
"zod": "^4.1.12" "zod": "^4.1.12"
}, },
"packageManager": "pnpm@10.10.0+sha512.d615db246fe70f25dcfea6d8d73dee782ce23e2245e3c4f6f888249fb568149318637dca73c2c5c8ef2a4ca0d5657fb9567188bfab47f566d1ee6ce987815c39" "packageManager": "pnpm@10.10.0+sha512.d615db246fe70f25dcfea6d8d73dee782ce23e2245e3c4f6f888249fb568149318637dca73c2c5c8ef2a4ca0d5657fb9567188bfab47f566d1ee6ce987815c39"
} }
+25 -517
View File
File diff suppressed because it is too large Load Diff
+5 -57
View File
@@ -586,12 +586,6 @@ version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b"
[[package]]
name = "byteorder-lite"
version = "0.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8f1fe948ff07f4bd06c30984e69f5b4899c516a3ef74f34df92a2df2ab535495"
[[package]] [[package]]
name = "bytes" name = "bytes"
version = "1.10.1" version = "1.10.1"
@@ -701,7 +695,7 @@ dependencies = [
[[package]] [[package]]
name = "cc-switch" name = "cc-switch"
version = "3.9.0-3" version = "3.8.2"
dependencies = [ dependencies = [
"anyhow", "anyhow",
"async-stream", "async-stream",
@@ -2220,7 +2214,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cc50b891e4acf8fe0e71ef88ec43ad82ee07b3810ad09de10f1d01f072ed4b98" checksum = "cc50b891e4acf8fe0e71ef88ec43ad82ee07b3810ad09de10f1d01f072ed4b98"
dependencies = [ dependencies = [
"byteorder", "byteorder",
"png 0.17.16", "png",
] ]
[[package]] [[package]]
@@ -2336,19 +2330,6 @@ dependencies = [
"icu_properties", "icu_properties",
] ]
[[package]]
name = "image"
version = "0.25.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "529feb3e6769d234375c4cf1ee2ce713682b8e76538cb13f9fc23e1400a591e7"
dependencies = [
"bytemuck",
"byteorder-lite",
"moxcms",
"num-traits",
"png 0.18.0",
]
[[package]] [[package]]
name = "indexmap" name = "indexmap"
version = "1.9.3" version = "1.9.3"
@@ -2778,16 +2759,6 @@ dependencies = [
"windows-sys 0.59.0", "windows-sys 0.59.0",
] ]
[[package]]
name = "moxcms"
version = "0.7.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0fbdd3d7436f8b5e892b8b7ea114271ff0fa00bc5acae845d53b07d498616ef6"
dependencies = [
"num-traits",
"pxfm",
]
[[package]] [[package]]
name = "muda" name = "muda"
version = "0.17.1" version = "0.17.1"
@@ -2803,7 +2774,7 @@ dependencies = [
"objc2-core-foundation", "objc2-core-foundation",
"objc2-foundation 0.3.1", "objc2-foundation 0.3.1",
"once_cell", "once_cell",
"png 0.17.16", "png",
"serde", "serde",
"thiserror 2.0.17", "thiserror 2.0.17",
"windows-sys 0.60.2", "windows-sys 0.60.2",
@@ -3592,19 +3563,6 @@ dependencies = [
"miniz_oxide", "miniz_oxide",
] ]
[[package]]
name = "png"
version = "0.18.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "97baced388464909d42d89643fe4361939af9b7ce7a31ee32a168f832a70f2a0"
dependencies = [
"bitflags 2.9.4",
"crc32fast",
"fdeflate",
"flate2",
"miniz_oxide",
]
[[package]] [[package]]
name = "polling" name = "polling"
version = "3.11.0" version = "3.11.0"
@@ -3736,15 +3694,6 @@ dependencies = [
"syn 1.0.109", "syn 1.0.109",
] ]
[[package]]
name = "pxfm"
version = "0.1.25"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a3cbdf373972bf78df4d3b518d07003938e2c7d1fb5891e55f9cb6df57009d84"
dependencies = [
"num-traits",
]
[[package]] [[package]]
name = "quick-xml" name = "quick-xml"
version = "0.37.5" version = "0.37.5"
@@ -5038,7 +4987,6 @@ dependencies = [
"heck 0.5.0", "heck 0.5.0",
"http", "http",
"http-range", "http-range",
"image",
"jni", "jni",
"libc", "libc",
"log", "log",
@@ -5107,7 +5055,7 @@ dependencies = [
"ico", "ico",
"json-patch", "json-patch",
"plist", "plist",
"png 0.17.16", "png",
"proc-macro2", "proc-macro2",
"quote", "quote",
"semver", "semver",
@@ -5862,7 +5810,7 @@ dependencies = [
"objc2-core-graphics", "objc2-core-graphics",
"objc2-foundation 0.3.1", "objc2-foundation 0.3.1",
"once_cell", "once_cell",
"png 0.17.16", "png",
"serde", "serde",
"thiserror 2.0.17", "thiserror 2.0.17",
"windows-sys 0.59.0", "windows-sys 0.59.0",
+2 -2
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "cc-switch" name = "cc-switch"
version = "3.9.0-3" version = "3.8.2"
description = "All-in-One Assistant for Claude Code, Codex & Gemini CLI" description = "All-in-One Assistant for Claude Code, Codex & Gemini CLI"
authors = ["Jason Young"] authors = ["Jason Young"]
license = "MIT" license = "MIT"
@@ -26,7 +26,7 @@ serde_json = "1.0"
serde = { version = "1.0", features = ["derive"] } serde = { version = "1.0", features = ["derive"] }
log = "0.4" log = "0.4"
chrono = { version = "0.4", features = ["serde"] } chrono = { version = "0.4", features = ["serde"] }
tauri = { version = "2.8.2", features = ["tray-icon", "protocol-asset", "image-png"] } tauri = { version = "2.8.2", features = ["tray-icon", "protocol-asset"] }
tauri-plugin-log = "2" tauri-plugin-log = "2"
tauri-plugin-opener = "2" tauri-plugin-opener = "2"
tauri-plugin-process = "2" tauri-plugin-process = "2"
Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.7 KiB

+4 -71
View File
@@ -1,36 +1,16 @@
use crate::error::AppError; use crate::error::AppError;
use auto_launch::{AutoLaunch, AutoLaunchBuilder}; use auto_launch::{AutoLaunch, AutoLaunchBuilder};
/// 获取 macOS 上的 .app bundle 路径
/// 将 `/path/to/CC Switch.app/Contents/MacOS/CC Switch` 转换为 `/path/to/CC Switch.app`
#[cfg(target_os = "macos")]
fn get_macos_app_bundle_path(exe_path: &std::path::Path) -> Option<std::path::PathBuf> {
let path_str = exe_path.to_string_lossy();
// 查找 .app/Contents/MacOS/ 模式
if let Some(app_pos) = path_str.find(".app/Contents/MacOS/") {
let app_bundle_end = app_pos + 4; // ".app" 的结束位置
Some(std::path::PathBuf::from(&path_str[..app_bundle_end]))
} else {
None
}
}
/// 初始化 AutoLaunch 实例 /// 初始化 AutoLaunch 实例
fn get_auto_launch() -> Result<AutoLaunch, AppError> { fn get_auto_launch() -> Result<AutoLaunch, AppError> {
let app_name = "CC Switch"; let app_name = "CC Switch";
let exe_path = let app_path =
std::env::current_exe().map_err(|e| AppError::Message(format!("无法获取应用路径: {e}")))?; std::env::current_exe().map_err(|e| AppError::Message(format!("无法获取应用路径: {e}")))?;
// macOS 需要使用 .app bundle 路径,否则 AppleScript login item 会打开终端
#[cfg(target_os = "macos")]
let app_path = get_macos_app_bundle_path(&exe_path).unwrap_or(exe_path);
#[cfg(not(target_os = "macos"))]
let app_path = exe_path;
// 使用 AutoLaunchBuilder 消除平台差异 // 使用 AutoLaunchBuilder 消除平台差异
// macOS: 使用 AppleScript 方式(默认),需要 .app bundle 路径 // Windows/Linux: new() 接受 3 参数
// Windows/Linux: 使用注册表/XDG autostart // macOS: new() 接受 4 参数(含 hidden 参数)
// Builder 模式自动处理这些差异
let auto_launch = AutoLaunchBuilder::new() let auto_launch = AutoLaunchBuilder::new()
.set_app_name(app_name) .set_app_name(app_name)
.set_app_path(&app_path.to_string_lossy()) .set_app_path(&app_path.to_string_lossy())
@@ -67,50 +47,3 @@ pub fn is_auto_launch_enabled() -> Result<bool, AppError> {
.is_enabled() .is_enabled()
.map_err(|e| AppError::Message(format!("检查开机自启状态失败: {e}"))) .map_err(|e| AppError::Message(format!("检查开机自启状态失败: {e}")))
} }
#[cfg(test)]
mod tests {
use super::*;
#[cfg(target_os = "macos")]
#[test]
fn test_get_macos_app_bundle_path_valid() {
let exe_path = std::path::Path::new("/Applications/CC Switch.app/Contents/MacOS/CC Switch");
let result = get_macos_app_bundle_path(exe_path);
assert_eq!(
result,
Some(std::path::PathBuf::from("/Applications/CC Switch.app"))
);
}
#[cfg(target_os = "macos")]
#[test]
fn test_get_macos_app_bundle_path_with_spaces() {
let exe_path =
std::path::Path::new("/Users/test/My Apps/CC Switch.app/Contents/MacOS/CC Switch");
let result = get_macos_app_bundle_path(exe_path);
assert_eq!(
result,
Some(std::path::PathBuf::from(
"/Users/test/My Apps/CC Switch.app"
))
);
}
#[cfg(target_os = "macos")]
#[test]
fn test_get_macos_app_bundle_path_not_in_bundle() {
let exe_path = std::path::Path::new("/usr/local/bin/cc-switch");
let result = get_macos_app_bundle_path(exe_path);
assert_eq!(result, None);
}
#[cfg(target_os = "macos")]
#[test]
fn test_get_macos_app_bundle_path_dev_build() {
// 开发环境下的路径通常不在 .app bundle 内
let exe_path = std::path::Path::new("/Users/dev/project/target/debug/cc-switch");
let result = get_macos_app_bundle_path(exe_path);
assert_eq!(result, None);
}
}
-243
View File
@@ -7,64 +7,6 @@ use std::path::{Path, PathBuf};
use crate::config::{atomic_write, get_claude_mcp_path, get_default_claude_mcp_path}; use crate::config::{atomic_write, get_claude_mcp_path, get_default_claude_mcp_path};
use crate::error::AppError; use crate::error::AppError;
/// 需要在 Windows 上用 cmd /c 包装的命令
/// 这些命令在 Windows 上实际是 .cmd 批处理文件,需要通过 cmd /c 来执行
#[cfg(windows)]
const WINDOWS_WRAP_COMMANDS: &[&str] = &["npx", "npm", "yarn", "pnpm", "node", "bun", "deno"];
/// Windows 平台:将 `npx args...` 转换为 `cmd /c npx args...`
/// 解决 Claude Code /doctor 报告的 "Windows requires 'cmd /c' wrapper to execute npx" 警告
#[cfg(windows)]
fn wrap_command_for_windows(obj: &mut Map<String, Value>) {
// 只处理 stdio 类型(默认或显式)
let server_type = obj.get("type").and_then(|v| v.as_str()).unwrap_or("stdio");
if server_type != "stdio" {
return;
}
let Some(cmd) = obj.get("command").and_then(|v| v.as_str()) else {
return;
};
// 已经是 cmd 的不重复包装
if cmd.eq_ignore_ascii_case("cmd") || cmd.eq_ignore_ascii_case("cmd.exe") {
return;
}
// 提取命令名(去掉 .cmd 后缀和路径)
let cmd_name = Path::new(cmd)
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or(cmd);
let needs_wrap = WINDOWS_WRAP_COMMANDS
.iter()
.any(|&c| cmd_name.eq_ignore_ascii_case(c));
if !needs_wrap {
return;
}
// 构建新的 args: ["/c", "原命令", ...原args]
let original_args = obj
.get("args")
.and_then(|v| v.as_array())
.cloned()
.unwrap_or_default();
let mut new_args = vec![Value::String("/c".into()), Value::String(cmd.into())];
new_args.extend(original_args);
obj.insert("command".into(), Value::String("cmd".into()));
obj.insert("args".into(), Value::Array(new_args));
}
/// 非 Windows 平台无需处理
#[cfg(not(windows))]
fn wrap_command_for_windows(_obj: &mut Map<String, Value>) {
// 非 Windows 平台不做任何处理
}
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
pub struct McpStatus { pub struct McpStatus {
@@ -163,55 +105,6 @@ pub fn read_mcp_json() -> Result<Option<String>, AppError> {
Ok(Some(content)) Ok(Some(content))
} }
/// 在 ~/.claude.json 根对象写入 hasCompletedOnboarding=true(用于跳过 Claude Code 初次安装确认)
/// 仅增量写入该字段,其他字段保持不变
pub fn set_has_completed_onboarding() -> Result<bool, AppError> {
let path = user_config_path();
let mut root = if path.exists() {
read_json_value(&path)?
} else {
serde_json::json!({})
};
let obj = root
.as_object_mut()
.ok_or_else(|| AppError::Config("~/.claude.json 根必须是对象".into()))?;
let already = obj
.get("hasCompletedOnboarding")
.and_then(|v| v.as_bool())
.unwrap_or(false);
if already {
return Ok(false);
}
obj.insert("hasCompletedOnboarding".into(), Value::Bool(true));
write_json_value(&path, &root)?;
Ok(true)
}
/// 删除 ~/.claude.json 根对象的 hasCompletedOnboarding 字段(恢复 Claude Code 初次安装确认)
/// 仅增量删除该字段,其他字段保持不变
pub fn clear_has_completed_onboarding() -> Result<bool, AppError> {
let path = user_config_path();
if !path.exists() {
return Ok(false);
}
let mut root = read_json_value(&path)?;
let obj = root
.as_object_mut()
.ok_or_else(|| AppError::Config("~/.claude.json 根必须是对象".into()))?;
let existed = obj.remove("hasCompletedOnboarding").is_some();
if !existed {
return Ok(false);
}
write_json_value(&path, &root)?;
Ok(true)
}
pub fn upsert_mcp_server(id: &str, spec: Value) -> Result<bool, AppError> { pub fn upsert_mcp_server(id: &str, spec: Value) -> Result<bool, AppError> {
if id.trim().is_empty() { if id.trim().is_empty() {
return Err(AppError::InvalidInput("MCP 服务器 ID 不能为空".into())); return Err(AppError::InvalidInput("MCP 服务器 ID 不能为空".into()));
@@ -397,9 +290,6 @@ pub fn set_mcp_servers_map(
obj.remove("homepage"); obj.remove("homepage");
obj.remove("docs"); obj.remove("docs");
// Windows 平台自动包装 npx/npm 等命令为 cmd /c 格式
wrap_command_for_windows(&mut obj);
out.insert(id.clone(), Value::Object(obj)); out.insert(id.clone(), Value::Object(obj));
} }
@@ -413,136 +303,3 @@ pub fn set_mcp_servers_map(
write_json_value(&path, &root)?; write_json_value(&path, &root)?;
Ok(()) Ok(())
} }
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
/// 测试 Windows 命令包装功能
/// 由于使用条件编译,在非 Windows 平台上测试的是空函数
#[test]
fn test_wrap_command_for_windows_npx() {
let mut obj = json!({"command": "npx", "args": ["-y", "@upstash/context7-mcp"]})
.as_object()
.unwrap()
.clone();
wrap_command_for_windows(&mut obj);
#[cfg(windows)]
{
assert_eq!(obj["command"], "cmd");
assert_eq!(
obj["args"],
json!(["/c", "npx", "-y", "@upstash/context7-mcp"])
);
}
#[cfg(not(windows))]
{
// 非 Windows 平台不做任何处理
assert_eq!(obj["command"], "npx");
}
}
#[test]
fn test_wrap_command_for_windows_npm() {
let mut obj = json!({"command": "npm", "args": ["run", "start"]})
.as_object()
.unwrap()
.clone();
wrap_command_for_windows(&mut obj);
#[cfg(windows)]
{
assert_eq!(obj["command"], "cmd");
assert_eq!(obj["args"], json!(["/c", "npm", "run", "start"]));
}
}
#[test]
fn test_wrap_command_for_windows_already_cmd() {
// 已经是 cmd 的不应该重复包装
let mut obj = json!({"command": "cmd", "args": ["/c", "npx", "-y", "foo"]})
.as_object()
.unwrap()
.clone();
wrap_command_for_windows(&mut obj);
assert_eq!(obj["command"], "cmd");
// args 应该保持不变,不会变成 ["/c", "cmd", "/c", "npx", ...]
assert_eq!(obj["args"], json!(["/c", "npx", "-y", "foo"]));
}
#[test]
fn test_wrap_command_for_windows_http_type_skipped() {
// http 类型不应该被处理
let mut obj = json!({"type": "http", "url": "https://example.com/mcp"})
.as_object()
.unwrap()
.clone();
wrap_command_for_windows(&mut obj);
assert!(!obj.contains_key("command"));
assert_eq!(obj["url"], "https://example.com/mcp");
}
#[test]
fn test_wrap_command_for_windows_other_command_skipped() {
// 非目标命令(如 python)不应该被包装
let mut obj = json!({"command": "python", "args": ["server.py"]})
.as_object()
.unwrap()
.clone();
wrap_command_for_windows(&mut obj);
// python 不在 WINDOWS_WRAP_COMMANDS 列表中,不应该被包装
assert_eq!(obj["command"], "python");
assert_eq!(obj["args"], json!(["server.py"]));
}
#[test]
fn test_wrap_command_for_windows_no_args() {
// 没有 args 的情况
let mut obj = json!({"command": "npx"}).as_object().unwrap().clone();
wrap_command_for_windows(&mut obj);
#[cfg(windows)]
{
assert_eq!(obj["command"], "cmd");
assert_eq!(obj["args"], json!(["/c", "npx"]));
}
}
#[test]
fn test_wrap_command_for_windows_with_cmd_suffix() {
// 处理 npx.cmd 格式
let mut obj = json!({"command": "npx.cmd", "args": ["-y", "foo"]})
.as_object()
.unwrap()
.clone();
wrap_command_for_windows(&mut obj);
#[cfg(windows)]
{
assert_eq!(obj["command"], "cmd");
assert_eq!(obj["args"], json!(["/c", "npx.cmd", "-y", "foo"]));
}
}
#[test]
fn test_wrap_command_for_windows_case_insensitive() {
// 大小写不敏感
let mut obj = json!({"command": "NPX", "args": ["-y", "foo"]})
.as_object()
.unwrap()
.clone();
wrap_command_for_windows(&mut obj);
#[cfg(windows)]
{
assert_eq!(obj["command"], "cmd");
assert_eq!(obj["args"], json!(["/c", "NPX", "-y", "foo"]));
}
}
}
+10 -28
View File
@@ -1,6 +1,6 @@
//! 故障转移队列命令 //! 故障转移队列命令
//! //!
//! 管理代理模式下的故障转移队列(基于 providers 表的 in_failover_queue 字段) //! 管理代理模式下的故障转移队列
use crate::database::FailoverQueueItem; use crate::database::FailoverQueueItem;
use crate::provider::Provider; use crate::provider::Provider;
@@ -56,47 +56,29 @@ pub async fn remove_from_failover_queue(
.map_err(|e| e.to_string()) .map_err(|e| e.to_string())
} }
/// 获取指定应用的自动故障转移开关状态(从 proxy_config 表读取) /// 重新排序故障转移队列
#[tauri::command] #[tauri::command]
pub async fn get_auto_failover_enabled( pub async fn reorder_failover_queue(
state: tauri::State<'_, AppState>, state: tauri::State<'_, AppState>,
app_type: String, app_type: String,
) -> Result<bool, String> { provider_ids: Vec<String>,
) -> Result<(), String> {
state state
.db .db
.get_proxy_config_for_app(&app_type) .reorder_failover_queue(&app_type, &provider_ids)
.await
.map(|config| config.auto_failover_enabled)
.map_err(|e| e.to_string()) .map_err(|e| e.to_string())
} }
/// 设置指定应用的自动故障转移开关状态(写入 proxy_config 表) /// 设置故障转移队列项的启用状态
///
/// 注意:关闭故障转移时不会清除队列,队列内容会保留供下次开启时使用
#[tauri::command] #[tauri::command]
pub async fn set_auto_failover_enabled( pub async fn set_failover_item_enabled(
state: tauri::State<'_, AppState>, state: tauri::State<'_, AppState>,
app_type: String, app_type: String,
provider_id: String,
enabled: bool, enabled: bool,
) -> Result<(), String> { ) -> Result<(), String> {
log::info!(
"[Failover] Setting auto_failover_enabled: app_type='{app_type}', enabled={enabled}"
);
// 读取当前配置
let mut config = state
.db
.get_proxy_config_for_app(&app_type)
.await
.map_err(|e| e.to_string())?;
// 更新 auto_failover_enabled 字段
config.auto_failover_enabled = enabled;
// 写回数据库
state state
.db .db
.update_proxy_config_for_app(config) .set_failover_item_enabled(&app_type, &provider_id, enabled)
.await
.map_err(|e| e.to_string()) .map_err(|e| e.to_string())
} }
+11 -50
View File
@@ -4,12 +4,6 @@ use crate::init_status::InitErrorPayload;
use tauri::AppHandle; use tauri::AppHandle;
use tauri_plugin_opener::OpenerExt; use tauri_plugin_opener::OpenerExt;
#[cfg(target_os = "windows")]
use std::os::windows::process::CommandExt;
#[cfg(target_os = "windows")]
const CREATE_NO_WINDOW: u32 = 0x08000000;
/// 打开外部链接 /// 打开外部链接
#[tauri::command] #[tauri::command]
pub async fn open_external(app: AppHandle, url: String) -> Result<bool, String> { pub async fn open_external(app: AppHandle, url: String) -> Result<bool, String> {
@@ -148,16 +142,11 @@ fn extract_version(raw: &str) -> String {
fn try_get_version(tool: &str) -> (Option<String>, Option<String>) { fn try_get_version(tool: &str) -> (Option<String>, Option<String>) {
use std::process::Command; use std::process::Command;
#[cfg(target_os = "windows")] let output = if cfg!(target_os = "windows") {
let output = {
Command::new("cmd") Command::new("cmd")
.args(["/C", &format!("{tool} --version")]) .args(["/C", &format!("{tool} --version")])
.creation_flags(CREATE_NO_WINDOW)
.output() .output()
}; } else {
#[cfg(not(target_os = "windows"))]
let output = {
Command::new("sh") Command::new("sh")
.arg("-c") .arg("-c")
.arg(format!("{tool} --version")) .arg(format!("{tool} --version"))
@@ -166,17 +155,11 @@ fn try_get_version(tool: &str) -> (Option<String>, Option<String>) {
match output { match output {
Ok(out) => { Ok(out) => {
let stdout = String::from_utf8_lossy(&out.stdout).trim().to_string();
let stderr = String::from_utf8_lossy(&out.stderr).trim().to_string();
if out.status.success() { if out.status.success() {
let raw = if stdout.is_empty() { &stderr } else { &stdout }; let raw = String::from_utf8_lossy(&out.stdout).trim().to_string();
if raw.is_empty() { (Some(extract_version(&raw)), None)
(None, Some("未安装或无法执行".to_string()))
} else {
(Some(extract_version(raw)), None)
}
} else { } else {
let err = if stderr.is_empty() { stdout } else { stderr }; let err = String::from_utf8_lossy(&out.stderr).trim().to_string();
( (
None, None,
Some(if err.is_empty() { Some(if err.is_empty() {
@@ -248,39 +231,17 @@ fn scan_cli_version(tool: &str) -> (Option<String>, Option<String>) {
if tool_path.exists() { if tool_path.exists() {
// 构建 PATH 环境变量,确保 node 可被找到 // 构建 PATH 环境变量,确保 node 可被找到
let current_path = std::env::var("PATH").unwrap_or_default(); let current_path = std::env::var("PATH").unwrap_or_default();
#[cfg(target_os = "windows")]
let new_path = format!("{};{}", path.display(), current_path);
#[cfg(not(target_os = "windows"))]
let new_path = format!("{}:{}", path.display(), current_path); let new_path = format!("{}:{}", path.display(), current_path);
#[cfg(target_os = "windows")] let output = Command::new(&tool_path)
let output = { .arg("--version")
// 使用 cmd /C 包装执行,确保子进程也在隐藏的控制台中运行 .env("PATH", &new_path)
Command::new("cmd") .output();
.args(["/C", &format!("\"{}\" --version", tool_path.display())])
.env("PATH", &new_path)
.creation_flags(CREATE_NO_WINDOW)
.output()
};
#[cfg(not(target_os = "windows"))]
let output = {
Command::new(&tool_path)
.arg("--version")
.env("PATH", &new_path)
.output()
};
if let Ok(out) = output { if let Ok(out) = output {
let stdout = String::from_utf8_lossy(&out.stdout).trim().to_string();
let stderr = String::from_utf8_lossy(&out.stderr).trim().to_string();
if out.status.success() { if out.status.success() {
let raw = if stdout.is_empty() { &stderr } else { &stdout }; let raw = String::from_utf8_lossy(&out.stdout).trim().to_string();
if !raw.is_empty() { return (Some(extract_version(&raw)), None);
return (Some(extract_version(raw)), None);
}
} }
} }
} }
-12
View File
@@ -34,15 +34,3 @@ pub async fn apply_claude_plugin_config(official: bool) -> Result<bool, String>
pub async fn is_claude_plugin_applied() -> Result<bool, String> { pub async fn is_claude_plugin_applied() -> Result<bool, String> {
crate::claude_plugin::is_claude_config_applied().map_err(|e| e.to_string()) crate::claude_plugin::is_claude_config_applied().map_err(|e| e.to_string())
} }
/// Claude Code:跳过初次安装确认(写入 ~/.claude.json 的 hasCompletedOnboarding=true
#[tauri::command]
pub async fn apply_claude_onboarding_skip() -> Result<bool, String> {
crate::claude_mcp::set_has_completed_onboarding().map_err(|e| e.to_string())
}
/// Claude Code:恢复初次安装确认(删除 ~/.claude.json 的 hasCompletedOnboarding 字段)
#[tauri::command]
pub async fn clear_claude_onboarding_skip() -> Result<bool, String> {
crate::claude_mcp::clear_has_completed_onboarding().map_err(|e| e.to_string())
}
-94
View File
@@ -229,97 +229,3 @@ pub fn update_providers_sort_order(
let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?; let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?;
ProviderService::update_sort_order(state.inner(), app_type, updates).map_err(|e| e.to_string()) ProviderService::update_sort_order(state.inner(), app_type, updates).map_err(|e| e.to_string())
} }
// ============================================================================
// 统一供应商(Universal Provider)命令
// ============================================================================
use crate::provider::UniversalProvider;
use std::collections::HashMap;
use tauri::{AppHandle, Emitter};
/// 统一供应商同步完成事件的 payload
#[derive(Clone, serde::Serialize)]
pub struct UniversalProviderSyncedEvent {
/// 操作类型: "upsert" | "delete" | "sync"
pub action: String,
/// 统一供应商 ID
pub id: String,
}
/// 发送统一供应商同步事件,通知前端刷新供应商列表
fn emit_universal_provider_synced(app: &AppHandle, action: &str, id: &str) {
let _ = app.emit(
"universal-provider-synced",
UniversalProviderSyncedEvent {
action: action.to_string(),
id: id.to_string(),
},
);
}
/// 获取所有统一供应商
#[tauri::command]
pub fn get_universal_providers(
state: State<'_, AppState>,
) -> Result<HashMap<String, UniversalProvider>, String> {
ProviderService::list_universal(state.inner()).map_err(|e| e.to_string())
}
/// 获取单个统一供应商
#[tauri::command]
pub fn get_universal_provider(
state: State<'_, AppState>,
id: String,
) -> Result<Option<UniversalProvider>, String> {
ProviderService::get_universal(state.inner(), &id).map_err(|e| e.to_string())
}
/// 添加或更新统一供应商
#[tauri::command]
pub fn upsert_universal_provider(
app: AppHandle,
state: State<'_, AppState>,
provider: UniversalProvider,
) -> Result<bool, String> {
let id = provider.id.clone();
let result =
ProviderService::upsert_universal(state.inner(), provider).map_err(|e| e.to_string())?;
// 发送事件通知前端刷新
emit_universal_provider_synced(&app, "upsert", &id);
Ok(result)
}
/// 删除统一供应商
#[tauri::command]
pub fn delete_universal_provider(
app: AppHandle,
state: State<'_, AppState>,
id: String,
) -> Result<bool, String> {
let result =
ProviderService::delete_universal(state.inner(), &id).map_err(|e| e.to_string())?;
// 发送事件通知前端刷新
emit_universal_provider_synced(&app, "delete", &id);
Ok(result)
}
/// 同步统一供应商到各应用(手动触发)
#[tauri::command]
pub fn sync_universal_provider(
app: AppHandle,
state: State<'_, AppState>,
id: String,
) -> Result<bool, String> {
let result =
ProviderService::sync_universal_to_apps(state.inner(), &id).map_err(|e| e.to_string())?;
// 发送事件通知前端刷新
emit_universal_provider_synced(&app, "sync", &id);
Ok(result)
}
+3 -147
View File
@@ -6,12 +6,12 @@ use crate::proxy::types::*;
use crate::proxy::{CircuitBreakerConfig, CircuitBreakerStats}; use crate::proxy::{CircuitBreakerConfig, CircuitBreakerStats};
use crate::store::AppState; use crate::store::AppState;
/// 启动代理服务器(仅启动服务,不接管 Live 配置) /// 启动代理服务器( Live 配置接管
#[tauri::command] #[tauri::command]
pub async fn start_proxy_server( pub async fn start_proxy_with_takeover(
state: tauri::State<'_, AppState>, state: tauri::State<'_, AppState>,
) -> Result<ProxyServerInfo, String> { ) -> Result<ProxyServerInfo, String> {
state.proxy_service.start().await state.proxy_service.start_with_takeover().await
} }
/// 停止代理服务器(恢复 Live 配置) /// 停止代理服务器(恢复 Live 配置)
@@ -20,27 +20,6 @@ pub async fn stop_proxy_with_restore(state: tauri::State<'_, AppState>) -> Resul
state.proxy_service.stop_with_restore().await state.proxy_service.stop_with_restore().await
} }
/// 获取各应用接管状态
#[tauri::command]
pub async fn get_proxy_takeover_status(
state: tauri::State<'_, AppState>,
) -> Result<ProxyTakeoverStatus, String> {
state.proxy_service.get_takeover_status().await
}
/// 为指定应用开启/关闭接管
#[tauri::command]
pub async fn set_proxy_takeover_for_app(
state: tauri::State<'_, AppState>,
app_type: String,
enabled: bool,
) -> Result<(), String> {
state
.proxy_service
.set_takeover_for_app(&app_type, enabled)
.await
}
/// 获取代理服务器状态 /// 获取代理服务器状态
#[tauri::command] #[tauri::command]
pub async fn get_proxy_status(state: tauri::State<'_, AppState>) -> Result<ProxyStatus, String> { pub async fn get_proxy_status(state: tauri::State<'_, AppState>) -> Result<ProxyStatus, String> {
@@ -62,63 +41,6 @@ pub async fn update_proxy_config(
state.proxy_service.update_config(&config).await state.proxy_service.update_config(&config).await
} }
// ==================== Global & Per-App Config ====================
/// 获取全局代理配置
///
/// 返回统一的全局配置字段(代理开关、监听地址、端口、日志开关)
#[tauri::command]
pub async fn get_global_proxy_config(
state: tauri::State<'_, AppState>,
) -> Result<GlobalProxyConfig, String> {
let db = &state.db;
db.get_global_proxy_config()
.await
.map_err(|e| e.to_string())
}
/// 更新全局代理配置
///
/// 更新统一的全局配置字段,会同时更新三行(claude/codex/gemini
#[tauri::command]
pub async fn update_global_proxy_config(
state: tauri::State<'_, AppState>,
config: GlobalProxyConfig,
) -> Result<(), String> {
let db = &state.db;
db.update_global_proxy_config(config)
.await
.map_err(|e| e.to_string())
}
/// 获取指定应用的代理配置
///
/// 返回应用级配置(enabled、auto_failover、超时、熔断器等)
#[tauri::command]
pub async fn get_proxy_config_for_app(
state: tauri::State<'_, AppState>,
app_type: String,
) -> Result<AppProxyConfig, String> {
let db = &state.db;
db.get_proxy_config_for_app(&app_type)
.await
.map_err(|e| e.to_string())
}
/// 更新指定应用的代理配置
///
/// 更新应用级配置(enabled、auto_failover、超时、熔断器等)
#[tauri::command]
pub async fn update_proxy_config_for_app(
state: tauri::State<'_, AppState>,
config: AppProxyConfig,
) -> Result<(), String> {
let db = &state.db;
db.update_proxy_config_for_app(config)
.await
.map_err(|e| e.to_string())
}
/// 检查代理服务器是否正在运行 /// 检查代理服务器是否正在运行
#[tauri::command] #[tauri::command]
pub async fn is_proxy_running(state: tauri::State<'_, AppState>) -> Result<bool, String> { pub async fn is_proxy_running(state: tauri::State<'_, AppState>) -> Result<bool, String> {
@@ -160,13 +82,8 @@ pub async fn get_provider_health(
} }
/// 重置熔断器 /// 重置熔断器
///
/// 重置后会检查是否应该切回队列中优先级更高的供应商:
/// 1. 检查自动故障转移是否开启
/// 2. 如果恢复的供应商在队列中优先级更高(queue_order 更小),则自动切换
#[tauri::command] #[tauri::command]
pub async fn reset_circuit_breaker( pub async fn reset_circuit_breaker(
app_handle: tauri::AppHandle,
state: tauri::State<'_, AppState>, state: tauri::State<'_, AppState>,
provider_id: String, provider_id: String,
app_type: String, app_type: String,
@@ -183,67 +100,6 @@ pub async fn reset_circuit_breaker(
.reset_provider_circuit_breaker(&provider_id, &app_type) .reset_provider_circuit_breaker(&provider_id, &app_type)
.await?; .await?;
// 3. 检查是否应该切回优先级更高的供应商(从 proxy_config 表读取)
// 只有当该应用已被代理接管(enabled=true)且开启了自动故障转移时才执行
let (app_enabled, auto_failover_enabled) = match db.get_proxy_config_for_app(&app_type).await {
Ok(config) => (config.enabled, config.auto_failover_enabled),
Err(e) => {
log::error!("[{app_type}] Failed to read proxy_config: {e}, defaulting to disabled");
(false, false)
}
};
if app_enabled && auto_failover_enabled && state.proxy_service.is_running().await {
// 获取当前供应商 ID
let current_id = db
.get_current_provider(&app_type)
.map_err(|e| e.to_string())?;
if let Some(current_id) = current_id {
// 获取故障转移队列
let queue = db
.get_failover_queue(&app_type)
.map_err(|e| e.to_string())?;
// 找到恢复的供应商和当前供应商在队列中的位置(使用 sort_index
let restored_order = queue
.iter()
.find(|item| item.provider_id == provider_id)
.and_then(|item| item.sort_index);
let current_order = queue
.iter()
.find(|item| item.provider_id == current_id)
.and_then(|item| item.sort_index);
// 如果恢复的供应商优先级更高(sort_index 更小),则切换
if let (Some(restored), Some(current)) = (restored_order, current_order) {
if restored < current {
log::info!(
"[Recovery] 供应商 {provider_id} 已恢复且优先级更高 (P{restored} vs P{current}),自动切换"
);
// 获取供应商名称用于日志和事件
let provider_name = db
.get_all_providers(&app_type)
.ok()
.and_then(|providers| providers.get(&provider_id).map(|p| p.name.clone()))
.unwrap_or_else(|| provider_id.clone());
// 创建故障转移切换管理器并执行切换
let switch_manager =
crate::proxy::failover_switch::FailoverSwitchManager::new(db.clone());
if let Err(e) = switch_manager
.try_switch(Some(&app_handle), &app_type, &provider_id, &provider_name)
.await
{
log::error!("[Recovery] 自动切换失败: {e}");
}
}
}
}
}
Ok(()) Ok(())
} }
+3 -1
View File
@@ -52,7 +52,9 @@ pub async fn stream_check_all_providers(
} }
if let Ok(queue) = state.db.get_failover_queue(app_type.as_str()) { if let Ok(queue) = state.db.get_failover_queue(app_type.as_str()) {
for item in queue { for item in queue {
ids.insert(item.provider_id); if item.enabled {
ids.insert(item.provider_id);
}
} }
} }
Some(ids) Some(ids)
+2 -3
View File
@@ -19,10 +19,9 @@ pub fn get_usage_summary(
#[tauri::command] #[tauri::command]
pub fn get_usage_trends( pub fn get_usage_trends(
state: State<'_, AppState>, state: State<'_, AppState>,
start_date: Option<i64>, days: u32,
end_date: Option<i64>,
) -> Result<Vec<DailyStats>, AppError> { ) -> Result<Vec<DailyStats>, AppError> {
state.db.get_daily_trends(start_date, end_date) state.db.get_daily_trends(days)
} }
/// 获取 Provider 统计 /// 获取 Provider 统计
+22 -23
View File
@@ -13,8 +13,6 @@ use std::fs;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use tempfile::NamedTempFile; use tempfile::NamedTempFile;
const CC_SWITCH_SQL_EXPORT_HEADER: &str = "-- CC Switch SQLite 导出";
impl Database { impl Database {
/// 导出为 SQLite 兼容的 SQL 文本 /// 导出为 SQLite 兼容的 SQL 文本
pub fn export_sql(&self, target_path: &Path) -> Result<(), AppError> { pub fn export_sql(&self, target_path: &Path) -> Result<(), AppError> {
@@ -38,8 +36,7 @@ impl Database {
} }
let sql_raw = fs::read_to_string(source_path).map_err(|e| AppError::io(source_path, e))?; let sql_raw = fs::read_to_string(source_path).map_err(|e| AppError::io(source_path, e))?;
let sql_content = sql_raw.trim_start_matches('\u{feff}'); let sql_content = Self::sanitize_import_sql(&sql_raw);
Self::validate_cc_switch_sql_export(sql_content)?;
// 导入前备份现有数据库 // 导入前备份现有数据库
let backup_path = self.backup_database_file()?; let backup_path = self.backup_database_file()?;
@@ -54,7 +51,7 @@ impl Database {
Connection::open(&temp_path).map_err(|e| AppError::Database(e.to_string()))?; Connection::open(&temp_path).map_err(|e| AppError::Database(e.to_string()))?;
temp_conn temp_conn
.execute_batch(sql_content) .execute_batch(&sql_content)
.map_err(|e| AppError::Database(format!("执行 SQL 导入失败: {e}")))?; .map_err(|e| AppError::Database(format!("执行 SQL 导入失败: {e}")))?;
// 补齐缺失表/索引并进行基础校验 // 补齐缺失表/索引并进行基础校验
@@ -96,17 +93,26 @@ impl Database {
Ok(snapshot) Ok(snapshot)
} }
fn validate_cc_switch_sql_export(sql: &str) -> Result<(), AppError> { /// 移除 SQLite 保留对象相关语句(如 sqlite_sequence),避免导入报错
let trimmed = sql.trim_start(); fn sanitize_import_sql(sql: &str) -> String {
if trimmed.starts_with(CC_SWITCH_SQL_EXPORT_HEADER) { let mut cleaned = String::new();
return Ok(()); let lower_keyword = "sqlite_sequence";
for stmt in sql.split(';') {
let trimmed = stmt.trim();
if trimmed.is_empty() {
continue;
}
if trimmed.to_ascii_lowercase().contains(lower_keyword) {
continue;
}
cleaned.push_str(trimmed);
cleaned.push_str(";\n");
} }
Err(AppError::localized( cleaned
"backup.sql.invalid_format",
"仅支持导入由 CC Switch 导出的 SQL 备份文件。",
"Only SQL backups exported by CC Switch are supported.",
))
} }
/// 生成一致性快照备份,返回备份文件路径(不存在主库时返回 None) /// 生成一致性快照备份,返回备份文件路径(不存在主库时返回 None)
@@ -123,15 +129,8 @@ impl Database {
fs::create_dir_all(&backup_dir).map_err(|e| AppError::io(&backup_dir, e))?; fs::create_dir_all(&backup_dir).map_err(|e| AppError::io(&backup_dir, e))?;
let base_id = format!("db_backup_{}", Utc::now().format("%Y%m%d_%H%M%S")); let backup_id = format!("db_backup_{}", Utc::now().format("%Y%m%d_%H%M%S"));
let mut backup_id = base_id.clone(); let backup_path = backup_dir.join(format!("{backup_id}.db"));
let mut backup_path = backup_dir.join(format!("{backup_id}.db"));
let mut counter = 1;
while backup_path.exists() {
backup_id = format!("{base_id}_{counter}");
backup_path = backup_dir.join(format!("{backup_id}.db"));
counter += 1;
}
{ {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
+130 -35
View File
@@ -1,32 +1,36 @@
//! 故障转移队列 DAO //! 故障转移队列 DAO
//! //!
//! 管理代理模式下的故障转移队列(基于 providers 表的 in_failover_queue 字段) //! 管理代理模式下的故障转移队列
use crate::database::{lock_conn, Database}; use crate::database::{lock_conn, Database};
use crate::error::AppError; use crate::error::AppError;
use crate::provider::Provider; use crate::provider::Provider;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::time::{SystemTime, UNIX_EPOCH};
/// 故障转移队列条目(简化版,用于前端展示) /// 故障转移队列条目
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
pub struct FailoverQueueItem { pub struct FailoverQueueItem {
pub provider_id: String, pub provider_id: String,
pub provider_name: String, pub provider_name: String,
pub sort_index: Option<usize>, pub queue_order: i32,
pub enabled: bool,
pub created_at: i64,
} }
impl Database { impl Database {
/// 获取故障转移队列(按 sort_index 排序) /// 获取故障转移队列(按 queue_order 排序)
pub fn get_failover_queue(&self, app_type: &str) -> Result<Vec<FailoverQueueItem>, AppError> { pub fn get_failover_queue(&self, app_type: &str) -> Result<Vec<FailoverQueueItem>, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
let mut stmt = conn let mut stmt = conn
.prepare( .prepare(
"SELECT id, name, sort_index "SELECT fq.provider_id, p.name, fq.queue_order, fq.enabled, fq.created_at
FROM providers FROM failover_queue fq
WHERE app_type = ?1 AND in_failover_queue = 1 JOIN providers p ON fq.provider_id = p.id AND fq.app_type = p.app_type
ORDER BY COALESCE(sort_index, 999999), id ASC", WHERE fq.app_type = ?1
ORDER BY fq.queue_order ASC",
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
@@ -35,7 +39,9 @@ impl Database {
Ok(FailoverQueueItem { Ok(FailoverQueueItem {
provider_id: row.get(0)?, provider_id: row.get(0)?,
provider_name: row.get(1)?, provider_name: row.get(1)?,
sort_index: row.get(2)?, queue_order: row.get(2)?,
enabled: row.get(3)?,
created_at: row.get(4)?,
}) })
}) })
.map_err(|e| AppError::Database(e.to_string()))? .map_err(|e| AppError::Database(e.to_string()))?
@@ -47,23 +53,43 @@ impl Database {
/// 获取故障转移队列中的供应商(完整 Provider 信息,按顺序) /// 获取故障转移队列中的供应商(完整 Provider 信息,按顺序)
pub fn get_failover_providers(&self, app_type: &str) -> Result<Vec<Provider>, AppError> { pub fn get_failover_providers(&self, app_type: &str) -> Result<Vec<Provider>, AppError> {
let queue = self.get_failover_queue(app_type)?;
let all_providers = self.get_all_providers(app_type)?; let all_providers = self.get_all_providers(app_type)?;
let result: Vec<Provider> = all_providers let mut result = Vec::new();
.into_values() for item in queue {
.filter(|p| p.in_failover_queue) if item.enabled {
.collect(); if let Some(provider) = all_providers.get(&item.provider_id) {
result.push(provider.clone());
}
}
}
Ok(result) Ok(result)
} }
/// 添加供应商到故障转移队列 /// 添加供应商到故障转移队列末尾
pub fn add_to_failover_queue(&self, app_type: &str, provider_id: &str) -> Result<(), AppError> { pub fn add_to_failover_queue(&self, app_type: &str, provider_id: &str) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
// 获取当前最大 queue_order
let max_order: i32 = conn
.query_row(
"SELECT COALESCE(MAX(queue_order), 0) FROM failover_queue WHERE app_type = ?1",
[app_type],
|row| row.get(0),
)
.map_err(|e| AppError::Database(e.to_string()))?;
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64;
conn.execute( conn.execute(
"UPDATE providers SET in_failover_queue = 1 WHERE id = ?1 AND app_type = ?2", "INSERT OR IGNORE INTO failover_queue (app_type, provider_id, queue_order, enabled, created_at)
rusqlite::params![provider_id, app_type], VALUES (?1, ?2, ?3, 1, ?4)",
rusqlite::params![app_type, provider_id, max_order + 1, now],
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
@@ -78,22 +104,90 @@ impl Database {
) -> Result<(), AppError> { ) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
// 1. 从队列中移除 // 获取被删除项的 queue_order
let removed_order: Option<i32> = conn
.query_row(
"SELECT queue_order FROM failover_queue WHERE app_type = ?1 AND provider_id = ?2",
[app_type, provider_id],
|row| row.get(0),
)
.ok();
// 删除该项
conn.execute( conn.execute(
"UPDATE providers SET in_failover_queue = 0 WHERE id = ?1 AND app_type = ?2", "DELETE FROM failover_queue WHERE app_type = ?1 AND provider_id = ?2",
rusqlite::params![provider_id, app_type], [app_type, provider_id],
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
// 2. 清除该供应商的健康状态(退出队列后不再需要健康监控 // 重新排序后面的项(填补空隙
if let Some(order) = removed_order {
conn.execute(
"UPDATE failover_queue
SET queue_order = queue_order - 1
WHERE app_type = ?1 AND queue_order > ?2",
rusqlite::params![app_type, order],
)
.map_err(|e| AppError::Database(e.to_string()))?;
}
Ok(())
}
/// 重新排序故障转移队列
/// provider_ids: 按新顺序排列的 provider_id 列表
pub fn reorder_failover_queue(
&self,
app_type: &str,
provider_ids: &[String],
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
// 使用事务确保原子性
conn.execute("BEGIN TRANSACTION", [])
.map_err(|e| AppError::Database(e.to_string()))?;
let result = (|| {
for (index, provider_id) in provider_ids.iter().enumerate() {
conn.execute(
"UPDATE failover_queue
SET queue_order = ?3
WHERE app_type = ?1 AND provider_id = ?2",
rusqlite::params![app_type, provider_id, (index + 1) as i32],
)
.map_err(|e| AppError::Database(e.to_string()))?;
}
Ok(())
})();
match result {
Ok(_) => {
conn.execute("COMMIT", [])
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
Err(e) => {
conn.execute("ROLLBACK", []).ok();
Err(e)
}
}
}
/// 设置故障转移队列中供应商的启用状态
pub fn set_failover_item_enabled(
&self,
app_type: &str,
provider_id: &str,
enabled: bool,
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
conn.execute( conn.execute(
"DELETE FROM provider_health WHERE provider_id = ?1 AND app_type = ?2", "UPDATE failover_queue SET enabled = ?3 WHERE app_type = ?1 AND provider_id = ?2",
rusqlite::params![provider_id, app_type], rusqlite::params![app_type, provider_id, enabled],
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
log::info!("已从故障转移队列移除供应商 {provider_id} ({app_type}), 并清除其健康状态");
Ok(()) Ok(())
} }
@@ -101,11 +195,8 @@ impl Database {
pub fn clear_failover_queue(&self, app_type: &str) -> Result<(), AppError> { pub fn clear_failover_queue(&self, app_type: &str) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
conn.execute( conn.execute("DELETE FROM failover_queue WHERE app_type = ?1", [app_type])
"UPDATE providers SET in_failover_queue = 0 WHERE app_type = ?1", .map_err(|e| AppError::Database(e.to_string()))?;
[app_type],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(()) Ok(())
} }
@@ -118,15 +209,15 @@ impl Database {
) -> Result<bool, AppError> { ) -> Result<bool, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
let in_queue: bool = conn let count: i32 = conn
.query_row( .query_row(
"SELECT in_failover_queue FROM providers WHERE id = ?1 AND app_type = ?2", "SELECT COUNT(*) FROM failover_queue WHERE app_type = ?1 AND provider_id = ?2",
rusqlite::params![provider_id, app_type], [app_type, provider_id],
|row| row.get(0), |row| row.get(0),
) )
.unwrap_or(false); .map_err(|e| AppError::Database(e.to_string()))?;
Ok(in_queue) Ok(count > 0)
} }
/// 获取可添加到故障转移队列的供应商(不在队列中的) /// 获取可添加到故障转移队列的供应商(不在队列中的)
@@ -135,10 +226,14 @@ impl Database {
app_type: &str, app_type: &str,
) -> Result<Vec<Provider>, AppError> { ) -> Result<Vec<Provider>, AppError> {
let all_providers = self.get_all_providers(app_type)?; let all_providers = self.get_all_providers(app_type)?;
let queue = self.get_failover_queue(app_type)?;
let queue_ids: std::collections::HashSet<_> =
queue.iter().map(|item| &item.provider_id).collect();
let available: Vec<Provider> = all_providers let available: Vec<Provider> = all_providers
.into_values() .into_values()
.filter(|p| !p.in_failover_queue) .filter(|p| !queue_ids.contains(&p.id))
.collect(); .collect();
Ok(available) Ok(available)
-1
View File
@@ -10,7 +10,6 @@ pub mod proxy;
pub mod settings; pub mod settings;
pub mod skills; pub mod skills;
pub mod stream_check; pub mod stream_check;
pub mod universal_providers;
// 所有 DAO 方法都通过 Database impl 提供,无需单独导出 // 所有 DAO 方法都通过 Database impl 提供,无需单独导出
// 导出 FailoverQueueItem 供外部使用 // 导出 FailoverQueueItem 供外部使用
+11 -19
View File
@@ -17,7 +17,7 @@ impl Database {
) -> Result<IndexMap<String, Provider>, AppError> { ) -> Result<IndexMap<String, Provider>, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
let mut stmt = conn.prepare( let mut stmt = conn.prepare(
"SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, in_failover_queue "SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta
FROM providers WHERE app_type = ?1 FROM providers WHERE app_type = ?1
ORDER BY COALESCE(sort_index, 999999), created_at ASC, id ASC" ORDER BY COALESCE(sort_index, 999999), created_at ASC, id ASC"
).map_err(|e| AppError::Database(e.to_string()))?; ).map_err(|e| AppError::Database(e.to_string()))?;
@@ -35,7 +35,6 @@ impl Database {
let icon: Option<String> = row.get(8)?; let icon: Option<String> = row.get(8)?;
let icon_color: Option<String> = row.get(9)?; let icon_color: Option<String> = row.get(9)?;
let meta_str: String = row.get(10)?; let meta_str: String = row.get(10)?;
let in_failover_queue: bool = row.get(11)?;
let settings_config = let settings_config =
serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null); serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null);
@@ -55,7 +54,6 @@ impl Database {
meta: Some(meta), meta: Some(meta),
icon, icon,
icon_color, icon_color,
in_failover_queue,
}, },
)) ))
}) })
@@ -131,7 +129,7 @@ impl Database {
) -> Result<Option<Provider>, AppError> { ) -> Result<Option<Provider>, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
let result = conn.query_row( let result = conn.query_row(
"SELECT name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, in_failover_queue "SELECT name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta
FROM providers WHERE id = ?1 AND app_type = ?2", FROM providers WHERE id = ?1 AND app_type = ?2",
params![id, app_type], params![id, app_type],
|row| { |row| {
@@ -145,7 +143,6 @@ impl Database {
let icon: Option<String> = row.get(7)?; let icon: Option<String> = row.get(7)?;
let icon_color: Option<String> = row.get(8)?; let icon_color: Option<String> = row.get(8)?;
let meta_str: String = row.get(9)?; let meta_str: String = row.get(9)?;
let in_failover_queue: bool = row.get(10)?;
let settings_config = serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null); let settings_config = serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null);
let meta: ProviderMeta = serde_json::from_str(&meta_str).unwrap_or_default(); let meta: ProviderMeta = serde_json::from_str(&meta_str).unwrap_or_default();
@@ -162,7 +159,6 @@ impl Database {
meta: Some(meta), meta: Some(meta),
icon, icon,
icon_color, icon_color,
in_failover_queue,
}) })
}, },
); );
@@ -188,18 +184,17 @@ impl Database {
let mut meta_clone = provider.meta.clone().unwrap_or_default(); let mut meta_clone = provider.meta.clone().unwrap_or_default();
let endpoints = std::mem::take(&mut meta_clone.custom_endpoints); let endpoints = std::mem::take(&mut meta_clone.custom_endpoints);
// 检查是否存在(用于判断新增/更新,以及保留 is_current 和 in_failover_queue // 检查是否存在(用于判断新增/更新,以及保留 is_current
let existing: Option<(bool, bool)> = tx let existing: Option<bool> = tx
.query_row( .query_row(
"SELECT is_current, in_failover_queue FROM providers WHERE id = ?1 AND app_type = ?2", "SELECT is_current FROM providers WHERE id = ?1 AND app_type = ?2",
params![provider.id, app_type], params![provider.id, app_type],
|row| Ok((row.get(0)?, row.get(1)?)), |row| row.get(0),
) )
.ok(); .ok();
let is_update = existing.is_some(); let is_update = existing.is_some();
let (is_current, in_failover_queue) = let is_current = existing.unwrap_or(false);
existing.unwrap_or((false, provider.in_failover_queue));
if is_update { if is_update {
// 更新模式:使用 UPDATE 避免触发 ON DELETE CASCADE // 更新模式:使用 UPDATE 避免触发 ON DELETE CASCADE
@@ -215,9 +210,8 @@ impl Database {
icon = ?8, icon = ?8,
icon_color = ?9, icon_color = ?9,
meta = ?10, meta = ?10,
is_current = ?11, is_current = ?11
in_failover_queue = ?12 WHERE id = ?12 AND app_type = ?13",
WHERE id = ?13 AND app_type = ?14",
params![ params![
provider.name, provider.name,
serde_json::to_string(&provider.settings_config).unwrap(), serde_json::to_string(&provider.settings_config).unwrap(),
@@ -230,7 +224,6 @@ impl Database {
provider.icon_color, provider.icon_color,
serde_json::to_string(&meta_clone).unwrap(), serde_json::to_string(&meta_clone).unwrap(),
is_current, is_current,
in_failover_queue,
provider.id, provider.id,
app_type, app_type,
], ],
@@ -241,8 +234,8 @@ impl Database {
tx.execute( tx.execute(
"INSERT INTO providers ( "INSERT INTO providers (
id, app_type, name, settings_config, website_url, category, id, app_type, name, settings_config, website_url, category,
created_at, sort_index, notes, icon, icon_color, meta, is_current, in_failover_queue created_at, sort_index, notes, icon, icon_color, meta, is_current
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)", ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13)",
params![ params![
provider.id, provider.id,
app_type, app_type,
@@ -257,7 +250,6 @@ impl Database {
provider.icon_color, provider.icon_color,
serde_json::to_string(&meta_clone).unwrap(), serde_json::to_string(&meta_clone).unwrap(),
is_current, is_current,
in_failover_queue,
], ],
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
+81 -327
View File
@@ -8,254 +8,63 @@ use crate::proxy::types::*;
use super::super::{lock_conn, Database}; use super::super::{lock_conn, Database};
impl Database { impl Database {
// ==================== Global Proxy Config ==================== // ==================== Proxy Config ====================
/// 获取全局代理配置(统一字段) /// 获取代理配置
///
/// 从 claude 行读取(三行镜像一致)
pub async fn get_global_proxy_config(&self) -> Result<GlobalProxyConfig, AppError> {
// 使用 block 限制 conn 的作用域,避免跨 await 持有锁
let result = {
let conn = lock_conn!(self.conn);
conn.query_row(
"SELECT proxy_enabled, listen_address, listen_port, enable_logging
FROM proxy_config WHERE app_type = 'claude'",
[],
|row| {
Ok(GlobalProxyConfig {
proxy_enabled: row.get::<_, i32>(0)? != 0,
listen_address: row.get(1)?,
listen_port: row.get::<_, i32>(2)? as u16,
enable_logging: row.get::<_, i32>(3)? != 0,
})
},
)
};
// conn 已在 block 结束时释放
match result {
Ok(config) => Ok(config),
Err(rusqlite::Error::QueryReturnedNoRows) => {
// 如果不存在,创建默认配置
self.init_proxy_config_rows().await?;
Ok(GlobalProxyConfig {
proxy_enabled: false,
listen_address: "127.0.0.1".to_string(),
listen_port: 5000,
enable_logging: true,
})
}
Err(e) => Err(AppError::Database(e.to_string())),
}
}
/// 更新全局代理配置(镜像写三行)
pub async fn update_global_proxy_config(
&self,
config: GlobalProxyConfig,
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
conn.execute(
"UPDATE proxy_config SET
proxy_enabled = ?1,
listen_address = ?2,
listen_port = ?3,
enable_logging = ?4,
updated_at = datetime('now')",
rusqlite::params![
if config.proxy_enabled { 1 } else { 0 },
config.listen_address,
config.listen_port as i32,
if config.enable_logging { 1 } else { 0 },
],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
/// 获取应用级代理配置
pub async fn get_proxy_config_for_app(
&self,
app_type: &str,
) -> Result<AppProxyConfig, AppError> {
// 使用 block 限制 conn 的作用域,避免跨 await 持有锁
let app_type_owned = app_type.to_string();
let result = {
let conn = lock_conn!(self.conn);
conn.query_row(
"SELECT app_type, enabled, auto_failover_enabled,
max_retries, streaming_first_byte_timeout, streaming_idle_timeout, non_streaming_timeout,
circuit_failure_threshold, circuit_success_threshold, circuit_timeout_seconds,
circuit_error_rate_threshold, circuit_min_requests
FROM proxy_config WHERE app_type = ?1",
[app_type],
|row| {
Ok(AppProxyConfig {
app_type: row.get(0)?,
enabled: row.get::<_, i32>(1)? != 0,
auto_failover_enabled: row.get::<_, i32>(2)? != 0,
max_retries: row.get::<_, i32>(3)? as u32,
streaming_first_byte_timeout: row.get::<_, i32>(4)? as u32,
streaming_idle_timeout: row.get::<_, i32>(5)? as u32,
non_streaming_timeout: row.get::<_, i32>(6)? as u32,
circuit_failure_threshold: row.get::<_, i32>(7)? as u32,
circuit_success_threshold: row.get::<_, i32>(8)? as u32,
circuit_timeout_seconds: row.get::<_, i32>(9)? as u32,
circuit_error_rate_threshold: row.get(10)?,
circuit_min_requests: row.get::<_, i32>(11)? as u32,
})
},
)
};
// conn 已在 block 结束时释放
match result {
Ok(config) => Ok(config),
Err(rusqlite::Error::QueryReturnedNoRows) => {
// 如果不存在,创建默认配置
self.init_proxy_config_rows().await?;
Ok(AppProxyConfig {
app_type: app_type_owned,
enabled: false,
auto_failover_enabled: false,
max_retries: 3,
streaming_first_byte_timeout: 30,
streaming_idle_timeout: 60,
non_streaming_timeout: 300,
circuit_failure_threshold: 5,
circuit_success_threshold: 2,
circuit_timeout_seconds: 60,
circuit_error_rate_threshold: 0.5,
circuit_min_requests: 10,
})
}
Err(e) => Err(AppError::Database(e.to_string())),
}
}
/// 更新应用级代理配置
pub async fn update_proxy_config_for_app(
&self,
config: AppProxyConfig,
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
conn.execute(
"UPDATE proxy_config SET
enabled = ?2,
auto_failover_enabled = ?3,
max_retries = ?4,
streaming_first_byte_timeout = ?5,
streaming_idle_timeout = ?6,
non_streaming_timeout = ?7,
circuit_failure_threshold = ?8,
circuit_success_threshold = ?9,
circuit_timeout_seconds = ?10,
circuit_error_rate_threshold = ?11,
circuit_min_requests = ?12,
updated_at = datetime('now')
WHERE app_type = ?1",
rusqlite::params![
config.app_type,
if config.enabled { 1 } else { 0 },
if config.auto_failover_enabled { 1 } else { 0 },
config.max_retries as i32,
config.streaming_first_byte_timeout as i32,
config.streaming_idle_timeout as i32,
config.non_streaming_timeout as i32,
config.circuit_failure_threshold as i32,
config.circuit_success_threshold as i32,
config.circuit_timeout_seconds as i32,
config.circuit_error_rate_threshold,
config.circuit_min_requests as i32,
],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
/// 初始化 proxy_config 表的三行数据
async fn init_proxy_config_rows(&self) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
for app_type in &["claude", "codex", "gemini"] {
conn.execute(
"INSERT OR IGNORE INTO proxy_config (app_type) VALUES (?1)",
[app_type],
)
.map_err(|e| AppError::Database(e.to_string()))?;
}
Ok(())
}
// ==================== Legacy Proxy Config (兼容旧代码) ====================
/// 获取代理配置(兼容旧接口,返回 claude 行的配置)
pub async fn get_proxy_config(&self) -> Result<ProxyConfig, AppError> { pub async fn get_proxy_config(&self) -> Result<ProxyConfig, AppError> {
// 使用 block 限制 conn 的作用域,避免跨 await 持有锁 // 在一个作用域内获取锁并查询,确保锁在await之前释放
let result = { let result = {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
conn.query_row( conn.query_row(
"SELECT listen_address, listen_port, max_retries, "SELECT enabled, listen_address, listen_port, max_retries,
enable_logging, request_timeout, enable_logging, live_takeover_active
streaming_first_byte_timeout, streaming_idle_timeout, non_streaming_timeout FROM proxy_config WHERE id = 1",
FROM proxy_config WHERE app_type = 'claude'",
[], [],
|row| { |row| {
Ok(ProxyConfig { Ok(ProxyConfig {
listen_address: row.get(0)?, enabled: row.get::<_, i32>(0)? != 0,
listen_port: row.get::<_, i32>(1)? as u16, listen_address: row.get(1)?,
max_retries: row.get::<_, i32>(2)? as u8, listen_port: row.get::<_, i32>(2)? as u16,
request_timeout: 300, // 废弃字段,返回默认值 max_retries: row.get::<_, i32>(3)? as u8,
enable_logging: row.get::<_, i32>(3)? != 0, request_timeout: row.get::<_, i32>(4)? as u64,
live_takeover_active: false, // 废弃字段 enable_logging: row.get::<_, i32>(5)? != 0,
streaming_first_byte_timeout: row.get::<_, i32>(4).unwrap_or(30) as u64, live_takeover_active: row.get::<_, i32>(6).unwrap_or(0) != 0,
streaming_idle_timeout: row.get::<_, i32>(5).unwrap_or(60) as u64,
non_streaming_timeout: row.get::<_, i32>(6).unwrap_or(300) as u64,
}) })
}, },
) )
}; }; // conn锁在这里释放
// conn 已在 block 结束时释放
match result { match result {
Ok(config) => Ok(config), Ok(config) => Ok(config),
Err(rusqlite::Error::QueryReturnedNoRows) => { Err(rusqlite::Error::QueryReturnedNoRows) => {
// 如果不存在,初始化默认配置 // 如果不存在,插入默认配置
self.init_proxy_config_rows().await?; let default_config = ProxyConfig::default();
Ok(ProxyConfig::default()) self.update_proxy_config(default_config.clone()).await?;
Ok(default_config)
} }
Err(e) => Err(AppError::Database(e.to_string())), Err(e) => Err(AppError::Database(e.to_string())),
} }
} }
/// 更新代理配置(兼容旧接口,更新所有三行的公共字段) /// 更新代理配置
pub async fn update_proxy_config(&self, config: ProxyConfig) -> Result<(), AppError> { pub async fn update_proxy_config(&self, config: ProxyConfig) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
// 更新所有三行的公共字段
conn.execute( conn.execute(
"UPDATE proxy_config SET "INSERT OR REPLACE INTO proxy_config
listen_address = ?1, (id, enabled, listen_address, listen_port, max_retries, request_timeout, enable_logging, live_takeover_active, target_app, created_at, updated_at)
listen_port = ?2, VALUES (1, ?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8,
max_retries = ?3, COALESCE((SELECT created_at FROM proxy_config WHERE id = 1), datetime('now')),
enable_logging = ?4, datetime('now'))",
streaming_first_byte_timeout = ?5,
streaming_idle_timeout = ?6,
non_streaming_timeout = ?7,
updated_at = datetime('now')",
rusqlite::params![ rusqlite::params![
if config.enabled { 1 } else { 0 },
config.listen_address, config.listen_address,
config.listen_port as i32, config.listen_port as i32,
config.max_retries as i32, config.max_retries as i32,
config.request_timeout as i32,
if config.enable_logging { 1 } else { 0 }, if config.enable_logging { 1 } else { 0 },
config.streaming_first_byte_timeout as i32, if config.live_takeover_active { 1 } else { 0 },
config.streaming_idle_timeout as i32, "claude", // 兼容旧字段,写入默认值
config.non_streaming_timeout as i32,
], ],
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
@@ -263,26 +72,28 @@ impl Database {
Ok(()) Ok(())
} }
/// 设置 Live 接管状态(兼容旧版本,更新 enabled 字段) /// 设置 Live 接管状态
pub async fn set_live_takeover_active(&self, _active: bool) -> Result<(), AppError> { pub async fn set_live_takeover_active(&self, active: bool) -> Result<(), AppError> {
// 不再使用此字段,由 enabled 字段替代 let conn = lock_conn!(self.conn);
// 保留空实现以兼容旧代码 conn.execute(
"UPDATE proxy_config SET live_takeover_active = ?1, updated_at = datetime('now') WHERE id = 1",
rusqlite::params![if active { 1 } else { 0 }],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(()) Ok(())
} }
/// 检查是否处于 Live 接管模式 /// 检查是否处于 Live 接管模式
///
/// 检查是否有任一 app 的 enabled = true
pub async fn is_live_takeover_active(&self) -> Result<bool, AppError> { pub async fn is_live_takeover_active(&self) -> Result<bool, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
let count: i64 = conn let active: i32 = conn
.query_row( .query_row(
"SELECT COUNT(*) FROM proxy_config WHERE enabled = 1", "SELECT COALESCE(live_takeover_active, 0) FROM proxy_config WHERE id = 1",
[], [],
|row| row.get(0), |row| row.get(0),
) )
.map_err(|e| AppError::Database(e.to_string()))?; .unwrap_or(0);
Ok(count > 0) Ok(active != 0)
} }
// ==================== Provider Health ==================== // ==================== Provider Health ====================
@@ -293,45 +104,28 @@ impl Database {
provider_id: &str, provider_id: &str,
app_type: &str, app_type: &str,
) -> Result<ProviderHealth, AppError> { ) -> Result<ProviderHealth, AppError> {
let result = { let conn = lock_conn!(self.conn);
let conn = lock_conn!(self.conn);
conn.query_row( conn.query_row(
"SELECT provider_id, app_type, is_healthy, consecutive_failures, "SELECT provider_id, app_type, is_healthy, consecutive_failures,
last_success_at, last_failure_at, last_error, updated_at last_success_at, last_failure_at, last_error, updated_at
FROM provider_health FROM provider_health
WHERE provider_id = ?1 AND app_type = ?2", WHERE provider_id = ?1 AND app_type = ?2",
rusqlite::params![provider_id, app_type], rusqlite::params![provider_id, app_type],
|row| { |row| {
Ok(ProviderHealth { Ok(ProviderHealth {
provider_id: row.get(0)?, provider_id: row.get(0)?,
app_type: row.get(1)?, app_type: row.get(1)?,
is_healthy: row.get::<_, i64>(2)? != 0, is_healthy: row.get::<_, i64>(2)? != 0,
consecutive_failures: row.get::<_, i64>(3)? as u32, consecutive_failures: row.get::<_, i64>(3)? as u32,
last_success_at: row.get(4)?, last_success_at: row.get(4)?,
last_failure_at: row.get(5)?, last_failure_at: row.get(5)?,
last_error: row.get(6)?, last_error: row.get(6)?,
updated_at: row.get(7)?, updated_at: row.get(7)?,
}) })
}, },
) )
}; .map_err(|e| AppError::Database(e.to_string()))
match result {
Ok(health) => Ok(health),
// 缺少记录时视为健康(关闭后清空状态,再次打开时默认正常)
Err(rusqlite::Error::QueryReturnedNoRows) => Ok(ProviderHealth {
provider_id: provider_id.to_string(),
app_type: app_type.to_string(),
is_healthy: true,
consecutive_failures: 0,
last_success_at: None,
last_failure_at: None,
last_error: None,
updated_at: chrono::Utc::now().to_rfc3339(),
}),
Err(e) => Err(AppError::Database(e.to_string())),
}
} }
/// 更新Provider健康状态 /// 更新Provider健康状态
@@ -436,20 +230,6 @@ impl Database {
Ok(()) Ok(())
} }
/// 清空指定应用的健康状态(关闭单个代理时使用)
pub async fn clear_provider_health_for_app(&self, app_type: &str) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
conn.execute(
"DELETE FROM provider_health WHERE app_type = ?1",
[app_type],
)
.map_err(|e| AppError::Database(e.to_string()))?;
log::debug!("Cleared provider health records for app {app_type}");
Ok(())
}
/// 清空所有Provider健康状态(代理停止时调用) /// 清空所有Provider健康状态(代理停止时调用)
pub async fn clear_all_provider_health(&self) -> Result<(), AppError> { pub async fn clear_all_provider_health(&self) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
@@ -461,22 +241,19 @@ impl Database {
Ok(()) Ok(())
} }
// ==================== Circuit Breaker Config (Legacy Compatibility) ==================== // ==================== Circuit Breaker Config ====================
/// 获取熔断器配置(兼容旧接口,从 claude 行读取) /// 获取熔断器配置
///
/// 熔断器配置已合并到 proxy_config 表,每 app 独立
/// 此方法保留用于兼容旧代码,建议使用 get_proxy_config_for_app
pub async fn get_circuit_breaker_config( pub async fn get_circuit_breaker_config(
&self, &self,
) -> Result<crate::proxy::circuit_breaker::CircuitBreakerConfig, AppError> { ) -> Result<crate::proxy::circuit_breaker::CircuitBreakerConfig, AppError> {
// 使用 block 限制 conn 的作用域,避免跨 await 持有锁 let conn = lock_conn!(self.conn);
let result = {
let conn = lock_conn!(self.conn); let config = conn
conn.query_row( .query_row(
"SELECT circuit_failure_threshold, circuit_success_threshold, circuit_timeout_seconds, "SELECT failure_threshold, success_threshold, timeout_seconds,
circuit_error_rate_threshold, circuit_min_requests error_rate_threshold, min_requests
FROM proxy_config WHERE app_type = 'claude'", FROM circuit_breaker_config WHERE id = 1",
[], [],
|row| { |row| {
Ok(crate::proxy::circuit_breaker::CircuitBreakerConfig { Ok(crate::proxy::circuit_breaker::CircuitBreakerConfig {
@@ -488,39 +265,27 @@ impl Database {
}) })
}, },
) )
}; .map_err(|e| AppError::Database(e.to_string()))?;
// conn 已在 block 结束时释放
match result { Ok(config)
Ok(config) => Ok(config),
Err(rusqlite::Error::QueryReturnedNoRows) => {
// 如果不存在,初始化默认配置
self.init_proxy_config_rows().await?;
Ok(crate::proxy::circuit_breaker::CircuitBreakerConfig::default())
}
Err(e) => Err(AppError::Database(e.to_string())),
}
} }
/// 更新熔断器配置(兼容旧接口,更新所有三行) /// 更新熔断器配置
///
/// 熔断器配置已合并到 proxy_config 表
/// 此方法保留用于兼容旧代码,建议使用 update_proxy_config_for_app
pub async fn update_circuit_breaker_config( pub async fn update_circuit_breaker_config(
&self, &self,
config: &crate::proxy::circuit_breaker::CircuitBreakerConfig, config: &crate::proxy::circuit_breaker::CircuitBreakerConfig,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
// 更新所有三行的熔断器配置
conn.execute( conn.execute(
"UPDATE proxy_config SET "UPDATE circuit_breaker_config
circuit_failure_threshold = ?1, SET failure_threshold = ?1,
circuit_success_threshold = ?2, success_threshold = ?2,
circuit_timeout_seconds = ?3, timeout_seconds = ?3,
circuit_error_rate_threshold = ?4, error_rate_threshold = ?4,
circuit_min_requests = ?5, min_requests = ?5,
updated_at = datetime('now')", updated_at = CURRENT_TIMESTAMP
WHERE id = 1",
rusqlite::params![ rusqlite::params![
config.failure_threshold as i32, config.failure_threshold as i32,
config.success_threshold as i32, config.success_threshold as i32,
@@ -556,17 +321,6 @@ impl Database {
Ok(()) Ok(())
} }
/// 检查是否存在任意 Live 配置备份
pub async fn has_any_live_backup(&self) -> Result<bool, AppError> {
let conn = lock_conn!(self.conn);
let count: i64 = conn
.query_row("SELECT COUNT(*) FROM proxy_live_backup", [], |row| {
row.get(0)
})
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(count > 0)
}
/// 获取 Live 配置备份 /// 获取 Live 配置备份
pub async fn get_live_backup(&self, app_type: &str) -> Result<Option<LiveBackup>, AppError> { pub async fn get_live_backup(&self, app_type: &str) -> Result<Option<LiveBackup>, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
-66
View File
@@ -62,70 +62,4 @@ impl Database {
Ok(()) Ok(())
} }
} }
// --- 代理接管状态管理(已废弃,使用 proxy_config.enabled 替代)---
/// 获取指定应用的代理接管状态
///
/// **已废弃**: 请使用 `proxy_config.enabled` 字段替代
/// 此方法仅用于数据库迁移时读取旧数据
#[deprecated(since = "3.9.0", note = "使用 get_proxy_config_for_app().enabled 替代")]
pub fn get_proxy_takeover_enabled(&self, app_type: &str) -> Result<bool, AppError> {
let key = format!("proxy_takeover_{app_type}");
match self.get_setting(&key)? {
Some(value) => Ok(value == "true"),
None => Ok(false),
}
}
/// 设置指定应用的代理接管状态
///
/// **已废弃**: 请使用 `proxy_config.enabled` 字段替代
#[deprecated(
since = "3.9.0",
note = "使用 update_proxy_config_for_app() 修改 enabled 字段"
)]
pub fn set_proxy_takeover_enabled(
&self,
app_type: &str,
enabled: bool,
) -> Result<(), AppError> {
let key = format!("proxy_takeover_{app_type}");
let value = if enabled { "true" } else { "false" };
self.set_setting(&key, value)
}
/// 检查是否有任一应用开启了代理接管
///
/// **已废弃**: 请使用 `is_live_takeover_active()` 替代
#[deprecated(since = "3.9.0", note = "使用 is_live_takeover_active() 替代")]
pub fn has_any_proxy_takeover(&self) -> Result<bool, AppError> {
let conn = lock_conn!(self.conn);
let count: i64 = conn
.query_row(
"SELECT COUNT(*) FROM settings WHERE key LIKE 'proxy_takeover_%' AND value = 'true'",
[],
|row| row.get(0),
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(count > 0)
}
/// 清除所有代理接管状态(将所有 proxy_takeover_* 设置为 false
///
/// **已废弃**: settings 表不再用于存储代理状态
#[deprecated(
since = "3.9.0",
note = "使用 update_proxy_config_for_app() 清除各应用的 enabled 字段"
)]
pub fn clear_all_proxy_takeover(&self) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
conn.execute(
"UPDATE settings SET value = 'false' WHERE key LIKE 'proxy_takeover_%'",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
log::info!("已清除所有代理接管状态");
Ok(())
}
} }
@@ -1,74 +0,0 @@
//! 统一供应商 (Universal Provider) DAO
//!
//! 提供统一供应商的 CRUD 操作。
use crate::database::{lock_conn, to_json_string, Database};
use crate::error::AppError;
use crate::provider::UniversalProvider;
use std::collections::HashMap;
/// 统一供应商的 Settings Key
const UNIVERSAL_PROVIDERS_KEY: &str = "universal_providers";
impl Database {
/// 获取所有统一供应商
pub fn get_all_universal_providers(
&self,
) -> Result<HashMap<String, UniversalProvider>, AppError> {
let conn = lock_conn!(self.conn);
let mut stmt = conn
.prepare("SELECT value FROM settings WHERE key = ?")
.map_err(|e| AppError::Database(e.to_string()))?;
let result: Option<String> = stmt
.query_row([UNIVERSAL_PROVIDERS_KEY], |row| row.get(0))
.ok();
match result {
Some(json) => serde_json::from_str(&json)
.map_err(|e| AppError::Database(format!("解析统一供应商数据失败: {e}"))),
None => Ok(HashMap::new()),
}
}
/// 获取单个统一供应商
pub fn get_universal_provider(&self, id: &str) -> Result<Option<UniversalProvider>, AppError> {
let providers = self.get_all_universal_providers()?;
Ok(providers.get(id).cloned())
}
/// 保存统一供应商(添加或更新)
pub fn save_universal_provider(&self, provider: &UniversalProvider) -> Result<(), AppError> {
let mut providers = self.get_all_universal_providers()?;
providers.insert(provider.id.clone(), provider.clone());
self.save_all_universal_providers(&providers)
}
/// 删除统一供应商
pub fn delete_universal_provider(&self, id: &str) -> Result<bool, AppError> {
let mut providers = self.get_all_universal_providers()?;
let existed = providers.remove(id).is_some();
if existed {
self.save_all_universal_providers(&providers)?;
}
Ok(existed)
}
/// 保存所有统一供应商(内部方法)
fn save_all_universal_providers(
&self,
providers: &HashMap<String, UniversalProvider>,
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
let json = to_json_string(providers)?;
conn.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES (?, ?)",
[UNIVERSAL_PROVIDERS_KEY, &json],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
}
File diff suppressed because it is too large Load Diff
+7 -54
View File
@@ -53,6 +53,7 @@ const LEGACY_SCHEMA_SQL: &str = r#"
#[derive(Debug)] #[derive(Debug)]
struct ColumnInfo { struct ColumnInfo {
name: String,
r#type: String, r#type: String,
notnull: i64, notnull: i64,
default: Option<String>, default: Option<String>,
@@ -64,9 +65,10 @@ fn get_column_info(conn: &Connection, table: &str, column: &str) -> ColumnInfo {
.expect("prepare pragma"); .expect("prepare pragma");
let mut rows = stmt.query([]).expect("query pragma"); let mut rows = stmt.query([]).expect("query pragma");
while let Some(row) = rows.next().expect("read row") { while let Some(row) = rows.next().expect("read row") {
let column_name: String = row.get(1).expect("name"); let name: String = row.get(1).expect("name");
if column_name.eq_ignore_ascii_case(column) { if name.eq_ignore_ascii_case(column) {
return ColumnInfo { return ColumnInfo {
name,
r#type: row.get::<_, String>(2).expect("type"), r#type: row.get::<_, String>(2).expect("type"),
notnull: row.get::<_, i64>(3).expect("notnull"), notnull: row.get::<_, i64>(3).expect("notnull"),
default: row.get::<_, Option<String>>(4).ok().flatten(), default: row.get::<_, Option<String>>(4).ok().flatten(),
@@ -199,53 +201,6 @@ fn migration_aligns_column_defaults_and_types() {
); );
} }
#[test]
fn create_tables_repairs_legacy_proxy_config_singleton_to_per_app() {
let conn = Connection::open_in_memory().expect("open memory db");
// 模拟测试版 v2user_version=2,但 proxy_config 仍是单例结构(无 app_type
Database::set_user_version(&conn, 2).expect("set user_version");
conn.execute_batch(
r#"
CREATE TABLE proxy_config (
id INTEGER PRIMARY KEY,
enabled INTEGER NOT NULL DEFAULT 0,
listen_address TEXT NOT NULL DEFAULT '127.0.0.1',
listen_port INTEGER NOT NULL DEFAULT 5000,
max_retries INTEGER NOT NULL DEFAULT 3,
request_timeout INTEGER NOT NULL DEFAULT 300,
enable_logging INTEGER NOT NULL DEFAULT 1,
target_app TEXT NOT NULL DEFAULT 'claude',
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
INSERT INTO proxy_config (id, enabled) VALUES (1, 1);
"#,
)
.expect("seed legacy proxy_config");
Database::create_tables_on_conn(&conn).expect("create tables should repair proxy_config");
assert!(
Database::has_column(&conn, "proxy_config", "app_type").expect("check app_type"),
"proxy_config should be migrated to per-app structure"
);
let count: i32 = conn
.query_row("SELECT COUNT(*) FROM proxy_config", [], |r| r.get(0))
.expect("count rows");
assert_eq!(count, 3, "per-app proxy_config should have 3 rows");
// 新结构下应能按 app_type 查询
let _: i32 = conn
.query_row(
"SELECT COUNT(*) FROM proxy_config WHERE app_type = 'claude'",
[],
|r| r.get(0),
)
.expect("query by app_type");
}
#[test] #[test]
fn dry_run_does_not_write_to_disk() { fn dry_run_does_not_write_to_disk() {
// Create minimal valid config for migration // Create minimal valid config for migration
@@ -290,14 +245,12 @@ fn dry_run_validates_schema_compatibility() {
meta: None, meta: None,
icon: None, icon: None,
icon_color: None, icon_color: None,
in_failover_queue: false,
}, },
); );
let manager = ProviderManager { let mut manager = ProviderManager::default();
providers, manager.providers = providers;
current: "test-provider".to_string(), manager.current = "test-provider".to_string();
};
let mut apps = HashMap::new(); let mut apps = HashMap::new();
apps.insert("claude".to_string(), manager); apps.insert("claude".to_string(), manager);
-1
View File
@@ -132,7 +132,6 @@ pub(crate) fn build_provider_from_request(
meta, meta,
icon: request.icon.clone(), icon: request.icon.clone(),
icon_color: None, icon_color: None,
in_failover_queue: false,
}; };
Ok(provider) Ok(provider)
+3 -3
View File
@@ -375,7 +375,7 @@ fn test_parse_prompt_deeplink() {
assert_eq!(request.name.unwrap(), "test"); assert_eq!(request.name.unwrap(), "test");
assert_eq!(request.content.unwrap(), content_b64); assert_eq!(request.content.unwrap(), content_b64);
assert_eq!(request.description.unwrap(), "desc"); assert_eq!(request.description.unwrap(), "desc");
assert!(request.enabled.unwrap()); assert_eq!(request.enabled.unwrap(), true);
} }
#[test] #[test]
@@ -391,13 +391,13 @@ fn test_parse_mcp_deeplink() {
assert_eq!(request.resource, "mcp"); assert_eq!(request.resource, "mcp");
assert_eq!(request.apps.unwrap(), "claude,codex"); assert_eq!(request.apps.unwrap(), "claude,codex");
assert_eq!(request.config.unwrap(), config_b64); assert_eq!(request.config.unwrap(), config_b64);
assert!(request.enabled.unwrap()); assert_eq!(request.enabled.unwrap(), true);
} }
#[test] #[test]
fn test_parse_skill_deeplink() { fn test_parse_skill_deeplink() {
let url = "ccswitch://v1/import?resource=skill&repo=owner/repo&directory=skills&branch=dev"; let url = "ccswitch://v1/import?resource=skill&repo=owner/repo&directory=skills&branch=dev";
let request = parse_deeplink_url(url).unwrap(); let request = parse_deeplink_url(&url).unwrap();
assert_eq!(request.resource, "skill"); assert_eq!(request.resource, "skill");
assert_eq!(request.repo.unwrap(), "owner/repo"); assert_eq!(request.repo.unwrap(), "owner/repo");
-4
View File
@@ -52,10 +52,6 @@ pub enum AppError {
}, },
#[error("数据库错误: {0}")] #[error("数据库错误: {0}")]
Database(String), Database(String),
#[error("所有供应商已熔断,无可用渠道")]
AllProvidersCircuitOpen,
#[error("未配置供应商")]
NoProvidersConfigured,
} }
impl AppError { impl AppError {
+4 -11
View File
@@ -42,8 +42,8 @@ fn write_json_value(path: &Path, value: &Value) -> Result<(), AppError> {
/// ///
/// 执行反向格式转换以保持与统一 MCP 结构的兼容性: /// 执行反向格式转换以保持与统一 MCP 结构的兼容性:
/// - httpUrl → url + type: "http" /// - httpUrl → url + type: "http"
/// - 仅有 url 字段 → 补齐 type: "sse"Gemini 以字段名推断传输类型) /// - 仅有 url 字段 → 保持不变(SSE 类型)
/// - 仅有 command 字段 → 补齐 type: "stdio" /// - 仅有 command 字段 → 保持不变(stdio 类型)
pub fn read_mcp_servers_map() -> Result<std::collections::HashMap<String, Value>, AppError> { pub fn read_mcp_servers_map() -> Result<std::collections::HashMap<String, Value>, AppError> {
let path = user_config_path(); let path = user_config_path();
if !path.exists() { if !path.exists() {
@@ -65,15 +65,8 @@ pub fn read_mcp_servers_map() -> Result<std::collections::HashMap<String, Value>
obj.insert("url".to_string(), http_url); obj.insert("url".to_string(), http_url);
obj.insert("type".to_string(), Value::String("http".to_string())); obj.insert("type".to_string(), Value::String("http".to_string()));
} }
// 如果有 url 但没有 type,不添加 type(默认为 SSE
// Gemini CLI 不使用 type 字段:这里补齐成统一结构,便于校验与导入 // 如果有 command 但没有 type,不添加 type(默认为 stdio
if obj.get("type").is_none() {
if obj.contains_key("command") {
obj.insert("type".to_string(), Value::String("stdio".to_string()));
} else if obj.contains_key("url") {
obj.insert("type".to_string(), Value::String("sse".to_string()));
}
}
} }
} }
+101 -150
View File
@@ -48,9 +48,8 @@ use tauri_plugin_deep_link::DeepLinkExt;
use tauri_plugin_dialog::{DialogExt, MessageDialogButtons, MessageDialogKind}; use tauri_plugin_dialog::{DialogExt, MessageDialogButtons, MessageDialogKind};
use std::sync::Arc; use std::sync::Arc;
#[cfg(target_os = "macos")]
use tauri::image::Image;
use tauri::tray::{TrayIconBuilder, TrayIconEvent}; use tauri::tray::{TrayIconBuilder, TrayIconEvent};
#[cfg(target_os = "macos")]
use tauri::RunEvent; use tauri::RunEvent;
use tauri::{Emitter, Manager}; use tauri::{Emitter, Manager};
@@ -135,19 +134,6 @@ async fn update_tray_menu(
} }
} }
#[cfg(target_os = "macos")]
fn macos_tray_icon() -> Option<Image<'static>> {
const ICON_BYTES: &[u8] = include_bytes!("../icons/tray/macos/statusbar_template_3x.png");
match Image::from_bytes(ICON_BYTES) {
Ok(icon) => Some(icon),
Err(err) => {
log::warn!("Failed to load macOS tray icon: {err}");
None
}
}
}
#[cfg_attr(mobile, tauri::mobile_entry_point)] #[cfg_attr(mobile, tauri::mobile_entry_point)]
pub fn run() { pub fn run() {
let mut builder = tauri::Builder::default(); let mut builder = tauri::Builder::default();
@@ -223,6 +209,44 @@ pub fn run() {
log::warn!("初始化 Updater 插件失败,已跳过:{e}"); log::warn!("初始化 Updater 插件失败,已跳过:{e}");
} }
} }
#[cfg(target_os = "macos")]
{
// 设置 macOS 标题栏背景色为主界面蓝色
if let Some(window) = app.get_webview_window("main") {
use objc2::rc::Retained;
use objc2::runtime::AnyObject;
use objc2_app_kit::NSColor;
match window.ns_window() {
Ok(ns_window_ptr) => {
if let Some(ns_window) =
unsafe { Retained::retain(ns_window_ptr as *mut AnyObject) }
{
// 使用与主界面 banner 相同的蓝色 #3498db
// #3498db = RGB(52, 152, 219)
let bg_color = unsafe {
NSColor::colorWithRed_green_blue_alpha(
52.0 / 255.0, // R: 52
152.0 / 255.0, // G: 152
219.0 / 255.0, // B: 219
1.0, // Alpha: 1.0
)
};
unsafe {
use objc2::msg_send;
let _: () =
msg_send![&*ns_window, setBackgroundColor: &*bg_color];
}
} else {
log::warn!("Failed to retain NSWindow reference");
}
}
Err(e) => log::warn!("Failed to get NSWindow pointer: {e}"),
}
}
}
// 初始化日志 // 初始化日志
if cfg!(debug_assertions) { if cfg!(debug_assertions) {
app.handle().plugin( app.handle().plugin(
@@ -480,26 +504,11 @@ pub fn run() {
}) })
.show_menu_on_left_click(true); .show_menu_on_left_click(true);
// 使用平台对应的托盘图标(macOS 使用模板图标适配深浅色) // 统一使用应用默认图标;待托盘模板图标就绪后再启用
#[cfg(target_os = "macos")] if let Some(icon) = app.default_window_icon() {
{ tray_builder = tray_builder.icon(icon.clone());
if let Some(icon) = macos_tray_icon() { } else {
tray_builder = tray_builder.icon(icon).icon_as_template(true); log::warn!("Failed to get default window icon for tray");
} else if let Some(icon) = app.default_window_icon() {
log::warn!("Falling back to default window icon for tray");
tray_builder = tray_builder.icon(icon.clone());
} else {
log::warn!("Failed to load macOS tray icon for tray");
}
}
#[cfg(not(target_os = "macos"))]
{
if let Some(icon) = app.default_window_icon() {
tray_builder = tray_builder.icon(icon.clone());
} else {
log::warn!("Failed to get default window icon for tray");
}
} }
let _tray = tray_builder.build(app)?; let _tray = tray_builder.build(app)?;
@@ -516,33 +525,49 @@ pub fn run() {
} }
} }
// 异常退出恢复 + 代理状态自动恢复 // 异常退出恢复 + 自动启动代理服务器
let app_handle = app.handle().clone(); let app_handle = app.handle().clone();
tauri::async_runtime::spawn(async move { tauri::async_runtime::spawn(async move {
let state = app_handle.state::<AppState>(); let state = app_handle.state::<AppState>();
// 检查是否有 Live 备份(表示上次异常退出时可能处于接管状态) // 1. 检测异常退出并恢复 Live 配置
let has_backups = match state.db.has_any_live_backup().await { match state.db.is_live_takeover_active().await {
Ok(v) => v, Ok(true) => {
Err(e) => { // 接管标志为 true 但代理未运行 → 上次异常退出
log::error!("检查 Live 备份失败: {e}"); if !state.proxy_service.is_running().await {
false log::warn!("检测到上次异常退出,正在恢复 Live 配置...");
if let Err(e) = state.proxy_service.recover_from_crash().await {
log::error!("恢复 Live 配置失败: {e}");
} else {
log::info!("Live 配置已从异常退出中恢复");
}
}
} }
}; Ok(false) => {
// 检查 Live 配置是否仍处于被接管状态(包含占位符) // 正常状态,无需恢复
let live_taken_over = state.proxy_service.detect_takeover_in_live_configs(); }
Err(e) => {
if has_backups || live_taken_over { log::error!("检查接管状态失败: {e}");
log::warn!("检测到上次异常退出(存在接管残留),正在恢复 Live 配置...");
if let Err(e) = state.proxy_service.recover_from_crash().await {
log::error!("恢复 Live 配置失败: {e}");
} else {
log::info!("Live 配置已恢复");
} }
} }
// 检查 settings 表中的代理状态,自动恢复代理服务 // 2. 自动启动代理服务器(如果配置为启用)
restore_proxy_state_on_startup(&state).await; match state.db.get_proxy_config().await {
Ok(config) => {
if config.enabled {
log::info!("代理服务配置为启用,正在启动...");
match state.proxy_service.start_with_takeover().await {
Ok(info) => log::info!(
"代理服务器自动启动成功: {}:{}",
info.address,
info.port
),
Err(e) => log::error!("代理服务器自动启动失败: {e}"),
}
}
}
Err(e) => log::error!("启动时获取代理配置失败: {e}"),
}
}); });
Ok(()) Ok(())
@@ -580,8 +605,6 @@ pub fn run() {
commands::read_claude_plugin_config, commands::read_claude_plugin_config,
commands::apply_claude_plugin_config, commands::apply_claude_plugin_config,
commands::is_claude_plugin_applied, commands::is_claude_plugin_applied,
commands::apply_claude_onboarding_skip,
commands::clear_claude_onboarding_skip,
// Claude MCP management // Claude MCP management
commands::get_claude_mcp_status, commands::get_claude_mcp_status,
commands::read_claude_mcp_config, commands::read_claude_mcp_config,
@@ -596,7 +619,7 @@ pub fn run() {
commands::upsert_mcp_server_in_config, commands::upsert_mcp_server_in_config,
commands::delete_mcp_server_in_config, commands::delete_mcp_server_in_config,
commands::set_mcp_enabled, commands::set_mcp_enabled,
// Unified MCP management // v3.7.0: Unified MCP management
commands::get_mcp_servers, commands::get_mcp_servers,
commands::upsert_mcp_server, commands::upsert_mcp_server,
commands::delete_mcp_server, commands::delete_mcp_server,
@@ -649,18 +672,11 @@ pub fn run() {
commands::set_auto_launch, commands::set_auto_launch,
commands::get_auto_launch_status, commands::get_auto_launch_status,
// Proxy server management // Proxy server management
commands::start_proxy_server, commands::start_proxy_with_takeover,
commands::stop_proxy_with_restore, commands::stop_proxy_with_restore,
commands::get_proxy_takeover_status,
commands::set_proxy_takeover_for_app,
commands::get_proxy_status, commands::get_proxy_status,
commands::get_proxy_config, commands::get_proxy_config,
commands::update_proxy_config, commands::update_proxy_config,
// Global & Per-App Config
commands::get_global_proxy_config,
commands::update_global_proxy_config,
commands::get_proxy_config_for_app,
commands::update_proxy_config_for_app,
commands::is_proxy_running, commands::is_proxy_running,
commands::is_live_takeover_active, commands::is_live_takeover_active,
commands::switch_proxy_provider, commands::switch_proxy_provider,
@@ -675,8 +691,8 @@ pub fn run() {
commands::get_available_providers_for_failover, commands::get_available_providers_for_failover,
commands::add_to_failover_queue, commands::add_to_failover_queue,
commands::remove_from_failover_queue, commands::remove_from_failover_queue,
commands::get_auto_failover_enabled, commands::reorder_failover_queue,
commands::set_auto_failover_enabled, commands::set_failover_item_enabled,
// Usage statistics // Usage statistics
commands::get_usage_summary, commands::get_usage_summary,
commands::get_usage_trends, commands::get_usage_trends,
@@ -694,12 +710,6 @@ pub fn run() {
commands::get_stream_check_config, commands::get_stream_check_config,
commands::save_stream_check_config, commands::save_stream_check_config,
commands::get_tool_versions, commands::get_tool_versions,
// Universal Provider management
commands::get_universal_providers,
commands::get_universal_provider,
commands::upsert_universal_provider,
commands::delete_universal_provider,
commands::sync_universal_provider,
]); ]);
let app = builder let app = builder
@@ -814,91 +824,32 @@ pub fn run() {
/// ///
/// 在应用退出前检查代理服务器状态,如果正在运行则停止代理并恢复 Live 配置。 /// 在应用退出前检查代理服务器状态,如果正在运行则停止代理并恢复 Live 配置。
/// 确保 Claude Code/Codex/Gemini 的配置不会处于损坏状态。 /// 确保 Claude Code/Codex/Gemini 的配置不会处于损坏状态。
/// 使用 stop_with_restore_keep_state 保留 settings 表中的代理状态,下次启动时自动恢复。
pub async fn cleanup_before_exit(app_handle: &tauri::AppHandle) { pub async fn cleanup_before_exit(app_handle: &tauri::AppHandle) {
if let Some(state) = app_handle.try_state::<store::AppState>() { if let Some(state) = app_handle.try_state::<store::AppState>() {
let proxy_service = &state.proxy_service; let proxy_service = &state.proxy_service;
// 退出时也需要兜底:代理可能已崩溃/未运行,但 Live 接管残留仍在(占位符/备份)。 // 检查代理是否在运行
let has_backups = match state.db.has_any_live_backup().await {
Ok(v) => v,
Err(e) => {
log::error!("退出时检查 Live 备份失败: {e}");
false
}
};
let live_taken_over = proxy_service.detect_takeover_in_live_configs();
let needs_restore = has_backups || live_taken_over;
if needs_restore {
log::info!("检测到接管残留,开始恢复 Live 配置(保留代理状态)...");
// 使用 keep_state 版本,保留 settings 表中的代理状态
if let Err(e) = proxy_service.stop_with_restore_keep_state().await {
log::error!("退出时恢复 Live 配置失败: {e}");
} else {
log::info!("已恢复 Live 配置(代理状态已保留,下次启动将自动恢复)");
}
return;
}
// 非接管模式:代理在运行则仅停止代理
if proxy_service.is_running().await { if proxy_service.is_running().await {
log::info!("检测到代理服务器正在运行,开始停止..."); log::info!("检测到代理服务器正在运行,开始清理...");
if let Err(e) = proxy_service.stop().await {
log::error!("退出时停止代理失败: {e}");
}
log::info!("代理服务器清理完成");
}
}
}
// ============================================================ // 检查是否处于 Live 接管模式
// 启动时恢复代理状态 if let Ok(is_takeover) = state.db.is_live_takeover_active().await {
// ============================================================ if is_takeover {
// 接管模式:停止并恢复配置
/// 启动时根据 proxy_config 表中的代理状态自动恢复代理服务 if let Err(e) = proxy_service.stop_with_restore().await {
/// log::error!("退出时恢复 Live 配置失败: {e}");
/// 检查 `proxy_config.enabled` 字段,如果有任一应用的状态为 `true`, } else {
/// 则自动启动代理服务并接管对应应用的 Live 配置。 log::info!("已恢复 Live 配置");
async fn restore_proxy_state_on_startup(state: &store::AppState) { }
// 收集需要恢复接管的应用列表(从 proxy_config.enabled 读取) } else {
let mut apps_to_restore = Vec::new(); // 非接管模式:仅停止代理
for app_type in ["claude", "codex", "gemini"] { if let Err(e) = proxy_service.stop().await {
if let Ok(config) = state.db.get_proxy_config_for_app(app_type).await { log::error!("退出时停止代理失败: {e}");
if config.enabled { }
apps_to_restore.push(app_type);
}
}
}
if apps_to_restore.is_empty() {
log::debug!("启动时无需恢复代理状态");
return;
}
log::info!("检测到上次代理状态需要恢复,应用列表: {apps_to_restore:?}");
// 逐个恢复接管状态
for app_type in apps_to_restore {
match state
.proxy_service
.set_takeover_for_app(app_type, true)
.await
{
Ok(()) => {
log::info!("✓ 已恢复 {app_type} 的代理接管状态");
}
Err(e) => {
log::error!("✗ 恢复 {app_type} 的代理接管状态失败: {e}");
// 失败时清除该应用的状态,避免下次启动再次尝试
if let Err(clear_err) = state
.proxy_service
.set_takeover_for_app(app_type, false)
.await
{
log::error!("清除 {app_type} 代理状态失败: {clear_err}");
} }
} }
log::info!("代理服务器清理完成");
} }
} }
} }
-15
View File
@@ -8,12 +8,6 @@ use crate::error::AppError;
use super::validation::{extract_server_spec, validate_server_spec}; use super::validation::{extract_server_spec, validate_server_spec};
fn should_sync_claude_mcp() -> bool {
// Claude 未安装/未初始化时:通常 ~/.claude 目录与 ~/.claude.json 都不存在。
// 按用户偏好:此时跳过写入/删除,不创建任何文件或目录。
crate::config::get_claude_config_dir().exists() || crate::config::get_claude_mcp_path().exists()
}
/// 返回已启用的 MCP 服务器(过滤 enabled==true /// 返回已启用的 MCP 服务器(过滤 enabled==true
fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> { fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> {
let mut out = HashMap::new(); let mut out = HashMap::new();
@@ -39,9 +33,6 @@ fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> {
/// 将 config.json 中 enabled==true 的项投影写入 ~/.claude.json /// 将 config.json 中 enabled==true 的项投影写入 ~/.claude.json
pub fn sync_enabled_to_claude(config: &MultiAppConfig) -> Result<(), AppError> { pub fn sync_enabled_to_claude(config: &MultiAppConfig) -> Result<(), AppError> {
if !should_sync_claude_mcp() {
return Ok(());
}
let enabled = collect_enabled_servers(&config.mcp.claude); let enabled = collect_enabled_servers(&config.mcp.claude);
crate::claude_mcp::set_mcp_servers_map(&enabled) crate::claude_mcp::set_mcp_servers_map(&enabled)
} }
@@ -116,9 +107,6 @@ pub fn sync_single_server_to_claude(
id: &str, id: &str,
server_spec: &Value, server_spec: &Value,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
if !should_sync_claude_mcp() {
return Ok(());
}
// 读取现有的 MCP 配置 // 读取现有的 MCP 配置
let current = crate::claude_mcp::read_mcp_servers_map()?; let current = crate::claude_mcp::read_mcp_servers_map()?;
@@ -132,9 +120,6 @@ pub fn sync_single_server_to_claude(
/// 从 Claude live 配置中移除单个 MCP 服务器 /// 从 Claude live 配置中移除单个 MCP 服务器
pub fn remove_server_from_claude(id: &str) -> Result<(), AppError> { pub fn remove_server_from_claude(id: &str) -> Result<(), AppError> {
if !should_sync_claude_mcp() {
return Ok(());
}
// 读取现有的 MCP 配置 // 读取现有的 MCP 配置
let mut current = crate::claude_mcp::read_mcp_servers_map()?; let mut current = crate::claude_mcp::read_mcp_servers_map()?;
+8 -35
View File
@@ -13,12 +13,6 @@ use crate::error::AppError;
use super::validation::{extract_server_spec, validate_server_spec}; use super::validation::{extract_server_spec, validate_server_spec};
fn should_sync_codex_mcp() -> bool {
// Codex 未安装/未初始化时:~/.codex 目录不存在。
// 按用户偏好:目录缺失时跳过写入/删除,不创建任何文件或目录。
crate::codex_config::get_codex_config_dir().exists()
}
/// 返回已启用的 MCP 服务器(过滤 enabled==true /// 返回已启用的 MCP 服务器(过滤 enabled==true
fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> { fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> {
let mut out = HashMap::new(); let mut out = HashMap::new();
@@ -279,9 +273,6 @@ pub fn import_from_codex(config: &mut MultiAppConfig) -> Result<usize, AppError>
/// - 仅更新 `mcp_servers` 表,保留其它键 /// - 仅更新 `mcp_servers` 表,保留其它键
/// - 仅写入启用项;无启用项时清理 mcp_servers 表 /// - 仅写入启用项;无启用项时清理 mcp_servers 表
pub fn sync_enabled_to_codex(config: &MultiAppConfig) -> Result<(), AppError> { pub fn sync_enabled_to_codex(config: &MultiAppConfig) -> Result<(), AppError> {
if !should_sync_codex_mcp() {
return Ok(());
}
use toml_edit::{Item, Table}; use toml_edit::{Item, Table};
// 1) 收集启用项(Codex 维度) // 1) 收集启用项(Codex 维度)
@@ -348,9 +339,6 @@ pub fn sync_single_server_to_codex(
id: &str, id: &str,
server_spec: &Value, server_spec: &Value,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
if !should_sync_codex_mcp() {
return Ok(());
}
use toml_edit::Item; use toml_edit::Item;
// 读取现有的 config.toml // 读取现有的 config.toml
@@ -359,14 +347,9 @@ pub fn sync_single_server_to_codex(
let mut doc = if config_path.exists() { let mut doc = if config_path.exists() {
let content = let content =
std::fs::read_to_string(&config_path).map_err(|e| AppError::io(&config_path, e))?; std::fs::read_to_string(&config_path).map_err(|e| AppError::io(&config_path, e))?;
// 尝试解析现有配置,如果失败则创建新文档(容错处理) content
match content.parse::<toml_edit::DocumentMut>() { .parse::<toml_edit::DocumentMut>()
Ok(doc) => doc, .map_err(|e| AppError::McpValidation(format!("解析 Codex config.toml 失败: {e}")))?
Err(e) => {
log::warn!("解析 Codex config.toml 失败: {e},将创建新配置");
toml_edit::DocumentMut::new()
}
}
} else { } else {
toml_edit::DocumentMut::new() toml_edit::DocumentMut::new()
}; };
@@ -393,8 +376,7 @@ pub fn sync_single_server_to_codex(
doc["mcp_servers"][id] = Item::Table(toml_table); doc["mcp_servers"][id] = Item::Table(toml_table);
// 写回文件 // 写回文件
let new_text = doc.to_string(); std::fs::write(&config_path, doc.to_string()).map_err(|e| AppError::io(&config_path, e))?;
crate::config::write_text_file(&config_path, &new_text)?;
Ok(()) Ok(())
} }
@@ -402,9 +384,6 @@ pub fn sync_single_server_to_codex(
/// 从 Codex live 配置中移除单个 MCP 服务器 /// 从 Codex live 配置中移除单个 MCP 服务器
/// 从正确的 [mcp_servers] 表中删除,同时清理可能存在于错误位置 [mcp.servers] 的数据 /// 从正确的 [mcp_servers] 表中删除,同时清理可能存在于错误位置 [mcp.servers] 的数据
pub fn remove_server_from_codex(id: &str) -> Result<(), AppError> { pub fn remove_server_from_codex(id: &str) -> Result<(), AppError> {
if !should_sync_codex_mcp() {
return Ok(());
}
let config_path = crate::codex_config::get_codex_config_path(); let config_path = crate::codex_config::get_codex_config_path();
if !config_path.exists() { if !config_path.exists() {
@@ -414,14 +393,9 @@ pub fn remove_server_from_codex(id: &str) -> Result<(), AppError> {
let content = let content =
std::fs::read_to_string(&config_path).map_err(|e| AppError::io(&config_path, e))?; std::fs::read_to_string(&config_path).map_err(|e| AppError::io(&config_path, e))?;
// 尝试解析现有配置,如果失败则直接返回(无法删除不存在的内容) let mut doc = content
let mut doc = match content.parse::<toml_edit::DocumentMut>() { .parse::<toml_edit::DocumentMut>()
Ok(doc) => doc, .map_err(|e| AppError::McpValidation(format!("解析 Codex config.toml 失败: {e}")))?;
Err(e) => {
log::warn!("解析 Codex config.toml 失败: {e},跳过删除操作");
return Ok(());
}
};
// 从正确的位置删除:[mcp_servers] // 从正确的位置删除:[mcp_servers]
if let Some(mcp_servers) = doc.get_mut("mcp_servers").and_then(|s| s.as_table_mut()) { if let Some(mcp_servers) = doc.get_mut("mcp_servers").and_then(|s| s.as_table_mut()) {
@@ -438,8 +412,7 @@ pub fn remove_server_from_codex(id: &str) -> Result<(), AppError> {
} }
// 写回文件 // 写回文件
let new_text = doc.to_string(); std::fs::write(&config_path, doc.to_string()).map_err(|e| AppError::io(&config_path, e))?;
crate::config::write_text_file(&config_path, &new_text)?;
Ok(()) Ok(())
} }
-15
View File
@@ -8,12 +8,6 @@ use crate::error::AppError;
use super::validation::{extract_server_spec, validate_server_spec}; use super::validation::{extract_server_spec, validate_server_spec};
fn should_sync_gemini_mcp() -> bool {
// Gemini 未安装/未初始化时:~/.gemini 目录不存在。
// 按用户偏好:目录缺失时跳过写入/删除,不创建任何文件或目录。
crate::gemini_config::get_gemini_dir().exists()
}
/// 返回已启用的 MCP 服务器(过滤 enabled==true /// 返回已启用的 MCP 服务器(过滤 enabled==true
fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> { fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> {
let mut out = HashMap::new(); let mut out = HashMap::new();
@@ -39,9 +33,6 @@ fn collect_enabled_servers(cfg: &McpConfig) -> HashMap<String, Value> {
/// 将 config.json 中 Gemini 的 enabled==true 项写入 Gemini MCP 配置 /// 将 config.json 中 Gemini 的 enabled==true 项写入 Gemini MCP 配置
pub fn sync_enabled_to_gemini(config: &MultiAppConfig) -> Result<(), AppError> { pub fn sync_enabled_to_gemini(config: &MultiAppConfig) -> Result<(), AppError> {
if !should_sync_gemini_mcp() {
return Ok(());
}
let enabled = collect_enabled_servers(&config.mcp.gemini); let enabled = collect_enabled_servers(&config.mcp.gemini);
crate::gemini_mcp::set_mcp_servers_map(&enabled) crate::gemini_mcp::set_mcp_servers_map(&enabled)
} }
@@ -112,9 +103,6 @@ pub fn sync_single_server_to_gemini(
id: &str, id: &str,
server_spec: &Value, server_spec: &Value,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
if !should_sync_gemini_mcp() {
return Ok(());
}
// 读取现有的 MCP 配置 // 读取现有的 MCP 配置
let mut current = crate::gemini_mcp::read_mcp_servers_map()?; let mut current = crate::gemini_mcp::read_mcp_servers_map()?;
@@ -127,9 +115,6 @@ pub fn sync_single_server_to_gemini(
/// 从 Gemini live 配置中移除单个 MCP 服务器 /// 从 Gemini live 配置中移除单个 MCP 服务器
pub fn remove_server_from_gemini(id: &str) -> Result<(), AppError> { pub fn remove_server_from_gemini(id: &str) -> Result<(), AppError> {
if !should_sync_gemini_mcp() {
return Ok(());
}
// 读取现有的 MCP 配置 // 读取现有的 MCP 配置
let mut current = crate::gemini_mcp::read_mcp_servers_map()?; let mut current = crate::gemini_mcp::read_mcp_servers_map()?;
-287
View File
@@ -36,10 +36,6 @@ pub struct Provider {
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "iconColor")] #[serde(rename = "iconColor")]
pub icon_color: Option<String>, pub icon_color: Option<String>,
/// 是否加入故障转移队列
#[serde(default)]
#[serde(rename = "inFailoverQueue")]
pub in_failover_queue: bool,
} }
impl Provider { impl Provider {
@@ -62,7 +58,6 @@ impl Provider {
meta: None, meta: None,
icon: None, icon: None,
icon_color: None, icon_color: None,
in_failover_queue: false,
} }
} }
} }
@@ -173,285 +168,3 @@ impl ProviderManager {
&self.providers &self.providers
} }
} }
// ============================================================================
// 统一供应商(Universal Provider- 跨应用共享配置
// ============================================================================
/// 统一供应商的应用启用状态
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct UniversalProviderApps {
#[serde(default)]
pub claude: bool,
#[serde(default)]
pub codex: bool,
#[serde(default)]
pub gemini: bool,
}
/// Claude 模型配置
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ClaudeModelConfig {
/// 主模型
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
/// Haiku 默认模型
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "haikuModel")]
pub haiku_model: Option<String>,
/// Sonnet 默认模型
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "sonnetModel")]
pub sonnet_model: Option<String>,
/// Opus 默认模型
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "opusModel")]
pub opus_model: Option<String>,
}
/// Codex 模型配置
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct CodexModelConfig {
/// 模型名称
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
/// 推理强度
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "reasoningEffort")]
pub reasoning_effort: Option<String>,
}
/// Gemini 模型配置
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct GeminiModelConfig {
/// 模型名称
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
}
/// 各应用的模型配置
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct UniversalProviderModels {
#[serde(skip_serializing_if = "Option::is_none")]
pub claude: Option<ClaudeModelConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub codex: Option<CodexModelConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub gemini: Option<GeminiModelConfig>,
}
/// 统一供应商(跨应用共享配置)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UniversalProvider {
/// 唯一标识
pub id: String,
/// 供应商名称
pub name: String,
/// 供应商类型(如 "newapi", "custom"
#[serde(rename = "providerType")]
pub provider_type: String,
/// 应用启用状态
pub apps: UniversalProviderApps,
/// API 基础地址
#[serde(rename = "baseUrl")]
pub base_url: String,
/// API 密钥
#[serde(rename = "apiKey")]
pub api_key: String,
/// 各应用的模型配置
#[serde(default)]
pub models: UniversalProviderModels,
/// 网站链接
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "websiteUrl")]
pub website_url: Option<String>,
/// 备注信息
#[serde(skip_serializing_if = "Option::is_none")]
pub notes: Option<String>,
/// 图标名称
#[serde(skip_serializing_if = "Option::is_none")]
pub icon: Option<String>,
/// 图标颜色
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "iconColor")]
pub icon_color: Option<String>,
/// 元数据
#[serde(skip_serializing_if = "Option::is_none")]
pub meta: Option<ProviderMeta>,
/// 创建时间戳
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "createdAt")]
pub created_at: Option<i64>,
/// 排序索引
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "sortIndex")]
pub sort_index: Option<usize>,
}
impl UniversalProvider {
/// 创建新的统一供应商
pub fn new(
id: String,
name: String,
provider_type: String,
base_url: String,
api_key: String,
) -> Self {
Self {
id,
name,
provider_type,
apps: UniversalProviderApps::default(),
base_url,
api_key,
models: UniversalProviderModels::default(),
website_url: None,
notes: None,
icon: None,
icon_color: None,
meta: None,
created_at: Some(chrono::Utc::now().timestamp_millis()),
sort_index: None,
}
}
/// 生成 Claude 供应商配置
pub fn to_claude_provider(&self) -> Option<Provider> {
if !self.apps.claude {
return None;
}
let models = self.models.claude.as_ref();
let model = models
.and_then(|m| m.model.clone())
.unwrap_or_else(|| "claude-sonnet-4-20250514".to_string());
let haiku = models
.and_then(|m| m.haiku_model.clone())
.unwrap_or_else(|| model.clone());
let sonnet = models
.and_then(|m| m.sonnet_model.clone())
.unwrap_or_else(|| model.clone());
let opus = models
.and_then(|m| m.opus_model.clone())
.unwrap_or_else(|| model.clone());
let settings_config = serde_json::json!({
"env": {
"ANTHROPIC_BASE_URL": self.base_url,
"ANTHROPIC_AUTH_TOKEN": self.api_key,
"ANTHROPIC_MODEL": model,
"ANTHROPIC_DEFAULT_HAIKU_MODEL": haiku,
"ANTHROPIC_DEFAULT_SONNET_MODEL": sonnet,
"ANTHROPIC_DEFAULT_OPUS_MODEL": opus,
}
});
Some(Provider {
id: format!("universal-claude-{}", self.id),
name: self.name.clone(),
settings_config,
website_url: self.website_url.clone(),
category: Some("aggregator".to_string()),
created_at: self.created_at,
sort_index: self.sort_index,
notes: self.notes.clone(),
meta: self.meta.clone(),
icon: self.icon.clone(),
icon_color: self.icon_color.clone(),
in_failover_queue: false,
})
}
/// 生成 Codex 供应商配置
pub fn to_codex_provider(&self) -> Option<Provider> {
if !self.apps.codex {
return None;
}
let models = self.models.codex.as_ref();
let model = models
.and_then(|m| m.model.clone())
.unwrap_or_else(|| "gpt-4o".to_string());
let reasoning_effort = models
.and_then(|m| m.reasoning_effort.clone())
.unwrap_or_else(|| "high".to_string());
// 确保 base_url 以 /v1 结尾(Codex 使用 OpenAI 兼容 API
let codex_base_url = if self.base_url.ends_with("/v1") {
self.base_url.clone()
} else {
format!("{}/v1", self.base_url.trim_end_matches('/'))
};
// 生成 Codex 的 config.toml 内容
let config_toml = format!(
r#"model_provider = "newapi"
model = "{model}"
model_reasoning_effort = "{reasoning_effort}"
disable_response_storage = true
[model_providers.newapi]
name = "NewAPI"
base_url = "{codex_base_url}"
wire_api = "responses"
requires_openai_auth = true"#
);
let settings_config = serde_json::json!({
"auth": {
"OPENAI_API_KEY": self.api_key
},
"config": config_toml
});
Some(Provider {
id: format!("universal-codex-{}", self.id),
name: self.name.clone(),
settings_config,
website_url: self.website_url.clone(),
category: Some("aggregator".to_string()),
created_at: self.created_at,
sort_index: self.sort_index,
notes: self.notes.clone(),
meta: self.meta.clone(),
icon: self.icon.clone(),
icon_color: self.icon_color.clone(),
in_failover_queue: false,
})
}
/// 生成 Gemini 供应商配置
pub fn to_gemini_provider(&self) -> Option<Provider> {
if !self.apps.gemini {
return None;
}
let models = self.models.gemini.as_ref();
let model = models
.and_then(|m| m.model.clone())
.unwrap_or_else(|| "gemini-2.5-pro".to_string());
let settings_config = serde_json::json!({
"env": {
"GOOGLE_GEMINI_BASE_URL": self.base_url,
"GEMINI_API_KEY": self.api_key,
"GEMINI_MODEL": model,
}
});
Some(Provider {
id: format!("universal-gemini-{}", self.id),
name: self.name.clone(),
settings_config,
website_url: self.website_url.clone(),
category: Some("aggregator".to_string()),
created_at: self.created_at,
sort_index: self.sort_index,
notes: self.notes.clone(),
meta: self.meta.clone(),
icon: self.icon.clone(),
icon_color: self.icon_color.clone(),
in_failover_queue: false,
})
}
}
-303
View File
@@ -1,303 +0,0 @@
//! 请求体过滤模块
//!
//! 过滤不应透传到上游的私有参数,防止内部信息泄露。
//!
//! ## 过滤规则
//! - 以 `_` 开头的字段被视为私有参数,会被递归过滤
//! - 支持白名单机制,允许透传特定的 `_` 前缀字段
//! - 支持嵌套对象和数组的深度过滤
//!
//! ## 使用场景
//! - `_internal_id`: 内部追踪 ID
//! - `_debug_mode`: 调试标记
//! - `_session_token`: 会话令牌
//! - `_client_version`: 客户端版本
use serde_json::Value;
use std::collections::HashSet;
/// 过滤私有参数(以 `_` 开头的字段)
///
/// 递归遍历 JSON 结构,移除所有以下划线开头的字段。
///
/// # Arguments
/// * `body` - 原始请求体
///
/// # Returns
/// 过滤后的请求体
///
/// # Example
/// ```ignore
/// let input = json!({
/// "model": "claude-3",
/// "_internal_id": "abc123",
/// "messages": [{"role": "user", "content": "hello", "_token": "secret"}]
/// });
/// let output = filter_private_params(input);
/// // output 中不包含 _internal_id 和 _token
/// ```
#[cfg(test)]
pub fn filter_private_params(body: Value) -> Value {
filter_private_params_with_whitelist(body, &[])
}
/// 过滤私有参数(支持白名单)
///
/// 递归遍历 JSON 结构,移除所有以下划线开头的字段,
/// 但保留白名单中指定的字段。
///
/// # Arguments
/// * `body` - 原始请求体
/// * `whitelist` - 白名单字段列表(不过滤这些字段)
///
/// # Returns
/// 过滤后的请求体
///
/// # Example
/// ```ignore
/// let input = json!({
/// "model": "claude-3",
/// "_metadata": {"key": "value"}, // 白名单中,保留
/// "_internal_id": "abc123" // 不在白名单中,过滤
/// });
/// let output = filter_private_params_with_whitelist(input, &["_metadata"]);
/// // output 包含 _metadata,不包含 _internal_id
/// ```
pub fn filter_private_params_with_whitelist(body: Value, whitelist: &[String]) -> Value {
let whitelist_set: HashSet<&str> = whitelist.iter().map(|s| s.as_str()).collect();
filter_recursive_with_whitelist(body, &mut Vec::new(), &whitelist_set)
}
/// 递归过滤实现
#[cfg(test)]
fn filter_recursive(value: Value, removed_keys: &mut Vec<String>) -> Value {
filter_recursive_with_whitelist(value, removed_keys, &HashSet::new())
}
/// 递归过滤实现(支持白名单)
fn filter_recursive_with_whitelist(
value: Value,
removed_keys: &mut Vec<String>,
whitelist: &HashSet<&str>,
) -> Value {
match value {
Value::Object(map) => {
let filtered: serde_json::Map<String, Value> = map
.into_iter()
.filter_map(|(key, val)| {
// 以 _ 开头且不在白名单中的字段被过滤
if key.starts_with('_') && !whitelist.contains(key.as_str()) {
removed_keys.push(key);
None
} else {
Some((
key,
filter_recursive_with_whitelist(val, removed_keys, whitelist),
))
}
})
.collect();
// 仅在有过滤时记录日志(避免每次请求都打印)
if !removed_keys.is_empty() {
log::debug!("[BodyFilter] 过滤私有参数: {removed_keys:?}");
removed_keys.clear();
}
Value::Object(filtered)
}
Value::Array(arr) => Value::Array(
arr.into_iter()
.map(|v| filter_recursive_with_whitelist(v, removed_keys, whitelist))
.collect(),
),
other => other,
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_filter_top_level_private_params() {
let input = json!({
"model": "claude-3",
"_internal_id": "abc123",
"_debug": true,
"max_tokens": 1024
});
let output = filter_private_params(input);
assert!(output.get("model").is_some());
assert!(output.get("max_tokens").is_some());
assert!(output.get("_internal_id").is_none());
assert!(output.get("_debug").is_none());
}
#[test]
fn test_filter_nested_private_params() {
let input = json!({
"model": "claude-3",
"messages": [
{
"role": "user",
"content": "hello",
"_session_token": "secret"
}
],
"metadata": {
"user_id": "user-1",
"_tracking_id": "track-1"
}
});
let output = filter_private_params(input);
// 顶级字段保留
assert!(output.get("model").is_some());
assert!(output.get("messages").is_some());
assert!(output.get("metadata").is_some());
// messages 数组中的私有参数被过滤
let messages = output.get("messages").unwrap().as_array().unwrap();
assert!(messages[0].get("role").is_some());
assert!(messages[0].get("content").is_some());
assert!(messages[0].get("_session_token").is_none());
// metadata 对象中的私有参数被过滤
let metadata = output.get("metadata").unwrap();
assert!(metadata.get("user_id").is_some());
assert!(metadata.get("_tracking_id").is_none());
}
#[test]
fn test_filter_deeply_nested() {
let input = json!({
"level1": {
"level2": {
"level3": {
"keep": "value",
"_remove": "secret"
}
}
}
});
let output = filter_private_params(input);
let level3 = output
.get("level1")
.unwrap()
.get("level2")
.unwrap()
.get("level3")
.unwrap();
assert!(level3.get("keep").is_some());
assert!(level3.get("_remove").is_none());
}
#[test]
fn test_filter_array_of_objects() {
let input = json!({
"items": [
{"id": 1, "_secret": "a"},
{"id": 2, "_secret": "b"},
{"id": 3, "_secret": "c"}
]
});
let output = filter_private_params(input);
let items = output.get("items").unwrap().as_array().unwrap();
for item in items {
assert!(item.get("id").is_some());
assert!(item.get("_secret").is_none());
}
}
#[test]
fn test_no_private_params() {
let input = json!({
"model": "claude-3",
"messages": [{"role": "user", "content": "hello"}]
});
let output = filter_private_params(input.clone());
// 无私有参数时,输出应与输入相同
assert_eq!(input, output);
}
#[test]
fn test_empty_object() {
let input = json!({});
let output = filter_private_params(input);
assert_eq!(output, json!({}));
}
#[test]
fn test_primitive_values() {
// 原始值不应被修改
assert_eq!(filter_private_params(json!(42)), json!(42));
assert_eq!(filter_private_params(json!("string")), json!("string"));
assert_eq!(filter_private_params(json!(true)), json!(true));
assert_eq!(filter_private_params(json!(null)), json!(null));
}
#[test]
fn test_whitelist_preserves_private_params() {
let input = json!({
"model": "claude-3",
"_metadata": {"key": "value"},
"_internal_id": "abc123",
"_stream_options": {"include_usage": true}
});
let whitelist = vec!["_metadata".to_string(), "_stream_options".to_string()];
let output = filter_private_params_with_whitelist(input, &whitelist);
// 白名单中的字段保留
assert!(output.get("_metadata").is_some());
assert!(output.get("_stream_options").is_some());
// 不在白名单中的私有字段被过滤
assert!(output.get("_internal_id").is_none());
// 普通字段保留
assert!(output.get("model").is_some());
}
#[test]
fn test_whitelist_nested() {
let input = json!({
"data": {
"_allowed": "keep",
"_forbidden": "remove",
"normal": "value"
}
});
let whitelist = vec!["_allowed".to_string()];
let output = filter_private_params_with_whitelist(input, &whitelist);
let data = output.get("data").unwrap();
assert!(data.get("_allowed").is_some());
assert!(data.get("_forbidden").is_none());
assert!(data.get("normal").is_some());
}
#[test]
fn test_empty_whitelist_same_as_default() {
let input = json!({
"model": "claude-3",
"_internal_id": "abc123"
});
let output1 = filter_private_params(input.clone());
let output2 = filter_private_params_with_whitelist(input, &[]);
assert_eq!(output1, output2);
}
}
+46 -134
View File
@@ -78,16 +78,6 @@ pub struct CircuitBreaker {
half_open_requests: Arc<AtomicU32>, half_open_requests: Arc<AtomicU32>,
} }
/// 熔断器放行结果
///
/// `used_half_open_permit` 表示本次放行是否占用了 HalfOpen 探测名额。
/// 调用方应在请求结束后把该值传回 `record_success` / `record_failure` 用于正确释放名额。
#[derive(Debug, Clone, Copy)]
pub struct AllowResult {
pub allowed: bool,
pub used_half_open_permit: bool,
}
impl CircuitBreaker { impl CircuitBreaker {
/// 创建新的熔断器 /// 创建新的熔断器
pub fn new(config: CircuitBreakerConfig) -> Self { pub fn new(config: CircuitBreakerConfig) -> Self {
@@ -140,16 +130,13 @@ impl CircuitBreaker {
} }
/// 检查是否允许请求通过 /// 检查是否允许请求通过
pub async fn allow_request(&self) -> AllowResult { pub async fn allow_request(&self) -> bool {
let state = *self.state.read().await; let state = *self.state.read().await;
let config = self.config.read().await;
match state { match state {
CircuitState::Closed => AllowResult { CircuitState::Closed => true,
allowed: true,
used_half_open_permit: false,
},
CircuitState::Open => { CircuitState::Open => {
let config = self.config.read().await;
// 检查是否应该尝试半开 // 检查是否应该尝试半开
if let Some(opened_at) = *self.last_opened_at.read().await { if let Some(opened_at) = *self.last_opened_at.read().await {
if opened_at.elapsed().as_secs() >= config.timeout_seconds { if opened_at.elapsed().as_secs() >= config.timeout_seconds {
@@ -158,47 +145,52 @@ impl CircuitBreaker {
"Circuit breaker transitioning from Open to HalfOpen (timeout reached)" "Circuit breaker transitioning from Open to HalfOpen (timeout reached)"
); );
self.transition_to_half_open().await; self.transition_to_half_open().await;
// 增加计数,确保 record_success/record_failure 减计数时不会下溢
// 转换后按当前状态决定是否需要获取 HalfOpen 探测名额 self.half_open_requests.fetch_add(1, Ordering::SeqCst);
let current_state = *self.state.read().await; return true;
return match current_state {
CircuitState::Closed => AllowResult {
allowed: true,
used_half_open_permit: false,
},
CircuitState::HalfOpen => self.allow_half_open_probe(),
CircuitState::Open => AllowResult {
allowed: false,
used_half_open_permit: false,
},
};
} }
} }
false
}
CircuitState::HalfOpen => {
// 半开状态限流:只允许有限请求通过进行探测
// 默认最多允许 1 个请求(可在配置中扩展)
let max_half_open_requests = 1u32;
let current = self.half_open_requests.fetch_add(1, Ordering::SeqCst);
AllowResult { if current < max_half_open_requests {
allowed: false, log::debug!(
used_half_open_permit: false, "Circuit breaker HalfOpen: allowing probe request ({}/{})",
current + 1,
max_half_open_requests
);
true
} else {
// 超过限额,回退计数,拒绝请求
self.half_open_requests.fetch_sub(1, Ordering::SeqCst);
log::debug!(
"Circuit breaker HalfOpen: rejecting request (limit reached: {max_half_open_requests})"
);
false
} }
} }
CircuitState::HalfOpen => self.allow_half_open_probe(),
} }
} }
/// 记录成功 /// 记录成功
pub async fn record_success(&self, used_half_open_permit: bool) { pub async fn record_success(&self) {
let state = *self.state.read().await; let state = *self.state.read().await;
let config = self.config.read().await; let config = self.config.read().await;
if used_half_open_permit {
self.release_half_open_permit();
}
// 重置失败计数 // 重置失败计数
self.consecutive_failures.store(0, Ordering::SeqCst); self.consecutive_failures.store(0, Ordering::SeqCst);
self.total_requests.fetch_add(1, Ordering::SeqCst); self.total_requests.fetch_add(1, Ordering::SeqCst);
match state { match state {
CircuitState::HalfOpen => { CircuitState::HalfOpen => {
// 释放 in-flight 名额(探测请求结束)
self.half_open_requests.fetch_sub(1, Ordering::SeqCst);
let successes = self.consecutive_successes.fetch_add(1, Ordering::SeqCst) + 1; let successes = self.consecutive_successes.fetch_add(1, Ordering::SeqCst) + 1;
log::debug!( log::debug!(
"Circuit breaker HalfOpen: {} consecutive successes (threshold: {})", "Circuit breaker HalfOpen: {} consecutive successes (threshold: {})",
@@ -220,14 +212,10 @@ impl CircuitBreaker {
} }
/// 记录失败 /// 记录失败
pub async fn record_failure(&self, used_half_open_permit: bool) { pub async fn record_failure(&self) {
let state = *self.state.read().await; let state = *self.state.read().await;
let config = self.config.read().await; let config = self.config.read().await;
if used_half_open_permit {
self.release_half_open_permit();
}
// 更新计数器 // 更新计数器
let failures = self.consecutive_failures.fetch_add(1, Ordering::SeqCst) + 1; let failures = self.consecutive_failures.fetch_add(1, Ordering::SeqCst) + 1;
self.total_requests.fetch_add(1, Ordering::SeqCst); self.total_requests.fetch_add(1, Ordering::SeqCst);
@@ -246,6 +234,9 @@ impl CircuitBreaker {
// 检查是否应该打开熔断器 // 检查是否应该打开熔断器
match state { match state {
CircuitState::HalfOpen => { CircuitState::HalfOpen => {
// 释放 in-flight 名额(探测请求结束)
self.half_open_requests.fetch_sub(1, Ordering::SeqCst);
// HalfOpen 状态下失败,立即转为 Open // HalfOpen 状态下失败,立即转为 Open
log::warn!("Circuit breaker HalfOpen probe failed, transitioning to Open"); log::warn!("Circuit breaker HalfOpen probe failed, transitioning to Open");
drop(config); drop(config);
@@ -316,56 +307,6 @@ impl CircuitBreaker {
self.transition_to_closed().await; self.transition_to_closed().await;
} }
fn allow_half_open_probe(&self) -> AllowResult {
// 半开状态限流:只允许有限请求通过进行探测
// 默认最多允许 1 个请求(可在配置中扩展)
let max_half_open_requests = 1u32;
let current = self.half_open_requests.fetch_add(1, Ordering::SeqCst);
if current < max_half_open_requests {
log::debug!(
"Circuit breaker HalfOpen: allowing probe request ({}/{})",
current + 1,
max_half_open_requests
);
AllowResult {
allowed: true,
used_half_open_permit: true,
}
} else {
// 超过限额,回退计数,拒绝请求
self.half_open_requests.fetch_sub(1, Ordering::SeqCst);
log::debug!(
"Circuit breaker HalfOpen: rejecting request (limit reached: {max_half_open_requests})"
);
AllowResult {
allowed: false,
used_half_open_permit: false,
}
}
}
fn release_half_open_permit(&self) {
let mut current = self.half_open_requests.load(Ordering::SeqCst);
loop {
if current == 0 {
// 理论上不应该发生:说明调用方传入的 used_half_open_permit 与实际占用不一致
log::debug!("Circuit breaker HalfOpen permit already released (counter=0)");
return;
}
match self.half_open_requests.compare_exchange(
current,
current - 1,
Ordering::SeqCst,
Ordering::SeqCst,
) {
Ok(_) => return,
Err(actual) => current = actual,
}
}
}
/// 转换到打开状态 /// 转换到打开状态
async fn transition_to_open(&self) { async fn transition_to_open(&self) {
*self.state.write().await = CircuitState::Open; *self.state.write().await = CircuitState::Open;
@@ -376,12 +317,7 @@ impl CircuitBreaker {
/// 转换到半开状态 /// 转换到半开状态
async fn transition_to_half_open(&self) { async fn transition_to_half_open(&self) {
let mut state = self.state.write().await; *self.state.write().await = CircuitState::HalfOpen;
if *state != CircuitState::Open {
return;
}
*state = CircuitState::HalfOpen;
self.consecutive_successes.store(0, Ordering::SeqCst); self.consecutive_successes.store(0, Ordering::SeqCst);
// 重置半开状态的请求限流计数 // 重置半开状态的请求限流计数
self.half_open_requests.store(0, Ordering::SeqCst); self.half_open_requests.store(0, Ordering::SeqCst);
@@ -423,16 +359,16 @@ mod tests {
// 初始状态应该是关闭 // 初始状态应该是关闭
assert_eq!(breaker.get_state().await, CircuitState::Closed); assert_eq!(breaker.get_state().await, CircuitState::Closed);
assert!(breaker.allow_request().await.allowed); assert!(breaker.allow_request().await);
// 记录 3 次失败 // 记录 3 次失败
for _ in 0..3 { for _ in 0..3 {
breaker.record_failure(false).await; breaker.record_failure().await;
} }
// 应该转换到打开状态 // 应该转换到打开状态
assert_eq!(breaker.get_state().await, CircuitState::Open); assert_eq!(breaker.get_state().await, CircuitState::Open);
assert!(!breaker.allow_request().await.allowed); assert!(!breaker.allow_request().await);
} }
#[tokio::test] #[tokio::test]
@@ -445,8 +381,8 @@ mod tests {
let breaker = CircuitBreaker::new(config); let breaker = CircuitBreaker::new(config);
// 打开熔断器 // 打开熔断器
breaker.record_failure(false).await; breaker.record_failure().await;
breaker.record_failure(false).await; breaker.record_failure().await;
assert_eq!(breaker.get_state().await, CircuitState::Open); assert_eq!(breaker.get_state().await, CircuitState::Open);
// 手动转换到半开状态 // 手动转换到半开状态
@@ -454,37 +390,13 @@ mod tests {
assert_eq!(breaker.get_state().await, CircuitState::HalfOpen); assert_eq!(breaker.get_state().await, CircuitState::HalfOpen);
// 记录 2 次成功 // 记录 2 次成功
breaker.record_success(false).await; breaker.record_success().await;
breaker.record_success(false).await; breaker.record_success().await;
// 应该转换到关闭状态 // 应该转换到关闭状态
assert_eq!(breaker.get_state().await, CircuitState::Closed); assert_eq!(breaker.get_state().await, CircuitState::Closed);
} }
#[tokio::test]
async fn test_half_open_transition_does_not_reset_inflight_permit() {
let config = CircuitBreakerConfig {
timeout_seconds: 0,
..Default::default()
};
let breaker = CircuitBreaker::new(config);
// 进入 Open,然后由于 timeout_seconds=0allow_request 会立即切换到 HalfOpen 并占用探测名额
breaker.transition_to_open().await;
let first = breaker.allow_request().await;
assert!(first.allowed);
assert!(first.used_half_open_permit);
assert_eq!(breaker.get_state().await, CircuitState::HalfOpen);
// 模拟并发下的“重复 HalfOpen 转换调用”,不应重置 in-flight 计数
breaker.transition_to_half_open().await;
// 由于名额仍被占用,第二次请求应被拒绝
let second = breaker.allow_request().await;
assert!(!second.allowed);
assert!(!second.used_half_open_permit);
}
#[tokio::test] #[tokio::test]
async fn test_circuit_breaker_reset() { async fn test_circuit_breaker_reset() {
let config = CircuitBreakerConfig { let config = CircuitBreakerConfig {
@@ -494,13 +406,13 @@ mod tests {
let breaker = CircuitBreaker::new(config); let breaker = CircuitBreaker::new(config);
// 打开熔断器 // 打开熔断器
breaker.record_failure(false).await; breaker.record_failure().await;
breaker.record_failure(false).await; breaker.record_failure().await;
assert_eq!(breaker.get_state().await, CircuitState::Open); assert_eq!(breaker.get_state().await, CircuitState::Open);
// 重置 // 重置
breaker.reset().await; breaker.reset().await;
assert_eq!(breaker.get_state().await, CircuitState::Closed); assert_eq!(breaker.get_state().await, CircuitState::Closed);
assert!(breaker.allow_request().await.allowed); assert!(breaker.allow_request().await);
} }
} }
-12
View File
@@ -23,12 +23,6 @@ pub enum ProxyError {
#[error("无可用的Provider")] #[error("无可用的Provider")]
NoAvailableProvider, NoAvailableProvider,
#[error("所有供应商已熔断,无可用渠道")]
AllProvidersCircuitOpen,
#[error("未配置供应商")]
NoProvidersConfigured,
#[allow(dead_code)] #[allow(dead_code)]
#[error("Provider不健康: {0}")] #[error("Provider不健康: {0}")]
ProviderUnhealthy(String), ProviderUnhealthy(String),
@@ -117,12 +111,6 @@ impl IntoResponse for ProxyError {
ProxyError::NoAvailableProvider => { ProxyError::NoAvailableProvider => {
(StatusCode::SERVICE_UNAVAILABLE, self.to_string()) (StatusCode::SERVICE_UNAVAILABLE, self.to_string())
} }
ProxyError::AllProvidersCircuitOpen => {
(StatusCode::SERVICE_UNAVAILABLE, self.to_string())
}
ProxyError::NoProvidersConfigured => {
(StatusCode::SERVICE_UNAVAILABLE, self.to_string())
}
ProxyError::ProviderUnhealthy(_) => { ProxyError::ProviderUnhealthy(_) => {
(StatusCode::SERVICE_UNAVAILABLE, self.to_string()) (StatusCode::SERVICE_UNAVAILABLE, self.to_string())
} }
-8
View File
@@ -27,12 +27,6 @@ pub fn map_proxy_error_to_status(error: &ProxyError) -> u16 {
// 无可用 Provider503 Service Unavailable // 无可用 Provider503 Service Unavailable
ProxyError::NoAvailableProvider => 503, ProxyError::NoAvailableProvider => 503,
// 所有供应商已熔断:503 Service Unavailable
ProxyError::AllProvidersCircuitOpen => 503,
// 未配置供应商:503 Service Unavailable
ProxyError::NoProvidersConfigured => 503,
// 重试耗尽:503 Service Unavailable // 重试耗尽:503 Service Unavailable
ProxyError::MaxRetriesExceeded => 503, ProxyError::MaxRetriesExceeded => 503,
@@ -63,8 +57,6 @@ pub fn get_error_message(error: &ProxyError) -> String {
ProxyError::Timeout(msg) => format!("请求超时: {msg}"), ProxyError::Timeout(msg) => format!("请求超时: {msg}"),
ProxyError::ForwardFailed(msg) => format!("转发失败: {msg}"), ProxyError::ForwardFailed(msg) => format!("转发失败: {msg}"),
ProxyError::NoAvailableProvider => "无可用 Provider".to_string(), ProxyError::NoAvailableProvider => "无可用 Provider".to_string(),
ProxyError::AllProvidersCircuitOpen => "所有供应商已熔断,无可用渠道".to_string(),
ProxyError::NoProvidersConfigured => "未配置供应商".to_string(),
ProxyError::MaxRetriesExceeded => "所有 Provider 都失败,重试耗尽".to_string(), ProxyError::MaxRetriesExceeded => "所有 Provider 都失败,重试耗尽".to_string(),
ProxyError::ProviderUnhealthy(msg) => format!("Provider 不健康: {msg}"), ProxyError::ProviderUnhealthy(msg) => format!("Provider 不健康: {msg}"),
ProxyError::DatabaseError(msg) => format!("数据库错误: {msg}"), ProxyError::DatabaseError(msg) => format!("数据库错误: {msg}"),
-15
View File
@@ -81,21 +81,6 @@ impl FailoverSwitchManager {
provider_id: &str, provider_id: &str,
provider_name: &str, provider_name: &str,
) -> Result<bool, AppError> { ) -> Result<bool, AppError> {
// 检查该应用是否已被代理接管(enabled=true
// 只有被接管的应用才允许执行故障转移切换
let app_enabled = match self.db.get_proxy_config_for_app(app_type).await {
Ok(config) => config.enabled,
Err(e) => {
log::warn!("[Failover] 无法读取 {app_type} 配置: {e},跳过切换");
return Ok(false);
}
};
if !app_enabled {
log::info!("[Failover] {app_type} 未被代理接管(enabled=false),跳过切换");
return Ok(false);
}
log::info!("[Failover] 开始切换供应商: {app_type} -> {provider_name} ({provider_id})"); log::info!("[Failover] 开始切换供应商: {app_type} -> {provider_name} ({provider_id})");
// 1. 更新数据库 is_current // 1. 更新数据库 is_current
+114 -303
View File
@@ -1,9 +1,8 @@
//! 请求转发器 //! 请求转发器
//! //!
//! 负责将请求转发到上游Provider,支持故障转移 //! 负责将请求转发到上游Provider,支持重试和故障转移
use super::{ use super::{
body_filter::filter_private_params_with_whitelist,
error::*, error::*,
failover_switch::FailoverSwitchManager, failover_switch::FailoverSwitchManager,
provider_router::ProviderRouter, provider_router::ProviderRouter,
@@ -18,119 +17,33 @@ use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use tokio::sync::RwLock; use tokio::sync::RwLock;
/// Headers 黑名单 - 不透传到上游的 Headers
///
/// 参考 Claude Code Hub 设计,过滤以下类别:
/// 1. 认证类(会被覆盖)
/// 2. 连接类(由 HTTP 客户端管理)
/// 3. 代理转发类
/// 4. CDN/云服务商特定头
/// 5. 请求追踪类
/// 6. 浏览器特定头(可能被上游检测)
///
/// 注意:客户端 IP 类(x-forwarded-for, x-real-ip)默认透传
const HEADER_BLACKLIST: &[&str] = &[
// 认证类(会被覆盖)
"authorization",
"x-api-key",
// 连接类
"host",
"content-length",
"connection",
"transfer-encoding",
// 编码类(会被覆盖为 identity)
"accept-encoding",
// 代理转发类(保留 x-forwarded-for 和 x-real-ip
"x-forwarded-host",
"x-forwarded-port",
"x-forwarded-proto",
"forwarded",
// CDN/云服务商特定头
"cf-connecting-ip",
"cf-ipcountry",
"cf-ray",
"cf-visitor",
"true-client-ip",
"fastly-client-ip",
"x-azure-clientip",
"x-azure-fdid",
"x-azure-ref",
"akamai-origin-hop",
"x-akamai-config-log-detail",
// 请求追踪类
"x-request-id",
"x-correlation-id",
"x-trace-id",
"x-amzn-trace-id",
"x-b3-traceid",
"x-b3-spanid",
"x-b3-parentspanid",
"x-b3-sampled",
"traceparent",
"tracestate",
// 浏览器特定头(可能被上游检测为非 CLI 请求)
"sec-fetch-mode",
"sec-fetch-site",
"sec-fetch-dest",
"sec-ch-ua",
"sec-ch-ua-mobile",
"sec-ch-ua-platform",
"accept-language",
// anthropic-beta 单独处理,避免重复
"anthropic-beta",
// 客户端 IP 单独处理(默认透传)
"x-forwarded-for",
"x-real-ip",
];
pub struct ForwardResult {
pub response: Response,
pub provider: Provider,
}
pub struct ForwardError {
pub error: ProxyError,
pub provider: Option<Provider>,
}
pub struct RequestForwarder { pub struct RequestForwarder {
client: Client, client: Client,
/// 共享的 ProviderRouter(持有熔断器状态) /// 共享的 ProviderRouter(持有熔断器状态)
router: Arc<ProviderRouter>, router: Arc<ProviderRouter>,
/// 单个 Provider 内的最大重试次数
max_retries: u8,
status: Arc<RwLock<ProxyStatus>>, status: Arc<RwLock<ProxyStatus>>,
current_providers: Arc<RwLock<std::collections::HashMap<String, (String, String)>>>, current_providers: Arc<RwLock<std::collections::HashMap<String, (String, String)>>>,
/// 故障转移切换管理器 /// 故障转移切换管理器
failover_manager: Arc<FailoverSwitchManager>, failover_manager: Arc<FailoverSwitchManager>,
/// AppHandle,用于发射事件和更新托盘 /// AppHandle,用于发射事件和更新托盘
app_handle: Option<tauri::AppHandle>, app_handle: Option<tauri::AppHandle>,
/// 请求开始时的"当前供应商 ID"(用于判断是否需要同步 UI/托盘)
current_provider_id_at_start: String,
} }
impl RequestForwarder { impl RequestForwarder {
#[allow(clippy::too_many_arguments)]
pub fn new( pub fn new(
router: Arc<ProviderRouter>, router: Arc<ProviderRouter>,
non_streaming_timeout: u64, timeout_secs: u64,
max_retries: u8,
status: Arc<RwLock<ProxyStatus>>, status: Arc<RwLock<ProxyStatus>>,
current_providers: Arc<RwLock<std::collections::HashMap<String, (String, String)>>>, current_providers: Arc<RwLock<std::collections::HashMap<String, (String, String)>>>,
failover_manager: Arc<FailoverSwitchManager>, failover_manager: Arc<FailoverSwitchManager>,
app_handle: Option<tauri::AppHandle>, app_handle: Option<tauri::AppHandle>,
current_provider_id_at_start: String,
_streaming_first_byte_timeout: u64,
_streaming_idle_timeout: u64,
) -> Self { ) -> Self {
// 全局超时设置为 1800 秒(30 分钟),确保业务层超时配置能正常工作
// 参考 Claude Code Hub 的 undici 全局超时设计
const GLOBAL_TIMEOUT_SECS: u64 = 1800;
let mut client_builder = Client::builder(); let mut client_builder = Client::builder();
if non_streaming_timeout > 0 { if timeout_secs > 0 {
// 使用配置的非流式超时 client_builder = client_builder.timeout(Duration::from_secs(timeout_secs));
client_builder = client_builder.timeout(Duration::from_secs(non_streaming_timeout));
} else {
// 禁用超时时使用全局超时作为保底
client_builder = client_builder.timeout(Duration::from_secs(GLOBAL_TIMEOUT_SECS));
} }
let client = client_builder let client = client_builder
@@ -140,14 +53,69 @@ impl RequestForwarder {
Self { Self {
client, client,
router, router,
max_retries,
status, status,
current_providers, current_providers,
failover_manager, failover_manager,
app_handle, app_handle,
current_provider_id_at_start,
} }
} }
/// 对单个 Provider 执行请求(带重试)
///
/// 在同一个 Provider 上最多重试 max_retries 次,使用指数退避
async fn forward_with_provider_retry(
&self,
provider: &Provider,
endpoint: &str,
body: &Value,
headers: &axum::http::HeaderMap,
adapter: &dyn ProviderAdapter,
) -> Result<Response, ProxyError> {
let mut last_error = None;
for attempt in 0..=self.max_retries {
if attempt > 0 {
// 指数退避:100ms, 200ms, 400ms, ...
let delay_ms = 100 * 2u64.pow(attempt as u32 - 1);
log::info!(
"[{}] 重试第 {}/{} 次(等待 {}ms",
adapter.name(),
attempt,
self.max_retries,
delay_ms
);
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
match self
.forward(provider, endpoint, body, headers, adapter)
.await
{
Ok(response) => return Ok(response),
Err(e) => {
let category = self.categorize_proxy_error(&e);
// 只有可重试的错误才继续重试
if category == ErrorCategory::NonRetryable {
return Err(e);
}
log::debug!(
"[{}] Provider {} 第 {} 次请求失败: {}",
adapter.name(),
provider.name,
attempt + 1,
e
);
last_error = Some(e);
}
}
}
Err(last_error.unwrap_or(ProxyError::MaxRetriesExceeded))
}
/// 转发请求(带故障转移) /// 转发请求(带故障转移)
/// ///
/// # Arguments /// # Arguments
@@ -163,16 +131,13 @@ impl RequestForwarder {
body: Value, body: Value,
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
providers: Vec<Provider>, providers: Vec<Provider>,
) -> Result<ForwardResult, ForwardError> { ) -> Result<Response, ProxyError> {
// 获取适配器 // 获取适配器
let adapter = get_adapter(app_type); let adapter = get_adapter(app_type);
let app_type_str = app_type.as_str(); let app_type_str = app_type.as_str();
if providers.is_empty() { if providers.is_empty() {
return Err(ForwardError { return Err(ProxyError::NoAvailableProvider);
error: ProxyError::NoAvailableProvider,
provider: None,
});
} }
log::info!( log::info!(
@@ -182,27 +147,17 @@ impl RequestForwarder {
); );
let mut last_error = None; let mut last_error = None;
let mut last_provider = None; let mut failover_happened = false;
let mut attempted_providers = 0usize; let mut attempted_providers = 0usize;
// 单 Provider 场景下跳过熔断器检查(故障转移关闭时)
let bypass_circuit_breaker = providers.len() == 1;
// 依次尝试每个供应商 // 依次尝试每个供应商
for provider in providers.iter() { for provider in providers.iter() {
// 发起请求前先获取熔断器放行许可(HalfOpen 会占用探测名额) // 发起请求前先获取熔断器放行许可(HalfOpen 会占用探测名额)
// 单 Provider 场景下跳过此检查,避免熔断器阻塞所有请求 if !self
let (allowed, used_half_open_permit) = if bypass_circuit_breaker { .router
(true, false) .allow_provider_request(&provider.id, app_type_str)
} else { .await
let permit = self {
.router
.allow_provider_request(&provider.id, app_type_str)
.await;
(permit.allowed, permit.used_half_open_permit)
};
if !allowed {
log::debug!( log::debug!(
"[{}] Provider {} 熔断器拒绝本次请求,跳过", "[{}] Provider {} 熔断器拒绝本次请求,跳过",
app_type_str, app_type_str,
@@ -212,6 +167,9 @@ impl RequestForwarder {
} }
attempted_providers += 1; attempted_providers += 1;
if attempted_providers > 1 {
failover_happened = true;
}
log::info!( log::info!(
"[{}] 尝试 {}/{} - 使用Provider: {} (sort_index: {})", "[{}] 尝试 {}/{} - 使用Provider: {} (sort_index: {})",
@@ -233,9 +191,9 @@ impl RequestForwarder {
let start = Instant::now(); let start = Instant::now();
// 转发请求(每个 Provider 只尝试一次,重试由客户端控制 // 转发请求(带单 Provider 内重试
match self match self
.forward(provider, endpoint, &body, &headers, adapter.as_ref()) .forward_with_provider_retry(provider, endpoint, &body, &headers, adapter.as_ref())
.await .await
{ {
Ok(response) => { Ok(response) => {
@@ -244,13 +202,7 @@ impl RequestForwarder {
// 成功:记录成功并更新熔断器 // 成功:记录成功并更新熔断器
if let Err(e) = self if let Err(e) = self
.router .router
.record_result( .record_result(&provider.id, app_type_str, true, None)
&provider.id,
app_type_str,
used_half_open_permit,
true,
None,
)
.await .await
{ {
log::warn!("Failed to record success: {e}"); log::warn!("Failed to record success: {e}");
@@ -270,18 +222,16 @@ impl RequestForwarder {
let mut status = self.status.write().await; let mut status = self.status.write().await;
status.success_requests += 1; status.success_requests += 1;
status.last_error = None; status.last_error = None;
let should_switch = if failover_happened {
self.current_provider_id_at_start.as_str() != provider.id.as_str();
if should_switch {
status.failover_count += 1; status.failover_count += 1;
log::info!( log::info!(
"[{}] 代理目标已切换到 Provider: {} (耗时: {}ms)", "[{}] 故障转移成功!切换到 Provider: {} (耗时: {}ms)",
app_type_str, app_type_str,
provider.name, provider.name,
latency latency
); );
// 异步触发供应商切换,更新 UI/托盘,并把“当前供应商”同步为实际使用的 provider // 异步触发供应商切换,更新 UI 和托盘菜单
let fm = self.failover_manager.clone(); let fm = self.failover_manager.clone();
let ah = self.app_handle.clone(); let ah = self.app_handle.clone();
let pid = provider.id.clone(); let pid = provider.id.clone();
@@ -310,10 +260,7 @@ impl RequestForwarder {
latency latency
); );
return Ok(ForwardResult { return Ok(response);
response,
provider: provider.clone(),
});
} }
Err(e) => { Err(e) => {
let latency = start.elapsed().as_millis() as u64; let latency = start.elapsed().as_millis() as u64;
@@ -321,13 +268,7 @@ impl RequestForwarder {
// 失败:记录失败并更新熔断器 // 失败:记录失败并更新熔断器
if let Err(record_err) = self if let Err(record_err) = self
.router .router
.record_result( .record_result(&provider.id, app_type_str, false, Some(e.to_string()))
&provider.id,
app_type_str,
used_half_open_permit,
false,
Some(e.to_string()),
)
.await .await
{ {
log::warn!("Failed to record failure: {record_err}"); log::warn!("Failed to record failure: {record_err}");
@@ -354,7 +295,6 @@ impl RequestForwarder {
); );
last_error = Some(e); last_error = Some(e);
last_provider = Some(provider.clone());
// 继续尝试下一个供应商 // 继续尝试下一个供应商
continue; continue;
} }
@@ -376,10 +316,7 @@ impl RequestForwarder {
provider.name, provider.name,
e e
); );
return Err(ForwardError { return Err(e);
error: e,
provider: Some(provider.clone()),
});
} }
} }
} }
@@ -397,10 +334,7 @@ impl RequestForwarder {
(status.success_requests as f32 / status.total_requests as f32) * 100.0; (status.success_requests as f32 / status.total_requests as f32) * 100.0;
} }
} }
return Err(ForwardError { return Err(ProxyError::NoAvailableProvider);
error: ProxyError::NoAvailableProvider,
provider: None,
});
} }
// 所有供应商都失败了 // 所有供应商都失败了
@@ -420,10 +354,7 @@ impl RequestForwarder {
providers.len() providers.len()
); );
Err(ForwardError { Err(last_error.unwrap_or(ProxyError::MaxRetriesExceeded))
error: last_error.unwrap_or(ProxyError::MaxRetriesExceeded),
provider: last_provider,
})
} }
/// 转发单个请求(使用适配器) /// 转发单个请求(使用适配器)
@@ -439,19 +370,12 @@ impl RequestForwarder {
let base_url = adapter.extract_base_url(provider)?; let base_url = adapter.extract_base_url(provider)?;
log::info!("[{}] base_url: {}", adapter.name(), base_url); log::info!("[{}] base_url: {}", adapter.name(), base_url);
// 使用适配器构建 URL
let url = adapter.build_url(&base_url, endpoint);
// 检查是否需要格式转换 // 检查是否需要格式转换
let needs_transform = adapter.needs_transform(provider); let needs_transform = adapter.needs_transform(provider);
let effective_endpoint =
if needs_transform && adapter.name() == "Claude" && endpoint == "/v1/messages" {
"/v1/chat/completions"
} else {
endpoint
};
// 使用适配器构建 URL
let url = adapter.build_url(&base_url, effective_endpoint);
// 记录原始请求 JSON // 记录原始请求 JSON
log::info!( log::info!(
"[{}] ====== 请求开始 ======\n>>> 原始请求 JSON:\n{}", "[{}] ====== 请求开始 ======\n>>> 原始请求 JSON:\n{}",
@@ -459,23 +383,10 @@ impl RequestForwarder {
serde_json::to_string_pretty(body).unwrap_or_else(|_| body.to_string()) serde_json::to_string_pretty(body).unwrap_or_else(|_| body.to_string())
); );
// 应用模型映射(独立于格式转换)
let (mapped_body, _original_model, mapped_model) =
super::model_mapper::apply_model_mapping(body.clone(), provider);
if let Some(ref mapped) = mapped_model {
log::info!(
"[{}] >>> 模型映射后的请求 JSON:\n{}",
adapter.name(),
serde_json::to_string_pretty(&mapped_body).unwrap_or_default()
);
log::info!("[{}] 模型已映射到: {}", adapter.name(), mapped);
}
// 转换请求体(如果需要) // 转换请求体(如果需要)
let request_body = if needs_transform { let request_body = if needs_transform {
log::info!("[{}] 转换请求格式 (Anthropic → OpenAI)", adapter.name()); log::info!("[{}] 转换请求格式 (Anthropic → OpenAI)", adapter.name());
let transformed = adapter.transform_request(mapped_body, provider)?; let transformed = adapter.transform_request(body.clone(), provider)?;
log::info!( log::info!(
"[{}] >>> 转换后的请求 JSON:\n{}", "[{}] >>> 转换后的请求 JSON:\n{}",
adapter.name(), adapter.name(),
@@ -483,31 +394,9 @@ impl RequestForwarder {
); );
transformed transformed
} else { } else {
mapped_body body.clone()
}; };
// 过滤私有参数(以 `_` 开头的字段),防止内部信息泄露到上游
// 默认使用空白名单,过滤所有 _ 前缀字段
let filtered_body = filter_private_params_with_whitelist(request_body, &[]);
// ========== 请求体日志(截断显示) ==========
let body_str = serde_json::to_string_pretty(&filtered_body)
.unwrap_or_else(|_| filtered_body.to_string());
let body_preview = if body_str.len() > 2000 {
format!(
"{}...\n[截断,总长度: {} 字符]",
&body_str[..2000],
body_str.len()
)
} else {
body_str
};
log::info!(
"[{}] ====== 最终请求体 ======\n{}",
adapter.name(),
body_preview
);
log::info!( log::info!(
"[{}] 转发请求: {} -> {}", "[{}] 转发请求: {} -> {}",
adapter.name(), adapter.name(),
@@ -518,73 +407,28 @@ impl RequestForwarder {
// 构建请求 // 构建请求
let mut request = self.client.post(&url); let mut request = self.client.post(&url);
// ========== 详细 Headers 日志 ========== // 只透传必要的 Headers(白名单模式)
log::info!("[{}] ====== 客户端原始 Headers ======", adapter.name()); let allowed_headers = [
for (key, value) in headers { "accept",
log::info!( "user-agent",
"[{}] {}: {:?}", "x-request-id",
adapter.name(), "x-stainless-arch",
key.as_str(), "x-stainless-lang",
value.to_str().unwrap_or("<binary>") "x-stainless-os",
); "x-stainless-package-version",
} "x-stainless-runtime",
"x-stainless-runtime-version",
// 过滤黑名单 Headers,保护隐私并避免冲突 ];
let mut filtered_headers: Vec<String> = Vec::new();
let mut passed_headers: Vec<(String, String)> = Vec::new();
for (key, value) in headers { for (key, value) in headers {
let key_str = key.as_str().to_lowercase(); let key_str = key.as_str().to_lowercase();
if HEADER_BLACKLIST.contains(&key_str.as_str()) { if allowed_headers.contains(&key_str.as_str()) {
filtered_headers.push(key_str); request = request.header(key, value);
continue;
}
let value_str = value.to_str().unwrap_or("<binary>").to_string();
passed_headers.push((key.as_str().to_string(), value_str.clone()));
request = request.header(key, value);
}
if !filtered_headers.is_empty() {
log::info!(
"[{}] ====== 被过滤的 Headers ({}) ======",
adapter.name(),
filtered_headers.len()
);
for h in &filtered_headers {
log::info!("[{}] - {}", adapter.name(), h);
} }
} }
// 处理 anthropic-beta Header(透传) // 确保 Content-Type 是 json
// 参考 Claude Code Hub 的实现,直接透传客户端的 beta 标记 request = request.header("Content-Type", "application/json");
if let Some(beta) = headers.get("anthropic-beta") {
if let Ok(beta_str) = beta.to_str() {
request = request.header("anthropic-beta", beta_str);
passed_headers.push(("anthropic-beta".to_string(), beta_str.to_string()));
log::info!("[{}] 透传 anthropic-beta: {}", adapter.name(), beta_str);
}
}
// 客户端 IP 透传(默认开启)
if let Some(xff) = headers.get("x-forwarded-for") {
if let Ok(xff_str) = xff.to_str() {
request = request.header("x-forwarded-for", xff_str);
passed_headers.push(("x-forwarded-for".to_string(), xff_str.to_string()));
log::debug!("[{}] 透传 x-forwarded-for: {}", adapter.name(), xff_str);
}
}
if let Some(real_ip) = headers.get("x-real-ip") {
if let Ok(real_ip_str) = real_ip.to_str() {
request = request.header("x-real-ip", real_ip_str);
passed_headers.push(("x-real-ip".to_string(), real_ip_str.to_string()));
log::debug!("[{}] 透传 x-real-ip: {}", adapter.name(), real_ip_str);
}
}
// 禁用压缩,避免 gzip 流式响应解析错误
// 参考 CCH: undici 在连接提前关闭时会对不完整的 gzip 流抛出错误
request = request.header("accept-encoding", "identity");
passed_headers.push(("accept-encoding".to_string(), "identity".to_string()));
// 使用适配器添加认证头 // 使用适配器添加认证头
if let Some(auth) = adapter.extract_auth(provider) { if let Some(auth) = adapter.extract_auth(provider) {
@@ -595,15 +439,6 @@ impl RequestForwarder {
auth.masked_key() auth.masked_key()
); );
request = adapter.add_auth_headers(request, &auth); request = adapter.add_auth_headers(request, &auth);
// 记录认证头(脱敏)
passed_headers.push((
"authorization".to_string(),
format!("Bearer {}...", &auth.api_key[..8.min(auth.api_key.len())]),
));
passed_headers.push((
"x-api-key".to_string(),
format!("{}...", &auth.api_key[..8.min(auth.api_key.len())]),
));
} else { } else {
log::error!( log::error!(
"[{}] 未找到 API KeyProvider: {}", "[{}] 未找到 API KeyProvider: {}",
@@ -612,34 +447,9 @@ impl RequestForwarder {
); );
} }
// anthropic-version 透传:优先使用客户端的版本号
// 参考 Claude Code Hub:透传客户端值而非固定版本
if let Some(version) = headers.get("anthropic-version") {
if let Ok(version_str) = version.to_str() {
// 覆盖适配器设置的默认版本
request = request.header("anthropic-version", version_str);
passed_headers.push(("anthropic-version".to_string(), version_str.to_string()));
log::info!(
"[{}] 透传 anthropic-version: {}",
adapter.name(),
version_str
);
}
}
// ========== 最终发送的 Headers 日志 ==========
log::info!(
"[{}] ====== 最终发送的 Headers ({}) ======",
adapter.name(),
passed_headers.len()
);
for (k, v) in &passed_headers {
log::info!("[{}] {}: {}", adapter.name(), k, v);
}
// 发送请求 // 发送请求
log::info!("[{}] 发送请求到: {}", adapter.name(), url); log::info!("[{}] 发送请求到: {}", adapter.name(), url);
let response = request.json(&filtered_body).send().await.map_err(|e| { let response = request.json(&request_body).send().await.map_err(|e| {
log::error!("[{}] 请求失败: {}", adapter.name(), e); log::error!("[{}] 请求失败: {}", adapter.name(), e);
if e.is_timeout() { if e.is_timeout() {
ProxyError::Timeout(format!("请求超时: {e}")) ProxyError::Timeout(format!("请求超时: {e}"))
@@ -673,6 +483,12 @@ impl RequestForwarder {
} }
} }
/// 分类ProxyError
///
/// 决定哪些错误应该触发故障转移到下一个 Provider
///
/// 设计原则:既然用户配置了多个供应商,就应该让所有供应商都尝试一遍。
/// 只有明确是客户端中断的情况才不重试。
fn categorize_proxy_error(&self, error: &ProxyError) -> ErrorCategory { fn categorize_proxy_error(&self, error: &ProxyError) -> ErrorCategory {
match error { match error {
// 网络和上游错误:都应该尝试下一个供应商 // 网络和上游错误:都应该尝试下一个供应商
@@ -683,14 +499,9 @@ impl RequestForwarder {
// 原因:不同供应商有不同的限制和认证,一个供应商的 4xx 错误 // 原因:不同供应商有不同的限制和认证,一个供应商的 4xx 错误
// 不代表其他供应商也会失败 // 不代表其他供应商也会失败
ProxyError::UpstreamError { .. } => ErrorCategory::Retryable, ProxyError::UpstreamError { .. } => ErrorCategory::Retryable,
// Provider 级配置/转换问题:换一个 Provider 可能就能成功
ProxyError::ConfigError(_) => ErrorCategory::Retryable,
ProxyError::TransformError(_) => ErrorCategory::Retryable,
ProxyError::AuthError(_) => ErrorCategory::Retryable,
ProxyError::StreamIdleTimeout(_) => ErrorCategory::Retryable,
// 无可用供应商:所有供应商都试过了,无法重试 // 无可用供应商:所有供应商都试过了,无法重试
ProxyError::NoAvailableProvider => ErrorCategory::NonRetryable, ProxyError::NoAvailableProvider => ErrorCategory::NonRetryable,
// 其他错误(数据库/内部错误等):不是供应商能解决的问题 // 其他错误(配置错误、数据库错误等):不是供应商问题,无需重试
_ => ErrorCategory::NonRetryable, _ => ErrorCategory::NonRetryable,
} }
} }
+9 -33
View File
@@ -31,26 +31,13 @@ pub struct UsageParserConfig {
// 模型提取器实现 // 模型提取器实现
// ============================================================================ // ============================================================================
/// Claude 流式响应模型提取(优先使用 usage.model /// Claude 流式响应模型提取(直接使用请求模型
fn claude_model_extractor(events: &[Value], request_model: &str) -> String { fn claude_model_extractor(_events: &[Value], request_model: &str) -> String {
// 首先尝试从解析的 usage 中获取模型
if let Some(usage) = TokenUsage::from_claude_stream_events(events) {
if let Some(model) = usage.model {
return model;
}
}
request_model.to_string() request_model.to_string()
} }
/// OpenAI Chat Completions 流式响应模型提取(优先使用 usage.model /// OpenAI Chat Completions 流式响应模型提取
fn openai_model_extractor(events: &[Value], request_model: &str) -> String { fn openai_model_extractor(events: &[Value], request_model: &str) -> String {
// 首先尝试从解析的 usage 中获取模型
if let Some(usage) = TokenUsage::from_openai_stream_events(events) {
if let Some(model) = usage.model {
return model;
}
}
// 回退:从事件中直接提取
events events
.iter() .iter()
.find_map(|e| e.get("model")?.as_str()) .find_map(|e| e.get("model")?.as_str())
@@ -58,15 +45,8 @@ fn openai_model_extractor(events: &[Value], request_model: &str) -> String {
.to_string() .to_string()
} }
/// Codex 智能流式响应模型提取(自动检测格式) /// Codex Responses API 流式响应模型提取
fn codex_auto_model_extractor(events: &[Value], request_model: &str) -> String { fn codex_model_extractor(events: &[Value], request_model: &str) -> String {
// 首先尝试从解析的 usage 中获取模型
if let Some(usage) = TokenUsage::from_codex_stream_events_auto(events) {
if let Some(model) = usage.model {
return model;
}
}
// 回退:从 response.completed 事件中提取
events events
.iter() .iter()
.find_map(|e| { .find_map(|e| {
@@ -76,10 +56,6 @@ fn codex_auto_model_extractor(events: &[Value], request_model: &str) -> String {
None None
} }
}) })
.or_else(|| {
// 再回退:从 OpenAI 格式事件中提取
events.iter().find_map(|e| e.get("model")?.as_str())
})
.unwrap_or(request_model) .unwrap_or(request_model)
.to_string() .to_string()
} }
@@ -115,11 +91,11 @@ pub const OPENAI_PARSER_CONFIG: UsageParserConfig = UsageParserConfig {
app_type_str: "codex", app_type_str: "codex",
}; };
/// Codex 智能解析配置(自动检测 OpenAI 或 Codex 格式 /// Codex Responses API 解析配置(用于 /v1/responses
pub const CODEX_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { pub const CODEX_PARSER_CONFIG: UsageParserConfig = UsageParserConfig {
stream_parser: TokenUsage::from_codex_stream_events_auto, stream_parser: TokenUsage::from_codex_stream_events,
response_parser: TokenUsage::from_codex_response_auto, response_parser: TokenUsage::from_codex_response,
model_extractor: codex_auto_model_extractor, model_extractor: codex_model_extractor,
app_type_str: "codex", app_type_str: "codex",
}; };
+11 -107
View File
@@ -5,44 +5,27 @@
use crate::app_config::AppType; use crate::app_config::AppType;
use crate::provider::Provider; use crate::provider::Provider;
use crate::proxy::{ use crate::proxy::{
extract_session_id, forwarder::RequestForwarder, server::ProxyState, types::AppProxyConfig, forwarder::RequestForwarder, server::ProxyState, types::ProxyConfig, ProxyError,
ProxyError,
}; };
use axum::http::HeaderMap;
use std::time::Instant; use std::time::Instant;
/// 流式超时配置
#[derive(Debug, Clone, Copy)]
pub struct StreamingTimeoutConfig {
/// 首字节超时(秒),0 表示禁用
pub first_byte_timeout: u64,
/// 静默期超时(秒),0 表示禁用
pub idle_timeout: u64,
}
/// 请求上下文 /// 请求上下文
/// ///
/// 贯穿整个请求生命周期,包含: /// 贯穿整个请求生命周期,包含:
/// - 计时信息 /// - 计时信息
/// - 应用级代理配置per-app /// - 代理配置
/// - 选中的 Provider 列表(用于故障转移) /// - 选中的 Provider 列表(用于故障转移)
/// - 请求模型名称 /// - 请求模型名称
/// - 日志标签 /// - 日志标签
/// - Session ID(用于日志关联)
pub struct RequestContext { pub struct RequestContext {
/// 请求开始时间 /// 请求开始时间
pub start_time: Instant, pub start_time: Instant,
/// 应用级代理配置(per-app,包含重试次数和超时配置) /// 代理配置快照
pub app_config: AppProxyConfig, pub config: ProxyConfig,
/// 选中的 Provider(故障转移链的第一个) /// 选中的 Provider(故障转移链的第一个)
pub provider: Provider, pub provider: Provider,
/// 完整的 Provider 列表(用于故障转移) /// 完整的 Provider 列表(用于故障转移)
providers: Vec<Provider>, providers: Vec<Provider>,
/// 请求开始时的"当前供应商"(用于判断是否需要同步 UI/托盘)
///
/// 这里使用本地 settings 的设备级 current provider。
/// 代理模式下如果实际使用的 provider 与此不一致,会触发切换以确保 UI 始终准确。
pub current_provider_id: String,
/// 请求中的模型名称 /// 请求中的模型名称
pub request_model: String, pub request_model: String,
/// 日志标签(如 "Claude"、"Codex"、"Gemini" /// 日志标签(如 "Claude"、"Codex"、"Gemini"
@@ -52,8 +35,6 @@ pub struct RequestContext {
/// 应用类型(预留,目前通过 app_type_str 使用) /// 应用类型(预留,目前通过 app_type_str 使用)
#[allow(dead_code)] #[allow(dead_code)]
pub app_type: AppType, pub app_type: AppType,
/// Session ID(从客户端请求提取或新生成)
pub session_id: String,
} }
impl RequestContext { impl RequestContext {
@@ -62,7 +43,6 @@ impl RequestContext {
/// # Arguments /// # Arguments
/// * `state` - 代理服务器状态 /// * `state` - 代理服务器状态
/// * `body` - 请求体 JSON /// * `body` - 请求体 JSON
/// * `headers` - 请求头(用于提取 Session ID
/// * `app_type` - 应用类型 /// * `app_type` - 应用类型
/// * `tag` - 日志标签 /// * `tag` - 日志标签
/// * `app_type_str` - 应用类型字符串 /// * `app_type_str` - 应用类型字符串
@@ -72,22 +52,12 @@ impl RequestContext {
pub async fn new( pub async fn new(
state: &ProxyState, state: &ProxyState,
body: &serde_json::Value, body: &serde_json::Value,
headers: &HeaderMap,
app_type: AppType, app_type: AppType,
tag: &'static str, tag: &'static str,
app_type_str: &'static str, app_type_str: &'static str,
) -> Result<Self, ProxyError> { ) -> Result<Self, ProxyError> {
let start_time = Instant::now(); let start_time = Instant::now();
let config = state.config.read().await.clone();
// 从数据库读取应用级代理配置(per-app)
let app_config = state
.db
.get_proxy_config_for_app(app_type_str)
.await
.map_err(|e| ProxyError::DatabaseError(e.to_string()))?;
let current_provider_id =
crate::settings::get_current_provider(&app_type).unwrap_or_default();
// 从请求体提取模型名称 // 从请求体提取模型名称
let request_model = body let request_model = body
@@ -96,31 +66,13 @@ impl RequestContext {
.unwrap_or("unknown") .unwrap_or("unknown")
.to_string(); .to_string();
// 提取 Session ID
let session_result = extract_session_id(headers, body, app_type_str);
let session_id = session_result.session_id.clone();
log::debug!(
"[{}] Session ID: {} (from {:?}, client_provided: {})",
tag,
session_id,
session_result.source,
session_result.client_provided
);
// 使用共享的 ProviderRouter 选择 Provider(熔断器状态跨请求保持) // 使用共享的 ProviderRouter 选择 Provider(熔断器状态跨请求保持)
// 注意:只在这里调用一次,结果传递给 forwarder,避免重复消耗 HalfOpen 名额 // 注意:只在这里调用一次,结果传递给 forwarder,避免重复消耗 HalfOpen 名额
let providers = state let providers = state
.provider_router .provider_router
.select_providers(app_type_str) .select_providers(app_type_str)
.await .await
.map_err(|e| match e { .map_err(|e| ProxyError::DatabaseError(e.to_string()))?;
crate::error::AppError::AllProvidersCircuitOpen => {
ProxyError::AllProvidersCircuitOpen
}
crate::error::AppError::NoProvidersConfigured => ProxyError::NoProvidersConfigured,
_ => ProxyError::DatabaseError(e.to_string()),
})?;
let provider = providers let provider = providers
.first() .first()
@@ -128,25 +80,22 @@ impl RequestContext {
.ok_or(ProxyError::NoAvailableProvider)?; .ok_or(ProxyError::NoAvailableProvider)?;
log::info!( log::info!(
"[{}] Provider: {}, model: {}, failover chain: {} providers, session: {}", "[{}] Provider: {}, model: {}, failover chain: {} providers",
tag, tag,
provider.name, provider.name,
request_model, request_model,
providers.len(), providers.len()
session_id
); );
Ok(Self { Ok(Self {
start_time, start_time,
app_config, config,
provider, provider,
providers, providers,
current_provider_id,
request_model, request_model,
tag, tag,
app_type_str, app_type_str,
app_type, app_type,
session_id,
}) })
} }
@@ -175,38 +124,15 @@ impl RequestContext {
/// 创建 RequestForwarder /// 创建 RequestForwarder
/// ///
/// 使用共享的 ProviderRouter,确保熔断器状态跨请求保持 /// 使用共享的 ProviderRouter,确保熔断器状态跨请求保持
///
/// 配置生效规则:
/// - 故障转移开启:超时配置正常生效(0 表示禁用超时)
/// - 故障转移关闭:超时配置不生效(全部传入 0)
pub fn create_forwarder(&self, state: &ProxyState) -> RequestForwarder { pub fn create_forwarder(&self, state: &ProxyState) -> RequestForwarder {
let (non_streaming_timeout, first_byte_timeout, idle_timeout) =
if self.app_config.auto_failover_enabled {
// 故障转移开启:使用配置的值(0 = 禁用超时)
(
self.app_config.non_streaming_timeout as u64,
self.app_config.streaming_first_byte_timeout as u64,
self.app_config.streaming_idle_timeout as u64,
)
} else {
// 故障转移关闭:不启用超时配置
log::info!(
"[{}] Failover disabled, timeout configs are bypassed",
self.tag
);
(0, 0, 0)
};
RequestForwarder::new( RequestForwarder::new(
state.provider_router.clone(), state.provider_router.clone(),
non_streaming_timeout, self.config.request_timeout,
self.config.max_retries,
state.status.clone(), state.status.clone(),
state.current_providers.clone(), state.current_providers.clone(),
state.failover_manager.clone(), state.failover_manager.clone(),
state.app_handle.clone(), state.app_handle.clone(),
self.current_provider_id.clone(),
first_byte_timeout,
idle_timeout,
) )
} }
@@ -222,26 +148,4 @@ impl RequestContext {
pub fn latency_ms(&self) -> u64 { pub fn latency_ms(&self) -> u64 {
self.start_time.elapsed().as_millis() as u64 self.start_time.elapsed().as_millis() as u64
} }
/// 获取流式超时配置
///
/// 配置生效规则:
/// - 故障转移开启:返回配置的值(0 表示禁用超时检查)
/// - 故障转移关闭:返回 0(禁用超时检查)
#[inline]
pub fn streaming_timeout_config(&self) -> StreamingTimeoutConfig {
if self.app_config.auto_failover_enabled {
// 故障转移开启:使用配置的值(0 = 禁用超时)
StreamingTimeoutConfig {
first_byte_timeout: self.app_config.streaming_first_byte_timeout as u64,
idle_timeout: self.app_config.streaming_idle_timeout as u64,
}
} else {
// 故障转移关闭:禁用流式超时检查
StreamingTimeoutConfig {
first_byte_timeout: 0,
idle_timeout: 0,
}
}
}
} }
+41 -72
View File
@@ -5,7 +5,7 @@
//! 重构后的结构: //! 重构后的结构:
//! - 通用逻辑提取到 `handler_context` 和 `response_processor` 模块 //! - 通用逻辑提取到 `handler_context` 和 `response_processor` 模块
//! - 各 handler 只保留独特的业务逻辑 //! - 各 handler 只保留独特的业务逻辑
//! - Claude 的格式转换逻辑保留在此文件(用于 OpenRouter 旧接口回退 //! - Claude 的格式转换逻辑保留在此文件(独有功能
use super::{ use super::{
error_mapper::{get_error_message, map_proxy_error_to_status}, error_mapper::{get_error_message, map_proxy_error_to_status},
@@ -54,24 +54,34 @@ pub async fn get_status(State(state): State<ProxyState>) -> Result<Json<ProxySta
/// 处理 /v1/messages 请求(Claude API /// 处理 /v1/messages 请求(Claude API
/// ///
/// Claude 处理器包含独特的格式转换逻辑: /// Claude 处理器包含独特的格式转换逻辑:
/// - 过去用于 OpenRouter 的 OpenAI Chat Completions 兼容接口(Anthropic OpenAI 转换) /// - 当使用 OpenRouter 等中转服务时,需要将 Anthropic 格式转换为 OpenAI 格式
/// - 现在 OpenRouter 已推出 Claude Code 兼容接口,默认不再启用该转换(逻辑保留以备回退) /// - 响应需要从 OpenAI 格式转回 Anthropic 格式
pub async fn handle_messages( pub async fn handle_messages(
State(state): State<ProxyState>, State(state): State<ProxyState>,
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
Json(body): Json<Value>, Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> { ) -> Result<axum::response::Response, ProxyError> {
let mut ctx = let ctx = RequestContext::new(&state, &body, AppType::Claude, "Claude", "claude").await?;
RequestContext::new(&state, &body, &headers, AppType::Claude, "Claude", "claude").await?;
// 检查是否需要格式转换(OpenRouter 等中转服务)
let adapter = get_adapter(&AppType::Claude);
let needs_transform = adapter.needs_transform(&ctx.provider);
let is_stream = body let is_stream = body
.get("stream") .get("stream")
.and_then(|s| s.as_bool()) .and_then(|s| s.as_bool())
.unwrap_or(false); .unwrap_or(false);
log::info!(
"[Claude] Provider: {}, needs_transform: {}, is_stream: {}",
ctx.provider.name,
needs_transform,
is_stream
);
// 转发请求 // 转发请求
let forwarder = ctx.create_forwarder(&state); let forwarder = ctx.create_forwarder(&state);
let result = match forwarder let response = match forwarder
.forward_with_retry( .forward_with_retry(
&AppType::Claude, &AppType::Claude,
"/v1/messages", "/v1/messages",
@@ -81,30 +91,13 @@ pub async fn handle_messages(
) )
.await .await
{ {
Ok(result) => result, Ok(resp) => resp,
Err(mut err) => { Err(e) => {
if let Some(provider) = err.provider.take() { log_forward_error(&state, &ctx, is_stream, &e);
ctx.provider = provider; return Err(e);
}
log_forward_error(&state, &ctx, is_stream, &err.error);
return Err(err.error);
} }
}; };
ctx.provider = result.provider;
let response = result.response;
// 检查是否需要格式转换(OpenRouter 等中转服务)
let adapter = get_adapter(&AppType::Claude);
let needs_transform = adapter.needs_transform(&ctx.provider);
log::info!(
"[Claude] Provider: {}, needs_transform: {}, is_stream: {}",
ctx.provider.name,
needs_transform,
is_stream
);
let status = response.status(); let status = response.status();
log::info!("[Claude] 上游响应状态: {status}"); log::info!("[Claude] 上游响应状态: {status}");
@@ -119,7 +112,7 @@ pub async fn handle_messages(
/// Claude 格式转换处理(独有逻辑) /// Claude 格式转换处理(独有逻辑)
/// ///
/// 处理 OpenRouter 旧 OpenAI 兼容接口的回退方案(当前默认不启用) /// 处理 OpenRouter 等需要格式转换的中转服务
async fn handle_claude_transform( async fn handle_claude_transform(
response: reqwest::Response, response: reqwest::Response,
ctx: &RequestContext, ctx: &RequestContext,
@@ -171,14 +164,10 @@ async fn handle_claude_transform(
}) })
}; };
// 获取流式超时配置
let timeout_config = ctx.streaming_timeout_config();
let logged_stream = create_logged_passthrough_stream( let logged_stream = create_logged_passthrough_stream(
sse_stream, sse_stream,
"Claude/OpenRouter", "Claude/OpenRouter",
Some(usage_collector), Some(usage_collector),
timeout_config,
); );
let mut headers = axum::http::HeaderMap::new(); let mut headers = axum::http::HeaderMap::new();
@@ -306,8 +295,7 @@ pub async fn handle_chat_completions(
) -> Result<axum::response::Response, ProxyError> { ) -> Result<axum::response::Response, ProxyError> {
log::info!("[Codex] ====== /v1/chat/completions 请求开始 ======"); log::info!("[Codex] ====== /v1/chat/completions 请求开始 ======");
let mut ctx = let ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?;
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
let is_stream = body let is_stream = body
.get("stream") .get("stream")
@@ -321,7 +309,7 @@ pub async fn handle_chat_completions(
); );
let forwarder = ctx.create_forwarder(&state); let forwarder = ctx.create_forwarder(&state);
let result = match forwarder let response = match forwarder
.forward_with_retry( .forward_with_retry(
&AppType::Codex, &AppType::Codex,
"/v1/chat/completions", "/v1/chat/completions",
@@ -331,19 +319,13 @@ pub async fn handle_chat_completions(
) )
.await .await
{ {
Ok(result) => result, Ok(resp) => resp,
Err(mut err) => { Err(e) => {
if let Some(provider) = err.provider.take() { log_forward_error(&state, &ctx, is_stream, &e);
ctx.provider = provider; return Err(e);
}
log_forward_error(&state, &ctx, is_stream, &err.error);
return Err(err.error);
} }
}; };
ctx.provider = result.provider;
let response = result.response;
log::info!("[Codex] 上游响应状态: {}", response.status()); log::info!("[Codex] 上游响应状态: {}", response.status());
process_response(response, &ctx, &state, &OPENAI_PARSER_CONFIG).await process_response(response, &ctx, &state, &OPENAI_PARSER_CONFIG).await
@@ -355,8 +337,7 @@ pub async fn handle_responses(
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
Json(body): Json<Value>, Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> { ) -> Result<axum::response::Response, ProxyError> {
let mut ctx = let ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?;
RequestContext::new(&state, &body, &headers, AppType::Codex, "Codex", "codex").await?;
let is_stream = body let is_stream = body
.get("stream") .get("stream")
@@ -364,7 +345,7 @@ pub async fn handle_responses(
.unwrap_or(false); .unwrap_or(false);
let forwarder = ctx.create_forwarder(&state); let forwarder = ctx.create_forwarder(&state);
let result = match forwarder let response = match forwarder
.forward_with_retry( .forward_with_retry(
&AppType::Codex, &AppType::Codex,
"/v1/responses", "/v1/responses",
@@ -374,19 +355,13 @@ pub async fn handle_responses(
) )
.await .await
{ {
Ok(result) => result, Ok(resp) => resp,
Err(mut err) => { Err(e) => {
if let Some(provider) = err.provider.take() { log_forward_error(&state, &ctx, is_stream, &e);
ctx.provider = provider; return Err(e);
}
log_forward_error(&state, &ctx, is_stream, &err.error);
return Err(err.error);
} }
}; };
ctx.provider = result.provider;
let response = result.response;
log::info!("[Codex] 上游响应状态: {}", response.status()); log::info!("[Codex] 上游响应状态: {}", response.status());
process_response(response, &ctx, &state, &CODEX_PARSER_CONFIG).await process_response(response, &ctx, &state, &CODEX_PARSER_CONFIG).await
@@ -404,7 +379,7 @@ pub async fn handle_gemini(
Json(body): Json<Value>, Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> { ) -> Result<axum::response::Response, ProxyError> {
// Gemini 的模型名称在 URI 中 // Gemini 的模型名称在 URI 中
let mut ctx = RequestContext::new(&state, &body, &headers, AppType::Gemini, "Gemini", "gemini") let ctx = RequestContext::new(&state, &body, AppType::Gemini, "Gemini", "gemini")
.await? .await?
.with_model_from_uri(&uri); .with_model_from_uri(&uri);
@@ -422,7 +397,7 @@ pub async fn handle_gemini(
.unwrap_or(false); .unwrap_or(false);
let forwarder = ctx.create_forwarder(&state); let forwarder = ctx.create_forwarder(&state);
let result = match forwarder let response = match forwarder
.forward_with_retry( .forward_with_retry(
&AppType::Gemini, &AppType::Gemini,
endpoint, endpoint,
@@ -432,19 +407,13 @@ pub async fn handle_gemini(
) )
.await .await
{ {
Ok(result) => result, Ok(resp) => resp,
Err(mut err) => { Err(e) => {
if let Some(provider) = err.provider.take() { log_forward_error(&state, &ctx, is_stream, &e);
ctx.provider = provider; return Err(e);
}
log_forward_error(&state, &ctx, is_stream, &err.error);
return Err(err.error);
} }
}; };
ctx.provider = result.provider;
let response = result.response;
log::info!("[Gemini] 上游响应状态: {}", response.status()); log::info!("[Gemini] 上游响应状态: {}", response.status());
process_response(response, &ctx, &state, &GEMINI_PARSER_CONFIG).await process_response(response, &ctx, &state, &GEMINI_PARSER_CONFIG).await
@@ -468,7 +437,7 @@ fn log_forward_error(
let request_id = uuid::Uuid::new_v4().to_string(); let request_id = uuid::Uuid::new_v4().to_string();
if let Err(e) = logger.log_error_with_context( if let Err(e) = logger.log_error_with_context(
request_id, request_id.clone(),
ctx.provider.id.clone(), ctx.provider.id.clone(),
ctx.app_type_str.to_string(), ctx.app_type_str.to_string(),
ctx.request_model.clone(), ctx.request_model.clone(),
@@ -476,7 +445,7 @@ fn log_forward_error(
error_message, error_message,
ctx.latency_ms(), ctx.latency_ms(),
is_streaming, is_streaming,
Some(ctx.session_id.clone()), Some(request_id),
None, None,
) { ) {
log::warn!("记录失败请求日志失败: {e}"); log::warn!("记录失败请求日志失败: {e}");
+1 -5
View File
@@ -2,7 +2,6 @@
//! //!
//! 提供本地HTTP代理服务,支持多Provider故障转移和请求透传 //! 提供本地HTTP代理服务,支持多Provider故障转移和请求透传
pub mod body_filter;
pub mod circuit_breaker; pub mod circuit_breaker;
pub mod error; pub mod error;
pub mod error_mapper; pub mod error_mapper;
@@ -12,7 +11,6 @@ pub mod handler_config;
pub mod handler_context; pub mod handler_context;
mod handlers; mod handlers;
mod health; mod health;
pub mod model_mapper;
pub mod provider_router; pub mod provider_router;
pub mod providers; pub mod providers;
pub mod response_handler; pub mod response_handler;
@@ -34,9 +32,7 @@ pub use provider_router::ProviderRouter;
#[allow(unused_imports)] #[allow(unused_imports)]
pub use response_handler::{NonStreamHandler, ResponseType, StreamHandler}; pub use response_handler::{NonStreamHandler, ResponseType, StreamHandler};
#[allow(unused_imports)] #[allow(unused_imports)]
pub use session::{ pub use session::{ClientFormat, ProxySession};
extract_session_id, ClientFormat, ProxySession, SessionIdResult, SessionIdSource,
};
#[allow(unused_imports)] #[allow(unused_imports)]
pub use types::{ProxyConfig, ProxyServerInfo, ProxyStatus}; pub use types::{ProxyConfig, ProxyServerInfo, ProxyStatus};
-264
View File
@@ -1,264 +0,0 @@
//! 模型映射模块
//!
//! 在请求转发前,根据 Provider 配置替换请求中的模型名称
use crate::provider::Provider;
use serde_json::Value;
/// 模型映射配置
pub struct ModelMapping {
pub haiku_model: Option<String>,
pub sonnet_model: Option<String>,
pub opus_model: Option<String>,
pub default_model: Option<String>,
pub reasoning_model: Option<String>,
}
impl ModelMapping {
/// 从 Provider 配置中提取模型映射
pub fn from_provider(provider: &Provider) -> Self {
let env = provider.settings_config.get("env");
Self {
haiku_model: env
.and_then(|e| e.get("ANTHROPIC_DEFAULT_HAIKU_MODEL"))
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(String::from),
sonnet_model: env
.and_then(|e| e.get("ANTHROPIC_DEFAULT_SONNET_MODEL"))
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(String::from),
opus_model: env
.and_then(|e| e.get("ANTHROPIC_DEFAULT_OPUS_MODEL"))
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(String::from),
default_model: env
.and_then(|e| e.get("ANTHROPIC_MODEL"))
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(String::from),
reasoning_model: env
.and_then(|e| e.get("ANTHROPIC_REASONING_MODEL"))
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(String::from),
}
}
/// 检查是否配置了任何模型映射
pub fn has_mapping(&self) -> bool {
self.haiku_model.is_some()
|| self.sonnet_model.is_some()
|| self.opus_model.is_some()
|| self.default_model.is_some()
}
/// 根据原始模型名称获取映射后的模型
pub fn map_model(&self, original_model: &str, has_thinking: bool) -> String {
let model_lower = original_model.to_lowercase();
// 1. thinking 模式优先使用推理模型
if has_thinking {
if let Some(ref m) = self.reasoning_model {
return m.clone();
}
}
// 2. 按模型类型匹配
if model_lower.contains("haiku") {
if let Some(ref m) = self.haiku_model {
return m.clone();
}
}
if model_lower.contains("opus") {
if let Some(ref m) = self.opus_model {
return m.clone();
}
}
if model_lower.contains("sonnet") {
if let Some(ref m) = self.sonnet_model {
return m.clone();
}
}
// 3. 默认模型
if let Some(ref m) = self.default_model {
return m.clone();
}
// 4. 无映射,保持原样
original_model.to_string()
}
}
/// 检测请求是否启用了 thinking 模式
pub fn has_thinking_enabled(body: &Value) -> bool {
body.get("thinking")
.and_then(|v| v.as_object())
.and_then(|o| o.get("type"))
.and_then(|t| t.as_str())
== Some("enabled")
}
/// 对请求体应用模型映射
///
/// 返回 (映射后的请求体, 原始模型名, 映射后模型名)
pub fn apply_model_mapping(
mut body: Value,
provider: &Provider,
) -> (Value, Option<String>, Option<String>) {
let mapping = ModelMapping::from_provider(provider);
// 如果没有配置映射,直接返回
if !mapping.has_mapping() {
let original = body.get("model").and_then(|m| m.as_str()).map(String::from);
return (body, original, None);
}
// 提取原始模型名
let original_model = body.get("model").and_then(|m| m.as_str()).map(String::from);
if let Some(ref original) = original_model {
let has_thinking = has_thinking_enabled(&body);
let mapped = mapping.map_model(original, has_thinking);
if mapped != *original {
log::info!("[ModelMapper] 模型映射: {original} → {mapped}");
body["model"] = serde_json::json!(mapped);
return (body, Some(original.clone()), Some(mapped));
}
}
(body, original_model, None)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn create_provider_with_mapping() -> Provider {
Provider {
id: "test".to_string(),
name: "Test".to_string(),
settings_config: json!({
"env": {
"ANTHROPIC_MODEL": "default-model",
"ANTHROPIC_DEFAULT_HAIKU_MODEL": "haiku-mapped",
"ANTHROPIC_DEFAULT_SONNET_MODEL": "sonnet-mapped",
"ANTHROPIC_DEFAULT_OPUS_MODEL": "opus-mapped",
"ANTHROPIC_REASONING_MODEL": "reasoning-model"
}
}),
website_url: None,
category: None,
created_at: None,
sort_index: None,
notes: None,
meta: None,
icon: None,
icon_color: None,
in_failover_queue: false,
}
}
fn create_provider_without_mapping() -> Provider {
Provider {
id: "test".to_string(),
name: "Test".to_string(),
settings_config: json!({}),
website_url: None,
category: None,
created_at: None,
sort_index: None,
notes: None,
meta: None,
icon: None,
icon_color: None,
in_failover_queue: false,
}
}
#[test]
fn test_sonnet_mapping() {
let provider = create_provider_with_mapping();
let body = json!({"model": "claude-sonnet-4-5-20250929"});
let (result, original, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "sonnet-mapped");
assert_eq!(original, Some("claude-sonnet-4-5-20250929".to_string()));
assert_eq!(mapped, Some("sonnet-mapped".to_string()));
}
#[test]
fn test_haiku_mapping() {
let provider = create_provider_with_mapping();
let body = json!({"model": "claude-haiku-4-5"});
let (result, _, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "haiku-mapped");
assert_eq!(mapped, Some("haiku-mapped".to_string()));
}
#[test]
fn test_opus_mapping() {
let provider = create_provider_with_mapping();
let body = json!({"model": "claude-opus-4-5"});
let (result, _, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "opus-mapped");
assert_eq!(mapped, Some("opus-mapped".to_string()));
}
#[test]
fn test_thinking_mode() {
let provider = create_provider_with_mapping();
let body = json!({
"model": "claude-sonnet-4-5",
"thinking": {"type": "enabled"}
});
let (result, _, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "reasoning-model");
assert_eq!(mapped, Some("reasoning-model".to_string()));
}
#[test]
fn test_thinking_disabled() {
let provider = create_provider_with_mapping();
let body = json!({
"model": "claude-sonnet-4-5",
"thinking": {"type": "disabled"}
});
let (result, _, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "sonnet-mapped");
assert_eq!(mapped, Some("sonnet-mapped".to_string()));
}
#[test]
fn test_unknown_model_uses_default() {
let provider = create_provider_with_mapping();
let body = json!({"model": "some-unknown-model"});
let (result, _, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "default-model");
assert_eq!(mapped, Some("default-model".to_string()));
}
#[test]
fn test_no_mapping_configured() {
let provider = create_provider_without_mapping();
let body = json!({"model": "claude-sonnet-4-5"});
let (result, original, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "claude-sonnet-4-5");
assert_eq!(original, Some("claude-sonnet-4-5".to_string()));
assert!(mapped.is_none());
}
#[test]
fn test_case_insensitive() {
let provider = create_provider_with_mapping();
let body = json!({"model": "Claude-SONNET-4-5"});
let (result, _, mapped) = apply_model_mapping(body, &provider);
assert_eq!(result["model"], "sonnet-mapped");
assert_eq!(mapped, Some("sonnet-mapped".to_string()));
}
}
+81 -188
View File
@@ -5,7 +5,7 @@
use crate::database::Database; use crate::database::Database;
use crate::error::AppError; use crate::error::AppError;
use crate::provider::Provider; use crate::provider::Provider;
use crate::proxy::circuit_breaker::{AllowResult, CircuitBreaker, CircuitBreakerConfig}; use crate::proxy::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig};
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use tokio::sync::RwLock; use tokio::sync::RwLock;
@@ -30,107 +30,84 @@ impl ProviderRouter {
/// 选择可用的供应商(支持故障转移) /// 选择可用的供应商(支持故障转移)
/// ///
/// 返回按优先级排序的可用供应商列表: /// 返回按优先级排序的可用供应商列表:
/// - 故障转移关闭时:仅返回当前供应商 /// 1. 当前供应商(is_current=true)始终第一位
/// - 故障转移开启时:完全按照故障转移队列顺序返回,忽略当前供应商设置 /// 2. 故障转移队列中的其他供应商(按 queue_order 排序)
/// 3. 只返回熔断器未打开的供应商
pub async fn select_providers(&self, app_type: &str) -> Result<Vec<Provider>, AppError> { pub async fn select_providers(&self, app_type: &str) -> Result<Vec<Provider>, AppError> {
let mut result = Vec::new(); let mut result = Vec::new();
let mut total_providers = 0usize; let all_providers = self.db.get_all_providers(app_type)?;
let mut circuit_open_count = 0usize;
// 检查该应用的自动故障转移开关是否开启(从 proxy_config 表读取) // 1. 当前供应商始终第一位
let auto_failover_enabled = match self.db.get_proxy_config_for_app(app_type).await { if let Some(current_id) = self.db.get_current_provider(app_type)? {
Ok(config) => { if let Some(current) = all_providers.get(&current_id) {
let enabled = config.auto_failover_enabled; let circuit_key = format!("{}:{}", app_type, current.id);
log::info!("[{app_type}] Failover enabled from proxy_config: {enabled}");
enabled
}
Err(e) => {
log::error!(
"[{app_type}] Failed to read proxy_config for auto_failover_enabled: {e}, defaulting to disabled"
);
false
}
};
if auto_failover_enabled {
// 故障转移开启:使用 in_failover_queue 标记的供应商,按 sort_index 排序
let failover_providers = self.db.get_failover_providers(app_type)?;
total_providers = failover_providers.len();
log::debug!("[{app_type}] Found {total_providers} failover queue provider(s)");
log::info!(
"[{app_type}] Failover enabled, using queue order ({total_providers} items)"
);
for provider in failover_providers {
// 检查熔断器状态
let circuit_key = format!("{}:{}", app_type, provider.id);
let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; let breaker = self.get_or_create_circuit_breaker(&circuit_key).await;
let state = breaker.get_state().await;
if breaker.is_available().await { if breaker.is_available().await {
log::debug!(
"[{}] Queue provider available: {} ({}) (state: {:?})",
app_type,
provider.name,
provider.id,
state
);
log::info!( log::info!(
"[{}] Queue provider available: {} ({}) at sort_index {:?}", "[{}] Current provider available: {} ({})",
app_type,
provider.name,
provider.id,
provider.sort_index
);
result.push(provider);
} else {
circuit_open_count += 1;
log::debug!(
"[{}] Queue provider {} circuit breaker open (state: {:?}), skipping",
app_type,
provider.name,
state
);
}
}
} else {
// 故障转移关闭:仅使用当前供应商,跳过熔断器检查
// 原因:单 Provider 场景下,熔断器打开会导致所有请求失败,用户体验差
log::info!("[{app_type}] Failover disabled, using current provider only (circuit breaker bypassed)");
if let Some(current_id) = self.db.get_current_provider(app_type)? {
if let Some(current) = self.db.get_provider_by_id(&current_id, app_type)? {
log::info!(
"[{}] Current provider: {} ({})",
app_type, app_type,
current.name, current.name,
current.id current.id
); );
total_providers = 1; result.push(current.clone());
result.push(current);
} else { } else {
log::debug!( log::warn!(
"[{app_type}] Current provider id {current_id} not found in database" "[{}] Current provider {} circuit breaker open, checking failover queue",
app_type,
current.name
);
}
}
}
// 2. 获取故障转移队列中的供应商
let queue = self.db.get_failover_queue(app_type)?;
for item in queue {
// 跳过已添加的当前供应商
if result.iter().any(|p| p.id == item.provider_id) {
continue;
}
// 跳过禁用的队列项
if !item.enabled {
continue;
}
// 获取供应商信息
if let Some(provider) = all_providers.get(&item.provider_id) {
// 检查熔断器状态
let circuit_key = format!("{}:{}", app_type, provider.id);
let breaker = self.get_or_create_circuit_breaker(&circuit_key).await;
if breaker.is_available().await {
log::info!(
"[{}] Failover provider available: {} ({}) at queue position {}",
app_type,
provider.name,
provider.id,
item.queue_order
);
result.push(provider.clone());
} else {
log::debug!(
"[{}] Failover provider {} circuit breaker open, skipping",
app_type,
provider.name
); );
} }
} else {
log::debug!("[{app_type}] No current provider configured");
} }
} }
if result.is_empty() { if result.is_empty() {
// 区分两种情况:全部熔断 vs 未配置供应商 return Err(AppError::Config(format!(
if total_providers > 0 && circuit_open_count == total_providers { "No available provider for {app_type} (all circuit breakers open or no providers configured)"
log::warn!("[{app_type}] 所有 {total_providers} 个供应商均已熔断,无可用渠道"); )));
return Err(AppError::AllProvidersCircuitOpen);
} else {
log::warn!("[{app_type}] 未配置供应商或故障转移队列为空");
return Err(AppError::NoProvidersConfigured);
}
} }
log::info!( log::info!(
"[{}] Provider chain: {} provider(s) available", "[{}] Failover chain: {} provider(s) available",
app_type, app_type,
result.len() result.len()
); );
@@ -146,7 +123,7 @@ impl ProviderRouter {
/// ///
/// 注意:调用方必须在请求结束后通过 `record_result()` 释放 HalfOpen 名额, /// 注意:调用方必须在请求结束后通过 `record_result()` 释放 HalfOpen 名额,
/// 否则会导致该 Provider 长时间无法进入探测状态。 /// 否则会导致该 Provider 长时间无法进入探测状态。
pub async fn allow_provider_request(&self, provider_id: &str, app_type: &str) -> AllowResult { pub async fn allow_provider_request(&self, provider_id: &str, app_type: &str) -> bool {
let circuit_key = format!("{app_type}:{provider_id}"); let circuit_key = format!("{app_type}:{provider_id}");
let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; let breaker = self.get_or_create_circuit_breaker(&circuit_key).await;
breaker.allow_request().await breaker.allow_request().await
@@ -157,30 +134,22 @@ impl ProviderRouter {
&self, &self,
provider_id: &str, provider_id: &str,
app_type: &str, app_type: &str,
used_half_open_permit: bool,
success: bool, success: bool,
error_msg: Option<String>, error_msg: Option<String>,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
// 1. 按应用独立获取熔断器配置(用于更新健康状态和判断是否禁用) // 1. 获取熔断器配置(用于更新健康状态和判断是否禁用)
let failure_threshold = match self.db.get_proxy_config_for_app(app_type).await { let config = self.db.get_circuit_breaker_config().await.ok();
Ok(app_config) => app_config.circuit_failure_threshold, let failure_threshold = config.map(|c| c.failure_threshold).unwrap_or(5);
Err(e) => {
log::warn!(
"Failed to load circuit config for {app_type}, using default threshold: {e}"
);
5 // 默认值
}
};
// 2. 更新熔断器状态 // 2. 更新熔断器状态
let circuit_key = format!("{app_type}:{provider_id}"); let circuit_key = format!("{app_type}:{provider_id}");
let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; let breaker = self.get_or_create_circuit_breaker(&circuit_key).await;
if success { if success {
breaker.record_success(used_half_open_permit).await; breaker.record_success().await;
log::debug!("Provider {provider_id} request succeeded"); log::debug!("Provider {provider_id} request succeeded");
} else { } else {
breaker.record_failure(used_half_open_permit).await; breaker.record_failure().await;
log::warn!( log::warn!(
"Provider {} request failed: {}", "Provider {} request failed: {}",
provider_id, provider_id,
@@ -267,34 +236,12 @@ impl ProviderRouter {
return breaker.clone(); return breaker.clone();
} }
// 从 key 中提取 app_type (格式: "app_type:provider_id") // 从数据库加载配置
let app_type = key.split(':').next().unwrap_or("claude"); let config = self
.db
// 按应用独立读取熔断器配置 .get_circuit_breaker_config()
let config = match self.db.get_proxy_config_for_app(app_type).await { .await
Ok(app_config) => { .unwrap_or_default();
log::debug!(
"Loading circuit breaker config for {key} (app={app_type}): \
failure_threshold={}, success_threshold={}, timeout={}s",
app_config.circuit_failure_threshold,
app_config.circuit_success_threshold,
app_config.circuit_timeout_seconds
);
crate::proxy::circuit_breaker::CircuitBreakerConfig {
failure_threshold: app_config.circuit_failure_threshold,
success_threshold: app_config.circuit_success_threshold,
timeout_seconds: app_config.circuit_timeout_seconds as u64,
error_rate_threshold: app_config.circuit_error_rate_threshold,
min_requests: app_config.circuit_min_requests,
}
}
Err(e) => {
log::warn!(
"Failed to load circuit breaker config for {key} (app={app_type}): {e}, using default"
);
crate::proxy::circuit_breaker::CircuitBreakerConfig::default()
}
};
log::debug!("Creating new circuit breaker for {key} with config: {config:?}"); log::debug!("Creating new circuit breaker for {key} with config: {config:?}");
@@ -316,68 +263,16 @@ mod tests {
let db = Arc::new(Database::memory().unwrap()); let db = Arc::new(Database::memory().unwrap());
let router = ProviderRouter::new(db); let router = ProviderRouter::new(db);
// 测试创建熔断器
let breaker = router.get_or_create_circuit_breaker("claude:test").await; let breaker = router.get_or_create_circuit_breaker("claude:test").await;
assert!(breaker.allow_request().await.allowed); assert!(breaker.allow_request().await);
} }
#[tokio::test] #[tokio::test]
async fn test_failover_disabled_uses_current_provider() { async fn select_providers_does_not_consume_half_open_permit() {
let db = Arc::new(Database::memory().unwrap());
let provider_a =
Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None);
let provider_b =
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
db.save_provider("claude", &provider_a).unwrap();
db.save_provider("claude", &provider_b).unwrap();
db.set_current_provider("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap();
let router = ProviderRouter::new(db.clone());
let providers = router.select_providers("claude").await.unwrap();
assert_eq!(providers.len(), 1);
assert_eq!(providers[0].id, "a");
}
#[tokio::test]
async fn test_failover_enabled_uses_queue_order() {
let db = Arc::new(Database::memory().unwrap());
// 设置 sort_index 来控制顺序:b=1, a=2
let mut provider_a =
Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None);
provider_a.sort_index = Some(2);
let mut provider_b =
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
provider_b.sort_index = Some(1);
db.save_provider("claude", &provider_a).unwrap();
db.save_provider("claude", &provider_b).unwrap();
db.set_current_provider("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap();
db.add_to_failover_queue("claude", "a").unwrap();
// 启用自动故障转移(使用新的 proxy_config API
let mut config = db.get_proxy_config_for_app("claude").await.unwrap();
config.auto_failover_enabled = true;
db.update_proxy_config_for_app(config).await.unwrap();
let router = ProviderRouter::new(db.clone());
let providers = router.select_providers("claude").await.unwrap();
assert_eq!(providers.len(), 2);
// 按 sort_index 排序:b(1) 在前,a(2) 在后
assert_eq!(providers[0].id, "b");
assert_eq!(providers[1].id, "a");
}
#[tokio::test]
async fn test_select_providers_does_not_consume_half_open_permit() {
let db = Arc::new(Database::memory().unwrap()); let db = Arc::new(Database::memory().unwrap());
// 配置:让熔断器 Open 后立刻进入 HalfOpentimeout_seconds=0),并用 1 次失败就打开熔断器
db.update_circuit_breaker_config(&CircuitBreakerConfig { db.update_circuit_breaker_config(&CircuitBreakerConfig {
failure_threshold: 1, failure_threshold: 1,
timeout_seconds: 0, timeout_seconds: 0,
@@ -386,6 +281,7 @@ mod tests {
.await .await
.unwrap(); .unwrap();
// 准备 2 个 ProviderA(当前)+ B(队列)
let provider_a = let provider_a =
Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None); Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None);
let provider_b = let provider_b =
@@ -393,25 +289,22 @@ mod tests {
db.save_provider("claude", &provider_a).unwrap(); db.save_provider("claude", &provider_a).unwrap();
db.save_provider("claude", &provider_b).unwrap(); db.save_provider("claude", &provider_b).unwrap();
db.set_current_provider("claude", "a").unwrap();
db.add_to_failover_queue("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap(); db.add_to_failover_queue("claude", "b").unwrap();
// 启用自动故障转移(使用新的 proxy_config API
let mut config = db.get_proxy_config_for_app("claude").await.unwrap();
config.auto_failover_enabled = true;
db.update_proxy_config_for_app(config).await.unwrap();
let router = ProviderRouter::new(db.clone()); let router = ProviderRouter::new(db.clone());
// 让 B 进入 Open 状态(failure_threshold=1
router router
.record_result("b", "claude", false, false, Some("fail".to_string())) .record_result("b", "claude", false, Some("fail".to_string()))
.await .await
.unwrap(); .unwrap();
// select_providers 只做“可用性判断”,不应占用 HalfOpen 探测名额
let providers = router.select_providers("claude").await.unwrap(); let providers = router.select_providers("claude").await.unwrap();
assert_eq!(providers.len(), 2); assert_eq!(providers.len(), 2);
assert!(router.allow_provider_request("b", "claude").await.allowed); // 如果 select_providers 错误地消耗了 HalfOpen 名额,这里会返回 false(被限流拒绝)
assert!(router.allow_provider_request("b", "claude").await);
} }
} }
+1 -1
View File
@@ -87,7 +87,7 @@ pub trait ProviderAdapter: Send + Sync {
/// 是否需要格式转换 /// 是否需要格式转换
/// ///
/// 默认返回 `false`(透传模式)。 /// 默认返回 `false`(透传模式)。
/// 仅当供应商需要格式转换时(如 Claude + OpenRouter 旧 OpenAI 兼容接口)才返回 `true`。 /// 仅当供应商需要格式转换时(如 Claude + OpenRouter)才返回 `true`。
/// ///
/// # Arguments /// # Arguments
/// * `provider` - Provider 配置 /// * `provider` - Provider 配置
+19 -73
View File
@@ -5,7 +5,7 @@
//! ## 认证模式 //! ## 认证模式
//! - **Claude**: Anthropic 官方 API (x-api-key + anthropic-version) //! - **Claude**: Anthropic 官方 API (x-api-key + anthropic-version)
//! - **ClaudeAuth**: 中转服务 (仅 Bearer 认证,无 x-api-key) //! - **ClaudeAuth**: 中转服务 (仅 Bearer 认证,无 x-api-key)
//! - **OpenRouter**: 已支持 Claude Code 兼容接口,默认透传(保留旧转换逻辑备用) //! - **OpenRouter**: 需要 Anthropic ↔ OpenAI 格式转换
use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType}; use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType};
use crate::provider::Provider; use crate::provider::Provider;
@@ -28,8 +28,10 @@ impl ClaudeAdapter {
/// - Claude: 默认 Anthropic 官方 /// - Claude: 默认 Anthropic 官方
pub fn provider_type(&self, provider: &Provider) -> ProviderType { pub fn provider_type(&self, provider: &Provider) -> ProviderType {
// 检测 OpenRouter // 检测 OpenRouter
if self.is_openrouter(provider) { if let Ok(base_url) = self.extract_base_url(provider) {
return ProviderType::OpenRouter; if base_url.contains("openrouter.ai") {
return ProviderType::OpenRouter;
}
} }
// 检测 ClaudeAuth (仅 Bearer 认证) // 检测 ClaudeAuth (仅 Bearer 认证)
@@ -48,24 +50,6 @@ impl ClaudeAdapter {
false false
} }
/// 检测 OpenRouter 是否启用兼容模式
fn is_openrouter_compat_enabled(&self, provider: &Provider) -> bool {
if !self.is_openrouter(provider) {
return false;
}
let raw = provider.settings_config.get("openrouter_compat_mode");
match raw {
Some(serde_json::Value::Bool(enabled)) => *enabled,
Some(serde_json::Value::Number(num)) => num.as_i64().unwrap_or(0) != 0,
Some(serde_json::Value::String(value)) => {
let normalized = value.trim().to_lowercase();
normalized == "true" || normalized == "1"
}
_ => true,
}
}
/// 检测是否为仅 Bearer 认证模式 /// 检测是否为仅 Bearer 认证模式
fn is_bearer_only_mode(&self, provider: &Provider) -> bool { fn is_bearer_only_mode(&self, provider: &Provider) -> bool {
// 检查 settings_config 中的 auth_mode // 检查 settings_config 中的 auth_mode
@@ -103,14 +87,6 @@ impl ClaudeAdapter {
log::debug!("[Claude] 使用 ANTHROPIC_AUTH_TOKEN"); log::debug!("[Claude] 使用 ANTHROPIC_AUTH_TOKEN");
return Some(key.to_string()); return Some(key.to_string());
} }
if let Some(key) = env
.get("ANTHROPIC_API_KEY")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
{
log::debug!("[Claude] 使用 ANTHROPIC_API_KEY");
return Some(key.to_string());
}
// OpenRouter key // OpenRouter key
if let Some(key) = env if let Some(key) = env
.get("OPENROUTER_API_KEY") .get("OPENROUTER_API_KEY")
@@ -210,13 +186,12 @@ impl ProviderAdapter for ClaudeAdapter {
} }
fn build_url(&self, base_url: &str, endpoint: &str) -> String { fn build_url(&self, base_url: &str, endpoint: &str) -> String {
// NOTE: // OpenRouter 使用 /v1/chat/completions
// 过去 OpenRouter 只有 OpenAI Chat Completions 兼容接口,需要把 Claude 的 `/v1/messages` if base_url.contains("openrouter.ai") {
// 映射到 `/v1/chat/completions`,并做 Anthropic ↔ OpenAI 的格式转换。 return format!("{}/v1/chat/completions", base_url.trim_end_matches('/'));
// }
// 现在 OpenRouter 已推出 Claude Code 兼容接口,因此默认直接透传 endpoint。
// 如需回退旧逻辑,可在 forwarder 中根据 needs_transform 改写 endpoint。
// Anthropic 直连
format!( format!(
"{}/{}", "{}/{}",
base_url.trim_end_matches('/'), base_url.trim_end_matches('/'),
@@ -232,24 +207,19 @@ impl ProviderAdapter for ClaudeAdapter {
.header("x-api-key", &auth.api_key) .header("x-api-key", &auth.api_key)
.header("anthropic-version", "2023-06-01"), .header("anthropic-version", "2023-06-01"),
// ClaudeAuth 中转服务: 仅 Bearer,无 x-api-key // ClaudeAuth 中转服务: 仅 Bearer,无 x-api-key
AuthStrategy::ClaudeAuth => request AuthStrategy::ClaudeAuth => {
.header("Authorization", format!("Bearer {}", auth.api_key)) request.header("Authorization", format!("Bearer {}", auth.api_key))
.header("anthropic-version", "2023-06-01"), }
// OpenRouter: Bearer // OpenRouter: Bearer
AuthStrategy::Bearer => request AuthStrategy::Bearer => {
.header("Authorization", format!("Bearer {}", auth.api_key)) request.header("Authorization", format!("Bearer {}", auth.api_key))
.header("anthropic-version", "2023-06-01"), }
_ => request, _ => request,
} }
} }
fn needs_transform(&self, _provider: &Provider) -> bool { fn needs_transform(&self, provider: &Provider) -> bool {
// NOTE: self.is_openrouter(provider)
// OpenRouter 已推出 Claude Code 兼容接口(可直接处理 `/v1/messages`),默认不再启用
// Anthropic ↔ OpenAI 的格式转换。
//
// 如果未来需要回退到旧的 OpenAI Chat Completions 方案,可恢复下面这行:
self.is_openrouter_compat_enabled(_provider)
} }
fn transform_request( fn transform_request(
@@ -283,7 +253,6 @@ mod tests {
meta: None, meta: None,
icon: None, icon: None,
icon_color: None, icon_color: None,
in_failover_queue: false,
} }
} }
@@ -315,21 +284,6 @@ mod tests {
assert_eq!(auth.strategy, AuthStrategy::Anthropic); assert_eq!(auth.strategy, AuthStrategy::Anthropic);
} }
#[test]
fn test_extract_auth_anthropic_api_key() {
let adapter = ClaudeAdapter::new();
let provider = create_provider(json!({
"env": {
"ANTHROPIC_BASE_URL": "https://api.anthropic.com",
"ANTHROPIC_API_KEY": "sk-ant-test-key"
}
}));
let auth = adapter.extract_auth(&provider).unwrap();
assert_eq!(auth.api_key, "sk-ant-test-key");
assert_eq!(auth.strategy, AuthStrategy::Anthropic);
}
#[test] #[test]
fn test_extract_auth_openrouter() { fn test_extract_auth_openrouter() {
let adapter = ClaudeAdapter::new(); let adapter = ClaudeAdapter::new();
@@ -424,7 +378,7 @@ mod tests {
fn test_build_url_openrouter() { fn test_build_url_openrouter() {
let adapter = ClaudeAdapter::new(); let adapter = ClaudeAdapter::new();
let url = adapter.build_url("https://openrouter.ai/api", "/v1/messages"); let url = adapter.build_url("https://openrouter.ai/api", "/v1/messages");
assert_eq!(url, "https://openrouter.ai/api/v1/messages"); assert_eq!(url, "https://openrouter.ai/api/v1/chat/completions");
} }
#[test] #[test]
@@ -444,13 +398,5 @@ mod tests {
} }
})); }));
assert!(adapter.needs_transform(&openrouter_provider)); assert!(adapter.needs_transform(&openrouter_provider));
let openrouter_disabled = create_provider(json!({
"env": {
"ANTHROPIC_BASE_URL": "https://openrouter.ai/api"
},
"openrouter_compat_mode": false
}));
assert!(!adapter.needs_transform(&openrouter_disabled));
} }
} }
-1
View File
@@ -174,7 +174,6 @@ mod tests {
meta: None, meta: None,
icon: None, icon: None,
icon_color: None, icon_color: None,
in_failover_queue: false,
} }
} }
-1
View File
@@ -250,7 +250,6 @@ mod tests {
meta: None, meta: None,
icon: None, icon: None,
icon_color: None, icon_color: None,
in_failover_queue: false,
} }
} }
+4 -9
View File
@@ -48,21 +48,17 @@ pub enum ProviderType {
Gemini, Gemini,
/// Google Gemini CLI (OAuth Bearer) /// Google Gemini CLI (OAuth Bearer)
GeminiCli, GeminiCli,
/// OpenRouter(已支持 Claude Code 兼容接口,默认透传;保留旧转换逻辑备用) /// OpenRouter (需要 Anthropic ↔ OpenAI 格式转换)
OpenRouter, OpenRouter,
} }
impl ProviderType { impl ProviderType {
/// 是否需要格式转换 /// 是否需要格式转换
/// ///
/// 过去 OpenRouter 需要将 Anthropic 格式转换为 OpenAI 格式 /// OpenRouter 需要将 Anthropic 格式转换为 OpenAI 格式
/// 现在默认关闭转换(因为 OpenRouter 已支持 Claude Code 兼容接口)。
#[allow(dead_code)] #[allow(dead_code)]
pub fn needs_transform(&self) -> bool { pub fn needs_transform(&self) -> bool {
match self { matches!(self, ProviderType::OpenRouter)
ProviderType::OpenRouter => false,
_ => false,
}
} }
/// 获取默认端点 /// 获取默认端点
@@ -209,7 +205,6 @@ mod tests {
meta: None, meta: None,
icon: None, icon: None,
icon_color: None, icon_color: None,
in_failover_queue: false,
} }
} }
@@ -220,7 +215,7 @@ mod tests {
assert!(!ProviderType::Codex.needs_transform()); assert!(!ProviderType::Codex.needs_transform());
assert!(!ProviderType::Gemini.needs_transform()); assert!(!ProviderType::Gemini.needs_transform());
assert!(!ProviderType::GeminiCli.needs_transform()); assert!(!ProviderType::GeminiCli.needs_transform());
assert!(!ProviderType::OpenRouter.needs_transform()); assert!(ProviderType::OpenRouter.needs_transform());
} }
#[test] #[test]
@@ -394,7 +394,6 @@ mod tests {
meta: None, meta: None,
icon: None, icon: None,
icon_color: None, icon_color: None,
in_failover_queue: false,
} }
} }
+14 -126
View File
@@ -3,11 +3,8 @@
//! 统一处理流式和非流式 API 响应 //! 统一处理流式和非流式 API 响应
use super::{ use super::{
handler_config::UsageParserConfig, handler_config::UsageParserConfig, handler_context::RequestContext, server::ProxyState,
handler_context::{RequestContext, StreamingTimeoutConfig}, usage::parser::TokenUsage, ProxyError,
server::ProxyState,
usage::parser::TokenUsage,
ProxyError,
}; };
use axum::response::Response; use axum::response::Response;
use bytes::Bytes; use bytes::Bytes;
@@ -20,7 +17,6 @@ use std::{
atomic::{AtomicBool, Ordering}, atomic::{AtomicBool, Ordering},
Arc, Arc,
}, },
time::Duration,
}; };
use tokio::sync::Mutex; use tokio::sync::Mutex;
@@ -64,12 +60,8 @@ pub async fn handle_streaming(
// 创建使用量收集器 // 创建使用量收集器
let usage_collector = create_usage_collector(ctx, state, status.as_u16(), parser_config); let usage_collector = create_usage_collector(ctx, state, status.as_u16(), parser_config);
// 获取流式超时配置 // 创建带日志的透传流
let timeout_config = ctx.streaming_timeout_config(); let logged_stream = create_logged_passthrough_stream(stream, ctx.tag, Some(usage_collector));
// 创建带日志和超时的透传流
let logged_stream =
create_logged_passthrough_stream(stream, ctx.tag, Some(usage_collector), timeout_config);
let body = axum::body::Body::from_stream(logged_stream); let body = axum::body::Body::from_stream(logged_stream);
builder.body(body).unwrap() builder.body(body).unwrap()
@@ -101,30 +93,13 @@ pub async fn handle_non_streaming(
// 解析使用量 // 解析使用量
if let Some(usage) = (parser_config.response_parser)(&json_value) { if let Some(usage) = (parser_config.response_parser)(&json_value) {
// 优先使用 usage 中解析出的模型名称,其次使用响应中的 model 字段,最后回退到请求模型
let model = if let Some(ref m) = usage.model {
m.clone()
} else if let Some(m) = json_value.get("model").and_then(|m| m.as_str()) {
m.to_string()
} else {
ctx.request_model.clone()
};
spawn_log_usage(state, ctx, usage, &model, status.as_u16(), false);
} else {
let model = json_value let model = json_value
.get("model") .get("model")
.and_then(|m| m.as_str()) .and_then(|m| m.as_str())
.unwrap_or(&ctx.request_model) .unwrap_or(&ctx.request_model);
.to_string();
spawn_log_usage( spawn_log_usage(state, ctx, usage, model, status.as_u16(), false);
state, } else {
ctx,
TokenUsage::default(),
&model,
status.as_u16(),
false,
);
log::debug!( log::debug!(
"[{}] 未能解析 usage 信息,跳过记录", "[{}] 未能解析 usage 信息,跳过记录",
parser_config.app_type_str parser_config.app_type_str
@@ -136,14 +111,6 @@ pub async fn handle_non_streaming(
ctx.tag, ctx.tag,
body_bytes.len() body_bytes.len()
); );
spawn_log_usage(
state,
ctx,
TokenUsage::default(),
&ctx.request_model,
status.as_u16(),
false,
);
} }
log::info!("[{}] ====== 请求结束 ======", ctx.tag); log::info!("[{}] ====== 请求结束 ======", ctx.tag);
@@ -264,7 +231,6 @@ fn create_usage_collector(
let start_time = ctx.start_time; let start_time = ctx.start_time;
let stream_parser = parser_config.stream_parser; let stream_parser = parser_config.stream_parser;
let model_extractor = parser_config.model_extractor; let model_extractor = parser_config.model_extractor;
let session_id = ctx.session_id.clone();
SseUsageCollector::new(start_time, move |events, first_token_ms| { SseUsageCollector::new(start_time, move |events, first_token_ms| {
if let Some(usage) = stream_parser(&events) { if let Some(usage) = stream_parser(&events) {
@@ -273,7 +239,6 @@ fn create_usage_collector(
let state = state.clone(); let state = state.clone();
let provider_id = provider_id.clone(); let provider_id = provider_id.clone();
let session_id = session_id.clone();
tokio::spawn(async move { tokio::spawn(async move {
log_usage_internal( log_usage_internal(
@@ -286,32 +251,10 @@ fn create_usage_collector(
first_token_ms, first_token_ms,
true, // is_streaming true, // is_streaming
status_code, status_code,
Some(session_id),
) )
.await; .await;
}); });
} else { } else {
let model = model_extractor(&events, &request_model);
let latency_ms = start_time.elapsed().as_millis() as u64;
let state = state.clone();
let provider_id = provider_id.clone();
let session_id = session_id.clone();
tokio::spawn(async move {
log_usage_internal(
&state,
&provider_id,
app_type_str,
&model,
TokenUsage::default(),
latency_ms,
first_token_ms,
true, // is_streaming
status_code,
Some(session_id),
)
.await;
});
log::debug!("[{tag}] 流式响应缺少 usage 统计,跳过消费记录"); log::debug!("[{tag}] 流式响应缺少 usage 统计,跳过消费记录");
} }
}) })
@@ -331,7 +274,6 @@ fn spawn_log_usage(
let app_type_str = ctx.app_type_str.to_string(); let app_type_str = ctx.app_type_str.to_string();
let model = model.to_string(); let model = model.to_string();
let latency_ms = ctx.latency_ms(); let latency_ms = ctx.latency_ms();
let session_id = ctx.session_id.clone();
tokio::spawn(async move { tokio::spawn(async move {
log_usage_internal( log_usage_internal(
@@ -344,7 +286,6 @@ fn spawn_log_usage(
None, None,
is_streaming, is_streaming,
status_code, status_code,
Some(session_id),
) )
.await; .await;
}); });
@@ -362,7 +303,6 @@ async fn log_usage_internal(
first_token_ms: Option<u64>, first_token_ms: Option<u64>,
is_streaming: bool, is_streaming: bool,
status_code: u16, status_code: u16,
session_id: Option<String>,
) { ) {
use super::usage::logger::UsageLogger; use super::usage::logger::UsageLogger;
@@ -386,15 +326,6 @@ async fn log_usage_internal(
let request_id = uuid::Uuid::new_v4().to_string(); let request_id = uuid::Uuid::new_v4().to_string();
log::debug!(
"[{app_type}] 记录请求日志: id={request_id}, provider={provider_id}, model={model}, streaming={is_streaming}, status={status_code}, latency_ms={latency_ms}, first_token_ms={first_token_ms:?}, session={}, input={}, output={}, cache_read={}, cache_creation={}",
session_id.as_deref().unwrap_or("none"),
usage.input_tokens,
usage.output_tokens,
usage.cache_read_tokens,
usage.cache_creation_tokens
);
if let Err(e) = logger.log_with_calculation( if let Err(e) = logger.log_with_calculation(
request_id, request_id,
provider_id.to_string(), provider_id.to_string(),
@@ -405,7 +336,7 @@ async fn log_usage_internal(
latency_ms, latency_ms,
first_token_ms, first_token_ms,
status_code, status_code,
session_id, None,
None, // provider_type None, // provider_type
is_streaming, is_streaming,
) { ) {
@@ -413,60 +344,21 @@ async fn log_usage_internal(
} }
} }
/// 创建带日志记录和超时控制的透传流 /// 创建带日志记录的透传流
pub fn create_logged_passthrough_stream( pub fn create_logged_passthrough_stream(
stream: impl Stream<Item = Result<Bytes, std::io::Error>> + Send + 'static, stream: impl Stream<Item = Result<Bytes, std::io::Error>> + Send + 'static,
tag: &'static str, tag: &'static str,
usage_collector: Option<SseUsageCollector>, usage_collector: Option<SseUsageCollector>,
timeout_config: StreamingTimeoutConfig,
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send { ) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
async_stream::stream! { async_stream::stream! {
let mut buffer = String::new(); let mut buffer = String::new();
let mut collector = usage_collector; let mut collector = usage_collector;
let mut is_first_chunk = true;
// 超时配置
let first_byte_timeout = if timeout_config.first_byte_timeout > 0 {
Some(Duration::from_secs(timeout_config.first_byte_timeout))
} else {
None
};
let idle_timeout = if timeout_config.idle_timeout > 0 {
Some(Duration::from_secs(timeout_config.idle_timeout))
} else {
None
};
tokio::pin!(stream); tokio::pin!(stream);
loop { while let Some(chunk) = stream.next().await {
// 选择超时时间:首字节超时或静默期超时 match chunk {
let timeout_duration = if is_first_chunk { Ok(bytes) => {
first_byte_timeout
} else {
idle_timeout
};
let chunk_result = match timeout_duration {
Some(duration) => {
match tokio::time::timeout(duration, stream.next()).await {
Ok(Some(chunk)) => Some(chunk),
Ok(None) => None, // 流结束
Err(_) => {
// 超时
let timeout_type = if is_first_chunk { "首字节" } else { "静默期" };
log::error!("[{tag}] 流式响应{}超时 ({}秒)", timeout_type, duration.as_secs());
yield Err(std::io::Error::other(format!("流式响应{timeout_type}超时")));
break;
}
}
}
None => stream.next().await, // 无超时限制
};
match chunk_result {
Some(Ok(bytes)) => {
is_first_chunk = false;
let text = String::from_utf8_lossy(&bytes); let text = String::from_utf8_lossy(&bytes);
buffer.push_str(&text); buffer.push_str(&text);
@@ -502,15 +394,11 @@ pub fn create_logged_passthrough_stream(
yield Ok(bytes); yield Ok(bytes);
} }
Some(Err(e)) => { Err(e) => {
log::error!("[{tag}] 流错误: {e}"); log::error!("[{tag}] 流错误: {e}");
yield Err(std::io::Error::other(e.to_string())); yield Err(std::io::Error::other(e.to_string()));
break; break;
} }
None => {
// 流正常结束
break;
}
} }
} }
-7
View File
@@ -191,23 +191,16 @@ impl ProxyServer {
.route("/v1/messages", post(handlers::handle_messages)) .route("/v1/messages", post(handlers::handle_messages))
.route("/claude/v1/messages", post(handlers::handle_messages)) .route("/claude/v1/messages", post(handlers::handle_messages))
// OpenAI Chat Completions API (Codex CLI,支持带前缀和不带前缀) // OpenAI Chat Completions API (Codex CLI,支持带前缀和不带前缀)
.route("/chat/completions", post(handlers::handle_chat_completions))
.route( .route(
"/v1/chat/completions", "/v1/chat/completions",
post(handlers::handle_chat_completions), post(handlers::handle_chat_completions),
) )
.route(
"/v1/v1/chat/completions",
post(handlers::handle_chat_completions),
)
.route( .route(
"/codex/v1/chat/completions", "/codex/v1/chat/completions",
post(handlers::handle_chat_completions), post(handlers::handle_chat_completions),
) )
// OpenAI Responses API (Codex CLI,支持带前缀和不带前缀) // OpenAI Responses API (Codex CLI,支持带前缀和不带前缀)
.route("/responses", post(handlers::handle_responses))
.route("/v1/responses", post(handlers::handle_responses)) .route("/v1/responses", post(handlers::handle_responses))
.route("/v1/v1/responses", post(handlers::handle_responses))
.route("/codex/v1/responses", post(handlers::handle_responses)) .route("/codex/v1/responses", post(handlers::handle_responses))
// Gemini API (支持带前缀和不带前缀) // Gemini API (支持带前缀和不带前缀)
.route("/v1beta/*path", post(handlers::handle_gemini)) .route("/v1beta/*path", post(handlers::handle_gemini))
-269
View File
@@ -1,15 +1,7 @@
//! Proxy Session - 请求会话管理 //! Proxy Session - 请求会话管理
//! //!
//! 为每个代理请求创建会话上下文,在整个请求生命周期中跟踪状态和元数据。 //! 为每个代理请求创建会话上下文,在整个请求生命周期中跟踪状态和元数据。
//!
//! ## Session ID 提取
//!
//! 支持从客户端请求中提取 Session ID,用于关联同一对话的多个请求:
//! - Claude: 从 `metadata.user_id` (格式: `user_xxx_session_yyy`) 或 `metadata.session_id` 提取
//! - Codex: 从 `previous_response_id` 或 headers 中的 `session_id` 提取
//! - 其他: 生成新的 UUID
use axum::http::HeaderMap;
use std::time::Instant; use std::time::Instant;
use uuid::Uuid; use uuid::Uuid;
@@ -184,179 +176,6 @@ impl ProxySession {
} }
} }
// ============================================================================
// Session ID 提取器
// ============================================================================
/// Session ID 来源
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SessionIdSource {
/// 从 metadata.user_id 提取 (Claude)
MetadataUserId,
/// 从 metadata.session_id 提取
MetadataSessionId,
/// 从 headers 提取 (Codex)
Header,
/// 从 previous_response_id 提取 (Codex)
PreviousResponseId,
/// 新生成
Generated,
}
/// Session ID 提取结果
#[derive(Debug, Clone)]
pub struct SessionIdResult {
/// 提取或生成的 Session ID
pub session_id: String,
/// Session ID 来源
pub source: SessionIdSource,
/// 是否为客户端提供的 ID(非新生成)
pub client_provided: bool,
}
/// 从请求中提取或生成 Session ID
///
/// 轻量化实现,仅提取 session_id 用于日志记录,不做复杂的 Session 管理。
///
/// ## 提取优先级
///
/// ### Claude 请求
/// 1. `metadata.user_id` (格式: `user_xxx_session_yyy`) → 提取 `yyy` 部分
/// 2. `metadata.session_id` → 直接使用
/// 3. 生成新 UUID
///
/// ### Codex 请求
/// 1. Headers: `session_id` 或 `x-session-id`
/// 2. `metadata.session_id`
/// 3. `previous_response_id` (对话延续)
/// 4. 生成新 UUID
///
/// ## 示例
///
/// ```ignore
/// let result = extract_session_id(&headers, &body, "claude");
/// println!("Session ID: {} (from {:?})", result.session_id, result.source);
/// ```
pub fn extract_session_id(
headers: &HeaderMap,
body: &serde_json::Value,
client_format: &str,
) -> SessionIdResult {
// Codex 请求特殊处理
if client_format == "codex" || client_format == "openai" {
if let Some(result) = extract_codex_session(headers, body) {
return result;
}
}
// Claude 请求:从 metadata 提取
if let Some(result) = extract_from_metadata(body) {
return result;
}
// 兜底:生成新 Session ID
generate_new_session_id()
}
/// 提取 Codex Session ID
fn extract_codex_session(headers: &HeaderMap, body: &serde_json::Value) -> Option<SessionIdResult> {
// 1. 从 headers 提取
for header_name in &["session_id", "x-session-id"] {
if let Some(value) = headers.get(*header_name) {
if let Ok(session_id) = value.to_str() {
// Codex Session ID 通常较长(UUID 格式)
if session_id.len() > 20 {
return Some(SessionIdResult {
session_id: format!("codex_{session_id}"),
source: SessionIdSource::Header,
client_provided: true,
});
}
}
}
}
// 2. 从 body.metadata.session_id 提取
if let Some(session_id) = body
.get("metadata")
.and_then(|m| m.get("session_id"))
.and_then(|v| v.as_str())
{
if session_id.len() > 10 {
return Some(SessionIdResult {
session_id: format!("codex_{session_id}"),
source: SessionIdSource::MetadataSessionId,
client_provided: true,
});
}
}
// 3. 从 previous_response_id 提取(对话延续)
if let Some(prev_id) = body.get("previous_response_id").and_then(|v| v.as_str()) {
if prev_id.len() > 10 {
return Some(SessionIdResult {
session_id: format!("codex_{prev_id}"),
source: SessionIdSource::PreviousResponseId,
client_provided: true,
});
}
}
None
}
/// 从 metadata 提取 Session ID (Claude)
fn extract_from_metadata(body: &serde_json::Value) -> Option<SessionIdResult> {
let metadata = body.get("metadata")?;
// 1. 从 metadata.user_id 提取(格式: user_xxx_session_yyy
if let Some(user_id) = metadata.get("user_id").and_then(|v| v.as_str()) {
if let Some(session_id) = parse_session_from_user_id(user_id) {
return Some(SessionIdResult {
session_id,
source: SessionIdSource::MetadataUserId,
client_provided: true,
});
}
}
// 2. 直接从 metadata.session_id 提取
if let Some(session_id) = metadata.get("session_id").and_then(|v| v.as_str()) {
if !session_id.is_empty() {
return Some(SessionIdResult {
session_id: session_id.to_string(),
source: SessionIdSource::MetadataSessionId,
client_provided: true,
});
}
}
None
}
/// 从 user_id 解析 session_id
///
/// 格式: `user_identifier_session_actual_session_id`
fn parse_session_from_user_id(user_id: &str) -> Option<String> {
// 查找 "_session_" 分隔符
if let Some(pos) = user_id.find("_session_") {
let session_id = &user_id[pos + 9..]; // "_session_" 长度为 9
if !session_id.is_empty() {
return Some(session_id.to_string());
}
}
None
}
/// 生成新的 Session ID
fn generate_new_session_id() -> SessionIdResult {
SessionIdResult {
session_id: Uuid::new_v4().to_string(),
source: SessionIdSource::Generated,
client_provided: false,
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -476,92 +295,4 @@ mod tests {
assert_eq!(ClientFormat::GeminiCli.as_str(), "gemini_cli"); assert_eq!(ClientFormat::GeminiCli.as_str(), "gemini_cli");
assert_eq!(ClientFormat::Unknown.as_str(), "unknown"); assert_eq!(ClientFormat::Unknown.as_str(), "unknown");
} }
// ========== Session ID 提取测试 ==========
#[test]
fn test_extract_session_from_claude_metadata_user_id() {
let headers = HeaderMap::new();
let body = json!({
"model": "claude-3-5-sonnet",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {
"user_id": "user_john_doe_session_abc123def456"
}
});
let result = extract_session_id(&headers, &body, "claude");
assert_eq!(result.session_id, "abc123def456");
assert_eq!(result.source, SessionIdSource::MetadataUserId);
assert!(result.client_provided);
}
#[test]
fn test_extract_session_from_claude_metadata_session_id() {
let headers = HeaderMap::new();
let body = json!({
"model": "claude-3-5-sonnet",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {
"session_id": "my-session-123"
}
});
let result = extract_session_id(&headers, &body, "claude");
assert_eq!(result.session_id, "my-session-123");
assert_eq!(result.source, SessionIdSource::MetadataSessionId);
assert!(result.client_provided);
}
#[test]
fn test_extract_session_from_codex_previous_response_id() {
let headers = HeaderMap::new();
let body = json!({
"input": "Write a function",
"previous_response_id": "resp_abc123def456789"
});
let result = extract_session_id(&headers, &body, "codex");
assert_eq!(result.session_id, "codex_resp_abc123def456789");
assert_eq!(result.source, SessionIdSource::PreviousResponseId);
assert!(result.client_provided);
}
#[test]
fn test_extract_session_generates_new_when_not_found() {
let headers = HeaderMap::new();
let body = json!({
"model": "claude-3-5-sonnet",
"messages": [{"role": "user", "content": "Hello"}]
});
let result = extract_session_id(&headers, &body, "claude");
assert!(!result.session_id.is_empty());
assert_eq!(result.source, SessionIdSource::Generated);
assert!(!result.client_provided);
}
#[test]
fn test_parse_session_from_user_id() {
assert_eq!(
parse_session_from_user_id("user_john_session_abc123"),
Some("abc123".to_string())
);
assert_eq!(
parse_session_from_user_id("my_app_session_xyz789"),
Some("xyz789".to_string())
);
// 注意: "_session_" 是分隔符,所以下面的字符串会匹配
assert_eq!(
parse_session_from_user_id("no_session_marker"),
Some("marker".to_string())
);
// 没有 "_session_" 分隔符的情况
assert_eq!(parse_session_from_user_id("user_john_abc123"), None);
assert_eq!(parse_session_from_user_id("_session_"), None);
}
} }
+4 -77
View File
@@ -3,54 +3,33 @@ use serde::{Deserialize, Serialize};
/// 代理服务器配置 /// 代理服务器配置
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProxyConfig { pub struct ProxyConfig {
/// 是否启用代理服务
pub enabled: bool,
/// 监听地址 /// 监听地址
pub listen_address: String, pub listen_address: String,
/// 监听端口 /// 监听端口
pub listen_port: u16, pub listen_port: u16,
/// 最大重试次数 /// 最大重试次数
pub max_retries: u8, pub max_retries: u8,
/// 请求超时时间(秒)- 已废弃,保留兼容 /// 请求超时时间(秒)
pub request_timeout: u64, pub request_timeout: u64,
/// 是否启用日志 /// 是否启用日志
pub enable_logging: bool, pub enable_logging: bool,
/// 是否正在接管 Live 配置 /// 是否正在接管 Live 配置
#[serde(default)] #[serde(default)]
pub live_takeover_active: bool, pub live_takeover_active: bool,
/// 流式首字超时(秒)- 等待首个数据块的最大时间
#[serde(default = "default_streaming_first_byte_timeout")]
pub streaming_first_byte_timeout: u64,
/// 流式静默超时(秒)- 两个数据块之间的最大间隔
#[serde(default = "default_streaming_idle_timeout")]
pub streaming_idle_timeout: u64,
/// 非流式总超时(秒)- 非流式请求的总超时时间
#[serde(default = "default_non_streaming_timeout")]
pub non_streaming_timeout: u64,
}
fn default_streaming_first_byte_timeout() -> u64 {
30
}
fn default_streaming_idle_timeout() -> u64 {
60
}
fn default_non_streaming_timeout() -> u64 {
600
} }
impl Default for ProxyConfig { impl Default for ProxyConfig {
fn default() -> Self { fn default() -> Self {
Self { Self {
enabled: false,
listen_address: "127.0.0.1".to_string(), listen_address: "127.0.0.1".to_string(),
listen_port: 15721, // 使用较少占用的高位端口 listen_port: 15721, // 使用较少占用的高位端口
max_retries: 3, max_retries: 3,
request_timeout: 300, request_timeout: 300,
enable_logging: true, enable_logging: true,
live_takeover_active: false, live_takeover_active: false,
streaming_first_byte_timeout: 30,
streaming_idle_timeout: 60,
non_streaming_timeout: 600,
} }
} }
} }
@@ -107,14 +86,6 @@ pub struct ProxyServerInfo {
pub started_at: String, pub started_at: String,
} }
/// 各应用的接管状态(是否改写该应用的 Live 配置指向本地代理)
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ProxyTakeoverStatus {
pub claude: bool,
pub codex: bool,
pub gemini: bool,
}
/// API 格式类型(预留,当前不需要格式转换) /// API 格式类型(预留,当前不需要格式转换)
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[allow(dead_code)] #[allow(dead_code)]
@@ -147,47 +118,3 @@ pub struct LiveBackup {
/// 备份时间 /// 备份时间
pub backed_up_at: String, pub backed_up_at: String,
} }
/// 全局代理配置(统一字段,三行镜像)
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct GlobalProxyConfig {
/// 代理总开关
pub proxy_enabled: bool,
/// 监听地址
pub listen_address: String,
/// 监听端口
pub listen_port: u16,
/// 是否启用日志
pub enable_logging: bool,
}
/// 应用级代理配置(每个 app 独立)
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AppProxyConfig {
/// 应用类型 (claude/codex/gemini)
pub app_type: String,
/// 该 app 代理启用开关
pub enabled: bool,
/// 该 app 自动故障转移开关
pub auto_failover_enabled: bool,
/// 最大重试次数
pub max_retries: u32,
/// 流式首字超时(秒)
pub streaming_first_byte_timeout: u32,
/// 流式静默超时(秒)
pub streaming_idle_timeout: u32,
/// 非流式总超时(秒)
pub non_streaming_timeout: u32,
/// 熔断失败阈值
pub circuit_failure_threshold: u32,
/// 熔断恢复阈值
pub circuit_success_threshold: u32,
/// 熔断恢复等待时间(秒)
pub circuit_timeout_seconds: u32,
/// 错误率阈值
pub circuit_error_rate_threshold: f64,
/// 计算错误率的最小请求数
pub circuit_min_requests: u32,
}
+5 -13
View File
@@ -35,11 +35,6 @@ impl CostCalculator {
/// - `usage`: Token 使用量 /// - `usage`: Token 使用量
/// - `pricing`: 模型定价 /// - `pricing`: 模型定价
/// - `cost_multiplier`: 成本倍数 (provider 自定义) /// - `cost_multiplier`: 成本倍数 (provider 自定义)
///
/// # 计算逻辑
/// - input_cost: (input_tokens - cache_read_tokens) × 输入价格
/// - cache_read_cost: cache_read_tokens × 缓存读取价格
/// - 这样避免缓存部分被重复计费
pub fn calculate( pub fn calculate(
usage: &TokenUsage, usage: &TokenUsage,
pricing: &ModelPricing, pricing: &ModelPricing,
@@ -47,10 +42,7 @@ impl CostCalculator {
) -> CostBreakdown { ) -> CostBreakdown {
let million = Decimal::from(1_000_000); let million = Decimal::from(1_000_000);
// 计算实际需要按输入价格计费的 token 数(减去缓存命中部分) let input_cost = Decimal::from(usage.input_tokens) * pricing.input_cost_per_million
let billable_input_tokens = usage.input_tokens.saturating_sub(usage.cache_read_tokens);
let input_cost = Decimal::from(billable_input_tokens) * pricing.input_cost_per_million
/ million / million
* cost_multiplier; * cost_multiplier;
let output_cost = Decimal::from(usage.output_tokens) * pricing.output_cost_per_million let output_cost = Decimal::from(usage.output_tokens) * pricing.output_cost_per_million
@@ -121,8 +113,8 @@ mod tests {
let cost = CostCalculator::calculate(&usage, &pricing, multiplier); let cost = CostCalculator::calculate(&usage, &pricing, multiplier);
// input: (1000 - 200) * 3.0 / 1M = 0.0024 (只计算非缓存部分) // input: 1000 * 3.0 / 1M = 0.003
assert_eq!(cost.input_cost, Decimal::from_str("0.0024").unwrap()); assert_eq!(cost.input_cost, Decimal::from_str("0.003").unwrap());
// output: 500 * 15.0 / 1M = 0.0075 // output: 500 * 15.0 / 1M = 0.0075
assert_eq!(cost.output_cost, Decimal::from_str("0.0075").unwrap()); assert_eq!(cost.output_cost, Decimal::from_str("0.0075").unwrap());
// cache_read: 200 * 0.3 / 1M = 0.00006 // cache_read: 200 * 0.3 / 1M = 0.00006
@@ -132,8 +124,8 @@ mod tests {
cost.cache_creation_cost, cost.cache_creation_cost,
Decimal::from_str("0.000375").unwrap() Decimal::from_str("0.000375").unwrap()
); );
// total: 0.0024 + 0.0075 + 0.00006 + 0.000375 = 0.010335 // total: 0.003 + 0.0075 + 0.00006 + 0.000375 = 0.010935
assert_eq!(cost.total_cost, Decimal::from_str("0.010335").unwrap()); assert_eq!(cost.total_cost, Decimal::from_str("0.010935").unwrap());
} }
#[test] #[test]
+21 -368
View File
@@ -34,12 +34,6 @@ impl TokenUsage {
/// 从 Claude API 非流式响应解析 /// 从 Claude API 非流式响应解析
pub fn from_claude_response(body: &Value) -> Option<Self> { pub fn from_claude_response(body: &Value) -> Option<Self> {
let usage = body.get("usage")?; let usage = body.get("usage")?;
// 提取响应中的模型名称
let model = body
.get("model")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
Some(Self { Some(Self {
input_tokens: usage.get("input_tokens")?.as_u64()? as u32, input_tokens: usage.get("input_tokens")?.as_u64()? as u32,
output_tokens: usage.get("output_tokens")?.as_u64()? as u32, output_tokens: usage.get("output_tokens")?.as_u64()? as u32,
@@ -51,7 +45,7 @@ impl TokenUsage {
.get("cache_creation_input_tokens") .get("cache_creation_input_tokens")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0) as u32, .unwrap_or(0) as u32,
model, model: None,
}) })
} }
@@ -59,20 +53,11 @@ impl TokenUsage {
#[allow(dead_code)] #[allow(dead_code)]
pub fn from_claude_stream_events(events: &[Value]) -> Option<Self> { pub fn from_claude_stream_events(events: &[Value]) -> Option<Self> {
let mut usage = Self::default(); let mut usage = Self::default();
let mut model: Option<String> = None;
for event in events { for event in events {
if let Some(event_type) = event.get("type").and_then(|v| v.as_str()) { if let Some(event_type) = event.get("type").and_then(|v| v.as_str()) {
match event_type { match event_type {
"message_start" => { "message_start" => {
// 从 message_start 提取模型名称
if model.is_none() {
if let Some(message) = event.get("message") {
if let Some(m) = message.get("model").and_then(|v| v.as_str()) {
model = Some(m.to_string());
}
}
}
if let Some(msg_usage) = event.get("message").and_then(|m| m.get("usage")) { if let Some(msg_usage) = event.get("message").and_then(|m| m.get("usage")) {
// 从 message_start 获取 input_tokens(原生 Claude API // 从 message_start 获取 input_tokens(原生 Claude API
if let Some(input) = if let Some(input) =
@@ -117,7 +102,6 @@ impl TokenUsage {
} }
if usage.input_tokens > 0 || usage.output_tokens > 0 { if usage.input_tokens > 0 || usage.output_tokens > 0 {
usage.model = model;
Some(usage) Some(usage)
} else { } else {
None None
@@ -157,32 +141,18 @@ impl TokenUsage {
return None; return None;
} }
// 提取响应中的模型名称
let model = body
.get("model")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
let cached_tokens = usage
.get("cache_read_input_tokens")
.and_then(|v| v.as_u64())
.or_else(|| {
usage
.get("input_tokens_details")
.and_then(|d| d.get("cached_tokens"))
.and_then(|v| v.as_u64())
})
.unwrap_or(0) as u32;
Some(Self { Some(Self {
input_tokens: input_tokens? as u32, input_tokens: input_tokens? as u32,
output_tokens: output_tokens? as u32, output_tokens: output_tokens? as u32,
cache_read_tokens: cached_tokens, cache_read_tokens: usage
.get("cache_read_input_tokens")
.and_then(|v| v.as_u64())
.unwrap_or(0) as u32,
cache_creation_tokens: usage cache_creation_tokens: usage
.get("cache_creation_input_tokens") .get("cache_creation_input_tokens")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0) as u32, .unwrap_or(0) as u32,
model, model: None,
}) })
} }
@@ -196,27 +166,16 @@ impl TokenUsage {
let input_tokens = usage.get("input_tokens")?.as_u64()? as u32; let input_tokens = usage.get("input_tokens")?.as_u64()? as u32;
let output_tokens = usage.get("output_tokens")?.as_u64()? as u32; let output_tokens = usage.get("output_tokens")?.as_u64()? as u32;
// 获取 cached_tokens (可能在 cache_read_input_tokens 或 input_tokens_details 中) // 获取 cached_tokens (可能在 input_tokens_details 中)
let cached_tokens = usage let cached_tokens = usage
.get("cache_read_input_tokens") .get("input_tokens_details")
.and_then(|d| d.get("cached_tokens"))
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.or_else(|| {
usage
.get("input_tokens_details")
.and_then(|d| d.get("cached_tokens"))
.and_then(|v| v.as_u64())
})
.unwrap_or(0) as u32; .unwrap_or(0) as u32;
// 调整 input_tokens: 减去 cached_tokens // 调整 input_tokens: 减去 cached_tokens
let adjusted_input = input_tokens.saturating_sub(cached_tokens); let adjusted_input = input_tokens.saturating_sub(cached_tokens);
// 提取响应中的模型名称
let model = body
.get("model")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
Some(Self { Some(Self {
input_tokens: adjusted_input, input_tokens: adjusted_input,
output_tokens, output_tokens,
@@ -225,7 +184,7 @@ impl TokenUsage {
.get("cache_creation_input_tokens") .get("cache_creation_input_tokens")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0) as u32, .unwrap_or(0) as u32,
model, model: None,
}) })
} }
@@ -239,7 +198,7 @@ impl TokenUsage {
if event_type == "response.completed" { if event_type == "response.completed" {
if let Some(response) = event.get("response") { if let Some(response) = event.get("response") {
log::debug!("[Codex] 找到 response.completed 事件,解析 usage"); log::debug!("[Codex] 找到 response.completed 事件,解析 usage");
return Self::from_codex_response_adjusted(response); return Self::from_codex_response(response);
} }
} }
} }
@@ -248,51 +207,6 @@ impl TokenUsage {
None None
} }
/// 智能 Codex 响应解析 - 自动检测 OpenAI 或 Codex 格式
///
/// Codex 支持两种 API 格式:
/// - `/v1/responses`: 使用 input_tokens/output_tokens
/// - `/v1/chat/completions`: 使用 prompt_tokens/completion_tokens (OpenAI 格式)
///
/// 注意:记录原始 input_tokens,费用计算时再减去 cached_tokens
pub fn from_codex_response_auto(body: &Value) -> Option<Self> {
let usage = body.get("usage")?;
// 检测格式:OpenAI 使用 prompt_tokensCodex 使用 input_tokens
if usage.get("prompt_tokens").is_some() {
log::debug!("[Codex] 检测到 OpenAI 格式 (prompt_tokens)");
Self::from_openai_response(body)
} else if usage.get("input_tokens").is_some() {
log::debug!("[Codex] 检测到 Codex 格式 (input_tokens)");
// 使用非调整版本,记录原始 input_tokens
Self::from_codex_response(body)
} else {
log::debug!("[Codex] 无法识别响应格式,usage: {usage:?}");
None
}
}
/// 智能 Codex 流式响应解析 - 自动检测 OpenAI 或 Codex 格式
pub fn from_codex_stream_events_auto(events: &[Value]) -> Option<Self> {
log::debug!("[Codex] 智能解析流式事件,共 {} 个事件", events.len());
// 先尝试 Codex Responses API 格式 (response.completed 事件)
for event in events {
if let Some(event_type) = event.get("type").and_then(|v| v.as_str()) {
if event_type == "response.completed" {
if let Some(response) = event.get("response") {
log::debug!("[Codex] 找到 response.completed 事件");
return Self::from_codex_response_auto(response);
}
}
}
}
// 回退到 OpenAI Chat Completions 格式 (最后一个 chunk 包含 usage)
log::debug!("[Codex] 尝试 OpenAI 流式格式");
Self::from_openai_stream_events(events)
}
/// 从 OpenAI Chat Completions API 响应解析 (prompt_tokens, completion_tokens) /// 从 OpenAI Chat Completions API 响应解析 (prompt_tokens, completion_tokens)
pub fn from_openai_response(body: &Value) -> Option<Self> { pub fn from_openai_response(body: &Value) -> Option<Self> {
let usage = body.get("usage")?; let usage = body.get("usage")?;
@@ -308,18 +222,12 @@ impl TokenUsage {
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0) as u32; .unwrap_or(0) as u32;
// 提取响应中的模型名称
let model = body
.get("model")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
Some(Self { Some(Self {
input_tokens: prompt_tokens as u32, input_tokens: prompt_tokens as u32,
output_tokens: completion_tokens as u32, output_tokens: completion_tokens as u32,
cache_read_tokens: cached_tokens, cache_read_tokens: cached_tokens,
cache_creation_tokens: 0, cache_creation_tokens: 0,
model, model: None,
}) })
} }
@@ -348,16 +256,9 @@ impl TokenUsage {
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.map(|s| s.to_string()); .map(|s| s.to_string());
let prompt_tokens = usage.get("promptTokenCount")?.as_u64()? as u32;
let total_tokens = usage.get("totalTokenCount")?.as_u64()? as u32;
// 输出 tokens = 总 tokens - 输入 tokens
// 这包含了 candidatesTokenCount + thoughtsTokenCount
let output_tokens = total_tokens.saturating_sub(prompt_tokens);
Some(Self { Some(Self {
input_tokens: prompt_tokens, input_tokens: usage.get("promptTokenCount")?.as_u64()? as u32,
output_tokens, output_tokens: usage.get("candidatesTokenCount")?.as_u64()? as u32,
cache_read_tokens: usage cache_read_tokens: usage
.get("cachedContentTokenCount") .get("cachedContentTokenCount")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
@@ -371,25 +272,20 @@ impl TokenUsage {
#[allow(dead_code)] #[allow(dead_code)]
pub fn from_gemini_stream_chunks(chunks: &[Value]) -> Option<Self> { pub fn from_gemini_stream_chunks(chunks: &[Value]) -> Option<Self> {
let mut total_input = 0u32; let mut total_input = 0u32;
let mut total_tokens = 0u32; let mut total_output = 0u32;
let mut total_cache_read = 0u32; let mut total_cache_read = 0u32;
let mut model: Option<String> = None; let mut model: Option<String> = None;
for chunk in chunks { for chunk in chunks {
if let Some(usage) = chunk.get("usageMetadata") { if let Some(usage) = chunk.get("usageMetadata") {
// 输入 tokens (通常在所有 chunk 中保持不变)
total_input = usage total_input = usage
.get("promptTokenCount") .get("promptTokenCount")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0) as u32; .unwrap_or(0) as u32;
total_output += usage
// 总 tokens (包含输入 + 输出 + 思考) .get("candidatesTokenCount")
total_tokens = usage
.get("totalTokenCount")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0) as u32; .unwrap_or(0) as u32;
// 缓存读取 tokens
total_cache_read = usage total_cache_read = usage
.get("cachedContentTokenCount") .get("cachedContentTokenCount")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
@@ -404,9 +300,6 @@ impl TokenUsage {
} }
} }
// 输出 tokens = 总 tokens - 输入 tokens
let total_output = total_tokens.saturating_sub(total_input);
if total_input > 0 || total_output > 0 { if total_input > 0 || total_output > 0 {
Some(Self { Some(Self {
input_tokens: total_input, input_tokens: total_input,
@@ -429,7 +322,6 @@ mod tests {
#[test] #[test]
fn test_claude_response_parsing() { fn test_claude_response_parsing() {
let response = json!({ let response = json!({
"model": "claude-sonnet-4-20250514",
"usage": { "usage": {
"input_tokens": 100, "input_tokens": 100,
"output_tokens": 50, "output_tokens": 50,
@@ -443,26 +335,6 @@ mod tests {
assert_eq!(usage.output_tokens, 50); assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.cache_read_tokens, 20); assert_eq!(usage.cache_read_tokens, 20);
assert_eq!(usage.cache_creation_tokens, 10); assert_eq!(usage.cache_creation_tokens, 10);
assert_eq!(usage.model, Some("claude-sonnet-4-20250514".to_string()));
}
#[test]
fn test_claude_response_parsing_no_model() {
let response = json!({
"usage": {
"input_tokens": 100,
"output_tokens": 50,
"cache_read_input_tokens": 20,
"cache_creation_input_tokens": 10
}
});
let usage = TokenUsage::from_claude_response(&response).unwrap();
assert_eq!(usage.input_tokens, 100);
assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.cache_read_tokens, 20);
assert_eq!(usage.cache_creation_tokens, 10);
assert_eq!(usage.model, None);
} }
#[test] #[test]
@@ -471,7 +343,6 @@ mod tests {
json!({ json!({
"type": "message_start", "type": "message_start",
"message": { "message": {
"model": "claude-sonnet-4-20250514",
"usage": { "usage": {
"input_tokens": 100, "input_tokens": 100,
"cache_read_input_tokens": 20, "cache_read_input_tokens": 20,
@@ -492,36 +363,6 @@ mod tests {
assert_eq!(usage.output_tokens, 50); assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.cache_read_tokens, 20); assert_eq!(usage.cache_read_tokens, 20);
assert_eq!(usage.cache_creation_tokens, 10); assert_eq!(usage.cache_creation_tokens, 10);
assert_eq!(usage.model, Some("claude-sonnet-4-20250514".to_string()));
}
#[test]
fn test_claude_stream_parsing_no_model() {
let events = vec![
json!({
"type": "message_start",
"message": {
"usage": {
"input_tokens": 100,
"cache_read_input_tokens": 20,
"cache_creation_input_tokens": 10
}
}
}),
json!({
"type": "message_delta",
"usage": {
"output_tokens": 50
}
}),
];
let usage = TokenUsage::from_claude_stream_events(&events).unwrap();
assert_eq!(usage.input_tokens, 100);
assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.cache_read_tokens, 20);
assert_eq!(usage.cache_creation_tokens, 10);
assert_eq!(usage.model, None);
} }
#[test] #[test]
@@ -545,18 +386,15 @@ mod tests {
let response = json!({ let response = json!({
"modelVersion": "gemini-3-pro-high", "modelVersion": "gemini-3-pro-high",
"usageMetadata": { "usageMetadata": {
"promptTokenCount": 8383, "promptTokenCount": 100,
"candidatesTokenCount": 50, "candidatesTokenCount": 50,
"thoughtsTokenCount": 114,
"totalTokenCount": 8547,
"cachedContentTokenCount": 20 "cachedContentTokenCount": 20
} }
}); });
let usage = TokenUsage::from_gemini_response(&response).unwrap(); let usage = TokenUsage::from_gemini_response(&response).unwrap();
assert_eq!(usage.input_tokens, 8383); assert_eq!(usage.input_tokens, 100);
// output_tokens = totalTokenCount - promptTokenCount = 8547 - 8383 = 164 assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.output_tokens, 164);
assert_eq!(usage.cache_read_tokens, 20); assert_eq!(usage.cache_read_tokens, 20);
assert_eq!(usage.cache_creation_tokens, 0); assert_eq!(usage.cache_creation_tokens, 0);
assert_eq!(usage.model, Some("gemini-3-pro-high".to_string())); assert_eq!(usage.model, Some("gemini-3-pro-high".to_string()));
@@ -568,78 +406,19 @@ mod tests {
let response = json!({ let response = json!({
"usageMetadata": { "usageMetadata": {
"promptTokenCount": 100, "promptTokenCount": 100,
"totalTokenCount": 150, "candidatesTokenCount": 50,
"cachedContentTokenCount": 20 "cachedContentTokenCount": 20
} }
}); });
let usage = TokenUsage::from_gemini_response(&response).unwrap(); let usage = TokenUsage::from_gemini_response(&response).unwrap();
assert_eq!(usage.input_tokens, 100); assert_eq!(usage.input_tokens, 100);
// output_tokens = totalTokenCount - promptTokenCount = 150 - 100 = 50
assert_eq!(usage.output_tokens, 50); assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.cache_read_tokens, 20); assert_eq!(usage.cache_read_tokens, 20);
assert_eq!(usage.cache_creation_tokens, 0); assert_eq!(usage.cache_creation_tokens, 0);
assert_eq!(usage.model, None); assert_eq!(usage.model, None);
} }
#[test]
fn test_gemini_response_with_thoughts() {
// 测试包含 thoughtsTokenCount 的实际响应
// 这是用户报告的真实场景
let response = json!({
"candidates": [
{
"content": {
"parts": [
{
"text": "",
"thoughtSignature": "EvcECvQE..."
}
],
"role": "model"
},
"finishReason": "STOP"
}
],
"modelVersion": "gemini-3-pro-high",
"responseId": "yupTafqLDu-PjMcPhrOx4QQ",
"usageMetadata": {
"candidatesTokenCount": 50,
"promptTokenCount": 8383,
"thoughtsTokenCount": 114,
"totalTokenCount": 8547
}
});
let usage = TokenUsage::from_gemini_response(&response).unwrap();
assert_eq!(usage.input_tokens, 8383);
// output_tokens = totalTokenCount - promptTokenCount
// = 8547 - 8383 = 164 (包含 candidatesTokenCount 50 + thoughtsTokenCount 114)
assert_eq!(usage.output_tokens, 164);
assert_eq!(usage.cache_read_tokens, 0);
assert_eq!(usage.cache_creation_tokens, 0);
assert_eq!(usage.model, Some("gemini-3-pro-high".to_string()));
}
#[test]
fn test_codex_response_parsing_cached_tokens_in_details() {
let response = json!({
"usage": {
"input_tokens": 1000,
"output_tokens": 500,
"input_tokens_details": {
"cached_tokens": 300
}
}
});
let usage = TokenUsage::from_codex_response(&response).unwrap();
// 非调整模式:input_tokens 保持原值,但应记录缓存命中
assert_eq!(usage.input_tokens, 1000);
assert_eq!(usage.output_tokens, 500);
assert_eq!(usage.cache_read_tokens, 300);
}
#[test] #[test]
fn test_codex_response_adjusted() { fn test_codex_response_adjusted() {
let response = json!({ let response = json!({
@@ -675,22 +454,6 @@ mod tests {
assert_eq!(usage.cache_read_tokens, 0); assert_eq!(usage.cache_read_tokens, 0);
} }
#[test]
fn test_codex_response_adjusted_cache_read_input_tokens() {
let response = json!({
"usage": {
"input_tokens": 1000,
"output_tokens": 500,
"cache_read_input_tokens": 200
}
});
let usage = TokenUsage::from_codex_response_adjusted(&response).unwrap();
assert_eq!(usage.input_tokens, 800);
assert_eq!(usage.output_tokens, 500);
assert_eq!(usage.cache_read_tokens, 200);
}
#[test] #[test]
fn test_codex_response_adjusted_saturating_sub() { fn test_codex_response_adjusted_saturating_sub() {
// 测试 cached_tokens > input_tokens 的边界情况 // 测试 cached_tokens > input_tokens 的边界情况
@@ -718,7 +481,6 @@ mod tests {
json!({ json!({
"type": "message_start", "type": "message_start",
"message": { "message": {
"model": "claude-sonnet-4-20250514",
"usage": { "usage": {
"input_tokens": 0, "input_tokens": 0,
"output_tokens": 0 "output_tokens": 0
@@ -740,7 +502,6 @@ mod tests {
let usage = TokenUsage::from_claude_stream_events(&events).unwrap(); let usage = TokenUsage::from_claude_stream_events(&events).unwrap();
assert_eq!(usage.input_tokens, 150); assert_eq!(usage.input_tokens, 150);
assert_eq!(usage.output_tokens, 75); assert_eq!(usage.output_tokens, 75);
assert_eq!(usage.model, Some("claude-sonnet-4-20250514".to_string()));
} }
#[test] #[test]
@@ -751,7 +512,6 @@ mod tests {
json!({ json!({
"type": "message_start", "type": "message_start",
"message": { "message": {
"model": "claude-sonnet-4-20250514",
"usage": { "usage": {
"input_tokens": 200, "input_tokens": 200,
"cache_read_input_tokens": 50 "cache_read_input_tokens": 50
@@ -770,112 +530,5 @@ mod tests {
assert_eq!(usage.input_tokens, 200); assert_eq!(usage.input_tokens, 200);
assert_eq!(usage.output_tokens, 100); assert_eq!(usage.output_tokens, 100);
assert_eq!(usage.cache_read_tokens, 50); assert_eq!(usage.cache_read_tokens, 50);
assert_eq!(usage.model, Some("claude-sonnet-4-20250514".to_string()));
}
// ============================================================================
// 智能 Codex 解析测试
// ============================================================================
#[test]
fn test_codex_response_auto_openai_format() {
// OpenAI 格式 (prompt_tokens/completion_tokens)
let response = json!({
"model": "gpt-4o",
"usage": {
"prompt_tokens": 1000,
"completion_tokens": 500,
"prompt_tokens_details": {
"cached_tokens": 200
}
}
});
let usage = TokenUsage::from_codex_response_auto(&response).unwrap();
assert_eq!(usage.input_tokens, 1000);
assert_eq!(usage.output_tokens, 500);
assert_eq!(usage.cache_read_tokens, 200);
assert_eq!(usage.model, Some("gpt-4o".to_string()));
}
#[test]
fn test_codex_response_auto_codex_format() {
// Codex 格式 (input_tokens/output_tokens)
let response = json!({
"model": "o3",
"usage": {
"input_tokens": 1000,
"output_tokens": 500,
"input_tokens_details": {
"cached_tokens": 300
}
}
});
let usage = TokenUsage::from_codex_response_auto(&response).unwrap();
// 记录原始 input_tokens,不调整
assert_eq!(usage.input_tokens, 1000);
assert_eq!(usage.output_tokens, 500);
assert_eq!(usage.cache_read_tokens, 300);
assert_eq!(usage.model, Some("o3".to_string()));
}
#[test]
fn test_codex_stream_events_auto_codex_format() {
// Codex Responses API 流式格式 (response.completed 事件)
let events = vec![
json!({
"type": "response.created",
"response": {
"id": "resp_123"
}
}),
json!({
"type": "response.completed",
"response": {
"model": "o3",
"usage": {
"input_tokens": 1000,
"output_tokens": 500,
"input_tokens_details": {
"cached_tokens": 200
}
}
}
}),
];
let usage = TokenUsage::from_codex_stream_events_auto(&events).unwrap();
// 记录原始 input_tokens,不调整
assert_eq!(usage.input_tokens, 1000);
assert_eq!(usage.output_tokens, 500);
assert_eq!(usage.cache_read_tokens, 200);
assert_eq!(usage.model, Some("o3".to_string()));
}
#[test]
fn test_codex_stream_events_auto_openai_format() {
// OpenAI Chat Completions 流式格式 (最后一个 chunk 包含 usage)
let events = vec![
json!({
"id": "chatcmpl-123",
"model": "gpt-4o",
"choices": [{"delta": {"content": "Hello"}}]
}),
json!({
"id": "chatcmpl-123",
"model": "gpt-4o",
"choices": [{"delta": {}}],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50
}
}),
];
let usage = TokenUsage::from_codex_stream_events_auto(&events).unwrap();
assert_eq!(usage.input_tokens, 100);
assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.model, Some("gpt-4o".to_string()));
} }
} }
+1 -3
View File
@@ -146,9 +146,7 @@ impl ConfigService {
let cfg_text = settings.get("config").and_then(Value::as_str); let cfg_text = settings.get("config").and_then(Value::as_str);
crate::codex_config::write_codex_live_atomic(auth, cfg_text)?; crate::codex_config::write_codex_live_atomic(auth, cfg_text)?;
// 注意:MCP 同步在 v3.7.0 中已通过 McpService 进行,不再在此调用 crate::mcp::sync_enabled_to_codex(config)?;
// sync_enabled_to_codex 使用旧的 config.mcp.codex 结构,在新架构中为空
// MCP 的启用/禁用应通过 McpService::toggle_app 进行
let cfg_text_after = crate::codex_config::read_and_validate_codex_config_text()?; let cfg_text_after = crate::codex_config::read_and_validate_codex_config_text()?;
if let Some(manager) = config.get_manager_mut(&AppType::Codex) { if let Some(manager) = config.get_manager_mut(&AppType::Codex) {
+9 -64
View File
@@ -17,27 +17,8 @@ impl McpService {
/// 添加或更新 MCP 服务器 /// 添加或更新 MCP 服务器
pub fn upsert_server(state: &AppState, server: McpServer) -> Result<(), AppError> { pub fn upsert_server(state: &AppState, server: McpServer) -> Result<(), AppError> {
// 读取旧状态:用于处理“编辑时取消勾选某个应用”的场景(需要从对应 live 配置中移除)
let prev_apps = state
.db
.get_all_mcp_servers()?
.get(&server.id)
.map(|s| s.apps.clone())
.unwrap_or_default();
state.db.save_mcp_server(&server)?; state.db.save_mcp_server(&server)?;
// 处理禁用:若旧版本启用但新版本取消,则需要从该应用的 live 配置移除
if prev_apps.claude && !server.apps.claude {
Self::remove_server_from_app(state, &server.id, &AppType::Claude)?;
}
if prev_apps.codex && !server.apps.codex {
Self::remove_server_from_app(state, &server.id, &AppType::Codex)?;
}
if prev_apps.gemini && !server.apps.gemini {
Self::remove_server_from_app(state, &server.id, &AppType::Gemini)?;
}
// 同步到各个启用的应用 // 同步到各个启用的应用
Self::sync_server_to_apps(state, &server)?; Self::sync_server_to_apps(state, &server)?;
@@ -209,22 +190,10 @@ impl McpService {
// 如果有导入的服务器,保存到数据库 // 如果有导入的服务器,保存到数据库
if count > 0 { if count > 0 {
if let Some(servers) = &temp_config.mcp.servers { if let Some(servers) = &temp_config.mcp.servers {
let mut existing = state.db.get_all_mcp_servers()?;
for server in servers.values() { for server in servers.values() {
// 已存在:仅启用 Claude,不覆盖其他字段(与导入模块语义保持一致) state.db.save_mcp_server(server)?;
let to_save = if let Some(existing_server) = existing.get(&server.id) { // 同步到 Claude live 配置
let mut merged = existing_server.clone(); Self::sync_server_to_apps(state, server)?;
merged.apps.claude = true;
merged
} else {
server.clone()
};
state.db.save_mcp_server(&to_save)?;
existing.insert(to_save.id.clone(), to_save.clone());
// 同步到对应应用 live 配置
Self::sync_server_to_apps(state, &to_save)?;
} }
} }
} }
@@ -243,22 +212,10 @@ impl McpService {
// 如果有导入的服务器,保存到数据库 // 如果有导入的服务器,保存到数据库
if count > 0 { if count > 0 {
if let Some(servers) = &temp_config.mcp.servers { if let Some(servers) = &temp_config.mcp.servers {
let mut existing = state.db.get_all_mcp_servers()?;
for server in servers.values() { for server in servers.values() {
// 已存在:仅启用 Codex,不覆盖其他字段(与导入模块语义保持一致) state.db.save_mcp_server(server)?;
let to_save = if let Some(existing_server) = existing.get(&server.id) { // 同步到 Codex live 配置
let mut merged = existing_server.clone(); Self::sync_server_to_apps(state, server)?;
merged.apps.codex = true;
merged
} else {
server.clone()
};
state.db.save_mcp_server(&to_save)?;
existing.insert(to_save.id.clone(), to_save.clone());
// 同步到对应应用 live 配置
Self::sync_server_to_apps(state, &to_save)?;
} }
} }
} }
@@ -277,22 +234,10 @@ impl McpService {
// 如果有导入的服务器,保存到数据库 // 如果有导入的服务器,保存到数据库
if count > 0 { if count > 0 {
if let Some(servers) = &temp_config.mcp.servers { if let Some(servers) = &temp_config.mcp.servers {
let mut existing = state.db.get_all_mcp_servers()?;
for server in servers.values() { for server in servers.values() {
// 已存在:仅启用 Gemini,不覆盖其他字段(与导入模块语义保持一致) state.db.save_mcp_server(server)?;
let to_save = if let Some(existing_server) = existing.get(&server.id) { // 同步到 Gemini live 配置
let mut merged = existing_server.clone(); Self::sync_server_to_apps(state, server)?;
merged.apps.gemini = true;
merged
} else {
server.clone()
};
state.db.save_mcp_server(&to_save)?;
existing.insert(to_save.id.clone(), to_save.clone());
// 同步到对应应用 live 配置
Self::sync_server_to_apps(state, &to_save)?;
} }
} }
} }
+14 -170
View File
@@ -144,29 +144,9 @@ impl ProviderService {
state.db.save_provider(app_type.as_str(), &provider)?; state.db.save_provider(app_type.as_str(), &provider)?;
if is_current { if is_current {
// 如果代理接管模式处于激活状态,并且代理服务正在运行: write_live_snapshot(&app_type, &provider)?;
// - 不写 Live 配置(否则会破坏接管) // Sync MCP
// - 仅更新 Live 备份(保证关闭代理时能恢复到最新配置) McpService::sync_all_enabled(state)?;
let is_app_taken_over =
futures::executor::block_on(state.db.get_live_backup(app_type.as_str()))
.ok()
.flatten()
.is_some();
let is_proxy_running = futures::executor::block_on(state.proxy_service.is_running());
let should_skip_live_write = is_app_taken_over && is_proxy_running;
if should_skip_live_write {
futures::executor::block_on(
state
.proxy_service
.update_live_backup_from_provider(app_type.as_str(), &provider),
)
.map_err(|e| AppError::Message(format!("更新 Live 备份失败: {e}")))?;
} else {
write_live_snapshot(&app_type, &provider)?;
// Sync MCP
McpService::sync_all_enabled(state)?;
}
} }
Ok(true) Ok(true)
@@ -211,18 +191,12 @@ impl ProviderService {
// Check if proxy takeover mode is active AND proxy server is actually running // Check if proxy takeover mode is active AND proxy server is actually running
// Both conditions must be true to use hot-switch mode // Both conditions must be true to use hot-switch mode
// Use blocking wait since this is a sync function // Use blocking wait since this is a sync function
let is_app_taken_over = let is_takeover_flag =
futures::executor::block_on(state.db.get_live_backup(app_type.as_str())) futures::executor::block_on(state.db.is_live_takeover_active()).unwrap_or(false);
.ok()
.flatten()
.is_some();
let is_proxy_running = futures::executor::block_on(state.proxy_service.is_running()); let is_proxy_running = futures::executor::block_on(state.proxy_service.is_running());
let live_taken_over = state
.proxy_service
.detect_takeover_in_live_config_for_app(&app_type);
// Hot-switch only when BOTH: this app is taken over AND proxy server is actually running // Hot-switch only when BOTH: takeover flag is set AND proxy server is actually running
let should_hot_switch = (is_app_taken_over || live_taken_over) && is_proxy_running; let should_hot_switch = is_takeover_flag && is_proxy_running;
if should_hot_switch { if should_hot_switch {
// Proxy takeover mode: hot-switch only, don't write Live config // Proxy takeover mode: hot-switch only, don't write Live config
@@ -257,6 +231,13 @@ impl ProviderService {
} }
// Normal mode: full switch with Live config write // Normal mode: full switch with Live config write
// Also clear stale takeover flag if proxy is not running but flag was set
if is_takeover_flag && !is_proxy_running {
log::warn!("检测到代理接管标志残留(代理已停止),清除标志并执行正常切换");
// Clear stale takeover flag
let _ = futures::executor::block_on(state.db.set_live_takeover_active(false));
}
Self::switch_normal(state, app_type, id, &providers) Self::switch_normal(state, app_type, id, &providers)
} }
@@ -693,140 +674,3 @@ pub struct ProviderSortUpdate {
#[serde(rename = "sortIndex")] #[serde(rename = "sortIndex")]
pub sort_index: usize, pub sort_index: usize,
} }
// ============================================================================
// 统一供应商(Universal Provider)服务方法
// ============================================================================
use crate::provider::UniversalProvider;
use std::collections::HashMap;
impl ProviderService {
/// 获取所有统一供应商
pub fn list_universal(
state: &AppState,
) -> Result<HashMap<String, UniversalProvider>, AppError> {
state.db.get_all_universal_providers()
}
/// 获取单个统一供应商
pub fn get_universal(
state: &AppState,
id: &str,
) -> Result<Option<UniversalProvider>, AppError> {
state.db.get_universal_provider(id)
}
/// 添加或更新统一供应商(不自动同步,需手动调用 sync_universal_to_apps
pub fn upsert_universal(
state: &AppState,
provider: UniversalProvider,
) -> Result<bool, AppError> {
// 保存统一供应商
state.db.save_universal_provider(&provider)?;
Ok(true)
}
/// 删除统一供应商
pub fn delete_universal(state: &AppState, id: &str) -> Result<bool, AppError> {
// 获取统一供应商(用于删除生成的子供应商)
let provider = state.db.get_universal_provider(id)?;
// 删除统一供应商
state.db.delete_universal_provider(id)?;
// 删除生成的子供应商
if let Some(p) = provider {
if p.apps.claude {
let claude_id = format!("universal-claude-{id}");
let _ = state.db.delete_provider("claude", &claude_id);
}
if p.apps.codex {
let codex_id = format!("universal-codex-{id}");
let _ = state.db.delete_provider("codex", &codex_id);
}
if p.apps.gemini {
let gemini_id = format!("universal-gemini-{id}");
let _ = state.db.delete_provider("gemini", &gemini_id);
}
}
Ok(true)
}
/// 同步统一供应商到各应用
pub fn sync_universal_to_apps(state: &AppState, id: &str) -> Result<bool, AppError> {
let provider = state
.db
.get_universal_provider(id)?
.ok_or_else(|| AppError::Message(format!("统一供应商 {id} 不存在")))?;
// 同步到 Claude
if let Some(mut claude_provider) = provider.to_claude_provider() {
// 合并已有配置
if let Some(existing) = state.db.get_provider_by_id(&claude_provider.id, "claude")? {
let mut merged = existing.settings_config.clone();
Self::merge_json(&mut merged, &claude_provider.settings_config);
claude_provider.settings_config = merged;
}
state.db.save_provider("claude", &claude_provider)?;
} else {
// 如果禁用了 Claude,删除对应的子供应商
let claude_id = format!("universal-claude-{id}");
let _ = state.db.delete_provider("claude", &claude_id);
}
// 同步到 Codex
if let Some(mut codex_provider) = provider.to_codex_provider() {
// 合并已有配置
if let Some(existing) = state.db.get_provider_by_id(&codex_provider.id, "codex")? {
let mut merged = existing.settings_config.clone();
Self::merge_json(&mut merged, &codex_provider.settings_config);
codex_provider.settings_config = merged;
}
state.db.save_provider("codex", &codex_provider)?;
} else {
let codex_id = format!("universal-codex-{id}");
let _ = state.db.delete_provider("codex", &codex_id);
}
// 同步到 Gemini
if let Some(mut gemini_provider) = provider.to_gemini_provider() {
// 合并已有配置
if let Some(existing) = state.db.get_provider_by_id(&gemini_provider.id, "gemini")? {
let mut merged = existing.settings_config.clone();
Self::merge_json(&mut merged, &gemini_provider.settings_config);
gemini_provider.settings_config = merged;
}
state.db.save_provider("gemini", &gemini_provider)?;
} else {
let gemini_id = format!("universal-gemini-{id}");
let _ = state.db.delete_provider("gemini", &gemini_id);
}
Ok(true)
}
/// 递归合并 JSONbase 为底,patch 覆盖同名字段
fn merge_json(base: &mut serde_json::Value, patch: &serde_json::Value) {
use serde_json::Value;
match (base, patch) {
(Value::Object(base_map), Value::Object(patch_map)) => {
for (k, v_patch) in patch_map {
match base_map.get_mut(k) {
Some(v_base) => Self::merge_json(v_base, v_patch),
None => {
base_map.insert(k.clone(), v_patch.clone());
}
}
}
}
// 其它类型:直接覆盖
(base_val, patch_val) => {
*base_val = patch_val.clone();
}
}
}
}
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -77,7 +77,7 @@ impl Default for SkillStore {
SkillRepo { SkillRepo {
owner: "ComposioHQ".to_string(), owner: "ComposioHQ".to_string(),
name: "awesome-claude-skills".to_string(), name: "awesome-claude-skills".to_string(),
branch: "master".to_string(), branch: "main".to_string(),
enabled: true, enabled: true,
}, },
SkillRepo { SkillRepo {
+3 -48
View File
@@ -3,7 +3,6 @@
//! 使用流式 API 进行快速健康检查,只需接收首个 chunk 即判定成功。 //! 使用流式 API 进行快速健康检查,只需接收首个 chunk 即判定成功。
use futures::StreamExt; use futures::StreamExt;
use regex::Regex;
use reqwest::Client; use reqwest::Client;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::json; use serde_json::json;
@@ -142,17 +141,15 @@ impl StreamCheckService {
.build() .build()
.map_err(|e| AppError::Message(format!("创建客户端失败: {e}")))?; .map_err(|e| AppError::Message(format!("创建客户端失败: {e}")))?;
let model_to_test = Self::resolve_test_model(app_type, provider, config);
let result = match app_type { let result = match app_type {
AppType::Claude => { AppType::Claude => {
Self::check_claude_stream(&client, &base_url, &auth, &model_to_test).await Self::check_claude_stream(&client, &base_url, &auth, &config.claude_model).await
} }
AppType::Codex => { AppType::Codex => {
Self::check_codex_stream(&client, &base_url, &auth, &model_to_test).await Self::check_codex_stream(&client, &base_url, &auth, &config.codex_model).await
} }
AppType::Gemini => { AppType::Gemini => {
Self::check_gemini_stream(&client, &base_url, &auth, &model_to_test).await Self::check_gemini_stream(&client, &base_url, &auth, &config.gemini_model).await
} }
}; };
@@ -382,48 +379,6 @@ impl StreamCheckService {
AppError::Message(e.to_string()) AppError::Message(e.to_string())
} }
} }
fn resolve_test_model(
app_type: &AppType,
provider: &Provider,
config: &StreamCheckConfig,
) -> String {
match app_type {
AppType::Claude => Self::extract_env_model(provider, "ANTHROPIC_MODEL")
.unwrap_or_else(|| config.claude_model.clone()),
AppType::Codex => {
Self::extract_codex_model(provider).unwrap_or_else(|| config.codex_model.clone())
}
AppType::Gemini => Self::extract_env_model(provider, "GEMINI_MODEL")
.unwrap_or_else(|| config.gemini_model.clone()),
}
}
fn extract_env_model(provider: &Provider, key: &str) -> Option<String> {
provider
.settings_config
.get("env")
.and_then(|env| env.get(key))
.and_then(|value| value.as_str())
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
fn extract_codex_model(provider: &Provider) -> Option<String> {
let config_text = provider
.settings_config
.get("config")
.and_then(|value| value.as_str())?;
if config_text.trim().is_empty() {
return None;
}
let re = Regex::new(r#"^model\s*=\s*["']([^"']+)["']"#).ok()?;
re.captures(config_text)
.and_then(|caps| caps.get(1))
.map(|m| m.as_str().trim().to_string())
.filter(|value| !value.is_empty())
}
} }
#[cfg(test)] #[cfg(test)]
+228 -155
View File
@@ -4,7 +4,7 @@
use crate::database::{lock_conn, Database}; use crate::database::{lock_conn, Database};
use crate::error::AppError; use crate::error::AppError;
use chrono::{Local, TimeZone}; use chrono::{Duration, Utc};
use rusqlite::{params, Connection, OptionalExtension}; use rusqlite::{params, Connection, OptionalExtension};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::Value; use serde_json::Value;
@@ -181,63 +181,29 @@ impl Database {
Ok(result) Ok(result)
} }
/// 获取每日趋势(滑动窗口,<=24h 按小时,>24h 按天,窗口与汇总一致) /// 获取每日趋势
pub fn get_daily_trends( pub fn get_daily_trends(&self, days: u32) -> Result<Vec<DailyStats>, AppError> {
&self,
start_date: Option<i64>,
end_date: Option<i64>,
) -> Result<Vec<DailyStats>, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
let end_ts = end_date.unwrap_or_else(|| Local::now().timestamp()); if days <= 1 {
let mut start_ts = start_date.unwrap_or_else(|| end_ts - 24 * 60 * 60); let sql = "SELECT
strftime('%Y-%m-%dT%H:00:00Z', datetime(created_at, 'unixepoch')) as bucket,
COUNT(*) as request_count,
COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost,
COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens,
COALESCE(SUM(input_tokens), 0) as total_input_tokens,
COALESCE(SUM(output_tokens), 0) as total_output_tokens,
COALESCE(SUM(cache_creation_tokens), 0) as total_cache_creation_tokens,
COALESCE(SUM(cache_read_tokens), 0) as total_cache_read_tokens
FROM proxy_request_logs
WHERE created_at >= strftime('%s', 'now', '-1 day')
GROUP BY bucket
ORDER BY bucket ASC";
if start_ts >= end_ts { let mut stmt = conn.prepare(sql)?;
start_ts = end_ts - 24 * 60 * 60; let rows = stmt.query_map([], |row| {
} Ok(DailyStats {
date: row.get(0)?,
let duration = end_ts - start_ts;
let bucket_seconds: i64 = if duration <= 24 * 60 * 60 {
60 * 60
} else {
24 * 60 * 60
};
let mut bucket_count: i64 = if duration <= 0 {
1
} else {
((duration as f64) / bucket_seconds as f64).ceil() as i64
};
// 固定 24 小时窗口为 24 个小时桶,避免浮点误差
if bucket_seconds == 60 * 60 {
bucket_count = 24;
}
if bucket_count < 1 {
bucket_count = 1;
}
let sql = "
SELECT
CAST((created_at - ?1) / ?3 AS INTEGER) as bucket_idx,
COUNT(*) as request_count,
COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost,
COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens,
COALESCE(SUM(input_tokens), 0) as total_input_tokens,
COALESCE(SUM(output_tokens), 0) as total_output_tokens,
COALESCE(SUM(cache_creation_tokens), 0) as total_cache_creation_tokens,
COALESCE(SUM(cache_read_tokens), 0) as total_cache_read_tokens
FROM proxy_request_logs
WHERE created_at >= ?1 AND created_at <= ?2
GROUP BY bucket_idx
ORDER BY bucket_idx ASC";
let mut stmt = conn.prepare(sql)?;
let rows = stmt.query_map(params![start_ts, end_ts, bucket_seconds], |row| {
Ok((
row.get::<_, i64>(0)?,
DailyStats {
date: String::new(),
request_count: row.get::<_, i64>(1)? as u64, request_count: row.get::<_, i64>(1)? as u64,
total_cost: format!("{:.6}", row.get::<_, f64>(2)?), total_cost: format!("{:.6}", row.get::<_, f64>(2)?),
total_tokens: row.get::<_, i64>(3)? as u64, total_tokens: row.get::<_, i64>(3)? as u64,
@@ -245,50 +211,99 @@ impl Database {
total_output_tokens: row.get::<_, i64>(5)? as u64, total_output_tokens: row.get::<_, i64>(5)? as u64,
total_cache_creation_tokens: row.get::<_, i64>(6)? as u64, total_cache_creation_tokens: row.get::<_, i64>(6)? as u64,
total_cache_read_tokens: row.get::<_, i64>(7)? as u64, total_cache_read_tokens: row.get::<_, i64>(7)? as u64,
}, })
)) })?;
})?;
let mut map: HashMap<i64, DailyStats> = HashMap::new(); let mut buckets: HashMap<String, DailyStats> = HashMap::new();
for row in rows { for row in rows {
let (mut bucket_idx, stat) = row?; let stat = row?;
if bucket_idx < 0 { buckets.insert(stat.date.clone(), stat);
continue;
} }
if bucket_idx >= bucket_count {
bucket_idx = bucket_count - 1; let mut stats = Vec::new();
let today = Utc::now().date_naive();
for hour in 0..24 {
let bucket = today
.and_hms_opt(hour, 0, 0)
.unwrap()
.format("%Y-%m-%dT%H:00:00Z")
.to_string();
if let Some(stat) = buckets.remove(&bucket) {
stats.push(stat);
} else {
stats.push(DailyStats {
date: bucket,
request_count: 0,
total_cost: "0.000000".to_string(),
total_tokens: 0,
total_input_tokens: 0,
total_output_tokens: 0,
total_cache_creation_tokens: 0,
total_cache_read_tokens: 0,
});
}
} }
map.insert(bucket_idx, stat); Ok(stats)
} else {
let sql = "SELECT
date(created_at, 'unixepoch') as bucket,
COUNT(*) as request_count,
COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost,
COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens,
COALESCE(SUM(input_tokens), 0) as total_input_tokens,
COALESCE(SUM(output_tokens), 0) as total_output_tokens,
COALESCE(SUM(cache_creation_tokens), 0) as total_cache_creation_tokens,
COALESCE(SUM(cache_read_tokens), 0) as total_cache_read_tokens
FROM proxy_request_logs
WHERE created_at >= strftime('%s', 'now', ?)
GROUP BY bucket
ORDER BY bucket ASC";
let mut stmt = conn.prepare(sql)?;
let rows = stmt.query_map([format!("-{days} days")], |row| {
Ok(DailyStats {
date: row.get(0)?,
request_count: row.get::<_, i64>(1)? as u64,
total_cost: format!("{:.6}", row.get::<_, f64>(2)?),
total_tokens: row.get::<_, i64>(3)? as u64,
total_input_tokens: row.get::<_, i64>(4)? as u64,
total_output_tokens: row.get::<_, i64>(5)? as u64,
total_cache_creation_tokens: row.get::<_, i64>(6)? as u64,
total_cache_read_tokens: row.get::<_, i64>(7)? as u64,
})
})?;
let mut map = HashMap::new();
for row in rows {
let stat = row?;
map.insert(stat.date.clone(), stat);
}
let mut stats = Vec::new();
let start_day =
Utc::now().date_naive() - Duration::days((days.saturating_sub(1)) as i64);
for i in 0..days {
let day = start_day + Duration::days(i as i64);
let key = day.format("%Y-%m-%d").to_string();
if let Some(stat) = map.remove(&key) {
stats.push(stat);
} else {
stats.push(DailyStats {
date: key,
request_count: 0,
total_cost: "0.000000".to_string(),
total_tokens: 0,
total_input_tokens: 0,
total_output_tokens: 0,
total_cache_creation_tokens: 0,
total_cache_read_tokens: 0,
});
}
}
Ok(stats)
} }
let mut stats = Vec::with_capacity(bucket_count as usize);
for i in 0..bucket_count {
let bucket_start_ts = start_ts + i * bucket_seconds;
let bucket_start = Local
.timestamp_opt(bucket_start_ts, 0)
.single()
.unwrap_or_else(Local::now);
let date = bucket_start.format("%Y-%m-%dT%H:%M:%S").to_string();
if let Some(mut stat) = map.remove(&i) {
stat.date = date;
stats.push(stat);
} else {
stats.push(DailyStats {
date,
request_count: 0,
total_cost: "0.000000".to_string(),
total_tokens: 0,
total_input_tokens: 0,
total_output_tokens: 0,
total_cache_creation_tokens: 0,
total_cache_read_tokens: 0,
});
}
}
Ok(stats)
} }
/// 获取 Provider 统计 /// 获取 Provider 统计
@@ -602,7 +617,7 @@ impl Database {
"SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) "SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0)
FROM proxy_request_logs FROM proxy_request_logs
WHERE provider_id = ? AND app_type = ? WHERE provider_id = ? AND app_type = ?
AND date(datetime(created_at, 'unixepoch', 'localtime')) = date('now', 'localtime')", AND date(created_at, 'unixepoch') = date('now')",
params![provider_id, app_type], params![provider_id, app_type],
|row| row.get(0), |row| row.get(0),
) )
@@ -614,7 +629,7 @@ impl Database {
"SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) "SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0)
FROM proxy_request_logs FROM proxy_request_logs
WHERE provider_id = ? AND app_type = ? WHERE provider_id = ? AND app_type = ?
AND strftime('%Y-%m', datetime(created_at, 'unixepoch', 'localtime')) = strftime('%Y-%m', 'now', 'localtime')", AND strftime('%Y-%m', created_at, 'unixepoch') = strftime('%Y-%m', 'now')",
params![provider_id, app_type], params![provider_id, app_type],
|row| row.get(0), |row| row.get(0),
) )
@@ -798,46 +813,89 @@ impl Database {
} }
} }
/// 标准化模型名称:去除供应商前缀并将点号替换为短横线
/// 例如:anthropic/claude-haiku-4.5 → claude-haiku-4-5
fn normalize_model_id(model_id: &str) -> String {
// 1. 去除供应商前缀(如 anthropic/、openai/
let stripped = if let Some(pos) = model_id.find('/') {
&model_id[pos + 1..]
} else {
model_id
};
// 2. 将点号替换为短横线(如 claude-haiku-4.5 → claude-haiku-4-5
stripped.replace('.', "-")
}
pub(crate) fn find_model_pricing_row( pub(crate) fn find_model_pricing_row(
conn: &Connection, conn: &Connection,
model_id: &str, model_id: &str,
) -> Result<Option<(String, String, String, String)>, AppError> { ) -> Result<Option<(String, String, String, String)>, AppError> {
// 1) 去除供应商前缀(/ 之前)与冒号后缀(: 之后),例如 moonshotai/kimi-k2-0905:exa → kimi-k2-0905 // 0. 标准化模型名称(去除前缀 + 点号转短横线)
let without_prefix = model_id // 例如:anthropic/claude-haiku-4.5 → claude-haiku-4-5
.rsplit_once('/') let normalized = normalize_model_id(model_id);
.map(|(_, rest)| rest)
.unwrap_or(model_id);
let cleaned = without_prefix
.split(':')
.next()
.map(str::trim)
.unwrap_or(without_prefix);
// 2) 精确匹配清洗后的名称 // 1. 精确匹配(先尝试原始名称,再尝试标准化后的名称
let exact = conn for id in [model_id, normalized.as_str()] {
.query_row( let exact = conn
"SELECT input_cost_per_million, output_cost_per_million, .query_row(
cache_read_cost_per_million, cache_creation_cost_per_million "SELECT input_cost_per_million, output_cost_per_million,
FROM model_pricing cache_read_cost_per_million, cache_creation_cost_per_million
WHERE model_id = ?1", FROM model_pricing
[cleaned], WHERE model_id = ?1",
|row| { [id],
Ok(( |row| {
row.get::<_, String>(0)?, Ok((
row.get::<_, String>(1)?, row.get::<_, String>(0)?,
row.get::<_, String>(2)?, row.get::<_, String>(1)?,
row.get::<_, String>(3)?, row.get::<_, String>(2)?,
)) row.get::<_, String>(3)?,
}, ))
) },
.optional() )
.map_err(|e| AppError::Database(format!("查询模型定价失败: {e}")))?; .optional()
.map_err(|e| AppError::Database(format!("查询模型定价失败: {e}")))?;
if exact.is_none() { if exact.is_some() {
log::warn!("模型 {model_id}(清洗后: {cleaned})未找到定价信息,成本将记录为 0"); if id != model_id {
log::info!("模型 {model_id} 标准化后精确匹配到: {id}");
}
return Ok(exact);
}
} }
Ok(exact) // 2. 逐步删除后缀匹配(claude-haiku-4-5-20250929 → claude-haiku-4-5 → claude-haiku-4 → claude-haiku
// 使用标准化后的名称进行后缀匹配
let mut current = normalized;
while let Some(pos) = current.rfind('-') {
current = current[..pos].to_string();
let result = conn
.query_row(
"SELECT input_cost_per_million, output_cost_per_million,
cache_read_cost_per_million, cache_creation_cost_per_million
FROM model_pricing
WHERE model_id = ?1",
[&current],
|row| {
Ok((
row.get::<_, String>(0)?,
row.get::<_, String>(1)?,
row.get::<_, String>(2)?,
row.get::<_, String>(3)?,
))
},
)
.optional()
.map_err(|e| AppError::Database(format!("查询模型定价失败: {e}")))?;
if result.is_some() {
log::info!("模型 {model_id} 通过删除后缀匹配到: {current}");
return Ok(result);
}
}
log::warn!("模型 {model_id} 未找到定价信息,成本将记录为 0");
Ok(None)
} }
#[cfg(test)] #[cfg(test)]
@@ -917,39 +975,54 @@ mod tests {
let db = Database::memory()?; let db = Database::memory()?;
let conn = lock_conn!(db.conn); let conn = lock_conn!(db.conn);
// 准备额外定价数据,覆盖前缀/后缀清洗场景 // 测试精确匹配
conn.execute( let result = find_model_pricing_row(&conn, "claude-sonnet-4-5")?;
"INSERT OR REPLACE INTO model_pricing ( assert!(result.is_some(), "应该能精确匹配 claude-sonnet-4-5");
model_id, display_name, input_cost_per_million, output_cost_per_million,
cache_read_cost_per_million, cache_creation_cost_per_million
) VALUES (?, ?, ?, ?, ?, ?)",
params![
"claude-haiku-4.5",
"Claude Haiku 4.5",
"1.0",
"2.0",
"0.0",
"0.0"
],
)?;
// 测试精确匹配(seed_model_pricing 已预置 claude-sonnet-4-5-20250929 // 测试带供应商前缀的模型名称(anthropic/claude-haiku-4.5 → claude-haiku-4-5
let result = find_model_pricing_row(&conn, "claude-sonnet-4-5-20250929")?;
assert!(
result.is_some(),
"应该能精确匹配 claude-sonnet-4-5-20250929"
);
// 清洗:去除前缀和冒号后缀
let result = find_model_pricing_row(&conn, "anthropic/claude-haiku-4.5")?; let result = find_model_pricing_row(&conn, "anthropic/claude-haiku-4.5")?;
assert!( assert!(
result.is_some(), result.is_some(),
"带前缀的模型 anthropic/claude-haiku-4.5 应能匹配到 claude-haiku-4.5" "应该能匹配带前缀的模型 anthropic/claude-haiku-4.5"
); );
let result = find_model_pricing_row(&conn, "moonshotai/kimi-k2-0905:exa")?;
// 测试带供应商前缀 + 点号的模型名称
let result = find_model_pricing_row(&conn, "anthropic/claude-sonnet-4.5")?;
assert!( assert!(
result.is_some(), result.is_some(),
"带前缀+冒号后缀的模型应清洗后匹配到 kimi-k2-0905" "应该能匹配带前缀的模型 anthropic/claude-sonnet-4.5"
);
// 测试逐步删除后缀匹配 - 日期后缀
let result = find_model_pricing_row(&conn, "claude-sonnet-4-5-20241022")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 claude-sonnet-4-5-20241022"
);
// 测试逐步删除后缀匹配 - 多个后缀
let result = find_model_pricing_row(&conn, "claude-haiku-4-5-20240229-preview")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 claude-haiku-4-5-20240229-preview"
);
// 测试 GPT 模型
let result = find_model_pricing_row(&conn, "gpt-5-2024-11-20")?;
assert!(result.is_some(), "应该能通过删除后缀匹配 gpt-5-2024-11-20");
// 测试 Gemini 模型
let result = find_model_pricing_row(&conn, "gemini-2.5-flash-exp")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 gemini-2.5-flash-exp"
);
// 测试 claude-sonnet-4-5 命名格式
let result = find_model_pricing_row(&conn, "claude-sonnet-4-5-20250929")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 claude-sonnet-4-5-20250929"
); );
// 测试不存在的模型 // 测试不存在的模型
-8
View File
@@ -31,9 +31,6 @@ pub struct AppSettings {
/// 是否启用 Claude 插件联动 /// 是否启用 Claude 插件联动
#[serde(default)] #[serde(default)]
pub enable_claude_plugin_integration: bool, pub enable_claude_plugin_integration: bool,
/// 是否跳过 Claude Code 初次安装确认
#[serde(default = "default_true")]
pub skip_claude_onboarding: bool,
/// 是否开机自启 /// 是否开机自启
#[serde(default)] #[serde(default)]
pub launch_on_startup: bool, pub launch_on_startup: bool,
@@ -68,17 +65,12 @@ fn default_minimize_to_tray_on_close() -> bool {
true true
} }
fn default_true() -> bool {
true
}
impl Default for AppSettings { impl Default for AppSettings {
fn default() -> Self { fn default() -> Self {
Self { Self {
show_in_tray: true, show_in_tray: true,
minimize_to_tray_on_close: true, minimize_to_tray_on_close: true,
enable_claude_plugin_integration: false, enable_claude_plugin_integration: false,
skip_claude_onboarding: true,
launch_on_startup: false, launch_on_startup: false,
language: None, language: None,
claude_config_dir: None, claude_config_dir: None,
+1 -1
View File
@@ -1,7 +1,7 @@
{ {
"$schema": "https://schema.tauri.app/config/2", "$schema": "https://schema.tauri.app/config/2",
"productName": "CC Switch", "productName": "CC Switch",
"version": "3.9.0-3", "version": "3.8.2",
"identifier": "com.ccswitch.desktop", "identifier": "com.ccswitch.desktop",
"build": { "build": {
"frontendDist": "../dist", "frontendDist": "../dist",
+1 -3
View File
@@ -4,9 +4,7 @@
"windows": [ "windows": [
{ {
"label": "main", "label": "main",
"titleBarStyle": "Visible", "titleBarStyle": "Visible"
"minWidth": 900,
"minHeight": 600
} }
] ]
} }
+16 -89
View File
@@ -76,8 +76,19 @@ fn sync_codex_provider_writes_auth_and_config() {
let mut config = MultiAppConfig::default(); let mut config = MultiAppConfig::default();
// 注意:v3.7.0 后 MCP 同步由 McpService 独立处理,不再通过 provider 切换触发 // 添加入测 MCP 启用项,确保 sync_enabled_to_codex 会写入 TOML
// 此测试仅验证 auth.json 和 config.toml 基础配置的写入 config.mcp.codex.servers.insert(
"echo-server".into(),
json!({
"id": "echo-server",
"enabled": true,
"server": {
"type": "stdio",
"command": "echo",
"args": ["hello"]
}
}),
);
let provider_config = json!({ let provider_config = json!({
"auth": { "auth": {
@@ -122,10 +133,9 @@ fn sync_codex_provider_writes_auth_and_config() {
); );
let toml_text = fs::read_to_string(&config_path).expect("read config.toml"); let toml_text = fs::read_to_string(&config_path).expect("read config.toml");
// 验证基础配置正确写入
assert!( assert!(
toml_text.contains("base_url"), toml_text.contains("command = \"echo\""),
"config.toml should contain base_url from provider config" "config.toml should contain serialized enabled MCP server"
); );
// 当前供应商应同步最新 config 文本 // 当前供应商应同步最新 config 文本
@@ -144,12 +154,6 @@ fn sync_enabled_to_codex_writes_enabled_servers() {
let _guard = test_mutex().lock().expect("acquire test mutex"); let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs(); reset_test_fs();
// 模拟 Codex 已安装/已初始化:存在 ~/.codex 目录
let path = cc_switch_lib::get_codex_config_path();
if let Some(parent) = path.parent() {
fs::create_dir_all(parent).expect("create codex dir");
}
let mut config = MultiAppConfig::default(); let mut config = MultiAppConfig::default();
config.mcp.codex.servers.insert( config.mcp.codex.servers.insert(
"stdio-enabled".into(), "stdio-enabled".into(),
@@ -166,6 +170,7 @@ fn sync_enabled_to_codex_writes_enabled_servers() {
cc_switch_lib::sync_enabled_to_codex(&config).expect("sync codex"); cc_switch_lib::sync_enabled_to_codex(&config).expect("sync codex");
let path = cc_switch_lib::get_codex_config_path();
assert!(path.exists(), "config.toml should be created"); assert!(path.exists(), "config.toml should be created");
let text = fs::read_to_string(&path).expect("read config.toml"); let text = fs::read_to_string(&path).expect("read config.toml");
assert!( assert!(
@@ -589,11 +594,6 @@ command = "echo"
fn sync_claude_enabled_mcp_projects_to_user_config() { fn sync_claude_enabled_mcp_projects_to_user_config() {
let _guard = test_mutex().lock().expect("acquire test mutex"); let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs(); reset_test_fs();
let home = ensure_test_home();
// 模拟 Claude 已安装/已初始化:存在 ~/.claude 目录
fs::create_dir_all(home.join(".claude")).expect("create claude dir");
let mut config = MultiAppConfig::default(); let mut config = MultiAppConfig::default();
config.mcp.claude.servers.insert( config.mcp.claude.servers.insert(
@@ -993,76 +993,3 @@ fn export_sql_returns_error_for_invalid_path() {
other => panic!("expected IoContext or Io error, got {other:?}"), other => panic!("expected IoContext or Io error, got {other:?}"),
} }
} }
#[test]
fn import_sql_rejects_non_cc_switch_backup() {
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
let state = create_test_state().expect("create test state");
let import_path = home.join("not-cc-switch.sql");
fs::write(&import_path, "CREATE TABLE x (id INTEGER);").expect("write import sql");
let err = state
.db
.import_sql(&import_path)
.expect_err("non-cc-switch sql should be rejected");
match err {
AppError::Localized { key, .. } => {
assert_eq!(key, "backup.sql.invalid_format");
}
other => panic!("expected Localized error, got {other:?}"),
}
}
#[test]
fn import_sql_accepts_cc_switch_exported_backup() {
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
// Create a database with some data and export it.
let mut config = MultiAppConfig::default();
{
let manager = config
.get_manager_mut(&AppType::Claude)
.expect("claude manager");
manager.current = "test-provider".to_string();
manager.providers.insert(
"test-provider".to_string(),
Provider::with_id(
"test-provider".to_string(),
"Test Provider".to_string(),
json!({"env": {"ANTHROPIC_API_KEY": "test-key"}}),
None,
),
);
}
let state = create_test_state_with_config(&config).expect("create test state");
let export_path = home.join("cc-switch-export.sql");
state
.db
.export_sql(&export_path)
.expect("export should succeed");
// Reset database, then import into a fresh one.
reset_test_fs();
let state = create_test_state().expect("create test state");
state
.db
.import_sql(&export_path)
.expect("import should succeed");
let providers = state
.db
.get_all_providers(AppType::Claude.as_str())
.expect("load providers");
assert!(
providers.contains_key("test-provider"),
"imported providers should contain test-provider"
);
}
-308
View File
@@ -246,311 +246,3 @@ fn set_mcp_enabled_for_codex_writes_live_config() {
"codex config should include the enabled server definition" "codex config should include the enabled server definition"
); );
} }
#[test]
fn enabling_codex_mcp_skips_when_codex_dir_missing() {
use support::create_test_state;
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
// 确认 Codex 配置目录不存在(模拟“未安装/未运行过 Codex CLI”)
assert!(
!home.join(".codex").exists(),
"~/.codex should not exist in fresh test environment"
);
let state = create_test_state().expect("create test state");
// 先插入一个未启用 Codex 的 MCP 服务器(避免 upsert 触发同步)
McpService::upsert_server(
&state,
McpServer {
id: "codex-server".to_string(),
name: "Codex Server".to_string(),
server: json!({
"type": "stdio",
"command": "echo"
}),
apps: McpApps {
claude: false,
codex: false,
gemini: false,
},
description: None,
homepage: None,
docs: None,
tags: Vec::new(),
},
)
.expect("insert server without syncing");
// 启用 Codex:目录缺失时应跳过写入(不创建 ~/.codex/config.toml
McpService::toggle_app(&state, "codex-server", AppType::Codex, true)
.expect("toggle codex should succeed even when ~/.codex is missing");
assert!(
!home.join(".codex").exists(),
"~/.codex should still not exist after skipped sync"
);
}
#[test]
fn upsert_mcp_server_disabling_app_removes_from_claude_live_config() {
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
// 模拟 Claude 已安装/已初始化:存在 ~/.claude 目录
fs::create_dir_all(home.join(".claude")).expect("create ~/.claude dir");
// 先创建一个启用 Claude 的 MCP 服务器
let state = support::create_test_state().expect("create test state");
McpService::upsert_server(
&state,
McpServer {
id: "echo".to_string(),
name: "echo".to_string(),
server: json!({
"type": "stdio",
"command": "echo"
}),
apps: McpApps {
claude: true,
codex: false,
gemini: false,
},
description: None,
homepage: None,
docs: None,
tags: Vec::new(),
},
)
.expect("upsert should sync to Claude live config");
// 确认已写入 ~/.claude.json
let mcp_path = get_claude_mcp_path();
let text = fs::read_to_string(&mcp_path).expect("read ~/.claude.json");
let v: serde_json::Value = serde_json::from_str(&text).expect("parse ~/.claude.json");
assert!(
v.pointer("/mcpServers/echo").is_some(),
"echo should exist in Claude live config after enabling"
);
// 再次 upsert:取消勾选 Claudeapps.claude=false),应从 Claude live 配置中移除
McpService::upsert_server(
&state,
McpServer {
id: "echo".to_string(),
name: "echo".to_string(),
server: json!({
"type": "stdio",
"command": "echo"
}),
apps: McpApps {
claude: false,
codex: false,
gemini: false,
},
description: None,
homepage: None,
docs: None,
tags: Vec::new(),
},
)
.expect("upsert disabling app should remove from Claude live config");
let text = fs::read_to_string(&mcp_path).expect("read ~/.claude.json after disable");
let v: serde_json::Value = serde_json::from_str(&text).expect("parse ~/.claude.json");
assert!(
v.pointer("/mcpServers/echo").is_none(),
"echo should be removed from Claude live config after disabling"
);
}
#[test]
fn import_mcp_from_multiple_apps_merges_enabled_flags() {
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
// 1) Claude: ~/.claude.json
let mcp_path = get_claude_mcp_path();
let claude_json = json!({
"mcpServers": {
"shared": {
"type": "stdio",
"command": "echo"
}
}
});
fs::write(
&mcp_path,
serde_json::to_string_pretty(&claude_json).expect("serialize claude mcp"),
)
.expect("seed ~/.claude.json");
// 2) Codex: ~/.codex/config.toml
let codex_dir = home.join(".codex");
fs::create_dir_all(&codex_dir).expect("create codex dir");
fs::write(
codex_dir.join("config.toml"),
r#"[mcp_servers.shared]
type = "stdio"
command = "echo"
"#,
)
.expect("seed ~/.codex/config.toml");
let state = support::create_test_state().expect("create test state");
McpService::import_from_claude(&state).expect("import from claude");
McpService::import_from_codex(&state).expect("import from codex");
let servers = state.db.get_all_mcp_servers().expect("get all mcp servers");
let entry = servers.get("shared").expect("shared server exists");
assert!(entry.apps.claude, "shared should enable Claude");
assert!(entry.apps.codex, "shared should enable Codex");
}
#[test]
fn import_mcp_from_gemini_sse_url_only_is_valid() {
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
// Gemini MCP 位于 ~/.gemini/settings.json
let gemini_dir = home.join(".gemini");
fs::create_dir_all(&gemini_dir).expect("create gemini dir");
let settings_path = gemini_dir.join("settings.json");
// Gemini SSE:只包含 urlGemini 不使用 type 字段)
let gemini_settings = json!({
"mcpServers": {
"sse-server": {
"url": "https://example.com/sse"
}
}
});
fs::write(
&settings_path,
serde_json::to_string_pretty(&gemini_settings).expect("serialize gemini settings"),
)
.expect("seed ~/.gemini/settings.json");
let state = support::create_test_state().expect("create test state");
let changed = McpService::import_from_gemini(&state).expect("import from gemini");
assert!(changed > 0, "should import at least 1 server");
let servers = state.db.get_all_mcp_servers().expect("get all mcp servers");
let entry = servers.get("sse-server").expect("sse-server exists");
assert!(entry.apps.gemini, "imported server should enable Gemini");
assert_eq!(
entry.server.get("type").and_then(|v| v.as_str()),
Some("sse"),
"Gemini url-only server should be normalized to type=sse in unified structure"
);
}
#[test]
fn enabling_gemini_mcp_skips_when_gemini_dir_missing() {
use support::create_test_state;
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
// 确认 Gemini 配置目录不存在(模拟“未安装/未运行过 Gemini CLI”)
assert!(
!home.join(".gemini").exists(),
"~/.gemini should not exist in fresh test environment"
);
let state = create_test_state().expect("create test state");
// 先插入一个未启用 Gemini 的 MCP 服务器(避免 upsert 触发同步)
McpService::upsert_server(
&state,
McpServer {
id: "gemini-server".to_string(),
name: "Gemini Server".to_string(),
server: json!({
"type": "sse",
"url": "https://example.com/sse"
}),
apps: McpApps {
claude: false,
codex: false,
gemini: false,
},
description: None,
homepage: None,
docs: None,
tags: Vec::new(),
},
)
.expect("insert server without syncing");
// 启用 Gemini:目录缺失时应跳过写入(不创建 ~/.gemini/settings.json
McpService::toggle_app(&state, "gemini-server", AppType::Gemini, true)
.expect("toggle gemini should succeed even when ~/.gemini is missing");
assert!(
!home.join(".gemini").exists(),
"~/.gemini should still not exist after skipped sync"
);
}
#[test]
fn enabling_claude_mcp_skips_when_claude_config_absent() {
use support::create_test_state;
let _guard = test_mutex().lock().expect("acquire test mutex");
reset_test_fs();
let home = ensure_test_home();
// 确认 Claude 相关目录/文件都不存在(模拟“未安装/未运行过 Claude”)
assert!(
!home.join(".claude").exists(),
"~/.claude should not exist in fresh test environment"
);
assert!(
!home.join(".claude.json").exists(),
"~/.claude.json should not exist in fresh test environment"
);
let state = create_test_state().expect("create test state");
// 先插入一个未启用 Claude 的 MCP 服务器(避免 upsert 触发同步)
McpService::upsert_server(
&state,
McpServer {
id: "claude-server".to_string(),
name: "Claude Server".to_string(),
server: json!({
"type": "stdio",
"command": "echo"
}),
apps: McpApps {
claude: false,
codex: false,
gemini: false,
},
description: None,
homepage: None,
docs: None,
tags: Vec::new(),
},
)
.expect("insert server without syncing");
// 启用 Claude:配置缺失时应跳过写入(不创建 ~/.claude.json
McpService::toggle_app(&state, "claude-server", AppType::Claude, true)
.expect("toggle claude should succeed even when ~/.claude is missing");
assert!(
!home.join(".claude.json").exists(),
"~/.claude.json should still not exist after skipped sync"
);
}
-2
View File
@@ -49,7 +49,6 @@ pub fn test_mutex() -> &'static Mutex<()> {
} }
/// 创建测试用的 AppState,包含一个空的数据库 /// 创建测试用的 AppState,包含一个空的数据库
#[allow(dead_code)]
pub fn create_test_state() -> Result<AppState, Box<dyn std::error::Error>> { pub fn create_test_state() -> Result<AppState, Box<dyn std::error::Error>> {
let db = Arc::new(Database::init()?); let db = Arc::new(Database::init()?);
let proxy_service = ProxyService::new(db.clone()); let proxy_service = ProxyService::new(db.clone());
@@ -57,7 +56,6 @@ pub fn create_test_state() -> Result<AppState, Box<dyn std::error::Error>> {
} }
/// 创建测试用的 AppState,并从 MultiAppConfig 迁移数据 /// 创建测试用的 AppState,并从 MultiAppConfig 迁移数据
#[allow(dead_code)]
pub fn create_test_state_with_config( pub fn create_test_state_with_config(
config: &MultiAppConfig, config: &MultiAppConfig,
) -> Result<AppState, Box<dyn std::error::Error>> { ) -> Result<AppState, Box<dyn std::error::Error>> {
+118 -258
View File
@@ -1,9 +1,7 @@
import { useEffect, useMemo, useState, useRef } from "react"; import { useEffect, useMemo, useState, useRef } from "react";
import { useTranslation } from "react-i18next"; import { useTranslation } from "react-i18next";
import { motion, AnimatePresence } from "framer-motion";
import { toast } from "sonner"; import { toast } from "sonner";
import { invoke } from "@tauri-apps/api/core"; import { invoke } from "@tauri-apps/api/core";
import { useQueryClient } from "@tanstack/react-query";
import { import {
Plus, Plus,
Settings, Settings,
@@ -26,7 +24,6 @@ import {
import { checkAllEnvConflicts, checkEnvConflicts } from "@/lib/api/env"; import { checkAllEnvConflicts, checkEnvConflicts } from "@/lib/api/env";
import { useProviderActions } from "@/hooks/useProviderActions"; import { useProviderActions } from "@/hooks/useProviderActions";
import { useProxyStatus } from "@/hooks/useProxyStatus"; import { useProxyStatus } from "@/hooks/useProxyStatus";
import { useLastValidValue } from "@/hooks/useLastValidValue";
import { extractErrorMessage } from "@/utils/errorUtils"; import { extractErrorMessage } from "@/utils/errorUtils";
import { cn } from "@/lib/utils"; import { cn } from "@/lib/utils";
import { AppSwitcher } from "@/components/AppSwitcher"; import { AppSwitcher } from "@/components/AppSwitcher";
@@ -44,25 +41,12 @@ import PromptPanel from "@/components/prompts/PromptPanel";
import { SkillsPage } from "@/components/skills/SkillsPage"; import { SkillsPage } from "@/components/skills/SkillsPage";
import { DeepLinkImportDialog } from "@/components/DeepLinkImportDialog"; import { DeepLinkImportDialog } from "@/components/DeepLinkImportDialog";
import { AgentsPanel } from "@/components/agents/AgentsPanel"; import { AgentsPanel } from "@/components/agents/AgentsPanel";
import { UniversalProviderPanel } from "@/components/universal";
import { Button } from "@/components/ui/button"; import { Button } from "@/components/ui/button";
type View = type View = "providers" | "settings" | "prompts" | "skills" | "mcp" | "agents";
| "providers"
| "settings"
| "prompts"
| "skills"
| "mcp"
| "agents"
| "universal";
const DRAG_BAR_HEIGHT = 28; // px
const HEADER_HEIGHT = 64; // px
const CONTENT_TOP_OFFSET = DRAG_BAR_HEIGHT + HEADER_HEIGHT;
function App() { function App() {
const { t } = useTranslation(); const { t } = useTranslation();
const queryClient = useQueryClient();
const [activeApp, setActiveApp] = useState<AppId>("claude"); const [activeApp, setActiveApp] = useState<AppId>("claude");
const [currentView, setCurrentView] = useState<View>("providers"); const [currentView, setCurrentView] = useState<View>("providers");
@@ -74,10 +58,6 @@ function App() {
const [envConflicts, setEnvConflicts] = useState<EnvConflict[]>([]); const [envConflicts, setEnvConflicts] = useState<EnvConflict[]>([]);
const [showEnvBanner, setShowEnvBanner] = useState(false); const [showEnvBanner, setShowEnvBanner] = useState(false);
// 使用 Hook 保存最后有效值,用于动画退出期间保持内容显示
const effectiveEditingProvider = useLastValidValue(editingProvider);
const effectiveUsageProvider = useLastValidValue(usageProvider);
const promptPanelRef = useRef<any>(null); const promptPanelRef = useRef<any>(null);
const mcpPanelRef = useRef<any>(null); const mcpPanelRef = useRef<any>(null);
const skillsPageRef = useRef<any>(null); const skillsPageRef = useRef<any>(null);
@@ -85,20 +65,7 @@ function App() {
"bg-orange-500 hover:bg-orange-600 dark:bg-orange-500 dark:hover:bg-orange-600 text-white shadow-lg shadow-orange-500/30 dark:shadow-orange-500/40 rounded-full w-8 h-8"; "bg-orange-500 hover:bg-orange-600 dark:bg-orange-500 dark:hover:bg-orange-600 text-white shadow-lg shadow-orange-500/30 dark:shadow-orange-500/40 rounded-full w-8 h-8";
// 获取代理服务状态 // 获取代理服务状态
const { const { isRunning: isProxyRunning, isTakeoverActive } = useProxyStatus();
isRunning: isProxyRunning,
takeoverStatus,
status: proxyStatus,
} = useProxyStatus();
// 当前应用的代理是否开启
const isCurrentAppTakeoverActive = takeoverStatus?.[activeApp] || false;
// 当前应用代理实际使用的供应商 ID(从 active_targets 中获取)
const activeProviderId = useMemo(() => {
const target = proxyStatus?.active_targets?.find(
(t) => t.app_type === activeApp,
);
return target?.provider_id;
}, [proxyStatus?.active_targets, activeApp]);
// 获取供应商列表,当代理服务运行时自动刷新 // 获取供应商列表,当代理服务运行时自动刷新
const { data, isLoading, refetch } = useProvidersQuery(activeApp, { const { data, isLoading, refetch } = useProvidersQuery(activeApp, {
@@ -142,38 +109,6 @@ function App() {
}; };
}, [activeApp, refetch]); }, [activeApp, refetch]);
// 监听统一供应商同步事件,刷新所有应用的供应商列表
useEffect(() => {
let unsubscribe: (() => void) | undefined;
const setupListener = async () => {
try {
const { listen } = await import("@tauri-apps/api/event");
unsubscribe = await listen("universal-provider-synced", async () => {
// 统一供应商同步后刷新所有应用的供应商列表
// 使用 invalidateQueries 使所有 providers 查询失效
await queryClient.invalidateQueries({ queryKey: ["providers"] });
// 同时更新托盘菜单
try {
await providersApi.updateTrayMenu();
} catch (error) {
console.error("[App] Failed to update tray menu", error);
}
});
} catch (error) {
console.error(
"[App] Failed to subscribe universal-provider-synced event",
error,
);
}
};
setupListener();
return () => {
unsubscribe?.();
};
}, [queryClient]);
// 应用启动时检测所有应用的环境变量冲突 // 应用启动时检测所有应用的环境变量冲突
useEffect(() => { useEffect(() => {
const checkEnvOnStartup = async () => { const checkEnvOnStartup = async () => {
@@ -251,21 +186,6 @@ function App() {
checkEnvOnSwitch(); checkEnvOnSwitch();
}, [activeApp]); }, [activeApp]);
useEffect(() => {
const handleGlobalShortcut = (event: KeyboardEvent) => {
if (event.key !== "," || !(event.metaKey || event.ctrlKey)) {
return;
}
event.preventDefault();
setCurrentView("settings");
};
window.addEventListener("keydown", handleGlobalShortcut);
return () => {
window.removeEventListener("keydown", handleGlobalShortcut);
};
}, []);
// 打开网站链接 // 打开网站链接
const handleOpenWebsite = async (url: string) => { const handleOpenWebsite = async (url: string) => {
try { try {
@@ -348,20 +268,7 @@ function App() {
// 导入配置成功后刷新 // 导入配置成功后刷新
const handleImportSuccess = async () => { const handleImportSuccess = async () => {
try { await refetch();
// 导入会影响所有应用的供应商数据:刷新所有 providers 缓存
await queryClient.invalidateQueries({
queryKey: ["providers"],
refetchType: "all",
});
await queryClient.refetchQueries({
queryKey: ["providers"],
type: "all",
});
} catch (error) {
console.error("[App] Failed to refresh providers after import", error);
await refetch();
}
try { try {
await providersApi.updateTrayMenu(); await providersApi.updateTrayMenu();
} catch (error) { } catch (error) {
@@ -370,115 +277,79 @@ function App() {
}; };
const renderContent = () => { const renderContent = () => {
const content = (() => { switch (currentView) {
switch (currentView) { case "settings":
case "settings": return (
return ( <SettingsPage
<SettingsPage open={true}
open={true} onOpenChange={() => setCurrentView("providers")}
onOpenChange={() => setCurrentView("providers")} onImportSuccess={handleImportSuccess}
onImportSuccess={handleImportSuccess} />
/> );
); case "prompts":
case "prompts": return (
return ( <PromptPanel
<PromptPanel ref={promptPanelRef}
ref={promptPanelRef} open={true}
open={true} onOpenChange={() => setCurrentView("providers")}
onOpenChange={() => setCurrentView("providers")} appId={activeApp}
appId={activeApp} />
/> );
); case "skills":
case "skills": return (
return ( <SkillsPage
<SkillsPage ref={skillsPageRef}
ref={skillsPageRef} onClose={() => setCurrentView("providers")}
onClose={() => setCurrentView("providers")} initialApp={activeApp}
initialApp={activeApp} />
/> );
); case "mcp":
case "mcp": return (
return ( <UnifiedMcpPanel
<UnifiedMcpPanel ref={mcpPanelRef}
ref={mcpPanelRef} onOpenChange={() => setCurrentView("providers")}
onOpenChange={() => setCurrentView("providers")} />
/> );
); case "agents":
case "agents": return <AgentsPanel onOpenChange={() => setCurrentView("providers")} />;
return ( default:
<AgentsPanel onOpenChange={() => setCurrentView("providers")} /> return (
); <div className="mx-auto max-w-[56rem] px-5 flex flex-col h-[calc(100vh-8rem)] overflow-hidden">
case "universal": {/* 独立滚动容器 - 解决 Linux/Ubuntu 下 DndContext 与滚轮事件冲突 */}
return ( <div className="flex-1 overflow-y-auto overflow-x-hidden pb-12 px-1">
<div className="mx-auto max-w-[56rem] px-5 pt-4"> <div className="space-y-4">
<UniversalProviderPanel /> <ProviderList
</div> providers={providers}
); currentProviderId={currentProviderId}
default: appId={activeApp}
return ( isLoading={isLoading}
<div className="mx-auto max-w-[56rem] px-5 flex flex-col h-[calc(100vh-8rem)] overflow-hidden"> isProxyRunning={isProxyRunning}
{/* 独立滚动容器 - 解决 Linux/Ubuntu 下 DndContext 与滚轮事件冲突 */} isProxyTakeover={isProxyRunning && isTakeoverActive}
<div className="flex-1 overflow-y-auto overflow-x-hidden pb-12 px-1"> onSwitch={switchProvider}
<AnimatePresence mode="wait"> onEdit={setEditingProvider}
<motion.div onDelete={setConfirmDelete}
key={activeApp} onDuplicate={handleDuplicateProvider}
initial={{ opacity: 0 }} onConfigureUsage={setUsageProvider}
animate={{ opacity: 1 }} onOpenWebsite={handleOpenWebsite}
exit={{ opacity: 0 }} onCreate={() => setIsAddOpen(true)}
transition={{ duration: 0.15 }} />
className="space-y-4"
>
<ProviderList
providers={providers}
currentProviderId={currentProviderId}
appId={activeApp}
isLoading={isLoading}
isProxyRunning={isProxyRunning}
isProxyTakeover={
isProxyRunning && isCurrentAppTakeoverActive
}
activeProviderId={activeProviderId}
onSwitch={switchProvider}
onEdit={setEditingProvider}
onDelete={setConfirmDelete}
onDuplicate={handleDuplicateProvider}
onConfigureUsage={setUsageProvider}
onOpenWebsite={handleOpenWebsite}
onCreate={() => setIsAddOpen(true)}
/>
</motion.div>
</AnimatePresence>
</div> </div>
</div> </div>
); </div>
} );
})(); }
return (
<AnimatePresence mode="wait">
<motion.div
key={currentView}
initial={{ opacity: 0 }}
animate={{ opacity: 1 }}
exit={{ opacity: 0 }}
transition={{ duration: 0.2 }}
>
{content}
</motion.div>
</AnimatePresence>
);
}; };
return ( return (
<div <div
className="flex flex-col h-screen overflow-hidden bg-background text-foreground selection:bg-primary/30" className="flex min-h-screen flex-col bg-background text-foreground selection:bg-primary/30"
style={{ overflowX: "hidden", paddingTop: CONTENT_TOP_OFFSET }} style={{ overflowX: "hidden" }}
> >
{/* 全局拖拽区域(顶部 28px),避免上边框无法拖动 */} {/* 全局拖拽区域(顶部 4px),避免上边框无法拖动 */}
<div <div
className="fixed top-0 left-0 right-0 z-[60]" className="fixed top-0 left-0 right-0 h-4 z-[60]"
data-tauri-drag-region data-tauri-drag-region
style={{ WebkitAppRegion: "drag", height: DRAG_BAR_HEIGHT } as any} style={{ WebkitAppRegion: "drag" } as any}
/> />
{/* 环境变量警告横幅 */} {/* 环境变量警告横幅 */}
{showEnvBanner && envConflicts.length > 0 && ( {showEnvBanner && envConflicts.length > 0 && (
@@ -508,18 +379,13 @@ function App() {
)} )}
<header <header
className="fixed z-50 w-full transition-all duration-300 bg-background/80 backdrop-blur-md" className="fixed top-0 z-50 w-full py-3 bg-background/80 backdrop-blur-md transition-all duration-300"
data-tauri-drag-region data-tauri-drag-region
style={ style={{ WebkitAppRegion: "drag" } as any}
{
WebkitAppRegion: "drag",
top: DRAG_BAR_HEIGHT,
height: HEADER_HEIGHT,
} as any
}
> >
<div className="h-4 w-full" aria-hidden data-tauri-drag-region />
<div <div
className="mx-auto flex h-full max-w-[56rem] flex-wrap items-center justify-between gap-2 px-6" className="mx-auto max-w-[56rem] px-6 flex flex-wrap items-center justify-between gap-2"
data-tauri-drag-region data-tauri-drag-region
style={{ WebkitAppRegion: "drag" } as any} style={{ WebkitAppRegion: "drag" } as any}
> >
@@ -535,7 +401,7 @@ function App() {
onClick={() => setCurrentView("providers")} onClick={() => setCurrentView("providers")}
className="mr-2 rounded-lg" className="mr-2 rounded-lg"
> >
<ArrowLeft className="w-4 h-4" /> <ArrowLeft className="h-4 w-4" />
</Button> </Button>
<h1 className="text-lg font-semibold"> <h1 className="text-lg font-semibold">
{currentView === "settings" && t("settings.title")} {currentView === "settings" && t("settings.title")}
@@ -544,10 +410,6 @@ function App() {
{currentView === "skills" && t("skills.title")} {currentView === "skills" && t("skills.title")}
{currentView === "mcp" && t("mcp.unifiedPanel.title")} {currentView === "mcp" && t("mcp.unifiedPanel.title")}
{currentView === "agents" && t("agents.title")} {currentView === "agents" && t("agents.title")}
{currentView === "universal" &&
t("universalProvider.title", {
defaultValue: "统一供应商",
})}
</h1> </h1>
</div> </div>
) : ( ) : (
@@ -559,7 +421,7 @@ function App() {
rel="noreferrer" rel="noreferrer"
className={cn( className={cn(
"text-xl font-semibold transition-colors", "text-xl font-semibold transition-colors",
isProxyRunning && isCurrentAppTakeoverActive isProxyRunning && isTakeoverActive
? "text-emerald-500 hover:text-emerald-600 dark:text-emerald-400 dark:hover:text-emerald-300" ? "text-emerald-500 hover:text-emerald-600 dark:text-emerald-400 dark:hover:text-emerald-300"
: "text-blue-500 hover:text-blue-600 dark:text-blue-400 dark:hover:text-blue-300", : "text-blue-500 hover:text-blue-600 dark:text-blue-400 dark:hover:text-blue-300",
)} )}
@@ -573,7 +435,7 @@ function App() {
title={t("common.settings")} title={t("common.settings")}
className="hover:bg-black/5 dark:hover:bg-white/5" className="hover:bg-black/5 dark:hover:bg-white/5"
> >
<Settings className="w-4 h-4" /> <Settings className="h-4 w-4" />
</Button> </Button>
</div> </div>
<UpdateBadge onClick={() => setCurrentView("settings")} /> <UpdateBadge onClick={() => setCurrentView("settings")} />
@@ -582,27 +444,27 @@ function App() {
</div> </div>
<div <div
className="flex items-center gap-2 h-[32px]" className="flex items-center gap-2"
style={{ WebkitAppRegion: "no-drag" } as any} style={{ WebkitAppRegion: "no-drag" } as any}
> >
{currentView === "prompts" && ( {currentView === "prompts" && (
<Button <Button
size="icon" size="icon"
onClick={() => promptPanelRef.current?.openAdd()} onClick={() => promptPanelRef.current?.openAdd()}
className={`ml-auto ${addActionButtonClass}`} className={addActionButtonClass}
title={t("prompts.add")} title={t("prompts.add")}
> >
<Plus className="w-5 h-5" /> <Plus className="h-5 w-5" />
</Button> </Button>
)} )}
{currentView === "mcp" && ( {currentView === "mcp" && (
<Button <Button
size="icon" size="icon"
onClick={() => mcpPanelRef.current?.openAdd()} onClick={() => mcpPanelRef.current?.openAdd()}
className={`ml-auto ${addActionButtonClass}`} className={addActionButtonClass}
title={t("mcp.unifiedPanel.addServer")} title={t("mcp.unifiedPanel.addServer")}
> >
<Plus className="w-5 h-5" /> <Plus className="h-5 w-5" />
</Button> </Button>
)} )}
{currentView === "skills" && ( {currentView === "skills" && (
@@ -613,7 +475,7 @@ function App() {
onClick={() => skillsPageRef.current?.refresh()} onClick={() => skillsPageRef.current?.refresh()}
className="hover:bg-black/5 dark:hover:bg-white/5" className="hover:bg-black/5 dark:hover:bg-white/5"
> >
<RefreshCw className="w-4 h-4 mr-2" /> <RefreshCw className="h-4 w-4 mr-2" />
{t("skills.refresh")} {t("skills.refresh")}
</Button> </Button>
<Button <Button
@@ -622,45 +484,41 @@ function App() {
onClick={() => skillsPageRef.current?.openRepoManager()} onClick={() => skillsPageRef.current?.openRepoManager()}
className="hover:bg-black/5 dark:hover:bg-white/5" className="hover:bg-black/5 dark:hover:bg-white/5"
> >
<Settings className="w-4 h-4 mr-2" /> <Settings className="h-4 w-4 mr-2" />
{t("skills.repoManager")} {t("skills.repoManager")}
</Button> </Button>
</> </>
)} )}
{currentView === "providers" && ( {currentView === "providers" && (
<> <>
<ProxyToggle activeApp={activeApp} /> <ProxyToggle />
<AppSwitcher activeApp={activeApp} onSwitch={setActiveApp} /> <AppSwitcher activeApp={activeApp} onSwitch={setActiveApp} />
<div className="flex items-center gap-1 p-1 bg-muted rounded-xl"> <div className="bg-muted p-1 rounded-xl flex items-center gap-1">
<Button {hasSkillsSupport && (
variant="ghost" <Button
size="sm" variant="ghost"
onClick={() => setCurrentView("skills")} size="sm"
className={cn( onClick={() => setCurrentView("skills")}
"text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5", className="text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5"
"transition-all duration-200 ease-in-out overflow-hidden", title={t("skills.manage")}
hasSkillsSupport >
? "opacity-100 w-8 scale-100 px-2" <Wrench className="h-4 w-4" />
: "opacity-0 w-0 scale-75 pointer-events-none px-0 -ml-1", </Button>
)} )}
title={t("skills.manage")}
>
<Wrench className="flex-shrink-0 w-4 h-4" />
</Button>
{/* TODO: Agents 功能开发中,暂时隐藏入口 */} {/* TODO: Agents 功能开发中,暂时隐藏入口 */}
{/* {isClaudeApp && ( {/* {isClaudeApp && (
<Button <Button
variant="ghost" variant="ghost"
size="sm" size="sm"
onClick={() => setCurrentView("agents")} onClick={() => setCurrentView("agents")}
className="text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5" className="text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5"
title="Agents" title="Agents"
> >
<Bot className="w-4 h-4" /> <Bot className="h-4 w-4" />
</Button> </Button>
)} */} )} */}
<Button <Button
variant="ghost" variant="ghost"
size="sm" size="sm"
@@ -668,7 +526,7 @@ function App() {
className="text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5" className="text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5"
title={t("prompts.manage")} title={t("prompts.manage")}
> >
<Book className="w-4 h-4" /> <Book className="h-4 w-4" />
</Button> </Button>
<Button <Button
variant="ghost" variant="ghost"
@@ -677,7 +535,7 @@ function App() {
className="text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5" className="text-muted-foreground hover:text-foreground hover:bg-black/5 dark:hover:bg-white/5"
title={t("mcp.title")} title={t("mcp.title")}
> >
<Server className="w-4 h-4" /> <Server className="h-4 w-4" />
</Button> </Button>
</div> </div>
@@ -686,7 +544,7 @@ function App() {
size="icon" size="icon"
className={`ml-2 ${addActionButtonClass}`} className={`ml-2 ${addActionButtonClass}`}
> >
<Plus className="w-5 h-5" /> <Plus className="h-5 w-5" />
</Button> </Button>
</> </>
)} )}
@@ -694,8 +552,13 @@ function App() {
</div> </div>
</header> </header>
<main className="flex-1 pb-12 animate-fade-in "> <main
<div className="pb-12">{renderContent()}</div> className={`flex-1 overflow-y-auto pb-12 animate-fade-in scroll-overlay ${
currentView === "providers" ? "pt-24" : "pt-20"
}`}
style={{ overflowX: "hidden" }}
>
{renderContent()}
</main> </main>
<AddProviderDialog <AddProviderDialog
@@ -707,7 +570,7 @@ function App() {
<EditProviderDialog <EditProviderDialog
open={Boolean(editingProvider)} open={Boolean(editingProvider)}
provider={effectiveEditingProvider} provider={editingProvider}
onOpenChange={(open) => { onOpenChange={(open) => {
if (!open) { if (!open) {
setEditingProvider(null); setEditingProvider(null);
@@ -715,19 +578,16 @@ function App() {
}} }}
onSubmit={handleEditProvider} onSubmit={handleEditProvider}
appId={activeApp} appId={activeApp}
isProxyTakeover={isProxyRunning && isCurrentAppTakeoverActive}
/> />
{effectiveUsageProvider && ( {usageProvider && (
<UsageScriptModal <UsageScriptModal
provider={effectiveUsageProvider} provider={usageProvider}
appId={activeApp} appId={activeApp}
isOpen={Boolean(usageProvider)} isOpen={Boolean(usageProvider)}
onClose={() => setUsageProvider(null)} onClose={() => setUsageProvider(null)}
onSave={(script) => { onSave={(script) => {
if (usageProvider) { void saveUsageScript(usageProvider, script);
void saveUsageScript(usageProvider, script);
}
}} }}
/> />
)} )}
+4 -4
View File
@@ -24,11 +24,11 @@ export function AppSwitcher({ activeApp, onSwitch }: AppSwitcherProps) {
}; };
return ( return (
<div className="inline-flex bg-muted rounded-xl p-1 gap-1"> <div className="inline-flex bg-muted rounded-lg p-1 gap-1">
<button <button
type="button" type="button"
onClick={() => handleSwitch("claude")} onClick={() => handleSwitch("claude")}
className={`group inline-flex items-center gap-2 px-3 h-8 rounded-md text-sm font-medium transition-all duration-200 ${ className={`group inline-flex items-center gap-2 px-3 py-2 rounded-md text-sm font-medium transition-all duration-200 ${
activeApp === "claude" activeApp === "claude"
? "bg-background text-foreground shadow-sm" ? "bg-background text-foreground shadow-sm"
: "text-muted-foreground hover:text-foreground hover:bg-background/50" : "text-muted-foreground hover:text-foreground hover:bg-background/50"
@@ -50,7 +50,7 @@ export function AppSwitcher({ activeApp, onSwitch }: AppSwitcherProps) {
<button <button
type="button" type="button"
onClick={() => handleSwitch("codex")} onClick={() => handleSwitch("codex")}
className={`group inline-flex items-center gap-2 px-3 h-8 rounded-md text-sm font-medium transition-all duration-200 ${ className={`group inline-flex items-center gap-2 px-3 py-2 rounded-md text-sm font-medium transition-all duration-200 ${
activeApp === "codex" activeApp === "codex"
? "bg-background text-foreground shadow-sm" ? "bg-background text-foreground shadow-sm"
: "text-muted-foreground hover:text-foreground hover:bg-background/50" : "text-muted-foreground hover:text-foreground hover:bg-background/50"
@@ -72,7 +72,7 @@ export function AppSwitcher({ activeApp, onSwitch }: AppSwitcherProps) {
<button <button
type="button" type="button"
onClick={() => handleSwitch("gemini")} onClick={() => handleSwitch("gemini")}
className={`group inline-flex items-center gap-2 px-3 h-8 rounded-md text-sm font-medium transition-all duration-200 ${ className={`group inline-flex items-center gap-2 px-3 py-2 rounded-md text-sm font-medium transition-all duration-200 ${
activeApp === "gemini" activeApp === "gemini"
? "bg-background text-foreground shadow-sm" ? "bg-background text-foreground shadow-sm"
: "text-muted-foreground hover:text-foreground hover:bg-background/50" : "text-muted-foreground hover:text-foreground hover:bg-background/50"
+1 -1
View File
@@ -56,7 +56,7 @@ export function UpdateBadge({ className = "", onClick }: UpdateBadgeProps) {
" "
aria-label={t("common.close")} aria-label={t("common.close")}
> >
<X className="w-3 h-3 text-muted-foreground" /> <X className="w-3 h-3 text-gray-400 dark:text-gray-500" />
</button> </button>
</div> </div>
); );
+3 -3
View File
@@ -114,7 +114,7 @@ const UsageFooter: React.FC<UsageFooterProps> = ({
{/* 第一行:更新时间和刷新按钮 */} {/* 第一行:更新时间和刷新按钮 */}
<div className="flex items-center gap-2 justify-end"> <div className="flex items-center gap-2 justify-end">
{/* 上次查询时间 */} {/* 上次查询时间 */}
<span className="text-[10px] text-muted-foreground/70 flex items-center gap-1"> <span className="text-[10px] text-gray-400 dark:text-gray-500 flex items-center gap-1">
<Clock size={10} /> <Clock size={10} />
{lastQueriedAt {lastQueriedAt
? formatRelativeTime(lastQueriedAt, now, t) ? formatRelativeTime(lastQueriedAt, now, t)
@@ -128,7 +128,7 @@ const UsageFooter: React.FC<UsageFooterProps> = ({
refetch(); refetch();
}} }}
disabled={loading} disabled={loading}
className="p-1 rounded hover:bg-muted transition-colors disabled:opacity-50 flex-shrink-0 text-muted-foreground" className="p-1 rounded hover:bg-gray-100 dark:hover:bg-gray-800 transition-colors disabled:opacity-50 flex-shrink-0 text-gray-400 dark:text-gray-500"
title={t("usage.refreshUsage")} title={t("usage.refreshUsage")}
> >
<RefreshCw size={12} className={loading ? "animate-spin" : ""} /> <RefreshCw size={12} className={loading ? "animate-spin" : ""} />
@@ -191,7 +191,7 @@ const UsageFooter: React.FC<UsageFooterProps> = ({
<div className="flex items-center gap-2"> <div className="flex items-center gap-2">
{/* 自动查询时间提示 */} {/* 自动查询时间提示 */}
{lastQueriedAt && ( {lastQueriedAt && (
<span className="text-[10px] text-muted-foreground/70 flex items-center gap-1"> <span className="text-[10px] text-gray-400 dark:text-gray-500 flex items-center gap-1">
<Clock size={10} /> <Clock size={10} />
{formatRelativeTime(lastQueriedAt, now, t)} {formatRelativeTime(lastQueriedAt, now, t)}
</span> </span>
+35 -47
View File
@@ -1,6 +1,5 @@
import React from "react"; import React from "react";
import { createPortal } from "react-dom"; import { createPortal } from "react-dom";
import { motion, AnimatePresence } from "framer-motion";
import { ArrowLeft } from "lucide-react"; import { ArrowLeft } from "lucide-react";
import { Button } from "@/components/ui/button"; import { Button } from "@/components/ui/button";
@@ -33,57 +32,46 @@ export const FullScreenPanel: React.FC<FullScreenPanelProps> = ({
}; };
}, [isOpen]); }, [isOpen]);
if (!isOpen) return null;
return createPortal( return createPortal(
<AnimatePresence> <div
{isOpen && ( className="fixed inset-0 z-[60] flex flex-col"
<motion.div style={{ backgroundColor: "hsl(var(--background))" }}
initial={{ opacity: 0 }} >
animate={{ opacity: 1 }} {/* Header */}
exit={{ opacity: 0 }} <div
transition={{ duration: 0.2 }} className="flex-shrink-0 py-3 border-b border-border-default"
className="fixed inset-0 z-[60] flex flex-col" style={{ backgroundColor: "hsl(var(--background))" }}
>
<div className="h-4 w-full" data-tauri-drag-region />
<div className="mx-auto max-w-[56rem] px-6 flex items-center gap-4">
<Button type="button" variant="outline" size="icon" onClick={onClose}>
<ArrowLeft className="h-4 w-4" />
</Button>
<h2 className="text-lg font-semibold text-foreground">{title}</h2>
</div>
</div>
{/* Content */}
<div className="flex-1 overflow-y-auto scroll-overlay">
<div className="mx-auto max-w-[56rem] px-6 py-6 space-y-6 w-full">
{children}
</div>
</div>
{/* Footer */}
{footer && (
<div
className="flex-shrink-0 py-4 border-t border-border-default"
style={{ backgroundColor: "hsl(var(--background))" }} style={{ backgroundColor: "hsl(var(--background))" }}
> >
{/* Header */} <div className="mx-auto max-w-[56rem] px-6 flex items-center justify-end gap-3">
<div {footer}
className="flex-shrink-0 py-3 border-b border-border-default"
style={{ backgroundColor: "hsl(var(--background))" }}
>
<div className="h-4 w-full" data-tauri-drag-region />
<div className="mx-auto max-w-[56rem] px-6 flex items-center gap-4">
<Button
type="button"
variant="outline"
size="icon"
onClick={onClose}
>
<ArrowLeft className="h-4 w-4" />
</Button>
<h2 className="text-lg font-semibold text-foreground">{title}</h2>
</div>
</div> </div>
</div>
{/* Content */}
<div className="flex-1 overflow-y-auto scroll-overlay">
<div className="mx-auto max-w-[56rem] px-6 py-6 space-y-6 w-full">
{children}
</div>
</div>
{/* Footer */}
{footer && (
<div
className="flex-shrink-0 py-4 border-t border-border-default"
style={{ backgroundColor: "hsl(var(--background))" }}
>
<div className="mx-auto max-w-[56rem] px-6 flex items-center justify-end gap-3">
{footer}
</div>
</div>
)}
</motion.div>
)} )}
</AnimatePresence>, </div>,
document.body, document.body,
); );
}; };
+3 -3
View File
@@ -191,14 +191,14 @@ export function EnvWarningBanner({
<div className="flex-1 min-w-0"> <div className="flex-1 min-w-0">
<label <label
htmlFor={key} htmlFor={key}
className="block text-sm font-medium text-foreground cursor-pointer" className="block text-sm font-medium text-gray-900 dark:text-gray-100 cursor-pointer"
> >
{conflict.varName} {conflict.varName}
</label> </label>
<p className="text-xs text-muted-foreground mt-1 break-all"> <p className="text-xs text-gray-600 dark:text-gray-400 mt-1 break-all">
{t("env.field.value")}: {conflict.varValue} {t("env.field.value")}: {conflict.varValue}
</p> </p>
<p className="text-xs text-muted-foreground mt-1"> <p className="text-xs text-gray-500 dark:text-gray-500 mt-1">
{t("env.field.source")}:{" "} {t("env.field.source")}:{" "}
{getSourceDescription(conflict)} {getSourceDescription(conflict)}
</p> </p>
+15 -15
View File
@@ -239,7 +239,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
<div className="flex-1 overflow-y-auto px-6 py-4 space-y-4"> <div className="flex-1 overflow-y-auto px-6 py-4 space-y-4">
{/* Hint */} {/* Hint */}
<div className="rounded-lg border border-border-default bg-gray-100/50 dark:bg-gray-800/50 p-3"> <div className="rounded-lg border border-border-default bg-gray-100/50 dark:bg-gray-800/50 p-3">
<p className="text-sm text-muted-foreground"> <p className="text-sm text-gray-500 dark:text-gray-400">
{t("mcp.wizard.hint")} {t("mcp.wizard.hint")}
</p> </p>
</div> </div>
@@ -248,7 +248,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
<div className="space-y-4 min-h-[400px]"> <div className="space-y-4 min-h-[400px]">
{/* Type */} {/* Type */}
<div> <div>
<label className="mb-2 block text-sm font-medium text-foreground"> <label className="mb-2 block text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.wizard.type")} <span className="text-red-500">*</span> {t("mcp.wizard.type")} <span className="text-red-500">*</span>
</label> </label>
<div className="flex gap-4"> <div className="flex gap-4">
@@ -262,7 +262,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
} }
className="w-4 h-4 accent-blue-500" className="w-4 h-4 accent-blue-500"
/> />
<span className="text-sm text-foreground"> <span className="text-sm text-gray-900 dark:text-gray-100">
{t("mcp.wizard.typeStdio")} {t("mcp.wizard.typeStdio")}
</span> </span>
</label> </label>
@@ -276,7 +276,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
} }
className="w-4 h-4 accent-blue-500" className="w-4 h-4 accent-blue-500"
/> />
<span className="text-sm text-foreground"> <span className="text-sm text-gray-900 dark:text-gray-100">
{t("mcp.wizard.typeHttp")} {t("mcp.wizard.typeHttp")}
</span> </span>
</label> </label>
@@ -290,7 +290,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
} }
className="w-4 h-4 accent-blue-500" className="w-4 h-4 accent-blue-500"
/> />
<span className="text-sm text-foreground"> <span className="text-sm text-gray-900 dark:text-gray-100">
{t("mcp.wizard.typeSse")} {t("mcp.wizard.typeSse")}
</span> </span>
</label> </label>
@@ -299,7 +299,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
{/* Title */} {/* Title */}
<div> <div>
<label className="mb-1 block text-sm font-medium text-foreground"> <label className="mb-1 block text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.form.title")} <span className="text-red-500">*</span> {t("mcp.form.title")} <span className="text-red-500">*</span>
</label> </label>
<Input <Input
@@ -317,7 +317,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
<> <>
{/* Command */} {/* Command */}
<div> <div>
<label className="mb-1 block text-sm font-medium text-foreground"> <label className="mb-1 block text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.wizard.command")}{" "} {t("mcp.wizard.command")}{" "}
<span className="text-red-500">*</span> <span className="text-red-500">*</span>
</label> </label>
@@ -333,7 +333,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
{/* Args */} {/* Args */}
<div> <div>
<label className="mb-1 block text-sm font-medium text-foreground"> <label className="mb-1 block text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.wizard.args")} {t("mcp.wizard.args")}
</label> </label>
<textarea <textarea
@@ -341,13 +341,13 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
onChange={(e) => setWizardArgs(e.target.value)} onChange={(e) => setWizardArgs(e.target.value)}
placeholder={t("mcp.wizard.argsPlaceholder")} placeholder={t("mcp.wizard.argsPlaceholder")}
rows={3} rows={3}
className="w-full rounded-md border border-border-default bg-background px-3 py-2 text-sm font-mono text-foreground placeholder:text-muted-foreground focus:outline-none focus:ring-2 focus:ring-blue-500/20 resize-y" className="w-full rounded-md border border-border-default bg-white dark:bg-gray-800 px-3 py-2 text-sm font-mono text-gray-900 dark:text-gray-100 placeholder:text-gray-400 dark:placeholder:text-gray-500 focus:outline-none focus:ring-2 focus:ring-blue-500/20 resize-y"
/> />
</div> </div>
{/* Env */} {/* Env */}
<div> <div>
<label className="mb-1 block text-sm font-medium text-foreground"> <label className="mb-1 block text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.wizard.env")} {t("mcp.wizard.env")}
</label> </label>
<textarea <textarea
@@ -355,7 +355,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
onChange={(e) => setWizardEnv(e.target.value)} onChange={(e) => setWizardEnv(e.target.value)}
placeholder={t("mcp.wizard.envPlaceholder")} placeholder={t("mcp.wizard.envPlaceholder")}
rows={3} rows={3}
className="w-full rounded-md border border-border-default bg-background px-3 py-2 text-sm font-mono text-foreground placeholder:text-muted-foreground focus:outline-none focus:ring-2 focus:ring-blue-500/20 resize-y" className="w-full rounded-md border border-border-default bg-white dark:bg-gray-800 px-3 py-2 text-sm font-mono text-gray-900 dark:text-gray-100 placeholder:text-gray-400 dark:placeholder:text-gray-500 focus:outline-none focus:ring-2 focus:ring-blue-500/20 resize-y"
/> />
</div> </div>
</> </>
@@ -366,7 +366,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
<> <>
{/* URL */} {/* URL */}
<div> <div>
<label className="mb-1 block text-sm font-medium text-foreground"> <label className="mb-1 block text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.wizard.url")}{" "} {t("mcp.wizard.url")}{" "}
<span className="text-red-500">*</span> <span className="text-red-500">*</span>
</label> </label>
@@ -382,7 +382,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
{/* Headers */} {/* Headers */}
<div> <div>
<label className="mb-1 block text-sm font-medium text-foreground"> <label className="mb-1 block text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.wizard.headers")} {t("mcp.wizard.headers")}
</label> </label>
<textarea <textarea
@@ -390,7 +390,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
onChange={(e) => setWizardHeaders(e.target.value)} onChange={(e) => setWizardHeaders(e.target.value)}
placeholder={t("mcp.wizard.headersPlaceholder")} placeholder={t("mcp.wizard.headersPlaceholder")}
rows={3} rows={3}
className="w-full rounded-md border border-border-default bg-background px-3 py-2 text-sm font-mono text-foreground placeholder:text-muted-foreground focus:outline-none focus:ring-2 focus:ring-blue-500/20 resize-y" className="w-full rounded-md border border-border-default bg-white dark:bg-gray-800 px-3 py-2 text-sm font-mono text-gray-900 dark:text-gray-100 placeholder:text-gray-400 dark:placeholder:text-gray-500 focus:outline-none focus:ring-2 focus:ring-blue-500/20 resize-y"
/> />
</div> </div>
</> </>
@@ -404,7 +404,7 @@ const McpWizardModal: React.FC<McpWizardModalProps> = ({
wizardUrl || wizardUrl ||
wizardHeaders) && ( wizardHeaders) && (
<div className="space-y-2 border-t border-border-default pt-4"> <div className="space-y-2 border-t border-border-default pt-4">
<h3 className="text-sm font-medium text-foreground"> <h3 className="text-sm font-medium text-gray-900 dark:text-gray-100">
{t("mcp.wizard.preview")} {t("mcp.wizard.preview")}
</h3> </h3>
<pre className="overflow-x-auto rounded-lg bg-gray-100 dark:bg-gray-800 p-3 text-xs font-mono text-gray-700 dark:text-gray-300"> <pre className="overflow-x-auto rounded-lg bg-gray-100 dark:bg-gray-800 p-3 text-xs font-mono text-gray-700 dark:text-gray-300">
+13 -11
View File
@@ -129,18 +129,18 @@ const UnifiedMcpPanel = React.forwardRef<
{/* Content - Scrollable */} {/* Content - Scrollable */}
<div className="flex-1 overflow-y-auto overflow-x-hidden pb-24"> <div className="flex-1 overflow-y-auto overflow-x-hidden pb-24">
{isLoading ? ( {isLoading ? (
<div className="text-center py-12 text-muted-foreground"> <div className="text-center py-12 text-gray-500 dark:text-gray-400">
{t("mcp.loading")} {t("mcp.loading")}
</div> </div>
) : serverEntries.length === 0 ? ( ) : serverEntries.length === 0 ? (
<div className="text-center py-12"> <div className="text-center py-12">
<div className="w-16 h-16 mx-auto mb-4 bg-muted rounded-full flex items-center justify-center"> <div className="w-16 h-16 mx-auto mb-4 bg-gray-100 dark:bg-gray-800 rounded-full flex items-center justify-center">
<Server size={24} className="text-muted-foreground" /> <Server size={24} className="text-gray-400 dark:text-gray-500" />
</div> </div>
<h3 className="text-lg font-medium text-foreground mb-2"> <h3 className="text-lg font-medium text-gray-900 dark:text-gray-100 mb-2">
{t("mcp.unifiedPanel.noServers")} {t("mcp.unifiedPanel.noServers")}
</h3> </h3>
<p className="text-muted-foreground text-sm"> <p className="text-gray-500 dark:text-gray-400 text-sm">
{t("mcp.emptyDescription")} {t("mcp.emptyDescription")}
</p> </p>
</div> </div>
@@ -237,7 +237,9 @@ const UnifiedMcpListItem: React.FC<UnifiedMcpListItemProps> = ({
{/* 左侧:服务器信息 */} {/* 左侧:服务器信息 */}
<div className="flex-1 min-w-0"> <div className="flex-1 min-w-0">
<div className="flex items-center gap-2 mb-1"> <div className="flex items-center gap-2 mb-1">
<h3 className="font-medium text-foreground">{name}</h3> <h3 className="font-medium text-gray-900 dark:text-gray-100">
{name}
</h3>
{docsUrl && ( {docsUrl && (
<Button <Button
type="button" type="button"
@@ -251,12 +253,12 @@ const UnifiedMcpListItem: React.FC<UnifiedMcpListItemProps> = ({
)} )}
</div> </div>
{description && ( {description && (
<p className="text-sm text-muted-foreground line-clamp-2"> <p className="text-sm text-gray-500 dark:text-gray-400 line-clamp-2">
{description} {description}
</p> </p>
)} )}
{!description && tags && tags.length > 0 && ( {!description && tags && tags.length > 0 && (
<p className="text-xs text-muted-foreground/70 truncate"> <p className="text-xs text-gray-400 dark:text-gray-500 truncate">
{tags.join(", ")} {tags.join(", ")}
</p> </p>
)} )}
@@ -267,7 +269,7 @@ const UnifiedMcpListItem: React.FC<UnifiedMcpListItemProps> = ({
<div className="flex items-center justify-between gap-3"> <div className="flex items-center justify-between gap-3">
<label <label
htmlFor={`${id}-claude`} htmlFor={`${id}-claude`}
className="text-sm text-foreground/80 cursor-pointer" className="text-sm text-gray-700 dark:text-gray-300 cursor-pointer"
> >
{t("mcp.unifiedPanel.apps.claude")} {t("mcp.unifiedPanel.apps.claude")}
</label> </label>
@@ -283,7 +285,7 @@ const UnifiedMcpListItem: React.FC<UnifiedMcpListItemProps> = ({
<div className="flex items-center justify-between gap-3"> <div className="flex items-center justify-between gap-3">
<label <label
htmlFor={`${id}-codex`} htmlFor={`${id}-codex`}
className="text-sm text-foreground/80 cursor-pointer" className="text-sm text-gray-700 dark:text-gray-300 cursor-pointer"
> >
{t("mcp.unifiedPanel.apps.codex")} {t("mcp.unifiedPanel.apps.codex")}
</label> </label>
@@ -299,7 +301,7 @@ const UnifiedMcpListItem: React.FC<UnifiedMcpListItemProps> = ({
<div className="flex items-center justify-between gap-3"> <div className="flex items-center justify-between gap-3">
<label <label
htmlFor={`${id}-gemini`} htmlFor={`${id}-gemini`}
className="text-sm text-foreground/80 cursor-pointer" className="text-sm text-gray-700 dark:text-gray-300 cursor-pointer"
> >
{t("mcp.unifiedPanel.apps.gemini")} {t("mcp.unifiedPanel.apps.gemini")}
</label> </label>
+4 -2
View File
@@ -36,9 +36,11 @@ const PromptListItem: React.FC<PromptListItemProps> = ({
</div> </div>
<div className="flex-1 min-w-0"> <div className="flex-1 min-w-0">
<h3 className="font-medium text-foreground mb-1">{prompt.name}</h3> <h3 className="font-medium text-gray-900 dark:text-gray-100 mb-1">
{prompt.name}
</h3>
{prompt.description && ( {prompt.description && (
<p className="text-sm text-muted-foreground truncate"> <p className="text-sm text-gray-500 dark:text-gray-400 truncate">
{prompt.description} {prompt.description}
</p> </p>
)} )}
+8 -5
View File
@@ -108,18 +108,21 @@ const PromptPanel = React.forwardRef<PromptPanelHandle, PromptPanelProps>(
<div className="flex-1 overflow-y-auto pb-16"> <div className="flex-1 overflow-y-auto pb-16">
{loading ? ( {loading ? (
<div className="text-center py-12 text-muted-foreground"> <div className="text-center py-12 text-gray-500 dark:text-gray-400">
{t("prompts.loading")} {t("prompts.loading")}
</div> </div>
) : promptEntries.length === 0 ? ( ) : promptEntries.length === 0 ? (
<div className="text-center py-12"> <div className="text-center py-12">
<div className="w-16 h-16 mx-auto mb-4 bg-muted rounded-full flex items-center justify-center"> <div className="w-16 h-16 mx-auto mb-4 bg-gray-100 dark:bg-gray-800 rounded-full flex items-center justify-center">
<FileText size={24} className="text-muted-foreground" /> <FileText
size={24}
className="text-gray-400 dark:text-gray-500"
/>
</div> </div>
<h3 className="text-lg font-medium text-foreground mb-2"> <h3 className="text-lg font-medium text-gray-900 dark:text-gray-100 mb-2">
{t("prompts.empty")} {t("prompts.empty")}
</h3> </h3>
<p className="text-muted-foreground text-sm"> <p className="text-gray-500 dark:text-gray-400 text-sm">
{t("prompts.emptyDescription")} {t("prompts.emptyDescription")}
</p> </p>
</div> </div>
+35 -121
View File
@@ -1,23 +1,17 @@
import { useCallback, useState } from "react"; import { useCallback } from "react";
import { useTranslation } from "react-i18next"; import { useTranslation } from "react-i18next";
import { Plus } from "lucide-react"; import { Plus } from "lucide-react";
import { toast } from "sonner";
import { Button } from "@/components/ui/button"; import { Button } from "@/components/ui/button";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { FullScreenPanel } from "@/components/common/FullScreenPanel"; import { FullScreenPanel } from "@/components/common/FullScreenPanel";
import type { Provider, CustomEndpoint, UniversalProvider } from "@/types"; import type { Provider, CustomEndpoint } from "@/types";
import type { AppId } from "@/lib/api"; import type { AppId } from "@/lib/api";
import { universalProvidersApi } from "@/lib/api";
import { import {
ProviderForm, ProviderForm,
type ProviderFormValues, type ProviderFormValues,
} from "@/components/providers/forms/ProviderForm"; } from "@/components/providers/forms/ProviderForm";
import { UniversalProviderFormModal } from "@/components/universal/UniversalProviderFormModal";
import { UniversalProviderPanel } from "@/components/universal";
import { providerPresets } from "@/config/claudeProviderPresets"; import { providerPresets } from "@/config/claudeProviderPresets";
import { codexProviderPresets } from "@/config/codexProviderPresets"; import { codexProviderPresets } from "@/config/codexProviderPresets";
import { geminiProviderPresets } from "@/config/geminiProviderPresets"; import { geminiProviderPresets } from "@/config/geminiProviderPresets";
import type { UniversalProviderPreset } from "@/config/universalProviderPresets";
interface AddProviderDialogProps { interface AddProviderDialogProps {
open: boolean; open: boolean;
@@ -33,46 +27,6 @@ export function AddProviderDialog({
onSubmit, onSubmit,
}: AddProviderDialogProps) { }: AddProviderDialogProps) {
const { t } = useTranslation(); const { t } = useTranslation();
const [activeTab, setActiveTab] = useState<"app-specific" | "universal">(
"app-specific",
);
const [universalFormOpen, setUniversalFormOpen] = useState(false);
const [selectedUniversalPreset, setSelectedUniversalPreset] =
useState<UniversalProviderPreset | null>(null);
// Handle universal provider save
const handleUniversalProviderSave = useCallback(
async (provider: UniversalProvider) => {
try {
await universalProvidersApi.upsert(provider);
toast.success(
t("universalProvider.addSuccess", {
defaultValue: "统一供应商添加成功",
}),
);
setUniversalFormOpen(false);
setSelectedUniversalPreset(null);
onOpenChange(false);
} catch (error) {
console.error(
"[AddProviderDialog] Failed to save universal provider",
error,
);
toast.error(
t("universalProvider.addFailed", {
defaultValue: "统一供应商添加失败",
}),
);
}
},
[t, onOpenChange],
);
// Close universal form and return to main dialog
const handleUniversalFormClose = useCallback(() => {
setUniversalFormOpen(false);
setSelectedUniversalPreset(null);
}, []);
const handleSubmit = useCallback( const handleSubmit = useCallback(
async (values: ProviderFormValues) => { async (values: ProviderFormValues) => {
@@ -202,86 +156,46 @@ export function AddProviderDialog({
[appId, onSubmit, onOpenChange], [appId, onSubmit, onOpenChange],
); );
// 动态 footer:根据当前 Tab 显示不同按钮 const submitLabel =
const footer = appId === "claude"
activeTab === "app-specific" ? ( ? t("provider.addClaudeProvider")
<> : appId === "codex"
<Button ? t("provider.addCodexProvider")
variant="outline" : t("provider.addGeminiProvider");
onClick={() => onOpenChange(false)}
className="border-border/20 hover:bg-accent hover:text-accent-foreground" const footer = (
> <>
{t("common.cancel")} <Button
</Button> variant="outline"
<Button onClick={() => onOpenChange(false)}
type="submit" className="border-border/20 hover:bg-accent hover:text-accent-foreground"
form="provider-form" >
className="bg-primary text-primary-foreground hover:bg-primary/90" {t("common.cancel")}
> </Button>
<Plus className="h-4 w-4 mr-2" /> <Button
{t("common.add")} type="submit"
</Button> form="provider-form"
</> className="bg-primary text-primary-foreground hover:bg-primary/90"
) : ( >
<> <Plus className="h-4 w-4 mr-2" />
<Button {t("common.add")}
variant="outline" </Button>
onClick={() => onOpenChange(false)} </>
className="border-border/20 hover:bg-accent hover:text-accent-foreground" );
>
{t("common.cancel")}
</Button>
<Button
onClick={() => setUniversalFormOpen(true)}
className="bg-primary text-primary-foreground hover:bg-primary/90"
>
<Plus className="h-4 w-4 mr-2" />
{t("universalProvider.add")}
</Button>
</>
);
return ( return (
<FullScreenPanel <FullScreenPanel
isOpen={open} isOpen={open}
title={t("provider.addNewProvider")} title={submitLabel}
onClose={() => onOpenChange(false)} onClose={() => onOpenChange(false)}
footer={footer} footer={footer}
> >
<Tabs <ProviderForm
value={activeTab} appId={appId}
onValueChange={(v) => setActiveTab(v as "app-specific" | "universal")} submitLabel={t("common.add")}
> onSubmit={handleSubmit}
<TabsList className="grid w-full grid-cols-2 mb-6"> onCancel={() => onOpenChange(false)}
<TabsTrigger value="app-specific"> showButtons={false}
{t(`apps.${appId}`)} {t("provider.tabProvider")}
</TabsTrigger>
<TabsTrigger value="universal">
{t("provider.tabUniversal")}
</TabsTrigger>
</TabsList>
<TabsContent value="app-specific" className="mt-0">
<ProviderForm
appId={appId}
submitLabel={t("common.add")}
onSubmit={handleSubmit}
onCancel={() => onOpenChange(false)}
showButtons={false}
/>
</TabsContent>
<TabsContent value="universal" className="mt-0">
<UniversalProviderPanel />
</TabsContent>
</Tabs>
{/* Universal Provider Form Modal */}
<UniversalProviderFormModal
isOpen={universalFormOpen}
onClose={handleUniversalFormClose}
onSave={handleUniversalProviderSave}
initialPreset={selectedUniversalPreset}
/> />
</FullScreenPanel> </FullScreenPanel>
); );
@@ -16,7 +16,6 @@ interface EditProviderDialogProps {
onOpenChange: (open: boolean) => void; onOpenChange: (open: boolean) => void;
onSubmit: (provider: Provider) => Promise<void> | void; onSubmit: (provider: Provider) => Promise<void> | void;
appId: AppId; appId: AppId;
isProxyTakeover?: boolean; // 代理接管模式下不读取 live(避免显示被接管后的代理配置)
} }
export function EditProviderDialog({ export function EditProviderDialog({
@@ -25,7 +24,6 @@ export function EditProviderDialog({
onOpenChange, onOpenChange,
onSubmit, onSubmit,
appId, appId,
isProxyTakeover = false,
}: EditProviderDialogProps) { }: EditProviderDialogProps) {
const { t } = useTranslation(); const { t } = useTranslation();
@@ -52,16 +50,6 @@ export function EditProviderDialog({
return; return;
} }
// 代理接管模式:Live 配置已被代理改写,读取 live 会导致编辑界面展示代理地址/占位符等内容
// 因此直接回退到 SSOT(数据库)配置,避免用户困惑与误保存
if (isProxyTakeover) {
if (!cancelled) {
setLiveSettings(null);
setHasLoadedLive(true);
}
return;
}
try { try {
const currentId = await providersApi.getCurrent(appId); const currentId = await providersApi.getCurrent(appId);
if (currentId && provider.id === currentId) { if (currentId && provider.id === currentId) {
@@ -94,7 +82,7 @@ export function EditProviderDialog({
return () => { return () => {
cancelled = true; cancelled = true;
}; };
}, [open, provider?.id, appId, hasLoadedLive, isProxyTakeover]); // 只依赖 provider.id,不依赖整个 provider 对象 }, [open, provider?.id, appId, hasLoadedLive]); // 只依赖 provider.id,不依赖整个 provider 对象
const initialSettingsConfig = useMemo(() => { const initialSettingsConfig = useMemo(() => {
return (liveSettings ?? provider?.settingsConfig ?? {}) as Record< return (liveSettings ?? provider?.settingsConfig ?? {}) as Record<
@@ -1,34 +0,0 @@
import { cn } from "@/lib/utils";
import { useTranslation } from "react-i18next";
interface FailoverPriorityBadgeProps {
priority: number; // 1, 2, 3, ...
className?: string;
}
/**
*
*
*/
export function FailoverPriorityBadge({
priority,
className,
}: FailoverPriorityBadgeProps) {
const { t } = useTranslation();
return (
<div
className={cn(
"inline-flex items-center px-1.5 py-0.5 rounded text-xs font-semibold",
"bg-emerald-500/10 text-emerald-600 dark:text-emerald-400",
className,
)}
title={t("failover.priority.tooltip", {
priority,
defaultValue: `故障转移优先级 ${priority}`,
})}
>
P{priority}
</div>
);
}
@@ -1,7 +1,6 @@
import React from "react"; import React from "react";
import { cn } from "@/lib/utils"; import { cn } from "@/lib/utils";
import type { HealthStatus } from "@/lib/api/model-test"; import type { HealthStatus } from "@/lib/api/model-test";
import { useTranslation } from "react-i18next";
interface HealthStatusIndicatorProps { interface HealthStatusIndicatorProps {
status: HealthStatus; status: HealthStatus;
@@ -12,20 +11,17 @@ interface HealthStatusIndicatorProps {
const statusConfig = { const statusConfig = {
operational: { operational: {
color: "bg-emerald-500", color: "bg-emerald-500",
labelKey: "health.operational", label: "正常",
labelFallback: "正常",
textColor: "text-emerald-600 dark:text-emerald-400", textColor: "text-emerald-600 dark:text-emerald-400",
}, },
degraded: { degraded: {
color: "bg-yellow-500", color: "bg-yellow-500",
labelKey: "health.degraded", label: "降级",
labelFallback: "降级",
textColor: "text-yellow-600 dark:text-yellow-400", textColor: "text-yellow-600 dark:text-yellow-400",
}, },
failed: { failed: {
color: "bg-red-500", color: "bg-red-500",
labelKey: "health.failed", label: "失败",
labelFallback: "失败",
textColor: "text-red-600 dark:text-red-400", textColor: "text-red-600 dark:text-red-400",
}, },
}; };
@@ -35,15 +31,13 @@ export const HealthStatusIndicator: React.FC<HealthStatusIndicatorProps> = ({
responseTimeMs, responseTimeMs,
className, className,
}) => { }) => {
const { t } = useTranslation();
const config = statusConfig[status]; const config = statusConfig[status];
const label = t(config.labelKey, { defaultValue: config.labelFallback });
return ( return (
<div className={cn("flex items-center gap-2", className)}> <div className={cn("flex items-center gap-2", className)}>
<div className={cn("w-2 h-2 rounded-full", config.color)} /> <div className={cn("w-2 h-2 rounded-full", config.color)} />
<span className={cn("text-xs font-medium", config.textColor)}> <span className={cn("text-xs font-medium", config.textColor)}>
{label} {config.label}
{responseTimeMs !== undefined && ` (${responseTimeMs}ms)`} {responseTimeMs !== undefined && ` (${responseTimeMs}ms)`}
</span> </span>
</div> </div>
+23 -78
View File
@@ -5,7 +5,6 @@ import {
Edit, Edit,
Loader2, Loader2,
Play, Play,
Plus,
TestTube2, TestTube2,
Trash2, Trash2,
} from "lucide-react"; } from "lucide-react";
@@ -23,10 +22,6 @@ interface ProviderActionsProps {
onTest?: () => void; onTest?: () => void;
onConfigureUsage: () => void; onConfigureUsage: () => void;
onDelete: () => void; onDelete: () => void;
// 故障转移相关
isAutoFailoverEnabled?: boolean;
isInFailoverQueue?: boolean;
onToggleFailover?: (enabled: boolean) => void;
} }
export function ProviderActions({ export function ProviderActions({
@@ -39,88 +34,38 @@ export function ProviderActions({
onTest, onTest,
onConfigureUsage, onConfigureUsage,
onDelete, onDelete,
// 故障转移相关
isAutoFailoverEnabled = false,
isInFailoverQueue = false,
onToggleFailover,
}: ProviderActionsProps) { }: ProviderActionsProps) {
const { t } = useTranslation(); const { t } = useTranslation();
const iconButtonClass = "h-8 w-8 p-1"; const iconButtonClass = "h-8 w-8 p-1";
// 故障转移模式下的按钮逻辑
const isFailoverMode = isAutoFailoverEnabled && onToggleFailover;
// 处理主按钮点击
const handleMainButtonClick = () => {
if (isFailoverMode) {
// 故障转移模式:切换队列状态
onToggleFailover(!isInFailoverQueue);
} else {
// 普通模式:切换供应商
onSwitch();
}
};
// 主按钮的状态和样式
const getMainButtonState = () => {
if (isFailoverMode) {
// 故障转移模式
if (isInFailoverQueue) {
return {
disabled: false,
variant: "secondary" as const,
className:
"bg-blue-100 text-blue-600 hover:bg-blue-200 dark:bg-blue-900/50 dark:text-blue-400 dark:hover:bg-blue-900/70",
icon: <Check className="h-4 w-4" />,
text: t("failover.inQueue", { defaultValue: "已加入" }),
};
}
return {
disabled: false,
variant: "default" as const,
className:
"bg-blue-500 hover:bg-blue-600 dark:bg-blue-600 dark:hover:bg-blue-700",
icon: <Plus className="h-4 w-4" />,
text: t("failover.addQueue", { defaultValue: "加入" }),
};
}
// 普通模式
if (isCurrent) {
return {
disabled: true,
variant: "secondary" as const,
className:
"bg-gray-200 text-muted-foreground hover:bg-gray-200 hover:text-muted-foreground dark:bg-gray-700 dark:hover:bg-gray-700",
icon: <Check className="h-4 w-4" />,
text: t("provider.inUse"),
};
}
return {
disabled: false,
variant: "default" as const,
className: isProxyTakeover
? "bg-emerald-500 hover:bg-emerald-600 dark:bg-emerald-600 dark:hover:bg-emerald-700"
: "",
icon: <Play className="h-4 w-4" />,
text: t("provider.enable"),
};
};
const buttonState = getMainButtonState();
return ( return (
<div className="flex items-center gap-1.5"> <div className="flex items-center gap-1.5">
<Button <Button
size="sm" size="sm"
variant={buttonState.variant} variant={isCurrent ? "secondary" : "default"}
onClick={handleMainButtonClick} onClick={onSwitch}
disabled={buttonState.disabled} disabled={isCurrent}
className={cn("w-[4.5rem] px-2.5", buttonState.className)} className={cn(
"w-[4.5rem] px-2.5",
isCurrent &&
"bg-gray-200 text-muted-foreground hover:bg-gray-200 hover:text-muted-foreground dark:bg-gray-700 dark:hover:bg-gray-700",
// 代理接管模式下启用按钮使用绿色
!isCurrent &&
isProxyTakeover &&
"bg-emerald-500 hover:bg-emerald-600 dark:bg-emerald-600 dark:hover:bg-emerald-700",
)}
> >
{buttonState.icon} {isCurrent ? (
{buttonState.text} <>
<Check className="h-4 w-4" />
{t("provider.inUse")}
</>
) : (
<>
<Play className="h-4 w-4" />
{t("provider.enable")}
</>
)}
</Button> </Button>
<div className="flex items-center gap-1"> <div className="flex items-center gap-1">
+15 -50
View File
@@ -12,7 +12,6 @@ import { ProviderActions } from "@/components/providers/ProviderActions";
import { ProviderIcon } from "@/components/ProviderIcon"; import { ProviderIcon } from "@/components/ProviderIcon";
import UsageFooter from "@/components/UsageFooter"; import UsageFooter from "@/components/UsageFooter";
import { ProviderHealthBadge } from "@/components/providers/ProviderHealthBadge"; import { ProviderHealthBadge } from "@/components/providers/ProviderHealthBadge";
import { FailoverPriorityBadge } from "@/components/providers/FailoverPriorityBadge";
import { useProviderHealth } from "@/lib/query/failover"; import { useProviderHealth } from "@/lib/query/failover";
import { useUsageQuery } from "@/lib/query/queries"; import { useUsageQuery } from "@/lib/query/queries";
@@ -37,12 +36,6 @@ interface ProviderCardProps {
isProxyRunning: boolean; isProxyRunning: boolean;
isProxyTakeover?: boolean; // 代理接管模式(Live配置已被接管,切换为热切换) isProxyTakeover?: boolean; // 代理接管模式(Live配置已被接管,切换为热切换)
dragHandleProps?: DragHandleProps; dragHandleProps?: DragHandleProps;
// 故障转移相关
isAutoFailoverEnabled?: boolean; // 是否开启自动故障转移
failoverPriority?: number; // 故障转移优先级(1 = P1, 2 = P2, ...
isInFailoverQueue?: boolean; // 是否在故障转移队列中
onToggleFailover?: (enabled: boolean) => void; // 切换故障转移队列
activeProviderId?: string; // 代理当前实际使用的供应商 ID(用于故障转移模式下标注绿色边框)
} }
const extractApiUrl = (provider: Provider, fallbackText: string) => { const extractApiUrl = (provider: Provider, fallbackText: string) => {
@@ -95,12 +88,6 @@ export function ProviderCard({
isProxyRunning, isProxyRunning,
isProxyTakeover = false, isProxyTakeover = false,
dragHandleProps, dragHandleProps,
// 故障转移相关
isAutoFailoverEnabled = false,
failoverPriority,
isInFailoverQueue = false,
onToggleFailover,
activeProviderId,
}: ProviderCardProps) { }: ProviderCardProps) {
const { t } = useTranslation(); const { t } = useTranslation();
@@ -161,32 +148,21 @@ export function ProviderCard({
onOpenWebsite(displayUrl); onOpenWebsite(displayUrl);
}; };
// 判断是否是"当前使用中"的供应商
// - 故障转移模式:代理实际使用的供应商(activeProviderId
// - 代理接管模式(非故障转移):isCurrent
// - 普通模式:isCurrent
const isActiveProvider = isAutoFailoverEnabled
? activeProviderId === provider.id
: isCurrent;
// 判断是否使用绿色(代理接管模式)还是蓝色(普通模式)
const shouldUseGreen = isProxyTakeover && isActiveProvider;
const shouldUseBlue = !isProxyTakeover && isActiveProvider;
return ( return (
<div <div
className={cn( className={cn(
"relative overflow-hidden rounded-xl border border-border p-4 transition-all duration-300", "relative overflow-hidden rounded-xl border border-border p-4 transition-all duration-300",
"bg-card text-card-foreground group", "bg-card text-card-foreground group",
// hover 时的边框效果 // 代理接管模式下 hover 使用绿色边框,否则使用蓝色
isAutoFailoverEnabled || isProxyTakeover isProxyTakeover
? "hover:border-emerald-500/50" ? "hover:border-emerald-500/50"
: "hover:border-border-active", : "hover:border-border-active",
// 当前激活的供应商边框样式 // 代理接管模式下当前供应商使用绿色边框
shouldUseGreen && isProxyTakeover && isCurrent
"border-emerald-500/60 shadow-sm shadow-emerald-500/10", ? "border-emerald-500/60 shadow-sm shadow-emerald-500/10"
shouldUseBlue && "border-blue-500/60 shadow-sm shadow-blue-500/10", : isCurrent
!isActiveProvider && "hover:shadow-sm", ? "border-primary/50 shadow-sm"
: "hover:shadow-sm",
dragHandleProps?.isDragging && dragHandleProps?.isDragging &&
"cursor-grabbing border-primary shadow-lg scale-105 z-10", "cursor-grabbing border-primary shadow-lg scale-105 z-10",
)} )}
@@ -194,11 +170,11 @@ export function ProviderCard({
<div <div
className={cn( className={cn(
"absolute inset-0 bg-gradient-to-r to-transparent transition-opacity duration-500 pointer-events-none", "absolute inset-0 bg-gradient-to-r to-transparent transition-opacity duration-500 pointer-events-none",
// 代理接管模式使用绿色渐变,普通模式使用蓝色渐变 // 代理接管模式使用绿色渐变,否则使用蓝色主色调
shouldUseGreen && "from-emerald-500/10", isProxyTakeover && isCurrent
shouldUseBlue && "from-blue-500/10", ? "from-emerald-500/10"
!isActiveProvider && "from-primary/10", : "from-primary/10",
isActiveProvider ? "opacity-100" : "opacity-0", isCurrent ? "opacity-100" : "opacity-0",
)} )}
/> />
<div className="relative flex flex-col gap-4 sm:flex-row sm:items-center sm:justify-between"> <div className="relative flex flex-col gap-4 sm:flex-row sm:items-center sm:justify-between">
@@ -233,20 +209,13 @@ export function ProviderCard({
{provider.name} {provider.name}
</h3> </h3>
{/* 健康状态徽章 */} {/* 健康状态徽章和优先级 */}
{isProxyRunning && isInFailoverQueue && health && ( {isProxyRunning && health && (
<ProviderHealthBadge <ProviderHealthBadge
consecutiveFailures={health.consecutive_failures} consecutiveFailures={health.consecutive_failures}
/> />
)} )}
{/* 故障转移优先级徽章 */}
{isAutoFailoverEnabled &&
isInFailoverQueue &&
failoverPriority && (
<FailoverPriorityBadge priority={failoverPriority} />
)}
{provider.category === "third_party" && {provider.category === "third_party" &&
provider.meta?.isPartner && ( provider.meta?.isPartner && (
<span <span
@@ -339,10 +308,6 @@ export function ProviderCard({
onTest={onTest ? () => onTest(provider) : undefined} onTest={onTest ? () => onTest(provider) : undefined}
onConfigureUsage={() => onConfigureUsage(provider)} onConfigureUsage={() => onConfigureUsage(provider)}
onDelete={() => onDelete(provider)} onDelete={() => onDelete(provider)}
// 故障转移相关
isAutoFailoverEnabled={isAutoFailoverEnabled}
isInFailoverQueue={isInFailoverQueue}
onToggleFailover={onToggleFailover}
/> />
</div> </div>
</div> </div>
@@ -1,6 +1,5 @@
import { cn } from "@/lib/utils"; import { cn } from "@/lib/utils";
import { ProviderHealthStatus } from "@/types/proxy"; import { ProviderHealthStatus } from "@/types/proxy";
import { useTranslation } from "react-i18next";
interface ProviderHealthBadgeProps { interface ProviderHealthBadgeProps {
consecutiveFailures: number; consecutiveFailures: number;
@@ -15,14 +14,11 @@ export function ProviderHealthBadge({
consecutiveFailures, consecutiveFailures,
className, className,
}: ProviderHealthBadgeProps) { }: ProviderHealthBadgeProps) {
const { t } = useTranslation();
// 根据失败次数计算状态 // 根据失败次数计算状态
const getStatus = () => { const getStatus = () => {
if (consecutiveFailures === 0) { if (consecutiveFailures === 0) {
return { return {
labelKey: "health.operational", label: "正常",
labelFallback: "正常",
status: ProviderHealthStatus.Healthy, status: ProviderHealthStatus.Healthy,
color: "bg-green-500", color: "bg-green-500",
// 使用更深/柔和的背景色,去除可能的白色内容感 // 使用更深/柔和的背景色,去除可能的白色内容感
@@ -31,8 +27,7 @@ export function ProviderHealthBadge({
}; };
} else if (consecutiveFailures < 5) { } else if (consecutiveFailures < 5) {
return { return {
labelKey: "health.degraded", label: "降级",
labelFallback: "降级",
status: ProviderHealthStatus.Degraded, status: ProviderHealthStatus.Degraded,
color: "bg-yellow-500", color: "bg-yellow-500",
bgColor: "bg-yellow-500/10", bgColor: "bg-yellow-500/10",
@@ -40,8 +35,7 @@ export function ProviderHealthBadge({
}; };
} else { } else {
return { return {
labelKey: "health.circuitOpen", label: "熔断",
labelFallback: "熔断",
status: ProviderHealthStatus.Failed, status: ProviderHealthStatus.Failed,
color: "bg-red-500", color: "bg-red-500",
bgColor: "bg-red-500/10", bgColor: "bg-red-500/10",
@@ -51,9 +45,6 @@ export function ProviderHealthBadge({
}; };
const statusConfig = getStatus(); const statusConfig = getStatus();
const label = t(statusConfig.labelKey, {
defaultValue: statusConfig.labelFallback,
});
return ( return (
<div <div
@@ -63,13 +54,10 @@ export function ProviderHealthBadge({
statusConfig.textColor, statusConfig.textColor,
className, className,
)} )}
title={t("health.consecutiveFailures", { title={`连续失败 ${consecutiveFailures}`}
count: consecutiveFailures,
defaultValue: `连续失败 ${consecutiveFailures}`,
})}
> >
<div className={cn("w-2 h-2 rounded-full", statusConfig.color)} /> <div className={cn("w-2 h-2 rounded-full", statusConfig.color)} />
<span>{label}</span> <span>{statusConfig.label}</span>
</div> </div>
); );
} }
+11 -218
View File
@@ -5,31 +5,13 @@ import {
useSortable, useSortable,
verticalListSortingStrategy, verticalListSortingStrategy,
} from "@dnd-kit/sortable"; } from "@dnd-kit/sortable";
import { import type { CSSProperties } from "react";
useEffect,
useMemo,
useRef,
useState,
type CSSProperties,
} from "react";
import { AnimatePresence, motion } from "framer-motion";
import { Search, X } from "lucide-react";
import { useTranslation } from "react-i18next";
import type { Provider } from "@/types"; import type { Provider } from "@/types";
import type { AppId } from "@/lib/api"; import type { AppId } from "@/lib/api";
import { useDragSort } from "@/hooks/useDragSort"; import { useDragSort } from "@/hooks/useDragSort";
import { useStreamCheck } from "@/hooks/useStreamCheck"; import { useStreamCheck } from "@/hooks/useStreamCheck";
import { ProviderCard } from "@/components/providers/ProviderCard"; import { ProviderCard } from "@/components/providers/ProviderCard";
import { ProviderEmptyState } from "@/components/providers/ProviderEmptyState"; import { ProviderEmptyState } from "@/components/providers/ProviderEmptyState";
import {
useAutoFailoverEnabled,
useFailoverQueue,
useAddToFailoverQueue,
useRemoveFromFailoverQueue,
} from "@/lib/query/failover";
import { useCallback } from "react";
import { Input } from "@/components/ui/input";
import { Button } from "@/components/ui/button";
interface ProviderListProps { interface ProviderListProps {
providers: Record<string, Provider>; providers: Record<string, Provider>;
@@ -45,7 +27,6 @@ interface ProviderListProps {
isLoading?: boolean; isLoading?: boolean;
isProxyRunning?: boolean; // 代理服务运行状态 isProxyRunning?: boolean; // 代理服务运行状态
isProxyTakeover?: boolean; // 代理接管模式(Live配置已被接管) isProxyTakeover?: boolean; // 代理接管模式(Live配置已被接管)
activeProviderId?: string; // 代理当前实际使用的供应商 ID(用于故障转移模式下标注绿色边框)
} }
export function ProviderList({ export function ProviderList({
@@ -60,11 +41,9 @@ export function ProviderList({
onOpenWebsite, onOpenWebsite,
onCreate, onCreate,
isLoading = false, isLoading = false,
isProxyRunning = false, isProxyRunning = false, // 默认值为 false
isProxyTakeover = false, isProxyTakeover = false, // 默认值为 false
activeProviderId,
}: ProviderListProps) { }: ProviderListProps) {
const { t } = useTranslation();
const { sortedProviders, sensors, handleDragEnd } = useDragSort( const { sortedProviders, sensors, handleDragEnd } = useDragSort(
providers, providers,
appId, appId,
@@ -73,103 +52,17 @@ export function ProviderList({
// 流式健康检查 // 流式健康检查
const { checkProvider, isChecking } = useStreamCheck(appId); const { checkProvider, isChecking } = useStreamCheck(appId);
// 故障转移相关
const { data: isAutoFailoverEnabled } = useAutoFailoverEnabled(appId);
const { data: failoverQueue } = useFailoverQueue(appId);
const addToQueue = useAddToFailoverQueue();
const removeFromQueue = useRemoveFromFailoverQueue();
// 联动状态:只有当前应用开启代理接管且故障转移开启时才启用故障转移模式
const isFailoverModeActive =
isProxyTakeover === true && isAutoFailoverEnabled === true;
// 计算供应商在故障转移队列中的优先级(基于 sortIndex 排序)
const getFailoverPriority = useCallback(
(providerId: string): number | undefined => {
if (!isFailoverModeActive || !failoverQueue) return undefined;
const index = failoverQueue.findIndex(
(item) => item.providerId === providerId,
);
return index >= 0 ? index + 1 : undefined;
},
[isFailoverModeActive, failoverQueue],
);
// 判断供应商是否在故障转移队列中
const isInFailoverQueue = useCallback(
(providerId: string): boolean => {
if (!isFailoverModeActive || !failoverQueue) return false;
return failoverQueue.some((item) => item.providerId === providerId);
},
[isFailoverModeActive, failoverQueue],
);
// 切换供应商的故障转移队列状态
const handleToggleFailover = useCallback(
(providerId: string, enabled: boolean) => {
if (enabled) {
addToQueue.mutate({ appType: appId, providerId });
} else {
removeFromQueue.mutate({ appType: appId, providerId });
}
},
[appId, addToQueue, removeFromQueue],
);
const handleTest = (provider: Provider) => { const handleTest = (provider: Provider) => {
checkProvider(provider.id, provider.name); checkProvider(provider.id, provider.name);
}; };
const [searchTerm, setSearchTerm] = useState("");
const [isSearchOpen, setIsSearchOpen] = useState(false);
const searchInputRef = useRef<HTMLInputElement>(null);
useEffect(() => {
const handleKeyDown = (event: KeyboardEvent) => {
const key = event.key.toLowerCase();
if ((event.metaKey || event.ctrlKey) && key === "f") {
event.preventDefault();
setIsSearchOpen(true);
return;
}
if (key === "escape") {
setIsSearchOpen(false);
}
};
window.addEventListener("keydown", handleKeyDown);
return () => window.removeEventListener("keydown", handleKeyDown);
}, []);
useEffect(() => {
if (isSearchOpen) {
const frame = requestAnimationFrame(() => {
searchInputRef.current?.focus();
searchInputRef.current?.select();
});
return () => cancelAnimationFrame(frame);
}
}, [isSearchOpen]);
const filteredProviders = useMemo(() => {
const keyword = searchTerm.trim().toLowerCase();
if (!keyword) return sortedProviders;
return sortedProviders.filter((provider) => {
const fields = [provider.name, provider.notes, provider.websiteUrl];
return fields.some((field) =>
field?.toString().toLowerCase().includes(keyword),
);
});
}, [searchTerm, sortedProviders]);
if (isLoading) { if (isLoading) {
return ( return (
<div className="space-y-3"> <div className="space-y-3">
{[0, 1, 2].map((index) => ( {[0, 1, 2].map((index) => (
<div <div
key={index} key={index}
className="w-full border border-dashed rounded-lg h-28 border-muted-foreground/40 bg-muted/40" className="h-28 w-full rounded-lg border border-dashed border-muted-foreground/40 bg-muted/40"
/> />
))} ))}
</div> </div>
@@ -180,18 +73,21 @@ export function ProviderList({
return <ProviderEmptyState onCreate={onCreate} />; return <ProviderEmptyState onCreate={onCreate} />;
} }
const renderProviderList = () => ( return (
<DndContext <DndContext
sensors={sensors} sensors={sensors}
collisionDetection={closestCenter} collisionDetection={closestCenter}
onDragEnd={handleDragEnd} onDragEnd={handleDragEnd}
> >
<SortableContext <SortableContext
items={filteredProviders.map((provider) => provider.id)} items={sortedProviders.map((provider) => provider.id)}
strategy={verticalListSortingStrategy} strategy={verticalListSortingStrategy}
> >
<div className="space-y-3"> <div
{filteredProviders.map((provider) => ( className="space-y-3 animate-slide-up"
style={{ animationDelay: "0.1s" }}
>
{sortedProviders.map((provider) => (
<SortableProviderCard <SortableProviderCard
key={provider.id} key={provider.id}
provider={provider} provider={provider}
@@ -207,98 +103,12 @@ export function ProviderList({
isTesting={isChecking(provider.id)} isTesting={isChecking(provider.id)}
isProxyRunning={isProxyRunning} isProxyRunning={isProxyRunning}
isProxyTakeover={isProxyTakeover} isProxyTakeover={isProxyTakeover}
// 故障转移相关:联动状态
isAutoFailoverEnabled={isFailoverModeActive}
failoverPriority={getFailoverPriority(provider.id)}
isInFailoverQueue={isInFailoverQueue(provider.id)}
onToggleFailover={(enabled) =>
handleToggleFailover(provider.id, enabled)
}
activeProviderId={activeProviderId}
/> />
))} ))}
</div> </div>
</SortableContext> </SortableContext>
</DndContext> </DndContext>
); );
return (
<div className="mt-4 space-y-4">
<AnimatePresence>
{isSearchOpen && (
<motion.div
key="provider-search"
initial={{ opacity: 0, y: -8, scale: 0.98 }}
animate={{ opacity: 1, y: 0, scale: 1 }}
exit={{ opacity: 0, y: -8, scale: 0.98 }}
transition={{ duration: 0.18, ease: "easeOut" }}
className="fixed left-1/2 top-[6.5rem] z-40 w-[min(90vw,26rem)] -translate-x-1/2 sm:right-6 sm:left-auto sm:translate-x-0"
>
<div className="p-4 space-y-3 border shadow-md rounded-2xl border-white/10 bg-background/95 shadow-black/20 backdrop-blur-md">
<div className="relative flex items-center gap-2">
<Search className="absolute w-4 h-4 -translate-y-1/2 pointer-events-none left-3 top-1/2 text-muted-foreground" />
<Input
ref={searchInputRef}
value={searchTerm}
onChange={(event) => setSearchTerm(event.target.value)}
placeholder={t("provider.searchPlaceholder", {
defaultValue: "Search name, notes, or URL...",
})}
aria-label={t("provider.searchAriaLabel", {
defaultValue: "Search providers",
})}
className="pr-16 pl-9"
/>
{searchTerm && (
<Button
variant="ghost"
size="sm"
className="absolute text-xs -translate-y-1/2 right-11 top-1/2"
onClick={() => setSearchTerm("")}
>
{t("common.clear", { defaultValue: "Clear" })}
</Button>
)}
<Button
variant="ghost"
size="icon"
className="ml-auto"
onClick={() => setIsSearchOpen(false)}
aria-label={t("provider.searchCloseAriaLabel", {
defaultValue: "Close provider search",
})}
>
<X className="w-4 h-4" />
</Button>
</div>
<div className="flex flex-wrap items-center justify-between gap-2 text-[11px] text-muted-foreground">
<span>
{t("provider.searchScopeHint", {
defaultValue: "Matches provider name, notes, and URL.",
})}
</span>
<span>
{t("provider.searchCloseHint", {
defaultValue: "Press Esc to close",
})}
</span>
</div>
</div>
</motion.div>
)}
</AnimatePresence>
{filteredProviders.length === 0 ? (
<div className="px-6 py-8 text-sm text-center border border-dashed rounded-lg border-border text-muted-foreground">
{t("provider.noSearchResults", {
defaultValue: "No providers match your search.",
})}
</div>
) : (
renderProviderList()
)}
</div>
);
} }
interface SortableProviderCardProps { interface SortableProviderCardProps {
@@ -315,12 +125,6 @@ interface SortableProviderCardProps {
isTesting: boolean; isTesting: boolean;
isProxyRunning: boolean; isProxyRunning: boolean;
isProxyTakeover: boolean; isProxyTakeover: boolean;
// 故障转移相关
isAutoFailoverEnabled: boolean;
failoverPriority?: number;
isInFailoverQueue: boolean;
onToggleFailover: (enabled: boolean) => void;
activeProviderId?: string;
} }
function SortableProviderCard({ function SortableProviderCard({
@@ -337,11 +141,6 @@ function SortableProviderCard({
isTesting, isTesting,
isProxyRunning, isProxyRunning,
isProxyTakeover, isProxyTakeover,
isAutoFailoverEnabled,
failoverPriority,
isInFailoverQueue,
onToggleFailover,
activeProviderId,
}: SortableProviderCardProps) { }: SortableProviderCardProps) {
const { const {
setNodeRef, setNodeRef,
@@ -380,12 +179,6 @@ function SortableProviderCard({
listeners, listeners,
isDragging, isDragging,
}} }}
// 故障转移相关
isAutoFailoverEnabled={isAutoFailoverEnabled}
failoverPriority={failoverPriority}
isInFailoverQueue={isInFailoverQueue}
onToggleFailover={onToggleFailover}
activeProviderId={activeProviderId}
/> />
</div> </div>
); );
@@ -30,13 +30,16 @@ const ApiKeyInput: React.FC<ApiKeyInputProps> = ({
const inputClass = `w-full px-3 py-2 pr-10 border rounded-lg text-sm transition-colors ${ const inputClass = `w-full px-3 py-2 pr-10 border rounded-lg text-sm transition-colors ${
disabled disabled
? "bg-muted border-border-default text-muted-foreground cursor-not-allowed" ? "bg-gray-100 dark:bg-gray-800 border-border-default text-gray-400 dark:text-gray-500 cursor-not-allowed"
: "border-border-default bg-background text-foreground focus:outline-none focus:ring-2 focus:ring-blue-500/20 dark:focus:ring-blue-400/20" : "border-border-default dark:bg-gray-800 dark:text-gray-100 focus:outline-none focus:ring-2 focus:ring-blue-500/20 dark:focus:ring-blue-400/20"
}`; }`;
return ( return (
<div className="space-y-2"> <div className="space-y-2">
<label htmlFor={id} className="block text-sm font-medium text-foreground"> <label
htmlFor={id}
className="block text-sm font-medium text-gray-900 dark:text-gray-100"
>
{label} {required && "*"} {label} {required && "*"}
</label> </label>
<div className="relative"> <div className="relative">
@@ -55,7 +58,7 @@ const ApiKeyInput: React.FC<ApiKeyInputProps> = ({
<button <button
type="button" type="button"
onClick={toggleShowKey} onClick={toggleShowKey}
className="absolute inset-y-0 right-0 flex items-center pr-3 text-muted-foreground hover:text-foreground transition-colors" className="absolute inset-y-0 right-0 flex items-center pr-3 text-gray-500 dark:text-gray-400 hover:text-gray-900 dark:hover:text-gray-100 transition-colors"
aria-label={showKey ? t("apiKeyInput.hide") : t("apiKeyInput.show")} aria-label={showKey ? t("apiKeyInput.hide") : t("apiKeyInput.show")}
> >
{showKey ? <EyeOff size={16} /> : <Eye size={16} />} {showKey ? <EyeOff size={16} /> : <Eye size={16} />}

Some files were not shown because too many files have changed in this diff Show More