diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index d2cea7e1f..2765dc703 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -103,19 +103,12 @@ jobs: - platform: macos-latest target: aarch64-apple-darwin name: macOS-arm64 - build_cli: true - platform: macos-15-intel target: x86_64-apple-darwin name: macOS-x64 - build_cli: true - platform: windows-2022 target: x86_64-pc-windows-msvc name: Windows-x64 - build_cli: true - - platform: ubuntu-22.04 - target: x86_64-unknown-linux-gnu - name: Linux-x64-cli - build_cli: true runs-on: ${{ matrix.platform }} env: @@ -751,66 +744,6 @@ jobs: printf 'Staged macOS assets under %s:\n' "$STAGING_DIR" find "$STAGING_DIR" -maxdepth 1 -type f | sort - - name: Build lime-cli release binary - if: matrix.build_cli && (matrix.platform == 'windows-2022' || matrix.platform == 'ubuntu-22.04' || (startsWith(matrix.platform, 'macos') && (steps.build_macos_primary.outcome == 'success' || steps.build_macos_retry.outcome == 'success'))) - shell: bash - env: - CARGO_PROFILE_RELEASE_LTO: "off" - CARGO_PROFILE_RELEASE_CODEGEN_UNITS: 32 - CARGO_INCREMENTAL: 0 - CARGO_TARGET_DIR: src-tauri/target - SCCACHE_GHA_ENABLED: "true" - RUSTC_WRAPPER: sccache - run: | - set -euxo pipefail - cargo build --manifest-path src-tauri/Cargo.toml -p lime-cli --release --target "${{ matrix.target }}" - - - name: Package lime-cli release asset - id: package_lime_cli - if: matrix.build_cli && (matrix.platform == 'windows-2022' || matrix.platform == 'ubuntu-22.04' || (startsWith(matrix.platform, 'macos') && (steps.build_macos_primary.outcome == 'success' || steps.build_macos_retry.outcome == 'success'))) - shell: bash - run: | - set -euxo pipefail - VERSION="${{ github.event.inputs.version || github.ref_name }}" - VERSION="${VERSION#v}" - - metadata="$( - node packages/lime-cli-npm/scripts/build-release.js \ - --target-triple "${{ matrix.target }}" \ - --version "$VERSION" \ - --out-dir "packages/lime-cli-npm/dist" \ - --json - )" - - echo "$metadata" - asset_path="$(node -e 'const data = JSON.parse(process.argv[1]); process.stdout.write(data.archivePath);' "$metadata")" - echo "asset_path=$asset_path" >> "$GITHUB_OUTPUT" - - - name: Stage lime-cli release asset - if: matrix.build_cli && (matrix.platform == 'windows-2022' || matrix.platform == 'ubuntu-22.04' || (startsWith(matrix.platform, 'macos') && (steps.build_macos_primary.outcome == 'success' || steps.build_macos_retry.outcome == 'success'))) - shell: bash - run: | - set -euxo pipefail - ASSET_PATH="${{ steps.package_lime_cli.outputs.asset_path }}" - STAGING_DIR="release-assets/${{ matrix.target }}" - - if [ -z "$ASSET_PATH" ] || [ ! -f "$ASSET_PATH" ]; then - echo "lime-cli asset missing: $ASSET_PATH" >&2 - exit 1 - fi - - node -e " - const fs = require('fs'); - const path = require('path'); - const src = process.argv[1]; - const destDir = process.argv[2]; - fs.mkdirSync(destDir, { recursive: true }); - const dest = path.join(destDir, path.basename(src)); - fs.copyFileSync(src, dest); - process.stdout.write(dest); - " "$ASSET_PATH" "$STAGING_DIR" - echo - - name: Upload staged release assets artifact if: always() uses: actions/upload-artifact@v4 @@ -935,6 +868,18 @@ jobs: printf 'Generated updater manifest:\n' cat release-updater/latest.json + - name: Prepare GitHub release upload assets + env: + RELEASE_TAG: ${{ github.event.inputs.version || github.ref_name }} + shell: bash + run: | + set -euo pipefail + node scripts/prepare-github-release-assets.mjs \ + --assets-dir release-assets \ + --out-dir release-github-assets \ + --version "$RELEASE_TAG" \ + --extra-asset release-updater/latest.json + - name: Upload staged release assets to GitHub Release env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} @@ -944,24 +889,114 @@ jobs: TAG="${{ github.event.inputs.version || github.ref_name }}" assets=() - while IFS= read -r asset; do - [ -n "$asset" ] || continue + while IFS= read -r -d '' asset; do assets+=("$asset") - done < <(find "release-assets" -type f ! -name "latest*.json" | sort) - assets+=("release-updater/latest.json") + done < <(find "release-github-assets" -type f -print0 | sort -z) if [ "${#assets[@]}" -eq 0 ]; then - echo "No staged release assets found" >&2 + echo "No prepared GitHub release assets found" >&2 exit 1 fi - printf 'Uploading staged release assets:\n' + printf 'Uploading GitHub release assets:\n' printf ' - %s\n' "${assets[@]}" gh release upload "$TAG" "${assets[@]}" \ --repo "$GITHUB_REPOSITORY" \ --clobber + - name: Publish GitHub Release + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + RELEASE_TAG: ${{ github.event.inputs.version || github.ref_name }} + shell: bash + run: | + set -euo pipefail + gh release edit "$RELEASE_TAG" \ + --repo "$GITHUB_REPOSITORY" \ + --draft=false \ + --latest + + publish_updater_assets_r2: + name: Publish updater assets to Cloudflare R2 + needs: publish_release_assets + if: needs.publish_release_assets.result == 'success' + continue-on-error: true + runs-on: ubuntu-22.04 + + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + ref: ${{ github.event.inputs.source_ref || github.ref }} + + - name: Setup Node.js + uses: actions/setup-node@v4 + with: + node-version: "22" + + - name: Download staged release assets + uses: actions/download-artifact@v4 + with: + pattern: release-assets-* + path: release-assets + merge-multiple: true + + - name: Prepare updater release notes + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + RELEASE_TAG: ${{ github.event.inputs.version || github.ref_name }} + shell: bash + run: | + set -euo pipefail + NOTES_FILE="$RUNNER_TEMP/release-notes.md" + + if gh release view "$RELEASE_TAG" --repo "$GITHUB_REPOSITORY" --json body -q .body > "$NOTES_FILE" && [ -s "$NOTES_FILE" ]; then + echo "Loaded release notes from GitHub Release" + exit 0 + fi + + if [ -f RELEASE_NOTES.md ]; then + cp RELEASE_NOTES.md "$NOTES_FILE" + echo "Loaded release notes from RELEASE_NOTES.md" + exit 0 + fi + + PREV_TAG="$(git tag --sort=-v:refname | grep -v "^${RELEASE_TAG}$" | head -n 1)" + if [ -z "$PREV_TAG" ]; then + CHANGELOG="$(git log --pretty=format:"- %s (%h)" "${RELEASE_TAG}" 2>/dev/null || git log --pretty=format:"- %s (%h)")" + else + CHANGELOG="$(git log --pretty=format:"- %s (%h)" "${PREV_TAG}..${RELEASE_TAG}" 2>/dev/null || git log --pretty=format:"- %s (%h)" "${PREV_TAG}..HEAD")" + fi + + { + echo "## Lime ${RELEASE_TAG}" + echo + echo "### Changes" + echo "${CHANGELOG}" + } > "$NOTES_FILE" + echo "Generated release notes from git history" + + - name: Generate stable updater manifest + env: + LIME_UPDATES_BASE_URL: ${{ secrets.LIME_UPDATES_BASE_URL || vars.LIME_UPDATES_BASE_URL }} + RELEASE_TAG: ${{ github.event.inputs.version || github.ref_name }} + shell: bash + run: | + set -euo pipefail + BASE_URL="${LIME_UPDATES_BASE_URL:-https://updates.limecloud.com}" + node scripts/release-updater-manifest.mjs \ + --assets-dir release-assets \ + --out-dir release-updater \ + --version "$RELEASE_TAG" \ + --base-url "$BASE_URL" \ + --notes-file "$RUNNER_TEMP/release-notes.md" \ + --channel stable + + printf 'Generated updater manifest for R2:\n' + cat release-updater/latest.json + - name: Upload updater assets to Cloudflare R2 env: CLOUDFLARE_ACCOUNT_ID: ${{ secrets.CLOUDFLARE_ACCOUNT_ID }} @@ -979,6 +1014,7 @@ jobs: const args = [ 'npx wrangler@latest r2 object put', JSON.stringify(`${process.env.BUCKET}/${item.key}`), + '--remote', '--file', JSON.stringify(item.file), '--content-type', @@ -1006,7 +1042,13 @@ jobs: BUCKET="${LIME_RELEASES_R2_BUCKET:-lime-releases}" KEEP="${LIME_R2_KEEP_RELEASES:-3}" + if ! npx wrangler@latest r2 object --help | grep -q "object list"; then + echo "::warning::wrangler r2 object list is unavailable; skipping old R2 updater cleanup" + exit 0 + fi + npx wrangler@latest r2 object list "$BUCKET" \ + --remote \ --prefix "lime/stable/v" \ --json > "$RUNNER_TEMP/r2-objects.json" @@ -1027,5 +1069,112 @@ jobs: cat "$RUNNER_TEMP/r2-delete-keys.txt" while IFS= read -r key; do [ -n "$key" ] || continue - npx wrangler@latest r2 object delete "${BUCKET}/${key}" + npx wrangler@latest r2 object delete "${BUCKET}/${key}" --remote done < "$RUNNER_TEMP/r2-delete-keys.txt" + + publish_lime_cli_assets: + name: Publish lime-cli release assets + needs: publish_release_assets + if: needs.publish_release_assets.result == 'success' + continue-on-error: true + strategy: + fail-fast: false + matrix: + include: + - platform: macos-latest + target: aarch64-apple-darwin + name: macOS-arm64-cli + - platform: macos-15-intel + target: x86_64-apple-darwin + name: macOS-x64-cli + - platform: windows-2022 + target: x86_64-pc-windows-msvc + name: Windows-x64-cli + - platform: ubuntu-22.04 + target: x86_64-unknown-linux-gnu + name: Linux-x64-cli + + runs-on: ${{ matrix.platform }} + + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + ref: ${{ github.event.inputs.source_ref || github.ref }} + + - name: Setup Node.js + uses: actions/setup-node@v4 + with: + node-version: "22" + + - name: Setup Rust + uses: dtolnay/rust-toolchain@stable + with: + targets: ${{ matrix.target }} + + - name: Setup sccache + uses: mozilla-actions/sccache-action@v0.0.9 + + - name: Setup Rust cache + uses: Swatinem/rust-cache@v2 + with: + workspaces: src-tauri + shared-key: "lime-cli-${{ matrix.target }}" + cache-on-failure: true + cache-all-crates: true + save-if: ${{ github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/') }} + + - name: Build lime-cli release binary + shell: bash + env: + CARGO_PROFILE_RELEASE_LTO: "off" + CARGO_PROFILE_RELEASE_CODEGEN_UNITS: 32 + CARGO_INCREMENTAL: 0 + CARGO_TARGET_DIR: src-tauri/target + SCCACHE_GHA_ENABLED: "true" + RUSTC_WRAPPER: sccache + run: | + set -euxo pipefail + cargo build --manifest-path src-tauri/Cargo.toml -p lime-cli --release --target "${{ matrix.target }}" + + - name: Package lime-cli release asset + id: package_lime_cli + shell: bash + run: | + set -euxo pipefail + VERSION="${{ github.event.inputs.version || github.ref_name }}" + VERSION="${VERSION#v}" + + metadata="$( + node packages/lime-cli-npm/scripts/build-release.js \ + --target-triple "${{ matrix.target }}" \ + --version "$VERSION" \ + --out-dir "packages/lime-cli-npm/dist" \ + --json + )" + + echo "$metadata" + asset_path="$(node -e 'const data = JSON.parse(process.argv[1]); process.stdout.write(data.archivePath);' "$metadata")" + echo "asset_path=$asset_path" >> "$GITHUB_OUTPUT" + + - name: Upload lime-cli asset to GitHub Release + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + RELEASE_TAG: ${{ github.event.inputs.version || github.ref_name }} + shell: bash + run: | + set -euo pipefail + ASSET_PATH="${{ steps.package_lime_cli.outputs.asset_path }}" + if [ -z "$ASSET_PATH" ] || [ ! -f "$ASSET_PATH" ]; then + echo "lime-cli asset missing: $ASSET_PATH" >&2 + exit 1 + fi + + gh release upload "$RELEASE_TAG" "$ASSET_PATH" \ + --repo "$GITHUB_REPOSITORY" \ + --clobber + + - name: Show sccache stats + if: always() + run: sccache --show-stats diff --git a/README.md b/README.md index 69c9acbdb..5e60631ff 100644 --- a/README.md +++ b/README.md @@ -86,7 +86,7 @@ Lime 是一个基于 Tauri 的桌面应用,面向创作者、内容团队与 - 基于 Aster Agent Runtime - 支持会话、流式执行、技能调用、子任务接力与长期运行 -- 底层保留多 Provider 接入、凭证池、路由与协议兼容能力 +- 底层保留多 Provider API Key 接入、路由与协议兼容能力 --- diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index 6a0689fb8..d985fee6f 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,109 +1,96 @@ -## Lime v1.21.0 +## Lime v1.22.0 -发布日期:`2026-04-28` +发布日期:`2026-04-29` ### 发布概览 -- 本次发布目标 tag 为 `v1.21.0`。 -- 本次发布聚焦稳定版自动更新与 R2 分发链路、OEM 云端商业闭环、工作台首页与侧栏体验、资源管理器、Provider 模型管理收口、主题外观与设置页更新。 -- 本轮待递交内容覆盖 Rust 后端、Tauri update command、DevBridge / mock、发布工作流、前端 Workspace / Settings / Provider Pool / Resource Manager、测试覆盖、版本锁文件与临时产物清理。 +- 本次发布目标 tag 为 `v1.22.0`。 +- 本次发布聚焦稳定版 GitHub Release / R2 分发链路收口、lime-cli 独立产物发布、Provider / Credential 旧路径清退、云端用户中心商业边界收口,以及 Agent 会话恢复与模型选择体验稳定性。 +- 本轮待递交内容覆盖 Rust 后端、Tauri 配置、发布工作流、release asset 脚本、Provider / Model / Credential 治理、前端 Workspace / Settings / Provider API Key 主路径、测试覆盖、版本锁文件与执行计划文档。 ### 重点更新 -#### 1. 稳定版更新与 R2 发布链路 +#### 1. 版本号同步到 v1.22.0 -- `.github/workflows/release.yml` 补齐稳定版 updater 发布门禁,要求签名密钥与更新地址就绪后再生成 updater artifacts。 -- 发布流程会规范化 Tauri updater 公钥,并在仅产出 sidecar `.sig` 时由发布脚本生成稳定版 `latest.json`。 -- 新增 `scripts/release-updater-manifest.mjs`,聚合各平台 `latest.json`,生成统一 `latest.json`、版本化清单、R2 上传计划与 manifest metadata。 -- 新增 `scripts/plan-r2-release-cleanup.mjs`,按稳定版本窗口规划旧 R2 updater 产物清理,避免发布桶无限增长。 -- `scripts/release-updater-manifest.test.mjs` 覆盖平台缺失、版本不匹配、同名跨平台 artifact 与旧版本清理保护逻辑。 -- `src-tauri/src/commands/update_cmd.rs` 与 `src-tauri/crates/services/src/update_check_service.rs` 切到静态清单检查 + Tauri updater 安装主链,并补齐 semver 比较与缓存兜底。 - -#### 2. 版本号同步到 v1.21.0 - -- 应用版本已同步为 `1.21.0`: +- 应用版本已同步为 `1.22.0`: - `package.json` - `package-lock.json` - `src-tauri/Cargo.toml` - `src-tauri/Cargo.lock` - `src-tauri/tauri.conf.json` - `src-tauri/tauri.conf.headless.json` -- `packages/lime-cli-npm/package.json` 与 `packages/lime-cli-npm/README.md` 已同步到 `1.21.0`,保持 CLI wrapper 与桌面 release 版本一致。 -- 浏览器模式默认 mock 的 update current version 已同步为 `1.21.0`。 +- `packages/lime-cli-npm/package.json` 与 `packages/lime-cli-npm/README.md` 已同步到 `1.22.0`,保持 CLI wrapper 与桌面 release 版本一致。 +- 浏览器模式默认 mock 的 update current version 已同步为 `1.22.0`。 +- GitHub release asset staging 测试中的当前发布样例已同步到 `v1.22.0`。 -#### 3. OEM 云端商业闭环 +#### 2. 稳定版发布与 R2 分发链路 -- 新增 `src/lib/oemCloudPaymentReturn.ts`,统一生成、解析、暂存并分发 `lime://payment/return` 支付回跳事件。 -- `useDeepLink` 识别支付回跳 deep link,直接分发云端商业刷新事件,不再走旧 `handle_deep_link` 命令分支。 -- `useOemCloudAccess` 接入云端激活、支付配置、套餐订单、充值订单、账本、积分余额与访问令牌刷新主链。 -- 套餐购买和积分充值 checkout 支持 HTTPS bridge 回跳 URL,支付完成后自动刷新云端权益、积分余额与订单 watcher。 -- `docs/exec-plans/oem-cloud-commerce-loop-progress.md` 记录当前阶段、已清退的旧支付配置入口与下一轮真实渠道沙箱验证计划。 +- `.github/workflows/release.yml` 将桌面应用构建、GitHub Release 资产发布、R2 updater 发布和 lime-cli 资产发布拆成更清晰的阶段。 +- 新增 `scripts/prepare-github-release-assets.mjs`,在上传 GitHub Release 前统一整理资产名,避免 macOS `Lime.app.tar.gz` / `.sig` 同名跨架构冲突。 +- GitHub Release 上传改为使用 `release-github-assets` 暂存目录,并在资产上传后显式发布 release、标记 latest。 +- R2 updater 发布改为独立 job,从 GitHub Release 或 `RELEASE_NOTES.md` 准备 updater release notes,再生成稳定版 manifest。 +- Cloudflare R2 上传 / 列表 / 删除命令补齐 `--remote`,并在 wrangler 不支持 `r2 object list` 时跳过旧版本清理而不是阻塞发布。 +- lime-cli release binary 与 npm wrapper 资产改为独立矩阵 job 发布,保留 macOS / Windows / Linux CLI 产物,不再耦合桌面安装包矩阵。 -#### 4. 工作台首页、侧栏与导航体验 +#### 3. Provider / Credential 旧路径清退 -- `AppSidebar` 增加最近对话 / 归档会话架、分页加载、归档切换、外观切换、账户菜单与折叠态细节。 -- 新增 `src/components/app-sidebar/AppSidebarConversationShelf.tsx`,把会话架从侧栏主体中拆出,降低侧栏单体复杂度。 -- 工作台首页空态升级为“先开始这一轮 / 继续这轮 / 直接开工”入口,强化任务起手、推荐模板与继续上下文。 -- Workspace / Task Center / ChatNavbar / Inputbar / EmptyState / Team Workspace 等主路径继续收口视觉状态、运行时状态与回归断言。 -- `src/lib/windowControls.ts`、窗口 chrome 与主窗口启动链路继续补齐 macOS / headless 场景下的窗口控制一致性。 +- 清退旧 Provider Pool 页面、凭证卡片、Credential 表单、OAuth / Kiro / Antigravity / Claude OAuth / usage 等旧命令与服务路径。 +- Rust 后端删除旧 credential crate、provider pool DAO / service、Kiro credential handler、旧 provider converter / translator / fingerprint 模型等 dead surface。 +- 前端保留当前 API Key Provider 设置主路径,并继续收口模型启用、模型能力、Prompt Cache 与 companion provider 概览口径。 +- `agentCommandCatalog`、`legacySurfaceCatalog`、DevBridge mock 与相关测试同步更新,避免已删除命令继续作为 current surface 出现。 +- 模型资源索引删除旧 Antigravity / Kiro / Codex alias/provider 静态入口,减少 provider 真相源分叉。 -#### 5. 资源管理器 +#### 4. Agent 会话恢复与工作台稳定性 -- 新增 `src/features/resource-manager/`,提供资源管理器页面、侧栏、工具栏、预览面板、Inspector、搜索与导航意图。 -- 支持图片、文本、Markdown、PDF、Office、音视频、数据文件、压缩包与系统委托类型的分层预览渲染。 -- 支持资源下载、复制、系统打开、Finder 揭示、聊天位置与项目资源上下文回跳。 -- 补齐 `ResourceManagerPage`、资源预览搜索、会话状态和导航意图测试。 +- 会话切换 / 恢复详情默认按 `historyLimit: 40` 拉取近期历史,完整历史加载仍通过显式 `historyLimit: 0` 入口完成。 +- `useAsterAgentChat` 回归断言已同步新的 session detail 拉取参数,覆盖 stop refresh、timeline cache hydrate、workspace guard 与 stale 快照刷新路径。 +- 工作台消息流、模型选择、Provider selector、Team Workspace、artifact / saved content 展示继续保持与 runtime execution metadata 对齐。 +- `ModelSelector`、`useConfiguredProviders`、`useProviderModels`、Prompt Cache 支持判断与 companion provider overview 补齐回归覆盖。 -#### 6. Provider Pool 与设置页收口 +#### 5. 云端用户中心与商业边界 -- Provider 模型管理改为“启用的模型”左侧列表 + 添加模型面板,删除旧 API Key 列表 / Provider 表单 / 模型列表拆分组件。 -- 新增 `ModelProviderList`、`ModelAddPanel`、`providerConfigUtils` 与连接测试类型,统一 Provider 配置工具函数与 UI 入口。 -- 设置页 Provider、About、Developer、Experimental、Appearance、Channels 与 Automation 页面继续收口布局、状态展示和回归断言。 -- Prompt Cache 与 Anthropic-compatible 能力口径更新,避免把显式 `cache_control` 能力误显示为自动缓存。 +- 新增 `docs/exec-plans/cloud-commerce-user-center-boundary.md`,明确套餐购买、支付、账单、用量明细统一收敛到 `limecore` 用户中心网页。 +- Lime 客户端云端服务设置面继续收口为会话状态、当前套餐、积分余额、待支付提醒与用户中心跳转入口。 +- 客户端移除直接创建套餐 / 充值订单的旧处理面,避免本地商业工作台与用户中心形成双轨。 +- `useOemCloudAccess` 与 OEM cloud / LimeHub provider 同步测试继续覆盖登录态、权益摘要、API Key 与回跳刷新路径。 -#### 7. App Update 前端与 mock +#### 6. 文档、治理与回归 -- `src/lib/api/appUpdate.ts` 扩展 release notes URL / pubDate / 错误信息字段。 -- `src/lib/tauri-mock/core.ts` 补齐 `check_update`、`check_for_updates` 与下载无更新态 mock,浏览器模式不再落入 unknown command。 -- About 设置页更新检查、下载失败、诊断错误和版本展示补齐测试覆盖。 - -#### 8. 文档、治理与临时产物 - -- `README.md` 更新产品定位文案:从“本地优先的 AI API Proxy 桌面应用”收敛为 AI Agent 创作工作台。 -- `src/lib/governance/legacySurfaceCatalog.json` 与测试补充新的 legacy surface 口径。 -- 删除根目录临时调试产物:`monitor.sh`、`network-before.md`、`post-hmr-state.png`、`tmp-e2e-home.png`、`knip.governance.json`。 -- 新增 `theme-scope-messages-ocean.png` 作为本轮主题视觉验证产物。 +- `docs/aiprompts/` 下 Provider、Credential Pool、Services、Hooks、Components、Overview 等导航文档同步当前 provider / credential / model registry 事实源。 +- `docs/content/03.providers/1.overview.md` 与 `src/components/provider-pool/api-key/README.md` 更新当前 Provider 配置入口说明。 +- `scripts/release-updater-manifest.test.mjs` 增加 GitHub release asset staging 覆盖,保护 macOS 同名 updater bundle 重命名逻辑。 +- `src-tauri/proptest-regressions/` 已纳入本轮待递交范围,保留 property test 回归种子。 ### 待递交范围确认 -- 版本与发布:版本文件、lockfile、CLI wrapper、release workflow、R2 updater manifest / cleanup 脚本与测试。 -- Rust 主链:update command、update service、window chrome、runner/app 模块、DevBridge dispatcher、tray 事件与菜单处理。 -- 前端主链:AppSidebar、Workspace、Task Center、EmptyState、Inputbar、Team Workspace、Settings、Provider Pool、Resource Manager、MCP、Memory、SceneApps、Resources。 -- 商业闭环:OEM cloud control plane API、支付回跳 deep link、权益 / 积分 / 账本刷新、订单 watcher 与执行计划文档。 -- 验证与治理:新增/更新测试、legacy catalog、release updater contract、删除临时调试文件与旧 Provider Pool 组件。 +- 版本与发布:版本文件、lockfile、Tauri 配置、CLI wrapper、release workflow、GitHub release asset staging 脚本与测试。 +- Rust 主链:Provider / Credential / Server / Services / Agent / DevBridge / model registry / router / websocket 相关 current surface 收口。 +- 前端主链:Agent Chat Workspace、MessageList、ModelSelector、Settings Provider、API Key Provider、Provider hooks、mock 与治理目录册。 +- 商业边界:云端用户中心执行计划、OEM cloud access / LimeHub provider sync、设置页云端服务入口。 +- 验证与治理:新增/更新测试、legacy catalog、release updater contract、删除旧 Provider Pool / credential / Kiro / Antigravity 等 dead surface。 ### 校验状态 - 已通过: - `npm run verify:app-version` - `cargo fmt --manifest-path "src-tauri/Cargo.toml" --all` - - `cargo test --manifest-path "src-tauri/Cargo.toml"` — 1112 passed / 0 failed / 2 ignored + - `cargo test --manifest-path "src-tauri/Cargo.toml"` — 1070 passed / 0 failed / 2 ignored - `cargo clippy --manifest-path "src-tauri/Cargo.toml" --all-targets --all-features` - `npm run lint` - - `npm test` — 44 个 Vitest smart 批次通过 + - `npx vitest run "src/components/agent/chat/hooks/useAsterAgentChat.test.tsx"` + - `npm test` — 43 个 Vitest smart 批次通过 - `npm run test:contracts` - - `npm run smoke:agent-runtime-tool-surface` - - `npm run smoke:agent-runtime-tool-surface-page` - `git diff --check` - `cargo test` 通过,当前存在 1 条预存 warning: - `write_auxiliary_runtime_projection_fixture` 的 `dead_code` -- `cargo clippy` 通过,当前存在 4 条预存 warning: +- `cargo clippy` 通过,当前存在 6 条 warning: - `crates/services/src/aster_session_store.rs` 的 `manual_repeat_n` - - `crates/skills/src/lime_llm_provider.rs` 的 `too_many_arguments` + - `crates/skills/src/lime_llm_provider.rs` 的 2 处 `too_many_arguments` + - `crates/agent/src/request_tool_policy.rs` 的 `too_many_arguments` - `crates/agent/src/session_execution_runtime.rs` 的 `needless_lifetimes` - `src/services/runtime_evidence_pack_service.rs` 的 `dead_code` -- GUI 主路径补充复测已通过:`smoke:agent-runtime-tool-surface` 与 `smoke:agent-runtime-tool-surface-page` 均确认 Harness 入口在执行态可见,修复此前等待 Harness 按钮超时的问题。 +- GUI 主路径未额外执行 `npm run verify:gui-smoke`;本轮发布收口以版本、发布链路、Provider / Credential 治理和前端 / Rust 回归为主要风险覆盖。 --- -**完整变更**: `v1.20.0` -> `v1.21.0` +**完整变更**: `v1.21.0` -> `v1.22.0` diff --git a/docs/aiprompts/README.md b/docs/aiprompts/README.md index 985f8b07f..8198dfab4 100644 --- a/docs/aiprompts/README.md +++ b/docs/aiprompts/README.md @@ -55,7 +55,7 @@ ### Provider 与数据 - `providers.md` - Provider 接入与认证方式 -- `credential-pool.md` - 凭证池、运行时路径、调试入口 +- `credential-pool.md` - 凭证池退役边界、启动清理与守卫 - `converter.md` - 协议转换与兼容层 - `database.md` - 数据库层与持久化 diff --git a/docs/aiprompts/aster-integration.md b/docs/aiprompts/aster-integration.md index 154aca4c1..242167b4f 100644 --- a/docs/aiprompts/aster-integration.md +++ b/docs/aiprompts/aster-integration.md @@ -2,7 +2,7 @@ ## 集成状态 ✅ -Lime 已完整集成 aster-rust 框架,包括凭证池桥接。 +Lime 已完整集成 aster-rust 框架。Provider 配置桥接已收敛到 API Key Provider。 ## 当前事实源 @@ -18,13 +18,12 @@ Lime 已完整集成 aster-rust 框架,包括凭证池桥接。 - `aster_state.rs` - Agent 状态管理 - `aster_agent.rs` - Agent 包装器 - `event_converter.rs` - 事件转换器 -- `credential_bridge.rs` - 凭证池桥接 +- `credential_bridge.rs` - API Key Provider 桥接 **Tauri 命令** (`src-tauri/src/commands/aster_agent_cmd.rs`): - `aster_agent_init` - 初始化 Agent - `aster_agent_configure_provider` - 手动配置 Provider -- `aster_agent_configure_from_pool` - 从凭证池配置 Provider(推荐) - `aster_agent_status` - 获取状态 - `agent_runtime_submit_turn` - 统一提交 turn - `agent_runtime_interrupt_turn` - 统一中断 turn @@ -59,9 +58,9 @@ Lime 已完整集成 aster-rust 框架,包括凭证池桥接。 │ │ │ │ │ ▼ ▼ │ │ ┌─────────────────────────────────────┐ │ -│ │ Lime 凭证池 │ │ -│ │ - ProviderPoolService │ │ +│ │ Lime Provider 配置 │ │ │ │ - ApiKeyProviderService │ │ +│ │ - ModelRegistryService │ │ │ └─────────────────────────────────────┘ │ └─────────────────────────────────────────────────────────────────┘ │ @@ -75,7 +74,7 @@ Lime 已完整集成 aster-rust 框架,包括凭证池桥接。 └─────────────────────────────────────────────────────────────────┘ ``` -## 凭证池桥接 +## Provider 桥接 ### 支持的凭证类型映射 @@ -83,12 +82,11 @@ Lime 已完整集成 aster-rust 框架,包括凭证池桥接。 | -------------------------- | -------------- | | OpenAIKey | openai | | ClaudeKey / AnthropicKey | anthropic | -| KiroOAuth | bedrock | -| GeminiOAuth / GeminiApiKey | google | +| GeminiApiKey | google | | VertexKey | gcpvertexai | -| CodexOAuth | codex | -| ClaudeOAuth | anthropic | -| AntigravityOAuth | google | +| Codex API Key | codex | + +Kiro / Gemini OAuth / Codex OAuth / Claude OAuth / Antigravity OAuth 均已退役,不再作为 Aster Provider 配置来源。 ### 使用方式 @@ -127,4 +125,4 @@ await submitAgentRuntimeTurn({ - [overview.md](overview.md) - 项目架构 - [providers.md](providers.md) - Provider 系统 -- [credential-pool.md](credential-pool.md) - 凭证池管理 +- [credential-pool.md](credential-pool.md) - 凭证池退役说明 diff --git a/docs/aiprompts/commands.md b/docs/aiprompts/commands.md index b0adadc86..dc5dde429 100644 --- a/docs/aiprompts/commands.md +++ b/docs/aiprompts/commands.md @@ -324,7 +324,7 @@ Companion 桌宠链路同样遵循这条路径。当前主入口为 `src/lib/api Lime 主应用会在本地维护 `ws://127.0.0.1:45554/companion/pet` 的桌宠 companion 入口。前端如需感知桌宠连接状态,应继续通过 `companion-pet-status` 事件监听统一状态,不要在页面或 Hook 里自行直连本地 `WebSocket`。 -如果 companion 协议继续扩展,也应优先延续“Lime 做宿主、桌宠只收脱敏派生状态”的边界。例如 provider 凭证池相关能力,允许 Lime 通过 `companion_send_pet_command` 下发诸如 `pet.provider_overview` 这类脱敏摘要,并允许桌宠通过 `pet.open_provider_settings` 请求 Lime 聚焦主窗口并跳到 `设置 -> AI 服务商`,或通过 `pet.request_provider_overview_sync` 请求 Lime 立即重发最新的脱敏摘要;桌宠交互增强能力也应继续走这条主链,例如双击 / 三击桌宠后发出 `pet.request_pet_cheer`、`pet.request_pet_next_step`,或通过 `pet.request_chat_reply` 携带用户输入文本,请求 Lime 代为调用当前可聊天模型,再统一回写 `pet.show_bubble`;但不允许桌宠直接读取凭证文件、数据库或内部 `/v1/credentials/*` 完整凭证接口。 +如果 companion 协议继续扩展,也应优先延续“Lime 做宿主、桌宠只收脱敏派生状态”的边界。例如 provider 概览相关能力,允许 Lime 通过 `companion_send_pet_command` 下发诸如 `pet.provider_overview` 这类脱敏摘要,并允许桌宠通过 `pet.open_provider_settings` 请求 Lime 聚焦主窗口并跳到 `设置 -> AI 服务商`,或通过 `pet.request_provider_overview_sync` 请求 Lime 立即重发最新的脱敏摘要;桌宠交互增强能力也应继续走这条主链,例如双击 / 三击桌宠后发出 `pet.request_pet_cheer`、`pet.request_pet_next_step`,或通过 `pet.request_chat_reply` 携带用户输入文本,请求 Lime 代为调用当前可聊天模型,再统一回写 `pet.show_bubble`;但不允许桌宠直接读取凭证文件、数据库或内部 `/v1/credentials/*` 完整凭证接口。 ## 命令契约的五个事实源 diff --git a/docs/aiprompts/components.md b/docs/aiprompts/components.md index 6c2df71a4..1bfffdd88 100644 --- a/docs/aiprompts/components.md +++ b/docs/aiprompts/components.md @@ -9,7 +9,7 @@ React 组件层,使用 TailwindCSS 和 shadcn/ui。 ``` src/components/ ├── ui/ # 基础 UI 组件 (shadcn/ui) -├── provider-pool/ # 凭证池管理 +├── provider-pool/api-key/ # API Key Provider 设置(目录名保留历史路径) ├── flow-monitor/ # 流量监控 ├── general-chat/ # 兼容画布桥接(非对话主入口) ├── terminal/ # 已移除的历史终端页面路径(不要恢复) @@ -39,24 +39,17 @@ export function AppSidebar() { } ``` -### ProviderPool +### ApiKeyProviderSection -凭证池管理组件。 +API Key Provider 设置组件,是当前设置页的 Provider 配置入口。 ```tsx -// src/components/provider-pool/ProviderPoolPanel.tsx -export function ProviderPoolPanel() { - const { credentials, addCredential, removeCredential } = useProviderPool(); - - return ( -
- - -
- ); -} +// src/components/provider-pool/api-key/ApiKeyProviderSection.tsx +import { ApiKeyProviderSection } from "@/components/provider-pool/api-key"; ``` +旧 `ProviderPoolPage`、凭证卡片、OAuth 表单与 `useProviderPool` 已退役,不允许重新作为设置入口。 + ### FlowMonitor 流量监控组件。 diff --git a/docs/aiprompts/credential-pool.md b/docs/aiprompts/credential-pool.md index a44676110..98440d874 100644 --- a/docs/aiprompts/credential-pool.md +++ b/docs/aiprompts/credential-pool.md @@ -1,229 +1,64 @@ -# 凭证池管理 +# 凭证池退役说明 -## 概述 +## 当前状态 -凭证池管理系统实现多凭证轮询负载均衡、健康检查和自动 Token 刷新。 +凭证池管理系统已退役,分类为 `dead`。 -## 核心组件 +后续 Provider 凭证与模型选择只允许收敛到以下 `current` 主路径: -``` -src-tauri/src/ -├── credential/ # 凭证池核心 -│ ├── mod.rs -│ ├── pool.rs # 凭证池实现 -│ └── health.rs # 健康检查 -└── services/ - ├── provider_pool_service.rs # 池服务 - └── token_cache_service.rs # Token 缓存 +- API Key Provider:应用内 Provider 配置、连接测试、模型发现与默认模型选择 +- configured providers:用户已配置 Provider 的事实源 +- 模型注册表:Provider / model 目录与协议能力事实源 +- 协议转换器:请求协议适配能力,例如 coding plan 仍可使用 `openai_to_antigravity` + +旧凭证池不再提供多凭证轮询、OAuth 登录、本地 CLI 凭证导入、健康检查、Token 自动刷新、使用量统计或管理页面。 + +## 已下线范围 + +- 前端凭证池管理页、旧凭证卡片、OAuth / 本地凭证表单、`useProviderPool` 与 `providerPool` API 网关 +- Rust `ProviderPoolService`、`TokenCacheService`、Kiro 事件服务、OAuth 命令与旧 provider pool 命令 +- Kiro / Qwen / Antigravity / Codex OAuth / Claude OAuth / Gemini OAuth 这类登录型或本地 CLI 凭证运行时 +- `provider_pool_credentials` 的运行时读取、健康写回与 fallback 选择 + +## 数据处理 + +启动期迁移会清理 Lime 管理的旧凭证池数据: + +- 清空 `provider_pool_credentials` +- 删除 Lime 应用数据目录下托管的 `credentials/` 副本 + +这只处理 Lime 自己复制和管理过的数据,不删除用户外部 CLI 原始目录,例如 `~/.codex`、`~/.gemini` 或其它第三方工具目录。 + +## 保留边界 + +以下内容不是凭证池功能,仍可继续演进: + +- API Key Provider 中的 OpenAI、Anthropic、Gemini API Key、OpenRouter、GitHub、Azure 等配置 +- 模型名或模型系列中出现的 `codex`、`gemini`、`qwen` 等字符串 +- 协议转换器,尤其是 `src-tauri/crates/providers/src/converter/openai_to_antigravity.rs` +- server 内部短期用于桥接 API Key Provider 的兼容 DTO;它只能承载 current API Key Provider 数据,不代表凭证池恢复 + +## 守卫 + +旧 UI / Hook / API 文件路径已登记到 `src/lib/governance/legacySurfaceCatalog.json`,不允许重新接回前端入口。 + +涉及 Provider 或命令边界时,至少执行: + +```bash +npm run test:contracts +npm run governance:legacy-report ``` -## 凭证池架构 +如果改动影响设置页或 Agent 运行主路径,再补: -``` -┌─────────────────────────────────────────────────────────────────┐ -│ ProviderPoolService │ -│ ┌─────────────────────────────────────────────────────────────┐│ -│ │ Credential Pool ││ -│ │ ┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐ ││ -│ │ │ Cred 1 │ │ Cred 2 │ │ Cred 3 │ │ Cred N │ ││ -│ │ │ Healthy │ │ Healthy │ │ Expired │ │ Healthy │ ││ -│ │ └────┬────┘ └────┬────┘ └────┬────┘ └────┬────┘ ││ -│ │ │ │ │ │ ││ -│ │ └────────────┴────────────┴────────────┘ ││ -│ │ │ ││ -│ │ Round Robin ││ -│ └─────────────────────────┼───────────────────────────────────┘│ -│ │ │ -│ ┌─────────────────────────┼───────────────────────────────────┐│ -│ │ Health Checker (定时任务) ││ -│ │ - Token 过期检查 ││ -│ │ - 自动刷新 ││ -│ │ - 不健康凭证剔除 ││ -│ └─────────────────────────────────────────────────────────────┘│ -└─────────────────────────────────────────────────────────────────┘ -``` - -## 负载均衡策略 - -### Round Robin (轮询) - -```rust -pub struct RoundRobinPool { - credentials: Vec, - current_index: AtomicUsize, -} - -impl RoundRobinPool { - pub fn next(&self) -> Option<&CredentialEntry> { - let healthy: Vec<_> = self.credentials - .iter() - .filter(|c| c.is_healthy()) - .collect(); - - if healthy.is_empty() { - return None; - } - - let index = self.current_index - .fetch_add(1, Ordering::Relaxed) % healthy.len(); - Some(healthy[index]) - } -} -``` - -### 权重轮询 (可选) - -```rust -pub struct WeightedPool { - credentials: Vec<(CredentialEntry, u32)>, // (凭证, 权重) -} -``` - -## 健康检查 - -### 检查项目 - -| 检查项 | 说明 | 频率 | -| ---------- | ------------------ | ------------ | -| Token 过期 | 检查 expires_at | 每次请求前 | -| Token 刷新 | 尝试刷新过期 Token | Token 过期时 | -| API 可用性 | 发送测试请求 | 定时 (5分钟) | - -### 健康状态 - -```rust -pub enum HealthStatus { - Healthy, // 健康 - TokenExpired, // Token 过期 - TokenRefreshing, // 正在刷新 - RefreshFailed(String), // 刷新失败 - Unhealthy(String), // 不健康 - Disabled, // 已禁用 -} -``` - -### 自动恢复 - -```rust -// 健康检查任务 -async fn health_check_task(pool: Arc) { - loop { - for credential in pool.credentials() { - match credential.health_status() { - HealthStatus::TokenExpired => { - // 尝试刷新 - if let Err(e) = pool.refresh_token(&credential).await { - credential.set_status(HealthStatus::RefreshFailed(e)); - } - } - HealthStatus::RefreshFailed(_) => { - // 重试刷新 (最多 3 次) - if credential.retry_count() < 3 { - pool.retry_refresh(&credential).await; - } - } - _ => {} - } - } - - tokio::time::sleep(Duration::from_secs(300)).await; - } -} -``` - -## Token 缓存 - -### 缓存策略 - -```rust -pub struct TokenCacheService { - cache: DashMap, -} - -struct CachedToken { - access_token: String, - expires_at: i64, - refresh_token: String, -} - -impl TokenCacheService { - pub async fn get_or_refresh(&self, credential_id: &str) -> Result { - if let Some(cached) = self.cache.get(credential_id) { - if !cached.is_expired() { - return Ok(cached.access_token.clone()); - } - } - - // 刷新并缓存 - let new_token = self.refresh(credential_id).await?; - self.cache.insert(credential_id.to_string(), new_token.clone()); - Ok(new_token.access_token) - } -} -``` - -### 数据库持久化 - -```sql -CREATE TABLE token_cache ( - credential_id TEXT PRIMARY KEY, - access_token TEXT NOT NULL, - refresh_token TEXT NOT NULL, - expires_at INTEGER NOT NULL, - updated_at INTEGER NOT NULL -); -``` - -## 凭证生命周期 - -``` -┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐ -│ 上传 │ ──▶ │ 验证 │ ──▶ │ 激活 │ ──▶ │ 使用中 │ -└─────────┘ └─────────┘ └─────────┘ └────┬────┘ - │ - ┌────────────────────────────────┘ - │ - ▼ - ┌─────────┐ ┌─────────┐ ┌─────────┐ - │ 过期 │ ──▶ │ 刷新 │ ──▶ │ 恢复 │ - └─────────┘ └────┬────┘ └─────────┘ - │ - ▼ (失败) - ┌─────────┐ - │ 禁用 │ - └─────────┘ -``` - -## API 接口 - -### Tauri Commands - -```rust -#[tauri::command] -async fn add_credential(provider: String, path: String) -> Result<()>; - -#[tauri::command] -async fn remove_credential(id: String) -> Result<()>; - -#[tauri::command] -async fn list_credentials() -> Result>; - -#[tauri::command] -async fn refresh_credential(id: String) -> Result<()>; - -#[tauri::command] -async fn get_pool_status() -> Result; +```bash +npm run verify:local +npm run verify:gui-smoke ``` ## 相关文档 -- [providers.md](providers.md) - Provider 系统 -- [services.md](services.md) - 业务服务 -- [database.md](database.md) - 数据库层 - -## 运行时路径与调试 - -- 凭证文件默认存放在应用数据目录下的 `lime/credentials/` -- `~/Library/Application Support/lime/credentials/` 只作为 macOS 示例,Windows 必须使用对应的应用数据目录 -- `request_logs`、日志目录等运行时路径也应通过统一 `app_paths` / 系统目录 API 获取,不要在实现里写死 -- 需要排查 Kiro 凭证加载时,可使用 `debug_kiro_credentials` 对应命令进行诊断;具体命令边界以 `docs/aiprompts/commands.md` 和 Rust 注册表为准 +- [providers.md](providers.md) - Provider current 主路径 +- [commands.md](commands.md) - Tauri 命令边界 +- [database.md](database.md) - 数据库层与启动迁移 +- [converter.md](converter.md) - 协议转换 diff --git a/docs/aiprompts/database.md b/docs/aiprompts/database.md index 1adc44be0..73a3f2cce 100644 --- a/docs/aiprompts/database.md +++ b/docs/aiprompts/database.md @@ -103,4 +103,4 @@ pub fn run_migrations(conn: &Connection) -> Result<()> { ## 相关文档 - [services.md](services.md) - 业务服务 -- [credential-pool.md](credential-pool.md) - 凭证池管理 +- [credential-pool.md](credential-pool.md) - 凭证池退役说明 diff --git a/docs/aiprompts/hooks.md b/docs/aiprompts/hooks.md index 02786951b..6c9090657 100644 --- a/docs/aiprompts/hooks.md +++ b/docs/aiprompts/hooks.md @@ -9,8 +9,7 @@ ``` src/hooks/ ├── index.ts # 导出入口 -├── useProviderPool.ts # 凭证池管理 -├── useOAuthCredentials.ts # OAuth 凭证 +├── useConfiguredProviders.ts # 已配置 Provider 读取 ├── useFlowEvents.ts # 流量事件 ├── useMcpServers.ts # MCP 服务器 ├── useDeepLink.ts # Deep Link 处理 @@ -23,6 +22,7 @@ src/hooks/ - 历史 `useTauri.ts` 兼容聚合层已删除,不要重新引入新的“大一统 API Hook”。 - Agent 工作台统一走 `src/components/agent/chat/hooks/index.ts` 暴露的 `useAgentChatUnified`,底层实现委托 `useAsterAgentChat`。 - 历史 `@/hooks/useUnifiedChat` 与 `src/lib/api/unified-chat.ts` 已删除,不要重建 compat Hook / API。 +- 凭证池 Hook `useProviderPool` 已删除;Provider 配置只允许走 API Key Provider / configured providers 主路径。 ## 核心 Hooks @@ -40,37 +40,9 @@ src/hooks/ - Hook 实现:`src/components/agent/chat/hooks/useAsterAgentChat.ts` - API 封装:`src/lib/api/agentRuntime.ts` -### useProviderPool +### useConfiguredProviders -```typescript -export function useProviderPool() { - const [credentials, setCredentials] = useState([]); - const [loading, setLoading] = useState(false); - - const refresh = async () => { - setLoading(true); - const list = await invoke("list_credentials"); - setCredentials(list); - setLoading(false); - }; - - const addCredential = async (provider: string, path: string) => { - await invoke("add_credential", { provider, filePath: path }); - await refresh(); - }; - - const removeCredential = async (id: string) => { - await invoke("remove_credential", { id }); - await refresh(); - }; - - useEffect(() => { - refresh(); - }, []); - - return { credentials, loading, addCredential, removeCredential, refresh }; -} -``` +Provider 列表与默认模型读取应优先复用 `src/lib/api/modelRegistry.ts`、`src/lib/api/appConfigTypes.ts` 和现有 configured provider Hook,不要重新创建凭证池 Hook 或 OAuth Hook。 ### useFlowEvents @@ -117,7 +89,7 @@ export function useDeepLink() { ### 命名约定 - 以 `use` 开头 -- 描述功能: `useProviderPool`, `useFlowEvents` +- 描述功能: `useConfiguredProviders`, `useFlowEvents` ### 返回值 diff --git a/docs/aiprompts/overview.md b/docs/aiprompts/overview.md index 548c3edf4..749fdbce2 100644 --- a/docs/aiprompts/overview.md +++ b/docs/aiprompts/overview.md @@ -8,7 +8,7 @@ Lime 是一个以创作为中心的本地优先 AI Agent 交互工作台,基 1. **产品层**:Workspace、通用工作区 / Harness、Agent 对话、Skills、Artifact/Canvas、记忆与风格 2. **能力层**:MCP、浏览器运行时、终端、插件、批量/心跳、Claw 渠道 -3. **基础设施层**:Aster Agent、Provider 凭证池、协议兼容、路由、服务器、数据库与监控 +3. **基础设施层**:Aster Agent、API Key Provider、模型注册表、协议兼容、路由、服务器、数据库与监控 其中,Provider 接入、协议兼容与运行时服务共同构成底层能力底座。 @@ -90,11 +90,11 @@ lime/ | 模块 | 说明 | |------|------| | `src-tauri/src/agent/` | Aster Agent 集成、会话历史、线程读模型、工具注册与流式桥接 | -| `providers/` | LLM Provider 认证和 API 实现 | +| `providers/` | LLM Provider API 实现与协议兼容 | | `services/` | 业务服务层 | | `converter/` | 协议转换与兼容层 | | `server/` | HTTP API 服务器 | -| `credential/` | 凭证池管理 | +| `api_key_provider` | API Key Provider 配置与凭证选择 | | `flow_monitor/` | 流量监控 | | `database/` | 数据持久化与 DAO | @@ -169,10 +169,10 @@ lime/ │ │ │ ▼ ▼ ▼ ┌─────────────────────────────────────────────────────────────────┐ -│ Provider Pool Service │ +│ API Key Provider / Model Registry │ │ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │ -│ │ 凭证轮询 │ │ 健康检查 │ │ Token 刷新 │ │ -│ │ (负载均衡) │ │ (自动剔除) │ │ (OAuth) │ │ +│ │ Provider │ │ 模型目录 │ │ 协议能力 │ │ +│ │ 配置选择 │ │ 解析 │ │ 判断 │ │ │ └──────┬──────┘ └──────┬──────┘ └───────────┬─────────────┘ │ └─────────┼────────────────┼─────────────────────┼────────────────┘ │ │ │ @@ -219,9 +219,10 @@ lime/ - `write_file`、画布联动与主题工作流负责把过程沉淀成交付物 ### 7. 多 Provider 与兼容层 -- OAuth 与 API Key Provider 并存 -- 凭证池、模型路由、协议兼容与 HTTP Server 作为底层支撑 +- Provider 配置统一走 API Key Provider / configured providers +- 模型注册表、协议兼容与 HTTP Server 作为底层支撑 - Prompt Cache 等运行时能力按 ProviderType 判断;`anthropic-compatible` 只表示 Anthropic wire format 兼容,不等于自动 Prompt Cache 能力 +- 旧凭证池、OAuth 与本地 CLI credential runtime 已退役;Antigravity 仅保留协议转换器能力 ### 8. 本地优先与可扩展 - 桌面应用、本地工作区、插件与外部工具扩展 @@ -241,7 +242,7 @@ lime/ ### 基础设施 - [providers.md](providers.md) - Provider 系统 -- [credential-pool.md](credential-pool.md) - 凭证池管理 +- [credential-pool.md](credential-pool.md) - 凭证池退役说明 - [converter.md](converter.md) - 协议转换 - [server.md](server.md) - HTTP 服务器 diff --git a/docs/aiprompts/providers.md b/docs/aiprompts/providers.md index 48d16fe63..04a28a7aa 100644 --- a/docs/aiprompts/providers.md +++ b/docs/aiprompts/providers.md @@ -2,7 +2,7 @@ ## 概述 -Provider 系统负责与各 LLM 服务商的认证和 API 交互。支持 OAuth 和 API Key 两种认证方式。 +Provider 系统负责与各 LLM 服务商的 API 交互。当前认证事实源是 API Key Provider / configured providers;旧 OAuth 与本地 CLI 凭证池运行时已退役。 如果需求同时涉及“候选模型解析、OEM 与本地 provider 协同、自动与设置平衡、成本/限额事件”,继续补读: @@ -16,15 +16,11 @@ src-tauri/src/providers/ ├── mod.rs # 模块入口和 Provider 枚举 ├── traits.rs # Provider trait 定义 ├── error.rs # 错误类型 -├── kiro.rs # Kiro/CodeWhisperer OAuth -├── gemini.rs # Gemini OAuth -├── qwen.rs # Qwen OAuth -├── antigravity.rs # Antigravity OAuth -├── claude_oauth.rs # Claude OAuth +├── gemini.rs # Gemini API Key 请求支持 +├── antigravity.rs # Antigravity 协议兼容支持 ├── claude_custom.rs # Claude API Key ├── openai_custom.rs # OpenAI API Key -├── codex.rs # Codex Provider -├── iflow.rs # iFlow Provider +├── codex.rs # OpenAI Responses / Codex 兼容请求支持 ├── vertex.rs # Vertex AI Provider └── tests.rs # 单元测试 ``` @@ -33,15 +29,10 @@ src-tauri/src/providers/ ```rust pub enum ProviderType { - Kiro, // Kiro/CodeWhisperer OAuth - Gemini, // Google Gemini OAuth - Qwen, // 通义千问 OAuth - Antigravity, // Antigravity (Gemini CLI) OAuth - ClaudeOAuth, // Claude OAuth ClaudeCustom, // Claude API Key OpenAICustom, // OpenAI API Key - Codex, // Codex - IFlow, // iFlow + GeminiApiKey, // Gemini API Key + Codex, // OpenAI Responses / Codex 兼容 API Key Vertex, // Vertex AI } ``` @@ -67,40 +58,11 @@ pub trait Provider: Send + Sync { } ``` -## OAuth Provider 实现 +## 已退役 Provider -### Kiro Provider +Kiro / Qwen / Antigravity OAuth / Codex OAuth / Claude OAuth / Gemini OAuth 都属于旧凭证池功能,分类为 `dead`。不得重新接回设置页、Tauri 命令、Token 刷新任务或运行时 fallback。 -```rust -// 凭证文件结构 -struct KiroCredential { - access_token: String, - refresh_token: String, - expires_at: i64, - client_id: Option, // 从 clientIdHash 合并 - client_secret: Option, // 从 clientIdHash 合并 -} - -// Token 刷新流程 -1. 检查 expires_at 是否过期 -2. 使用 refresh_token 请求新 token -3. 更新凭证文件 -``` - -### Gemini Provider - -```rust -// OAuth 端点 -const AUTH_URL: &str = "https://accounts.google.com/o/oauth2/v2/auth"; -const TOKEN_URL: &str = "https://oauth2.googleapis.com/token"; - -// 凭证文件结构 -struct GeminiCredential { - access_token: String, - refresh_token: String, - expires_at: i64, -} -``` +Antigravity 的协议转换能力不是凭证池功能,`openai_to_antigravity` converter 仍可用于 coding plan。 ## API Key Provider 实现 @@ -163,64 +125,16 @@ Lime 当前把 Prompt Cache 能力视为 **Provider 显式声明优先、类型 2. 上游服务是否真的声明支持 Anthropic Automatic Prompt Caching 3. 响应 usage 中是否存在 `cache_creation_input_tokens` / `cache_read_input_tokens` / `cached_input_tokens` -## 凭证管理策略 - -### 方案 B: 独立副本策略 - -``` -原始凭证文件 (用户上传) - │ - ▼ -┌─────────────────────────────────────┐ -│ 合并 clientIdHash 中的 │ -│ client_id / client_secret │ -└─────────────────────────────────────┘ - │ - ▼ -副本凭证文件 (credentials/ 目录) - │ - ▼ -独立刷新和管理 -``` - -优点: -- 每个副本完全独立 -- 支持多账号场景 -- 不影响原始文件 - -## 健康检查 - -```rust -// 健康检查逻辑 -async fn health_check(&self, credential: &CredentialData) -> HealthStatus { - // 1. 检查 Token 是否过期 - if self.is_token_expired(credential) { - return HealthStatus::TokenExpired; - } - - // 2. 尝试刷新 Token - if let Err(e) = self.refresh_token(credential).await { - return HealthStatus::RefreshFailed(e); - } - - // 3. 发送测试请求 - match self.send_test_request(credential).await { - Ok(_) => HealthStatus::Healthy, - Err(e) => HealthStatus::Unhealthy(e), - } -} -``` - ## 添加新 Provider 1. 在 `providers/` 创建新模块文件 2. 实现 `Provider` trait 3. 在 `ProviderType` 枚举添加新类型 -4. 在 `ProviderPoolService` 注册健康检查 -5. 更新前端 Provider 选择器 +4. 同步 API Key Provider schema、模型注册表与连接测试 +5. 更新前端 Provider 选择器与文档 ## 相关文档 -- [credential-pool.md](credential-pool.md) - 凭证池管理 +- [credential-pool.md](credential-pool.md) - 凭证池退役说明 - [converter.md](converter.md) - 协议转换 - [server.md](server.md) - HTTP 服务器 diff --git a/docs/aiprompts/services.md b/docs/aiprompts/services.md index b2d33045f..23933c0ae 100644 --- a/docs/aiprompts/services.md +++ b/docs/aiprompts/services.md @@ -9,8 +9,8 @@ ``` src-tauri/src/services/ ├── mod.rs # 模块入口 -├── provider_pool_service.rs # 凭证池服务 -├── token_cache_service.rs # Token 缓存 +├── api_key_provider_service.rs # API Key Provider 服务 +├── model_registry_service.rs # 模型注册表服务 ├── mcp_service.rs # MCP 服务器管理 ├── prompt_service.rs # Prompt 管理 ├── skill_service.rs # 技能管理 @@ -23,45 +23,14 @@ src-tauri/src/services/ > 注意:`general_chat/` 兼容壳已删除。 > 新功能与新治理都应直接落到 `agent_runtime_*` 与现役 `agent/chat` 体系,不要重新引回旧入口。 -> `ProviderPoolService::select_credential_with_fallback_legacy` 也已删除,凭证选择统一走现役 `select_credential_with_fallback`。 +> 旧 `ProviderPoolService`、`TokenCacheService` 与 OAuth/local CLI credential runtime 已退役。凭证选择统一走 `ApiKeyProviderService`。 -### ProviderPoolService +### ApiKeyProviderService ```rust -pub struct ProviderPoolService { - pools: HashMap, - health_checker: HealthChecker, -} - -impl ProviderPoolService { - /// 获取下一个可用凭证 - pub async fn next_credential(&self, provider: ProviderType) -> Option; - - /// 添加凭证到池 - pub async fn add_credential(&self, credential: Credential) -> Result<()>; - - /// 移除凭证 - pub async fn remove_credential(&self, id: &str) -> Result<()>; - - /// 启动健康检查 - pub fn start_health_check(&self); -} -``` - -### TokenCacheService - -```rust -pub struct TokenCacheService { - cache: DashMap, - db: Arc, -} - -impl TokenCacheService { - /// 获取或刷新 Token - pub async fn get_or_refresh(&self, credential_id: &str) -> Result; - - /// 使 Token 失效 - pub async fn invalidate(&self, credential_id: &str); +impl ApiKeyProviderService { + /// 选择当前 Provider 可用的 API Key 配置 + pub async fn select_credential_for_provider(&self, provider_id: &str) -> Result; } ``` @@ -87,24 +56,10 @@ impl McpService { ## 服务注入 ```rust -// 在 main.rs 中初始化 -let pool_service = Arc::new(ProviderPoolService::new()); -let token_cache = Arc::new(TokenCacheService::new(db.clone())); - -app.manage(pool_service); -app.manage(token_cache); - -// 在命令中使用 -#[tauri::command] -async fn add_credential( - pool: State<'_, Arc>, - // ... -) -> Result<(), String> { - pool.add_credential(credential).await -} +// 服务由 bootstrap / state 注入,命令层不再管理凭证池服务。 ``` ## 相关文档 - [commands.md](commands.md) - Tauri 命令 -- [credential-pool.md](credential-pool.md) - 凭证池管理 +- [credential-pool.md](credential-pool.md) - 凭证池退役说明 diff --git a/docs/content/03.providers/1.overview.md b/docs/content/03.providers/1.overview.md index e8532cfab..d999a1f81 100644 --- a/docs/content/03.providers/1.overview.md +++ b/docs/content/03.providers/1.overview.md @@ -22,17 +22,17 @@ Provider 文档正在按 LimeNext V2 current 连接入口重建。 先补可用 API Key,再读取真实模型目录、做连接验证,并校准默认模型与兼容协议 - `中转服务` 对应 Lime Connect。 - 浏览已验证中转商,获取 API Key 后通过 `lime://connect` 深链一键带入凭证池 + 浏览已验证中转商,获取 API Key 后通过 `lime://connect` 深链接入 API Key Provider - `语音服务` 单独管理语音相关 Provider -- `OAuth 凭证` - 管理 Kiro、Gemini、Antigravity、Codex、Claude OAuth 等登录型凭证 +- `协议转换` + 保留请求协议适配能力,例如 coding plan 可继续使用 Antigravity 协议转换器 ## 当前建议 1. 默认先从 `服务商` 分类补可用 API Key -2. 只有需要登录型凭证时,再进入 `OAuth 凭证` -3. 需要浏览中转商时,再进入 `中转服务` +2. 需要浏览中转商时,再进入 `中转服务` +3. 不再配置 Kiro、Gemini OAuth、Antigravity OAuth、Codex OAuth、Claude OAuth 等登录型凭证 4. 不再以旧 YAML 配置文档作为日常配置入口 ## 后续说明 diff --git a/docs/exec-plans/README.md b/docs/exec-plans/README.md index fa3d06f7a..5ec2059ca 100644 --- a/docs/exec-plans/README.md +++ b/docs/exec-plans/README.md @@ -27,6 +27,7 @@ - 参考运行时主链总计划:`docs/exec-plans/upstream-runtime-alignment-plan.md` - 参考运行时主链进度日志:`docs/exec-plans/upstream-runtime-alignment-progress.md` - Provider 模型能力 taxonomy 进度日志:`docs/exec-plans/provider-model-taxonomy-progress.md` +- 云端套餐与支付边界收口计划:`docs/exec-plans/cloud-commerce-user-center-boundary.md` - `@` 命令本地执行纠偏计划:`docs/exec-plans/at-command-local-execution-alignment-plan.md` - LimeNext 总实施计划(`legacy current reference`,当前主规划已切到 `docs/roadmap/limenextv2/README.md`):`docs/exec-plans/limenext-plan.md` - LimeNext 推进日志:`docs/exec-plans/limenext-progress.md` diff --git a/docs/exec-plans/cloud-commerce-user-center-boundary.md b/docs/exec-plans/cloud-commerce-user-center-boundary.md new file mode 100644 index 000000000..589bc65aa --- /dev/null +++ b/docs/exec-plans/cloud-commerce-user-center-boundary.md @@ -0,0 +1,25 @@ +# 云端套餐与支付边界收口执行计划 + +## 主目标 + +把套餐购买、支付、账单、用量明细统一收敛到 `limecore` 用户中心网页;Lime 客户端只保留会话状态、当前套餐、积分余额、待支付提醒和跳转入口。 + +## 事实源 + +- `current`:`limecore` control-plane 与 `apps/user-center-web` 的 `/pricing`、`/billing`、`/subscription`、`/credits` +- `current`:Lime 客户端 `useOemCloudAccess` 负责登录态、权益摘要、API Key、支付回跳同步 +- `dead`:Lime 客户端内置套餐卡、充值包卡、用量图、账单表、直接创建购买订单的本地商业工作台 + +## 本轮进度 + +- 已将 Lime 设置页云端服务面收口为摘要卡 + 用户中心入口。 +- 已从客户端设置页移除本地套餐购买、积分充值、用量明细和账单表渲染。 +- 已从 `useOemCloudAccess` 返回面移除客户端直接创建套餐/充值订单的处理器。 +- 已在 `limecore` 用户中心 `/billing` 补齐用量与账单 tab,使用真实 usage / credits / billing dashboard 和真实 checkout。 +- 已将 `limecore` 用户中心 `/pricing` 文案从“模型价格”收敛为“套餐与价格”。 +- 已让 Lime 客户端入口直达 `/billing?tab=usage` 与 `/billing?tab=billing`,并让用户中心 tab 与 URL 查询参数同步。 +- 已从用户中心主导航移除独立“账单管理”菜单,保留 `/subscription` 直接路由作为旧链接可达页面,主导航收敛到 `/billing`。 + +## 下一刀 + +继续验证两个仓库的定向类型检查与回归测试;如果后续发现客户端仍存在套餐购买 UI 文案或旧 testid,应继续按 `dead` 删除,不再迁回客户端。 diff --git a/docs/ops.md b/docs/ops.md index 06963ed98..50d5b0cdf 100644 --- a/docs/ops.md +++ b/docs/ops.md @@ -25,7 +25,7 @@ ## 数据与日志位置 - SQLite 数据库:`~/.lime/lime.db` -- 凭证池副本目录:macOS `~/Library/Application Support/lime/credentials/`,Linux `~/.local/share/lime/credentials/`,Windows `%APPDATA%\\lime\\credentials\\` +- 旧凭证池副本目录:macOS `~/Library/Application Support/lime/credentials/`,Linux `~/.local/share/lime/credentials/`,Windows `%APPDATA%\\lime\\credentials\\`;启动期会清理 Lime 管理的副本 - OAuth/Token 目录(默认):`~/.lime/auth/` - 日志目录:`~/.lime/logs/` - 请求日志目录:`~/.lime/request_logs/` @@ -41,7 +41,7 @@ - 复制以下路径: - 配置文件(macOS: `~/Library/Application Support/lime/config.yaml`,Linux: `~/.config/lime/config.yaml`,Windows: `%APPDATA%\\lime\\config.yaml`) - 配置备份文件:`config.yaml.backup` - - 凭证池副本目录(macOS: `~/Library/Application Support/lime/credentials/`,Linux: `~/.local/share/lime/credentials/`,Windows: `%APPDATA%\\lime\\credentials\\`) + - 旧凭证池副本目录(macOS: `~/Library/Application Support/lime/credentials/`,Linux: `~/.local/share/lime/credentials/`,Windows: `%APPDATA%\\lime\\credentials\\`) - `~/.lime/lime.db` - `~/.lime/auth/`(如需要保留 OAuth/Token) - `~/.lime/logs/`、`~/.lime/request_logs/`(如需保留日志) diff --git a/homebrew/Casks/lime.rb b/homebrew/Casks/lime.rb index a6b302271..13dad443b 100644 --- a/homebrew/Casks/lime.rb +++ b/homebrew/Casks/lime.rb @@ -12,7 +12,7 @@ cask "lime" do end name "Lime" - desc "AI 代理服务桌面应用 - 多 Provider 凭证池管理" + desc "AI 代理服务桌面应用 - 多 Provider API Key 管理" homepage "https://github.com/aiclientproxy/lime" livecheck do diff --git a/package-lock.json b/package-lock.json index ac9a8356e..2c0f4ae2d 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,12 +1,12 @@ { "name": "lime", - "version": "1.21.0", + "version": "1.22.0", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "lime", - "version": "1.21.0", + "version": "1.22.0", "dependencies": { "@babel/standalone": "^7.29.0", "@fabianlars/tauri-plugin-oauth": "^2", diff --git a/package.json b/package.json index 87d283c79..30ba6c241 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "lime", "private": true, - "version": "1.21.0", + "version": "1.22.0", "type": "module", "engines": { "node": ">=22.0.0" diff --git a/packages/lime-cli-npm/README.md b/packages/lime-cli-npm/README.md index 13209c7f4..54057c5c2 100644 --- a/packages/lime-cli-npm/README.md +++ b/packages/lime-cli-npm/README.md @@ -112,7 +112,7 @@ npm run build:release -- \ ```bash npm run build:release -- \ --target-triple "aarch64-apple-darwin" \ - --version "1.21.0" \ + --version "1.22.0" \ --out-dir "./dist" ``` diff --git a/packages/lime-cli-npm/package.json b/packages/lime-cli-npm/package.json index d0141c84e..5b44bc88f 100644 --- a/packages/lime-cli-npm/package.json +++ b/packages/lime-cli-npm/package.json @@ -1,6 +1,6 @@ { "name": "@limecloud/lime-cli", - "version": "1.21.0", + "version": "1.22.0", "description": "Lime 官方任务 CLI", "bin": { "lime": "scripts/run.js" diff --git a/scripts/prepare-github-release-assets.mjs b/scripts/prepare-github-release-assets.mjs new file mode 100644 index 000000000..fb24e783c --- /dev/null +++ b/scripts/prepare-github-release-assets.mjs @@ -0,0 +1,213 @@ +#!/usr/bin/env node + +import fs from "node:fs"; +import path from "node:path"; +import process from "node:process"; +import { fileURLToPath } from "node:url"; + +function parseArgs(argv) { + const args = { + extraAsset: [], + }; + + for (let index = 0; index < argv.length; index += 1) { + const item = argv[index]; + if (!item.startsWith("--")) { + continue; + } + + const key = item.slice(2); + const next = argv[index + 1]; + const value = !next || next.startsWith("--") ? "true" : next; + if (value !== "true") { + index += 1; + } + + if (key === "extra-asset") { + args.extraAsset.push(value); + } else { + args[key] = value; + } + } + + return args; +} + +function normalizeVersion(value) { + return String(value || "") + .trim() + .replace(/^v/, ""); +} + +function listFilesRecursive(root) { + if (!fs.existsSync(root)) { + return []; + } + + return fs + .readdirSync(root, { withFileTypes: true }) + .flatMap((entry) => { + const filePath = path.join(root, entry.name); + if (entry.isDirectory()) { + return listFilesRecursive(filePath); + } + return entry.isFile() ? [filePath] : []; + }) + .sort(); +} + +function isPathInside(parent, child) { + const relative = path.relative(parent, child); + return ( + Boolean(relative) && + !relative.startsWith("..") && + !path.isAbsolute(relative) + ); +} + +function targetFromAssetPath(assetsDir, filePath) { + if (!isPathInside(assetsDir, filePath)) { + return ""; + } + + return path.relative(assetsDir, filePath).split(path.sep)[0] || ""; +} + +function duplicateTargetLabel(target) { + if (target === "aarch64-apple-darwin") { + return "macos-arm64"; + } + if (target === "x86_64-apple-darwin") { + return "macos-x64"; + } + if (target === "x86_64-pc-windows-msvc") { + return "windows-x64"; + } + if (target === "x86_64-unknown-linux-gnu") { + return "linux-x64"; + } + return target.replace(/[^A-Za-z0-9._-]+/g, "-") || "asset"; +} + +function macUpdaterAssetName(basename, target, version) { + let arch = ""; + if (target === "aarch64-apple-darwin") { + arch = "aarch64"; + } else if (target === "x86_64-apple-darwin") { + arch = "x64"; + } + + if (!arch) { + return ""; + } + if (basename === "Lime.app.tar.gz") { + return `Lime_${version}_${arch}.app.tar.gz`; + } + if (basename === "Lime.app.tar.gz.sig") { + return `Lime_${version}_${arch}.app.tar.gz.sig`; + } + return ""; +} + +function githubAssetName(filePath, context) { + const basename = path.basename(filePath); + const target = targetFromAssetPath(context.assetsDir, filePath); + const duplicateCount = context.basenameCounts.get(basename) || 0; + + if (duplicateCount <= 1) { + return basename; + } + + const macName = macUpdaterAssetName(basename, target, context.version); + if (macName) { + return macName; + } + + return `${duplicateTargetLabel(target)}-${basename}`; +} + +function prepareGitHubReleaseAssets(options) { + const assetsDir = path.resolve(options.assetsDir || "release-assets"); + const outDir = path.resolve(options.outDir || "release-github-assets"); + const version = normalizeVersion(options.version); + const extraAssets = (options.extraAssets || []).map((item) => + path.resolve(item), + ); + + if (!version) { + throw new Error("version is required"); + } + if (!fs.existsSync(assetsDir)) { + throw new Error(`assets directory is missing: ${assetsDir}`); + } + + const releaseAssetFiles = listFilesRecursive(assetsDir).filter( + (filePath) => !/^latest.*\.json$/i.test(path.basename(filePath)), + ); + const inputFiles = [...releaseAssetFiles, ...extraAssets].sort(); + + for (const filePath of inputFiles) { + if (!fs.existsSync(filePath) || !fs.statSync(filePath).isFile()) { + throw new Error(`release asset is missing: ${filePath}`); + } + } + + const basenameCounts = new Map(); + for (const filePath of inputFiles) { + const basename = path.basename(filePath); + basenameCounts.set(basename, (basenameCounts.get(basename) || 0) + 1); + } + + fs.rmSync(outDir, { recursive: true, force: true }); + fs.mkdirSync(outDir, { recursive: true }); + + const usedNames = new Set(); + const copied = []; + for (const filePath of inputFiles) { + const name = githubAssetName(filePath, { + assetsDir, + basenameCounts, + version, + }); + if (usedNames.has(name)) { + throw new Error(`duplicate GitHub release asset name: ${name}`); + } + usedNames.add(name); + + const destination = path.join(outDir, name); + fs.copyFileSync(filePath, destination); + copied.push({ + name, + source: filePath, + destination, + }); + } + + return copied.sort((left, right) => left.name.localeCompare(right.name)); +} + +function main() { + const args = parseArgs(process.argv.slice(2)); + const copied = prepareGitHubReleaseAssets({ + assetsDir: args["assets-dir"], + extraAssets: args.extraAsset, + outDir: args["out-dir"], + version: + args.version || process.env.RELEASE_TAG || process.env.GITHUB_REF_NAME, + }); + + console.log("Prepared GitHub release upload assets:"); + for (const item of copied) { + console.log(` - ${path.relative(process.cwd(), item.destination)}`); + } +} + +const isCli = + process.argv[1] && + path.resolve(process.argv[1]) === fileURLToPath(import.meta.url); + +if (isCli) { + main(); +} + +export { prepareGitHubReleaseAssets }; diff --git a/scripts/release-updater-manifest.test.mjs b/scripts/release-updater-manifest.test.mjs index a14dbd67b..40d0b2d5f 100644 --- a/scripts/release-updater-manifest.test.mjs +++ b/scripts/release-updater-manifest.test.mjs @@ -4,6 +4,7 @@ import path from "node:path"; import { describe, expect, it } from "vitest"; import { planR2ReleaseCleanup } from "./plan-r2-release-cleanup.mjs"; +import { prepareGitHubReleaseAssets } from "./prepare-github-release-assets.mjs"; import { collectUpdaterManifest, writeOutputs, @@ -215,9 +216,7 @@ describe("release updater manifest", () => { path.join(assetsDir, "aarch64-apple-darwin", "Lime.app.tar.gz.sig"), encodedSignature, ); - writeFile( - path.join(assetsDir, "x86_64-apple-darwin", "Lime.app.tar.gz"), - ); + writeFile(path.join(assetsDir, "x86_64-apple-darwin", "Lime.app.tar.gz")); writeFile( path.join(assetsDir, "x86_64-apple-darwin", "Lime.app.tar.gz.sig"), rawSignature, @@ -243,11 +242,7 @@ describe("release updater manifest", () => { baseUrl: "https://updates.limecloud.com", channel: "stable", notes: "notes", - requiredPlatforms: [ - "darwin-aarch64", - "darwin-x86_64", - "windows-x86_64", - ], + requiredPlatforms: ["darwin-aarch64", "darwin-x86_64", "windows-x86_64"], version: "v1.20.0", }); @@ -292,3 +287,63 @@ describe("R2 release cleanup", () => { expect(plan.protectedVersions).toContain("1.16.0"); }); }); + +describe("GitHub release asset staging", () => { + it("同名 macOS updater 资产上传 GitHub Release 前应重命名", () => { + const root = fs.mkdtempSync( + path.join(os.tmpdir(), "lime-github-release-assets-"), + ); + const assetsDir = path.join(root, "release-assets"); + const outDir = path.join(root, "release-github-assets"); + const latestPath = path.join(root, "release-updater", "latest.json"); + + writeFile(path.join(assetsDir, "aarch64-apple-darwin", "Lime.app.tar.gz")); + writeFile( + path.join(assetsDir, "aarch64-apple-darwin", "Lime.app.tar.gz.sig"), + "arm-sig", + ); + writeFile( + path.join(assetsDir, "aarch64-apple-darwin", "Lime_1.22.0_aarch64.dmg"), + ); + writeFile(path.join(assetsDir, "x86_64-apple-darwin", "Lime.app.tar.gz")); + writeFile( + path.join(assetsDir, "x86_64-apple-darwin", "Lime.app.tar.gz.sig"), + "x64-sig", + ); + writeFile( + path.join(assetsDir, "x86_64-apple-darwin", "Lime_1.22.0_x64.dmg"), + ); + writeFile(latestPath, "{}"); + + const copied = prepareGitHubReleaseAssets({ + assetsDir, + extraAssets: [latestPath], + outDir, + version: "v1.22.0", + }); + + expect(copied.map((item) => item.name).sort()).toEqual( + [ + "Lime_1.22.0_aarch64.app.tar.gz", + "Lime_1.22.0_aarch64.app.tar.gz.sig", + "Lime_1.22.0_aarch64.dmg", + "Lime_1.22.0_x64.app.tar.gz", + "Lime_1.22.0_x64.app.tar.gz.sig", + "Lime_1.22.0_x64.dmg", + "latest.json", + ].sort(), + ); + expect( + fs.readFileSync( + path.join(outDir, "Lime_1.22.0_aarch64.app.tar.gz.sig"), + "utf8", + ), + ).toBe("arm-sig"); + expect( + fs.readFileSync( + path.join(outDir, "Lime_1.22.0_x64.app.tar.gz.sig"), + "utf8", + ), + ).toBe("x64-sig"); + }); +}); diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 8d012b36f..f2402f870 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -165,16 +165,6 @@ version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" -[[package]] -name = "aead" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" -dependencies = [ - "crypto-common", - "generic-array", -] - [[package]] name = "aes" version = "0.8.4" @@ -1862,30 +1852,6 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" -[[package]] -name = "chacha20" -version = "0.9.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818" -dependencies = [ - "cfg-if", - "cipher", - "cpufeatures", -] - -[[package]] -name = "chacha20poly1305" -version = "0.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "10cd79432192d1c0f4e1a0fef9527696cc039165d729fb41b3f4f4f354c2dc35" -dependencies = [ - "aead", - "chacha20", - "cipher", - "poly1305", - "zeroize", -] - [[package]] name = "chrono" version = "0.4.43" @@ -1918,7 +1884,6 @@ checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ "crypto-common", "inout", - "zeroize", ] [[package]] @@ -2403,7 +2368,6 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" dependencies = [ "generic-array", - "rand_core 0.6.4", "typenum", ] @@ -5101,7 +5065,7 @@ dependencies = [ [[package]] name = "lime" -version = "1.21.0" +version = "1.22.0" dependencies = [ "anyhow", "arboard", @@ -5132,7 +5096,6 @@ dependencies = [ "lime-browser-runtime", "lime-config", "lime-core", - "lime-credential", "lime-embedding", "lime-gateway", "lime-infra", @@ -5205,7 +5168,7 @@ dependencies = [ [[package]] name = "lime-agent" -version = "1.21.0" +version = "1.22.0" dependencies = [ "anyhow", "aster-core", @@ -5234,7 +5197,7 @@ dependencies = [ [[package]] name = "lime-browser-runtime" -version = "1.21.0" +version = "1.22.0" dependencies = [ "chrono", "futures", @@ -5251,7 +5214,7 @@ dependencies = [ [[package]] name = "lime-cli" -version = "1.21.0" +version = "1.22.0" dependencies = [ "clap", "lime-core", @@ -5263,7 +5226,7 @@ dependencies = [ [[package]] name = "lime-config" -version = "1.21.0" +version = "1.22.0" dependencies = [ "async-trait", "lime-core", @@ -5279,7 +5242,7 @@ dependencies = [ [[package]] name = "lime-core" -version = "1.21.0" +version = "1.22.0" dependencies = [ "aster-models", "async-trait", @@ -5317,28 +5280,6 @@ dependencies = [ "zip 0.6.6", ] -[[package]] -name = "lime-credential" -version = "1.21.0" -dependencies = [ - "axum 0.7.9", - "base64 0.22.1", - "chacha20poly1305", - "chrono", - "dashmap 5.5.3", - "lime-core", - "lime-infra", - "proptest", - "rand 0.8.5", - "reqwest 0.12.28", - "serde", - "serde_json", - "sha2", - "tempfile", - "tokio", - "tracing", -] - [[package]] name = "lime-embedding" version = "0.1.0" @@ -5354,7 +5295,7 @@ dependencies = [ [[package]] name = "lime-gateway" -version = "1.21.0" +version = "1.22.0" dependencies = [ "aes", "axum 0.7.9", @@ -5384,7 +5325,7 @@ dependencies = [ [[package]] name = "lime-infra" -version = "1.21.0" +version = "1.22.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -5404,7 +5345,7 @@ dependencies = [ [[package]] name = "lime-mcp" -version = "1.21.0" +version = "1.22.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -5420,7 +5361,7 @@ dependencies = [ [[package]] name = "lime-media-runtime" -version = "1.21.0" +version = "1.22.0" dependencies = [ "axum 0.7.9", "chrono", @@ -5451,7 +5392,7 @@ dependencies = [ [[package]] name = "lime-processor" -version = "1.21.0" +version = "1.22.0" dependencies = [ "async-trait", "lime-core", @@ -5470,7 +5411,7 @@ dependencies = [ [[package]] name = "lime-providers" -version = "1.21.0" +version = "1.22.0" dependencies = [ "anyhow", "async-stream", @@ -5525,7 +5466,7 @@ dependencies = [ [[package]] name = "lime-server" -version = "1.21.0" +version = "1.22.0" dependencies = [ "aster-core", "async-stream", @@ -5540,7 +5481,6 @@ dependencies = [ "lime-agent", "lime-config", "lime-core", - "lime-credential", "lime-infra", "lime-processor", "lime-providers", @@ -5570,7 +5510,7 @@ dependencies = [ [[package]] name = "lime-server-utils" -version = "1.21.0" +version = "1.22.0" dependencies = [ "axum 0.7.9", "futures", @@ -5585,7 +5525,7 @@ dependencies = [ [[package]] name = "lime-services" -version = "1.21.0" +version = "1.22.0" dependencies = [ "anyhow", "aster-core", @@ -5628,7 +5568,7 @@ dependencies = [ [[package]] name = "lime-skills" -version = "1.21.0" +version = "1.22.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -5646,7 +5586,7 @@ dependencies = [ [[package]] name = "lime-websocket" -version = "1.21.0" +version = "1.22.0" dependencies = [ "axum 0.7.9", "chrono", @@ -6715,12 +6655,6 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" -[[package]] -name = "opaque-debug" -version = "0.3.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" - [[package]] name = "open" version = "5.3.3" @@ -7399,17 +7333,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "poly1305" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf" -dependencies = [ - "cpufeatures", - "opaque-debug", - "universal-hash", -] - [[package]] name = "portable-atomic" version = "1.13.1" @@ -11016,16 +10939,6 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" -[[package]] -name = "universal-hash" -version = "0.5.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" -dependencies = [ - "crypto-common", - "subtle", -] - [[package]] name = "unsafe-libyaml" version = "0.2.11" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index c4cdfab94..42c87fa89 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,10 +1,15 @@ [workspace] members = ["crates/*"] -exclude = ["crates/aster", "crates/aster-models", "crates/aster-rust"] +exclude = [ + "crates/aster", + "crates/aster-models", + "crates/aster-rust", + "crates/credential", +] resolver = "2" [workspace.package] -version = "1.21.0" +version = "1.22.0" edition = "2021" authors = ["coso"] repository = "https://github.com/aiclientproxy/lime" @@ -17,7 +22,6 @@ lime-config = { path = "crates/config" } lime-infra = { path = "crates/infra" } lime-providers = { path = "crates/providers" } lime-services = { path = "crates/services" } -lime-credential = { path = "crates/credential" } lime-websocket = { path = "crates/websocket" } lime-processor = { path = "crates/processor" } lime-server-utils = { path = "crates/server-utils" } @@ -189,7 +193,7 @@ version = "2.4" [package] name = "lime" -version = "1.21.0" +version = "1.22.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" @@ -211,7 +215,6 @@ lime-config.workspace = true lime-infra.workspace = true lime-providers.workspace = true lime-services.workspace = true -lime-credential.workspace = true lime-websocket.workspace = true lime-processor.workspace = true lime-server-utils.workspace = true diff --git a/src-tauri/crates/agent/src/aster_state.rs b/src-tauri/crates/agent/src/aster_state.rs index e7e7aeb4a..5ac4effa4 100644 --- a/src-tauri/crates/agent/src/aster_state.rs +++ b/src-tauri/crates/agent/src/aster_state.rs @@ -2,7 +2,7 @@ //! //! 管理 Aster Agent 实例和相关状态 //! 提供 Tauri 应用与 Aster 框架的桥接 -//! 支持从 Lime 凭证池自动选择凭证 +//! 支持从 Lime API Key Provider 自动选择凭证 //! //! ## 重要:SessionStore 注入 //! @@ -116,7 +116,7 @@ pub struct RuntimeInterruptMarker { pub struct ProviderConfig { /// Provider 名称 (openai, anthropic, google, ollama 等) pub provider_name: String, - /// Provider 选择器(优先保留前端 provider_id / pool provider_type) + /// Provider 选择器(优先保留前端 provider_id / API Key Provider 类型) pub provider_selector: Option, /// 模型名称 pub model_name: String, @@ -124,12 +124,10 @@ pub struct ProviderConfig { pub api_key: Option, /// Base URL (可选,用于自定义端点) pub base_url: Option, - /// 凭证 UUID(来自凭证池,用于记录使用和健康状态) + /// 凭证 UUID(来自 API Key Provider,用于记录使用和健康状态) pub credential_uuid: Option, /// 是否强制 OpenAI provider 使用 Responses API pub force_responses_api: bool, - /// OAuth/本地 Provider 需要的凭证文件路径 - pub credential_path: Option, /// 当前回合是否需要用 toolshim 兼容无原生 tools 的模型 pub toolshim: bool, /// toolshim 解释器模型(可与实际回复模型不同) @@ -306,7 +304,6 @@ impl AsterAgentState { .clone() .unwrap_or_else(|| format!("manual:{session_id}")), force_responses_api: config.force_responses_api, - credential_path: config.credential_path.clone(), toolshim: config.toolshim, toolshim_model: config.toolshim_model.clone(), }) @@ -339,13 +336,13 @@ impl AsterAgentState { Ok(()) } - /// 从凭证池配置 Provider + /// 从 API Key Provider 配置 Provider /// - /// 自动从 Lime 凭证池选择可用凭证并配置 Aster Provider + /// 自动从 Lime API Key Provider 选择可用凭证并配置 Aster Provider /// /// # 参数 /// - `db`: 数据库连接 - /// - `provider_type`: Provider 类型 (openai, anthropic, kiro 等) + /// - `provider_type`: Provider 类型 (openai, anthropic, google 等) /// - `model`: 模型名称 /// - `session_id`: 会话 ID pub async fn configure_provider_from_pool( @@ -358,12 +355,12 @@ impl AsterAgentState { // 确保 Agent 已初始化(使用带数据库的版本) self.init_agent_with_db(db).await?; - // 从凭证池选择凭证并获取配置 + // 从 API Key Provider 选择凭证并获取配置 let aster_config = self .credential_bridge .select_and_configure(db, provider_type, model) .await - .map_err(|e| format!("从凭证池选择凭证失败: {e}"))?; + .map_err(|e| format!("从 API Key Provider 选择凭证失败: {e}"))?; // 创建 Provider let provider = create_aster_provider(&aster_config) @@ -388,7 +385,6 @@ impl AsterAgentState { base_url: aster_config.base_url.clone(), credential_uuid: Some(aster_config.credential_uuid.clone()), force_responses_api: aster_config.force_responses_api, - credential_path: aster_config.credential_path.clone(), toolshim: aster_config.toolshim, toolshim_model: aster_config.toolshim_model.clone(), }; @@ -408,7 +404,7 @@ impl AsterAgentState { } tracing::info!( - "[AsterAgent] 从凭证池配置 Provider 成功: {} / {} (凭证: {})", + "[AsterAgent] 从 API Key Provider 配置 Provider 成功: {} / {} (凭证: {})", aster_config.provider_name, aster_config.model_name, aster_config.credential_uuid @@ -450,7 +446,7 @@ impl AsterAgentState { /// 清除当前 Provider 配置 /// - /// 用于切换凭证后重置状态,下次对话时会重新从凭证池选择凭证 + /// 用于切换凭证后重置状态,下次对话时会重新从 API Key Provider 选择凭证 pub async fn clear_provider_config(&self) { let mut config_guard = self.current_provider_config.write().await; *config_guard = None; @@ -853,7 +849,6 @@ mod tests { base_url: None, credential_uuid: None, force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -878,7 +873,6 @@ mod tests { base_url: None, credential_uuid: None, force_responses_api: true, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -894,7 +888,7 @@ mod tests { } #[test] - fn test_provider_config_detects_kiro_provider_session_token_capability() { + fn test_provider_config_treats_retired_kiro_as_history_replay_only() { let config = ProviderConfig { provider_name: "kiro".to_string(), provider_selector: Some("kiro".to_string()), @@ -903,14 +897,13 @@ mod tests { base_url: None, credential_uuid: None, force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; assert_eq!( config.provider_continuation_capability(), - ProviderContinuationCapability::ProviderSessionToken + ProviderContinuationCapability::HistoryReplayOnly ); assert_eq!( config.provider_continuation_state(), diff --git a/src-tauri/crates/agent/src/credential_bridge.rs b/src-tauri/crates/agent/src/credential_bridge.rs index 263de8c0b..9275ac211 100644 --- a/src-tauri/crates/agent/src/credential_bridge.rs +++ b/src-tauri/crates/agent/src/credential_bridge.rs @@ -1,28 +1,21 @@ -//! 凭证池桥接模块 +//! API Key Provider 桥接模块 //! -//! 将 Lime 凭证池与 Aster Provider 系统连接 -//! 支持从凭证池自动选择凭证并配置 Aster Provider +//! 将 Lime API Key Provider 主路径与 Aster Provider 系统连接。 //! //! ## 功能 -//! - 从凭证池选择可用凭证 +//! - 从 API Key Provider 选择可用凭证 //! - 将凭证转换为 Aster Provider 配置 -//! - 支持 OAuth 和 API Key 两种凭证类型 -//! - 自动刷新过期的 OAuth Token //! - 智能拆分 base_url 为 host + path,避免路径重复(如智谱 /v4/v1 问题) use aster::model::ModelConfig; use aster::providers::base::Provider; use lime_core::database::dao::api_key_provider::{infer_managed_runtime_spec, ApiProviderType}; use lime_core::database::DbConnection; -use lime_core::models::provider_pool_model::{ - CredentialData, PoolProviderType, ProviderCredential, -}; +use lime_core::models::provider_pool_model::{CredentialData, ProviderCredential}; use lime_core::models::provider_type::is_custom_provider_id; use lime_services::api_key_provider_service::ApiKeyProviderService; -use lime_services::provider_pool_service::ProviderPoolService; use std::sync::Arc; -use crate::kiro_provider_adapter::LimeKiroProvider; use crate::provider_safety::wrap_provider_with_safety; /// 凭证桥接错误 @@ -71,19 +64,16 @@ pub struct AsterProviderConfig { pub credential_uuid: String, /// 是否强制 OpenAI provider 使用 Responses API(用于 Codex 等兼容链路) pub force_responses_api: bool, - /// OAuth/本地 Provider 需要的凭证文件路径 - pub credential_path: Option, /// 当前回合是否启用 toolshim pub toolshim: bool, /// toolshim 解释器模型 pub toolshim_model: Option, } -/// 凭证池桥接器 +/// Provider 凭证桥接器 /// -/// 负责从 Lime 凭证池选择凭证并转换为 Aster Provider 配置 +/// 负责从 Lime API Key Provider 选择凭证并转换为 Aster Provider 配置 pub struct CredentialBridge { - pool_service: ProviderPoolService, api_key_service: ApiKeyProviderService, } @@ -96,7 +86,6 @@ impl Default for CredentialBridge { impl CredentialBridge { pub fn new() -> Self { Self { - pool_service: ProviderPoolService::new(), api_key_service: ApiKeyProviderService::new(), } } @@ -106,7 +95,7 @@ impl CredentialBridge { .filter(|value| !value.is_empty()) } - /// 从凭证池选择凭证并创建 Aster Provider 配置 + /// 从 API Key Provider 选择凭证并创建 Aster Provider 配置 /// /// # 参数 /// - `db`: 数据库连接 @@ -121,19 +110,9 @@ impl CredentialBridge { provider_type: &str, model: &str, ) -> Result { - // 1. 从凭证池选择凭证 - // 将 provider_type 同时作为 provider_id_hint 传递,支持 60+ API Key Provider - // 例如 "deepseek", "moonshot", "qwen" 等 let credential = self - .pool_service - .select_credential_with_fallback( - db, - &self.api_key_service, - provider_type, - Some(model), - Some(provider_type), // 传递 provider_id_hint 支持智能降级 - None, - ) + .api_key_service + .select_credential_for_provider(db, provider_type, Some(provider_type), None) .await .map_err(CredentialBridgeError::DatabaseError)? .ok_or_else(|| { @@ -223,22 +202,6 @@ impl CredentialBridge { false, ), - // Kiro OAuth - 需要获取 access_token - CredentialData::KiroOAuth { creds_file_path } => { - let token = self - .get_kiro_token(creds_file_path, db, &credential.uuid) - .await?; - ("kiro".to_string(), Some(token), None, false) - } - - // Gemini OAuth - CredentialData::GeminiOAuth { - creds_file_path, .. - } => { - let token = self.get_oauth_token(creds_file_path).await?; - ("google".to_string(), Some(token), None, false) - } - // Gemini API Key CredentialData::GeminiApiKey { api_key, base_url, .. @@ -259,33 +222,11 @@ impl CredentialBridge { false, ), - // Codex OAuth - CredentialData::CodexOAuth { - creds_file_path, - api_base_url, - } => { - let token = self.get_codex_token(creds_file_path).await?; - ( - // 统一走 OpenAI provider,保证 tools/stream 事件链路一致 - "openai".to_string(), - Some(token), - api_base_url.clone(), - true, - ) - } - - // Claude OAuth - CredentialData::ClaudeOAuth { creds_file_path } => { - let token = self.get_oauth_token(creds_file_path).await?; - ("anthropic".to_string(), Some(token), None, false) - } - - // Antigravity OAuth - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => { - let token = self.get_oauth_token(creds_file_path).await?; - ("google".to_string(), Some(token), None, false) + unsupported => { + return Err(CredentialBridgeError::UnsupportedCredentialType(format!( + "凭证池/OAuth 凭证已退役,当前只支持 API Key Provider 凭证: {:?}", + unsupported + ))); } }; @@ -297,88 +238,11 @@ impl CredentialBridge { base_url, credential_uuid: credential.uuid.clone(), force_responses_api, - credential_path: match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => Some(creds_file_path.clone()), - _ => None, - }, toolshim: false, toolshim_model: None, }) } - /// 获取 Kiro OAuth Token - async fn get_kiro_token( - &self, - creds_path: &str, - _db: &DbConnection, - _uuid: &str, - ) -> Result { - use lime_providers::providers::KiroProvider; - - let mut provider = KiroProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("加载 Kiro 凭证失败: {e}")) - })?; - - // 检查 token 是否过期,如果过期则刷新 - if provider.is_token_expired() { - tracing::info!("[CredentialBridge] Kiro token 已过期,尝试刷新"); - self.pool_service - .refresh_kiro_token(creds_path) - .await - .map_err(CredentialBridgeError::TokenRefreshFailed)?; - - // 重新加载凭证 - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("重新加载凭证失败: {e}")) - })?; - } - - provider.credentials.access_token.ok_or_else(|| { - CredentialBridgeError::TokenRefreshFailed("缺少 access_token".to_string()) - }) - } - - /// 获取通用 OAuth Token - async fn get_oauth_token(&self, creds_path: &str) -> Result { - let content = std::fs::read_to_string(creds_path).map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("读取凭证文件失败: {e}")) - })?; - - let creds: serde_json::Value = serde_json::from_str(&content) - .map_err(|e| CredentialBridgeError::TokenRefreshFailed(format!("解析凭证失败: {e}")))?; - - creds["access_token"] - .as_str() - .map(String::from) - .ok_or_else(|| { - CredentialBridgeError::TokenRefreshFailed("凭证中缺少 access_token".to_string()) - }) - } - - /// 获取 Codex OAuth Token - async fn get_codex_token(&self, creds_path: &str) -> Result { - use lime_providers::providers::CodexProvider; - - let mut provider = CodexProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("加载 Codex 凭证失败: {e}")) - })?; - - provider.ensure_valid_token().await.map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("获取 Codex token 失败: {e}")) - }) - } - /// 记录凭证使用 pub fn record_usage(&self, db: &DbConnection, uuid: &str) -> Result<(), CredentialBridgeError> { if let Some(api_key_id) = self.resolve_fallback_api_key_id(uuid) { @@ -388,9 +252,8 @@ impl CredentialBridge { .map_err(CredentialBridgeError::DatabaseError); } - self.pool_service - .record_usage(db, uuid) - .map_err(CredentialBridgeError::DatabaseError) + tracing::debug!("[CredentialBridge] 忽略已退役的凭证池使用记录: {}", uuid); + Ok(()) } /// 标记凭证为健康 @@ -404,9 +267,9 @@ impl CredentialBridge { return Ok(()); } - self.pool_service - .mark_healthy(db, uuid, model) - .map_err(CredentialBridgeError::DatabaseError) + let _ = (db, model); + tracing::debug!("[CredentialBridge] 忽略已退役的凭证池健康标记: {}", uuid); + Ok(()) } /// 标记凭证为不健康 @@ -420,9 +283,9 @@ impl CredentialBridge { return Ok(()); } - self.pool_service - .mark_unhealthy(db, uuid, error) - .map_err(CredentialBridgeError::DatabaseError) + let _ = (db, error); + tracing::debug!("[CredentialBridge] 忽略已退役的凭证池失败标记: {}", uuid); + Ok(()) } } @@ -443,28 +306,6 @@ pub async fn create_aster_provider( ); } - if config.provider_name == "kiro" { - let model_config = build_provider_model_config(config)?; - - let credential_path = config.credential_path.clone().ok_or_else(|| { - CredentialBridgeError::ProviderCreationFailed( - "Kiro provider 缺少 credential_path".to_string(), - ) - })?; - - let provider = LimeKiroProvider::new(credential_path, model_config).map_err(|error| { - CredentialBridgeError::ProviderCreationFailed(format!( - "创建 Kiro Provider 失败: {}", - error - )) - })?; - - return Ok(wrap_provider_with_safety( - Arc::new(provider), - disable_default_fast_model, - )); - } - // 设置环境变量 set_provider_env_vars(config); @@ -515,7 +356,7 @@ fn is_first_party_openai_base_url(base_url: &str) -> bool { } fn is_first_party_anthropic_selector(selector: &str) -> bool { - matches!(selector, "anthropic" | "claude" | "claude_oauth") + matches!(selector, "anthropic" | "claude") } fn is_first_party_anthropic_base_url(base_url: &str) -> bool { @@ -714,28 +555,6 @@ fn set_provider_env_vars(config: &AsterProviderConfig) { } } -/// Provider 类型映射 -/// -/// 将 Lime PoolProviderType 映射到 Aster Provider 名称 -pub fn map_pool_type_to_aster(pool_type: &PoolProviderType) -> &'static str { - match pool_type { - PoolProviderType::Kiro => "kiro", - PoolProviderType::Gemini => "google", - PoolProviderType::Antigravity => "google", - PoolProviderType::OpenAI => "openai", - PoolProviderType::Claude => "anthropic", - PoolProviderType::Anthropic => "anthropic", - PoolProviderType::AnthropicCompatible => "anthropic", - PoolProviderType::Vertex => "gcpvertexai", - PoolProviderType::GeminiApiKey => "google", - PoolProviderType::Codex => "codex", - PoolProviderType::ClaudeOAuth => "anthropic", - PoolProviderType::AzureOpenai => "azure", - PoolProviderType::AwsBedrock => "bedrock", - PoolProviderType::Ollama => "ollama", - } -} - /// 将 provider_type 字符串映射到 Aster Provider 名称 /// /// 支持 60+ API Key Provider,包括 deepseek, moonshot, qwen 等 @@ -750,9 +569,8 @@ fn map_provider_type_to_aster(provider_type: &str) -> &'static str { "anthropic" | "claude" => "anthropic", "google" | "gemini" => "google", "bedrock" => "bedrock", - "kiro" | "codewhisperer" => "kiro", "gcpvertexai" | "vertex" => "gcpvertexai", - "codex" => "codex", + "codex" => "openai", "azure" | "azure-openai" => "azure", "ollama" => "ollama", @@ -789,17 +607,6 @@ fn map_provider_type_to_aster_with_api_type( mod tests { use super::*; - #[test] - fn test_map_pool_type_to_aster() { - assert_eq!(map_pool_type_to_aster(&PoolProviderType::OpenAI), "openai"); - assert_eq!( - map_pool_type_to_aster(&PoolProviderType::Claude), - "anthropic" - ); - assert_eq!(map_pool_type_to_aster(&PoolProviderType::Gemini), "google"); - assert_eq!(map_pool_type_to_aster(&PoolProviderType::Kiro), "kiro"); - } - #[test] fn test_map_provider_type_to_aster_with_api_type() { assert_eq!( @@ -836,7 +643,6 @@ mod tests { base_url: Some("https://example.com/openai".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: true, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -871,7 +677,6 @@ mod tests { base_url: Some("https://api.openai.com/v1".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: true, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -908,7 +713,6 @@ mod tests { base_url: Some("https://open.bigmodel.cn/api/paas/v4".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -926,7 +730,6 @@ mod tests { base_url: Some("https://api.openai.com/v1".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -944,7 +747,6 @@ mod tests { base_url: Some("https://example.com/openai".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: true, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -962,7 +764,6 @@ mod tests { base_url: Some("https://token-plan-cn.xiaomimimo.com/anthropic".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -980,7 +781,6 @@ mod tests { base_url: Some("https://api.anthropic.com".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -1029,7 +829,6 @@ mod tests { base_url: Some("https://open.bigmodel.cn/api/anthropic".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -1064,7 +863,6 @@ mod tests { base_url: Some("https://api.anthropic.com".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: false, - credential_path: None, toolshim: false, toolshim_model: None, }; @@ -1088,7 +886,6 @@ mod tests { base_url: Some("http://127.0.0.1:11434".to_string()), credential_uuid: "test-uuid".to_string(), force_responses_api: false, - credential_path: None, toolshim: true, toolshim_model: Some("glm-5.1:cloud".to_string()), }; diff --git a/src-tauri/crates/agent/src/kiro_provider_adapter.rs b/src-tauri/crates/agent/src/kiro_provider_adapter.rs deleted file mode 100644 index bba1be1e3..000000000 --- a/src-tauri/crates/agent/src/kiro_provider_adapter.rs +++ /dev/null @@ -1,360 +0,0 @@ -use anyhow::anyhow; -use aster::conversation::message::{Message, MessageContent}; -use aster::model::ModelConfig; -use aster::providers::base::{ - ConfigKey, MessageStream, ModelInfo, Provider, ProviderMetadata, ProviderUsage, Usage, -}; -use aster::providers::errors::ProviderError; -use aster::providers::formats::openai::{ - format_messages, format_tools, response_to_streaming_message, -}; -use aster::providers::utils::ImageFormat; -use aster::session_context::current_turn_context; -use async_stream::try_stream; -use async_trait::async_trait; -use futures::{pin_mut, StreamExt}; -use lime_core::models::openai::ChatCompletionRequest; -use lime_providers::providers::{KiroProvider, TokenManager}; -use lime_providers::streaming::converter::{StreamConverter, StreamFormat as LimeStreamFormat}; -use rmcp::model::{Role, Tool}; -use serde_json::json; -use uuid::Uuid; - -const KIRO_PROVIDER_NAME: &str = "kiro"; - -pub(crate) struct LimeKiroProvider { - credential_path: String, - model: ModelConfig, - name: String, -} - -impl LimeKiroProvider { - pub(crate) fn new( - credential_path: impl Into, - model: ModelConfig, - ) -> Result { - let credential_path = credential_path.into(); - if credential_path.trim().is_empty() { - return Err(ProviderError::ExecutionError( - "Kiro provider 缺少 credential_path".to_string(), - )); - } - - Ok(Self { - credential_path, - model, - name: KIRO_PROVIDER_NAME.to_string(), - }) - } - - async fn load_provider(&self) -> Result { - let mut provider = KiroProvider::new(); - provider - .load_credentials_from_path(&self.credential_path) - .await - .map_err(|error| { - ProviderError::Authentication(format!("加载 Kiro 凭证失败: {}", error)) - })?; - - provider.ensure_valid_token().await.map_err(|error| { - ProviderError::Authentication(format!("刷新 Kiro Token 失败: {}", error)) - })?; - - Ok(provider) - } - - fn normalize_optional_text(value: Option<&str>) -> Option { - value - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_string) - } - - fn resolve_conversation_id_from_turn_context() -> Option { - let turn_context = current_turn_context()?; - let provider_continuation = turn_context - .metadata - .get("provider_continuation")? - .as_object()?; - - if provider_continuation - .get("enabled") - .and_then(serde_json::Value::as_bool) - != Some(true) - { - return None; - } - - if provider_continuation - .get("kind") - .and_then(serde_json::Value::as_str) - .map(str::trim) - != Some("provider_session_token") - { - return None; - } - - for key in [ - "session_token", - "sessionToken", - "provider_session_token", - "providerSessionToken", - "conversation_id", - "conversationId", - ] { - if let Some(value) = Self::normalize_optional_text( - provider_continuation - .get(key) - .and_then(serde_json::Value::as_str), - ) { - return Some(value); - } - } - - None - } - - fn resolve_or_create_conversation_id() -> String { - Self::resolve_conversation_id_from_turn_context() - .unwrap_or_else(|| Uuid::new_v4().to_string()) - } - - fn build_chat_request( - model_config: &ModelConfig, - system: &str, - messages: &[Message], - tools: &[Tool], - stream: bool, - ) -> Result { - let mut openai_messages = vec![json!({ - "role": "system", - "content": system, - })]; - openai_messages.extend(format_messages(messages, &ImageFormat::OpenAi)); - - let tools_payload = format_tools(tools) - .map_err(|error| ProviderError::ExecutionError(error.to_string()))?; - - let mut payload = json!({ - "model": model_config.model_name, - "messages": openai_messages, - "stream": stream, - }); - - if !tools_payload.is_empty() { - payload["tools"] = json!(tools_payload); - } - - serde_json::from_value(payload).map_err(|error| { - ProviderError::ExecutionError(format!( - "构造 Kiro ChatCompletionRequest 失败: {}", - error - )) - }) - } - - fn attach_conversation_id(mut message: Message, conversation_id: &str) -> Message { - if message.role == Role::Assistant { - message.id = Some(conversation_id.to_string()); - } - message - } - - fn push_or_merge_content(target: &mut Vec, content: MessageContent) { - match (target.last_mut(), &content) { - (Some(MessageContent::Text(existing)), MessageContent::Text(incoming)) => { - existing.text.push_str(&incoming.text); - } - (Some(MessageContent::Thinking(existing)), MessageContent::Thinking(incoming)) => { - existing.thinking.push_str(&incoming.thinking); - } - _ => target.push(content), - } - } - - fn merge_message_chunk(target: &mut Message, chunk: Message) { - if target.id.is_none() { - target.id = chunk.id.clone(); - } - - for content in chunk.content { - Self::push_or_merge_content(&mut target.content, content); - } - } - - async fn stream_with_model_and_conversation( - &self, - model_config: &ModelConfig, - system: &str, - messages: &[Message], - tools: &[Tool], - conversation_id: String, - ) -> Result { - let request = Self::build_chat_request(model_config, system, messages, tools, true)?; - let provider = self.load_provider().await?; - let source_stream = provider - .call_api_stream_with_conversation_id(&request, Some(&conversation_id)) - .await - .map_err(Self::map_lime_provider_error)?; - - let model_name = model_config.model_name.clone(); - let openai_line_stream = Box::pin(try_stream! { - let mut source_stream = source_stream; - let mut converter = StreamConverter::with_model( - LimeStreamFormat::AwsEventStream, - LimeStreamFormat::OpenAiSse, - &model_name, - ); - - while let Some(chunk) = source_stream.next().await { - let chunk = chunk.map_err(|error| anyhow!(error.to_string()))?; - for event in converter.convert(&chunk) { - for line in event.lines() { - let line = line.trim_end_matches('\r'); - if !line.is_empty() { - yield line.to_string(); - } - } - } - } - - for event in converter.finish() { - for line in event.lines() { - let line = line.trim_end_matches('\r'); - if !line.is_empty() { - yield line.to_string(); - } - } - } - }); - - Ok(Box::pin(try_stream! { - let message_stream = response_to_streaming_message(openai_line_stream); - pin_mut!(message_stream); - - while let Some(item) = message_stream.next().await { - let (message, usage) = item.map_err(|error| { - ProviderError::RequestFailed(format!("解析 Kiro 流式响应失败: {}", error)) - })?; - - let message = message.map(|message| { - Self::attach_conversation_id(message, &conversation_id) - }); - - yield (message, usage); - } - })) - } - - fn map_lime_provider_error(error: lime_providers::providers::ProviderError) -> ProviderError { - match error { - lime_providers::providers::ProviderError::AuthenticationError(details) => { - ProviderError::Authentication(details) - } - lime_providers::providers::ProviderError::RateLimitError(details) => { - ProviderError::RateLimitExceeded { - details, - retry_delay: None, - } - } - lime_providers::providers::ProviderError::ServerError(details) => { - ProviderError::ServerError(details) - } - lime_providers::providers::ProviderError::RequestError(details) => { - ProviderError::RequestFailed(details) - } - lime_providers::providers::ProviderError::ParseError(details) - | lime_providers::providers::ProviderError::ConfigurationError(details) - | lime_providers::providers::ProviderError::Unknown(details) - | lime_providers::providers::ProviderError::TokenExpired(details) - | lime_providers::providers::ProviderError::NetworkError(details) => { - ProviderError::ExecutionError(details) - } - } - } -} - -#[async_trait] -impl Provider for LimeKiroProvider { - fn metadata() -> ProviderMetadata - where - Self: Sized, - { - ProviderMetadata::with_models( - KIRO_PROVIDER_NAME, - "Kiro", - "Lime 本地 Kiro/CodeWhisperer Provider 适配器", - "claude-sonnet-4-5", - vec![ModelInfo::new("claude-sonnet-4-5", 200_000)], - "", - vec![ConfigKey::new("KIRO_CREDENTIAL_PATH", true, true, None)], - ) - } - - fn get_name(&self) -> &str { - &self.name - } - - async fn complete_with_model( - &self, - model_config: &ModelConfig, - system: &str, - messages: &[Message], - tools: &[Tool], - ) -> Result<(Message, ProviderUsage), ProviderError> { - let conversation_id = Self::resolve_or_create_conversation_id(); - let mut stream = self - .stream_with_model_and_conversation( - model_config, - system, - messages, - tools, - conversation_id.clone(), - ) - .await?; - - let mut final_message = Message::assistant().with_id(conversation_id.clone()); - let mut final_usage = None; - - while let Some(item) = stream.next().await { - let (message, usage) = item?; - if let Some(message) = message { - Self::merge_message_chunk(&mut final_message, message); - } - if usage.is_some() { - final_usage = usage; - } - } - - let usage = final_usage.unwrap_or_else(|| { - ProviderUsage::new(model_config.model_name.clone(), Usage::default()) - }); - - Ok((final_message, usage)) - } - - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - - async fn stream( - &self, - system: &str, - messages: &[Message], - tools: &[Tool], - ) -> Result { - let conversation_id = Self::resolve_or_create_conversation_id(); - self.stream_with_model_and_conversation( - &self.model, - system, - messages, - tools, - conversation_id, - ) - .await - } - - fn supports_streaming(&self) -> bool { - true - } -} diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs index d2cd2cd3f..e11ec0b4a 100644 --- a/src-tauri/crates/agent/src/lib.rs +++ b/src-tauri/crates/agent/src/lib.rs @@ -21,7 +21,6 @@ pub mod durable_memory_fs; pub mod event_converter; pub mod filesystem_event_protocol; pub mod hooks; -mod kiro_provider_adapter; pub mod lsp_bridge; pub mod mcp_bridge; pub mod prompt; diff --git a/src-tauri/crates/agent/src/provider_continuation_state.rs b/src-tauri/crates/agent/src/provider_continuation_state.rs index 0c1035829..88e0700a6 100644 --- a/src-tauri/crates/agent/src/provider_continuation_state.rs +++ b/src-tauri/crates/agent/src/provider_continuation_state.rs @@ -19,11 +19,6 @@ fn is_openai_responses_model(model_name: &str) -> bool { normalized.starts_with("gpt-5") && normalized.contains("codex") } -fn is_kiro_session_provider(candidate: &str) -> bool { - let normalized = candidate.trim().to_ascii_lowercase(); - normalized.contains("kiro") || normalized.contains("codewhisperer") -} - #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum ProviderContinuationCapability { @@ -64,14 +59,6 @@ pub fn resolve_provider_continuation_capability( return ProviderContinuationCapability::PreviousResponseId; } - if provider_candidates - .iter() - .flatten() - .any(|candidate| is_kiro_session_provider(candidate)) - { - return ProviderContinuationCapability::ProviderSessionToken; - } - ProviderContinuationCapability::HistoryReplayOnly } @@ -224,10 +211,10 @@ mod tests { } #[test] - fn test_resolve_provider_continuation_capability_detects_kiro_provider_session_token() { + fn test_resolve_provider_continuation_capability_treats_retired_kiro_as_history_replay_only() { assert_eq!( resolve_provider_continuation_capability("kiro", Some("kiro"), "claude-3.7", false), - ProviderContinuationCapability::ProviderSessionToken + ProviderContinuationCapability::HistoryReplayOnly ); assert_eq!( resolve_provider_continuation_capability( @@ -236,7 +223,7 @@ mod tests { "claude-3.7", false ), - ProviderContinuationCapability::ProviderSessionToken + ProviderContinuationCapability::HistoryReplayOnly ); } diff --git a/src-tauri/crates/agent/src/request_tool_policy.rs b/src-tauri/crates/agent/src/request_tool_policy.rs index 9b40034e9..5a5d6b75a 100644 --- a/src-tauri/crates/agent/src/request_tool_policy.rs +++ b/src-tauri/crates/agent/src/request_tool_policy.rs @@ -1621,6 +1621,7 @@ pub async fn execute_web_search_preflight_if_needed( message_text: &str, working_directory: Option<&Path>, cancel_token: Option, + turn_context: Option, policy: &RequestToolPolicy, tracker: &mut WebSearchExecutionTracker, ) -> Result { @@ -1687,6 +1688,7 @@ pub async fn execute_web_search_preflight_if_needed( let session_id = session_id.to_string(); let working_directory = working_directory.clone(); let cancel_token = cancel_token.clone(); + let turn_context = turn_context.clone(); async move { let query = planned.query.clone(); let params = serde_json::json!({ "query": query }); @@ -1694,12 +1696,13 @@ pub async fn execute_web_search_preflight_if_needed( if let Some(token) = cancel_token { context = context.with_cancellation_token(token); } - let result = { + let result = aster::session_context::with_turn_context(turn_context, async { let registry = registry_arc.read().await; registry .execute(&preflight_tool_name, params, &context, None) .await - }; + }) + .await; match result { Ok(tool_result) => PreflightSearchOutcome { index: planned.index, @@ -1825,6 +1828,7 @@ where &message_text, working_directory, cancel_token.clone(), + session_config.turn_context.clone(), request_tool_policy, &mut web_search_tracker, ) @@ -1839,6 +1843,7 @@ where &message_text, working_directory, cancel_token.clone(), + session_config.turn_context.clone(), request_tool_policy, &mut optional_preflight_tracker, ), @@ -2192,12 +2197,76 @@ mod tests { use aster::providers::base::{Provider, ProviderMetadata, ProviderUsage, Usage}; use aster::providers::errors::ProviderError; use aster::session::{SessionManager, SessionType, TurnContextOverride, TurnStatus}; + use aster::tools::{PermissionCheckResult, Tool, ToolError, ToolResult}; use async_trait::async_trait; use std::collections::HashMap; use std::path::PathBuf; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; + struct TurnContextGatedWebSearchTool; + + #[async_trait] + impl Tool for TurnContextGatedWebSearchTool { + fn name(&self) -> &str { + "WebSearch" + } + + fn description(&self) -> &str { + "测试用 WebSearch 工具" + } + + fn input_schema(&self) -> serde_json::Value { + serde_json::json!({ + "type": "object", + "properties": { + "query": { "type": "string" } + }, + "required": ["query"] + }) + } + + async fn check_permissions( + &self, + _params: &serde_json::Value, + _context: &ToolContext, + ) -> PermissionCheckResult { + let allowed = aster::session_context::current_turn_context() + .as_ref() + .is_some_and(|turn_context| { + ["web_search_enabled", "webSearchEnabled"] + .iter() + .any(|key| { + turn_context + .metadata + .get(*key) + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + }) + }); + + if allowed { + PermissionCheckResult::allow() + } else { + PermissionCheckResult::ask("WebSearch 需要联网确认。") + } + } + + async fn execute( + &self, + params: serde_json::Value, + _context: &ToolContext, + ) -> Result { + let query = params + .get("query") + .and_then(serde_json::Value::as_str) + .unwrap_or_default(); + Ok(ToolResult::success(format!( + "预检索测试结果:https://example.com/search?q={query}" + ))) + } + } + struct ContextLengthExceededProvider; #[async_trait] @@ -2533,6 +2602,53 @@ mod tests { ); } + #[tokio::test] + async fn web_search_preflight_uses_turn_context_for_permission_check() { + let agent = Agent::new(); + { + let registry_arc = agent.tool_registry().clone(); + let mut registry = registry_arc.write().await; + registry.register(Box::new(TurnContextGatedWebSearchTool)); + } + + let policy = resolve_request_tool_policy_with_mode( + Some(true), + Some(RequestToolPolicyMode::Required), + false, + ); + let mut metadata = HashMap::new(); + metadata.insert("webSearchEnabled".to_string(), serde_json::json!(true)); + let turn_context = TurnContextOverride { + metadata, + ..TurnContextOverride::default() + }; + let mut tracker = WebSearchExecutionTracker::default(); + + let execution = execute_web_search_preflight_if_needed( + &agent, + "session-web-preflight-permission", + "继续", + None, + None, + Some(turn_context), + &policy, + &mut tracker, + ) + .await + .expect("预调用应继承 turn context 并免确认执行"); + + assert!(execution + .system_prompt_appendix + .as_deref() + .unwrap_or_default() + .contains("预检索测试结果")); + assert!(execution.events.iter().any(|event| matches!( + event, + RuntimeAgentEvent::ToolEnd { result, .. } if result.success + ))); + assert!(tracker.validate_web_search_requirement(&policy).is_ok()); + } + #[test] fn discarded_optional_preflight_attempt_should_not_force_synthesis_retry() { let diagnostics = StreamEventDiagnostics::default(); diff --git a/src-tauri/crates/agent/src/runtime_queue.rs b/src-tauri/crates/agent/src/runtime_queue.rs index 169a30478..0d06900a9 100644 --- a/src-tauri/crates/agent/src/runtime_queue.rs +++ b/src-tauri/crates/agent/src/runtime_queue.rs @@ -30,6 +30,7 @@ fn emit_runtime_queue_event( fn spawn_runtime_turn_task( session_id: String, + event_name: String, context: C, executor: RuntimeQueueExecutor, emitter: RuntimeQueueEventEmitter, @@ -38,6 +39,8 @@ fn spawn_runtime_turn_task( C: Clone + Send + Sync + 'static, { let thread_name = format!("lime-runtime-turn-{}", session_id); + let event_name_for_thread = event_name.clone(); + let emitter_for_thread = emitter.clone(); let spawn_result = std::thread::Builder::new() .name(thread_name) .stack_size(RUNTIME_TURN_THREAD_STACK_SIZE) @@ -52,9 +55,12 @@ fn spawn_runtime_turn_task( { Ok(runtime) => runtime, Err(error) => { - tracing::error!( - "[AsterAgent][Queue] 创建 runtime turn 专用运行时失败: {}", - error + let message = format!("创建 runtime turn 专用运行时失败: {error}"); + tracing::error!("[AsterAgent][Queue] {}", message); + emit_runtime_queue_event( + &emitter_for_thread, + &event_name_for_thread, + RuntimeAgentEvent::Error { message }, ); return; } @@ -62,27 +68,31 @@ fn spawn_runtime_turn_task( runtime.block_on(async move { let result = executor(context.clone(), payload).await; + if let Err(error) = result { + tracing::warn!("[AsterAgent][Queue] 队列任务执行失败: {}", error); + emit_runtime_queue_event( + &emitter_for_thread, + &event_name_for_thread, + RuntimeAgentEvent::Error { message: error }, + ); + } if let Err(error) = continue_runtime_queue_after_turn( session_id, context.clone(), executor.clone(), - emitter.clone(), + emitter_for_thread.clone(), ) .await { tracing::warn!("[AsterAgent][Queue] 调度下一条排队 turn 失败: {}", error); } - if let Err(error) = result { - tracing::warn!("[AsterAgent][Queue] 队列任务执行失败: {}", error); - } }); }); if let Err(error) = spawn_result { - tracing::error!( - "[AsterAgent][Queue] 启动 runtime turn 专用线程失败: {}", - error - ); + let message = format!("启动 runtime turn 专用线程失败: {error}"); + tracing::error!("[AsterAgent][Queue] {}", message); + emit_runtime_queue_event(&emitter, &event_name, RuntimeAgentEvent::Error { message }); } } @@ -138,6 +148,7 @@ where spawn_runtime_turn_task( session_id, + event_name, context, executor, emitter, @@ -192,7 +203,14 @@ where .map_err(|error| format!("提交 runtime queue turn 失败: {error}"))? { RuntimeQueueSubmitResult::StartNow => { - spawn_runtime_turn_task(session_id, context, executor, emitter, queued_task.payload); + spawn_runtime_turn_task( + session_id, + queued_task.event_name, + context, + executor, + emitter, + queued_task.payload, + ); Ok(()) } RuntimeQueueSubmitResult::Busy => Err("当前会话仍在生成,无法立即开始执行".to_string()), diff --git a/src-tauri/crates/core/src/agent/types.rs b/src-tauri/crates/core/src/agent/types.rs index 6bc14db5b..f1518527e 100644 --- a/src-tauri/crates/core/src/agent/types.rs +++ b/src-tauri/crates/core/src/agent/types.rs @@ -29,8 +29,6 @@ pub enum ProviderType { Qwen, /// Codex (OpenAI 兼容) Codex, - /// Antigravity (OpenAI 兼容) - Antigravity, /// iFlow (OpenAI 兼容) IFlow, } @@ -47,7 +45,6 @@ impl ProviderType { "openai" => Self::OpenAI, "qwen" => Self::Qwen, "codex" => Self::Codex, - "antigravity" => Self::Antigravity, "iflow" => Self::IFlow, _ => Self::OpenAI, // 默认使用 OpenAI 协议 } @@ -116,7 +113,7 @@ impl ProviderType { pub fn is_openai_compatible(&self) -> bool { matches!( self, - Self::OpenAI | Self::Qwen | Self::Codex | Self::Antigravity | Self::IFlow | Self::Kiro + Self::OpenAI | Self::Qwen | Self::Codex | Self::IFlow | Self::Kiro ) } } diff --git a/src-tauri/crates/core/src/config/tests.rs b/src-tauri/crates/core/src/config/tests.rs index b31a85ad8..5a5f823ee 100644 --- a/src-tauri/crates/core/src/config/tests.rs +++ b/src-tauri/crates/core/src/config/tests.rs @@ -2507,7 +2507,6 @@ fn arb_valid_provider_type() -> impl Strategy { Just("gemini".to_string()), Just("openai".to_string()), Just("claude".to_string()), - Just("antigravity".to_string()), Just("vertex".to_string()), Just("gemini_api_key".to_string()), Just("codex".to_string()), diff --git a/src-tauri/crates/core/src/config/types.rs b/src-tauri/crates/core/src/config/types.rs index 3a1fc7c1b..db5f57f1d 100644 --- a/src-tauri/crates/core/src/config/types.rs +++ b/src-tauri/crates/core/src/config/types.rs @@ -1577,7 +1577,7 @@ pub struct RoutingConfig { } fn default_provider() -> String { - "kiro".to_string() + "openai".to_string() } impl Default for RoutingConfig { @@ -1982,51 +1982,6 @@ impl Default for ModelsConfig { }, ); - // Antigravity - providers.insert( - "antigravity".to_string(), - ProviderModelsConfig { - label: "Antigravity".to_string(), - models: vec![ - ModelInfo { - id: "gemini-3-pro-preview".to_string(), - name: None, - enabled: true, - }, - ModelInfo { - id: "gemini-3-pro-image-preview".to_string(), - name: None, - enabled: true, - }, - ModelInfo { - id: "gemini-3-flash-preview".to_string(), - name: None, - enabled: true, - }, - ModelInfo { - id: "gemini-2.5-computer-use-preview-10-2025".to_string(), - name: None, - enabled: true, - }, - ModelInfo { - id: "gemini-claude-sonnet-4-5".to_string(), - name: None, - enabled: true, - }, - ModelInfo { - id: "gemini-claude-sonnet-4-5-thinking".to_string(), - name: None, - enabled: true, - }, - ModelInfo { - id: "gemini-claude-opus-4-5-thinking".to_string(), - name: None, - enabled: true, - }, - ], - }, - ); - // Submodel providers.insert( "submodel".to_string(), diff --git a/src-tauri/crates/core/src/connect/README.md b/src-tauri/crates/core/src/connect/README.md index 0b5cf826b..4b69fc19b 100644 --- a/src-tauri/crates/core/src/connect/README.md +++ b/src-tauri/crates/core/src/connect/README.md @@ -6,7 +6,7 @@ Lime Connect 模块,实现中转商生态合作方案。 通过 Deep Link 协议实现一键配置功能,支持中转商品牌展示。 -API Key 直接集成到凭证池系统,无需单独存储。 +API Key 直接集成到 API Key Provider 主路径,不再进入旧凭证池系统。 ## 功能概述 @@ -43,7 +43,7 @@ API Key 直接集成到凭证池系统,无需单独存储。 - Requirements 1.x - Deep Link 协议处理 - Requirements 2.x - 中转商注册表管理 -- Requirements 4.x - API Key 存储(已集成到凭证池系统) +- Requirements 4.x - API Key 存储(已集成到 API Key Provider) - Requirements 5.3 - 统计回调(Webhook) ## 更新提醒 diff --git a/src-tauri/crates/core/src/connect/mod.rs b/src-tauri/crates/core/src/connect/mod.rs index 7571cde15..3c49baf9a 100644 --- a/src-tauri/crates/core/src/connect/mod.rs +++ b/src-tauri/crates/core/src/connect/mod.rs @@ -20,7 +20,7 @@ //! let registry = RelayRegistry::new(cache_path); //! let relay_info = registry.get(&payload.relay); //! -//! // API Key 直接保存到凭证池系统(通过 ProviderPoolService) +//! // API Key 通过 API Key Provider 主路径保存 //! ``` // 子模块声明 diff --git a/src-tauri/crates/core/src/credential/health.rs b/src-tauri/crates/core/src/credential/health.rs deleted file mode 100644 index 64669d172..000000000 --- a/src-tauri/crates/core/src/credential/health.rs +++ /dev/null @@ -1,426 +0,0 @@ -//! 凭证健康检查器实现 -//! -//! 提供凭证健康状态检查和自动更新功能 - -use super::pool::{CredentialPool, PoolError}; -use super::types::{Credential, CredentialStatus}; -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::time::Duration; - -/// 健康状态 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub enum HealthStatus { - /// 健康 - Healthy, - /// 不健康 - Unhealthy { - /// 不健康原因 - reason: String, - /// 连续失败次数 - consecutive_failures: u32, - }, - /// 未知(未检查过) - Unknown, -} - -/// 健康检查结果 -#[derive(Debug, Clone)] -pub struct HealthCheckResult { - /// 凭证 ID - pub credential_id: String, - /// 健康状态 - pub status: HealthStatus, - /// 检查时间 - pub checked_at: DateTime, - /// 检查延迟(毫秒) - pub latency_ms: Option, -} - -/// 健康检查配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HealthCheckConfig { - /// 检查间隔 - pub check_interval: Duration, - /// 连续失败阈值(达到此值标记为不健康) - pub failure_threshold: u32, - /// 恢复阈值(连续成功此次数后恢复为健康) - pub recovery_threshold: u32, -} - -impl Default for HealthCheckConfig { - fn default() -> Self { - Self { - check_interval: Duration::from_secs(60), - failure_threshold: 3, - recovery_threshold: 1, - } - } -} - -/// 健康检查器 - 管理凭证健康状态 -pub struct HealthChecker { - /// 配置 - config: HealthCheckConfig, -} - -impl HealthChecker { - /// 创建新的健康检查器 - pub fn new(config: HealthCheckConfig) -> Self { - Self { config } - } - - /// 使用默认配置创建健康检查器 - pub fn with_defaults() -> Self { - Self::new(HealthCheckConfig::default()) - } - - /// 获取配置 - pub fn config(&self) -> &HealthCheckConfig { - &self.config - } - - /// 获取失败阈值 - pub fn failure_threshold(&self) -> u32 { - self.config.failure_threshold - } - - /// 获取恢复阈值 - pub fn recovery_threshold(&self) -> u32 { - self.config.recovery_threshold - } - - /// 检查单个凭证的健康状态 - /// - /// 根据凭证的统计信息判断健康状态 - pub fn check(&self, credential: &Credential) -> HealthCheckResult { - let status = self.evaluate_health(credential); - - HealthCheckResult { - credential_id: credential.id.clone(), - status, - checked_at: Utc::now(), - latency_ms: if credential.stats.successful_requests > 0 { - Some(credential.stats.avg_latency_ms as u64) - } else { - None - }, - } - } - - /// 评估凭证健康状态 - fn evaluate_health(&self, credential: &Credential) -> HealthStatus { - // 如果已经被标记为不健康,返回当前状态 - if let CredentialStatus::Unhealthy { reason } = &credential.status { - return HealthStatus::Unhealthy { - reason: reason.clone(), - consecutive_failures: credential.stats.consecutive_failures, - }; - } - - // 检查连续失败次数 - if credential.stats.consecutive_failures >= self.config.failure_threshold { - return HealthStatus::Unhealthy { - reason: format!( - "连续失败 {} 次(阈值: {})", - credential.stats.consecutive_failures, self.config.failure_threshold - ), - consecutive_failures: credential.stats.consecutive_failures, - }; - } - - // 如果没有请求记录,状态未知 - if credential.stats.total_requests == 0 { - return HealthStatus::Unknown; - } - - HealthStatus::Healthy - } - - /// 记录凭证使用失败并更新健康状态 - /// - /// 如果连续失败次数达到阈值,自动标记为不健康 - /// - /// # 返回 - /// - `true` 如果凭证被标记为不健康 - /// - `false` 如果凭证仍然健康 - pub fn record_failure( - &self, - pool: &CredentialPool, - credential_id: &str, - ) -> Result { - // 记录失败 - pool.record_failure(credential_id)?; - - // 获取更新后的凭证 - let credential = pool - .get(credential_id) - .ok_or_else(|| PoolError::CredentialNotFound(credential_id.to_string()))?; - - // 检查是否需要标记为不健康 - if credential.stats.consecutive_failures >= self.config.failure_threshold { - let reason = format!("连续认证失败 {} 次", credential.stats.consecutive_failures); - pool.mark_unhealthy(credential_id, reason)?; - return Ok(true); - } - - Ok(false) - } - - /// 记录凭证使用成功并更新健康状态 - /// - /// 如果凭证之前不健康,成功后会恢复为健康状态 - /// - /// # 返回 - /// - `true` 如果凭证从不健康恢复为健康 - /// - `false` 如果凭证状态未改变 - pub fn record_success( - &self, - pool: &CredentialPool, - credential_id: &str, - latency_ms: u64, - ) -> Result { - // 获取当前状态 - let was_unhealthy = pool - .get(credential_id) - .map(|c| matches!(c.status, CredentialStatus::Unhealthy { .. })) - .unwrap_or(false); - - // 记录成功 - pool.record_success(credential_id, latency_ms)?; - - // 如果之前不健康,恢复为健康 - if was_unhealthy { - pool.mark_active(credential_id)?; - return Ok(true); - } - - Ok(false) - } - - /// 批量检查凭证池中所有凭证的健康状态 - pub fn check_all(&self, pool: &CredentialPool) -> Vec { - pool.all().iter().map(|cred| self.check(cred)).collect() - } - - /// 获取池中不健康的凭证数量 - pub fn unhealthy_count(&self, pool: &CredentialPool) -> usize { - pool.all() - .iter() - .filter(|cred| matches!(self.check(cred).status, HealthStatus::Unhealthy { .. })) - .count() - } - - /// 检查凭证是否应该被标记为不健康 - /// - /// 这是一个纯函数,不会修改任何状态 - pub fn should_mark_unhealthy(&self, consecutive_failures: u32) -> bool { - consecutive_failures >= self.config.failure_threshold - } - - /// 尝试恢复不健康的凭证 - /// - /// 将所有不健康的凭证恢复为活跃状态(用于手动恢复) - pub fn recover_all(&self, pool: &CredentialPool) -> Vec { - let mut recovered = Vec::new(); - - for cred in pool.all() { - if matches!(cred.status, CredentialStatus::Unhealthy { .. }) - && pool.mark_active(&cred.id).is_ok() - { - recovered.push(cred.id.clone()); - } - } - - recovered - } -} - -impl Default for HealthChecker { - fn default() -> Self { - Self::with_defaults() - } -} - -#[cfg(test)] -mod health_tests { - use super::*; - use crate::credential::CredentialData; - use crate::ProviderType; - - fn create_test_credential(id: &str) -> Credential { - Credential::new( - id.to_string(), - ProviderType::Kiro, - CredentialData::ApiKey { - key: format!("key-{id}"), - base_url: None, - }, - ) - } - - #[test] - fn test_health_checker_new() { - let checker = HealthChecker::with_defaults(); - assert_eq!(checker.failure_threshold(), 3); - assert_eq!(checker.recovery_threshold(), 1); - } - - #[test] - fn test_health_checker_custom_config() { - let config = HealthCheckConfig { - check_interval: Duration::from_secs(30), - failure_threshold: 5, - recovery_threshold: 2, - }; - let checker = HealthChecker::new(config); - assert_eq!(checker.failure_threshold(), 5); - assert_eq!(checker.recovery_threshold(), 2); - } - - #[test] - fn test_check_healthy_credential() { - let checker = HealthChecker::with_defaults(); - let mut cred = create_test_credential("test-1"); - - // 记录一些成功请求 - cred.stats.record_success(100); - cred.stats.record_success(150); - - let result = checker.check(&cred); - assert_eq!(result.credential_id, "test-1"); - assert!(matches!(result.status, HealthStatus::Healthy)); - } - - #[test] - fn test_check_unknown_credential() { - let checker = HealthChecker::with_defaults(); - let cred = create_test_credential("test-1"); - - // 没有任何请求记录 - let result = checker.check(&cred); - assert!(matches!(result.status, HealthStatus::Unknown)); - } - - #[test] - fn test_check_unhealthy_credential() { - let checker = HealthChecker::with_defaults(); - let mut cred = create_test_credential("test-1"); - - // 记录 3 次连续失败 - cred.stats.record_failure(); - cred.stats.record_failure(); - cred.stats.record_failure(); - - let result = checker.check(&cred); - assert!(matches!(result.status, HealthStatus::Unhealthy { .. })); - - if let HealthStatus::Unhealthy { - consecutive_failures, - .. - } = result.status - { - assert_eq!(consecutive_failures, 3); - } - } - - #[test] - fn test_record_failure_marks_unhealthy() { - let checker = HealthChecker::with_defaults(); - let pool = CredentialPool::new(ProviderType::Kiro); - pool.add(create_test_credential("test-1")).unwrap(); - - // 前两次失败不应标记为不健康 - assert!(!checker.record_failure(&pool, "test-1").unwrap()); - assert!(!checker.record_failure(&pool, "test-1").unwrap()); - - // 第三次失败应标记为不健康 - assert!(checker.record_failure(&pool, "test-1").unwrap()); - - // 验证状态 - let cred = pool.get("test-1").unwrap(); - assert!(matches!(cred.status, CredentialStatus::Unhealthy { .. })); - } - - #[test] - fn test_record_success_recovers_unhealthy() { - let checker = HealthChecker::with_defaults(); - let pool = CredentialPool::new(ProviderType::Kiro); - pool.add(create_test_credential("test-1")).unwrap(); - - // 标记为不健康 - pool.mark_unhealthy("test-1", "test reason".to_string()) - .unwrap(); - - // 记录成功应恢复 - let recovered = checker.record_success(&pool, "test-1", 100).unwrap(); - assert!(recovered); - - // 验证状态 - let cred = pool.get("test-1").unwrap(); - assert!(matches!(cred.status, CredentialStatus::Active)); - } - - #[test] - fn test_check_all() { - let checker = HealthChecker::with_defaults(); - let pool = CredentialPool::new(ProviderType::Kiro); - - pool.add(create_test_credential("cred-1")).unwrap(); - pool.add(create_test_credential("cred-2")).unwrap(); - pool.add(create_test_credential("cred-3")).unwrap(); - - let results = checker.check_all(&pool); - assert_eq!(results.len(), 3); - } - - #[test] - fn test_unhealthy_count() { - let checker = HealthChecker::with_defaults(); - let pool = CredentialPool::new(ProviderType::Kiro); - - pool.add(create_test_credential("cred-1")).unwrap(); - pool.add(create_test_credential("cred-2")).unwrap(); - pool.add(create_test_credential("cred-3")).unwrap(); - - // 标记一个为不健康 - pool.mark_unhealthy("cred-2", "test".to_string()).unwrap(); - - assert_eq!(checker.unhealthy_count(&pool), 1); - } - - #[test] - fn test_should_mark_unhealthy() { - let checker = HealthChecker::with_defaults(); - - assert!(!checker.should_mark_unhealthy(0)); - assert!(!checker.should_mark_unhealthy(1)); - assert!(!checker.should_mark_unhealthy(2)); - assert!(checker.should_mark_unhealthy(3)); - assert!(checker.should_mark_unhealthy(4)); - } - - #[test] - fn test_recover_all() { - let checker = HealthChecker::with_defaults(); - let pool = CredentialPool::new(ProviderType::Kiro); - - pool.add(create_test_credential("cred-1")).unwrap(); - pool.add(create_test_credential("cred-2")).unwrap(); - pool.add(create_test_credential("cred-3")).unwrap(); - - // 标记两个为不健康 - pool.mark_unhealthy("cred-1", "test".to_string()).unwrap(); - pool.mark_unhealthy("cred-3", "test".to_string()).unwrap(); - - let recovered = checker.recover_all(&pool); - assert_eq!(recovered.len(), 2); - assert!(recovered.contains(&"cred-1".to_string())); - assert!(recovered.contains(&"cred-3".to_string())); - - // 验证所有凭证都是活跃状态 - for cred in pool.all() { - assert!(matches!(cred.status, CredentialStatus::Active)); - } - } -} diff --git a/src-tauri/crates/core/src/credential/mod.rs b/src-tauri/crates/core/src/credential/mod.rs deleted file mode 100644 index 8e5cc8eb0..000000000 --- a/src-tauri/crates/core/src/credential/mod.rs +++ /dev/null @@ -1,15 +0,0 @@ -//! 凭证池核心类型和独立逻辑 -//! -//! 包含凭证类型定义、凭证池管理、健康检查和风控模块。 -//! 负载均衡器(balancer)、配额管理(quota)和同步服务(sync) -//! 因依赖 infra crate 保留在主 crate 中。 - -pub mod health; -pub mod pool; -pub mod risk; -pub mod types; - -pub use health::{HealthCheckConfig, HealthCheckResult, HealthChecker, HealthStatus}; -pub use pool::{CredentialPool, PoolError, PoolStatus}; -pub use risk::{CooldownConfig, RateLimitEvent, RateLimitStats, RiskController, RiskLevel}; -pub use types::{Credential, CredentialData, CredentialStats, CredentialStatus}; diff --git a/src-tauri/crates/core/src/credential/pool.rs b/src-tauri/crates/core/src/credential/pool.rs deleted file mode 100644 index b45066b47..000000000 --- a/src-tauri/crates/core/src/credential/pool.rs +++ /dev/null @@ -1,416 +0,0 @@ -//! 凭证池实现 -//! -//! 使用 DashMap 实现线程安全的凭证池管理 - -use super::types::{Credential, CredentialStatus}; -use crate::ProviderType; -use chrono::{DateTime, Duration, Utc}; -use dashmap::DashMap; -use serde::{Deserialize, Serialize}; -use std::sync::atomic::{AtomicUsize, Ordering}; - -/// 凭证池 - 管理同一 Provider 的多个凭证 -pub struct CredentialPool { - /// 所属 Provider 类型 - provider: ProviderType, - /// 凭证存储(id -> Credential) - pub credentials: DashMap, - /// 轮询索引(用于负载均衡) - round_robin_index: AtomicUsize, -} - -/// 凭证池状态 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct PoolStatus { - /// Provider 类型 - pub provider: ProviderType, - /// 总凭证数 - pub total: usize, - /// 活跃凭证数 - pub active: usize, - /// 冷却中凭证数 - pub cooldown: usize, - /// 不健康凭证数 - pub unhealthy: usize, - /// 已禁用凭证数 - pub disabled: usize, -} - -/// 凭证池错误 -#[derive(Debug, Clone, PartialEq)] -pub enum PoolError { - /// 凭证已存在 - CredentialExists(String), - /// 凭证不存在 - CredentialNotFound(String), - /// 凭证池为空 - EmptyPool, - /// 所有凭证不可用 - NoAvailableCredential, -} - -impl std::fmt::Display for PoolError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - PoolError::CredentialExists(id) => write!(f, "凭证已存在: {id}"), - PoolError::CredentialNotFound(id) => write!(f, "凭证不存在: {id}"), - PoolError::EmptyPool => write!(f, "凭证池为空"), - PoolError::NoAvailableCredential => write!(f, "没有可用的凭证"), - } - } -} - -impl std::error::Error for PoolError {} - -impl CredentialPool { - /// 创建新的凭证池 - pub fn new(provider: ProviderType) -> Self { - Self { - provider, - credentials: DashMap::new(), - round_robin_index: AtomicUsize::new(0), - } - } - - /// 获取 Provider 类型 - pub fn provider(&self) -> ProviderType { - self.provider - } - - /// 获取凭证池大小 - pub fn len(&self) -> usize { - self.credentials.len() - } - - /// 检查凭证池是否为空 - pub fn is_empty(&self) -> bool { - self.credentials.is_empty() - } - - /// 添加凭证到池中 - /// - /// # 错误 - /// - 如果凭证 ID 已存在,返回 `PoolError::CredentialExists` - pub fn add(&self, credential: Credential) -> Result<(), PoolError> { - if self.credentials.contains_key(&credential.id) { - return Err(PoolError::CredentialExists(credential.id.clone())); - } - self.credentials.insert(credential.id.clone(), credential); - Ok(()) - } - - /// 从池中移除凭证 - /// - /// # 错误 - /// - 如果凭证不存在,返回 `PoolError::CredentialNotFound` - pub fn remove(&self, id: &str) -> Result { - self.credentials - .remove(id) - .map(|(_, cred)| cred) - .ok_or_else(|| PoolError::CredentialNotFound(id.to_string())) - } - - /// 获取凭证(只读) - pub fn get(&self, id: &str) -> Option { - self.credentials.get(id).map(|r| r.value().clone()) - } - - /// 获取所有凭证 ID - pub fn ids(&self) -> Vec { - self.credentials.iter().map(|r| r.key().clone()).collect() - } - - /// 获取所有凭证 - pub fn all(&self) -> Vec { - self.credentials.iter().map(|r| r.value().clone()).collect() - } - - /// 检查凭证是否存在 - pub fn contains(&self, id: &str) -> bool { - self.credentials.contains_key(id) - } - - /// 获取池状态 - pub fn status(&self) -> PoolStatus { - let mut active = 0; - let mut cooldown = 0; - let mut unhealthy = 0; - let mut disabled = 0; - - for entry in self.credentials.iter() { - match &entry.value().status { - CredentialStatus::Active => active += 1, - CredentialStatus::Cooldown { .. } => cooldown += 1, - CredentialStatus::Unhealthy { .. } => unhealthy += 1, - CredentialStatus::Disabled => disabled += 1, - } - } - - PoolStatus { - provider: self.provider, - total: self.credentials.len(), - active, - cooldown, - unhealthy, - disabled, - } - } - - /// 获取活跃凭证数量 - pub fn active_count(&self) -> usize { - self.credentials - .iter() - .filter(|r| r.value().is_available()) - .count() - } - - /// 标记凭证为冷却状态 - pub fn mark_cooldown(&self, id: &str, duration: Duration) -> Result<(), PoolError> { - let mut entry = self - .credentials - .get_mut(id) - .ok_or_else(|| PoolError::CredentialNotFound(id.to_string()))?; - - entry.status = CredentialStatus::Cooldown { - until: Utc::now() + duration, - }; - Ok(()) - } - - /// 标记凭证为不健康状态 - pub fn mark_unhealthy(&self, id: &str, reason: String) -> Result<(), PoolError> { - let mut entry = self - .credentials - .get_mut(id) - .ok_or_else(|| PoolError::CredentialNotFound(id.to_string()))?; - - entry.status = CredentialStatus::Unhealthy { reason }; - Ok(()) - } - - /// 恢复凭证为活跃状态 - pub fn mark_active(&self, id: &str) -> Result<(), PoolError> { - let mut entry = self - .credentials - .get_mut(id) - .ok_or_else(|| PoolError::CredentialNotFound(id.to_string()))?; - - entry.status = CredentialStatus::Active; - Ok(()) - } - - /// 更新过期的冷却状态 - /// 将冷却期已过的凭证恢复为活跃状态 - pub fn refresh_cooldowns(&self) { - let now = Utc::now(); - for mut entry in self.credentials.iter_mut() { - if let CredentialStatus::Cooldown { until } = &entry.status { - if *until <= now { - entry.status = CredentialStatus::Active; - } - } - } - } - - /// 获取下一个可用凭证(轮询策略) - /// - /// # 错误 - /// - 如果池为空,返回 `PoolError::EmptyPool` - /// - 如果没有可用凭证,返回 `PoolError::NoAvailableCredential` - pub fn next_available(&self) -> Result { - if self.credentials.is_empty() { - return Err(PoolError::EmptyPool); - } - - // 先刷新冷却状态 - self.refresh_cooldowns(); - - // 收集所有活跃凭证 - let active_creds: Vec<_> = self - .credentials - .iter() - .filter(|r| r.value().is_available()) - .map(|r| r.value().clone()) - .collect(); - - if active_creds.is_empty() { - return Err(PoolError::NoAvailableCredential); - } - - // 轮询选择 - let index = self.round_robin_index.fetch_add(1, Ordering::SeqCst) % active_creds.len(); - Ok(active_creds[index].clone()) - } - - /// 获取最早恢复时间(当所有凭证都在冷却时) - pub fn earliest_recovery(&self) -> Option> { - self.credentials - .iter() - .filter_map(|r| { - if let CredentialStatus::Cooldown { until } = &r.value().status { - Some(*until) - } else { - None - } - }) - .min() - } - - /// 记录凭证使用成功 - pub fn record_success(&self, id: &str, latency_ms: u64) -> Result<(), PoolError> { - let mut entry = self - .credentials - .get_mut(id) - .ok_or_else(|| PoolError::CredentialNotFound(id.to_string()))?; - - entry.mark_used(); - entry.stats.record_success(latency_ms); - Ok(()) - } - - /// 记录凭证使用失败 - pub fn record_failure(&self, id: &str) -> Result<(), PoolError> { - let mut entry = self - .credentials - .get_mut(id) - .ok_or_else(|| PoolError::CredentialNotFound(id.to_string()))?; - - entry.mark_used(); - entry.stats.record_failure(); - Ok(()) - } -} - -#[cfg(test)] -mod pool_tests { - use super::*; - use crate::credential::CredentialData; - - fn create_test_credential(id: &str) -> Credential { - Credential::new( - id.to_string(), - ProviderType::Kiro, - CredentialData::ApiKey { - key: format!("key-{id}"), - base_url: None, - }, - ) - } - - #[test] - fn test_pool_new() { - let pool = CredentialPool::new(ProviderType::Kiro); - assert_eq!(pool.provider(), ProviderType::Kiro); - assert!(pool.is_empty()); - assert_eq!(pool.len(), 0); - } - - #[test] - fn test_pool_add() { - let pool = CredentialPool::new(ProviderType::Kiro); - let cred = create_test_credential("test-1"); - - assert!(pool.add(cred.clone()).is_ok()); - assert_eq!(pool.len(), 1); - assert!(pool.contains("test-1")); - - // 重复添加应失败 - let result = pool.add(cred); - assert!(matches!(result, Err(PoolError::CredentialExists(_)))); - } - - #[test] - fn test_pool_remove() { - let pool = CredentialPool::new(ProviderType::Kiro); - let cred = create_test_credential("test-1"); - - pool.add(cred).unwrap(); - assert_eq!(pool.len(), 1); - - let removed = pool.remove("test-1").unwrap(); - assert_eq!(removed.id, "test-1"); - assert!(pool.is_empty()); - - // 移除不存在的凭证应失败 - let result = pool.remove("test-1"); - assert!(matches!(result, Err(PoolError::CredentialNotFound(_)))); - } - - #[test] - fn test_pool_get() { - let pool = CredentialPool::new(ProviderType::Kiro); - let cred = create_test_credential("test-1"); - - pool.add(cred).unwrap(); - - let retrieved = pool.get("test-1"); - assert!(retrieved.is_some()); - assert_eq!(retrieved.unwrap().id, "test-1"); - - assert!(pool.get("nonexistent").is_none()); - } - - #[test] - fn test_pool_status() { - let pool = CredentialPool::new(ProviderType::Kiro); - - pool.add(create_test_credential("active-1")).unwrap(); - pool.add(create_test_credential("active-2")).unwrap(); - pool.add(create_test_credential("cooldown-1")).unwrap(); - pool.add(create_test_credential("unhealthy-1")).unwrap(); - - pool.mark_cooldown("cooldown-1", Duration::hours(1)) - .unwrap(); - pool.mark_unhealthy("unhealthy-1", "test reason".to_string()) - .unwrap(); - - let status = pool.status(); - assert_eq!(status.total, 4); - assert_eq!(status.active, 2); - assert_eq!(status.cooldown, 1); - assert_eq!(status.unhealthy, 1); - assert_eq!(status.disabled, 0); - } - - #[test] - fn test_pool_next_available_empty() { - let pool = CredentialPool::new(ProviderType::Kiro); - let result = pool.next_available(); - assert!(matches!(result, Err(PoolError::EmptyPool))); - } - - #[test] - fn test_pool_next_available_all_cooldown() { - let pool = CredentialPool::new(ProviderType::Kiro); - pool.add(create_test_credential("cred-1")).unwrap(); - pool.mark_cooldown("cred-1", Duration::hours(1)).unwrap(); - - let result = pool.next_available(); - assert!(matches!(result, Err(PoolError::NoAvailableCredential))); - } - - #[test] - fn test_pool_record_success() { - let pool = CredentialPool::new(ProviderType::Kiro); - pool.add(create_test_credential("test-1")).unwrap(); - - pool.record_success("test-1", 100).unwrap(); - - let cred = pool.get("test-1").unwrap(); - assert_eq!(cred.stats.total_requests, 1); - assert_eq!(cred.stats.successful_requests, 1); - assert!(cred.last_used.is_some()); - } - - #[test] - fn test_pool_record_failure() { - let pool = CredentialPool::new(ProviderType::Kiro); - pool.add(create_test_credential("test-1")).unwrap(); - - pool.record_failure("test-1").unwrap(); - - let cred = pool.get("test-1").unwrap(); - assert_eq!(cred.stats.total_requests, 1); - assert_eq!(cred.stats.consecutive_failures, 1); - } -} diff --git a/src-tauri/crates/core/src/credential/risk.rs b/src-tauri/crates/core/src/credential/risk.rs deleted file mode 100644 index c5e433dbe..000000000 --- a/src-tauri/crates/core/src/credential/risk.rs +++ /dev/null @@ -1,586 +0,0 @@ -//! 风控模块 -//! -//! 提供限流检测、冷却期管理和风险评估功能。 -//! -//! ## 功能 -//! -//! - **限流检测**: 检测 API 返回的限流错误(429、rate limit) -//! - **冷却期管理**: 自动计算和管理凭证冷却时间 -//! - **风险评估**: 根据历史数据评估凭证风险等级 - -use chrono::{DateTime, Duration, Utc}; -use dashmap::DashMap; -use serde::{Deserialize, Serialize}; -use std::collections::VecDeque; -use std::sync::atomic::{AtomicU64, Ordering}; - -/// 风险等级 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum RiskLevel { - /// 低风险 - 正常使用 - Low, - /// 中风险 - 接近限流阈值 - Medium, - /// 高风险 - 频繁触发限流 - High, - /// 危险 - 需要立即冷却 - Critical, -} - -impl RiskLevel { - /// 获取风险等级对应的冷却时间倍数 - pub fn cooldown_multiplier(&self) -> f64 { - match self { - RiskLevel::Low => 1.0, - RiskLevel::Medium => 1.5, - RiskLevel::High => 2.0, - RiskLevel::Critical => 3.0, - } - } - - /// 获取风险等级描述 - pub fn description(&self) -> &'static str { - match self { - RiskLevel::Low => "正常", - RiskLevel::Medium => "接近限流", - RiskLevel::High => "频繁限流", - RiskLevel::Critical => "需要冷却", - } - } -} - -/// 限流事件 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RateLimitEvent { - /// 凭证 ID - pub credential_id: String, - /// 事件时间 - pub timestamp: DateTime, - /// HTTP 状态码 - pub status_code: Option, - /// 错误消息 - pub error_message: Option, - /// 建议的重试时间(秒) - pub retry_after_secs: Option, -} - -impl RateLimitEvent { - /// 创建新的限流事件 - pub fn new(credential_id: String) -> Self { - Self { - credential_id, - timestamp: Utc::now(), - status_code: None, - error_message: None, - retry_after_secs: None, - } - } - - /// 设置状态码 - pub fn with_status_code(mut self, code: u16) -> Self { - self.status_code = Some(code); - self - } - - /// 设置错误消息 - pub fn with_error_message(mut self, message: String) -> Self { - self.error_message = Some(message); - self - } - - /// 设置重试时间 - pub fn with_retry_after(mut self, secs: u64) -> Self { - self.retry_after_secs = Some(secs); - self - } -} - -/// 冷却配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CooldownConfig { - /// 基础冷却时间(秒) - pub base_cooldown_secs: u64, - /// 最大冷却时间(秒) - pub max_cooldown_secs: u64, - /// 冷却时间增长因子(指数退避) - pub backoff_factor: f64, - /// 限流事件窗口大小(保留最近 N 个事件) - pub event_window_size: usize, - /// 限流事件时间窗口(秒)- 只统计此时间内的事件 - pub event_time_window_secs: u64, - /// 触发中风险的限流次数阈值 - pub medium_risk_threshold: u32, - /// 触发高风险的限流次数阈值 - pub high_risk_threshold: u32, - /// 触发危险的限流次数阈值 - pub critical_risk_threshold: u32, -} - -impl Default for CooldownConfig { - fn default() -> Self { - Self { - base_cooldown_secs: 60, // 1 分钟 - max_cooldown_secs: 3600, // 1 小时 - backoff_factor: 2.0, // 指数退避因子 - event_window_size: 100, // 保留最近 100 个事件 - event_time_window_secs: 3600, // 1 小时内的事件 - medium_risk_threshold: 3, // 3 次限流 -> 中风险 - high_risk_threshold: 5, // 5 次限流 -> 高风险 - critical_risk_threshold: 10, // 10 次限流 -> 危险 - } - } -} - -/// 凭证风控状态 -#[derive(Debug)] -struct CredentialRiskState { - /// 限流事件历史 - events: VecDeque, - /// 连续限流次数 - consecutive_rate_limits: AtomicU64, - /// 当前冷却结束时间 - cooldown_until: Option>, - /// 上次限流时间 - last_rate_limit: Option>, -} - -impl CredentialRiskState { - fn new() -> Self { - Self { - events: VecDeque::new(), - consecutive_rate_limits: AtomicU64::new(0), - cooldown_until: None, - last_rate_limit: None, - } - } -} - -/// 风控控制器 -/// -/// 管理凭证的限流检测和冷却期 -pub struct RiskController { - /// 配置 - config: CooldownConfig, - /// 各凭证的风控状态 - states: DashMap, -} - -impl RiskController { - /// 创建新的风控控制器 - pub fn new(config: CooldownConfig) -> Self { - Self { - config, - states: DashMap::new(), - } - } - - /// 使用默认配置创建 - pub fn with_defaults() -> Self { - Self::new(CooldownConfig::default()) - } - - /// 获取配置 - pub fn config(&self) -> &CooldownConfig { - &self.config - } - - /// 记录限流事件 - /// - /// # 返回 - /// 建议的冷却时间(秒) - pub fn record_rate_limit(&self, event: RateLimitEvent) -> u64 { - let credential_id = event.credential_id.clone(); - let retry_after = event.retry_after_secs; - - let mut state = self - .states - .entry(credential_id.clone()) - .or_insert_with(CredentialRiskState::new); - - // 更新连续限流次数 - state.consecutive_rate_limits.fetch_add(1, Ordering::SeqCst); - state.last_rate_limit = Some(Utc::now()); - - // 添加事件到历史 - state.events.push_back(event); - - // 清理过期事件 - self.cleanup_old_events(&mut state); - - // 计算冷却时间 - let cooldown_secs = self.calculate_cooldown(&state, retry_after); - - // 设置冷却结束时间 - state.cooldown_until = Some(Utc::now() + Duration::seconds(cooldown_secs as i64)); - - cooldown_secs - } - - /// 记录成功请求(重置连续限流计数) - pub fn record_success(&self, credential_id: &str) { - if let Some(state) = self.states.get_mut(credential_id) { - state.consecutive_rate_limits.store(0, Ordering::SeqCst); - } - } - - /// 获取凭证的风险等级 - pub fn get_risk_level(&self, credential_id: &str) -> RiskLevel { - let state = match self.states.get(credential_id) { - Some(s) => s, - None => return RiskLevel::Low, - }; - - let recent_count = self.count_recent_events(&state); - - if recent_count >= self.config.critical_risk_threshold { - RiskLevel::Critical - } else if recent_count >= self.config.high_risk_threshold { - RiskLevel::High - } else if recent_count >= self.config.medium_risk_threshold { - RiskLevel::Medium - } else { - RiskLevel::Low - } - } - - /// 检查凭证是否在冷却中 - pub fn is_in_cooldown(&self, credential_id: &str) -> bool { - self.states - .get(credential_id) - .and_then(|state| state.cooldown_until) - .map(|until| Utc::now() < until) - .unwrap_or(false) - } - - /// 获取凭证的冷却结束时间 - pub fn get_cooldown_until(&self, credential_id: &str) -> Option> { - self.states - .get(credential_id) - .and_then(|state| state.cooldown_until) - .filter(|until| Utc::now() < *until) - } - - /// 获取凭证的剩余冷却时间(秒) - pub fn get_remaining_cooldown_secs(&self, credential_id: &str) -> Option { - self.get_cooldown_until(credential_id).map(|until| { - let remaining = until - Utc::now(); - remaining.num_seconds().max(0) as u64 - }) - } - - /// 手动清除凭证的冷却状态 - pub fn clear_cooldown(&self, credential_id: &str) { - if let Some(mut state) = self.states.get_mut(credential_id) { - state.cooldown_until = None; - state.consecutive_rate_limits.store(0, Ordering::SeqCst); - } - } - - /// 获取所有处于冷却中的凭证 ID - pub fn get_cooling_credentials(&self) -> Vec { - let now = Utc::now(); - self.states - .iter() - .filter(|entry| { - entry - .value() - .cooldown_until - .map(|until| now < until) - .unwrap_or(false) - }) - .map(|entry| entry.key().clone()) - .collect() - } - - /// 获取凭证的限流事件统计 - pub fn get_event_stats(&self, credential_id: &str) -> Option { - self.states.get(credential_id).map(|state| { - let recent_count = self.count_recent_events(&state); - let consecutive = state.consecutive_rate_limits.load(Ordering::SeqCst); - - RateLimitStats { - total_events: state.events.len(), - recent_events: recent_count as usize, - consecutive_rate_limits: consecutive, - last_rate_limit: state.last_rate_limit, - cooldown_until: state.cooldown_until, - risk_level: self.get_risk_level(credential_id), - } - }) - } - - /// 检测响应是否为限流错误 - pub fn is_rate_limit_error(status_code: u16, body: Option<&str>) -> bool { - // HTTP 429 Too Many Requests - if status_code == 429 { - return true; - } - - // 检查响应体中的限流关键词 - if let Some(body) = body { - let body_lower = body.to_lowercase(); - if body_lower.contains("rate limit") - || body_lower.contains("rate_limit") - || body_lower.contains("ratelimit") - || body_lower.contains("too many requests") - || body_lower.contains("quota exceeded") - || body_lower.contains("resource_exhausted") - { - return true; - } - } - - false - } - - /// 从响应头解析 Retry-After - pub fn parse_retry_after(header_value: &str) -> Option { - // 尝试解析为秒数 - if let Ok(secs) = header_value.parse::() { - return Some(secs); - } - - // 尝试解析为 HTTP 日期格式 - if let Ok(date) = DateTime::parse_from_rfc2822(header_value) { - let until = date.with_timezone(&Utc); - let now = Utc::now(); - if until > now { - return Some((until - now).num_seconds() as u64); - } - } - - None - } - - /// 清理过期事件 - fn cleanup_old_events(&self, state: &mut CredentialRiskState) { - let cutoff = Utc::now() - Duration::seconds(self.config.event_time_window_secs as i64); - - // 移除过期事件 - while let Some(front) = state.events.front() { - if front.timestamp < cutoff { - state.events.pop_front(); - } else { - break; - } - } - - // 限制事件数量 - while state.events.len() > self.config.event_window_size { - state.events.pop_front(); - } - } - - /// 统计最近的限流事件数 - fn count_recent_events(&self, state: &CredentialRiskState) -> u32 { - let cutoff = Utc::now() - Duration::seconds(self.config.event_time_window_secs as i64); - state - .events - .iter() - .filter(|e| e.timestamp >= cutoff) - .count() as u32 - } - - /// 计算冷却时间 - fn calculate_cooldown(&self, state: &CredentialRiskState, retry_after: Option) -> u64 { - // 如果有 Retry-After,优先使用 - if let Some(retry) = retry_after { - return retry.min(self.config.max_cooldown_secs); - } - - // 使用指数退避计算冷却时间 - let consecutive = state.consecutive_rate_limits.load(Ordering::SeqCst); - let base = self.config.base_cooldown_secs as f64; - let factor = self.config.backoff_factor; - - // cooldown = base * factor^(consecutive - 1) - let cooldown = if consecutive > 0 { - base * factor.powi((consecutive - 1) as i32) - } else { - base - }; - - // 根据风险等级调整 - let risk_level = self.get_risk_level_from_state(state); - let adjusted = cooldown * risk_level.cooldown_multiplier(); - - // 限制在最大值内 - (adjusted as u64).min(self.config.max_cooldown_secs) - } - - /// 从状态计算风险等级 - fn get_risk_level_from_state(&self, state: &CredentialRiskState) -> RiskLevel { - let recent_count = self.count_recent_events(state); - - if recent_count >= self.config.critical_risk_threshold { - RiskLevel::Critical - } else if recent_count >= self.config.high_risk_threshold { - RiskLevel::High - } else if recent_count >= self.config.medium_risk_threshold { - RiskLevel::Medium - } else { - RiskLevel::Low - } - } -} - -impl Default for RiskController { - fn default() -> Self { - Self::with_defaults() - } -} - -/// 限流事件统计 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RateLimitStats { - /// 总事件数 - pub total_events: usize, - /// 最近事件数(时间窗口内) - pub recent_events: usize, - /// 连续限流次数 - pub consecutive_rate_limits: u64, - /// 上次限流时间 - pub last_rate_limit: Option>, - /// 冷却结束时间 - pub cooldown_until: Option>, - /// 风险等级 - pub risk_level: RiskLevel, -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_risk_controller_new() { - let controller = RiskController::with_defaults(); - assert_eq!(controller.config().base_cooldown_secs, 60); - } - - #[test] - fn test_record_rate_limit() { - let controller = RiskController::with_defaults(); - let event = RateLimitEvent::new("cred-1".to_string()).with_status_code(429); - - let cooldown = controller.record_rate_limit(event); - assert!(cooldown >= 60); // 至少是基础冷却时间 - - assert!(controller.is_in_cooldown("cred-1")); - assert_eq!(controller.get_risk_level("cred-1"), RiskLevel::Low); - } - - #[test] - fn test_risk_level_escalation() { - let controller = RiskController::with_defaults(); - - // 记录多次限流事件 - for i in 0..5 { - let event = RateLimitEvent::new("cred-1".to_string()) - .with_status_code(429) - .with_error_message(format!("Rate limit {i}")); - controller.record_rate_limit(event); - } - - // 应该达到高风险 - assert_eq!(controller.get_risk_level("cred-1"), RiskLevel::High); - } - - #[test] - fn test_record_success_resets_consecutive() { - let controller = RiskController::with_defaults(); - - // 记录限流 - let event = RateLimitEvent::new("cred-1".to_string()); - controller.record_rate_limit(event); - - // 记录成功 - controller.record_success("cred-1"); - - // 连续计数应该重置 - let stats = controller.get_event_stats("cred-1").unwrap(); - assert_eq!(stats.consecutive_rate_limits, 0); - } - - #[test] - fn test_clear_cooldown() { - let controller = RiskController::with_defaults(); - - let event = RateLimitEvent::new("cred-1".to_string()); - controller.record_rate_limit(event); - - assert!(controller.is_in_cooldown("cred-1")); - - controller.clear_cooldown("cred-1"); - - assert!(!controller.is_in_cooldown("cred-1")); - } - - #[test] - fn test_is_rate_limit_error() { - assert!(RiskController::is_rate_limit_error(429, None)); - assert!(RiskController::is_rate_limit_error( - 200, - Some("rate limit exceeded") - )); - assert!(RiskController::is_rate_limit_error( - 500, - Some("RESOURCE_EXHAUSTED") - )); - assert!(!RiskController::is_rate_limit_error(200, Some("success"))); - } - - #[test] - fn test_parse_retry_after() { - assert_eq!(RiskController::parse_retry_after("60"), Some(60)); - assert_eq!(RiskController::parse_retry_after("3600"), Some(3600)); - assert!(RiskController::parse_retry_after("invalid").is_none()); - } - - #[test] - fn test_retry_after_priority() { - let controller = RiskController::with_defaults(); - - // 使用 retry_after 的事件 - let event = RateLimitEvent::new("cred-1".to_string()).with_retry_after(120); - - let cooldown = controller.record_rate_limit(event); - assert_eq!(cooldown, 120); // 应该使用 retry_after 的值 - } - - #[test] - fn test_exponential_backoff() { - let controller = RiskController::with_defaults(); - - // 第一次限流 - let event1 = RateLimitEvent::new("cred-1".to_string()); - let cooldown1 = controller.record_rate_limit(event1); - - // 第二次限流(应该更长) - let event2 = RateLimitEvent::new("cred-1".to_string()); - let cooldown2 = controller.record_rate_limit(event2); - - assert!(cooldown2 > cooldown1); - } - - #[test] - fn test_get_cooling_credentials() { - let controller = RiskController::with_defaults(); - - controller.record_rate_limit(RateLimitEvent::new("cred-1".to_string())); - controller.record_rate_limit(RateLimitEvent::new("cred-2".to_string())); - - let cooling = controller.get_cooling_credentials(); - assert_eq!(cooling.len(), 2); - assert!(cooling.contains(&"cred-1".to_string())); - assert!(cooling.contains(&"cred-2".to_string())); - } - - #[test] - fn test_risk_level_cooldown_multiplier() { - assert_eq!(RiskLevel::Low.cooldown_multiplier(), 1.0); - assert_eq!(RiskLevel::Medium.cooldown_multiplier(), 1.5); - assert_eq!(RiskLevel::High.cooldown_multiplier(), 2.0); - assert_eq!(RiskLevel::Critical.cooldown_multiplier(), 3.0); - } -} diff --git a/src-tauri/crates/core/src/credential/types.rs b/src-tauri/crates/core/src/credential/types.rs deleted file mode 100644 index c7c49f15d..000000000 --- a/src-tauri/crates/core/src/credential/types.rs +++ /dev/null @@ -1,243 +0,0 @@ -//! 凭证相关类型定义 -//! -//! 定义凭证、凭证数据、凭证状态等核心类型 - -use crate::ProviderType; -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; - -/// 凭证 - 表示单个 API 凭证 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -pub struct Credential { - /// 唯一标识符 - pub id: String, - /// 所属 Provider 类型 - pub provider: ProviderType, - /// 凭证数据 - pub data: CredentialData, - /// 创建时间 - pub created_at: DateTime, - /// 最后使用时间 - pub last_used: Option>, - /// 当前状态 - pub status: CredentialStatus, - /// 统计信息 - pub stats: CredentialStats, - /// Per-Key 代理 URL(覆盖全局代理) - #[serde(default, skip_serializing_if = "Option::is_none")] - pub proxy_url: Option, -} - -impl Credential { - /// 创建新凭证 - pub fn new(id: String, provider: ProviderType, data: CredentialData) -> Self { - Self { - id, - provider, - data, - created_at: Utc::now(), - last_used: None, - status: CredentialStatus::Active, - stats: CredentialStats::default(), - proxy_url: None, - } - } - - /// 创建带代理的凭证 - pub fn with_proxy(mut self, proxy_url: Option) -> Self { - self.proxy_url = proxy_url; - self - } - - /// 设置代理 URL - pub fn set_proxy_url(&mut self, proxy_url: Option) { - self.proxy_url = proxy_url; - } - - /// 获取代理 URL - pub fn proxy_url(&self) -> Option<&str> { - self.proxy_url.as_deref() - } - - /// 检查凭证是否可用(活跃状态) - pub fn is_available(&self) -> bool { - matches!(self.status, CredentialStatus::Active) - } - - /// 更新最后使用时间 - pub fn mark_used(&mut self) { - self.last_used = Some(Utc::now()); - } -} - -/// 凭证数据 - 不同 Provider 有不同的凭证格式 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum CredentialData { - /// OAuth 凭证(用于 Kiro、Gemini、Qwen 等) - OAuth { - access_token: String, - refresh_token: Option, - expires_at: Option>, - }, - /// API Key 凭证(用于 OpenAI、Claude 等) - ApiKey { - key: String, - base_url: Option, - }, -} - -/// 凭证状态 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -#[serde(tag = "status", rename_all = "snake_case")] -pub enum CredentialStatus { - /// 活跃可用 - Active, - /// 冷却中(配额超限等) - Cooldown { - /// 冷却结束时间 - until: DateTime, - }, - /// 不健康(连续失败) - Unhealthy { - /// 不健康原因 - reason: String, - }, - /// 已禁用 - Disabled, -} - -/// 凭证统计信息 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)] -pub struct CredentialStats { - /// 总请求数 - pub total_requests: u64, - /// 成功请求数 - pub successful_requests: u64, - /// 连续失败次数 - pub consecutive_failures: u32, - /// 平均延迟(毫秒) - pub avg_latency_ms: f64, -} - -impl CredentialStats { - /// 记录成功请求 - pub fn record_success(&mut self, latency_ms: u64) { - self.total_requests += 1; - self.successful_requests += 1; - self.consecutive_failures = 0; - - // 更新平均延迟(移动平均) - let n = self.successful_requests as f64; - self.avg_latency_ms = self.avg_latency_ms * (n - 1.0) / n + latency_ms as f64 / n; - } - - /// 记录失败请求 - pub fn record_failure(&mut self) { - self.total_requests += 1; - self.consecutive_failures += 1; - } - - /// 获取成功率 - pub fn success_rate(&self) -> f64 { - if self.total_requests == 0 { - 1.0 - } else { - self.successful_requests as f64 / self.total_requests as f64 - } - } -} - -#[cfg(test)] -mod type_tests { - use super::*; - - #[test] - fn test_credential_new() { - let cred = Credential::new( - "test-id".to_string(), - ProviderType::Kiro, - CredentialData::ApiKey { - key: "test-key".to_string(), - base_url: None, - }, - ); - - assert_eq!(cred.id, "test-id"); - assert_eq!(cred.provider, ProviderType::Kiro); - assert!(cred.is_available()); - assert!(cred.last_used.is_none()); - } - - #[test] - fn test_credential_is_available() { - let mut cred = Credential::new( - "test".to_string(), - ProviderType::Gemini, - CredentialData::ApiKey { - key: "key".to_string(), - base_url: None, - }, - ); - - assert!(cred.is_available()); - - cred.status = CredentialStatus::Cooldown { - until: Utc::now() + chrono::Duration::hours(1), - }; - assert!(!cred.is_available()); - - cred.status = CredentialStatus::Unhealthy { - reason: "test".to_string(), - }; - assert!(!cred.is_available()); - - cred.status = CredentialStatus::Disabled; - assert!(!cred.is_available()); - } - - #[test] - fn test_credential_stats_success() { - let mut stats = CredentialStats::default(); - - stats.record_success(100); - assert_eq!(stats.total_requests, 1); - assert_eq!(stats.successful_requests, 1); - assert_eq!(stats.consecutive_failures, 0); - assert!((stats.avg_latency_ms - 100.0).abs() < 0.001); - - stats.record_success(200); - assert_eq!(stats.total_requests, 2); - assert!((stats.avg_latency_ms - 150.0).abs() < 0.001); - } - - #[test] - fn test_credential_stats_failure() { - let mut stats = CredentialStats::default(); - - stats.record_failure(); - assert_eq!(stats.total_requests, 1); - assert_eq!(stats.successful_requests, 0); - assert_eq!(stats.consecutive_failures, 1); - - stats.record_failure(); - assert_eq!(stats.consecutive_failures, 2); - - stats.record_success(100); - assert_eq!(stats.consecutive_failures, 0); - } - - #[test] - fn test_credential_stats_success_rate() { - let mut stats = CredentialStats::default(); - - // 空统计应返回 1.0 - assert!((stats.success_rate() - 1.0).abs() < 0.001); - - stats.record_success(100); - assert!((stats.success_rate() - 1.0).abs() < 0.001); - - stats.record_failure(); - assert!((stats.success_rate() - 0.5).abs() < 0.001); - } -} diff --git a/src-tauri/crates/core/src/database/README.md b/src-tauri/crates/core/src/database/README.md index 9a502d818..2b1ca2e02 100644 --- a/src-tauri/crates/core/src/database/README.md +++ b/src-tauri/crates/core/src/database/README.md @@ -18,8 +18,8 @@ ### 核心表 - `api_key_providers` - API Key Provider 配置 -- `api_keys` - API Key 条目(已迁移到 provider_pool_credentials) -- `provider_pool_credentials` - 凭证池(统一管理所有凭证) +- `api_key_providers` - API Key Provider 配置主表 +- `provider_pool_credentials` - 旧凭证池表;启动期清空,仅保留历史迁移边界 - `providers` - Provider 配置 - `settings` - 应用设置 @@ -43,7 +43,7 @@ | `dao/api_key_provider.rs` | API Key Provider DAO | | `dao/mcp.rs` | MCP 服务器 DAO | | `dao/prompts.rs` | 提示词 DAO | -| `dao/provider_pool.rs` | 凭证池 DAO | +| `dao/provider_pool.rs` | 旧凭证池 DAO;不得作为运行时选择入口 | | `dao/providers.rs` | Provider DAO | | `dao/skills.rs` | 技能 DAO | @@ -51,7 +51,7 @@ ### API Keys 迁移 -`migrate_api_keys_to_pool()` 函数将 `api_keys` 表中的数据迁移到 `provider_pool_credentials` 表: +旧 `migrate_api_keys_to_pool()` 只属于历史迁移链。当前启动期会清理 `provider_pool_credentials`,运行时不再读取该表选择凭证。 - 根据 provider_type 自动转换为对应的 CredentialData 类型 - 保留使用统计和错误计数 diff --git a/src-tauri/crates/core/src/database/dao/mod.rs b/src-tauri/crates/core/src/database/dao/mod.rs index 9fbdf3981..3fd879e48 100644 --- a/src-tauri/crates/core/src/database/dao/mod.rs +++ b/src-tauri/crates/core/src/database/dao/mod.rs @@ -16,7 +16,6 @@ pub mod mcp; pub mod orchestrator; pub mod persona_dao; pub mod prompts; -pub mod provider_pool; pub mod providers; pub mod publish_config_dao; pub mod skills; diff --git a/src-tauri/crates/core/src/database/dao/provider_pool.rs b/src-tauri/crates/core/src/database/dao/provider_pool.rs deleted file mode 100644 index f0a820e58..000000000 --- a/src-tauri/crates/core/src/database/dao/provider_pool.rs +++ /dev/null @@ -1,504 +0,0 @@ -//! Provider Pool 数据访问对象 -//! -//! 提供凭证池的 CRUD 操作。 - -use crate::models::provider_pool_model::{ - CachedTokenInfo, CredentialData, CredentialSource, PoolProviderType, ProviderCredential, - ProviderPools, -}; -use chrono::{DateTime, TimeZone, Utc}; -use rusqlite::{params, Connection}; - -pub struct ProviderPoolDao; - -impl ProviderPoolDao { - /// 获取所有凭证 - pub fn get_all(conn: &Connection) -> Result, rusqlite::Error> { - let mut stmt = conn.prepare( - "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, - last_used, last_error_time, last_error_message, last_health_check_time, - last_health_check_model, created_at, updated_at, source, proxy_url - FROM provider_pool_credentials - ORDER BY provider_type, created_at ASC", - )?; - - let rows = stmt.query_map([], Self::row_to_credential)?; - - let mut credentials = Vec::new(); - for cred in rows.flatten() { - credentials.push(cred); - } - Ok(credentials) - } - - /// 获取指定类型的凭证 - pub fn get_by_type( - conn: &Connection, - provider_type: &PoolProviderType, - ) -> Result, rusqlite::Error> { - let mut stmt = conn.prepare( - "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, - last_used, last_error_time, last_error_message, last_health_check_time, - last_health_check_model, created_at, updated_at, source, proxy_url - FROM provider_pool_credentials - WHERE provider_type = ?1 - ORDER BY created_at ASC", - )?; - - let rows = stmt.query_map([provider_type.to_string()], |row| { - Self::row_to_credential(row) - })?; - - let mut credentials = Vec::new(); - for cred in rows.flatten() { - credentials.push(cred); - } - Ok(credentials) - } - - /// 获取指定 UUID 的凭证 - pub fn get_by_uuid( - conn: &Connection, - uuid: &str, - ) -> Result, rusqlite::Error> { - let mut stmt = conn.prepare( - "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, - last_used, last_error_time, last_error_message, last_health_check_time, - last_health_check_model, created_at, updated_at, source, proxy_url - FROM provider_pool_credentials - WHERE uuid = ?1", - )?; - - let mut rows = stmt.query([uuid])?; - if let Some(row) = rows.next()? { - Ok(Some(Self::row_to_credential(row)?)) - } else { - Ok(None) - } - } - - /// 根据名称获取凭证 - pub fn get_by_name( - conn: &Connection, - name: &str, - ) -> Result, rusqlite::Error> { - let mut stmt = conn.prepare( - "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, - last_used, last_error_time, last_error_message, last_health_check_time, - last_health_check_model, created_at, updated_at, source, proxy_url - FROM provider_pool_credentials - WHERE name = ?1", - )?; - - let mut rows = stmt.query([name])?; - if let Some(row) = rows.next()? { - Ok(Some(Self::row_to_credential(row)?)) - } else { - Ok(None) - } - } - - /// 获取所有凭证按类型分组 - pub fn get_grouped(conn: &Connection) -> Result { - let all = Self::get_all(conn)?; - let mut grouped: ProviderPools = std::collections::HashMap::new(); - - for cred in all { - grouped.entry(cred.provider_type).or_default().push(cred); - } - - Ok(grouped) - } - - /// 插入新凭证 - pub fn insert(conn: &Connection, cred: &ProviderCredential) -> Result<(), rusqlite::Error> { - let credential_json = - serde_json::to_string(&cred.credential).unwrap_or_else(|_| "{}".to_string()); - let not_supported_models_json = - serde_json::to_string(&cred.not_supported_models).unwrap_or_else(|_| "[]".to_string()); - let supported_models_json = - serde_json::to_string(&cred.supported_models).unwrap_or_else(|_| "[]".to_string()); - let source_str = match cred.source { - CredentialSource::Manual => "manual", - CredentialSource::Imported => "imported", - CredentialSource::Private => "private", - }; - - conn.execute( - "INSERT INTO provider_pool_credentials - (uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, - last_used, last_error_time, last_error_message, last_health_check_time, - last_health_check_model, created_at, updated_at, source, proxy_url) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21)", - params![ - cred.uuid, - cred.provider_type.to_string(), - credential_json, - cred.name, - cred.is_healthy, - cred.is_disabled, - cred.check_health, - cred.check_model_name, - not_supported_models_json, - supported_models_json, - cred.usage_count, - cred.error_count, - cred.last_used.map(|t| t.timestamp()), - cred.last_error_time.map(|t| t.timestamp()), - cred.last_error_message, - cred.last_health_check_time.map(|t| t.timestamp()), - cred.last_health_check_model, - cred.created_at.timestamp(), - cred.updated_at.timestamp(), - source_str, - cred.proxy_url, - ], - )?; - Ok(()) - } - - /// 更新凭证 - pub fn update(conn: &Connection, cred: &ProviderCredential) -> Result<(), rusqlite::Error> { - let credential_json = - serde_json::to_string(&cred.credential).unwrap_or_else(|_| "{}".to_string()); - let not_supported_models_json = - serde_json::to_string(&cred.not_supported_models).unwrap_or_else(|_| "[]".to_string()); - let supported_models_json = - serde_json::to_string(&cred.supported_models).unwrap_or_else(|_| "[]".to_string()); - - conn.execute( - "UPDATE provider_pool_credentials SET - provider_type = ?2, credential_data = ?3, name = ?4, is_healthy = ?5, - is_disabled = ?6, check_health = ?7, check_model_name = ?8, - not_supported_models = ?9, supported_models = ?10, usage_count = ?11, error_count = ?12, - last_used = ?13, last_error_time = ?14, last_error_message = ?15, - last_health_check_time = ?16, last_health_check_model = ?17, updated_at = ?18, proxy_url = ?19 - WHERE uuid = ?1", - params![ - cred.uuid, - cred.provider_type.to_string(), - credential_json, - cred.name, - cred.is_healthy, - cred.is_disabled, - cred.check_health, - cred.check_model_name, - not_supported_models_json, - supported_models_json, - cred.usage_count, - cred.error_count, - cred.last_used.map(|t| t.timestamp()), - cred.last_error_time.map(|t| t.timestamp()), - cred.last_error_message, - cred.last_health_check_time.map(|t| t.timestamp()), - cred.last_health_check_model, - cred.updated_at.timestamp(), - cred.proxy_url, - ], - )?; - Ok(()) - } - - /// 删除凭证 - pub fn delete(conn: &Connection, uuid: &str) -> Result { - let affected = conn.execute( - "DELETE FROM provider_pool_credentials WHERE uuid = ?1", - [uuid], - )?; - Ok(affected > 0) - } - - /// 更新健康状态 - #[allow(clippy::too_many_arguments)] - pub fn update_health_status( - conn: &Connection, - uuid: &str, - is_healthy: bool, - error_count: u32, - last_error_time: Option>, - last_error_message: Option<&str>, - last_health_check_time: Option>, - last_health_check_model: Option<&str>, - ) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE provider_pool_credentials SET - is_healthy = ?2, error_count = ?3, last_error_time = ?4, - last_error_message = ?5, last_health_check_time = ?6, - last_health_check_model = ?7, updated_at = ?8 - WHERE uuid = ?1", - params![ - uuid, - is_healthy, - error_count, - last_error_time.map(|t| t.timestamp()), - last_error_message, - last_health_check_time.map(|t| t.timestamp()), - last_health_check_model, - Utc::now().timestamp(), - ], - )?; - Ok(()) - } - - /// 更新使用统计 - pub fn update_usage( - conn: &Connection, - uuid: &str, - usage_count: u64, - last_used: DateTime, - ) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE provider_pool_credentials SET - usage_count = ?2, last_used = ?3, updated_at = ?4 - WHERE uuid = ?1", - params![ - uuid, - usage_count, - last_used.timestamp(), - Utc::now().timestamp() - ], - )?; - Ok(()) - } - - /// 重置凭证计数器 - pub fn reset_counters(conn: &Connection, uuid: &str) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE provider_pool_credentials SET - usage_count = 0, error_count = 0, is_healthy = 1, - last_error_time = NULL, last_error_message = NULL, updated_at = ?2 - WHERE uuid = ?1", - params![uuid, Utc::now().timestamp()], - )?; - Ok(()) - } - - /// 重置指定类型的所有凭证健康状态 - pub fn reset_health_by_type( - conn: &Connection, - provider_type: &PoolProviderType, - ) -> Result { - let affected = conn.execute( - "UPDATE provider_pool_credentials SET - is_healthy = 1, error_count = 0, last_error_time = NULL, - last_error_message = NULL, updated_at = ?2 - WHERE provider_type = ?1", - params![provider_type.to_string(), Utc::now().timestamp()], - )?; - Ok(affected) - } - - /// 从数据库行转换为 ProviderCredential - fn row_to_credential(row: &rusqlite::Row) -> Result { - let uuid: String = row.get(0)?; - let provider_type_str: String = row.get(1)?; - let credential_json: String = row.get(2)?; - let name: Option = row.get(3)?; - let is_healthy: bool = row.get(4)?; - let is_disabled: bool = row.get(5)?; - let check_health: bool = row.get(6)?; - let check_model_name: Option = row.get(7)?; - let not_supported_models_json: Option = row.get(8)?; - let supported_models_json: Option = row.get(9)?; - let usage_count: u64 = row.get::<_, i64>(10)? as u64; - let error_count: u32 = row.get::<_, i32>(11)? as u32; - let last_used_ts: Option = row.get(12)?; - let last_error_time_ts: Option = row.get(13)?; - let last_error_message: Option = row.get(14)?; - let last_health_check_time_ts: Option = row.get(15)?; - let last_health_check_model: Option = row.get(16)?; - let created_at_ts: i64 = row.get(17)?; - let updated_at_ts: i64 = row.get(18)?; - let source_str: Option = row.get(19).ok(); - let proxy_url: Option = row.get(20).ok(); - - let provider_type: PoolProviderType = - provider_type_str.parse().unwrap_or(PoolProviderType::Kiro); - - let credential: CredentialData = serde_json::from_str(&credential_json).map_err(|e| { - rusqlite::Error::FromSqlConversionFailure(2, rusqlite::types::Type::Text, Box::new(e)) - })?; - - let not_supported_models: Vec = not_supported_models_json - .and_then(|s| serde_json::from_str(&s).ok()) - .unwrap_or_default(); - - let supported_models: Vec = supported_models_json - .and_then(|s| serde_json::from_str(&s).ok()) - .unwrap_or_default(); - - let source = match source_str.as_deref() { - Some("imported") => CredentialSource::Imported, - Some("private") => CredentialSource::Private, - _ => CredentialSource::Manual, - }; - - Ok(ProviderCredential { - uuid, - provider_type, - credential, - name, - is_healthy, - is_disabled, - check_health, - check_model_name, - not_supported_models, - supported_models, - usage_count, - error_count, - last_used: last_used_ts.and_then(|ts| Utc.timestamp_opt(ts, 0).single()), - last_error_time: last_error_time_ts.and_then(|ts| Utc.timestamp_opt(ts, 0).single()), - last_error_message, - last_health_check_time: last_health_check_time_ts - .and_then(|ts| Utc.timestamp_opt(ts, 0).single()), - last_health_check_model, - created_at: Utc - .timestamp_opt(created_at_ts, 0) - .single() - .unwrap_or_default(), - updated_at: Utc - .timestamp_opt(updated_at_ts, 0) - .single() - .unwrap_or_default(), - cached_token: None, // 从 get_token_cache 单独获取 - source, - proxy_url, - prompt_cache_mode_override: None, - }) - } - - // ==================== Token 缓存操作 ==================== - - /// 获取凭证的 Token 缓存信息 - pub fn get_token_cache( - conn: &Connection, - uuid: &str, - ) -> Result, rusqlite::Error> { - let mut stmt = conn.prepare( - "SELECT cached_access_token, cached_refresh_token, token_expiry_time, - last_refresh_time, refresh_error_count, last_refresh_error - FROM provider_pool_credentials - WHERE uuid = ?1", - )?; - - let mut rows = stmt.query([uuid])?; - if let Some(row) = rows.next()? { - let access_token: Option = row.get(0)?; - let refresh_token: Option = row.get(1)?; - let expiry_time_str: Option = row.get(2)?; - let last_refresh_str: Option = row.get(3)?; - let refresh_error_count: i32 = row.get::<_, Option>(4)?.unwrap_or(0); - let last_refresh_error: Option = row.get(5)?; - - // 如果没有缓存的 token,返回 None - if access_token.is_none() { - return Ok(None); - } - - let expiry_time = expiry_time_str - .and_then(|s| DateTime::parse_from_rfc3339(&s).ok()) - .map(|dt| dt.with_timezone(&Utc)); - - let last_refresh = last_refresh_str - .and_then(|s| DateTime::parse_from_rfc3339(&s).ok()) - .map(|dt| dt.with_timezone(&Utc)); - - Ok(Some(CachedTokenInfo { - access_token, - refresh_token, - expiry_time, - last_refresh, - refresh_error_count: refresh_error_count as u32, - last_refresh_error, - })) - } else { - Ok(None) - } - } - - /// 更新凭证的 Token 缓存 - pub fn update_token_cache( - conn: &Connection, - uuid: &str, - token_info: &CachedTokenInfo, - ) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE provider_pool_credentials SET - cached_access_token = ?2, - cached_refresh_token = ?3, - token_expiry_time = ?4, - last_refresh_time = ?5, - refresh_error_count = ?6, - last_refresh_error = ?7, - updated_at = ?8 - WHERE uuid = ?1", - params![ - uuid, - token_info.access_token, - token_info.refresh_token, - token_info.expiry_time.map(|t| t.to_rfc3339()), - token_info.last_refresh.map(|t| t.to_rfc3339()), - token_info.refresh_error_count as i32, - token_info.last_refresh_error, - Utc::now().timestamp(), - ], - )?; - Ok(()) - } - - /// 清除凭证的 Token 缓存 - pub fn clear_token_cache(conn: &Connection, uuid: &str) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE provider_pool_credentials SET - cached_access_token = NULL, - cached_refresh_token = NULL, - token_expiry_time = NULL, - last_refresh_time = NULL, - refresh_error_count = 0, - last_refresh_error = NULL, - updated_at = ?2 - WHERE uuid = ?1", - params![uuid, Utc::now().timestamp()], - )?; - Ok(()) - } - - /// 记录 Token 刷新错误 - pub fn record_token_refresh_error( - conn: &Connection, - uuid: &str, - error_message: &str, - ) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE provider_pool_credentials SET - refresh_error_count = COALESCE(refresh_error_count, 0) + 1, - last_refresh_error = ?2, - updated_at = ?3 - WHERE uuid = ?1", - params![uuid, error_message, Utc::now().timestamp()], - )?; - Ok(()) - } - - /// 重置 Token 刷新错误计数 - #[allow(dead_code)] - pub fn reset_token_refresh_errors( - conn: &Connection, - uuid: &str, - ) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE provider_pool_credentials SET - refresh_error_count = 0, - last_refresh_error = NULL, - updated_at = ?2 - WHERE uuid = ?1", - params![uuid, Utc::now().timestamp()], - )?; - Ok(()) - } -} diff --git a/src-tauri/crates/core/src/database/startup_migrations.rs b/src-tauri/crates/core/src/database/startup_migrations.rs index b71b26aeb..7a2e102ef 100644 --- a/src-tauri/crates/core/src/database/startup_migrations.rs +++ b/src-tauri/crates/core/src/database/startup_migrations.rs @@ -1,3 +1,5 @@ +use std::path::Path; + use rusqlite::Connection; use super::{migration, migration_v2, migration_v3, migration_v4, migration_v5, migration_v6}; @@ -62,8 +64,7 @@ fn run_nonfatal_count_migration( fn run_provider_pool_startup_migrations(conn: &Connection) { run_provider_id_migration(conn); migration::check_model_registry_version(conn); - run_api_keys_to_pool_migration(conn); - run_legacy_api_key_cleanup(conn); + run_retired_provider_pool_cleanup(conn); } fn run_provider_id_migration(conn: &Connection) { @@ -78,22 +79,83 @@ fn run_provider_id_migration(conn: &Connection) { ); } -fn run_api_keys_to_pool_migration(conn: &Connection) { - run_nonfatal_count_migration( +#[derive(Debug, Clone, Copy)] +struct RetiredProviderPoolCleanup { + rows_deleted: usize, + managed_files_deleted: usize, +} + +fn run_retired_provider_pool_cleanup(conn: &Connection) { + run_nonfatal_logged_startup_migration( conn, - "API Key 迁移失败", - migration::migrate_api_keys_to_pool, - |_, count| format!("[数据库] 已将 {} 条 API Key 迁移到凭证池", count), + "凭证池退役清理失败", + |tx| { + let rows_deleted = clear_provider_pool_credentials(tx)?; + let managed_files_deleted = remove_managed_provider_pool_credential_files()?; + Ok(RetiredProviderPoolCleanup { + rows_deleted, + managed_files_deleted, + }) + }, + |_, result| { + if result.rows_deleted == 0 && result.managed_files_deleted == 0 { + return None; + } + Some(format!( + "[数据库] 凭证池已退役,清理 {} 条旧凭证记录和 {} 个托管凭证文件", + result.rows_deleted, result.managed_files_deleted + )) + }, ); } -fn run_legacy_api_key_cleanup(conn: &Connection) { - run_nonfatal_count_migration( - conn, - "旧 API Key 凭证清理失败", - migration::cleanup_legacy_api_key_credentials, - |_, count| format!("[数据库] 已清理 {} 条旧 API Key 凭证", count), - ); +fn clear_provider_pool_credentials(conn: &Connection) -> Result { + conn.execute("DELETE FROM provider_pool_credentials", []) + .map_err(|error| error.to_string()) +} + +fn count_managed_files(path: &Path) -> usize { + let Ok(metadata) = std::fs::symlink_metadata(path) else { + return 0; + }; + + let file_type = metadata.file_type(); + if file_type.is_symlink() || metadata.is_file() { + return 1; + } + + if !metadata.is_dir() { + return 0; + } + + let Ok(entries) = std::fs::read_dir(path) else { + return 0; + }; + + entries + .filter_map(Result::ok) + .map(|entry| count_managed_files(&entry.path())) + .sum() +} + +fn remove_managed_provider_pool_credential_files() -> Result { + let data_dir = crate::app_paths::preferred_data_dir().map_err(|error| error.to_string())?; + let credentials_dir = data_dir.join("credentials"); + + let metadata = match std::fs::symlink_metadata(&credentials_dir) { + Ok(metadata) => metadata, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(0), + Err(error) => return Err(error.to_string()), + }; + + let deleted_count = count_managed_files(&credentials_dir); + if metadata.file_type().is_symlink() || metadata.is_file() { + std::fs::remove_file(&credentials_dir).map_err(|error| error.to_string())?; + } else if metadata.is_dir() { + std::fs::remove_dir_all(&credentials_dir).map_err(|error| error.to_string())?; + } + + Ok(deleted_count) } fn run_mcp_startup_migrations(conn: &Connection) { diff --git a/src-tauri/crates/core/src/lib.rs b/src-tauri/crates/core/src/lib.rs index b40d6e45e..25e2166ec 100644 --- a/src-tauri/crates/core/src/lib.rs +++ b/src-tauri/crates/core/src/lib.rs @@ -12,7 +12,6 @@ //! - `config`: 配置管理(类型、YAML、热重载、导入导出) //! - `connect`: Deep Link 协议和中转商注册表 //! - `middleware`: HTTP 中间件(认证、限速) -//! - `orchestrator`: 模型选择编排器 //! - `plugin`: 插件系统(加载、管理、UI、安装) //! - `session`: 会话管理(限速、粘性路由) //! - `session_files`: 会话文件存储 @@ -48,9 +47,6 @@ pub mod general_chat; // 路由系统 pub mod router; -// 凭证池核心(types, pool, health, risk) -pub mod credential; - // 请求处理器核心类型(context, error) pub mod processor; pub mod provider_prompt_cache_support; diff --git a/src-tauri/crates/core/src/models/codewhisperer.rs b/src-tauri/crates/core/src/models/codewhisperer.rs deleted file mode 100644 index 98197e0ab..000000000 --- a/src-tauri/crates/core/src/models/codewhisperer.rs +++ /dev/null @@ -1,173 +0,0 @@ -//! CodeWhisperer/Kiro API 数据模型 -//! -//! 支持标准工具和特殊工具类型(如 web_search)。 -//! -//! # 更新日志 -//! -//! - 2025-12-27: 添加 CWWebSearchTool 支持,修复 Issue #49 -use serde::{Deserialize, Serialize}; - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CodeWhispererRequest { - pub conversation_state: ConversationState, - #[serde(skip_serializing_if = "Option::is_none")] - pub profile_arn: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ConversationState { - pub chat_trigger_type: String, - pub conversation_id: String, - pub current_message: CurrentMessage, - #[serde(skip_serializing_if = "Option::is_none")] - pub history: Option>, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CurrentMessage { - pub user_input_message: UserInputMessage, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct UserInputMessage { - pub content: String, - pub model_id: String, - pub origin: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub images: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub user_input_message_context: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct UserInputMessageContext { - #[serde(skip_serializing_if = "Option::is_none")] - pub tools: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_results: Option>, -} - -/// CodeWhisperer 工具项 -/// -/// 支持两种类型: -/// - 标准工具(带 tool_specification) -/// - 联网搜索工具(仅 type 字段) -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(untagged)] -pub enum CWToolItem { - /// 标准工具定义 - Standard(CWTool), - /// 联网搜索工具 - WebSearch(CWWebSearchTool), -} - -/// 标准工具定义 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWTool { - pub tool_specification: ToolSpecification, -} - -/// 联网搜索工具 -/// -/// Codex/Kiro API 支持的特殊工具类型,用于联网搜索。 -/// 格式:`{"type": "web_search"}` -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWWebSearchTool { - #[serde(rename = "type")] - pub tool_type: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ToolSpecification { - pub name: String, - pub description: String, - pub input_schema: InputSchema, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct InputSchema { - pub json: serde_json::Value, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWToolResult { - pub content: Vec, - pub status: String, - pub tool_use_id: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWTextContent { - pub text: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWImage { - pub format: String, - pub source: CWImageSource, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWImageSource { - pub bytes: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(untagged)] -pub enum HistoryItem { - User(UserHistoryItem), - Assistant(AssistantHistoryItem), -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct UserHistoryItem { - pub user_input_message: UserInputMessage, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct AssistantHistoryItem { - pub assistant_response_message: AssistantResponseMessage, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct AssistantResponseMessage { - pub content: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_uses: Option>, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWToolUse { - pub input: serde_json::Value, - pub name: String, - pub tool_use_id: String, -} - -// Response types -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWStreamEvent { - #[serde(skip_serializing_if = "Option::is_none")] - pub assistant_response_event: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct AssistantResponseEvent { - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_use: Option, -} diff --git a/src-tauri/crates/core/src/models/kiro_fingerprint.rs b/src-tauri/crates/core/src/models/kiro_fingerprint.rs deleted file mode 100644 index ef00fd915..000000000 --- a/src-tauri/crates/core/src/models/kiro_fingerprint.rs +++ /dev/null @@ -1,199 +0,0 @@ -//! Kiro 凭证指纹绑定模型 -//! -//! 为每个 Kiro 凭证存储独立的 Machine ID,实现多账号指纹隔离。 - -#![allow(dead_code)] - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::fs; -use std::path::PathBuf; - -/// Kiro 凭证指纹绑定 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct KiroFingerprintBinding { - /// 凭证 UUID - pub credential_uuid: String, - /// 绑定的 Machine ID - pub machine_id: String, - /// 创建时间 - pub created_at: DateTime, - /// 最后切换时间 - pub last_switched_at: Option>, -} - -/// 指纹绑定存储 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct KiroFingerprintStore { - /// 凭证 UUID -> 指纹绑定 - pub bindings: HashMap, -} - -impl KiroFingerprintStore { - /// 获取存储文件路径 - pub fn get_storage_path() -> Result { - let app_data_dir = dirs::data_dir() - .ok_or_else(|| "无法获取应用数据目录".to_string())? - .join("lime"); - - // 确保目录存在 - if !app_data_dir.exists() { - fs::create_dir_all(&app_data_dir).map_err(|e| format!("创建应用数据目录失败: {e}"))?; - } - - Ok(app_data_dir.join("kiro_fingerprints.json")) - } - - /// 从文件加载 - pub fn load() -> Result { - let path = Self::get_storage_path()?; - - if !path.exists() { - return Ok(Self::default()); - } - - let content = - fs::read_to_string(&path).map_err(|e| format!("读取指纹存储文件失败: {e}"))?; - - serde_json::from_str(&content).map_err(|e| format!("解析指纹存储文件失败: {e}")) - } - - /// 保存到文件 - pub fn save(&self) -> Result<(), String> { - let path = Self::get_storage_path()?; - let content = - serde_json::to_string_pretty(self).map_err(|e| format!("序列化指纹存储失败: {e}"))?; - - fs::write(&path, content).map_err(|e| format!("写入指纹存储文件失败: {e}")) - } - - /// 获取凭证的指纹绑定 - pub fn get_binding(&self, credential_uuid: &str) -> Option<&KiroFingerprintBinding> { - self.bindings.get(credential_uuid) - } - - /// 获取或创建凭证的指纹绑定 - /// - /// 如果凭证没有绑定指纹,会基于凭证信息生成一个新的 Machine ID - pub fn get_or_create_binding( - &mut self, - credential_uuid: &str, - profile_arn: Option<&str>, - client_id: Option<&str>, - ) -> Result<&KiroFingerprintBinding, String> { - if !self.bindings.contains_key(credential_uuid) { - // 生成基于凭证的 Machine ID - let machine_id = generate_stable_machine_id(credential_uuid, profile_arn, client_id); - - let binding = KiroFingerprintBinding { - credential_uuid: credential_uuid.to_string(), - machine_id, - created_at: Utc::now(), - last_switched_at: None, - }; - - self.bindings.insert(credential_uuid.to_string(), binding); - self.save()?; - } - - Ok(self.bindings.get(credential_uuid).unwrap()) - } - - /// 更新最后切换时间 - pub fn update_last_switched(&mut self, credential_uuid: &str) -> Result<(), String> { - if let Some(binding) = self.bindings.get_mut(credential_uuid) { - binding.last_switched_at = Some(Utc::now()); - self.save()?; - } - Ok(()) - } - - /// 删除凭证的指纹绑定 - pub fn remove_binding(&mut self, credential_uuid: &str) -> Result<(), String> { - self.bindings.remove(credential_uuid); - self.save() - } -} - -/// 生成稳定的 Machine ID -/// -/// 基于凭证信息生成一个稳定的 UUID 格式 Machine ID。 -/// 同一凭证每次生成的 Machine ID 相同,确保账号身份一致。 -fn generate_stable_machine_id( - credential_uuid: &str, - profile_arn: Option<&str>, - client_id: Option<&str>, -) -> String { - use sha2::{Digest, Sha256}; - - // 使用凭证相关信息作为种子 - let seed = format!( - "kiro_fingerprint:{}:{}:{}", - credential_uuid, - profile_arn.unwrap_or(""), - client_id.unwrap_or("") - ); - - let mut hasher = Sha256::new(); - hasher.update(seed.as_bytes()); - let result = hasher.finalize(); - - // 将哈希结果转换为 UUID 格式 - let hex = format!("{result:x}"); - format!( - "{}-{}-{}-{}-{}", - &hex[0..8], - &hex[8..12], - &hex[12..16], - &hex[16..20], - &hex[20..32] - ) -} - -/// 切换到本地的结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SwitchToLocalResult { - /// 是否成功 - pub success: bool, - /// 结果消息 - pub message: String, - /// 是否需要用户操作(如需管理员权限) - pub requires_action: bool, - /// 切换的 Machine ID - pub machine_id: Option, - /// 是否需要重启 Kiro IDE - pub requires_kiro_restart: bool, -} - -impl SwitchToLocalResult { - pub fn success(message: impl Into, machine_id: String) -> Self { - Self { - success: true, - message: message.into(), - requires_action: false, - machine_id: Some(machine_id), - requires_kiro_restart: true, - } - } - - pub fn error(message: impl Into) -> Self { - Self { - success: false, - message: message.into(), - requires_action: false, - machine_id: None, - requires_kiro_restart: false, - } - } - - pub fn requires_admin(message: impl Into) -> Self { - Self { - success: false, - message: message.into(), - requires_action: true, - machine_id: None, - requires_kiro_restart: false, - } - } -} diff --git a/src-tauri/crates/core/src/models/mod.rs b/src-tauri/crates/core/src/models/mod.rs index 4f37a9cd5..6b0014f4f 100644 --- a/src-tauri/crates/core/src/models/mod.rs +++ b/src-tauri/crates/core/src/models/mod.rs @@ -5,9 +5,7 @@ pub mod anthropic; pub mod app_type; pub mod client_type; -pub mod codewhisperer; pub mod injection_types; -pub mod kiro_fingerprint; pub mod machine_id; pub mod mcp_model; pub mod model_registry; @@ -24,8 +22,6 @@ pub mod vertex_model; pub use anthropic::*; pub use app_type::AppType; pub use client_type::{select_provider, ClientType}; -#[allow(unused_imports)] -pub use codewhisperer::*; pub use injection_types::{InjectionMode, InjectionRule}; pub use mcp_model::McpServer; #[allow(unused_imports)] diff --git a/src-tauri/crates/core/src/models/provider_pool_model.rs b/src-tauri/crates/core/src/models/provider_pool_model.rs index e8c72fea5..fdd1b5c12 100644 --- a/src-tauri/crates/core/src/models/provider_pool_model.rs +++ b/src-tauri/crates/core/src/models/provider_pool_model.rs @@ -8,8 +8,6 @@ use serde::{Deserialize, Serialize}; use std::collections::HashMap; use uuid::Uuid; -use super::provider_type::ANTIGRAVITY_MODELS_FALLBACK; - /// 凭证来源枚举 /// 用于标识凭证是如何添加到凭证池的 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)] @@ -55,11 +53,6 @@ pub enum CredentialData { project_id: Option, }, - /// Antigravity OAuth 凭证(文件路径)- Google 内部 Gemini 3 Pro - AntigravityOAuth { - creds_file_path: String, - project_id: Option, - }, /// OpenAI API Key 凭证 OpenAIKey { api_key: String, @@ -116,11 +109,6 @@ impl CredentialData { format!("Gemini OAuth: {}", mask_path(creds_file_path)) } - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => { - format!("Antigravity OAuth: {}", mask_path(creds_file_path)) - } CredentialData::OpenAIKey { api_key, .. } => { format!("OpenAI: {}", mask_key(api_key)) } @@ -153,8 +141,6 @@ impl CredentialData { match self { CredentialData::KiroOAuth { .. } => PoolProviderType::Kiro, CredentialData::GeminiOAuth { .. } => PoolProviderType::Gemini, - - CredentialData::AntigravityOAuth { .. } => PoolProviderType::Antigravity, CredentialData::OpenAIKey { .. } => PoolProviderType::OpenAI, CredentialData::ClaudeKey { .. } => PoolProviderType::Claude, CredentialData::VertexKey { .. } => PoolProviderType::Vertex, @@ -326,7 +312,7 @@ impl ProviderCredential { /// 检查两个来源的排除列表: /// 1. `not_supported_models` - 通用的不支持模型列表(精确匹配) /// 2. `excluded_models` - 来自 CredentialData::GeminiApiKey 的排除列表(支持通配符) - /// 3. Antigravity 凭证只支持特定的模型列表 + /// 3. API Key 凭证可声明额外的模型排除列表 pub fn supports_model(&self, model: &str) -> bool { // 检查通用的不支持模型列表(精确匹配) if self.not_supported_models.contains(&model.to_string()) { @@ -345,13 +331,6 @@ impl ProviderCredential { } } - // Antigravity 凭证只支持特定的模型 - // 使用 providers::antigravity 中定义的模型列表(fallback) - // 实际模型列表由 models/aliases/antigravity.json 定义 - if let CredentialData::AntigravityOAuth { .. } = &self.credential { - return ANTIGRAVITY_MODELS_FALLBACK.contains(&model); - } - true } @@ -538,7 +517,6 @@ pub fn get_default_check_model(provider_type: PoolProviderType) -> &'static str PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929", // Anthropic 兼容格式使用相同的健康检查模型 PoolProviderType::AnthropicCompatible => "claude-sonnet-4-5-20250929", - PoolProviderType::Antigravity => "gemini-3-pro-preview", PoolProviderType::Vertex => "gemini-2.0-flash", PoolProviderType::GeminiApiKey => "gemini-2.5-flash", PoolProviderType::Codex => "gpt-4o-mini", @@ -590,7 +568,6 @@ fn get_credential_type(cred: &CredentialData) -> String { match cred { CredentialData::KiroOAuth { .. } => "kiro_oauth".to_string(), CredentialData::GeminiOAuth { .. } => "gemini_oauth".to_string(), - CredentialData::AntigravityOAuth { .. } => "antigravity_oauth".to_string(), CredentialData::OpenAIKey { .. } => "openai_key".to_string(), CredentialData::ClaudeKey { .. } => "claude_key".to_string(), CredentialData::VertexKey { .. } => "vertex_key".to_string(), @@ -608,9 +585,6 @@ pub fn get_oauth_creds_path(cred: &CredentialData) -> Option { CredentialData::GeminiOAuth { creds_file_path, .. } => Some(creds_file_path.clone()), - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => Some(creds_file_path.clone()), CredentialData::CodexOAuth { creds_file_path, .. } => Some(creds_file_path.clone()), diff --git a/src-tauri/crates/core/src/models/provider_type.rs b/src-tauri/crates/core/src/models/provider_type.rs index cdc1f10ae..4da93b178 100644 --- a/src-tauri/crates/core/src/models/provider_type.rs +++ b/src-tauri/crates/core/src/models/provider_type.rs @@ -23,7 +23,6 @@ pub enum ProviderType { /// Anthropic 兼容格式(支持 system 数组格式等变体) #[serde(rename = "anthropic_compatible")] AnthropicCompatible, - Antigravity, Vertex, #[serde(rename = "gemini_api_key")] GeminiApiKey, @@ -55,7 +54,6 @@ impl std::fmt::Display for ProviderType { ProviderType::Claude => write!(f, "claude"), ProviderType::ClaudeOAuth => write!(f, "claude_oauth"), ProviderType::AnthropicCompatible => write!(f, "anthropic_compatible"), - ProviderType::Antigravity => write!(f, "antigravity"), ProviderType::Vertex => write!(f, "vertex"), ProviderType::GeminiApiKey => write!(f, "gemini_api_key"), ProviderType::Codex => write!(f, "codex"), @@ -80,7 +78,6 @@ impl std::str::FromStr for ProviderType { "anthropic_compatible" | "anthropic-compatible" => { Ok(ProviderType::AnthropicCompatible) } - "antigravity" => Ok(ProviderType::Antigravity), "vertex" => Ok(ProviderType::Vertex), "gemini_api_key" => Ok(ProviderType::GeminiApiKey), "codex" => Ok(ProviderType::Codex), @@ -113,26 +110,6 @@ impl std::str::FromStr for ProviderType { } } -/// Antigravity 支持的模型列表(fallback,当无法从 models 仓库获取时使用) -pub const ANTIGRAVITY_MODELS_FALLBACK: &[&str] = &[ - "gemini-2.5-computer-use-preview-10-2025", - "gemini-3-pro-image-preview", - "gemini-3-pro-preview", - "gemini-3-flash-preview", - "gemini-2.5-flash-preview", - "gemini-2.5-flash", - "gemini-2.5-pro", - "gemini-3-flash", - "gemini-3-pro-high", - "gemini-3-pro-low", - "gemini-claude-sonnet-4-5", - "gemini-claude-sonnet-4-5-thinking", - "gemini-claude-opus-4-5-thinking", - "claude-sonnet-4-5", - "claude-sonnet-4-5-thinking", - "claude-opus-4-5-thinking", -]; - #[cfg(test)] mod tests { use super::*; diff --git a/src-tauri/crates/core/src/orchestrator/pool_builder.rs b/src-tauri/crates/core/src/orchestrator/pool_builder.rs index 56704afd6..03d6234a5 100644 --- a/src-tauri/crates/core/src/orchestrator/pool_builder.rs +++ b/src-tauri/crates/core/src/orchestrator/pool_builder.rs @@ -16,7 +16,6 @@ pub enum ProviderType { Kiro, Azure, Bedrock, - Antigravity, Custom, } @@ -30,7 +29,6 @@ impl ProviderType { "kiro" | "codewhisperer" => Some(ProviderType::Kiro), "azure" => Some(ProviderType::Azure), "bedrock" => Some(ProviderType::Bedrock), - "antigravity" => Some(ProviderType::Antigravity), _ => Some(ProviderType::Custom), } } @@ -44,7 +42,6 @@ impl ProviderType { ProviderType::Kiro => "Kiro", ProviderType::Azure => "Azure", ProviderType::Bedrock => "Bedrock", - ProviderType::Antigravity => "Antigravity", ProviderType::Custom => "Custom", } } @@ -243,47 +240,6 @@ pub fn builtin_provider_definitions() -> Vec { ], default_base_url: None, }, - // Antigravity (Google Cloud Code Assist) - ProviderDefinition { - provider_type: ProviderType::Antigravity, - display_name: "Antigravity".to_string(), - families: vec![ - // Max 等级:Gemini 3 Pro 和 Claude Opus - ModelFamily { - name: "gemini-3-pro".to_string(), - pattern: "gemini-3-pro*".to_string(), - tier: 3, - description: Some("Gemini 3 Pro via Antigravity".to_string()), - }, - ModelFamily { - name: "opus".to_string(), - pattern: "*opus*".to_string(), - tier: 3, - description: Some("Claude Opus via Antigravity".to_string()), - }, - // Pro 等级:Claude Sonnet 和 Gemini 2.5 - ModelFamily { - name: "sonnet".to_string(), - pattern: "*sonnet*".to_string(), - tier: 2, - description: Some("Claude Sonnet via Antigravity".to_string()), - }, - ModelFamily { - name: "gemini-2.5".to_string(), - pattern: "gemini-2.5*".to_string(), - tier: 2, - description: Some("Gemini 2.5 via Antigravity".to_string()), - }, - // Mini 等级:Flash 模型 - ModelFamily { - name: "gemini-3-flash".to_string(), - pattern: "gemini-3-flash*".to_string(), - tier: 1, - description: Some("Gemini 3 Flash via Antigravity".to_string()), - }, - ], - default_base_url: None, - }, ] } @@ -460,7 +416,7 @@ pub struct CredentialInfo { pub id: String, /// Provider 类型(用于模型分类) pub provider_type: ProviderType, - /// 原始 Provider 类型字符串(用于前端识别,如 "antigravity"、"kiro" 等) + /// 原始 Provider 类型字符串(用于前端识别,如 "kiro" 等) pub original_provider_type: Option, /// 支持的模型列表 pub supported_models: Vec, @@ -541,7 +497,7 @@ impl DynamicPoolBuilder { .or_else(|| metadata.as_ref().and_then(|m| m.family.clone())); // 构建 AvailableModel - // 优先使用原始 provider 类型(如 "antigravity"),否则使用枚举名称 + // 优先使用原始 provider 类型,否则使用枚举名称 let provider_type_str = credential .original_provider_type .clone() diff --git a/src-tauri/crates/core/src/router/rules.rs b/src-tauri/crates/core/src/router/rules.rs index 2a57cabf9..29e6ee14e 100644 --- a/src-tauri/crates/core/src/router/rules.rs +++ b/src-tauri/crates/core/src/router/rules.rs @@ -92,9 +92,9 @@ mod tests { #[test] fn test_route_returns_default() { - let router = Router::new(ProviderType::Antigravity); + let router = Router::new(ProviderType::Gemini); let result = router.route("any-model"); - assert_eq!(result.provider, Some(ProviderType::Antigravity)); + assert_eq!(result.provider, Some(ProviderType::Gemini)); assert!(result.is_default); } diff --git a/src-tauri/crates/core/src/websocket/types.rs b/src-tauri/crates/core/src/websocket/types.rs index fe0f14032..5a9ca58c9 100644 --- a/src-tauri/crates/core/src/websocket/types.rs +++ b/src-tauri/crates/core/src/websocket/types.rs @@ -69,12 +69,6 @@ pub enum WsMessage { Ping { timestamp: i64 }, /// 心跳响应 Pong { timestamp: i64 }, - /// 订阅 Kiro 凭证状态事件 - SubscribeKiroEvents, - /// 取消订阅 Kiro 凭证状态事件 - UnsubscribeKiroEvents, - /// Kiro 凭证状态事件通知 - KiroCredentialEvent(WsKiroEvent), } /// WebSocket API 请求 @@ -321,78 +315,3 @@ pub struct WsStatsSnapshot { pub total_messages: u64, pub total_errors: u64, } - -/// WebSocket Kiro 凭证事件 -/// -/// 用于通过 WebSocket 推送 Kiro 凭证状态变化 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "event_type", rename_all = "snake_case")] -pub enum WsKiroEvent { - /// 凭证状态更新 - CredentialStatusUpdate { - uuid: String, - is_healthy: bool, - is_disabled: bool, - error_count: u32, - health_score: Option, - last_used: Option>, - }, - /// 凭证刷新开始 - RefreshStarted { - uuid: String, - credential_name: Option, - }, - /// 凭证刷新成功 - RefreshSuccess { - uuid: String, - credential_name: Option, - new_token_info: KiroTokenInfo, - }, - /// 凭证刷新失败 - RefreshFailed { - uuid: String, - credential_name: Option, - error: String, - error_code: Option, - }, - /// 凭证健康检查完成 - HealthCheckCompleted { - uuid: String, - credential_name: Option, - is_healthy: bool, - health_score: Option, - last_check: DateTime, - }, - /// 凭证池统计更新 - PoolStatsUpdate { - total_credentials: u32, - healthy_credentials: u32, - available_credentials: u32, - average_health_score: Option, - last_rotation: Option>, - }, - /// 凭证轮换事件 - CredentialRotated { - from_uuid: Option, - to_uuid: String, - reason: String, - rotation_time: DateTime, - }, - /// 凭证自动禁用事件 - CredentialAutoDisabled { - uuid: String, - credential_name: Option, - reason: String, - error_type: String, - disable_time: DateTime, - }, -} - -/// Kiro Token 信息(用于刷新成功事件) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct KiroTokenInfo { - pub expires_at: DateTime, - pub auth_method: String, - pub provider: String, - pub region: String, -} diff --git a/src-tauri/crates/credential/Cargo.toml b/src-tauri/crates/credential/Cargo.toml deleted file mode 100644 index 4c90fa05b..000000000 --- a/src-tauri/crates/credential/Cargo.toml +++ /dev/null @@ -1,41 +0,0 @@ -[package] -name = "lime-credential" -version.workspace = true -edition.workspace = true -authors.workspace = true - -[dependencies] -lime-core.workspace = true -lime-infra.workspace = true - -# 序列化 -serde.workspace = true -serde_json.workspace = true - -# 异步运行时 -tokio.workspace = true - -# 日志 -tracing.workspace = true - -# HTTP 服务器(AllCredentialsExhaustedError 的 IntoResponse) -axum.workspace = true - -# HTTP 客户端 -reqwest.workspace = true - -# 时间 -chrono.workspace = true - -# 并发 -dashmap.workspace = true - -# 加密 -chacha20poly1305 = "0.10" -base64.workspace = true -rand.workspace = true -sha2.workspace = true - -[dev-dependencies] -proptest.workspace = true -tempfile.workspace = true diff --git a/src-tauri/crates/credential/src/balancer.rs b/src-tauri/crates/credential/src/balancer.rs deleted file mode 100644 index b9c30ed10..000000000 --- a/src-tauri/crates/credential/src/balancer.rs +++ /dev/null @@ -1,567 +0,0 @@ -//! 负载均衡器实现 -//! -//! 提供轮询负载均衡策略,支持凭证冷却和自动恢复 - -use chrono::{DateTime, Duration, Utc}; -use dashmap::DashMap; -use lime_core::credential::health::{HealthCheckConfig, HealthChecker}; -use lime_core::credential::pool::{CredentialPool, PoolError}; -use lime_core::credential::types::Credential; -use lime_core::ProviderType; -use lime_infra::ProxyClientFactory; -use reqwest::Client; -use serde::{Deserialize, Serialize}; -use std::sync::atomic::{AtomicUsize, Ordering}; -use std::sync::Arc; - -/// 负载均衡策略 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] -#[serde(rename_all = "snake_case")] -pub enum BalanceStrategy { - /// 轮询策略(默认) - #[default] - RoundRobin, - /// 最少使用策略 - LeastUsed, - /// 随机策略 - Random, -} - -/// 冷却信息 -#[derive(Debug, Clone)] -pub struct CooldownInfo { - /// 冷却结束时间 - pub until: DateTime, - /// 冷却原因 - pub reason: String, -} - -/// 凭证选择结果 - 包含凭证和对应的 HTTP 客户端 -#[derive(Debug)] -pub struct CredentialSelection { - /// 选中的凭证 - pub credential: Credential, - /// 配置了代理的 HTTP 客户端 - pub client: Client, -} - -/// 负载均衡器 - 管理多个 Provider 的凭证池 -pub struct LoadBalancer { - /// 负载均衡策略 - strategy: BalanceStrategy, - /// 各 Provider 的凭证池 - pools: DashMap>, - /// 轮询索引(每个 Provider 独立) - round_robin_indices: DashMap, - /// 健康检查器 - health_checker: HealthChecker, - /// 代理客户端工厂 - proxy_factory: ProxyClientFactory, -} - -impl LoadBalancer { - /// 创建新的负载均衡器 - pub fn new(strategy: BalanceStrategy) -> Self { - Self { - strategy, - pools: DashMap::new(), - round_robin_indices: DashMap::new(), - health_checker: HealthChecker::with_defaults(), - proxy_factory: ProxyClientFactory::new(), - } - } - - /// 创建使用轮询策略的负载均衡器 - pub fn round_robin() -> Self { - Self::new(BalanceStrategy::RoundRobin) - } - - /// 创建带自定义健康检查配置的负载均衡器 - pub fn with_health_config(strategy: BalanceStrategy, health_config: HealthCheckConfig) -> Self { - Self { - strategy, - pools: DashMap::new(), - round_robin_indices: DashMap::new(), - health_checker: HealthChecker::new(health_config), - proxy_factory: ProxyClientFactory::new(), - } - } - - /// 创建带全局代理的负载均衡器 - pub fn with_global_proxy(mut self, proxy_url: Option) -> Self { - self.proxy_factory = self.proxy_factory.with_global_proxy(proxy_url); - self - } - - /// 设置全局代理 - pub fn set_global_proxy(&mut self, proxy_url: Option) { - self.proxy_factory = ProxyClientFactory::new().with_global_proxy(proxy_url); - } - - /// 获取代理客户端工厂 - pub fn proxy_factory(&self) -> &ProxyClientFactory { - &self.proxy_factory - } - - /// 获取健康检查器 - pub fn health_checker(&self) -> &HealthChecker { - &self.health_checker - } - - /// 获取当前策略 - pub fn strategy(&self) -> BalanceStrategy { - self.strategy - } - - /// 设置负载均衡策略 - pub fn set_strategy(&mut self, strategy: BalanceStrategy) { - self.strategy = strategy; - } - - /// 注册凭证池 - pub fn register_pool(&self, pool: Arc) { - let provider = pool.provider(); - self.pools.insert(provider, pool); - self.round_robin_indices - .insert(provider, AtomicUsize::new(0)); - } - - /// 获取凭证池 - pub fn get_pool(&self, provider: ProviderType) -> Option> { - self.pools.get(&provider).map(|r| r.value().clone()) - } - - /// 移除凭证池 - pub fn remove_pool(&self, provider: ProviderType) -> Option> { - self.round_robin_indices.remove(&provider); - self.pools.remove(&provider).map(|(_, pool)| pool) - } - - /// 获取所有已注册的 Provider - pub fn providers(&self) -> Vec { - self.pools.iter().map(|r| *r.key()).collect() - } - - /// 选择下一个可用凭证(使用当前策略) - pub fn select(&self, provider: ProviderType) -> Result { - let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?; - pool.refresh_cooldowns(); - match self.strategy { - BalanceStrategy::RoundRobin => self.select_round_robin(&pool, provider), - BalanceStrategy::LeastUsed => self.select_least_used(&pool), - BalanceStrategy::Random => self.select_random(&pool), - } - } - - /// 选择下一个可用凭证并创建配置了代理的 HTTP 客户端 - pub fn select_with_client( - &self, - provider: ProviderType, - ) -> Result { - let credential = self.select(provider)?; - let client = self - .proxy_factory - .create_client(credential.proxy_url()) - .map_err(|e| PoolError::CredentialNotFound(format!("代理配置错误: {e}")))?; - Ok(CredentialSelection { credential, client }) - } - - /// 为指定凭证创建配置了代理的 HTTP 客户端 - pub fn create_client_for_credential( - &self, - credential: &Credential, - ) -> Result { - self.proxy_factory - .create_client(credential.proxy_url()) - .map_err(|e| PoolError::CredentialNotFound(format!("代理配置错误: {e}"))) - } - - /// 选择下一个可用凭证,支持代理失败时的故障转移 - pub fn select_with_failover( - &self, - provider: ProviderType, - max_attempts: Option, - ) -> Result { - let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?; - pool.refresh_cooldowns(); - - let active_count = pool.active_count(); - if active_count == 0 { - return Err(PoolError::NoAvailableCredential); - } - - let attempts = max_attempts.unwrap_or(active_count).min(active_count); - let mut last_error = None; - let mut tried_ids = std::collections::HashSet::new(); - - for _ in 0..attempts { - let credential = match self.select(provider) { - Ok(cred) => cred, - Err(e) => { - last_error = Some(e); - break; - } - }; - - if tried_ids.contains(&credential.id) { - continue; - } - tried_ids.insert(credential.id.clone()); - - match self.proxy_factory.create_client(credential.proxy_url()) { - Ok(client) => { - return Ok(CredentialSelection { credential, client }); - } - Err(e) => { - tracing::warn!( - credential_id = %credential.id, - proxy_url = ?credential.proxy_url(), - error = %e, - "代理连接失败,尝试下一个凭证" - ); - last_error = Some(PoolError::CredentialNotFound(format!( - "凭证 {} 的代理配置错误: {}", - credential.id, e - ))); - } - } - } - - Err(last_error.unwrap_or(PoolError::NoAvailableCredential)) - } - - /// 报告代理连接失败并尝试故障转移 - pub fn failover_on_proxy_error( - &self, - provider: ProviderType, - failed_credential_id: &str, - ) -> Result { - let _ = self.report(provider, failed_credential_id, false, 0); - tracing::warn!( - credential_id = %failed_credential_id, - provider = %provider, - "代理连接失败,执行故障转移" - ); - self.select_with_client(provider) - } - - /// 轮询选择凭证 - fn select_round_robin( - &self, - pool: &CredentialPool, - provider: ProviderType, - ) -> Result { - let active_creds: Vec = pool - .all() - .into_iter() - .filter(|c| c.is_available()) - .collect(); - - if active_creds.is_empty() { - return Err(PoolError::NoAvailableCredential); - } - - let index_entry = self - .round_robin_indices - .entry(provider) - .or_insert_with(|| AtomicUsize::new(0)); - - let index = index_entry.fetch_add(1, Ordering::SeqCst) % active_creds.len(); - Ok(active_creds[index].clone()) - } - - /// 最少使用选择凭证 - fn select_least_used(&self, pool: &CredentialPool) -> Result { - pool.all() - .into_iter() - .filter(|c| c.is_available()) - .min_by_key(|c| c.stats.total_requests) - .ok_or(PoolError::NoAvailableCredential) - } - - /// 随机选择凭证 - fn select_random(&self, pool: &CredentialPool) -> Result { - let active_creds: Vec = pool - .all() - .into_iter() - .filter(|c| c.is_available()) - .collect(); - - if active_creds.is_empty() { - return Err(PoolError::NoAvailableCredential); - } - - let now = Utc::now().timestamp_nanos_opt().unwrap_or(0) as usize; - let index = now % active_creds.len(); - Ok(active_creds[index].clone()) - } - - /// 标记凭证为冷却状态 - pub fn mark_cooldown( - &self, - provider: ProviderType, - credential_id: &str, - duration: Duration, - ) -> Result<(), PoolError> { - let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?; - pool.mark_cooldown(credential_id, duration) - } - - /// 恢复凭证为活跃状态 - pub fn mark_active( - &self, - provider: ProviderType, - credential_id: &str, - ) -> Result<(), PoolError> { - let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?; - pool.mark_active(credential_id) - } - - /// 刷新所有池的冷却状态 - pub fn refresh_all_cooldowns(&self) { - for pool in self.pools.iter() { - pool.refresh_cooldowns(); - } - } - - /// 报告凭证使用结果 - pub fn report( - &self, - provider: ProviderType, - credential_id: &str, - success: bool, - latency_ms: u64, - ) -> Result { - let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?; - if success { - self.health_checker - .record_success(&pool, credential_id, latency_ms) - } else { - self.health_checker.record_failure(&pool, credential_id) - } - } - - /// 获取 Provider 的最早恢复时间 - pub fn earliest_recovery(&self, provider: ProviderType) -> Option> { - self.pools - .get(&provider) - .and_then(|pool| pool.earliest_recovery()) - } - - /// 检查是否有可用凭证 - pub fn has_available(&self, provider: ProviderType) -> bool { - self.pools - .get(&provider) - .map(|pool| pool.active_count() > 0) - .unwrap_or(false) - } - - /// 获取活跃凭证数量 - pub fn active_count(&self, provider: ProviderType) -> usize { - self.pools - .get(&provider) - .map(|pool| pool.active_count()) - .unwrap_or(0) - } -} - -impl Default for LoadBalancer { - fn default() -> Self { - Self::round_robin() - } -} - -#[cfg(test)] -mod balancer_tests { - use super::*; - use lime_core::credential::types::CredentialData; - - fn create_test_credential(id: &str, provider: ProviderType) -> Credential { - Credential::new( - id.to_string(), - provider, - CredentialData::ApiKey { - key: format!("key-{id}"), - base_url: None, - }, - ) - } - - #[test] - fn test_load_balancer_new() { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - assert_eq!(lb.strategy(), BalanceStrategy::RoundRobin); - assert!(lb.providers().is_empty()); - } - - #[test] - fn test_load_balancer_register_pool() { - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool.clone()); - assert!(lb.providers().contains(&ProviderType::Kiro)); - assert!(lb.get_pool(ProviderType::Kiro).is_some()); - } - - #[test] - fn test_load_balancer_select_round_robin() { - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - pool.add(create_test_credential("cred-2", ProviderType::Kiro)) - .unwrap(); - pool.add(create_test_credential("cred-3", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool); - - let c1 = lb.select(ProviderType::Kiro).unwrap(); - let c2 = lb.select(ProviderType::Kiro).unwrap(); - let c3 = lb.select(ProviderType::Kiro).unwrap(); - let ids: std::collections::HashSet<_> = [c1.id, c2.id, c3.id].into_iter().collect(); - assert_eq!(ids.len(), 3); - } - - #[test] - fn test_load_balancer_select_empty_pool() { - let lb = LoadBalancer::round_robin(); - let result = lb.select(ProviderType::Kiro); - assert!(matches!(result, Err(PoolError::EmptyPool))); - } - - #[test] - fn test_load_balancer_cooldown() { - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - pool.add(create_test_credential("cred-2", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool); - - lb.mark_cooldown(ProviderType::Kiro, "cred-1", Duration::hours(1)) - .unwrap(); - assert_eq!(lb.active_count(ProviderType::Kiro), 1); - - let selected = lb.select(ProviderType::Kiro).unwrap(); - assert_eq!(selected.id, "cred-2"); - } - - #[test] - fn test_load_balancer_report() { - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool.clone()); - - let changed = lb.report(ProviderType::Kiro, "cred-1", true, 100).unwrap(); - assert!(!changed); - let cred = pool.get("cred-1").unwrap(); - assert_eq!(cred.stats.total_requests, 1); - assert_eq!(cred.stats.successful_requests, 1); - - let changed = lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap(); - assert!(!changed); - let cred = pool.get("cred-1").unwrap(); - assert_eq!(cred.stats.total_requests, 2); - assert_eq!(cred.stats.consecutive_failures, 1); - } - - #[test] - fn test_load_balancer_auto_unhealthy() { - use lime_core::credential::types::CredentialStatus; - - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool.clone()); - - assert!(!lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap()); - assert!(!lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap()); - assert!(lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap()); - - let cred = pool.get("cred-1").unwrap(); - assert!(matches!(cred.status, CredentialStatus::Unhealthy { .. })); - } - - #[test] - fn test_load_balancer_auto_recovery() { - use lime_core::credential::types::CredentialStatus; - - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool.clone()); - - pool.mark_unhealthy("cred-1", "test".to_string()).unwrap(); - let recovered = lb.report(ProviderType::Kiro, "cred-1", true, 100).unwrap(); - assert!(recovered); - - let cred = pool.get("cred-1").unwrap(); - assert!(matches!(cred.status, CredentialStatus::Active)); - } - - #[test] - fn test_load_balancer_cooldown_recovery() { - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool.clone()); - - { - let mut entry = pool.credentials.get_mut("cred-1").unwrap(); - entry.status = lime_core::credential::types::CredentialStatus::Cooldown { - until: Utc::now() - Duration::seconds(1), - }; - } - - let cred = pool.get("cred-1").unwrap(); - assert!(matches!( - cred.status, - lime_core::credential::types::CredentialStatus::Cooldown { .. } - )); - - let selected = lb.select(ProviderType::Kiro).unwrap(); - assert_eq!(selected.id, "cred-1"); - - let cred = pool.get("cred-1").unwrap(); - assert!(matches!( - cred.status, - lime_core::credential::types::CredentialStatus::Active - )); - } - - #[test] - fn test_load_balancer_earliest_recovery() { - let lb = LoadBalancer::round_robin(); - let pool = Arc::new(CredentialPool::new(ProviderType::Kiro)); - pool.add(create_test_credential("cred-1", ProviderType::Kiro)) - .unwrap(); - pool.add(create_test_credential("cred-2", ProviderType::Kiro)) - .unwrap(); - lb.register_pool(pool); - - assert!(lb.earliest_recovery(ProviderType::Kiro).is_none()); - - lb.mark_cooldown(ProviderType::Kiro, "cred-1", Duration::hours(2)) - .unwrap(); - lb.mark_cooldown(ProviderType::Kiro, "cred-2", Duration::hours(1)) - .unwrap(); - - let recovery = lb.earliest_recovery(ProviderType::Kiro); - assert!(recovery.is_some()); - - let expected = Utc::now() + Duration::hours(1); - let diff = (recovery.unwrap() - expected).num_seconds().abs(); - assert!( - diff < 5, - "Recovery time should be approximately 1 hour from now" - ); - } -} diff --git a/src-tauri/crates/credential/src/encryption.rs b/src-tauri/crates/credential/src/encryption.rs deleted file mode 100644 index ce87831bc..000000000 --- a/src-tauri/crates/credential/src/encryption.rs +++ /dev/null @@ -1,284 +0,0 @@ -//! ChaCha20-Poly1305 AEAD 加密模块 -//! -//! 提供凭证加密/解密功能: -//! - ChaCha20-Poly1305 认证加密(防篡改) -//! - 随机 nonce(每次加密生成新的 12 字节 nonce) -//! - 密钥派生(SHA-256) -//! - 格式:enc2:base64(nonce || ciphertext || tag) - -use base64::{engine::general_purpose::STANDARD as BASE64, Engine}; -use chacha20poly1305::{ - aead::{Aead, KeyInit, OsRng}, - ChaCha20Poly1305, Nonce, -}; -use sha2::{Digest, Sha256}; - -/// 加密前缀标识 -const ENCRYPTED_PREFIX: &str = "enc2:"; - -/// Nonce 长度(12 字节) -const NONCE_SIZE: usize = 12; - -/// 加密器 -pub struct Encryptor { - cipher: ChaCha20Poly1305, -} - -impl Encryptor { - /// 从密码/密钥创建加密器 - /// - /// 使用 SHA-256 将任意长度的密钥派生为 256-bit 密钥 - pub fn new(key: &str) -> Self { - let derived_key = Self::derive_key(key); - let cipher = ChaCha20Poly1305::new(&derived_key.into()); - Self { cipher } - } - - /// 从原始 32 字节密钥创建加密器 - pub fn from_raw_key(key: &[u8; 32]) -> Self { - let cipher = ChaCha20Poly1305::new(key.into()); - Self { cipher } - } - - /// 使用 SHA-256 派生 256-bit 密钥 - fn derive_key(password: &str) -> [u8; 32] { - let mut hasher = Sha256::new(); - hasher.update(password.as_bytes()); - let result = hasher.finalize(); - let mut key = [0u8; 32]; - key.copy_from_slice(&result); - key - } - - /// 加密明文 - /// - /// 返回格式:enc2:base64(nonce || ciphertext) - pub fn encrypt(&self, plaintext: &str) -> Result { - use chacha20poly1305::aead::AeadCore; - - // 生成随机 nonce - let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng); - - // 加密 - let ciphertext = self - .cipher - .encrypt(&nonce, plaintext.as_bytes()) - .map_err(|_| EncryptionError::EncryptionFailed)?; - - // 组合 nonce + ciphertext - let mut combined = Vec::with_capacity(NONCE_SIZE + ciphertext.len()); - combined.extend_from_slice(&nonce); - combined.extend_from_slice(&ciphertext); - - // Base64 编码并添加前缀 - Ok(format!("{}{}", ENCRYPTED_PREFIX, BASE64.encode(&combined))) - } - - /// 解密密文 - /// - /// 输入格式:enc2:base64(nonce || ciphertext) - pub fn decrypt(&self, encrypted: &str) -> Result { - // 检查前缀 - let encoded = encrypted - .strip_prefix(ENCRYPTED_PREFIX) - .ok_or(EncryptionError::InvalidFormat)?; - - // Base64 解码 - let combined = BASE64 - .decode(encoded) - .map_err(|_| EncryptionError::InvalidBase64)?; - - // 分离 nonce 和 ciphertext - if combined.len() < NONCE_SIZE { - return Err(EncryptionError::InvalidFormat); - } - - let (nonce_bytes, ciphertext) = combined.split_at(NONCE_SIZE); - let nonce = Nonce::from_slice(nonce_bytes); - - // 解密 - let plaintext = self - .cipher - .decrypt(nonce, ciphertext) - .map_err(|_| EncryptionError::DecryptionFailed)?; - - String::from_utf8(plaintext).map_err(|_| EncryptionError::InvalidUtf8) - } - - /// 检查文本是否已加密 - pub fn is_encrypted(text: &str) -> bool { - text.starts_with(ENCRYPTED_PREFIX) - } - - /// 加密(如果尚未加密) - pub fn encrypt_if_needed(&self, text: &str) -> Result { - if Self::is_encrypted(text) { - Ok(text.to_string()) - } else { - self.encrypt(text) - } - } - - /// 解密(如果已加密) - pub fn decrypt_if_needed(&self, text: &str) -> Result { - if Self::is_encrypted(text) { - self.decrypt(text) - } else { - Ok(text.to_string()) - } - } -} - -/// 加密错误 -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum EncryptionError { - /// 加密失败 - EncryptionFailed, - /// 解密失败(密钥错误或数据被篡改) - DecryptionFailed, - /// 无效的格式(缺少 enc2: 前缀) - InvalidFormat, - /// 无效的 Base64 编码 - InvalidBase64, - /// 无效的 UTF-8 - InvalidUtf8, -} - -impl std::fmt::Display for EncryptionError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - Self::EncryptionFailed => write!(f, "加密失败"), - Self::DecryptionFailed => write!(f, "解密失败:密钥错误或数据被篡改"), - Self::InvalidFormat => write!(f, "无效的加密格式"), - Self::InvalidBase64 => write!(f, "无效的 Base64 编码"), - Self::InvalidUtf8 => write!(f, "无效的 UTF-8 编码"), - } - } -} - -impl std::error::Error for EncryptionError {} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_encrypt_decrypt_roundtrip() { - let enc = Encryptor::new("test-password"); - let plaintext = "sk-abc123-secret-api-key"; - let encrypted = enc.encrypt(plaintext).unwrap(); - assert!(encrypted.starts_with(ENCRYPTED_PREFIX)); - let decrypted = enc.decrypt(&encrypted).unwrap(); - assert_eq!(decrypted, plaintext); - } - - #[test] - fn test_different_nonces() { - let enc = Encryptor::new("test-password"); - let plaintext = "same-plaintext"; - let encrypted1 = enc.encrypt(plaintext).unwrap(); - let encrypted2 = enc.encrypt(plaintext).unwrap(); - assert_ne!(encrypted1, encrypted2); - // 两者都能正确解密 - assert_eq!(enc.decrypt(&encrypted1).unwrap(), plaintext); - assert_eq!(enc.decrypt(&encrypted2).unwrap(), plaintext); - } - - #[test] - fn test_wrong_key_fails() { - let enc1 = Encryptor::new("correct-password"); - let enc2 = Encryptor::new("wrong-password"); - let encrypted = enc1.encrypt("secret").unwrap(); - assert_eq!( - enc2.decrypt(&encrypted), - Err(EncryptionError::DecryptionFailed) - ); - } - - #[test] - fn test_is_encrypted() { - assert!(Encryptor::is_encrypted("enc2:abc123")); - assert!(!Encryptor::is_encrypted("plain-text")); - assert!(!Encryptor::is_encrypted("enc1:old-format")); - assert!(!Encryptor::is_encrypted("")); - } - - #[test] - fn test_encrypt_if_needed_already_encrypted() { - let enc = Encryptor::new("key"); - let already = "enc2:already-encrypted-data"; - let result = enc.encrypt_if_needed(already).unwrap(); - assert_eq!(result, already); - } - - #[test] - fn test_decrypt_if_needed_not_encrypted() { - let enc = Encryptor::new("key"); - let plain = "not-encrypted"; - let result = enc.decrypt_if_needed(plain).unwrap(); - assert_eq!(result, plain); - } - - #[test] - fn test_invalid_format() { - let enc = Encryptor::new("key"); - assert_eq!( - enc.decrypt("no-prefix"), - Err(EncryptionError::InvalidFormat) - ); - } - - #[test] - fn test_invalid_base64() { - let enc = Encryptor::new("key"); - assert_eq!( - enc.decrypt("enc2:!!!invalid-base64!!!"), - Err(EncryptionError::InvalidBase64) - ); - } - - #[test] - fn test_tampered_data() { - let enc = Encryptor::new("key"); - let encrypted = enc.encrypt("secret").unwrap(); - // 篡改密文中的一个字符 - let encoded = encrypted.strip_prefix(ENCRYPTED_PREFIX).unwrap(); - let mut bytes = BASE64.decode(encoded).unwrap(); - if let Some(last) = bytes.last_mut() { - *last ^= 0xFF; - } - let tampered = format!("{}{}", ENCRYPTED_PREFIX, BASE64.encode(&bytes)); - assert_eq!( - enc.decrypt(&tampered), - Err(EncryptionError::DecryptionFailed) - ); - } - - #[test] - fn test_empty_string() { - let enc = Encryptor::new("key"); - let encrypted = enc.encrypt("").unwrap(); - assert_eq!(enc.decrypt(&encrypted).unwrap(), ""); - } - - #[test] - fn test_unicode_content() { - let enc = Encryptor::new("密钥"); - let plaintext = "你好世界 🌍 こんにちは"; - let encrypted = enc.encrypt(plaintext).unwrap(); - assert_eq!(enc.decrypt(&encrypted).unwrap(), plaintext); - } - - #[test] - fn test_from_raw_key() { - let raw_key: [u8; 32] = [ - 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, - 0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, - 0x1d, 0x1e, 0x1f, 0x20, - ]; - let enc = Encryptor::from_raw_key(&raw_key); - let plaintext = "raw-key-test"; - let encrypted = enc.encrypt(plaintext).unwrap(); - assert_eq!(enc.decrypt(&encrypted).unwrap(), plaintext); - } -} diff --git a/src-tauri/crates/credential/src/lib.rs b/src-tauri/crates/credential/src/lib.rs deleted file mode 100644 index aa59782c0..000000000 --- a/src-tauri/crates/credential/src/lib.rs +++ /dev/null @@ -1,22 +0,0 @@ -//! 凭证池管理 crate -//! -//! 提供负载均衡、配额管理和凭证同步功能 -//! -//! ## 模块结构 -//! -//! - `balancer` - 负载均衡策略(轮询、最少使用、随机) -//! - `quota` - 配额超限检测、自动切换和冷却恢复 -//! - `sync` - 凭证与 YAML 配置文件的同步 - -mod balancer; -pub mod encryption; -mod quota; -mod sync; - -// 重新导出 -pub use balancer::{BalanceStrategy, CooldownInfo, CredentialSelection, LoadBalancer}; -pub use quota::{ - create_shared_quota_manager, start_quota_cleanup_task, AllCredentialsExhaustedError, - QuotaAutoSwitchResult, QuotaExceededRecord, QuotaManager, -}; -pub use sync::{CredentialSyncService, SyncError}; diff --git a/src-tauri/crates/credential/src/quota.rs b/src-tauri/crates/credential/src/quota.rs deleted file mode 100644 index 26f1bbf98..000000000 --- a/src-tauri/crates/credential/src/quota.rs +++ /dev/null @@ -1,890 +0,0 @@ -//! 配额管理器实现 -//! -//! 提供配额超限检测、自动切换和冷却恢复功能 - -use chrono::{DateTime, Duration, Utc}; -use dashmap::DashMap; -use lime_core::config::QuotaExceededConfig; -use lime_infra::resilience::{QUOTA_EXCEEDED_KEYWORDS, QUOTA_EXCEEDED_STATUS_CODES}; -use serde::{Deserialize, Serialize}; -use std::sync::Arc; - -/// 配额超限记录 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct QuotaExceededRecord { - /// 凭证 ID - pub credential_id: String, - /// 超限时间 - pub exceeded_at: DateTime, - /// 冷却结束时间 - pub cooldown_until: DateTime, - /// 超限原因 - pub reason: String, -} - -/// 配额管理器 -#[derive(Debug)] -pub struct QuotaManager { - /// 配额超限配置 - config: QuotaExceededConfig, - /// 超限凭证记录(credential_id -> record) - exceeded_credentials: DashMap, -} - -impl QuotaManager { - /// 创建新的配额管理器 - pub fn new(config: QuotaExceededConfig) -> Self { - Self { - config, - exceeded_credentials: DashMap::new(), - } - } - - /// 使用默认配置创建配额管理器 - pub fn with_defaults() -> Self { - Self::new(QuotaExceededConfig::default()) - } - - /// 获取配置 - pub fn config(&self) -> &QuotaExceededConfig { - &self.config - } - - /// 更新配置 - pub fn set_config(&mut self, config: QuotaExceededConfig) { - self.config = config; - } - - /// 获取冷却时长 - pub fn cooldown_duration(&self) -> Duration { - Duration::seconds(self.config.cooldown_seconds as i64) - } - - /// 标记凭证为配额超限 - pub fn mark_quota_exceeded(&self, credential_id: &str, reason: &str) -> QuotaExceededRecord { - let now = Utc::now(); - let cooldown_until = now + self.cooldown_duration(); - - let record = QuotaExceededRecord { - credential_id: credential_id.to_string(), - exceeded_at: now, - cooldown_until, - reason: reason.to_string(), - }; - - self.exceeded_credentials - .insert(credential_id.to_string(), record.clone()); - - tracing::info!( - credential_id = %credential_id, - cooldown_until = %cooldown_until, - reason = %reason, - "凭证配额超限,已标记冷却" - ); - - record - } - - /// 检查凭证是否可用(未超限或已过冷却期) - pub fn is_available(&self, credential_id: &str) -> bool { - match self.exceeded_credentials.get(credential_id) { - Some(record) => { - let now = Utc::now(); - if now >= record.cooldown_until { - drop(record); - self.exceeded_credentials.remove(credential_id); - true - } else { - false - } - } - None => true, - } - } - - /// 获取凭证的冷却结束时间 - pub fn get_cooldown_until(&self, credential_id: &str) -> Option> { - self.exceeded_credentials - .get(credential_id) - .map(|r| r.cooldown_until) - } - - /// 设置凭证的冷却结束时间(用于测试) - pub fn set_cooldown_until(&self, credential_id: &str, until: DateTime) { - if let Some(mut record) = self.exceeded_credentials.get_mut(credential_id) { - record.cooldown_until = until; - } - } - - /// 获取凭证的超限记录 - pub fn get_record(&self, credential_id: &str) -> Option { - self.exceeded_credentials - .get(credential_id) - .map(|r| r.clone()) - } - - /// 清理过期的冷却记录 - pub fn cleanup_expired(&self) -> usize { - let now = Utc::now(); - let mut cleaned = 0; - - let expired_ids: Vec = self - .exceeded_credentials - .iter() - .filter(|r| now >= r.cooldown_until) - .map(|r| r.credential_id.clone()) - .collect(); - - for id in expired_ids { - self.exceeded_credentials.remove(&id); - cleaned += 1; - tracing::debug!(credential_id = %id, "凭证冷却期已过,已恢复可用"); - } - - if cleaned > 0 { - tracing::info!(count = cleaned, "已清理过期的配额超限记录"); - } - - cleaned - } - - /// 手动恢复凭证(移除冷却状态) - pub fn restore_credential(&self, credential_id: &str) -> bool { - self.exceeded_credentials.remove(credential_id).is_some() - } - - /// 获取所有处于冷却期的凭证 ID - pub fn get_exceeded_credentials(&self) -> Vec { - self.exceeded_credentials - .iter() - .map(|r| r.credential_id.clone()) - .collect() - } - - /// 获取超限凭证数量 - pub fn exceeded_count(&self) -> usize { - self.exceeded_credentials.len() - } - - /// 检查是否为配额超限错误 - pub fn is_quota_exceeded_error(status_code: Option, error_message: &str) -> bool { - if let Some(code) = status_code { - if QUOTA_EXCEEDED_STATUS_CODES.contains(&code) { - return true; - } - } - - let error_lower = error_message.to_lowercase(); - for keyword in QUOTA_EXCEEDED_KEYWORDS { - if error_lower.contains(keyword) { - return true; - } - } - - false - } - - /// 获取预览模型名称 - pub fn get_preview_model(&self, model: &str) -> Option { - if !self.config.switch_preview_model { - return None; - } - if Self::is_preview_model(model) { - return None; - } - Some(format!("{model}-preview")) - } - - /// 检查模型是否为预览版本 - pub fn is_preview_model(model: &str) -> bool { - model.ends_with("-preview") || model.contains("-preview-") - } - - /// 获取原始模型名称(从预览版本) - pub fn get_original_model(model: &str) -> Option { - if !Self::is_preview_model(model) { - return None; - } - model.find("-preview").map(|pos| model[..pos].to_string()) - } - - /// 检查是否启用自动切换项目 - pub fn is_switch_project_enabled(&self) -> bool { - self.config.switch_project - } - - /// 检查是否启用预览模型回退 - pub fn is_switch_preview_model_enabled(&self) -> bool { - self.config.switch_preview_model - } - - /// 获取最早的恢复时间 - pub fn earliest_recovery(&self) -> Option> { - self.exceeded_credentials - .iter() - .map(|r| r.cooldown_until) - .min() - } - - /// 获取剩余冷却时间(秒) - pub fn remaining_cooldown_seconds(&self, credential_id: &str) -> Option { - self.exceeded_credentials.get(credential_id).map(|r| { - let now = Utc::now(); - (r.cooldown_until - now).num_seconds() - }) - } -} - -impl Default for QuotaManager { - fn default() -> Self { - Self::with_defaults() - } -} - -/// 创建共享的配额管理器 -pub fn create_shared_quota_manager(config: QuotaExceededConfig) -> Arc { - Arc::new(QuotaManager::new(config)) -} - -/// 启动配额管理器的定期清理任务 -pub fn start_quota_cleanup_task( - manager: Arc, - interval_secs: u64, -) -> tokio::task::JoinHandle<()> { - tokio::spawn(async move { - let mut interval = tokio::time::interval(std::time::Duration::from_secs(interval_secs)); - loop { - interval.tick().await; - let cleaned = manager.cleanup_expired(); - if cleaned > 0 { - tracing::debug!(cleaned_count = cleaned, "定期清理配额超限记录完成"); - } - } - }) -} - -/// 配额自动切换结果 -#[derive(Debug, Clone)] -pub struct QuotaAutoSwitchResult { - /// 是否成功切换 - pub switched: bool, - /// 新的凭证 ID(如果切换成功) - pub new_credential_id: Option, - /// 是否使用了预览模型 - pub used_preview_model: bool, - /// 预览模型名称(如果使用了预览模型) - pub preview_model: Option, - /// 消息 - pub message: String, -} - -impl QuotaAutoSwitchResult { - pub fn switched(new_credential_id: String) -> Self { - let message = format!("已切换到凭证: {new_credential_id}"); - Self { - switched: true, - new_credential_id: Some(new_credential_id), - used_preview_model: false, - preview_model: None, - message, - } - } - - pub fn preview_model(model: String) -> Self { - let message = format!("已切换到预览模型: {model}"); - Self { - switched: false, - new_credential_id: None, - used_preview_model: true, - preview_model: Some(model), - message, - } - } - - pub fn not_switched(message: &str) -> Self { - Self { - switched: false, - new_credential_id: None, - used_preview_model: false, - preview_model: None, - message: message.to_string(), - } - } - - pub fn all_exhausted(earliest_recovery: Option>) -> Self { - let message = match earliest_recovery { - Some(time) => format!("所有凭证配额超限,最早恢复时间: {time}"), - None => "所有凭证配额超限,无可用凭证".to_string(), - }; - Self { - switched: false, - new_credential_id: None, - used_preview_model: false, - preview_model: None, - message, - } - } -} - -/// 所有凭证耗尽错误 -#[derive(Debug, Clone)] -pub struct AllCredentialsExhaustedError { - /// 最早恢复时间 - pub earliest_recovery: Option>, - /// 重试等待秒数(用于 Retry-After 头) - pub retry_after_seconds: Option, - /// 错误消息 - pub message: String, -} - -impl AllCredentialsExhaustedError { - pub fn new(earliest_recovery: Option>) -> Self { - let retry_after_seconds = earliest_recovery.map(|time| { - let now = Utc::now(); - if time > now { - (time - now).num_seconds().max(0) as u64 - } else { - 0 - } - }); - - let message = match earliest_recovery { - Some(time) => format!( - "所有凭证配额超限,最早恢复时间: {}", - time.format("%Y-%m-%d %H:%M:%S UTC") - ), - None => "所有凭证配额超限,无可用凭证".to_string(), - }; - - Self { - earliest_recovery, - retry_after_seconds, - message, - } - } - - pub fn status_code(&self) -> u16 { - 503 - } - - pub fn retry_after_header(&self) -> Option { - self.retry_after_seconds.map(|s| s.to_string()) - } -} - -impl std::fmt::Display for AllCredentialsExhaustedError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.message) - } -} - -impl std::error::Error for AllCredentialsExhaustedError {} - -/// 实现 IntoResponse 以便在 axum 处理器中直接返回 503 响应 -impl axum::response::IntoResponse for AllCredentialsExhaustedError { - fn into_response(self) -> axum::response::Response { - use axum::http::{header, StatusCode}; - use axum::Json; - - let json_body = serde_json::json!({ - "error": { - "message": self.message, - "type": "all_credentials_exhausted", - "code": 503, - "retry_after_seconds": self.retry_after_seconds - } - }); - - let mut response = (StatusCode::SERVICE_UNAVAILABLE, Json(json_body)).into_response(); - - if let Some(retry_after) = self.retry_after_header() { - if let Ok(header_value) = retry_after.parse() { - response - .headers_mut() - .insert(header::RETRY_AFTER, header_value); - } - } - - response - } -} - -impl QuotaManager { - /// 处理配额超限并尝试自动切换 - pub fn handle_quota_exceeded( - &self, - failed_credential_id: &str, - model: &str, - available_credential_ids: &[String], - error_message: &str, - ) -> QuotaAutoSwitchResult { - self.mark_quota_exceeded(failed_credential_id, error_message); - - if self.config.switch_project { - for cred_id in available_credential_ids { - if cred_id != failed_credential_id && self.is_available(cred_id) { - tracing::info!( - from_credential = %failed_credential_id, - to_credential = %cred_id, - "配额超限,自动切换凭证" - ); - return QuotaAutoSwitchResult::switched(cred_id.clone()); - } - } - } - - if self.config.switch_preview_model { - if let Some(preview) = self.get_preview_model(model) { - tracing::info!( - original_model = %model, - preview_model = %preview, - "配额超限,切换到预览模型" - ); - return QuotaAutoSwitchResult::preview_model(preview); - } - } - - let earliest = self.earliest_recovery(); - tracing::warn!( - credential_id = %failed_credential_id, - earliest_recovery = ?earliest, - "所有凭证配额超限" - ); - QuotaAutoSwitchResult::all_exhausted(earliest) - } - - /// 选择下一个可用凭证 - pub fn select_available_credential( - &self, - available_credential_ids: &[String], - ) -> Option { - for cred_id in available_credential_ids { - if self.is_available(cred_id) { - return Some(cred_id.clone()); - } - } - None - } - - /// 过滤出可用的凭证 ID 列表 - pub fn filter_available_credentials(&self, credential_ids: &[String]) -> Vec { - credential_ids - .iter() - .filter(|id| self.is_available(id)) - .cloned() - .collect() - } - - /// 检查是否所有凭证都已耗尽 - pub fn check_all_exhausted( - &self, - credential_ids: &[String], - ) -> Result<(), AllCredentialsExhaustedError> { - let available = self.filter_available_credentials(credential_ids); - if available.is_empty() { - Err(AllCredentialsExhaustedError::new(self.earliest_recovery())) - } else { - Ok(()) - } - } - - /// 获取所有凭证耗尽时的错误响应 - pub fn get_exhausted_error(&self) -> AllCredentialsExhaustedError { - AllCredentialsExhaustedError::new(self.earliest_recovery()) - } -} - -#[cfg(test)] -mod unit_tests { - use super::*; - - #[test] - fn test_quota_auto_switch_result_switched() { - let result = QuotaAutoSwitchResult::switched("cred-2".to_string()); - assert!(result.switched); - assert_eq!(result.new_credential_id, Some("cred-2".to_string())); - assert!(!result.used_preview_model); - assert!(result.preview_model.is_none()); - } - - #[test] - fn test_quota_auto_switch_result_preview_model() { - let result = QuotaAutoSwitchResult::preview_model("gemini-2.5-pro-preview".to_string()); - assert!(!result.switched); - assert!(result.new_credential_id.is_none()); - assert!(result.used_preview_model); - assert_eq!( - result.preview_model, - Some("gemini-2.5-pro-preview".to_string()) - ); - } - - #[test] - fn test_quota_auto_switch_result_not_switched() { - let result = QuotaAutoSwitchResult::not_switched("No available credentials"); - assert!(!result.switched); - assert!(result.new_credential_id.is_none()); - assert!(!result.used_preview_model); - assert!(result.preview_model.is_none()); - } - - #[test] - fn test_quota_auto_switch_result_all_exhausted() { - let result = QuotaAutoSwitchResult::all_exhausted(None); - assert!(!result.switched); - assert!(result.new_credential_id.is_none()); - assert!(!result.used_preview_model); - assert!(result.message.contains("无可用凭证")); - } - - #[test] - fn test_handle_quota_exceeded_switch_project() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: false, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - let available = vec![ - "cred-1".to_string(), - "cred-2".to_string(), - "cred-3".to_string(), - ]; - let result = manager.handle_quota_exceeded( - "cred-1", - "gemini-2.5-pro", - &available, - "Rate limit exceeded", - ); - assert!(result.switched); - assert_eq!(result.new_credential_id, Some("cred-2".to_string())); - assert!(!result.used_preview_model); - } - - #[test] - fn test_handle_quota_exceeded_switch_preview_model() { - let config = QuotaExceededConfig { - switch_project: false, - switch_preview_model: true, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - let available = vec!["cred-1".to_string()]; - let result = manager.handle_quota_exceeded( - "cred-1", - "gemini-2.5-pro", - &available, - "Rate limit exceeded", - ); - assert!(!result.switched); - assert!(result.used_preview_model); - assert_eq!( - result.preview_model, - Some("gemini-2.5-pro-preview".to_string()) - ); - } - - #[test] - fn test_handle_quota_exceeded_all_exhausted() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: false, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - manager.mark_quota_exceeded("cred-1", "test"); - manager.mark_quota_exceeded("cred-2", "test"); - let available = vec!["cred-1".to_string(), "cred-2".to_string()]; - let result = manager.handle_quota_exceeded( - "cred-1", - "gemini-2.5-pro", - &available, - "Rate limit exceeded", - ); - assert!(!result.switched); - assert!(!result.used_preview_model); - assert!(result.message.contains("所有凭证配额超限")); - } - - #[test] - fn test_select_available_credential() { - let manager = QuotaManager::with_defaults(); - manager.mark_quota_exceeded("cred-1", "test"); - let available = vec![ - "cred-1".to_string(), - "cred-2".to_string(), - "cred-3".to_string(), - ]; - let selected = manager.select_available_credential(&available); - assert_eq!(selected, Some("cred-2".to_string())); - } - - #[test] - fn test_filter_available_credentials() { - let manager = QuotaManager::with_defaults(); - manager.mark_quota_exceeded("cred-1", "test"); - manager.mark_quota_exceeded("cred-3", "test"); - let all = vec![ - "cred-1".to_string(), - "cred-2".to_string(), - "cred-3".to_string(), - "cred-4".to_string(), - ]; - let available = manager.filter_available_credentials(&all); - assert_eq!(available, vec!["cred-2".to_string(), "cred-4".to_string()]); - } - - #[test] - fn test_all_credentials_exhausted_error() { - let error = AllCredentialsExhaustedError::new(None); - assert_eq!(error.status_code(), 503); - assert!(error.retry_after_header().is_none()); - assert!(error.message.contains("无可用凭证")); - } - - #[test] - fn test_all_credentials_exhausted_error_with_recovery() { - let recovery_time = Utc::now() + Duration::seconds(300); - let error = AllCredentialsExhaustedError::new(Some(recovery_time)); - assert_eq!(error.status_code(), 503); - assert!(error.retry_after_header().is_some()); - let retry_after = error.retry_after_seconds.unwrap(); - assert!(retry_after > 0); - assert!(retry_after <= 300); - } - - #[test] - fn test_check_all_exhausted_has_available() { - let manager = QuotaManager::with_defaults(); - manager.mark_quota_exceeded("cred-1", "test"); - let all = vec!["cred-1".to_string(), "cred-2".to_string()]; - let result = manager.check_all_exhausted(&all); - assert!(result.is_ok()); - } - - #[test] - fn test_check_all_exhausted_none_available() { - let manager = QuotaManager::with_defaults(); - manager.mark_quota_exceeded("cred-1", "test"); - manager.mark_quota_exceeded("cred-2", "test"); - let all = vec!["cred-1".to_string(), "cred-2".to_string()]; - let result = manager.check_all_exhausted(&all); - assert!(result.is_err()); - let error = result.unwrap_err(); - assert_eq!(error.status_code(), 503); - assert!(error.earliest_recovery.is_some()); - } - - #[test] - fn test_get_exhausted_error() { - let manager = QuotaManager::with_defaults(); - manager.mark_quota_exceeded("cred-1", "test"); - let error = manager.get_exhausted_error(); - assert_eq!(error.status_code(), 503); - assert!(error.earliest_recovery.is_some()); - } - - #[test] - fn test_quota_manager_new() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - assert_eq!(manager.config().cooldown_seconds, 300); - assert!(manager.config().switch_project); - assert!(manager.config().switch_preview_model); - assert_eq!(manager.exceeded_count(), 0); - } - - #[test] - fn test_quota_manager_mark_exceeded() { - let manager = QuotaManager::with_defaults(); - let record = manager.mark_quota_exceeded("cred-1", "Rate limit exceeded"); - assert_eq!(record.credential_id, "cred-1"); - assert_eq!(record.reason, "Rate limit exceeded"); - assert!(record.cooldown_until > Utc::now()); - assert_eq!(manager.exceeded_count(), 1); - } - - #[test] - fn test_quota_manager_is_available() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 1, - }; - let manager = QuotaManager::new(config); - assert!(manager.is_available("cred-1")); - manager.mark_quota_exceeded("cred-1", "test"); - assert!(!manager.is_available("cred-1")); - std::thread::sleep(std::time::Duration::from_secs(2)); - assert!(manager.is_available("cred-1")); - } - - #[test] - fn test_quota_manager_cleanup_expired() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 0, - }; - let manager = QuotaManager::new(config); - manager.mark_quota_exceeded("cred-1", "test"); - manager.mark_quota_exceeded("cred-2", "test"); - manager.mark_quota_exceeded("cred-3", "test"); - assert_eq!(manager.exceeded_count(), 3); - std::thread::sleep(std::time::Duration::from_millis(100)); - let cleaned = manager.cleanup_expired(); - assert_eq!(cleaned, 3); - assert_eq!(manager.exceeded_count(), 0); - } - - #[test] - fn test_quota_manager_restore_credential() { - let manager = QuotaManager::with_defaults(); - manager.mark_quota_exceeded("cred-1", "test"); - assert!(!manager.is_available("cred-1")); - let restored = manager.restore_credential("cred-1"); - assert!(restored); - assert!(manager.is_available("cred-1")); - let restored = manager.restore_credential("cred-1"); - assert!(!restored); - } - - #[test] - fn test_quota_manager_is_quota_exceeded_error() { - assert!(QuotaManager::is_quota_exceeded_error(Some(429), "")); - assert!(QuotaManager::is_quota_exceeded_error( - Some(400), - "Rate limit exceeded" - )); - assert!(QuotaManager::is_quota_exceeded_error( - Some(400), - "Quota exceeded for this API" - )); - assert!(QuotaManager::is_quota_exceeded_error( - Some(400), - "Too many requests" - )); - assert!(!QuotaManager::is_quota_exceeded_error( - Some(400), - "Bad Request" - )); - assert!(!QuotaManager::is_quota_exceeded_error( - Some(500), - "Internal Server Error" - )); - } - - #[test] - fn test_quota_manager_get_preview_model() { - let manager = QuotaManager::with_defaults(); - assert_eq!( - manager.get_preview_model("gemini-2.5-pro"), - Some("gemini-2.5-pro-preview".to_string()) - ); - assert_eq!( - manager.get_preview_model("claude-3-opus"), - Some("claude-3-opus-preview".to_string()) - ); - assert_eq!(manager.get_preview_model("gemini-2.5-pro-preview"), None); - assert_eq!( - manager.get_preview_model("claude-3-opus-preview-20240101"), - None - ); - } - - #[test] - fn test_quota_manager_get_preview_model_disabled() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: false, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - assert_eq!(manager.get_preview_model("gemini-2.5-pro"), None); - } - - #[test] - fn test_is_preview_model() { - assert!(QuotaManager::is_preview_model("gemini-2.5-pro-preview")); - assert!(QuotaManager::is_preview_model( - "claude-3-opus-preview-20240101" - )); - assert!(QuotaManager::is_preview_model("gpt-4-preview")); - assert!(!QuotaManager::is_preview_model("gemini-2.5-pro")); - assert!(!QuotaManager::is_preview_model("claude-3-opus")); - assert!(!QuotaManager::is_preview_model("gpt-4")); - } - - #[test] - fn test_get_original_model() { - assert_eq!( - QuotaManager::get_original_model("gemini-2.5-pro-preview"), - Some("gemini-2.5-pro".to_string()) - ); - assert_eq!( - QuotaManager::get_original_model("claude-3-opus-preview-20240101"), - Some("claude-3-opus".to_string()) - ); - assert_eq!( - QuotaManager::get_original_model("gpt-4-preview"), - Some("gpt-4".to_string()) - ); - assert_eq!(QuotaManager::get_original_model("gemini-2.5-pro"), None); - assert_eq!(QuotaManager::get_original_model("claude-3-opus"), None); - } - - #[test] - fn test_quota_manager_earliest_recovery() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - assert!(manager.earliest_recovery().is_none()); - manager.mark_quota_exceeded("cred-1", "test"); - let recovery = manager.earliest_recovery(); - assert!(recovery.is_some()); - } - - #[test] - fn test_quota_manager_remaining_cooldown_seconds() { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - assert!(manager.remaining_cooldown_seconds("cred-1").is_none()); - manager.mark_quota_exceeded("cred-1", "test"); - let remaining = manager.remaining_cooldown_seconds("cred-1"); - assert!(remaining.is_some()); - assert!(remaining.unwrap() > 0); - assert!(remaining.unwrap() <= 300); - } - - #[test] - fn test_all_credentials_exhausted_into_response() { - use axum::http::{header, StatusCode}; - use axum::response::IntoResponse; - - let error = AllCredentialsExhaustedError::new(None); - let response = error.into_response(); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); - assert!(response.headers().get(header::RETRY_AFTER).is_none()); - - let recovery_time = Utc::now() + Duration::seconds(300); - let error = AllCredentialsExhaustedError::new(Some(recovery_time)); - let response = error.into_response(); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); - let retry_after = response.headers().get(header::RETRY_AFTER); - assert!(retry_after.is_some()); - let retry_value: u64 = retry_after.unwrap().to_str().unwrap().parse().unwrap(); - assert!(retry_value > 0); - assert!(retry_value <= 300); - } -} diff --git a/src-tauri/crates/credential/src/sync.rs b/src-tauri/crates/credential/src/sync.rs deleted file mode 100644 index f76665d5c..000000000 --- a/src-tauri/crates/credential/src/sync.rs +++ /dev/null @@ -1,628 +0,0 @@ -//! 凭证同步服务 -//! -//! 负责将凭证池变更同步到 YAML 配置文件 -//! 实现凭证的添加、删除、更新操作与配置文件的同步 - -use lime_core::config::{ - expand_tilde, ApiKeyEntry, Config, ConfigError, ConfigManager, CredentialEntry, YamlService, -}; -use lime_core::models::provider_pool_model::{ - CredentialData, PoolProviderType, ProviderCredential, -}; -use std::path::PathBuf; -use std::sync::{Arc, RwLock}; - -/// 凭证同步服务错误类型 -#[derive(Debug, Clone)] -pub enum SyncError { - /// 配置错误 - ConfigError(String), - /// IO 错误 - IoError(String), - /// 凭证不存在 - CredentialNotFound(String), - /// 无效的凭证类型 - InvalidCredentialType(String), -} - -impl std::fmt::Display for SyncError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - SyncError::ConfigError(msg) => write!(f, "配置错误: {msg}"), - SyncError::IoError(msg) => write!(f, "IO 错误: {msg}"), - SyncError::CredentialNotFound(id) => write!(f, "凭证不存在: {id}"), - SyncError::InvalidCredentialType(msg) => write!(f, "无效的凭证类型: {msg}"), - } - } -} - -impl std::error::Error for SyncError {} - -impl From for SyncError { - fn from(err: ConfigError) -> Self { - SyncError::ConfigError(err.to_string()) - } -} - -impl From for SyncError { - fn from(err: std::io::Error) -> Self { - SyncError::IoError(err.to_string()) - } -} - -/// 凭证同步服务 -pub struct CredentialSyncService { - /// 配置管理器 - config_manager: Arc>, -} - -impl CredentialSyncService { - /// 创建新的凭证同步服务 - pub fn new(config_manager: Arc>) -> Self { - Self { config_manager } - } - - /// 获取当前配置 - fn get_config(&self) -> Result { - let manager = self - .config_manager - .read() - .map_err(|e| SyncError::ConfigError(format!("获取配置锁失败: {e}")))?; - Ok(manager.config().clone()) - } - - /// 更新配置并保存 - fn update_config(&self, config: Config) -> Result<(), SyncError> { - let mut manager = self - .config_manager - .write() - .map_err(|e| SyncError::ConfigError(format!("获取配置写锁失败: {e}")))?; - - let config_path = manager.config_path().to_path_buf(); - manager.set_config(config.clone()); - YamlService::save_preserve_comments(&config_path, &config)?; - Ok(()) - } - - /// 获取 auth_dir 的绝对路径 - pub fn get_auth_dir(&self) -> Result { - let config = self.get_config()?; - Ok(expand_tilde(&config.auth_dir)) - } - - /// 确保 auth_dir 目录存在 - pub fn ensure_auth_dir(&self) -> Result { - let auth_dir = self.get_auth_dir()?; - std::fs::create_dir_all(&auth_dir)?; - Ok(auth_dir) - } - - /// 添加凭证并同步到配置 - pub fn add_credential(&self, credential: &ProviderCredential) -> Result<(), SyncError> { - let mut config = self.get_config()?; - - match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - let token_file = - self.save_oauth_token_file(creds_file_path, &credential.uuid, "kiro")?; - let entry = CredentialEntry { - id: credential.uuid.clone(), - token_file, - disabled: credential.is_disabled, - proxy_url: None, - }; - config.credential_pool.kiro.push(entry); - } - CredentialData::GeminiOAuth { - creds_file_path, .. - } => { - let token_file = - self.save_oauth_token_file(creds_file_path, &credential.uuid, "gemini")?; - let entry = CredentialEntry { - id: credential.uuid.clone(), - token_file, - disabled: credential.is_disabled, - proxy_url: None, - }; - config.credential_pool.gemini.push(entry); - } - CredentialData::AntigravityOAuth { .. } => { - return Err(SyncError::InvalidCredentialType( - "Antigravity 凭证暂不支持同步到配置".to_string(), - )); - } - CredentialData::OpenAIKey { api_key, base_url } => { - let entry = ApiKeyEntry { - id: credential.uuid.clone(), - api_key: api_key.clone(), - base_url: base_url.clone(), - disabled: credential.is_disabled, - proxy_url: None, - }; - config.credential_pool.openai.push(entry); - } - CredentialData::ClaudeKey { api_key, base_url } => { - let entry = ApiKeyEntry { - id: credential.uuid.clone(), - api_key: api_key.clone(), - base_url: base_url.clone(), - disabled: credential.is_disabled, - proxy_url: None, - }; - config.credential_pool.claude.push(entry); - } - CredentialData::VertexKey { - api_key, - base_url, - model_aliases, - } => { - use lime_core::models::vertex_model::VertexModelAlias; - let models: Vec = model_aliases - .iter() - .map(|(alias, name)| VertexModelAlias { - alias: alias.clone(), - name: name.clone(), - }) - .collect(); - let entry = lime_core::models::vertex_model::VertexApiKeyEntry { - id: credential.uuid.clone(), - api_key: api_key.clone(), - base_url: base_url.clone(), - models, - proxy_url: None, - disabled: credential.is_disabled, - }; - config.credential_pool.vertex_api_keys.push(entry); - } - CredentialData::GeminiApiKey { - api_key, - base_url, - excluded_models, - } => { - use lime_core::config::GeminiApiKeyEntry; - let entry = GeminiApiKeyEntry { - id: credential.uuid.clone(), - api_key: api_key.clone(), - base_url: base_url.clone(), - proxy_url: None, - excluded_models: excluded_models.clone(), - disabled: credential.is_disabled, - }; - config.credential_pool.gemini_api_keys.push(entry); - } - CredentialData::CodexOAuth { .. } => { - return Err(SyncError::InvalidCredentialType( - "Codex 凭证暂不支持同步到配置".to_string(), - )); - } - CredentialData::ClaudeOAuth { .. } => { - return Err(SyncError::InvalidCredentialType( - "Claude OAuth 凭证暂不支持同步到配置".to_string(), - )); - } - CredentialData::AnthropicKey { api_key, base_url } => { - let entry = ApiKeyEntry { - id: credential.uuid.clone(), - api_key: api_key.clone(), - base_url: base_url.clone(), - disabled: credential.is_disabled, - proxy_url: None, - }; - config.credential_pool.claude.push(entry); - } - } - - self.update_config(config) - } - - /// 保存 OAuth token 文件到 auth_dir - fn save_oauth_token_file( - &self, - source_path: &str, - credential_id: &str, - provider: &str, - ) -> Result { - let auth_dir = self.ensure_auth_dir()?; - let provider_dir = auth_dir.join(provider); - std::fs::create_dir_all(&provider_dir)?; - - let token_filename = format!("{credential_id}.json"); - let token_path = provider_dir.join(&token_filename); - - let source = expand_tilde(source_path); - if source.exists() { - std::fs::copy(&source, &token_path)?; - } - - Ok(format!("{provider}/{token_filename}")) - } - - /// 删除 OAuth token 文件 - fn delete_oauth_token_file(&self, token_file: &str) -> Result<(), SyncError> { - let auth_dir = self.get_auth_dir()?; - let token_path = auth_dir.join(token_file); - if token_path.exists() { - std::fs::remove_file(&token_path)?; - } - Ok(()) - } - - /// 删除凭证并同步到配置 - pub fn remove_credential( - &self, - provider_type: PoolProviderType, - credential_id: &str, - ) -> Result<(), SyncError> { - let mut config = self.get_config()?; - let mut found = false; - - match provider_type { - PoolProviderType::Kiro => { - if let Some(pos) = config - .credential_pool - .kiro - .iter() - .position(|e| e.id == credential_id) - { - let entry = config.credential_pool.kiro.remove(pos); - self.delete_oauth_token_file(&entry.token_file)?; - found = true; - } - } - PoolProviderType::Gemini => { - if let Some(pos) = config - .credential_pool - .gemini - .iter() - .position(|e| e.id == credential_id) - { - let entry = config.credential_pool.gemini.remove(pos); - self.delete_oauth_token_file(&entry.token_file)?; - found = true; - } - } - PoolProviderType::OpenAI => { - if let Some(pos) = config - .credential_pool - .openai - .iter() - .position(|e| e.id == credential_id) - { - config.credential_pool.openai.remove(pos); - found = true; - } - } - PoolProviderType::Claude => { - if let Some(pos) = config - .credential_pool - .claude - .iter() - .position(|e| e.id == credential_id) - { - config.credential_pool.claude.remove(pos); - found = true; - } - } - PoolProviderType::Antigravity => { - return Err(SyncError::InvalidCredentialType( - "Antigravity 凭证暂不支持同步到配置".to_string(), - )); - } - PoolProviderType::Vertex => { - if let Some(pos) = config - .credential_pool - .vertex_api_keys - .iter() - .position(|e| e.id == credential_id) - { - config.credential_pool.vertex_api_keys.remove(pos); - found = true; - } - } - PoolProviderType::GeminiApiKey => { - if let Some(pos) = config - .credential_pool - .gemini_api_keys - .iter() - .position(|e| e.id == credential_id) - { - config.credential_pool.gemini_api_keys.remove(pos); - found = true; - } - } - PoolProviderType::Codex => { - return Err(SyncError::InvalidCredentialType( - "Codex 凭证暂不支持同步到配置".to_string(), - )); - } - PoolProviderType::ClaudeOAuth => { - return Err(SyncError::InvalidCredentialType( - "Claude OAuth 凭证暂不支持同步到配置".to_string(), - )); - } - PoolProviderType::AnthropicCompatible => { - return Err(SyncError::InvalidCredentialType( - "Anthropic Compatible 凭证暂不支持同步到配置".to_string(), - )); - } - PoolProviderType::Anthropic - | PoolProviderType::AzureOpenai - | PoolProviderType::AwsBedrock - | PoolProviderType::Ollama => { - return Err(SyncError::InvalidCredentialType( - "API Key Provider 凭证不支持同步到配置".to_string(), - )); - } - } - - if !found { - return Err(SyncError::CredentialNotFound(credential_id.to_string())); - } - - self.update_config(config) - } - - /// 更新凭证并同步到配置 - pub fn update_credential(&self, credential: &ProviderCredential) -> Result<(), SyncError> { - let mut config = self.get_config()?; - let mut found = false; - - match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - if let Some(entry) = config - .credential_pool - .kiro - .iter_mut() - .find(|e| e.id == credential.uuid) - { - entry.disabled = credential.is_disabled; - let new_token_file = - self.save_oauth_token_file(creds_file_path, &credential.uuid, "kiro")?; - entry.token_file = new_token_file; - found = true; - } - } - CredentialData::GeminiOAuth { - creds_file_path, .. - } => { - if let Some(entry) = config - .credential_pool - .gemini - .iter_mut() - .find(|e| e.id == credential.uuid) - { - entry.disabled = credential.is_disabled; - let new_token_file = - self.save_oauth_token_file(creds_file_path, &credential.uuid, "gemini")?; - entry.token_file = new_token_file; - found = true; - } - } - CredentialData::AntigravityOAuth { .. } => { - return Err(SyncError::InvalidCredentialType( - "Antigravity 凭证暂不支持同步到配置".to_string(), - )); - } - CredentialData::OpenAIKey { api_key, base_url } => { - if let Some(entry) = config - .credential_pool - .openai - .iter_mut() - .find(|e| e.id == credential.uuid) - { - entry.api_key = api_key.clone(); - entry.base_url = base_url.clone(); - entry.disabled = credential.is_disabled; - found = true; - } - } - CredentialData::ClaudeKey { api_key, base_url } => { - if let Some(entry) = config - .credential_pool - .claude - .iter_mut() - .find(|e| e.id == credential.uuid) - { - entry.api_key = api_key.clone(); - entry.base_url = base_url.clone(); - entry.disabled = credential.is_disabled; - found = true; - } - } - CredentialData::VertexKey { - api_key, - base_url, - model_aliases, - } => { - if let Some(entry) = config - .credential_pool - .vertex_api_keys - .iter_mut() - .find(|e| e.id == credential.uuid) - { - use lime_core::models::vertex_model::VertexModelAlias; - entry.api_key = api_key.clone(); - entry.base_url = base_url.clone(); - entry.models = model_aliases - .iter() - .map(|(alias, name)| VertexModelAlias { - alias: alias.clone(), - name: name.clone(), - }) - .collect(); - entry.disabled = credential.is_disabled; - found = true; - } - } - CredentialData::GeminiApiKey { - api_key, - base_url, - excluded_models, - } => { - if let Some(entry) = config - .credential_pool - .gemini_api_keys - .iter_mut() - .find(|e| e.id == credential.uuid) - { - entry.api_key = api_key.clone(); - entry.base_url = base_url.clone(); - entry.excluded_models = excluded_models.clone(); - entry.disabled = credential.is_disabled; - found = true; - } - } - CredentialData::CodexOAuth { .. } => { - return Err(SyncError::InvalidCredentialType( - "Codex 凭证暂不支持同步到配置".to_string(), - )); - } - CredentialData::ClaudeOAuth { .. } => { - return Err(SyncError::InvalidCredentialType( - "Claude OAuth 凭证暂不支持同步到配置".to_string(), - )); - } - CredentialData::AnthropicKey { api_key, base_url } => { - if let Some(entry) = config - .credential_pool - .claude - .iter_mut() - .find(|e| e.id == credential.uuid) - { - entry.api_key = api_key.clone(); - entry.base_url = base_url.clone(); - entry.disabled = credential.is_disabled; - found = true; - } - } - } - - if !found { - return Err(SyncError::CredentialNotFound(credential.uuid.clone())); - } - - self.update_config(config) - } - - /// 从配置加载凭证到池中 - pub fn load_from_config(&self) -> Result, SyncError> { - let config = self.get_config()?; - let auth_dir = self.get_auth_dir()?; - let mut credentials = Vec::new(); - - // 加载 Kiro 凭证 - for entry in &config.credential_pool.kiro { - let token_path = auth_dir.join(&entry.token_file); - let mut cred = ProviderCredential::new( - PoolProviderType::Kiro, - CredentialData::KiroOAuth { - creds_file_path: token_path.to_string_lossy().to_string(), - }, - ); - cred.uuid = entry.id.clone(); - cred.is_disabled = entry.disabled; - credentials.push(cred); - } - - // 加载 Gemini 凭证 - for entry in &config.credential_pool.gemini { - let token_path = auth_dir.join(&entry.token_file); - let mut cred = ProviderCredential::new( - PoolProviderType::Gemini, - CredentialData::GeminiOAuth { - creds_file_path: token_path.to_string_lossy().to_string(), - project_id: None, - }, - ); - cred.uuid = entry.id.clone(); - cred.is_disabled = entry.disabled; - credentials.push(cred); - } - - // 加载 OpenAI 凭证 - for entry in &config.credential_pool.openai { - let mut cred = ProviderCredential::new( - PoolProviderType::OpenAI, - CredentialData::OpenAIKey { - api_key: entry.api_key.clone(), - base_url: entry.base_url.clone(), - }, - ); - cred.uuid = entry.id.clone(); - cred.is_disabled = entry.disabled; - credentials.push(cred); - } - - // 加载 Claude 凭证 - for entry in &config.credential_pool.claude { - let mut cred = ProviderCredential::new( - PoolProviderType::Claude, - CredentialData::ClaudeKey { - api_key: entry.api_key.clone(), - base_url: entry.base_url.clone(), - }, - ); - cred.uuid = entry.id.clone(); - cred.is_disabled = entry.disabled; - credentials.push(cred); - } - - // 加载 Vertex AI 凭证 - for entry in &config.credential_pool.vertex_api_keys { - let model_aliases: std::collections::HashMap = entry - .models - .iter() - .map(|m| (m.alias.clone(), m.name.clone())) - .collect(); - let mut cred = ProviderCredential::new( - PoolProviderType::Vertex, - CredentialData::VertexKey { - api_key: entry.api_key.clone(), - base_url: entry.base_url.clone(), - model_aliases, - }, - ); - cred.uuid = entry.id.clone(); - cred.is_disabled = entry.disabled; - credentials.push(cred); - } - - // 加载 Gemini API Key 凭证 - for entry in &config.credential_pool.gemini_api_keys { - let mut cred = ProviderCredential::new( - PoolProviderType::GeminiApiKey, - CredentialData::GeminiApiKey { - api_key: entry.api_key.clone(), - base_url: entry.base_url.clone(), - excluded_models: entry.excluded_models.clone(), - }, - ); - cred.uuid = entry.id.clone(); - cred.is_disabled = entry.disabled; - credentials.push(cred); - } - - Ok(credentials) - } - - /// 获取 OAuth token 文件的完整路径 - pub fn get_token_file_path(&self, token_file: &str) -> Result { - let auth_dir = self.get_auth_dir()?; - Ok(auth_dir.join(token_file)) - } - - /// 读取 OAuth token 文件内容 - pub fn read_token_file(&self, token_file: &str) -> Result { - let path = self.get_token_file_path(token_file)?; - std::fs::read_to_string(&path).map_err(SyncError::from) - } - - /// 写入 OAuth token 文件内容 - pub fn write_token_file(&self, token_file: &str, content: &str) -> Result<(), SyncError> { - let path = self.get_token_file_path(token_file)?; - if let Some(parent) = path.parent() { - std::fs::create_dir_all(parent)?; - } - std::fs::write(&path, content).map_err(SyncError::from) - } -} diff --git a/src-tauri/crates/infra/src/resilience/timeout.rs b/src-tauri/crates/infra/src/resilience/timeout.rs index c866f4c46..fb13b4db5 100644 --- a/src-tauri/crates/infra/src/resilience/timeout.rs +++ b/src-tauri/crates/infra/src/resilience/timeout.rs @@ -22,8 +22,8 @@ pub struct TimeoutConfig { impl Default for TimeoutConfig { fn default() -> Self { Self { - request_timeout_ms: 120_000, // 2 分钟 - stream_idle_timeout_ms: 30_000, // 30 秒 + request_timeout_ms: 120_000, // 2 分钟 + stream_idle_timeout_ms: 120_000, // 2 分钟,兼容推理模型长时间思考 } } } @@ -434,7 +434,7 @@ mod unit_tests { fn test_timeout_config_default() { let config = TimeoutConfig::default(); assert_eq!(config.request_timeout_ms, 120_000); - assert_eq!(config.stream_idle_timeout_ms, 30_000); + assert_eq!(config.stream_idle_timeout_ms, 120_000); assert!(config.has_request_timeout()); assert!(config.has_stream_idle_timeout()); } diff --git a/src-tauri/crates/processor/src/processor.rs b/src-tauri/crates/processor/src/processor.rs index fa5f1d13a..495c2618c 100644 --- a/src-tauri/crates/processor/src/processor.rs +++ b/src-tauri/crates/processor/src/processor.rs @@ -19,7 +19,6 @@ use lime_core::plugin::PluginManager; use lime_core::router::{ModelMapper, Router}; use lime_core::ProviderType; use lime_infra::{Failover, Injector, Retrier, StatsAggregator, TimeoutController, TokenTracker}; -use lime_services::provider_pool_service::ProviderPoolService; use parking_lot::RwLock as ParkingLotRwLock; use std::sync::Arc; use tokio::sync::RwLock; @@ -46,8 +45,6 @@ pub struct RequestProcessor { pub stats: Arc>, /// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) pub tokens: Arc>, - /// 凭证池服务 - pub pool_service: Arc, /// 热重载协调锁(避免配置更新期间请求读取不一致的配置) pub reload_lock: Arc>, /// 提示路由器 @@ -68,7 +65,6 @@ impl RequestProcessor { plugins: Arc, stats: Arc>, tokens: Arc>, - pool_service: Arc, ) -> Self { Self { router, @@ -80,7 +76,6 @@ impl RequestProcessor { plugins, stats, tokens, - pool_service, reload_lock: Arc::new(RwLock::new(())), hint_router: Arc::new(RwLock::new(lime_core::router::HintRouter::default())), conversation_trimmer: Arc::new(crate::conversation_manager::ConversationTrimmer::new( @@ -90,7 +85,7 @@ impl RequestProcessor { } /// 使用默认配置创建请求处理器 - pub fn with_defaults(pool_service: Arc) -> Self { + pub fn with_defaults() -> Self { Self { router: Arc::new(RwLock::new(Self::create_router_with_defaults())), mapper: Arc::new(RwLock::new(ModelMapper::new())), @@ -101,7 +96,6 @@ impl RequestProcessor { plugins: Arc::new(PluginManager::with_defaults()), stats: Arc::new(ParkingLotRwLock::new(StatsAggregator::with_defaults())), tokens: Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults())), - pool_service, reload_lock: Arc::new(RwLock::new(())), hint_router: Arc::new(RwLock::new(lime_core::router::HintRouter::default())), conversation_trimmer: Arc::new(crate::conversation_manager::ConversationTrimmer::new( @@ -121,7 +115,6 @@ impl RequestProcessor { /// 使用共享的统计和 Token 追踪器创建请求处理器 pub fn with_shared_telemetry( - pool_service: Arc, stats: Arc>, tokens: Arc>, ) -> Self { @@ -135,7 +128,6 @@ impl RequestProcessor { plugins: Arc::new(PluginManager::with_defaults()), stats, tokens, - pool_service, reload_lock: Arc::new(RwLock::new(())), hint_router: Arc::new(RwLock::new(lime_core::router::HintRouter::default())), conversation_trimmer: Arc::new(crate::conversation_manager::ConversationTrimmer::new( diff --git a/src-tauri/crates/processor/src/steps/provider.rs b/src-tauri/crates/processor/src/steps/provider.rs index d47b0144a..485c05e16 100644 --- a/src-tauri/crates/processor/src/steps/provider.rs +++ b/src-tauri/crates/processor/src/steps/provider.rs @@ -12,7 +12,6 @@ use lime_infra::resilience::{FailoverManager, TimeoutError}; use lime_infra::{ Failover, FailoverConfig, Retrier, RetryConfig, TimeoutConfig, TimeoutController, }; -use lime_services::provider_pool_service::ProviderPoolService; use std::future::Future; use std::sync::Arc; @@ -72,7 +71,6 @@ pub struct ProviderStep { retrier: Arc, failover: Arc, timeout: Arc, - pool_service: Arc, } impl ProviderStep { @@ -80,22 +78,19 @@ impl ProviderStep { retrier: Arc, failover: Arc, timeout: Arc, - pool_service: Arc, ) -> Self { Self { retrier, failover, timeout, - pool_service, } } - pub fn with_defaults(pool_service: Arc) -> Self { + pub fn with_defaults() -> Self { Self { retrier: Arc::new(Retrier::with_defaults()), failover: Arc::new(Failover::new(FailoverConfig::default())), timeout: Arc::new(TimeoutController::with_defaults()), - pool_service, } } @@ -103,13 +98,11 @@ impl ProviderStep { retry_config: RetryConfig, failover_config: FailoverConfig, timeout_config: TimeoutConfig, - pool_service: Arc, ) -> Self { Self { retrier: Arc::new(Retrier::new(retry_config)), failover: Arc::new(Failover::new(failover_config)), timeout: Arc::new(TimeoutController::new(timeout_config)), - pool_service, } } @@ -122,9 +115,6 @@ impl ProviderStep { pub fn timeout(&self) -> &TimeoutController { &self.timeout } - pub fn pool_service(&self) -> &ProviderPoolService { - &self.pool_service - } /// 带重试执行 Provider 调用 pub async fn execute_with_retry( @@ -388,23 +378,20 @@ mod tests { #[tokio::test] async fn test_provider_step_new() { - let pool_service = Arc::new(ProviderPoolService::new()); - let step = ProviderStep::with_defaults(pool_service); + let step = ProviderStep::with_defaults(); assert_eq!(step.name(), "provider"); assert!(step.is_enabled()); } #[tokio::test] async fn test_provider_step_execute() { - let pool_service = Arc::new(ProviderPoolService::new()); - let step = ProviderStep::with_defaults(pool_service); + let step = ProviderStep::with_defaults(); let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string()); let mut payload = serde_json::json!({"model": "claude-sonnet-4-5"}); assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); } #[tokio::test] async fn test_provider_step_with_config() { - let pool_service = Arc::new(ProviderPoolService::new()); let retry_config = RetryConfig::new(5, 500, 10000); let failover_config = FailoverConfig::new(true, true); let timeout_config = TimeoutConfig::new(60000, 15000); @@ -413,7 +400,6 @@ mod tests { retry_config.clone(), failover_config.clone(), timeout_config.clone(), - pool_service, ); assert_eq!(step.retrier().config().max_retries, 5); @@ -462,8 +448,7 @@ mod tests { #[test] fn test_is_retryable_status() { - let pool_service = Arc::new(ProviderPoolService::new()); - let step = ProviderStep::with_defaults(pool_service); + let step = ProviderStep::with_defaults(); assert!(step.is_retryable_status(408)); assert!(step.is_retryable_status(429)); @@ -481,8 +466,7 @@ mod tests { #[tokio::test] async fn test_execute_with_retry_success() { - let pool_service = Arc::new(ProviderPoolService::new()); - let step = ProviderStep::with_defaults(pool_service); + let step = ProviderStep::with_defaults(); let mut ctx = RequestContext::new("test-model".to_string()); let result = step @@ -502,8 +486,7 @@ mod tests { #[tokio::test] async fn test_execute_with_retry_non_retryable_error() { - let pool_service = Arc::new(ProviderPoolService::new()); - let step = ProviderStep::with_defaults(pool_service); + let step = ProviderStep::with_defaults(); let mut ctx = RequestContext::new("test-model".to_string()); let result = step @@ -520,8 +503,7 @@ mod tests { #[tokio::test] async fn test_handle_failover() { - let pool_service = Arc::new(ProviderPoolService::new()); - let step = ProviderStep::with_defaults(pool_service); + let step = ProviderStep::with_defaults(); let mut ctx = RequestContext::new("test-model".to_string()); ctx.set_provider(ProviderType::Kiro); @@ -539,8 +521,7 @@ mod tests { #[tokio::test] async fn test_handle_failover_no_alternative() { - let pool_service = Arc::new(ProviderPoolService::new()); - let step = ProviderStep::with_defaults(pool_service); + let step = ProviderStep::with_defaults(); let mut ctx = RequestContext::new("test-model".to_string()); ctx.set_provider(ProviderType::Kiro); @@ -553,13 +534,11 @@ mod tests { #[tokio::test] async fn test_execute_with_timeout_success() { - let pool_service = Arc::new(ProviderPoolService::new()); let timeout_config = TimeoutConfig::new(5000, 1000); let step = ProviderStep::with_config( RetryConfig::default(), FailoverConfig::default(), timeout_config, - pool_service, ); let ctx = RequestContext::new("test-model".to_string()); @@ -579,13 +558,11 @@ mod tests { #[tokio::test] async fn test_execute_with_timeout_timeout() { - let pool_service = Arc::new(ProviderPoolService::new()); let timeout_config = TimeoutConfig::new(50, 0); let step = ProviderStep::with_config( RetryConfig::default(), FailoverConfig::default(), timeout_config, - pool_service, ); let ctx = RequestContext::new("test-model".to_string()); diff --git a/src-tauri/crates/providers/src/converter/README.md b/src-tauri/crates/providers/src/converter/README.md index 13676b14e..04be65d2d 100644 --- a/src-tauri/crates/providers/src/converter/README.md +++ b/src-tauri/crates/providers/src/converter/README.md @@ -4,15 +4,11 @@ ## 架构说明 -协议转换模块,实现不同 LLM API 格式之间的转换。 -支持 OpenAI、Claude、CodeWhisperer、Antigravity 等格式。 +协议转换模块,实现 current LLM API 格式之间的转换。旧 CodeWhisperer / Kiro 转换面已退役。 ## 文件索引 - `mod.rs` - 模块入口 -- `protocol_selector.rs` - 协议选择器 -- `openai_to_cw.rs` - OpenAI → CodeWhisperer 转换(支持 web_search 工具) -- `cw_to_openai.rs` - CodeWhisperer → OpenAI 转换 - `anthropic_to_openai.rs` - Anthropic → OpenAI 转换 - `openai_to_antigravity.rs` - OpenAI → Antigravity (Gemini CLI) 转换 - `reasoning_handler.rs` - 推理内容处理器(DeepSeek/OpenAI o1 等) diff --git a/src-tauri/crates/providers/src/converter/cw_to_openai.rs b/src-tauri/crates/providers/src/converter/cw_to_openai.rs deleted file mode 100644 index 53178da69..000000000 --- a/src-tauri/crates/providers/src/converter/cw_to_openai.rs +++ /dev/null @@ -1,141 +0,0 @@ -//! CodeWhisperer 响应转换为 OpenAI 格式 -#![allow(dead_code)] - -use lime_core::models::codewhisperer::*; -use lime_core::models::openai::*; -use std::time::{SystemTime, UNIX_EPOCH}; -use uuid::Uuid; - -/// 将 CodeWhisperer 流式事件转换为 OpenAI 格式 -pub fn convert_cw_event_to_openai_chunk( - event: &CWStreamEvent, - model: &str, - response_id: &str, -) -> Option { - let created = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - if let Some(resp_event) = &event.assistant_response_event { - // 文本内容 - if let Some(content) = &resp_event.content { - return Some(ChatCompletionChunk { - id: response_id.to_string(), - object: "chat.completion.chunk".to_string(), - created, - model: model.to_string(), - choices: vec![StreamChoice { - index: 0, - delta: StreamDelta { - role: Some("assistant".to_string()), - content: Some(content.clone()), - tool_calls: None, - reasoning_content: None, - }, - finish_reason: None, - }], - }); - } - - // Tool use - if let Some(tool_use) = &resp_event.tool_use { - return Some(ChatCompletionChunk { - id: response_id.to_string(), - object: "chat.completion.chunk".to_string(), - created, - model: model.to_string(), - choices: vec![StreamChoice { - index: 0, - delta: StreamDelta { - role: Some("assistant".to_string()), - content: None, - tool_calls: Some(vec![ToolCall { - id: tool_use.tool_use_id.clone(), - call_type: "function".to_string(), - function: FunctionCall { - name: tool_use.name.clone(), - arguments: serde_json::to_string(&tool_use.input) - .unwrap_or_default(), - }, - }]), - reasoning_content: None, - }, - finish_reason: None, - }], - }); - } - } - - None -} - -/// 创建完成的 OpenAI 响应 -pub fn create_openai_response( - content: &str, - tool_calls: Option>, - model: &str, - prompt_tokens: u32, - completion_tokens: u32, -) -> ChatCompletionResponse { - let created = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - let finish_reason = if tool_calls.is_some() { - "tool_calls" - } else { - "stop" - }; - - ChatCompletionResponse { - id: format!("chatcmpl-{}", Uuid::new_v4()), - object: "chat.completion".to_string(), - created, - model: model.to_string(), - choices: vec![Choice { - index: 0, - message: ResponseMessage { - role: "assistant".to_string(), - content: if content.is_empty() { - None - } else { - Some(content.to_string()) - }, - tool_calls, - }, - finish_reason: finish_reason.to_string(), - }], - usage: Usage { - prompt_tokens, - completion_tokens, - total_tokens: prompt_tokens + completion_tokens, - }, - } -} - -/// 创建流式结束 chunk -pub fn create_stream_end_chunk(model: &str, response_id: &str) -> ChatCompletionChunk { - let created = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - ChatCompletionChunk { - id: response_id.to_string(), - object: "chat.completion.chunk".to_string(), - created, - model: model.to_string(), - choices: vec![StreamChoice { - index: 0, - delta: StreamDelta { - role: None, - content: None, - tool_calls: None, - reasoning_content: None, - }, - finish_reason: Some("stop".to_string()), - }], - } -} diff --git a/src-tauri/crates/providers/src/converter/mod.rs b/src-tauri/crates/providers/src/converter/mod.rs index 4a73a5795..2c7d0aee4 100644 --- a/src-tauri/crates/providers/src/converter/mod.rs +++ b/src-tauri/crates/providers/src/converter/mod.rs @@ -1,19 +1,10 @@ pub mod anthropic_to_openai; -pub mod cw_to_openai; pub mod openai_to_antigravity; -pub mod openai_to_cw; -pub mod protocol_selector; pub mod reasoning_handler; #[allow(unused_imports)] pub use anthropic_to_openai::*; #[allow(unused_imports)] -pub use cw_to_openai::*; -#[allow(unused_imports)] pub use openai_to_antigravity::*; #[allow(unused_imports)] -pub use openai_to_cw::*; -#[allow(unused_imports)] -pub use protocol_selector::*; -#[allow(unused_imports)] pub use reasoning_handler::*; diff --git a/src-tauri/crates/providers/src/converter/openai_to_cw.rs b/src-tauri/crates/providers/src/converter/openai_to_cw.rs deleted file mode 100644 index 66bef74df..000000000 --- a/src-tauri/crates/providers/src/converter/openai_to_cw.rs +++ /dev/null @@ -1,507 +0,0 @@ -//! OpenAI 格式转换为 CodeWhisperer 格式 -//! -//! 支持标准工具和特殊工具类型(如 web_search)。 -//! -//! # 更新日志 -//! -//! - 2025-12-27: 添加 web_search 工具支持,修复 Issue #49 - -#![allow(dead_code)] - -use lime_core::models::codewhisperer::*; -use lime_core::models::openai::*; -use std::collections::HashMap; -use uuid::Uuid; - -/// 模型映射表 -/// -/// 参考 AIClient-2-API 的 provider-models.js 和 claude-kiro.js -/// 支持的模型列表: -/// - claude-opus-4-5, claude-opus-4-5-20251101 -/// - claude-haiku-4-5, claude-haiku-4-5-20251001 -/// - claude-sonnet-4-5, claude-sonnet-4-5-20250929 -/// - claude-sonnet-4-20250514 -/// - claude-3-7-sonnet-20250219, claude-3-5-sonnet-20241022, claude-3-5-sonnet-latest -pub fn get_model_map() -> HashMap<&'static str, &'static str> { - let mut map = HashMap::new(); - // Opus 4.5 系列 - map.insert("claude-opus-4-5", "claude-opus-4.5"); - map.insert("claude-opus-4-5-20251101", "claude-opus-4.5"); - // Haiku 4.5 系列 - map.insert("claude-haiku-4-5", "claude-haiku-4.5"); - map.insert("claude-haiku-4-5-20251001", "claude-haiku-4.5"); - // Sonnet 4.5 系列 - map.insert("claude-sonnet-4-5", "CLAUDE_SONNET_4_5_20250929_V1_0"); - map.insert( - "claude-sonnet-4-5-20250929", - "CLAUDE_SONNET_4_5_20250929_V1_0", - ); - // Sonnet 4 系列 - map.insert("claude-sonnet-4-20250514", "CLAUDE_SONNET_4_20250514_V1_0"); - // Sonnet 3.7/3.5 系列(兼容旧版本) - map.insert( - "claude-3-7-sonnet-20250219", - "CLAUDE_3_7_SONNET_20250219_V1_0", - ); - map.insert( - "claude-3-5-sonnet-20241022", - "CLAUDE_3_7_SONNET_20250219_V1_0", - ); - map.insert( - "claude-3-5-sonnet-latest", - "CLAUDE_3_7_SONNET_20250219_V1_0", - ); - map -} - -/// 获取 Kiro 支持的模型列表 -/// -/// 返回所有支持的模型 ID,用于前端展示和健康检查 -pub fn get_supported_models() -> Vec<&'static str> { - vec![ - "claude-opus-4-5", - "claude-opus-4-5-20251101", - "claude-haiku-4-5", - "claude-haiku-4-5-20251001", - "claude-sonnet-4-5", - "claude-sonnet-4-5-20250929", - "claude-sonnet-4-20250514", - "claude-3-7-sonnet-20250219", - ] -} - -pub const DEFAULT_MODEL: &str = "CLAUDE_SONNET_4_5_20250929_V1_0"; - -/// 预处理消息:合并连续的 tool 消息到前一个 assistant 消息后的 user 消息 -fn preprocess_messages(messages: &[&ChatMessage]) -> Vec { - let mut result: Vec = Vec::new(); - let mut pending_tool_results: Vec = Vec::new(); - - for msg in messages { - match msg.role.as_str() { - "tool" => { - // 收集 tool 结果 - let content = msg.get_content_text(); - let tool_id = msg.tool_call_id.clone().unwrap_or_default(); - pending_tool_results.push(CWToolResult { - content: vec![CWTextContent { text: content }], - status: "success".to_string(), - tool_use_id: tool_id, - }); - } - "user" => { - // 如果有待处理的 tool results,合并到这个 user 消息 - let content = msg.get_content_text(); - let mut tool_results = pending_tool_results.clone(); - pending_tool_results.clear(); - - // 去重 tool_results - let mut seen_ids = std::collections::HashSet::new(); - tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - result.push(ProcessedMessage { - role: "user".to_string(), - content, - tool_calls: None, - tool_results: if tool_results.is_empty() { - None - } else { - Some(tool_results) - }, - }); - } - "assistant" => { - // 如果有待处理的 tool results,先创建一个 user 消息 - if !pending_tool_results.is_empty() { - let mut tool_results = pending_tool_results.clone(); - pending_tool_results.clear(); - - // 去重 tool_results - let mut seen_ids = std::collections::HashSet::new(); - tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - result.push(ProcessedMessage { - role: "user".to_string(), - content: "Tool results provided.".to_string(), - tool_calls: None, - tool_results: Some(tool_results), - }); - } - - let content = msg.get_content_text(); - let tool_calls = msg.tool_calls.as_ref().map(|calls| { - calls - .iter() - .map(|tc| CWToolUse { - input: serde_json::from_str(&tc.function.arguments) - .unwrap_or(serde_json::json!({})), - name: tc.function.name.clone(), - tool_use_id: tc.id.clone(), - }) - .collect() - }); - - result.push(ProcessedMessage { - role: "assistant".to_string(), - content, - tool_calls, - tool_results: None, - }); - } - _ => {} - } - } - - // 处理末尾的 tool results - if !pending_tool_results.is_empty() { - let mut tool_results = pending_tool_results; - let mut seen_ids = std::collections::HashSet::new(); - tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - result.push(ProcessedMessage { - role: "user".to_string(), - content: "Tool results provided.".to_string(), - tool_calls: None, - tool_results: Some(tool_results), - }); - } - - result -} - -#[derive(Debug, Clone)] -struct ProcessedMessage { - role: String, - content: String, - tool_calls: Option>, - tool_results: Option>, -} - -/// 将 OpenAI ChatCompletionRequest 转换为 CodeWhisperer 请求 -pub fn convert_openai_to_codewhisperer( - request: &ChatCompletionRequest, - profile_arn: Option, -) -> CodeWhispererRequest { - let model_map = get_model_map(); - let cw_model = model_map - .get(request.model.as_str()) - .map(|s| s.to_string()) - .unwrap_or_else(|| DEFAULT_MODEL.to_string()); - - let conversation_id = Uuid::new_v4().to_string(); - - // 提取 system prompt 和消息 - let mut system_prompt = String::new(); - let mut raw_messages: Vec<&ChatMessage> = Vec::new(); - - for msg in &request.messages { - if msg.role == "system" { - system_prompt = msg.get_content_text(); - } else { - raw_messages.push(msg); - } - } - - // 预处理消息:合并 tool 消息 - let messages = preprocess_messages(&raw_messages); - - // 构建历史记录 - let mut history: Vec = Vec::new(); - let mut start_idx = 0; - - // 处理 system prompt - 合并到第一条用户消息 - if !system_prompt.is_empty() && !messages.is_empty() && messages[0].role == "user" { - let first_content = &messages[0].content; - let combined = format!("{system_prompt}\n\n{first_content}"); - - let mut user_msg = UserInputMessage { - content: combined, - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }; - - // 如果第一条消息有 tool_results,也要包含 - if let Some(ref tool_results) = messages[0].tool_results { - user_msg.user_input_message_context = Some(UserInputMessageContext { - tools: None, - tool_results: Some(tool_results.clone()), - }); - } - - history.push(HistoryItem::User(UserHistoryItem { - user_input_message: user_msg, - })); - start_idx = 1; - } - - // 处理历史消息(除最后一条) - for msg in messages - .iter() - .take(messages.len().saturating_sub(1)) - .skip(start_idx) - { - match msg.role.as_str() { - "user" => { - let content = if msg.content.is_empty() { - if msg.tool_results.is_some() { - "Tool results provided.".to_string() - } else { - "Continue".to_string() - } - } else { - msg.content.clone() - }; - - let mut user_msg = UserInputMessage { - content, - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }; - - if let Some(ref tool_results) = msg.tool_results { - user_msg.user_input_message_context = Some(UserInputMessageContext { - tools: None, - tool_results: Some(tool_results.clone()), - }); - } - - history.push(HistoryItem::User(UserHistoryItem { - user_input_message: user_msg, - })); - } - "assistant" => { - let content = if msg.content.is_empty() { - "I understand.".to_string() - } else { - msg.content.clone() - }; - - history.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content, - tool_uses: msg.tool_calls.clone(), - }, - })); - } - _ => {} - } - } - - // 修复历史记录交替顺序 - let history = fix_history_alternation(history, &cw_model); - - // 构建当前消息 - let (current_content, current_tool_results) = if let Some(last_msg) = messages.last() { - if last_msg.role == "assistant" { - ("Continue".to_string(), None) - } else { - let content = if last_msg.content.is_empty() { - if last_msg.tool_results.is_some() { - "Tool results provided.".to_string() - } else { - "Continue".to_string() - } - } else { - last_msg.content.clone() - }; - (content, last_msg.tool_results.clone()) - } - } else { - ("Continue".to_string(), None) - }; - - // 构建 tools - let tools = request.tools.as_ref().map(|tools| { - let mut cw_tools: Vec = Vec::new(); - let mut function_count = 0; - - for t in tools.iter() { - match t { - // 标准函数工具 - Tool::Function { function } => { - // 限制最多 50 个函数工具 - if function_count >= 50 { - continue; - } - function_count += 1; - - let params = function - .parameters - .clone() - .unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})); - - let desc = function - .description - .clone() - .unwrap_or_else(|| format!("Tool: {}", function.name)); - - cw_tools.push(CWToolItem::Standard(CWTool { - tool_specification: ToolSpecification { - name: function.name.clone(), - // P1 安全修复:使用字符边界安全的截断,防止 UTF-8 panic - description: if desc.len() > 500 { - let truncated: String = desc.chars().take(497).collect(); - format!("{truncated}...") - } else { - desc - }, - input_schema: InputSchema { json: params }, - }, - })); - } - // 联网搜索工具(Codex 格式) - Tool::WebSearch => { - tracing::info!("[CW_TOOLS] 添加 web_search 工具"); - cw_tools.push(CWToolItem::WebSearch(CWWebSearchTool { - tool_type: "web_search".to_string(), - })); - } - // 联网搜索工具(Claude Code 格式) - Tool::WebSearch20250305 => { - tracing::info!("[CW_TOOLS] 添加 web_search 工具 (from web_search_20250305)"); - cw_tools.push(CWToolItem::WebSearch(CWWebSearchTool { - tool_type: "web_search".to_string(), - })); - } - } - } - - cw_tools - }); - - let user_input_message_context = if tools.is_some() || current_tool_results.is_some() { - Some(UserInputMessageContext { - tools, - tool_results: current_tool_results, - }) - } else { - None - }; - - CodeWhispererRequest { - conversation_state: ConversationState { - chat_trigger_type: "MANUAL".to_string(), - conversation_id, - current_message: CurrentMessage { - user_input_message: UserInputMessage { - content: current_content, - model_id: cw_model, - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context, - }, - }, - history: if history.is_empty() { - None - } else { - Some(history) - }, - }, - profile_arn, - } -} - -/// 修复历史记录,确保 user/assistant 严格交替 -fn fix_history_alternation(history: Vec, model_id: &str) -> Vec { - if history.is_empty() { - return history; - } - - let mut fixed: Vec = Vec::new(); - - for item in history { - match &item { - HistoryItem::User(user_item) => { - // 如果上一条也是 user,合并 tool_results 或插入占位 assistant - if let Some(HistoryItem::User(last_user)) = fixed.last_mut() { - // 尝试合并 tool_results - let has_tool_results = user_item - .user_input_message - .user_input_message_context - .as_ref() - .map(|ctx| ctx.tool_results.is_some()) - .unwrap_or(false); - - if has_tool_results { - // 合并 tool_results 到上一个 user 消息 - let new_results = user_item - .user_input_message - .user_input_message_context - .as_ref() - .and_then(|ctx| ctx.tool_results.clone()) - .unwrap_or_default(); - - if let Some(ref mut ctx) = - last_user.user_input_message.user_input_message_context - { - if let Some(ref mut existing) = ctx.tool_results { - existing.extend(new_results); - } else { - ctx.tool_results = Some(new_results); - } - } else { - last_user.user_input_message.user_input_message_context = - Some(UserInputMessageContext { - tools: None, - tool_results: Some(new_results), - }); - } - continue; - } else { - // 插入占位 assistant - fixed.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content: "I understand.".to_string(), - tool_uses: None, - }, - })); - } - } - fixed.push(item); - } - HistoryItem::Assistant(_) => { - // 如果上一条也是 assistant,插入占位 user - if let Some(HistoryItem::Assistant(_)) = fixed.last() { - fixed.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: "Continue".to_string(), - model_id: model_id.to_string(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }, - })); - } - // 如果历史为空,先插入 user - if fixed.is_empty() { - fixed.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: "Continue".to_string(), - model_id: model_id.to_string(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }, - })); - } - fixed.push(item); - } - } - } - - // 确保以 assistant 结尾 - if let Some(HistoryItem::User(_)) = fixed.last() { - fixed.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content: "I understand.".to_string(), - tool_uses: None, - }, - })); - } - - fixed -} diff --git a/src-tauri/crates/providers/src/converter/protocol_selector.rs b/src-tauri/crates/providers/src/converter/protocol_selector.rs deleted file mode 100644 index dcb448ede..000000000 --- a/src-tauri/crates/providers/src/converter/protocol_selector.rs +++ /dev/null @@ -1,233 +0,0 @@ -//! 协议选择器 - 智能选择最优协议转换路径 -//! -//! 根据源协议、目标 Provider 和请求特征,选择最优的协议转换路径。 - -#![allow(dead_code)] - -use lime_core::models::provider_pool_model::PoolProviderType; - -/// 协议类型 -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum Protocol { - /// OpenAI Chat Completions API - OpenAI, - /// Anthropic Messages API (Claude) - Anthropic, - /// CodeWhisperer API (Kiro) - CodeWhisperer, - /// Gemini API (Google) - Gemini, - /// Antigravity API (Google Internal) - Antigravity, -} - -impl Protocol { - pub fn as_str(&self) -> &'static str { - match self { - Protocol::OpenAI => "openai", - Protocol::Anthropic => "anthropic", - Protocol::CodeWhisperer => "codewhisperer", - Protocol::Gemini => "gemini", - Protocol::Antigravity => "antigravity", - } - } -} - -/// 转换路径 -#[derive(Debug, Clone)] -pub struct ConversionPath { - /// 源协议 - pub source: Protocol, - /// 目标协议 - pub target: Protocol, - /// 是否需要转换 - pub needs_conversion: bool, - /// 转换复杂度 (0-10, 越低越好) - pub complexity: u8, -} - -/// 协议选择器 -pub struct ProtocolSelector; - -impl ProtocolSelector { - /// 获取 Provider 的原生协议 - pub fn native_protocol(provider: PoolProviderType) -> Protocol { - match provider { - PoolProviderType::Kiro => Protocol::CodeWhisperer, - PoolProviderType::Gemini => Protocol::Gemini, - PoolProviderType::OpenAI => Protocol::OpenAI, - PoolProviderType::Claude => Protocol::Anthropic, - PoolProviderType::ClaudeOAuth => Protocol::Anthropic, // Claude OAuth uses Anthropic protocol - // Anthropic 兼容格式使用 Anthropic 协议 - PoolProviderType::AnthropicCompatible => Protocol::Anthropic, - PoolProviderType::Antigravity => Protocol::Antigravity, - PoolProviderType::Vertex => Protocol::Gemini, // Vertex AI uses Gemini protocol - PoolProviderType::GeminiApiKey => Protocol::Gemini, // Gemini API Key uses Gemini protocol - PoolProviderType::Codex => Protocol::OpenAI, // Codex uses OpenAI protocol - // API Key Provider 类型 - PoolProviderType::Anthropic => Protocol::Anthropic, - PoolProviderType::AzureOpenai => Protocol::OpenAI, - PoolProviderType::AwsBedrock => Protocol::Anthropic, - PoolProviderType::Ollama => Protocol::OpenAI, - } - } - - /// 选择最优转换路径 - pub fn select_path( - source_protocol: Protocol, - target_provider: PoolProviderType, - ) -> ConversionPath { - let target_protocol = Self::native_protocol(target_provider); - - // 如果源和目标协议相同,无需转换 - if source_protocol == target_protocol { - return ConversionPath { - source: source_protocol, - target: target_protocol, - needs_conversion: false, - complexity: 0, - }; - } - - // 计算转换复杂度 - let complexity = Self::calculate_complexity(source_protocol, target_protocol); - - ConversionPath { - source: source_protocol, - target: target_protocol, - needs_conversion: true, - complexity, - } - } - - /// 计算转换复杂度 - fn calculate_complexity(source: Protocol, target: Protocol) -> u8 { - match (source, target) { - // OpenAI <-> Anthropic: 中等复杂度 - (Protocol::OpenAI, Protocol::Anthropic) => 3, - (Protocol::Anthropic, Protocol::OpenAI) => 3, - - // OpenAI <-> CodeWhisperer: 较高复杂度(需要处理历史格式) - (Protocol::OpenAI, Protocol::CodeWhisperer) => 5, - (Protocol::CodeWhisperer, Protocol::OpenAI) => 5, - - // OpenAI <-> Gemini/Antigravity: 中等复杂度 - (Protocol::OpenAI, Protocol::Gemini) => 4, - (Protocol::OpenAI, Protocol::Antigravity) => 4, - (Protocol::Gemini, Protocol::OpenAI) => 4, - (Protocol::Antigravity, Protocol::OpenAI) => 4, - - // Anthropic <-> CodeWhisperer: 较高复杂度 - (Protocol::Anthropic, Protocol::CodeWhisperer) => 6, - (Protocol::CodeWhisperer, Protocol::Anthropic) => 6, - - // Anthropic <-> Gemini/Antigravity: 中等复杂度 - (Protocol::Anthropic, Protocol::Gemini) => 5, - (Protocol::Anthropic, Protocol::Antigravity) => 5, - (Protocol::Gemini, Protocol::Anthropic) => 5, - (Protocol::Antigravity, Protocol::Anthropic) => 5, - - // Gemini <-> Antigravity: 低复杂度(格式相似) - (Protocol::Gemini, Protocol::Antigravity) => 1, - (Protocol::Antigravity, Protocol::Gemini) => 1, - - // 其他情况 - _ => 7, - } - } - - /// 检查是否支持直接转换 - pub fn supports_direct_conversion(source: Protocol, target: Protocol) -> bool { - matches!( - (source, target), - (Protocol::OpenAI, Protocol::Anthropic) - | (Protocol::Anthropic, Protocol::OpenAI) - | (Protocol::OpenAI, Protocol::CodeWhisperer) - | (Protocol::CodeWhisperer, Protocol::OpenAI) - | (Protocol::OpenAI, Protocol::Gemini) - | (Protocol::OpenAI, Protocol::Antigravity) - | (Protocol::Gemini, Protocol::OpenAI) - | (Protocol::Antigravity, Protocol::OpenAI) - | (Protocol::Gemini, Protocol::Antigravity) - | (Protocol::Antigravity, Protocol::Gemini) - ) - } - - /// 获取推荐的中间协议(用于不支持直接转换的情况) - pub fn intermediate_protocol(source: Protocol, target: Protocol) -> Option { - // 大多数情况下,OpenAI 是最好的中间协议 - if !Self::supports_direct_conversion(source, target) - && source != Protocol::OpenAI - && target != Protocol::OpenAI - { - return Some(Protocol::OpenAI); - } - None - } - - /// 获取 Provider 支持的输入协议列表 - pub fn supported_input_protocols(_provider: PoolProviderType) -> Vec { - // 所有 Provider 都支持 OpenAI 和 Anthropic 协议输入 - vec![Protocol::OpenAI, Protocol::Anthropic] - } - - /// 检查请求是否需要特殊处理 - pub fn needs_special_handling( - source: Protocol, - target_provider: PoolProviderType, - has_tools: bool, - has_images: bool, - ) -> bool { - // 工具调用在某些转换中需要特殊处理 - if has_tools { - matches!( - (source, target_provider), - (Protocol::OpenAI, PoolProviderType::Kiro) - | (Protocol::Anthropic, PoolProviderType::Kiro) - ) - } else if has_images { - // 图片在某些 Provider 中需要特殊处理 - match target_provider { - PoolProviderType::Kiro => true, // Kiro 不支持图片 - _ => false, - } - } else { - false - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_native_protocol() { - assert_eq!( - ProtocolSelector::native_protocol(PoolProviderType::Kiro), - Protocol::CodeWhisperer - ); - assert_eq!( - ProtocolSelector::native_protocol(PoolProviderType::OpenAI), - Protocol::OpenAI - ); - assert_eq!( - ProtocolSelector::native_protocol(PoolProviderType::Claude), - Protocol::Anthropic - ); - } - - #[test] - fn test_select_path_no_conversion() { - let path = ProtocolSelector::select_path(Protocol::OpenAI, PoolProviderType::OpenAI); - assert!(!path.needs_conversion); - assert_eq!(path.complexity, 0); - } - - #[test] - fn test_select_path_with_conversion() { - let path = ProtocolSelector::select_path(Protocol::OpenAI, PoolProviderType::Kiro); - assert!(path.needs_conversion); - assert_eq!(path.target, Protocol::CodeWhisperer); - } -} diff --git a/src-tauri/crates/providers/src/lib.rs b/src-tauri/crates/providers/src/lib.rs index b4313dfdb..54a862373 100644 --- a/src-tauri/crates/providers/src/lib.rs +++ b/src-tauri/crates/providers/src/lib.rs @@ -3,8 +3,8 @@ //! 包含所有 Provider 实现、协议转换、流式传输等核心业务模块。 //! //! ## 模块结构 -//! - `providers`: Provider 实现(Kiro、Gemini、Claude、OpenAI、Vertex 等) -//! - `converter`: 协议转换(OpenAI ↔ CW、OpenAI ↔ Antigravity 等) +//! - `providers`: Provider 实现(API Key Provider 所需的 Gemini、Claude、OpenAI、Vertex 等) +//! - `converter`: 协议转换(OpenAI ↔ Antigravity 等) //! - `streaming`: 流式传输管理 //! - `translator`: 请求/响应翻译层 //! - `stream`: 流事件解析和生成 diff --git a/src-tauri/crates/providers/src/providers/README.md b/src-tauri/crates/providers/src/providers/README.md index c0a78a2a2..9cfaaf8c2 100644 --- a/src-tauri/crates/providers/src/providers/README.md +++ b/src-tauri/crates/providers/src/providers/README.md @@ -4,26 +4,22 @@ ## 架构说明 -各 LLM Provider 的认证和 API 实现。 -支持 OAuth 和 API Key 两种认证方式。 +各 LLM Provider 的 API 实现。认证事实源已经收敛到 API Key Provider;旧 OAuth / 本地 CLI 凭证 provider runtime 已退役。 ## 文件索引 - `mod.rs` - 模块入口和 Provider 枚举 - `traits.rs` - Provider trait 定义 - `error.rs` - 错误类型定义 -- `kiro.rs` - Kiro/CodeWhisperer OAuth 认证 -- `gemini.rs` - Gemini OAuth 认证 -- `qwen.rs` - Qwen OAuth 认证 -- `antigravity.rs` - Antigravity OAuth 认证 -- `claude_oauth.rs` - Claude OAuth 认证 +- `gemini.rs` - Gemini API Key 请求支持 - `claude_custom.rs` - Claude API Key 认证 - `openai_custom.rs` - OpenAI API Key 认证 -- `codex.rs` - Codex Provider -- `iflow.rs` - iFlow Provider +- `codex.rs` - OpenAI Responses / Codex 兼容请求支持 - `vertex.rs` - Vertex AI Provider - `tests.rs` - 单元测试 +Kiro、Antigravity OAuth、Claude OAuth、Gemini OAuth、Codex OAuth 属于旧凭证池功能,不应重新加入本目录编译图。Antigravity 仅保留 `converter/openai_to_antigravity.rs` 协议转换能力。 + ## 更新提醒 任何文件变更后,请更新此文档和相关的上级文档。 diff --git a/src-tauri/crates/providers/src/providers/antigravity.rs b/src-tauri/crates/providers/src/providers/antigravity.rs deleted file mode 100644 index 6903c2d0a..000000000 --- a/src-tauri/crates/providers/src/providers/antigravity.rs +++ /dev/null @@ -1,2689 +0,0 @@ -//! Antigravity Provider - Google 内部 Gemini 3 Pro 接口 -//! -//! 支持 Gemini 3 Pro 等高级模型,通过 Google 内部 API 访问。 - -#![allow(dead_code)] - -use super::traits::{CredentialProvider, ProviderResult}; -use async_trait::async_trait; -use reqwest::Client; -use serde::{Deserialize, Serialize}; -use std::error::Error; -use std::path::PathBuf; -use std::sync::Arc; -use tokio::sync::oneshot; -use uuid::Uuid; - -// ============================================================================ -// Antigravity API 错误类型 -// ============================================================================ - -/// Antigravity API 错误 -/// -/// 携带 HTTP 状态码,便于调用方透传给客户端 -#[derive(Debug, Clone)] -pub struct AntigravityApiError { - /// HTTP 状态码 - pub status_code: u16, - /// 错误消息 - pub message: String, - /// 原始响应体(如果有) - pub body: Option, -} - -impl AntigravityApiError { - /// 创建新的 API 错误 - pub fn new(status_code: u16, message: impl Into) -> Self { - Self { - status_code, - message: message.into(), - body: None, - } - } - - /// 创建带响应体的 API 错误 - pub fn with_body( - status_code: u16, - message: impl Into, - body: impl Into, - ) -> Self { - Self { - status_code, - message: message.into(), - body: Some(body.into()), - } - } - - /// 是否是可重试的错误(429 或 5xx) - pub fn is_retryable(&self) -> bool { - self.status_code == 429 || (self.status_code >= 500 && self.status_code < 600) - } - - /// 是否是权限错误(401 或 403) - pub fn is_auth_error(&self) -> bool { - self.status_code == 401 || self.status_code == 403 - } - - /// 是否是配额耗尽错误(429) - pub fn is_rate_limit(&self) -> bool { - self.status_code == 429 - } - - /// 获取用户友好的错误消息 - pub fn user_message(&self) -> String { - match self.status_code { - 401 => format!("认证失败,请重新登录: {}", self.message), - 403 => format!("权限不足: {}", self.message), - 429 => format!("请求过于频繁,请稍后重试: {}", self.message), - 500..=599 => format!("服务器错误 ({}): {}", self.status_code, self.message), - _ => format!("API 错误 ({}): {}", self.status_code, self.message), - } - } -} - -impl std::fmt::Display for AntigravityApiError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "HTTP {} - {}", self.status_code, self.message) - } -} - -impl std::error::Error for AntigravityApiError {} - -// Constants -// 正确的 Cloud Code API 端点(参考 Antigravity-Manager) -const ANTIGRAVITY_BASE_URL_PROD: &str = "https://cloudcode-pa.googleapis.com"; -const ANTIGRAVITY_BASE_URL_DAILY: &str = "https://daily-cloudcode-pa.sandbox.googleapis.com"; -const ANTIGRAVITY_BASE_URL_AUTOPUSH: &str = "https://autopush-cloudcode-pa.sandbox.googleapis.com"; -const ANTIGRAVITY_API_VERSION: &str = "v1internal"; -const CREDENTIALS_DIR: &str = ".antigravity"; -const CREDENTIALS_FILE: &str = "oauth_creds.json"; - -// OAuth credentials - 从环境变量读取 -// 必须通过环境变量配置: -// - ANTIGRAVITY_OAUTH_CLIENT_ID -// - ANTIGRAVITY_OAUTH_CLIENT_SECRET -fn oauth_client_id() -> String { - std::env::var("ANTIGRAVITY_OAUTH_CLIENT_ID") - .expect("ANTIGRAVITY_OAUTH_CLIENT_ID environment variable must be set") -} - -fn oauth_client_secret() -> String { - std::env::var("ANTIGRAVITY_OAUTH_CLIENT_SECRET") - .expect("ANTIGRAVITY_OAUTH_CLIENT_SECRET environment variable must be set") -} - -// OAuth scopes -const OAUTH_SCOPES: &[&str] = &[ - "https://www.googleapis.com/auth/cloud-platform", - "https://www.googleapis.com/auth/userinfo.email", - "https://www.googleapis.com/auth/userinfo.profile", - "https://www.googleapis.com/auth/cclog", - "https://www.googleapis.com/auth/experimentsandconfigs", -]; - -// Token 刷新提前量(秒) -const REFRESH_SKEW: i64 = 3000; - -// Token 即将过期的阈值(秒)- 10 分钟 -const TOKEN_EXPIRING_SOON_THRESHOLD: i64 = 600; - -/// Token 验证结果 -/// Requirements: 1.1, 1.2, 1.3, 1.4 -#[derive(Debug, Clone, PartialEq)] -pub enum TokenValidationResult { - /// Token 有效,包含剩余有效时间(秒) - Valid { expires_in_secs: i64 }, - /// Token 即将过期(少于 10 分钟),需要主动刷新 - ExpiringSoon { expires_in_secs: i64 }, - /// Token 已过期 - Expired, - /// Token 无效(缺失、为空或格式错误) - Invalid { reason: String }, -} - -impl TokenValidationResult { - /// 是否需要刷新 Token - pub fn needs_refresh(&self) -> bool { - matches!( - self, - TokenValidationResult::ExpiringSoon { .. } - | TokenValidationResult::Expired - | TokenValidationResult::Invalid { .. } - ) - } - - /// 是否可以使用(有效或即将过期但仍可用) - pub fn is_usable(&self) -> bool { - matches!( - self, - TokenValidationResult::Valid { .. } | TokenValidationResult::ExpiringSoon { .. } - ) - } -} - -/// Token 刷新错误类型 -/// Requirements: 2.1 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub enum TokenRefreshError { - /// OAuth invalid_grant 错误 - 需要用户重新授权 - InvalidGrant { message: String }, - /// 网络错误 - 可以重试 - NetworkError { message: String }, - /// 服务器错误 (5xx) - 可以重试 - ServerError { message: String }, - /// 未知错误 - Unknown { message: String }, -} - -impl TokenRefreshError { - /// 是否需要用户重新授权 - pub fn requires_reauth(&self) -> bool { - matches!(self, TokenRefreshError::InvalidGrant { .. }) - } - - /// 是否可以重试 - pub fn is_retryable(&self) -> bool { - matches!( - self, - TokenRefreshError::NetworkError { .. } | TokenRefreshError::ServerError { .. } - ) - } - - /// 获取用户友好的错误消息 - pub fn user_message(&self) -> String { - match self { - TokenRefreshError::InvalidGrant { .. } => { - "Antigravity 授权已过期,请重新登录授权".to_string() - } - TokenRefreshError::NetworkError { message } => { - format!("网络连接失败: {message}") - } - TokenRefreshError::ServerError { message } => { - format!("Google 服务暂时不可用: {message}") - } - TokenRefreshError::Unknown { message } => { - format!("Token 刷新失败: {message}") - } - } - } -} - -impl std::fmt::Display for TokenRefreshError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.user_message()) - } -} - -impl std::error::Error for TokenRefreshError {} - -/// Antigravity 支持的模型列表(fallback,当无法从 models 仓库获取时使用) -pub const ANTIGRAVITY_MODELS_FALLBACK: &[&str] = &[ - "gemini-2.5-computer-use-preview-10-2025", - "gemini-3-pro-image-preview", - "gemini-3-pro-preview", - "gemini-3-flash-preview", - "gemini-2.5-flash-preview", - "gemini-2.5-flash", - "gemini-2.5-pro", - "gemini-3-flash", - "gemini-3-pro-high", - "gemini-3-pro-low", - "gemini-claude-sonnet-4-5", - "gemini-claude-sonnet-4-5-thinking", - "gemini-claude-opus-4-5-thinking", - "claude-sonnet-4-5", - "claude-sonnet-4-5-thinking", - "claude-opus-4-5-thinking", -]; - -/// 模型别名映射(fallback,当无法从 models 仓库获取时使用) -/// 格式:用户友好名称 -> 内部 API 名称 -pub const ANTIGRAVITY_ALIAS_FALLBACK: &[(&str, &str)] = &[ - // 需要映射的模型 - ("gemini-2.5-computer-use-preview-10-2025", "rev19-uic3-1p"), - ("gemini-3-pro-image-preview", "gemini-3-pro-image"), - // Gemini 3 preview 模型映射到正式名称 - ("gemini-3-flash-preview", "gemini-3-flash"), - ("gemini-3-pro-preview", "gemini-3-pro-high"), - // Gemini 2.5 preview 模型映射 - ("gemini-2.5-flash-preview", "gemini-2.5-flash"), - // Claude via Antigravity - ("gemini-claude-sonnet-4-5", "claude-sonnet-4-5"), - ( - "gemini-claude-sonnet-4-5-thinking", - "claude-sonnet-4-5-thinking", - ), - ( - "gemini-claude-opus-4-5-thinking", - "claude-opus-4-5-thinking", - ), -]; - -/// 模型别名映射(用户友好名称 -> 内部名称) -/// 使用 fallback 映射,当无法从 ModelRegistryService 获取时使用 -fn alias_to_model_name(model: &str) -> String { - for (alias, internal) in ANTIGRAVITY_ALIAS_FALLBACK { - if *alias == model { - return internal.to_string(); - } - } - model.to_string() -} - -/// 内部模型名称 -> 用户友好名称 -#[allow(dead_code)] -fn model_name_to_alias(model: &str) -> String { - for (alias, internal) in ANTIGRAVITY_ALIAS_FALLBACK { - if *internal == model { - return alias.to_string(); - } - } - model.to_string() -} - -/// 生成随机请求 ID -fn generate_request_id() -> String { - format!("agent-{}", Uuid::new_v4()) -} - -/// 生成随机会话 ID -fn generate_session_id() -> String { - // 使用 UUID 的一部分作为随机数 - let uuid = Uuid::new_v4(); - let bytes = uuid.as_bytes(); - let n: u64 = u64::from_le_bytes([ - bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7], - ]) % 9_000_000_000_000_000_000; - format!("-{n}") -} - -/// 生成随机项目 ID -fn generate_project_id() -> String { - let adjectives = ["useful", "bright", "swift", "calm", "bold"]; - let nouns = ["fuze", "wave", "spark", "flow", "core"]; - let uuid = Uuid::new_v4(); - let bytes = uuid.as_bytes(); - let adj = adjectives[(bytes[0] as usize) % adjectives.len()]; - let noun = nouns[(bytes[1] as usize) % nouns.len()]; - let random_part: String = uuid.to_string()[..5].to_lowercase(); - format!("{adj}-{noun}-{random_part}") -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AntigravityCredentials { - pub access_token: Option, - pub refresh_token: Option, - pub token_type: Option, - /// 过期时间戳(毫秒)- 兼容旧格式 - #[serde(skip_serializing_if = "Option::is_none")] - pub expiry_date: Option, - /// 过期时间(RFC3339 格式)- - pub expire: Option, - pub scope: Option, - /// 最后刷新时间(RFC3339 格式) - #[serde(skip_serializing_if = "Option::is_none")] - pub last_refresh: Option, - /// 凭证类型标识 - #[serde(default = "default_antigravity_type", rename = "type")] - pub cred_type: String, - /// Token 有效期(秒) - #[serde(skip_serializing_if = "Option::is_none")] - pub expires_in: Option, - /// Token 获取时间戳(毫秒)- - #[serde(skip_serializing_if = "Option::is_none")] - pub timestamp: Option, - /// 是否启用 - - #[serde(skip_serializing_if = "Option::is_none")] - pub enable: Option, - /// 项目 ID - #[serde(skip_serializing_if = "Option::is_none")] - pub project_id: Option, - /// 用户邮箱 - #[serde(skip_serializing_if = "Option::is_none")] - pub email: Option, -} - -fn default_antigravity_type() -> String { - "antigravity".to_string() -} - -impl Default for AntigravityCredentials { - fn default() -> Self { - Self { - access_token: None, - refresh_token: None, - token_type: Some("Bearer".to_string()), - expiry_date: None, - expire: None, - scope: None, - last_refresh: None, - cred_type: default_antigravity_type(), - expires_in: None, - timestamp: None, - enable: None, - project_id: None, - email: None, - } - } -} - -/// Antigravity Provider -pub struct AntigravityProvider { - pub credentials: AntigravityCredentials, - pub project_id: Option, - pub client: Client, - pub base_urls: Vec, - pub available_models: Vec, -} - -impl Default for AntigravityProvider { - fn default() -> Self { - Self { - credentials: AntigravityCredentials::default(), - project_id: None, - client: Client::builder() - .timeout(std::time::Duration::from_secs(120)) - .build() - .unwrap_or_else(|_| Client::new()), - // 只使用生产环境和 daily 环境(参考 Antigravity-Manager) - // 沙盒环境(autopush)需要特殊许可证,不适合普通用户 - base_urls: vec![ - ANTIGRAVITY_BASE_URL_PROD.to_string(), - ANTIGRAVITY_BASE_URL_DAILY.to_string(), - ], - available_models: ANTIGRAVITY_MODELS_FALLBACK - .iter() - .map(|s| s.to_string()) - .collect(), - } - } -} - -impl AntigravityProvider { - pub fn new() -> Self { - Self::default() - } - - pub fn default_creds_path() -> PathBuf { - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(CREDENTIALS_DIR) - .join(CREDENTIALS_FILE) - } - - pub async fn load_credentials(&mut self) -> Result<(), Box> { - let path = Self::default_creds_path(); - - if tokio::fs::try_exists(&path).await.unwrap_or(false) { - let content = tokio::fs::read_to_string(&path).await?; - let creds: AntigravityCredentials = serde_json::from_str(&content)?; - self.credentials = creds; - } - - Ok(()) - } - - async fn resolve_credentials_path(path: &str) -> PathBuf { - let explicit_path = PathBuf::from(path); - if tokio::fs::try_exists(&explicit_path).await.unwrap_or(false) { - return explicit_path; - } - - let default_path = Self::default_creds_path(); - if explicit_path != default_path - && tokio::fs::try_exists(&default_path).await.unwrap_or(false) - { - tracing::warn!( - "[ANTIGRAVITY] 显式凭证路径不存在,回退到默认路径: missing={}, fallback={}", - explicit_path.display(), - default_path.display() - ); - return default_path; - } - - explicit_path - } - - pub async fn load_credentials_from_path( - &mut self, - path: &str, - ) -> Result<(), Box> { - let resolved_path = Self::resolve_credentials_path(path).await; - let content = tokio::fs::read_to_string(&resolved_path).await?; - - // 尝试解析为单个凭证对象 - if let Ok(creds) = serde_json::from_str::(&content) { - self.credentials = creds; - // 如果凭证中有 project_id,设置到 provider - if let Some(ref pid) = self.credentials.project_id { - self.project_id = Some(pid.clone()); - } - return Ok(()); - } - - // 尝试解析为数组格式(兼容 antigravity2api-nodejs 的 accounts.json) - if let Ok(creds_array) = serde_json::from_str::>(&content) { - // 找到第一个启用的凭证 - if let Some(creds) = creds_array.into_iter().find(|c| c.enable != Some(false)) { - self.credentials = creds; - // 如果凭证中有 project_id,设置到 provider - if let Some(ref pid) = self.credentials.project_id { - self.project_id = Some(pid.clone()); - } - return Ok(()); - } - return Err("凭证文件中没有可用的账号(所有账号都被禁用)".into()); - } - - Err("无法解析凭证文件,请确保是有效的 JSON 格式".into()) - } - - pub async fn save_credentials(&self) -> Result<(), Box> { - let path = Self::default_creds_path(); - if let Some(parent) = path.parent() { - tokio::fs::create_dir_all(parent).await?; - } - let content = serde_json::to_string_pretty(&self.credentials)?; - tokio::fs::write(&path, content).await?; - Ok(()) - } - - pub fn is_token_valid(&self) -> bool { - if self.credentials.access_token.is_none() { - return false; - } - - // 检查是否被禁用 - if self.credentials.enable == Some(false) { - return false; - } - - // 优先检查 RFC3339 格式的过期时间 - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) { - let now = chrono::Utc::now(); - // Token valid if more than 5 minutes until expiry - return expires > now + chrono::Duration::minutes(5); - } - } - - // 兼容旧的毫秒时间戳格式 - if let Some(expiry) = self.credentials.expiry_date { - let now = chrono::Utc::now().timestamp_millis(); - // Token valid if more than 5 minutes until expiry - return expiry > now + 300_000; - } - - // 兼容 antigravity2api-nodejs 格式:timestamp + expires_in - if let (Some(timestamp), Some(expires_in)) = - (self.credentials.timestamp, self.credentials.expires_in) - { - let expiry = timestamp + (expires_in * 1000); - let now = chrono::Utc::now().timestamp_millis(); - // Token valid if more than 5 minutes until expiry - return expiry > now + 300_000; - } - - true - } - - pub fn is_token_expiring_soon(&self) -> bool { - // 优先检查 RFC3339 格式的过期时间 - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) { - let now = chrono::Utc::now(); - let refresh_skew = chrono::Duration::seconds(REFRESH_SKEW); - return expires <= now + refresh_skew; - } - } - - // 兼容旧的毫秒时间戳格式 - if let Some(expiry) = self.credentials.expiry_date { - let now = chrono::Utc::now().timestamp_millis(); - let refresh_skew_ms = REFRESH_SKEW * 1000; - return expiry <= now + refresh_skew_ms; - } - - // 兼容 antigravity2api-nodejs 格式:timestamp + expires_in - if let (Some(timestamp), Some(expires_in)) = - (self.credentials.timestamp, self.credentials.expires_in) - { - let expiry = timestamp + (expires_in * 1000); - let now = chrono::Utc::now().timestamp_millis(); - let refresh_skew_ms = REFRESH_SKEW * 1000; - return expiry <= now + refresh_skew_ms; - } - - true - } - - /// 验证 Token 状态(支持多种时间格式) - /// Requirements: 1.1, 1.2, 1.3, 1.4 - pub fn validate_token(&self) -> TokenValidationResult { - // 检查 access_token 是否存在且非空 - match &self.credentials.access_token { - None => { - return TokenValidationResult::Invalid { - reason: "access_token 缺失".to_string(), - }; - } - Some(token) if token.trim().is_empty() => { - return TokenValidationResult::Invalid { - reason: "access_token 为空".to_string(), - }; - } - _ => {} - } - - // 检查是否被禁用 - if self.credentials.enable == Some(false) { - return TokenValidationResult::Invalid { - reason: "凭证已被禁用".to_string(), - }; - } - - // 检查 refresh_token 是否存在(用于后续刷新) - if self.credentials.refresh_token.is_none() { - return TokenValidationResult::Invalid { - reason: "refresh_token 缺失,无法刷新".to_string(), - }; - } - - let now = chrono::Utc::now(); - let now_millis = now.timestamp_millis(); - - // 尝试解析过期时间(支持多种格式) - let expires_in_secs: Option = { - // 优先检查 RFC3339 格式 - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) { - Some((expires.timestamp_millis() - now_millis) / 1000) - } else { - // RFC3339 解析失败,尝试其他格式 - None - } - } else { - None - } - } - .or_else(|| { - // 兼容毫秒时间戳格式 - self.credentials - .expiry_date - .map(|expiry| (expiry - now_millis) / 1000) - }) - .or_else(|| { - // 兼容 timestamp + expires_in 格式 - match (self.credentials.timestamp, self.credentials.expires_in) { - (Some(timestamp), Some(expires_in)) => { - let expiry = timestamp + (expires_in * 1000); - Some((expiry - now_millis) / 1000) - } - _ => None, - } - }); - - match expires_in_secs { - Some(secs) if secs <= 0 => TokenValidationResult::Expired, - Some(secs) if secs <= TOKEN_EXPIRING_SOON_THRESHOLD => { - TokenValidationResult::ExpiringSoon { - expires_in_secs: secs, - } - } - Some(secs) => TokenValidationResult::Valid { - expires_in_secs: secs, - }, - None => { - // 无法解析过期时间,视为已过期(Requirements: 1.4) - TokenValidationResult::Invalid { - reason: "无法解析过期时间格式".to_string(), - } - } - } - } - - /// 分类 Token 刷新错误 - /// Requirements: 2.1 - fn classify_refresh_error(status: u16, body: &str) -> TokenRefreshError { - // 检查是否是 invalid_grant 错误 - if status == 400 && body.contains("invalid_grant") { - return TokenRefreshError::InvalidGrant { - message: "Refresh token 已失效或被撤销".to_string(), - }; - } - - // 服务器错误 (5xx) - if status >= 500 { - return TokenRefreshError::ServerError { - message: format!("HTTP {status}: {body}"), - }; - } - - // 其他客户端错误 - if status >= 400 { - return TokenRefreshError::Unknown { - message: format!("HTTP {status}: {body}"), - }; - } - - TokenRefreshError::Unknown { - message: body.to_string(), - } - } - - /// 带重试的 Token 刷新 - /// Requirements: 2.2, 2.3 - pub async fn refresh_token_with_retry( - &mut self, - max_retries: u32, - ) -> Result { - let refresh_token = self - .credentials - .refresh_token - .as_ref() - .ok_or_else(|| TokenRefreshError::InvalidGrant { - message: "No refresh token available".to_string(), - })? - .clone(); - - let client_id = oauth_client_id(); - let client_secret = oauth_client_secret(); - let params = [ - ("client_id", client_id.as_str()), - ("client_secret", client_secret.as_str()), - ("refresh_token", refresh_token.as_str()), - ("grant_type", "refresh_token"), - ]; - - let mut last_error: Option = None; - let mut retry_count = 0; - - while retry_count <= max_retries { - if retry_count > 0 { - // 指数退避:100ms, 200ms, 400ms, ... - let delay_ms = 100 * (1 << (retry_count - 1)); - tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await; - tracing::info!( - "[Antigravity] Token 刷新重试 {}/{}, 延迟 {}ms", - retry_count, - max_retries, - delay_ms - ); - } - - let result = self - .client - .post("https://oauth2.googleapis.com/token") - .form(¶ms) - .send() - .await; - - match result { - Ok(resp) => { - let status = resp.status(); - if status.is_success() { - // 成功,解析响应 - match resp.json::().await { - Ok(data) => { - let new_token = data["access_token"].as_str().ok_or_else(|| { - TokenRefreshError::Unknown { - message: "响应中缺少 access_token".to_string(), - } - })?; - - self.credentials.access_token = Some(new_token.to_string()); - - // 更新过期时间 - if let Some(expires_in) = data["expires_in"].as_i64() { - let now = chrono::Utc::now(); - let expires_at = now + chrono::Duration::seconds(expires_in); - self.credentials.expire = Some(expires_at.to_rfc3339()); - self.credentials.expiry_date = - Some(expires_at.timestamp_millis()); - self.credentials.expires_in = Some(expires_in); - self.credentials.timestamp = Some(now.timestamp_millis()); - } - - // 更新 refresh_token(如果返回了新的) - if let Some(new_refresh) = data["refresh_token"].as_str() { - self.credentials.refresh_token = Some(new_refresh.to_string()); - } - - self.credentials.last_refresh = - Some(chrono::Utc::now().to_rfc3339()); - - // 保存凭证 - if let Err(e) = self.save_credentials().await { - tracing::warn!("[Antigravity] 保存凭证失败: {}", e); - } - - return Ok(new_token.to_string()); - } - Err(e) => { - last_error = Some(TokenRefreshError::Unknown { - message: format!("解析响应失败: {e}"), - }); - } - } - } else { - // 请求失败 - let status_code = status.as_u16(); - let body = resp.text().await.unwrap_or_default(); - let error = Self::classify_refresh_error(status_code, &body); - - // invalid_grant 不重试 - if error.requires_reauth() { - return Err(error); - } - - // 可重试的错误 - if error.is_retryable() { - last_error = Some(error); - retry_count += 1; - continue; - } - - return Err(error); - } - } - Err(e) => { - // 网络错误,可重试 - last_error = Some(TokenRefreshError::NetworkError { - message: e.to_string(), - }); - retry_count += 1; - continue; - } - } - - retry_count += 1; - } - - // 所有重试都失败 - Err(last_error.unwrap_or_else(|| TokenRefreshError::Unknown { - message: "Token 刷新失败,已达到最大重试次数".to_string(), - })) - } - - pub async fn refresh_token(&mut self) -> Result> { - let refresh_token = self - .credentials - .refresh_token - .as_ref() - .ok_or("No refresh token available")?; - - let client_id = oauth_client_id(); - let client_secret = oauth_client_secret(); - let params = [ - ("client_id", client_id.as_str()), - ("client_secret", client_secret.as_str()), - ("refresh_token", refresh_token.as_str()), - ("grant_type", "refresh_token"), - ]; - - let resp = self - .client - .post("https://oauth2.googleapis.com/token") - .form(¶ms) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(format!("Token refresh failed: {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - - let new_token = data["access_token"] - .as_str() - .ok_or("No access token in response")?; - - self.credentials.access_token = Some(new_token.to_string()); - - // 更新过期时间(同时保存多种格式以兼容) - if let Some(expires_in) = data["expires_in"].as_i64() { - let now = chrono::Utc::now(); - let expires_at = now + chrono::Duration::seconds(expires_in); - - // RFC3339 格式 - self.credentials.expire = Some(expires_at.to_rfc3339()); - // 毫秒时间戳格式 - self.credentials.expiry_date = Some(expires_at.timestamp_millis()); - // antigravity2api-nodejs 格式 - self.credentials.expires_in = Some(expires_in); - self.credentials.timestamp = Some(now.timestamp_millis()); - } - - // 如果返回了新的 refresh_token,也更新它 - if let Some(new_refresh) = data["refresh_token"].as_str() { - self.credentials.refresh_token = Some(new_refresh.to_string()); - } - - // 更新最后刷新时间(RFC3339 格式) - self.credentials.last_refresh = Some(chrono::Utc::now().to_rfc3339()); - - // Save refreshed credentials - self.save_credentials().await?; - - Ok(new_token.to_string()) - } - - /// 调用 Antigravity API(内部方法) - /// - /// 返回 `AntigravityApiError` 以便调用方获取 HTTP 状态码 - async fn call_api_internal( - &self, - base_url: &str, - method: &str, - body: &serde_json::Value, - ) -> Result { - let token = self - .credentials - .access_token - .as_ref() - .ok_or_else(|| AntigravityApiError::new(401, "No access token"))?; - - let url = format!("{base_url}/{ANTIGRAVITY_API_VERSION}:{method}"); - - // 打印详细的请求信息 - eprintln!("========== [ANTIGRAVITY_API] 请求详情 =========="); - eprintln!("[ANTIGRAVITY_API] URL: {url}"); - eprintln!("[ANTIGRAVITY_API] Method: {method}"); - eprintln!( - "[ANTIGRAVITY_API] Token (前20字符): {}...", - &token[..token.len().min(20)] - ); - eprintln!( - "[ANTIGRAVITY_API] 请求体: {}", - serde_json::to_string_pretty(body).unwrap_or_default() - ); - - let resp = self - .client - .post(&url) - .header("Authorization", format!("Bearer {token}")) - .header("Content-Type", "application/json") - .header("User-Agent", "antigravity/1.11.9 windows/amd64") - .json(body) - .send() - .await - .map_err(|e| { - eprintln!("[ANTIGRAVITY_API] 网络错误: {e}"); - AntigravityApiError::new(503, format!("Network error: {e}")) - })?; - - let status = resp.status(); - let status_code = status.as_u16(); - eprintln!("[ANTIGRAVITY_API] 响应状态码: {status}"); - - if !status.is_success() { - let body_text = resp.text().await.unwrap_or_default(); - eprintln!("[ANTIGRAVITY_API] 错误响应体: {body_text}"); - eprintln!("========== [ANTIGRAVITY_API] 请求失败 =========="); - return Err(AntigravityApiError::with_body( - status_code, - format!("API call failed: {status}"), - body_text, - )); - } - - let response_text = resp - .text() - .await - .map_err(|e| AntigravityApiError::new(500, format!("Failed to read response: {e}")))?; - eprintln!("[ANTIGRAVITY_API] 响应体: {response_text}"); - - let data: serde_json::Value = serde_json::from_str(&response_text) - .map_err(|e| AntigravityApiError::new(500, format!("Failed to parse response: {e}")))?; - - eprintln!("========== [ANTIGRAVITY_API] 请求成功 =========="); - Ok(data) - } - - /// 调用 API,支持多环境降级 - /// - /// 返回 `AntigravityApiError` 以便调用方获取 HTTP 状态码并透传给客户端。 - /// 只在特定错误(429 配额耗尽、5xx 服务器错误)时才降级到备用端点, - /// 403 权限错误等不应该降级,因为换端点也没用。 - pub async fn call_api( - &self, - method: &str, - body: &serde_json::Value, - ) -> Result { - let mut last_error: Option = None; - - for (idx, base_url) in self.base_urls.iter().enumerate() { - match self.call_api_internal(base_url, method, body).await { - Ok(data) => return Ok(data), - Err(e) => { - // 使用 AntigravityApiError 的方法判断是否可重试 - let should_fallback = e.is_retryable(); - - if should_fallback && idx + 1 < self.base_urls.len() { - tracing::warn!( - "[Antigravity] {} 返回可重试错误 (HTTP {}), 尝试下一个端点", - base_url, - e.status_code - ); - last_error = Some(e); - continue; - } - - // 403、401 等权限错误直接返回,不降级 - tracing::warn!( - "[Antigravity] {} 失败 (HTTP {}): {}", - base_url, - e.status_code, - e.message - ); - return Err(e); - } - } - } - - Err(last_error - .unwrap_or_else(|| AntigravityApiError::new(503, "All Antigravity base URLs failed"))) - } - - /// 发现项目 ID - pub async fn discover_project(&mut self) -> Result> { - if let Some(ref project_id) = self.project_id { - return Ok(project_id.clone()); - } - - let body = serde_json::json!({ - "cloudaicompanionProject": "", - "metadata": { - "ideType": "IDE_UNSPECIFIED", - "platform": "PLATFORM_UNSPECIFIED", - "pluginType": "GEMINI", - "duetProject": "" - } - }); - - let resp = self.call_api("loadCodeAssist", &body).await?; - - if let Some(project) = resp["cloudaicompanionProject"].as_str() { - if !project.is_empty() { - self.project_id = Some(project.to_string()); - return Ok(project.to_string()); - } - } - - // Need to onboard - let onboard_body = serde_json::json!({ - "tierId": "free-tier", - "cloudaicompanionProject": "", - "metadata": { - "ideType": "IDE_UNSPECIFIED", - "platform": "PLATFORM_UNSPECIFIED", - "pluginType": "GEMINI", - "duetProject": "" - } - }); - - let mut lro_resp = self.call_api("onboardUser", &onboard_body).await?; - - // Poll until done - for _ in 0..30 { - if lro_resp["done"].as_bool().unwrap_or(false) { - break; - } - tokio::time::sleep(tokio::time::Duration::from_secs(2)).await; - lro_resp = self.call_api("onboardUser", &onboard_body).await?; - } - - let project_id = lro_resp["response"]["cloudaicompanionProject"]["id"] - .as_str() - .unwrap_or("") - .to_string(); - - if project_id.is_empty() { - // 生成一个随机项目 ID 作为后备 - let fallback = generate_project_id(); - self.project_id = Some(fallback.clone()); - return Ok(fallback); - } - - self.project_id = Some(project_id.clone()); - Ok(project_id) - } - - /// 获取可用模型列表 - pub async fn fetch_available_models( - &mut self, - ) -> Result, Box> { - let body = serde_json::json!({}); - - match self.call_api("fetchAvailableModels", &body).await { - Ok(resp) => { - if let Some(models) = resp["models"].as_object() { - self.available_models = models - .keys() - .filter_map(|name| { - let alias = model_name_to_alias(name); - if alias.is_empty() { - None - } else { - Some(alias.to_string()) - } - }) - .collect(); - } - } - Err(e) => { - tracing::warn!( - "[Antigravity] Failed to fetch models: {}, using defaults", - e - ); - } - } - - Ok(self.available_models.clone()) - } - - /// 生成内容(非流式) - /// - /// 返回 `AntigravityApiError` 以便调用方获取 HTTP 状态码并透传给客户端。 - pub async fn generate_content( - &self, - model: &str, - request_body: &serde_json::Value, - ) -> Result { - eprintln!("========== [ANTIGRAVITY_GENERATE] 开始生成内容 =========="); - eprintln!("[ANTIGRAVITY_GENERATE] 模型: {model}"); - eprintln!( - "[ANTIGRAVITY_GENERATE] 请求体: {}", - serde_json::to_string_pretty(request_body).unwrap_or_default() - ); - - let project_id = self.project_id.clone().unwrap_or_else(generate_project_id); - eprintln!("[ANTIGRAVITY_GENERATE] 项目ID: {project_id}"); - - let actual_model = alias_to_model_name(model); - eprintln!("[ANTIGRAVITY_GENERATE] 实际模型名: {actual_model}"); - - let payload = self.build_antigravity_request(&actual_model, &project_id, request_body); - eprintln!( - "[ANTIGRAVITY_GENERATE] 构建的 payload: {}", - serde_json::to_string_pretty(&payload).unwrap_or_default() - ); - - eprintln!("[ANTIGRAVITY_GENERATE] 调用 call_api..."); - let resp = self.call_api("generateContent", &payload).await?; - eprintln!("[ANTIGRAVITY_GENERATE] call_api 返回成功"); - - // 转换为 Gemini 格式响应 - let result = self.to_gemini_response(&resp); - eprintln!("========== [ANTIGRAVITY_GENERATE] 生成内容完成 =========="); - Ok(result) - } - - /// 构建 Antigravity 请求 - /// - /// 注意:此方法用于简单的非流式请求。 - /// 对于完整的 OpenAI 格式转换,请使用 `convert_openai_to_antigravity_with_context`。 - fn build_antigravity_request( - &self, - model: &str, - project_id: &str, - request_body: &serde_json::Value, - ) -> serde_json::Value { - let mut payload = request_body.clone(); - - // 设置基本字段 - payload["model"] = serde_json::json!(model); - payload["userAgent"] = serde_json::json!("antigravity"); - payload["project"] = serde_json::json!(project_id); - payload["requestId"] = serde_json::json!(generate_request_id()); - - // 确保 request 对象存在 - if payload.get("request").is_none() { - payload["request"] = serde_json::json!({}); - } - - // 设置会话 ID - payload["request"]["sessionId"] = serde_json::json!(generate_session_id()); - - // 添加默认安全设置(如果不存在) - if payload - .get("request") - .and_then(|r| r.get("safetySettings")) - .is_none() - { - payload["request"]["safetySettings"] = serde_json::json!([ - {"category": "HARM_CATEGORY_HARASSMENT", "threshold": "OFF"}, - {"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "OFF"}, - {"category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", "threshold": "OFF"}, - {"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "OFF"}, - {"category": "HARM_CATEGORY_CIVIC_INTEGRITY", "threshold": "BLOCK_NONE"} - ]); - } - - payload - } - - /// 转换为 Gemini 格式响应 - fn to_gemini_response(&self, antigravity_resp: &serde_json::Value) -> serde_json::Value { - let mut response = serde_json::json!({}); - - if let Some(candidates) = antigravity_resp.get("candidates") { - response["candidates"] = candidates.clone(); - } - - if let Some(usage) = antigravity_resp.get("usageMetadata") { - response["usageMetadata"] = usage.clone(); - } - - if let Some(feedback) = antigravity_resp.get("promptFeedback") { - response["promptFeedback"] = feedback.clone(); - } - - response - } - - /// 检查模型是否支持 - pub fn supports_model(&self, model: &str) -> bool { - self.available_models.iter().any(|m| m == model) - } -} - -// ============================================================================ -// OAuth 登录功能 -// ============================================================================ - -/// OAuth 回调结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct OAuthCallbackResult { - pub code: String, - pub state: String, -} - -/// OAuth 登录成功后的凭证信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AntigravityOAuthResult { - pub credentials: AntigravityCredentials, - pub creds_file_path: String, -} - -/// 生成 OAuth 授权 URL -pub fn generate_auth_url(port: u16, state: &str) -> String { - let scopes = OAUTH_SCOPES.join(" "); - let redirect_uri = format!("http://localhost:{port}/oauth-callback"); - let client_id = oauth_client_id(); - - let params = [ - ("access_type", "offline"), - ("client_id", client_id.as_str()), - ("prompt", "consent"), - ("redirect_uri", &redirect_uri), - ("response_type", "code"), - ("scope", &scopes), - ("state", state), - ]; - - let query = params - .iter() - .map(|(k, v)| format!("{}={}", k, urlencoding::encode(v))) - .collect::>() - .join("&"); - - format!("https://accounts.google.com/o/oauth2/v2/auth?{query}") -} - -/// 用授权码交换 Token -pub async fn exchange_code_for_token( - client: &Client, - code: &str, - redirect_uri: &str, -) -> Result> { - let client_id = oauth_client_id(); - let client_secret = oauth_client_secret(); - let params = [ - ("code", code), - ("client_id", client_id.as_str()), - ("client_secret", client_secret.as_str()), - ("redirect_uri", redirect_uri), - ("grant_type", "authorization_code"), - ]; - - let resp = client - .post("https://oauth2.googleapis.com/token") - .form(¶ms) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(format!("Token 交换失败: {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - Ok(data) -} - -/// 获取用户邮箱 -pub async fn fetch_user_email( - client: &Client, - access_token: &str, -) -> Result, Box> { - let resp = client - .get("https://www.googleapis.com/oauth2/v2/userinfo") - .header("Authorization", format!("Bearer {access_token}")) - .send() - .await?; - - if resp.status().is_success() { - let data: serde_json::Value = resp.json().await?; - Ok(data["email"].as_str().map(|s| s.to_string())) - } else { - Ok(None) - } -} - -/// 获取项目 ID(验证账号资格) -/// 返回值说明: -/// - Ok(Some(FetchedProjectId::HasProject(id))) - 有资格,且有 projectId -/// - Ok(Some(FetchedProjectId::NoProject)) - 有资格,但 projectId 为空(需要生成随机 ID) -/// - Ok(None) - 无资格(字段不存在,即 undefined) -/// - Err(_) - 请求失败 -#[derive(Debug, Clone)] -pub enum FetchedProjectId { - /// 有 projectId - HasProject(String), - /// projectId 为空字符串(有资格但无 projectId) - NoProject, -} - -pub async fn fetch_project_id_for_oauth( - client: &Client, - access_token: &str, -) -> Result, Box> { - tracing::info!("[Antigravity OAuth] 正在获取 projectId..."); - - let resp = client - .post("https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:loadCodeAssist") - .header("Authorization", format!("Bearer {access_token}")) - .header("User-Agent", "antigravity/1.11.9 windows/amd64") - .header("Content-Type", "application/json") - .json(&serde_json::json!({ "metadata": { "ideType": "ANTIGRAVITY" } })) - .send() - .await?; - - let status = resp.status(); - tracing::info!("[Antigravity OAuth] loadCodeAssist 响应状态: {}", status); - - if status.is_success() { - let body_text = resp.text().await?; - tracing::info!("[Antigravity OAuth] loadCodeAssist 响应体: {}", body_text); - - let data: serde_json::Value = serde_json::from_str(&body_text)?; - // 检查字段是否存在 - // - 如果字段不存在(undefined)-> None(无资格) - // - 如果字段存在但为空字符串 -> Some(NoProject)(有资格但无 projectId) - // - 如果字段存在且有值 -> Some(HasProject(id))(有资格且有 projectId) - match data.get("cloudaicompanionProject") { - None => { - tracing::warn!("[Antigravity OAuth] cloudaicompanionProject 字段不存在"); - Ok(None) // 字段不存在,无资格 - } - Some(value) => { - if value.is_null() { - tracing::warn!("[Antigravity OAuth] cloudaicompanionProject 为 null"); - Ok(None) // null 也视为无资格 - } else if let Some(s) = value.as_str() { - if s.is_empty() { - tracing::info!("[Antigravity OAuth] cloudaicompanionProject 为空字符串,有资格但无 projectId"); - Ok(Some(FetchedProjectId::NoProject)) // 空字符串,有资格但无 projectId - } else { - tracing::info!("[Antigravity OAuth] 获取到 project_id: {}", s); - Ok(Some(FetchedProjectId::HasProject(s.to_string()))) // 有 projectId - } - } else { - tracing::warn!( - "[Antigravity OAuth] cloudaicompanionProject 不是字符串类型: {:?}", - value - ); - Ok(None) // 非字符串类型,视为无资格 - } - } - } - } else { - let body = resp.text().await.unwrap_or_default(); - tracing::error!( - "[Antigravity OAuth] loadCodeAssist 请求失败: {} - {}", - status, - body - ); - Err(format!("loadCodeAssist 请求失败: {status} - {body}").into()) - } -} - -/// OAuth 成功页面 HTML -const OAUTH_SUCCESS_HTML: &str = r#" - - - - 授权成功 - - - -
-

✓ 授权成功

-

账号已添加到 Lime

- -

可以关闭此页面

-
- -"#; - -/// OAuth 失败页面 HTML -const OAUTH_ERROR_HTML: &str = r#" - - - - 授权失败 - - - -
-

✗ 授权失败

-

ERROR_PLACEHOLDER

-

请关闭此页面后重试

-
- -"#; - -/// OAuth 授权 URL 结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct OAuthAuthUrlResult { - pub auth_url: String, - pub port: u16, - pub state: String, -} - -/// 启动 OAuth 服务器并返回授权 URL(不打开浏览器) -/// 服务器会在后台等待回调,成功后返回凭证 -pub async fn start_oauth_server_and_get_url( - skip_project_id_fetch: bool, -) -> Result< - ( - String, - impl std::future::Future>>, - ), - Box, -> { - use axum::{extract::Query, response::Html, routing::get, Router}; - use std::collections::HashMap; - use tokio::net::TcpListener; - - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build()?; - - // 生成随机 state - let state = Uuid::new_v4().to_string(); - let state_clone = state.clone(); - - // 创建 channel 用于接收回调结果 - let (tx, rx) = oneshot::channel::>(); - let tx = Arc::new(tokio::sync::Mutex::new(Some(tx))); - - // 绑定到随机端口 - let listener = TcpListener::bind("127.0.0.1:0").await?; - let port = listener.local_addr()?.port(); - - let redirect_uri = format!("http://localhost:{port}/oauth-callback"); - let redirect_uri_clone = redirect_uri.clone(); - - // 生成授权 URL - let auth_url = generate_auth_url(port, &state); - - tracing::info!( - "[Antigravity OAuth] 服务器启动在端口 {}, 授权 URL: {}", - port, - auth_url - ); - - // 构建路由 - let app = Router::new().route( - "/oauth-callback", - get(move |Query(params): Query>| { - let tx = tx.clone(); - let client = client.clone(); - let state_expected = state_clone.clone(); - let redirect_uri = redirect_uri_clone.clone(); - - async move { - let code = params.get("code"); - let returned_state = params.get("state"); - let error = params.get("error"); - - // 检查错误 - if let Some(err) = error { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", err); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(format!("OAuth 错误: {err}"))); - } - return Html(html); - } - - // 检查 state - if returned_state.map(|s| s.as_str()) != Some(&state_expected) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "State 验证失败"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("State 验证失败".to_string())); - } - return Html(html); - } - - // 检查 code - let code = match code { - Some(c) => c, - None => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "未收到授权码"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("未收到授权码".to_string())); - } - return Html(html); - } - }; - - // 交换 Token - let token_result = exchange_code_for_token(&client, code, &redirect_uri).await; - let token_data = match token_result { - Ok(data) => data, - Err(e) => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &e.to_string()); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - }; - - let access_token = token_data["access_token"].as_str().unwrap_or_default(); - let refresh_token = token_data["refresh_token"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expires_in"].as_i64(); - - // 获取用户邮箱 - let email = fetch_user_email(&client, access_token).await.ok().flatten(); - - // 获取项目 ID - let project_id = if skip_project_id_fetch { - tracing::info!("[Antigravity OAuth] 跳过 projectId 获取,使用随机生成的 ID"); - Some(generate_project_id()) - } else { - match fetch_project_id_for_oauth(&client, access_token).await { - Ok(Some(FetchedProjectId::HasProject(pid))) => Some(pid), - Ok(Some(FetchedProjectId::NoProject)) => { - tracing::info!("[Antigravity OAuth] projectId 为空,使用随机生成的 ID"); - Some(generate_project_id()) - } - Ok(None) => { - tracing::warn!("[Antigravity OAuth] 无法获取 projectId(字段不存在),使用随机生成的 ID"); - Some(generate_project_id()) - } - Err(e) => { - tracing::warn!("[Antigravity OAuth] 获取 projectId 失败: {}, 使用随机 ID", e); - Some(generate_project_id()) - } - } - }; - - // 构建凭证 - let now = chrono::Utc::now(); - let credentials = AntigravityCredentials { - access_token: Some(access_token.to_string()), - refresh_token, - token_type: Some("Bearer".to_string()), - expiry_date: expires_in.map(|e| (now + chrono::Duration::seconds(e)).timestamp_millis()), - expire: expires_in.map(|e| (now + chrono::Duration::seconds(e)).to_rfc3339()), - scope: Some(OAUTH_SCOPES.join(" ")), - last_refresh: Some(now.to_rfc3339()), - cred_type: "antigravity".to_string(), - expires_in, - timestamp: Some(now.timestamp_millis()), - enable: Some(true), - project_id, - email: email.clone(), - }; - - // 保存凭证到应用数据目录 - let creds_dir = dirs::data_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join("lime") - .join("credentials") - .join("antigravity"); - - if let Err(e) = std::fs::create_dir_all(&creds_dir) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("创建目录失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - - // 使用 UUID 作为文件名 - let file_name = format!("{}.json", Uuid::new_v4()); - let creds_path = creds_dir.join(&file_name); - - let creds_json = match serde_json::to_string_pretty(&credentials) { - Ok(json) => json, - Err(e) => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("序列化失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - }; - - if let Err(e) = std::fs::write(&creds_path, &creds_json) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("保存凭证失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - - let creds_path_str = creds_path.to_string_lossy().to_string(); - tracing::info!("[Antigravity OAuth] 凭证已保存到: {}", creds_path_str); - - // 发送成功结果 - let result = AntigravityOAuthResult { - credentials, - creds_file_path: creds_path_str, - }; - - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Ok(result)); - } - - let email_display = email.unwrap_or_else(|| "未知邮箱".to_string()); - let html = OAUTH_SUCCESS_HTML.replace("EMAIL_PLACEHOLDER", &email_display); - Html(html) - } - }), - ); - - // 创建等待回调的 Future - let wait_future = async move { - // 启动服务器 - let server = axum::serve(listener, app); - - // 同时运行服务器和等待回调结果 - tokio::select! { - result = async { - tokio::time::timeout( - std::time::Duration::from_secs(300), - rx - ).await - } => { - match result { - Ok(Ok(Ok(r))) => Ok(r), - Ok(Ok(Err(e))) => Err(e.into()), - Ok(Err(_)) => Err("OAuth 回调通道关闭".into()), - Err(_) => Err("OAuth 登录超时(5分钟)".into()), - } - } - server_result = server => { - match server_result { - Ok(_) => Err("服务器意外关闭".into()), - Err(e) => Err(format!("服务器错误: {e}").into()), - } - } - } - }; - - Ok((auth_url, wait_future)) -} - -/// 启动 OAuth 登录流程(使用指定端口) -/// 用于配合 get_oauth_auth_url 使用 -pub async fn start_oauth_login_with_port( - port: u16, - state: String, - skip_project_id_fetch: bool, -) -> Result> { - use axum::{extract::Query, response::Html, routing::get, Router}; - use std::collections::HashMap; - use tokio::net::TcpListener; - - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build()?; - - let state_clone = state.clone(); - - // 创建 channel 用于接收回调结果 - let (tx, rx) = oneshot::channel::>(); - let tx = Arc::new(tokio::sync::Mutex::new(Some(tx))); - - // 绑定到指定端口 - let listener = TcpListener::bind(format!("127.0.0.1:{port}")).await?; - - let redirect_uri = format!("http://localhost:{port}/oauth-callback"); - let redirect_uri_clone = redirect_uri.clone(); - - // 构建路由 - let app = Router::new().route( - "/oauth-callback", - get(move |Query(params): Query>| { - let tx = tx.clone(); - let client = client.clone(); - let state_expected = state_clone.clone(); - let redirect_uri = redirect_uri_clone.clone(); - - async move { - let code = params.get("code"); - let returned_state = params.get("state"); - let error = params.get("error"); - - // 检查错误 - if let Some(err) = error { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", err); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(format!("OAuth 错误: {err}"))); - } - return Html(html); - } - - // 检查 state - if returned_state.map(|s| s.as_str()) != Some(&state_expected) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "State 验证失败"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("State 验证失败".to_string())); - } - return Html(html); - } - - // 检查 code - let code = match code { - Some(c) => c, - None => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "未收到授权码"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("未收到授权码".to_string())); - } - return Html(html); - } - }; - - // 交换 Token - let token_result = exchange_code_for_token(&client, code, &redirect_uri).await; - let token_data = match token_result { - Ok(data) => data, - Err(e) => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &e.to_string()); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - }; - - let access_token = token_data["access_token"].as_str().unwrap_or_default(); - let refresh_token = token_data["refresh_token"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expires_in"].as_i64(); - - // 获取用户邮箱 - let email = fetch_user_email(&client, access_token).await.ok().flatten(); - - // 获取项目 ID - let project_id = if skip_project_id_fetch { - tracing::info!("[Antigravity OAuth] 跳过 projectId 获取,使用随机生成的 ID"); - Some(generate_project_id()) - } else { - match fetch_project_id_for_oauth(&client, access_token).await { - Ok(Some(FetchedProjectId::HasProject(pid))) => Some(pid), - Ok(Some(FetchedProjectId::NoProject)) => { - tracing::info!("[Antigravity OAuth] projectId 为空,使用随机生成的 ID"); - Some(generate_project_id()) - } - Ok(None) => { - tracing::warn!("[Antigravity OAuth] 无法获取 projectId(字段不存在),使用随机生成的 ID"); - Some(generate_project_id()) - } - Err(e) => { - tracing::warn!("[Antigravity OAuth] 获取 projectId 失败: {}, 使用随机 ID", e); - Some(generate_project_id()) - } - } - }; - - // 构建凭证 - let now = chrono::Utc::now(); - let credentials = AntigravityCredentials { - access_token: Some(access_token.to_string()), - refresh_token, - token_type: Some("Bearer".to_string()), - expiry_date: expires_in.map(|e| (now + chrono::Duration::seconds(e)).timestamp_millis()), - expire: expires_in.map(|e| (now + chrono::Duration::seconds(e)).to_rfc3339()), - scope: Some(OAUTH_SCOPES.join(" ")), - last_refresh: Some(now.to_rfc3339()), - cred_type: "antigravity".to_string(), - expires_in, - timestamp: Some(now.timestamp_millis()), - enable: Some(true), - project_id, - email: email.clone(), - }; - - // 保存凭证到应用数据目录 - let creds_dir = dirs::data_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join("lime") - .join("credentials") - .join("antigravity"); - - if let Err(e) = std::fs::create_dir_all(&creds_dir) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("创建目录失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - - // 使用 UUID 作为文件名 - let file_name = format!("{}.json", Uuid::new_v4()); - let creds_path = creds_dir.join(&file_name); - - let creds_json = match serde_json::to_string_pretty(&credentials) { - Ok(json) => json, - Err(e) => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("序列化失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - }; - - if let Err(e) = std::fs::write(&creds_path, &creds_json) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("保存凭证失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - - let creds_path_str = creds_path.to_string_lossy().to_string(); - tracing::info!("[Antigravity OAuth] 凭证已保存到: {}", creds_path_str); - - // 发送成功结果 - let result = AntigravityOAuthResult { - credentials, - creds_file_path: creds_path_str, - }; - - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Ok(result)); - } - - let email_display = email.unwrap_or_else(|| "未知邮箱".to_string()); - let html = OAUTH_SUCCESS_HTML.replace("EMAIL_PLACEHOLDER", &email_display); - Html(html) - } - }), - ); - - // 启动服务器 - let server = axum::serve(listener, app); - - // 同时运行服务器和等待回调结果 - tokio::select! { - result = async { - tokio::time::timeout( - std::time::Duration::from_secs(300), - rx - ).await - } => { - match result { - Ok(Ok(Ok(r))) => Ok(r), - Ok(Ok(Err(e))) => Err(e.into()), - Ok(Err(_)) => Err("OAuth 回调通道关闭".into()), - Err(_) => Err("OAuth 登录超时(5分钟)".into()), - } - } - server_result = server => { - match server_result { - Ok(_) => Err("服务器意外关闭".into()), - Err(e) => Err(format!("服务器错误: {e}").into()), - } - } - } -} - -/// 启动 OAuth 登录流程 -/// 返回 (auth_url, credentials_file_path) -pub async fn start_oauth_login( - skip_project_id_fetch: bool, -) -> Result> { - use axum::{extract::Query, response::Html, routing::get, Router}; - use std::collections::HashMap; - use tokio::net::TcpListener; - - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build()?; - - // 生成随机 state - let state = Uuid::new_v4().to_string(); - let state_clone = state.clone(); - - // 创建 channel 用于接收回调结果 - let (tx, rx) = oneshot::channel::>(); - let tx = Arc::new(tokio::sync::Mutex::new(Some(tx))); - - // 绑定到随机端口 - let listener = TcpListener::bind("127.0.0.1:0").await?; - let port = listener.local_addr()?.port(); - - let redirect_uri = format!("http://localhost:{port}/oauth-callback"); - let redirect_uri_clone = redirect_uri.clone(); - - // 构建路由 - let app = Router::new().route( - "/oauth-callback", - get(move |Query(params): Query>| { - let tx = tx.clone(); - let client = client.clone(); - let state_expected = state_clone.clone(); - let redirect_uri = redirect_uri_clone.clone(); - - async move { - let code = params.get("code"); - let returned_state = params.get("state"); - let error = params.get("error"); - - // 检查错误 - if let Some(err) = error { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", err); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(format!("OAuth 错误: {err}"))); - } - return Html(html); - } - - // 检查 state - if returned_state.map(|s| s.as_str()) != Some(&state_expected) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "State 验证失败"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("State 验证失败".to_string())); - } - return Html(html); - } - - // 检查 code - let code = match code { - Some(c) => c, - None => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "未收到授权码"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("未收到授权码".to_string())); - } - return Html(html); - } - }; - - // 交换 Token - let token_result = exchange_code_for_token(&client, code, &redirect_uri).await; - let token_data = match token_result { - Ok(data) => data, - Err(e) => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &e.to_string()); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - }; - - let access_token = token_data["access_token"].as_str().unwrap_or_default(); - let refresh_token = token_data["refresh_token"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expires_in"].as_i64(); - - // 获取用户邮箱 - let email = fetch_user_email(&client, access_token).await.ok().flatten(); - - // 获取项目 ID - // 参考 antigravity2api-nodejs 的逻辑: - // - projectId === undefined -> 无资格(但我们改为使用随机 ID,因为很多账号都没有 projectId) - // - projectId === "" -> 有资格但无 projectId,使用随机生成的 - // - projectId 有值 -> 有资格且有 projectId - let project_id = if skip_project_id_fetch { - tracing::info!("[Antigravity OAuth] 跳过 projectId 获取,使用随机生成的 ID"); - Some(generate_project_id()) - } else { - match fetch_project_id_for_oauth(&client, access_token).await { - Ok(Some(FetchedProjectId::HasProject(pid))) => { - // 有资格且有 projectId - Some(pid) - } - Ok(Some(FetchedProjectId::NoProject)) => { - // 有资格但 projectId 为空,使用随机生成的 - tracing::info!("[Antigravity OAuth] projectId 为空,使用随机生成的 ID"); - Some(generate_project_id()) - } - Ok(None) => { - // 字段不存在,也使用随机 ID(很多账号都是这种情况) - tracing::warn!("[Antigravity OAuth] 无法获取 projectId(字段不存在),使用随机生成的 ID"); - Some(generate_project_id()) - } - Err(e) => { - tracing::warn!("[Antigravity OAuth] 获取 projectId 失败: {}, 使用随机 ID", e); - Some(generate_project_id()) - } - } - }; - - // 构建凭证 - let now = chrono::Utc::now(); - let credentials = AntigravityCredentials { - access_token: Some(access_token.to_string()), - refresh_token, - token_type: Some("Bearer".to_string()), - expiry_date: expires_in.map(|e| (now + chrono::Duration::seconds(e)).timestamp_millis()), - expire: expires_in.map(|e| (now + chrono::Duration::seconds(e)).to_rfc3339()), - scope: Some(OAUTH_SCOPES.join(" ")), - last_refresh: Some(now.to_rfc3339()), - cred_type: "antigravity".to_string(), - expires_in, - timestamp: Some(now.timestamp_millis()), - enable: Some(true), - project_id, - email: email.clone(), - }; - - // 保存凭证到应用数据目录 - let creds_dir = dirs::data_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join("lime") - .join("credentials") - .join("antigravity"); - - if let Err(e) = std::fs::create_dir_all(&creds_dir) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("创建目录失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - - // 使用 UUID 作为文件名 - let file_name = format!("{}.json", Uuid::new_v4()); - let creds_path = creds_dir.join(&file_name); - - let creds_json = match serde_json::to_string_pretty(&credentials) { - Ok(json) => json, - Err(e) => { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("序列化失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - }; - - if let Err(e) = std::fs::write(&creds_path, &creds_json) { - let html = OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &format!("保存凭证失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - - let creds_path_str = creds_path.to_string_lossy().to_string(); - tracing::info!("[Antigravity OAuth] 凭证已保存到: {}", creds_path_str); - - // 发送成功结果 - let result = AntigravityOAuthResult { - credentials, - creds_file_path: creds_path_str, - }; - - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Ok(result)); - } - - let email_display = email.unwrap_or_else(|| "未知邮箱".to_string()); - let html = OAUTH_SUCCESS_HTML.replace("EMAIL_PLACEHOLDER", &email_display); - Html(html) - } - }), - ); - - // 生成授权 URL - let auth_url = generate_auth_url(port, &state); - - // 打开浏览器 - tracing::info!("[Antigravity OAuth] 打开浏览器进行授权: {}", auth_url); - if let Err(e) = open::that(&auth_url) { - tracing::warn!("[Antigravity OAuth] 无法自动打开浏览器: {}", e); - } - - // 启动服务器 - let server = axum::serve(listener, app); - - // 同时运行服务器和等待回调结果 - tokio::select! { - // 等待回调结果(带超时) - result = async { - tokio::time::timeout( - std::time::Duration::from_secs(300), - rx - ).await - } => { - match result { - Ok(Ok(Ok(r))) => Ok(r), - Ok(Ok(Err(e))) => Err(e.into()), - Ok(Err(_)) => Err("OAuth 回调通道关闭".into()), - Err(_) => Err("OAuth 登录超时(5分钟)".into()), - } - } - // 服务器运行(不会主动结束,除非出错) - server_result = server => { - match server_result { - Ok(_) => Err("服务器意外关闭".into()), - Err(e) => Err(format!("服务器错误: {e}").into()), - } - } - } -} - -// ============================================================================ -// CredentialProvider Trait 实现 -// ============================================================================ - -#[async_trait] -impl CredentialProvider for AntigravityProvider { - async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()> { - AntigravityProvider::load_credentials_from_path(self, path).await - } - - async fn save_credentials(&self) -> ProviderResult<()> { - AntigravityProvider::save_credentials(self).await - } - - fn is_token_valid(&self) -> bool { - AntigravityProvider::is_token_valid(self) - } - - fn is_token_expiring_soon(&self) -> bool { - AntigravityProvider::is_token_expiring_soon(self) - } - - async fn refresh_token(&mut self) -> ProviderResult { - AntigravityProvider::refresh_token(self).await - } - - fn get_access_token(&self) -> Option<&str> { - self.credentials.access_token.as_deref() - } - - fn provider_type(&self) -> &'static str { - "antigravity" - } -} - -// ============================================================================ -// StreamingProvider Trait 实现 -// ============================================================================ - -use crate::converter::openai_to_antigravity::convert_openai_to_antigravity_with_context; -use crate::providers::ProviderError; -use crate::streaming::traits::{ - reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider, -}; -use lime_core::models::openai::ChatCompletionRequest; - -#[async_trait] -impl StreamingProvider for AntigravityProvider { - /// 发起流式 API 调用 - /// - /// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。 - /// Antigravity 使用 Gemini 流式格式。 - /// - /// # 需求覆盖 - /// - 需求 1.4: AntigravityProvider 流式支持 - async fn call_api_stream( - &self, - request: &ChatCompletionRequest, - ) -> Result { - tracing::info!("[ANTIGRAVITY_STREAM] ========== call_api_stream 开始 =========="); - - let token = self - .credentials - .access_token - .as_ref() - .ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?; - - tracing::info!("[ANTIGRAVITY_STREAM] Token 长度: {} 字符", token.len()); - - let project_id = self.project_id.clone().unwrap_or_else(generate_project_id); - let actual_model = alias_to_model_name(&request.model); - - tracing::info!( - "[ANTIGRAVITY_STREAM] project_id={}, request.model={}, actual_model={}", - project_id, - request.model, - actual_model - ); - - // 使用统一的转换函数构建请求体 - let payload = convert_openai_to_antigravity_with_context(request, &project_id); - - tracing::info!( - "[ANTIGRAVITY_STREAM] 请求体 (完整): {}", - serde_json::to_string_pretty(&payload).unwrap_or_default() - ); - - // 尝试多个 base URL - let mut last_error: Option = None; - - for base_url in &self.base_urls { - let url = format!("{base_url}/{ANTIGRAVITY_API_VERSION}:streamGenerateContent"); - - eprintln!("[ANTIGRAVITY_STREAM] ========== 发起 HTTP 请求 =========="); - eprintln!("[ANTIGRAVITY_STREAM] URL: {url}"); - eprintln!("[ANTIGRAVITY_STREAM] Model: {actual_model}"); - eprintln!( - "[ANTIGRAVITY_STREAM] Token 前20字符: {}...", - &token[..20.min(token.len())] - ); - tracing::info!( - "[ANTIGRAVITY_STREAM] ========== 发起 HTTP 请求 ==========\n URL: {}\n Model: {}\n Method: POST", - url, - actual_model - ); - - let result = self - .client - .post(&url) - .header("Authorization", format!("Bearer {token}")) - .header("Content-Type", "application/json") - .header("Accept", "text/event-stream") - .header("User-Agent", "antigravity/1.11.9 windows/amd64") - .json(&payload) - .send() - .await; - - match result { - Ok(resp) => { - let status = resp.status(); - eprintln!("[ANTIGRAVITY_STREAM] HTTP 响应状态: {status}"); - tracing::info!("[ANTIGRAVITY_STREAM] HTTP 响应状态: {}", status); - - if status.is_success() { - eprintln!("[ANTIGRAVITY_STREAM] ✓ 流式响应成功建立"); - tracing::info!("[ANTIGRAVITY_STREAM] ✓ 流式响应成功建立,返回流"); - return Ok(reqwest_stream_to_stream_response(resp)); - } else { - let body = resp.text().await.unwrap_or_default(); - eprintln!( - "[ANTIGRAVITY_STREAM] ✗ 请求失败\n Base URL: {}\n Status: {}\n Body: {}", - base_url, - status, - &body[..body.len().min(500)] - ); - tracing::error!( - "[ANTIGRAVITY_STREAM] ✗ 请求失败\n Base URL: {}\n Status: {}\n Body: {}", - base_url, - status, - body - ); - last_error = Some(ProviderError::from_http_status(status.as_u16(), &body)); - } - } - Err(e) => { - eprintln!( - "[ANTIGRAVITY_STREAM] ✗ 连接失败\n Base URL: {base_url}\n Error: {e}" - ); - tracing::error!( - "[ANTIGRAVITY_STREAM] ✗ 连接失败\n Base URL: {}\n Error: {}", - base_url, - e - ); - last_error = Some(ProviderError::from_reqwest_error(&e)); - } - } - } - - tracing::error!("[ANTIGRAVITY_STREAM] 所有 base URL 都失败了"); - Err(last_error.unwrap_or_else(|| { - ProviderError::NetworkError("All Antigravity base URLs failed".to_string()) - })) - } - - fn supports_streaming(&self) -> bool { - self.credentials.access_token.is_some() && self.credentials.enable != Some(false) - } - - fn provider_name(&self) -> &'static str { - "AntigravityProvider" - } - - fn stream_format(&self) -> StreamFormat { - StreamFormat::GeminiStream - } -} - -// ==================== 测试模块 ==================== - -#[cfg(test)] -mod tests { - use super::*; - use proptest::prelude::*; - use std::ffi::OsString; - use std::sync::{Mutex, OnceLock}; - use tempfile::tempdir; - - fn env_lock() -> &'static Mutex<()> { - static LOCK: OnceLock> = OnceLock::new(); - LOCK.get_or_init(|| Mutex::new(())) - } - - struct EnvGuard { - values: Vec<(&'static str, Option)>, - } - - impl EnvGuard { - fn set(entries: &[(&'static str, OsString)]) -> Self { - let mut values = Vec::new(); - for (key, value) in entries { - values.push((*key, std::env::var_os(key))); - std::env::set_var(key, value); - } - Self { values } - } - } - - impl Drop for EnvGuard { - fn drop(&mut self) { - for (key, previous) in self.values.drain(..) { - if let Some(value) = previous { - std::env::set_var(key, value); - } else { - std::env::remove_var(key); - } - } - } - } - - // 辅助函数:检查是否为 Valid 状态 - fn is_valid(result: &TokenValidationResult) -> bool { - matches!(result, TokenValidationResult::Valid { .. }) - } - - // 辅助函数:检查是否为 ExpiringSoon 状态 - fn is_expiring_soon(result: &TokenValidationResult) -> bool { - matches!(result, TokenValidationResult::ExpiringSoon { .. }) - } - - // 辅助函数:检查是否为 Expired 状态 - fn is_expired(result: &TokenValidationResult) -> bool { - matches!(result, TokenValidationResult::Expired) - } - - // 辅助函数:检查是否为 Invalid 状态 - fn is_invalid(result: &TokenValidationResult) -> bool { - matches!(result, TokenValidationResult::Invalid { .. }) - } - - // ==================== Property 1: Token 过期时间解析正确性 ==================== - // Feature: antigravity-token-refresh, Property 1: Token 过期时间解析正确性 - // Validates: Requirements 1.1, 1.3 - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 1: 对于任何有效的过期时间(RFC3339 格式),validate_token() 应正确判断状态 - #[test] - fn prop_validate_token_rfc3339_format( - expires_in_secs in -3600i64..7200i64, // -1小时到2小时 (边界 ±2s 容差避免时间竞争) - ) { - let now = chrono::Utc::now(); - let expires_at = now + chrono::Duration::seconds(expires_in_secs); - - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - provider.credentials.refresh_token = Some("test_refresh".to_string()); - provider.credentials.expire = Some(expires_at.to_rfc3339()); - - let result = provider.validate_token(); - - if expires_in_secs <= -2 { - prop_assert!(is_expired(&result), "Expected Expired for expires_in_secs={}", expires_in_secs); - } else if (2..=TOKEN_EXPIRING_SOON_THRESHOLD - 2).contains(&expires_in_secs) { - prop_assert!(is_expiring_soon(&result), "Expected ExpiringSoon for expires_in_secs={}", expires_in_secs); - } else if expires_in_secs > TOKEN_EXPIRING_SOON_THRESHOLD + 2 { - prop_assert!(is_valid(&result), "Expected Valid for expires_in_secs={}", expires_in_secs); - } - // 边界值 ±2s 内跳过断言,避免时间竞争 - } - - /// Property 1: 对于任何有效的过期时间(毫秒时间戳格式),validate_token() 应正确判断状态 - #[test] - fn prop_validate_token_timestamp_format( - expires_in_secs in -3600i64..7200i64, - ) { - let now = chrono::Utc::now(); - let expires_at = now + chrono::Duration::seconds(expires_in_secs); - - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - provider.credentials.refresh_token = Some("test_refresh".to_string()); - provider.credentials.expiry_date = Some(expires_at.timestamp_millis()); - - let result = provider.validate_token(); - - if expires_in_secs <= -2 { - prop_assert!(is_expired(&result), "Expected Expired for expires_in_secs={}", expires_in_secs); - } else if (2..=TOKEN_EXPIRING_SOON_THRESHOLD - 2).contains(&expires_in_secs) { - prop_assert!(is_expiring_soon(&result), "Expected ExpiringSoon for expires_in_secs={}", expires_in_secs); - } else if expires_in_secs > TOKEN_EXPIRING_SOON_THRESHOLD + 2 { - prop_assert!(is_valid(&result), "Expected Valid for expires_in_secs={}", expires_in_secs); - } - // 边界值 ±2s 内跳过断言,避免时间竞争 - } - - /// Property 1: 对于任何有效的过期时间(timestamp + expires_in 格式),validate_token() 应正确判断状态 - #[test] - fn prop_validate_token_expires_in_format( - expires_in_secs in 1i64..7200i64, // 只测试正数,因为这个格式不支持负数 - ) { - let now = chrono::Utc::now(); - - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - provider.credentials.refresh_token = Some("test_refresh".to_string()); - provider.credentials.timestamp = Some(now.timestamp_millis()); - provider.credentials.expires_in = Some(expires_in_secs); - - let result = provider.validate_token(); - - // 由于时间精度问题,允许 2 秒的误差 - if expires_in_secs <= 1 { - prop_assert!(is_expired(&result) || is_expiring_soon(&result), "Expected Expired or ExpiringSoon for expires_in_secs={}", expires_in_secs); - } else if (2..=TOKEN_EXPIRING_SOON_THRESHOLD - 2).contains(&expires_in_secs) { - prop_assert!(is_expiring_soon(&result), "Expected ExpiringSoon for expires_in_secs={}", expires_in_secs); - } else if expires_in_secs > TOKEN_EXPIRING_SOON_THRESHOLD + 2 { - prop_assert!(is_valid(&result), "Expected Valid for expires_in_secs={}", expires_in_secs); - } - // 边界值 ±2s 内跳过断言,避免时间竞争 - } - } - - // ==================== Property 2: 空 Token 检测 ==================== - // Feature: antigravity-token-refresh, Property 2: 空 Token 检测 - // Validates: Requirements 1.2 - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 2: 对于任何空或仅包含空白字符的 token,validate_token() 应返回 Invalid - #[test] - fn prop_validate_token_empty_detection( - whitespace in "[ \t\n\r]*", // 生成各种空白字符组合 - ) { - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = Some(whitespace); - provider.credentials.refresh_token = Some("test_refresh".to_string()); - - let result = provider.validate_token(); - - prop_assert!(is_invalid(&result), "Expected Invalid for empty/whitespace token"); - } - } - - /// Property 2: None token 应返回 Invalid - #[test] - fn test_validate_token_none() { - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = None; - provider.credentials.refresh_token = Some("test_refresh".to_string()); - - let result = provider.validate_token(); - - assert!( - matches!(result, TokenValidationResult::Invalid { reason } if reason.contains("缺失")) - ); - } - - /// Property 2: 缺少 refresh_token 应返回 Invalid - #[test] - fn test_validate_token_no_refresh_token() { - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - provider.credentials.refresh_token = None; - - let result = provider.validate_token(); - - assert!( - matches!(result, TokenValidationResult::Invalid { reason } if reason.contains("refresh_token")) - ); - } - - /// Property 2: 禁用的凭证应返回 Invalid - #[test] - fn test_validate_token_disabled() { - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - provider.credentials.refresh_token = Some("test_refresh".to_string()); - provider.credentials.enable = Some(false); - - let result = provider.validate_token(); - - assert!( - matches!(result, TokenValidationResult::Invalid { reason } if reason.contains("禁用")) - ); - } - - // ==================== TokenRefreshError 测试 ==================== - - #[test] - fn test_classify_refresh_error_invalid_grant() { - let error = - AntigravityProvider::classify_refresh_error(400, r#"{"error": "invalid_grant"}"#); - assert!(matches!(error, TokenRefreshError::InvalidGrant { .. })); - assert!(error.requires_reauth()); - assert!(!error.is_retryable()); - } - - #[test] - fn test_classify_refresh_error_server_error() { - let error = AntigravityProvider::classify_refresh_error(500, "Internal Server Error"); - assert!(matches!(error, TokenRefreshError::ServerError { .. })); - assert!(!error.requires_reauth()); - assert!(error.is_retryable()); - } - - #[test] - fn test_classify_refresh_error_unknown() { - let error = AntigravityProvider::classify_refresh_error(403, "Forbidden"); - assert!(matches!(error, TokenRefreshError::Unknown { .. })); - assert!(!error.requires_reauth()); - assert!(!error.is_retryable()); - } - - // ==================== TokenValidationResult 方法测试 ==================== - - #[test] - fn test_token_validation_result_needs_refresh() { - assert!(!TokenValidationResult::Valid { - expires_in_secs: 3600 - } - .needs_refresh()); - assert!(TokenValidationResult::ExpiringSoon { - expires_in_secs: 300 - } - .needs_refresh()); - assert!(TokenValidationResult::Expired.needs_refresh()); - assert!(TokenValidationResult::Invalid { - reason: "test".to_string() - } - .needs_refresh()); - } - - #[test] - fn test_token_validation_result_is_usable() { - assert!(TokenValidationResult::Valid { - expires_in_secs: 3600 - } - .is_usable()); - assert!(TokenValidationResult::ExpiringSoon { - expires_in_secs: 300 - } - .is_usable()); - assert!(!TokenValidationResult::Expired.is_usable()); - assert!(!TokenValidationResult::Invalid { - reason: "test".to_string() - } - .is_usable()); - } - - // ==================== Property 5: 重试次数限制 ==================== - // Feature: antigravity-token-refresh, Property 5: 重试次数限制 - // Validates: Requirements 2.2 - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 5: 对于任何 HTTP 状态码,错误分类应正确识别可重试错误 - #[test] - fn prop_classify_error_retryable( - status in 100u16..600u16, - ) { - let error = AntigravityProvider::classify_refresh_error(status, "test error"); - - // 5xx 错误应该是可重试的 - if status >= 500 { - prop_assert!(error.is_retryable(), "5xx errors should be retryable"); - } - - // 400 + invalid_grant 不应该重试 - if status == 400 { - let invalid_grant_error = AntigravityProvider::classify_refresh_error(400, "invalid_grant"); - prop_assert!(!invalid_grant_error.is_retryable(), "invalid_grant should not be retryable"); - prop_assert!(invalid_grant_error.requires_reauth(), "invalid_grant should require reauth"); - } - } - - /// Property 5: 对于任何错误类型,user_message 应返回非空字符串 - #[test] - fn prop_error_user_message_not_empty( - status in 100u16..600u16, - body in ".*", - ) { - let error = AntigravityProvider::classify_refresh_error(status, &body); - let message = error.user_message(); - prop_assert!(!message.is_empty(), "User message should not be empty"); - } - } - - /// 测试 TokenRefreshError 的 Display 实现 - #[test] - fn test_token_refresh_error_display() { - let errors = vec![ - TokenRefreshError::InvalidGrant { - message: "test".to_string(), - }, - TokenRefreshError::NetworkError { - message: "test".to_string(), - }, - TokenRefreshError::ServerError { - message: "test".to_string(), - }, - TokenRefreshError::Unknown { - message: "test".to_string(), - }, - ]; - - for error in errors { - let display = format!("{error}"); - assert!(!display.is_empty()); - } - } - - /// 测试缺少 refresh_token 时 refresh_token_with_retry 应返回 InvalidGrant 错误 - #[tokio::test] - async fn test_refresh_token_with_retry_no_refresh_token() { - let mut provider = AntigravityProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - provider.credentials.refresh_token = None; - - let result = provider.refresh_token_with_retry(3).await; - - assert!(result.is_err()); - let error = result.unwrap_err(); - assert!(error.requires_reauth()); - } - - #[tokio::test] - async fn load_credentials_from_path_prefers_existing_explicit_path() { - let _guard = env_lock().lock().expect("env lock"); - let temp = tempdir().expect("create tempdir"); - let _env = EnvGuard::set(&[("HOME", temp.path().as_os_str().to_os_string())]); - let explicit_path = temp.path().join("explicit.json"); - std::fs::write( - &explicit_path, - r#"{"access_token":"explicit_token","refresh_token":"refresh","project_id":"explicit-project"}"#, - ) - .expect("write explicit creds"); - - let mut provider = AntigravityProvider::new(); - provider - .load_credentials_from_path(explicit_path.to_string_lossy().as_ref()) - .await - .expect("load explicit creds"); - - assert_eq!( - provider.credentials.access_token.as_deref(), - Some("explicit_token") - ); - assert_eq!(provider.project_id.as_deref(), Some("explicit-project")); - } - - #[tokio::test] - async fn load_credentials_from_path_falls_back_to_default_path_when_explicit_missing() { - let _guard = env_lock().lock().expect("env lock"); - let temp = tempdir().expect("create tempdir"); - let _env = EnvGuard::set(&[("HOME", temp.path().as_os_str().to_os_string())]); - let default_path = temp.path().join(".antigravity").join("oauth_creds.json"); - std::fs::create_dir_all(default_path.parent().expect("default parent")) - .expect("create default dir"); - std::fs::write( - &default_path, - r#"{"access_token":"fallback_token","refresh_token":"refresh","project_id":"fallback-project"}"#, - ) - .expect("write fallback creds"); - - let mut provider = AntigravityProvider::new(); - provider - .load_credentials_from_path(temp.path().join("missing.json").to_string_lossy().as_ref()) - .await - .expect("load fallback creds"); - - assert_eq!( - provider.credentials.access_token.as_deref(), - Some("fallback_token") - ); - assert_eq!(provider.project_id.as_deref(), Some("fallback-project")); - } -} diff --git a/src-tauri/crates/providers/src/providers/claude_oauth.rs b/src-tauri/crates/providers/src/providers/claude_oauth.rs deleted file mode 100644 index c534ddf24..000000000 --- a/src-tauri/crates/providers/src/providers/claude_oauth.rs +++ /dev/null @@ -1,964 +0,0 @@ -//! Claude OAuth Provider -//! -//! 实现 Anthropic Claude OAuth 认证流程,与 claude-relay-service 对齐。 -//! -//! ## 支持的授权方式 -//! -//! 1. **标准 OAuth 流程** - 使用官方 redirect_uri,用户需手动复制授权码 -//! 2. **Cookie 自动授权** - 使用 sessionKey 自动完成整个 OAuth 流程 -//! 3. **Setup Token** - 只需推理权限,无 refresh_token -//! - -#![allow(dead_code)] -//! ## 主要功能 -//! -//! - Token 刷新和重试机制 -//! - 统一凭证格式 -//! - 组织信息获取 - -use super::error::{ - create_auth_error, create_config_error, create_token_refresh_error, ProviderError, -}; -use reqwest::Client; -use serde::{Deserialize, Serialize}; -use std::error::Error; -use std::path::PathBuf; - -// OAuth 端点和凭证 - 与 claude-relay-service 完全一致 -const CLAUDE_AUTH_URL: &str = "https://claude.ai/oauth/authorize"; -const CLAUDE_TOKEN_URL: &str = "https://console.anthropic.com/v1/oauth/token"; -const CLAUDE_CLIENT_ID: &str = "9d1c250a-e61b-44d9-88ed-5944d1962f5e"; -// 使用 Anthropic 官方 redirect_uri(用户需手动复制授权码) -const CLAUDE_REDIRECT_URI: &str = "https://console.anthropic.com/oauth/code/callback"; -// OAuth scopes - 与 claude-relay-service 一致 -const CLAUDE_SCOPES: &str = "org:create_api_key user:profile user:inference"; -// Setup Token 只需要推理权限 -const CLAUDE_SCOPES_SETUP: &str = "user:inference"; - -/// Claude OAuth 凭证存储 -/// -/// 与 CLIProxyAPI 的 ClaudeTokenStorage 格式兼容 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ClaudeOAuthCredentials { - /// 访问令牌 - #[serde(default, skip_serializing_if = "Option::is_none")] - pub access_token: Option, - /// 刷新令牌 - #[serde(default, skip_serializing_if = "Option::is_none")] - pub refresh_token: Option, - /// 用户邮箱 - #[serde(default, skip_serializing_if = "Option::is_none")] - pub email: Option, - /// 过期时间(RFC3339 格式) - #[serde(default, skip_serializing_if = "Option::is_none")] - pub expire: Option, - /// 最后刷新时间(RFC3339 格式) - #[serde(default, skip_serializing_if = "Option::is_none")] - pub last_refresh: Option, - /// 凭证类型标识 - #[serde(default = "default_claude_type", rename = "type")] - pub cred_type: String, -} - -fn default_claude_type() -> String { - "claude_oauth".to_string() -} - -impl Default for ClaudeOAuthCredentials { - fn default() -> Self { - Self { - access_token: None, - refresh_token: None, - email: None, - expire: None, - last_refresh: None, - cred_type: default_claude_type(), - } - } -} - -/// PKCE codes for OAuth2 authorization -#[derive(Debug, Clone)] -pub struct PKCECodes { - /// Cryptographically random string for code verification - pub code_verifier: String, - /// SHA256 hash of code_verifier, base64url-encoded - pub code_challenge: String, -} - -impl PKCECodes { - /// Generate new PKCE codes - pub fn generate() -> Result> { - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - use rand::RngCore; - use sha2::{Digest, Sha256}; - - let mut bytes = [0u8; 32]; - rand::thread_rng().fill_bytes(&mut bytes); - let code_verifier = URL_SAFE_NO_PAD.encode(bytes); - - let mut hasher = Sha256::new(); - hasher.update(code_verifier.as_bytes()); - let hash = hasher.finalize(); - let code_challenge = URL_SAFE_NO_PAD.encode(hash); - - Ok(Self { - code_verifier, - code_challenge, - }) - } -} - -/// Claude OAuth Provider -/// -/// 处理 Anthropic Claude 的 OAuth 认证和 API 调用 -pub struct ClaudeOAuthProvider { - /// OAuth 凭证 - pub credentials: ClaudeOAuthCredentials, - /// HTTP 客户端 - pub client: Client, - /// 凭证文件路径 - pub creds_path: Option, -} - -impl Default for ClaudeOAuthProvider { - fn default() -> Self { - Self { - credentials: ClaudeOAuthCredentials::default(), - client: Client::new(), - creds_path: None, - } - } -} - -impl ClaudeOAuthProvider { - /// 创建新的 ClaudeOAuthProvider 实例 - pub fn new() -> Self { - Self::default() - } - - /// 使用自定义 HTTP 客户端创建 - pub fn with_client(client: Client) -> Self { - Self { - client, - ..Self::default() - } - } - - /// 获取默认凭证文件路径 - pub fn default_creds_path() -> PathBuf { - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".claude") - .join("oauth_creds.json") - } - - /// 从默认路径加载凭证 - pub async fn load_credentials(&mut self) -> Result<(), Box> { - let path = Self::default_creds_path(); - self.load_credentials_from_path_internal(&path).await - } - - /// 从指定路径加载凭证 - pub async fn load_credentials_from_path( - &mut self, - path: &str, - ) -> Result<(), Box> { - let path = PathBuf::from(path); - self.load_credentials_from_path_internal(&path).await - } - - async fn load_credentials_from_path_internal( - &mut self, - path: &PathBuf, - ) -> Result<(), Box> { - if tokio::fs::try_exists(&path).await.unwrap_or(false) { - let content = tokio::fs::read_to_string(&path).await?; - let creds: ClaudeOAuthCredentials = serde_json::from_str(&content)?; - tracing::info!( - "[CLAUDE_OAUTH] 凭证已加载: has_access={}, has_refresh={}, email={:?}", - creds.access_token.is_some(), - creds.refresh_token.is_some(), - creds.email - ); - self.credentials = creds; - self.creds_path = Some(path.clone()); - } else { - tracing::warn!("[CLAUDE_OAUTH] 凭证文件不存在: {:?}", path); - } - Ok(()) - } - - /// 保存凭证到文件 - pub async fn save_credentials(&self) -> Result<(), Box> { - let path = self - .creds_path - .clone() - .unwrap_or_else(Self::default_creds_path); - - if let Some(parent) = path.parent() { - tokio::fs::create_dir_all(parent).await?; - } - - let content = serde_json::to_string_pretty(&self.credentials)?; - tokio::fs::write(&path, content).await?; - tracing::info!("[CLAUDE_OAUTH] 凭证已保存到 {:?}", path); - Ok(()) - } - - /// 检查 Token 是否有效 - pub fn is_token_valid(&self) -> bool { - if self.credentials.access_token.is_none() { - return false; - } - - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) { - let now = chrono::Utc::now(); - return expires > now + chrono::Duration::minutes(5); - } - } - - true - } - - /// 刷新 Token - 与 CLIProxyAPI 对齐,使用 JSON 格式 - pub async fn refresh_token(&mut self) -> Result> { - let refresh_token = self - .credentials - .refresh_token - .as_ref() - .ok_or_else(|| create_config_error("没有可用的 refresh_token"))?; - - tracing::info!("[CLAUDE_OAUTH] 正在刷新 Token"); - - // 与 CLIProxyAPI 对齐:使用 JSON 格式请求体 - let body = serde_json::json!({ - "client_id": CLAUDE_CLIENT_ID, - "grant_type": "refresh_token", - "refresh_token": refresh_token - }); - - let resp = self - .client - .post(CLAUDE_TOKEN_URL) - .header("Content-Type", "application/json") - .header("Accept", "application/json") - .json(&body) - .send() - .await - .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; - - if !resp.status().is_success() { - let status = resp.status().as_u16(); - let body = resp.text().await.unwrap_or_default(); - tracing::error!("[CLAUDE_OAUTH] Token 刷新失败: {} - {}", status, body); - self.mark_invalid(); - return Err(create_token_refresh_error(status, &body, "CLAUDE_OAUTH")); - } - - let data: serde_json::Value = resp - .json() - .await - .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; - - let new_access_token = data["access_token"] - .as_str() - .ok_or_else(|| create_auth_error("响应中没有 access_token"))? - .to_string(); - - self.credentials.access_token = Some(new_access_token.clone()); - - if let Some(rt) = data["refresh_token"].as_str() { - self.credentials.refresh_token = Some(rt.to_string()); - } - - // 从响应中提取用户邮箱 - if let Some(email) = data["account"]["email_address"].as_str() { - self.credentials.email = Some(email.to_string()); - } - - // 更新过期时间 - let expires_in = data["expires_in"].as_i64().unwrap_or(3600); - let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in); - self.credentials.expire = Some(expires_at.to_rfc3339()); - self.credentials.last_refresh = Some(chrono::Utc::now().to_rfc3339()); - - self.save_credentials().await?; - - tracing::info!("[CLAUDE_OAUTH] Token 刷新成功"); - Ok(new_access_token) - } - - /// 带重试机制的 Token 刷新 - pub async fn refresh_token_with_retry( - &mut self, - max_retries: u32, - ) -> Result> { - let mut last_error = None; - - for attempt in 0..max_retries { - if attempt > 0 { - let delay = std::time::Duration::from_secs(1 << attempt); - tracing::info!("[CLAUDE_OAUTH] 第 {} 次重试,等待 {:?}", attempt + 1, delay); - tokio::time::sleep(delay).await; - } - - match self.refresh_token().await { - Ok(token) => return Ok(token), - Err(e) => { - tracing::warn!( - "[CLAUDE_OAUTH] Token 刷新第 {} 次尝试失败: {}", - attempt + 1, - e - ); - last_error = Some(e); - } - } - } - - self.mark_invalid(); - tracing::error!("[CLAUDE_OAUTH] Token 刷新在 {} 次尝试后失败", max_retries); - Err(last_error.unwrap_or_else(|| create_auth_error("Token 刷新失败,请重新登录"))) - } - - /// 确保 Token 有效,必要时自动刷新 - pub async fn ensure_valid_token(&mut self) -> Result> { - if !self.is_token_valid() { - tracing::info!("[CLAUDE_OAUTH] Token 需要刷新"); - self.refresh_token_with_retry(3).await - } else { - self.credentials - .access_token - .clone() - .ok_or_else(|| create_config_error("没有可用的 access_token")) - } - } - - /// 标记凭证为无效 - pub fn mark_invalid(&mut self) { - tracing::warn!("[CLAUDE_OAUTH] 标记凭证为无效"); - self.credentials.access_token = None; - self.credentials.expire = None; - } - - /// 获取 OAuth 授权 URL - pub fn get_auth_url(&self) -> &'static str { - CLAUDE_AUTH_URL - } - - /// 获取 OAuth Token URL - pub fn get_token_url(&self) -> &'static str { - CLAUDE_TOKEN_URL - } - - /// 获取 OAuth Client ID - pub fn get_client_id(&self) -> &'static str { - CLAUDE_CLIENT_ID - } - - /// 获取 redirect URI(官方 Anthropic 回调地址) - pub fn get_redirect_uri(&self) -> &'static str { - CLAUDE_REDIRECT_URI - } - - /// 获取 OAuth scopes - pub fn get_scopes(&self) -> &'static str { - CLAUDE_SCOPES - } -} - -// ============================================================================ -// OAuth 登录功能 -// ============================================================================ - -use uuid::Uuid; - -/// OAuth 登录成功后的凭证信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ClaudeOAuthResult { - pub credentials: ClaudeOAuthCredentials, - pub creds_file_path: String, -} - -/// 生成 Claude OAuth 授权 URL(使用官方 redirect_uri) -/// -/// 用户需要: -/// 1. 打开此 URL 进行授权 -/// 2. 授权后从浏览器地址栏复制授权码 -/// 3. 将授权码粘贴回应用 -pub fn generate_claude_auth_url(state: &str, code_challenge: &str) -> String { - let params = [ - ("code", "true"), - ("client_id", CLAUDE_CLIENT_ID), - ("response_type", "code"), - ("redirect_uri", CLAUDE_REDIRECT_URI), - ("scope", CLAUDE_SCOPES), - ("state", state), - ("code_challenge", code_challenge), - ("code_challenge_method", "S256"), - ]; - - let query = params - .iter() - .map(|(k, v)| format!("{}={}", k, urlencoding::encode(v))) - .collect::>() - .join("&"); - - format!("{CLAUDE_AUTH_URL}?{query}") -} - -/// 生成 Setup Token 授权 URL(只需要推理权限) -pub fn generate_claude_setup_token_auth_url(state: &str, code_challenge: &str) -> String { - let params = [ - ("code", "true"), - ("client_id", CLAUDE_CLIENT_ID), - ("response_type", "code"), - ("redirect_uri", CLAUDE_REDIRECT_URI), - ("scope", CLAUDE_SCOPES_SETUP), - ("state", state), - ("code_challenge", code_challenge), - ("code_challenge_method", "S256"), - ]; - - let query = params - .iter() - .map(|(k, v)| format!("{}={}", k, urlencoding::encode(v))) - .collect::>() - .join("&"); - - format!("{CLAUDE_AUTH_URL}?{query}") -} - -/// 用授权码交换 Token(使用官方 redirect_uri) -pub async fn exchange_claude_code_for_token( - client: &Client, - code: &str, - code_verifier: &str, - state: &str, -) -> Result> { - // 清理授权码,移除 URL 片段(与 claude-relay-service 一致) - let cleaned_code = code.split('#').next().unwrap_or(code); - let cleaned_code = cleaned_code.split('&').next().unwrap_or(cleaned_code); - - let body = serde_json::json!({ - "grant_type": "authorization_code", - "client_id": CLAUDE_CLIENT_ID, - "code": cleaned_code, - "redirect_uri": CLAUDE_REDIRECT_URI, - "code_verifier": code_verifier, - "state": state - }); - - tracing::info!( - "[CLAUDE_OAUTH] 正在交换授权码,code 长度: {}", - cleaned_code.len() - ); - - let resp = client - .post(CLAUDE_TOKEN_URL) - .header("Content-Type", "application/json") - .header("Accept", "application/json") - .header("User-Agent", "claude-cli/1.0.56 (external, cli)") - .json(&body) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - tracing::error!("[CLAUDE_OAUTH] Token 交换失败: {} - {}", status, body); - return Err(format!("Token 交换失败: {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - tracing::info!("[CLAUDE_OAUTH] Token 交换成功"); - Ok(data) -} - -/// OAuth 参数(用于手动授权码流程) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ClaudeOAuthParams { - /// 授权 URL - pub auth_url: String, - /// PKCE code_verifier(需要保存用于后续交换 token) - pub code_verifier: String, - /// state 参数 - pub state: String, - /// code_challenge - pub code_challenge: String, -} - -/// 生成 OAuth 授权参数(不启动服务器) -/// -/// 返回授权 URL 和 PKCE 参数,用户需要: -/// 1. 打开 auth_url 进行授权 -/// 2. 授权后从页面复制授权码 -/// 3. 调用 exchange_claude_authorization_code 交换 token -pub fn generate_claude_oauth_params() -> Result> { - let pkce_codes = PKCECodes::generate()?; - let state = Uuid::new_v4().to_string(); - - let auth_url = generate_claude_auth_url(&state, &pkce_codes.code_challenge); - - tracing::info!( - "[CLAUDE_OAUTH] 生成授权参数,state: {}, auth_url: {}", - state, - auth_url - ); - - Ok(ClaudeOAuthParams { - auth_url, - code_verifier: pkce_codes.code_verifier, - state, - code_challenge: pkce_codes.code_challenge, - }) -} - -/// 生成 Setup Token 授权参数 -pub fn generate_claude_setup_token_params( -) -> Result> { - let pkce_codes = PKCECodes::generate()?; - let state = Uuid::new_v4().to_string(); - - let auth_url = generate_claude_setup_token_auth_url(&state, &pkce_codes.code_challenge); - - tracing::info!("[CLAUDE_OAUTH] 生成 Setup Token 授权参数,state: {}", state); - - Ok(ClaudeOAuthParams { - auth_url, - code_verifier: pkce_codes.code_verifier, - state, - code_challenge: pkce_codes.code_challenge, - }) -} - -/// 解析授权码(支持完整 URL 或直接授权码) -pub fn parse_claude_authorization_code( - input: &str, -) -> Result> { - let trimmed = input.trim(); - - // 情况1: 完整 URL - if trimmed.starts_with("http://") || trimmed.starts_with("https://") { - if let Ok(url) = url::Url::parse(trimmed) { - if let Some(code) = url - .query_pairs() - .find(|(k, _)| k == "code") - .map(|(_, v)| v.to_string()) - { - return Ok(code); - } - } - return Err("回调 URL 中未找到授权码 (code 参数)".into()); - } - - // 情况2: 直接授权码(可能包含 URL fragments) - let cleaned = trimmed.split('#').next().unwrap_or(trimmed); - let cleaned = cleaned.split('&').next().unwrap_or(cleaned); - - if cleaned.len() < 10 { - return Err("授权码格式无效,请确保复制了完整的授权码".into()); - } - - Ok(cleaned.to_string()) -} - -/// 使用授权码交换 Token 并保存凭证 -pub async fn exchange_claude_authorization_code( - authorization_code: &str, - code_verifier: &str, - state: &str, -) -> Result> { - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build()?; - - // 解析授权码 - let code = parse_claude_authorization_code(authorization_code)?; - - // 交换 Token - let token_data = exchange_claude_code_for_token(&client, &code, code_verifier, state).await?; - - let access_token = token_data["access_token"].as_str().unwrap_or_default(); - let refresh_token = token_data["refresh_token"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expires_in"].as_i64(); - - // 从响应中提取用户邮箱 - let email = token_data["account"]["email_address"] - .as_str() - .map(|s| s.to_string()); - - // 构建凭证 - let now = chrono::Utc::now(); - let credentials = ClaudeOAuthCredentials { - access_token: Some(access_token.to_string()), - refresh_token, - email: email.clone(), - expire: expires_in.map(|e| (now + chrono::Duration::seconds(e)).to_rfc3339()), - last_refresh: Some(now.to_rfc3339()), - cred_type: "claude_oauth".to_string(), - }; - - // 保存凭证到应用数据目录 - let creds_dir = dirs::data_dir() - .unwrap_or_else(|| std::path::PathBuf::from(".")) - .join("lime") - .join("credentials") - .join("claude_oauth"); - - std::fs::create_dir_all(&creds_dir)?; - - // 生成唯一文件名 - let uuid = Uuid::new_v4().to_string(); - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs(); - let filename = format!("claude_oauth_{}_{}.json", &uuid[..8], timestamp); - let creds_file_path = creds_dir.join(&filename); - - // 保存凭证 - let creds_json = serde_json::to_string_pretty(&credentials)?; - std::fs::write(&creds_file_path, &creds_json)?; - - tracing::info!( - "[CLAUDE_OAUTH] 凭证已保存到: {:?}, email: {:?}", - creds_file_path, - email - ); - - Ok(ClaudeOAuthResult { - credentials, - creds_file_path: creds_file_path.to_string_lossy().to_string(), - }) -} - -/// 启动 Claude OAuth 登录流程(打开浏览器,返回授权参数) -/// -/// 新流程: -/// 1. 生成授权参数 -/// 2. 打开浏览器 -/// 3. 返回参数供后续使用(用户需手动输入授权码) -pub async fn start_claude_oauth_login() -> Result> { - let params = generate_claude_oauth_params()?; - - tracing::info!("[CLAUDE_OAUTH] 打开浏览器进行授权: {}", params.auth_url); - - // 打开浏览器 - if let Err(e) = open::that(¶ms.auth_url) { - tracing::warn!( - "[CLAUDE_OAUTH] 无法打开浏览器: {}. 请手动打开 URL: {}", - e, - params.auth_url - ); - } - - Ok(params) -} - -// ============================================================================ -// Cookie 自动授权功能(参考 claude-relay-service 实现) -// ============================================================================ - -/// Cookie 自动授权配置 -const CLAUDE_AI_URL: &str = "https://claude.ai"; -const CLAUDE_ORGANIZATIONS_URL: &str = "https://claude.ai/api/organizations"; - -/// 组织信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct OrganizationInfo { - pub uuid: String, - pub capabilities: Vec, -} - -/// Cookie 自动授权结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CookieOAuthResult { - pub credentials: ClaudeOAuthCredentials, - pub creds_file_path: String, - pub organization_uuid: Option, - pub capabilities: Vec, -} - -/// 构建带 Cookie 的请求头 -fn build_cookie_headers(session_key: &str) -> reqwest::header::HeaderMap { - let mut headers = reqwest::header::HeaderMap::new(); - headers.insert("Accept", "application/json".parse().unwrap()); - headers.insert("Accept-Language", "en-US,en;q=0.9".parse().unwrap()); - headers.insert("Cache-Control", "no-cache".parse().unwrap()); - headers.insert( - "Cookie", - format!("sessionKey={session_key}").parse().unwrap(), - ); - headers.insert("Origin", CLAUDE_AI_URL.parse().unwrap()); - headers.insert("Referer", format!("{CLAUDE_AI_URL}/new").parse().unwrap()); - headers.insert( - "User-Agent", - "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36" - .parse() - .unwrap(), - ); - headers -} - -/// 使用 Cookie 获取组织信息 -async fn get_organization_info( - client: &Client, - session_key: &str, -) -> Result> { - let headers = build_cookie_headers(session_key); - - tracing::info!("[CLAUDE_OAUTH] 使用 Cookie 获取组织信息"); - - let resp = client - .get(CLAUDE_ORGANIZATIONS_URL) - .headers(headers) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - if status.as_u16() == 403 || status.as_u16() == 401 { - return Err("Cookie 授权失败:无效的 sessionKey 或已过期".into()); - } - if status.as_u16() == 302 { - return Err("请求被 Cloudflare 拦截,请稍后重试".into()); - } - return Err(format!("获取组织信息失败:HTTP {status}").into()); - } - - let data: serde_json::Value = resp.json().await?; - - if !data.is_array() { - return Err("获取组织信息失败:响应格式无效".into()); - } - - let orgs = data.as_array().unwrap(); - - // 找到具有 chat 能力且能力最多的组织 - let mut best_org: Option = None; - let mut max_capabilities = 0; - - for org in orgs { - let capabilities: Vec = org["capabilities"] - .as_array() - .map(|arr| { - arr.iter() - .filter_map(|v| v.as_str().map(|s| s.to_string())) - .collect() - }) - .unwrap_or_default(); - - // 必须有 chat 能力 - if !capabilities.contains(&"chat".to_string()) { - continue; - } - - // 选择能力最多的组织 - if capabilities.len() > max_capabilities { - if let Some(uuid) = org["uuid"].as_str() { - best_org = Some(OrganizationInfo { - uuid: uuid.to_string(), - capabilities: capabilities.clone(), - }); - max_capabilities = capabilities.len(); - } - } - } - - best_org.ok_or_else(|| "未找到具有 chat 能力的组织".into()) -} - -/// 使用 Cookie 自动获取授权码 -async fn authorize_with_cookie( - client: &Client, - session_key: &str, - organization_uuid: &str, - scope: &str, -) -> Result<(String, String, String), Box> { - // 生成 PKCE 参数 - let pkce_codes = PKCECodes::generate()?; - let state = Uuid::new_v4().to_string(); - - // 构建授权 URL - let authorize_url = format!("https://claude.ai/v1/oauth/{organization_uuid}/authorize"); - - // 构建请求 payload - let payload = serde_json::json!({ - "response_type": "code", - "client_id": CLAUDE_CLIENT_ID, - "organization_uuid": organization_uuid, - "redirect_uri": CLAUDE_REDIRECT_URI, - "scope": scope, - "state": state, - "code_challenge": pkce_codes.code_challenge, - "code_challenge_method": "S256" - }); - - let mut headers = build_cookie_headers(session_key); - headers.insert("Content-Type", "application/json".parse().unwrap()); - - tracing::info!("[CLAUDE_OAUTH] 使用 Cookie 请求授权,scope: {}", scope); - - let resp = client - .post(&authorize_url) - .headers(headers) - .json(&payload) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - if status.as_u16() == 403 || status.as_u16() == 401 { - return Err("Cookie 授权失败:无效的 sessionKey 或已过期".into()); - } - if status.as_u16() == 302 { - return Err("请求被 Cloudflare 拦截,请稍后重试".into()); - } - let body = resp.text().await.unwrap_or_default(); - return Err(format!("授权请求失败:HTTP {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - - // 从响应中获取 redirect_uri - let redirect_uri = data["redirect_uri"] - .as_str() - .ok_or("授权响应中未找到 redirect_uri")?; - - tracing::info!( - "[CLAUDE_OAUTH] 获取到 redirect_uri: {}...", - &redirect_uri[..redirect_uri.len().min(80)] - ); - - // 解析 redirect_uri 获取授权码 - let url = url::Url::parse(redirect_uri)?; - let authorization_code = url - .query_pairs() - .find(|(k, _)| k == "code") - .map(|(_, v)| v.to_string()) - .ok_or("redirect_uri 中未找到授权码")?; - - tracing::info!( - "[CLAUDE_OAUTH] 通过 Cookie 获取授权码成功,长度: {}", - authorization_code.len() - ); - - Ok((authorization_code, pkce_codes.code_verifier, state)) -} - -/// 完整的 Cookie 自动授权流程 -/// -/// 参考 claude-relay-service 的 oauthWithCookie 实现 -/// -/// # 参数 -/// - `session_key`: 从浏览器 Cookie 中获取的 sessionKey -/// - `is_setup_token`: 是否为 Setup Token 模式(只需要推理权限) -/// -/// # 返回 -/// - 成功时返回凭证信息和组织信息 -pub async fn oauth_with_cookie( - session_key: &str, - is_setup_token: bool, -) -> Result> { - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .redirect(reqwest::redirect::Policy::none()) // 禁止自动重定向 - .build()?; - - tracing::info!( - "[CLAUDE_OAUTH] 开始 Cookie 自动授权流程,is_setup_token: {}", - is_setup_token - ); - - // 步骤1:获取组织信息 - tracing::info!("[CLAUDE_OAUTH] 步骤 1/3: 获取组织信息..."); - let org_info = get_organization_info(&client, session_key).await?; - tracing::info!( - "[CLAUDE_OAUTH] 找到组织: uuid={}, capabilities={:?}", - org_info.uuid, - org_info.capabilities - ); - - // 步骤2:确定 scope 并获取授权码 - let scope = if is_setup_token { - CLAUDE_SCOPES_SETUP - } else { - "user:profile user:inference" - }; - - tracing::info!("[CLAUDE_OAUTH] 步骤 2/3: 获取授权码..."); - let (authorization_code, code_verifier, state) = - authorize_with_cookie(&client, session_key, &org_info.uuid, scope).await?; - - // 步骤3:交换 Token - tracing::info!("[CLAUDE_OAUTH] 步骤 3/3: 交换 Token..."); - let token_data = - exchange_claude_code_for_token(&client, &authorization_code, &code_verifier, &state) - .await?; - - let access_token = token_data["access_token"].as_str().unwrap_or_default(); - let refresh_token = if is_setup_token { - None // Setup Token 没有 refresh_token - } else { - token_data["refresh_token"].as_str().map(|s| s.to_string()) - }; - let expires_in = token_data["expires_in"].as_i64(); - - // 从响应中提取用户邮箱 - let email = token_data["account"]["email_address"] - .as_str() - .map(|s| s.to_string()); - - // 构建凭证 - let now = chrono::Utc::now(); - let credentials = ClaudeOAuthCredentials { - access_token: Some(access_token.to_string()), - refresh_token, - email: email.clone(), - expire: expires_in.map(|e| (now + chrono::Duration::seconds(e)).to_rfc3339()), - last_refresh: Some(now.to_rfc3339()), - cred_type: if is_setup_token { - "claude_setup_token".to_string() - } else { - "claude_oauth".to_string() - }, - }; - - // 保存凭证到应用数据目录 - let creds_dir = dirs::data_dir() - .unwrap_or_else(|| std::path::PathBuf::from(".")) - .join("lime") - .join("credentials") - .join("claude_oauth"); - - std::fs::create_dir_all(&creds_dir)?; - - // 生成唯一文件名 - let uuid = Uuid::new_v4().to_string(); - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs(); - let token_type = if is_setup_token { "setup" } else { "oauth" }; - let filename = format!("claude_{}_{}_{}.json", token_type, &uuid[..8], timestamp); - let creds_file_path = creds_dir.join(&filename); - - // 保存凭证 - let creds_json = serde_json::to_string_pretty(&credentials)?; - std::fs::write(&creds_file_path, &creds_json)?; - - tracing::info!( - "[CLAUDE_OAUTH] Cookie 自动授权成功,凭证已保存到: {:?}, email: {:?}", - creds_file_path, - email - ); - - Ok(CookieOAuthResult { - credentials, - creds_file_path: creds_file_path.to_string_lossy().to_string(), - organization_uuid: Some(org_info.uuid), - capabilities: org_info.capabilities, - }) -} diff --git a/src-tauri/crates/providers/src/providers/codex.rs b/src-tauri/crates/providers/src/providers/codex.rs index 0bca39a25..66ca5ccf1 100644 --- a/src-tauri/crates/providers/src/providers/codex.rs +++ b/src-tauri/crates/providers/src/providers/codex.rs @@ -1,60 +1,20 @@ -//! OpenAI Codex OAuth Provider +//! Codex API Key Provider //! -//! Implements OAuth authentication flow for OpenAI Codex API. -//! Supports PKCE (Proof Key for Code Exchange) for secure authentication. +//! Codex OAuth、本地 token refresh 与凭证池导入已退役;本模块只保留 +//! API Key Provider 仍使用的 Responses API URL、请求转换和模型识别能力。 #![allow(dead_code)] -use super::error::{ - create_auth_error, create_config_error, create_token_refresh_error, ProviderError, -}; use lime_core::api_host_utils::is_openai_responses_endpoint; use reqwest::Client; use serde::{Deserialize, Serialize}; use std::error::Error; -use std::path::PathBuf; -// OAuth Constants -const OPENAI_AUTH_URL: &str = "https://auth.openai.com/oauth/authorize"; -const OPENAI_TOKEN_URL: &str = "https://auth.openai.com/oauth/token"; -const OPENAI_CLIENT_ID: &str = "app_EMoamEEZ73f0CkXaXp7hrann"; -const DEFAULT_CALLBACK_PORT: u16 = 1455; -const CODEX_API_BASE_URL: &str = "https://chatgpt.com/backend-api/codex"; const DEFAULT_API_BASE_URL: &str = "https://api.openai.com"; -/// Codex OAuth credentials storage -/// -/// Stores OAuth tokens and user information for Codex authentication. -/// Compatible with CLIProxyAPI's CodexTokenStorage format and Codex CLI official format. -/// -/// Supports multiple field name formats: -/// - snake_case: `refresh_token`, `access_token`, `id_token`, `account_id`, `last_refresh` -/// - camelCase: `refreshToken`, `accessToken`, `idToken`, `accountId`, `lastRefresh` -/// -/// 同时兼容 Codex CLI 的 API Key 登录格式: -/// - `api_key` / `apiKey` -/// - `api_base_url` / `apiBaseUrl` #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CodexCredentials { - /// JWT ID token containing user claims - #[serde(default, skip_serializing_if = "Option::is_none", alias = "idToken")] - pub id_token: Option, - /// OAuth2 access token for API access - #[serde( - default, - skip_serializing_if = "Option::is_none", - alias = "accessToken" - )] - pub access_token: Option, - /// Refresh token for obtaining new access tokens - #[serde( - default, - skip_serializing_if = "Option::is_none", - alias = "refreshToken" - )] - pub refresh_token: Option, - /// API Key(Codex CLI 支持通过 API Key 登录) - /// 支持字段名: api_key, apiKey, OPENAI_API_KEY + /// API Key(支持 Codex CLI API Key JSON 字段名,便于现有测试和手工构造复用) #[serde( default, skip_serializing_if = "Option::is_none", @@ -65,31 +25,9 @@ pub struct CodexCredentials { /// API Base URL(可选) #[serde(default, skip_serializing_if = "Option::is_none", alias = "apiBaseUrl")] pub api_base_url: Option, - /// OpenAI account identifier - #[serde(default, skip_serializing_if = "Option::is_none", alias = "accountId")] - pub account_id: Option, - /// Timestamp of last token refresh - #[serde( - default, - skip_serializing_if = "Option::is_none", - alias = "lastRefresh" - )] - pub last_refresh: Option, - /// User email address - #[serde(default, skip_serializing_if = "Option::is_none")] - pub email: Option, - /// Authentication provider type (always "codex") + /// Provider 类型标记 #[serde(default = "default_type")] pub r#type: String, - /// Token expiration timestamp (RFC3339 format) - /// Supports: `expired`, `expires_at`, `expiresAt` - #[serde( - default, - skip_serializing_if = "Option::is_none", - alias = "expired", - alias = "expiresAt" - )] - pub expires_at: Option, } fn default_type() -> String { @@ -99,296 +37,16 @@ fn default_type() -> String { impl Default for CodexCredentials { fn default() -> Self { Self { - id_token: None, - access_token: None, - refresh_token: None, api_key: None, api_base_url: None, - account_id: None, - last_refresh: None, - email: None, r#type: default_type(), - expires_at: None, } } } -/// PKCE codes for OAuth2 authorization -#[derive(Debug, Clone)] -pub struct PKCECodes { - /// Cryptographically random string for code verification - pub code_verifier: String, - /// SHA256 hash of code_verifier, base64url-encoded - pub code_challenge: String, -} - -impl PKCECodes { - /// Generate new PKCE codes - pub fn generate() -> Result> { - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - use rand::RngCore; - use sha2::{Digest, Sha256}; - - // Generate 96 random bytes for code verifier - let mut bytes = [0u8; 96]; - rand::thread_rng().fill_bytes(&mut bytes); - let code_verifier = URL_SAFE_NO_PAD.encode(bytes); - - // Generate code challenge using S256 method - let mut hasher = Sha256::new(); - hasher.update(code_verifier.as_bytes()); - let hash = hasher.finalize(); - let code_challenge = URL_SAFE_NO_PAD.encode(hash); - - Ok(Self { - code_verifier, - code_challenge, - }) - } -} - -/// OAuth callback result -#[derive(Debug, Clone)] -pub struct OAuthCallbackResult { - /// Authorization code from OAuth callback - pub code: String, - /// State parameter for CSRF protection - pub state: String, - /// Error message if authentication failed - pub error: Option, -} - -/// OAuth server for handling OAuth callbacks -pub struct OAuthServer { - port: u16, - shutdown_tx: Option>, -} - -impl OAuthServer { - /// Create a new OAuth server on the specified port - pub fn new(port: u16) -> Self { - Self { - port, - shutdown_tx: None, - } - } - - /// Start the OAuth server and wait for a callback - /// - /// Returns the authorization code and state from the OAuth callback. - /// The server will automatically shut down after receiving a callback or timeout. - pub async fn wait_for_callback( - &mut self, - timeout: std::time::Duration, - ) -> Result> { - use axum::{extract::Query, response::Html, routing::get, Router}; - use std::collections::HashMap; - use tokio::sync::oneshot; - - let (result_tx, result_rx) = oneshot::channel::(); - let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); - self.shutdown_tx = Some(shutdown_tx); - - // Wrap result_tx in Arc for sharing across requests - let result_tx = std::sync::Arc::new(tokio::sync::Mutex::new(Some(result_tx))); - - let result_tx_clone = result_tx.clone(); - let callback_handler = move |Query(params): Query>| { - let result_tx = result_tx_clone.clone(); - async move { - let code = params.get("code").cloned().unwrap_or_default(); - let state = params.get("state").cloned().unwrap_or_default(); - let error = params.get("error").cloned(); - - let result = OAuthCallbackResult { - code, - state, - error: error.clone(), - }; - - // Send result (ignore if already sent) - if let Some(tx) = result_tx.lock().await.take() { - let _ = tx.send(result); - } - - // Return success HTML - if error.is_some() { - Html(OAUTH_ERROR_HTML.to_string()) - } else { - Html(OAUTH_SUCCESS_HTML.to_string()) - } - } - }; - - let app = Router::new().route("/auth/callback", get(callback_handler)); - - let addr = std::net::SocketAddr::from(([127, 0, 0, 1], self.port)); - let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { - if e.kind() == std::io::ErrorKind::AddrInUse { - format!( - "Port {} is already in use. Please close any application using this port.", - self.port - ) - } else { - format!("Failed to bind to port {}: {}", self.port, e) - } - })?; - - tracing::info!( - "[CODEX] OAuth server listening on http://127.0.0.1:{}", - self.port - ); - - // Spawn server with graceful shutdown - let server = axum::serve(listener, app).with_graceful_shutdown(async move { - let _ = shutdown_rx.await; - }); - - tokio::spawn(async move { - if let Err(e) = server.await { - tracing::error!("[CODEX] OAuth server error: {}", e); - } - }); - - // Wait for callback with timeout - let result = tokio::time::timeout(timeout, result_rx).await; - - // Trigger shutdown - if let Some(tx) = self.shutdown_tx.take() { - let _ = tx.send(()); - } - - match result { - Ok(Ok(callback_result)) => { - if let Some(ref error) = callback_result.error { - Err(format!("OAuth error: {error}").into()) - } else { - Ok(callback_result) - } - } - Ok(Err(_)) => Err("OAuth callback channel closed unexpectedly".into()), - Err(_) => { - Err("OAuth callback timeout - no response received within the time limit".into()) - } - } - } -} - -// HTML templates for OAuth callback responses -const OAUTH_SUCCESS_HTML: &str = r#" - - - Authentication Successful - - - -
-
- -
-

Authentication Successful!

-

You can close this window and return to Lime.

-
- -"#; - -const OAUTH_ERROR_HTML: &str = r#" - - - Authentication Failed - - - -
-
- -
-

Authentication Failed

-

Please close this window and try again.

-
- -"#; - -/// Codex OAuth Provider -/// -/// Handles OAuth authentication and API calls for OpenAI Codex. pub struct CodexProvider { - /// OAuth credentials pub credentials: CodexCredentials, - /// HTTP client for API requests pub client: Client, - /// Path to credentials file - pub creds_path: Option, - /// OAuth callback port - pub callback_port: u16, } impl Default for CodexProvider { @@ -396,19 +54,15 @@ impl Default for CodexProvider { Self { credentials: CodexCredentials::default(), client: Client::new(), - creds_path: None, - callback_port: DEFAULT_CALLBACK_PORT, } } } impl CodexProvider { - /// Create a new CodexProvider instance pub fn new() -> Self { Self::default() } - /// Create a new CodexProvider with a custom HTTP client pub fn with_client(client: Client) -> Self { Self { client, @@ -416,46 +70,32 @@ impl CodexProvider { } } - /// Get the default credentials file path - pub fn default_creds_path() -> PathBuf { - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".codex") - .join("auth.json") + pub fn with_api_key(api_key: impl Into, api_base_url: Option) -> Self { + Self { + credentials: CodexCredentials { + api_key: Some(api_key.into()), + api_base_url, + ..Default::default() + }, + client: Client::new(), + } } - /// Get the OAuth authorization URL - pub fn get_auth_url(&self) -> &'static str { - OPENAI_AUTH_URL - } - - /// Get the OAuth token URL - pub fn get_token_url(&self) -> &'static str { - OPENAI_TOKEN_URL - } - - /// Get the OAuth client ID - pub fn get_client_id(&self) -> &'static str { - OPENAI_CLIENT_ID - } - - /// Get the redirect URI for OAuth callback - pub fn get_redirect_uri(&self) -> String { - format!("http://localhost:{}/auth/callback", self.callback_port) - } - - /// Get the API base URL - pub fn get_api_base_url(&self) -> &'static str { - CODEX_API_BASE_URL - } - - /// 获取已配置的 API Key(trim 后的非空值) fn get_api_key(&self) -> Option<&str> { self.credentials .api_key .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()) + .map(str::trim) + .filter(|value| !value.is_empty()) + } + + fn api_base_url(&self) -> &str { + self.credentials + .api_base_url + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(DEFAULT_API_BASE_URL) } pub fn build_responses_url(base_url: &str) -> String { @@ -465,10 +105,9 @@ impl CodexProvider { return base.to_string(); } - // 规则说明: - // - 如果 base_url 以 /v1 结尾:直接拼 /responses - // - 如果 base_url 只有域名(path 为空或 /):拼 /v1/responses(OpenAI 标准) - // - 如果 base_url 已包含路径前缀(如 https://yunyi.cfd/codex):认为前缀已包含路由信息,拼 /responses + // - base_url 以 /v1 结尾:直接拼 /responses + // - base_url 只有域名:拼 /v1/responses(OpenAI 标准) + // - base_url 已包含路径前缀:认为前缀已包含路由信息,拼 /responses if base.ends_with("/v1") { return format!("{base}/responses"); } @@ -481,628 +120,17 @@ impl CodexProvider { return format!("{base}/responses"); } - // 兜底:保持旧行为 format!("{base}/v1/responses") } - /// Load credentials from the default path - pub async fn load_credentials(&mut self) -> Result<(), Box> { - let path = Self::default_creds_path(); - self.load_credentials_from_path_internal(&path).await - } - - /// Load credentials from a specific path - pub async fn load_credentials_from_path( - &mut self, - path: &str, - ) -> Result<(), Box> { - let path = PathBuf::from(path); - self.load_credentials_from_path_internal(&path).await - } - - async fn load_credentials_from_path_internal( - &mut self, - path: &PathBuf, - ) -> Result<(), Box> { - if tokio::fs::try_exists(&path).await.unwrap_or(false) { - let content = tokio::fs::read_to_string(&path).await?; - - // 尝试解析凭证文件 - let creds: CodexCredentials = serde_json::from_str(&content).map_err(|e| { - tracing::error!("[CODEX] 凭证文件解析失败: {}. 文件路径: {:?}", e, path); - format!("凭证文件格式错误: {e}") - })?; - - // 检查关键字段 - let has_api_key = creds - .api_key - .as_deref() - .map(|s| !s.trim().is_empty()) - .unwrap_or(false); - if creds.refresh_token.is_none() && !has_api_key { - tracing::warn!( - "[CODEX] 凭证文件缺少 refresh_token/api_key 字段。支持的字段名: refresh_token, refreshToken, api_key, apiKey" - ); - // 打印文件中的顶级字段名,帮助调试 - if let Ok(json_value) = serde_json::from_str::(&content) { - if let Some(obj) = json_value.as_object() { - let keys: Vec<&String> = obj.keys().collect(); - tracing::info!("[CODEX] 凭证文件包含的字段: {:?}", keys); - } - } - } - - tracing::info!( - "[CODEX] 凭证加载成功: has_access={}, has_refresh={}, has_api_key={}, email={:?}, path={:?}", - creds.access_token.is_some(), - creds.refresh_token.is_some(), - has_api_key, - creds.email, - path - ); - self.credentials = creds; - self.creds_path = Some(path.clone()); - } else { - tracing::warn!("[CODEX] 凭证文件不存在: {:?}", path); - return Err(format!("凭证文件不存在: {path:?}").into()); - } - Ok(()) - } - - /// Save credentials to file - pub async fn save_credentials(&self) -> Result<(), Box> { - let path = self - .creds_path - .clone() - .unwrap_or_else(Self::default_creds_path); - - if let Some(parent) = path.parent() { - tokio::fs::create_dir_all(parent).await?; - } - - let content = serde_json::to_string_pretty(&self.credentials)?; - tokio::fs::write(&path, content).await?; - tracing::info!("[CODEX] Credentials saved to {:?}", path); - Ok(()) - } - - /// Check if the access token is expired - pub fn is_token_expired(&self) -> bool { - // API Key 模式:不涉及过期概念 - if self.get_api_key().is_some() { - return false; - } - - if let Some(expires_str) = &self.credentials.expires_at { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expires_str) { - let now = chrono::Utc::now(); - // Consider expired if less than 5 minutes remaining - return expires < now + chrono::Duration::minutes(5); - } - } - // If no expiry info, assume expired to be safe - true - } - - /// Check if credentials are valid (has access token and not expired) - pub fn is_valid(&self) -> bool { - if self.get_api_key().is_some() { - return true; - } - self.credentials.access_token.is_some() && !self.is_token_expired() - } - - /// Generate the OAuth authorization URL with PKCE - pub fn generate_auth_url( - &self, - state: &str, - pkce_codes: &PKCECodes, - ) -> Result> { - let params = [ - ("client_id", OPENAI_CLIENT_ID), - ("response_type", "code"), - ("redirect_uri", &self.get_redirect_uri()), - // 使用基础 scope,与 CLIProxyAPI 保持一致 - // chatgpt.com/backend-api/codex 端点不需要 api.responses.write 等额外权限 - ("scope", "openid email profile offline_access"), - ("state", state), - ("code_challenge", &pkce_codes.code_challenge), - ("code_challenge_method", "S256"), - ("prompt", "login"), - ("id_token_add_organizations", "true"), - ("codex_cli_simplified_flow", "true"), - ]; - - let query = serde_urlencoded::to_string(params)?; - Ok(format!("{OPENAI_AUTH_URL}?{query}")) - } - - /// Generate a random state string for CSRF protection - pub fn generate_state() -> Result> { - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - use rand::RngCore; - - let mut bytes = [0u8; 32]; - rand::thread_rng().fill_bytes(&mut bytes); - Ok(URL_SAFE_NO_PAD.encode(bytes)) - } - - /// Exchange authorization code for tokens - pub async fn exchange_code_for_tokens( - &mut self, - code: &str, - pkce_codes: &PKCECodes, - ) -> Result<(), Box> { - let params = [ - ("grant_type", "authorization_code"), - ("client_id", OPENAI_CLIENT_ID), - ("code", code), - ("redirect_uri", &self.get_redirect_uri()), - ("code_verifier", &pkce_codes.code_verifier), - ]; - - let resp = self - .client - .post(OPENAI_TOKEN_URL) - .header("Content-Type", "application/x-www-form-urlencoded") - .header("Accept", "application/json") - .form(¶ms) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(format!("Token exchange failed: {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - - // Parse token response - let access_token = data["access_token"] - .as_str() - .ok_or("No access_token in response")? - .to_string(); - let refresh_token = data["refresh_token"].as_str().map(|s| s.to_string()); - let id_token = data["id_token"].as_str().map(|s| s.to_string()); - let expires_in = data["expires_in"].as_i64().unwrap_or(3600); - - // Parse ID token to extract user info - let (account_id, email) = if let Some(ref id_token) = id_token { - parse_jwt_claims(id_token) - } else { - (None, None) - }; - - // Calculate expiration time - let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in); - - self.credentials = CodexCredentials { - id_token, - access_token: Some(access_token), - refresh_token, - api_key: None, - api_base_url: None, - account_id, - last_refresh: Some(chrono::Utc::now().to_rfc3339()), - email, - r#type: "codex".to_string(), - expires_at: Some(expires_at.to_rfc3339()), - }; - - // Save credentials - self.save_credentials().await?; - - tracing::info!( - "[CODEX] Token exchange successful, email={:?}", - self.credentials.email - ); - Ok(()) - } - - /// Refresh the access token using the refresh token - /// - /// Supports three authentication modes (in priority order): - /// 1. **API Key Mode**: Returns the API key directly (no refresh needed) - /// 2. **OAuth Mode**: Refreshes the access token using the refresh token - /// 3. **Access Token Mode**: Returns the existing access token (may be expired) - /// - /// # Returns - /// * `Ok(String)` - The access token or API key - /// * `Err` - If no credentials are available - /// - /// # Examples - /// ```ignore - /// // API Key mode - /// provider.credentials.api_key = Some("sk-test".to_string()); - /// let token = provider.refresh_token().await?; // Returns "sk-test" - /// - /// // OAuth mode - /// provider.credentials.refresh_token = Some("refresh_token".to_string()); - /// let token = provider.refresh_token().await?; // Refreshes and returns new access_token - /// - /// // Access Token mode (fallback) - /// provider.credentials.access_token = Some("access_token".to_string()); - /// let token = provider.refresh_token().await?; // Returns "access_token" (with warning) - /// ``` - pub async fn refresh_token(&mut self) -> Result> { - // 1. API Key 模式无需刷新(优先级最高) - if let Some(api_key) = self.get_api_key() { - return Ok(api_key.to_string()); - } - - // 2. 无 refresh_token 时的降级处理 - if self.credentials.refresh_token.is_none() { - // 2a. 有 access_token:返回(可能过期,由上层处理) - if let Some(ref access_token) = self.credentials.access_token { - tracing::warn!("[CODEX] 没有 refresh_token,返回现有 access_token(可能已过期)"); - return Ok(access_token.clone()); - } - - // 2b. 无任何凭证:清晰的错误指导 - return Err(create_config_error( - "没有可用的认证凭证。请配置以下任一方式:\n\ - 1. API Key 模式:在凭证文件中添加 api_key/apiKey 字段\n\ - 2. OAuth 模式:使用 OAuth 登录获取 refresh_token\n\ - 3. Access Token 模式:在凭证文件中添加 access_token/accessToken 字段", - )); - } - - // 3. OAuth 刷新流程(标准流程) - let refresh_token = - self.credentials.refresh_token.as_ref().ok_or_else(|| { - create_config_error("OAuth 刷新令牌不可用 (refresh_token is None)") - })?; - - tracing::info!("[CODEX] 正在刷新 access token"); - - let params = [ - ("client_id", OPENAI_CLIENT_ID), - ("grant_type", "refresh_token"), - ("refresh_token", refresh_token.as_str()), - // 使用基础 scope,与 CLIProxyAPI 保持一致 - ("scope", "openid profile email"), - ]; - - let resp = self - .client - .post(OPENAI_TOKEN_URL) - .header("Content-Type", "application/x-www-form-urlencoded") - .header("Accept", "application/json") - .form(¶ms) - .send() - .await - .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; - - if !resp.status().is_success() { - let status = resp.status().as_u16(); - let body = resp.text().await.unwrap_or_default(); - tracing::error!("[CODEX] Token refresh failed: {} - {}", status, body); - - // Mark credentials as invalid on refresh failure - self.mark_invalid(); - - return Err(create_token_refresh_error(status, &body, "CODEX")); - } - - let data: serde_json::Value = resp - .json() - .await - .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; - - // Update credentials - let new_access_token = data["access_token"] - .as_str() - .ok_or_else(|| create_auth_error("响应中没有 access_token"))? - .to_string(); - - self.credentials.access_token = Some(new_access_token.clone()); - - if let Some(rt) = data["refresh_token"].as_str() { - self.credentials.refresh_token = Some(rt.to_string()); - } - - if let Some(id_token) = data["id_token"].as_str() { - self.credentials.id_token = Some(id_token.to_string()); - let (account_id, email) = parse_jwt_claims(id_token); - if account_id.is_some() { - self.credentials.account_id = account_id; - } - if email.is_some() { - self.credentials.email = email; - } - } - - let expires_in = data["expires_in"].as_i64().unwrap_or(3600); - let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in); - self.credentials.expires_at = Some(expires_at.to_rfc3339()); - self.credentials.last_refresh = Some(chrono::Utc::now().to_rfc3339()); - - // Save updated credentials - self.save_credentials().await?; - - tracing::info!("[CODEX] Token refresh successful"); - Ok(new_access_token) - } - - /// Refresh token with retry mechanism - /// - /// Attempts to refresh the token up to `max_retries` times with linear backoff (1s, 2s, 3s). - /// Marks credentials as invalid if all retries fail. - /// - /// # Arguments - /// * `max_retries` - Maximum number of retry attempts (typically 3) - /// - /// # Returns - /// * `Ok(String)` - The new access token on success - /// * `Err` - Error if all retries fail - pub async fn refresh_token_with_retry( - &mut self, - max_retries: u32, - ) -> Result> { - let mut last_error = None; - - for attempt in 0..max_retries { - if attempt > 0 { - // Linear backoff: 1s, 2s, 3s, ... (as per Requirements 8.2) - let delay = std::time::Duration::from_secs((attempt) as u64); - tracing::info!( - "[CODEX] Retry attempt {}/{} after {:?}", - attempt + 1, - max_retries, - delay - ); - tokio::time::sleep(delay).await; - } - - match self.refresh_token().await { - Ok(token) => { - if attempt > 0 { - tracing::info!( - "[CODEX] Token refresh succeeded on attempt {}", - attempt + 1 - ); - } - return Ok(token); - } - Err(e) => { - tracing::warn!( - "[CODEX] Token refresh attempt {}/{} failed: {}", - attempt + 1, - max_retries, - e - ); - last_error = Some(e); - } - } - } - - // All retries failed - mark as invalid - self.mark_invalid(); - tracing::error!( - "[CODEX] Token refresh failed after {} attempts", - max_retries - ); - - Err(last_error.unwrap_or_else(|| create_auth_error("Token 刷新失败,请重新登录"))) - } - - /// Check if token needs refresh (expiring within the specified duration) - pub fn needs_refresh(&self, lead_time: chrono::Duration) -> bool { - // API Key 模式无需刷新 - if self.get_api_key().is_some() { - return false; - } - - if self.credentials.access_token.is_none() { - return true; - } - - if let Some(expires_str) = &self.credentials.expires_at { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expires_str) { - let now = chrono::Utc::now(); - return expires < now + lead_time; - } - } - - // If no expiry info, assume needs refresh - true - } - - /// Ensure token is valid, refreshing if necessary - /// - /// This is the recommended method to call before making API requests. - /// It will automatically refresh the token if it's expired or about to expire. - pub async fn ensure_valid_token(&mut self) -> Result> { - // 兼容 Codex CLI 的 API Key 登录:auth.json 只有 api_key,没有 refresh_token - if let Some(api_key) = self.get_api_key() { - return Ok(api_key.to_string()); - } - - // Refresh if token expires within 5 minutes - let lead_time = chrono::Duration::minutes(5); - - if self.needs_refresh(lead_time) { - tracing::info!("[CODEX] Token needs refresh, attempting refresh with retry"); - self.refresh_token_with_retry(3).await - } else { - self.credentials - .access_token - .clone() - .ok_or_else(|| create_config_error("没有可用的 access_token")) - } - } - - /// Mark credentials as invalid (e.g., after refresh failure) - pub fn mark_invalid(&mut self) { - tracing::warn!("[CODEX] Marking credentials as invalid"); - self.credentials.access_token = None; - self.credentials.expires_at = None; - } - - /// Get the access token, refreshing if necessary - pub async fn get_access_token(&mut self) -> Result> { - // API Key 模式直接返回 - if let Some(api_key) = self.get_api_key() { - return Ok(api_key.to_string()); - } - - if self.is_token_expired() { - self.refresh_token().await?; - } - self.credentials - .access_token - .clone() - .ok_or_else(|| create_config_error("没有可用的 access_token")) - } - - /// Perform OAuth login flow - /// - /// Opens a browser for OAuth authentication and waits for the callback. - /// Returns the email of the authenticated user on success. - pub async fn oauth_login(&mut self) -> Result> { - tracing::info!("[CODEX] Starting OAuth login flow"); - - // Generate PKCE codes and state - let pkce_codes = PKCECodes::generate()?; - let state = Self::generate_state()?; - - // Generate authorization URL - let auth_url = self.generate_auth_url(&state, &pkce_codes)?; - - // Start OAuth server - let mut oauth_server = OAuthServer::new(self.callback_port); - - // Open browser - tracing::info!("[CODEX] Opening browser for authentication"); - if let Err(e) = open::that(&auth_url) { - tracing::warn!( - "[CODEX] Failed to open browser: {}. Please open the URL manually.", - e - ); - println!("Please open the following URL in your browser:\n{auth_url}"); - } - - // Wait for callback (5 minute timeout) - let timeout = std::time::Duration::from_secs(300); - let callback_result = oauth_server.wait_for_callback(timeout).await?; - - // Verify state - if callback_result.state != state { - return Err("OAuth state mismatch - possible CSRF attack".into()); - } - - // Exchange code for tokens - self.exchange_code_for_tokens(&callback_result.code, &pkce_codes) - .await?; - - let email = self - .credentials - .email - .clone() - .unwrap_or_else(|| "unknown".to_string()); - tracing::info!("[CODEX] OAuth login successful for {}", email); - - Ok(email) - } - - /// Perform OAuth login without opening browser (for headless/SSH environments) - /// - /// Returns the authorization URL that the user should open manually. - pub fn start_oauth_login( - &self, - ) -> Result<(String, PKCECodes, String), Box> { - let pkce_codes = PKCECodes::generate()?; - let state = Self::generate_state()?; - let auth_url = self.generate_auth_url(&state, &pkce_codes)?; - Ok((auth_url, pkce_codes, state)) - } - - /// Complete OAuth login after receiving callback - pub async fn complete_oauth_login( - &mut self, - code: &str, - pkce_codes: &PKCECodes, - expected_state: &str, - received_state: &str, - ) -> Result> { - // Verify state - if received_state != expected_state { - return Err("OAuth state mismatch - possible CSRF attack".into()); - } - - // Exchange code for tokens - self.exchange_code_for_tokens(code, pkce_codes).await?; - - let email = self - .credentials - .email - .clone() - .unwrap_or_else(|| "unknown".to_string()); - tracing::info!("[CODEX] OAuth login completed for {}", email); - - Ok(email) - } - - /// Call the Codex API for chat completions - /// - /// Routes GPT model requests through the Codex OAuth endpoint. - /// The request should be in OpenAI chat completion format. - pub async fn call_api( + pub(crate) async fn call_api( &self, request: &serde_json::Value, ) -> Result> { - enum AuthMode { - ApiKey, - OAuth, - } - - let (token, mode) = match self.get_api_key() { - Some(api_key) => (api_key, AuthMode::ApiKey), - None => ( - self.credentials - .access_token - .as_deref() - .ok_or("No access token or api_key available")?, - AuthMode::OAuth, - ), - }; - - // Build the Codex API URL - let url = match mode { - AuthMode::ApiKey => { - let has_custom_base_url = self - .credentials - .api_base_url - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()) - .is_some(); - - let base_url = self - .credentials - .api_base_url - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()) - .unwrap_or(DEFAULT_API_BASE_URL); - - // Warn if API key doesn't look like OpenAI format but no custom base URL is set - if !has_custom_base_url && !token.starts_with("sk-") { - tracing::warn!( - "[CODEX] API key does not appear to be an OpenAI key (doesn't start with 'sk-'), \ - but no api_base_url is configured. Requests will be sent to {}. \ - If you're using a third-party API provider, please add 'api_base_url' to ~/.codex/auth.json", - DEFAULT_API_BASE_URL - ); - } - - Self::build_responses_url(base_url) - } - AuthMode::OAuth => format!("{CODEX_API_BASE_URL}/responses"), - }; - - // Transform OpenAI chat completion request to Codex format + let token = self + .get_api_key() + .ok_or("Codex API Key 未配置,请通过 API Key Provider 配置。")?; + let url = Self::build_responses_url(self.api_base_url()); let codex_request = transform_to_codex_format(request)?; let mut req = self @@ -1115,18 +143,13 @@ impl CodexProvider { .header("Openai-Beta", "responses=experimental") .json(&codex_request); - // 部分三方 Codex 代理(如 Yunyi)会依赖 Codex CLI 的特征 headers; - // 仅在 OAuth 模式或显式配置了自定义 base_url 时附加,避免影响 OpenAI 官方 Key 模式。 - let should_add_codex_cli_headers = matches!(mode, AuthMode::OAuth) - || (matches!(mode, AuthMode::ApiKey) - && self - .credentials - .api_base_url - .as_deref() - .map(|s| !s.trim().is_empty()) - .unwrap_or(false)); - - if should_add_codex_cli_headers { + if self + .credentials + .api_base_url + .as_deref() + .map(|value| !value.trim().is_empty()) + .unwrap_or(false) + { req = req .header("Version", "0.21.0") .header( @@ -1136,31 +159,18 @@ impl CodexProvider { .header("Originator", "codex_cli_rs") .header("Session_id", uuid::Uuid::new_v4().to_string()) .header("Conversation_id", uuid::Uuid::new_v4().to_string()); - - // 仅在 account_id 非空时添加 Chatgpt-Account-Id header - // 参考 CLIProxyAPI: 空值时不发送此 header - if let Some(account_id) = self.credentials.account_id.as_deref() { - if !account_id.trim().is_empty() { - req = req.header("Chatgpt-Account-Id", account_id); - } - } } - let resp = req.send().await?; - - Ok(resp) + Ok(req.send().await?) } - /// Call the Codex API with streaming response - pub async fn call_api_stream( + pub(crate) async fn call_api_stream( &self, request: &serde_json::Value, ) -> Result> { - // Same as call_api - Codex always returns SSE stream self.call_api(request).await } - /// Check if this provider supports the given model pub fn supports_model(model: &str) -> bool { let model_lower = model.to_lowercase(); model_lower.starts_with("gpt-") @@ -1171,69 +181,6 @@ impl CodexProvider { } } -/// Parse JWT token to extract account_id and email -/// -/// Extracts user information from the JWT ID token returned by OpenAI OAuth. -/// The account_id is extracted from the `chatgpt_account_id` field in the -/// `https://api.openai.com/auth` claim, which is required for Codex API calls. -fn parse_jwt_claims(token: &str) -> (Option, Option) { - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - - let parts: Vec<&str> = token.split('.').collect(); - if parts.len() != 3 { - tracing::warn!( - "[CODEX] Invalid JWT token format: expected 3 parts, got {}", - parts.len() - ); - return (None, None); - } - - // Decode payload (second part) - JWT uses URL-safe base64 without padding - let payload = match URL_SAFE_NO_PAD.decode(parts[1]) { - Ok(bytes) => bytes, - Err(_) => { - // Try with padding added - let padded = format!("{}{}", parts[1], "=".repeat((4 - parts[1].len() % 4) % 4)); - match base64::engine::general_purpose::URL_SAFE.decode(&padded) { - Ok(bytes) => bytes, - Err(e) => { - tracing::warn!("[CODEX] Failed to decode JWT payload: {}", e); - return (None, None); - } - } - } - }; - - let claims: serde_json::Value = match serde_json::from_slice(&payload) { - Ok(v) => v, - Err(e) => { - tracing::warn!("[CODEX] Failed to parse JWT claims: {}", e); - return (None, None); - } - }; - - // Extract email from standard claim - let email = claims["email"].as_str().map(|s| s.to_string()); - - // Extract account_id from OpenAI-specific claims - // Priority: chatgpt_account_id > user_id > sub - // The chatgpt_account_id is the correct field for Codex API calls - let auth_info = &claims["https://api.openai.com/auth"]; - let account_id = auth_info["chatgpt_account_id"] - .as_str() - .or_else(|| auth_info["user_id"].as_str()) - .or_else(|| claims["sub"].as_str()) - .map(|s| s.to_string()); - - tracing::debug!( - "[CODEX] JWT parsed: email={:?}, account_id={:?}", - email, - account_id - ); - - (account_id, email) -} - fn extract_text_fragments(content: &serde_json::Value) -> Vec { if let Some(text) = content .as_str() @@ -1287,16 +234,13 @@ fn resolve_codex_instructions( ]) } -/// Transform OpenAI chat completion request to Codex format -/// 参考 CLIProxyAPI: internal/translator/codex/openai/chat-completions/codex_openai_request.go +/// Transform OpenAI chat completion request to Codex Responses format. fn transform_to_codex_format( request: &serde_json::Value, ) -> Result> { let model = request["model"].as_str().unwrap_or("gpt-4o"); let messages = request["messages"].as_array(); - // 注意:stream 参数被忽略,Codex API 强制要求 stream = true - // Build input array from messages let mut input = Vec::new(); let mut system_instructions: Vec = Vec::new(); @@ -1307,7 +251,6 @@ fn transform_to_codex_format( match role { "system" => { - // system 消息统一转换到 instructions,避免污染用户输入。 system_instructions.extend(extract_text_fragments(content)); } "user" => { @@ -1343,7 +286,6 @@ fn transform_to_codex_format( } } "assistant" => { - // Assistant message content let content_parts = if let Some(text) = content.as_str() { vec![serde_json::json!({"type": "output_text", "text": text})] } else if let Some(arr) = content.as_array() { @@ -1366,7 +308,6 @@ fn transform_to_codex_format( })); } - // Handle tool calls for assistant messages if let Some(tool_calls) = msg["tool_calls"].as_array() { for tc in tool_calls { if tc["type"].as_str() == Some("function") { @@ -1381,7 +322,6 @@ fn transform_to_codex_format( } } "tool" => { - // Tool results let tool_call_id = msg["tool_call_id"].as_str().unwrap_or(""); let output = content.as_str().unwrap_or(""); input.push(serde_json::json!({ @@ -1395,13 +335,10 @@ fn transform_to_codex_format( } } - // Build the Codex request - 参考 CLIProxyAPI 的必需字段 - // 注意:Codex API 要求 stream 必须为 true - // 参考 CLIProxyAPI: internal/runtime/executor/codex_executor.go 第 107 行 let mut codex_request = serde_json::json!({ "model": model, "input": input, - "stream": true, // Codex API 强制要求 stream = true + "stream": true, "store": false, "parallel_tool_calls": true, "reasoning": { @@ -1415,7 +352,6 @@ fn transform_to_codex_format( codex_request["instructions"] = serde_json::json!(instructions); } - // 处理可选参数:temperature, max_tokens (-> max_output_tokens), top_p if let Some(temp) = request.get("temperature") { codex_request["temperature"] = temp.clone(); } @@ -1426,7 +362,6 @@ fn transform_to_codex_format( codex_request["top_p"] = top_p.clone(); } - // Build tools array if present if let Some(tools) = request["tools"].as_array() { let codex_tools: Vec = tools .iter() @@ -1441,7 +376,6 @@ fn transform_to_codex_format( "parameters": func["parameters"] })) } else if !tool_type.is_empty() { - // Pass through built-in tools directly Some(tool.clone()) } else { None @@ -1454,7 +388,6 @@ fn transform_to_codex_format( } } - // Handle tool_choice if let Some(tool_choice) = request.get("tool_choice") { if let Some(tc_str) = tool_choice.as_str() { codex_request["tool_choice"] = serde_json::json!(tc_str); @@ -1471,12 +404,10 @@ fn transform_to_codex_format( } } - // Handle reasoning effort if let Some(reasoning_effort) = request["reasoning_effort"].as_str() { codex_request["reasoning"]["effort"] = serde_json::json!(reasoning_effort); } - // Handle response_format for Structured Outputs if let Some(rf) = request.get("response_format") { let rf_type = rf["type"].as_str().unwrap_or(""); match rf_type { @@ -1518,78 +449,11 @@ mod tests { #[test] fn test_codex_credentials_default() { let creds = CodexCredentials::default(); - assert!(creds.access_token.is_none()); - assert!(creds.refresh_token.is_none()); assert!(creds.api_key.is_none()); + assert!(creds.api_base_url.is_none()); assert_eq!(creds.r#type, "codex"); } - #[test] - fn test_codex_credentials_serialization() { - let creds = CodexCredentials { - access_token: Some("test_token".to_string()), - refresh_token: Some("test_refresh".to_string()), - email: Some("test@example.com".to_string()), - ..Default::default() - }; - - let json = serde_json::to_string(&creds).unwrap(); - assert!(json.contains("test_token")); - assert!(json.contains("test@example.com")); - - let parsed: CodexCredentials = serde_json::from_str(&json).unwrap(); - assert_eq!(parsed.access_token, creds.access_token); - assert_eq!(parsed.email, creds.email); - } - - #[test] - fn test_codex_credentials_camel_case_alias() { - // 测试 camelCase 字段名的支持(Codex CLI 官方格式) - let json = r#"{ - "idToken": "test_id_token", - "accessToken": "test_access_token", - "refreshToken": "test_refresh_token", - "accountId": "test_account_id", - "lastRefresh": "2024-01-01T00:00:00Z", - "email": "test@example.com", - "type": "codex", - "expiresAt": "2024-12-31T23:59:59Z" - }"#; - - let creds: CodexCredentials = serde_json::from_str(json).unwrap(); - assert_eq!(creds.id_token, Some("test_id_token".to_string())); - assert_eq!(creds.access_token, Some("test_access_token".to_string())); - assert_eq!(creds.refresh_token, Some("test_refresh_token".to_string())); - assert_eq!(creds.account_id, Some("test_account_id".to_string())); - assert_eq!(creds.last_refresh, Some("2024-01-01T00:00:00Z".to_string())); - assert_eq!(creds.email, Some("test@example.com".to_string())); - assert_eq!(creds.expires_at, Some("2024-12-31T23:59:59Z".to_string())); - } - - #[test] - fn test_codex_credentials_snake_case() { - // 测试 snake_case 字段名的支持(CLIProxyAPI 格式) - let json = r#"{ - "id_token": "test_id_token", - "access_token": "test_access_token", - "refresh_token": "test_refresh_token", - "account_id": "test_account_id", - "last_refresh": "2024-01-01T00:00:00Z", - "email": "test@example.com", - "type": "codex", - "expired": "2024-12-31T23:59:59Z" - }"#; - - let creds: CodexCredentials = serde_json::from_str(json).unwrap(); - assert_eq!(creds.id_token, Some("test_id_token".to_string())); - assert_eq!(creds.access_token, Some("test_access_token".to_string())); - assert_eq!(creds.refresh_token, Some("test_refresh_token".to_string())); - assert_eq!(creds.account_id, Some("test_account_id".to_string())); - assert_eq!(creds.last_refresh, Some("2024-01-01T00:00:00Z".to_string())); - assert_eq!(creds.email, Some("test@example.com".to_string())); - assert_eq!(creds.expires_at, Some("2024-12-31T23:59:59Z".to_string())); - } - #[test] fn test_codex_credentials_api_key_fields() { let json = r#"{ @@ -1616,36 +480,11 @@ mod tests { ); } - #[test] - fn test_codex_credentials_expires_at_alias() { - // 测试 expires_at 字段的多种别名 - let json1 = r#"{"expired": "2024-12-31T23:59:59Z"}"#; - let json2 = r#"{"expires_at": "2024-12-31T23:59:59Z"}"#; - let json3 = r#"{"expiresAt": "2024-12-31T23:59:59Z"}"#; - - let creds1: CodexCredentials = serde_json::from_str(json1).unwrap(); - let creds2: CodexCredentials = serde_json::from_str(json2).unwrap(); - let creds3: CodexCredentials = serde_json::from_str(json3).unwrap(); - - assert_eq!(creds1.expires_at, Some("2024-12-31T23:59:59Z".to_string())); - assert_eq!(creds2.expires_at, Some("2024-12-31T23:59:59Z".to_string())); - assert_eq!(creds3.expires_at, Some("2024-12-31T23:59:59Z".to_string())); - } - - #[test] - fn test_pkce_generation() { - let pkce = PKCECodes::generate().unwrap(); - assert!(!pkce.code_verifier.is_empty()); - assert!(!pkce.code_challenge.is_empty()); - // Verifier should be 128 chars (96 bytes base64 encoded) - assert_eq!(pkce.code_verifier.len(), 128); - } - #[test] fn test_codex_provider_default() { let provider = CodexProvider::new(); - assert_eq!(provider.callback_port, DEFAULT_CALLBACK_PORT); - assert!(provider.credentials.access_token.is_none()); + assert!(provider.credentials.api_key.is_none()); + assert_eq!(provider.api_base_url(), DEFAULT_API_BASE_URL); } #[test] @@ -1672,126 +511,18 @@ mod tests { ); } - #[tokio::test] - async fn test_ensure_valid_token_prefers_api_key() { - let mut provider = CodexProvider::new(); - provider.credentials.api_key = Some("sk-test".to_string()); - - let token = provider.ensure_valid_token().await.unwrap(); - assert_eq!(token, "sk-test"); - } - - #[test] - fn test_generate_auth_url() { - let provider = CodexProvider::new(); - let pkce = PKCECodes::generate().unwrap(); - let state = "test_state"; - - let url = provider.generate_auth_url(state, &pkce).unwrap(); - assert!(url.starts_with(OPENAI_AUTH_URL)); - assert!(url.contains("client_id=")); - assert!(url.contains("code_challenge=")); - assert!(url.contains("state=test_state")); - } - - #[test] - fn test_parse_jwt_claims_with_sub() { - // Mock JWT with only sub claim (fallback case) - // Header: {"alg":"RS256","typ":"JWT"} - // Payload: {"email":"test@example.com","sub":"user123"} - let mock_jwt = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJlbWFpbCI6InRlc3RAZXhhbXBsZS5jb20iLCJzdWIiOiJ1c2VyMTIzIn0.signature"; - - let (account_id, email) = parse_jwt_claims(mock_jwt); - assert_eq!(email, Some("test@example.com".to_string())); - assert_eq!(account_id, Some("user123".to_string())); - } - - #[test] - fn test_parse_jwt_claims_with_chatgpt_account_id() { - // Mock JWT with chatgpt_account_id in https://api.openai.com/auth claim - // This is the preferred field for Codex API calls - // Payload: {"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"chatgpt_account_id":"chatgpt_acc_123","user_id":"uid_456"}} - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - - let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#); - let payload = URL_SAFE_NO_PAD.encode(r#"{"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"chatgpt_account_id":"chatgpt_acc_123","user_id":"uid_456"}}"#); - let mock_jwt = format!("{header}.{payload}.signature"); - - let (account_id, email) = parse_jwt_claims(&mock_jwt); - assert_eq!(email, Some("test@example.com".to_string())); - // Should prefer chatgpt_account_id over user_id and sub - assert_eq!(account_id, Some("chatgpt_acc_123".to_string())); - } - - #[test] - fn test_parse_jwt_claims_with_user_id() { - // Mock JWT with user_id but no chatgpt_account_id - // Payload: {"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"user_id":"uid_456"}} - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - - let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#); - let payload = URL_SAFE_NO_PAD.encode(r#"{"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"user_id":"uid_456"}}"#); - let mock_jwt = format!("{header}.{payload}.signature"); - - let (account_id, email) = parse_jwt_claims(&mock_jwt); - assert_eq!(email, Some("test@example.com".to_string())); - // Should use user_id when chatgpt_account_id is not present - assert_eq!(account_id, Some("uid_456".to_string())); - } - - #[test] - fn test_parse_jwt_claims_invalid_token() { - // Invalid JWT format - let (account_id, email) = parse_jwt_claims("invalid.token"); - assert_eq!(account_id, None); - assert_eq!(email, None); - - // Empty token - let (account_id, email) = parse_jwt_claims(""); - assert_eq!(account_id, None); - assert_eq!(email, None); - } - - #[test] - fn test_is_token_expired() { - let mut provider = CodexProvider::new(); - - // No expiry - should be considered expired - assert!(provider.is_token_expired()); - - // API Key 模式 - 不应视为过期 - provider.credentials.api_key = Some("sk-test".to_string()); - assert!(!provider.is_token_expired()); - - // Expired token - provider.credentials.api_key = None; - provider.credentials.expires_at = Some("2020-01-01T00:00:00Z".to_string()); - assert!(provider.is_token_expired()); - - // Valid token (far future) - provider.credentials.expires_at = Some("2099-01-01T00:00:00Z".to_string()); - assert!(!provider.is_token_expired()); - } - #[test] fn test_supports_model() { - // GPT models should be supported assert!(CodexProvider::supports_model("gpt-4")); assert!(CodexProvider::supports_model("gpt-4o")); assert!(CodexProvider::supports_model("gpt-4-turbo")); - assert!(CodexProvider::supports_model("GPT-4")); // Case insensitive - - // O-series models should be supported + assert!(CodexProvider::supports_model("GPT-4")); assert!(CodexProvider::supports_model("o1")); assert!(CodexProvider::supports_model("o1-preview")); assert!(CodexProvider::supports_model("o3")); assert!(CodexProvider::supports_model("o4-mini")); - - // Codex models should be supported (contains "codex") assert!(CodexProvider::supports_model("codex-mini")); assert!(CodexProvider::supports_model("gpt-4-codex")); - - // Non-GPT models should not be supported assert!(!CodexProvider::supports_model("claude-3")); assert!(!CodexProvider::supports_model("gemini-pro")); assert!(!CodexProvider::supports_model("llama-2")); @@ -1812,13 +543,12 @@ mod tests { assert_eq!(result["model"], "gpt-4o"); assert_eq!(result["stream"], true); - // system 消息应映射到 instructions - assert!(result.get("instructions").is_some()); - let instructions = result["instructions"].as_str().unwrap(); - assert_eq!(instructions, "You are a helpful assistant."); + assert_eq!( + result["instructions"].as_str(), + Some("You are a helpful assistant.") + ); let input = result["input"].as_array().unwrap(); - // system message 不应污染输入,保留 user 消息即可 assert_eq!(input.len(), 1); assert_eq!(input[0]["role"], "user"); } @@ -1896,501 +626,4 @@ mod tests { assert_eq!(result["max_output_tokens"], 1000); assert_eq!(result["top_p"], 0.9); } - - #[tokio::test] - async fn test_refresh_token_with_only_access_token() { - // 场景:只有 access_token(无 refresh_token 和 api_key) - let mut provider = CodexProvider::new(); - provider.credentials.access_token = Some("test_access_token".to_string()); - provider.credentials.refresh_token = None; - provider.credentials.api_key = None; - - let result = provider.refresh_token().await; - assert!(result.is_ok()); - assert_eq!(result.unwrap(), "test_access_token"); - } - - #[tokio::test] - async fn test_refresh_token_with_no_credentials() { - // 场景:无任何凭证(api_key、refresh_token、access_token 均为 None) - let mut provider = CodexProvider::new(); - provider.credentials.api_key = None; - provider.credentials.refresh_token = None; - provider.credentials.access_token = None; - - let result = provider.refresh_token().await; - assert!(result.is_err()); - let error_msg = result.unwrap_err().to_string(); - assert!(error_msg.contains("没有可用的认证凭证")); - assert!(error_msg.contains("API Key 模式")); - assert!(error_msg.contains("OAuth 模式")); - assert!(error_msg.contains("Access Token 模式")); - } - - #[tokio::test] - async fn test_api_key_priority_over_refresh_token() { - // 场景:同时有 api_key 和 refresh_token - let mut provider = CodexProvider::new(); - provider.credentials.api_key = Some("sk-test-api-key".to_string()); - provider.credentials.refresh_token = Some("test_refresh_token".to_string()); - provider.credentials.access_token = Some("test_access_token".to_string()); - - let result = provider.refresh_token().await; - assert!(result.is_ok()); - // 应该返回 API Key(优先级最高) - assert_eq!(result.unwrap(), "sk-test-api-key"); - } - - #[tokio::test] - async fn test_refresh_token_with_expired_access_token() { - // 场景:只有 access_token(已过期) - let mut provider = CodexProvider::new(); - provider.credentials.access_token = Some("expired_access_token".to_string()); - provider.credentials.expires_at = Some("2020-01-01T00:00:00Z".to_string()); - provider.credentials.refresh_token = None; - provider.credentials.api_key = None; - - let result = provider.refresh_token().await; - assert!(result.is_ok()); - // 应该返回 access_token(即使已过期,由上层处理) - assert_eq!(result.unwrap(), "expired_access_token"); - } -} - -// ============================================================================ -// OAuth 登录功能(参考 Antigravity 实现) -// ============================================================================ - -use once_cell::sync::Lazy; -use std::sync::Arc; -use tokio::sync::{oneshot, RwLock}; -use uuid::Uuid; - -/// 全局 Codex OAuth 服务器状态 -/// 用于在重新打开授权对话框时关闭之前的服务器 -static CODEX_OAUTH_SERVER_SHUTDOWN: Lazy>>> = - Lazy::new(|| RwLock::new(None)); - -/// 停止之前运行的 Codex OAuth 服务器(如果有) -pub async fn stop_codex_oauth_server() { - let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; - if let Some(shutdown_tx) = guard.take() { - tracing::info!("[Codex OAuth] 关闭之前的 OAuth 服务器"); - let _ = shutdown_tx.send(()); - // 给服务器一些时间来关闭 - tokio::time::sleep(std::time::Duration::from_millis(100)).await; - } -} - -/// OAuth 登录成功后的凭证信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CodexOAuthResult { - pub credentials: CodexCredentials, - pub creds_file_path: String, -} - -/// OpenAI OAuth 固定回调端口(必须与 client_id 注册的回调地址一致) -const OPENAI_OAUTH_CALLBACK_PORT: u16 = 1455; - -/// OpenAI OAuth 固定回调路径(必须与 client_id 注册的回调地址一致) -const OPENAI_OAUTH_CALLBACK_PATH: &str = "/auth/callback"; - -/// 生成 OAuth 授权 URL(用于外部浏览器登录) -/// -/// 注意:OpenAI OAuth 要求 redirect_uri 必须是预先注册的固定地址 -/// Codex CLI 的 client_id 只注册了 http://localhost:1455/auth/callback -pub fn generate_codex_auth_url(state: &str, code_challenge: &str) -> String { - let redirect_uri = - format!("http://localhost:{OPENAI_OAUTH_CALLBACK_PORT}{OPENAI_OAUTH_CALLBACK_PATH}"); - - let params = [ - ("client_id", OPENAI_CLIENT_ID), - ("response_type", "code"), - ("redirect_uri", redirect_uri.as_str()), - // 使用基础 scope,与 CLIProxyAPI 保持一致 - ("scope", "openid email profile offline_access"), - ("state", state), - ("code_challenge", code_challenge), - ("code_challenge_method", "S256"), - ("prompt", "login"), - ("id_token_add_organizations", "true"), - ("codex_cli_simplified_flow", "true"), - ]; - - let query = params - .iter() - .map(|(k, v)| format!("{}={}", k, urlencoding::encode(v))) - .collect::>() - .join("&"); - - format!("{OPENAI_AUTH_URL}?{query}") -} - -/// 用授权码交换 Token -pub async fn exchange_codex_code_for_token( - client: &Client, - code: &str, - code_verifier: &str, - redirect_uri: &str, -) -> Result> { - let params = [ - ("grant_type", "authorization_code"), - ("client_id", OPENAI_CLIENT_ID), - ("code", code), - ("redirect_uri", redirect_uri), - ("code_verifier", code_verifier), - ]; - - let resp = client - .post(OPENAI_TOKEN_URL) - .header("Content-Type", "application/x-www-form-urlencoded") - .header("Accept", "application/json") - .form(¶ms) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(format!("Token 交换失败: {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - Ok(data) -} - -/// OAuth 成功页面 HTML -const CODEX_OAUTH_SUCCESS_HTML: &str = r#" - - - - 授权成功 - - - -
-

✓ 授权成功

-

Codex 账号已添加到 Lime

- -

可以关闭此页面

-
- -"#; - -/// OAuth 失败页面 HTML -const CODEX_OAUTH_ERROR_HTML: &str = r#" - - - - 授权失败 - - - -
-

✗ 授权失败

-

ERROR_PLACEHOLDER

-

请关闭此页面后重试

-
- -"#; - -/// 启动 OAuth 服务器并返回授权 URL(不打开浏览器) -/// 服务器会在后台等待回调,成功后返回凭证 -/// -/// 注意:OpenAI OAuth 要求使用固定的回调地址 http://localhost:1455/auth/callback -pub async fn start_codex_oauth_server_and_get_url() -> Result< - ( - String, - impl std::future::Future>>, - ), - Box, -> { - use axum::{extract::Query, response::Html, routing::get, Router}; - use std::collections::HashMap; - use tokio::net::TcpListener; - - // 首先停止之前可能运行的 OAuth 服务器 - stop_codex_oauth_server().await; - - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build()?; - - // 生成 PKCE codes - let pkce_codes = PKCECodes::generate()?; - let code_verifier = pkce_codes.code_verifier.clone(); - let code_challenge = pkce_codes.code_challenge.clone(); - - // 生成随机 state - let state = Uuid::new_v4().to_string(); - let state_clone = state.clone(); - - // 创建 channel 用于接收回调结果 - let (tx, rx) = oneshot::channel::>(); - let tx = Arc::new(tokio::sync::Mutex::new(Some(tx))); - - // 创建 shutdown channel - let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); - - // 保存 shutdown sender 到全局状态 - { - let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; - *guard = Some(shutdown_tx); - } - - // 使用固定端口 1455(OpenAI OAuth 要求) - let port = OPENAI_OAUTH_CALLBACK_PORT; - let listener = TcpListener::bind(format!("127.0.0.1:{port}")).await.map_err(|e| { - if e.kind() == std::io::ErrorKind::AddrInUse { - format!( - "端口 {port} 已被占用。OpenAI OAuth 要求使用固定端口 1455,请关闭占用该端口的应用后重试。" - ) - } else { - format!("绑定端口 {port} 失败: {e}") - } - })?; - - let redirect_uri = - format!("http://localhost:{OPENAI_OAUTH_CALLBACK_PORT}{OPENAI_OAUTH_CALLBACK_PATH}"); - let redirect_uri_clone = redirect_uri.clone(); - - // 生成授权 URL(不再传入 port 参数) - let auth_url = generate_codex_auth_url(&state, &code_challenge); - - tracing::info!( - "[Codex OAuth] 服务器启动在端口 {}, 授权 URL: {}", - port, - auth_url - ); - - // 构建路由(使用固定的回调路径 /auth/callback) - let app = Router::new().route( - OPENAI_OAUTH_CALLBACK_PATH, - get(move |Query(params): Query>| { - let tx = tx.clone(); - let client = client.clone(); - let state_expected = state_clone.clone(); - let redirect_uri = redirect_uri_clone.clone(); - let code_verifier = code_verifier.clone(); - - async move { - let code = params.get("code"); - let returned_state = params.get("state"); - let error = params.get("error"); - - // 检查错误 - if let Some(err) = error { - let html = CODEX_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", err); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(format!("OAuth 错误: {err}"))); - } - return Html(html); - } - - // 检查 state - if returned_state.map(|s| s.as_str()) != Some(&state_expected) { - let html = - CODEX_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "State 验证失败"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("State 验证失败".to_string())); - } - return Html(html); - } - - // 检查 code - let code = match code { - Some(c) => c, - None => { - let html = - CODEX_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", "未收到授权码"); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err("未收到授权码".to_string())); - } - return Html(html); - } - }; - - // 交换 Token - let token_result = - exchange_codex_code_for_token(&client, code, &code_verifier, &redirect_uri) - .await; - let token_data = match token_result { - Ok(data) => data, - Err(e) => { - let html = - CODEX_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &e.to_string()); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(e.to_string())); - } - return Html(html); - } - }; - - let access_token = token_data["access_token"].as_str().unwrap_or_default(); - let refresh_token = token_data["refresh_token"].as_str().map(|s| s.to_string()); - let id_token = token_data["id_token"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expires_in"].as_i64(); - - // 解析 ID Token 获取用户信息 - let (account_id, email) = if let Some(ref id_token) = id_token { - parse_jwt_claims(id_token) - } else { - (None, None) - }; - - // 构建凭证 - let now = chrono::Utc::now(); - let credentials = CodexCredentials { - id_token, - access_token: Some(access_token.to_string()), - refresh_token, - api_key: None, - api_base_url: None, - account_id, - last_refresh: Some(now.to_rfc3339()), - email: email.clone(), - r#type: "codex".to_string(), - expires_at: expires_in - .map(|e| (now + chrono::Duration::seconds(e)).to_rfc3339()), - }; - - // 保存凭证到应用数据目录 - let creds_dir = dirs::data_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join("lime") - .join("credentials") - .join("codex"); - - if let Err(e) = std::fs::create_dir_all(&creds_dir) { - let html = CODEX_OAUTH_ERROR_HTML - .replace("ERROR_PLACEHOLDER", &format!("创建目录失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(format!("创建目录失败: {e}"))); - } - return Html(html); - } - - // 生成唯一文件名 - let uuid = Uuid::new_v4().to_string(); - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - let filename = format!("codex_{}_{}.json", &uuid[..8], timestamp); - let creds_file_path = creds_dir.join(&filename); - - // 保存凭证 - let creds_json = match serde_json::to_string_pretty(&credentials) { - Ok(json) => json, - Err(e) => { - let html = CODEX_OAUTH_ERROR_HTML - .replace("ERROR_PLACEHOLDER", &format!("序列化凭证失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(format!("序列化凭证失败: {e}"))); - } - return Html(html); - } - }; - - if let Err(e) = std::fs::write(&creds_file_path, &creds_json) { - let html = CODEX_OAUTH_ERROR_HTML - .replace("ERROR_PLACEHOLDER", &format!("保存凭证失败: {e}")); - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Err(format!("保存凭证失败: {e}"))); - } - return Html(html); - } - - tracing::info!("[Codex OAuth] 凭证已保存到: {:?}", creds_file_path); - - // 发送成功结果 - let result = CodexOAuthResult { - credentials, - creds_file_path: creds_file_path.to_string_lossy().to_string(), - }; - - if let Some(sender) = tx.lock().await.take() { - let _ = sender.send(Ok(result)); - } - - // 返回成功页面 - let html = CODEX_OAUTH_SUCCESS_HTML.replace( - "EMAIL_PLACEHOLDER", - &email.unwrap_or_else(|| "未知邮箱".to_string()), - ); - Html(html) - } - }), - ); - - // 启动服务器(支持优雅关闭) - let server = axum::serve(listener, app).with_graceful_shutdown(async move { - let _ = shutdown_rx.await; - tracing::info!("[Codex OAuth] 服务器收到关闭信号"); - }); - - // 创建等待 future - let wait_future = async move { - // 设置超时(5 分钟) - let timeout = tokio::time::timeout(std::time::Duration::from_secs(300), async { - // 启动服务器(在后台运行) - tokio::spawn(async move { - if let Err(e) = server.await { - tracing::error!("[Codex OAuth] 服务器错误: {}", e); - } - }); - - // 等待回调结果 - match rx.await { - Ok(result) => result.map_err(|e| { - Box::new(std::io::Error::other(e)) as Box - }), - Err(_) => Err("OAuth 回调通道关闭".into()), - } - }); - - match timeout.await { - Ok(result) => { - // 成功或失败后都清理全局状态 - let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; - *guard = None; - result - } - Err(_) => { - // 超时后也清理全局状态 - let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; - *guard = None; - Err("OAuth 登录超时(5分钟)".into()) - } - } - }; - - Ok((auth_url, wait_future)) -} - -/// 启动 Codex OAuth 登录流程(自动打开浏览器) -pub async fn start_codex_oauth_login() -> Result> { - let (auth_url, wait_future) = start_codex_oauth_server_and_get_url().await?; - - tracing::info!("[Codex OAuth] 打开浏览器进行授权: {}", auth_url); - - // 打开浏览器 - if let Err(e) = open::that(&auth_url) { - tracing::warn!("[Codex OAuth] 无法打开浏览器: {}. 请手动打开 URL.", e); - } - - // 等待回调 - wait_future.await } diff --git a/src-tauri/crates/providers/src/providers/gemini.rs b/src-tauri/crates/providers/src/providers/gemini.rs index 2141ec1ba..29bd54f19 100644 --- a/src-tauri/crates/providers/src/providers/gemini.rs +++ b/src-tauri/crates/providers/src/providers/gemini.rs @@ -1,485 +1,14 @@ -//! Gemini CLI OAuth Provider +//! Gemini API Key Provider //! -//! 实现 Google Gemini OAuth 认证流程,与 CLIProxyAPI 对齐。 -//! 支持 Token 刷新、重试机制和统一凭证格式。 +//! Gemini OAuth 已随凭证池退役;本模块只保留 API Key Provider 主路径。 -#![allow(dead_code)] - -use super::error::{ - create_auth_error, create_config_error, create_token_refresh_error, ProviderError, -}; -use super::traits::{CredentialProvider, ProviderResult}; -use async_trait::async_trait; use reqwest::Client; -use serde::{Deserialize, Serialize}; use std::error::Error; -use std::path::PathBuf; - -// Constants - 与 CLIProxyAPI 对齐 -const CODE_ASSIST_ENDPOINT: &str = "https://cloudcode-pa.googleapis.com"; -const CODE_ASSIST_API_VERSION: &str = "v1internal"; -const CREDENTIALS_DIR: &str = ".gemini"; -const CREDENTIALS_FILE: &str = "oauth_creds.json"; - -// OAuth 端点 -const GEMINI_TOKEN_URL: &str = "https://oauth2.googleapis.com/token"; - -// Gemini CLI OAuth 配置 - 与 CLIProxyAPI 对齐 -const DEFAULT_GEMINI_OAUTH_CLIENT_ID: &str = - "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com"; -const DEFAULT_GEMINI_OAUTH_CLIENT_SECRET: &str = "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"; - -// OAuth 凭证 - 优先从环境变量读取,否则使用硬编码的默认值 -fn get_oauth_client_id() -> String { - std::env::var("GEMINI_OAUTH_CLIENT_ID") - .unwrap_or_else(|_| DEFAULT_GEMINI_OAUTH_CLIENT_ID.to_string()) -} - -fn get_oauth_client_secret() -> String { - std::env::var("GEMINI_OAUTH_CLIENT_SECRET") - .unwrap_or_else(|_| DEFAULT_GEMINI_OAUTH_CLIENT_SECRET.to_string()) -} - -#[allow(dead_code)] -pub const GEMINI_MODELS: &[&str] = &[ - "gemini-2.5-flash", - "gemini-2.5-flash-lite", - "gemini-2.5-pro", - "gemini-2.5-pro-preview-06-05", - "gemini-2.5-flash-preview-09-2025", - "gemini-3-pro-preview", -]; - -/// Gemini OAuth 凭证存储 -/// -/// 与 CLIProxyAPI 的 GeminiTokenStorage 格式兼容 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GeminiCredentials { - /// 访问令牌 - pub access_token: Option, - /// 刷新令牌 - pub refresh_token: Option, - /// 令牌类型 - pub token_type: Option, - /// 过期时间戳(毫秒)- 兼容旧格式 - #[serde(skip_serializing_if = "Option::is_none")] - pub expiry_date: Option, - /// 过期时间(RFC3339 格式)- 新格式 - #[serde(skip_serializing_if = "Option::is_none")] - pub expire: Option, - /// OAuth 作用域 - pub scope: Option, - /// 用户邮箱 - #[serde(skip_serializing_if = "Option::is_none")] - pub email: Option, - /// 最后刷新时间(RFC3339 格式) - #[serde(skip_serializing_if = "Option::is_none")] - pub last_refresh: Option, - /// 凭证类型标识 - #[serde(default = "default_gemini_type", rename = "type")] - pub cred_type: String, - /// 嵌套的 token 对象(兼容 CLIProxyAPI 格式) - #[serde(skip_serializing_if = "Option::is_none")] - pub token: Option, -} - -fn default_gemini_type() -> String { - "gemini".to_string() -} - -/// 嵌套的 Token 信息(兼容 CLIProxyAPI 格式) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GeminiTokenInfo { - pub access_token: Option, - pub refresh_token: Option, - pub token_uri: Option, - pub client_id: Option, - pub client_secret: Option, - pub scopes: Option>, -} - -impl Default for GeminiCredentials { - fn default() -> Self { - Self { - access_token: None, - refresh_token: None, - token_type: Some("Bearer".to_string()), - expiry_date: None, - expire: None, - scope: None, - email: None, - last_refresh: None, - cred_type: default_gemini_type(), - token: None, - } - } -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GeminiContent { - pub role: String, - pub parts: Vec, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GeminiPart { - #[serde(skip_serializing_if = "Option::is_none")] - pub text: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct GeminiRequest { - pub model: String, - pub project: String, - pub request: GeminiRequestBody, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct GeminiRequestBody { - pub contents: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - pub system_instruction: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub generation_config: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct GeminiGenerationConfig { - #[serde(skip_serializing_if = "Option::is_none")] - pub temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub max_output_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub top_p: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub top_k: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GeminiResponse { - pub candidates: Option>, - #[serde(rename = "usageMetadata")] - pub usage_metadata: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct GeminiCandidate { - pub content: Option, - pub finish_reason: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct GeminiUsageMetadata { - pub prompt_token_count: Option, - pub candidates_token_count: Option, - pub total_token_count: Option, -} - -pub struct GeminiProvider { - pub credentials: GeminiCredentials, - pub project_id: Option, - pub client: Client, -} - -impl Default for GeminiProvider { - fn default() -> Self { - Self { - credentials: GeminiCredentials::default(), - project_id: None, - client: Client::new(), - } - } -} - -impl GeminiProvider { - pub fn new() -> Self { - Self::default() - } - - pub fn default_creds_path() -> PathBuf { - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(CREDENTIALS_DIR) - .join(CREDENTIALS_FILE) - } - - pub async fn load_credentials(&mut self) -> Result<(), Box> { - let path = Self::default_creds_path(); - - if tokio::fs::try_exists(&path).await.unwrap_or(false) { - let content = tokio::fs::read_to_string(&path).await?; - let creds: GeminiCredentials = serde_json::from_str(&content)?; - self.credentials = creds; - } - - Ok(()) - } - - pub async fn load_credentials_from_path( - &mut self, - path: &str, - ) -> Result<(), Box> { - let content = tokio::fs::read_to_string(path).await?; - let creds: GeminiCredentials = serde_json::from_str(&content)?; - self.credentials = creds; - Ok(()) - } - - pub async fn save_credentials(&self) -> Result<(), Box> { - let path = Self::default_creds_path(); - if let Some(parent) = path.parent() { - tokio::fs::create_dir_all(parent).await?; - } - let content = serde_json::to_string_pretty(&self.credentials)?; - tokio::fs::write(&path, content).await?; - Ok(()) - } - - /// 检查 Token 是否有效 - pub fn is_token_valid(&self) -> bool { - if self.credentials.access_token.is_none() { - return false; - } - - // 优先检查 RFC3339 格式的过期时间 - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) { - let now = chrono::Utc::now(); - // Token 有效期需要超过 5 分钟 - return expires > now + chrono::Duration::minutes(5); - } - } - - // 兼容旧的毫秒时间戳格式 - if let Some(expiry) = self.credentials.expiry_date { - let now = chrono::Utc::now().timestamp_millis(); - return expiry > now + 300_000; - } - - true - } - - /// 刷新 Token - pub async fn refresh_token(&mut self) -> Result> { - let refresh_token = self - .credentials - .refresh_token - .as_ref() - .ok_or_else(|| create_config_error("没有可用的 refresh_token"))?; - - let client_id = get_oauth_client_id(); - let client_secret = get_oauth_client_secret(); - - tracing::info!("[GEMINI] 正在刷新 Token"); - - let params = [ - ("client_id", client_id.as_str()), - ("client_secret", client_secret.as_str()), - ("refresh_token", refresh_token.as_str()), - ("grant_type", "refresh_token"), - ]; - - let resp = self - .client - .post(GEMINI_TOKEN_URL) - .header("Content-Type", "application/x-www-form-urlencoded") - .header("Accept", "application/json") - .form(¶ms) - .send() - .await - .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; - - if !resp.status().is_success() { - let status = resp.status().as_u16(); - let body = resp.text().await.unwrap_or_default(); - tracing::error!("[GEMINI] Token 刷新失败: {} - {}", status, body); - return Err(create_token_refresh_error(status, &body, "GEMINI")); - } - - let data: serde_json::Value = resp - .json() - .await - .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; - - let new_token = data["access_token"] - .as_str() - .ok_or_else(|| create_auth_error("响应中没有 access_token"))?; - - self.credentials.access_token = Some(new_token.to_string()); - - // 更新过期时间(同时保存两种格式以兼容) - if let Some(expires_in) = data["expires_in"].as_i64() { - let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in); - self.credentials.expire = Some(expires_at.to_rfc3339()); - self.credentials.expiry_date = Some(expires_at.timestamp_millis()); - } - - // 更新最后刷新时间 - self.credentials.last_refresh = Some(chrono::Utc::now().to_rfc3339()); - - // 保存刷新后的凭证 - self.save_credentials().await?; - - tracing::info!("[GEMINI] Token 刷新成功"); - Ok(new_token.to_string()) - } - - /// 带重试机制的 Token 刷新 - /// - /// 最多重试 `max_retries` 次,使用指数退避策略 - pub async fn refresh_token_with_retry( - &mut self, - max_retries: u32, - ) -> Result> { - let mut last_error = None; - - for attempt in 0..max_retries { - if attempt > 0 { - // 指数退避: 1s, 2s, 4s, ... - let delay = std::time::Duration::from_secs(1 << attempt); - tracing::info!("[GEMINI] 第 {} 次重试,等待 {:?}", attempt + 1, delay); - tokio::time::sleep(delay).await; - } - - match self.refresh_token().await { - Ok(token) => return Ok(token), - Err(e) => { - tracing::warn!("[GEMINI] Token 刷新第 {} 次尝试失败: {}", attempt + 1, e); - last_error = Some(e); - } - } - } - - tracing::error!("[GEMINI] Token 刷新在 {} 次尝试后失败", max_retries); - Err(last_error.unwrap_or_else(|| create_auth_error("Token 刷新失败,请重新登录"))) - } - - /// 确保 Token 有效,必要时自动刷新 - pub async fn ensure_valid_token(&mut self) -> Result> { - if !self.is_token_valid() { - tracing::info!("[GEMINI] Token 需要刷新"); - self.refresh_token_with_retry(3).await - } else { - self.credentials - .access_token - .clone() - .ok_or_else(|| "没有可用的 access_token".into()) - } - } - - pub fn get_api_url(&self, action: &str) -> String { - format!("{CODE_ASSIST_ENDPOINT}/{CODE_ASSIST_API_VERSION}:{action}") - } - - pub async fn call_api( - &self, - action: &str, - body: &serde_json::Value, - ) -> Result> { - let token = self - .credentials - .access_token - .as_ref() - .ok_or("No access token")?; - - let url = self.get_api_url(action); - - let resp = self - .client - .post(&url) - .header("Authorization", format!("Bearer {token}")) - .header("Content-Type", "application/json") - .json(body) - .send() - .await?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(format!("API call failed: {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - Ok(data) - } - - pub async fn discover_project(&mut self) -> Result> { - if let Some(ref project_id) = self.project_id { - return Ok(project_id.clone()); - } - - let body = serde_json::json!({ - "cloudaicompanionProject": "", - "metadata": { - "ideType": "IDE_UNSPECIFIED", - "platform": "PLATFORM_UNSPECIFIED", - "pluginType": "GEMINI", - "duetProject": "" - } - }); - - let resp = self.call_api("loadCodeAssist", &body).await?; - - if let Some(project) = resp["cloudaicompanionProject"].as_str() { - if !project.is_empty() { - self.project_id = Some(project.to_string()); - return Ok(project.to_string()); - } - } - - // Need to onboard - let onboard_body = serde_json::json!({ - "tierId": "free-tier", - "cloudaicompanionProject": "", - "metadata": { - "ideType": "IDE_UNSPECIFIED", - "platform": "PLATFORM_UNSPECIFIED", - "pluginType": "GEMINI", - "duetProject": "" - } - }); - - let mut lro_resp = self.call_api("onboardUser", &onboard_body).await?; - - // Poll until done - for _ in 0..30 { - if lro_resp["done"].as_bool().unwrap_or(false) { - break; - } - tokio::time::sleep(tokio::time::Duration::from_secs(2)).await; - lro_resp = self.call_api("onboardUser", &onboard_body).await?; - } - - let project_id = lro_resp["response"]["cloudaicompanionProject"]["id"] - .as_str() - .unwrap_or("") - .to_string(); - - if project_id.is_empty() { - return Err("Failed to discover project ID".into()); - } - - self.project_id = Some(project_id.clone()); - Ok(project_id) - } -} - -// ============ Gemini API Key Provider ============ /// Default Gemini API base URL pub const GEMINI_API_BASE_URL: &str = "https://generativelanguage.googleapis.com"; -/// Gemini API Key Provider for multi-account load balancing -/// -/// This provider supports: -/// - Multiple API keys with round-robin load balancing -/// - Per-key custom base URLs -/// - Model exclusion filtering (to be implemented in task 11.2) +/// Gemini API Key Provider for multi-account load balancing. #[derive(Debug, Clone)] pub struct GeminiApiKeyCredential { /// Credential ID @@ -497,7 +26,7 @@ pub struct GeminiApiKeyCredential { } impl GeminiApiKeyCredential { - /// Create a new Gemini API Key credential + /// Create a new Gemini API Key credential. pub fn new(id: String, api_key: String) -> Self { Self { id, @@ -509,46 +38,44 @@ impl GeminiApiKeyCredential { } } - /// Set custom base URL + /// Set custom base URL. pub fn with_base_url(mut self, base_url: Option) -> Self { self.base_url = base_url; self } - /// Set excluded models + /// Set excluded models. pub fn with_excluded_models(mut self, excluded_models: Vec) -> Self { self.excluded_models = excluded_models; self } - /// Set proxy URL + /// Set proxy URL. pub fn with_proxy_url(mut self, proxy_url: Option) -> Self { self.proxy_url = proxy_url; self } - /// Set disabled state + /// Set disabled state. pub fn with_disabled(mut self, disabled: bool) -> Self { self.disabled = disabled; self } - /// Get the effective base URL (custom or default) + /// Get the effective base URL (custom or default). pub fn get_base_url(&self) -> &str { self.base_url.as_deref().unwrap_or(GEMINI_API_BASE_URL) } - /// Check if this credential is available (not disabled) + /// Check if this credential is available (not disabled). pub fn is_available(&self) -> bool { !self.disabled } - /// Check if this credential supports the given model - /// Returns false if the model matches any exclusion pattern + /// Check if this credential supports the given model. pub fn supports_model(&self, model: &str) -> bool { !self.excluded_models.iter().any(|pattern| { if pattern.contains('*') { - // Simple wildcard matching let pattern = pattern.replace('*', ".*"); regex::Regex::new(&format!("^{pattern}$")) .map(|re| re.is_match(model)) @@ -559,16 +86,13 @@ impl GeminiApiKeyCredential { }) } - /// Build the API URL for a given model and action + /// Build the API URL for a given model and action. pub fn build_api_url(&self, model: &str, action: &str) -> String { format!("{}/v1beta/models/{}:{}", self.get_base_url(), model, action) } } -/// Gemini API Key Provider -/// -/// Manages multiple Gemini API keys with load balancing support. -/// Integrates with the credential pool system for round-robin selection. +/// Gemini API Key Provider. pub struct GeminiApiKeyProvider { /// HTTP client pub client: Client, @@ -581,19 +105,19 @@ impl Default for GeminiApiKeyProvider { } impl GeminiApiKeyProvider { - /// Create a new Gemini API Key provider + /// Create a new Gemini API Key provider. pub fn new() -> Self { Self { client: Client::new(), } } - /// Create a provider with a custom HTTP client + /// Create a provider with a custom HTTP client. pub fn with_client(client: Client) -> Self { Self { client } } - /// Make a generateContent request using the given credential + /// Make a generateContent request using the given credential. pub async fn generate_content( &self, credential: &GeminiApiKeyCredential, @@ -621,7 +145,7 @@ impl GeminiApiKeyProvider { Ok(data) } - /// Make a streamGenerateContent request using the given credential + /// Make a streamGenerateContent request using the given credential. pub async fn stream_generate_content( &self, credential: &GeminiApiKeyCredential, @@ -651,7 +175,7 @@ impl GeminiApiKeyProvider { Ok(resp) } - /// List available models using the given credential + /// List available models using the given credential. pub async fn list_models( &self, credential: &GeminiApiKeyCredential, @@ -721,14 +245,9 @@ mod gemini_api_key_tests { "gemini-*-preview".to_string(), ]); - // Exact match exclusion assert!(!cred.supports_model("gemini-2.5-pro")); - - // Wildcard exclusion assert!(!cred.supports_model("gemini-3-preview")); assert!(!cred.supports_model("gemini-2.5-preview")); - - // Not excluded assert!(cred.supports_model("gemini-2.5-flash")); assert!(cred.supports_model("gemini-2.0-flash")); } @@ -755,660 +274,3 @@ mod gemini_api_key_tests { let _provider = GeminiApiKeyProvider::new(); } } - -// ============================================================================ -// Gemini OAuth 登录功能 -// ============================================================================ - -use std::sync::Arc; -use tokio::sync::oneshot; -use uuid::Uuid; - -// Gemini CLI OAuth 配置 - 与 claude-relay-service 对齐 -pub const GEMINI_OAUTH_CLIENT_ID: &str = - "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com"; -pub const GEMINI_OAUTH_CLIENT_SECRET: &str = "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"; -pub const GEMINI_OAUTH_SCOPES: &[&str] = &["https://www.googleapis.com/auth/cloud-platform"]; -pub const GEMINI_OAUTH_REDIRECT_URI: &str = "https://codeassist.google.com/authcode"; - -/// OAuth 登录成功后的凭证信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GeminiOAuthResult { - pub credentials: GeminiCredentials, - pub creds_file_path: String, -} - -/// 生成 PKCE code_verifier 和 code_challenge -fn generate_pkce() -> (String, String) { - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - use sha2::{Digest, Sha256}; - - // 生成 43-128 字符的随机字符串作为 code_verifier - let code_verifier: String = (0..64) - .map(|_| { - let idx = rand::random::() % 66; - let chars = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-._~"; - chars[idx as usize] as char - }) - .collect(); - - // 计算 code_challenge = BASE64URL(SHA256(code_verifier)) - let mut hasher = Sha256::new(); - hasher.update(code_verifier.as_bytes()); - let hash = hasher.finalize(); - let code_challenge = URL_SAFE_NO_PAD.encode(hash); - - (code_verifier, code_challenge) -} - -/// 生成 OAuth 授权 URL(使用 PKCE) -pub fn generate_gemini_auth_url(state: &str, code_challenge: &str) -> String { - let scopes = GEMINI_OAUTH_SCOPES.join(" "); - - let params = [ - ("access_type", "offline"), - ("client_id", GEMINI_OAUTH_CLIENT_ID), - ("code_challenge", code_challenge), - ("code_challenge_method", "S256"), - ("prompt", "select_account"), - ("redirect_uri", GEMINI_OAUTH_REDIRECT_URI), - ("response_type", "code"), - ("scope", &scopes), - ("state", state), - ]; - - let query = params - .iter() - .map(|(k, v)| format!("{}={}", k, urlencoding::encode(v))) - .collect::>() - .join("&"); - - format!("https://accounts.google.com/o/oauth2/v2/auth?{query}") -} - -/// Gemini OAuth 会话信息(用于存储 PKCE code_verifier) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GeminiOAuthSession { - pub session_id: String, - pub code_verifier: String, - pub state: String, - pub created_at: i64, -} - -/// 生成 Gemini OAuth 授权 URL 和会话信息 -/// -/// 返回 (auth_url, session) 元组 -/// - auth_url: 用户需要在浏览器中打开的授权 URL -/// - session: 包含 code_verifier 的会话信息,用于后续交换 token -pub fn generate_gemini_auth_url_with_session() -> (String, GeminiOAuthSession) { - let (code_verifier, code_challenge) = generate_pkce(); - let state = Uuid::new_v4().to_string(); - let session_id = Uuid::new_v4().to_string(); - - let auth_url = generate_gemini_auth_url(&state, &code_challenge); - - let session = GeminiOAuthSession { - session_id, - code_verifier, - state, - created_at: chrono::Utc::now().timestamp(), - }; - - (auth_url, session) -} - -/// 用授权码交换 Token 并创建凭证 -/// -/// 完整流程: -/// 1. 用 code + code_verifier 交换 tokens -/// 2. 获取用户邮箱 -/// 3. 获取项目 ID -/// 4. 保存凭证到文件 -pub async fn exchange_gemini_code_and_create_credentials( - code: &str, - code_verifier: &str, -) -> Result> { - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build()?; - - tracing::info!("[Gemini OAuth] 正在用授权码交换 Token..."); - - // 交换 Token - let token_data = exchange_gemini_code_for_token(&client, code, code_verifier).await?; - - let access_token = token_data["access_token"] - .as_str() - .ok_or("响应中没有 access_token")? - .to_string(); - let refresh_token = token_data["refresh_token"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expires_in"].as_i64(); - - tracing::info!("[Gemini OAuth] Token 交换成功"); - - // 获取用户邮箱 - let email = fetch_gemini_user_email(&client, &access_token) - .await - .ok() - .flatten(); - - tracing::info!("[Gemini OAuth] 用户邮箱: {:?}", email); - - // 获取项目 ID - let _project_id = fetch_gemini_project_id(&client, &access_token) - .await - .ok() - .flatten(); - - // 构建凭证 - let now = chrono::Utc::now(); - let expires_at = expires_in.map(|secs| now + chrono::Duration::seconds(secs)); - - let credentials = GeminiCredentials { - access_token: Some(access_token), - refresh_token, - token_type: Some("Bearer".to_string()), - expiry_date: expires_at.map(|t| t.timestamp_millis()), - expire: expires_at.map(|t| t.to_rfc3339()), - scope: Some(GEMINI_OAUTH_SCOPES.join(" ")), - email, - last_refresh: Some(now.to_rfc3339()), - cred_type: "gemini".to_string(), - token: None, - }; - - // 保存凭证到文件 - let file_path = save_gemini_credentials_to_file(&credentials).await?; - - tracing::info!("[Gemini OAuth] 凭证已保存到: {}", file_path); - - Ok(GeminiOAuthResult { - credentials, - creds_file_path: file_path, - }) -} - -/// 用授权码交换 Token(使用 PKCE) -pub async fn exchange_gemini_code_for_token( - client: &Client, - code: &str, - code_verifier: &str, -) -> Result> { - let params = [ - ("code", code), - ("client_id", GEMINI_OAUTH_CLIENT_ID), - ("client_secret", GEMINI_OAUTH_CLIENT_SECRET), - ("code_verifier", code_verifier), - ("redirect_uri", GEMINI_OAUTH_REDIRECT_URI), - ("grant_type", "authorization_code"), - ]; - - let resp = client.post(GEMINI_TOKEN_URL).form(¶ms).send().await?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(format!("Token 交换失败: {status} - {body}").into()); - } - - let data: serde_json::Value = resp.json().await?; - Ok(data) -} - -/// 获取用户邮箱 -pub async fn fetch_gemini_user_email( - client: &Client, - access_token: &str, -) -> Result, Box> { - let resp = client - .get("https://www.googleapis.com/oauth2/v2/userinfo") - .header("Authorization", format!("Bearer {access_token}")) - .send() - .await?; - - if resp.status().is_success() { - let data: serde_json::Value = resp.json().await?; - Ok(data["email"].as_str().map(|s| s.to_string())) - } else { - Ok(None) - } -} - -/// 获取项目 ID(通过 loadCodeAssist 接口) -pub async fn fetch_gemini_project_id( - client: &Client, - access_token: &str, -) -> Result, Box> { - tracing::info!("[Gemini OAuth] 正在获取 projectId..."); - - let resp = client - .post(format!( - "{CODE_ASSIST_ENDPOINT}/{CODE_ASSIST_API_VERSION}:loadCodeAssist" - )) - .header("Authorization", format!("Bearer {access_token}")) - .header("Content-Type", "application/json") - .json(&serde_json::json!({ - "cloudaicompanionProject": "", - "metadata": { - "ideType": "IDE_UNSPECIFIED", - "platform": "PLATFORM_UNSPECIFIED", - "pluginType": "GEMINI", - "duetProject": "" - } - })) - .send() - .await?; - - let status = resp.status(); - tracing::info!("[Gemini OAuth] loadCodeAssist 响应状态: {}", status); - - if status.is_success() { - let data: serde_json::Value = resp.json().await?; - if let Some(project) = data["cloudaicompanionProject"].as_str() { - if !project.is_empty() { - tracing::info!("[Gemini OAuth] 获取到 projectId: {}", project); - return Ok(Some(project.to_string())); - } - } - tracing::info!("[Gemini OAuth] cloudaicompanionProject 为空"); - Ok(None) - } else { - let body = resp.text().await.unwrap_or_default(); - tracing::warn!( - "[Gemini OAuth] loadCodeAssist 请求失败: {} - {}", - status, - body - ); - Ok(None) - } -} - -/// OAuth 成功页面 HTML -const GEMINI_OAUTH_SUCCESS_HTML: &str = r#" - - - - 授权成功 - - - -
-

✓ 授权成功

-

Gemini 账号已添加到 Lime

- -

可以关闭此页面

-
- -"#; - -/// OAuth 失败页面 HTML -const GEMINI_OAUTH_ERROR_HTML: &str = r#" - - - - 授权失败 - - - -
-

✗ 授权失败

-

ERROR_PLACEHOLDER

-

请关闭此页面后重试

-
- -"#; - -/// 保存 Gemini 凭证到文件 -async fn save_gemini_credentials_to_file( - credentials: &GeminiCredentials, -) -> Result> { - // 生成唯一文件名 - let uuid = Uuid::new_v4().to_string(); - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap() - .as_secs(); - - let filename = format!("gemini_{}_{}_gemini.json", &uuid[..8], timestamp); - - // 获取凭证存储目录 - let credentials_dir = dirs::data_dir() - .ok_or("无法获取应用数据目录")? - .join("lime") - .join("credentials"); - - // 确保目录存在 - tokio::fs::create_dir_all(&credentials_dir).await?; - - let file_path = credentials_dir.join(&filename); - - // 写入凭证 - let content = serde_json::to_string_pretty(credentials)?; - tokio::fs::write(&file_path, content).await?; - - Ok(file_path.to_string_lossy().to_string()) -} - -/// 启动 OAuth 服务器并返回授权 URL(不打开浏览器) -/// 服务器会在后台等待回调,成功后返回凭证 -pub async fn start_gemini_oauth_server_and_get_url() -> Result< - ( - String, - impl std::future::Future< - Output = Result>, - >, - ), - Box, -> { - use axum::{extract::Query, response::Html, routing::get, Router}; - use std::collections::HashMap; - use tokio::net::TcpListener; - - let client = Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build()?; - - // 生成 PKCE - let (code_verifier, code_challenge) = generate_pkce(); - - // 生成随机 state - let state = Uuid::new_v4().to_string(); - let state_clone = state.clone(); - let code_verifier_clone = code_verifier.clone(); - - // 创建 channel 用于接收回调结果 - let (tx, rx) = oneshot::channel::>(); - let tx = Arc::new(tokio::sync::Mutex::new(Some(tx))); - - // 尝试绑定到多个端口 - let ports_to_try = [11451, 11452, 11453, 11454, 11455, 0]; - let mut listener = None; - let mut bound_port = 0; - - for port in ports_to_try { - match TcpListener::bind(format!("127.0.0.1:{port}")).await { - Ok(l) => { - bound_port = l.local_addr()?.port(); - listener = Some(l); - tracing::info!("[Gemini OAuth] 成功绑定到端口 {}", bound_port); - break; - } - Err(e) => { - tracing::warn!("[Gemini OAuth] 端口 {} 绑定失败: {}", port, e); - continue; - } - } - } - - let listener = listener.ok_or("无法绑定到任何可用端口")?; - - // 生成授权 URL - let auth_url = generate_gemini_auth_url(&state, &code_challenge); - - tracing::info!( - "[Gemini OAuth] 服务器启动在端口 {}, 授权 URL: {}", - bound_port, - auth_url - ); - - // 构建路由 - let app = Router::new().route( - "/oauth-callback", - get(move |Query(params): Query>| { - let tx = tx.clone(); - let client = client.clone(); - let state_expected = state_clone.clone(); - let code_verifier = code_verifier_clone.clone(); - - async move { - let code = params.get("code"); - let returned_state = params.get("state"); - let error = params.get("error"); - - // 检查错误 - if let Some(err) = error { - let error_desc = params - .get("error_description") - .map(|s| s.as_str()) - .unwrap_or("未知错误"); - let error_msg = format!("{err}: {error_desc}"); - tracing::error!("[Gemini OAuth] 授权失败: {}", error_msg); - - if let Some(tx) = tx.lock().await.take() { - let _ = tx.send(Err(error_msg.clone())); - } - - let html = GEMINI_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &error_msg); - return Html(html); - } - - // 验证 state - if returned_state.map(|s| s.as_str()) != Some(&state_expected) { - let error_msg = "State 验证失败"; - tracing::error!("[Gemini OAuth] {}", error_msg); - - if let Some(tx) = tx.lock().await.take() { - let _ = tx.send(Err(error_msg.to_string())); - } - - let html = GEMINI_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", error_msg); - return Html(html); - } - - // 获取授权码 - let code = match code { - Some(c) => c, - None => { - let error_msg = "未收到授权码"; - tracing::error!("[Gemini OAuth] {}", error_msg); - - if let Some(tx) = tx.lock().await.take() { - let _ = tx.send(Err(error_msg.to_string())); - } - - let html = GEMINI_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", error_msg); - return Html(html); - } - }; - - tracing::info!("[Gemini OAuth] 收到授权码,正在交换 Token..."); - - // 交换 Token - let token_result = - exchange_gemini_code_for_token(&client, code, &code_verifier).await; - - match token_result { - Ok(token_data) => { - let access_token = token_data["access_token"] - .as_str() - .unwrap_or("") - .to_string(); - let refresh_token = - token_data["refresh_token"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expires_in"].as_i64(); - - // 获取用户邮箱 - let email = fetch_gemini_user_email(&client, &access_token) - .await - .ok() - .flatten(); - - // 获取项目 ID - let project_id = fetch_gemini_project_id(&client, &access_token) - .await - .ok() - .flatten(); - - // 构建凭证 - let now = chrono::Utc::now(); - let expires_at = - expires_in.map(|secs| now + chrono::Duration::seconds(secs)); - - let credentials = GeminiCredentials { - access_token: Some(access_token), - refresh_token, - token_type: Some("Bearer".to_string()), - expiry_date: expires_at.map(|t| t.timestamp_millis()), - expire: expires_at.map(|t| t.to_rfc3339()), - scope: Some(GEMINI_OAUTH_SCOPES.join(" ")), - email: email.clone(), - last_refresh: Some(now.to_rfc3339()), - cred_type: "gemini".to_string(), - token: None, - }; - - // 保存凭证到文件 - match save_gemini_credentials_to_file(&credentials).await { - Ok(file_path) => { - tracing::info!("[Gemini OAuth] 凭证已保存到: {}", file_path); - - let result = GeminiOAuthResult { - credentials: credentials.clone(), - creds_file_path: file_path, - }; - - if let Some(tx) = tx.lock().await.take() { - let _ = tx.send(Ok(result)); - } - - let email_display = email.unwrap_or_else(|| "未知邮箱".to_string()); - let project_display = project_id - .map(|p| format!("

Project ID: {p}

")) - .unwrap_or_default(); - let html = GEMINI_OAUTH_SUCCESS_HTML - .replace("EMAIL_PLACEHOLDER", &email_display) - .replace( - "\n", - &format!("{project_display}\n"), - ); - Html(html) - } - Err(e) => { - let error_msg = format!("保存凭证失败: {e}"); - tracing::error!("[Gemini OAuth] {}", error_msg); - - if let Some(tx) = tx.lock().await.take() { - let _ = tx.send(Err(error_msg.clone())); - } - - let html = GEMINI_OAUTH_ERROR_HTML - .replace("ERROR_PLACEHOLDER", &error_msg); - Html(html) - } - } - } - Err(e) => { - let error_msg = format!("Token 交换失败: {e}"); - tracing::error!("[Gemini OAuth] {}", error_msg); - - if let Some(tx) = tx.lock().await.take() { - let _ = tx.send(Err(error_msg.clone())); - } - - let html = GEMINI_OAUTH_ERROR_HTML.replace("ERROR_PLACEHOLDER", &error_msg); - Html(html) - } - } - } - }), - ); - - // 启动服务器 - let server_future = async move { - axum::serve(listener, app) - .await - .map_err(|e| format!("服务器错误: {e}")) - }; - - // 启动服务器任务 - tokio::spawn(server_future); - - // 返回授权 URL 和等待结果的 Future - let wait_future = async move { - match rx.await { - Ok(result) => result.map_err(|e| e.into()), - Err(_) => Err("OAuth 回调通道关闭".into()), - } - }; - - Ok((auth_url, wait_future)) -} - -/// 启动 Gemini OAuth 登录流程(自动打开浏览器) -pub async fn start_gemini_oauth_login( -) -> Result> { - let (auth_url, wait_future) = start_gemini_oauth_server_and_get_url().await?; - - // 打开浏览器 - tracing::info!("[Gemini OAuth] 正在打开浏览器..."); - if let Err(e) = open::that(&auth_url) { - tracing::warn!("[Gemini OAuth] 无法自动打开浏览器: {}", e); - } - - // 等待回调 - wait_future.await -} - -// ============================================================================ -// CredentialProvider Trait 实现 -// ============================================================================ - -#[async_trait] -impl CredentialProvider for GeminiProvider { - async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()> { - GeminiProvider::load_credentials_from_path(self, path).await - } - - async fn save_credentials(&self) -> ProviderResult<()> { - GeminiProvider::save_credentials(self).await - } - - fn is_token_valid(&self) -> bool { - GeminiProvider::is_token_valid(self) - } - - fn is_token_expiring_soon(&self) -> bool { - // Gemini 使用与 is_token_valid 相同的逻辑,但阈值为 10 分钟 - if self.credentials.access_token.is_none() { - return true; - } - - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) { - let now = chrono::Utc::now(); - return expires <= now + chrono::Duration::minutes(10); - } - } - - if let Some(expiry) = self.credentials.expiry_date { - let now = chrono::Utc::now().timestamp_millis(); - return expiry <= now + 600_000; // 10 分钟 - } - - false - } - - async fn refresh_token(&mut self) -> ProviderResult { - GeminiProvider::refresh_token(self).await - } - - fn get_access_token(&self) -> Option<&str> { - self.credentials.access_token.as_deref() - } - - fn provider_type(&self) -> &'static str { - "gemini" - } -} diff --git a/src-tauri/crates/providers/src/providers/kiro.rs b/src-tauri/crates/providers/src/providers/kiro.rs deleted file mode 100644 index 4cbaf08d2..000000000 --- a/src-tauri/crates/providers/src/providers/kiro.rs +++ /dev/null @@ -1,1382 +0,0 @@ -//! Kiro/CodeWhisperer Provider - -#![allow(dead_code)] - -// 使用新的 translator 模块替代旧的 converter -use crate::providers::traits::{CredentialProvider, ProviderResult}; -use crate::translator::kiro::anthropic::request::convert_anthropic_to_codewhisperer; -use crate::translator::kiro::openai::request::convert_openai_to_codewhisperer_with_conversation_id; -use async_trait::async_trait; -use lime_core::models::anthropic::AnthropicMessagesRequest; -use lime_core::models::openai::*; -use reqwest::Client; -use serde::{Deserialize, Serialize}; -use std::error::Error; -use std::path::PathBuf; - -const MAX_KIRO_DEBUG_REQUEST_FILES: usize = 200; - -async fn prune_kiro_debug_request_files(debug_dir: &std::path::Path) { - let mut entries = match tokio::fs::read_dir(debug_dir).await { - Ok(entries) => entries, - Err(_) => return, - }; - - let mut files: Vec<(PathBuf, std::time::SystemTime)> = Vec::new(); - while let Ok(Some(entry)) = entries.next_entry().await { - let file_name = entry.file_name(); - let file_name = file_name.to_string_lossy(); - if !file_name.starts_with("cw_request_") || !file_name.ends_with(".json") { - continue; - } - - let modified = entry - .metadata() - .await - .ok() - .and_then(|metadata| metadata.modified().ok()) - .unwrap_or(std::time::SystemTime::UNIX_EPOCH); - files.push((entry.path(), modified)); - } - - if files.len() <= MAX_KIRO_DEBUG_REQUEST_FILES { - return; - } - - files.sort_by_key(|(_, modified)| *modified); - let overflow = files.len().saturating_sub(MAX_KIRO_DEBUG_REQUEST_FILES); - for (path, _) in files.into_iter().take(overflow) { - let _ = tokio::fs::remove_file(path).await; - } -} - -/// 根据凭证信息生成唯一的 Machine ID -/// -/// 采用静态 UUID 方案:每个凭证生成固定的 Machine ID,不随时间变化 -/// 优先级:uuid > profileArn > clientId > 系统硬件 ID -/// -/// 这是目前最稳定的方案,与 AIClient-2-API 的实现完全一致 -pub fn generate_machine_id_from_credentials( - profile_arn: Option<&str>, - client_id: Option<&str>, -) -> String { - generate_machine_id_from_credentials_with_uuid(None, profile_arn, client_id) -} - -/// 带 UUID 参数的 Machine ID 生成函数(与 AIClient-2-API 完全一致) -/// -/// 优先级:uuid > profileArn > clientId > 默认值 -/// 生成静态的 SHA256 哈希,不包含时间因子 -pub fn generate_machine_id_from_credentials_with_uuid( - uuid: Option<&str>, - profile_arn: Option<&str>, - client_id: Option<&str>, -) -> String { - use sha2::{Digest, Sha256}; - - // 优先级:uuid > profileArn > clientId > 默认值(与 AIClient-2-API 一致) - let unique_key = uuid - .filter(|s| !s.is_empty()) - .or(profile_arn.filter(|s| !s.is_empty())) - .or(client_id.filter(|s| !s.is_empty())) - .unwrap_or("KIRO_DEFAULT_MACHINE"); - - // 静态哈希,不添加时间因子(与 AIClient-2-API 保持一致) - let mut hasher = Sha256::new(); - hasher.update(unique_key.as_bytes()); - let result = hasher.finalize(); - format!("{result:x}") -} - -/// 获取系统运行时信息 -/// -/// 返回真实的操作系统名称和版本,用于构建更真实的 User-Agent -fn get_system_runtime_info() -> (String, String) { - let os_name = if cfg!(target_os = "macos") { - // macOS: 获取真实版本号 - let version = std::process::Command::new("sw_vers") - .arg("-productVersion") - .output() - .ok() - .and_then(|o| String::from_utf8(o.stdout).ok()) - .map(|s| s.trim().to_string()) - .unwrap_or_else(|| "14.0".to_string()); - format!("macos#{version}") - } else if cfg!(target_os = "linux") { - // Linux: 获取内核版本 - let version = std::process::Command::new("uname") - .arg("-r") - .output() - .ok() - .and_then(|o| String::from_utf8(o.stdout).ok()) - .map(|s| s.trim().to_string()) - .unwrap_or_else(|| "5.15.0".to_string()); - format!("linux#{version}") - } else if cfg!(target_os = "windows") { - // Windows: 使用固定版本(实际应该获取真实版本) - "windows#10.0".to_string() - } else { - "other#1.0".to_string() - }; - - // Node.js 版本模拟(Kiro IDE 使用 Electron,内置 Node.js) - // 使用常见的 LTS 版本 - let node_version = "20.18.0".to_string(); - - (os_name, node_version) -} - -/// 生成设备指纹 (Machine ID 的 SHA256) - 保留用于兼容 -/// -/// 与 Kiro IDE 保持一致的指纹生成方式(参考 Kir-Manager): -/// - macOS: 使用 IOPlatformUUID(硬件级别唯一标识) -/// - Linux: 使用 /etc/machine-id -/// - Windows: 使用 WMI 获取系统 UUID -/// -/// 最终返回 SHA256 哈希后的 64 字符十六进制字符串 -fn get_device_fingerprint() -> String { - use sha2::{Digest, Sha256}; - - let raw_id = - get_raw_machine_id().unwrap_or_else(|| "00000000-0000-0000-0000-000000000000".to_string()); - - // 使用 SHA256 生成 64 字符的十六进制指纹 - let mut hasher = Sha256::new(); - hasher.update(raw_id.as_bytes()); - let result = hasher.finalize(); - format!("{result:x}") -} - -/// 获取原始 Machine ID(未哈希) -fn get_raw_machine_id() -> Option { - use std::process::Command; - - if cfg!(target_os = "macos") { - // macOS: 使用 ioreg 获取 IOPlatformUUID - Command::new("ioreg") - .args(["-rd1", "-c", "IOPlatformExpertDevice"]) - .output() - .ok() - .and_then(|o| String::from_utf8(o.stdout).ok()) - .and_then(|s| { - s.lines() - .find(|l| l.contains("IOPlatformUUID")) - .and_then(|l| l.split('=').nth(1)) - .map(|s| s.trim().trim_matches('"').to_lowercase()) - }) - } else if cfg!(target_os = "linux") { - // Linux: 读取 /etc/machine-id 或 /var/lib/dbus/machine-id - std::fs::read_to_string("/etc/machine-id") - .or_else(|_| std::fs::read_to_string("/var/lib/dbus/machine-id")) - .ok() - .map(|s| s.trim().to_lowercase()) - } else if cfg!(target_os = "windows") { - // Windows: 使用 wmic 获取系统 UUID - #[cfg(target_os = "windows")] - { - use std::os::windows::process::CommandExt; - Command::new("wmic") - .args(["csproduct", "get", "UUID"]) - .creation_flags(0x08000000) // CREATE_NO_WINDOW - .output() - .ok() - .and_then(|o| String::from_utf8(o.stdout).ok()) - .and_then(|s| { - s.lines() - .skip(1) // 跳过表头 - .find(|l| !l.trim().is_empty()) - .map(|s| s.trim().to_lowercase()) - }) - } - #[cfg(not(target_os = "windows"))] - None - } else { - None - } -} - -/// 获取 Kiro IDE 版本号 -/// -/// 尝试从 Kiro.app 的 Info.plist 读取实际版本,失败时使用默认值 -fn get_kiro_version() -> String { - use std::process::Command; - - if cfg!(target_os = "macos") { - // 尝试从 Kiro.app 读取版本 - let kiro_paths = [ - "/Applications/Kiro.app/Contents/Info.plist", - // 用户目录下的安装 - &format!( - "{}/Applications/Kiro.app/Contents/Info.plist", - dirs::home_dir() - .map(|p| p.to_string_lossy().to_string()) - .unwrap_or_default() - ), - ]; - - for plist_path in &kiro_paths { - if let Ok(output) = Command::new("defaults") - .args(["read", plist_path, "CFBundleShortVersionString"]) - .output() - { - if let Ok(version) = String::from_utf8(output.stdout) { - let version = version.trim(); - if !version.is_empty() { - return version.to_string(); - } - } - } - } - } - - // 默认版本号 - "0.1.25".to_string() -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct KiroCredentials { - pub access_token: Option, - pub refresh_token: Option, - pub client_id: Option, - pub client_secret: Option, - pub profile_arn: Option, - /// 过期时间(支持 RFC3339 格式和时间戳格式) - pub expires_at: Option, - /// 过期时间(RFC3339 格式)- 与 CLIProxyAPI 兼容 - #[serde(skip_serializing_if = "Option::is_none")] - pub expire: Option, - pub region: Option, - pub auth_method: Option, - pub client_id_hash: Option, - /// 最后刷新时间(RFC3339 格式) - #[serde(skip_serializing_if = "Option::is_none")] - pub last_refresh: Option, - /// 凭证类型标识 - #[serde(default = "default_kiro_type", rename = "type")] - pub cred_type: String, -} - -fn default_kiro_type() -> String { - "kiro".to_string() -} - -impl Default for KiroCredentials { - fn default() -> Self { - Self { - access_token: None, - refresh_token: None, - client_id: None, - client_secret: None, - profile_arn: None, - expires_at: None, - expire: None, - region: Some("us-east-1".to_string()), - auth_method: Some("social".to_string()), - client_id_hash: None, - last_refresh: None, - cred_type: default_kiro_type(), - } - } -} - -pub struct KiroProvider { - pub credentials: KiroCredentials, - pub client: Client, - /// 当前加载的凭证文件路径 - pub creds_path: Option, -} - -impl Default for KiroProvider { - fn default() -> Self { - // 创建带超时配置的 HTTP 客户端 - // 参考 AIClient-2-API: AXIOS_TIMEOUT: 300000 (5分钟) - let client = Client::builder() - .connect_timeout(std::time::Duration::from_secs(30)) // 连接超时 30 秒 - .timeout(std::time::Duration::from_secs(300)) // 总超时 5 分钟 - .build() - .unwrap_or_else(|_| Client::new()); - - Self { - credentials: KiroCredentials::default(), - client, - creds_path: None, - } - } -} - -impl Clone for KiroProvider { - fn clone(&self) -> Self { - Self { - credentials: self.credentials.clone(), - client: reqwest::Client::new(), - creds_path: self.creds_path.clone(), - } - } -} - -impl KiroProvider { - pub fn new() -> Self { - Self::default() - } - - pub fn default_creds_path() -> PathBuf { - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".aws") - .join("sso") - .join("cache") - .join("kiro-auth-token.json") - } - - pub async fn load_credentials(&mut self) -> Result<(), Box> { - let path = Self::default_creds_path(); - let dir = path.parent().ok_or("Invalid path: no parent directory")?; - - let mut merged = KiroCredentials::default(); - - // 读取主凭证文件 - if tokio::fs::try_exists(&path).await.unwrap_or(false) { - let content = tokio::fs::read_to_string(&path).await?; - let creds: KiroCredentials = serde_json::from_str(&content)?; - tracing::info!( - "[KIRO] Main file loaded: has_access={}, has_refresh={}, has_client_id={}, auth_method={:?}", - creds.access_token.is_some(), - creds.refresh_token.is_some(), - creds.client_id.is_some(), - creds.auth_method - ); - merge_credentials(&mut merged, &creds); - } - - // 如果有 clientIdHash,尝试加载对应的 client_id 和 client_secret - if let Some(hash) = &merged.client_id_hash { - let hash_file_path = dir.join(format!("{hash}.json")); - tracing::info!( - "[KIRO] 检查 clientIdHash 文件: {}", - hash_file_path.display() - ); - if tokio::fs::try_exists(&hash_file_path) - .await - .unwrap_or(false) - { - if let Ok(content) = tokio::fs::read_to_string(&hash_file_path).await { - if let Ok(creds) = serde_json::from_str::(&content) { - tracing::info!( - "[KIRO] Hash file {:?}: has_client_id={}, has_client_secret={}", - hash_file_path.file_name(), - creds.client_id.is_some(), - creds.client_secret.is_some() - ); - merge_credentials(&mut merged, &creds); - } else { - tracing::error!( - "[KIRO] 无法解析 clientIdHash 文件: {}", - hash_file_path.display() - ); - } - } else { - tracing::error!( - "[KIRO] 无法读取 clientIdHash 文件: {}", - hash_file_path.display() - ); - } - } else { - tracing::warn!( - "[KIRO] clientIdHash {} 指向的文件不存在: {}", - hash, - hash_file_path.display() - ); - } - } else { - tracing::info!("[KIRO] 没有 clientIdHash 字段"); - } - - // 安全修复:不再遍历目录中其他 JSON 文件,避免串凭证/串账号风险 - // 只信任主凭证文件和 clientIdHash 指向的文件 - - tracing::info!( - "[KIRO] Final merged: has_access={}, has_refresh={}, has_client_id={}, has_client_secret={}, auth_method={:?}", - merged.access_token.is_some(), - merged.refresh_token.is_some(), - merged.client_id.is_some(), - merged.client_secret.is_some(), - merged.auth_method - ); - - self.credentials = merged; - self.creds_path = Some(path); - - // 加载完成后,智能检测并更新认证方式(如果需要) - let detected_auth_method = self.detect_auth_method(); - if self.credentials.auth_method.as_deref().unwrap_or("social") != detected_auth_method { - tracing::info!( - "[KIRO] 加载后检测到需要调整认证方式为: {}", - detected_auth_method - ); - self.set_auth_method(&detected_auth_method); - } - - Ok(()) - } - - /// 从指定路径加载凭证 - /// - /// 副本文件应包含完整的 client_id/client_secret(在复制时已合并)。 - /// 如果副本文件中没有,会尝试从 clientIdHash 文件读取作为回退。 - pub async fn load_credentials_from_path( - &mut self, - path: &str, - ) -> Result<(), Box> { - let path = std::path::PathBuf::from(path); - - let mut merged = KiroCredentials::default(); - - // 读取主凭证文件 - if tokio::fs::try_exists(&path).await.unwrap_or(false) { - let content = tokio::fs::read_to_string(&path).await?; - let creds: KiroCredentials = serde_json::from_str(&content)?; - tracing::info!( - "[KIRO] 加载凭证文件 {:?}: has_access={}, has_refresh={}, has_client_id={}, has_client_secret={}, auth_method={:?}", - path, - creds.access_token.is_some(), - creds.refresh_token.is_some(), - creds.client_id.is_some(), - creds.client_secret.is_some(), - creds.auth_method - ); - merge_credentials(&mut merged, &creds); - } else { - return Err(format!("凭证文件不存在: {path:?}").into()); - } - - // 如果副本文件中已有 client_id/client_secret,直接使用(方案B:完全独立) - if merged.client_id.is_some() && merged.client_secret.is_some() { - tracing::info!("[KIRO] 副本文件包含完整的 client_id/client_secret,无需读取外部文件"); - } else { - // 回退:尝试从外部文件读取(兼容旧的副本文件) - let aws_sso_cache_dir = dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".aws") - .join("sso") - .join("cache"); - - let mut found_credentials = false; - - // 方式1:如果有 clientIdHash,尝试从对应文件读取 - if let Some(hash) = &merged.client_id_hash.clone() { - tracing::info!( - "[KIRO] 副本文件缺少 client_id/client_secret,尝试从 clientIdHash 文件读取" - ); - let hash_file_path = aws_sso_cache_dir.join(format!("{hash}.json")); - - if tokio::fs::try_exists(&hash_file_path) - .await - .unwrap_or(false) - { - if let Ok(content) = tokio::fs::read_to_string(&hash_file_path).await { - if let Ok(json_value) = serde_json::from_str::(&content) - { - if merged.client_id.is_none() { - merged.client_id = json_value - .get("clientId") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - } - if merged.client_secret.is_none() { - merged.client_secret = json_value - .get("clientSecret") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - } - if merged.client_id.is_some() && merged.client_secret.is_some() { - found_credentials = true; - tracing::info!( - "[KIRO] 从 clientIdHash 文件补充: has_client_id={}, has_client_secret={}", - merged.client_id.is_some(), - merged.client_secret.is_some() - ); - } - } - } - } - } - - // 方式2:如果没有 clientIdHash 或未找到,扫描目录中的其他 JSON 文件 - if !found_credentials - && tokio::fs::try_exists(&aws_sso_cache_dir) - .await - .unwrap_or(false) - { - tracing::info!("[KIRO] 扫描 .aws/sso/cache 目录查找 client_id/client_secret"); - if let Ok(mut entries) = tokio::fs::read_dir(&aws_sso_cache_dir).await { - while let Ok(Some(entry)) = entries.next_entry().await { - let file_path = entry.path(); - if file_path.extension().map(|e| e == "json").unwrap_or(false) { - let file_name = - file_path.file_name().and_then(|n| n.to_str()).unwrap_or(""); - // 跳过主凭证文件和备份文件 - if file_name.starts_with("kiro-auth-token") { - continue; - } - if let Ok(content) = tokio::fs::read_to_string(&file_path).await { - if let Ok(json_value) = - serde_json::from_str::(&content) - { - let has_client_id = json_value - .get("clientId") - .and_then(|v| v.as_str()) - .is_some(); - let has_client_secret = json_value - .get("clientSecret") - .and_then(|v| v.as_str()) - .is_some(); - if has_client_id && has_client_secret { - merged.client_id = json_value - .get("clientId") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - merged.client_secret = json_value - .get("clientSecret") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - found_credentials = true; - tracing::info!( - "[KIRO] 从 {} 补充 client_id/client_secret", - file_name - ); - break; - } - } - } - } - } - } - } - - if !found_credentials { - tracing::warn!("[KIRO] 未找到 client_id/client_secret,将使用 social 认证"); - } - } - - tracing::info!( - "[KIRO] 最终凭证状态: has_access={}, has_refresh={}, has_client_id={}, has_client_secret={}, auth_method={:?}", - merged.access_token.is_some(), - merged.refresh_token.is_some(), - merged.client_id.is_some(), - merged.client_secret.is_some(), - merged.auth_method - ); - - self.credentials = merged; - self.creds_path = Some(path); - - // 加载完成后,智能检测并更新认证方式(如果需要) - let detected_auth_method = self.detect_auth_method(); - if self.credentials.auth_method.as_deref().unwrap_or("social") != detected_auth_method { - tracing::info!( - "[KIRO] 从路径加载后检测到需要调整认证方式为: {}", - detected_auth_method - ); - self.set_auth_method(&detected_auth_method); - } - - Ok(()) - } - - pub fn get_base_url(&self) -> String { - let region = self.credentials.region.as_deref().unwrap_or("us-east-1"); - format!("https://codewhisperer.{region}.amazonaws.com/generateAssistantResponse") - } - - pub fn get_refresh_url(&self) -> String { - let region = self.credentials.region.as_deref().unwrap_or("us-east-1"); - let auth_method = self - .credentials - .auth_method - .as_deref() - .unwrap_or("social") - .to_lowercase(); - - if auth_method == "idc" { - format!("https://oidc.{region}.amazonaws.com/token") - } else { - format!("https://prod.{region}.auth.desktop.kiro.dev/refreshToken") - } - } - - /// 构建健康检查使用的端点,与实际API调用保持一致 - pub fn get_health_check_url(&self) -> String { - // 重用基础URL逻辑,确保健康检查与实际API调用使用相同端点 - self.get_base_url() - } - - /// 从凭证文件中提取 region 信息的静态方法,供健康检查服务使用 - pub fn extract_region_from_creds(creds_content: &str) -> Result { - let creds: serde_json::Value = - serde_json::from_str(creds_content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let region = creds["region"].as_str().unwrap_or("us-east-1").to_string(); - - Ok(region) - } - - /// 构建健康检查端点的静态方法,供外部服务使用 - pub fn build_health_check_url(region: &str) -> String { - format!("https://codewhisperer.{region}.amazonaws.com/generateAssistantResponse") - } - - /// 检查 Token 是否已过期 - /// - /// 支持两种格式: - /// - RFC3339 格式(新格式,与 CLIProxyAPI 兼容) - /// - 时间戳格式(旧格式) - pub fn is_token_expired(&self) -> bool { - // 优先检查 RFC3339 格式的过期时间(新格式) - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) { - let now = chrono::Utc::now(); - // 提前5分钟判断为过期,避免边界情况 - return expires <= now + chrono::Duration::minutes(5); - } - } - - // 兼容旧的时间戳格式 - if let Some(expires_str) = &self.credentials.expires_at { - if let Ok(expires_timestamp) = expires_str.parse::() { - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs() as i64; - - // 提前5分钟判断为过期,避免边界情况 - return now >= (expires_timestamp - 300); - } - } - - // 如果没有过期时间信息,保守地认为可能需要刷新 - true - } - - /// 验证 refresh_token 的基本有效性 - pub fn validate_refresh_token(&self) -> Result<(), String> { - let refresh_token = self.credentials.refresh_token.as_ref() - .ok_or("缺少 refresh_token。\n💡 解决方案:\n1. 重新添加 OAuth 凭证\n2. 确保凭证文件包含完整的认证信息")?; - - // 基本格式验证 - if refresh_token.trim().is_empty() { - return Err("refresh_token 为空。\n💡 解决方案:\n1. 检查凭证文件是否损坏\n2. 重新生成 OAuth 凭证".to_string()); - } - - let token_len = refresh_token.len(); - - // 检测 refreshToken 是否被截断 - // 正常的 refreshToken 长度应该在 500+ 字符 - let is_truncated = - token_len < 100 || refresh_token.ends_with("...") || refresh_token.contains("..."); - - if is_truncated { - // 安全修复:不打印 token 内容,只打印长度 - tracing::error!("[KIRO] 检测到 refreshToken 被截断!长度: {}", token_len); - return Err(format!( - "refreshToken 已被截断(长度: {token_len} 字符)。\n\n⚠️ 这通常是 Kiro IDE 为了防止凭证被第三方工具使用而故意截断的。\n\n💡 解决方案:\n1. 使用 Kir-Manager 工具获取完整的凭证\n2. 或者使用其他方式获取未截断的凭证文件\n3. 正常的 refreshToken 长度应该在 500+ 字符" - )); - } - - // 检查是否看起来像有效的 token(简单的长度和格式检查) - if refresh_token.len() < 10 { - return Err("refresh_token 格式异常(长度过短)。\n💡 解决方案:\n1. 凭证文件可能已损坏\n2. 重新获取 OAuth 凭证".to_string()); - } - - Ok(()) - } - - /// 检测认证方式 - /// - /// 注意:不再自动降级!IdC 和 Social 的 refreshToken 不兼容, - /// 不能将 IdC 的 refreshToken 用于 Social 端点。 - pub fn detect_auth_method(&self) -> String { - // 直接返回配置中的认证方式,不做降级 - let auth_method = self.credentials.auth_method.as_deref().unwrap_or("social"); - tracing::debug!("[KIRO] 使用配置的认证方式: {}", auth_method); - auth_method.to_lowercase() - } - - /// 检查 IdC 认证配置是否完整 - pub fn is_idc_config_complete(&self) -> bool { - self.credentials.client_id.is_some() && self.credentials.client_secret.is_some() - } - - /// 更新认证方式到凭证中(仅在内存中,需要调用 save_credentials 持久化) - pub fn set_auth_method(&mut self, method: &str) { - let old_method = self.credentials.auth_method.as_deref().unwrap_or("social"); - if old_method != method { - tracing::info!("[KIRO] 认证方式从 {} 切换到 {}", old_method, method); - self.credentials.auth_method = Some(method.to_string()); - } - } - - pub async fn refresh_token(&mut self) -> Result> { - // 首先验证 refresh_token 的有效性 - self.validate_refresh_token()?; - - tracing::info!("[KIRO] 开始 Token 刷新流程"); - tracing::info!( - "[KIRO] 当前凭证状态: has_client_id={}, has_client_secret={}, auth_method={:?}", - self.credentials.client_id.is_some(), - self.credentials.client_secret.is_some(), - self.credentials.auth_method - ); - - // 先克隆必要的值,避免借用冲突 - let refresh_token = self - .credentials - .refresh_token - .as_ref() - .ok_or("No refresh token")? - .clone(); - - // 获取认证方式 - let auth_method = self.detect_auth_method(); - tracing::info!("[KIRO] 使用认证方式: {}", auth_method); - - // 检查 IdC 认证是否有完整配置 - if auth_method == "idc" && !self.is_idc_config_complete() { - let has_client_id = self.credentials.client_id.is_some(); - let has_client_secret = self.credentials.client_secret.is_some(); - - // IdC 认证缺少必要凭证,返回明确错误(不能降级到 social,因为 refreshToken 不兼容) - let missing = match (has_client_id, has_client_secret) { - (false, false) => "clientId 和 clientSecret", - (false, true) => "clientId", - (true, false) => "clientSecret", - _ => unreachable!(), - }; - - return Err(format!( - "IdC 认证配置不完整:缺少 {missing}。\n\n⚠️ 注意:IdC 凭证的 refreshToken 无法用于 Social 认证,必须提供完整的 IdC 配置。\n\n💡 解决方案:\n1. 删除当前凭证\n2. 重新从 Kiro IDE 获取最新的凭证文件(确保完成完整的 SSO 登录流程)\n3. 确保 ~/.aws/sso/cache/ 目录下有对应的 clientIdHash 文件\n4. 重新添加凭证到 Lime" - ).into()); - } - let refresh_url = self.get_refresh_url(); - - tracing::debug!( - "[KIRO] refresh_token: auth_method={}, refresh_url={}", - auth_method, - refresh_url - ); - tracing::debug!( - "[KIRO] has_client_id={}, has_client_secret={}", - self.credentials.client_id.is_some(), - self.credentials.client_secret.is_some() - ); - - // 获取设备指纹和版本号(用于 Social 认证的 User-Agent) - // 使用基于凭证的 Machine ID,确保每个账号有独立的指纹 - let machine_id = generate_machine_id_from_credentials( - self.credentials.profile_arn.as_deref(), - self.credentials.client_id.as_deref(), - ); - let kiro_version = get_kiro_version(); - - let resp = if auth_method == "idc" { - // IdC 认证使用 JSON 格式(参考 Kir-Manager 实现) - let client_id = self - .credentials - .client_id - .as_ref() - .ok_or("IdC 认证配置错误:缺少 client_id。建议删除后重新添加 OAuth 凭证")?; - let client_secret = self - .credentials - .client_secret - .as_ref() - .ok_or("IdC 认证配置错误:缺少 client_secret。建议删除后重新添加 OAuth 凭证")?; - - // 使用 JSON 格式发送请求(与 Kir-Manager 保持一致) - let body = serde_json::json!({ - "refreshToken": &refresh_token, - "clientId": client_id, - "clientSecret": client_secret, - "grantType": "refresh_token" - }); - - tracing::debug!("[KIRO] IdC 刷新请求体已构建"); - - // IdC 认证的 Headers(参考 Kir-Manager) - self.client - .post(&refresh_url) - .header("Content-Type", "application/json") - .header("Host", "oidc.us-east-1.amazonaws.com") - .header( - "x-amz-user-agent", - format!("aws-sdk-js/3.738.0 ua/2.1 os/other lang/js api/sso-oidc#3.738.0 m/E KiroIDE-{kiro_version}-{machine_id}"), - ) - .header("User-Agent", "node") - .header("Accept", "*/*") - .header("Connection", "close") - .json(&body) - .send() - .await? - } else { - // Social 认证使用简单的 JSON 格式(参考 Kir-Manager) - let body = serde_json::json!({ "refreshToken": &refresh_token }); - - // Social 认证的 Headers(参考 Kir-Manager) - self.client - .post(&refresh_url) - .header("User-Agent", format!("KiroIDE-{kiro_version}-{machine_id}")) - .header("Accept", "application/json, text/plain, */*") - .header("Accept-Encoding", "br, gzip, deflate") - .header("Content-Type", "application/json") - .header("Accept-Language", "*") - .header("Sec-Fetch-Mode", "cors") - .header("Connection", "close") - .json(&body) - .send() - .await? - }; - - tracing::info!("[KIRO] Token 刷新响应状态: {}", resp.status()); - - if !resp.status().is_success() { - let status = resp.status(); - let body_text = resp.text().await.unwrap_or_default(); - - tracing::warn!("[KIRO] Token 刷新失败: {} - {}", status, body_text); - - // 根据具体的HTTP状态码提供更友好的错误信息 - let error_msg = match status.as_u16() { - 401 => { - if body_text.contains("Bad credentials") || body_text.contains("invalid") { - format!("OAuth 凭证已过期或无效,需要重新认证。\n💡 解决方案:\n1. 删除当前 OAuth 凭证\n2. 重新添加 OAuth 凭证\n3. 确保使用最新的凭证文件\n\n技术详情:{status} {body_text}") - } else { - format!("认证失败,Token 可能已过期。\n💡 解决方案:\n1. 检查 AWS 账户状态\n2. 重新生成 OAuth 凭证\n3. 确保凭证文件格式正确\n\n技术详情:{status} {body_text}") - } - } - 403 => format!("权限不足,无法刷新 Token。\n💡 解决方案:\n1. 检查 AWS 账户权限\n2. 确保 OAuth 应用配置正确\n3. 联系管理员检查权限设置\n\n技术详情:{status} {body_text}"), - 429 => format!("请求过于频繁,已被限流。\n💡 解决方案:\n1. 等待 5-10 分钟后重试\n2. 减少 Token 刷新频率\n3. 检查是否有其他程序在同时使用\n\n技术详情:{status} {body_text}"), - 500..=599 => format!("服务器错误,AWS OAuth 服务暂时不可用。\n💡 解决方案:\n1. 稍后重试(通常几分钟后恢复)\n2. 检查 AWS 服务状态页面\n3. 如持续失败,联系 AWS 支持\n\n技术详情:{status} {body_text}"), - _ => format!("Token 刷新失败。\n💡 解决方案:\n1. 检查网络连接\n2. 确认凭证文件完整性\n3. 尝试重新添加凭证\n\n技术详情:{status} {body_text}") - }; - - return Err(error_msg.into()); - } - - let data: serde_json::Value = resp.json().await?; - - // AWS OIDC returns snake_case, social endpoint returns camelCase - let new_token = data["accessToken"] - .as_str() - .or_else(|| data["access_token"].as_str()) - .ok_or("No access token in response")?; - - self.credentials.access_token = Some(new_token.to_string()); - - // Handle both camelCase and snake_case response formats - if let Some(rt) = data["refreshToken"] - .as_str() - .or_else(|| data["refresh_token"].as_str()) - { - self.credentials.refresh_token = Some(rt.to_string()); - } - if let Some(arn) = data["profileArn"].as_str() { - self.credentials.profile_arn = Some(arn.to_string()); - } - - // 更新过期时间(如果响应中包含) - if let Some(expires_in) = data["expiresIn"] - .as_i64() - .or_else(|| data["expires_in"].as_i64()) - { - let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in); - self.credentials.expire = Some(expires_at.to_rfc3339()); - // 同时更新旧格式以保持兼容 - self.credentials.expires_at = Some(expires_at.timestamp().to_string()); - } - - // 更新最后刷新时间(RFC3339 格式) - self.credentials.last_refresh = Some(chrono::Utc::now().to_rfc3339()); - - // 保存更新后的凭证到文件 - self.save_credentials().await?; - - Ok(new_token.to_string()) - } - - pub async fn save_credentials(&self) -> Result<(), Box> { - // 使用加载时的路径或默认路径 - let path = self - .creds_path - .clone() - .unwrap_or_else(Self::default_creds_path); - - // 读取现有文件内容 - let mut existing: serde_json::Value = if tokio::fs::try_exists(&path).await.unwrap_or(false) - { - let content = tokio::fs::read_to_string(&path).await?; - serde_json::from_str(&content).unwrap_or(serde_json::json!({})) - } else { - serde_json::json!({}) - }; - - // 更新字段 - if let Some(token) = &self.credentials.access_token { - existing["accessToken"] = serde_json::json!(token); - } - if let Some(token) = &self.credentials.refresh_token { - existing["refreshToken"] = serde_json::json!(token); - } - if let Some(arn) = &self.credentials.profile_arn { - existing["profileArn"] = serde_json::json!(arn); - } - - // 添加统一凭证格式字段(与 CLIProxyAPI 兼容) - existing["type"] = serde_json::json!(self.credentials.cred_type); - if let Some(expire) = &self.credentials.expire { - existing["expire"] = serde_json::json!(expire); - } - if let Some(last_refresh) = &self.credentials.last_refresh { - existing["lastRefresh"] = serde_json::json!(last_refresh); - } - - // 写回文件 - let content = serde_json::to_string_pretty(&existing)?; - tokio::fs::write(&path, content).await?; - - Ok(()) - } - - /// 检查 token 是否即将过期(10 分钟内) - /// - /// 支持两种格式: - /// - RFC3339 格式(新格式,与 CLIProxyAPI 兼容) - /// - 时间戳格式(旧格式) - pub fn is_token_expiring_soon(&self) -> bool { - // 优先检查 RFC3339 格式的过期时间(新格式) - if let Some(expire_str) = &self.credentials.expire { - if let Ok(expiry) = chrono::DateTime::parse_from_rfc3339(expire_str) { - let now = chrono::Utc::now(); - let threshold = now + chrono::Duration::minutes(10); - return expiry < threshold; - } - } - - // 兼容旧格式(expires_at 可能是 RFC3339 或时间戳) - if let Some(expires_at) = &self.credentials.expires_at { - // 尝试解析为 RFC3339 - if let Ok(expiry) = chrono::DateTime::parse_from_rfc3339(expires_at) { - let now = chrono::Utc::now(); - let threshold = now + chrono::Duration::minutes(10); - return expiry < threshold; - } - // 尝试解析为时间戳 - if let Ok(expires_timestamp) = expires_at.parse::() { - let now = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs() as i64; - return now >= (expires_timestamp - 600); // 10 分钟 = 600 秒 - } - } - // 如果没有过期时间,假设不需要刷新 - false - } - - pub async fn call_api( - &self, - request: &ChatCompletionRequest, - ) -> Result> { - self.call_api_with_conversation_id(request, None).await - } - - pub async fn call_api_with_conversation_id( - &self, - request: &ChatCompletionRequest, - conversation_id: Option<&str>, - ) -> Result> { - let token = self - .credentials - .access_token - .as_ref() - .ok_or("No access token")?; - - let profile_arn = if self.credentials.auth_method.as_deref() == Some("social") { - self.credentials.profile_arn.clone() - } else { - None - }; - - let cw_request = convert_openai_to_codewhisperer_with_conversation_id( - request, - profile_arn.clone(), - conversation_id, - ); - let url = self.get_base_url(); - - // 安全修复:仅在 LIME_DEBUG=1 时写入请求调试文件,兼容旧的 PROXYCAST_DEBUG。 - let debug_enabled = - lime_core::env_compat::bool_var(&["LIME_DEBUG", "PROXYCAST_DEBUG"]).unwrap_or(false); - if debug_enabled { - if let Ok(json_str) = serde_json::to_string_pretty(&cw_request) { - let debug_dir = lime_core::app_paths::resolve_logs_dir() - .unwrap_or_else(|_| std::env::temp_dir().join("lime").join("logs")); - let uuid_prefix = uuid::Uuid::new_v4() - .to_string() - .split('-') - .next() - .unwrap_or("unknown") - .to_string(); - let debug_path = debug_dir.join(format!("cw_request_{uuid_prefix}.json")); - let _ = tokio::fs::create_dir_all(&debug_dir).await; - if tokio::fs::write(&debug_path, &json_str).await.is_ok() { - prune_kiro_debug_request_files(&debug_dir).await; - tracing::debug!("[CW_REQ] Request saved to {:?}", debug_path); - } - } - } - - // 记录历史消息数量和 tool_results 情况(不落盘) - let history_len = cw_request - .conversation_state - .history - .as_ref() - .map(|h| h.len()) - .unwrap_or(0); - let current_has_tools = cw_request - .conversation_state - .current_message - .user_input_message - .user_input_message_context - .as_ref() - .map(|ctx| ctx.tool_results.as_ref().map(|tr| tr.len()).unwrap_or(0)) - .unwrap_or(0); - tracing::info!( - "[CW_REQ] history={} current_tool_results={}", - history_len, - current_has_tools - ); - - // 生成基于凭证的唯一 Machine ID(关键改进:每个账号独立指纹) - let machine_id = generate_machine_id_from_credentials( - profile_arn.as_deref(), - self.credentials.client_id.as_deref(), - ); - let kiro_version = get_kiro_version(); - let (os_name, node_version) = get_system_runtime_info(); - - tracing::debug!( - "[KIRO_FINGERPRINT] machine_id={} (based on profile_arn={}, client_id={})", - &machine_id[..16], - profile_arn.is_some(), - self.credentials.client_id.is_some() - ); - - let resp = self - .client - .post(&url) - .header("Authorization", format!("Bearer {token}")) - .header("Content-Type", "application/json") - .header("Accept", "application/json") - .header("amz-sdk-invocation-id", uuid::Uuid::new_v4().to_string()) - .header("amz-sdk-request", "attempt=1; max=1") - .header("x-amzn-kiro-agent-mode", "vibe") - // 关键指纹头:使用基于凭证的唯一 Machine ID - .header( - "x-amz-user-agent", - format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}"), - ) - .header( - "user-agent", - format!( - "aws-sdk-js/1.0.0 ua/2.1 os/{os_name} lang/js md/nodejs#{node_version} api/codewhispererruntime#1.0.0 m/E KiroIDE-{kiro_version}-{machine_id}" - ), - ) - // 添加 Connection: close 避免连接复用被检测 - .header("Connection", "close") - .json(&cw_request) - .send() - .await?; - - Ok(resp) - } - - pub async fn call_api_stream_with_conversation_id( - &self, - request: &ChatCompletionRequest, - conversation_id: Option<&str>, - ) -> Result { - let token = self - .credentials - .access_token - .as_ref() - .ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?; - - let profile_arn = if self.credentials.auth_method.as_deref() == Some("social") { - self.credentials.profile_arn.clone() - } else { - None - }; - - let cw_request = convert_openai_to_codewhisperer_with_conversation_id( - request, - profile_arn.clone(), - conversation_id, - ); - let url = self.get_base_url(); - - // 生成基于凭证的唯一 Machine ID - let machine_id = generate_machine_id_from_credentials( - profile_arn.as_deref(), - self.credentials.client_id.as_deref(), - ); - let kiro_version = get_kiro_version(); - let (os_name, node_version) = get_system_runtime_info(); - - tracing::info!( - "[KIRO_STREAM] 发起流式请求: url={} machine_id={}...", - url, - &machine_id[..16] - ); - - let resp = self - .client - .post(&url) - .header("Authorization", format!("Bearer {token}")) - .header("Content-Type", "application/json") - .header("Accept", "application/vnd.amazon.eventstream") - .header("amz-sdk-invocation-id", uuid::Uuid::new_v4().to_string()) - .header("amz-sdk-request", "attempt=1; max=1") - .header("x-amzn-kiro-agent-mode", "vibe") - .header( - "x-amz-user-agent", - format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}"), - ) - .header( - "user-agent", - format!( - "aws-sdk-js/1.0.0 ua/2.1 os/{os_name} lang/js md/nodejs#{node_version} api/codewhispererruntime#1.0.0 m/E KiroIDE-{kiro_version}-{machine_id}" - ), - ) - // 注意:不要设置 Connection: close,否则会导致流式响应无法工作 - .json(&cw_request) - .send() - .await - .map_err(|e| { - tracing::error!("[KIRO_STREAM] 请求发送失败: {}", e); - ProviderError::from_reqwest_error(&e) - })?; - - tracing::info!("[KIRO_STREAM] 收到响应: status={}", resp.status()); - - // 检查响应状态 - let status = resp.status(); - if !status.is_success() { - let body = resp.text().await.unwrap_or_default(); - tracing::error!("[KIRO_STREAM] 请求失败: {} - {}", status, body); - return Err(ProviderError::from_http_status(status.as_u16(), &body)); - } - - tracing::info!("[KIRO_STREAM] 流式响应开始: status={}", status); - - // 将 reqwest 响应转换为 StreamResponse - Ok(reqwest_stream_to_stream_response(resp)) - } -} - -fn merge_credentials(target: &mut KiroCredentials, source: &KiroCredentials) { - if source.access_token.is_some() { - target.access_token = source.access_token.clone(); - } - if source.refresh_token.is_some() { - target.refresh_token = source.refresh_token.clone(); - } - if source.client_id.is_some() { - target.client_id = source.client_id.clone(); - } - if source.client_secret.is_some() { - target.client_secret = source.client_secret.clone(); - } - if source.profile_arn.is_some() { - target.profile_arn = source.profile_arn.clone(); - } - if source.expires_at.is_some() { - target.expires_at = source.expires_at.clone(); - } - if source.expire.is_some() { - target.expire = source.expire.clone(); - } - if source.region.is_some() { - target.region = source.region.clone(); - } - if source.auth_method.is_some() { - target.auth_method = source.auth_method.clone(); - } - if source.client_id_hash.is_some() { - target.client_id_hash = source.client_id_hash.clone(); - } - if source.last_refresh.is_some() { - target.last_refresh = source.last_refresh.clone(); - } - // cred_type 使用默认值,不需要合并 -} - -// ============================================================================ -// CredentialProvider Trait 实现 -// ============================================================================ - -#[async_trait] -impl CredentialProvider for KiroProvider { - async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()> { - // 调用已有的实现 - KiroProvider::load_credentials_from_path(self, path).await - } - - async fn save_credentials(&self) -> ProviderResult<()> { - KiroProvider::save_credentials(self).await - } - - fn is_token_valid(&self) -> bool { - !self.is_token_expired() - } - - fn is_token_expiring_soon(&self) -> bool { - KiroProvider::is_token_expiring_soon(self) - } - - async fn refresh_token(&mut self) -> ProviderResult { - KiroProvider::refresh_token(self).await - } - - fn get_access_token(&self) -> Option<&str> { - self.credentials.access_token.as_deref() - } - - fn provider_type(&self) -> &'static str { - "kiro" - } -} - -// ============================================================================ -// StreamingProvider Trait 实现 -// ============================================================================ - -use crate::providers::ProviderError; -use crate::streaming::traits::{ - reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider, -}; - -#[async_trait] -impl StreamingProvider for KiroProvider { - /// 发起流式 API 调用 - /// - /// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。 - /// Kiro/CodeWhisperer 使用 AWS Event Stream 格式。 - /// - /// # 需求覆盖 - /// - 需求 1.1: KiroProvider 流式支持 - async fn call_api_stream( - &self, - request: &ChatCompletionRequest, - ) -> Result { - self.call_api_stream_with_conversation_id(request, None) - .await - } - - fn supports_streaming(&self) -> bool { - true - } - - fn provider_name(&self) -> &'static str { - "KiroProvider" - } - - fn stream_format(&self) -> StreamFormat { - StreamFormat::AwsEventStream - } -} - -// ============================================================================ -// Anthropic 格式直接支持 -// ============================================================================ - -impl KiroProvider { - /// 直接处理 Anthropic 格式的流式请求 - /// - /// 绕过 OpenAI 中间格式,直接从 Anthropic → CodeWhisperer - /// 这样可以保留 Anthropic 特有的字段(如 tool_choice) - pub async fn call_api_stream_anthropic( - &self, - request: &AnthropicMessagesRequest, - ) -> Result { - let token = self - .credentials - .access_token - .as_ref() - .ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?; - - let profile_arn = if self.credentials.auth_method.as_deref() == Some("social") { - self.credentials.profile_arn.clone() - } else { - None - }; - - // 直接转换 Anthropic → CodeWhisperer(不经过 OpenAI) - let cw_request = convert_anthropic_to_codewhisperer(request, profile_arn.clone()); - let url = self.get_base_url(); - - // 生成基于凭证的唯一 Machine ID - let machine_id = generate_machine_id_from_credentials( - profile_arn.as_deref(), - self.credentials.client_id.as_deref(), - ); - let kiro_version = get_kiro_version(); - let (os_name, node_version) = get_system_runtime_info(); - - tracing::info!( - "[KIRO_STREAM_ANTHROPIC] 直接 Anthropic→CodeWhisperer 流式请求: url={} machine_id={}...", - url, - &machine_id[..16] - ); - - let resp = self - .client - .post(&url) - .header("Authorization", format!("Bearer {token}")) - .header("Content-Type", "application/json") - .header("Accept", "application/vnd.amazon.eventstream") - .header("amz-sdk-invocation-id", uuid::Uuid::new_v4().to_string()) - .header("amz-sdk-request", "attempt=1; max=1") - .header("x-amzn-kiro-agent-mode", "vibe") - .header( - "x-amz-user-agent", - format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}"), - ) - .header( - "user-agent", - format!( - "aws-sdk-js/1.0.0 ua/2.1 os/{os_name} lang/js md/nodejs#{node_version} api/codewhispererruntime#1.0.0 m/E KiroIDE-{kiro_version}-{machine_id}" - ), - ) - .json(&cw_request) - .send() - .await - .map_err(|e| { - tracing::error!("[KIRO_STREAM_ANTHROPIC] 请求发送失败: {}", e); - ProviderError::from_reqwest_error(&e) - })?; - - tracing::info!("[KIRO_STREAM_ANTHROPIC] 收到响应: status={}", resp.status()); - - // 检查响应状态 - let status = resp.status(); - if !status.is_success() { - let body = resp.text().await.unwrap_or_default(); - tracing::error!("[KIRO_STREAM_ANTHROPIC] 请求失败: {} - {}", status, body); - return Err(ProviderError::from_http_status(status.as_u16(), &body)); - } - - tracing::info!("[KIRO_STREAM_ANTHROPIC] 流式响应开始: status={}", status); - - // 将 reqwest 响应转换为 StreamResponse - Ok(reqwest_stream_to_stream_response(resp)) - } -} diff --git a/src-tauri/crates/providers/src/providers/mod.rs b/src-tauri/crates/providers/src/providers/mod.rs index 328401ec3..99a1fb19e 100644 --- a/src-tauri/crates/providers/src/providers/mod.rs +++ b/src-tauri/crates/providers/src/providers/mod.rs @@ -1,40 +1,21 @@ -pub mod antigravity; pub mod claude_custom; -pub mod claude_oauth; pub mod codex; pub mod error; pub mod gemini; -pub mod kiro; pub mod novita; pub mod openai_custom; -pub mod traits; pub mod vertex; #[cfg(test)] mod tests; -// Trait exports -#[allow(unused_imports)] -pub use traits::{CredentialProvider, ProviderResult, TokenManager}; - -#[allow(unused_imports)] -pub use antigravity::AntigravityApiError; -#[allow(unused_imports)] -pub use antigravity::AntigravityProvider; -#[allow(unused_imports)] -pub use antigravity::ANTIGRAVITY_MODELS_FALLBACK; -#[allow(unused_imports)] pub use claude_custom::{ClaudeCustomProvider, PromptCacheMode}; #[allow(unused_imports)] -pub use claude_oauth::ClaudeOAuthProvider; -#[allow(unused_imports)] pub use codex::CodexProvider; #[allow(unused_imports)] pub use error::ProviderError; #[allow(unused_imports)] -pub use gemini::{GeminiApiKeyCredential, GeminiApiKeyProvider, GeminiProvider}; -#[allow(unused_imports)] -pub use kiro::KiroProvider; +pub use gemini::{GeminiApiKeyCredential, GeminiApiKeyProvider}; #[allow(unused_imports)] pub use novita::{ NovitaProvider, NOVITA_API_BASE_URL, NOVITA_DEFAULT_MODEL, NOVITA_EMBEDDING_MODEL, diff --git a/src-tauri/crates/providers/src/providers/tests.rs b/src-tauri/crates/providers/src/providers/tests.rs index dbfb6fd7f..4d6f862db 100644 --- a/src-tauri/crates/providers/src/providers/tests.rs +++ b/src-tauri/crates/providers/src/providers/tests.rs @@ -2,200 +2,14 @@ //! //! 使用 proptest 进行属性测试 -use chrono::{Duration, Utc}; use proptest::prelude::*; use crate::providers::codex::CodexProvider; use crate::providers::vertex::VertexProvider; -/// Generate a random lead time in minutes (1 to 30 minutes) -fn arb_lead_time_mins() -> impl Strategy { - 1i64..30i64 -} - -/// 生成不会与 lead_time 边界冲突的时间偏移 -/// 避免 time_offset_secs 恰好等于 lead_time_mins * 60 的情况 -#[allow(dead_code)] -fn arb_time_offset_avoiding_boundary(lead_time_mins: i64) -> impl Strategy { - let boundary = lead_time_mins * 60; - // 生成不等于边界值的时间偏移 - (-3600i64..7200i64).prop_filter("避免边界值", move |&offset| offset != boundary) -} - proptest! { #![proptest_config(ProptestConfig::with_cases(100))] - /// **Feature: cliproxyapi-parity, Property 2: Token Refresh Timing** - /// *For any* stored OAuth token with expiration time T, the refresh mechanism - /// SHALL be triggered before time T. - /// **Validates: Requirements 1.2, 2.2** - /// - /// This test verifies that: - /// 1. When token expires within lead_time, needs_refresh returns true - /// 2. When token expires after lead_time, needs_refresh returns false - /// 3. When no token exists, needs_refresh returns true - /// 4. When no expiry info exists, needs_refresh returns true - #[test] - fn test_codex_token_refresh_timing( - lead_time_mins in arb_lead_time_mins(), - time_offset_secs in -3600i64..7200i64, - ) { - let lead_time = Duration::minutes(lead_time_mins); - let lead_time_secs = lead_time_mins * 60; - - // 跳过边界条件,因为时间精度问题可能导致不确定行为 - prop_assume!(time_offset_secs != lead_time_secs); - - let mut provider = CodexProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - - // Set expiration time relative to now - let now = Utc::now(); - let expires_at = now + Duration::seconds(time_offset_secs); - provider.credentials.expires_at = Some(expires_at.to_rfc3339()); - - let needs_refresh = provider.needs_refresh(lead_time); - - // Token should need refresh if it expires within lead_time from now - // i.e., expires_at < now + lead_time - // i.e., time_offset_secs < lead_time_secs - let expected_needs_refresh = time_offset_secs < lead_time_secs; - - prop_assert_eq!( - needs_refresh, - expected_needs_refresh, - "Codex: Token with expiry in {} seconds should {} refresh with lead time of {} minutes", - time_offset_secs, - if expected_needs_refresh { "need" } else { "not need" }, - lead_time_mins - ); - } - - /// **Feature: cliproxyapi-parity, Property 2: Token Refresh Timing** - /// Test that Codex OAuth tokens trigger refresh before expiration (second test) - /// **Validates: Requirements 1.2, 2.2** - #[test] - fn test_codex_token_refresh_timing_second( - lead_time_mins in arb_lead_time_mins(), - time_offset_secs in -3600i64..7200i64, - ) { - let lead_time = Duration::minutes(lead_time_mins); - let lead_time_secs = lead_time_mins * 60; - - // 跳过边界条件,因为时间精度问题可能导致不确定行为 - prop_assume!(time_offset_secs != lead_time_secs); - - let mut provider = CodexProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - - // Set expiration time relative to now - let now = Utc::now(); - let expires_at = now + Duration::seconds(time_offset_secs); - provider.credentials.expires_at = Some(expires_at.to_rfc3339()); - - let needs_refresh = provider.needs_refresh(lead_time); - - // Token should need refresh if it expires within lead_time from now - let expected_needs_refresh = time_offset_secs < lead_time_secs; - - prop_assert_eq!( - needs_refresh, - expected_needs_refresh, - "Codex: Token with expiry in {} seconds should {} refresh with lead time of {} minutes", - time_offset_secs, - if expected_needs_refresh { "need" } else { "not need" }, - lead_time_mins - ); - } - - /// **Feature: cliproxyapi-parity, Property 2: Token Refresh Timing** - /// Test that missing access token always triggers refresh - /// **Validates: Requirements 1.2, 2.2** - #[test] - fn test_codex_missing_token_needs_refresh( - lead_time_mins in arb_lead_time_mins(), - ) { - let lead_time = Duration::minutes(lead_time_mins); - - let provider = CodexProvider::new(); - // No access token set - - let needs_refresh = provider.needs_refresh(lead_time); - - prop_assert!( - needs_refresh, - "Codex: Missing access token should always need refresh" - ); - } - - /// **Feature: cliproxyapi-parity, Property 2: Token Refresh Timing** - /// Test that missing expiry info triggers refresh - /// **Validates: Requirements 1.2, 2.2** - #[test] - fn test_codex_missing_expiry_needs_refresh( - lead_time_mins in arb_lead_time_mins(), - ) { - let lead_time = Duration::minutes(lead_time_mins); - - let mut provider = CodexProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - // No expires_at set - - let needs_refresh = provider.needs_refresh(lead_time); - - prop_assert!( - needs_refresh, - "Codex: Missing expiry info should always need refresh" - ); - } - - /// **Feature: cliproxyapi-parity, Property 2: Token Refresh Timing** - /// Test that refresh is triggered strictly before expiration time - /// This ensures the invariant: if needs_refresh(lead_time) is false, - /// then the token will not expire within lead_time duration - /// **Validates: Requirements 1.2, 2.2** - #[test] - fn test_refresh_timing_invariant( - lead_time_mins in arb_lead_time_mins(), - extra_buffer_secs in 1i64..60i64, - ) { - let lead_time = Duration::minutes(lead_time_mins); - let lead_time_secs = lead_time_mins * 60; - - let mut provider = CodexProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - - // Set expiration time to exactly lead_time + extra_buffer from now - // This should NOT need refresh - let now = Utc::now(); - let expires_at = now + Duration::seconds(lead_time_secs + extra_buffer_secs); - provider.credentials.expires_at = Some(expires_at.to_rfc3339()); - - let needs_refresh = provider.needs_refresh(lead_time); - - prop_assert!( - !needs_refresh, - "Token expiring in {} seconds (lead_time={} mins + {} secs buffer) should not need refresh", - lead_time_secs + extra_buffer_secs, - lead_time_mins, - extra_buffer_secs - ); - - // Verify the invariant: if needs_refresh is false, the token expires after lead_time - // Note: is_token_expired() uses a hardcoded 5-minute buffer, which is different from needs_refresh - // So we verify the actual expiration time instead - if let Some(expires_str) = &provider.credentials.expires_at { - if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expires_str) { - let expires_utc = expires.with_timezone(&Utc); - let now = Utc::now(); - prop_assert!( - expires_utc >= now + lead_time, - "Token should not expire within lead_time when needs_refresh is false" - ); - } - } - } - /// **Feature: cliproxyapi-parity, Property 3: Provider Routing Correctness** /// *For any* request with model name M and provider type P, the router SHALL select /// a credential of type P that supports model M. @@ -853,44 +667,3 @@ proptest! { } } } - -#[cfg(test)] -mod unit_tests { - use super::*; - - #[test] - fn test_codex_needs_refresh_boundary() { - let mut provider = CodexProvider::new(); - provider.credentials.access_token = Some("test_token".to_string()); - - let lead_time = Duration::minutes(5); - - // Token expiring well after lead_time - should NOT need refresh - // Use a large buffer to avoid timing issues - let now = Utc::now(); - let expires_at = now + Duration::minutes(10); - provider.credentials.expires_at = Some(expires_at.to_rfc3339()); - assert!( - !provider.needs_refresh(lead_time), - "Token expiring in 10 mins should not need refresh with 5 min lead time" - ); - - // Token expiring well before lead_time - should need refresh - let now = Utc::now(); - let expires_at = now + Duration::minutes(2); - provider.credentials.expires_at = Some(expires_at.to_rfc3339()); - assert!( - provider.needs_refresh(lead_time), - "Token expiring in 2 mins should need refresh with 5 min lead time" - ); - - // Token already expired - should need refresh - let now = Utc::now(); - let expires_at = now - Duration::minutes(1); - provider.credentials.expires_at = Some(expires_at.to_rfc3339()); - assert!( - provider.needs_refresh(lead_time), - "Expired token should need refresh" - ); - } -} diff --git a/src-tauri/crates/providers/src/providers/traits.rs b/src-tauri/crates/providers/src/providers/traits.rs deleted file mode 100644 index 70c922347..000000000 --- a/src-tauri/crates/providers/src/providers/traits.rs +++ /dev/null @@ -1,184 +0,0 @@ -//! Provider Trait 定义 -//! -//! 统一的 Provider 接口,用于凭证管理和 Token 生命周期管理。 - -#![allow(dead_code)] - -use async_trait::async_trait; -use std::error::Error; - -/// Provider 结果类型别名(与现有方法签名兼容) -pub type ProviderResult = Result>; - -/// 凭证管理 Trait -/// -/// 定义所有 OAuth Provider 必须实现的凭证管理接口 -#[async_trait] -pub trait CredentialProvider: Send + Sync { - /// 从指定路径加载凭证 - async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()>; - - /// 保存凭证到文件 - async fn save_credentials(&self) -> ProviderResult<()>; - - /// 检查 Token 是否有效(未过期) - fn is_token_valid(&self) -> bool; - - /// 检查 Token 是否即将过期(通常提前 5 分钟) - fn is_token_expiring_soon(&self) -> bool; - - /// 刷新 Token - /// - /// 返回新的 access_token - async fn refresh_token(&mut self) -> ProviderResult; - - /// 获取当前 access_token - fn get_access_token(&self) -> Option<&str>; - - /// 获取 Provider 类型名称 - fn provider_type(&self) -> &'static str; -} - -/// Token 管理辅助 Trait -/// -/// 提供带重试的 Token 刷新功能 -#[async_trait] -pub trait TokenManager: CredentialProvider { - /// 带重试的 Token 刷新 - /// - /// # Arguments - /// * `max_retries` - 最大重试次数 - /// * `retry_delay_ms` - 重试间隔(毫秒) - async fn refresh_token_with_retry( - &mut self, - max_retries: u32, - retry_delay_ms: u64, - ) -> ProviderResult { - let mut last_error = None; - - for attempt in 0..=max_retries { - match self.refresh_token().await { - Ok(token) => return Ok(token), - Err(e) => { - tracing::warn!( - "[{}] Token refresh attempt {} failed: {}", - self.provider_type(), - attempt + 1, - e - ); - last_error = Some(e); - if attempt < max_retries { - tokio::time::sleep(tokio::time::Duration::from_millis(retry_delay_ms)) - .await; - } - } - } - } - - Err(last_error.unwrap_or_else(|| "Token refresh failed".into())) - } - - /// 确保 Token 有效(如需要则刷新) - async fn ensure_valid_token(&mut self) -> ProviderResult { - if !self.is_token_valid() || self.is_token_expiring_soon() { - self.refresh_token().await - } else { - self.get_access_token() - .map(|s| s.to_string()) - .ok_or_else(|| "No access token available".into()) - } - } -} - -// 为所有实现了 CredentialProvider 的类型自动实现 TokenManager -impl TokenManager for T {} - -#[cfg(test)] -mod tests { - use super::*; - - // Mock Provider for testing - struct MockProvider { - token: Option, - valid: bool, - expiring_soon: bool, - refresh_count: u32, - } - - #[async_trait] - impl CredentialProvider for MockProvider { - async fn load_credentials_from_path(&mut self, _path: &str) -> ProviderResult<()> { - Ok(()) - } - - async fn save_credentials(&self) -> ProviderResult<()> { - Ok(()) - } - - fn is_token_valid(&self) -> bool { - self.valid - } - - fn is_token_expiring_soon(&self) -> bool { - self.expiring_soon - } - - async fn refresh_token(&mut self) -> ProviderResult { - self.refresh_count += 1; - self.token = Some(format!("new_token_{}", self.refresh_count)); - self.valid = true; - self.expiring_soon = false; - Ok(self.token.clone().unwrap()) - } - - fn get_access_token(&self) -> Option<&str> { - self.token.as_deref() - } - - fn provider_type(&self) -> &'static str { - "mock" - } - } - - #[tokio::test] - async fn test_ensure_valid_token_when_valid() { - let mut provider = MockProvider { - token: Some("existing_token".to_string()), - valid: true, - expiring_soon: false, - refresh_count: 0, - }; - - let token = provider.ensure_valid_token().await.unwrap(); - assert_eq!(token, "existing_token"); - assert_eq!(provider.refresh_count, 0); - } - - #[tokio::test] - async fn test_ensure_valid_token_when_expiring() { - let mut provider = MockProvider { - token: Some("old_token".to_string()), - valid: true, - expiring_soon: true, - refresh_count: 0, - }; - - let token = provider.ensure_valid_token().await.unwrap(); - assert_eq!(token, "new_token_1"); - assert_eq!(provider.refresh_count, 1); - } - - #[tokio::test] - async fn test_ensure_valid_token_when_invalid() { - let mut provider = MockProvider { - token: Some("invalid_token".to_string()), - valid: false, - expiring_soon: false, - refresh_count: 0, - }; - - let token = provider.ensure_valid_token().await.unwrap(); - assert_eq!(token, "new_token_1"); - assert_eq!(provider.refresh_count, 1); - } -} diff --git a/src-tauri/crates/providers/src/stream/pipeline.rs b/src-tauri/crates/providers/src/stream/pipeline.rs index 0238e5959..313142b08 100644 --- a/src-tauri/crates/providers/src/stream/pipeline.rs +++ b/src-tauri/crates/providers/src/stream/pipeline.rs @@ -7,7 +7,7 @@ //! ```ignore //! use lime::stream::pipeline::{StreamPipeline, PipelineConfig}; //! -//! let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); +//! let config = PipelineConfig::aws_event_stream_to_anthropic("claude-sonnet-4-5".to_string()); //! let pipeline = StreamPipeline::new(config); //! //! // 处理字节流 @@ -23,8 +23,8 @@ use futures::{Stream, StreamExt}; /// 后端类型 #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum BackendType { - /// Kiro/CodeWhisperer (AWS Event Stream) - Kiro, + /// AWS Event Stream + AwsEventStream, /// OpenAI (SSE) OpenAi, /// Anthropic (SSE) @@ -54,20 +54,20 @@ pub struct PipelineConfig { } impl PipelineConfig { - /// 创建 Kiro → Anthropic 配置 - pub fn kiro_to_anthropic(model: String) -> Self { + /// 创建 AWS Event Stream → Anthropic 配置 + pub fn aws_event_stream_to_anthropic(model: String) -> Self { Self { - backend: BackendType::Kiro, + backend: BackendType::AwsEventStream, frontend: FrontendType::Anthropic, model, message_id: None, } } - /// 创建 Kiro → OpenAI 配置 - pub fn kiro_to_openai(model: String) -> Self { + /// 创建 AWS Event Stream → OpenAI 配置 + pub fn aws_event_stream_to_openai(model: String) -> Self { Self { - backend: BackendType::Kiro, + backend: BackendType::AwsEventStream, frontend: FrontendType::OpenAi, model, message_id: None, @@ -102,7 +102,7 @@ impl SseGenerator { pub struct StreamPipeline { /// 配置 config: PipelineConfig, - /// AWS Event Stream 解析器(用于 Kiro 后端) + /// AWS Event Stream 解析器 aws_parser: Option, /// SSE 生成器 generator: SseGenerator, @@ -112,7 +112,9 @@ impl StreamPipeline { /// 创建新的管道 pub fn new(config: PipelineConfig) -> Self { let aws_parser = match config.backend { - BackendType::Kiro => Some(AwsEventStreamParser::with_model(config.model.clone())), + BackendType::AwsEventStream => { + Some(AwsEventStreamParser::with_model(config.model.clone())) + } _ => None, }; @@ -256,26 +258,26 @@ mod tests { use super::*; #[test] - fn test_pipeline_config_kiro_to_anthropic() { - let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); - assert_eq!(config.backend, BackendType::Kiro); + fn test_pipeline_config_aws_event_stream_to_anthropic() { + let config = PipelineConfig::aws_event_stream_to_anthropic("claude-sonnet-4-5".to_string()); + assert_eq!(config.backend, BackendType::AwsEventStream); assert_eq!(config.frontend, FrontendType::Anthropic); assert_eq!(config.model, "claude-sonnet-4-5"); } #[test] - fn test_pipeline_config_kiro_to_openai() { - let config = PipelineConfig::kiro_to_openai("gpt-4".to_string()); - assert_eq!(config.backend, BackendType::Kiro); + fn test_pipeline_config_aws_event_stream_to_openai() { + let config = PipelineConfig::aws_event_stream_to_openai("gpt-4".to_string()); + assert_eq!(config.backend, BackendType::AwsEventStream); assert_eq!(config.frontend, FrontendType::OpenAi); } #[test] fn test_pipeline_process_content() { - let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); + let config = PipelineConfig::aws_event_stream_to_anthropic("claude-sonnet-4-5".to_string()); let mut pipeline = StreamPipeline::new(config); - // 模拟 Kiro 内容事件 + // 模拟 AWS Event Stream 内容事件 let bytes = br#"{"content":"Hello"}"#; let sse = pipeline.process_chunk(bytes); @@ -288,7 +290,7 @@ mod tests { #[test] fn test_pipeline_process_tool_use() { - let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); + let config = PipelineConfig::aws_event_stream_to_anthropic("claude-sonnet-4-5".to_string()); let mut pipeline = StreamPipeline::new(config); // 工具调用开始 @@ -312,10 +314,10 @@ mod tests { #[test] fn test_pipeline_openai_output() { - let config = PipelineConfig::kiro_to_openai("gpt-4".to_string()); + let config = PipelineConfig::aws_event_stream_to_openai("gpt-4".to_string()); let mut pipeline = StreamPipeline::new(config); - // 模拟 Kiro 内容事件 + // 模拟 AWS Event Stream 内容事件 let bytes = br#"{"content":"Hello"}"#; let sse = pipeline.process_chunk(bytes); diff --git a/src-tauri/crates/providers/src/streaming/traits.rs b/src-tauri/crates/providers/src/streaming/traits.rs index 78a78d03d..63b568c54 100644 --- a/src-tauri/crates/providers/src/streaming/traits.rs +++ b/src-tauri/crates/providers/src/streaming/traits.rs @@ -4,10 +4,8 @@ //! //! # 需求覆盖 //! -//! - 需求 1.1: KiroProvider 流式支持 //! - 需求 1.2: ClaudeCustomProvider 流式支持 //! - 需求 1.3: OpenAICustomProvider 流式支持 -//! - 需求 1.4: AntigravityProvider 流式支持 #![allow(dead_code)] @@ -86,13 +84,13 @@ pub trait StreamingProvider: Send + Sync { /// 定义不同 Provider 使用的流式响应格式。 #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum StreamFormat { - /// AWS Event Stream 格式(Kiro/CodeWhisperer 使用) + /// AWS Event Stream 格式 AwsEventStream, /// Anthropic SSE 格式(Claude 使用) AnthropicSse, /// OpenAI SSE 格式(OpenAI 兼容 API 使用) OpenAiSse, - /// Gemini 流式格式(Antigravity/Gemini 使用) + /// Gemini 流式格式 GeminiStream, } diff --git a/src-tauri/crates/providers/src/translator/kiro/anthropic/mod.rs b/src-tauri/crates/providers/src/translator/kiro/anthropic/mod.rs deleted file mode 100644 index 430de7866..000000000 --- a/src-tauri/crates/providers/src/translator/kiro/anthropic/mod.rs +++ /dev/null @@ -1,13 +0,0 @@ -//! Anthropic 协议 → Kiro 后端转换 -//! -//! 处理 Anthropic Messages API 格式与 CodeWhisperer 格式之间的转换。 -//! -//! # 重要说明 -//! -//! 这是 Claude Code 使用的协议,是最核心的转换路径。 - -pub mod request; -pub mod response; - -pub use request::AnthropicRequestTranslator; -pub use response::AnthropicResponseTranslator; diff --git a/src-tauri/crates/providers/src/translator/kiro/anthropic/request.rs b/src-tauri/crates/providers/src/translator/kiro/anthropic/request.rs deleted file mode 100644 index b9e6c2d0f..000000000 --- a/src-tauri/crates/providers/src/translator/kiro/anthropic/request.rs +++ /dev/null @@ -1,736 +0,0 @@ -//! Anthropic 请求直接转换为 CodeWhisperer 请求 -//! -//! 直接将 Anthropic MessagesRequest 转换为 CodeWhisperer API 格式, -//! 无需经过 OpenAI 中间格式,减少转换开销。 - -use crate::translator::kiro::openai::request::{get_model_map, DEFAULT_MODEL}; -use crate::translator::traits::{RequestTranslator, TranslateError}; -use lime_core::models::anthropic::*; -use lime_core::models::codewhisperer::*; -use std::collections::HashSet; -use uuid::Uuid; - -/// Anthropic 到 Kiro 请求转换器 -#[derive(Debug, Clone)] -pub struct AnthropicRequestTranslator { - /// 可选的 Profile ARN (AWS CodeWhisperer) - pub profile_arn: Option, -} - -impl Default for AnthropicRequestTranslator { - fn default() -> Self { - Self::new() - } -} - -impl AnthropicRequestTranslator { - /// 创建新的转换器 - pub fn new() -> Self { - Self { profile_arn: None } - } - - /// 使用 Profile ARN 创建转换器 - pub fn with_profile_arn(profile_arn: String) -> Self { - Self { - profile_arn: Some(profile_arn), - } - } -} - -impl RequestTranslator for AnthropicRequestTranslator { - type Input = AnthropicMessagesRequest; - type Output = CodeWhispererRequest; - type Error = TranslateError; - - fn translate_request(&self, request: Self::Input) -> Result { - Ok(convert_anthropic_to_codewhisperer( - &request, - self.profile_arn.clone(), - )) - } -} - -// ============================================================================ -// 内部类型 -// ============================================================================ - -#[derive(Debug, Clone)] -struct ProcessedMessage { - role: String, - content: String, - tool_uses: Option>, - tool_results: Option>, - images: Option>, -} - -// ============================================================================ -// 转换函数 -// ============================================================================ - -/// 将 Anthropic MessagesRequest 直接转换为 CodeWhisperer 请求 -pub fn convert_anthropic_to_codewhisperer( - request: &AnthropicMessagesRequest, - profile_arn: Option, -) -> CodeWhispererRequest { - let model_map = get_model_map(); - let cw_model = model_map - .get(request.model.as_str()) - .map(|s| s.to_string()) - .unwrap_or_else(|| DEFAULT_MODEL.to_string()); - - let conversation_id = Uuid::new_v4().to_string(); - - // 提取 system prompt - let mut system_prompt = extract_system_text(&request.system); - - // 处理 tool_choice: required - CodeWhisperer 不支持此参数,通过 prompt 注入强制 - if is_tool_choice_required(&request.tool_choice) && request.tools.is_some() { - let tool_instruction = "\n\n[CRITICAL INSTRUCTION] You MUST use one of the provided tools to respond. Do NOT respond with plain text. Call a tool function immediately."; - system_prompt.push_str(tool_instruction); - tracing::info!("[KIRO_TRANSLATE] tool_choice=required detected in Anthropic request, injected tool instruction"); - } - - // 预处理消息 - let messages = preprocess_anthropic_messages(&request.messages); - - // 构建历史记录 - let mut history: Vec = Vec::new(); - let mut start_idx = 0; - - // 处理 system prompt - 合并到第一条用户消息 - if !system_prompt.is_empty() && !messages.is_empty() && messages[0].role == "user" { - let first_content = &messages[0].content; - let combined = format!("{system_prompt}\n\n{first_content}"); - - let mut user_msg = UserInputMessage { - content: combined, - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: messages[0].images.clone(), // 传递图片 - user_input_message_context: None, - }; - - if let Some(ref tool_results) = messages[0].tool_results { - user_msg.user_input_message_context = Some(UserInputMessageContext { - tools: None, - tool_results: Some(tool_results.clone()), - }); - } - - history.push(HistoryItem::User(UserHistoryItem { - user_input_message: user_msg, - })); - start_idx = 1; - } - - // 处理历史消息(除最后一条) - for msg in messages - .iter() - .take(messages.len().saturating_sub(1)) - .skip(start_idx) - { - match msg.role.as_str() { - "user" => { - let content = if msg.content.is_empty() { - if msg.tool_results.is_some() { - "Tool results provided.".to_string() - } else { - "Continue".to_string() - } - } else { - msg.content.clone() - }; - - let mut user_msg = UserInputMessage { - content, - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: msg.images.clone(), // 传递图片 - user_input_message_context: None, - }; - - if let Some(ref tool_results) = msg.tool_results { - user_msg.user_input_message_context = Some(UserInputMessageContext { - tools: None, - tool_results: Some(tool_results.clone()), - }); - } - - history.push(HistoryItem::User(UserHistoryItem { - user_input_message: user_msg, - })); - } - "assistant" => { - let content = if msg.content.is_empty() { - "I understand.".to_string() - } else { - msg.content.clone() - }; - - history.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content, - tool_uses: msg.tool_uses.clone(), - }, - })); - } - _ => {} - } - } - - // 修复历史记录交替顺序 - let history = fix_history_alternation(history, &cw_model); - - // 构建当前消息 - let (current_content, current_tool_results, current_images) = - if let Some(last_msg) = messages.last() { - if last_msg.role == "assistant" { - ("Continue".to_string(), None, None) - } else { - let content = if last_msg.content.is_empty() { - if last_msg.tool_results.is_some() { - "Tool results provided.".to_string() - } else { - "Continue".to_string() - } - } else { - last_msg.content.clone() - }; - ( - content, - last_msg.tool_results.clone(), - last_msg.images.clone(), - ) - } - } else { - ("Continue".to_string(), None, None) - }; - - // 构建 tools - let tools = convert_anthropic_tools(&request.tools); - - let user_input_message_context = if tools.is_some() || current_tool_results.is_some() { - Some(UserInputMessageContext { - tools, - tool_results: current_tool_results, - }) - } else { - None - }; - - // 记录图片信息 - if let Some(ref imgs) = current_images { - tracing::info!( - "[KIRO_TRANSLATE] Current message contains {} image(s)", - imgs.len() - ); - } - - CodeWhispererRequest { - conversation_state: ConversationState { - chat_trigger_type: "MANUAL".to_string(), - conversation_id, - current_message: CurrentMessage { - user_input_message: UserInputMessage { - content: current_content, - model_id: cw_model, - origin: "AI_EDITOR".to_string(), - images: current_images, // 传递当前消息的图片 - user_input_message_context, - }, - }, - history: if history.is_empty() { - None - } else { - Some(history) - }, - }, - profile_arn, - } -} - -/// 提取 system prompt 文本 -fn extract_system_text(system: &Option) -> String { - match system { - Some(serde_json::Value::String(s)) => s.clone(), - Some(serde_json::Value::Array(arr)) => arr - .iter() - .filter_map(|item| { - if item.get("type") == Some(&serde_json::Value::String("text".to_string())) { - item.get("text") - .and_then(|t| t.as_str()) - .map(|s| s.to_string()) - } else { - None - } - }) - .collect::>() - .join("\n"), - _ => String::new(), - } -} - -/// 预处理 Anthropic 消息 -fn preprocess_anthropic_messages(messages: &[AnthropicMessage]) -> Vec { - let mut result: Vec = Vec::new(); - - for msg in messages { - let processed = convert_anthropic_message(msg); - result.extend(processed); - } - - // 合并连续的 user 消息中的 tool_results - let mut merged: Vec = Vec::new(); - let mut pending_tool_results: Vec = Vec::new(); - - for msg in result { - if msg.role == "user" { - if let Some(ref tr) = msg.tool_results { - pending_tool_results.extend(tr.clone()); - } - if !msg.content.is_empty() || msg.tool_results.is_none() { - // 去重 tool_results - let mut seen_ids = HashSet::new(); - pending_tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - merged.push(ProcessedMessage { - role: msg.role, - content: if msg.content.is_empty() && !pending_tool_results.is_empty() { - "Tool results provided.".to_string() - } else { - msg.content - }, - tool_uses: None, - tool_results: if pending_tool_results.is_empty() { - None - } else { - Some(pending_tool_results.clone()) - }, - images: msg.images, // 保留图片 - }); - pending_tool_results.clear(); - } - } else { - // 如果有待处理的 tool_results,先创建 user 消息 - if !pending_tool_results.is_empty() { - let mut seen_ids = HashSet::new(); - pending_tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - merged.push(ProcessedMessage { - role: "user".to_string(), - content: "Tool results provided.".to_string(), - tool_uses: None, - tool_results: Some(pending_tool_results.clone()), - images: None, - }); - pending_tool_results.clear(); - } - merged.push(msg); - } - } - - // 处理末尾的 tool_results - if !pending_tool_results.is_empty() { - let mut seen_ids = HashSet::new(); - pending_tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - merged.push(ProcessedMessage { - role: "user".to_string(), - content: "Tool results provided.".to_string(), - tool_uses: None, - tool_results: Some(pending_tool_results), - images: None, - }); - } - - merged -} - -/// 转换单条 Anthropic 消息 -fn convert_anthropic_message(msg: &AnthropicMessage) -> Vec { - let mut result: Vec = Vec::new(); - - match &msg.content { - serde_json::Value::String(s) => { - result.push(ProcessedMessage { - role: msg.role.clone(), - content: s.clone(), - tool_uses: None, - tool_results: None, - images: None, - }); - } - serde_json::Value::Array(parts) => { - let mut text_parts: Vec = Vec::new(); - let mut tool_uses: Vec = Vec::new(); - let mut tool_results: Vec = Vec::new(); - let mut images: Vec = Vec::new(); - - for part in parts { - let part_type = part.get("type").and_then(|t| t.as_str()).unwrap_or(""); - - match part_type { - "text" => { - if let Some(text) = part.get("text").and_then(|t| t.as_str()) { - text_parts.push(text.to_string()); - } - } - "image" => { - // 处理 Anthropic 格式的图片 - // { "type": "image", "source": { "type": "base64", "media_type": "image/jpeg", "data": "..." } } - if let Some(source) = part.get("source") { - let source_type = - source.get("type").and_then(|t| t.as_str()).unwrap_or(""); - if source_type == "base64" { - let media_type = source - .get("media_type") - .and_then(|m| m.as_str()) - .unwrap_or("image/jpeg"); - let data = - source.get("data").and_then(|d| d.as_str()).unwrap_or(""); - - if !data.is_empty() { - // 从 media_type 提取格式 (image/jpeg -> jpeg) - let format = - media_type.split('/').nth(1).unwrap_or("jpeg").to_string(); - - images.push(CWImage { - format, - source: CWImageSource { - bytes: data.to_string(), - }, - }); - tracing::debug!( - "[KIRO_TRANSLATE] Converted image: media_type={}", - media_type - ); - } - } - } - } - "tool_use" => { - let default_id = format!("toolu_{}", &Uuid::new_v4().to_string()[..8]); - let id = part - .get("id") - .and_then(|i| i.as_str()) - .unwrap_or(&default_id); - let name = part.get("name").and_then(|n| n.as_str()).unwrap_or(""); - let input = part.get("input").cloned().unwrap_or(serde_json::json!({})); - - tool_uses.push(CWToolUse { - tool_use_id: id.to_string(), - name: name.to_string(), - input, - }); - } - "tool_result" => { - let tool_use_id = part - .get("tool_use_id") - .and_then(|i| i.as_str()) - .unwrap_or(""); - let content_text = extract_tool_result_content(part.get("content")); - let is_error = part - .get("is_error") - .and_then(|e| e.as_bool()) - .unwrap_or(false); - - tool_results.push(CWToolResult { - tool_use_id: tool_use_id.to_string(), - content: vec![CWTextContent { text: content_text }], - status: if is_error { - "error".to_string() - } else { - "success".to_string() - }, - }); - } - _ => {} - } - } - - // 处理 assistant 消息 - if msg.role == "assistant" { - result.push(ProcessedMessage { - role: "assistant".to_string(), - content: text_parts.join(""), - tool_uses: if tool_uses.is_empty() { - None - } else { - Some(tool_uses) - }, - tool_results: None, - images: None, // assistant 消息不包含图片 - }); - } - // 处理 user 消息 - else if msg.role == "user" { - // 先添加 tool results - if !tool_results.is_empty() { - result.push(ProcessedMessage { - role: "user".to_string(), - content: String::new(), - tool_uses: None, - tool_results: Some(tool_results), - images: None, - }); - } - - // 添加文本内容和图片 - if !text_parts.is_empty() || !images.is_empty() { - result.push(ProcessedMessage { - role: "user".to_string(), - content: text_parts.join(""), - tool_uses: None, - tool_results: None, - images: if images.is_empty() { - None - } else { - Some(images) - }, - }); - } - } - } - _ => {} - } - - result -} - -/// 提取 tool_result 内容 -fn extract_tool_result_content(content: Option<&serde_json::Value>) -> String { - match content { - Some(serde_json::Value::String(s)) => s.clone(), - Some(serde_json::Value::Array(arr)) => arr - .iter() - .filter_map(|item| { - if item.get("type") == Some(&serde_json::Value::String("text".to_string())) { - item.get("text") - .and_then(|t| t.as_str()) - .map(|s| s.to_string()) - } else { - None - } - }) - .collect::>() - .join("\n"), - _ => String::new(), - } -} - -/// 转换 Anthropic tools 为 CodeWhisperer tools -fn convert_anthropic_tools(tools: &Option>) -> Option> { - tools.as_ref().map(|tools| { - let mut cw_tools: Vec = Vec::new(); - let mut function_count = 0; - - for t in tools.iter() { - // 处理特殊工具类型 - if t.name == "web_search" || t.name == "web_search_20250305" { - cw_tools.push(CWToolItem::WebSearch(CWWebSearchTool { - tool_type: "web_search".to_string(), - })); - continue; - } - - // 限制最多 50 个函数工具 - if function_count >= 50 { - continue; - } - function_count += 1; - - let params = t - .input_schema - .clone() - .unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})); - - let desc = t - .description - .clone() - .unwrap_or_else(|| format!("Tool: {}", t.name)); - - cw_tools.push(CWToolItem::Standard(CWTool { - tool_specification: ToolSpecification { - name: t.name.clone(), - description: if desc.len() > 500 { - let truncated: String = desc.chars().take(497).collect(); - format!("{truncated}...") - } else { - desc - }, - input_schema: InputSchema { json: params }, - }, - })); - } - - cw_tools - }) -} - -/// 修复历史记录,确保 user/assistant 严格交替 -fn fix_history_alternation(history: Vec, model_id: &str) -> Vec { - if history.is_empty() { - return history; - } - - let mut fixed: Vec = Vec::new(); - - for item in history { - match &item { - HistoryItem::User(user_item) => { - if let Some(HistoryItem::User(last_user)) = fixed.last_mut() { - let has_tool_results = user_item - .user_input_message - .user_input_message_context - .as_ref() - .map(|ctx| ctx.tool_results.is_some()) - .unwrap_or(false); - - if has_tool_results { - let new_results = user_item - .user_input_message - .user_input_message_context - .as_ref() - .and_then(|ctx| ctx.tool_results.clone()) - .unwrap_or_default(); - - if let Some(ref mut ctx) = - last_user.user_input_message.user_input_message_context - { - if let Some(ref mut existing) = ctx.tool_results { - existing.extend(new_results); - } else { - ctx.tool_results = Some(new_results); - } - } else { - last_user.user_input_message.user_input_message_context = - Some(UserInputMessageContext { - tools: None, - tool_results: Some(new_results), - }); - } - continue; - } else { - fixed.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content: "I understand.".to_string(), - tool_uses: None, - }, - })); - } - } - fixed.push(item); - } - HistoryItem::Assistant(_) => { - if let Some(HistoryItem::Assistant(_)) = fixed.last() { - fixed.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: "Continue".to_string(), - model_id: model_id.to_string(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }, - })); - } - if fixed.is_empty() { - fixed.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: "Continue".to_string(), - model_id: model_id.to_string(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }, - })); - } - fixed.push(item); - } - } - } - - // 确保以 assistant 结尾 - if let Some(HistoryItem::User(_)) = fixed.last() { - fixed.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content: "I understand.".to_string(), - tool_uses: None, - }, - })); - } - - fixed -} - -/// 检查 tool_choice 是否为 required -/// -/// Anthropic tool_choice 可以是: -/// - {"type": "any"} - 必须调用工具 -/// - {"type": "tool", "name": "xxx"} - 必须调用指定工具 -fn is_tool_choice_required(tool_choice: &Option) -> bool { - match tool_choice { - Some(serde_json::Value::Object(obj)) => { - if let Some(serde_json::Value::String(t)) = obj.get("type") { - t == "any" || t == "tool" - } else { - false - } - } - // OpenAI 风格的 "required" 字符串 - Some(serde_json::Value::String(s)) => s == "required" || s == "any", - _ => false, - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_convert_simple_request() { - let request = AnthropicMessagesRequest { - model: "claude-sonnet-4-5".to_string(), - messages: vec![AnthropicMessage { - role: "user".to_string(), - content: serde_json::json!("Hello"), - }], - system: None, - max_tokens: Some(1024), - stream: true, - temperature: None, - tools: None, - tool_choice: None, - }; - - let translator = AnthropicRequestTranslator::new(); - let result = translator.translate_request(request); - assert!(result.is_ok()); - - let cw_request = result.unwrap(); - assert_eq!( - cw_request - .conversation_state - .current_message - .user_input_message - .model_id, - "CLAUDE_SONNET_4_5_20250929_V1_0" - ); - } - - #[test] - fn test_extract_system_text_string() { - let system = Some(serde_json::json!("You are a helpful assistant.")); - let text = extract_system_text(&system); - assert_eq!(text, "You are a helpful assistant."); - } - - #[test] - fn test_extract_system_text_array() { - let system = Some(serde_json::json!([ - {"type": "text", "text": "Line 1"}, - {"type": "text", "text": "Line 2"} - ])); - let text = extract_system_text(&system); - assert_eq!(text, "Line 1\nLine 2"); - } -} diff --git a/src-tauri/crates/providers/src/translator/kiro/anthropic/response.rs b/src-tauri/crates/providers/src/translator/kiro/anthropic/response.rs deleted file mode 100644 index 0b952b067..000000000 --- a/src-tauri/crates/providers/src/translator/kiro/anthropic/response.rs +++ /dev/null @@ -1,189 +0,0 @@ -//! Kiro 响应转换为 Anthropic SSE 格式 -//! -//! 将 `StreamEvent` 转换为 Anthropic Messages API 流式响应格式。 -//! 这是 Claude Code 使用的协议。 - -use crate::stream::{AnthropicSseGenerator, StreamEvent}; -use crate::translator::traits::{ResponseTranslator, SseResponseTranslator}; - -/// Anthropic 响应转换器 -/// -/// 将 `StreamEvent` 转换为 Anthropic SSE 格式 -#[derive(Debug)] -pub struct AnthropicResponseTranslator { - /// SSE 生成器 - generator: AnthropicSseGenerator, -} - -impl Default for AnthropicResponseTranslator { - fn default() -> Self { - Self::new("unknown".to_string()) - } -} - -impl AnthropicResponseTranslator { - /// 创建新的转换器 - pub fn new(model: String) -> Self { - Self { - generator: AnthropicSseGenerator::new(model), - } - } - - /// 使用指定的消息 ID 创建转换器 - pub fn with_id(id: String, model: String) -> Self { - Self { - generator: AnthropicSseGenerator::with_id(id, model), - } - } - - /// 获取消息 ID - pub fn message_id(&self) -> &str { - self.generator.message_id() - } - - /// 获取模型名称 - pub fn model(&self) -> &str { - self.generator.model() - } -} - -impl ResponseTranslator for AnthropicResponseTranslator { - type Output = Vec; - - fn translate_event(&mut self, event: &StreamEvent) -> Option { - let events = self.generator.generate(event); - if events.is_empty() { - None - } else { - Some(events) - } - } - - fn finalize(&mut self) -> Vec { - Vec::new() // Anthropic 生成器在 MessageStop 时已经发送了所有结束事件 - } - - fn reset(&mut self) { - self.generator = AnthropicSseGenerator::new("unknown".to_string()); - } -} - -impl SseResponseTranslator for AnthropicResponseTranslator { - fn translate_to_sse(&mut self, event: &StreamEvent) -> Vec { - self.generator.generate(event) - } - - fn finalize_sse(&mut self) -> Vec { - Vec::new() // Anthropic 生成器在 MessageStop 时已经发送了所有结束事件 - } - - fn reset(&mut self) { - self.generator = AnthropicSseGenerator::new("unknown".to_string()); - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::stream::{ContentBlockType, StopReason}; - - #[test] - fn test_translate_message_start() { - let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); - - let event = StreamEvent::MessageStart { - id: "msg_123".to_string(), - model: "claude-3-sonnet".to_string(), - }; - - let sse = translator.translate_event(&event); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert!(!sse.is_empty()); - assert!(sse[0].starts_with("event: message_start\ndata: ")); - } - - #[test] - fn test_translate_text_content() { - let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); - - // 先发送 message_start - let _ = translator.translate_event(&StreamEvent::MessageStart { - id: "msg_123".to_string(), - model: "claude-3-sonnet".to_string(), - }); - - // 内容块开始 - let sse = translator.translate_event(&StreamEvent::ContentBlockStart { - index: 0, - block_type: ContentBlockType::Text, - }); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert!(sse[0].contains("content_block_start")); - - // 文本增量 - let sse = translator.translate_event(&StreamEvent::TextDelta { - text: "Hello".to_string(), - }); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert!(sse[0].contains("content_block_delta")); - assert!(sse[0].contains("Hello")); - } - - #[test] - fn test_translate_tool_use() { - let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); - - // 发送 message_start - let _ = translator.translate_event(&StreamEvent::MessageStart { - id: "msg_123".to_string(), - model: "claude-3-sonnet".to_string(), - }); - - // 工具调用内容块开始 - let sse = translator.translate_event(&StreamEvent::ContentBlockStart { - index: 1, - block_type: ContentBlockType::ToolUse { - id: "tool_abc".to_string(), - name: "read_file".to_string(), - }, - }); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert!(sse[0].contains("content_block_start")); - assert!(sse[0].contains("tool_use")); - - // 工具参数增量 - let sse = translator.translate_event(&StreamEvent::ToolUseInputDelta { - id: "tool_abc".to_string(), - partial_json: "{\"path\":".to_string(), - }); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert!(sse[0].contains("input_json_delta")); - } - - #[test] - fn test_translate_message_stop() { - let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); - - // 发送 message_start - let _ = translator.translate_event(&StreamEvent::MessageStart { - id: "msg_123".to_string(), - model: "claude-3-sonnet".to_string(), - }); - - // 消息结束 - let sse = translator.translate_event(&StreamEvent::MessageStop { - stop_reason: StopReason::EndTurn, - }); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert_eq!(sse.len(), 2); - assert!(sse[0].contains("message_delta")); - assert!(sse[0].contains("end_turn")); - assert!(sse[1].contains("message_stop")); - } -} diff --git a/src-tauri/crates/providers/src/translator/kiro/mod.rs b/src-tauri/crates/providers/src/translator/kiro/mod.rs deleted file mode 100644 index 590b4c6bf..000000000 --- a/src-tauri/crates/providers/src/translator/kiro/mod.rs +++ /dev/null @@ -1,43 +0,0 @@ -//! Kiro/CodeWhisperer 后端协议转换 -//! -//! 处理与 AWS CodeWhisperer (Kiro) 后端的协议转换。 -//! -//! # 子模块 -//! -//! - `openai`: OpenAI 前端协议支持 -//! - `anthropic`: Anthropic 前端协议支持 -//! -//! # 调用链 -//! -//! ## OpenAI 协议 -//! ```text -//! OpenAI ChatCompletionRequest -//! → [openai/request.rs] translate_request -//! → CodeWhispererRequest -//! → [backends/kiro.rs] call_stream -//! → AWS Event Stream bytes -//! → [stream/parsers/aws_event_stream.rs] parse -//! → StreamEvent -//! → [openai/response.rs] translate_event -//! → OpenAI SSE -//! ``` -//! -//! ## Anthropic 协议 -//! ```text -//! Anthropic MessagesRequest -//! → [anthropic/request.rs] translate_request -//! → CodeWhispererRequest -//! → [backends/kiro.rs] call_stream -//! → AWS Event Stream bytes -//! → [stream/parsers/aws_event_stream.rs] parse -//! → StreamEvent -//! → [anthropic/response.rs] translate_event -//! → Anthropic SSE -//! ``` - -pub mod anthropic; -pub mod openai; - -// 重新导出常用类型 -pub use anthropic::{AnthropicRequestTranslator, AnthropicResponseTranslator}; -pub use openai::{OpenAiRequestTranslator, OpenAiResponseTranslator}; diff --git a/src-tauri/crates/providers/src/translator/kiro/openai/mod.rs b/src-tauri/crates/providers/src/translator/kiro/openai/mod.rs deleted file mode 100644 index 6bd8bc2e4..000000000 --- a/src-tauri/crates/providers/src/translator/kiro/openai/mod.rs +++ /dev/null @@ -1,9 +0,0 @@ -//! OpenAI 协议 → Kiro 后端转换 -//! -//! 处理 OpenAI ChatCompletion API 格式与 CodeWhisperer 格式之间的转换。 - -pub mod request; -pub mod response; - -pub use request::OpenAiRequestTranslator; -pub use response::OpenAiResponseTranslator; diff --git a/src-tauri/crates/providers/src/translator/kiro/openai/request.rs b/src-tauri/crates/providers/src/translator/kiro/openai/request.rs deleted file mode 100644 index a24b1d832..000000000 --- a/src-tauri/crates/providers/src/translator/kiro/openai/request.rs +++ /dev/null @@ -1,674 +0,0 @@ -//! OpenAI 请求转换为 CodeWhisperer 请求 -//! -//! 将 OpenAI ChatCompletionRequest 转换为 CodeWhisperer API 格式。 -//! -//! # 模型映射 -//! -//! - claude-opus-4-5 → claude-opus-4.5 -//! - claude-sonnet-4-5 → CLAUDE_SONNET_4_5_20250929_V1_0 -//! - claude-sonnet-4-20250514 → CLAUDE_SONNET_4_20250514_V1_0 -//! - claude-haiku-4-5 → claude-haiku-4.5 - -use crate::translator::traits::{RequestTranslator, TranslateError}; -use lime_core::models::codewhisperer::*; -use lime_core::models::openai::*; -use std::collections::{HashMap, HashSet}; -use uuid::Uuid; - -/// OpenAI 到 Kiro 请求转换器 -#[derive(Debug, Clone)] -pub struct OpenAiRequestTranslator { - /// 可选的 Profile ARN (AWS CodeWhisperer) - pub profile_arn: Option, -} - -impl Default for OpenAiRequestTranslator { - fn default() -> Self { - Self::new() - } -} - -impl OpenAiRequestTranslator { - /// 创建新的转换器 - pub fn new() -> Self { - Self { profile_arn: None } - } - - /// 使用 Profile ARN 创建转换器 - pub fn with_profile_arn(profile_arn: String) -> Self { - Self { - profile_arn: Some(profile_arn), - } - } -} - -impl RequestTranslator for OpenAiRequestTranslator { - type Input = ChatCompletionRequest; - type Output = CodeWhispererRequest; - type Error = TranslateError; - - fn translate_request(&self, request: Self::Input) -> Result { - Ok(convert_openai_to_codewhisperer( - &request, - self.profile_arn.clone(), - )) - } -} - -// ============================================================================ -// 模型映射 -// ============================================================================ - -/// 模型映射表 -pub fn get_model_map() -> HashMap<&'static str, &'static str> { - let mut map = HashMap::new(); - // Opus 4.5 系列 - map.insert("claude-opus-4-5", "claude-opus-4.5"); - map.insert("claude-opus-4-5-20251101", "claude-opus-4.5"); - // Haiku 4.5 系列 - map.insert("claude-haiku-4-5", "claude-haiku-4.5"); - map.insert("claude-haiku-4-5-20251001", "claude-haiku-4.5"); - // Sonnet 4.5 系列 - map.insert("claude-sonnet-4-5", "CLAUDE_SONNET_4_5_20250929_V1_0"); - map.insert( - "claude-sonnet-4-5-20250929", - "CLAUDE_SONNET_4_5_20250929_V1_0", - ); - // Sonnet 4 系列 - map.insert("claude-sonnet-4-20250514", "CLAUDE_SONNET_4_20250514_V1_0"); - // Sonnet 3.7/3.5 系列(兼容旧版本) - map.insert( - "claude-3-7-sonnet-20250219", - "CLAUDE_3_7_SONNET_20250219_V1_0", - ); - map.insert( - "claude-3-5-sonnet-20241022", - "CLAUDE_3_7_SONNET_20250219_V1_0", - ); - map.insert( - "claude-3-5-sonnet-latest", - "CLAUDE_3_7_SONNET_20250219_V1_0", - ); - map -} - -/// 获取支持的模型列表 -pub fn get_supported_models() -> Vec<&'static str> { - vec![ - "claude-opus-4-5", - "claude-opus-4-5-20251101", - "claude-haiku-4-5", - "claude-haiku-4-5-20251001", - "claude-sonnet-4-5", - "claude-sonnet-4-5-20250929", - "claude-sonnet-4-20250514", - "claude-3-7-sonnet-20250219", - ] -} - -/// 默认模型 -pub const DEFAULT_MODEL: &str = "CLAUDE_SONNET_4_5_20250929_V1_0"; - -// ============================================================================ -// 内部类型 -// ============================================================================ - -#[derive(Debug, Clone)] -struct ProcessedMessage { - role: String, - content: String, - tool_calls: Option>, - tool_results: Option>, - images: Option>, -} - -// ============================================================================ -// 转换函数 -// ============================================================================ - -/// 将 OpenAI ChatCompletionRequest 转换为 CodeWhisperer 请求 -pub fn convert_openai_to_codewhisperer( - request: &ChatCompletionRequest, - profile_arn: Option, -) -> CodeWhispererRequest { - convert_openai_to_codewhisperer_with_conversation_id(request, profile_arn, None) -} - -pub fn convert_openai_to_codewhisperer_with_conversation_id( - request: &ChatCompletionRequest, - profile_arn: Option, - conversation_id: Option<&str>, -) -> CodeWhispererRequest { - let model_map = get_model_map(); - let cw_model = model_map - .get(request.model.as_str()) - .map(|s| s.to_string()) - .unwrap_or_else(|| DEFAULT_MODEL.to_string()); - - let conversation_id = conversation_id - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_string) - .unwrap_or_else(|| Uuid::new_v4().to_string()); - - // 提取 system prompt 和消息 - let mut system_prompt = String::new(); - let mut raw_messages: Vec<&ChatMessage> = Vec::new(); - - for msg in &request.messages { - if msg.role == "system" { - system_prompt = msg.get_content_text(); - } else { - raw_messages.push(msg); - } - } - - // 调试日志:打印 tool_choice 和 tools 信息 - tracing::info!( - "[KIRO_TRANSLATE] 收到请求: tool_choice={:?}, has_tools={}, tools_count={}", - request.tool_choice, - request.tools.is_some(), - request.tools.as_ref().map(|t| t.len()).unwrap_or(0) - ); - - // 处理 tool_choice: required - CodeWhisperer 不支持此参数,通过 prompt 注入强制 - if is_tool_choice_required(&request.tool_choice) && request.tools.is_some() { - let tool_instruction = "\n\n[CRITICAL INSTRUCTION] You MUST use one of the provided tools to respond. Do NOT respond with plain text. Call a tool function immediately."; - system_prompt.push_str(tool_instruction); - tracing::info!("[KIRO_TRANSLATE] tool_choice=required detected, injected tool instruction"); - } - - // 预处理消息:合并 tool 消息 - let messages = preprocess_messages(&raw_messages); - - // 构建历史记录 - let mut history: Vec = Vec::new(); - let mut start_idx = 0; - - // 处理 system prompt - 合并到第一条用户消息 - if !system_prompt.is_empty() && !messages.is_empty() && messages[0].role == "user" { - let first_content = &messages[0].content; - let combined = format!("{system_prompt}\n\n{first_content}"); - - let mut user_msg = UserInputMessage { - content: combined, - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: messages[0].images.clone(), // 传递图片 - user_input_message_context: None, - }; - - if let Some(ref tool_results) = messages[0].tool_results { - user_msg.user_input_message_context = Some(UserInputMessageContext { - tools: None, - tool_results: Some(tool_results.clone()), - }); - } - - history.push(HistoryItem::User(UserHistoryItem { - user_input_message: user_msg, - })); - start_idx = 1; - } - - // 处理历史消息(除最后一条) - for msg in messages - .iter() - .take(messages.len().saturating_sub(1)) - .skip(start_idx) - { - match msg.role.as_str() { - "user" => { - let content = if msg.content.is_empty() { - if msg.tool_results.is_some() { - "Tool results provided.".to_string() - } else { - "Continue".to_string() - } - } else { - msg.content.clone() - }; - - let mut user_msg = UserInputMessage { - content, - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: msg.images.clone(), // 传递图片 - user_input_message_context: None, - }; - - if let Some(ref tool_results) = msg.tool_results { - user_msg.user_input_message_context = Some(UserInputMessageContext { - tools: None, - tool_results: Some(tool_results.clone()), - }); - } - - history.push(HistoryItem::User(UserHistoryItem { - user_input_message: user_msg, - })); - } - "assistant" => { - let content = if msg.content.is_empty() { - "I understand.".to_string() - } else { - msg.content.clone() - }; - - history.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content, - tool_uses: msg.tool_calls.clone(), - }, - })); - } - _ => {} - } - } - - // 修复历史记录交替顺序 - let history = fix_history_alternation(history, &cw_model); - - // 构建当前消息 - let (current_content, current_tool_results, current_images) = - if let Some(last_msg) = messages.last() { - if last_msg.role == "assistant" { - ("Continue".to_string(), None, None) - } else { - let content = if last_msg.content.is_empty() { - if last_msg.tool_results.is_some() { - "Tool results provided.".to_string() - } else { - "Continue".to_string() - } - } else { - last_msg.content.clone() - }; - ( - content, - last_msg.tool_results.clone(), - last_msg.images.clone(), - ) - } - } else { - ("Continue".to_string(), None, None) - }; - - // 构建 tools - let tools = convert_tools(&request.tools); - - let user_input_message_context = if tools.is_some() || current_tool_results.is_some() { - Some(UserInputMessageContext { - tools, - tool_results: current_tool_results, - }) - } else { - None - }; - - // 记录图片信息 - if let Some(ref imgs) = current_images { - tracing::info!( - "[KIRO_TRANSLATE] OpenAI current message contains {} image(s)", - imgs.len() - ); - } - - CodeWhispererRequest { - conversation_state: ConversationState { - chat_trigger_type: "MANUAL".to_string(), - conversation_id, - current_message: CurrentMessage { - user_input_message: UserInputMessage { - content: current_content, - model_id: cw_model, - origin: "AI_EDITOR".to_string(), - images: current_images, // 传递当前消息的图片 - user_input_message_context, - }, - }, - history: if history.is_empty() { - None - } else { - Some(history) - }, - }, - profile_arn, - } -} - -/// 预处理消息:合并连续的 tool 消息到前一个 assistant 消息后的 user 消息 -fn preprocess_messages(messages: &[&ChatMessage]) -> Vec { - let mut result: Vec = Vec::new(); - let mut pending_tool_results: Vec = Vec::new(); - - for msg in messages { - match msg.role.as_str() { - "tool" => { - let content = msg.get_content_text(); - let tool_id = msg.tool_call_id.clone().unwrap_or_default(); - pending_tool_results.push(CWToolResult { - content: vec![CWTextContent { text: content }], - status: "success".to_string(), - tool_use_id: tool_id, - }); - } - "user" => { - let content = msg.get_content_text(); - let mut tool_results = pending_tool_results.clone(); - pending_tool_results.clear(); - - // 去重 tool_results - let mut seen_ids = HashSet::new(); - tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - // 提取图片 - let raw_images = msg.get_images(); - let images = if raw_images.is_empty() { - None - } else { - Some( - raw_images - .into_iter() - .map(|(format, data)| CWImage { - format, - source: CWImageSource { bytes: data }, - }) - .collect(), - ) - }; - - if images.is_some() { - tracing::debug!( - "[KIRO_TRANSLATE] OpenAI user message contains {} image(s)", - images.as_ref().map(|v: &Vec| v.len()).unwrap_or(0) - ); - } - - result.push(ProcessedMessage { - role: "user".to_string(), - content, - tool_calls: None, - tool_results: if tool_results.is_empty() { - None - } else { - Some(tool_results) - }, - images, - }); - } - "assistant" => { - // 如果有待处理的 tool results,先创建一个 user 消息 - if !pending_tool_results.is_empty() { - let mut tool_results = pending_tool_results.clone(); - pending_tool_results.clear(); - - let mut seen_ids = HashSet::new(); - tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - result.push(ProcessedMessage { - role: "user".to_string(), - content: "Tool results provided.".to_string(), - tool_calls: None, - tool_results: Some(tool_results), - images: None, - }); - } - - let content = msg.get_content_text(); - let tool_calls = msg.tool_calls.as_ref().map(|calls| { - calls - .iter() - .map(|tc| CWToolUse { - input: serde_json::from_str(&tc.function.arguments) - .unwrap_or(serde_json::json!({})), - name: tc.function.name.clone(), - tool_use_id: tc.id.clone(), - }) - .collect() - }); - - result.push(ProcessedMessage { - role: "assistant".to_string(), - content, - tool_calls, - tool_results: None, - images: None, // assistant 消息不包含图片 - }); - } - _ => {} - } - } - - // 处理末尾的 tool results - if !pending_tool_results.is_empty() { - let mut tool_results = pending_tool_results; - let mut seen_ids = HashSet::new(); - tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); - - result.push(ProcessedMessage { - role: "user".to_string(), - content: "Tool results provided.".to_string(), - tool_calls: None, - tool_results: Some(tool_results), - images: None, - }); - } - - result -} - -/// 转换工具列表 -fn convert_tools(tools: &Option>) -> Option> { - tools.as_ref().map(|tools| { - let mut cw_tools: Vec = Vec::new(); - let mut function_count = 0; - - for t in tools.iter() { - match t { - Tool::Function { function } => { - // 限制最多 50 个函数工具 - if function_count >= 50 { - continue; - } - function_count += 1; - - let params = function - .parameters - .clone() - .unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})); - - let desc = function - .description - .clone() - .unwrap_or_else(|| format!("Tool: {}", function.name)); - - cw_tools.push(CWToolItem::Standard(CWTool { - tool_specification: ToolSpecification { - name: function.name.clone(), - description: if desc.len() > 500 { - let truncated: String = desc.chars().take(497).collect(); - format!("{truncated}...") - } else { - desc - }, - input_schema: InputSchema { json: params }, - }, - })); - } - Tool::WebSearch | Tool::WebSearch20250305 => { - cw_tools.push(CWToolItem::WebSearch(CWWebSearchTool { - tool_type: "web_search".to_string(), - })); - } - } - } - - cw_tools - }) -} - -/// 修复历史记录,确保 user/assistant 严格交替 -fn fix_history_alternation(history: Vec, model_id: &str) -> Vec { - if history.is_empty() { - return history; - } - - let mut fixed: Vec = Vec::new(); - - for item in history { - match &item { - HistoryItem::User(user_item) => { - if let Some(HistoryItem::User(last_user)) = fixed.last_mut() { - let has_tool_results = user_item - .user_input_message - .user_input_message_context - .as_ref() - .map(|ctx| ctx.tool_results.is_some()) - .unwrap_or(false); - - if has_tool_results { - let new_results = user_item - .user_input_message - .user_input_message_context - .as_ref() - .and_then(|ctx| ctx.tool_results.clone()) - .unwrap_or_default(); - - if let Some(ref mut ctx) = - last_user.user_input_message.user_input_message_context - { - if let Some(ref mut existing) = ctx.tool_results { - existing.extend(new_results); - } else { - ctx.tool_results = Some(new_results); - } - } else { - last_user.user_input_message.user_input_message_context = - Some(UserInputMessageContext { - tools: None, - tool_results: Some(new_results), - }); - } - continue; - } else { - fixed.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content: "I understand.".to_string(), - tool_uses: None, - }, - })); - } - } - fixed.push(item); - } - HistoryItem::Assistant(_) => { - if let Some(HistoryItem::Assistant(_)) = fixed.last() { - fixed.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: "Continue".to_string(), - model_id: model_id.to_string(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }, - })); - } - if fixed.is_empty() { - fixed.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: "Continue".to_string(), - model_id: model_id.to_string(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }, - })); - } - fixed.push(item); - } - } - } - - // 确保以 assistant 结尾 - if let Some(HistoryItem::User(_)) = fixed.last() { - fixed.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content: "I understand.".to_string(), - tool_uses: None, - }, - })); - } - - fixed -} - -/// 检查 tool_choice 是否为 required -/// -/// tool_choice 可以是: -/// - "required" 字符串 -/// - {"type": "any"} 或类似结构 -fn is_tool_choice_required(tool_choice: &Option) -> bool { - match tool_choice { - Some(serde_json::Value::String(s)) => s == "required" || s == "any", - Some(serde_json::Value::Object(obj)) => { - // 检查 {"type": "any"} 或 {"type": "tool", ...} - if let Some(serde_json::Value::String(t)) = obj.get("type") { - t == "any" || t == "tool" - } else { - false - } - } - _ => false, - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_mapping() { - let map = get_model_map(); - assert_eq!(map.get("claude-opus-4-5"), Some(&"claude-opus-4.5")); - assert_eq!( - map.get("claude-sonnet-4-5"), - Some(&"CLAUDE_SONNET_4_5_20250929_V1_0") - ); - } - - #[test] - fn test_convert_simple_request() { - let request = ChatCompletionRequest { - model: "claude-sonnet-4-5".to_string(), - messages: vec![ChatMessage { - role: "user".to_string(), - content: Some(MessageContent::Text("Hello".to_string())), - tool_calls: None, - tool_call_id: None, - reasoning_content: None, - }], - tools: None, - stream: false, - max_tokens: None, - temperature: None, - top_p: None, - tool_choice: None, - reasoning_effort: None, - }; - - let translator = OpenAiRequestTranslator::new(); - let result = translator.translate_request(request); - assert!(result.is_ok()); - - let cw_request = result.unwrap(); - assert_eq!( - cw_request - .conversation_state - .current_message - .user_input_message - .model_id, - "CLAUDE_SONNET_4_5_20250929_V1_0" - ); - } -} diff --git a/src-tauri/crates/providers/src/translator/kiro/openai/response.rs b/src-tauri/crates/providers/src/translator/kiro/openai/response.rs deleted file mode 100644 index a04def48d..000000000 --- a/src-tauri/crates/providers/src/translator/kiro/openai/response.rs +++ /dev/null @@ -1,130 +0,0 @@ -//! Kiro 响应转换为 OpenAI SSE 格式 -//! -//! 将 `StreamEvent` 转换为 OpenAI Chat Completions 流式响应格式。 - -use crate::stream::{OpenAiSseGenerator, StreamEvent}; -use crate::translator::traits::{ResponseTranslator, SseResponseTranslator}; - -/// OpenAI 响应转换器 -/// -/// 将 `StreamEvent` 转换为 OpenAI SSE 格式 -#[derive(Debug)] -pub struct OpenAiResponseTranslator { - /// SSE 生成器 - generator: OpenAiSseGenerator, -} - -impl Default for OpenAiResponseTranslator { - fn default() -> Self { - Self::new("unknown".to_string()) - } -} - -impl OpenAiResponseTranslator { - /// 创建新的转换器 - pub fn new(model: String) -> Self { - Self { - generator: OpenAiSseGenerator::new(model), - } - } - - /// 使用指定的响应 ID 创建转换器 - pub fn with_id(id: String, model: String) -> Self { - Self { - generator: OpenAiSseGenerator::with_id(id, model), - } - } - - /// 获取响应 ID - pub fn response_id(&self) -> &str { - self.generator.response_id() - } -} - -impl ResponseTranslator for OpenAiResponseTranslator { - type Output = String; - - fn translate_event(&mut self, event: &StreamEvent) -> Option { - self.generator.generate(event) - } - - fn finalize(&mut self) -> Vec { - vec![self.generator.generate_done()] - } - - fn reset(&mut self) { - self.generator = OpenAiSseGenerator::new("unknown".to_string()); - } -} - -impl SseResponseTranslator for OpenAiResponseTranslator { - fn translate_to_sse(&mut self, event: &StreamEvent) -> Vec { - self.generator.generate(event).into_iter().collect() - } - - fn finalize_sse(&mut self) -> Vec { - vec![self.generator.generate_done()] - } - - fn reset(&mut self) { - self.generator = OpenAiSseGenerator::new("unknown".to_string()); - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::stream::StopReason; - - #[test] - fn test_translate_text_delta() { - let mut translator = OpenAiResponseTranslator::new("gpt-4".to_string()); - - let event = StreamEvent::TextDelta { - text: "Hello".to_string(), - }; - - let sse = translator.translate_event(&event); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert!(sse.starts_with("data: ")); - assert!(sse.contains("\"content\":\"Hello\"")); - } - - #[test] - fn test_translate_tool_use() { - let mut translator = OpenAiResponseTranslator::new("gpt-4".to_string()); - - // 工具调用开始 - let event = StreamEvent::ToolUseStart { - id: "call_123".to_string(), - name: "read_file".to_string(), - }; - let sse = translator.translate_event(&event); - assert!(sse.is_some()); - assert!(sse.unwrap().contains("\"tool_calls\"")); - - // 工具参数增量 - let event = StreamEvent::ToolUseInputDelta { - id: "call_123".to_string(), - partial_json: "{\"path\":".to_string(), - }; - let sse = translator.translate_event(&event); - assert!(sse.is_some()); - } - - #[test] - fn test_translate_message_stop() { - let mut translator = OpenAiResponseTranslator::new("gpt-4".to_string()); - - let event = StreamEvent::MessageStop { - stop_reason: StopReason::EndTurn, - }; - - let sse = translator.translate_event(&event); - assert!(sse.is_some()); - let sse = sse.unwrap(); - assert!(sse.contains("\"finish_reason\":\"stop\"")); - assert!(sse.contains("[DONE]")); - } -} diff --git a/src-tauri/crates/providers/src/translator/mod.rs b/src-tauri/crates/providers/src/translator/mod.rs index c069a3f10..8b50bbf2c 100644 --- a/src-tauri/crates/providers/src/translator/mod.rs +++ b/src-tauri/crates/providers/src/translator/mod.rs @@ -1,52 +1,17 @@ //! 协议转换层 //! -//! 处理不同前端协议(OpenAI、Anthropic、Gemini CLI)与不同后端(Kiro、Codex、Claude) -//! 之间的请求和响应格式转换。 +//! 处理 current Provider 所需的请求和响应格式转换。 //! //! # 架构设计 //! //! ```text //! translator/ -//! ├── traits.rs # 转换器 trait 定义 -//! └── kiro/ # Kiro/CodeWhisperer 后端 -//! ├── openai/ # OpenAI 前端协议 -//! │ ├── request.rs # OpenAI → Kiro 请求 -//! │ └── response.rs # StreamEvent → OpenAI SSE -//! └── anthropic/ # Anthropic 前端协议 -//! ├── request.rs # Anthropic → Kiro 请求 -//! └── response.rs # StreamEvent → Anthropic SSE -//! ``` -//! -//! # 使用示例 -//! -//! ```ignore -//! use lime::translator::kiro::{ -//! AnthropicRequestTranslator, AnthropicResponseTranslator, -//! }; -//! use lime::translator::traits::RequestTranslator; -//! -//! // 请求转换 -//! let translator = AnthropicRequestTranslator::new(); -//! let cw_request = translator.translate_request(anthropic_request)?; -//! -//! // 响应转换 -//! let mut response_translator = AnthropicResponseTranslator::new(model); -//! for event in stream_events { -//! let sse_events = response_translator.translate_to_sse(&event); -//! for sse in sse_events { -//! // 发送 SSE 到客户端 -//! } -//! } +//! └── traits.rs # 转换器 trait 定义 //! ``` -pub mod kiro; pub mod traits; // 重新导出核心类型 -pub use kiro::{ - AnthropicRequestTranslator, AnthropicResponseTranslator, OpenAiRequestTranslator, - OpenAiResponseTranslator, -}; pub use traits::{ RequestTranslator, ResponseTranslator, SseResponseTranslator, TranslateError, TranslateErrorKind, diff --git a/src-tauri/crates/providers/src/translator/traits.rs b/src-tauri/crates/providers/src/translator/traits.rs index cae70847a..37ebbc2c7 100644 --- a/src-tauri/crates/providers/src/translator/traits.rs +++ b/src-tauri/crates/providers/src/translator/traits.rs @@ -17,7 +17,7 @@ use crate::stream::StreamEvent; /// # 类型参数 /// /// - `Input`: 前端请求类型(如 OpenAI ChatCompletionRequest) -/// - `Output`: 后端请求类型(如 CodeWhispererRequest) +/// - `Output`: 后端请求类型 /// - `Error`: 转换错误类型 pub trait RequestTranslator { /// 前端请求类型 diff --git a/src-tauri/crates/server/Cargo.toml b/src-tauri/crates/server/Cargo.toml index d1d3621ee..c25819211 100644 --- a/src-tauri/crates/server/Cargo.toml +++ b/src-tauri/crates/server/Cargo.toml @@ -9,7 +9,6 @@ lime-config.workspace = true lime-infra.workspace = true lime-providers.workspace = true lime-services.workspace = true -lime-credential.workspace = true lime-websocket.workspace = true lime-processor.workspace = true lime-server-utils.workspace = true diff --git a/src-tauri/crates/server/src/handlers/api.rs b/src-tauri/crates/server/src/handlers/api.rs index bb1306354..4d050e960 100644 --- a/src-tauri/crates/server/src/handlers/api.rs +++ b/src-tauri/crates/server/src/handlers/api.rs @@ -44,11 +44,9 @@ use lime_core::models::anthropic::AnthropicMessagesRequest; use lime_core::models::openai::{ChatCompletionRequest, ContentPart, MessageContent}; use lime_core::ProviderType; use lime_processor::RequestContext; -use lime_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; use lime_providers::streaming::StreamFormat as StreamingFormat; use lime_server_utils::{ - build_anthropic_response, build_anthropic_stream_response, build_error_response_with_meta, - build_gateway_error_json, message_content_len, parse_cw_response, safe_truncate, + build_error_response_with_meta, build_gateway_error_json, message_content_len, }; use super::{call_provider_anthropic, call_provider_openai}; @@ -57,7 +55,7 @@ async fn select_credential_for_request( state: &AppState, request_id: Option<&str>, selected_provider: &str, - model: &str, + _model: &str, client_type: &ClientType, explicit_provider_id: Option<&str>, log_prefix: &str, @@ -74,13 +72,14 @@ async fn select_credential_for_request( if let Some(explicit_provider_id) = explicit_provider_id { eprintln!("[{log_prefix}] 使用 X-Provider-Id 指定的 provider: {explicit_provider_id}"); let cred = state - .pool_service - .select_credential_with_client_check( + .api_key_service + .select_credential_for_provider( db, explicit_provider_id, - Some(model), + Some(explicit_provider_id), Some(client_type), ) + .await .ok() .flatten(); @@ -112,14 +111,18 @@ async fn select_credential_for_request( if !state.allow_provider_fallback { eprintln!( - "[{log_prefix}] 已禁用自动降级(retry.auto_switch_provider=false),仅从 Provider Pool 选择" + "[{log_prefix}] 已禁用自动降级(retry.auto_switch_provider=false),仅从 API Key Provider 选择" ); - return match state.pool_service.select_credential_with_client_check( - db, - selected_provider, - Some(model), - Some(client_type), - ) { + return match state + .api_key_service + .select_credential_for_provider( + db, + selected_provider, + Some(selected_provider), + Some(client_type), + ) + .await + { Ok(cred) => { if cred.is_some() { eprintln!("[{log_prefix}] 找到凭证: provider={selected_provider}"); @@ -139,12 +142,10 @@ async fn select_credential_for_request( let provider_id_hint = selected_provider.to_lowercase(); match state - .pool_service - .select_credential_with_fallback( + .api_key_service + .select_credential_for_provider( db, - &state.api_key_service, selected_provider, - Some(model), Some(provider_id_hint.as_str()), Some(client_type), ) @@ -2312,7 +2313,7 @@ pub async fn chat_completions( } } - // 如果找到凭证池中的凭证,使用它 + // 如果找到 API Key Provider 凭证,使用它 if let Some(cred) = credential { eprintln!( "[CHAT_COMPLETIONS] 使用凭证: type={}, name={:?}, uuid={}", @@ -2323,7 +2324,7 @@ pub async fn chat_completions( state.logs.write().await.add( "info", &format!( - "[ROUTE] Using pool credential: type={} name={:?} uuid={}", + "[ROUTE] Using API Key Provider credential: type={} name={:?} uuid={}", cred.provider_type, cred.name, &cred.uuid[..8] @@ -2411,367 +2412,37 @@ pub async fn chat_completions( ); } - // 回退到旧的单凭证模式(仅当允许自动降级且选择的 Provider 是 Kiro 时) - // 其余情况(含禁用自动降级)直接返回无可用凭证错误 + // 凭证池和旧 Kiro 单凭证模式已退役,未找到 API Key Provider 时直接返回错误。 // **Validates: Requirements 3.2** - if !state.allow_provider_fallback || effective_provider.to_lowercase() != "kiro" { - let reason = if !state.allow_provider_fallback { - "auto fallback disabled by retry.auto_switch_provider=false" - } else { - "legacy mode only supports Kiro" - }; - state.logs.write().await.add( - "error", - &format!( - "[ROUTE] No pool credential found for '{effective_provider}' (client_type={client_type}), {reason}" - ), - ); - let message = if !state.allow_provider_fallback { - format!( - "没有找到可用的 '{}' 凭证(已禁用自动降级)。请在凭证池中添加对应的凭证。", - effective_provider - ) - } else { - format!( - "没有找到可用的 '{}' 凭证。请在凭证池中添加对应的凭证。", - effective_provider - ) - }; - return build_error_response_with_meta( - StatusCode::SERVICE_UNAVAILABLE.as_u16(), - &message, - Some(&ctx.request_id), - Some(&effective_provider), - Some(GatewayErrorCode::NoCredentials), - ); - } - + let reason = if !state.allow_provider_fallback { + "auto fallback disabled by retry.auto_switch_provider=false" + } else { + "legacy credential pool fallback retired" + }; state.logs.write().await.add( - "debug", - &format!("[ROUTE] No pool credential found for '{effective_provider}', using legacy mode"), + "error", + &format!( + "[ROUTE] No API Key Provider credential found for '{effective_provider}' (client_type={client_type}), {reason}" + ), ); - - // 启动 Flow 捕获(legacy mode) - - // 使用实际的 provider ID 构建 Flow Metadata - let _provider_type = effective_provider - .parse::() - .unwrap_or(ProviderType::OpenAI); - - // 检查是否需要拦截请求(legacy mode) - // **Validates: Requirements 2.1, 2.3, 2.5** - - // 检查是否需要刷新 token(无 token 或即将过期) - { - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - let needs_refresh = - kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); - if needs_refresh { - if let Err(e) = kiro.refresh_token().await { - state - .logs - .write() - .await - .add("error", &format!("Token refresh failed: {e}")); - // 标记 Flow 失败 - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ).into_response(); - } - } - } - - let kiro = state.kiro.read().await; - - match kiro.call_api(&request).await { - Ok(resp) => { - let status = resp.status(); - if status.is_success() { - match resp.text().await { - Ok(body) => { - let parsed = parse_cw_response(&body); - let has_tool_calls = !parsed.tool_calls.is_empty(); - - state.logs.write().await.add( - "info", - &format!( - "Request completed: content_len={}, tool_calls={}", - parsed.content.len(), - parsed.tool_calls.len() - ), - ); - - // 构建消息 - let message = if has_tool_calls { - serde_json::json!({ - "role": "assistant", - "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, - "tool_calls": parsed.tool_calls.iter().map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": "function", - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - }).collect::>() - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": parsed.content - }) - }; - - // 估算 Token 数量(基于字符数,约 4 字符 = 1 token) - let estimated_output_tokens = (parsed.content.len() / 4) as u32; - // 估算输入 Token(基于请求消息) - let estimated_input_tokens = request - .messages - .iter() - .map(|m| { - let content_len = match &m.content { - Some(c) => message_content_len(c), - None => 0, - }; - content_len / 4 - }) - .sum::() - as u32; - - let response = serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } - }], - "usage": { - "prompt_tokens": estimated_input_tokens, - "completion_tokens": estimated_output_tokens, - "total_tokens": estimated_input_tokens + estimated_output_tokens - } - }); - // 记录成功请求统计 - record_request_telemetry( - &state, - &ctx, - lime_infra::telemetry::RequestStatus::Success, - None, - ); - // 记录 Token 使用量 - record_token_usage( - &state, - &ctx, - Some(estimated_input_tokens), - Some(estimated_output_tokens), - ); - // 完成 Flow 捕获并检查响应拦截 - // **Validates: Requirements 2.1, 2.5** - let response = Json(response).into_response(); - return attach_route_debug_headers( - finalize_replayable_response( - response, - &mut idempotency_guard, - &mut dedup_guard, - &mut cache_guard, - &ctx.request_id, - ) - .await, - &selected_provider, - &effective_provider, - &ctx.resolved_model, - ); - } - Err(e) => { - // 记录失败请求统计 - record_request_telemetry( - &state, - &ctx, - lime_infra::telemetry::RequestStatus::Failed, - Some(e.to_string()), - ); - // 标记 Flow 失败 - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } else if status.as_u16() == 403 || status.as_u16() == 402 { - // Token 过期或账户问题,尝试重新加载凭证并刷新 - drop(kiro); - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - state.logs.write().await.add( - "warn", - &format!( - "[AUTH] Got {}, reloading credentials and attempting token refresh...", - status.as_u16() - ), - ); - - // 先重新加载凭证文件(可能用户换了账户) - if let Err(e) = kiro.load_credentials().await { - state.logs.write().await.add( - "error", - &format!("[AUTH] Failed to reload credentials: {e}"), - ); - } - - match kiro.refresh_token().await { - Ok(_) => { - state - .logs - .write() - .await - .add("info", "[AUTH] Token refreshed successfully after reload"); - // 重试请求 - drop(kiro); - let kiro = state.kiro.read().await; - match kiro.call_api(&request).await { - Ok(retry_resp) => { - if retry_resp.status().is_success() { - match retry_resp.text().await { - Ok(body) => { - let parsed = parse_cw_response(&body); - let has_tool_calls = !parsed.tool_calls.is_empty(); - - let message = if has_tool_calls { - serde_json::json!({ - "role": "assistant", - "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, - "tool_calls": parsed.tool_calls.iter().map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": "function", - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - }).collect::>() - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": parsed.content - }) - }; - - let response = serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } - }], - "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 - } - }); - // 完成 Flow 捕获并检查响应拦截(重试成功) - // **Validates: Requirements 2.1, 2.5** - let response = Json(response).into_response(); - return attach_route_debug_headers( - finalize_replayable_response( - response, - &mut idempotency_guard, - &mut dedup_guard, - &mut cache_guard, - &ctx.request_id, - ) - .await, - &selected_provider, - &effective_provider, - &ctx.resolved_model, - ); - } - Err(e) => { - // 标记 Flow 失败 - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ).into_response(); - } - } - } - let body = retry_resp.text().await.unwrap_or_default(); - // 标记 Flow 失败(重试失败) - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), - ).into_response() - } - Err(e) => { - // 标记 Flow 失败 - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("[AUTH] Token refresh failed: {e}")); - // 标记 Flow 失败 - ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ) - .into_response() - } - } - } else { - let body = resp.text().await.unwrap_or_default(); - state.logs.write().await.add( - "error", - &format!("Upstream error {}: {}", status, safe_truncate(&body, 200)), - ); - // 标记 Flow 失败 - ( - StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})) - ).into_response() - } - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("API call failed: {e}")); - // 标记 Flow 失败 - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } + let message = if !state.allow_provider_fallback { + format!( + "没有找到可用的 '{}' API Key Provider 凭证(已禁用自动降级)。", + effective_provider + ) + } else { + format!( + "没有找到可用的 '{}' API Key Provider 凭证。", + effective_provider + ) + }; + build_error_response_with_meta( + StatusCode::SERVICE_UNAVAILABLE.as_u16(), + &message, + Some(&ctx.request_id), + Some(&effective_provider), + Some(GatewayErrorCode::NoCredentials), + ) } pub async fn anthropic_messages( @@ -3087,12 +2758,12 @@ pub async fn anthropic_messages( } } - // 如果找到凭证池中的凭证,使用它 + // 如果找到 API Key Provider 凭证,使用它 if let Some(cred) = credential { state.logs.write().await.add( "info", &format!( - "[ROUTE] Using pool credential: type={} name={:?} uuid={}", + "[ROUTE] Using API Key Provider credential: type={} name={:?} uuid={}", cred.provider_type, cred.name, &cred.uuid[..8] @@ -3181,386 +2852,42 @@ pub async fn anthropic_messages( ); } - // 回退到旧的单凭证模式(仅当允许自动降级且选择的 Provider 是 Kiro 时) - // 其余情况(含禁用自动降级)直接返回无可用凭证错误 + // 凭证池和旧 Kiro 单凭证模式已退役,未找到 API Key Provider 时直接返回错误。 // **Validates: Requirements 3.2** - if !state.allow_provider_fallback || effective_provider.to_lowercase() != "kiro" { - let reason = if !state.allow_provider_fallback { - "auto fallback disabled by retry.auto_switch_provider=false" - } else { - "legacy mode only supports Kiro" - }; - state.logs.write().await.add( - "error", - &format!( - "[ROUTE] No pool credential found for '{effective_provider}' (client_type={client_type}), {reason}" - ), - ); - let message = if !state.allow_provider_fallback { - format!( - "没有找到可用的 '{}' 凭证(已禁用自动降级)。请在凭证池中添加对应的凭证。", - effective_provider - ) - } else { - format!( - "没有找到可用的 '{}' 凭证。请在凭证池中添加对应的凭证。", - effective_provider - ) - }; - let body = build_gateway_error_json( - StatusCode::SERVICE_UNAVAILABLE.as_u16(), - &message, - Some(&ctx.request_id), - Some(&effective_provider), - Some(GatewayErrorCode::NoCredentials), - ); - return ( - StatusCode::SERVICE_UNAVAILABLE, - Json(serde_json::json!({ "type": "error", "error": body["error"].clone() })), - ) - .into_response(); - } - + let reason = if !state.allow_provider_fallback { + "auto fallback disabled by retry.auto_switch_provider=false" + } else { + "legacy credential pool fallback retired" + }; state.logs.write().await.add( - "debug", - &format!("[ROUTE] No pool credential found for '{effective_provider}', using legacy mode"), - ); - - // 启动 Flow 捕获(legacy mode) - - // 使用实际的 provider ID 构建 Flow Metadata - let _provider_type = effective_provider - .parse::() - .unwrap_or(ProviderType::OpenAI); - - // 检查是否需要拦截请求(legacy mode) - // **Validates: Requirements 2.1, 2.3, 2.5** - - // 检查是否需要刷新 token(无 token 或即将过期) - { - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - let needs_refresh = - kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); - if needs_refresh { - state.logs.write().await.add( - "info", - "[AUTH] No access token or token expiring soon, attempting refresh...", - ); - if let Err(e) = kiro.refresh_token().await { - state - .logs - .write() - .await - .add("error", &format!("[AUTH] Token refresh failed: {e}")); - // 标记 Flow 失败 - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ) - .into_response(); - } - state - .logs - .write() - .await - .add("info", "[AUTH] Token refreshed successfully"); - } - } - - // 转换为 OpenAI 格式 - let openai_request = convert_anthropic_to_openai(&request); - - // 记录转换后的请求信息 - state.logs.write().await.add( - "debug", + "error", &format!( - "[CONVERT] OpenAI format: messages={} tools={} stream={}", - openai_request.messages.len(), - openai_request.tools.as_ref().map(|t| t.len()).unwrap_or(0), - openai_request.stream + "[ROUTE] No API Key Provider credential found for '{effective_provider}' (client_type={client_type}), {reason}" ), ); - - let kiro = state.kiro.read().await; - - match kiro.call_api(&openai_request).await { - Ok(resp) => { - let status = resp.status(); - state - .logs - .write() - .await - .add("info", &format!("[RESP] Upstream status: {status}")); - - if status.is_success() { - match resp.bytes().await { - Ok(bytes) => { - // 使用 lossy 转换,避免无效 UTF-8 导致崩溃 - let body = String::from_utf8_lossy(&bytes).to_string(); - - // 记录原始响应长度 - state.logs.write().await.add( - "debug", - &format!("[RESP] Raw body length: {} bytes", bytes.len()), - ); - - // 保存原始响应到文件用于调试 - let request_id = uuid::Uuid::new_v4().to_string()[..8].to_string(); - state.logs.read().await.log_raw_response(&request_id, &body); - state.logs.write().await.add( - "debug", - &format!("[RESP] Raw response saved to raw_response_{request_id}.txt"), - ); - - // 记录响应的前200字符用于调试(减少日志量) - let preview: String = - body.chars().filter(|c| !c.is_control()).take(200).collect(); - state - .logs - .write() - .await - .add("debug", &format!("[RESP] Body preview: {preview}")); - - let parsed = parse_cw_response(&body); - - // 详细记录解析结果 - state.logs.write().await.add( - "info", - &format!( - "[RESP] Parsed: content_len={}, tool_calls={}, content_preview={}", - parsed.content.len(), - parsed.tool_calls.len(), - parsed.content.chars().take(100).collect::() - ), - ); - - // 记录 tool calls 详情 - for (i, tc) in parsed.tool_calls.iter().enumerate() { - state.logs.write().await.add( - "debug", - &format!( - "[RESP] Tool call {}: name={} id={}", - i, tc.function.name, tc.id - ), - ); - } - - // 如果请求流式响应,返回 SSE 格式 - if request.stream { - // 完成 Flow 捕获并检查响应拦截(流式) - // **Validates: Requirements 2.1, 2.5** - return build_anthropic_stream_response(&request.model, &parsed); - } - - // 完成 Flow 捕获并检查响应拦截(非流式) - // **Validates: Requirements 2.1, 2.5** - - // 非流式响应 - let response = build_anthropic_response(&request.model, &parsed); - return attach_route_debug_headers( - finalize_replayable_response( - response, - &mut idempotency_guard, - &mut dedup_guard, - &mut cache_guard, - &ctx.request_id, - ) - .await, - &selected_provider, - &effective_provider, - &ctx.resolved_model, - ); - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("[ERROR] Response body read failed: {e}")); - // 标记 Flow 失败 - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } else if status.as_u16() == 403 || status.as_u16() == 402 { - // Token 过期或账户问题,尝试重新加载凭证并刷新 - drop(kiro); - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - state.logs.write().await.add( - "warn", - &format!( - "[AUTH] Got {}, reloading credentials and attempting token refresh...", - status.as_u16() - ), - ); - - // 先重新加载凭证文件(可能用户换了账户) - if let Err(e) = kiro.load_credentials().await { - state.logs.write().await.add( - "error", - &format!("[AUTH] Failed to reload credentials: {e}"), - ); - } - - match kiro.refresh_token().await { - Ok(_) => { - state.logs.write().await.add( - "info", - "[AUTH] Token refreshed successfully, retrying request...", - ); - drop(kiro); - let kiro = state.kiro.read().await; - match kiro.call_api(&openai_request).await { - Ok(retry_resp) => { - let retry_status = retry_resp.status(); - state.logs.write().await.add( - "info", - &format!("[RETRY] Response status: {retry_status}"), - ); - if retry_resp.status().is_success() { - match retry_resp.bytes().await { - Ok(bytes) => { - let body = String::from_utf8_lossy(&bytes).to_string(); - let parsed = parse_cw_response(&body); - state.logs.write().await.add( - "info", - &format!( - "[RETRY] Success: content_len={}, tool_calls={}", - parsed.content.len(), parsed.tool_calls.len() - ), - ); - // 完成 Flow 捕获并检查响应拦截(重试成功) - // **Validates: Requirements 2.1, 2.5** - if request.stream { - return build_anthropic_stream_response( - &request.model, - &parsed, - ); - } - let response = - build_anthropic_response(&request.model, &parsed); - return attach_route_debug_headers( - finalize_replayable_response( - response, - &mut idempotency_guard, - &mut dedup_guard, - &mut cache_guard, - &ctx.request_id, - ) - .await, - &selected_provider, - &effective_provider, - &ctx.resolved_model, - ); - } - Err(e) => { - state.logs.write().await.add( - "error", - &format!("[RETRY] Body read failed: {e}"), - ); - // 标记 Flow 失败 - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(); - } - } - } - let body = retry_resp - .bytes() - .await - .map(|b| String::from_utf8_lossy(&b).to_string()) - .unwrap_or_default(); - state.logs.write().await.add( - "error", - &format!( - "[RETRY] Failed with status {retry_status}: {}", - safe_truncate(&body, 500) - ), - ); - // 标记 Flow 失败(重试失败) - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), - ) - .into_response() - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("[RETRY] Request failed: {e}")); - // 标记 Flow 失败 - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("[AUTH] Token refresh failed: {e}")); - // 标记 Flow 失败 - ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ) - .into_response() - } - } - } else { - let body = resp.text().await.unwrap_or_default(); - state.logs.write().await.add( - "error", - &format!( - "[ERROR] Upstream error HTTP {}: {}", - status, - safe_truncate(&body, 500) - ), - ); - // 标记 Flow 失败 - ( - StatusCode::from_u16(status.as_u16()) - .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json( - serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}}), - ), - ) - .into_response() - } - } - Err(e) => { - // 详细记录网络/连接错误 - let error_details = format!("{e:?}"); - state - .logs - .write() - .await - .add("error", &format!("[ERROR] Kiro API call failed: {e}")); - state.logs.write().await.add( - "debug", - &format!("[ERROR] Full error details: {error_details}"), - ); - // 标记 Flow 失败 - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } + let message = if !state.allow_provider_fallback { + format!( + "没有找到可用的 '{}' API Key Provider 凭证(已禁用自动降级)。", + effective_provider + ) + } else { + format!( + "没有找到可用的 '{}' API Key Provider 凭证。", + effective_provider + ) + }; + let body = build_gateway_error_json( + StatusCode::SERVICE_UNAVAILABLE.as_u16(), + &message, + Some(&ctx.request_id), + Some(&effective_provider), + Some(GatewayErrorCode::NoCredentials), + ); + ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ "type": "error", "error": body["error"].clone() })), + ) + .into_response() } // ============================================================================ @@ -3605,17 +2932,11 @@ fn should_use_true_streaming( ) -> bool { use lime_core::models::provider_pool_model::CredentialData; - // TODO: 当 StreamingProvider trait 实现后,根据凭证类型返回 true - // 目前所有 Provider 都使用伪流式模式 match &credential.credential { - // Kiro/CodeWhisperer - 需要实现 StreamingProvider - CredentialData::KiroOAuth { .. } => false, // Claude - 需要实现 StreamingProvider CredentialData::ClaudeKey { .. } => false, // OpenAI - 需要实现 StreamingProvider CredentialData::OpenAIKey { .. } => false, - // Antigravity - 需要实现 StreamingProvider - CredentialData::AntigravityOAuth { .. } => false, // 其他类型暂不支持流式 _ => false, } diff --git a/src-tauri/crates/server/src/handlers/credentials_api.rs b/src-tauri/crates/server/src/handlers/credentials_api.rs index 1fca43813..4d5bfd026 100644 --- a/src-tauri/crates/server/src/handlers/credentials_api.rs +++ b/src-tauri/crates/server/src/handlers/credentials_api.rs @@ -1,9 +1,6 @@ //! 凭证 API 端点(用于 aster Agent 集成) //! -//! 为 aster 子进程提供凭证查询接口,支持多种凭证类型: -//! - OAuth 凭证(Kiro, Gemini, Qwen, Antigravity 等) -//! - API Key Provider(OpenAI, Anthropic, Gemini API Key 等) -//! - OAuth 插件凭证(动态加载的第三方插件) +//! 为 aster 子进程提供 API Key Provider 查询接口。 //! //! 此 API 仅供内部使用,返回完整的凭证信息(包括未脱敏的 access_token)。 @@ -18,8 +15,6 @@ use serde::{Deserialize, Serialize}; use crate::AppState; use lime_core::database::dao::api_key_provider::ApiKeyProviderDao; -use lime_core::database::dao::provider_pool::ProviderPoolDao; -use lime_core::models::provider_pool_model::PoolProviderType; use super::api_key_provider_utils::{build_api_key_headers, collect_api_key_provider_ids}; @@ -32,8 +27,7 @@ pub struct SelectCredentialRequest { /// 指定模型(可选) #[serde(skip_serializing_if = "Option::is_none")] pub model: Option, - /// 凭证来源偏好(可选):oauth, api_key, plugin - /// 如果不指定,会按优先级自动选择 + /// 凭证来源偏好(可选):当前只支持 api_key #[serde(skip_serializing_if = "Option::is_none")] pub source_preference: Option, } @@ -42,12 +36,8 @@ pub struct SelectCredentialRequest { #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] #[serde(rename_all = "snake_case")] pub enum CredentialType { - /// OAuth 凭证(凭证池) - OAuth, /// API Key(API Key Provider) ApiKey, - /// OAuth 插件凭证 - Plugin, } /// 凭证信息响应 @@ -92,15 +82,7 @@ impl IntoResponse for CredentialApiError { /// POST /v1/credentials/select - 选择可用凭证 /// -/// 支持多种凭证来源: -/// 1. OAuth 凭证池(Kiro, Gemini, Qwen, Antigravity 等) -/// 2. API Key Provider(OpenAI, Anthropic, Gemini API Key 等) -/// 3. OAuth 插件凭证(动态加载的第三方插件) -/// -/// 选择优先级(如果未指定 source_preference): -/// 1. 首先尝试 OAuth 凭证池 -/// 2. 然后尝试 API Key Provider -/// 3. 最后尝试 OAuth 插件 +/// 只从 API Key Provider 主路径选择凭证。 pub async fn credentials_select( State(state): State, _headers: HeaderMap, @@ -119,16 +101,8 @@ pub async fn credentials_select( status_code: 503, })?; - // 根据 source_preference 决定选择策略 let source_pref = request.source_preference.as_deref(); - // 尝试从 OAuth 凭证池选择 - if source_pref.is_none() || source_pref == Some("oauth") { - if let Some(response) = try_select_oauth_credential(&state, db, &request).await? { - return Ok(Json(response)); - } - } - // 尝试从 API Key Provider 选择(智能降级) if source_pref.is_none() || source_pref == Some("api_key") { if let Some(response) = try_select_api_key_credential(&state, db, &request).await? { @@ -136,77 +110,17 @@ pub async fn credentials_select( } } - // 尝试从 OAuth 插件选择 - if source_pref.is_none() || source_pref == Some("plugin") { - if let Some(response) = try_select_plugin_credential(&state, &request).await? { - return Ok(Json(response)); - } - } - // 没有找到可用凭证 Err(CredentialApiError { error: "no_available_credentials".to_string(), message: format!( - "没有可用的 {} 凭证。您可以在 API Key Provider 中配置 API Key 作为降级选项。", + "没有可用的 {} API Key Provider 凭证。", request.provider_type ), status_code: 503, }) } -/// 尝试从 OAuth 凭证池选择凭证 -async fn try_select_oauth_credential( - state: &AppState, - db: &lime_core::database::DbConnection, - request: &SelectCredentialRequest, -) -> Result, CredentialApiError> { - // 使用 ProviderPoolService 智能选择凭证 - let credential = match state.pool_service.select_credential( - db, - &request.provider_type, - request.model.as_deref(), - ) { - Ok(Some(cred)) => cred, - Ok(None) => return Ok(None), - Err(_) => return Ok(None), - }; - - // 获取 access_token - let access_token = match credential - .cached_token - .as_ref() - .and_then(|cache| cache.access_token.clone()) - { - Some(token) => token, - None => return Ok(None), - }; - - // 根据 Provider 类型确定 base_url - let base_url = get_oauth_base_url(&credential.provider_type); - - let response = CredentialResponse { - uuid: credential.uuid.clone(), - provider_type: credential.provider_type.to_string(), - credential_type: CredentialType::OAuth, - access_token, - base_url, - expires_at: credential - .cached_token - .as_ref() - .and_then(|cache| cache.expiry_time), - name: credential.name.clone(), - extra_headers: None, - }; - - tracing::info!( - "[CREDENTIALS_API] OAuth 凭证选择成功: {} ({})", - response.name.as_deref().unwrap_or("未命名"), - response.uuid - ); - - Ok(Some(response)) -} - /// 尝试从 API Key Provider 选择凭证 async fn try_select_api_key_credential( state: &AppState, @@ -268,33 +182,8 @@ async fn try_select_api_key_credential( Ok(None) } -/// 尝试从 OAuth 插件选择凭证(已禁用 - 插件系统已移除) -async fn try_select_plugin_credential( - _state: &AppState, - _request: &SelectCredentialRequest, -) -> Result, CredentialApiError> { - // OAuth 插件系统已移除,直接返回 None - Ok(None) -} - -/// 根据 OAuth Provider 类型获取 base_url -fn get_oauth_base_url(provider_type: &PoolProviderType) -> String { - match provider_type { - PoolProviderType::Kiro => "https://api.anthropic.com".to_string(), - PoolProviderType::Gemini => "https://generativelanguage.googleapis.com".to_string(), - PoolProviderType::Antigravity => "https://api.anthropic.com".to_string(), - PoolProviderType::Vertex => "https://vertex-ai.googleapis.com".to_string(), - PoolProviderType::GeminiApiKey => "https://generativelanguage.googleapis.com".to_string(), - PoolProviderType::Codex => "https://api.openai.com/v1".to_string(), - PoolProviderType::ClaudeOAuth => "https://api.anthropic.com".to_string(), - _ => "https://api.openai.com/v1".to_string(), - } -} - /// GET /v1/credentials/{uuid}/token - 获取指定凭证的 Token /// -/// 支持多种凭证类型: -/// - OAuth 凭证池中的凭证 /// - API Key Provider 中的 API Key pub async fn credentials_get_token( State(state): State, @@ -309,12 +198,6 @@ pub async fn credentials_get_token( status_code: 503, })?; - // 首先尝试从 OAuth 凭证池查询 - if let Some(response) = try_get_oauth_token(&state, db, &uuid).await? { - return Ok(Json(response)); - } - - // 然后尝试从 API Key Provider 查询 if let Some(response) = try_get_api_key_token(&state, db, &uuid).await? { return Ok(Json(response)); } @@ -327,111 +210,6 @@ pub async fn credentials_get_token( }) } -/// 尝试从 OAuth 凭证池获取 Token -async fn try_get_oauth_token( - state: &AppState, - db: &lime_core::database::DbConnection, - uuid: &str, -) -> Result, CredentialApiError> { - // 查询凭证 - let credential = { - let conn = db.lock().map_err(|e| CredentialApiError { - error: "database_lock_error".to_string(), - message: format!("数据库锁定失败: {e}"), - status_code: 500, - })?; - - match ProviderPoolDao::get_by_uuid(&conn, uuid) { - Ok(Some(cred)) => cred, - Ok(None) => return Ok(None), - Err(_) => return Ok(None), - } - }; - - // 如果 Token 即将过期,尝试刷新 - let cached_token = if let Some(cache) = &credential.cached_token { - if let Some(expiry_time) = cache.expiry_time { - let now = chrono::Utc::now(); - let time_until_expiry = expiry_time - now; - - // 如果距离过期不到 30 分钟,尝试刷新 - if time_until_expiry < chrono::Duration::minutes(30) { - tracing::info!("[CREDENTIALS_API] Token 即将过期,尝试刷新: {}", uuid); - match state - .token_cache - .refresh_and_cache_with_events( - db, - uuid, - false, - Some(state.kiro_event_service.clone()), - ) - .await - { - Ok(new_token) => { - tracing::info!("[CREDENTIALS_API] Token 刷新成功: {}", uuid); - Some(new_token) - } - Err(e) => { - tracing::warn!("[CREDENTIALS_API] Token 刷新失败,使用现有 Token: {}", e); - cache.access_token.clone() - } - } - } else { - cache.access_token.clone() - } - } else { - cache.access_token.clone() - } - } else { - None - }; - - let access_token = match cached_token { - Some(token) => token, - None => return Ok(None), - }; - - // 根据 Provider 类型确定 base_url - let base_url = get_oauth_base_url(&credential.provider_type); - - // 重新查询凭证以获取更新后的 expires_at - let updated_credential = { - let conn = db.lock().map_err(|e| CredentialApiError { - error: "database_lock_error".to_string(), - message: format!("数据库锁定失败: {e}"), - status_code: 500, - })?; - - match ProviderPoolDao::get_by_uuid(&conn, uuid) { - Ok(Some(cred)) => cred, - Ok(None) => return Ok(None), - Err(_) => return Ok(None), - } - }; - - let response = CredentialResponse { - uuid: updated_credential.uuid.clone(), - provider_type: updated_credential.provider_type.to_string(), - credential_type: CredentialType::OAuth, - access_token, - base_url, - expires_at: updated_credential - .cached_token - .as_ref() - .and_then(|cache| cache.expiry_time), - name: updated_credential.name.clone(), - extra_headers: None, - }; - - tracing::info!( - "[CREDENTIALS_API] 返回 OAuth 凭证 Token: {} ({})", - response.name.as_deref().unwrap_or("未命名"), - response.uuid - ); - - Ok(Some(response)) -} - /// 尝试从 API Key Provider 获取 Token async fn try_get_api_key_token( state: &AppState, diff --git a/src-tauri/crates/server/src/handlers/image_handler.rs b/src-tauri/crates/server/src/handlers/image_handler.rs index 14a6535c5..ed92df3a0 100644 --- a/src-tauri/crates/server/src/handlers/image_handler.rs +++ b/src-tauri/crates/server/src/handlers/image_handler.rs @@ -1,19 +1,18 @@ //! 图像生成 API 处理器 //! //! 实现 OpenAI 兼容的 `/v1/images/generations` 端点, -//! 通过 Antigravity Provider 调用 Gemini 图像生成模型。 +//! 通过已配置的 API Key Provider 调用图像生成模型。 //! //! # 功能 //! - 接收 OpenAI 格式的图像生成请求 -//! - 转换为 Antigravity/Gemini 格式 -//! - 调用 Antigravity Provider +//! - 调用 API Key Provider //! - 返回 OpenAI 格式的响应 //! //! # 需求覆盖 //! - 需求 1.1: 实现 `/v1/images/generations` 端点 //! - 需求 4.1: 验证请求参数 -//! - 需求 4.2: 获取 Antigravity 凭证 -//! - 需求 4.3: 调用 Antigravity Provider +//! - 需求 4.2: 获取当前 API Key Provider +//! - 需求 4.3: 调用 Provider //! - 需求 4.4: 转换响应格式 use axum::{ @@ -27,11 +26,6 @@ use super::image_api_provider; use crate::handlers::verify_api_key; use crate::AppState; use lime_core::models::openai::ImageGenerationRequest; -use lime_core::models::provider_pool_model::CredentialData; -use lime_providers::converter::openai_to_antigravity::{ - convert_antigravity_image_response, convert_image_request_to_antigravity, -}; -use lime_providers::providers::AntigravityProvider; fn read_explicit_provider_id(headers: &HeaderMap) -> Option { headers @@ -128,10 +122,22 @@ pub async fn handle_image_generation( return (StatusCode::OK, Json(response)).into_response(); } Ok(None) => { - state.logs.write().await.add( - "debug", - "[IMAGE] 图片服务未命中当前路由,继续回退到 Antigravity 兼容链路", - ); + state + .logs + .write() + .await + .add("debug", "[IMAGE] 图片服务未命中当前 API Key Provider 路由"); + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "message": "No image-capable API Key Provider configured", + "type": "server_error", + "code": "no_image_provider" + } + })), + ) + .into_response(); } Err(error) => { state @@ -152,253 +158,4 @@ pub async fn handle_image_generation( .into_response(); } } - - // 获取 Antigravity 凭证 - let db = match &state.db { - Some(db) => db, - None => { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": "Database not available", - "type": "server_error" - } - })), - ) - .into_response(); - } - }; - - // 从凭证池获取 Antigravity 凭证 - let credential = match state - .pool_service - .select_credential(db, "antigravity", None) - { - Ok(Some(cred)) => cred, - Ok(None) => { - state - .logs - .write() - .await - .add("error", "[IMAGE] 没有可用的 Antigravity 凭证"); - return ( - StatusCode::SERVICE_UNAVAILABLE, - Json(serde_json::json!({ - "error": { - "message": "No Antigravity credentials available for image generation", - "type": "server_error", - "code": "no_credentials" - } - })), - ) - .into_response(); - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("[IMAGE] 获取凭证失败: {e}")); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": format!("Failed to get credentials: {}", e), - "type": "server_error" - } - })), - ) - .into_response(); - } - }; - - // 提取 Antigravity 凭证信息 - let (creds_file_path, project_id) = match &credential.credential { - CredentialData::AntigravityOAuth { - creds_file_path, - project_id, - } => (creds_file_path.clone(), project_id.clone()), - _ => { - state - .logs - .write() - .await - .add("error", "[IMAGE] 选中的凭证不是 Antigravity 类型"); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": "Selected credential is not Antigravity type", - "type": "server_error" - } - })), - ) - .into_response(); - } - }; - - // 创建 Antigravity Provider - let mut antigravity = AntigravityProvider::new(); - if let Err(e) = antigravity - .load_credentials_from_path(&creds_file_path) - .await - { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": format!("Failed to load Antigravity credentials: {}", e), - "type": "server_error" - } - })), - ) - .into_response(); - } - - // 验证并刷新 Token - let validation_result = antigravity.validate_token(); - if validation_result.needs_refresh() { - tracing::info!("[IMAGE] Token 需要刷新,开始刷新..."); - if let Err(refresh_error) = antigravity.refresh_token_with_retry(3).await { - tracing::error!("[IMAGE] Token 刷新失败: {:?}", refresh_error); - let _ = state.pool_service.mark_unhealthy_with_details( - db, - &credential.uuid, - &refresh_error, - ); - let (status, message) = if refresh_error.requires_reauth() { - (StatusCode::UNAUTHORIZED, refresh_error.user_message()) - } else { - ( - StatusCode::INTERNAL_SERVER_ERROR, - refresh_error.user_message(), - ) - }; - return ( - status, - Json(serde_json::json!({ - "error": { - "message": message, - "type": "authentication_error" - } - })), - ) - .into_response(); - } - } - - // 设置项目 ID - if let Some(pid) = project_id { - antigravity.project_id = Some(pid); - } else if let Err(e) = antigravity.discover_project().await { - tracing::warn!("[IMAGE] Failed to discover project: {}", e); - } - - let proj_id = antigravity.project_id.clone().unwrap_or_default(); - - // 转换请求为 Antigravity 格式 - let antigravity_request = convert_image_request_to_antigravity(&request, &proj_id); - - state.logs.write().await.add( - "debug", - &format!( - "[IMAGE] Antigravity 请求: model={}", - antigravity_request["model"].as_str().unwrap_or("unknown") - ), - ); - - // 调用 Antigravity API - 直接使用 call_api 而不是 generate_content - // 因为 generate_content 内部的 to_gemini_response 会丢失嵌套在 response 字段下的数据 - let model = antigravity_request["model"] - .as_str() - .unwrap_or("gemini-3-pro-image-preview"); - - eprintln!("[IMAGE] 调用 Antigravity API: model={model}"); - eprintln!( - "[IMAGE] 请求内容: {}", - serde_json::to_string_pretty(&antigravity_request).unwrap_or_default() - ); - - match antigravity - .call_api("generateContent", &antigravity_request) - .await - { - Ok(resp) => { - // 调试:打印原始响应 - eprintln!( - "[IMAGE] Antigravity 原始响应: {}", - serde_json::to_string_pretty(&resp).unwrap_or_default() - ); - state.logs.write().await.add( - "debug", - &format!( - "[IMAGE] Antigravity 原始响应: {}", - serde_json::to_string(&resp).unwrap_or_default() - ), - ); - - // 转换响应为 OpenAI 格式 - match convert_antigravity_image_response(&resp, &request.response_format) { - Ok(image_response) => { - // 记录成功 - let _ = state - .pool_service - .mark_healthy(db, &credential.uuid, Some(model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); - - state.logs.write().await.add( - "info", - &format!("[IMAGE] 图像生成成功: {} 张图片", image_response.data.len()), - ); - - (StatusCode::OK, Json(image_response)).into_response() - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("[IMAGE] 响应转换失败: {e}")); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": e, - "type": "server_error", - "code": "image_generation_failed" - } - })), - ) - .into_response() - } - } - } - Err(e) => { - let _ = state - .pool_service - .mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); - state - .logs - .write() - .await - .add("error", &format!("[IMAGE] Antigravity API 调用失败: {e}")); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": format!("Image generation failed: {}", e), - "type": "server_error", - "code": "api_error" - } - })), - ) - .into_response() - } - } } diff --git a/src-tauri/crates/server/src/handlers/kiro_credential.rs b/src-tauri/crates/server/src/handlers/kiro_credential.rs deleted file mode 100644 index f2080b69b..000000000 --- a/src-tauri/crates/server/src/handlers/kiro_credential.rs +++ /dev/null @@ -1,655 +0,0 @@ -//! Kiro凭证管理API处理器 -//! -//! 为kiro凭证池管理提供REST API端点,支持: -//! - 获取可用凭证列表 -//! - 智能选择凭证 -//! - 手动刷新凭证 -//! - 凭证状态查询 - -use axum::{ - extract::{Path, State}, - http::{HeaderMap, StatusCode}, - response::{IntoResponse, Response}, - Json, -}; -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; - -use crate::AppState; -use lime_core::database::dao::provider_pool::ProviderPoolDao; -use lime_core::models::provider_pool_model::{ - CachedTokenInfo, PoolProviderType, ProviderCredential, -}; - -/// 可用凭证信息 -#[derive(Debug, Clone, Serialize)] -pub struct AvailableCredential { - /// 凭证UUID - pub uuid: String, - /// 凭证名称 - pub name: String, - /// 是否可用 - pub available: bool, - /// Token过期时间 - pub expires_at: Option>, - /// 最后使用时间 - pub last_used: Option>, - /// 健康状态分数 (0-100) - pub health_score: f64, - /// 错误计数 - pub error_count: u32, - /// 最后错误信息 - pub last_error: Option, -} - -/// 获取可用凭证列表的响应 -#[derive(Debug, Serialize)] -pub struct AvailableCredentialsResponse { - /// 可用凭证列表 - pub credentials: Vec, - /// 总凭证数 - pub total: usize, - /// 可用凭证数 - pub available: usize, - /// 系统状态 - pub status: String, -} - -/// 选择凭证请求参数 -#[derive(Debug, Deserialize)] -pub struct SelectCredentialRequest { - /// 指定模型(可选) - pub model: Option, - /// 强制选择特定UUID(可选) - pub force_uuid: Option, -} - -/// 选择凭证响应 -#[derive(Debug, Serialize)] -pub struct SelectCredentialResponse { - /// 选中的凭证UUID - pub uuid: String, - /// 凭证名称 - pub name: String, - /// Access Token(脱敏显示) - pub access_token_preview: String, - /// Token过期时间 - pub expires_at: Option>, - /// 选择原因 - pub selection_reason: String, -} - -/// 刷新凭证响应 -#[derive(Debug, Serialize)] -pub struct RefreshCredentialResponse { - /// 凭证UUID - pub uuid: String, - /// 刷新是否成功 - pub success: bool, - /// 新的过期时间 - pub new_expires_at: Option>, - /// 刷新结果信息 - pub message: String, - /// 错误信息(如果有) - pub error: Option, -} - -/// API错误响应 -#[derive(Debug, Serialize)] -pub struct ApiError { - pub error: String, - pub message: String, - pub status_code: u16, -} - -impl IntoResponse for ApiError { - fn into_response(self) -> Response { - let status = - StatusCode::from_u16(self.status_code).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); - (status, Json(self)).into_response() - } -} - -/// GET /api/kiro/credentials/available - 获取可用凭证列表 -pub async fn get_available_credentials( - State(state): State, - _headers: HeaderMap, -) -> Result, ApiError> { - tracing::info!("[KIRO_API] 获取可用凭证列表请求"); - - let db = &state.db.as_ref().ok_or_else(|| ApiError { - error: "database_unavailable".to_string(), - message: "数据库连接不可用".to_string(), - status_code: 503, - })?; - - let _pool_service = &state.pool_service; - let token_cache = &state.token_cache; - - // 获取所有kiro凭证 - let credentials = { - let conn = db.lock().map_err(|e| ApiError { - error: "database_lock_error".to_string(), - message: format!("数据库锁定失败: {e}"), - status_code: 500, - })?; - - ProviderPoolDao::get_all(&conn) - .map_err(|e| ApiError { - error: "database_query_error".to_string(), - message: format!("查询凭证失败: {e}"), - status_code: 500, - })? - .into_iter() - .filter(|cred| cred.provider_type == PoolProviderType::Kiro) - .collect::>() - }; - - let mut available_credentials = Vec::new(); - let mut available_count = 0; - - for credential in &credentials { - // 获取凭证缓存状态 - let cache_status = token_cache - .get_cache_status(db, &credential.uuid) - .map_err(|e| ApiError { - error: "cache_query_error".to_string(), - message: format!("获取缓存状态失败: {e}"), - status_code: 500, - })?; - - // 计算健康状态分数 - let health_score = calculate_health_score(credential, cache_status.as_ref()); - - let is_available = health_score > 50.0; // 健康分数大于50认为可用 - if is_available { - available_count += 1; - } - - let available_cred = AvailableCredential { - uuid: credential.uuid.clone(), - name: credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()), - available: is_available, - expires_at: cache_status.as_ref().and_then(|c| c.expiry_time), - last_used: credential.last_used, - health_score, - error_count: cache_status - .as_ref() - .map(|c| c.refresh_error_count) - .unwrap_or(0), - last_error: cache_status.and_then(|c| c.last_refresh_error), - }; - - available_credentials.push(available_cred); - } - - // 按健康分数降序排列 - available_credentials.sort_by(|a, b| { - b.health_score - .partial_cmp(&a.health_score) - .unwrap_or(std::cmp::Ordering::Equal) - }); - - let response = AvailableCredentialsResponse { - credentials: available_credentials, - total: credentials.len(), - available: available_count, - status: if available_count > 0 { - "healthy".to_string() - } else { - "degraded".to_string() - }, - }; - - tracing::info!( - "[KIRO_API] 返回{}个凭证,其中{}个可用", - response.total, - response.available - ); - Ok(Json(response)) -} - -/// POST /api/kiro/credentials/select - 智能选择凭证 -pub async fn select_credential( - State(state): State, - _headers: HeaderMap, - Json(request): Json, -) -> Result, ApiError> { - tracing::info!( - "[KIRO_API] 选择凭证请求,模型: {:?}, 强制UUID: {:?}", - request.model, - request.force_uuid - ); - - let db = &state.db.as_ref().ok_or_else(|| ApiError { - error: "database_unavailable".to_string(), - message: "数据库连接不可用".to_string(), - status_code: 503, - })?; - - let selected_credential = if let Some(ref force_uuid) = request.force_uuid { - // 强制选择指定UUID - let conn = db.lock().map_err(|e| ApiError { - error: "database_lock_error".to_string(), - message: format!("数据库锁定失败: {e}"), - status_code: 500, - })?; - - ProviderPoolDao::get_by_uuid(&conn, force_uuid) - .map_err(|e| ApiError { - error: "database_query_error".to_string(), - message: format!("查询凭证失败: {e}"), - status_code: 500, - })? - .ok_or_else(|| ApiError { - error: "credential_not_found".to_string(), - message: format!("未找到UUID为{force_uuid}的凭证"), - status_code: 404, - })? - } else { - // 智能选择最优凭证(Kiro 是 OAuth-only,不支持降级到 API Key) - state - .pool_service - .select_credential(db, "kiro", request.model.as_deref()) - .map_err(|e| ApiError { - error: "selection_error".to_string(), - message: format!("凭证选择失败: {e}"), - status_code: 500, - })? - .ok_or_else(|| ApiError { - error: "no_available_credentials".to_string(), - message: "没有可用的 Kiro 凭证。Kiro 仅支持 OAuth 认证,无法降级到 API Key。" - .to_string(), - status_code: 503, - })? - }; - - // 脱敏显示token - let token_preview = if let Some(cached_token_info) = &selected_credential.cached_token { - if let Some(token) = &cached_token_info.access_token { - if token.len() > 20 { - format!("{}...{}", &token[..10], &token[token.len() - 10..]) - } else { - "***".to_string() - } - } else { - "无token".to_string() - } - } else { - "未缓存".to_string() - }; - - let selection_reason = if request.force_uuid.is_some() { - "手动指定".to_string() - } else { - "智能算法选择".to_string() - }; - - let expires_at = selected_credential - .cached_token - .as_ref() - .and_then(|cache| cache.expiry_time); - - let response = SelectCredentialResponse { - uuid: selected_credential.uuid.clone(), - name: selected_credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()), - access_token_preview: token_preview, - expires_at, - selection_reason, - }; - - tracing::info!( - "[KIRO_API] 选择凭证成功: {} ({})", - response.name, - response.uuid - ); - Ok(Json(response)) -} - -/// PUT /api/kiro/credentials/{uuid}/refresh - 手动刷新指定凭证 -pub async fn refresh_credential( - State(state): State, - Path(uuid): Path, - _headers: HeaderMap, -) -> Result, ApiError> { - tracing::info!("[KIRO_API] 刷新凭证请求: {}", uuid); - - let db = &state.db.as_ref().ok_or_else(|| ApiError { - error: "database_unavailable".to_string(), - message: "数据库连接不可用".to_string(), - status_code: 503, - })?; - - let token_cache = &state.token_cache; - - // 验证凭证存在且为kiro类型 - let credential = { - let conn = db.lock().map_err(|e| ApiError { - error: "database_lock_error".to_string(), - message: format!("数据库锁定失败: {e}"), - status_code: 500, - })?; - - let cred = ProviderPoolDao::get_by_uuid(&conn, &uuid) - .map_err(|e| ApiError { - error: "database_query_error".to_string(), - message: format!("查询凭证失败: {e}"), - status_code: 500, - })? - .ok_or_else(|| ApiError { - error: "credential_not_found".to_string(), - message: format!("未找到UUID为{uuid}的凭证"), - status_code: 404, - })?; - - if cred.provider_type.to_string() != PoolProviderType::Kiro.to_string() { - return Err(ApiError { - error: "invalid_credential_type".to_string(), - message: format!("凭证类型不是kiro: {}", cred.provider_type), - status_code: 400, - }); - } - - cred - }; - - // 执行强制刷新 - match token_cache - .refresh_and_cache_with_events(db, &uuid, true, Some(state.kiro_event_service.clone())) - .await - { - Ok(_new_token) => { - // 获取刷新后的缓存状态 - let cache_status = token_cache - .get_cache_status(db, &uuid) - .map_err(|e| ApiError { - error: "cache_query_error".to_string(), - message: format!("获取刷新后缓存状态失败: {e}"), - status_code: 500, - })?; - - let response = RefreshCredentialResponse { - uuid: uuid.clone(), - success: true, - new_expires_at: cache_status.as_ref().and_then(|c| c.expiry_time), - message: format!( - "凭证 {} 刷新成功", - credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()) - ), - error: None, - }; - - tracing::info!( - "[KIRO_API] 凭证刷新成功: {} ({})", - credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()), - uuid - ); - Ok(Json(response)) - } - Err(refresh_error) => { - let response = RefreshCredentialResponse { - uuid: uuid.clone(), - success: false, - new_expires_at: None, - message: format!( - "凭证 {} 刷新失败", - credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()) - ), - error: Some(refresh_error.clone()), - }; - - tracing::warn!( - "[KIRO_API] 凭证刷新失败: {} ({}): {}", - credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()), - uuid, - refresh_error - ); - Ok(Json(response)) - } - } -} - -/// GET /api/kiro/credentials/{uuid}/status - 获取凭证详细状态 -pub async fn get_credential_status( - State(state): State, - Path(uuid): Path, - _headers: HeaderMap, -) -> Result, ApiError> { - tracing::info!("[KIRO_API] 获取凭证状态: {}", uuid); - - let db = &state.db.as_ref().ok_or_else(|| ApiError { - error: "database_unavailable".to_string(), - message: "数据库连接不可用".to_string(), - status_code: 503, - })?; - - let token_cache = &state.token_cache; - - // 验证凭证存在 - let credential = { - let conn = db.lock().map_err(|e| ApiError { - error: "database_lock_error".to_string(), - message: format!("数据库锁定失败: {e}"), - status_code: 500, - })?; - - ProviderPoolDao::get_by_uuid(&conn, &uuid) - .map_err(|e| ApiError { - error: "database_query_error".to_string(), - message: format!("查询凭证失败: {e}"), - status_code: 500, - })? - .ok_or_else(|| ApiError { - error: "credential_not_found".to_string(), - message: format!("未找到UUID为{uuid}的凭证"), - status_code: 404, - })? - }; - - // 获取缓存状态 - let cache_status = token_cache - .get_cache_status(db, &uuid) - .map_err(|e| ApiError { - error: "cache_query_error".to_string(), - message: format!("获取缓存状态失败: {e}"), - status_code: 500, - })?; - - // 计算健康分数 - let health_score = calculate_health_score(&credential, cache_status.as_ref()); - - let mut status = serde_json::Map::new(); - status.insert( - "uuid".to_string(), - serde_json::Value::String(credential.uuid.clone()), - ); - status.insert( - "name".to_string(), - serde_json::Value::String( - credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()), - ), - ); - status.insert( - "provider_type".to_string(), - serde_json::Value::String(credential.provider_type.to_string()), - ); - status.insert( - "created_at".to_string(), - serde_json::Value::String(credential.created_at.to_rfc3339()), - ); - status.insert( - "last_used".to_string(), - credential - .last_used - .map(|dt| serde_json::Value::String(dt.to_rfc3339())) - .unwrap_or(serde_json::Value::Null), - ); - status.insert( - "health_score".to_string(), - serde_json::Value::Number( - serde_json::Number::from_f64(health_score).unwrap_or(serde_json::Number::from(0)), - ), - ); - status.insert( - "is_available".to_string(), - serde_json::Value::Bool(health_score > 50.0), - ); - - if let Some(cache) = cache_status { - status.insert( - "has_cached_token".to_string(), - serde_json::Value::Bool(cache.access_token.is_some()), - ); - status.insert( - "expires_at".to_string(), - cache - .expiry_time - .map(|dt| serde_json::Value::String(dt.to_rfc3339())) - .unwrap_or(serde_json::Value::Null), - ); - status.insert( - "last_refresh".to_string(), - cache - .last_refresh - .map(|dt| serde_json::Value::String(dt.to_rfc3339())) - .unwrap_or(serde_json::Value::Null), - ); - status.insert( - "refresh_error_count".to_string(), - serde_json::Value::Number(serde_json::Number::from(cache.refresh_error_count)), - ); - status.insert( - "last_refresh_error".to_string(), - cache - .last_refresh_error - .map(serde_json::Value::String) - .unwrap_or(serde_json::Value::Null), - ); - } else { - status.insert( - "has_cached_token".to_string(), - serde_json::Value::Bool(false), - ); - status.insert("expires_at".to_string(), serde_json::Value::Null); - status.insert("last_refresh".to_string(), serde_json::Value::Null); - status.insert( - "refresh_error_count".to_string(), - serde_json::Value::Number(serde_json::Number::from(0)), - ); - status.insert("last_refresh_error".to_string(), serde_json::Value::Null); - } - - tracing::info!( - "[KIRO_API] 返回凭证状态: {} (健康分数: {:.1})", - credential - .name - .clone() - .unwrap_or_else(|| "未命名".to_string()), - health_score - ); - Ok(Json(serde_json::Value::Object(status))) -} - -/// 计算凭证健康分数 -/// -/// 基于凭证的基本状态、缓存状态、错误计数等因素综合计算健康分数 -/// 分数范围: 0-100,分数越高表示凭证越健康 -fn calculate_health_score( - credential: &ProviderCredential, - cache_status: Option<&CachedTokenInfo>, -) -> f64 { - let mut score = 0.0; - - // 1. 基础健康状态 (40分) - if credential.is_healthy { - score += 40.0; - } else { - score -= 20.0; // 不健康严重扣分 - } - - // 2. 错误计数影响 (20分) - let error_count = credential.error_count; - if error_count == 0 { - score += 20.0; - } else if error_count <= 2 { - score += 10.0; // 少量错误,轻微扣分 - } else { - score -= error_count as f64 * 5.0; // 错误越多扣分越多 - } - - // 3. Token缓存状态 (25分) - if let Some(cache) = cache_status { - if cache.access_token.is_some() { - score += 15.0; // 有缓存token - - // 检查过期时间 - if let Some(expiry_time) = cache.expiry_time { - let now = chrono::Utc::now(); - let time_until_expiry = expiry_time - now; - - if time_until_expiry > chrono::Duration::hours(1) { - score += 10.0; // 距离过期还有较长时间 - } else if time_until_expiry > chrono::Duration::minutes(30) { - score += 5.0; // 距离过期还有一些时间 - } else if time_until_expiry <= chrono::Duration::zero() { - score -= 10.0; // 已过期 - } - } - } else { - score -= 5.0; // 没有缓存token - } - - // 刷新错误计数影响 - if cache.refresh_error_count == 0 { - // 无刷新错误,不加分不减分 - } else if cache.refresh_error_count <= 2 { - score -= cache.refresh_error_count as f64 * 2.0; // 少量刷新错误 - } else { - score -= cache.refresh_error_count as f64 * 5.0; // 大量刷新错误严重扣分 - } - } else { - score -= 10.0; // 完全没有缓存状态 - } - - // 4. 使用活跃度 (15分) - if let Some(last_used) = credential.last_used { - let now = chrono::Utc::now(); - let time_since_used = now - last_used; - - if time_since_used <= chrono::Duration::hours(1) { - score += 15.0; // 最近1小时内使用过 - } else if time_since_used <= chrono::Duration::hours(24) { - score += 10.0; // 最近24小时内使用过 - } else if time_since_used <= chrono::Duration::days(7) { - score += 5.0; // 最近一周内使用过 - } else { - score += 0.0; // 很久未使用,不扣分也不加分 - } - } else { - score -= 5.0; // 从未使用过 - } - - // 确保分数在0-100范围内 - score.max(0.0).min(100.0) -} diff --git a/src-tauri/crates/server/src/handlers/mod.rs b/src-tauri/crates/server/src/handlers/mod.rs index 1615f9871..eb4ddd5c5 100644 --- a/src-tauri/crates/server/src/handlers/mod.rs +++ b/src-tauri/crates/server/src/handlers/mod.rs @@ -8,7 +8,6 @@ pub mod chrome_bridge_ws; pub mod credentials_api; pub(crate) mod image_api_provider; pub mod image_handler; -pub mod kiro_credential; pub mod provider_calls; pub mod websocket; @@ -16,11 +15,5 @@ pub use api::*; pub use chrome_bridge_ws::*; pub use credentials_api::*; pub use image_handler::*; -// 避免 SelectCredentialRequest 歧义 glob re-export(credentials_api 和 kiro_credential 都定义了同名类型) -pub use kiro_credential::{ - get_available_credentials, get_credential_status, refresh_credential, select_credential, - AvailableCredential, AvailableCredentialsResponse, RefreshCredentialResponse, - SelectCredentialResponse, -}; pub use provider_calls::*; pub use websocket::*; diff --git a/src-tauri/crates/server/src/handlers/provider_calls.rs b/src-tauri/crates/server/src/handlers/provider_calls.rs index 746d5c597..388a3134f 100644 --- a/src-tauri/crates/server/src/handlers/provider_calls.rs +++ b/src-tauri/crates/server/src/handlers/provider_calls.rs @@ -10,15 +10,6 @@ //! - `StreamManager`: 管理流式请求的生命周期 //! - `StreamingProvider`: Provider 的流式 API 接口 //! - `FlowMonitor`: 实时捕获流式响应 -//! - `handle_kiro_stream()`: Kiro 凭证的真正流式处理(AWS Event Stream → Anthropic SSE) -//! -//! # Kiro 凭证流式处理 -//! -//! 当使用 Kiro 凭证且 `stream=true` 时,系统会: -//! 1. 调用 `KiroProvider.call_api_stream()` 获取 AWS Event Stream 格式的流式响应 -//! 2. 使用 `AwsEventStreamParser` 实时解析每个 JSON payload -//! 3. 使用 `AnthropicSseGenerator` 转换为 Anthropic SSE 格式 -//! 4. 通过 `FlowMonitor.process_chunk()` 记录每个 chunk //! //! # 错误处理 //! @@ -38,7 +29,7 @@ //! - 需求 5.1: 流式传输期间发生网络错误时,发出错误事件并以失败状态完成 flow //! - 需求 5.2: AWS Event Stream 解析失败时记录错误并继续处理后续 chunks //! - 需求 5.3: 将上游 Provider 返回的错误转发给客户端 -//! - 需求 6.1: 流式请求使用 handle_kiro_stream() +//! - 需求 6.1: 流式请求直接走 current API Key Provider 主链 //! - 需求 6.2: 非流式请求返回完整 JSON 响应 use axum::{ @@ -55,25 +46,31 @@ use lime_core::models::anthropic::AnthropicMessagesRequest; use lime_core::models::openai::ChatCompletionRequest; use lime_core::models::provider_pool_model::{CredentialData, ProviderCredential}; use lime_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; -use lime_providers::converter::openai_to_antigravity::{ - convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context, -}; use lime_providers::providers::{ - AntigravityProvider, ClaudeCustomProvider, CodexProvider, KiroProvider, OpenAICustomProvider, - PromptCacheMode, VertexProvider, + ClaudeCustomProvider, OpenAICustomProvider, PromptCacheMode, VertexProvider, }; use lime_providers::session::store_thought_signature; -use lime_providers::stream::{PipelineConfig, StreamPipeline}; use lime_providers::streaming::traits::StreamingProvider; use lime_providers::streaming::{ StreamConfig, StreamContext, StreamError, StreamFormat as StreamingFormat, StreamManager, StreamResponse, }; use lime_server_utils::{ - build_anthropic_response, build_anthropic_stream_response, build_error_response, - build_error_response_with_status, parse_cw_response, safe_truncate, CWParsedResponse, + build_anthropic_response, build_anthropic_stream_response, CWParsedResponse, }; +fn retired_credential_response(kind: &str) -> Response { + ( + StatusCode::GONE, + Json(serde_json::json!({ + "error": { + "message": format!("{kind} 已退役。请改用 API Key Provider / configured providers。") + } + })), + ) + .into_response() +} + /// 根据凭证调用 Provider (Anthropic 格式) /// /// # 参数 @@ -85,358 +82,13 @@ pub async fn call_provider_anthropic( state: &AppState, credential: &ProviderCredential, request: &AnthropicMessagesRequest, - flow_id: Option<&str>, + _flow_id: Option<&str>, ) -> Response { match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - // 如果是流式请求,使用真正的流式处理(需求 1.1, 6.1) - if request.stream { - return handle_kiro_stream(state, credential, request, flow_id).await; - } - - // 非流式请求,使用现有的 call_api() 方法(需求 6.1, 6.2, 6.3) - // 使用 TokenCacheService 获取有效 token - let db = match &state.db { - Some(db) => db, - None => { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": "Database not available"}})), - ) - .into_response(); - } - }; - // 获取缓存的 token - let token = match state - .token_cache - .get_valid_token(db, &credential.uuid) - .await - { - Ok(t) => t, - Err(e) => { - tracing::warn!("[POOL] Token cache miss, loading from source: {}", e); - // 回退到从源文件加载 - let mut kiro = KiroProvider::new(); - if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await { - // 记录凭证加载失败 - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to load Kiro credentials: {}", e)}})), - ) - .into_response(); - } - if let Err(e) = kiro.refresh_token().await { - // 记录 Token 刷新失败 - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Token refresh failed: {e}")), - ); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), - ) - .into_response(); - } - kiro.credentials.access_token.unwrap_or_default() - } - }; - // 使用获取到的 token 创建 KiroProvider - let mut kiro = KiroProvider::new(); - // 从源文件加载其他配置(region, profile_arn 等) - // 注意:必须先加载凭证文件,再设置 token,因为 load_credentials_from_path 会覆盖整个 credentials - let _ = kiro.load_credentials_from_path(creds_file_path).await; - // 使用缓存的 token 覆盖文件中的 token(缓存的 token 更新) - kiro.credentials.access_token = Some(token); - let openai_request = convert_anthropic_to_openai(request); - let resp = match kiro.call_api(&openai_request).await { - Ok(r) => r, - Err(e) => { - // 记录 API 调用失败 - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&e.to_string()), - ); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(); - } - }; - let status = resp.status(); - if status.is_success() { - match resp.bytes().await { - Ok(bytes) => { - let body = String::from_utf8_lossy(&bytes).to_string(); - let parsed = parse_cw_response(&body); - // 记录成功 - let _ = state.pool_service.mark_healthy( - db, - &credential.uuid, - Some(&request.model), - ); - let _ = state.pool_service.record_usage(db, &credential.uuid); - // 非流式请求返回完整 JSON 响应(需求 6.2) - build_anthropic_response(&request.model, &parsed) - } - Err(e) => { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&e.to_string()), - ); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } else if status.as_u16() == 401 || status.as_u16() == 403 { - // Token 过期,强制刷新并重试 - tracing::info!( - "[POOL] Got {}, forcing token refresh for {}", - status, - &credential.uuid[..8] - ); - let new_token = match state - .token_cache - .refresh_and_cache(db, &credential.uuid, true) - .await - { - Ok(t) => t, - Err(e) => { - // 记录 Token 刷新失败 - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Token refresh failed: {e}")), - ); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), - ) - .into_response(); - } - }; - // 使用新 token 重试 - kiro.credentials.access_token = Some(new_token); - match kiro.call_api(&openai_request).await { - Ok(retry_resp) => { - if retry_resp.status().is_success() { - match retry_resp.bytes().await { - Ok(bytes) => { - let body = String::from_utf8_lossy(&bytes).to_string(); - let parsed = parse_cw_response(&body); - // 记录重试成功 - let _ = state.pool_service.mark_healthy( - db, - &credential.uuid, - Some(&request.model), - ); - let _ = state.pool_service.record_usage(db, &credential.uuid); - // 非流式请求返回完整 JSON 响应(需求 6.2) - build_anthropic_response(&request.model, &parsed) - } - Err(e) => { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&e.to_string()), - ); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } else { - let body = retry_resp.text().await.unwrap_or_default(); - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Retry failed: {body}")), - ); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), - ) - .into_response() - } - } - Err(e) => { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&e.to_string()), - ); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } else { - let status_code = status.as_u16(); - let body = resp.text().await.unwrap_or_default(); - eprintln!("[PROVIDER_CALL] Kiro 请求失败: status={} body={}", status_code, &body[..body.len().min(500)]); - // 只有 5xx 错误才标记为不健康 - if status_code >= 500 { - let _ = state - .pool_service - .mark_unhealthy(db, &credential.uuid, Some(&body)); - } - // 转发上游的实际状态码 - ( - StatusCode::from_u16(status_code).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json(serde_json::json!({"error": {"message": body}})), - ) - .into_response() - } - } - CredentialData::GeminiOAuth { .. } => { - // Gemini OAuth 路由暂不支持 - ( - StatusCode::NOT_IMPLEMENTED, - Json(serde_json::json!({"error": {"message": "Gemini OAuth routing not yet implemented. Use /v1/messages with Gemini models instead."}})), - ) - .into_response() - } - CredentialData::AntigravityOAuth { - creds_file_path, - project_id, - } => { - let mut antigravity = AntigravityProvider::new(); - if let Err(e) = antigravity - .load_credentials_from_path(creds_file_path) - .await - { - // 记录凭证加载失败 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - } - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to load Antigravity credentials: {}", e)}})), - ) - .into_response(); - } - - // 使用新的 validate_token() 方法检查 Token 状态 - let validation_result = antigravity.validate_token(); - tracing::info!("[Antigravity] Token 验证结果: {:?}", validation_result); - - // 根据验证结果决定是否刷新 - if validation_result.needs_refresh() { - tracing::info!("[Antigravity] Token 需要刷新,开始刷新..."); - match antigravity.refresh_token_with_retry(3).await { - Ok(new_token) => { - tracing::info!("[Antigravity] Token 刷新成功,新 token 长度: {}", new_token.len()); - // 刷新成功,标记为健康 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( - db, - &credential.uuid, - None, - ); - } - } - Err(refresh_error) => { - tracing::error!("[Antigravity] Token 刷新失败: {:?}", refresh_error); - // 使用新的 mark_unhealthy_with_details 方法 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy_with_details( - db, - &credential.uuid, - &refresh_error, - ); - } - - // 根据错误类型返回不同的状态码和消息 - let (status, message) = if refresh_error.requires_reauth() { - (StatusCode::UNAUTHORIZED, refresh_error.user_message()) - } else { - (StatusCode::INTERNAL_SERVER_ERROR, refresh_error.user_message()) - }; - - return ( - status, - Json(serde_json::json!({"error": {"message": message}})), - ) - .into_response(); - } - } - } - - // 设置项目 ID - if let Some(pid) = project_id { - antigravity.project_id = Some(pid.clone()); - } else if let Err(e) = antigravity.discover_project().await { - tracing::warn!("[Antigravity] Failed to discover project: {}", e); - } - // 获取 project_id 用于请求 - let proj_id = antigravity.project_id.clone().unwrap_or_default(); - // 先转换为 OpenAI 格式,再转换为 Antigravity 格式 - let openai_request = convert_anthropic_to_openai(request); - let antigravity_request = convert_openai_to_antigravity_with_context(&openai_request, &proj_id); - match antigravity - .generate_content(&request.model, &antigravity_request) - .await - { - Ok(resp) => { - // 转换为 OpenAI 格式,再构建 Anthropic 响应 - let content = resp["candidates"][0]["content"]["parts"][0]["text"] - .as_str() - .unwrap_or(""); - let parsed = CWParsedResponse { - content: content.to_string(), - tool_calls: Vec::new(), - usage_credits: 0.0, - context_usage_percentage: 0.0, - }; - // 记录成功 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( - db, - &credential.uuid, - Some(&request.model), - ); - let _ = state.pool_service.record_usage(db, &credential.uuid); - } - if request.stream { - build_anthropic_stream_response(&request.model, &parsed) - } else { - build_anthropic_response(&request.model, &parsed) - } - } - Err(api_err) => { - // 记录 API 调用失败 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&api_err.message), - ); - } - - // 直接使用 AntigravityApiError 的状态码构建响应 - build_error_response_with_status(api_err.status_code, &api_err.to_string()) - } - } + CredentialData::KiroOAuth { .. } => { + retired_credential_response("Kiro OAuth/local CLI credential") } + CredentialData::GeminiOAuth { .. } => retired_credential_response("Gemini OAuth credential"), CredentialData::OpenAIKey { api_key, base_url } => { let openai = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone()); let openai_request = convert_anthropic_to_openai(request); @@ -463,13 +115,13 @@ pub async fn call_provider_anthropic( }; // 记录成功 if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); let _ = - state.pool_service.record_usage(db, &credential.uuid); + state.record_credential_usage(db, &credential.uuid); } if request.stream { build_anthropic_stream_response(&request.model, &parsed) @@ -480,7 +132,7 @@ pub async fn call_provider_anthropic( // 记录解析失败和原始响应 eprintln!("[PROVIDER_CALL] 解析 OpenAI 响应失败,原始响应: {}", &body); if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some("Failed to parse OpenAI response"), @@ -495,7 +147,7 @@ pub async fn call_provider_anthropic( } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&e.to_string()), @@ -515,7 +167,7 @@ pub async fn call_provider_anthropic( // 只有 5xx 错误才标记为不健康,4xx 错误(如模型不支持)不应该标记凭证为不健康 if status_code >= 500 { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&body), @@ -532,7 +184,7 @@ pub async fn call_provider_anthropic( } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&e.to_string()), @@ -605,12 +257,12 @@ pub async fn call_provider_anthropic( ); // 记录成功 if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.record_credential_usage(db, &credential.uuid); } // 透传流式响应,保持 SSE 格式 let stream = resp.bytes_stream(); @@ -645,12 +297,12 @@ pub async fn call_provider_anthropic( ); // 记录成功 if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.record_credential_usage(db, &credential.uuid); } Response::builder() .status(StatusCode::OK) @@ -673,7 +325,7 @@ pub async fn call_provider_anthropic( ), ); if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&body), @@ -693,7 +345,7 @@ pub async fn call_provider_anthropic( &format!("[CLAUDE] 读取响应失败: {e}"), ); if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&e.to_string()), @@ -709,7 +361,7 @@ pub async fn call_provider_anthropic( } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&e.to_string()), @@ -734,8 +386,8 @@ pub async fn call_provider_anthropic( Ok(body) => { if status.is_success() { if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy(db, &credential.uuid, Some(&request.model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.mark_credential_healthy(db, &credential.uuid, Some(&request.model)); + let _ = state.record_credential_usage(db, &credential.uuid); } Response::builder() .status(StatusCode::OK) @@ -746,14 +398,14 @@ pub async fn call_provider_anthropic( }) } else { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&body)); + let _ = state.mark_credential_unhealthy(db, &credential.uuid, Some(&body)); } (StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), Json(serde_json::json!({"error": {"message": body}}))).into_response() } } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); + let _ = state.mark_credential_unhealthy(db, &credential.uuid, Some(&e.to_string())); } (StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}}))).into_response() } @@ -761,7 +413,7 @@ pub async fn call_provider_anthropic( } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); + let _ = state.mark_credential_unhealthy(db, &credential.uuid, Some(&e.to_string())); } (StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}}))).into_response() } @@ -775,14 +427,9 @@ pub async fn call_provider_anthropic( ) .into_response() } - // 新增的凭证类型暂不支持 Anthropic 格式 - CredentialData::CodexOAuth { .. } - | CredentialData::ClaudeOAuth { .. } => { - ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({"error": {"message": "This credential type does not support Anthropic format yet"}})), - ) - .into_response() + CredentialData::CodexOAuth { .. } => retired_credential_response("Codex OAuth credential"), + CredentialData::ClaudeOAuth { .. } => { + retired_credential_response("Claude OAuth credential") } // Anthropic API Key - 根据 base_url 决定调用方式 CredentialData::AnthropicKey { api_key, base_url } => { @@ -823,12 +470,12 @@ pub async fn call_provider_anthropic( "[ANTHROPIC] 流式请求,透传 SSE 响应", ); if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.record_credential_usage(db, &credential.uuid); } let stream = resp.bytes_stream(); return Response::builder() @@ -853,12 +500,12 @@ pub async fn call_provider_anthropic( Ok(body) => { if status.is_success() { if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.record_credential_usage(db, &credential.uuid); } Response::builder() .status(StatusCode::OK) @@ -881,7 +528,7 @@ pub async fn call_provider_anthropic( ), ); if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&format!("API error: {status}")), @@ -904,7 +551,7 @@ pub async fn call_provider_anthropic( } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&format!("API call failed: {e}")), @@ -938,13 +585,12 @@ pub async fn call_provider_openai( // 调试:打印凭证类型 let cred_type = match &credential.credential { - CredentialData::KiroOAuth { .. } => "KiroOAuth", + CredentialData::KiroOAuth { .. } => "RetiredKiroOAuth", CredentialData::ClaudeKey { .. } => "ClaudeKey", CredentialData::OpenAIKey { .. } => "OpenAIKey", - CredentialData::GeminiOAuth { .. } => "GeminiOAuth", + CredentialData::GeminiOAuth { .. } => "RetiredGeminiOAuth", CredentialData::GeminiApiKey { .. } => "GeminiApiKey", CredentialData::VertexKey { .. } => "VertexKey", - CredentialData::AntigravityOAuth { .. } => "AntigravityOAuth", _ => "Other", }; tracing::info!( @@ -956,641 +602,10 @@ pub async fn call_provider_openai( ); match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - // 优先使用 token cache,避免每次都刷新 token - let db = match &state.db { - Some(db) => db, - None => { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": "Database not available"}})), - ) - .into_response(); - } - }; - - // 获取缓存的 token(自动处理过期和刷新) - let token = match state - .token_cache - .get_valid_token(db, &credential.uuid) - .await - { - Ok(t) => t, - Err(e) => { - tracing::warn!("[POOL] Token cache miss, loading from source: {}", e); - // 降级:从源文件加载并刷新 - let mut kiro = KiroProvider::new(); - if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to load Kiro credentials: {}", e)}})), - ) - .into_response(); - } - if let Err(e) = kiro.refresh_token().await { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Token refresh failed: {e}")), - ); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), - ) - .into_response(); - } - kiro.credentials.access_token.unwrap_or_default() - } - }; - - // 使用获取到的 token 创建 KiroProvider - let mut kiro = KiroProvider::new(); - // 从源文件加载其他配置(region, profile_arn 等) - // 注意:必须先加载凭证文件,再设置 token,因为 load_credentials_from_path 会覆盖整个 credentials - let _ = kiro.load_credentials_from_path(creds_file_path).await; - // 使用缓存的 token 覆盖文件中的 token(缓存的 token 更新) - kiro.credentials.access_token = Some(token); - - tracing::info!("[CALL_PROVIDER_OPENAI] request.stream = {}, model = {}", request.stream, request.model); - - // 检查是否为流式请求 - if request.stream { - // 流式请求处理 - tracing::info!("[OPENAI_STREAM] 处理流式请求, model={}", request.model); - match kiro.call_api_stream(request).await { - Ok(stream_response) => { - // 记录成功 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy(db, &credential.uuid, Some(&request.model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); - } - - tracing::info!("[OPENAI_STREAM] 开始转换流式响应"); - - // 使用新的统一流处理管道 (Kiro → OpenAI) - let config = PipelineConfig::kiro_to_openai(request.model.clone()); - let pipeline = std::sync::Arc::new(tokio::sync::Mutex::new( - StreamPipeline::new(config), - )); - - // 创建转换流 - let pipeline_for_stream = pipeline.clone(); - let pipeline_for_finalize = pipeline.clone(); - let final_stream = async_stream::stream! { - use futures::StreamExt; - - let mut stream_response = stream_response; - - while let Some(chunk_result) = stream_response.next().await { - match chunk_result { - Ok(bytes) => { - tracing::debug!( - "[OPENAI_STREAM] 收到 {} 字节数据", - bytes.len() - ); - - // 使用 Pipeline 处理 chunk - let sse_events = { - let mut pipeline_guard = pipeline_for_stream.lock().await; - pipeline_guard.process_chunk(&bytes) - }; - - tracing::debug!( - "[OPENAI_STREAM] 生成 {} 个 SSE 事件", - sse_events.len() - ); - - // yield 每个 SSE 事件 - for sse_str in sse_events { - yield Ok::(sse_str); - } - } - Err(e) => { - tracing::error!("[OPENAI_STREAM] 流式传输错误: {}", e); - yield Err(e); - return; - } - } - } - - tracing::info!("[OPENAI_STREAM] 流结束,生成 finalize 事件"); - - // 流结束,使用 Pipeline 生成结束事件 - let final_events = { - let mut pipeline_guard = pipeline_for_finalize.lock().await; - pipeline_guard.finish() - }; - - tracing::info!("[OPENAI_STREAM] finalize 生成 {} 个事件", final_events.len()); - - for sse_str in final_events { - yield Ok::(sse_str); - } - }; - - tracing::info!("[OPENAI_STREAM] 构建 SSE 响应"); - - // 转换为 Body 流 - let body_stream = final_stream.map(|result| -> Result { - match result { - Ok(event) => Ok(axum::body::Bytes::from(event)), - Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())), - } - }); - - // 构建 SSE 响应 - return Response::builder() - .status(StatusCode::OK) - .header(header::CONTENT_TYPE, "text/event-stream") - .header(header::CACHE_CONTROL, "no-cache") - .header(header::CONNECTION, "keep-alive") - .header(header::TRANSFER_ENCODING, "chunked") - .header("X-Accel-Buffering", "no") - .body(Body::from_stream(body_stream)) - .unwrap_or_else(|_| { - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json( - serde_json::json!({"error": {"message": "Failed to build streaming response"}}), - ), - ) - .into_response() - }); - } - Err(e) => { - // 记录请求错误 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); - } - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(); - } - } - } - - // 非流式请求处理 - match kiro.call_api(request).await { - Ok(resp) => { - let status = resp.status(); - if status.is_success() { - // 记录成功 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy(db, &credential.uuid, Some(&request.model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); - } - match resp.text().await { - Ok(body) => { - let parsed = parse_cw_response(&body); - let has_tool_calls = !parsed.tool_calls.is_empty(); - let message = if has_tool_calls { - serde_json::json!({ - "role": "assistant", - "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, - "tool_calls": parsed.tool_calls.iter().map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": "function", - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - }).collect::>() - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": parsed.content - }) - }; - Json(serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } - }], - "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 - } - })) - .into_response() - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } - } else { - // 记录 API 调用失败 - let body = resp.text().await.unwrap_or_default(); - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&format!("HTTP {}: {}", status, safe_truncate(&body, 100)))); - } - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": body}})), - ) - .into_response() - } - } - Err(e) => { - // 记录请求错误 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); - } - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() - } - } - } - CredentialData::GeminiOAuth { .. } => { - ( - StatusCode::NOT_IMPLEMENTED, - Json(serde_json::json!({"error": {"message": "Gemini OAuth routing not yet implemented."}})), - ) - .into_response() - } - CredentialData::AntigravityOAuth { creds_file_path, project_id } => { - eprintln!("\n========== [ANTIGRAVITY] 开始处理 Antigravity 请求 =========="); - eprintln!("[ANTIGRAVITY] 凭证文件: {creds_file_path}"); - eprintln!("[ANTIGRAVITY] 项目ID: {project_id:?}"); - eprintln!("[ANTIGRAVITY] 模型: {}", request.model); - eprintln!("[ANTIGRAVITY] 流式: {}", request.stream); - - let mut antigravity = AntigravityProvider::new(); - if let Err(e) = antigravity.load_credentials_from_path(creds_file_path).await { - eprintln!("[ANTIGRAVITY] 加载凭证失败: {e}"); - // 记录凭证加载失败 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - } - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to load Antigravity credentials: {}", e)}})), - ) - .into_response(); - } - eprintln!("[ANTIGRAVITY] 凭证加载成功"); - - // 使用新的 validate_token() 方法检查 Token 状态 - let validation_result = antigravity.validate_token(); - eprintln!("[ANTIGRAVITY] Token 验证结果: {validation_result:?}"); - eprintln!("[ANTIGRAVITY] needs_refresh() = {}", validation_result.needs_refresh()); - tracing::info!("[Antigravity] Token 验证结果: {:?}", validation_result); - - // 根据验证结果决定是否刷新 - if validation_result.needs_refresh() { - eprintln!("[ANTIGRAVITY] Token 需要刷新,开始刷新..."); - tracing::info!("[Antigravity] Token 需要刷新,开始刷新..."); - match antigravity.refresh_token_with_retry(3).await { - Ok(new_token) => { - eprintln!("[ANTIGRAVITY] Token 刷新成功,新 token 长度: {}", new_token.len()); - tracing::info!("[Antigravity] Token 刷新成功,新 token 长度: {}", new_token.len()); - // 刷新成功,标记为健康 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( - db, - &credential.uuid, - None, - ); - } - } - Err(refresh_error) => { - eprintln!("[ANTIGRAVITY] Token 刷新失败: {refresh_error:?}"); - tracing::error!("[Antigravity] Token 刷新失败: {:?}", refresh_error); - // 使用新的 mark_unhealthy_with_details 方法 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy_with_details( - db, - &credential.uuid, - &refresh_error, - ); - } - - // 根据错误类型返回不同的状态码和消息 - let (status, message) = if refresh_error.requires_reauth() { - (StatusCode::UNAUTHORIZED, refresh_error.user_message()) - } else { - (StatusCode::INTERNAL_SERVER_ERROR, refresh_error.user_message()) - }; - - return ( - status, - Json(serde_json::json!({"error": {"message": message}})), - ) - .into_response(); - } - } - } else { - eprintln!("[ANTIGRAVITY] Token 不需要刷新,继续使用现有 Token"); - } - - // 设置项目 ID - if let Some(pid) = project_id { - antigravity.project_id = Some(pid.clone()); - } else if let Err(e) = antigravity.discover_project().await { - tracing::warn!("[Antigravity] Failed to discover project: {}", e); - } - - tracing::info!("[ANTIGRAVITY] request.stream = {}, model = {}, project_id = {:?}", - request.stream, request.model, antigravity.project_id); - - // 检查是否为流式请求 - if request.stream { - tracing::info!("[ANTIGRAVITY_STREAM] ========== 开始处理流式请求 =========="); - tracing::info!("[ANTIGRAVITY_STREAM] model={}, has_token={}", - request.model, antigravity.credentials.access_token.is_some()); - - // 检查是否是图片生成模型 - // 注意:gemini-3-pro-image-preview 是支持图片理解的模型,不是图片生成模型 - // 只有明确的图片生成模型才需要走非流式路径 - let is_image_generation_model = request.model == "imagen" - || request.model.starts_with("imagen-") - || request.model.contains("image-generation"); - tracing::info!("[ANTIGRAVITY_STREAM] is_image_generation_model={}", is_image_generation_model); - - // 对于图片生成模型,使用非流式请求然后模拟流式返回 - if is_image_generation_model { - tracing::info!("[ANTIGRAVITY_STREAM] 图片生成模型,使用非流式请求"); - - // 获取 project_id 用于请求 - let proj_id = antigravity.project_id.clone().unwrap_or_default(); - // 转换请求格式 - 这已经是完整的 Antigravity 请求格式 - let antigravity_request = convert_openai_to_antigravity_with_context(request, &proj_id); - - // 直接调用 call_api,因为 antigravity_request 已经是完整格式 - match antigravity.call_api("generateContent", &antigravity_request).await { - Ok(resp) => { - let resp_str = serde_json::to_string_pretty(&resp).unwrap_or_default(); - if is_lime_debug_enabled() { - let debug_dir = lime_core::app_paths::resolve_logs_dir() - .unwrap_or_else(|_| std::env::temp_dir().join("lime").join("logs")); - let _ = std::fs::create_dir_all(&debug_dir); - let debug_file = debug_dir.join("antigravity_image_response.json"); - let _ = std::fs::write(&debug_file, &resp_str); - tracing::info!( - "[ANTIGRAVITY_STREAM] 原始响应已保存到: {:?}, 大小: {} bytes", - debug_file, - resp_str.len() - ); - eprintln!( - "[ANTIGRAVITY_STREAM] 原始响应已保存到: {:?}, 大小: {} bytes", - debug_file, - resp_str.len() - ); - } - - tracing::info!("[ANTIGRAVITY_STREAM] 图片生成完成,转换为流式响应"); - - // 将非流式响应转换为 OpenAI 格式 - let openai_response = convert_antigravity_to_openai_response(&resp, &request.model); - - let openai_str = serde_json::to_string_pretty(&openai_response).unwrap_or_default(); - if is_lime_debug_enabled() { - let debug_dir = lime_core::app_paths::resolve_logs_dir() - .unwrap_or_else(|_| std::env::temp_dir().join("lime").join("logs")); - let _ = std::fs::create_dir_all(&debug_dir); - let openai_debug_file = - debug_dir.join("antigravity_image_openai_response.json"); - let _ = std::fs::write(&openai_debug_file, &openai_str); - tracing::info!( - "[ANTIGRAVITY_STREAM] OpenAI 响应已保存到: {:?}, 大小: {} bytes", - openai_debug_file, - openai_str.len() - ); - eprintln!( - "[ANTIGRAVITY_STREAM] OpenAI 响应已保存到: {:?}, 大小: {} bytes", - openai_debug_file, - openai_str.len() - ); - } - - // 将非流式响应转换为流式 SSE 格式 - let model = request.model.clone(); - let chunk_id = format!("chatcmpl-{}", uuid::Uuid::new_v4()); - let created = chrono::Utc::now().timestamp(); - - // 提取内容 - let content = openai_response - .get("choices") - .and_then(|c| c.as_array()) - .and_then(|arr| arr.first()) - .and_then(|choice| choice.get("message")) - .and_then(|msg| msg.get("content")) - .and_then(|c| c.as_str()) - .unwrap_or(""); - - tracing::info!("[ANTIGRAVITY_STREAM] 图片内容长度: {} 字符", content.len()); - eprintln!("[ANTIGRAVITY_STREAM] 图片内容长度: {} 字符", content.len()); - - // 构建 SSE 事件 - let mut sse_events = String::new(); - - // 发送内容 chunk - if !content.is_empty() { - let chunk_response = serde_json::json!({ - "id": chunk_id, - "object": "chat.completion.chunk", - "created": created, - "model": model, - "choices": [{ - "index": 0, - "delta": { - "content": content - }, - "finish_reason": null - }] - }); - sse_events.push_str(&format!("data: {chunk_response}\n\n")); - } - - // 发送结束 chunk - let done_response = serde_json::json!({ - "id": chunk_id, - "object": "chat.completion.chunk", - "created": created, - "model": model, - "choices": [{ - "index": 0, - "delta": {}, - "finish_reason": "stop" - }] - }); - sse_events.push_str(&format!("data: {done_response}\n\n")); - sse_events.push_str("data: [DONE]\n\n"); - - return Response::builder() - .status(StatusCode::OK) - .header(header::CONTENT_TYPE, "text/event-stream") - .header(header::CACHE_CONTROL, "no-cache") - .header(header::CONNECTION, "keep-alive") - .body(Body::from(sse_events)) - .unwrap_or_else(|_| { - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": "Failed to build streaming response"}})), - ) - .into_response() - }); - } - Err(api_err) => { - tracing::error!("[ANTIGRAVITY_STREAM] 图片生成失败 (HTTP {}): {}", api_err.status_code, api_err.message); - // 直接使用 AntigravityApiError 的状态码构建响应 - return build_error_response_with_status(api_err.status_code, &api_err.to_string()); - } - } - } - - match antigravity.call_api_stream(request).await { - Ok(stream_response) => { - eprintln!("[ANTIGRAVITY_STREAM] ✓ 流式响应已建立"); - tracing::info!("[ANTIGRAVITY_STREAM] ✓ 流式响应已建立"); - - let model = request.model.clone(); - - // Antigravity 返回的是分片的 JSON,需要累积所有数据后解析 - // 使用 channel 来收集所有数据,然后一次性返回 - let (tx, rx) = tokio::sync::oneshot::channel::>(); - - // 在后台任务中收集所有数据 - let model_clone = model.clone(); - tokio::spawn(async move { - use futures::StreamExt; - let mut stream = stream_response; - let mut all_data = String::new(); - let mut chunk_count = 0u32; - - while let Some(result) = stream.next().await { - chunk_count += 1; - match result { - Ok(bytes) => { - let text = String::from_utf8_lossy(&bytes); - all_data.push_str(&text); - - if chunk_count <= 3 { - eprintln!("[ANTIGRAVITY_STREAM] 收集 chunk #{}: {} bytes", chunk_count, bytes.len()); - } else if chunk_count % 200 == 0 { - eprintln!("[ANTIGRAVITY_STREAM] 已收集 {} 个 chunk, 总大小: {} bytes", chunk_count, all_data.len()); - } - } - Err(e) => { - eprintln!("[ANTIGRAVITY_STREAM] chunk #{chunk_count} 错误: {e}"); - let _ = tx.send(Err(e.to_string())); - return; - } - } - } - - eprintln!("[ANTIGRAVITY_STREAM] 流结束,共收集 {} 个 chunk, 总大小: {} bytes", chunk_count, all_data.len()); - - // 尝试解析累积的 JSON 数据 - // Antigravity 返回格式: { "response": { "candidates": [...] } } - let result = parse_antigravity_accumulated_response(&all_data, &model_clone); - let _ = tx.send(result); - }); - - // 等待数据收集完成,然后构建 SSE 响应 - let sse_stream = async_stream::stream! { - match rx.await { - Ok(Ok(sse_content)) => { - // 返回累积的 SSE 事件 - yield Ok::<_, std::io::Error>(axum::body::Bytes::from(sse_content)); - } - Ok(Err(e)) => { - eprintln!("[ANTIGRAVITY_STREAM] 解析错误: {e}"); - let error_event = format!( - "data: {{\"error\": {{\"message\": \"{}\"}}}}\n\ndata: [DONE]\n\n", - e.replace("\"", "\\\"") - ); - yield Ok(axum::body::Bytes::from(error_event)); - } - Err(_) => { - eprintln!("[ANTIGRAVITY_STREAM] channel 接收错误"); - let error_event = "data: {\"error\": {\"message\": \"Internal error\"}}\n\ndata: [DONE]\n\n"; - yield Ok(axum::body::Bytes::from(error_event.to_string())); - } - } - }; - - return Response::builder() - .status(StatusCode::OK) - .header(header::CONTENT_TYPE, "text/event-stream") - .header(header::CACHE_CONTROL, "no-cache") - .header(header::CONNECTION, "keep-alive") - .header("X-Accel-Buffering", "no") - .body(Body::from_stream(sse_stream)) - .unwrap_or_else(|_| { - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json( - serde_json::json!({"error": {"message": "Failed to build streaming response"}}), - ), - ) - .into_response() - }); - } - Err(provider_err) => { - // call_api_stream 返回 ProviderError,使用字符串解析状态码 - return build_error_response(&provider_err.to_string()); - } - } - } - - // 非流式请求处理 - eprintln!("[ANTIGRAVITY_OPENAI] ========== 开始处理非流式请求 =========="); - eprintln!("[ANTIGRAVITY_OPENAI] 模型: {}", request.model); - - // 获取 project_id 用于请求 - let proj_id = antigravity.project_id.clone().unwrap_or_default(); - eprintln!("[ANTIGRAVITY_OPENAI] 项目ID: {proj_id}"); - - // 转换请求格式 - eprintln!("[ANTIGRAVITY_OPENAI] 开始转换请求格式..."); - let antigravity_request = convert_openai_to_antigravity_with_context(request, &proj_id); - eprintln!("[ANTIGRAVITY_OPENAI] 请求格式转换完成"); - - eprintln!("[ANTIGRAVITY_OPENAI] 调用 generate_content..."); - match antigravity.generate_content(&request.model, &antigravity_request).await { - Ok(resp) => { - eprintln!("[ANTIGRAVITY_OPENAI] generate_content 返回成功"); - let openai_response = convert_antigravity_to_openai_response(&resp, &request.model); - eprintln!("[ANTIGRAVITY_OPENAI] ========== 非流式请求处理完成 =========="); - Json(openai_response).into_response() - } - Err(api_err) => { - eprintln!("[ANTIGRAVITY_OPENAI] generate_content 失败 (HTTP {}): {}", api_err.status_code, api_err.message); - eprintln!("[ANTIGRAVITY_OPENAI] ========== 非流式请求处理失败 =========="); - - // 直接使用 AntigravityApiError 的状态码构建响应 - build_error_response_with_status(api_err.status_code, &api_err.to_string()) - } - } + CredentialData::KiroOAuth { .. } => { + retired_credential_response("Kiro OAuth/local CLI credential") } + CredentialData::GeminiOAuth { .. } => retired_credential_response("Gemini OAuth credential"), CredentialData::OpenAIKey { api_key, base_url } => { let openai = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone()); @@ -1860,12 +875,12 @@ pub async fn call_provider_openai( match openai.call_api_stream(request).await { Ok(stream_response) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.record_credential_usage(db, &credential.uuid); } let body_stream = @@ -1894,7 +909,7 @@ pub async fn call_provider_openai( } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&format!("Streaming API call failed: {e}")), @@ -1925,15 +940,15 @@ pub async fn call_provider_openai( // 非流式响应 if status.is_success() { if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.record_credential_usage(db, &credential.uuid); } } else if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&format!("API error: {status}")), @@ -1961,7 +976,7 @@ pub async fn call_provider_openai( } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&format!("API call failed: {e}")), @@ -1983,210 +998,9 @@ pub async fn call_provider_openai( .into_response() } } - // Codex OAuth 凭证处理 - CredentialData::CodexOAuth { - creds_file_path, - api_base_url, - } => { - // 加载 Codex 凭证 - let mut codex = CodexProvider::new(); - if let Err(e) = codex.load_credentials_from_path(creds_file_path).await { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to load Codex credentials: {}", e)}})), - ) - .into_response(); - } - - // 如果配置了自定义 API Base URL,覆盖凭证文件中的配置 - if let Some(base_url) = api_base_url { - if !base_url.trim().is_empty() { - codex.credentials.api_base_url = Some(base_url.clone()); - } - } - - // 确保 token 有效 - if let Err(e) = codex.ensure_valid_token().await { - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Codex token refresh failed: {}", e)}})), - ) - .into_response(); - } - - // 将 ChatCompletionRequest 转换为 serde_json::Value - let request_json = match serde_json::to_value(request) { - Ok(v) => v, - Err(e) => { - return ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({"error": {"message": format!("Failed to serialize request: {}", e)}})), - ) - .into_response(); - } - }; - - // 调用 Codex API - match codex.call_api(&request_json).await { - Ok(response) => { - let status = response.status(); - let headers = response.headers().clone(); - - // 检查是否为流式响应 - if request.stream { - // 流式响应:读取 Codex SSE 流,转换为 OpenAI SSE 格式 - // 参考 CLIProxyAPI: internal/translator/codex/openai/chat-completions/codex_openai_response.go - use std::sync::Arc; - use tokio::sync::Mutex; - - let bytes_stream = response.bytes_stream(); - - // 创建转换状态(包含缓冲区) - struct StreamState { - convert_state: CodexConvertState, - buffer: String, - } - - let state = Arc::new(Mutex::new(StreamState { - convert_state: CodexConvertState::default(), - buffer: String::new(), - })); - - let converted_stream = bytes_stream.map(move |result| { - let state = Arc::clone(&state); - async move { - match result { - Ok(bytes) => { - let chunk = String::from_utf8_lossy(&bytes); - let mut state = state.lock().await; - state.buffer.push_str(&chunk); - - let mut output = String::new(); - - // 处理缓冲区中的完整行 - while let Some(newline_pos) = state.buffer.find('\n') { - let line = state.buffer[..newline_pos].to_string(); - state.buffer = state.buffer[newline_pos + 1..].to_string(); - - if let Some(data) = line.strip_prefix("data: ") { - if let Ok(json) = serde_json::from_str::(data) { - if let Some(converted) = convert_codex_event_to_openai_sse_with_state( - &json, - &mut state.convert_state, - ) { - output.push_str(&format!("data: {converted}\n\n")); - } - } - } - } - - Ok::<_, std::io::Error>(bytes::Bytes::from(output)) - } - Err(e) => { - tracing::error!("[Codex] Stream error: {}", e); - Err(std::io::Error::other(e.to_string())) - } - } - } - }).buffer_unordered(1).filter_map(|result| async move { - match result { - Ok(bytes) if !bytes.is_empty() => Some(Ok(bytes)), - Ok(_) => None, - Err(e) => Some(Err(e)), - } - }); - - let body = Body::from_stream(converted_stream); - let mut response_builder = Response::builder() - .status(status) - .header(header::CONTENT_TYPE, "text/event-stream") - .header(header::CACHE_CONTROL, "no-cache") - .header(header::CONNECTION, "keep-alive"); - - for (key, value) in headers.iter() { - if key != header::CONTENT_TYPE - && key != header::TRANSFER_ENCODING - && key != header::CONTENT_LENGTH - { - response_builder = response_builder.header(key, value); - } - } - - response_builder.body(body).unwrap_or_else(|_| { - (StatusCode::INTERNAL_SERVER_ERROR, "Failed to build response") - .into_response() - }) - } else { - // 非流式响应:读取 SSE 流,解析 response.completed 事件,转换为 OpenAI 格式 - // 参考 CLIProxyAPI: internal/translator/codex/openai/chat-completions/codex_openai_response.go - match response.bytes().await { - Ok(body) => { - // 解析 SSE 数据,查找 response.completed 事件 - let body_str = String::from_utf8_lossy(&body); - let mut completed_data: Option = None; - - for line in body_str.lines() { - if let Some(data) = line.strip_prefix("data: ") { - if let Ok(json) = serde_json::from_str::(data) { - if json.get("type").and_then(|t| t.as_str()) == Some("response.completed") { - completed_data = Some(json); - break; - } - } - } - } - - match completed_data { - Some(codex_response) => { - // 转换为 OpenAI Chat Completions 格式 - let openai_response = convert_codex_to_openai_non_stream(&codex_response); - Response::builder() - .status(StatusCode::OK) - .header(header::CONTENT_TYPE, "application/json") - .body(Body::from(openai_response.to_string())) - .unwrap_or_else(|_| { - (StatusCode::INTERNAL_SERVER_ERROR, "Failed to build response") - .into_response() - }) - } - None => { - tracing::error!("[Codex] No response.completed event found in SSE stream"); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": "No response.completed event found in Codex response"}})), - ) - .into_response() - } - } - } - Err(e) => { - tracing::error!("[Codex] Failed to read response body: {}", e); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to read Codex response: {}", e)}})), - ) - .into_response() - } - } - } - } - Err(e) => { - tracing::error!("[Codex] API call failed: {}", e); - ( - StatusCode::BAD_GATEWAY, - Json(serde_json::json!({"error": {"message": format!("Codex API call failed: {}", e)}})), - ) - .into_response() - } - } - } - // 新增的凭证类型暂不支持 OpenAI 格式 + CredentialData::CodexOAuth { .. } => retired_credential_response("Codex OAuth credential"), CredentialData::ClaudeOAuth { .. } => { - ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({"error": {"message": "This credential type does not support OpenAI format yet"}})), - ) - .into_response() + retired_credential_response("Claude OAuth credential") } } } @@ -2206,12 +1020,8 @@ pub async fn call_provider_openai( /// 流式格式枚举 pub fn get_stream_format_for_credential(credential: &ProviderCredential) -> StreamingFormat { match &credential.credential { - CredentialData::KiroOAuth { .. } => StreamingFormat::AwsEventStream, CredentialData::ClaudeKey { .. } => StreamingFormat::AnthropicSse, CredentialData::OpenAIKey { .. } => StreamingFormat::OpenAiSse, - // TODO: 任务 6 完成后,将这些改为 GeminiStream - CredentialData::AntigravityOAuth { .. } => StreamingFormat::OpenAiSse, - CredentialData::GeminiOAuth { .. } => StreamingFormat::OpenAiSse, CredentialData::GeminiApiKey { .. } => StreamingFormat::OpenAiSse, CredentialData::VertexKey { .. } => StreamingFormat::OpenAiSse, _ => StreamingFormat::OpenAiSse, @@ -2580,307 +1390,6 @@ pub async fn monitor_client_disconnect(cancel_token: tokio_util::sync::Cancellat cancel_token.cancelled().await; } -// ============================================================================ -// Kiro 凭证真正流式响应处理 -// ============================================================================ - -/// Kiro 凭证流式响应处理 -/// -/// 实现真正的端到端流式传输,将 AWS Event Stream 格式转换为 Anthropic SSE 格式。 -/// -/// # 参数 -/// - `state`: 应用状态 -/// - `credential`: Kiro 凭证信息 -/// - `request`: Anthropic 格式请求 -/// - `flow_id`: Flow ID(可选,用于流式响应处理) -/// -/// # 需求覆盖 -/// - 需求 1.1: 使用 reqwest 的流式响应模式 -/// - 需求 1.2: 实时解析每个 JSON payload 并转换为 Anthropic SSE 事件 -/// - 需求 1.3: 立即发送 content_block_delta 事件给客户端 -/// - 需求 3.1: Flow Monitor 记录 chunk_count 大于 0 -/// - 需求 3.2: 调用 process_chunk 更新流重建器 -/// - 需求 3.3: 流完成时拥有完整的重建响应内容 -/// - 需求 4.4: 在流式请求前检查 Token 是否即将过期(10分钟内)并提前刷新 -pub async fn handle_kiro_stream( - state: &AppState, - credential: &ProviderCredential, - request: &AnthropicMessagesRequest, - flow_id: Option<&str>, -) -> Response { - tracing::info!( - "[KIRO_STREAM] handle_kiro_stream 被调用, model={}, flow_id={:?}", - request.model, - flow_id - ); - - // 提取凭证文件路径 - let creds_file_path = match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(), - _ => { - tracing::error!("[KIRO_STREAM] 无效的凭证类型"); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": "Invalid credential type for Kiro stream"}})), - ) - .into_response(); - } - }; - - tracing::info!("[KIRO_STREAM] 凭证文件路径: {}", creds_file_path); - - // 获取数据库连接 - let db = match &state.db { - Some(db) => db, - None => { - tracing::error!("[KIRO_STREAM] 数据库不可用"); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": "Database not available"}})), - ) - .into_response(); - } - }; - - // 获取有效 token(需求 4.4: 检查 Token 是否即将过期,10分钟内则提前刷新) - let token = match state - .token_cache - .ensure_token_valid_for_streaming(db, &credential.uuid, 10) - .await - { - Ok(t) => t, - Err(e) => { - tracing::warn!( - "[KIRO_STREAM] Token validation failed, loading from source: {}", - e - ); - // 回退到从源文件加载 - let mut kiro = KiroProvider::new(); - if let Err(e) = kiro.load_credentials_from_path(&creds_file_path).await { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to load Kiro credentials: {}", e)}})), - ) - .into_response(); - } - if let Err(e) = kiro.refresh_token().await { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Token refresh failed: {e}")), - ); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), - ) - .into_response(); - } - kiro.credentials.access_token.unwrap_or_default() - } - }; - - // 创建 KiroProvider 并设置 token - let mut kiro = KiroProvider::new(); - // 从源文件加载其他配置(region, profile_arn 等) - // 注意:必须先加载凭证文件,再设置 token,因为 load_credentials_from_path 会覆盖整个 credentials - let _ = kiro.load_credentials_from_path(&creds_file_path).await; - // 使用缓存的 token 覆盖文件中的 token(缓存的 token 更新) - kiro.credentials.access_token = Some(token); - - tracing::info!("[KIRO_STREAM] 准备调用 call_api_stream_anthropic (直接转换)"); - - // 调用流式 API - 直接使用 Anthropic 格式(需求 4.1, 4.2, 4.3: 401/403 错误重试逻辑) - let stream_response = match kiro.call_api_stream_anthropic(request).await { - Ok(stream) => { - tracing::info!("[KIRO_STREAM] call_api_stream 成功返回流"); - stream - } - Err(e) => { - tracing::error!("[KIRO_STREAM] call_api_stream 失败: {}", e); - // 检查是否是 401/403 错误或 Token 过期,需要刷新 token 重试(需求 4.1) - let needs_token_refresh = matches!( - &e, - lime_providers::providers::ProviderError::AuthenticationError(_) - | lime_providers::providers::ProviderError::TokenExpired(_) - ); - - if needs_token_refresh { - tracing::info!( - "[KIRO_STREAM] Got auth/token error ({}), forcing token refresh for {}", - e.short_message(), - &credential.uuid[..8] - ); - // 强制刷新 token(需求 4.1) - let new_token = match state - .token_cache - .refresh_and_cache(db, &credential.uuid, true) - .await - { - Ok(t) => t, - Err(refresh_err) => { - // 需求 4.3: Token 刷新失败时返回明确的错误信息 - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Token refresh failed: {refresh_err}")), - ); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({ - "error": { - "type": "authentication_error", - "message": format!("Token refresh failed: {}", refresh_err) - } - })), - ) - .into_response(); - } - }; - - // 使用新 token 重试(需求 4.2) - kiro.credentials.access_token = Some(new_token); - match kiro.call_api_stream_anthropic(request).await { - Ok(stream) => stream, - Err(retry_err) => { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&retry_err.to_string()), - ); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "type": "api_error", - "message": format!("Retry failed after token refresh: {}", retry_err) - } - })), - ) - .into_response(); - } - } - } else { - let _ = - state - .pool_service - .mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "type": "api_error", - "message": e.to_string() - } - })), - ) - .into_response(); - } - } - }; - - // 记录成功 - let _ = state - .pool_service - .mark_healthy(db, &credential.uuid, Some(&request.model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); - - tracing::info!( - "[KIRO_STREAM] 开始处理流式响应, model={}, flow_id={:?}", - request.model, - flow_id - ); - - // 使用新的统一流处理管道 (Kiro → Anthropic) - let config = PipelineConfig::kiro_to_anthropic(request.model.clone()); - let pipeline = std::sync::Arc::new(tokio::sync::Mutex::new(StreamPipeline::new(config))); - - let pipeline_clone = pipeline.clone(); - let pipeline_for_finalize = pipeline.clone(); - - let final_stream = async_stream::stream! { - use futures::StreamExt; - - let mut stream_response = stream_response; - - while let Some(chunk_result) = stream_response.next().await { - match chunk_result { - Ok(bytes) => { - tracing::info!( - "[KIRO_STREAM] 收到 {} 字节数据", - bytes.len() - ); - - let sse_strings = { - let mut pipeline_guard = pipeline_clone.lock().await; - pipeline_guard.process_chunk(&bytes) - }; - - tracing::info!( - "[KIRO_STREAM] 生成 {} 个 SSE 事件", - sse_strings.len() - ); - - for sse_str in sse_strings { - yield Ok::(sse_str); - } - } - Err(e) => { - tracing::error!("[KIRO_STREAM] 流式传输期间发生错误: {}", e); - yield Err(e); - return; - } - } - } - - tracing::info!("[KIRO_STREAM] 流结束,生成 finalize 事件"); - - let final_events = { - let mut pipeline_guard = pipeline_for_finalize.lock().await; - pipeline_guard.finish() - }; - - tracing::info!("[KIRO_STREAM] finalize 生成 {} 个事件", final_events.len()); - - for sse_str in final_events { - yield Ok::(sse_str); - } - }; - - tracing::info!("[KIRO_STREAM] 构建 SSE 响应"); - - // 转换为 Body 流 - let body_stream = final_stream.map(|result| -> Result { - match result { - Ok(event) => Ok(axum::body::Bytes::from(event)), - Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())), - } - }); - - // 构建 SSE 响应 - Response::builder() - .status(StatusCode::OK) - .header(header::CONTENT_TYPE, "text/event-stream") - .header(header::CACHE_CONTROL, "no-cache") - .header(header::CONNECTION, "keep-alive") - .header(header::TRANSFER_ENCODING, "chunked") - .header("X-Accel-Buffering", "no") - .body(Body::from_stream(body_stream)) - .unwrap_or_else(|_| { - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json( - serde_json::json!({"error": {"message": "Failed to build streaming response"}}), - ), - ) - .into_response() - }) -} - fn is_lime_debug_enabled() -> bool { lime_core::env_compat::bool_var(&["LIME_DEBUG", "PROXYCAST_DEBUG"]).unwrap_or(false) } diff --git a/src-tauri/crates/server/src/handlers/websocket.rs b/src-tauri/crates/server/src/handlers/websocket.rs index a0455b8a6..cc812368b 100644 --- a/src-tauri/crates/server/src/handlers/websocket.rs +++ b/src-tauri/crates/server/src/handlers/websocket.rs @@ -25,13 +25,7 @@ use lime_core::models::provider_pool_model::ProviderCredential; use lime_core::websocket::WsErrorCode; use lime_processor::RequestContext; use lime_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; -use lime_providers::converter::openai_to_antigravity::{ - convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context, -}; -use lime_providers::providers::{ - AntigravityProvider, ClaudeCustomProvider, KiroProvider, OpenAICustomProvider, PromptCacheMode, -}; -use lime_server_utils::parse_cw_response; +use lime_providers::providers::{ClaudeCustomProvider, OpenAICustomProvider, PromptCacheMode}; use lime_websocket::{ WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsMessage as WsProtoMessage, }; @@ -253,32 +247,6 @@ async fn handle_ws_message( "Invalid message type from client", ))), WsProtoMessage::Error(_) => None, - WsProtoMessage::SubscribeKiroEvents => { - // TODO: 实现Kiro事件订阅 - Some(WsProtoMessage::Response(WsApiResponse { - request_id: "subscribe_kiro_events".to_string(), - payload: serde_json::json!({ - "status": "subscribed", - "message": "Successfully subscribed to kiro events" - }), - })) - } - WsProtoMessage::UnsubscribeKiroEvents => { - // TODO: 实现Kiro事件取消订阅 - Some(WsProtoMessage::Response(WsApiResponse { - request_id: "unsubscribe_kiro_events".to_string(), - payload: serde_json::json!({ - "status": "unsubscribed", - "message": "Successfully unsubscribed from kiro events" - }), - })) - } - WsProtoMessage::KiroCredentialEvent(_) => { - // Kiro事件是服务端到客户端的消息,客户端不应该发送 - Some(WsProtoMessage::Error(WsError::invalid_message( - "KiroCredentialEvent messages are server-to-client only", - ))) - } } } @@ -420,11 +388,12 @@ async fn handle_ws_chat_completions( // 获取默认 provider let default_provider = state.default_provider.read().await.clone(); - // 尝试从凭证池中选择凭证(不降级,指定什么就用什么) + // 从 API Key Provider 主路径选择凭证。 let credential = match &state.db { Some(db) => state - .pool_service - .select_credential(db, &default_provider, Some(&request.model)) + .api_key_service + .select_credential_for_provider(db, &default_provider, Some(&default_provider), None) + .await .ok() .flatten(), None => None, @@ -446,9 +415,7 @@ async fn handle_ws_chat_completions( build_ws_gateway_error( Some(request_id.to_string()), GatewayErrorCode::NoCredentials, - format!( - "No available credentials for provider '{default_provider}'. Please add credentials in the Provider Pool." - ), + format!("No available API Key Provider credentials for provider '{default_provider}'."), ) } } @@ -486,11 +453,12 @@ async fn handle_ws_anthropic_messages( // 获取默认 provider let default_provider = state.default_provider.read().await.clone(); - // 尝试从凭证池中选择凭证(带智能降级) + // 从 API Key Provider 主路径选择凭证。 let credential = match &state.db { Some(db) => state - .pool_service - .select_credential(db, &default_provider, Some(&request.model)) + .api_key_service + .select_credential_for_provider(db, &default_provider, Some(&default_provider), None) + .await .ok() .flatten(), None => None, @@ -510,9 +478,7 @@ async fn handle_ws_anthropic_messages( build_ws_gateway_error( Some(request_id.to_string()), GatewayErrorCode::NoCredentials, - format!( - "No available credentials for provider '{default_provider}'. Please add credentials in the Provider Pool." - ), + format!("No available API Key Provider credentials for provider '{default_provider}'."), ) } } @@ -526,106 +492,11 @@ pub async fn call_provider_openai_for_ws( use lime_core::models::provider_pool_model::CredentialData; match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - let mut kiro = KiroProvider::new(); - if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await { - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - } - return Err(e.to_string()); - } - if let Err(e) = kiro.refresh_token().await { - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Token refresh failed: {e}")), - ); - } - return Err(e.to_string()); - } - - let resp = match kiro.call_api(request).await { - Ok(r) => r, - Err(e) => { - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&e.to_string()), - ); - } - return Err(e.to_string()); - } - }; - if resp.status().is_success() { - let body = resp.text().await.map_err(|e| e.to_string())?; - let parsed = parse_cw_response(&body); - let has_tool_calls = !parsed.tool_calls.is_empty(); - - // 记录成功 - if let Some(db) = &state.db { - let _ = - state - .pool_service - .mark_healthy(db, &credential.uuid, Some(&request.model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); - } - - let message = if has_tool_calls { - serde_json::json!({ - "role": "assistant", - "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, - "tool_calls": parsed.tool_calls.iter().map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": "function", - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - }).collect::>() - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": parsed.content - }) - }; - - Ok(serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } - }], - "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 - } - })) - } else { - let body = resp.text().await.unwrap_or_default(); - if let Some(db) = &state.db { - let _ = state - .pool_service - .mark_unhealthy(db, &credential.uuid, Some(&body)); - } - Err(format!("Upstream error: {body}")) - } + CredentialData::KiroOAuth { .. } + | CredentialData::GeminiOAuth { .. } + | CredentialData::CodexOAuth { .. } + | CredentialData::ClaudeOAuth { .. } => { + Err("凭证池/OAuth/local CLI credential 已退役,请改用 API Key Provider。".to_string()) } CredentialData::OpenAIKey { api_key, base_url } => { let provider = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone()); @@ -633,7 +504,7 @@ pub async fn call_provider_openai_for_ws( Ok(r) => r, Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&e.to_string()), @@ -646,10 +517,8 @@ pub async fn call_provider_openai_for_ws( // 记录成功 if let Some(db) = &state.db { let _ = - state - .pool_service - .mark_healthy(db, &credential.uuid, Some(&request.model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); + state.mark_credential_healthy(db, &credential.uuid, Some(&request.model)); + let _ = state.record_credential_usage(db, &credential.uuid); } resp.json::() .await @@ -657,9 +526,7 @@ pub async fn call_provider_openai_for_ws( } else { let body = resp.text().await.unwrap_or_default(); if let Some(db) = &state.db { - let _ = state - .pool_service - .mark_unhealthy(db, &credential.uuid, Some(&body)); + let _ = state.mark_credential_unhealthy(db, &credential.uuid, Some(&body)); } Err(format!("Upstream error: {body}")) } @@ -690,18 +557,18 @@ pub async fn call_provider_openai_for_ws( Ok(result) => { // 记录成功 if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( + let _ = state.mark_credential_healthy( db, &credential.uuid, Some(&request.model), ); - let _ = state.pool_service.record_usage(db, &credential.uuid); + let _ = state.record_credential_usage(db, &credential.uuid); } Ok(result) } Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&e.to_string()), @@ -711,99 +578,8 @@ pub async fn call_provider_openai_for_ws( } } } - CredentialData::AntigravityOAuth { - creds_file_path, - project_id, - } => { - let mut antigravity = AntigravityProvider::new(); - if let Err(e) = antigravity - .load_credentials_from_path(creds_file_path) - .await - { - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&format!("Failed to load credentials: {e}")), - ); - } - return Err(e.to_string()); - } - - // 使用新的 validate_token() 方法检查 Token 状态 - let validation_result = antigravity.validate_token(); - tracing::info!("[Antigravity WS] Token 验证结果: {:?}", validation_result); - - // 根据验证结果决定是否刷新 - if validation_result.needs_refresh() { - tracing::info!("[Antigravity WS] Token 需要刷新,开始刷新..."); - match antigravity.refresh_token_with_retry(3).await { - Ok(new_token) => { - tracing::info!( - "[Antigravity WS] Token 刷新成功,新 token 长度: {}", - new_token.len() - ); - // 刷新成功,标记为健康 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy(db, &credential.uuid, None); - } - } - Err(refresh_error) => { - tracing::error!("[Antigravity WS] Token 刷新失败: {:?}", refresh_error); - // 使用新的 mark_unhealthy_with_details 方法 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy_with_details( - db, - &credential.uuid, - &refresh_error, - ); - } - return Err(refresh_error.user_message()); - } - } - } - - // 设置项目 ID - if let Some(pid) = project_id { - antigravity.project_id = Some(pid.clone()); - } - let proj_id = antigravity.project_id.clone().unwrap_or_default(); - - let antigravity_request = convert_openai_to_antigravity_with_context(request, &proj_id); - match antigravity - .call_api("generateContent", &antigravity_request) - .await - { - Ok(resp) => { - // 记录成功 - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( - db, - &credential.uuid, - Some(&request.model), - ); - let _ = state.pool_service.record_usage(db, &credential.uuid); - } - Ok(convert_antigravity_to_openai_response( - &resp, - &request.model, - )) - } - Err(e) => { - if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( - db, - &credential.uuid, - Some(&e.to_string()), - ); - } - Err(e.to_string()) - } - } - } - // GeminiOAuth 和 QwenOAuth 暂不支持 WebSocket,需要使用 HTTP 端点 _ => Err( - "This credential type is not yet supported via WebSocket. Please use HTTP endpoints." + "This credential type is not supported via WebSocket. Please use API Key Provider credentials." .to_string(), ), } @@ -844,7 +620,7 @@ pub async fn call_provider_anthropic_for_ws( Ok(r) => r, Err(e) => { if let Some(db) = &state.db { - let _ = state.pool_service.mark_unhealthy( + let _ = state.mark_credential_unhealthy( db, &credential.uuid, Some(&e.to_string()), @@ -857,10 +633,8 @@ pub async fn call_provider_anthropic_for_ws( // 记录成功 if let Some(db) = &state.db { let _ = - state - .pool_service - .mark_healthy(db, &credential.uuid, Some(&request.model)); - let _ = state.pool_service.record_usage(db, &credential.uuid); + state.mark_credential_healthy(db, &credential.uuid, Some(&request.model)); + let _ = state.record_credential_usage(db, &credential.uuid); } resp.json::() .await @@ -868,9 +642,7 @@ pub async fn call_provider_anthropic_for_ws( } else { let body = resp.text().await.unwrap_or_default(); if let Some(db) = &state.db { - let _ = state - .pool_service - .mark_unhealthy(db, &credential.uuid, Some(&body)); + let _ = state.mark_credential_unhealthy(db, &credential.uuid, Some(&body)); } Err(format!("Upstream error: {body}")) } diff --git a/src-tauri/crates/server/src/lib.rs b/src-tauri/crates/server/src/lib.rs index 80b87232f..467899086 100644 --- a/src-tauri/crates/server/src/lib.rs +++ b/src-tauri/crates/server/src/lib.rs @@ -18,29 +18,14 @@ use lime_core::config::{ Config, ConfigChangeKind, ConfigManager, EndpointProvidersConfig, FileChangeEvent, FileWatcher, HotReloadManager, ReloadResult, }; -use lime_core::database::dao::provider_pool::ProviderPoolDao; use lime_core::database::DbConnection; use lime_core::logger::LogStore; use lime_core::models::anthropic::*; -use lime_core::models::openai::*; -use lime_core::models::provider_pool_model::CredentialData; -use lime_credential::CredentialSyncService; use lime_infra::injection::Injector; use lime_processor::{RequestContext, RequestProcessor}; -use lime_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; -use lime_providers::providers::antigravity::AntigravityProvider; use lime_providers::providers::claude_custom::ClaudeCustomProvider; -use lime_providers::providers::gemini::GeminiProvider; -use lime_providers::providers::kiro::KiroProvider; use lime_providers::providers::openai_custom::OpenAICustomProvider; -use lime_server_utils::{ - build_anthropic_response, build_anthropic_stream_response, build_error_response, - build_error_response_with_status, build_gemini_cli_request, build_gemini_native_request, - models, parse_cw_response, -}; -use lime_services::kiro_event_service::KiroEventService; -use lime_services::provider_pool_service::ProviderPoolService; -use lime_services::token_cache_service::TokenCacheService; +use lime_server_utils::models; use lime_websocket::{WsConfig, WsConnectionManager, WsStats}; use serde::{Deserialize, Serialize}; use std::path::PathBuf; @@ -285,8 +270,6 @@ pub struct ServerState { pub running: bool, pub requests: u64, pub start_time: Option, - pub kiro_provider: KiroProvider, - pub gemini_provider: GeminiProvider, pub openai_custom_provider: OpenAICustomProvider, pub claude_custom_provider: ClaudeCustomProvider, pub default_provider_ref: Arc>, @@ -308,8 +291,6 @@ pub struct ServerState { impl ServerState { pub fn new(config: Config) -> Self { - let kiro = KiroProvider::new(); - let gemini = GeminiProvider::new(); let openai_custom = OpenAICustomProvider::new(); let claude_custom = ClaudeCustomProvider::new(); let default_provider_ref = Arc::new(RwLock::new(config.default_provider.clone())); @@ -334,8 +315,6 @@ impl ServerState { running: false, requests: 0, start_time: None, - kiro_provider: kiro, - gemini_provider: gemini, openai_custom_provider: openai_custom, claude_custom_provider: claude_custom, default_provider_ref, @@ -390,12 +369,9 @@ impl ServerState { pub async fn start( &mut self, logs: Arc>, - pool_service: Arc, - token_cache: Arc, db: Option, ) -> Result<(), Box> { - self.start_with_telemetry(logs, pool_service, token_cache, db, None, None, None) - .await + self.start_with_telemetry(logs, db, None, None, None).await } /// 启动服务器(使用共享的遥测实例) @@ -405,8 +381,6 @@ impl ServerState { pub async fn start_with_telemetry( &mut self, logs: Arc>, - pool_service: Arc, - token_cache: Arc, db: Option, shared_stats: Option>>, shared_tokens: Option>>, @@ -414,8 +388,6 @@ impl ServerState { ) -> Result<(), Box> { self.start_with_telemetry_and_flow_monitor( logs, - pool_service, - token_cache, db, shared_stats, shared_tokens, @@ -431,8 +403,6 @@ impl ServerState { pub async fn start_with_telemetry_and_flow_monitor( &mut self, logs: Arc>, - pool_service: Arc, - token_cache: Arc, db: Option, shared_stats: Option>>, shared_tokens: Option>>, @@ -464,10 +434,6 @@ impl ServerState { let api_key = self.config.server.api_key.clone(); let default_provider_ref = self.default_provider_ref.clone(); - // 重新加载凭证 - let _ = self.kiro_provider.load_credentials().await; - let kiro = self.kiro_provider.clone(); - // 创建参数注入器 let injection_enabled = self.config.injection.enabled; let injector = Injector::with_rules( @@ -486,11 +452,10 @@ impl ServerState { // 创建请求处理器(在 spawn 之前创建,以便保存 router_ref) let processor = match (&shared_stats, &shared_tokens) { (Some(stats), Some(tokens)) => Arc::new(RequestProcessor::with_shared_telemetry( - pool_service.clone(), stats.clone(), tokens.clone(), )), - _ => Arc::new(RequestProcessor::with_defaults(pool_service.clone())), + _ => Arc::new(RequestProcessor::with_defaults()), }; // 从配置初始化 Router 的默认 Provider @@ -556,11 +521,8 @@ impl ServerState { port, &api_key, default_provider_ref, - kiro, logs, rx, - pool_service, - token_cache, db, injector, injection_enabled, @@ -608,12 +570,7 @@ pub struct AppState { pub api_key: String, pub base_url: String, pub default_provider: Arc>, - pub kiro: Arc>, pub logs: Arc>, - pub kiro_refresh_lock: Arc>, - pub gemini_refresh_lock: Arc>, - pub pool_service: Arc, - pub token_cache: Arc, pub db: Option, /// 参数注入器 pub injector: Arc>, @@ -638,8 +595,6 @@ pub struct AppState { /// Provider 维度模型配置(用于能力感知回退) pub provider_models: Arc>, - /// Kiro 事件服务 - pub kiro_event_service: Arc, /// API Key Provider 服务(用于智能降级) pub api_key_service: Arc, /// 速率限制器 @@ -657,6 +612,42 @@ pub struct AppState { pub sanitizer: Arc, } +impl AppState { + fn fallback_api_key_id<'a>(&self, uuid: &'a str) -> Option<&'a str> { + uuid.strip_prefix("fallback-") + .filter(|value| !value.is_empty()) + } + + pub fn record_credential_usage(&self, db: &DbConnection, uuid: &str) -> Result<(), String> { + if let Some(api_key_id) = self.fallback_api_key_id(uuid) { + return self.api_key_service.record_usage(db, api_key_id); + } + + tracing::debug!("[SERVER] 忽略已退役的凭证池使用记录: {}", uuid); + Ok(()) + } + + pub fn mark_credential_healthy( + &self, + _db: &DbConnection, + uuid: &str, + _model: Option<&str>, + ) -> Result<(), String> { + tracing::debug!("[SERVER] 忽略凭证健康写回: {}", uuid); + Ok(()) + } + + pub fn mark_credential_unhealthy( + &self, + _db: &DbConnection, + uuid: &str, + _error: Option<&str>, + ) -> Result<(), String> { + tracing::debug!("[SERVER] 忽略凭证失败写回: {}", uuid); + Ok(()) + } +} + /// 启动配置文件监控 /// /// 监控配置文件变化并触发热重载。 @@ -730,32 +721,7 @@ async fn start_config_watcher( let new_config = manager.config(); update_processor_config(&processor_clone, &new_config).await; - // 同步凭证池 - if let (Some(ref db), Some(ref cfg_manager)) = - (&db_clone, &config_manager_clone) - { - match sync_credential_pool_from_config(db, cfg_manager, &logs_clone) - .await - { - Ok(count) => { - tracing::info!( - "[HOT_RELOAD] 凭证池同步完成,共 {} 个凭证", - count - ); - logs_clone.write().await.add( - "info", - &format!("[HOT_RELOAD] 凭证池同步完成,共 {count} 个凭证"), - ); - } - Err(e) => { - tracing::warn!("[HOT_RELOAD] 凭证池同步失败: {}", e); - logs_clone - .write() - .await - .add("warn", &format!("[HOT_RELOAD] 凭证池同步失败: {e}")); - } - } - } + let _ = (&db_clone, &config_manager_clone); } ReloadResult::RolledBack { error, .. } => { tracing::warn!("[HOT_RELOAD] 配置热重载失败,已回滚: {}", error); @@ -866,58 +832,6 @@ async fn update_processor_config(processor: &RequestProcessor, config: &Config) tracing::info!("[HOT_RELOAD] 处理器配置更新完成"); } -/// 从配置同步凭证池 -/// -/// 当配置热重载成功后,从 YAML 配置中加载凭证并同步到数据库。 -/// -/// # 同步策略 -/// -/// - 从配置中加载所有凭证 -/// - 对于配置中存在但数据库中不存在的凭证,添加到数据库 -/// - 对于配置中存在且数据库中也存在的凭证,更新数据库中的记录 -/// - 对于数据库中存在但配置中不存在的凭证,保留(不删除,避免丢失运行时状态) -async fn sync_credential_pool_from_config( - db: &DbConnection, - config_manager: &Arc>, - _logs: &Arc>, -) -> Result { - // 创建凭证同步服务 - let sync_service = CredentialSyncService::new(config_manager.clone()); - - // 从配置加载凭证 - let credentials = sync_service.load_from_config().map_err(|e| e.to_string())?; - - let conn = lime_core::database::lock_db(db)?; - let mut synced_count = 0; - - for cred in &credentials { - // 检查凭证是否已存在 - let existing = - ProviderPoolDao::get_by_uuid(&conn, &cred.uuid).map_err(|e| e.to_string())?; - - if existing.is_some() { - // 更新现有凭证 - ProviderPoolDao::update(&conn, cred).map_err(|e| e.to_string())?; - tracing::debug!( - "[HOT_RELOAD] 更新凭证: {} ({})", - cred.uuid, - cred.provider_type - ); - } else { - // 添加新凭证 - ProviderPoolDao::insert(&conn, cred).map_err(|e| e.to_string())?; - tracing::debug!( - "[HOT_RELOAD] 添加凭证: {} ({})", - cred.uuid, - cred.provider_type - ); - } - synced_count += 1; - } - - Ok(synced_count) -} - /// 开发桥接启动回调类型 pub type DevBridgeCallback = Box; @@ -926,11 +840,8 @@ async fn run_server( port: u16, api_key: &str, default_provider: Arc>, - kiro: KiroProvider, logs: Arc>, shutdown: oneshot::Receiver<()>, - pool_service: Arc, - token_cache: Arc, db: Option, injector: Injector, injection_enabled: bool, @@ -955,11 +866,10 @@ async fn run_server( Some(p) => p, None => match (&shared_stats, &shared_tokens) { (Some(stats), Some(tokens)) => Arc::new(RequestProcessor::with_shared_telemetry( - pool_service.clone(), stats.clone(), tokens.clone(), )), - _ => Arc::new(RequestProcessor::with_defaults(pool_service.clone())), + _ => Arc::new(RequestProcessor::with_defaults()), }, }; @@ -1043,9 +953,6 @@ async fn run_server( .unwrap_or_default(), ); - // 创建 Kiro 事件服务 - let kiro_event_service = Arc::new(KiroEventService::new()); - // 创建 API Key Provider 服务 let api_key_service = Arc::new(lime_services::api_key_provider_service::ApiKeyProviderService::new()); @@ -1059,12 +966,7 @@ async fn run_server( api_key: api_key.to_string(), base_url, default_provider, - kiro: Arc::new(RwLock::new(kiro)), logs, - kiro_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), - gemini_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), - pool_service, - token_cache, db, injector: Arc::new(RwLock::new(injector)), injection_enabled: Arc::new(RwLock::new(injection_enabled)), @@ -1077,7 +979,6 @@ async fn run_server( amp_router, endpoint_providers, provider_models, - kiro_event_service, api_key_service, rate_limiter: Some(Arc::new( middleware::rate_limit::SlidingWindowRateLimiter::new( @@ -1114,25 +1015,6 @@ async fn run_server( // 设置请求体大小限制为 100MB,支持大型上下文请求(如 Claude Code 的 /compact 命令) let body_limit = 100 * 1024 * 1024; // 100MB - // Kiro凭证管理API路由 - let kiro_api_routes = Router::new() - .route( - "/api/kiro/credentials/available", - get(handlers::get_available_credentials), - ) - .route( - "/api/kiro/credentials/select", - post(handlers::select_credential), - ) - .route( - "/api/kiro/credentials/{uuid}/refresh", - axum::routing::put(handlers::refresh_credential), - ) - .route( - "/api/kiro/credentials/{uuid}/status", - get(handlers::get_credential_status), - ); - // 凭证 API 路由(用于 aster Agent 集成) let credentials_api_routes = Router::new() .route("/v1/credentials/select", post(handlers::credentials_select)) @@ -1204,8 +1086,6 @@ async fn run_server( "/lime-chrome-control/:lime_key", get(handlers::chrome_control_ws_upgrade), ) - // Kiro凭证管理API路由 - .merge(kiro_api_routes) // 凭证 API 路由(用于 aster Agent 集成) .merge(credentials_api_routes) .layer(cors_layer) @@ -1398,501 +1278,15 @@ async fn gemini_generate_content( .into_response(); } - let is_stream = method == "streamGenerateContent"; + let _ = request; - // 获取默认 provider - let default_provider = state.default_provider.read().await.clone(); - - // 尝试从凭证池中选择凭证(不降级,指定什么就用什么) - let credential = match &state.db { - Some(db) => state - .pool_service - .select_credential(db, &default_provider, Some(model)) - .ok() - .flatten(), - None => None, - }; - - let cred = match credential { - Some(c) => c, - None => { - return ( - StatusCode::SERVICE_UNAVAILABLE, - Json(serde_json::json!({ - "error": { - "message": format!("No available credentials for provider '{}'. Please add credentials in the Provider Pool.", default_provider) - } - })), - ) - .into_response(); - } - }; - - state.logs.write().await.add( - "info", - &format!( - "[GEMINI] 使用凭证: type={} name={:?} uuid={}", - cred.provider_type, - cred.name, - &cred.uuid[..8] - ), - ); - - // 调用 Antigravity Provider - match &cred.credential { - CredentialData::AntigravityOAuth { - creds_file_path, - project_id, - } => { - let mut antigravity = AntigravityProvider::new(); - if let Err(e) = antigravity - .load_credentials_from_path(creds_file_path) - .await - { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": format!("加载 Antigravity 凭证失败: {}", e) - } - })), - ) - .into_response(); + ( + StatusCode::GONE, + Json(serde_json::json!({ + "error": { + "message": "Gemini CLI OAuth 原生端点已退役。请通过 API Key Provider 使用 /v1/chat/completions 或 /v1/messages。" } - - // 使用新的 validate_token() 方法检查 Token 状态 - let validation_result = antigravity.validate_token(); - tracing::info!( - "[Antigravity Gemini] Token 验证结果: {:?}", - validation_result - ); - - // 根据验证结果决定是否刷新 - if validation_result.needs_refresh() { - tracing::info!("[Antigravity Gemini] Token 需要刷新,开始刷新..."); - match antigravity.refresh_token_with_retry(3).await { - Ok(new_token) => { - tracing::info!( - "[Antigravity Gemini] Token 刷新成功,新 token 长度: {}", - new_token.len() - ); - } - Err(refresh_error) => { - tracing::error!("[Antigravity Gemini] Token 刷新失败: {:?}", refresh_error); - - // 根据错误类型返回不同的状态码和消息 - let (status, message) = if refresh_error.requires_reauth() { - (StatusCode::UNAUTHORIZED, refresh_error.user_message()) - } else { - ( - StatusCode::INTERNAL_SERVER_ERROR, - refresh_error.user_message(), - ) - }; - - return ( - status, - Json(serde_json::json!({ - "error": { - "message": message - } - })), - ) - .into_response(); - } - } - } - - // 设置项目 ID - if let Some(pid) = project_id { - antigravity.project_id = Some(pid.clone()); - } else if antigravity.project_id.is_none() { - // 如果凭证中没有 project_id,尝试从 API 获取或生成随机 ID - if let Err(e) = antigravity.discover_project().await { - tracing::warn!("[Antigravity] 获取项目 ID 失败: {},使用随机生成的 ID", e); - // 生成随机项目 ID - let uuid = uuid::Uuid::new_v4(); - let bytes = uuid.as_bytes(); - let adjectives = ["useful", "bright", "swift", "calm", "bold"]; - let nouns = ["fuze", "wave", "spark", "flow", "core"]; - let adj = adjectives[(bytes[0] as usize) % adjectives.len()]; - let noun = nouns[(bytes[1] as usize) % nouns.len()]; - let random_part: String = uuid.to_string()[..5].to_lowercase(); - antigravity.project_id = Some(format!("{adj}-{noun}-{random_part}")); - } - } - - let proj_id = antigravity.project_id.clone().unwrap_or_else(|| { - // 最后的后备:生成随机 ID - let uuid = uuid::Uuid::new_v4(); - format!("lime-{}", &uuid.to_string()[..8]) - }); - - state - .logs - .write() - .await - .add("debug", &format!("[GEMINI] 使用 project_id: {proj_id}")); - - // 构建 Antigravity 请求体 - // 直接使用用户传入的 Gemini 格式请求,只添加必要的字段 - let antigravity_request = build_gemini_native_request(&request, model, &proj_id); - - state.logs.write().await.add( - "debug", - &format!( - "[GEMINI] 请求体: {}", - serde_json::to_string(&antigravity_request).unwrap_or_default() - ), - ); - - if is_stream { - // 流式响应 - 暂不支持,返回错误 - return ( - StatusCode::NOT_IMPLEMENTED, - Json(serde_json::json!({ - "error": { - "message": "流式响应暂不支持,请使用 generateContent" - } - })), - ) - .into_response(); - } - - // 非流式响应 - match antigravity - .call_api("generateContent", &antigravity_request) - .await - { - Ok(resp) => { - state.logs.write().await.add( - "info", - &format!( - "[GEMINI] 响应成功: {}", - serde_json::to_string(&resp) - .unwrap_or_default() - .chars() - .take(200) - .collect::() - ), - ); - - // 直接返回 Gemini 格式响应 - Json(resp).into_response() - } - Err(api_err) => { - state.logs.write().await.add( - "error", - &format!( - "[GEMINI] 请求失败 (HTTP {}): {}", - api_err.status_code, api_err.message - ), - ); - - // 直接使用 AntigravityApiError 的状态码构建响应 - build_error_response_with_status(api_err.status_code, &api_err.to_string()) - } - } - } - CredentialData::GeminiOAuth { - creds_file_path, - project_id, - } => { - // 使用 GeminiProvider 处理 Gemini CLI OAuth 凭证 - let mut gemini = GeminiProvider::new(); - if let Err(e) = gemini.load_credentials_from_path(creds_file_path).await { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": format!("加载 Gemini 凭证失败: {}", e) - } - })), - ) - .into_response(); - } - - // 检查并刷新 Token - if !gemini.is_token_valid() { - tracing::info!("[Gemini CLI] Token 需要刷新,开始刷新..."); - match gemini.refresh_token_with_retry(3).await { - Ok(new_token) => { - tracing::info!( - "[Gemini CLI] Token 刷新成功,新 token 长度: {}", - new_token.len() - ); - } - Err(refresh_error) => { - tracing::error!("[Gemini CLI] Token 刷新失败: {:?}", refresh_error); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({ - "error": { - "message": format!("Token 刷新失败: {}", refresh_error) - } - })), - ) - .into_response(); - } - } - } - - // 设置项目 ID - if let Some(pid) = project_id { - gemini.project_id = Some(pid.clone()); - } else if gemini.project_id.is_none() { - // 尝试从 API 获取项目 ID - if let Err(e) = gemini.discover_project().await { - tracing::warn!("[Gemini CLI] 获取项目 ID 失败: {},使用随机生成的 ID", e); - let uuid = uuid::Uuid::new_v4(); - let bytes = uuid.as_bytes(); - let adjectives = ["useful", "bright", "swift", "calm", "bold"]; - let nouns = ["fuze", "wave", "spark", "flow", "core"]; - let adj = adjectives[(bytes[0] as usize) % adjectives.len()]; - let noun = nouns[(bytes[1] as usize) % nouns.len()]; - let random_part: String = uuid.to_string()[..5].to_lowercase(); - gemini.project_id = Some(format!("{adj}-{noun}-{random_part}")); - } - } - - let proj_id = gemini.project_id.clone().unwrap_or_else(|| { - let uuid = uuid::Uuid::new_v4(); - format!("lime-{}", &uuid.to_string()[..8]) - }); - - state - .logs - .write() - .await - .add("debug", &format!("[GEMINI CLI] 使用 project_id: {proj_id}")); - - // 构建 Gemini CLI 请求体 - // Gemini CLI 使用 Cloud Code Assist 端点,不做模型名称映射 - let gemini_request = build_gemini_cli_request(&request, model, &proj_id); - - state.logs.write().await.add( - "debug", - &format!( - "[GEMINI CLI] 请求体: {}", - serde_json::to_string(&gemini_request).unwrap_or_default() - ), - ); - - if is_stream { - // 流式响应 - 暂不支持 - return ( - StatusCode::NOT_IMPLEMENTED, - Json(serde_json::json!({ - "error": { - "message": "Gemini CLI 流式响应暂不支持,请使用 generateContent" - } - })), - ) - .into_response(); - } - - // 非流式响应 - match gemini.call_api("generateContent", &gemini_request).await { - Ok(resp) => { - state.logs.write().await.add( - "info", - &format!( - "[GEMINI CLI] 响应成功: {}", - serde_json::to_string(&resp) - .unwrap_or_default() - .chars() - .take(200) - .collect::() - ), - ); - - // 直接返回 Gemini 格式响应 - Json(resp).into_response() - } - Err(api_err) => { - state - .logs - .write() - .await - .add("error", &format!("[GEMINI CLI] 请求失败: {api_err}")); - - build_error_response(&api_err.to_string()) - } - } - } - _ => ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({ - "error": { - "message": "Gemini 原生协议只支持 Antigravity 或 Gemini CLI OAuth 凭证" - } - })), - ) - .into_response(), - } -} - -/// 内部 Anthropic messages 处理 (使用默认 Kiro) -/// 预留:用于内部直接调用 Kiro API -#[allow(dead_code)] -async fn anthropic_messages_internal( - state: &AppState, - request: &AnthropicMessagesRequest, -) -> Response { - // 检查 token - { - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - let needs_refresh = - kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); - if needs_refresh { - if let Err(e) = kiro.refresh_token().await { - state - .logs - .write() - .await - .add("error", &format!("[AUTH] Token refresh failed: {e}")); - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ) - .into_response(); - } - } - } - - let openai_request = convert_anthropic_to_openai(request); - let kiro = state.kiro.read().await; - - match kiro.call_api(&openai_request).await { - Ok(resp) => { - let status = resp.status(); - if status.is_success() { - match resp.bytes().await { - Ok(bytes) => { - let body = String::from_utf8_lossy(&bytes).to_string(); - let parsed = parse_cw_response(&body); - if request.stream { - build_anthropic_stream_response(&request.model, &parsed) - } else { - build_anthropic_response(&request.model, &parsed) - } - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } - } else { - let body = resp.text().await.unwrap_or_default(); - ( - StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})), - ) - .into_response() - } - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } -} - -/// 内部 OpenAI chat completions 处理 (使用默认 Kiro) -/// 预留:用于内部直接调用 Kiro API -#[allow(dead_code)] -async fn chat_completions_internal(state: &AppState, request: &ChatCompletionRequest) -> Response { - { - let _guard = state.kiro_refresh_lock.lock().await; - let mut kiro = state.kiro.write().await; - let needs_refresh = - kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon(); - if needs_refresh { - if let Err(e) = kiro.refresh_token().await { - return ( - StatusCode::UNAUTHORIZED, - Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), - ) - .into_response(); - } - } - } - - let kiro = state.kiro.read().await; - match kiro.call_api(request).await { - Ok(resp) => { - let status = resp.status(); - if status.is_success() { - match resp.text().await { - Ok(body) => { - let parsed = parse_cw_response(&body); - let has_tool_calls = !parsed.tool_calls.is_empty(); - - let message = if has_tool_calls { - serde_json::json!({ - "role": "assistant", - "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, - "tool_calls": parsed.tool_calls.iter().map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": "function", - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - }).collect::>() - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": parsed.content - }) - }; - - let response = serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } - }], - "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 - } - }); - Json(response).into_response() - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } - } else { - let body = resp.text().await.unwrap_or_default(); - ( - StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})), - ) - .into_response() - } - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } + })), + ) + .into_response() } diff --git a/src-tauri/crates/services/src/api_key_provider_service.rs b/src-tauri/crates/services/src/api_key_provider_service.rs index d341e6c75..8b545e777 100644 --- a/src-tauri/crates/services/src/api_key_provider_service.rs +++ b/src-tauri/crates/services/src/api_key_provider_service.rs @@ -5,7 +5,10 @@ //! **Feature: provider-ui-refactor** //! **Validates: Requirements 7.3, 9.1, 9.2, 9.3** -use crate::provider_type_mapping::pool_provider_type_to_api_type; +use crate::provider_type_mapping::{ + api_provider_type_to_pool_type, is_custom_provider_id, pool_provider_type_to_api_type, + resolve_pool_provider_type_or_default, +}; use base64::{engine::general_purpose::STANDARD as BASE64, Engine}; use chrono::Utc; use lime_core::api_host_utils::{ @@ -2543,6 +2546,60 @@ impl ApiKeyProviderService { // ==================== 智能降级 ==================== + /// 从 API Key Provider 主路径选择凭证。 + /// + /// Provider Pool 已退役;运行时不再先读 `provider_pool_credentials`,只从 + /// `api_key_providers` / `api_keys` 选择可用凭证。 + pub async fn select_credential_for_provider( + &self, + db: &DbConnection, + provider_type: &str, + provider_id_hint: Option<&str>, + client_type: Option<&lime_core::models::client_type::ClientType>, + ) -> Result, String> { + let mut pool_type = resolve_pool_provider_type_or_default(provider_type); + let mut resolved_provider_id_hint = provider_id_hint; + + if is_custom_provider_id(provider_type) { + resolved_provider_id_hint = Some(provider_type); + } + + if let Some(custom_provider_id) = + resolved_provider_id_hint.filter(|id| is_custom_provider_id(id)) + { + match self.get_provider(db, custom_provider_id) { + Ok(Some(provider_with_keys)) => { + pool_type = + api_provider_type_to_pool_type(provider_with_keys.provider.provider_type); + tracing::debug!( + "[API_KEY_PROVIDER] custom provider '{}' 真实类型 {:?} -> {:?}", + custom_provider_id, + provider_with_keys.provider.provider_type, + pool_type + ); + } + Ok(None) => { + tracing::debug!( + "[API_KEY_PROVIDER] custom provider '{}' 不存在,继续使用解析类型 {:?}", + custom_provider_id, + pool_type + ); + } + Err(error) => { + tracing::warn!( + "[API_KEY_PROVIDER] 查询 custom provider '{}' 失败: {},继续使用解析类型 {:?}", + custom_provider_id, + error, + pool_type + ); + } + } + } + + self.get_fallback_credential(db, &pool_type, resolved_provider_id_hint, client_type) + .await + } + /// 根据 PoolProviderType 获取降级凭证 /// /// 用于智能降级场景:当 Provider Pool 无可用凭证时,自动从 API Key Provider 查找 diff --git a/src-tauri/crates/services/src/kiro_event_service.rs b/src-tauri/crates/services/src/kiro_event_service.rs deleted file mode 100644 index bfd492674..000000000 --- a/src-tauri/crates/services/src/kiro_event_service.rs +++ /dev/null @@ -1,409 +0,0 @@ -//! Kiro 凭证事件服务 -//! -//! 负责管理 Kiro 凭证相关的实时事件推送,包括: -//! - 凭证状态更新 -//! - Token 刷新事件 -//! - 健康检查结果 -//! - 凭证池统计 - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use tokio::sync::{broadcast, RwLock}; - -use lime_core::websocket::{KiroTokenInfo, WsKiroEvent}; - -/// Kiro 事件服务 -#[derive(Debug)] -pub struct KiroEventService { - /// 事件发送器 - event_sender: broadcast::Sender, - /// 凭证状态缓存 - credential_states: RwLock>, -} - -/// 缓存的凭证状态 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CachedCredentialState { - uuid: String, - is_healthy: bool, - is_disabled: bool, - error_count: u32, - health_score: Option, - last_used: Option>, - last_updated: DateTime, -} - -impl KiroEventService { - /// 创建新的 Kiro 事件服务 - pub fn new() -> Self { - let (event_sender, _) = broadcast::channel(1000); - Self { - event_sender, - credential_states: RwLock::new(HashMap::new()), - } - } - - /// 订阅 Kiro 事件 - pub fn subscribe(&self) -> broadcast::Receiver { - self.event_sender.subscribe() - } - - /// 发送凭证状态更新事件 - pub async fn emit_credential_status_update( - &self, - uuid: String, - is_healthy: bool, - is_disabled: bool, - error_count: u32, - health_score: Option, - last_used: Option>, - ) { - let now = Utc::now(); - - // 更新缓存 - { - let mut states = self.credential_states.write().await; - states.insert( - uuid.clone(), - CachedCredentialState { - uuid: uuid.clone(), - is_healthy, - is_disabled, - error_count, - health_score, - last_used, - last_updated: now, - }, - ); - } - - // 发送事件 - let event = WsKiroEvent::CredentialStatusUpdate { - uuid, - is_healthy, - is_disabled, - error_count, - health_score, - last_used, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send credential status update event: {}", e); - } - } - - /// 发送凭证刷新开始事件 - pub async fn emit_refresh_started(&self, uuid: String, credential_name: Option) { - let event = WsKiroEvent::RefreshStarted { - uuid, - credential_name, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send refresh started event: {}", e); - } - } - - /// 发送凭证刷新成功事件 - pub async fn emit_refresh_success( - &self, - uuid: String, - credential_name: Option, - expires_at: DateTime, - auth_method: String, - provider: String, - region: String, - ) { - let new_token_info = KiroTokenInfo { - expires_at, - auth_method, - provider, - region, - }; - - let event = WsKiroEvent::RefreshSuccess { - uuid, - credential_name, - new_token_info, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send refresh success event: {}", e); - } - } - - /// 发送凭证刷新失败事件 - pub async fn emit_refresh_failed( - &self, - uuid: String, - credential_name: Option, - error: String, - error_code: Option, - ) { - let event = WsKiroEvent::RefreshFailed { - uuid, - credential_name, - error, - error_code, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send refresh failed event: {}", e); - } - } - - /// 发送健康检查完成事件 - pub async fn emit_health_check_completed( - &self, - uuid: String, - credential_name: Option, - is_healthy: bool, - health_score: Option, - ) { - let last_check = Utc::now(); - let event = WsKiroEvent::HealthCheckCompleted { - uuid, - credential_name, - is_healthy, - health_score, - last_check, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send health check completed event: {}", e); - } - } - - /// 发送凭证池统计更新事件 - pub async fn emit_pool_stats_update( - &self, - total_credentials: u32, - healthy_credentials: u32, - available_credentials: u32, - average_health_score: Option, - last_rotation: Option>, - ) { - let event = WsKiroEvent::PoolStatsUpdate { - total_credentials, - healthy_credentials, - available_credentials, - average_health_score, - last_rotation, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send pool stats update event: {}", e); - } - } - - /// 发送凭证轮换事件 - pub async fn emit_credential_rotated( - &self, - from_uuid: Option, - to_uuid: String, - reason: String, - ) { - let rotation_time = Utc::now(); - let event = WsKiroEvent::CredentialRotated { - from_uuid, - to_uuid, - reason, - rotation_time, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send credential rotated event: {}", e); - } - } - - /// 发送凭证自动禁用事件 - pub async fn emit_credential_auto_disabled( - &self, - uuid: String, - credential_name: Option, - reason: String, - error_type: String, - ) { - let disable_time = Utc::now(); - let event = WsKiroEvent::CredentialAutoDisabled { - uuid, - credential_name, - reason, - error_type, - disable_time, - }; - - if let Err(e) = self.event_sender.send(event) { - tracing::debug!("Failed to send credential auto disabled event: {}", e); - } - } - - /// 获取当前活跃订阅者数量 - pub fn subscriber_count(&self) -> usize { - self.event_sender.receiver_count() - } - - /// 获取缓存的凭证状态 - pub async fn get_credential_state(&self, uuid: &str) -> Option { - let states = self.credential_states.read().await; - states.get(uuid).cloned() - } - - /// 获取所有凭证状态 - pub async fn get_all_credential_states(&self) -> Vec { - let states = self.credential_states.read().await; - states.values().cloned().collect() - } - - /// 清理过期的凭证状态缓存 - pub async fn cleanup_expired_states(&self, retention_hours: u64) { - let cutoff = Utc::now() - chrono::Duration::hours(retention_hours as i64); - let mut states = self.credential_states.write().await; - states.retain(|_, state| state.last_updated > cutoff); - } -} - -impl Default for KiroEventService { - fn default() -> Self { - Self::new() - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[tokio::test] - async fn test_credential_status_update_event() { - let service = KiroEventService::new(); - let mut receiver = service.subscribe(); - - // 发送凭证状态更新事件 - let uuid = "test-uuid".to_string(); - service - .emit_credential_status_update(uuid.clone(), true, false, 0, Some(85.5), None) - .await; - - // 验证事件被正确接收 - let event = receiver.recv().await.unwrap(); - match event { - WsKiroEvent::CredentialStatusUpdate { - uuid: event_uuid, - is_healthy, - health_score, - .. - } => { - assert_eq!(event_uuid, uuid); - assert!(is_healthy); - assert_eq!(health_score, Some(85.5)); - } - _ => panic!("Expected CredentialStatusUpdate event"), - } - - // 验证状态被缓存 - let cached_state = service.get_credential_state(&uuid).await.unwrap(); - assert_eq!(cached_state.uuid, uuid); - assert!(cached_state.is_healthy); - assert_eq!(cached_state.health_score, Some(85.5)); - } - - #[tokio::test] - async fn test_refresh_events() { - let service = KiroEventService::new(); - let mut receiver = service.subscribe(); - - let uuid = "test-uuid".to_string(); - let credential_name = Some("test-credential".to_string()); - - // 测试刷新开始事件 - service - .emit_refresh_started(uuid.clone(), credential_name.clone()) - .await; - - let event = receiver.recv().await.unwrap(); - match event { - WsKiroEvent::RefreshStarted { - uuid: event_uuid, - credential_name: event_name, - } => { - assert_eq!(event_uuid, uuid); - assert_eq!(event_name, credential_name); - } - _ => panic!("Expected RefreshStarted event"), - } - - // 测试刷新成功事件 - let expires_at = Utc::now() + chrono::Duration::hours(1); - service - .emit_refresh_success( - uuid.clone(), - credential_name.clone(), - expires_at, - "IdC".to_string(), - "BuilderId".to_string(), - "us-east-1".to_string(), - ) - .await; - - let event = receiver.recv().await.unwrap(); - match event { - WsKiroEvent::RefreshSuccess { - uuid: event_uuid, - new_token_info, - .. - } => { - assert_eq!(event_uuid, uuid); - assert_eq!(new_token_info.auth_method, "IdC"); - assert_eq!(new_token_info.provider, "BuilderId"); - } - _ => panic!("Expected RefreshSuccess event"), - } - } - - #[tokio::test] - async fn test_multiple_subscribers() { - let service = KiroEventService::new(); - let mut receiver1 = service.subscribe(); - let mut receiver2 = service.subscribe(); - - assert_eq!(service.subscriber_count(), 2); - - // 发送事件 - service - .emit_refresh_started("test-uuid".to_string(), None) - .await; - - // 两个订阅者都应该收到事件 - let event1 = receiver1.recv().await.unwrap(); - let event2 = receiver2.recv().await.unwrap(); - - matches!(event1, WsKiroEvent::RefreshStarted { .. }); - matches!(event2, WsKiroEvent::RefreshStarted { .. }); - } - - #[tokio::test] - async fn test_credential_state_caching() { - let service = KiroEventService::new(); - - // 添加多个凭证状态 - service - .emit_credential_status_update("uuid1".to_string(), true, false, 0, Some(90.0), None) - .await; - - service - .emit_credential_status_update("uuid2".to_string(), false, true, 5, Some(30.0), None) - .await; - - // 验证所有状态都被缓存 - let all_states = service.get_all_credential_states().await; - assert_eq!(all_states.len(), 2); - - // 验证可以按 UUID 获取特定状态 - let state1 = service.get_credential_state("uuid1").await.unwrap(); - assert_eq!(state1.health_score, Some(90.0)); - - let state2 = service.get_credential_state("uuid2").await.unwrap(); - assert_eq!(state2.error_count, 5); - } -} diff --git a/src-tauri/crates/services/src/lib.rs b/src-tauri/crates/services/src/lib.rs index b5aa9c70d..e95821884 100644 --- a/src-tauri/crates/services/src/lib.rs +++ b/src-tauri/crates/services/src/lib.rs @@ -7,7 +7,6 @@ //! - `file_browser_service` - 文件浏览服务 //! - `sysinfo_service` - 系统信息服务 //! - `update_check_service` - 更新检查服务 -//! - `usage_service` - 使用统计服务 #![allow(clippy::type_complexity)] #![allow(clippy::let_underscore_future)] @@ -33,17 +32,13 @@ //! - `material_service` - 素材服务 //! - `persona_service` - 人设服务 //! - `model_registry_service` - 模型注册服务 -//! - `model_service` - 模型服务 //! - `prompt_service` - Prompt 服务 //! - `mcp_service` - MCP 服务 //! - `aster_session_store` - Aster 会话存储 //! - `session_context_service` - 会话上下文服务 //! - `ai_summary_service` - AI 摘要服务 //! - `project_context_builder` - 项目上下文构建器 -//! - `kiro_event_service` - Kiro 事件服务 //! - `api_key_provider_service` - API Key Provider 服务 -//! - `provider_pool_service` - Provider 池服务 -//! - `token_cache_service` - Token 缓存服务 // 无外部依赖的服务 pub mod context_memory_service; @@ -52,7 +47,6 @@ pub mod screenshot_capture_service; pub mod screenshot_image_service; pub mod sysinfo_service; pub mod update_check_service; -pub mod usage_service; pub mod voice_asr_service; pub mod voice_command_service; pub mod voice_config_service; @@ -73,7 +67,6 @@ pub mod backup_service; pub mod material_service; pub mod mcp_service; pub mod model_registry_service; -pub mod model_service; pub mod persona_service; pub mod prompt_service; // 依赖其他 services 的服务 @@ -81,12 +74,7 @@ pub mod ai_summary_service; pub mod project_context_builder; pub mod session_context_service; -// 事件服务 -pub mod kiro_event_service; - // 依赖 providers 的服务 pub mod api_key_provider_service; -pub mod provider_pool_service; pub mod provider_type_mapping; -pub mod token_cache_service; pub mod video_generation_service; diff --git a/src-tauri/crates/services/src/model_registry_service.rs b/src-tauri/crates/services/src/model_registry_service.rs index d68eae24c..39ca4ae7b 100644 --- a/src-tauri/crates/services/src/model_registry_service.rs +++ b/src-tauri/crates/services/src/model_registry_service.rs @@ -1133,10 +1133,10 @@ impl ModelRegistryService { .then(a.display_name.cmp(&b.display_name)) }); - // 3. 加载别名配置 + // 3. 凭证池 Provider 别名已退役,保留空集合避免旧模型目录回流 let mut aliases = HashMap::new(); let aliases_dir = models_dir.join("aliases"); - let alias_files = ["kiro", "antigravity", "codex", "gemini"]; + let alias_files: [&str; 0] = []; for alias_name in alias_files { let alias_file = aliases_dir.join(format!("{alias_name}.json")); diff --git a/src-tauri/crates/services/src/model_service.rs b/src-tauri/crates/services/src/model_service.rs deleted file mode 100644 index d26816f40..000000000 --- a/src-tauri/crates/services/src/model_service.rs +++ /dev/null @@ -1,563 +0,0 @@ -//! 模型管理服务 -//! -//! 提供统一的模型获取、缓存和查询接口,支持从不同 Provider 获取模型列表。 - -use lime_core::database::dao::api_key_provider::{infer_managed_runtime_spec, ApiProviderType}; -use lime_core::database::dao::provider_pool::ProviderPoolDao; -use lime_core::database::DbConnection; -use lime_core::models::provider_pool_model::{ - CredentialData, PoolProviderType, ProviderCredential, -}; -use reqwest::Client; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::time::Duration; - -/// 模型信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelInfo { - /// 模型 ID - pub id: String, - /// 模型对象类型(通常是 "model") - pub object: String, - /// 拥有者(如 "anthropic", "google", "openai") - pub owned_by: String, - /// 创建时间(可选) - #[serde(skip_serializing_if = "Option::is_none")] - pub created: Option, -} - -/// /v1/models 接口的响应格式 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelsResponse { - pub object: String, - pub data: Vec, -} - -/// 模型服务 -pub struct ModelService { - /// HTTP 客户端 - client: Client, - /// 请求超时时间 - timeout: Duration, -} - -impl Default for ModelService { - fn default() -> Self { - Self::new() - } -} - -impl ModelService { - /// 创建新的模型服务实例 - pub fn new() -> Self { - Self { - client: Client::builder() - .timeout(Duration::from_secs(10)) - .build() - .unwrap_or_default(), - timeout: Duration::from_secs(10), - } - } - - /// 从凭证获取支持的模型列表 - /// - /// 根据凭证类型调用相应的 /v1/models 接口 - pub async fn fetch_models_for_credential( - &self, - credential: &ProviderCredential, - ) -> Result, String> { - tracing::info!( - "[MODEL_SERVICE] 获取凭证模型列表: uuid={}, provider_type={}", - credential.uuid, - credential.provider_type - ); - - match &credential.credential { - // Antigravity 使用固定的模型列表(从配置文件读取) - CredentialData::AntigravityOAuth { .. } => { - // Antigravity 不提供标准的 /v1/models 接口 - // 直接返回预定义的模型列表 - tracing::info!("[MODEL_SERVICE] Antigravity 使用预定义模型列表"); - Ok(self.get_default_models_for_provider(&credential.provider_type)) - } - - // OAuth 凭证:由于需要处理 Token 刷新等复杂逻辑,暂时使用默认模型列表 - // TODO: 未来可以通过 ProviderPoolService 来获取动态模型列表 - CredentialData::KiroOAuth { .. } - | CredentialData::GeminiOAuth { .. } - | CredentialData::CodexOAuth { .. } - | CredentialData::ClaudeOAuth { .. } => { - tracing::info!("[MODEL_SERVICE] OAuth 凭证使用默认模型列表"); - Ok(self.get_default_models_for_provider(&credential.provider_type)) - } - - // API Key 类型凭证:直接调用 Provider 的 API - CredentialData::OpenAIKey { base_url, api_key } => { - tracing::info!("[MODEL_SERVICE] 使用 OpenAI API Key"); - self.fetch_models_openai(base_url.as_deref(), api_key).await - } - CredentialData::ClaudeKey { base_url, api_key } => { - tracing::info!("[MODEL_SERVICE] 使用 Claude API Key"); - self.fetch_models_claude(base_url.as_deref(), api_key).await - } - CredentialData::AnthropicKey { base_url, api_key } => { - tracing::info!("[MODEL_SERVICE] 使用 Anthropic API Key"); - self.fetch_models_anthropic(base_url.as_deref(), api_key) - .await - } - CredentialData::GeminiApiKey { - api_key, base_url, .. - } => { - tracing::info!("[MODEL_SERVICE] 使用 Gemini API Key"); - self.fetch_models_gemini(base_url.as_deref(), api_key).await - } - CredentialData::VertexKey { .. } => { - tracing::info!("[MODEL_SERVICE] Vertex AI 使用固定模型列表"); - // Vertex AI 使用固定的模型列表 - Ok(self.get_default_models_for_provider(&credential.provider_type)) - } - } - } - - /// 获取 OpenAI 兼容 API 的模型列表 - async fn fetch_models_openai( - &self, - base_url: Option<&str>, - api_key: &str, - ) -> Result, String> { - let url = format!("{}/v1/models", base_url.unwrap_or("https://api.openai.com")); - - tracing::info!("[MODEL_SERVICE] 请求 OpenAI API 获取模型列表: url={}", url); - - let response = self - .client - .get(&url) - .header("Authorization", format!("Bearer {api_key}")) - .timeout(self.timeout) - .send() - .await - .map_err(|e| { - tracing::error!("[MODEL_SERVICE] OpenAI 请求失败: {}", e); - format!("请求失败: {e}") - })?; - - let status = response.status(); - tracing::info!("[MODEL_SERVICE] OpenAI 响应状态码: {}", status); - - if !status.is_success() { - let error_body = response.text().await.unwrap_or_default(); - tracing::error!( - "[MODEL_SERVICE] OpenAI HTTP 错误: status={}, body={}", - status, - error_body - ); - return Err(format!("HTTP 错误: {status}")); - } - - let response_text = response.text().await.map_err(|e| { - tracing::error!("[MODEL_SERVICE] 读取 OpenAI 响应体失败: {}", e); - format!("读取响应体失败: {e}") - })?; - - tracing::debug!("[MODEL_SERVICE] OpenAI 响应体: {}", response_text); - - let models_response: ModelsResponse = - serde_json::from_str(&response_text).map_err(|e| { - tracing::error!( - "[MODEL_SERVICE] 解析 OpenAI 响应失败: {}, 响应内容: {}", - e, - response_text - ); - format!("解析响应失败: {e}") - })?; - - let model_ids: Vec = models_response.data.into_iter().map(|m| m.id).collect(); - - tracing::info!("[MODEL_SERVICE] OpenAI 成功获取 {} 个模型", model_ids.len()); - - Ok(model_ids) - } - - /// 获取 Claude API 的模型列表 - async fn fetch_models_claude( - &self, - base_url: Option<&str>, - api_key: &str, - ) -> Result, String> { - // Claude API 使用 OpenAI 兼容格式 - self.fetch_models_openai(base_url, api_key).await - } - - /// 获取 Anthropic API 的模型列表 - async fn fetch_models_anthropic( - &self, - base_url: Option<&str>, - api_key: &str, - ) -> Result, String> { - // Anthropic API 使用 OpenAI 兼容格式 - let url = format!( - "{}/v1/models", - base_url.unwrap_or("https://api.anthropic.com") - ); - - tracing::info!( - "[MODEL_SERVICE] 请求 Anthropic API 获取模型列表: url={}", - url - ); - - let runtime_spec = infer_managed_runtime_spec( - ApiProviderType::Anthropic, - base_url.unwrap_or("https://api.anthropic.com"), - ); - let auth_value = runtime_spec - .auth_prefix - .map(|prefix| format!("{prefix} {api_key}")) - .unwrap_or_else(|| api_key.to_string()); - - let mut request = self - .client - .get(&url) - .header(runtime_spec.auth_header, auth_value) - .timeout(self.timeout); - - for (name, value) in runtime_spec.extra_headers { - request = request.header(*name, *value); - } - - let response = request.send().await.map_err(|e| { - tracing::error!("[MODEL_SERVICE] Anthropic 请求失败: {}", e); - format!("请求失败: {e}") - })?; - - let status = response.status(); - tracing::info!("[MODEL_SERVICE] Anthropic 响应状态码: {}", status); - - if !status.is_success() { - let error_body = response.text().await.unwrap_or_default(); - tracing::error!( - "[MODEL_SERVICE] Anthropic HTTP 错误: status={}, body={}", - status, - error_body - ); - return Err(format!("HTTP 错误: {status}")); - } - - let response_text = response.text().await.map_err(|e| { - tracing::error!("[MODEL_SERVICE] 读取 Anthropic 响应体失败: {}", e); - format!("读取响应体失败: {e}") - })?; - - tracing::debug!("[MODEL_SERVICE] Anthropic 响应体: {}", response_text); - - let models_response: ModelsResponse = - serde_json::from_str(&response_text).map_err(|e| { - tracing::error!( - "[MODEL_SERVICE] 解析 Anthropic 响应失败: {}, 响应内容: {}", - e, - response_text - ); - format!("解析响应失败: {e}") - })?; - - let model_ids: Vec = models_response.data.into_iter().map(|m| m.id).collect(); - - tracing::info!( - "[MODEL_SERVICE] Anthropic 成功获取 {} 个模型", - model_ids.len() - ); - - Ok(model_ids) - } - - /// 获取 Gemini API 的模型列表 - async fn fetch_models_gemini( - &self, - base_url: Option<&str>, - api_key: &str, - ) -> Result, String> { - let url = format!( - "{}/v1/models?key={}", - base_url.unwrap_or("https://generativelanguage.googleapis.com"), - api_key - ); - - tracing::info!("[MODEL_SERVICE] 请求 Gemini API 获取模型列表: url={}", url); - - let response = self - .client - .get(&url) - .timeout(self.timeout) - .send() - .await - .map_err(|e| { - tracing::error!("[MODEL_SERVICE] Gemini 请求失败: {}", e); - format!("请求失败: {e}") - })?; - - let status = response.status(); - tracing::info!("[MODEL_SERVICE] Gemini 响应状态码: {}", status); - - if !status.is_success() { - let error_body = response.text().await.unwrap_or_default(); - tracing::error!( - "[MODEL_SERVICE] Gemini HTTP 错误: status={}, body={}", - status, - error_body - ); - return Err(format!("HTTP 错误: {status}")); - } - - let response_text = response.text().await.map_err(|e| { - tracing::error!("[MODEL_SERVICE] 读取 Gemini 响应体失败: {}", e); - format!("读取响应体失败: {e}") - })?; - - tracing::debug!("[MODEL_SERVICE] Gemini 响应体: {}", response_text); - - // Gemini API 返回格式不同,需要特殊处理 - let response_json: serde_json::Value = - serde_json::from_str(&response_text).map_err(|e| { - tracing::error!( - "[MODEL_SERVICE] 解析 Gemini 响应失败: {}, 响应内容: {}", - e, - response_text - ); - format!("解析响应失败: {e}") - })?; - - let models = response_json - .get("models") - .and_then(|m| m.as_array()) - .ok_or_else(|| { - tracing::error!("[MODEL_SERVICE] Gemini 响应格式错误: 缺少 models 字段"); - "响应格式错误".to_string() - })?; - - let model_ids: Vec = models - .iter() - .filter_map(|m| m.get("name").and_then(|n| n.as_str())) - .map(|name| { - // Gemini API 返回的是 "models/gemini-pro",需要提取模型名 - name.strip_prefix("models/").unwrap_or(name).to_string() - }) - .collect(); - - tracing::info!( - "[MODEL_SERVICE] Gemini 成功获取 {} 个模型: {:?}", - model_ids.len(), - model_ids - ); - - Ok(model_ids) - } - - /// 获取 Provider 的默认模型列表(用于无法动态获取的情况) - pub fn get_default_models_for_provider(&self, provider_type: &PoolProviderType) -> Vec { - match provider_type { - PoolProviderType::Kiro => vec![ - "claude-sonnet-4-5".to_string(), - "claude-sonnet-4-5-20250929".to_string(), - "claude-3-7-sonnet-20250219".to_string(), - "claude-3-5-sonnet-latest".to_string(), - "claude-haiku-4-5".to_string(), - ], - PoolProviderType::Gemini => vec![ - // Gemini 3 系列 - "gemini-3-pro-preview".to_string(), - "gemini-3-flash-preview".to_string(), - // Gemini 2.5 系列 - "gemini-2.5-pro".to_string(), - "gemini-2.5-flash".to_string(), - "gemini-2.5-flash-lite".to_string(), - // Gemini 2.0 系列 - "gemini-2.0-flash".to_string(), - "gemini-2.0-flash-lite".to_string(), - // Gemini 1.5 系列(已弃用但仍可用) - "gemini-1.5-pro".to_string(), - "gemini-1.5-flash".to_string(), - ], - PoolProviderType::Antigravity => vec![ - "gemini-2.5-computer-use-preview-10-2025".to_string(), - "gemini-3-pro-image-preview".to_string(), - "gemini-3-pro-preview".to_string(), - "gemini-3-flash-preview".to_string(), - "gemini-2.5-flash-preview".to_string(), - "gemini-claude-sonnet-4-5".to_string(), - "gemini-claude-sonnet-4-5-thinking".to_string(), - "gemini-claude-opus-4-5-thinking".to_string(), - ], - PoolProviderType::OpenAI => vec![ - "gpt-4o".to_string(), - "gpt-4o-mini".to_string(), - "gpt-3.5-turbo".to_string(), - ], - PoolProviderType::Claude | PoolProviderType::Anthropic => vec![ - "claude-sonnet-4-5-20250929".to_string(), - "claude-3-5-sonnet-20241022".to_string(), - "claude-3-5-haiku-20241022".to_string(), - ], - PoolProviderType::GeminiApiKey => { - vec!["gemini-2.5-flash".to_string(), "gemini-2.5-pro".to_string()] - } - _ => vec![], - } - } - - /// 更新凭证的支持模型列表到数据库 - pub fn update_credential_models( - &self, - db: &DbConnection, - credential_uuid: &str, - models: Vec, - ) -> Result<(), String> { - let conn = db.lock().map_err(|e| e.to_string())?; - - // 序列化模型列表为 JSON - let models_json = serde_json::to_string(&models).map_err(|e| e.to_string())?; - - conn.execute( - "UPDATE provider_pool_credentials SET supported_models = ?1, updated_at = ?2 WHERE uuid = ?3", - rusqlite::params![models_json, chrono::Utc::now().timestamp(), credential_uuid], - ) - .map_err(|e| e.to_string())?; - - Ok(()) - } - - /// 获取凭证的支持模型列表(从数据库) - pub fn get_credential_models( - &self, - db: &DbConnection, - credential_uuid: &str, - ) -> Result, String> { - let conn = db.lock().map_err(|e| e.to_string())?; - - let mut stmt = conn - .prepare("SELECT supported_models FROM provider_pool_credentials WHERE uuid = ?1") - .map_err(|e| e.to_string())?; - - let models_json: Option = stmt.query_row([credential_uuid], |row| row.get(0)).ok(); - - match models_json { - Some(json) => serde_json::from_str(&json).map_err(|e| e.to_string()), - None => Ok(vec![]), - } - } - - /// 获取所有凭证的模型列表(按 Provider 类型分组) - pub fn get_all_models_by_provider( - &self, - db: &DbConnection, - ) -> Result>, String> { - let conn = db.lock().map_err(|e| e.to_string())?; - let credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - drop(conn); - - let mut models_by_provider: HashMap> = HashMap::new(); - - for cred in credentials { - if cred.is_disabled || !cred.is_healthy { - continue; - } - - let provider_key = cred.provider_type.to_string(); - - models_by_provider - .entry(provider_key) - .or_default() - .extend(cred.supported_models); - } - - // 去重 - for models in models_by_provider.values_mut() { - models.sort(); - models.dedup(); - } - - Ok(models_by_provider) - } - - /// 获取可用的所有模型列表(合并所有健康凭证的模型) - pub fn get_all_available_models(&self, db: &DbConnection) -> Result, String> { - let models_by_provider = self.get_all_models_by_provider(db)?; - - let mut all_models: Vec = models_by_provider.into_values().flatten().collect(); - - all_models.sort(); - all_models.dedup(); - - Ok(all_models) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use lime_core::database::dao::provider_pool::ProviderPoolDao; - use lime_core::database::schema; - use lime_core::models::provider_pool_model::{CredentialData, PoolProviderType}; - use rusqlite::Connection; - use std::sync::{Arc, Mutex}; - - fn setup_test_db() -> DbConnection { - let conn = Connection::open_in_memory().expect("open in-memory db"); - schema::create_tables(&conn).expect("create schema"); - Arc::new(Mutex::new(conn)) - } - - #[test] - fn test_get_default_models_for_provider() { - let service = ModelService::new(); - - let kiro_models = service.get_default_models_for_provider(&PoolProviderType::Kiro); - assert!(!kiro_models.is_empty()); - assert!(kiro_models.contains(&"claude-sonnet-4-5".to_string())); - - let gemini_models = service.get_default_models_for_provider(&PoolProviderType::Gemini); - assert!(!gemini_models.is_empty()); - assert!(gemini_models.contains(&"gemini-2.5-flash".to_string())); - } - - #[test] - fn test_get_all_available_models_uses_loaded_supported_models_without_relocking_db() { - let db = setup_test_db(); - let mut openai = ProviderCredential::new( - PoolProviderType::OpenAI, - CredentialData::OpenAIKey { - api_key: "sk-test".to_string(), - base_url: None, - }, - ); - openai.supported_models = vec!["gpt-4o".to_string(), "gpt-4.1".to_string()]; - - let mut gemini = ProviderCredential::new( - PoolProviderType::GeminiApiKey, - CredentialData::GeminiApiKey { - api_key: "gm-test".to_string(), - base_url: None, - excluded_models: Vec::new(), - }, - ); - gemini.supported_models = vec!["gemini-2.5-flash".to_string(), "gpt-4o".to_string()]; - - { - let conn = db.lock().expect("lock db for seed"); - ProviderPoolDao::insert(&conn, &openai).expect("insert openai credential"); - ProviderPoolDao::insert(&conn, &gemini).expect("insert gemini credential"); - } - - let models = ModelService::new() - .get_all_available_models(&db) - .expect("list available models"); - - assert_eq!( - models, - vec![ - "gemini-2.5-flash".to_string(), - "gpt-4.1".to_string(), - "gpt-4o".to_string(), - ] - ); - } -} diff --git a/src-tauri/crates/services/src/provider_pool_service.rs b/src-tauri/crates/services/src/provider_pool_service.rs deleted file mode 100644 index 640b16962..000000000 --- a/src-tauri/crates/services/src/provider_pool_service.rs +++ /dev/null @@ -1,2421 +0,0 @@ -//! Provider Pool 管理服务 -//! -//! 提供凭证池的选择、健康检测、负载均衡等功能。 - -#![allow(dead_code)] - -use crate::api_key_provider_service::ApiKeyProviderService; -use crate::provider_type_mapping::{ - api_provider_type_to_pool_type, is_custom_provider_id, parse_pool_provider_type, - resolve_pool_provider_type_or_default, -}; -use chrono::Utc; -use lime_core::database::dao::api_key_provider::{infer_managed_runtime_spec, ApiProviderType}; -use lime_core::database::dao::provider_pool::ProviderPoolDao; -use lime_core::database::DbConnection; -use lime_core::models::client_type::ClientType; -use lime_core::models::provider_pool_model::{ - get_default_check_model, get_oauth_creds_path, CredentialData, CredentialDisplay, - HealthCheckResult, OAuthStatus, PoolProviderType, PoolStats, ProviderCredential, - ProviderPoolOverview, -}; -use lime_providers::providers::antigravity::TokenRefreshError; -use lime_providers::providers::kiro::KiroProvider; -use reqwest::Client; -use serde::{Deserialize, Serialize}; - -/// 扩展 ProviderCredential 的客户端兼容性检查 -/// (此方法依赖 server::client_detector,不适合放在 core crate) -trait ProviderCredentialClientCompat { - fn is_compatible_with_client(&self, client_type: Option<&ClientType>) -> bool; -} - -impl ProviderCredentialClientCompat for ProviderCredential { - fn is_compatible_with_client(&self, client_type: Option<&ClientType>) -> bool { - if let Some(error_msg) = &self.last_error_message { - if error_msg.contains("only authorized for use with Claude Code") { - return matches!(client_type, Some(ClientType::ClaudeCode)); - } - } - true - } -} -use std::collections::{HashMap, HashSet}; -use std::path::Path; -use std::sync::atomic::AtomicUsize; -use std::time::Duration; - -/// 凭证健康信息 -/// Requirements: 3.1, 3.2 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CredentialHealthInfo { - /// 凭证 UUID - pub uuid: String, - /// 凭证名称 - pub name: Option, - /// Provider 类型 - pub provider_type: String, - /// 是否健康 - pub is_healthy: bool, - /// 最后错误信息 - pub last_error: Option, - /// 最后错误时间(RFC3339 格式) - pub last_error_time: Option, - /// 错误次数 - pub failure_count: u32, - /// 是否需要重新授权 - pub requires_reauth: bool, -} - -/// 凭证选择错误 -/// Requirements: 3.4 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub enum SelectionError { - /// 没有凭证 - NoCredentials, - /// 所有凭证都不健康 - AllUnhealthy { details: Vec }, - /// 模型不支持 - ModelNotSupported { model: String }, -} - -/// 凭证池管理服务 -pub struct ProviderPoolService { - /// HTTP 客户端(用于健康检测) - client: Client, - /// 轮询索引(按 provider_type 和可选的 model 分组) - round_robin_index: std::sync::RwLock>, - /// 最大错误次数(超过后标记为不健康) - max_error_count: u32, - /// 健康检查超时时间 - health_check_timeout: Duration, -} - -impl Default for ProviderPoolService { - fn default() -> Self { - Self::new() - } -} - -impl ProviderPoolService { - pub fn new() -> Self { - Self { - client: Client::builder() - .timeout(Duration::from_secs(30)) - .build() - .unwrap_or_default(), - round_robin_index: std::sync::RwLock::new(HashMap::new()), - max_error_count: 3, - health_check_timeout: Duration::from_secs(30), - } - } - - /// 获取所有凭证概览 - pub fn get_overview(&self, db: &DbConnection) -> Result, String> { - let conn = lime_core::database::lock_db(db)?; - let grouped = ProviderPoolDao::get_grouped(&conn).map_err(|e| e.to_string())?; - - let mut overview = Vec::new(); - for (provider_type, mut credentials) in grouped { - // 为每个凭证加载 token 缓存 - for cred in &mut credentials { - cred.cached_token = ProviderPoolDao::get_token_cache(&conn, &cred.uuid) - .ok() - .flatten(); - } - - let stats = PoolStats::from_credentials(&credentials); - let displays: Vec = credentials.iter().map(|c| c.into()).collect(); - - overview.push(ProviderPoolOverview { - provider_type: provider_type.to_string(), - stats, - credentials: displays, - }); - } - - // 按 provider_type 排序 - overview.sort_by(|a, b| a.provider_type.cmp(&b.provider_type)); - Ok(overview) - } - - /// 获取指定类型的凭证列表 - pub fn get_by_type( - &self, - db: &DbConnection, - provider_type: &str, - ) -> Result, String> { - let pt = parse_pool_provider_type(provider_type)?; - let conn = lime_core::database::lock_db(db)?; - let mut credentials = - ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?; - - // 为每个凭证加载 token 缓存 - for cred in &mut credentials { - cred.cached_token = ProviderPoolDao::get_token_cache(&conn, &cred.uuid) - .ok() - .flatten(); - } - - Ok(credentials.iter().map(|c| c.into()).collect()) - } - - /// 获取指定 UUID 的凭证 - pub fn get_by_uuid( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result, String> { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::get_by_uuid(&conn, uuid).map_err(|e| e.to_string()) - } - - /// 添加凭证 - pub fn add_credential( - &self, - db: &DbConnection, - provider_type: &str, - credential: CredentialData, - name: Option, - check_health: Option, - check_model_name: Option, - ) -> Result { - let pt = parse_pool_provider_type(provider_type)?; - - let mut cred = ProviderCredential::new(pt, credential); - cred.name = name; - cred.check_health = check_health.unwrap_or(true); - cred.check_model_name = check_model_name; - - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::insert(&conn, &cred).map_err(|e| e.to_string())?; - - Ok(cred) - } - - /// 更新凭证 - pub fn update_credential( - &self, - db: &DbConnection, - uuid: &str, - name: Option, - is_disabled: Option, - check_health: Option, - check_model_name: Option, - not_supported_models: Option>, - proxy_url: Option, - ) -> Result { - let conn = lime_core::database::lock_db(db)?; - let mut cred = ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))?; - - // 处理 name:空字符串表示清除,None 表示不修改 - if let Some(n) = name { - cred.name = if n.is_empty() { None } else { Some(n) }; - } - if let Some(d) = is_disabled { - cred.is_disabled = d; - } - if let Some(c) = check_health { - cred.check_health = c; - } - // 处理 check_model_name:空字符串表示清除,None 表示不修改 - if let Some(m) = check_model_name { - cred.check_model_name = if m.is_empty() { None } else { Some(m) }; - } - if let Some(models) = not_supported_models { - cred.not_supported_models = models; - } - // 处理 proxy_url:空字符串表示清除,None 表示不修改 - if let Some(p) = proxy_url { - cred.proxy_url = if p.is_empty() { None } else { Some(p) }; - } - cred.updated_at = Utc::now(); - - ProviderPoolDao::update(&conn, &cred).map_err(|e| e.to_string())?; - Ok(cred) - } - - /// 删除凭证 - pub fn delete_credential(&self, db: &DbConnection, uuid: &str) -> Result { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::delete(&conn, uuid).map_err(|e| e.to_string()) - } - - /// 选择一个可用的凭证(智能轮换策略) - /// - /// 增强版轮换策略,考虑以下因素: - /// - 健康状态:优先选择健康的凭证 - /// - 使用频率:优先选择使用次数较少的凭证 - /// - 错误率:避免选择错误次数过多的凭证 - /// - 冷却时间:避免短时间内重复使用同一凭证 - pub fn select_credential( - &self, - db: &DbConnection, - provider_type: &str, - model: Option<&str>, - ) -> Result, String> { - self.select_credential_with_client_check(db, provider_type, model, None) - } - - /// 选择凭证并检查客户端兼容性 - /// - /// 内部方法,支持客户端类型检查 - pub fn select_credential_with_client_check( - &self, - db: &DbConnection, - provider_type: &str, - model: Option<&str>, - client_type: Option<&lime_core::models::client_type::ClientType>, - ) -> Result, String> { - if is_custom_provider_id(provider_type) { - eprintln!("[SELECT_CREDENTIAL] custom provider '{provider_type}' 使用智能降级路径"); - return Ok(None); - } - - // 对于未知的 provider_type,直接返回 None(不是错误) - // 这样可以让 select_credential_with_fallback 继续尝试智能降级 - let pt: PoolProviderType = match parse_pool_provider_type(provider_type) { - Ok(pt) => pt, - Err(_) => { - eprintln!( - "[SELECT_CREDENTIAL] 未知的 provider_type '{provider_type}', 返回 None 以便智能降级" - ); - return Ok(None); - } - }; - let conn = lime_core::database::lock_db(db)?; - - // 获取凭证,对于 AI Provider 类型,也查找 Assistant 类型的凭证 - let mut credentials = - ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?; - eprintln!( - "[SELECT_CREDENTIAL] provider_type={}, pt={:?}, initial_count={}", - provider_type, - pt, - credentials.len() - ); - - // AI Provider 和 Assistant 共享凭证(都使用 AI Provider API) - if pt == PoolProviderType::Anthropic { - let assistant_creds = ProviderPoolDao::get_by_type(&conn, &PoolProviderType::Claude) - .map_err(|e| e.to_string())?; - eprintln!( - "[SELECT_CREDENTIAL] AI Provider: adding {} Assistant credentials", - assistant_creds.len() - ); - credentials.extend(assistant_creds); - } else if pt == PoolProviderType::Claude { - let ai_provider_creds = - ProviderPoolDao::get_by_type(&conn, &PoolProviderType::Anthropic) - .map_err(|e| e.to_string())?; - eprintln!( - "[SELECT_CREDENTIAL] Assistant: adding {} AI Provider credentials", - ai_provider_creds.len() - ); - credentials.extend(ai_provider_creds); - } - - let recoverable_antigravity: HashMap> = - if pt == PoolProviderType::Antigravity { - credentials - .iter() - .filter(|cred| Self::should_auto_recover_antigravity_credential(cred)) - .map(|cred| (cred.uuid.clone(), cred.check_model_name.clone())) - .collect() - } else { - HashMap::new() - }; - - drop(conn); - - if !recoverable_antigravity.is_empty() { - for (uuid, check_model_name) in &recoverable_antigravity { - self.mark_healthy(db, uuid, check_model_name.as_deref())?; - } - for cred in &mut credentials { - if let Some(check_model_name) = recoverable_antigravity.get(&cred.uuid) { - cred.mark_healthy(check_model_name.clone()); - } - } - } - - eprintln!( - "[SELECT_CREDENTIAL] total_credentials={}, model={:?}", - credentials.len(), - model - ); - - // 过滤可用的凭证 - let mut available: Vec<_> = credentials - .into_iter() - .filter(|c| { - let is_avail = c.is_available(); - if !is_avail { - eprintln!( - "[SELECT_CREDENTIAL] credential {} (type={}) is_available={} (is_healthy={}, is_disabled={}, error_count={}, last_error={:?})", - c.name.as_deref().unwrap_or("unnamed"), - c.provider_type, - is_avail, - c.is_healthy, - c.is_disabled, - c.error_count, - c.last_error_message - ); - } else { - eprintln!( - "[SELECT_CREDENTIAL] credential {} (type={}) is_available={}", - c.name.as_deref().unwrap_or("unnamed"), - c.provider_type, - is_avail - ); - } - is_avail - }) - .collect(); - - eprintln!( - "[SELECT_CREDENTIAL] after is_available filter: {}", - available.len() - ); - - // 如果指定了模型,进一步过滤支持该模型的凭证 - if let Some(m) = model { - available.retain(|c| { - let supports = c.supports_model(m); - eprintln!( - "[SELECT_CREDENTIAL] credential {} supports_model({})={}", - c.name.as_deref().unwrap_or("unnamed"), - m, - supports - ); - supports - }); - } - - // 过滤客户端兼容的凭证 - available.retain(|c| { - let compatible = c.is_compatible_with_client(client_type); - if !compatible { - eprintln!( - "[SELECT_CREDENTIAL] credential {} 不兼容客户端类型 {:?}", - c.name.as_deref().unwrap_or("unnamed"), - client_type - ); - } - compatible - }); - - eprintln!( - "[SELECT_CREDENTIAL] after client compatibility filter: {}", - available.len() - ); - - if available.is_empty() { - return Ok(None); - } - - // 如果只有一个可用凭证,直接返回 - if available.len() == 1 { - return Ok(Some(available.into_iter().next().unwrap())); - } - - // 智能选择:基于权重分数选择最优凭证 - let selected = self.select_best_credential_by_weight(&available); - - Ok(Some(selected)) - } - - fn should_auto_recover_antigravity_credential(cred: &ProviderCredential) -> bool { - let CredentialData::AntigravityOAuth { - creds_file_path, .. - } = &cred.credential - else { - return false; - }; - - if cred.is_healthy || Path::new(creds_file_path).exists() { - return false; - } - - let last_error = cred.last_error_message.as_deref().unwrap_or_default(); - let is_missing_path_failure = last_error.contains("Failed to load credentials") - || last_error.contains("No such file or directory"); - if !is_missing_path_failure { - return false; - } - - lime_providers::providers::antigravity::AntigravityProvider::default_creds_path().exists() - } - - /// 带智能降级的凭证选择 - /// - /// 当 Provider Pool 无可用凭证时,自动从 API Key Provider 降级查找 - /// - /// # 参数 - /// - `db`: 数据库连接 - /// - `api_key_service`: API Key Provider 服务 - /// - `provider_type`: Provider 类型字符串,如 "assistant", "openai" - /// - `model`: 可选的模型名称 - /// - `provider_id_hint`: 可选的 provider_id 提示,用于 60+ Provider 直接查找 - /// - `client_type`: 可选的客户端类型,用于凭证兼容性检查 - /// - /// # 返回 - /// - `Ok(Some(credential))`: 找到可用凭证(来自 Pool 或降级) - /// - `Ok(None)`: 没有找到任何可用凭证 - /// - `Err(e)`: 查询过程中发生错误 - pub async fn select_credential_with_fallback( - &self, - db: &DbConnection, - api_key_service: &ApiKeyProviderService, - provider_type: &str, - model: Option<&str>, - provider_id_hint: Option<&str>, - client_type: Option<&lime_core::models::client_type::ClientType>, - ) -> Result, String> { - eprintln!( - "[select_credential_with_fallback] 开始: provider_type={provider_type}, model={model:?}, provider_id_hint={provider_id_hint:?}" - ); - - // Step 1: 尝试从 Provider Pool 选择 (OAuth + API Key) - if let Some(cred) = - self.select_credential_with_client_check(db, provider_type, model, client_type)? - { - eprintln!( - "[select_credential_with_fallback] 从 Provider Pool 找到凭证: {:?}", - cred.name - ); - return Ok(Some(cred)); - } - eprintln!("[select_credential_with_fallback] Provider Pool 未找到凭证,尝试智能降级"); - - // Step 2: 智能降级到 API Key Provider - let mut pt = resolve_pool_provider_type_or_default(provider_type); - let mut resolved_provider_id_hint = provider_id_hint; - - // 对 custom-* 场景优先查询真实 Provider 类型,避免默认按 OpenAI 协议处理 - if is_custom_provider_id(provider_type) { - resolved_provider_id_hint = Some(provider_type); - } - - if let Some(custom_provider_id) = - resolved_provider_id_hint.filter(|id| is_custom_provider_id(id)) - { - match api_key_service.get_provider(db, custom_provider_id) { - Ok(Some(provider_with_keys)) => { - pt = api_provider_type_to_pool_type(provider_with_keys.provider.provider_type); - eprintln!( - "[select_credential_with_fallback] custom provider '{}' 真实类型 {:?} -> {:?}", - custom_provider_id, - provider_with_keys.provider.provider_type, - pt - ); - } - Ok(None) => { - eprintln!( - "[select_credential_with_fallback] custom provider '{custom_provider_id}' 不存在,继续使用解析类型 {pt:?}" - ); - } - Err(e) => { - eprintln!( - "[select_credential_with_fallback] 查询 custom provider '{custom_provider_id}' 失败: {e},继续使用解析类型 {pt:?}" - ); - } - } - } - - eprintln!( - "[select_credential_with_fallback] 解析 provider_type '{provider_type}' -> {pt:?}" - ); - - // 传入 provider_id_hint 支持 60+ Provider - eprintln!("[select_credential_with_fallback] 调用 get_fallback_credential"); - if let Some(cred) = api_key_service - .get_fallback_credential(db, &pt, resolved_provider_id_hint, client_type) - .await? - { - eprintln!( - "[select_credential_with_fallback] 智能降级成功: {:?}", - cred.name - ); - return Ok(Some(cred)); - } - - // Step 3: 都没有找到 - eprintln!( - "[select_credential_with_fallback] 未找到任何凭证 for provider_type='{provider_type}'" - ); - Ok(None) - } - - /// 基于权重分数选择最优凭证 - fn select_best_credential_by_weight( - &self, - credentials: &[ProviderCredential], - ) -> ProviderCredential { - let now = chrono::Utc::now(); - - let mut best_score = f64::MIN; - let mut best_credential = None; - - for cred in credentials { - let score = self.calculate_credential_score(cred, now, credentials); - if score > best_score { - best_score = score; - best_credential = Some(cred); - } - } - - best_credential.unwrap().clone() - } - - /// 计算凭证的综合分数(分数越高越优先) - fn calculate_credential_score( - &self, - cred: &ProviderCredential, - now: chrono::DateTime, - all_credentials: &[ProviderCredential], - ) -> f64 { - let mut score = 0.0; - - // 1. 健康状态权重 (40分) - if cred.is_healthy { - score += 40.0; - } else { - score -= 20.0; // 不健康的凭证严重扣分 - } - - // 2. 使用频率权重 (30分) - 使用次数越少分数越高 - let max_usage = all_credentials - .iter() - .map(|c| c.usage_count) - .max() - .unwrap_or(1); - if max_usage > 0 { - let usage_ratio = cred.usage_count as f64 / max_usage as f64; - score += 30.0 * (1.0 - usage_ratio); // 使用越少分数越高 - } else { - score += 30.0; // 如果都没使用过,给满分 - } - - // 3. 错误率权重 (20分) - 错误越少分数越高 - let total_requests = cred.usage_count + cred.error_count as u64; - if total_requests > 0 { - let error_ratio = cred.error_count as f64 / total_requests as f64; - score += 20.0 * (1.0 - error_ratio); // 错误率越低分数越高 - } else { - score += 20.0; // 没有历史记录给满分 - } - - // 4. 冷却时间权重 (10分) - 距离上次使用时间越长分数越高 - if let Some(last_used) = &cred.last_used { - let duration_since_last_use = now.signed_duration_since(*last_used); - let minutes_since_last_use = duration_since_last_use.num_minutes() as f64; - - // 超过5分钟的冷却时间给满分,否则按比例给分 - let cooldown_score = if minutes_since_last_use >= 5.0 { - 10.0 - } else { - 10.0 * (minutes_since_last_use / 5.0) - }; - score += cooldown_score; - } else { - score += 10.0; // 从未使用过给满分 - } - - score - } - - /// 记录凭证使用 - pub fn record_usage(&self, db: &DbConnection, uuid: &str) -> Result<(), String> { - let conn = lime_core::database::lock_db(db)?; - let cred = ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))?; - - ProviderPoolDao::update_usage(&conn, uuid, cred.usage_count + 1, Utc::now()) - .map_err(|e| e.to_string()) - } - - /// 标记凭证为健康 - pub fn mark_healthy( - &self, - db: &DbConnection, - uuid: &str, - check_model: Option<&str>, - ) -> Result<(), String> { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::update_health_status( - &conn, - uuid, - true, - 0, - None, - None, - Some(Utc::now()), - check_model, - ) - .map_err(|e| e.to_string()) - } - - /// 标记凭证为不健康 - pub fn mark_unhealthy( - &self, - db: &DbConnection, - uuid: &str, - error_message: Option<&str>, - ) -> Result<(), String> { - let conn = lime_core::database::lock_db(db)?; - let cred = ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))?; - - let new_error_count = cred.error_count + 1; - let is_healthy = new_error_count < self.max_error_count; - - ProviderPoolDao::update_health_status( - &conn, - uuid, - is_healthy, - new_error_count, - Some(Utc::now()), - error_message, - None, - None, - ) - .map_err(|e| e.to_string()) - } - - /// 重置凭证计数器 - pub fn reset_counters(&self, db: &DbConnection, uuid: &str) -> Result<(), String> { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::reset_counters(&conn, uuid).map_err(|e| e.to_string()) - } - - /// 重置指定类型的所有凭证健康状态 - pub fn reset_health_by_type( - &self, - db: &DbConnection, - provider_type: &str, - ) -> Result { - let pt = parse_pool_provider_type(provider_type)?; - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::reset_health_by_type(&conn, &pt).map_err(|e| e.to_string()) - } - - /// 获取凭证健康状态 - /// Requirements: 3.2 - pub fn get_credential_health( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result, String> { - let conn = lime_core::database::lock_db(db)?; - let cred = ProviderPoolDao::get_by_uuid(&conn, uuid).map_err(|e| e.to_string())?; - - Ok(cred.map(|c| CredentialHealthInfo { - uuid: c.uuid.clone(), - name: c.name.clone(), - provider_type: c.provider_type.to_string(), - is_healthy: c.is_healthy, - last_error: c.last_error_message.clone(), - last_error_time: c.last_error_time.map(|t| t.to_rfc3339()), - failure_count: c.error_count, - requires_reauth: c - .last_error_message - .as_ref() - .map(|e| e.contains("invalid_grant") || e.contains("重新授权")) - .unwrap_or(false), - })) - } - - /// 获取所有凭证的健康状态 - /// Requirements: 3.2 - pub fn get_all_credential_health( - &self, - db: &DbConnection, - ) -> Result, String> { - let conn = lime_core::database::lock_db(db)?; - let credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - - Ok(credentials - .into_iter() - .map(|c| CredentialHealthInfo { - uuid: c.uuid.clone(), - name: c.name.clone(), - provider_type: c.provider_type.to_string(), - is_healthy: c.is_healthy, - last_error: c.last_error_message.clone(), - last_error_time: c.last_error_time.map(|t| t.to_rfc3339()), - failure_count: c.error_count, - requires_reauth: c - .last_error_message - .as_ref() - .map(|e| e.contains("invalid_grant") || e.contains("重新授权")) - .unwrap_or(false), - }) - .collect()) - } - - /// 标记凭证为不健康(带详细错误信息) - /// Requirements: 3.1, 3.2 - pub fn mark_unhealthy_with_details( - &self, - db: &DbConnection, - uuid: &str, - error: &TokenRefreshError, - ) -> Result<(), String> { - let error_message = error.user_message(); - let requires_reauth = error.requires_reauth(); - - let conn = lime_core::database::lock_db(db)?; - let cred = ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))?; - - let new_error_count = cred.error_count + 1; - // 如果需要重新授权,直接标记为不健康 - let is_healthy = if requires_reauth { - false - } else { - new_error_count < self.max_error_count - }; - - let error_msg = if requires_reauth { - format!("[需要重新授权] {error_message}") - } else { - error_message - }; - - ProviderPoolDao::update_health_status( - &conn, - uuid, - is_healthy, - new_error_count, - Some(Utc::now()), - Some(&error_msg), - None, - None, - ) - .map_err(|e| e.to_string()) - } - - /// 选择一个健康的凭证 - /// Requirements: 2.4, 3.3, 3.4 - pub fn select_healthy_credential( - &self, - db: &DbConnection, - provider_type: &str, - model: Option<&str>, - ) -> Result { - let pt: PoolProviderType = provider_type - .parse() - .map_err(|_| SelectionError::NoCredentials)?; - let conn = lime_core::database::lock_db(db).map_err(|_| SelectionError::NoCredentials)?; - let credentials = - ProviderPoolDao::get_by_type(&conn, &pt).map_err(|_| SelectionError::NoCredentials)?; - drop(conn); - - if credentials.is_empty() { - return Err(SelectionError::NoCredentials); - } - - // 过滤可用的凭证(健康且未禁用) - let mut available: Vec<_> = credentials - .iter() - .filter(|c| c.is_available() && c.is_healthy) - .collect(); - - // 如果指定了模型,进一步过滤支持该模型的凭证 - if let Some(m) = model { - available.retain(|c| c.supports_model(m)); - if available.is_empty() { - // 检查是否有凭证支持该模型但不健康 - let unhealthy_supporting: Vec<_> = credentials - .iter() - .filter(|c| c.supports_model(m) && !c.is_healthy) - .collect(); - - if !unhealthy_supporting.is_empty() { - // 返回不健康凭证的详细信息 - let details: Vec = unhealthy_supporting - .into_iter() - .map(|c| CredentialHealthInfo { - uuid: c.uuid.clone(), - name: c.name.clone(), - provider_type: c.provider_type.to_string(), - is_healthy: c.is_healthy, - last_error: c.last_error_message.clone(), - last_error_time: c.last_error_time.map(|t| t.to_rfc3339()), - failure_count: c.error_count, - requires_reauth: c - .last_error_message - .as_ref() - .map(|e| e.contains("invalid_grant") || e.contains("重新授权")) - .unwrap_or(false), - }) - .collect(); - return Err(SelectionError::AllUnhealthy { details }); - } - - return Err(SelectionError::ModelNotSupported { - model: m.to_string(), - }); - } - } - - if available.is_empty() { - // 所有凭证都不健康 - let details: Vec = credentials - .iter() - .filter(|c| !c.is_healthy) - .map(|c| CredentialHealthInfo { - uuid: c.uuid.clone(), - name: c.name.clone(), - provider_type: c.provider_type.to_string(), - is_healthy: c.is_healthy, - last_error: c.last_error_message.clone(), - last_error_time: c.last_error_time.map(|t| t.to_rfc3339()), - failure_count: c.error_count, - requires_reauth: c - .last_error_message - .as_ref() - .map(|e| e.contains("invalid_grant") || e.contains("重新授权")) - .unwrap_or(false), - }) - .collect(); - return Err(SelectionError::AllUnhealthy { details }); - } - - // 使用轮询策略选择凭证 - let key = format!("{}:{}", provider_type, model.unwrap_or("*")); - let index = { - let indices = self.round_robin_index.read().unwrap(); - indices - .get(&key) - .map(|i| i.load(std::sync::atomic::Ordering::Relaxed)) - .unwrap_or(0) - }; - - let selected_index = index % available.len(); - let selected = available[selected_index].clone(); - - // 更新轮询索引 - { - let mut indices = self.round_robin_index.write().unwrap(); - indices - .entry(key) - .or_insert_with(|| AtomicUsize::new(0)) - .store(index + 1, std::sync::atomic::Ordering::Relaxed); - } - - Ok(selected) - } - - /// 执行单个凭证的健康检查 - /// - /// 如果遇到 401 错误,会自动尝试刷新 token 后重试 - pub async fn check_credential_health( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result { - let cred = { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))? - }; - - let check_model = cred - .check_model_name - .clone() - .unwrap_or_else(|| get_default_check_model(cred.provider_type).to_string()); - - let start = std::time::Instant::now(); - let result = self - .perform_health_check(&cred.credential, &check_model) - .await; - let duration_ms = start.elapsed().as_millis() as u64; - - match result { - Ok(_) => { - self.mark_healthy(db, uuid, Some(&check_model))?; - Ok(HealthCheckResult { - uuid: uuid.to_string(), - success: true, - model: Some(check_model), - message: Some("Health check passed".to_string()), - duration_ms, - }) - } - Err(e) => { - // 如果是 401 错误,尝试刷新 token 后重试 - if e.contains("401") || e.contains("Unauthorized") { - tracing::info!("[健康检查] 检测到 401 错误,尝试刷新 token: {}", uuid); - - // 尝试刷新 token - match self.refresh_credential_token(db, uuid).await { - Ok(_) => { - tracing::info!("[健康检查] Token 刷新成功,重新检查健康状态"); - - // 重新获取凭证(token 已更新) - let updated_cred = { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))? - }; - - // 重新执行健康检查 - let retry_start = std::time::Instant::now(); - let retry_result = self - .perform_health_check(&updated_cred.credential, &check_model) - .await; - let retry_duration_ms = retry_start.elapsed().as_millis() as u64; - - match retry_result { - Ok(_) => { - self.mark_healthy(db, uuid, Some(&check_model))?; - return Ok(HealthCheckResult { - uuid: uuid.to_string(), - success: true, - model: Some(check_model), - message: Some( - "Health check passed after token refresh".to_string(), - ), - duration_ms: duration_ms + retry_duration_ms, - }); - } - Err(retry_e) => { - tracing::warn!("[健康检查] Token 刷新后仍然失败: {}", retry_e); - self.mark_unhealthy(db, uuid, Some(&retry_e))?; - return Ok(HealthCheckResult { - uuid: uuid.to_string(), - success: false, - model: Some(check_model), - message: Some(retry_e), - duration_ms: duration_ms + retry_duration_ms, - }); - } - } - } - Err(refresh_err) => { - tracing::warn!("[健康检查] Token 刷新失败: {}", refresh_err); - // Token 刷新失败,返回原始错误 - self.mark_unhealthy(db, uuid, Some(&e))?; - return Ok(HealthCheckResult { - uuid: uuid.to_string(), - success: false, - model: Some(check_model), - message: Some(format!("{e} (Token 刷新失败: {refresh_err})")), - duration_ms, - }); - } - } - } - - self.mark_unhealthy(db, uuid, Some(&e))?; - Ok(HealthCheckResult { - uuid: uuid.to_string(), - success: false, - model: Some(check_model), - message: Some(e), - duration_ms, - }) - } - } - } - - /// 执行指定类型的所有凭证健康检查 - pub async fn check_type_health( - &self, - db: &DbConnection, - provider_type: &str, - ) -> Result, String> { - let pt = parse_pool_provider_type(provider_type)?; - let credentials = { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())? - }; - - let mut results = Vec::new(); - for cred in credentials { - if cred.is_disabled || !cred.check_health { - continue; - } - - let result = self.check_credential_health(db, &cred.uuid).await?; - results.push(result); - } - - Ok(results) - } - - /// 执行实际的健康检查请求 - async fn perform_health_check( - &self, - credential: &CredentialData, - model: &str, - ) -> Result<(), String> { - // 根据凭证类型构建测试请求 - match credential { - CredentialData::KiroOAuth { creds_file_path } => { - self.check_kiro_health(creds_file_path, model).await - } - CredentialData::GeminiOAuth { - creds_file_path, - project_id, - } => { - self.check_gemini_health(creds_file_path, project_id.as_deref(), model) - .await - } - CredentialData::AntigravityOAuth { - creds_file_path, - project_id, - } => { - self.check_antigravity_health(creds_file_path, project_id.as_deref(), model) - .await - } - CredentialData::OpenAIKey { api_key, base_url } => { - self.check_openai_health(api_key, base_url.as_deref(), model) - .await - } - CredentialData::ClaudeKey { api_key, base_url } => { - self.check_claude_health(api_key, base_url.as_deref(), model) - .await - } - CredentialData::VertexKey { - api_key, base_url, .. - } => { - self.check_vertex_health(api_key, base_url.as_deref(), model) - .await - } - CredentialData::GeminiApiKey { - api_key, base_url, .. - } => { - self.check_gemini_api_key_health(api_key, base_url.as_deref(), model) - .await - } - CredentialData::CodexOAuth { - creds_file_path, - api_base_url, - } => { - self.check_codex_health(creds_file_path, api_base_url.as_deref(), model) - .await - } - CredentialData::ClaudeOAuth { creds_file_path } => { - self.check_claude_oauth_health(creds_file_path, model).await - } - CredentialData::AnthropicKey { api_key, base_url } => { - // Anthropic API Key 使用与 Claude API Key 相同的健康检查逻辑 - self.check_claude_health(api_key, base_url.as_deref(), model) - .await - } - } - } - - /// 将技术错误转换为用户友好的错误信息 - fn format_user_friendly_error(&self, error: &str, provider_type: &str) -> String { - if error.contains("No client_id") { - format!("OAuth 配置不完整:缺少必要的认证参数。\n💡 解决方案:\n1. 检查 {provider_type} OAuth 凭证配置是否完整\n2. 如问题持续,建议删除后重新添加此凭证\n3. 或者切换到其他可用的凭证") - } else if error.contains("请求失败") || error.contains("error sending request") { - format!("网络连接失败,无法访问 {provider_type} 服务。\n💡 解决方案:\n1. 检查网络连接是否正常\n2. 确认防火墙或代理设置\n3. 稍后重试,如问题持续请联系网络管理员") - } else if error.contains("HTTP 401") || error.contains("HTTP 403") { - format!("{provider_type} 认证失败,凭证可能已过期或无效。\n💡 解决方案:\n1. 点击\"刷新\"按钮尝试更新 Token\n2. 如刷新失败,请删除后重新添加此凭证\n3. 检查账户权限是否正常") - } else if error.contains("HTTP 429") { - format!("{provider_type} 请求频率过高,已被限流。\n💡 解决方案:\n1. 稍等几分钟后再次尝试\n2. 考虑添加更多凭证分散负载") - } else if error.contains("HTTP 500") - || error.contains("HTTP 502") - || error.contains("HTTP 503") - { - format!("{provider_type} 服务暂时不可用。\n💡 解决方案:\n1. 这通常是服务提供方的临时问题\n2. 请稍后重试\n3. 如问题持续,可尝试其他凭证") - } else if error.contains("读取凭证文件失败") || error.contains("解析凭证失败") - { - "凭证文件损坏或不可读。\n💡 解决方案:\n1. 凭证文件可能已损坏\n2. 建议删除此凭证后重新添加\n3. 确保文件权限正确且格式为有效的 JSON".to_string() - } else { - // 对于其他未识别的错误,提供通用建议 - format!("操作失败:{error}\n💡 建议:\n1. 检查网络连接和凭证状态\n2. 尝试刷新 Token 或重新添加凭证\n3. 如问题持续,请联系技术支持") - } - } - - // Kiro OAuth 健康检查 - async fn check_kiro_health(&self, creds_path: &str, model: &str) -> Result<(), String> { - tracing::debug!("[KIRO HEALTH] 开始健康检查,凭证路径: {}", creds_path); - - // 使用 KiroProvider 加载凭证(包括 clientIdHash 文件) - let mut provider = KiroProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| self.format_user_friendly_error(&format!("加载凭证失败: {e}"), "Kiro"))?; - - let access_token = provider - .credentials - .access_token - .as_ref() - .ok_or_else(|| "凭证中缺少 access_token".to_string())?; - - let health_check_url = provider.get_health_check_url(); - - // 获取 modelId 映射 - let model_id = match model { - "claude-opus-4-5" | "claude-opus-4-5-20251101" => "claude-opus-4.5", - "claude-haiku-4-5" => "claude-haiku-4.5", - "claude-sonnet-4-5" | "claude-sonnet-4-5-20250929" => "CLAUDE_SONNET_4_5_20250929_V1_0", - "claude-sonnet-4-20250514" => "CLAUDE_SONNET_4_20250514_V1_0", - "claude-3-7-sonnet-20250219" => "CLAUDE_3_7_SONNET_20250219_V1_0", - _ => "claude-haiku-4.5", // 默认使用 haiku - }; - - tracing::debug!("[KIRO HEALTH] 健康检查 URL: {}", health_check_url); - tracing::debug!("[KIRO HEALTH] 使用模型: {} -> {}", model, model_id); - - // 构建与实际 API 调用相同格式的测试请求(参考 AIClient-2-API 实现) - let conversation_id = uuid::Uuid::new_v4().to_string(); - let mut request_body = serde_json::json!({ - "conversationState": { - "chatTriggerType": "MANUAL", - "conversationId": conversation_id, - "currentMessage": { - "userInputMessage": { - "content": "Say OK", - "modelId": model_id, - "origin": "AI_EDITOR" - } - } - } - }); - - // 如果是 social 认证方式,需要添加 profileArn - if provider.credentials.auth_method.as_deref() == Some("social") { - if let Some(profile_arn) = &provider.credentials.profile_arn { - request_body["profileArn"] = serde_json::json!(profile_arn); - } - } - - tracing::debug!("[KIRO HEALTH] 请求体已构建"); - - let response = self - .client - .post(&health_check_url) - .bearer_auth(access_token) - .header("Content-Type", "application/json") - .header("Accept", "application/json") - .header("x-amz-user-agent", "aws-sdk-js/1.0.7 KiroIDE-0.1.25") - .header("user-agent", "aws-sdk-js/1.0.7 ua/2.1 os/macos#14.0 lang/js md/nodejs#20.16.0 api/codewhispererstreaming#1.0.7 m/E KiroIDE-0.1.25") - .header("amz-sdk-invocation-id", uuid::Uuid::new_v4().to_string()) - .header("amz-sdk-request", "attempt=1; max=1") - .header("x-amzn-kiro-agent-mode", "vibe") - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| self.format_user_friendly_error(&format!("请求失败: {e}"), "Kiro"))?; - - let status = response.status(); - tracing::info!("[KIRO HEALTH] 响应状态: {}", status); - - if status.is_success() { - tracing::info!("[KIRO HEALTH] 健康检查成功"); - Ok(()) - } else { - let body_text = response.text().await.unwrap_or_default(); - tracing::warn!("[KIRO HEALTH] 健康检查失败: {} - {}", status, body_text); - let error_msg = format!("HTTP {status}: {body_text}"); - Err(self.format_user_friendly_error(&error_msg, "Kiro")) - } - } - - // Gemini OAuth 健康检查 - // 使用 cloudcode-pa.googleapis.com API(与 Gemini CLI 兼容) - // 使用 loadCodeAssist 接口进行健康检查,这是最简单可靠的方式 - async fn check_gemini_health( - &self, - creds_path: &str, - _project_id: Option<&str>, - _model: &str, - ) -> Result<(), String> { - let creds_content = - std::fs::read_to_string(creds_path).map_err(|e| format!("读取凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&creds_content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let access_token = creds["access_token"] - .as_str() - .ok_or_else(|| "凭证中缺少 access_token".to_string())?; - - // 使用 loadCodeAssist 接口进行健康检查 - // 这个接口用于获取项目信息,是最简单可靠的健康检查方式 - let url = "https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist"; - - let request_body = serde_json::json!({ - "cloudaicompanionProject": "", - "metadata": { - "ideType": "IDE_UNSPECIFIED", - "platform": "PLATFORM_UNSPECIFIED", - "pluginType": "GEMINI", - "duetProject": "" - } - }); - - let response = self - .client - .post(url) - .bearer_auth(access_token) - .header("Content-Type", "application/json") - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| format!("请求失败: {e}"))?; - - if response.status().is_success() { - Ok(()) - } else { - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - Err(format!("HTTP {status} - {body}")) - } - } - - // Antigravity OAuth 健康检查 - async fn check_antigravity_health( - &self, - creds_path: &str, - _project_id: Option<&str>, - _model: &str, - ) -> Result<(), String> { - let creds_content = - std::fs::read_to_string(creds_path).map_err(|e| format!("读取凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&creds_content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let access_token = creds["access_token"] - .as_str() - .ok_or_else(|| "凭证中缺少 access_token".to_string())?; - - // 使用 fetchAvailableModels 作为健康检查 - let url = - "https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:fetchAvailableModels"; - - let response = self - .client - .post(url) - .bearer_auth(access_token) - .header("User-Agent", "antigravity/1.11.5 windows/amd64") - .json(&serde_json::json!({})) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| format!("请求失败: {e}"))?; - - if response.status().is_success() { - Ok(()) - } else { - Err(format!("HTTP {}", response.status())) - } - } - - // OpenAI API 健康检查 - // 与 OpenAI Provider 保持一致的 URL 处理逻辑 - fn is_version_path_segment(segment: &str) -> bool { - segment.starts_with('v') - && segment.len() >= 2 - && segment[1..].chars().all(|c| c.is_ascii_digit()) - } - - fn build_openai_url_from_base(base_url: &str, endpoint: &str) -> String { - let base = base_url.trim_end_matches('/'); - let has_version = base - .rsplit('/') - .next() - .map(Self::is_version_path_segment) - .unwrap_or(false); - - if has_version { - format!("{base}/{endpoint}") - } else { - format!("{base}/v1/{endpoint}") - } - } - - fn parent_base_url(base_url: &str) -> Option { - let base = base_url.trim(); - if base.is_empty() { - return None; - } - - let mut url = reqwest::Url::parse(base) - .or_else(|_| reqwest::Url::parse(&format!("http://{base}"))) - .ok()?; - - let path = url.path().trim_end_matches('/'); - if path.is_empty() || path == "/" { - return None; - } - - let mut segments: Vec<&str> = path - .split('/') - .filter(|segment| !segment.is_empty()) - .collect(); - if segments.is_empty() { - return None; - } - segments.pop(); - - let new_path = if segments.is_empty() { - "/".to_string() - } else { - format!("/{}", segments.join("/")) - }; - - url.set_path(&new_path); - url.set_query(None); - url.set_fragment(None); - - Some(url.to_string().trim_end_matches('/').to_string()) - } - - fn push_openai_url_candidates(urls: &mut Vec, base_url: &str, endpoint: &str) { - if base_url.trim().is_empty() { - return; - } - - let primary = Self::build_openai_url_from_base(base_url, endpoint); - if !urls.iter().any(|url| url == &primary) { - urls.push(primary.clone()); - } - - if primary.contains("/v1/") { - let no_v1 = primary.replacen("/v1/", "/", 1); - if !urls.iter().any(|url| url == &no_v1) { - urls.push(no_v1); - } - } - } - - fn build_openai_health_check_urls(base_url: Option<&str>) -> Vec { - let raw_base = base_url - .map(str::trim) - .filter(|value| !value.is_empty()) - .unwrap_or("https://api.openai.com"); - let normalized_base = raw_base.trim_end_matches('/').to_string(); - - let mut urls = Vec::new(); - let mut visited = HashSet::new(); - visited.insert(normalized_base.clone()); - - Self::push_openai_url_candidates(&mut urls, &normalized_base, "chat/completions"); - - let mut current = normalized_base; - for _ in 0..6 { - let Some(parent) = Self::parent_base_url(¤t) else { - break; - }; - if !visited.insert(parent.clone()) { - break; - } - Self::push_openai_url_candidates(&mut urls, &parent, "chat/completions"); - current = parent; - } - - if urls.is_empty() { - urls.push("https://api.openai.com/v1/chat/completions".to_string()); - } - urls - } - - async fn check_openai_health( - &self, - api_key: &str, - base_url: Option<&str>, - model: &str, - ) -> Result<(), String> { - let urls = Self::build_openai_health_check_urls(base_url); - - let request_body = serde_json::json!({ - "model": model, - "messages": [{"role": "user", "content": "Say OK"}], - "max_tokens": 10 - }); - - let mut last_error: Option = None; - - for (index, url) in urls.iter().enumerate() { - tracing::debug!("[HEALTH_CHECK] OpenAI API URL: {}, model: {}", url, model); - - let response = match self - .client - .post(url) - .bearer_auth(api_key) - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - { - Ok(response) => response, - Err(error) => { - let message = format!("请求失败: {error}"); - last_error = Some(message.clone()); - if index + 1 < urls.len() { - tracing::warn!( - "[HEALTH_CHECK] OpenAI API URL {} 请求失败,继续尝试后续候选: {}", - url, - message - ); - continue; - } - return Err(message); - } - }; - - if response.status().is_success() { - return Ok(()); - } - - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - let message = format!( - "HTTP {} - {}", - status, - body.chars().take(200).collect::() - ); - last_error = Some(message.clone()); - - let can_retry_next_url = matches!( - status, - reqwest::StatusCode::NOT_FOUND | reqwest::StatusCode::METHOD_NOT_ALLOWED - ); - if can_retry_next_url && index + 1 < urls.len() { - tracing::warn!( - "[HEALTH_CHECK] OpenAI API URL {} 返回 {},尝试下一个候选 URL", - url, - status - ); - continue; - } - - return Err(message); - } - - Err(last_error.unwrap_or_else(|| "OpenAI 健康检查失败".to_string())) - } - - // Claude API 健康检查 - // 与 ClaudeCustomProvider 保持一致的 URL 处理逻辑 - async fn check_claude_health( - &self, - api_key: &str, - base_url: Option<&str>, - model: &str, - ) -> Result<(), String> { - // 与 ClaudeCustomProvider::get_base_url() 保持一致 - // base_url 应该不带 /v1,在这里拼接 - // 但为了兼容用户可能输入带 /v1 的情况,这里做智能处理 - let base = base_url.unwrap_or("https://api.anthropic.com"); - let base = base.trim_end_matches('/'); - - // 如果用户输入了带 /v1 的 URL,直接使用;否则拼接 /v1 - let url = if base.ends_with("/v1") { - format!("{base}/messages") - } else { - format!("{base}/v1/messages") - }; - - let request_body = serde_json::json!({ - "model": model, - "messages": [{"role": "user", "content": "Say OK"}], - "max_tokens": 10 - }); - - tracing::debug!("[HEALTH_CHECK] Claude API URL: {}, model: {}", url, model); - - let runtime_spec = infer_managed_runtime_spec(ApiProviderType::Anthropic, base); - let auth_value = runtime_spec - .auth_prefix - .map(|prefix| format!("{prefix} {api_key}")) - .unwrap_or_else(|| api_key.to_string()); - - let mut request = self - .client - .post(&url) - .header(runtime_spec.auth_header, auth_value) - .json(&request_body) - .timeout(self.health_check_timeout); - - if runtime_spec.protocol_family - == lime_core::database::dao::api_key_provider::ProviderProtocolFamily::Anthropic - && runtime_spec - .auth_header - .eq_ignore_ascii_case("Authorization") - { - request = request.header("x-api-key", api_key); - } - - for (name, value) in runtime_spec.extra_headers { - request = request.header(*name, *value); - } - - let response = request.send().await.map_err(|e| format!("请求失败: {e}"))?; - - if response.status().is_success() { - Ok(()) - } else { - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - Err(format!( - "HTTP {} - {}", - status, - body.chars().take(200).collect::() - )) - } - } - - // Vertex AI 健康检查 - async fn check_vertex_health( - &self, - api_key: &str, - base_url: Option<&str>, - model: &str, - ) -> Result<(), String> { - let base = base_url.unwrap_or("https://generativelanguage.googleapis.com/v1beta"); - let url = format!("{base}/models/{model}:generateContent"); - - let request_body = serde_json::json!({ - "contents": [{"role": "user", "parts": [{"text": "Say OK"}]}], - "generationConfig": {"maxOutputTokens": 10} - }); - - let response = self - .client - .post(&url) - .header("x-goog-api-key", api_key) - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| format!("请求失败: {e}"))?; - - if response.status().is_success() { - Ok(()) - } else { - Err(format!("HTTP {}", response.status())) - } - } - - // Gemini API Key 健康检查 - async fn check_gemini_api_key_health( - &self, - api_key: &str, - base_url: Option<&str>, - model: &str, - ) -> Result<(), String> { - let base = base_url.unwrap_or("https://generativelanguage.googleapis.com"); - let url = format!("{base}/v1beta/models/{model}:generateContent"); - - let request_body = serde_json::json!({ - "contents": [{"role": "user", "parts": [{"text": "Say OK"}]}], - "generationConfig": {"maxOutputTokens": 10} - }); - - let response = self - .client - .post(&url) - .header("x-goog-api-key", api_key) - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| format!("请求失败: {e}"))?; - - if response.status().is_success() { - Ok(()) - } else { - Err(format!("HTTP {}", response.status())) - } - } - - // Codex 健康检查 - // 支持 Yunyi 等代理使用 responses API 格式 - async fn check_codex_health( - &self, - creds_path: &str, - override_base_url: Option<&str>, - model: &str, - ) -> Result<(), String> { - use lime_providers::providers::codex::CodexProvider; - - let mut provider = CodexProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载 Codex 凭证失败: {e}"))?; - - let token = provider - .ensure_valid_token() - .await - .map_err(|e| format!("获取 Codex Token 失败: 配置错误,请检查凭证设置。详情:{e}"))?; - - // 优先使用 override_base_url(来自 CredentialData),其次使用凭证文件中的配置 - let base_url = override_base_url - .map(|s| s.trim()) - .filter(|s| !s.is_empty()) - .or_else(|| { - provider - .credentials - .api_base_url - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()) - }); - - // 检查是否使用 API Key 模式(如果有 api_key 且没有 refresh_token/access_token) - let is_api_key_mode = provider - .credentials - .api_key - .as_deref() - .map(|s| !s.trim().is_empty()) - .unwrap_or(false) - && provider.credentials.refresh_token.is_none(); - - // API Key 模式使用 chat/completions API,OAuth 模式使用 responses API - if is_api_key_mode && base_url.is_none() { - // API Key 直连 OpenAI:使用 chat/completions API - return self.check_openai_health(&token, None, model).await; - } - - // OAuth 模式或有自定义 base_url:使用 responses API - let url = match base_url { - Some(base) => CodexProvider::build_responses_url(base), - None => "https://api.openai.com/v1/responses".to_string(), - }; - - // Codex/Yunyi 使用 responses API 格式;云驿等代理要求 stream 必须为 true - let request_body = serde_json::json!({ - "model": model, - "input": [{ - "type": "message", - "role": "user", - "content": [{"type": "input_text", "text": "Say OK"}] - }], - "max_output_tokens": 10, - "stream": true - }); - - tracing::debug!( - "[HEALTH_CHECK] Codex responses API URL: {}, model: {}", - url, - model - ); - - let response = self - .client - .post(&url) - .bearer_auth(&token) - .header("Content-Type", "application/json") - .header("Accept", "text/event-stream") - .header("Openai-Beta", "responses=experimental") - .header("Originator", "codex_cli_rs") - .header("Session_id", uuid::Uuid::new_v4().to_string()) - .header("Conversation_id", uuid::Uuid::new_v4().to_string()) - .header( - "User-Agent", - "codex_cli_rs/0.77.0 (Lime health check; Mac OS; arm64)", - ) - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| format!("请求失败: {e}"))?; - - if response.status().is_success() { - Ok(()) - } else { - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - Err(format!( - "HTTP {} - {}", - status, - body.chars().take(200).collect::() - )) - } - } - - // Claude OAuth 健康检查 - async fn check_claude_oauth_health(&self, creds_path: &str, model: &str) -> Result<(), String> { - use lime_providers::providers::claude_oauth::ClaudeOAuthProvider; - - let mut provider = ClaudeOAuthProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载 Claude OAuth 凭证失败: {e}"))?; - - let token = provider - .ensure_valid_token() - .await - .map_err(|e| format!("获取 Claude OAuth Token 失败: {e}"))?; - - // 使用 Anthropic API 进行健康检查 - let url = "https://api.anthropic.com/v1/messages"; - let request_body = serde_json::json!({ - "model": model, - "messages": [{"role": "user", "content": "Say OK"}], - "max_tokens": 10 - }); - - let response = self - .client - .post(url) - .header("Authorization", format!("Bearer {token}")) - .header("anthropic-version", "2023-06-01") - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| format!("请求失败: {e}"))?; - - if response.status().is_success() { - Ok(()) - } else { - Err(format!("HTTP {}", response.status())) - } - } - - /// 获取 OAuth 凭证状态 - pub fn get_oauth_status( - &self, - creds_path: &str, - provider_type: &str, - ) -> Result { - let content = - std::fs::read_to_string(creds_path).map_err(|e| format!("读取凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {e}"))?; - - let has_api_key = creds - .get("apiKey") - .or_else(|| creds.get("api_key")) - .map(|v| v.as_str().is_some()) - .unwrap_or(false); - - let has_oauth_access_token = creds - .get("accessToken") - .or_else(|| creds.get("access_token")) - .map(|v| v.as_str().is_some()) - .unwrap_or(false); - - let has_access_token = has_oauth_access_token || has_api_key; - - let has_refresh_token = creds - .get("refreshToken") - .or_else(|| creds.get("refresh_token")) - .map(|v| v.as_str().is_some()) - .unwrap_or(false); - - // 检查 token 是否有效(根据 expiry_date 判断) - let (is_token_valid, expiry_info) = match provider_type { - "kiro" => { - let expires_at = creds - .get("expiresAt") - .or_else(|| creds.get("expires_at")) - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - // Kiro 没有标准的过期时间字段,假设有 access_token 就有效 - (has_access_token, expires_at) - } - - "codex" => { - // Codex: 兼容 OAuth token 或 Codex CLI 的 API Key 登录 - if has_api_key { - (true, None) - } else { - let expires_at = creds - .get("expiresAt") - .or_else(|| creds.get("expires_at")) - .or_else(|| creds.get("expired")) - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - (has_oauth_access_token, expires_at) - } - } - _ => (has_access_token, None), - }; - - Ok(OAuthStatus { - has_access_token, - has_refresh_token, - is_token_valid, - expiry_info, - creds_path: creds_path.to_string(), - }) - } - - /// 刷新 OAuth Token (Kiro) - /// - /// 使用副本文件中的凭证进行刷新,副本文件应包含完整的 client_id/client_secret。 - /// 支持多账号场景,每个副本文件完全独立。 - pub async fn refresh_kiro_token(&self, creds_path: &str) -> Result { - let mut provider = lime_providers::providers::kiro::KiroProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| self.format_user_friendly_error(&format!("加载凭证失败: {e}"), "Kiro"))?; - - // 使用副本文件中的凭证刷新 Token - provider - .refresh_token() - .await - .map_err(|e| self.format_user_friendly_error(&format!("刷新 Token 失败: {e}"), "Kiro")) - } - - /// 刷新 OAuth Token (Gemini) - pub async fn refresh_gemini_token(&self, creds_path: &str) -> Result { - let mut provider = lime_providers::providers::gemini::GeminiProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载凭证失败: {e}"))?; - provider - .refresh_token() - .await - .map_err(|e| format!("刷新 Token 失败: {e}")) - } - - /// 刷新 OAuth Token (Antigravity) - pub async fn refresh_antigravity_token(&self, creds_path: &str) -> Result { - let mut provider = lime_providers::providers::antigravity::AntigravityProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载凭证失败: {e}"))?; - provider - .refresh_token() - .await - .map_err(|e| format!("刷新 Token 失败: {e}")) - } - - /// 刷新凭证池中指定凭证的 OAuth Token - pub async fn refresh_credential_token( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result { - let cred = { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))? - }; - - match &cred.credential { - CredentialData::KiroOAuth { creds_file_path } => { - self.refresh_kiro_token(creds_file_path).await - } - CredentialData::GeminiOAuth { - creds_file_path, .. - } => self.refresh_gemini_token(creds_file_path).await, - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => self.refresh_antigravity_token(creds_file_path).await, - _ => Err("此凭证类型不支持 Token 刷新".to_string()), - } - } - - /// 获取凭证池中指定凭证的 OAuth 状态 - pub fn get_credential_oauth_status( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result { - let cred = { - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))? - }; - - let creds_path = get_oauth_creds_path(&cred.credential) - .ok_or_else(|| "此凭证类型不是 OAuth 凭证".to_string())?; - - self.get_oauth_status(&creds_path, &cred.provider_type.to_string()) - } - - /// 添加带来源的凭证 - pub fn add_credential_with_source( - &self, - db: &DbConnection, - provider_type: &str, - credential: CredentialData, - name: Option, - check_health: Option, - check_model_name: Option, - source: lime_core::models::provider_pool_model::CredentialSource, - ) -> Result { - let pt = parse_pool_provider_type(provider_type)?; - - let mut cred = ProviderCredential::new_with_source(pt, credential, source); - cred.name = name; - cred.check_health = check_health.unwrap_or(true); - cred.check_model_name = check_model_name; - - let conn = lime_core::database::lock_db(db)?; - ProviderPoolDao::insert(&conn, &cred).map_err(|e| e.to_string())?; - - Ok(cred) - } - - /// 迁移 Private 配置到凭证池 - /// - /// 从 providers 配置中读取单个凭证配置,迁移到凭证池中并标记为 Private 来源 - pub fn migrate_private_config( - &self, - db: &DbConnection, - config: &lime_core::config::Config, - ) -> Result { - use lime_core::config::expand_tilde; - use lime_core::models::provider_pool_model::CredentialSource; - - let mut result = MigrationResult::default(); - - // 迁移 Kiro 凭证 - if config.providers.kiro.enabled { - if let Some(creds_path) = &config.providers.kiro.credentials_path { - let expanded_path = expand_tilde(creds_path); - let expanded_path_str = expanded_path.to_string_lossy().to_string(); - if expanded_path.exists() { - // 检查是否已存在相同路径的凭证 - if !self.credential_exists_by_path(db, &expanded_path_str)? { - match self.add_credential_with_source( - db, - "kiro", - CredentialData::KiroOAuth { - creds_file_path: expanded_path_str.clone(), - }, - Some("Private Kiro".to_string()), - Some(true), - None, - CredentialSource::Private, - ) { - Ok(_) => result.migrated_count += 1, - Err(e) => result.errors.push(format!("Kiro: {e}")), - } - } else { - result.skipped_count += 1; - } - } - } - } - - // 迁移 Gemini 凭证 - if config.providers.gemini.enabled { - if let Some(creds_path) = &config.providers.gemini.credentials_path { - let expanded_path = expand_tilde(creds_path); - let expanded_path_str = expanded_path.to_string_lossy().to_string(); - if expanded_path.exists() { - if !self.credential_exists_by_path(db, &expanded_path_str)? { - match self.add_credential_with_source( - db, - "gemini", - CredentialData::GeminiOAuth { - creds_file_path: expanded_path_str.clone(), - project_id: config.providers.gemini.project_id.clone(), - }, - Some("Private Gemini".to_string()), - Some(true), - None, - CredentialSource::Private, - ) { - Ok(_) => result.migrated_count += 1, - Err(e) => result.errors.push(format!("Gemini: {e}")), - } - } else { - result.skipped_count += 1; - } - } - } - } - - // 迁移 OpenAI 凭证 - if config.providers.openai.enabled { - if let Some(api_key) = &config.providers.openai.api_key { - if !self.credential_exists_by_api_key(db, api_key)? { - match self.add_credential_with_source( - db, - "openai", - CredentialData::OpenAIKey { - api_key: api_key.clone(), - base_url: config.providers.openai.base_url.clone(), - }, - Some("Private OpenAI".to_string()), - Some(true), - None, - CredentialSource::Private, - ) { - Ok(_) => result.migrated_count += 1, - Err(e) => result.errors.push(format!("OpenAI: {e}")), - } - } else { - result.skipped_count += 1; - } - } - } - - // 迁移 Claude 凭证 - if config.providers.claude.enabled { - if let Some(api_key) = &config.providers.claude.api_key { - if !self.credential_exists_by_api_key(db, api_key)? { - match self.add_credential_with_source( - db, - "claude", - CredentialData::ClaudeKey { - api_key: api_key.clone(), - base_url: config.providers.claude.base_url.clone(), - }, - Some("Private Claude".to_string()), - Some(true), - None, - CredentialSource::Private, - ) { - Ok(_) => result.migrated_count += 1, - Err(e) => result.errors.push(format!("Claude: {e}")), - } - } else { - result.skipped_count += 1; - } - } - } - - Ok(result) - } - - /// 检查是否存在相同路径的凭证 - fn credential_exists_by_path(&self, db: &DbConnection, path: &str) -> Result { - let conn = lime_core::database::lock_db(db)?; - let all_creds = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - - for cred in all_creds { - if let Some(cred_path) = get_oauth_creds_path(&cred.credential) { - if cred_path == path { - return Ok(true); - } - } - } - Ok(false) - } - - /// 检查是否存在相同 API Key 的凭证 - fn credential_exists_by_api_key( - &self, - db: &DbConnection, - api_key: &str, - ) -> Result { - let conn = lime_core::database::lock_db(db)?; - let all_creds = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - - for cred in all_creds { - match &cred.credential { - CredentialData::OpenAIKey { api_key: key, .. } - | CredentialData::ClaudeKey { api_key: key, .. } => { - if key == api_key { - return Ok(true); - } - } - _ => {} - } - } - Ok(false) - } -} - -/// 迁移结果 -#[derive(Debug, Clone, Default)] -pub struct MigrationResult { - /// 成功迁移的凭证数量 - pub migrated_count: usize, - /// 跳过的凭证数量(已存在) - pub skipped_count: usize, - /// 错误信息列表 - pub errors: Vec, -} - -// ==================== 测试模块 ==================== - -#[cfg(test)] -mod tests { - use super::*; - use lime_core::database::dao::api_key_provider::ApiProviderType; - use lime_core::database::schema::create_tables; - use rusqlite::Connection; - use std::ffi::OsString; - use std::sync::{Arc, Mutex, OnceLock}; - use tempfile::tempdir; - - fn env_lock() -> &'static Mutex<()> { - static LOCK: OnceLock> = OnceLock::new(); - LOCK.get_or_init(|| Mutex::new(())) - } - - struct EnvGuard { - values: Vec<(&'static str, Option)>, - } - - impl EnvGuard { - fn set(entries: &[(&'static str, OsString)]) -> Self { - let mut values = Vec::new(); - for (key, value) in entries { - values.push((*key, std::env::var_os(key))); - std::env::set_var(key, value); - } - Self { values } - } - } - - impl Drop for EnvGuard { - fn drop(&mut self) { - for (key, previous) in self.values.drain(..) { - if let Some(value) = previous { - std::env::set_var(key, value); - } else { - std::env::remove_var(key); - } - } - } - } - - fn setup_test_db() -> DbConnection { - let conn = Connection::open_in_memory().expect("open in memory db"); - create_tables(&conn).expect("create tables"); - Arc::new(Mutex::new(conn)) - } - - // ==================== Property 3: 不健康凭证排除 ==================== - // Feature: antigravity-token-refresh, Property 3: 不健康凭证排除 - // Validates: Requirements 2.4, 3.3 - - #[test] - fn test_credential_health_info_creation() { - let info = CredentialHealthInfo { - uuid: "test-uuid".to_string(), - name: Some("Test Credential".to_string()), - provider_type: "antigravity".to_string(), - is_healthy: false, - last_error: Some("Token refresh failed".to_string()), - last_error_time: Some("2024-01-01T00:00:00Z".to_string()), - failure_count: 3, - requires_reauth: true, - }; - - assert_eq!(info.uuid, "test-uuid"); - assert!(!info.is_healthy); - assert!(info.requires_reauth); - assert_eq!(info.failure_count, 3); - } - - #[test] - fn test_selection_error_no_credentials() { - let error = SelectionError::NoCredentials; - // 验证可以序列化 - let json = serde_json::to_string(&error).unwrap(); - assert!(json.contains("NoCredentials")); - } - - #[test] - fn test_selection_error_all_unhealthy() { - let details = vec![CredentialHealthInfo { - uuid: "test-uuid".to_string(), - name: Some("Test".to_string()), - provider_type: "antigravity".to_string(), - is_healthy: false, - last_error: Some("invalid_grant".to_string()), - last_error_time: None, - failure_count: 1, - requires_reauth: true, - }]; - - let error = SelectionError::AllUnhealthy { details }; - let json = serde_json::to_string(&error).unwrap(); - assert!(json.contains("AllUnhealthy")); - assert!(json.contains("invalid_grant")); - } - - #[test] - fn test_selection_error_model_not_supported() { - let error = SelectionError::ModelNotSupported { - model: "gpt-5".to_string(), - }; - let json = serde_json::to_string(&error).unwrap(); - assert!(json.contains("ModelNotSupported")); - assert!(json.contains("gpt-5")); - } - - // ==================== Property 4: 健康状态记录完整性 ==================== - // Feature: antigravity-token-refresh, Property 4: 健康状态记录完整性 - // Validates: Requirements 3.2 - - #[test] - fn test_credential_health_info_requires_reauth_detection() { - // 测试 invalid_grant 检测 - let info_with_invalid_grant = CredentialHealthInfo { - uuid: "test".to_string(), - name: None, - provider_type: "antigravity".to_string(), - is_healthy: false, - last_error: Some("Token refresh failed: invalid_grant".to_string()), - last_error_time: Some(chrono::Utc::now().to_rfc3339()), - failure_count: 1, - requires_reauth: true, - }; - assert!(info_with_invalid_grant.requires_reauth); - - // 测试重新授权检测 - let info_with_reauth = CredentialHealthInfo { - uuid: "test".to_string(), - name: None, - provider_type: "antigravity".to_string(), - is_healthy: false, - last_error: Some("[需要重新授权] Token 已过期".to_string()), - last_error_time: Some(chrono::Utc::now().to_rfc3339()), - failure_count: 1, - requires_reauth: true, - }; - assert!(info_with_reauth.requires_reauth); - - // 测试普通错误不需要重新授权 - let info_normal_error = CredentialHealthInfo { - uuid: "test".to_string(), - name: None, - provider_type: "antigravity".to_string(), - is_healthy: false, - last_error: Some("Network error".to_string()), - last_error_time: Some(chrono::Utc::now().to_rfc3339()), - failure_count: 1, - requires_reauth: false, - }; - assert!(!info_normal_error.requires_reauth); - } - - #[test] - fn test_credential_health_info_serialization() { - let info = CredentialHealthInfo { - uuid: "test-uuid".to_string(), - name: Some("Test".to_string()), - provider_type: "antigravity".to_string(), - is_healthy: true, - last_error: None, - last_error_time: None, - failure_count: 0, - requires_reauth: false, - }; - - // 测试序列化 - let json = serde_json::to_string(&info).unwrap(); - assert!(json.contains("test-uuid")); - assert!(json.contains("antigravity")); - - // 测试反序列化 - let deserialized: CredentialHealthInfo = serde_json::from_str(&json).unwrap(); - assert_eq!(deserialized.uuid, info.uuid); - assert_eq!(deserialized.is_healthy, info.is_healthy); - } - - #[test] - fn test_api_provider_type_to_pool_type_mapping() { - assert_eq!( - api_provider_type_to_pool_type(ApiProviderType::Anthropic), - PoolProviderType::Claude - ); - assert_eq!( - api_provider_type_to_pool_type(ApiProviderType::AnthropicCompatible), - PoolProviderType::AnthropicCompatible - ); - assert_eq!( - api_provider_type_to_pool_type(ApiProviderType::Gemini), - PoolProviderType::GeminiApiKey - ); - assert_eq!( - api_provider_type_to_pool_type(ApiProviderType::Openai), - PoolProviderType::OpenAI - ); - } - - #[test] - fn test_build_openai_health_check_urls_supports_nested_base_path() { - let urls = ProviderPoolService::build_openai_health_check_urls(Some( - "http://127.0.0.1:3030/openai/v1", - )); - - assert!(urls.contains(&"http://127.0.0.1:3030/openai/v1/chat/completions".to_string())); - assert!(urls.contains(&"http://127.0.0.1:3030/openai/chat/completions".to_string())); - assert!(urls.contains(&"http://127.0.0.1:3030/v1/chat/completions".to_string())); - } - - #[test] - fn test_build_openai_health_check_urls_defaults_to_official_endpoint() { - let urls = ProviderPoolService::build_openai_health_check_urls(None); - assert_eq!(urls[0], "https://api.openai.com/v1/chat/completions"); - } - - #[test] - fn select_credential_should_auto_recover_antigravity_from_default_path() { - let _guard = env_lock().lock().expect("env lock"); - let temp = tempdir().expect("create tempdir"); - let _env = EnvGuard::set(&[("HOME", temp.path().as_os_str().to_os_string())]); - let default_path = temp.path().join(".antigravity").join("oauth_creds.json"); - std::fs::create_dir_all(default_path.parent().expect("default parent")) - .expect("create default dir"); - std::fs::write( - &default_path, - r#"{"access_token":"fallback_token","refresh_token":"refresh","project_id":"fallback-project"}"#, - ) - .expect("write default creds"); - - let db = setup_test_db(); - let service = ProviderPoolService::new(); - let inserted = service - .add_credential( - &db, - "antigravity", - CredentialData::AntigravityOAuth { - creds_file_path: temp - .path() - .join("missing.json") - .to_string_lossy() - .to_string(), - project_id: Some("fallback-project".to_string()), - }, - Some("Recovered Antigravity".to_string()), - Some(true), - None, - ) - .expect("insert credential"); - - for _ in 0..3 { - service - .mark_unhealthy( - &db, - &inserted.uuid, - Some("Failed to load credentials: No such file or directory (os error 2)"), - ) - .expect("mark unhealthy"); - } - - let selected = service - .select_credential(&db, "antigravity", Some("gemini-3-pro-image-preview")) - .expect("select credential") - .expect("recovered credential should be selectable"); - - assert_eq!(selected.uuid, inserted.uuid); - assert!(selected.is_healthy); - - let recovered = service - .get_by_uuid(&db, &inserted.uuid) - .expect("query credential") - .expect("credential should exist"); - assert!(recovered.is_healthy); - assert_eq!(recovered.error_count, 0); - assert!(recovered.last_error_message.is_none()); - } -} diff --git a/src-tauri/crates/services/src/provider_type_mapping.rs b/src-tauri/crates/services/src/provider_type_mapping.rs index cc61a8cd3..999b98478 100644 --- a/src-tauri/crates/services/src/provider_type_mapping.rs +++ b/src-tauri/crates/services/src/provider_type_mapping.rs @@ -1,7 +1,6 @@ //! Provider 类型映射与解析工具 //! -//! 统一 services 层中 PoolProviderType 与 ApiProviderType 的映射规则, -//! 避免 `provider_pool_service` 与 `api_key_provider_service` 规则漂移。 +//! 统一 API Key Provider 主路径中历史 ProviderType 与 ApiProviderType 的映射规则。 use lime_core::database::dao::api_key_provider::ApiProviderType; use lime_core::models::provider_pool_model::PoolProviderType; @@ -12,11 +11,6 @@ pub(crate) fn is_custom_provider_id(provider_type: &str) -> bool { core_is_custom_provider_id(provider_type) } -/// 解析 PoolProviderType -pub(crate) fn parse_pool_provider_type(provider_type: &str) -> Result { - provider_type.parse().map_err(|e: String| e) -} - /// 解析 PoolProviderType(失败时回退到 OpenAI) pub(crate) fn resolve_pool_provider_type_or_default(provider_type: &str) -> PoolProviderType { provider_type.parse().unwrap_or(PoolProviderType::OpenAI) @@ -47,9 +41,6 @@ pub(crate) fn pool_provider_type_to_api_type( PoolProviderType::GeminiApiKey => Some(ApiProviderType::Gemini), PoolProviderType::Vertex => Some(ApiProviderType::Vertexai), - // OAuth 类型 - 可降级到 API Key - PoolProviderType::Gemini => Some(ApiProviderType::Gemini), // Gemini OAuth → Gemini API Key - // API Key Provider 类型 - 直接映射 PoolProviderType::Anthropic => Some(ApiProviderType::Anthropic), PoolProviderType::AnthropicCompatible => Some(ApiProviderType::AnthropicCompatible), @@ -57,19 +48,19 @@ pub(crate) fn pool_provider_type_to_api_type( PoolProviderType::AwsBedrock => Some(ApiProviderType::AwsBedrock), PoolProviderType::Ollama => Some(ApiProviderType::Ollama), - // OAuth-only,无降级 - PoolProviderType::Kiro => None, - PoolProviderType::Codex => None, - PoolProviderType::ClaudeOAuth => None, - PoolProviderType::Antigravity => None, + // 已退役的凭证池 / OAuth 类型不再自动映射到 API Key Provider。 + PoolProviderType::Kiro + | PoolProviderType::Gemini + | PoolProviderType::Codex + | PoolProviderType::ClaudeOAuth => None, } } #[cfg(test)] mod tests { use super::{ - api_provider_type_to_pool_type, is_custom_provider_id, parse_pool_provider_type, - pool_provider_type_to_api_type, resolve_pool_provider_type_or_default, + api_provider_type_to_pool_type, is_custom_provider_id, pool_provider_type_to_api_type, + resolve_pool_provider_type_or_default, }; use lime_core::database::dao::api_key_provider::ApiProviderType; use lime_core::models::provider_pool_model::PoolProviderType; @@ -112,11 +103,6 @@ mod tests { #[test] fn test_pool_provider_type_parser_helpers() { - assert_eq!( - parse_pool_provider_type("openai").unwrap(), - PoolProviderType::OpenAI - ); - assert!(parse_pool_provider_type("not-exists").is_err()); assert_eq!( resolve_pool_provider_type_or_default("not-exists"), PoolProviderType::OpenAI diff --git a/src-tauri/crates/services/src/token_cache_service.rs b/src-tauri/crates/services/src/token_cache_service.rs deleted file mode 100644 index 6ceea3cd1..000000000 --- a/src-tauri/crates/services/src/token_cache_service.rs +++ /dev/null @@ -1,1063 +0,0 @@ -//! Token 缓存管理服务 -//! -//! 负责管理凭证池中 OAuth Token 的生命周期: -//! - 从源文件加载初始 Token -//! - 缓存刷新后的 Token 到数据库 - -#![allow(dead_code)] -//! - 按需刷新即将过期的 Token -//! - 处理 401/403 错误时的强制刷新 - -use crate::kiro_event_service::KiroEventService; -use chrono::Utc; -use dashmap::DashMap; -use lime_core::database::dao::provider_pool::ProviderPoolDao; -use lime_core::database::DbConnection; -use lime_core::models::provider_pool_model::{ - CachedTokenInfo, CredentialData, PoolProviderType, ProviderCredential, -}; -use lime_providers::providers::gemini::GeminiProvider; -use lime_providers::providers::kiro::KiroProvider; -use std::sync::Arc; -use tokio::sync::Mutex; - -/// Token 刷新错误类型 -#[derive(Debug, Clone, PartialEq)] -pub enum RefreshErrorType { - /// Token被截断或格式问题 - TokenTruncated, - /// Token格式异常(长度过短等) - TokenFormat, - /// 网络连接问题 - Network, - /// 服务不可用 - ServiceUnavailable, - /// 认证失败(401, 403等) - AuthenticationFailed, - /// 未知错误 - Unknown, -} - -/// Token 刷新错误分类结果 -#[derive(Debug, Clone)] -pub struct RefreshErrorClassification { - /// 错误类型 - pub error_type: RefreshErrorType, - /// 错误描述 - pub error_description: String, - /// 建议重试次数 - pub retry_count: u32, - /// 是否支持降级策略 - pub supports_fallback: bool, - /// 是否应该自动禁用凭证(永久性错误) - pub should_disable_credential: bool, -} - -/// Token 缓存服务 -pub struct TokenCacheService { - /// 每凭证一把锁,防止并发刷新 - locks: DashMap>>, -} - -impl Default for TokenCacheService { - fn default() -> Self { - Self::new() - } -} - -impl TokenCacheService { - pub fn new() -> Self { - Self { - locks: DashMap::new(), - } - } - - /// 获取有效的 Token(核心方法) - /// - /// 1. 检查数据库缓存是否有效 - /// 2. 如果缓存有效且未过期,直接返回 - /// 3. 如果缓存无效或即将过期,执行刷新 - /// 4. 如果刷新失败(如 refreshToken 被截断),尝试使用源文件中的 accessToken - pub async fn get_valid_token(&self, db: &DbConnection, uuid: &str) -> Result { - // 首先检查缓存 - let cached = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_token_cache(&conn, uuid).map_err(|e| e.to_string())? - }; - - // 缓存有效且未即将过期,直接返回 - if let Some(ref cache) = cached { - if cache.is_valid() && !cache.is_expiring_soon() { - if let Some(token) = &cache.access_token { - tracing::debug!( - "[TOKEN_CACHE] Using cached token for {}, expires at {:?}", - &uuid[..8], - cache.expiry_time - ); - return Ok(token.clone()); - } - } - } - - // 需要刷新(无缓存、已过期或即将过期) - match self.refresh_and_cache(db, uuid, false).await { - Ok(token) => Ok(token), - Err(refresh_error) => { - // 增强的错误处理机制 - 智能检测各种token问题 - let error_classification = self.classify_refresh_error(&refresh_error); - - tracing::warn!( - "[TOKEN_CACHE] Token 刷新失败,错误类型: {:?}, 详情: {}", - error_classification.error_type, - &refresh_error - ); - - match error_classification.error_type { - RefreshErrorType::TokenTruncated | RefreshErrorType::TokenFormat => { - tracing::warn!( - "[TOKEN_CACHE] 检测到 token 问题,尝试使用源文件中的 accessToken: {}", - &uuid[..8] - ); - - // 获取凭证信息 - let credential = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))? - }; - - // 尝试从源文件读取 accessToken - match self.read_token_from_source(&credential).await { - Ok(token_info) => { - if let Some(token) = token_info.access_token { - tracing::info!( - "[TOKEN_CACHE] 使用源文件中的 accessToken 作为降级方案: {}", - &uuid[..8] - ); - - // 缓存这个 token 但标记为降级状态 - let cache_info = CachedTokenInfo { - access_token: Some(token.clone()), - refresh_token: token_info.refresh_token, - expiry_time: None, // 无法确定过期时间 - last_refresh: Some(Utc::now()), - refresh_error_count: error_classification.retry_count, - last_refresh_error: Some(format!( - "{}(降级使用源文件 accessToken): {}", - error_classification.error_description, refresh_error - )), - }; - - // 缓存到数据库 - if let Ok(conn) = db.lock() { - let _ = ProviderPoolDao::update_token_cache( - &conn, - uuid, - &cache_info, - ); - } - - return Ok(token); - } - } - Err(e) => { - tracing::error!( - "[TOKEN_CACHE] 降级策略失败,无法从源文件读取 accessToken: {}", - e - ); - } - } - } - RefreshErrorType::Network | RefreshErrorType::ServiceUnavailable => { - tracing::warn!("[TOKEN_CACHE] 网络/服务问题,建议稍后重试: {}", &uuid[..8]); - // 可以考虑使用缓存中的过期 token 作为临时方案 - if let Some(cache) = cached { - if let Some(token) = cache.access_token { - tracing::info!( - "[TOKEN_CACHE] 网络问题时使用过期缓存 token: {}", - &uuid[..8] - ); - return Ok(token); - } - } - } - RefreshErrorType::AuthenticationFailed => { - tracing::error!("[TOKEN_CACHE] 认证失败,凭证可能已被撤销: {}", &uuid[..8]); - // 认证失败通常需要用户重新授权,不进行降级 - } - RefreshErrorType::Unknown => { - tracing::warn!("[TOKEN_CACHE] 未知错误类型,使用默认处理: {}", &uuid[..8]); - } - } - - // 更新错误计数 - if let Ok(conn) = db.lock() { - let _ = ProviderPoolDao::record_token_refresh_error( - &conn, - uuid, - &format!( - "{}(分类: {:?}): {}", - error_classification.error_description, - error_classification.error_type, - refresh_error - ), - ); - } - - // 返回分类后的错误信息 - Err(format!( - "{}: {}", - error_classification.error_description, refresh_error - )) - } - } - } - - /// 刷新 Token 并缓存到数据库(带事件发送) - /// - /// - force: 是否强制刷新(忽略缓存状态) - /// - kiro_event_service: 可选的事件服务,用于发送 Kiro 凭证刷新事件 - /// - /// 优化说明:添加了随机延迟机制,避免多个凭证同时刷新造成请求过于集中 - pub async fn refresh_and_cache_with_events( - &self, - db: &DbConnection, - uuid: &str, - force: bool, - kiro_event_service: Option>, - ) -> Result { - // 添加随机延迟,避免多个凭证同时刷新 - // 基于凭证UUID生成0-30秒的随机延迟,确保同一凭证的延迟时间一致但不同凭证间分散 - if !force { - let delay_ms = self.calculate_refresh_delay(uuid); - if delay_ms > 0 { - tracing::debug!( - "[TOKEN_CACHE] Adding {}ms delay before refreshing token for {}", - delay_ms, - &uuid[..8] - ); - tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await; - } - } - - // 获取该凭证的锁 - let lock = self - .locks - .entry(uuid.to_string()) - .or_insert_with(|| Arc::new(Mutex::new(()))) - .clone(); - - let _guard = lock.lock().await; - - // 双重检查:可能其他线程已完成刷新 - if !force { - let cached = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_token_cache(&conn, uuid).map_err(|e| e.to_string())? - }; - - if let Some(cache) = cached { - if cache.is_valid() && !cache.is_expiring_soon() { - if let Some(token) = cache.access_token { - tracing::debug!( - "[TOKEN_CACHE] Double-check: another thread refreshed for {}", - &uuid[..8] - ); - return Ok(token); - } - } - } - } - - // 获取凭证信息 - let credential = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))? - }; - - tracing::info!( - "[TOKEN_CACHE] Refreshing token for {} ({})", - &uuid[..8], - credential.provider_type - ); - - // 发送刷新开始事件(仅针对 Kiro 凭证) - if let Some(event_service) = &kiro_event_service { - if credential.provider_type == PoolProviderType::Kiro { - event_service - .emit_refresh_started(uuid.to_string(), credential.name.clone()) - .await; - } - } - - // 执行刷新 - match self.do_refresh(&credential).await { - Ok(token_info) => { - // 缓存到数据库 - { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::update_token_cache(&conn, uuid, &token_info) - .map_err(|e| e.to_string())?; - } - - let token = token_info - .access_token - .ok_or_else(|| "Refresh succeeded but no access_token".to_string())?; - - tracing::info!( - "[TOKEN_CACHE] Token refreshed and cached for {}, expires at {:?}", - &uuid[..8], - token_info.expiry_time - ); - - // 发送刷新成功事件(仅针对 Kiro 凭证) - if let Some(event_service) = &kiro_event_service { - if credential.provider_type == PoolProviderType::Kiro { - event_service - .emit_refresh_success( - uuid.to_string(), - credential.name.clone(), - token_info - .expiry_time - .unwrap_or_else(|| Utc::now() + chrono::Duration::hours(1)), - "IdC".to_string(), // 默认为IdC认证 - "BuilderId".to_string(), - "us-east-1".to_string(), - ) - .await; - } - } - - Ok(token) - } - Err(e) => { - // 记录刷新错误 - { - let conn = db.lock().map_err(|e| e.to_string())?; - let _ = ProviderPoolDao::record_token_refresh_error(&conn, uuid, &e); - } - - tracing::error!( - "[TOKEN_CACHE] Token refresh failed for {}: {}", - &uuid[..8], - e - ); - - // 分析错误并决定是否自动禁用凭证 - let error_classification = self.classify_refresh_error(&e); - - // 如果是永久性错误,自动禁用凭证 - if error_classification.should_disable_credential { - let disable_result = { - let conn = db.lock().map_err(|e| e.to_string())?; - // 简化禁用逻辑:直接在数据库中标记为禁用 - let sql = "UPDATE credentials SET is_disabled = true WHERE uuid = ?"; - conn.execute(sql, [&uuid]).map_err(|e| e.to_string()) - }; - - match disable_result { - Ok(_) => { - tracing::warn!( - "[TOKEN_CACHE] Auto-disabled credential {} due to permanent failure: {:?}", - &uuid[..8], - error_classification.error_type - ); - - // 发送凭证禁用事件 - if let Some(event_service) = &kiro_event_service { - if credential.provider_type == PoolProviderType::Kiro { - // 发送状态更新事件 - event_service - .emit_credential_status_update( - uuid.to_string(), - false, // is_healthy - true, // is_disabled - credential.error_count + 1, - Some(0.0), // health_score降为0 - None, - ) - .await; - - // 发送自动禁用事件 - event_service - .emit_credential_auto_disabled( - uuid.to_string(), - credential.name.clone(), - error_classification.error_description.clone(), - format!("{:?}", error_classification.error_type), - ) - .await; - } - } - } - Err(disable_err) => { - tracing::error!( - "[TOKEN_CACHE] Failed to auto-disable credential {}: {}", - &uuid[..8], - disable_err - ); - } - } - } - - // 发送刷新失败事件(仅针对 Kiro 凭证) - if let Some(event_service) = &kiro_event_service { - if credential.provider_type == PoolProviderType::Kiro { - event_service - .emit_refresh_failed( - uuid.to_string(), - credential.name.clone(), - e.clone(), - Some(format!("{:?}", error_classification.error_type)), - ) - .await; - } - } - - Err(e) - } - } - } - - /// 执行实际的 Token 刷新 - async fn do_refresh(&self, credential: &ProviderCredential) -> Result { - match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - self.refresh_kiro(creds_file_path).await - } - CredentialData::GeminiOAuth { - creds_file_path, .. - } => self.refresh_gemini(creds_file_path).await, - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => self.refresh_antigravity(creds_file_path).await, - CredentialData::OpenAIKey { api_key, .. } => { - // API Key 不需要刷新,直接返回 - Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, // 永不过期 - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::ClaudeKey { api_key, .. } => { - // API Key 不需要刷新,直接返回 - Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, // 永不过期 - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::VertexKey { api_key, .. } => { - // API Key 不需要刷新,直接返回 - Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, // 永不过期 - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::GeminiApiKey { api_key, .. } => { - // API Key 不需要刷新,直接返回 - Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, // 永不过期 - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::CodexOAuth { - creds_file_path, .. - } => self.refresh_codex(creds_file_path).await, - CredentialData::ClaudeOAuth { creds_file_path } => { - self.refresh_claude_oauth(creds_file_path).await - } - CredentialData::AnthropicKey { api_key, .. } => { - // API Key 不需要刷新,直接返回 - Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, // 永不过期 - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - } - } - - /// 刷新 Kiro Token - async fn refresh_kiro(&self, creds_path: &str) -> Result { - let mut provider = KiroProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载 Kiro 凭证失败: {e}"))?; - - let token = provider - .refresh_token() - .await - .map_err(|e| format!("刷新 Kiro Token 失败: {e}"))?; - - // Kiro token 通常 1 小时过期,我们假设 50 分钟 - let expiry_time = Utc::now() + chrono::Duration::minutes(50); - - Ok(CachedTokenInfo { - access_token: Some(token), - refresh_token: provider.credentials.refresh_token.clone(), - expiry_time: Some(expiry_time), - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - - /// 刷新 Gemini Token - async fn refresh_gemini(&self, creds_path: &str) -> Result { - let mut provider = GeminiProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载 Gemini 凭证失败: {e}"))?; - - let token = provider - .refresh_token() - .await - .map_err(|e| format!("刷新 Gemini Token 失败: {e}"))?; - - // Gemini token 通常 1 小时过期 - let expiry_time = provider - .credentials - .expiry_date - .map(|ts| chrono::DateTime::from_timestamp(ts, 0).unwrap_or_default()) - .unwrap_or_else(|| Utc::now() + chrono::Duration::minutes(50)); - - Ok(CachedTokenInfo { - access_token: Some(token), - refresh_token: provider.credentials.refresh_token.clone(), - expiry_time: Some(expiry_time), - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - - /// 刷新 Antigravity Token - async fn refresh_antigravity(&self, creds_path: &str) -> Result { - use lime_providers::providers::antigravity::AntigravityProvider; - - let mut provider = AntigravityProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载 Antigravity 凭证失败: {e}"))?; - - let token = provider - .refresh_token() - .await - .map_err(|e| format!("刷新 Antigravity Token 失败: {e}"))?; - - // Antigravity token 通常 1 小时过期 - let expiry_time = provider - .credentials - .expiry_date - .map(|ts| chrono::DateTime::from_timestamp_millis(ts).unwrap_or_default()) - .unwrap_or_else(|| Utc::now() + chrono::Duration::minutes(50)); - - Ok(CachedTokenInfo { - access_token: Some(token), - refresh_token: provider.credentials.refresh_token.clone(), - expiry_time: Some(expiry_time), - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - - /// 刷新 Codex Token - async fn refresh_codex(&self, creds_path: &str) -> Result { - use lime_providers::providers::codex::CodexProvider; - - let mut provider = CodexProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载 Codex 凭证失败: {e}"))?; - - let token = provider - .refresh_token_with_retry(3) - .await - .map_err(|e| format!("刷新 Codex Token 失败: {e}"))?; - - // 解析过期时间 - let expiry_time = provider - .credentials - .expires_at - .as_ref() - .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|| Utc::now() + chrono::Duration::minutes(50)); - - Ok(CachedTokenInfo { - access_token: Some(token), - refresh_token: provider.credentials.refresh_token.clone(), - expiry_time: Some(expiry_time), - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - - /// 刷新 Claude OAuth Token - async fn refresh_claude_oauth(&self, creds_path: &str) -> Result { - use lime_providers::providers::claude_oauth::ClaudeOAuthProvider; - - let mut provider = ClaudeOAuthProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| format!("加载 Claude OAuth 凭证失败: {e}"))?; - - let token = provider - .refresh_token_with_retry(3) - .await - .map_err(|e| format!("刷新 Claude OAuth Token 失败: {e}"))?; - - // 解析过期时间 - let expiry_time = provider - .credentials - .expire - .as_ref() - .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|| Utc::now() + chrono::Duration::minutes(50)); - - Ok(CachedTokenInfo { - access_token: Some(token), - refresh_token: provider.credentials.refresh_token.clone(), - expiry_time: Some(expiry_time), - last_refresh: Some(Utc::now()), - refresh_error_count: 0, - last_refresh_error: None, - }) - } - - /// 从源文件加载初始 Token(首次使用时) - pub async fn load_initial_token( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result { - let credential = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_by_uuid(&conn, uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("Credential not found: {uuid}"))? - }; - - // 尝试从源文件读取 token - let token_info = self.read_token_from_source(&credential).await?; - - // 缓存到数据库 - { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::update_token_cache(&conn, uuid, &token_info) - .map_err(|e| e.to_string())?; - } - - token_info - .access_token - .ok_or_else(|| "源文件中没有 access_token".to_string()) - } - - /// 从源文件读取 Token(不刷新) - async fn read_token_from_source( - &self, - credential: &ProviderCredential, - ) -> Result { - match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - let content = tokio::fs::read_to_string(creds_file_path) - .await - .map_err(|e| format!("读取 Kiro 凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let access_token = creds["accessToken"] - .as_str() - .or_else(|| creds["access_token"].as_str()) - .map(|s| s.to_string()); - let refresh_token = creds["refreshToken"] - .as_str() - .or_else(|| creds["refresh_token"].as_str()) - .map(|s| s.to_string()); - - Ok(CachedTokenInfo { - access_token, - refresh_token, - expiry_time: None, // Kiro 源文件通常没有过期时间 - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::GeminiOAuth { - creds_file_path, .. - } => { - let content = tokio::fs::read_to_string(creds_file_path) - .await - .map_err(|e| format!("读取 Gemini 凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let access_token = creds["access_token"].as_str().map(|s| s.to_string()); - let refresh_token = creds["refresh_token"].as_str().map(|s| s.to_string()); - let expiry_time = creds["expiry_date"] - .as_i64() - .and_then(|ts| chrono::DateTime::from_timestamp(ts, 0)); - - Ok(CachedTokenInfo { - access_token, - refresh_token, - expiry_time, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => { - let content = tokio::fs::read_to_string(creds_file_path) - .await - .map_err(|e| format!("读取 Antigravity 凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let access_token = creds["access_token"].as_str().map(|s| s.to_string()); - let refresh_token = creds["refresh_token"].as_str().map(|s| s.to_string()); - let expiry_time = creds["expiry_date"] - .as_i64() - .and_then(|ts| chrono::DateTime::from_timestamp(ts, 0)); - - Ok(CachedTokenInfo { - access_token, - refresh_token, - expiry_time, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::OpenAIKey { api_key, .. } => Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }), - CredentialData::ClaudeKey { api_key, .. } => Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }), - CredentialData::VertexKey { api_key, .. } => Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }), - CredentialData::GeminiApiKey { api_key, .. } => Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }), - CredentialData::CodexOAuth { - creds_file_path, .. - } => { - let content = tokio::fs::read_to_string(creds_file_path) - .await - .map_err(|e| format!("读取 Codex 凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let access_token = creds["access_token"].as_str().map(|s| s.to_string()); - let refresh_token = creds["refresh_token"].as_str().map(|s| s.to_string()); - let expiry_time = creds["expired"] - .as_str() - .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) - .map(|dt| dt.with_timezone(&Utc)); - - Ok(CachedTokenInfo { - access_token, - refresh_token, - expiry_time, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::ClaudeOAuth { creds_file_path } => { - let content = tokio::fs::read_to_string(creds_file_path) - .await - .map_err(|e| format!("读取 Claude OAuth 凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证失败: {e}"))?; - - let access_token = creds["access_token"].as_str().map(|s| s.to_string()); - let refresh_token = creds["refresh_token"].as_str().map(|s| s.to_string()); - let expiry_time = creds["expire"] - .as_str() - .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) - .map(|dt| dt.with_timezone(&Utc)); - - Ok(CachedTokenInfo { - access_token, - refresh_token, - expiry_time, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }) - } - CredentialData::AnthropicKey { api_key, .. } => Ok(CachedTokenInfo { - access_token: Some(api_key.clone()), - refresh_token: None, - expiry_time: None, - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }), - } - } - - /// 清除凭证的 Token 缓存 - pub fn clear_cache(&self, db: &DbConnection, uuid: &str) -> Result<(), String> { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::clear_token_cache(&conn, uuid).map_err(|e| e.to_string()) - } - - /// 检查凭证类型是否支持 Token 刷新 - pub fn supports_refresh(provider_type: PoolProviderType) -> bool { - matches!( - provider_type, - PoolProviderType::Kiro | PoolProviderType::Gemini - ) - } - - /// 获取凭证的缓存状态 - pub fn get_cache_status( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result, String> { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_token_cache(&conn, uuid).map_err(|e| e.to_string()) - } - - /// 计算刷新延迟时间(毫秒) - /// - /// 基于凭证UUID生成确定性但分散的延迟时间,避免多个凭证同时刷新 - /// 延迟范围:0-30秒,确保同一凭证每次的延迟一致 - fn calculate_refresh_delay(&self, uuid: &str) -> u64 { - use std::collections::hash_map::DefaultHasher; - use std::hash::{Hash, Hasher}; - - // 使用凭证UUID作为种子生成确定性的延迟 - let mut hasher = DefaultHasher::new(); - uuid.hash(&mut hasher); - let hash_value = hasher.finish(); - - // 生成0-30秒的延迟(转换为毫秒) - hash_value % 30000 - } - - /// 智能错误分类方法 - /// - /// 基于错误信息智能识别错误类型,提供针对性的处理建议 - fn classify_refresh_error(&self, error_message: &str) -> RefreshErrorClassification { - let error_lower = error_message.to_lowercase(); - - // Token 被截断问题检测(最严重的问题,优先检查) - if error_lower.contains("截断") || error_lower.contains("truncated") { - return RefreshErrorClassification { - error_type: RefreshErrorType::TokenTruncated, - error_description: "Token 被截断,需检查配置文件".to_string(), - retry_count: 1, - supports_fallback: true, - should_disable_credential: true, // 永久性问题,自动禁用 - }; - } - - // Token 格式问题检测 - if error_lower.contains("格式异常") - || error_lower.contains("长度过短") - || error_lower.contains("format") - || error_lower.contains("invalid") - || error_lower.contains("malformed") - { - return RefreshErrorClassification { - error_type: RefreshErrorType::TokenFormat, - error_description: "Token 格式异常,需重新配置".to_string(), - retry_count: 1, - supports_fallback: true, - should_disable_credential: true, // 配置问题,自动禁用 - }; - } - - // 认证失败检测 - if error_lower.contains("unauthorized") - || error_lower.contains("forbidden") - || error_lower.contains("401") - || error_lower.contains("403") - || error_lower.contains("认证失败") - || error_lower.contains("invalid_grant") - || error_lower.contains("access_denied") - || error_lower.contains("refresh_token") - || error_lower.contains("expired") - { - return RefreshErrorClassification { - error_type: RefreshErrorType::AuthenticationFailed, - error_description: "认证失败,凭证已过期或无效".to_string(), - retry_count: 0, // 不建议重试 - supports_fallback: false, - should_disable_credential: true, // 认证失效,自动禁用 - }; - } - - // 网络问题检测 - if error_lower.contains("network") - || error_lower.contains("connection") - || error_lower.contains("timeout") - || error_lower.contains("dns") - || error_lower.contains("connect") - || error_lower.contains("网络") - || error_lower.contains("连接") - { - return RefreshErrorClassification { - error_type: RefreshErrorType::Network, - error_description: "网络连接问题".to_string(), - retry_count: 3, - supports_fallback: true, - should_disable_credential: false, // 临时问题,不禁用 - }; - } - - // 服务不可用检测 - if error_lower.contains("service unavailable") - || error_lower.contains("502") - || error_lower.contains("503") - || error_lower.contains("504") - || error_lower.contains("internal server error") - || error_lower.contains("服务不可用") - { - return RefreshErrorClassification { - error_type: RefreshErrorType::ServiceUnavailable, - error_description: "服务暂时不可用".to_string(), - retry_count: 2, - supports_fallback: true, - should_disable_credential: false, // 临时问题,不禁用 - }; - } - - // 未知错误(默认分类) - RefreshErrorClassification { - error_type: RefreshErrorType::Unknown, - error_description: "未知错误".to_string(), - retry_count: 1, - supports_fallback: false, - should_disable_credential: false, // 未知错误暂不自动禁用 - } - } - - /// 刷新 Token 并缓存到数据库(兼容版本) - /// - /// - force: 是否强制刷新(忽略缓存状态) - /// - /// 此方法保持与旧版本的兼容性,不发送任何事件。 - /// 如需事件支持,请使用 refresh_and_cache_with_events 方法。 - pub async fn refresh_and_cache( - &self, - db: &DbConnection, - uuid: &str, - force: bool, - ) -> Result { - self.refresh_and_cache_with_events(db, uuid, force, None) - .await - } - - /// 检查 Token 是否即将过期并提前刷新(需求 4.4) - /// - /// 在流式请求前调用此方法,检查 Token 是否在指定分钟数内过期。 - /// 如果即将过期,则提前刷新 Token。 - /// - /// # 参数 - /// - `db`: 数据库连接 - /// - `uuid`: 凭证 UUID - /// - `minutes`: 检查的时间阈值(分钟),默认 10 分钟 - /// - /// # 返回 - /// - `Ok(token)`: 有效的 Token(可能是刷新后的新 Token) - /// - `Err(error)`: 获取或刷新 Token 失败 - pub async fn ensure_token_valid_for_streaming( - &self, - db: &DbConnection, - uuid: &str, - minutes: i64, - ) -> Result { - // 首先检查缓存 - let cached = { - let conn = db.lock().map_err(|e| e.to_string())?; - lime_core::database::dao::provider_pool::ProviderPoolDao::get_token_cache(&conn, uuid) - .map_err(|e| e.to_string())? - }; - - // 检查是否需要提前刷新(使用指定的分钟数阈值) - if let Some(ref cache) = cached { - if cache.is_valid() && !cache.is_expiring_within_minutes(minutes) { - if let Some(token) = &cache.access_token { - tracing::debug!( - "[TOKEN_CACHE] Token valid for streaming ({}min threshold) for {}, expires at {:?}", - minutes, - &uuid[..8], - cache.expiry_time - ); - return Ok(token.clone()); - } - } - - // Token 即将过期(在指定分钟数内),提前刷新 - if cache.is_expiring_within_minutes(minutes) { - tracing::info!( - "[TOKEN_CACHE] Token expiring within {}min for {}, proactively refreshing", - minutes, - &uuid[..8] - ); - } - } - - // 需要刷新(无缓存、已过期或即将过期) - self.refresh_and_cache(db, uuid, false).await - } -} diff --git a/src-tauri/crates/services/src/usage_service.rs b/src-tauri/crates/services/src/usage_service.rs deleted file mode 100644 index 692d3efd5..000000000 --- a/src-tauri/crates/services/src/usage_service.rs +++ /dev/null @@ -1,915 +0,0 @@ -//! Usage Service - Kiro 用量查询服务 -//! -//! 通过调用 AWS Q 的 getUsageLimits API 获取用户的用量信息。 -//! 参考 Kir-Manager 项目的 usage/usage.go 实现。 - -use reqwest::header::{HeaderMap, HeaderValue, USER_AGENT}; -use serde::{Deserialize, Serialize}; -use std::error::Error; -use uuid::Uuid; - -// ============================================================================ -// 常量定义 -// ============================================================================ - -/// API 端点 -pub const USAGE_LIMITS_URL: &str = "https://q.us-east-1.amazonaws.com/getUsageLimits"; - -/// Query 参数 -pub const ORIGIN_PARAM: &str = "AI_EDITOR"; -pub const RESOURCE_TYPE_PARAM: &str = "AGENTIC_REQUEST"; - -/// HTTP 请求超时(秒) -pub const HTTP_TIMEOUT_SECS: u64 = 10; - -// ============================================================================ -// API Response 数据模型 -// ============================================================================ - -/// API 响应结构 -#[derive(Debug, Clone, Deserialize, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct UsageLimitsResponse { - pub subscription_info: SubscriptionInfo, - pub usage_breakdown_list: Vec, -} - -/// 订阅信息结构 -#[derive(Debug, Clone, Deserialize, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct SubscriptionInfo { - pub subscription_title: String, - #[serde(rename = "type")] - pub subscription_type: String, -} - -/// 用量明细结构 -#[derive(Debug, Clone, Deserialize, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct UsageBreakdown { - pub usage_limit_with_precision: f64, - pub current_usage_with_precision: f64, - pub display_name: String, - #[serde(default)] - pub free_trial_info: Option, - #[serde(default)] - pub bonuses: Option>, -} - -/// 免费试用信息 -#[derive(Debug, Clone, Deserialize, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct FreeTrialInfo { - pub usage_limit_with_precision: f64, - pub current_usage_with_precision: f64, - pub free_trial_status: String, -} - -/// 奖励额度 -#[derive(Debug, Clone, Deserialize, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct Bonus { - pub bonus_code: String, - pub usage_limit: f64, - pub current_usage: f64, - pub status: String, -} - -// ============================================================================ -// 计算结果数据模型 -// ============================================================================ - -/// 计算后的用量信息 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -#[serde(rename_all = "camelCase")] -pub struct UsageInfo { - /// 订阅类型名称 - pub subscription_title: String, - /// 总额度 - pub usage_limit: f64, - /// 已使用 - pub current_usage: f64, - /// 余额 = usage_limit - current_usage - pub balance: f64, - /// 余额低于 20% - pub is_low_balance: bool, -} - -impl UsageInfo { - /// 创建空的 UsageInfo - pub fn empty() -> Self { - Self::default() - } -} - -// ============================================================================ -// URL 构造函数 -// ============================================================================ - -/// 构造 API 请求 URL -/// -/// **Property 4: Social Auth URL Construction** -/// **Property 5: IdC Auth URL Construction** -/// **Validates: Requirements 2.1, 2.2** -/// -/// - Social 认证: 包含 profileArn 参数 -/// - IdC 认证: 不包含 profileArn 参数 -pub fn build_usage_api_url( - auth_method: &str, - profile_arn: Option<&str>, -) -> Result> { - let mut url = url::Url::parse(USAGE_LIMITS_URL)?; - - { - let mut query = url.query_pairs_mut(); - query.append_pair("origin", ORIGIN_PARAM); - query.append_pair("resourceType", RESOURCE_TYPE_PARAM); - - // Property 4: Social Auth URL Construction - // 只有 social 类型才加入 profileArn - if auth_method == "social" { - match profile_arn { - Some(arn) if !arn.is_empty() => { - query.append_pair("profileArn", arn); - } - _ => { - return Err("social auth requires profileArn".into()); - } - } - } - // Property 5: IdC Auth URL Construction - // IdC 类型不包含 profileArn - } - - Ok(url.to_string()) -} - -// ============================================================================ -// 请求头构造函数 -// ============================================================================ - -/// 构造 API 请求头 -/// -/// **Property 6: User-Agent Header Format** -/// **Validates: Requirements 4.1, 4.2, 4.3** -/// -/// Headers: -/// - User-Agent: aws-sdk-js/1.0.0 ua/2.1 os/{os} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{version}-{machineId} -/// - x-amz-user-agent: aws-sdk-js/1.0.0 KiroIDE-{version}-{machineId} -/// - amz-sdk-invocation-id: UUID -/// - amz-sdk-request: attempt=1; max=1 -pub fn build_request_headers( - access_token: &str, - kiro_version: &str, - machine_id: &str, -) -> Result> { - let mut headers = HeaderMap::new(); - - // Authorization header - let auth_value = format!("Bearer {access_token}"); - headers.insert("Authorization", HeaderValue::from_str(&auth_value)?); - - // User-Agent header - // 格式: aws-sdk-js/1.0.0 ua/2.1 os/{os} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{version}-{machineId} - let os_name = std::env::consts::OS; - let user_agent = format!( - "aws-sdk-js/1.0.0 ua/2.1 os/{os_name} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{kiro_version}-{machine_id}" - ); - headers.insert(USER_AGENT, HeaderValue::from_str(&user_agent)?); - - // x-amz-user-agent header - let x_amz_user_agent = format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}"); - headers.insert( - "x-amz-user-agent", - HeaderValue::from_str(&x_amz_user_agent)?, - ); - - // amz-sdk-invocation-id: 每次请求随机生成 UUID - headers.insert( - "amz-sdk-invocation-id", - HeaderValue::from_str(&Uuid::new_v4().to_string())?, - ); - - // amz-sdk-request header - headers.insert( - "amz-sdk-request", - HeaderValue::from_static("attempt=1; max=1"), - ); - - // Connection header - headers.insert("Connection", HeaderValue::from_static("close")); - - Ok(headers) -} - -/// 构造 User-Agent 字符串(用于测试) -#[cfg(test)] -pub fn build_user_agent(kiro_version: &str, machine_id: &str) -> String { - let os_name = std::env::consts::OS; - format!( - "aws-sdk-js/1.0.0 ua/2.1 os/{os_name} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{kiro_version}-{machine_id}" - ) -} - -/// 构造 x-amz-user-agent 字符串(用于测试) -#[cfg(test)] -pub fn build_x_amz_user_agent(kiro_version: &str, machine_id: &str) -> String { - format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}") -} - -// ============================================================================ -// API 调用函数 -// ============================================================================ - -/// 调用 AWS Q getUsageLimits API 获取用量信息 -/// -/// **Validates: Requirements 1.1, 4.4** -/// -/// # Arguments -/// * `access_token` - Bearer token -/// * `auth_method` - 认证方式 ("social" 或 "idc") -/// * `profile_arn` - Social 认证需要的 profileArn -/// * `machine_id` - 设备 ID (SHA256 哈希) -/// * `kiro_version` - Kiro 版本号 -/// -/// # Returns -/// * `Ok(UsageInfo)` - 成功时返回计算后的用量信息 -/// * `Err` - 失败时返回错误 -pub async fn get_usage_limits( - access_token: &str, - auth_method: &str, - profile_arn: Option<&str>, - machine_id: &str, - kiro_version: &str, -) -> Result> { - // 验证参数 - if access_token.is_empty() { - return Err("invalid token: missing accessToken".into()); - } - - if machine_id.is_empty() { - return Err("invalid machineID: empty".into()); - } - - // 构造 URL - let url = build_usage_api_url(auth_method, profile_arn)?; - - // 构造请求头 - let headers = build_request_headers(access_token, kiro_version, machine_id)?; - - // 创建 HTTP 客户端(带超时) - let client = reqwest::Client::builder() - .timeout(std::time::Duration::from_secs(HTTP_TIMEOUT_SECS)) - .build()?; - - // 发送请求 - let response = client.get(&url).headers(headers).send().await?; - - // 检查状态码 - let status = response.status(); - if !status.is_success() { - let body = response.text().await.unwrap_or_default(); - return Err(format!("API request failed with status {status}: {body}").into()); - } - - // 解析响应 - let usage_response: UsageLimitsResponse = response.json().await?; - - // 计算余额并返回 - Ok(calculate_balance(&usage_response)) -} - -/// 安全地调用 API 获取用量信息 -/// -/// **Property 3: Error Handling Graceful Degradation** -/// **Validates: Requirements 1.4** -/// -/// 当发生任何错误时,返回空的 UsageInfo 而非 panic -pub async fn get_usage_limits_safe( - access_token: &str, - auth_method: &str, - profile_arn: Option<&str>, - machine_id: &str, - kiro_version: &str, -) -> UsageInfo { - match get_usage_limits( - access_token, - auth_method, - profile_arn, - machine_id, - kiro_version, - ) - .await - { - Ok(info) => info, - Err(e) => { - tracing::warn!("Failed to get usage limits: {}", e); - UsageInfo::empty() - } - } -} - -// ============================================================================ -// 余额计算函数 -// ============================================================================ - -/// 低余额阈值 (20%) -pub const LOW_BALANCE_THRESHOLD: f64 = 0.2; - -/// 从 API 响应计算余额 -/// -/// **Property 1: Balance Calculation Correctness** -/// **Validates: Requirements 1.2** -/// -/// 计算逻辑: -/// - 总额度 = Σ(usage_limit_with_precision + free_trial_info?.usage_limit_with_precision + Σ(bonuses[].usage_limit)) -/// - 总使用 = Σ(current_usage_with_precision + free_trial_info?.current_usage_with_precision + Σ(bonuses[].current_usage)) -/// - 余额 = 总额度 - 总使用 -pub fn calculate_balance(response: &UsageLimitsResponse) -> UsageInfo { - calculate_balance_with_threshold(response, LOW_BALANCE_THRESHOLD) -} - -/// 从 API 响应计算余额(使用指定阈值) -/// -/// threshold: 低余额阈值(0.0 ~ 1.0),例如 0.2 表示余额低于 20% 时为低余额 -pub fn calculate_balance_with_threshold( - response: &UsageLimitsResponse, - threshold: f64, -) -> UsageInfo { - let mut total_usage_limit = 0.0; - let mut total_current_usage = 0.0; - - for breakdown in &response.usage_breakdown_list { - // 基本额度 - total_usage_limit += breakdown.usage_limit_with_precision; - total_current_usage += breakdown.current_usage_with_precision; - - // 免费试用额度(如果存在) - if let Some(ref free_trial) = breakdown.free_trial_info { - total_usage_limit += free_trial.usage_limit_with_precision; - total_current_usage += free_trial.current_usage_with_precision; - } - - // 奖励额度(如果存在) - if let Some(ref bonuses) = breakdown.bonuses { - for bonus in bonuses { - total_usage_limit += bonus.usage_limit; - total_current_usage += bonus.current_usage; - } - } - } - - let balance = total_usage_limit - total_current_usage; - - // Property 2: Low Balance Detection - // Validates: Requirements 1.3 - // is_low_balance = (balance / total_usage_limit) < threshold - let is_low_balance = if total_usage_limit > 0.0 { - (balance / total_usage_limit) < threshold - } else { - false - }; - - UsageInfo { - subscription_title: response.subscription_info.subscription_title.clone(), - usage_limit: total_usage_limit, - current_usage: total_current_usage, - balance, - is_low_balance, - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use proptest::prelude::*; - use urlencoding; - - // ======================================================================== - // Arbitrary 生成器 - // ======================================================================== - - /// 生成有效的 Bonus - fn arb_bonus() -> impl Strategy { - ( - "[a-zA-Z0-9]{4,10}", // bonus_code - 0.0..1000.0f64, // usage_limit - 0.0..1000.0f64, // current_usage - prop_oneof!["ACTIVE", "EXPIRED", "PENDING"], - ) - .prop_map(|(bonus_code, usage_limit, current_usage, status)| Bonus { - bonus_code, - usage_limit, - current_usage, - status: status.to_string(), - }) - } - - /// 生成有效的 FreeTrialInfo - fn arb_free_trial_info() -> impl Strategy { - ( - 0.0..1000.0f64, // usage_limit_with_precision - 0.0..1000.0f64, // current_usage_with_precision - prop_oneof!["ACTIVE", "EXPIRED"], - ) - .prop_map(|(usage_limit, current_usage, status)| FreeTrialInfo { - usage_limit_with_precision: usage_limit, - current_usage_with_precision: current_usage, - free_trial_status: status.to_string(), - }) - } - - /// 生成有效的 UsageBreakdown - fn arb_usage_breakdown() -> impl Strategy { - ( - 0.0..1000.0f64, // usage_limit_with_precision - 0.0..1000.0f64, // current_usage_with_precision - "[a-zA-Z ]{5,20}", // display_name - prop::option::of(arb_free_trial_info()), // free_trial_info - prop::option::of(prop::collection::vec(arb_bonus(), 0..3)), // bonuses - ) - .prop_map( - |(usage_limit, current_usage, display_name, free_trial_info, bonuses)| { - UsageBreakdown { - usage_limit_with_precision: usage_limit, - current_usage_with_precision: current_usage, - display_name, - free_trial_info, - bonuses, - } - }, - ) - } - - /// 生成有效的 SubscriptionInfo - fn arb_subscription_info() -> impl Strategy { - ( - prop_oneof!["Free Tier", "Pro", "Enterprise"], - prop_oneof!["FREE", "PAID", "TRIAL"], - ) - .prop_map(|(title, sub_type)| SubscriptionInfo { - subscription_title: title.to_string(), - subscription_type: sub_type.to_string(), - }) - } - - /// 生成有效的 UsageLimitsResponse - fn arb_usage_limits_response() -> impl Strategy { - ( - arb_subscription_info(), - prop::collection::vec(arb_usage_breakdown(), 1..5), - ) - .prop_map( - |(subscription_info, usage_breakdown_list)| UsageLimitsResponse { - subscription_info, - usage_breakdown_list, - }, - ) - } - - // ======================================================================== - // Property 1: Balance Calculation Correctness - // **Feature: kiro-usage-api, Property 1: Balance Calculation Correctness** - // **Validates: Requirements 1.2** - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 1: 余额计算正确性 - /// - /// *For any* UsageLimitsResponse with valid usage breakdown data, - /// the calculated balance SHALL equal (total_usage_limit - total_current_usage), - /// where totals include base amounts, free trial amounts, and bonus amounts. - #[test] - fn prop_balance_calculation_correctness(response in arb_usage_limits_response()) { - let result = calculate_balance(&response); - - // 手动计算期望值 - let mut expected_limit = 0.0; - let mut expected_usage = 0.0; - - for breakdown in &response.usage_breakdown_list { - expected_limit += breakdown.usage_limit_with_precision; - expected_usage += breakdown.current_usage_with_precision; - - if let Some(ref ft) = breakdown.free_trial_info { - expected_limit += ft.usage_limit_with_precision; - expected_usage += ft.current_usage_with_precision; - } - - if let Some(ref bonuses) = breakdown.bonuses { - for bonus in bonuses { - expected_limit += bonus.usage_limit; - expected_usage += bonus.current_usage; - } - } - } - - let expected_balance = expected_limit - expected_usage; - - // 使用近似比较(浮点数精度问题) - let epsilon = 1e-10; - prop_assert!((result.usage_limit - expected_limit).abs() < epsilon, - "usage_limit mismatch: got {}, expected {}", result.usage_limit, expected_limit); - prop_assert!((result.current_usage - expected_usage).abs() < epsilon, - "current_usage mismatch: got {}, expected {}", result.current_usage, expected_usage); - prop_assert!((result.balance - expected_balance).abs() < epsilon, - "balance mismatch: got {}, expected {}", result.balance, expected_balance); - } - } - - // ======================================================================== - // Property 2: Low Balance Detection - // **Feature: kiro-usage-api, Property 2: Low Balance Detection** - // **Validates: Requirements 1.3** - // ======================================================================== - - /// 生成有效的 UsageInfo(直接生成,用于测试低余额检测) - fn arb_usage_info() -> impl Strategy { - // 生成 usage_limit 和 balance,确保 balance <= usage_limit - (0.01..1000.0f64).prop_flat_map(|usage_limit| { - // balance 可以是 0 到 usage_limit 之间的任意值 - (Just(usage_limit), 0.0..=usage_limit) - }) - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 2: 低余额检测 - /// - /// *For any* UsageInfo where balance/usage_limit < 0.2 (and usage_limit > 0), - /// is_low_balance SHALL be true; otherwise it SHALL be false. - #[test] - fn prop_low_balance_detection((usage_limit, balance) in arb_usage_info()) { - // 构造一个简单的响应来测试低余额检测 - let current_usage = usage_limit - balance; - let response = UsageLimitsResponse { - subscription_info: SubscriptionInfo { - subscription_title: "Test".to_string(), - subscription_type: "FREE".to_string(), - }, - usage_breakdown_list: vec![UsageBreakdown { - usage_limit_with_precision: usage_limit, - current_usage_with_precision: current_usage, - display_name: "Test".to_string(), - free_trial_info: None, - bonuses: None, - }], - }; - - let result = calculate_balance(&response); - - // 计算期望的 is_low_balance - let ratio = balance / usage_limit; - let expected_low_balance = ratio < LOW_BALANCE_THRESHOLD; - - prop_assert_eq!( - result.is_low_balance, - expected_low_balance, - "is_low_balance mismatch: got {}, expected {} (ratio: {}, threshold: {})", - result.is_low_balance, - expected_low_balance, - ratio, - LOW_BALANCE_THRESHOLD - ); - } - - /// Property 2 边界情况: 当 usage_limit 为 0 时,is_low_balance 应为 false - #[test] - fn prop_low_balance_zero_limit(current_usage in 0.0..100.0f64) { - let response = UsageLimitsResponse { - subscription_info: SubscriptionInfo { - subscription_title: "Test".to_string(), - subscription_type: "FREE".to_string(), - }, - usage_breakdown_list: vec![UsageBreakdown { - usage_limit_with_precision: 0.0, - current_usage_with_precision: current_usage, - display_name: "Test".to_string(), - free_trial_info: None, - bonuses: None, - }], - }; - - let result = calculate_balance(&response); - - // 当 usage_limit 为 0 时,is_low_balance 应为 false(避免除零) - prop_assert!(!result.is_low_balance, - "is_low_balance should be false when usage_limit is 0"); - } - } - - // ======================================================================== - // Property 4: Social Auth URL Construction - // **Feature: kiro-usage-api, Property 4: Social Auth URL Construction** - // **Validates: Requirements 2.1** - // ======================================================================== - - /// 生成有效的 profileArn - fn arb_profile_arn() -> impl Strategy { - "[a-zA-Z0-9:/-]{10,50}".prop_map(|s| format!("arn:aws:iam::{s}")) - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 4: Social Auth URL 构造 - /// - /// *For any* request with auth_method="social" and a valid profile_arn, - /// the request URL SHALL contain `profileArn={profile_arn}` as a query parameter. - #[test] - fn prop_social_auth_url_contains_profile_arn(profile_arn in arb_profile_arn()) { - let url = build_usage_api_url("social", Some(&profile_arn)).unwrap(); - - // URL 应该包含 profileArn 参数 - prop_assert!(url.contains("profileArn="), - "Social auth URL should contain profileArn parameter, got: {}", url); - - // URL 应该包含编码后的 profile_arn 值 - let encoded_arn = urlencoding::encode(&profile_arn); - prop_assert!(url.contains(&encoded_arn.to_string()), - "Social auth URL should contain encoded profileArn value '{}', got: {}", encoded_arn, url); - - // URL 应该包含基本参数 - prop_assert!(url.contains("origin=AI_EDITOR"), - "URL should contain origin parameter, got: {}", url); - prop_assert!(url.contains("resourceType=AGENTIC_REQUEST"), - "URL should contain resourceType parameter, got: {}", url); - } - - /// Property 4 边界情况: Social auth 缺少 profileArn 应返回错误 - #[test] - fn prop_social_auth_requires_profile_arn(_dummy in 0..10i32) { - // 测试 None - let result = build_usage_api_url("social", None); - prop_assert!(result.is_err(), "Social auth without profileArn should fail"); - - // 测试空字符串 - let result = build_usage_api_url("social", Some("")); - prop_assert!(result.is_err(), "Social auth with empty profileArn should fail"); - } - } - - // ======================================================================== - // Property 5: IdC Auth URL Construction - // **Feature: kiro-usage-api, Property 5: IdC Auth URL Construction** - // **Validates: Requirements 2.2** - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 5: IdC Auth URL 构造 - /// - /// *For any* request with auth_method="idc", - /// the request URL SHALL NOT contain `profileArn` as a query parameter. - #[test] - fn prop_idc_auth_url_no_profile_arn(profile_arn in arb_profile_arn()) { - // 即使提供了 profile_arn,IdC 认证也不应该包含它 - let url = build_usage_api_url("idc", Some(&profile_arn)).unwrap(); - - // URL 不应该包含 profileArn 参数 - prop_assert!(!url.contains("profileArn"), - "IdC auth URL should NOT contain profileArn parameter, got: {}", url); - - // URL 应该包含基本参数 - prop_assert!(url.contains("origin=AI_EDITOR"), - "URL should contain origin parameter, got: {}", url); - prop_assert!(url.contains("resourceType=AGENTIC_REQUEST"), - "URL should contain resourceType parameter, got: {}", url); - } - - /// Property 5: IdC auth 不需要 profileArn - #[test] - fn prop_idc_auth_works_without_profile_arn(_dummy in 0..10i32) { - // IdC 认证不需要 profileArn - let result = build_usage_api_url("idc", None); - prop_assert!(result.is_ok(), "IdC auth without profileArn should succeed"); - - let url = result.unwrap(); - prop_assert!(!url.contains("profileArn"), - "IdC auth URL should NOT contain profileArn parameter, got: {}", url); - } - } - - // ======================================================================== - // Property 3: Error Handling Graceful Degradation - // **Feature: kiro-usage-api, Property 3: Error Handling Graceful Degradation** - // **Validates: Requirements 1.4** - // ======================================================================== - - /// 生成各种错误输入场景 - #[derive(Debug, Clone)] - enum ErrorScenario { - EmptyToken, - EmptyMachineId, - MissingProfileArn, - EmptyProfileArn, - InvalidToken, - } - - /// 生成错误场景的策略 - fn arb_error_scenario() -> impl Strategy { - prop_oneof![ - Just(ErrorScenario::EmptyToken), - Just(ErrorScenario::EmptyMachineId), - Just(ErrorScenario::MissingProfileArn), - Just(ErrorScenario::EmptyProfileArn), - Just(ErrorScenario::InvalidToken), - ] - } - - /// 生成随机的有效 token(用于非空 token 场景) - fn arb_valid_token() -> impl Strategy { - "[a-zA-Z0-9]{20,50}".prop_map(|s| s) - } - - /// 生成随机的有效 machine_id(用于非空 machine_id 场景) - fn arb_valid_machine_id() -> impl Strategy { - "[a-f0-9]{32,64}".prop_map(|s| s) - } - - /// 生成随机的有效 kiro_version - fn arb_valid_kiro_version() -> impl Strategy { - "[0-9]{1,2}\\.[0-9]{1,2}\\.[0-9]{1,3}".prop_map(|s| s) - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 3: 错误处理优雅降级 - /// - /// *For any* error condition (network error, invalid response, missing token), - /// the safe wrapper function SHALL return an empty UsageInfo with zero values - /// instead of panicking. - #[test] - fn prop_error_handling_graceful_degradation( - scenario in arb_error_scenario(), - valid_token in arb_valid_token(), - valid_machine_id in arb_valid_machine_id(), - valid_version in arb_valid_kiro_version(), - ) { - // 创建 tokio runtime 来运行异步代码 - let rt = tokio::runtime::Runtime::new().unwrap(); - - let result = rt.block_on(async { - match scenario { - ErrorScenario::EmptyToken => { - // 空 token 应该导致错误 - get_usage_limits_safe( - "", - "social", - Some("arn:aws:test"), - &valid_machine_id, - &valid_version - ).await - } - ErrorScenario::EmptyMachineId => { - // 空 machine_id 应该导致错误 - get_usage_limits_safe( - &valid_token, - "social", - Some("arn:aws:test"), - "", - &valid_version - ).await - } - ErrorScenario::MissingProfileArn => { - // social 认证缺少 profileArn 应该导致错误 - get_usage_limits_safe( - &valid_token, - "social", - None, - &valid_machine_id, - &valid_version - ).await - } - ErrorScenario::EmptyProfileArn => { - // social 认证空 profileArn 应该导致错误 - get_usage_limits_safe( - &valid_token, - "social", - Some(""), - &valid_machine_id, - &valid_version - ).await - } - ErrorScenario::InvalidToken => { - // 无效 token 会导致网络错误(401/403) - get_usage_limits_safe( - "invalid_token_that_will_fail", - "idc", - None, - &valid_machine_id, - &valid_version - ).await - } - } - }); - - // 无论什么错误场景,safe 函数都应该返回空的 UsageInfo - prop_assert_eq!(result.usage_limit, 0.0, - "Error scenario {:?} should return zero usage_limit", scenario); - prop_assert_eq!(result.current_usage, 0.0, - "Error scenario {:?} should return zero current_usage", scenario); - prop_assert_eq!(result.balance, 0.0, - "Error scenario {:?} should return zero balance", scenario); - prop_assert!(!result.is_low_balance, - "Error scenario {:?} should return false is_low_balance", scenario); - prop_assert!(result.subscription_title.is_empty(), - "Error scenario {:?} should return empty subscription_title", scenario); - } - } - - // 保留原有的单元测试作为补充(快速验证) - #[tokio::test] - async fn test_error_handling_empty_token() { - let result = - get_usage_limits_safe("", "social", Some("arn:aws:test"), "machine123", "1.0.0").await; - assert_eq!(result.usage_limit, 0.0); - assert_eq!(result.balance, 0.0); - } - - #[tokio::test] - async fn test_error_handling_empty_machine_id() { - let result = - get_usage_limits_safe("token123", "social", Some("arn:aws:test"), "", "1.0.0").await; - assert_eq!(result.usage_limit, 0.0); - assert_eq!(result.balance, 0.0); - } - - // ======================================================================== - // Property 6: User-Agent Header Format - // **Feature: kiro-usage-api, Property 6: User-Agent Header Format** - // **Validates: Requirements 4.1, 4.2** - // ======================================================================== - - /// 生成有效的 Kiro 版本号 - fn arb_kiro_version() -> impl Strategy { - "[0-9]{1,2}\\.[0-9]{1,2}\\.[0-9]{1,3}".prop_map(|s| s) - } - - /// 生成有效的 Machine ID (SHA256 哈希) - fn arb_machine_id() -> impl Strategy { - "[a-f0-9]{64}".prop_map(|s| s) - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// Property 6: User-Agent 头格式 - /// - /// *For any* kiro_version and machine_id strings, - /// the User-Agent header SHALL match the format: - /// `aws-sdk-js/1.0.0 ua/2.1 os/{os} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{version}-{machineId}` - #[test] - fn prop_user_agent_header_format( - kiro_version in arb_kiro_version(), - machine_id in arb_machine_id() - ) { - let user_agent = build_user_agent(&kiro_version, &machine_id); - - // 验证格式各部分 - prop_assert!(user_agent.starts_with("aws-sdk-js/1.0.0 ua/2.1 os/"), - "User-Agent should start with 'aws-sdk-js/1.0.0 ua/2.1 os/', got: {}", user_agent); - - prop_assert!(user_agent.contains("lang/rust"), - "User-Agent should contain 'lang/rust', got: {}", user_agent); - - prop_assert!(user_agent.contains("api/codewhispererruntime#1.0.0"), - "User-Agent should contain 'api/codewhispererruntime#1.0.0', got: {}", user_agent); - - prop_assert!(user_agent.contains("m/N,E"), - "User-Agent should contain 'm/N,E', got: {}", user_agent); - - // 验证包含 KiroIDE-{version}-{machineId} - let kiro_suffix = format!("KiroIDE-{kiro_version}-{machine_id}"); - prop_assert!(user_agent.ends_with(&kiro_suffix), - "User-Agent should end with '{}', got: {}", kiro_suffix, user_agent); - } - - /// Property 6: x-amz-user-agent 头格式 - /// - /// *For any* kiro_version and machine_id strings, - /// the x-amz-user-agent header SHALL match the format: - /// `aws-sdk-js/1.0.0 KiroIDE-{version}-{machineId}` - #[test] - fn prop_x_amz_user_agent_header_format( - kiro_version in arb_kiro_version(), - machine_id in arb_machine_id() - ) { - let x_amz_user_agent = build_x_amz_user_agent(&kiro_version, &machine_id); - - // 验证格式 - let expected = format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}"); - prop_assert_eq!(x_amz_user_agent, expected, - "x-amz-user-agent format mismatch"); - } - } -} diff --git a/src-tauri/crates/services/src/voice_command_service.rs b/src-tauri/crates/services/src/voice_command_service.rs index 58c953057..f3fd6c137 100644 --- a/src-tauri/crates/services/src/voice_command_service.rs +++ b/src-tauri/crates/services/src/voice_command_service.rs @@ -87,7 +87,7 @@ pub async fn transcribe_audio( ); } } - return Err("未配置语音识别服务。请在设置 → 凭证池 → ASR 中添加讯飞、百度或 OpenAI Whisper 凭证。".to_string()); + return Err("未配置语音识别服务。请在设置 → Agent → 语音中添加讯飞、百度或 OpenAI Whisper 凭证。".to_string()); } Err(error) => { tracing::error!("[语音识别] 获取默认凭证失败: {}", error); diff --git a/src-tauri/crates/skills/src/lime_llm_provider.rs b/src-tauri/crates/skills/src/lime_llm_provider.rs index 64ab49516..003f4c55a 100644 --- a/src-tauri/crates/skills/src/lime_llm_provider.rs +++ b/src-tauri/crates/skills/src/lime_llm_provider.rs @@ -1,6 +1,6 @@ //! Lime LLM Provider 实现 //! -//! 使用 ProviderPoolService 选择凭证并调用 LLM API。 +//! 使用 API Key Provider 选择凭证并调用 LLM API。 //! trait 定义(LlmProvider, SkillError)已迁移到 lime-skills crate。 use std::sync::Arc; @@ -14,20 +14,17 @@ use lime_core::models::anthropic::AnthropicMessagesRequest; use lime_core::models::provider_pool_model::PoolProviderType; use lime_core::models::provider_pool_model::{CredentialData, ProviderCredential}; use lime_providers::providers::claude_custom::{ClaudeCustomProvider, PromptCacheMode}; -use lime_providers::providers::kiro::KiroProvider; +use lime_providers::providers::gemini::{GeminiApiKeyCredential, GeminiApiKeyProvider}; use lime_providers::providers::openai_custom::OpenAICustomProvider; use lime_services::api_key_provider_service::ApiKeyProviderService; -use lime_services::provider_pool_service::ProviderPoolService; use crate::{LlmProvider, SkillError}; /// Lime LLM Provider /// -/// 使用 ProviderPoolService 选择凭证并调用 LLM API。 +/// 使用 API Key Provider 选择凭证并调用 LLM API。 /// 实现 aster-rust 定义的 LlmProvider trait。 pub struct LimeLlmProvider { - /// 凭证池服务 - pool_service: Arc, /// API Key Provider 服务(用于智能降级) api_key_service: Arc, /// 数据库连接 @@ -40,16 +37,10 @@ impl LimeLlmProvider { /// 创建新的 LimeLlmProvider 实例 /// /// # Arguments - /// * `pool_service` - 凭证池服务 /// * `api_key_service` - API Key 服务 /// * `db` - 数据库连接 - pub fn new( - pool_service: Arc, - api_key_service: Arc, - db: DbConnection, - ) -> Self { + pub fn new(api_key_service: Arc, db: DbConnection) -> Self { Self { - pool_service, api_key_service, db, preferred_provider: None, @@ -59,18 +50,15 @@ impl LimeLlmProvider { /// 创建带有偏好 Provider 的实例 /// /// # Arguments - /// * `pool_service` - 凭证池服务 /// * `api_key_service` - API Key 服务 /// * `db` - 数据库连接 /// * `preferred_provider` - 偏好的 Provider 类型 pub fn with_preferred_provider( - pool_service: Arc, api_key_service: Arc, db: DbConnection, preferred_provider: String, ) -> Self { Self { - pool_service, api_key_service, db, preferred_provider: Some(preferred_provider), @@ -99,10 +87,8 @@ impl LimeLlmProvider { match provider.to_lowercase().as_str() { "openai" | "gpt" => Some(PoolProviderType::OpenAI), "anthropic" | "claude" => Some(PoolProviderType::Claude), - "gemini" | "google" => Some(PoolProviderType::Gemini), - "kiro" | "codewhisperer" => Some(PoolProviderType::Kiro), + "gemini" | "google" => Some(PoolProviderType::GeminiApiKey), "vertex" => Some(PoolProviderType::Vertex), - "codex" => Some(PoolProviderType::Codex), _ => None, } } @@ -125,10 +111,6 @@ impl LimeLlmProvider { model: &str, ) -> Result { match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - self.call_kiro_api(creds_file_path, system_prompt, user_message, model) - .await - } CredentialData::ClaudeKey { api_key, base_url } => { self.call_claude_api( api_key, @@ -178,77 +160,84 @@ impl LimeLlmProvider { ) .await } + CredentialData::GeminiApiKey { + api_key, + base_url, + excluded_models, + } => { + self.call_gemini_api( + &credential.uuid, + api_key, + base_url.as_deref(), + excluded_models, + system_prompt, + user_message, + model, + ) + .await + } _ => Err(SkillError::ProviderError(format!( - "不支持的凭证类型: {:?}", + "凭证池/OAuth 凭证已退役,当前只支持 API Key Provider 凭证: {:?}", credential.provider_type ))), } } - /// 调用 Kiro API - async fn call_kiro_api( + /// 调用 Gemini API Key Provider + async fn call_gemini_api( &self, - creds_file_path: &str, + credential_id: &str, + api_key: &str, + base_url: Option<&str>, + excluded_models: &[String], system_prompt: &str, user_message: &str, model: &str, ) -> Result { - use lime_core::models::anthropic::AnthropicMessage; - use lime_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; - use lime_providers::providers::traits::CredentialProvider; - use lime_server_utils::parse_cw_response; + let credential = + GeminiApiKeyCredential::new(credential_id.to_string(), api_key.to_string()) + .with_base_url(base_url.map(ToString::to_string)) + .with_excluded_models(excluded_models.to_vec()); - let mut kiro = KiroProvider::new(); - kiro.load_credentials_from_path(creds_file_path) - .await - .map_err(|e| SkillError::ProviderError(format!("加载 Kiro 凭证失败: {}", e)))?; - - // 确保 Token 有效 - if !kiro.is_token_valid() || kiro.is_token_expiring_soon() { - kiro.refresh_token() - .await - .map_err(|e| SkillError::ProviderError(format!("刷新 Token 失败: {}", e)))?; - } - - // 构建 Anthropic 请求 - let request = AnthropicMessagesRequest { - model: model.to_string(), - max_tokens: Some(4096), - system: Some(serde_json::Value::String(system_prompt.to_string())), - messages: vec![AnthropicMessage { - role: "user".to_string(), - content: serde_json::Value::String(user_message.to_string()), - }], - stream: false, - temperature: None, - tools: None, - tool_choice: None, - }; - - // 转换为 OpenAI 格式并调用 - let openai_request = convert_anthropic_to_openai(&request); - let resp = kiro - .call_api(&openai_request) - .await - .map_err(|e| SkillError::ProviderError(format!("Kiro API 调用失败: {}", e)))?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); + if !credential.supports_model(model) { return Err(SkillError::ProviderError(format!( - "Kiro API 返回错误: status={}, body={}", - status, body + "Gemini API Key 凭证不支持模型: {}", + model ))); } - let bytes = resp - .bytes() - .await - .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; - let body = String::from_utf8_lossy(&bytes).to_string(); - let parsed = parse_cw_response(&body); + let provider = GeminiApiKeyProvider::new(); + let prompt = if system_prompt.trim().is_empty() { + user_message.to_string() + } else { + format!("{}\n\n{}", system_prompt, user_message) + }; + let body = serde_json::json!({ + "contents": [ + { + "role": "user", + "parts": [{ "text": prompt }] + } + ], + "generationConfig": { + "maxOutputTokens": 4096 + } + }); - Ok(parsed.content) + let json = provider + .generate_content(&credential, model, &body) + .await + .map_err(|e| SkillError::ProviderError(format!("Gemini API 调用失败: {}", e)))?; + + let content = json["candidates"] + .as_array() + .and_then(|arr| arr.first()) + .and_then(|candidate| candidate["content"]["parts"].as_array()) + .and_then(|parts| parts.first()) + .and_then(|part| part["text"].as_str()) + .unwrap_or(""); + + Ok(content.to_string()) } /// 调用 Claude API @@ -400,13 +389,13 @@ impl LlmProvider for LimeLlmProvider { /// 调用 LLM 进行对话 /// /// # 实现说明 - /// 1. 使用 ProviderPoolService.select_credential_with_fallback() 选择凭证 + /// 1. 使用 API Key Provider 选择凭证 /// 2. 如果指定了 preferred_provider,优先选择该类型的凭证 /// 3. 如果指定了 model,传递给底层 provider /// 4. 如果没有可用凭证,返回 ProviderError /// /// # Requirements - /// - 1.2: 使用 ProviderPoolService 选择可用凭证 + /// - 1.2: 使用 API Key Provider 选择可用凭证 /// - 1.3: 优先选择指定 provider 类型的凭证 /// - 1.4: 将 model 参数传递给底层 provider /// - 1.5: 没有可用凭证时返回 ProviderError @@ -428,15 +417,13 @@ impl LlmProvider for LimeLlmProvider { model_name ); - // 使用 ProviderPoolService 选择凭证(Requirements 1.2, 1.3) + // 使用 API Key Provider 选择凭证(Requirements 1.2, 1.3) let credential = self - .pool_service - .select_credential_with_fallback( + .api_key_service + .select_credential_for_provider( &self.db, - &self.api_key_service, provider_type, - Some(model_name), - None, // provider_id_hint + Some(provider_type), None, // client_type ) .await @@ -463,16 +450,14 @@ impl LlmProvider for LimeLlmProvider { // 记录使用情况 match &result { Ok(_) => { - let _ = self.pool_service.record_usage(&self.db, &credential.uuid); - let _ = - self.pool_service - .mark_healthy(&self.db, &credential.uuid, Some(model_name)); + if let Some(api_key_id) = credential.uuid.strip_prefix("fallback-") { + let _ = self.api_key_service.record_usage(&self.db, api_key_id); + } } Err(e) => { - let _ = self.pool_service.mark_unhealthy( - &self.db, - &credential.uuid, - Some(&e.to_string()), + tracing::debug!( + "[LimeLlmProvider] 调用失败,API Key Provider 不写回凭证池健康状态: {}", + e ); } } @@ -521,23 +506,11 @@ mod tests { fn test_map_skill_provider_gemini() { assert_eq!( LimeLlmProvider::map_skill_provider_to_pool_type("gemini"), - Some(PoolProviderType::Gemini) + Some(PoolProviderType::GeminiApiKey) ); assert_eq!( LimeLlmProvider::map_skill_provider_to_pool_type("google"), - Some(PoolProviderType::Gemini) - ); - } - - #[test] - fn test_map_skill_provider_kiro() { - assert_eq!( - LimeLlmProvider::map_skill_provider_to_pool_type("kiro"), - Some(PoolProviderType::Kiro) - ); - assert_eq!( - LimeLlmProvider::map_skill_provider_to_pool_type("codewhisperer"), - Some(PoolProviderType::Kiro) + Some(PoolProviderType::GeminiApiKey) ); } diff --git a/src-tauri/crates/websocket/src/handler.rs b/src-tauri/crates/websocket/src/handler.rs index 9d5333441..0d458c66d 100644 --- a/src-tauri/crates/websocket/src/handler.rs +++ b/src-tauri/crates/websocket/src/handler.rs @@ -274,17 +274,6 @@ async fn handle_message( ))) } WsMessage::Error(_) => None, - WsMessage::SubscribeKiroEvents => { - // TODO: 实现Kiro事件订阅 - None - } - WsMessage::UnsubscribeKiroEvents => { - // TODO: 实现Kiro事件取消订阅 - None - } - WsMessage::KiroCredentialEvent(_) => Some(WsMessage::Error(WsError::invalid_message( - "KiroCredentialEvent messages are server-to-client only", - ))), } } diff --git a/src-tauri/crates/websocket/src/lib.rs b/src-tauri/crates/websocket/src/lib.rs index 48fd4d33e..b591c897a 100644 --- a/src-tauri/crates/websocket/src/lib.rs +++ b/src-tauri/crates/websocket/src/lib.rs @@ -20,8 +20,8 @@ pub mod stream; pub use handlers::RpcHandler; pub use lime_core::websocket::types; pub use lime_core::websocket::{ - KiroTokenInfo, WsApiRequest, WsApiResponse, WsConfig, WsConnection, WsEndpoint, WsError, - WsKiroEvent, WsMessage, WsStats, WsStatsSnapshot, WsStreamChunk, WsStreamEnd, + WsApiRequest, WsApiResponse, WsConfig, WsConnection, WsEndpoint, WsError, WsMessage, WsStats, + WsStatsSnapshot, WsStreamChunk, WsStreamEnd, }; pub use processor::MessageProcessor; pub use protocol::{GatewayRpcRequest, GatewayRpcResponse, RpcError, RpcMethod}; diff --git a/src-tauri/proptest-regressions/config/tests.txt b/src-tauri/proptest-regressions/config/tests.txt new file mode 100644 index 000000000..2cb6886f5 --- /dev/null +++ b/src-tauri/proptest-regressions/config/tests.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc d1c5121cbf2443b7b1b0aa0c62b5a79e69056478a04b9451979d3e7d94845c65 # shrinks to provider = "antigravity" diff --git a/src-tauri/resources/models/aliases/antigravity.json b/src-tauri/resources/models/aliases/antigravity.json deleted file mode 100644 index da39eef14..000000000 --- a/src-tauri/resources/models/aliases/antigravity.json +++ /dev/null @@ -1,59 +0,0 @@ -{ - "$schema": "../schema/alias.schema.json", - "provider": "antigravity", - "description": "Antigravity 中转服务的模型别名映射", - "models": [ - "gemini-2.5-computer-use-preview-10-2025", - "gemini-3-pro-image-preview", - "gemini-3-pro-preview", - "gemini-3-flash-preview", - "gemini-claude-sonnet-4-5", - "gemini-claude-sonnet-4-5-thinking", - "gemini-claude-opus-4-5-thinking" - ], - "aliases": { - "gemini-2.5-computer-use-preview-10-2025": { - "actual": "gemini-2.5-computer-use", - "internal_name": "rev19-uic3-1p", - "provider": "google", - "description": "Gemini 2.5 Computer Use 预览版" - }, - "gemini-3-pro-image-preview": { - "actual": "gemini-3-pro-image", - "internal_name": "gemini-3-pro-image", - "provider": "google", - "description": "Gemini 3 Pro 图像版" - }, - "gemini-3-pro-preview": { - "actual": "gemini-3-pro", - "internal_name": "gemini-3-pro-high", - "provider": "google", - "description": "Gemini 3 Pro 预览版" - }, - "gemini-3-flash-preview": { - "actual": "gemini-3-flash", - "internal_name": "gemini-3-flash", - "provider": "google", - "description": "Gemini 3 Flash 预览版" - }, - "gemini-claude-sonnet-4-5": { - "actual": "claude-sonnet-4-5-20250929", - "internal_name": "claude-sonnet-4-5", - "provider": "anthropic", - "description": "通过 Antigravity 访问的 Claude Sonnet 4.5" - }, - "gemini-claude-sonnet-4-5-thinking": { - "actual": "claude-sonnet-4-5-20250929", - "internal_name": "claude-sonnet-4-5-thinking", - "provider": "anthropic", - "description": "Claude Sonnet 4.5 思考模式" - }, - "gemini-claude-opus-4-5-thinking": { - "actual": "claude-opus-4-5-20251101", - "internal_name": "claude-opus-4-5-thinking", - "provider": "anthropic", - "description": "Claude Opus 4.5 思考模式" - } - }, - "updated_at": "2026-01-07T00:00:00Z" -} diff --git a/src-tauri/resources/models/aliases/codex.json b/src-tauri/resources/models/aliases/codex.json deleted file mode 100644 index 7ae999243..000000000 --- a/src-tauri/resources/models/aliases/codex.json +++ /dev/null @@ -1,24 +0,0 @@ -{ - "$schema": "../schema/alias.schema.json", - "provider": "codex", - "description": "OpenAI Codex CLI 支持的模型", - "models": [ - "gpt-5.3-codex", - "gpt-5.4" - ], - "aliases": { - "gpt-5.4": { - "actual": "gpt-5.4", - "internal_name": "gpt-5.4", - "provider": "openai", - "description": "最新前沿模型,跨知识、推理和编码的全面提升" - }, - "gpt-5.3-codex": { - "actual": "gpt-5.3-codex", - "internal_name": "gpt-5.3-codex", - "provider": "openai", - "description": "Codex 最新一代模型,编码与推理能力增强" - } - }, - "updated_at": "2026-02-11T00:00:00Z" -} \ No newline at end of file diff --git a/src-tauri/resources/models/aliases/gemini.json b/src-tauri/resources/models/aliases/gemini.json deleted file mode 100644 index 93d028fca..000000000 --- a/src-tauri/resources/models/aliases/gemini.json +++ /dev/null @@ -1,38 +0,0 @@ -{ - "$schema": "../schema/alias.schema.json", - "provider": "gemini", - "description": "Gemini CLI OAuth 服务的模型别名映射(基于 Cloud Code Assist)", - "models": [ - "gemini-3-pro-preview", - "gemini-3-flash-preview", - "gemini-flash-latest", - "gemini-pro-latest" - ], - "aliases": { - "gemini-3-pro-preview": { - "actual": "gemini-3-pro-preview", - "internal_name": "gemini-3-pro-preview", - "provider": "google", - "description": "Gemini 3 Pro 预览版" - }, - "gemini-3-flash-preview": { - "actual": "gemini-3-flash-preview", - "internal_name": "gemini-3-flash-preview", - "provider": "google", - "description": "Gemini 3 Flash 预览版" - }, - "gemini-flash-latest": { - "actual": "gemini-flash-latest", - "internal_name": "gemini-flash-latest", - "provider": "google", - "description": "Gemini Flash 最新版别名" - }, - "gemini-pro-latest": { - "actual": "gemini-pro-latest", - "internal_name": "gemini-pro-latest", - "provider": "google", - "description": "Gemini Pro 最新版别名" - } - }, - "updated_at": "2026-01-13T00:00:00Z" -} diff --git a/src-tauri/resources/models/aliases/kiro.json b/src-tauri/resources/models/aliases/kiro.json deleted file mode 100644 index 66dba5fdd..000000000 --- a/src-tauri/resources/models/aliases/kiro.json +++ /dev/null @@ -1,56 +0,0 @@ -{ - "$schema": "../schema/alias.schema.json", - "provider": "kiro", - "description": "Kiro/CodeWhisperer 服务的模型别名映射(基于 AWS Bedrock)", - "models": [ - "claude-opus-4-5-20251101", - "claude-haiku-4-5-20251001", - "claude-sonnet-4-5-20250929", - "claude-sonnet-4-20250514" - ], - "aliases": { - "claude-opus-4-5": { - "actual": "claude-opus-4-5-20251101", - "internal_name": "claude-opus-4.5", - "provider": "anthropic", - "description": "Claude Opus 4.5 via Kiro" - }, - "claude-opus-4-5-20251101": { - "actual": "claude-opus-4-5-20251101", - "internal_name": "claude-opus-4.5", - "provider": "anthropic", - "description": "Claude Opus 4.5 via Kiro" - }, - "claude-haiku-4-5": { - "actual": "claude-haiku-4-5-20251001", - "internal_name": "claude-haiku-4.5", - "provider": "anthropic", - "description": "Claude Haiku 4.5 via Kiro" - }, - "claude-haiku-4-5-20251001": { - "actual": "claude-haiku-4-5-20251001", - "internal_name": "claude-haiku-4.5", - "provider": "anthropic", - "description": "Claude Haiku 4.5 via Kiro" - }, - "claude-sonnet-4-5": { - "actual": "claude-sonnet-4-5-20250929", - "internal_name": "CLAUDE_SONNET_4_5_20250929_V1_0", - "provider": "anthropic", - "description": "Claude Sonnet 4.5 via Kiro" - }, - "claude-sonnet-4-5-20250929": { - "actual": "claude-sonnet-4-5-20250929", - "internal_name": "CLAUDE_SONNET_4_5_20250929_V1_0", - "provider": "anthropic", - "description": "Claude Sonnet 4.5 via Kiro" - }, - "claude-sonnet-4-20250514": { - "actual": "claude-sonnet-4-20250514", - "internal_name": "CLAUDE_SONNET_4_20250514_V1_0", - "provider": "anthropic", - "description": "Claude Sonnet 4 via Kiro" - } - }, - "updated_at": "2026-01-07T00:00:00Z" -} diff --git a/src-tauri/resources/models/index.json b/src-tauri/resources/models/index.json index 0c79deb02..03e57fff6 100644 --- a/src-tauri/resources/models/index.json +++ b/src-tauri/resources/models/index.json @@ -5,7 +5,6 @@ "abacus", "aihubmix", "alibaba", - "antigravity", "alibaba-cn", "amazon-bedrock", "anthropic", @@ -17,7 +16,6 @@ "chutes", "cloudflare-ai-gateway", "cloudflare-workers-ai", - "codex", "cohere", "cortecs", "deepinfra", @@ -80,7 +78,7 @@ "zhipuai", "zhipuai-coding-plan" ], - "total_models": 2041, + "total_models": 2030, "sources": { "models_dev": "https://models.dev/api.json", "manual": [] diff --git a/src-tauri/resources/models/providers/antigravity.json b/src-tauri/resources/models/providers/antigravity.json deleted file mode 100644 index b7b2b6b01..000000000 --- a/src-tauri/resources/models/providers/antigravity.json +++ /dev/null @@ -1,297 +0,0 @@ -{ - "$schema": "../schema/model.schema.json", - "provider": { - "id": "antigravity", - "name": "Antigravity" - }, - "models": [ - { - "id": "gemini-3-pro-preview", - "name": "Gemini 3 Pro Preview", - "family": "gemini-pro", - "tier": "max", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 1000000, - "max_output": 64000 - }, - "status": "active", - "release_date": "2025-11-18", - "is_latest": true - }, - { - "id": "gemini-3-flash-preview", - "name": "Gemini 3 Flash Preview", - "family": "gemini-flash", - "tier": "mini", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 1048576, - "max_output": 65536 - }, - "status": "active", - "release_date": "2025-12-17", - "is_latest": true - }, - { - "id": "gemini-3-pro-image-preview", - "name": "Gemini 3 Pro Image", - "family": "gemini-pro-image", - "tier": "max", - "capabilities": { - "vision": true, - "tools": false, - "streaming": true, - "json_mode": false, - "function_calling": false, - "reasoning": false - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 32000, - "max_output": 8192 - }, - "status": "active", - "release_date": "2025-12-01", - "is_latest": true - }, - { - "id": "gemini-2.5-pro", - "name": "Gemini 2.5 Pro", - "family": "gemini-pro", - "tier": "max", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 1048576, - "max_output": 65536 - }, - "status": "active", - "release_date": "2025-03-20", - "is_latest": false - }, - { - "id": "gemini-2.5-flash", - "name": "Gemini 2.5 Flash", - "family": "gemini-flash", - "tier": "mini", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 1048576, - "max_output": 65536 - }, - "status": "active", - "release_date": "2025-03-20", - "is_latest": false - }, - { - "id": "gemini-2.5-flash-lite", - "name": "Gemini 2.5 Flash Lite", - "family": "gemini-flash-lite", - "tier": "mini", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 1048576, - "max_output": 65536 - }, - "status": "active", - "release_date": "2025-06-17", - "is_latest": false - }, - { - "id": "gemini-2.5-computer-use-preview-10-2025", - "name": "Gemini 2.5 Computer Use", - "family": "gemini-computer-use", - "tier": "pro", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 128000, - "max_output": 8192 - }, - "status": "preview", - "release_date": "2025-10-01", - "is_latest": true - }, - { - "id": "gemini-2.0-flash", - "name": "Gemini 2.0 Flash", - "family": "gemini-flash", - "tier": "mini", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": false - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 1048576, - "max_output": 8192 - }, - "status": "active", - "release_date": "2024-12-11", - "is_latest": false - }, - { - "id": "gemini-claude-opus-4-5-thinking", - "name": "Claude Opus 4.5 Thinking", - "family": "claude-opus", - "tier": "max", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 200000, - "max_output": 32000 - }, - "status": "active", - "release_date": "2025-11-01", - "is_latest": true - }, - { - "id": "gemini-claude-sonnet-4-5-thinking", - "name": "Claude Sonnet 4.5 Thinking", - "family": "claude-sonnet", - "tier": "pro", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 200000, - "max_output": 64000 - }, - "status": "active", - "release_date": "2025-09-29", - "is_latest": true - }, - { - "id": "gemini-claude-sonnet-4-5", - "name": "Claude Sonnet 4.5", - "family": "claude-sonnet", - "tier": "pro", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": false - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 200000, - "max_output": 64000 - }, - "status": "active", - "release_date": "2025-09-29", - "is_latest": false - } - ], - "updated_at": "2026-01-12T00:00:00.000Z", - "source": "antigravity-manager" -} diff --git a/src-tauri/resources/models/providers/codex.json b/src-tauri/resources/models/providers/codex.json deleted file mode 100644 index ed207b13b..000000000 --- a/src-tauri/resources/models/providers/codex.json +++ /dev/null @@ -1,89 +0,0 @@ -{ - "$schema": "../schema/model.schema.json", - "provider": { - "id": "codex", - "name": "Codex" - }, - "models": [ - { - "id": "gpt-5.2", - "name": "GPT-5.2", - "family": "gpt-5", - "tier": "pro", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 400000, - "max_output": 128000 - }, - "status": "active", - "release_date": "2025-12-11", - "is_latest": true - }, - { - "id": "gpt-5.3-codex", - "name": "GPT-5.3 Codex", - "family": "gpt-5-codex", - "tier": "pro", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 272000, - "max_output": 128000 - }, - "status": "active", - "release_date": "2026-02-05", - "is_latest": true - }, - { - "id": "gpt-5.2-codex", - "name": "GPT-5.2 Codex", - "family": "gpt-5-codex", - "tier": "pro", - "capabilities": { - "vision": true, - "tools": true, - "streaming": true, - "json_mode": true, - "function_calling": true, - "reasoning": true - }, - "pricing": { - "input": 0, - "output": 0, - "currency": "USD" - }, - "limits": { - "context": 272000, - "max_output": 128000 - }, - "status": "active", - "release_date": "2025-12-11", - "is_latest": false - } - ], - "updated_at": "2026-02-11T00:00:00.000Z", - "source": "codex-cli" -} \ No newline at end of file diff --git a/src-tauri/src/README.md b/src-tauri/src/README.md index 811edfa82..a750003ee 100644 --- a/src-tauri/src/README.md +++ b/src-tauri/src/README.md @@ -12,8 +12,8 @@ Tauri 后端核心代码,处理系统级功能和 API 服务。 - `commands/` - Tauri 命令处理(前端调用入口) - `config/` - 配置管理(导入/导出/热重载) - `connect/` - Lime Connect 模块(中转商生态合作) -- `converter/` - 协议转换(OpenAI ↔ CW/Claude/Antigravity) -- `credential/` - 凭证池管理(负载均衡、健康检查) +- `converter/` - 协议转换(OpenAI / Claude / Antigravity) +- API Key Provider / configured providers - 当前 Provider 凭证与模型配置主路径 - `database/` - 数据库层(SQLite + DAO) - `errors/` - 错误类型定义(项目、人设、素材、模板、迁移错误) - `flow_monitor/` - LLM 流量监控(拦截、存储、查询) diff --git a/src-tauri/src/agent/README.md b/src-tauri/src/agent/README.md index 890a30b49..0cd6fcfbe 100644 --- a/src-tauri/src/agent/README.md +++ b/src-tauri/src/agent/README.md @@ -9,7 +9,7 @@ AI Agent 集成模块,基于 aster-rust 框架实现。 ### 设计决策 - **Aster 框架**:使用 aster-rust 框架获得多 Provider、工具系统、会话管理等能力 -- **凭证池桥接**:自动从 Lime 凭证池选择凭证配置 Aster Provider +- **Provider 桥接**:自动从 API Key Provider 选择凭证配置 Aster Provider - **流式响应**:通过 Tauri 事件系统向前端推送流式内容 - **Skills 集成**:自动加载 Lime Skills 到 aster-rust,使 AI 能够自动调用 - **多代理收敛**:旧 SubAgent scheduler Tauri 桥已删除,角色与 team runtime 纯逻辑统一收敛到 `crates/agent/src/subagent_scheduler.rs` @@ -66,13 +66,13 @@ AsterAgentState::reload_lime_skills(); - 再看 `docs/aiprompts/prompt-foundation.md` - 不要从本页示例反推新的 Tauri 主路径 -### 从凭证池配置(推荐) +### 从 API Key Provider 配置(推荐) ```rust // 初始化(同时加载 Skills) state.init_agent_with_db(&db).await?; -// 从凭证池自动选择凭证并配置 Provider +// 从 API Key Provider 自动选择凭证并配置 Provider let config = state .configure_provider_from_pool(&db, "openai", "gpt-4", &session_id) .await?; @@ -118,18 +118,18 @@ let stream = agent.reply(user_message, session_config, Some(cancel_token)).await |------|------| | `aster_agent_init` | 初始化 Agent | | `aster_agent_configure_provider` | 手动配置 Provider | -| `aster_agent_configure_from_pool` | 从凭证池配置 Provider(推荐) | +| `aster_agent_configure_provider` | 从 API Key Provider / configured providers 配置 Provider | | `agent_runtime_submit_turn` | 统一提交 turn | | `agent_runtime_interrupt_turn` | 统一中断 turn | | `agent_runtime_create/list/get/update/delete_session` | 统一会话管理 | | `agent_runtime_spawn_subagent / agent_runtime_send_subagent_input / agent_runtime_wait_subagents / agent_runtime_resume_subagent / agent_runtime_close_subagent` | subagent 控制面 | | `agent_runtime_respond_action` | 统一响应工具确认 / ask / elicitation | -## 凭证池桥接 +## Provider 桥接 `credential_bridge.rs` 在主 crate 中仅作为兼容导出,核心逻辑位于 `crates/agent/src/credential_bridge.rs`: -- 自动从凭证池选择可用凭证 +- 自动从 API Key Provider 选择可用凭证 - 支持 OAuth 和 API Key 两种凭证类型 - 自动刷新过期的 OAuth Token - 记录凭证使用和健康状态 diff --git a/src-tauri/src/agent/credential_bridge.rs b/src-tauri/src/agent/credential_bridge.rs index b6a09e1a8..5e02235f4 100644 --- a/src-tauri/src/agent/credential_bridge.rs +++ b/src-tauri/src/agent/credential_bridge.rs @@ -1,4 +1,4 @@ -//! 凭证池桥接模块(重导出层) +//! API Key Provider 桥接模块(重导出层) //! //! 纯逻辑已迁移到 `lime-agent` crate, //! 本模块仅保留兼容导出。 diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index af0d5fe26..44faee61b 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -11,10 +11,8 @@ use crate::commands::connect_cmd::ConnectStateWrapper; use crate::commands::context_memory::ContextMemoryServiceState; use crate::commands::machine_id_cmd::MachineIdState; use crate::commands::model_registry_cmd::ModelRegistryState; -use crate::commands::orchestrator_cmd::OrchestratorState; use crate::commands::plugin_cmd::PluginManagerState; use crate::commands::plugin_install_cmd::PluginInstallerState; -use crate::commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState}; use crate::commands::session_files_cmd::SessionFilesState; use crate::commands::skill_cmd::SkillServiceState; use crate::commands::webview_cmd::{ @@ -34,12 +32,10 @@ use lime_scheduler::AgentScheduler; use lime_server as server; use lime_services::api_key_provider_service::ApiKeyProviderService; use lime_services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; -use lime_services::provider_pool_service::ProviderPoolService; use lime_services::skill_service::SkillService; -use lime_services::token_cache_service::TokenCacheService; use lime_services::update_check_service::UpdateCheckServiceState; -use super::types::{AppState, LogState, TokenCacheServiceState}; +use super::types::{AppState, LogState}; pub use lime_core::app_bootstrap::{load_and_validate_config, ConfigError}; @@ -49,17 +45,13 @@ pub struct AppStates { pub logs: LogState, pub db: DbConnection, pub skill_service: SkillServiceState, - pub provider_pool_service: ProviderPoolServiceState, pub api_key_provider_service: ApiKeyProviderServiceState, - pub credential_sync_service: CredentialSyncServiceState, - pub token_cache_service: TokenCacheServiceState, pub machine_id_service: MachineIdState, pub plugin_manager: PluginManagerState, pub plugin_installer: PluginInstallerState, pub plugin_rpc_manager: crate::commands::plugin_rpc_cmd::PluginRpcManagerState, pub telemetry: crate::commands::telemetry_cmd::TelemetryState, pub aster_agent: AsterAgentState, - pub orchestrator: OrchestratorState, pub connect_state: ConnectStateWrapper, pub model_registry: ModelRegistryState, pub global_config_manager: GlobalConfigManagerState, @@ -143,9 +135,6 @@ pub fn init_states(config: &Config) -> Result { let skill_service = SkillService::new().map_err(|e| format!("SkillService 初始化失败: {e}"))?; let skill_service_state = SkillServiceState(Arc::new(skill_service)); - let provider_pool_service = ProviderPoolService::new(); - let provider_pool_service_state = ProviderPoolServiceState(Arc::new(provider_pool_service)); - let api_key_provider_service = ApiKeyProviderService::new(); match api_key_provider_service.migrate_legacy_api_key_encryption(&db) { Ok(0) => { @@ -164,11 +153,6 @@ pub fn init_states(config: &Config) -> Result { let api_key_provider_service_state = ApiKeyProviderServiceState(Arc::new(api_key_provider_service)); - let credential_sync_service_state = CredentialSyncServiceState(None); - - let token_cache_service = TokenCacheService::new(); - let token_cache_service_state = TokenCacheServiceState(Arc::new(token_cache_service)); - let machine_id_service = lime_services::machine_id_service::MachineIdService::new() .map_err(|e| format!("MachineIdService 初始化失败: {e}"))?; let machine_id_service_state: MachineIdState = Arc::new(RwLock::new(machine_id_service)); @@ -188,7 +172,6 @@ pub fn init_states(config: &Config) -> Result { // 其他状态 let aster_agent_state = AsterAgentState::new(); - let orchestrator_state = OrchestratorState::new(); // 初始化 Connect 状态(延迟初始化,在 setup hook 中完成) let connect_state = ConnectStateWrapper(Arc::new(RwLock::new(None))); @@ -270,17 +253,13 @@ pub fn init_states(config: &Config) -> Result { logs, db, skill_service: skill_service_state, - provider_pool_service: provider_pool_service_state, api_key_provider_service: api_key_provider_service_state, - credential_sync_service: credential_sync_service_state, - token_cache_service: token_cache_service_state, machine_id_service: machine_id_service_state, plugin_manager: plugin_manager_state, plugin_installer: plugin_installer_state, plugin_rpc_manager: plugin_rpc_manager_state, telemetry: telemetry_state, aster_agent: aster_agent_state, - orchestrator: orchestrator_state, connect_state, model_registry: model_registry_state, global_config_manager: global_config_manager_state, diff --git a/src-tauri/src/app/commands/api_test.rs b/src-tauri/src/app/commands/api_test.rs index a7c9ff165..5e3cd70fb 100644 --- a/src-tauri/src/app/commands/api_test.rs +++ b/src-tauri/src/app/commands/api_test.rs @@ -37,7 +37,7 @@ pub struct ApiCompatibilityResult { /// 检查 API 兼容性 #[tauri::command] pub async fn check_api_compatibility( - state: tauri::State<'_, AppState>, + _state: tauri::State<'_, AppState>, logs: tauri::State<'_, LogState>, provider: String, ) -> Result { @@ -49,37 +49,20 @@ pub async fn check_api_compatibility( &format!("[API检测] 开始检测 {provider_type} API 兼容性(代理工具调用能力测试)..."), ); - let s = state.read().await; let mut results: Vec = Vec::new(); let mut warnings: Vec = Vec::new(); // 代理工具调用主链需要的测试项目 let test_cases: Vec<(&str, &str)> = match provider_type { - ProviderType::Kiro => vec![ - ("claude-sonnet-4-5", "basic"), - ("claude-sonnet-4-5", "tool_call"), - ], - ProviderType::Gemini => vec![ - ("gemini-2.5-flash", "basic"), - ("gemini-2.5-flash", "tool_call"), - ], - ProviderType::Antigravity => vec![ - ("gemini-3-pro-preview", "basic"), - ("gemini-3-pro-preview", "tool_call"), - ], + ProviderType::Kiro + | ProviderType::Gemini + | ProviderType::GeminiApiKey + | ProviderType::Codex + | ProviderType::ClaudeOAuth => vec![], ProviderType::Vertex => vec![ ("gemini-2.0-flash", "basic"), ("gemini-2.0-flash", "tool_call"), ], - ProviderType::GeminiApiKey => vec![ - ("gemini-2.5-flash", "basic"), - ("gemini-2.5-flash", "tool_call"), - ], - ProviderType::Codex => vec![("gpt-4.1", "basic"), ("gpt-4.1", "tool_call")], - ProviderType::ClaudeOAuth => vec![ - ("claude-sonnet-4-5", "basic"), - ("claude-sonnet-4-5", "tool_call"), - ], // Anthropic 兼容格式 - 使用 Claude 相同的测试 ProviderType::AnthropicCompatible => vec![], ProviderType::OpenAI | ProviderType::Claude => vec![], @@ -95,7 +78,7 @@ pub async fn check_api_compatibility( let test_name = format!("{model} ({test_type})"); // 根据测试类型构建不同的请求 - let test_request = match test_type { + let _test_request = match test_type { "tool_call" => { // 测试 Tool Calls - 代理工具调用主链能力 crate::models::openai::ChatCompletionRequest { @@ -157,13 +140,8 @@ pub async fn check_api_compatibility( } }; - let result = match provider_type { - ProviderType::Kiro => s.kiro_provider.call_api(&test_request).await, - ProviderType::Gemini => { - Err("Gemini API compatibility check not yet implemented".into()) - } - _ => Err("Provider not supported for direct API check".into()), - }; + let result: Result> = + Err("Provider not supported for direct API check".into()); let time_ms = start.elapsed().as_millis() as u64; diff --git a/src-tauri/src/app/commands/gemini.rs b/src-tauri/src/app/commands/gemini.rs deleted file mode 100644 index 9ee31ac2b..000000000 --- a/src-tauri/src/app/commands/gemini.rs +++ /dev/null @@ -1,184 +0,0 @@ -//! Gemini Provider 命令 (Legacy) -//! -//! 包含 Gemini 凭证管理相关命令。 -//! 这些命令保留用于向后兼容,新代码应使用统一的 OAuth 命令。 - -use crate::app::commands::kiro::{CheckResult, EnvVariable}; -use crate::app::types::{AppState, LogState}; -use crate::app::utils::mask_token; -use crate::providers; - -/// Gemini 凭证状态 -#[derive(serde::Serialize)] -pub struct GeminiCredentialStatus { - pub loaded: bool, - pub has_access_token: bool, - pub has_refresh_token: bool, - pub expiry_date: Option, - pub is_valid: bool, - pub creds_path: String, -} - -/// 获取 Gemini 凭证状态 -#[tauri::command] -pub async fn get_gemini_credentials( - state: tauri::State<'_, AppState>, -) -> Result { - let s = state.read().await; - let creds = &s.gemini_provider.credentials; - let path = providers::gemini::GeminiProvider::default_creds_path(); - - Ok(GeminiCredentialStatus { - loaded: creds.access_token.is_some() || creds.refresh_token.is_some(), - has_access_token: creds.access_token.is_some(), - has_refresh_token: creds.refresh_token.is_some(), - expiry_date: creds.expiry_date, - is_valid: s.gemini_provider.is_token_valid(), - creds_path: path.to_string_lossy().to_string(), - }) -} - -/// 重新加载 Gemini 凭证 -#[tauri::command] -pub async fn reload_gemini_credentials( - state: tauri::State<'_, AppState>, - logs: tauri::State<'_, LogState>, -) -> Result { - let mut s = state.write().await; - logs.write().await.add("info", "[Gemini] 正在加载凭证..."); - s.gemini_provider - .load_credentials() - .await - .map_err(|e| e.to_string())?; - logs.write().await.add("info", "[Gemini] 凭证加载成功"); - Ok("Gemini credentials reloaded".to_string()) -} - -/// 刷新 Gemini Token -#[tauri::command] -pub async fn refresh_gemini_token( - state: tauri::State<'_, AppState>, - logs: tauri::State<'_, LogState>, -) -> Result { - let mut s = state.write().await; - logs.write().await.add("info", "[Gemini] 正在刷新 Token..."); - let result = s - .gemini_provider - .refresh_token() - .await - .map_err(|e| e.to_string()); - match &result { - Ok(_) => logs.write().await.add("info", "[Gemini] Token 刷新成功"), - Err(e) => logs - .write() - .await - .add("error", &format!("[Gemini] Token 刷新失败: {e}")), - } - result -} - -/// 获取 Gemini 环境变量 -#[tauri::command] -pub async fn get_gemini_env_variables( - state: tauri::State<'_, AppState>, -) -> Result, String> { - let s = state.read().await; - let creds = &s.gemini_provider.credentials; - let mut vars = Vec::new(); - - if let Some(token) = &creds.access_token { - vars.push(EnvVariable { - key: "GEMINI_ACCESS_TOKEN".to_string(), - value: token.clone(), - masked: mask_token(token), - }); - } - if let Some(token) = &creds.refresh_token { - vars.push(EnvVariable { - key: "GEMINI_REFRESH_TOKEN".to_string(), - value: token.clone(), - masked: mask_token(token), - }); - } - if let Some(expiry) = creds.expiry_date { - let expiry_str = expiry.to_string(); - vars.push(EnvVariable { - key: "GEMINI_EXPIRY_DATE".to_string(), - value: expiry_str.clone(), - masked: expiry_str, - }); - } - - Ok(vars) -} - -/// 获取 Gemini Token 文件哈希 -#[tauri::command] -pub async fn get_gemini_token_file_hash() -> Result { - let path = providers::gemini::GeminiProvider::default_creds_path(); - if !tokio::fs::try_exists(&path).await.unwrap_or(false) { - return Ok("".to_string()); - } - - let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; - let hash = format!("{:x}", md5::compute(&content)); - Ok(hash) -} - -/// 检查并重新加载 Gemini 凭证 -#[tauri::command] -pub async fn check_and_reload_gemini_credentials( - state: tauri::State<'_, AppState>, - logs: tauri::State<'_, LogState>, - last_hash: String, -) -> Result { - let path = providers::gemini::GeminiProvider::default_creds_path(); - - if !tokio::fs::try_exists(&path).await.unwrap_or(false) { - return Ok(CheckResult { - changed: false, - new_hash: "".to_string(), - reloaded: false, - }); - } - - let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; - let new_hash = format!("{:x}", md5::compute(&content)); - - if !last_hash.is_empty() && new_hash != last_hash { - logs.write() - .await - .add("info", "[Gemini][自动检测] 凭证文件已变化,正在重新加载..."); - - let mut s = state.write().await; - match s.gemini_provider.load_credentials().await { - Ok(_) => { - logs.write() - .await - .add("info", "[Gemini][自动检测] 凭证重新加载成功"); - Ok(CheckResult { - changed: true, - new_hash, - reloaded: true, - }) - } - Err(e) => { - logs.write().await.add( - "error", - &format!("[Gemini][自动检测] 凭证重新加载失败: {e}"), - ); - Ok(CheckResult { - changed: true, - new_hash, - reloaded: false, - }) - } - } - } else { - Ok(CheckResult { - changed: false, - new_hash, - reloaded: false, - }) - } -} diff --git a/src-tauri/src/app/commands/kiro.rs b/src-tauri/src/app/commands/kiro.rs deleted file mode 100644 index 58b9541b5..000000000 --- a/src-tauri/src/app/commands/kiro.rs +++ /dev/null @@ -1,231 +0,0 @@ -//! Kiro Provider 命令 (Legacy) -//! -//! 包含 Kiro 凭证管理相关命令。 -//! 这些命令保留用于向后兼容,新代码应使用统一的 OAuth 命令。 - -use crate::app::types::{AppState, LogState}; -use crate::app::utils::mask_token; -use crate::providers; - -/// Kiro 凭证状态 -#[derive(serde::Serialize)] -pub struct KiroCredentialStatus { - pub loaded: bool, - pub has_access_token: bool, - pub has_refresh_token: bool, - pub region: Option, - pub auth_method: Option, - pub expires_at: Option, - pub creds_path: String, -} - -/// 环境变量 -#[derive(serde::Serialize)] -pub struct EnvVariable { - pub key: String, - pub value: String, - pub masked: String, -} - -/// 检查结果 -#[derive(serde::Serialize)] -pub struct CheckResult { - pub changed: bool, - pub new_hash: String, - pub reloaded: bool, -} - -/// 刷新 Kiro Token -#[tauri::command] -pub async fn refresh_kiro_token( - state: tauri::State<'_, AppState>, - logs: tauri::State<'_, LogState>, -) -> Result { - let mut s = state.write().await; - logs.write().await.add("info", "Refreshing Kiro token..."); - let result = s - .kiro_provider - .refresh_token() - .await - .map_err(|e| e.to_string()); - match &result { - Ok(_) => logs - .write() - .await - .add("info", "Token refreshed successfully"), - Err(e) => logs - .write() - .await - .add("error", &format!("Token refresh failed: {e}")), - } - result -} - -/// 重新加载凭证 -#[tauri::command] -pub async fn reload_credentials( - state: tauri::State<'_, AppState>, - logs: tauri::State<'_, LogState>, -) -> Result { - let mut s = state.write().await; - logs.write().await.add("info", "Reloading credentials..."); - s.kiro_provider - .load_credentials() - .await - .map_err(|e| e.to_string())?; - logs.write().await.add("info", "Credentials reloaded"); - Ok("Credentials reloaded".to_string()) -} - -/// 获取 Kiro 凭证状态 -#[tauri::command] -pub async fn get_kiro_credentials( - state: tauri::State<'_, AppState>, -) -> Result { - let s = state.read().await; - let creds = &s.kiro_provider.credentials; - let path = providers::kiro::KiroProvider::default_creds_path(); - - Ok(KiroCredentialStatus { - loaded: creds.access_token.is_some() || creds.refresh_token.is_some(), - has_access_token: creds.access_token.is_some(), - has_refresh_token: creds.refresh_token.is_some(), - region: creds.region.clone(), - auth_method: creds.auth_method.clone(), - expires_at: creds.expires_at.clone(), - creds_path: path.to_string_lossy().to_string(), - }) -} - -/// 获取环境变量 -#[tauri::command] -pub async fn get_env_variables( - state: tauri::State<'_, AppState>, -) -> Result, String> { - let s = state.read().await; - let creds = &s.kiro_provider.credentials; - let mut vars = Vec::new(); - - // P0 安全修复:不再返回明文敏感凭证,仅返回 masked 版本 - if let Some(token) = &creds.access_token { - vars.push(EnvVariable { - key: "KIRO_ACCESS_TOKEN".to_string(), - value: String::new(), // 不返回明文 - masked: mask_token(token), - }); - } - if let Some(token) = &creds.refresh_token { - vars.push(EnvVariable { - key: "KIRO_REFRESH_TOKEN".to_string(), - value: String::new(), // 不返回明文 - masked: mask_token(token), - }); - } - if let Some(id) = &creds.client_id { - vars.push(EnvVariable { - key: "KIRO_CLIENT_ID".to_string(), - value: String::new(), // 不返回明文 - masked: mask_token(id), - }); - } - if let Some(secret) = &creds.client_secret { - vars.push(EnvVariable { - key: "KIRO_CLIENT_SECRET".to_string(), - value: String::new(), // 不返回明文 - masked: mask_token(secret), - }); - } - if let Some(arn) = &creds.profile_arn { - vars.push(EnvVariable { - key: "KIRO_PROFILE_ARN".to_string(), - value: arn.clone(), - masked: arn.clone(), - }); - } - if let Some(region) = &creds.region { - vars.push(EnvVariable { - key: "KIRO_REGION".to_string(), - value: region.clone(), - masked: region.clone(), - }); - } - if let Some(method) = &creds.auth_method { - vars.push(EnvVariable { - key: "KIRO_AUTH_METHOD".to_string(), - value: method.clone(), - masked: method.clone(), - }); - } - - Ok(vars) -} - -/// 获取 Token 文件哈希 -#[tauri::command] -pub async fn get_token_file_hash() -> Result { - let path = providers::kiro::KiroProvider::default_creds_path(); - if !tokio::fs::try_exists(&path).await.unwrap_or(false) { - return Ok("".to_string()); - } - - let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; - let hash = format!("{:x}", md5::compute(&content)); - Ok(hash) -} - -/// 检查凭证文件变化并自动重新加载 -#[tauri::command] -pub async fn check_and_reload_credentials( - state: tauri::State<'_, AppState>, - logs: tauri::State<'_, LogState>, - last_hash: String, -) -> Result { - let path = providers::kiro::KiroProvider::default_creds_path(); - - if !tokio::fs::try_exists(&path).await.unwrap_or(false) { - return Ok(CheckResult { - changed: false, - new_hash: "".to_string(), - reloaded: false, - }); - } - - let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; - let new_hash = format!("{:x}", md5::compute(&content)); - - if !last_hash.is_empty() && new_hash != last_hash { - logs.write() - .await - .add("info", "[自动检测] 凭证文件已变化,正在重新加载..."); - - let mut s = state.write().await; - match s.kiro_provider.load_credentials().await { - Ok(_) => { - logs.write() - .await - .add("info", "[自动检测] 凭证重新加载成功"); - Ok(CheckResult { - changed: true, - new_hash, - reloaded: true, - }) - } - Err(e) => { - logs.write() - .await - .add("error", &format!("[自动检测] 凭证重新加载失败: {e}")); - Ok(CheckResult { - changed: true, - new_hash, - reloaded: false, - }) - } - } - } else { - Ok(CheckResult { - changed: false, - new_hash, - reloaded: false, - }) - } -} diff --git a/src-tauri/src/app/commands/mod.rs b/src-tauri/src/app/commands/mod.rs index a1e0f94ac..b094932b0 100644 --- a/src-tauri/src/app/commands/mod.rs +++ b/src-tauri/src/app/commands/mod.rs @@ -5,8 +5,6 @@ //! ## 模块结构 //! - `server` - 服务器控制命令 //! - `config` - 配置管理命令 -//! - `kiro` - Kiro Provider 命令 (legacy) -//! - `gemini` - Gemini Provider 命令 (legacy) //! - `custom_providers` - 自定义 Provider 命令 (OpenAI/Claude Custom) //! - `logs` - 日志命令 //! - `api_test` - API 测试和兼容性检查命令 @@ -14,8 +12,6 @@ mod api_test; mod config; mod custom_providers; -mod gemini; -mod kiro; mod logs; mod server; @@ -23,7 +19,5 @@ mod server; pub use api_test::*; pub use config::*; pub use custom_providers::*; -pub use gemini::*; -pub use kiro::*; pub use logs::*; pub use server::*; diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 086ff4c93..6e0c42bf5 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -130,17 +130,13 @@ pub fn run() { logs, db, skill_service: skill_service_state, - provider_pool_service: provider_pool_service_state, api_key_provider_service: api_key_provider_service_state, - credential_sync_service: credential_sync_service_state, - token_cache_service: token_cache_service_state, machine_id_service: machine_id_service_state, plugin_manager: plugin_manager_state, plugin_installer: plugin_installer_state, plugin_rpc_manager: plugin_rpc_manager_state, telemetry: telemetry_state, aster_agent: aster_agent_state, - orchestrator: orchestrator_state, connect_state, model_registry: model_registry_state, global_config_manager: global_config_manager_state, @@ -161,7 +157,6 @@ pub fn run() { let state_clone = state.clone(); let logs_clone = logs.clone(); let db_clone = db.clone(); - let pool_service_clone = provider_pool_service_state.0.clone(); #[cfg(debug_assertions)] let api_key_provider_service_clone = api_key_provider_service_state.0.clone(); #[cfg(debug_assertions)] @@ -170,7 +165,6 @@ pub fn run() { let model_registry_clone = model_registry_state.clone(); #[cfg(debug_assertions)] let skill_service_clone = skill_service_state.0.clone(); - let token_cache_clone = token_cache_service_state.0.clone(); let shared_stats_clone = shared_stats.clone(); let shared_tokens_clone = shared_tokens.clone(); let shared_logger_clone = shared_logger.clone(); @@ -244,17 +238,13 @@ pub fn run() { .manage(logs) .manage(db) .manage(skill_service_state) - .manage(provider_pool_service_state) .manage(api_key_provider_service_state) - .manage(credential_sync_service_state) - .manage(token_cache_service_state) .manage(machine_id_service_state) .manage(telemetry_state) .manage(plugin_manager_state) .manage(plugin_installer_state) .manage(plugin_rpc_manager_state) .manage(aster_agent_state) - .manage(orchestrator_state) .manage(connect_state) .manage(model_registry_state) .manage(global_config_manager_state) @@ -502,7 +492,6 @@ pub fn run() { let server_state = state_clone.clone(); let logs = logs_clone.clone(); let db = Some(db_clone.clone()); - let pool_service = pool_service_clone.clone(); let api_key_provider_service = api_key_provider_service_clone.clone(); let connect_state = connect_state_clone.clone(); let model_registry = model_registry_clone.clone(); @@ -515,7 +504,6 @@ pub fn run() { server_state, logs, db, - pool_service, api_key_provider_service, connect_state, model_registry, @@ -742,78 +730,11 @@ pub fn run() { let state = state_clone.clone(); let logs = logs_clone.clone(); let db = db_clone.clone(); - let pool_service = pool_service_clone.clone(); - let token_cache = token_cache_clone.clone(); let shared_stats = shared_stats_clone.clone(); let shared_tokens = shared_tokens_clone.clone(); let shared_logger = shared_logger_clone.clone(); let app_handle = app.handle().clone(); tauri::async_runtime::spawn(async move { - let mut available_credentials = 0usize; - let mut total_credentials = 0usize; - - // 先加载凭证池中的凭证 - { - logs.write().await.add("info", "[启动] 正在加载凭证池..."); - - // 获取凭证池概览信息 - match pool_service.get_overview(&db) { - Ok(overview) => { - let mut loaded_types = Vec::new(); - for provider_overview in overview { - let enabled_credentials: Vec<_> = provider_overview - .credentials - .iter() - .filter(|credential| !credential.is_disabled) - .collect(); - let count = enabled_credentials.len(); - if count > 0 { - total_credentials += count; - available_credentials += enabled_credentials - .iter() - .filter(|credential| credential.is_healthy) - .count(); - let provider_name = - match provider_overview.provider_type.as_str() { - "kiro" => "Kiro", - "gemini" => "Gemini", - "antigravity" => "Antigravity", - "openai" => "OpenAI", - "claude" => "Claude", - "codex" => "Codex", - "claude_oauth" => "Claude OAuth", - _ => &provider_overview.provider_type, - }; - loaded_types.push(format!("{provider_name} ({count} 个)")); - } - } - - if loaded_types.is_empty() { - logs.write().await.add("warn", "[启动] 未找到任何可用凭证"); - } else { - let message = format!( - "[启动] 凭证已加载: {} (共 {} 个)", - loaded_types.join(", "), - total_credentials - ); - logs.write().await.add("info", &message); - } - } - Err(e) => { - logs.write() - .await - .add("warn", &format!("[启动] 获取凭证池信息失败: {e}")); - } - } - - // 兼容性:仍然尝试加载旧的 Kiro 凭证(如果存在) - let mut s = state.write().await; - if let Err(e) = s.kiro_provider.load_credentials().await { - logs.write() - .await - .add("debug", &format!("[启动] 旧版 Kiro 凭证加载失败: {e}")); - } - } // 启动服务器(使用共享的遥测实例) { let mut s = state.write().await; @@ -823,8 +744,6 @@ pub fn run() { match s .start_with_telemetry( logs.clone(), - pool_service, - token_cache, Some(db), Some(shared_stats), Some(shared_tokens), @@ -854,18 +773,12 @@ pub fn run() { let tray_guard = tray_state.0.read().await; if let Some(tray_manager) = tray_guard.as_ref() { let current_state = tray_manager.get_state().await; - let icon_status = if total_credentials == 0 || available_credentials == 0 { - TrayIconStatus::Error - } else if available_credentials < total_credentials { - TrayIconStatus::Warning - } else { - TrayIconStatus::Running - }; + let icon_status = TrayIconStatus::Running; let snapshot = TrayStateSnapshot { icon_status, - available_credentials, - total_credentials, + available_credentials: current_state.available_credentials, + total_credentials: current_state.total_credentials, today_requests: current_state.today_requests, auto_start_enabled: current_state.auto_start_enabled, current_model_provider_type: current_state.current_model_provider_type, @@ -1096,28 +1009,6 @@ pub fn run() { app_commands::get_endpoint_providers, app_commands::set_endpoint_provider, app_commands::update_provider_env_vars, - // Unified OAuth commands (new) - commands::oauth_cmd::get_oauth_credentials, - commands::oauth_cmd::reload_oauth_credentials, - commands::oauth_cmd::refresh_oauth_token, - commands::oauth_cmd::get_oauth_env_variables, - commands::oauth_cmd::get_oauth_token_file_hash, - commands::oauth_cmd::check_and_reload_oauth_credentials, - commands::oauth_cmd::get_all_oauth_credentials, - // Legacy Kiro commands (from app::commands, deprecated) - app_commands::refresh_kiro_token, - app_commands::reload_credentials, - app_commands::get_kiro_credentials, - app_commands::get_env_variables, - app_commands::get_token_file_hash, - app_commands::check_and_reload_credentials, - // Legacy Gemini commands (from app::commands, deprecated) - app_commands::get_gemini_credentials, - app_commands::reload_gemini_credentials, - app_commands::refresh_gemini_token, - app_commands::get_gemini_env_variables, - app_commands::get_gemini_token_file_hash, - app_commands::check_and_reload_gemini_credentials, // OpenAI Custom commands (from app::commands) app_commands::get_openai_custom_status, app_commands::set_openai_custom_config, @@ -1258,60 +1149,6 @@ pub fn run() { commands::sceneapp_cmd::sceneapp_get_scorecard, // Ecommerce Review Reply commands commands::ecommerce_review_reply_cmd::execute_ecommerce_review_reply, - // Provider Pool commands - commands::provider_pool_cmd::get_provider_pool_overview, - commands::provider_pool_cmd::get_provider_pool_credentials, - commands::provider_pool_cmd::add_provider_pool_credential, - commands::provider_pool_cmd::update_provider_pool_credential, - commands::provider_pool_cmd::delete_provider_pool_credential, - commands::provider_pool_cmd::toggle_provider_pool_credential, - commands::provider_pool_cmd::reset_provider_pool_credential, - commands::provider_pool_cmd::reset_provider_pool_health, - commands::provider_pool_cmd::check_provider_pool_credential_health, - commands::provider_pool_cmd::check_provider_pool_type_health, - commands::provider_pool_cmd::add_kiro_oauth_credential, - commands::provider_pool_cmd::add_kiro_from_json, - commands::provider_pool_cmd::add_gemini_oauth_credential, - commands::provider_pool_cmd::add_antigravity_oauth_credential, - commands::provider_pool_cmd::add_openai_key_credential, - commands::provider_pool_cmd::add_claude_key_credential, - commands::provider_pool_cmd::add_gemini_api_key_credential, - commands::provider_pool_cmd::add_codex_oauth_credential, - commands::provider_pool_cmd::add_claude_oauth_credential, - commands::provider_pool_cmd::refresh_pool_credential_token, - commands::provider_pool_cmd::get_pool_credential_oauth_status, - commands::provider_pool_cmd::debug_kiro_credentials, - commands::provider_pool_cmd::test_user_credentials, - commands::provider_pool_cmd::migrate_private_config_to_pool, - commands::provider_pool_cmd::start_antigravity_oauth_login, - commands::provider_pool_cmd::get_antigravity_auth_url_and_wait, - commands::provider_pool_cmd::get_codex_auth_url_and_wait, - commands::provider_pool_cmd::start_codex_oauth_login, - commands::provider_pool_cmd::get_claude_oauth_auth_url_and_wait, - commands::provider_pool_cmd::start_claude_oauth_login, - commands::provider_pool_cmd::exchange_claude_oauth_code, - commands::provider_pool_cmd::claude_oauth_with_cookie, - commands::provider_pool_cmd::get_gemini_auth_url_and_wait, - commands::provider_pool_cmd::start_gemini_oauth_login, - commands::provider_pool_cmd::exchange_gemini_code, - commands::provider_pool_cmd::get_kiro_credential_fingerprint, - commands::provider_pool_cmd::get_credential_health, - commands::provider_pool_cmd::get_all_credential_health, - // Kiro Builder ID 登录命令 - commands::provider_pool_cmd::start_kiro_builder_id_login, - commands::provider_pool_cmd::poll_kiro_builder_id_auth, - commands::provider_pool_cmd::cancel_kiro_builder_id_login, - commands::provider_pool_cmd::add_kiro_from_builder_id_auth, - // Kiro Social Auth 登录命令 (Google/GitHub) - commands::provider_pool_cmd::start_kiro_social_auth_login, - commands::provider_pool_cmd::exchange_kiro_social_auth_token, - commands::provider_pool_cmd::cancel_kiro_social_auth_login, - commands::provider_pool_cmd::start_kiro_social_auth_callback_server, - // Playwright 指纹浏览器登录命令 - commands::provider_pool_cmd::check_playwright_available, - commands::provider_pool_cmd::install_playwright, - commands::provider_pool_cmd::start_kiro_playwright_login, - commands::provider_pool_cmd::cancel_kiro_playwright_login, commands::browser_runtime_cmd::open_browser_runtime_debugger_window, commands::browser_runtime_cmd::close_browser_runtime_debugger_window, commands::browser_runtime_cmd::launch_browser_session, @@ -1370,8 +1207,6 @@ pub fn run() { commands::injection_cmd::update_injection_rule, // Hint route commands commands::security_perf_cmd::get_hint_routes, - // Usage commands - commands::usage_cmd::get_kiro_usage, // Tray commands commands::tray_cmd::sync_tray_model_shortcuts, // Plugin commands @@ -1413,8 +1248,6 @@ pub fn run() { commands::window_cmd::center_window, commands::window_cmd::toggle_fullscreen, commands::window_cmd::is_fullscreen, - // Auto fix commands - commands::auto_fix_cmd::auto_fix_configuration, // Machine ID commands commands::machine_id_cmd::get_current_machine_id, commands::machine_id_cmd::set_machine_id, @@ -1433,10 +1266,6 @@ pub fn run() { commands::machine_id_cmd::paste_machine_id_from_clipboard, commands::machine_id_cmd::get_system_info, commands::windows_startup_cmd::get_windows_startup_diagnostics, - // Kiro Local commands - commands::kiro_local::switch_kiro_to_local, - commands::kiro_local::get_kiro_fingerprint_info, - commands::kiro_local::get_local_kiro_credential_uuid, // Agent commands commands::agent_cmd::agent_start_process, commands::agent_cmd::agent_stop_process, @@ -1447,7 +1276,6 @@ pub fn run() { commands::aster_agent_cmd::command_api::provider_api::aster_agent_status, commands::aster_agent_cmd::command_api::provider_api::aster_agent_reset, commands::aster_agent_cmd::command_api::provider_api::aster_agent_configure_provider, - commands::aster_agent_cmd::command_api::provider_api::aster_agent_configure_from_pool, commands::aster_agent_cmd::command_api::runtime_api::agent_runtime_submit_turn, commands::aster_agent_cmd::command_api::runtime_api::agent_runtime_interrupt_turn, commands::aster_agent_cmd::command_api::runtime_api::agent_runtime_compact_session, @@ -1489,25 +1317,6 @@ pub fn run() { commands::models_cmd::toggle_model_enabled, commands::models_cmd::add_provider, commands::models_cmd::remove_provider, - // Orchestrator commands - commands::orchestrator_cmd::init_orchestrator, - commands::orchestrator_cmd::get_orchestrator_config, - commands::orchestrator_cmd::update_orchestrator_config, - commands::orchestrator_cmd::get_pool_stats, - commands::orchestrator_cmd::get_tier_models, - commands::orchestrator_cmd::get_all_models, - commands::orchestrator_cmd::update_orchestrator_credentials, - commands::orchestrator_cmd::add_orchestrator_credential, - commands::orchestrator_cmd::remove_orchestrator_credential, - commands::orchestrator_cmd::mark_credential_unhealthy, - commands::orchestrator_cmd::mark_credential_healthy, - commands::orchestrator_cmd::update_credential_load, - commands::orchestrator_cmd::select_model, - commands::orchestrator_cmd::quick_select_model, - commands::orchestrator_cmd::select_model_for_task, - commands::orchestrator_cmd::list_strategies, - commands::orchestrator_cmd::list_service_tiers, - commands::orchestrator_cmd::list_task_hints, // Connect commands // _Requirements: 1.4, 2.3, 4.1, 5.3_ commands::connect_cmd::handle_deep_link, @@ -1535,13 +1344,6 @@ pub fn run() { commands::model_registry_cmd::get_all_alias_configs, commands::model_registry_cmd::fetch_provider_models_from_api, commands::model_registry_cmd::fetch_provider_models_auto, - // Model Management commands (动态模型列表) - commands::model_cmd::get_credential_models, - commands::model_cmd::refresh_credential_models, - commands::model_cmd::get_all_models_by_provider, - commands::model_cmd::get_all_available_models, - commands::model_cmd::refresh_all_credential_models, - commands::model_cmd::get_default_models_for_provider, // WebSocket commands commands::websocket_cmd::get_websocket_status, commands::websocket_cmd::get_websocket_connections, diff --git a/src-tauri/src/app/state.rs b/src-tauri/src/app/state.rs index 3a685be21..2bde9e5d3 100644 --- a/src-tauri/src/app/state.rs +++ b/src-tauri/src/app/state.rs @@ -8,10 +8,8 @@ use tokio::sync::RwLock; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; use crate::commands::context_memory::ContextMemoryServiceState; use crate::commands::machine_id_cmd::MachineIdState; -use crate::commands::orchestrator_cmd::OrchestratorState; use crate::commands::plugin_cmd::PluginManagerState; use crate::commands::plugin_install_cmd::PluginInstallerState; -use crate::commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState}; use crate::commands::skill_cmd::SkillServiceState; use crate::config::{GlobalConfigManager, GlobalConfigManagerState}; use crate::database; @@ -20,11 +18,9 @@ use crate::telemetry; use lime_core::config::{Config, ConfigManager}; use lime_services::api_key_provider_service::ApiKeyProviderService; use lime_services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; -use lime_services::provider_pool_service::ProviderPoolService; use lime_services::skill_service::SkillService; -use lime_services::token_cache_service::TokenCacheService; -use super::types::{AppState, LogState, TokenCacheServiceState}; +use super::types::{AppState, LogState}; use crate::logger; use lime_server as server; @@ -47,14 +43,10 @@ pub fn init_global_config_manager(config: &Config) -> GlobalConfigManagerState { /// 初始化服务状态 pub struct ServiceStates { pub skill_service: SkillServiceState, - pub provider_pool_service: ProviderPoolServiceState, pub api_key_provider_service: ApiKeyProviderServiceState, - pub credential_sync_service: CredentialSyncServiceState, - pub token_cache_service: TokenCacheServiceState, pub machine_id_service: MachineIdState, pub plugin_manager: PluginManagerState, pub plugin_installer: PluginInstallerState, - pub orchestrator: OrchestratorState, pub context_memory_service: ContextMemoryServiceState, } @@ -64,22 +56,11 @@ pub fn init_service_states() -> ServiceStates { let skill_service = SkillService::new().expect("Failed to initialize SkillService"); let skill_service_state = SkillServiceState(Arc::new(skill_service)); - // Initialize ProviderPoolService - let provider_pool_service = ProviderPoolService::new(); - let provider_pool_service_state = ProviderPoolServiceState(Arc::new(provider_pool_service)); - // Initialize ApiKeyProviderService let api_key_provider_service = ApiKeyProviderService::new(); let api_key_provider_service_state = ApiKeyProviderServiceState(Arc::new(api_key_provider_service)); - // Initialize CredentialSyncService (optional) - let credential_sync_service_state = CredentialSyncServiceState(None); - - // Initialize TokenCacheService - let token_cache_service = TokenCacheService::new(); - let token_cache_service_state = TokenCacheServiceState(Arc::new(token_cache_service)); - // Initialize MachineIdService let machine_id_service = lime_services::machine_id_service::MachineIdService::new() .expect("Failed to initialize MachineIdService"); @@ -92,9 +73,6 @@ pub fn init_service_states() -> ServiceStates { // Initialize PluginInstaller let plugin_installer_state = init_plugin_installer(); - // Initialize Orchestrator State - let orchestrator_state = OrchestratorState::new(); - // Initialize ContextMemoryService let app_config = lime_core::config::load_config().unwrap_or_default(); let context_memory_config = build_context_memory_config(&app_config); @@ -104,14 +82,10 @@ pub fn init_service_states() -> ServiceStates { ServiceStates { skill_service: skill_service_state, - provider_pool_service: provider_pool_service_state, api_key_provider_service: api_key_provider_service_state, - credential_sync_service: credential_sync_service_state, - token_cache_service: token_cache_service_state, machine_id_service: machine_id_service_state, plugin_manager: plugin_manager_state, plugin_installer: plugin_installer_state, - orchestrator: orchestrator_state, context_memory_service: context_memory_service_state, } } diff --git a/src-tauri/src/app/types.rs b/src-tauri/src/app/types.rs index 2e7058141..eb6e85039 100644 --- a/src-tauri/src/app/types.rs +++ b/src-tauri/src/app/types.rs @@ -9,7 +9,6 @@ use tokio::sync::RwLock; use crate::logger; use crate::tray::TrayManager; use lime_server as server; -use lime_services::token_cache_service::TokenCacheService; use lime_core::event_emit::EventEmit; @@ -22,9 +21,6 @@ pub type AppState = Arc>; /// 日志状态类型别名 pub type LogState = Arc>; -/// TokenCacheService 状态封装 -pub struct TokenCacheServiceState(pub Arc); - /// TrayManager 状态封装 pub struct TrayManagerState(pub Arc>>>); diff --git a/src-tauri/src/commands/aster_agent_cmd/command_api.rs b/src-tauri/src/commands/aster_agent_cmd/command_api.rs index 97a66082b..38a738ac4 100644 --- a/src-tauri/src/commands/aster_agent_cmd/command_api.rs +++ b/src-tauri/src/commands/aster_agent_cmd/command_api.rs @@ -64,8 +64,7 @@ fn build_subagent_control_runtime( } pub(crate) use provider_api::{ - aster_agent_configure_from_pool, aster_agent_configure_provider, aster_agent_init, - aster_agent_reset, aster_agent_status, + aster_agent_configure_provider, aster_agent_init, aster_agent_reset, aster_agent_status, }; pub(crate) use runtime_api::{ agent_runtime_compact_session, agent_runtime_diff_file_checkpoint, diff --git a/src-tauri/src/commands/aster_agent_cmd/command_api/provider_api.rs b/src-tauri/src/commands/aster_agent_cmd/command_api/provider_api.rs index 7cf989a8d..d9239615f 100644 --- a/src-tauri/src/commands/aster_agent_cmd/command_api/provider_api.rs +++ b/src-tauri/src/commands/aster_agent_cmd/command_api/provider_api.rs @@ -74,7 +74,6 @@ pub async fn aster_agent_configure_provider( base_url: request.base_url, credential_uuid: None, force_responses_api: false, - credential_path: None, toolshim: matches!( request.tool_call_strategy, Some(RuntimeToolCallStrategy::ToolShim) @@ -104,42 +103,6 @@ pub async fn aster_agent_configure_provider( }) } -/// 从凭证池配置 Aster Agent 的 Provider -/// -/// 自动从 Lime 凭证池选择可用凭证并配置 Aster Provider -#[tauri::command] -pub async fn aster_agent_configure_from_pool( - state: State<'_, AsterAgentState>, - db: State<'_, DbConnection>, - request: ConfigureFromPoolRequest, - session_id: String, -) -> Result { - tracing::info!( - "[AsterAgent] 从凭证池配置 Provider: {} / {}", - request.provider_type, - request.model_name - ); - - let aster_config = state - .configure_provider_from_pool( - &db, - &request.provider_type, - &request.model_name, - &session_id, - ) - .await?; - persist_session_provider_routing(&session_id, &request.provider_type).await?; - - Ok(AsterAgentStatus { - initialized: true, - provider_configured: true, - provider_name: Some(aster_config.provider_name), - provider_selector: aster_config.provider_selector, - model_name: Some(aster_config.model_name), - credential_uuid: Some(aster_config.credential_uuid), - }) -} - /// 获取 Aster Agent 状态 #[tauri::command] pub async fn aster_agent_status( diff --git a/src-tauri/src/commands/aster_agent_cmd/dto.rs b/src-tauri/src/commands/aster_agent_cmd/dto.rs index c7f58284d..f28873657 100644 --- a/src-tauri/src/commands/aster_agent_cmd/dto.rs +++ b/src-tauri/src/commands/aster_agent_cmd/dto.rs @@ -12,7 +12,7 @@ pub struct AsterAgentStatus { #[serde(skip_serializing_if = "Option::is_none")] pub provider_selector: Option, pub model_name: Option, - /// 凭证 UUID(来自凭证池) + /// 凭证 UUID(来自 API Key Provider) #[serde(skip_serializing_if = "Option::is_none")] pub credential_uuid: Option, } @@ -36,15 +36,6 @@ pub struct ConfigureProviderRequest { pub toolshim_model: Option, } -/// 从凭证池配置 Provider 的请求 -#[derive(Debug, Deserialize)] -pub struct ConfigureFromPoolRequest { - /// Provider 类型 (openai, anthropic, kiro, gemini 等) - pub provider_type: String, - /// 模型名称 - pub model_name: String, -} - #[derive(Debug, Default, Deserialize)] #[serde(rename_all = "camelCase")] pub struct AgentRuntimeToolInventoryRequest { diff --git a/src-tauri/src/commands/aster_agent_cmd/mod.rs b/src-tauri/src/commands/aster_agent_cmd/mod.rs index b0a4990f0..dd101c9bd 100644 --- a/src-tauri/src/commands/aster_agent_cmd/mod.rs +++ b/src-tauri/src/commands/aster_agent_cmd/mod.rs @@ -77,7 +77,8 @@ use aster::session::{SessionType, SubagentSessionMetadata}; use aster::tools::task_output_tool::TaskOutputInput; use aster::tools::{ BashTool, PermissionBehavior, PermissionCheckResult, TaskManager, TaskOutputTool, TaskStopTool, - Tool, ToolContext, ToolError, ToolOptions, ToolResult, MAX_OUTPUT_LENGTH, + Tool, ToolContext, ToolError, ToolOptions, ToolResult, WebFetchTool, WebSearchTool, + MAX_OUTPUT_LENGTH, }; use async_trait::async_trait; use futures::{FutureExt, StreamExt}; @@ -346,8 +347,8 @@ pub(crate) use command_api::{ agent_runtime_replay_request, agent_runtime_resume_subagent, agent_runtime_resume_thread, agent_runtime_save_review_decision, agent_runtime_send_subagent_input, agent_runtime_spawn_subagent, agent_runtime_submit_turn, agent_runtime_update_session, - agent_runtime_wait_subagents, aster_agent_configure_from_pool, aster_agent_configure_provider, - aster_agent_init, aster_agent_reset, aster_agent_status, + agent_runtime_wait_subagents, aster_agent_configure_provider, aster_agent_init, + aster_agent_reset, aster_agent_status, }; pub(crate) use cover_skill_launch::{ append_cover_skill_launch_session_permissions, merge_system_prompt_with_cover_skill_launch, @@ -381,7 +382,7 @@ pub(crate) use dto::{ AgentRuntimeSubmitTurnRequest, AgentRuntimeThreadDiagnostics, AgentRuntimeThreadReadModel, AgentRuntimeToolInventoryRequest, AgentRuntimeUpdateSessionRequest, AgentRuntimeWaitSubagentsRequest, AgentRuntimeWaitSubagentsResponse, AsterAgentStatus, - AsterChatRequest, AutoContinuePayload, ConfigureFromPoolRequest, ConfigureProviderRequest, + AsterChatRequest, AutoContinuePayload, ConfigureProviderRequest, }; pub(crate) use form_skill_launch::{ append_form_skill_launch_session_permissions, merge_system_prompt_with_form_skill_launch, @@ -496,8 +497,9 @@ pub(crate) use tool_runtime::social_generate_cover_image_cmd; #[cfg(test)] #[allow(unused_imports)] pub(crate) use tool_runtime::{ - append_subagent_tool_scope_session_permissions, extract_runtime_subagent_result_text, - LimeBrowserMcpTool, SocialGenerateCoverImageTool, ToolSearchBridgeTool, + append_subagent_tool_scope_session_permissions, ensure_default_web_tools_registered, + extract_runtime_subagent_result_text, LimeBrowserMcpTool, SocialGenerateCoverImageTool, + ToolSearchBridgeTool, }; pub(crate) use tool_runtime::{apply_workspace_sandbox_permissions, ImageInput}; pub(crate) use tool_runtime::{ diff --git a/src-tauri/src/commands/aster_agent_cmd/request_model_resolution.rs b/src-tauri/src/commands/aster_agent_cmd/request_model_resolution.rs index c1af9ff4d..20688f9b0 100644 --- a/src-tauri/src/commands/aster_agent_cmd/request_model_resolution.rs +++ b/src-tauri/src/commands/aster_agent_cmd/request_model_resolution.rs @@ -31,7 +31,7 @@ struct ProviderResolutionContext { #[derive(Debug, Clone, PartialEq, Eq)] enum RuntimeProviderConfigurationStrategy { Manual { base_url: Option }, - CredentialPool, + ApiKeyProvider, } #[derive(Debug, Clone)] @@ -95,16 +95,14 @@ fn provider_alias_config_key(provider_key: &str) -> String { fn provider_registry_id_from_key(provider_key: &str) -> String { match normalize_identifier(provider_key).as_str() { "openai" => "openai".to_string(), - "anthropic" | "anthropic-compatible" | "claude" | "claude_oauth" => "anthropic".to_string(), + "anthropic" | "anthropic-compatible" | "claude" => "anthropic".to_string(), "gemini" | "gemini_api_key" => "gemini".to_string(), "azure-openai" => "openai".to_string(), "vertexai" => "google".to_string(), "ollama" => "ollama".to_string(), "fal" => "fal".to_string(), - "kiro" => "kiro".to_string(), "qwen" => "alibaba".to_string(), "codex" => "codex".to_string(), - "antigravity" => "antigravity".to_string(), "iflow" => "openai".to_string(), normalized => normalized.to_string(), } @@ -113,7 +111,7 @@ fn provider_registry_id_from_key(provider_key: &str) -> String { fn provider_type_from_key(provider_key: &str) -> Option { match normalize_identifier(provider_key).as_str() { "openai" | "iflow" => Some(ApiProviderType::Openai), - "anthropic" | "claude" | "claude_oauth" => Some(ApiProviderType::Anthropic), + "anthropic" | "claude" => Some(ApiProviderType::Anthropic), "anthropic-compatible" => Some(ApiProviderType::AnthropicCompatible), "gemini" | "gemini_api_key" => Some(ApiProviderType::Gemini), "azure-openai" => Some(ApiProviderType::AzureOpenai), @@ -165,7 +163,7 @@ fn resolve_runtime_provider_configuration_strategy( is_credentialless_local_provider || context.provider_type == Some(ApiProviderType::Ollama); if !should_use_manual_provider { - return RuntimeProviderConfigurationStrategy::CredentialPool; + return RuntimeProviderConfigurationStrategy::ApiKeyProvider; } let fallback_base_url = context.provider_type.map(|provider_type| { @@ -2384,7 +2382,7 @@ async fn build_runtime_request_provider_config_from_preference( let provider_strategy = resolve_runtime_provider_configuration_strategy(&context); let base_url = match provider_strategy { RuntimeProviderConfigurationStrategy::Manual { base_url } => base_url, - RuntimeProviderConfigurationStrategy::CredentialPool => None, + RuntimeProviderConfigurationStrategy::ApiKeyProvider => None, }; let model_meta = find_model_meta(&resolved_model, &catalog); let model_capabilities = model_meta @@ -2472,7 +2470,7 @@ pub(super) async fn resolve_runtime_provider_auth_recovery_config( let provider_strategy = resolve_runtime_provider_configuration_strategy(&context); let base_url = match provider_strategy { RuntimeProviderConfigurationStrategy::Manual { base_url } => base_url, - RuntimeProviderConfigurationStrategy::CredentialPool => None, + RuntimeProviderConfigurationStrategy::ApiKeyProvider => None, }; let model_meta = find_model_meta(&fallback_model, &catalog); let model_capabilities = model_meta diff --git a/src-tauri/src/commands/aster_agent_cmd/runtime_turn.rs b/src-tauri/src/commands/aster_agent_cmd/runtime_turn.rs index 81b46a998..8fb82c108 100644 --- a/src-tauri/src/commands/aster_agent_cmd/runtime_turn.rs +++ b/src-tauri/src/commands/aster_agent_cmd/runtime_turn.rs @@ -548,7 +548,7 @@ fn spawn_runtime_memory_capture_task( #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ProviderConfigApplyMode { Direct, - CredentialPool, + ApiKeyProvider, } fn normalize_provider_identity(value: &str) -> String { @@ -573,7 +573,7 @@ fn resolve_provider_config_apply_mode( return ProviderConfigApplyMode::Direct; } - ProviderConfigApplyMode::CredentialPool + ProviderConfigApplyMode::ApiKeyProvider } async fn apply_runtime_turn_provider_config( @@ -606,7 +606,6 @@ async fn apply_runtime_turn_provider_config( base_url: provider_config.base_url.clone(), credential_uuid: None, force_responses_api: false, - credential_path: None, toolshim: matches!( provider_config.tool_call_strategy, Some(RuntimeToolCallStrategy::ToolShim) @@ -629,7 +628,7 @@ async fn apply_runtime_turn_provider_config( ProviderConfigApplyMode::Direct => { state.configure_provider(config, session_id, db).await?; } - ProviderConfigApplyMode::CredentialPool => { + ProviderConfigApplyMode::ApiKeyProvider => { state .configure_provider_from_pool( db, diff --git a/src-tauri/src/commands/aster_agent_cmd/tests.rs b/src-tauri/src/commands/aster_agent_cmd/tests.rs index 0a2f01f61..00b2567fb 100644 --- a/src-tauri/src/commands/aster_agent_cmd/tests.rs +++ b/src-tauri/src/commands/aster_agent_cmd/tests.rs @@ -948,6 +948,41 @@ mod tests { assert!(allowed_registry.contains("WebSearch")); } + #[test] + fn test_default_web_tools_are_restored_after_fast_chat_prune() { + let disabled_policy = resolve_request_tool_policy(Some(false), false); + let allowed_policy = resolve_request_tool_policy(Some(true), false); + let mut registry = aster::tools::ToolRegistry::new(); + registry.register(Box::new(DummyTool::new( + "WebSearch", + "Web search", + serde_json::json!({"type": "object"}), + ))); + registry.register(Box::new(DummyTool::new( + "WebFetch", + "Web fetch", + serde_json::json!({"type": "object"}), + ))); + + prune_fast_chat_request_tool_policy_tools_from_registry( + &mut registry, + TurnExecutionProfile::FastChat, + &disabled_policy, + ); + assert!(!registry.contains("WebSearch")); + assert!(!registry.contains("WebFetch")); + + ensure_default_web_tools_registered(&mut registry); + prune_fast_chat_request_tool_policy_tools_from_registry( + &mut registry, + TurnExecutionProfile::FastChat, + &allowed_policy, + ); + + assert!(registry.contains("WebSearch")); + assert!(registry.contains("WebFetch")); + } + #[test] fn test_build_service_skill_launch_run_request_requires_attached_session() { let metadata = serde_json::json!({ diff --git a/src-tauri/src/commands/aster_agent_cmd/tool_runtime.rs b/src-tauri/src/commands/aster_agent_cmd/tool_runtime.rs index e8e00c116..65cb98fd1 100644 --- a/src-tauri/src/commands/aster_agent_cmd/tool_runtime.rs +++ b/src-tauri/src/commands/aster_agent_cmd/tool_runtime.rs @@ -108,6 +108,15 @@ fn unregister_named_tools(registry: &mut aster::tools::ToolRegistry, tool_names: } } +pub(crate) fn ensure_default_web_tools_registered(registry: &mut aster::tools::ToolRegistry) { + if !registry.contains("WebFetch") { + registry.register(Box::new(WebFetchTool::new())); + } + if !registry.contains("WebSearch") { + registry.register(Box::new(WebSearchTool::new())); + } +} + const FAST_CHAT_DISABLED_WEB_TOOL_PATTERNS: &[&str] = &["WebSearch", "web_search", "WebFetch", "web_fetch"]; const SUBAGENT_TOOL_SCOPE_DEFAULT_DENY_PRIORITY: i32 = 1298; @@ -517,6 +526,7 @@ pub(crate) async fn apply_workspace_sandbox_permissions( app_handle.clone(), config_manager.0.clone(), ); + ensure_default_web_tools_registered(&mut registry); workspace_tools::wrap_registry_native_tools_for_workspace_runtime(&mut registry); prune_fast_chat_request_tool_policy_tools_from_registry( &mut registry, diff --git a/src-tauri/src/commands/auto_fix_cmd.rs b/src-tauri/src/commands/auto_fix_cmd.rs deleted file mode 100644 index 774fe1300..000000000 --- a/src-tauri/src/commands/auto_fix_cmd.rs +++ /dev/null @@ -1,263 +0,0 @@ -//! 自动修复命令 -//! -//! 提供自动检测和修复常见配置问题的功能 - -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::database::DbConnection; -use crate::models::provider_pool_model::PoolProviderType; -use crate::{config, AppState, LogState, ProviderType}; -use serde::{Deserialize, Serialize}; -use tauri::State; - -/// 自动修复结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AutoFixResult { - pub issues_found: Vec, - pub fixes_applied: Vec, - pub warnings: Vec, -} - -/// 自动检测并修复配置问题 -#[tauri::command] -pub async fn auto_fix_configuration( - state: State<'_, AppState>, - logs: State<'_, LogState>, - db: State<'_, DbConnection>, -) -> Result { - let mut result = AutoFixResult { - issues_found: Vec::new(), - fixes_applied: Vec::new(), - warnings: Vec::new(), - }; - - logs.write() - .await - .add("info", "[自动修复] 开始检测配置问题..."); - - // 检查默认Provider配置 - if let Err(e) = fix_default_provider_issue(&state, &logs, &db, &mut result).await { - result.warnings.push(format!("修复默认Provider时出错: {e}")); - } - - // 检查凭证池状态 - if let Err(e) = check_credential_pool_issues(&db, &mut result).await { - result.warnings.push(format!("检查凭证池时出错: {e}")); - } - - logs.write().await.add( - "info", - &format!( - "[自动修复] 完成,发现 {} 个问题,修复 {} 个", - result.issues_found.len(), - result.fixes_applied.len() - ), - ); - - Ok(result) -} - -/// 修复默认Provider配置问题 -async fn fix_default_provider_issue( - state: &State<'_, AppState>, - logs: &State<'_, LogState>, - db: &State<'_, DbConnection>, - result: &mut AutoFixResult, -) -> Result<(), String> { - let current_default = { - let s = state.read().await; - s.config.default_provider.clone() - }; - - // 获取可用的凭证类型统计 - let credential_stats = get_credential_stats(db).await?; - - // 检查是否有Kiro凭证但默认Provider不是kiro - if credential_stats.kiro_count > 0 && current_default != "kiro" { - result.issues_found.push(format!( - "默认Provider设置为 '{}' 但有 {} 个Kiro凭证可用", - current_default, credential_stats.kiro_count - )); - - // 自动修复:设置默认Provider为kiro - if let Err(e) = set_default_provider_internal(state, logs, "kiro".to_string()).await { - result - .warnings - .push(format!("无法自动修复默认Provider: {e}")); - } else { - result - .fixes_applied - .push("默认Provider已自动设置为 'kiro'".to_string()); - logs.write() - .await - .add("info", "[自动修复] 默认Provider已设置为kiro"); - } - } - // 检查是否默认Provider指向的凭证类型不可用 - else if !is_provider_available(¤t_default, &credential_stats) { - result - .issues_found - .push(format!("默认Provider '{current_default}' 没有可用凭证")); - - // 寻找最佳替代Provider - if let Some(best_provider) = find_best_available_provider(&credential_stats) { - if let Err(e) = set_default_provider_internal(state, logs, best_provider.clone()).await - { - result - .warnings - .push(format!("无法自动修复默认Provider: {e}")); - } else { - result - .fixes_applied - .push(format!("默认Provider已自动设置为 '{best_provider}'")); - logs.write().await.add( - "info", - &format!("[自动修复] 默认Provider已设置为{best_provider}"), - ); - } - } else { - result - .warnings - .push("没有找到可用的Provider作为默认选择".to_string()); - } - } - - Ok(()) -} - -/// 获取凭证统计信息 -#[derive(Debug, Default)] -struct CredentialStats { - kiro_count: usize, - gemini_count: usize, - openai_count: usize, - claude_count: usize, - total_count: usize, -} - -async fn get_credential_stats(db: &State<'_, DbConnection>) -> Result { - let conn = db.lock().map_err(|e| e.to_string())?; - let mut stats = CredentialStats::default(); - - // 统计各类型凭证数量(只计算启用且健康的凭证) - let all_credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - - for cred in all_credentials - .iter() - .filter(|c| !c.is_disabled && c.is_healthy) - { - match cred.provider_type { - PoolProviderType::Kiro => stats.kiro_count += 1, - PoolProviderType::Gemini => stats.gemini_count += 1, - PoolProviderType::OpenAI => stats.openai_count += 1, - PoolProviderType::Claude => stats.claude_count += 1, - _ => {} - } - stats.total_count += 1; - } - - tracing::info!( - "[自动修复] 凭证统计: kiro={}, gemini={}, claude={}, openai={}, total={}", - stats.kiro_count, - stats.gemini_count, - stats.claude_count, - stats.openai_count, - stats.total_count - ); - - Ok(stats) -} - -/// 检查Provider是否有可用凭证 -fn is_provider_available(provider: &str, stats: &CredentialStats) -> bool { - match provider { - "kiro" => stats.kiro_count > 0, - "gemini" => stats.gemini_count > 0, - "openai" => stats.openai_count > 0, - "claude" => stats.claude_count > 0, - _ => false, - } -} - -/// 寻找最佳可用Provider -fn find_best_available_provider(stats: &CredentialStats) -> Option { - // 优先级:kiro > gemini > claude > openai - if stats.kiro_count > 0 { - Some("kiro".to_string()) - } else if stats.gemini_count > 0 { - Some("gemini".to_string()) - } else if stats.claude_count > 0 { - Some("claude".to_string()) - } else if stats.openai_count > 0 { - Some("openai".to_string()) - } else { - None - } -} - -/// 内部设置默认Provider函数 -async fn set_default_provider_internal( - state: &State<'_, AppState>, - _logs: &State<'_, LogState>, - provider: String, -) -> Result<(), String> { - // 验证provider - let provider_type: ProviderType = provider.parse().map_err(|e: String| e)?; - - let mut s = state.write().await; - s.config.default_provider = provider.clone(); - - // 同时更新运行中服务器的 default_provider_ref - { - let mut dp = s.default_provider_ref.write().await; - *dp = provider.clone(); - } - - // 同时更新运行中服务器的 router(如果服务器正在运行) - if let Some(router_ref) = &s.router_ref { - let mut router = router_ref.write().await; - router.set_default_provider(provider_type); - tracing::info!("[AUTO_FIX] 动态更新 Router 默认 Provider: {}", provider); - } - - config::save_config(&s.config).map_err(|e| e.to_string())?; - - Ok(()) -} - -/// 检查凭证池问题 -async fn check_credential_pool_issues( - db: &State<'_, DbConnection>, - result: &mut AutoFixResult, -) -> Result<(), String> { - let conn = db.lock().map_err(|e| e.to_string())?; - let credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - - // 检查是否有过期的token缓存 - let mut expired_tokens = 0; - for cred in &credentials { - if let Some(ref token_info) = cred.cached_token { - if let Some(expiry) = token_info.expiry_time { - if chrono::Utc::now() > expiry { - expired_tokens += 1; - } - } - } - } - - if expired_tokens > 0 { - result - .issues_found - .push(format!("发现 {expired_tokens} 个过期的token缓存")); - // 过期token会在使用时自动刷新,这里只是报告 - } - - // 检查是否有禁用的凭证 - let disabled_count = credentials.iter().filter(|c| c.is_disabled).count(); - if disabled_count > 0 { - result - .issues_found - .push(format!("有 {disabled_count} 个凭证被禁用")); - } - - Ok(()) -} diff --git a/src-tauri/src/commands/kiro_local.rs b/src-tauri/src/commands/kiro_local.rs deleted file mode 100644 index 783900453..000000000 --- a/src-tauri/src/commands/kiro_local.rs +++ /dev/null @@ -1,393 +0,0 @@ -//! Kiro 凭证本地切换命令 -//! -//! 将 Kiro 凭证切换到本地 IDE,同时切换设备指纹。 - -use crate::commands::provider_pool_cmd::ProviderPoolServiceState; -use crate::database::DbConnection; -use crate::models::kiro_fingerprint::{KiroFingerprintStore, SwitchToLocalResult}; -use crate::models::provider_pool_model::CredentialData; -use lime_services::machine_id_service::MachineIdService; -use serde::{Deserialize, Serialize}; -use sha2::{Digest, Sha256}; -use std::fs; -use std::path::PathBuf; -use tauri::State; - -/// Kiro auth token 文件格式 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -struct KiroAuthToken { - access_token: String, - refresh_token: String, - expires_at: String, - client_id_hash: String, - auth_method: String, - provider: String, - region: String, - #[serde(skip_serializing_if = "Option::is_none")] - client_id: Option, - #[serde(skip_serializing_if = "Option::is_none")] - client_secret: Option, -} - -/// 客户端注册文件格式 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -struct ClientRegistration { - client_id: String, - client_secret: String, - expires_at: String, - scopes: Vec, -} - -/// 获取 AWS SSO cache 目录 -fn get_aws_sso_cache_dir() -> Result { - let home = dirs::home_dir().ok_or_else(|| "无法获取用户主目录".to_string())?; - let cache_dir = home.join(".aws").join("sso").join("cache"); - - // 确保目录存在 - if !cache_dir.exists() { - fs::create_dir_all(&cache_dir).map_err(|e| format!("创建 AWS SSO cache 目录失败: {e}"))?; - } - - Ok(cache_dir) -} - -/// 计算 clientIdHash(备用方案,使用 SHA256 的前 40 位模拟 SHA1 格式) -fn calculate_client_id_hash() -> String { - let start_url = "https://view.awsapps.com/start"; - let json_str = format!("{{\"startUrl\":\"{start_url}\"}}"); - - let mut hasher = Sha256::new(); - hasher.update(json_str.as_bytes()); - let result = hasher.finalize(); - - // SHA1 是 40 位十六进制,取 SHA256 的前 20 字节(40 位十六进制) - format!("{result:x}")[..40].to_string() -} - -/// 切换 Kiro 凭证到本地 -/// -/// 1. 从凭证池读取指定凭证 -/// 2. 获取/生成绑定的 Machine ID -/// 3. 切换系统机器码 -/// 4. 写入凭证到 ~/.aws/sso/cache/kiro-auth-token.json -#[tauri::command] -pub async fn switch_kiro_to_local( - uuid: String, - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, -) -> Result { - tracing::info!("[KIRO_LOCAL] 开始切换凭证到本地: {}", uuid); - - // 1. 获取凭证信息 - let credential = pool_service - .0 - .get_by_uuid(&db, &uuid) - .map_err(|e| format!("获取凭证失败: {e}"))? - .ok_or_else(|| format!("找不到凭证: {uuid}"))?; - - // 检查是否为 Kiro 凭证 - let creds_file_path = match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(), - _ => return Err("只支持 Kiro OAuth 凭证".to_string()), - }; - - // 2. 读取凭证文件 - let creds_content = - fs::read_to_string(&creds_file_path).map_err(|e| format!("读取凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&creds_content).map_err(|e| format!("解析凭证文件失败: {e}"))?; - - // 3. 获取/生成绑定的 Machine ID - let mut fingerprint_store = - KiroFingerprintStore::load().map_err(|e| format!("加载指纹存储失败: {e}"))?; - - let profile_arn = creds.get("profileArn").and_then(|v| v.as_str()); - let client_id = creds.get("clientId").and_then(|v| v.as_str()); - - let binding = fingerprint_store - .get_or_create_binding(&uuid, profile_arn, client_id) - .map_err(|e| format!("获取指纹绑定失败: {e}"))?; - - let machine_id = binding.machine_id.clone(); - tracing::info!("[KIRO_LOCAL] 使用 Machine ID: {}", &machine_id[..8]); - - // 4. 切换系统机器码 - let machine_service = - MachineIdService::new().map_err(|e| format!("初始化机器码服务失败: {e}"))?; - - let machine_result = machine_service - .set_machine_id(&machine_id) - .await - .map_err(|e| format!("切换机器码失败: {e}"))?; - - if !machine_result.success { - if machine_result.requires_admin { - return Ok(SwitchToLocalResult::requires_admin(format!( - "需要管理员权限切换机器码: {}", - machine_result.message - ))); - } - return Ok(SwitchToLocalResult::error(format!( - "切换机器码失败: {}", - machine_result.message - ))); - } - - // 5. 准备 Kiro auth token 数据 - let access_token = creds - .get("accessToken") - .and_then(|v| v.as_str()) - .ok_or_else(|| "凭证文件缺少 accessToken".to_string())?; - - let refresh_token = creds - .get("refreshToken") - .and_then(|v| v.as_str()) - .ok_or_else(|| "凭证文件缺少 refreshToken".to_string())?; - - let expires_at = creds - .get("expiresAt") - .and_then(|v| v.as_str()) - .unwrap_or(""); - - let auth_method = creds - .get("authMethod") - .and_then(|v| v.as_str()) - .unwrap_or("social"); - - let provider = creds - .get("provider") - .and_then(|v| v.as_str()) - .unwrap_or("BuilderId"); - - let region = creds - .get("region") - .and_then(|v| v.as_str()) - .unwrap_or("us-east-1"); - - let client_id_hash = creds - .get("clientIdHash") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()) - .unwrap_or_else(calculate_client_id_hash); - - // 6. 写入 kiro-auth-token.json - let cache_dir = get_aws_sso_cache_dir()?; - let auth_token_path = cache_dir.join("kiro-auth-token.json"); - - let auth_token = KiroAuthToken { - access_token: access_token.to_string(), - refresh_token: refresh_token.to_string(), - expires_at: expires_at.to_string(), - client_id_hash: client_id_hash.clone(), - auth_method: auth_method.to_string(), - provider: provider.to_string(), - region: region.to_string(), - client_id: creds - .get("clientId") - .and_then(|v| v.as_str()) - .map(String::from), - client_secret: creds - .get("clientSecret") - .and_then(|v| v.as_str()) - .map(String::from), - }; - - let auth_token_json = serde_json::to_string_pretty(&auth_token) - .map_err(|e| format!("序列化 auth token 失败: {e}"))?; - - fs::write(&auth_token_path, &auth_token_json) - .map_err(|e| format!("写入 kiro-auth-token.json 失败: {e}"))?; - - tracing::info!("[KIRO_LOCAL] 已写入 kiro-auth-token.json"); - - // 7. 如果是 IdC 认证,写入客户端注册文件 - if auth_method.to_lowercase() == "idc" { - if let (Some(client_id), Some(client_secret)) = ( - creds.get("clientId").and_then(|v| v.as_str()), - creds.get("clientSecret").and_then(|v| v.as_str()), - ) { - let registration = ClientRegistration { - client_id: client_id.to_string(), - client_secret: client_secret.to_string(), - expires_at: expires_at.to_string(), - scopes: vec![ - "codewhisperer:completions".to_string(), - "codewhisperer:analysis".to_string(), - "codewhisperer:conversations".to_string(), - ], - }; - - let registration_path = cache_dir.join(format!("{client_id_hash}.json")); - let registration_json = serde_json::to_string_pretty(®istration) - .map_err(|e| format!("序列化客户端注册信息失败: {e}"))?; - - fs::write(®istration_path, ®istration_json) - .map_err(|e| format!("写入客户端注册文件失败: {e}"))?; - - tracing::info!( - "[KIRO_LOCAL] 已写入客户端注册文件: {}.json", - &client_id_hash[..8] - ); - } - } - - // 8. 更新最后切换时间 - fingerprint_store - .update_last_switched(&uuid) - .map_err(|e| format!("更新切换时间失败: {e}"))?; - - let credential_name = credential - .name - .clone() - .unwrap_or_else(|| uuid[..8].to_string()); - - Ok(SwitchToLocalResult::success( - format!( - "已切换到凭证 \"{}\",机器码: {}...\n请重启 Kiro IDE 使配置生效", - credential_name, - &machine_id[..8] - ), - machine_id, - )) -} - -/// 获取凭证的指纹信息 -#[tauri::command] -pub async fn get_kiro_fingerprint_info( - uuid: String, - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, -) -> Result { - // 获取凭证信息 - let credential = pool_service - .0 - .get_by_uuid(&db, &uuid) - .map_err(|e| format!("获取凭证失败: {e}"))? - .ok_or_else(|| format!("找不到凭证: {uuid}"))?; - - // 检查是否为 Kiro 凭证 - let creds_file_path = match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(), - _ => return Err("只支持 Kiro OAuth 凭证".to_string()), - }; - - // 读取凭证文件 - let creds_content = - fs::read_to_string(&creds_file_path).map_err(|e| format!("读取凭证文件失败: {e}"))?; - let creds: serde_json::Value = - serde_json::from_str(&creds_content).map_err(|e| format!("解析凭证文件失败: {e}"))?; - - // 获取指纹绑定 - let mut fingerprint_store = - KiroFingerprintStore::load().map_err(|e| format!("加载指纹存储失败: {e}"))?; - - let profile_arn = creds.get("profileArn").and_then(|v| v.as_str()); - let client_id = creds.get("clientId").and_then(|v| v.as_str()); - - let binding = fingerprint_store - .get_or_create_binding(&uuid, profile_arn, client_id) - .map_err(|e| format!("获取指纹绑定失败: {e}"))?; - - let auth_method = creds - .get("authMethod") - .and_then(|v| v.as_str()) - .unwrap_or("social"); - - let source = if profile_arn.is_some() { - "profileArn" - } else if client_id.is_some() { - "clientId" - } else { - "uuid" - }; - - Ok(KiroFingerprintInfo { - machine_id: binding.machine_id.clone(), - machine_id_short: binding.machine_id[..8].to_string(), - source: source.to_string(), - auth_method: auth_method.to_string(), - created_at: binding.created_at.to_rfc3339(), - last_switched_at: binding.last_switched_at.map(|t| t.to_rfc3339()), - }) -} - -/// 获取当前本地使用的 Kiro 凭证 UUID -/// -/// 读取 ~/.aws/sso/cache/kiro-auth-token.json,与凭证池中的凭证比较 -#[tauri::command] -pub async fn get_local_kiro_credential_uuid( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, -) -> Result, String> { - // 读取本地 kiro-auth-token.json - let cache_dir = get_aws_sso_cache_dir()?; - let auth_token_path = cache_dir.join("kiro-auth-token.json"); - - if !auth_token_path.exists() { - return Ok(None); - } - - let local_content = - fs::read_to_string(&auth_token_path).map_err(|e| format!("读取本地凭证文件失败: {e}"))?; - let local_creds: serde_json::Value = - serde_json::from_str(&local_content).map_err(|e| format!("解析本地凭证文件失败: {e}"))?; - - let local_access_token = local_creds.get("accessToken").and_then(|v| v.as_str()); - let local_refresh_token = local_creds.get("refreshToken").and_then(|v| v.as_str()); - - if local_access_token.is_none() && local_refresh_token.is_none() { - return Ok(None); - } - - // 获取所有 Kiro 凭证 - let overview = pool_service.0.get_overview(&db)?; - let kiro_pool = overview.iter().find(|p| p.provider_type == "kiro"); - - if let Some(pool) = kiro_pool { - for cred_display in &pool.credentials { - // 读取凭证文件并比较 - if let Ok(Some(cred)) = pool_service.0.get_by_uuid(&db, &cred_display.uuid) { - if let CredentialData::KiroOAuth { creds_file_path } = &cred.credential { - if let Ok(content) = fs::read_to_string(creds_file_path) { - if let Ok(creds) = serde_json::from_str::(&content) { - let access_token = creds.get("accessToken").and_then(|v| v.as_str()); - let refresh_token = creds.get("refreshToken").and_then(|v| v.as_str()); - - // 比较 token - let matches = matches!((local_access_token, access_token), (Some(l), Some(r)) if l == r) - || matches!( - (local_refresh_token, refresh_token), - (Some(l), Some(r)) if l == r - ); - - if matches { - return Ok(Some(cred_display.uuid.clone())); - } - } - } - } - } - } - } - - Ok(None) -} - -/// Kiro 指纹信息(用于前端显示) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct KiroFingerprintInfo { - /// 完整的 Machine ID - pub machine_id: String, - /// 简短的 Machine ID(前8位) - pub machine_id_short: String, - /// 指纹来源(profileArn/clientId/uuid) - pub source: String, - /// 认证方式 - pub auth_method: String, - /// 创建时间 - pub created_at: String, - /// 最后切换时间 - pub last_switched_at: Option, -} diff --git a/src-tauri/src/commands/memory_search_cmd.rs b/src-tauri/src/commands/memory_search_cmd.rs index 88a42e3cf..b8a8a0fdf 100644 --- a/src-tauri/src/commands/memory_search_cmd.rs +++ b/src-tauri/src/commands/memory_search_cmd.rs @@ -9,7 +9,6 @@ use lime_memory::models::{ }; use lime_memory::search; use lime_services::api_key_provider_service::ApiKeyProviderService; -use lime_services::provider_pool_service::ProviderPoolService; use rusqlite::params; use serde::{Deserialize, Serialize}; use serde_json; @@ -119,16 +118,13 @@ pub async fn unified_memory_semantic_search( tracing::info!("[Semantic Search] Query: {}", options.query); - let provider_pool_service = ProviderPoolService::new(); let api_key_service = ApiKeyProviderService::new(); - let credential = match provider_pool_service - .select_credential_with_fallback( + let credential = match api_key_service + .select_credential_for_provider( &db, - &api_key_service, "openai", None::<&str>, - None::<&str>, None::<&lime_core::models::client_type::ClientType>, ) .await @@ -188,18 +184,13 @@ pub async fn unified_memory_hybrid_search( options.semantic_weight ); - // Use provider pool system to get API key - let provider_pool_service = ProviderPoolService::new(); let api_key_service = ApiKeyProviderService::new(); - // Try to get credential from provider pool or fallback to API key provider - let credential = match provider_pool_service - .select_credential_with_fallback( + let credential = match api_key_service + .select_credential_for_provider( &db, - &api_key_service, "openai", None::<&str>, - None::<&str>, None::<&lime_core::models::client_type::ClientType>, ) .await @@ -228,7 +219,7 @@ pub async fn unified_memory_hybrid_search( } }; - tracing::debug!("[Hybrid Search] Using API key from provider pool"); + tracing::debug!("[Hybrid Search] Using API key from API Key Provider"); // Get query embedding let query_embedding = lime_embedding::get_embedding(&options.query, &api_key, None) diff --git a/src-tauri/src/commands/memory_search_cmd.rs.bak b/src-tauri/src/commands/memory_search_cmd.rs.bak deleted file mode 100644 index 12209c94e..000000000 --- a/src-tauri/src/commands/memory_search_cmd.rs.bak +++ /dev/null @@ -1,219 +0,0 @@ -//! Memory search commands -//! -//! Provides Tauri commands for semantic and hybrid search - -use crate::database::DbConnection; -use lime_memory::search; -use lime_memory::models::{UnifiedMemory, MemoryCategory}; -use lime_services::provider_pool_service::ProviderPoolService; -use lime_services::api_key_provider_service::ApiKeyProviderService; -use serde::{Deserialize, Serialize}; -use tauri::State; - -// ==================== Request Types ==================== - -/// Semantic search options -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SemanticSearchOptions { - /// Query text - pub query: String, - /// Category filter (optional) - pub category: Option, - /// Minimum similarity threshold (0.0-1.0, default 0.5) - pub min_similarity: f32, - /// Result limit (optional, default 50) - pub limit: Option, -} - -impl SemanticSearchOptions { - pub fn with_defaults(mut self) -> Self { - if self.min_similarity == 0.0 { - self.min_similarity = 0.5; - } - self - } -} - -/// Hybrid search options -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HybridSearchOptions { - /// Query text - pub query: String, - /// Category filter (optional) - pub category: Option, - /// Semantic search weight (0.0-1.0, default 0.6) - pub semantic_weight: f32, - /// Keyword search weight (automatically calculated as 1.0 - semantic_weight) - /// Minimum similarity threshold - pub min_similarity: f32, - /// Result limit (optional, default 50) - pub limit: Option, -} - -impl HybridSearchOptions { - pub fn with_defaults(mut self) -> Self { - if self.semantic_weight == 0.0 { - self.semantic_weight = 0.6; - } - if self.min_similarity == 0.0 { - self.min_similarity = 0.5; - } - self - } -} - -// ==================== Commands ==================== - -/// Semantic search (vector similarity) -#[tauri::command] -pub async fn unified_memory_semantic_search( - db: State<'_, DbConnection>, - options: SemanticSearchOptions, -) -> Result, String> { - let options = options.with_defaults(); - - tracing::info!( - "[Semantic Search] Query: {}, category: {:?}", - options.query, - options.category - ); - - // Use provider pool system to get API key - let provider_pool_service = ProviderPoolService::new(); - let api_key_service = ApiKeyProviderService::new(); - - // Try to get credential from provider pool or fallback to API key provider - let credential = match provider_pool_service - .select_credential_with_fallback( - &db, - &api_key_service, - "openai", - None::<&str>, - None::<&str>, - None::<&lime_core::models::client_type::ClientType>, - ) - .await - { - Ok(Some(cred)) => cred, - Ok(None) => { - return Err(String::from( - "没有可用的 OpenAI 凭证。请在设置中添加 OpenAI API Key。" - )); - } - Err(e) => return Err(format!("获取凭证失败: {}", e)), - }; - - // Extract API key from credential - let api_key = match credential.credential { - lime_core::models::provider_pool_model::CredentialData::OpenAIKey { - api_key, - .. - } => api_key, - lime_core::models::provider_pool_model::CredentialData::AnthropicKey { - api_key, - .. - } => api_key, - _ => { - return Err(String::from( - "语义搜索需要 OpenAI API Key 凭证。" - )); - } - }; - - tracing::debug!("[Semantic Search] Using API key from provider pool"); - - // Get query embedding - let query_embedding = lime_embedding::get_embedding(&options.query, &api_key, None).await - .map_err(|e| format!("Failed to get embedding: {}", e))?; - - // Execute semantic search - let results = { - let conn = db.lock().unwrap(); - search::semantic_search( - &*conn, - &query_embedding, - options.category.as_ref(), - options.min_similarity, - ) - .map_err(|e| format!("Semantic search failed: {}", e).to_string()) - }?; - - tracing::info!("[Semantic Search] Returning {} results", results.len()); - - Ok(results) -} - -/// Hybrid search (semantic + keyword) -#[tauri::command] -pub async fn unified_memory_hybrid_search( - -/// Hybrid search (semantic + keyword) -#[tauri::command] -pub async fn unified_memory_hybrid_search( - db: State<'_, DbConnection>, - options: HybridSearchOptions, -) -> Result, String> { - let options = options.with_defaults(); - - tracing::info!( - "[Hybrid Search] Query: {}, semantic_weight: {}", - options.query, - options.semantic_weight - ); - - // Use provider pool system to get API key - let provider_pool_service = ProviderPoolService::new(); - let api_key_service = ApiKeyProviderService::new(); - - // Try to get credential from provider pool or fallback to API key provider - let credential = match provider_pool_service - .select_credential_with_fallback( - &db, - &api_key_service, - "openai", - None::<&str>, - None::<&str>, - None::<&lime_core::models::client_type::ClientType>, - ) - .await - { - Ok(Some(cred)) => cred, - Ok(None) => { - return Err(String::from( - "没有可用的 OpenAI 凭证。请在设置中添加 OpenAI API Key。" - )); - } - Err(e) => return Err(format!("获取凭证失败: {}", e)), - }; - - // Extract API key from credential - let api_key = match credential.credential { - lime_core::models::provider_pool_model::CredentialData::OpenAIKey { - api_key, - .. - } => api_key, - lime_core::models::provider_pool_model::CredentialData::AnthropicKey { - api_key, - .. - } => api_key, - _ => { - return Err(String::from( - "语义搜索需要 OpenAI API Key 凭证。" - )); - } - }; - - tracing::debug!("[Hybrid Search] Using API key from provider pool"); - - // Get query embedding - let query_embedding = lime_embedding::get_embedding(&options.query, &api_key, None).await - .map_err(|e| format!("Failed to get embedding: {}", e))?; - - // For now, just return semantic search results - // TODO: Implement keyword search and merge with weights - let results = semantic_results; - - tracing::info!("[Hybrid Search] Returning {} results", results.len()); - - Ok(results) -} diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 80a833bd0..0fe7c8fcd 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -3,7 +3,6 @@ pub mod agent_cmd; pub mod api_key_provider_cmd; pub mod asr_cmd; pub mod aster_agent_cmd; -pub mod auto_fix_cmd; pub mod automation_cmd; pub mod auxiliary_model_selection; pub mod browser_connector_cmd; @@ -26,7 +25,6 @@ pub mod gateway_tunnel_cmd; pub mod image_search_cmd; pub mod image_upload_cmd; pub mod injection_cmd; -pub mod kiro_local; pub mod machine_id_cmd; pub mod material_cmd; pub mod mcp_cmd; @@ -35,18 +33,14 @@ pub mod memory_cmd; pub mod memory_feedback_cmd; pub mod memory_management_cmd; pub mod memory_search_cmd; -pub mod model_cmd; pub mod model_registry_cmd; pub mod models_cmd; -pub mod oauth_cmd; pub mod openclaw_cmd; -pub mod orchestrator_cmd; pub mod persona_cmd; pub mod plugin_cmd; pub mod plugin_install_cmd; pub mod plugin_rpc_cmd; pub mod prompt_cmd; -pub mod provider_pool_cmd; pub mod sceneapp_cmd; pub mod screenshot_cmd; pub mod security_perf_cmd; @@ -61,7 +55,6 @@ pub mod theme_context_cmd; pub mod tray_cmd; pub mod unified_memory_cmd; pub mod update_cmd; -pub mod usage_cmd; pub mod usage_stats_cmd; pub mod video_generation_cmd; pub mod voice_test_cmd; diff --git a/src-tauri/src/commands/model_cmd.rs b/src-tauri/src/commands/model_cmd.rs deleted file mode 100644 index d33d0ece5..000000000 --- a/src-tauri/src/commands/model_cmd.rs +++ /dev/null @@ -1,156 +0,0 @@ -//! 模型管理相关命令 - -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::database::DbConnection; -use lime_services::model_service::ModelService; -use std::collections::HashMap; -use tauri::State; - -/// 获取凭证支持的模型列表(从数据库缓存) -#[tauri::command] -pub fn get_credential_models( - db: State<'_, DbConnection>, - credential_uuid: String, -) -> Result, String> { - tracing::info!( - "[GET_CREDENTIAL_MODELS] 获取凭证模型列表: {}", - credential_uuid - ); - - let model_service = ModelService::new(); - model_service.get_credential_models(&db, &credential_uuid) -} - -/// 刷新凭证的模型列表(从 Provider API 重新获取) -#[tauri::command] -pub async fn refresh_credential_models( - db: State<'_, DbConnection>, - credential_uuid: String, -) -> Result, String> { - tracing::info!("[REFRESH_CREDENTIAL_MODELS] ========== 开始刷新凭证模型列表 =========="); - tracing::info!( - "[REFRESH_CREDENTIAL_MODELS] credential_uuid: {}", - credential_uuid - ); - - let model_service = ModelService::new(); - - // 从数据库获取凭证信息 - let credential = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_by_uuid(&conn, &credential_uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("凭证不存在: {credential_uuid}"))? - }; - - tracing::info!( - "[REFRESH_CREDENTIAL_MODELS] 凭证信息: provider_type={}, name={:?}", - credential.provider_type, - credential.name - ); - - // 从 Provider API 获取模型列表 - tracing::info!("[REFRESH_CREDENTIAL_MODELS] 开始从 Provider API 获取模型列表..."); - let models = model_service - .fetch_models_for_credential(&credential) - .await?; - - tracing::info!( - "[REFRESH_CREDENTIAL_MODELS] 成功获取 {} 个模型: {:?}", - models.len(), - models - ); - - // 更新到数据库 - tracing::info!("[REFRESH_CREDENTIAL_MODELS] 更新模型列表到数据库..."); - model_service.update_credential_models(&db, &credential_uuid, models.clone())?; - - tracing::info!("[REFRESH_CREDENTIAL_MODELS] ========== 刷新完成 =========="); - - Ok(models) -} - -/// 获取所有凭证的模型列表(按 Provider 类型分组) -#[tauri::command] -pub fn get_all_models_by_provider( - db: State<'_, DbConnection>, -) -> Result>, String> { - tracing::info!("[GET_ALL_MODELS_BY_PROVIDER] 获取所有 Provider 的模型列表"); - - let model_service = ModelService::new(); - model_service.get_all_models_by_provider(&db) -} - -/// 获取所有可用的模型列表(合并所有健康凭证的模型) -#[tauri::command] -pub fn get_all_available_models(db: State<'_, DbConnection>) -> Result, String> { - tracing::info!("[GET_ALL_AVAILABLE_MODELS] 获取所有可用模型"); - - let model_service = ModelService::new(); - model_service.get_all_available_models(&db) -} - -/// 批量刷新所有凭证的模型列表 -#[tauri::command] -pub async fn refresh_all_credential_models( - db: State<'_, DbConnection>, -) -> Result, String>>, String> { - tracing::info!("[REFRESH_ALL_CREDENTIAL_MODELS] 批量刷新所有凭证的模型列表"); - - let model_service = ModelService::new(); - - // 获取所有凭证 - let credentials = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())? - }; - - let mut results = HashMap::new(); - - for credential in credentials { - if credential.is_disabled { - tracing::debug!("[REFRESH_ALL] 跳过已禁用的凭证: {}", credential.uuid); - continue; - } - - tracing::info!( - "[REFRESH_ALL] 刷新凭证: {} ({})", - credential.uuid, - credential.provider_type - ); - - // 尝试获取模型列表 - let result = match model_service.fetch_models_for_credential(&credential).await { - Ok(models) => { - // 更新到数据库 - if let Err(e) = - model_service.update_credential_models(&db, &credential.uuid, models.clone()) - { - tracing::error!("[REFRESH_ALL] 更新数据库失败: {}", e); - Err(format!("更新数据库失败: {e}")) - } else { - tracing::info!("[REFRESH_ALL] 成功刷新 {} 个模型", models.len()); - Ok(models) - } - } - Err(e) => { - tracing::warn!("[REFRESH_ALL] 获取模型列表失败: {}", e); - Err(e) - } - }; - - results.insert(credential.uuid.clone(), result); - } - - Ok(results) -} - -/// 获取 Provider 的默认模型列表 -#[tauri::command] -pub fn get_default_models_for_provider(provider_type: String) -> Result, String> { - let pt: crate::models::provider_pool_model::PoolProviderType = - provider_type.parse().map_err(|e: String| e)?; - - let model_service = ModelService::new(); - Ok(model_service.get_default_models_for_provider(&pt)) -} diff --git a/src-tauri/src/commands/models_cmd.rs b/src-tauri/src/commands/models_cmd.rs index 802b34773..cb8b84ba9 100644 --- a/src-tauri/src/commands/models_cmd.rs +++ b/src-tauri/src/commands/models_cmd.rs @@ -59,19 +59,16 @@ pub struct SimpleProviderConfig { pub models: Vec, } -/// 需要使用别名配置的 Provider 列表 -const ALIAS_PROVIDERS: &[&str] = &["antigravity", "kiro", "codex", "gemini", "gemini_api_key"]; +/// 凭证池 Provider 别名已退役;模型列表统一走模型注册表。 +const ALIAS_PROVIDERS: &[&str] = &[]; /// 别名配置文件名映射(某些 Provider 共享同一个别名配置) fn get_alias_config_key(provider: &str) -> &str { - match provider { - "gemini_api_key" => "gemini", // Gemini API Key 使用 gemini 的别名配置 - _ => provider, - } + provider } /// 获取所有 Provider 的简化配置(用于前端下拉框) -/// 对于别名 Provider(antigravity、kiro、codex、gemini、gemini_api_key),优先使用别名配置中的模型列表 +/// 模型列表统一来自模型注册表。 #[tauri::command] pub async fn get_all_provider_models( app_state: State<'_, AppState>, diff --git a/src-tauri/src/commands/oauth_cmd.rs b/src-tauri/src/commands/oauth_cmd.rs deleted file mode 100644 index 15a023f3a..000000000 --- a/src-tauri/src/commands/oauth_cmd.rs +++ /dev/null @@ -1,422 +0,0 @@ -//! Unified OAuth Commands for Kiro/Gemini Providers -//! -//! This module consolidates the OAuth credential management commands -//! for OAuth providers into a single set of parameterized commands. - -use crate::providers; -use crate::AppState; -use crate::LogState; -use serde::{Deserialize, Serialize}; -use std::path::PathBuf; -use tauri::State; - -/// Supported OAuth provider types -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -#[serde(rename_all = "lowercase")] -pub enum OAuthProvider { - Kiro, - Gemini, -} - -impl OAuthProvider { - pub fn from_str(s: &str) -> Result { - match s.to_lowercase().as_str() { - "kiro" => Ok(OAuthProvider::Kiro), - "gemini" => Ok(OAuthProvider::Gemini), - _ => Err(format!("Unknown provider: {s}")), - } - } - - pub fn display_name(&self) -> &'static str { - match self { - OAuthProvider::Kiro => "Kiro", - OAuthProvider::Gemini => "Gemini", - } - } -} - -/// Unified credential status for all OAuth providers -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct OAuthCredentialStatus { - pub provider: String, - pub loaded: bool, - pub has_access_token: bool, - pub has_refresh_token: bool, - pub is_valid: bool, - pub expiry_info: Option, - pub creds_path: String, - /// Provider-specific additional info - pub extra: serde_json::Value, -} - -/// Environment variable representation -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct EnvVariable { - pub key: String, - pub value: String, - pub masked: String, -} - -/// Result of credential file change check -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CheckResult { - pub changed: bool, - pub new_hash: String, - pub reloaded: bool, -} - -fn mask_token(token: &str) -> String { - let chars: Vec = token.chars().collect(); - if chars.len() <= 12 { - "****".to_string() - } else { - let prefix: String = chars[..6].iter().collect(); - let suffix: String = chars[chars.len() - 4..].iter().collect(); - format!("{prefix}****{suffix}") - } -} - -fn get_creds_path(provider: &OAuthProvider) -> PathBuf { - match provider { - OAuthProvider::Kiro => providers::kiro::KiroProvider::default_creds_path(), - OAuthProvider::Gemini => providers::gemini::GeminiProvider::default_creds_path(), - } -} - -/// Get OAuth credentials status for a provider -#[tauri::command] -pub async fn get_oauth_credentials( - state: State<'_, AppState>, - provider: String, -) -> Result { - let provider_type = OAuthProvider::from_str(&provider)?; - let s = state.read().await; - let path = get_creds_path(&provider_type); - - match provider_type { - OAuthProvider::Kiro => { - let creds = &s.kiro_provider.credentials; - Ok(OAuthCredentialStatus { - provider: provider.clone(), - loaded: creds.access_token.is_some() || creds.refresh_token.is_some(), - has_access_token: creds.access_token.is_some(), - has_refresh_token: creds.refresh_token.is_some(), - is_valid: creds.access_token.is_some() && !s.kiro_provider.is_token_expiring_soon(), - expiry_info: creds.expires_at.clone(), - creds_path: path.to_string_lossy().to_string(), - extra: serde_json::json!({ - "region": creds.region, - "auth_method": creds.auth_method, - }), - }) - } - OAuthProvider::Gemini => { - let creds = &s.gemini_provider.credentials; - Ok(OAuthCredentialStatus { - provider: provider.clone(), - loaded: creds.access_token.is_some() || creds.refresh_token.is_some(), - has_access_token: creds.access_token.is_some(), - has_refresh_token: creds.refresh_token.is_some(), - is_valid: s.gemini_provider.is_token_valid(), - expiry_info: creds.expiry_date.map(|d| d.to_string()), - creds_path: path.to_string_lossy().to_string(), - extra: serde_json::json!({}), - }) - } - } -} - -/// Reload OAuth credentials from file -#[tauri::command] -pub async fn reload_oauth_credentials( - state: State<'_, AppState>, - logs: State<'_, LogState>, - provider: String, -) -> Result { - let provider_type = OAuthProvider::from_str(&provider)?; - let display_name = provider_type.display_name(); - - logs.write() - .await - .add("info", &format!("[{display_name}] 正在加载凭证...")); - - let mut s = state.write().await; - - let result = match provider_type { - OAuthProvider::Kiro => s.kiro_provider.load_credentials().await, - OAuthProvider::Gemini => s.gemini_provider.load_credentials().await, - }; - - match result { - Ok(_) => { - logs.write() - .await - .add("info", &format!("[{display_name}] 凭证加载成功")); - Ok(format!("{display_name} credentials reloaded")) - } - Err(e) => { - logs.write() - .await - .add("error", &format!("[{display_name}] 凭证加载失败: {e}")); - Err(e.to_string()) - } - } -} - -/// Refresh OAuth token for a provider -#[tauri::command] -pub async fn refresh_oauth_token( - state: State<'_, AppState>, - logs: State<'_, LogState>, - provider: String, -) -> Result { - let provider_type = OAuthProvider::from_str(&provider)?; - let display_name = provider_type.display_name(); - - logs.write() - .await - .add("info", &format!("[{display_name}] 正在刷新 Token...")); - - let mut s = state.write().await; - - let result = match provider_type { - OAuthProvider::Kiro => s.kiro_provider.refresh_token().await, - OAuthProvider::Gemini => s.gemini_provider.refresh_token().await, - }; - - match result { - Ok(_token) => { - logs.write() - .await - .add("info", &format!("[{display_name}] Token 刷新成功")); - // P0 安全修复:不返回明文 token - Ok("Token 刷新成功".to_string()) - } - Err(e) => { - logs.write() - .await - .add("error", &format!("[{display_name}] Token 刷新失败: {e}")); - Err(e.to_string()) - } - } -} - -/// Get environment variables for a provider -#[tauri::command] -pub async fn get_oauth_env_variables( - state: State<'_, AppState>, - provider: String, -) -> Result, String> { - let provider_type = OAuthProvider::from_str(&provider)?; - let s = state.read().await; - let mut vars = Vec::new(); - - match provider_type { - OAuthProvider::Kiro => { - let creds = &s.kiro_provider.credentials; - // P0 安全修复:不返回明文敏感凭证 - if let Some(token) = &creds.access_token { - vars.push(EnvVariable { - key: "KIRO_ACCESS_TOKEN".to_string(), - value: String::new(), - masked: mask_token(token), - }); - } - if let Some(token) = &creds.refresh_token { - vars.push(EnvVariable { - key: "KIRO_REFRESH_TOKEN".to_string(), - value: String::new(), - masked: mask_token(token), - }); - } - if let Some(id) = &creds.client_id { - vars.push(EnvVariable { - key: "KIRO_CLIENT_ID".to_string(), - value: String::new(), - masked: mask_token(id), - }); - } - if let Some(secret) = &creds.client_secret { - vars.push(EnvVariable { - key: "KIRO_CLIENT_SECRET".to_string(), - value: String::new(), - masked: mask_token(secret), - }); - } - if let Some(arn) = &creds.profile_arn { - vars.push(EnvVariable { - key: "KIRO_PROFILE_ARN".to_string(), - value: arn.clone(), - masked: arn.clone(), - }); - } - if let Some(region) = &creds.region { - vars.push(EnvVariable { - key: "KIRO_REGION".to_string(), - value: region.clone(), - masked: region.clone(), - }); - } - if let Some(method) = &creds.auth_method { - vars.push(EnvVariable { - key: "KIRO_AUTH_METHOD".to_string(), - value: method.clone(), - masked: method.clone(), - }); - } - } - OAuthProvider::Gemini => { - let creds = &s.gemini_provider.credentials; - // P0 安全修复:不返回明文敏感凭证 - if let Some(token) = &creds.access_token { - vars.push(EnvVariable { - key: "GEMINI_ACCESS_TOKEN".to_string(), - value: String::new(), - masked: mask_token(token), - }); - } - if let Some(token) = &creds.refresh_token { - vars.push(EnvVariable { - key: "GEMINI_REFRESH_TOKEN".to_string(), - value: String::new(), - masked: mask_token(token), - }); - } - if let Some(expiry) = creds.expiry_date { - let expiry_str = expiry.to_string(); - vars.push(EnvVariable { - key: "GEMINI_EXPIRY_DATE".to_string(), - value: expiry_str.clone(), - masked: expiry_str, - }); - } - } - } - - Ok(vars) -} - -/// Get token file hash for a provider -#[tauri::command] -pub async fn get_oauth_token_file_hash(provider: String) -> Result { - let provider_type = OAuthProvider::from_str(&provider)?; - let path = get_creds_path(&provider_type); - - if !tokio::fs::try_exists(&path).await.unwrap_or(false) { - return Ok("".to_string()); - } - - let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; - let hash = format!("{:x}", md5::compute(&content)); - Ok(hash) -} - -/// Check credential file changes and auto-reload -#[tauri::command] -pub async fn check_and_reload_oauth_credentials( - state: State<'_, AppState>, - logs: State<'_, LogState>, - provider: String, - last_hash: String, -) -> Result { - let provider_type = OAuthProvider::from_str(&provider)?; - let display_name = provider_type.display_name(); - let path = get_creds_path(&provider_type); - - if !tokio::fs::try_exists(&path).await.unwrap_or(false) { - return Ok(CheckResult { - changed: false, - new_hash: "".to_string(), - reloaded: false, - }); - } - - let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; - let new_hash = format!("{:x}", md5::compute(&content)); - - if !last_hash.is_empty() && new_hash != last_hash { - logs.write().await.add( - "info", - &format!("[{display_name}][自动检测] 凭证文件已变化,正在重新加载..."), - ); - - let mut s = state.write().await; - let result = match provider_type { - OAuthProvider::Kiro => s.kiro_provider.load_credentials().await, - OAuthProvider::Gemini => s.gemini_provider.load_credentials().await, - }; - - match result { - Ok(_) => { - logs.write().await.add( - "info", - &format!("[{display_name}][自动检测] 凭证重新加载成功"), - ); - Ok(CheckResult { - changed: true, - new_hash, - reloaded: true, - }) - } - Err(e) => { - logs.write().await.add( - "error", - &format!("[{display_name}][自动检测] 凭证重新加载失败: {e}"), - ); - Ok(CheckResult { - changed: true, - new_hash, - reloaded: false, - }) - } - } - } else { - Ok(CheckResult { - changed: false, - new_hash, - reloaded: false, - }) - } -} - -/// Get all OAuth providers status at once -#[tauri::command] -pub async fn get_all_oauth_credentials( - state: State<'_, AppState>, -) -> Result, String> { - let s = state.read().await; - let mut results = Vec::new(); - - // Kiro - let kiro_creds = &s.kiro_provider.credentials; - let kiro_path = providers::kiro::KiroProvider::default_creds_path(); - results.push(OAuthCredentialStatus { - provider: "kiro".to_string(), - loaded: kiro_creds.access_token.is_some() || kiro_creds.refresh_token.is_some(), - has_access_token: kiro_creds.access_token.is_some(), - has_refresh_token: kiro_creds.refresh_token.is_some(), - is_valid: kiro_creds.access_token.is_some() && !s.kiro_provider.is_token_expiring_soon(), - expiry_info: kiro_creds.expires_at.clone(), - creds_path: kiro_path.to_string_lossy().to_string(), - extra: serde_json::json!({ - "region": kiro_creds.region, - "auth_method": kiro_creds.auth_method, - }), - }); - - // Gemini - let gemini_creds = &s.gemini_provider.credentials; - let gemini_path = providers::gemini::GeminiProvider::default_creds_path(); - results.push(OAuthCredentialStatus { - provider: "gemini".to_string(), - loaded: gemini_creds.access_token.is_some() || gemini_creds.refresh_token.is_some(), - has_access_token: gemini_creds.access_token.is_some(), - has_refresh_token: gemini_creds.refresh_token.is_some(), - is_valid: s.gemini_provider.is_token_valid(), - expiry_info: gemini_creds.expiry_date.map(|d| d.to_string()), - creds_path: gemini_path.to_string_lossy().to_string(), - extra: serde_json::json!({}), - }); - - Ok(results) -} diff --git a/src-tauri/src/commands/orchestrator_cmd.rs b/src-tauri/src/commands/orchestrator_cmd.rs deleted file mode 100644 index 810f4f255..000000000 --- a/src-tauri/src/commands/orchestrator_cmd.rs +++ /dev/null @@ -1,516 +0,0 @@ -//! 模型编排器 Tauri 命令 -//! -//! 提供前端访问模型编排器的接口。 - -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::database::DbConnection; -use lime_core::orchestrator::{ - get_global_orchestrator, init_global_orchestrator, AvailableModel, CredentialInfo, - OrchestratorConfig, PoolStats, ProviderType, SelectionContext, SelectionResult, ServiceTier, - StrategyInfo, TaskHint, -}; -use serde::{Deserialize, Serialize}; -use tauri::State; -use tokio::sync::RwLock; - -/// 编排器状态 -pub struct OrchestratorState { - initialized: RwLock, -} - -impl OrchestratorState { - pub fn new() -> Self { - Self { - initialized: RwLock::new(false), - } - } -} - -impl Default for OrchestratorState { - fn default() -> Self { - Self::new() - } -} - -// ============================================================================ -// 初始化命令 -// ============================================================================ - -/// 初始化编排器 -#[tauri::command] -pub async fn init_orchestrator( - state: State<'_, OrchestratorState>, - db: State<'_, DbConnection>, -) -> Result<(), String> { - let mut initialized = state.initialized.write().await; - if *initialized { - return Ok(()); - } - - let orchestrator = init_global_orchestrator(); - *initialized = true; - - // 从数据库加载凭证并同步到 orchestrator - let credentials = { - let conn = db.lock().map_err(|e| format!("获取数据库连接失败: {e}"))?; - ProviderPoolDao::get_all(&conn).map_err(|e| format!("获取凭证列表失败: {e}"))? - }; - - // 转换凭证格式 - let cred_infos: Vec = credentials - .iter() - .filter(|c| !c.is_disabled && c.is_healthy) - .map(|c| { - // 从 credential 中提取支持的模型列表 - let supported_models = extract_supported_models(&c.credential); - // 保存原始的 provider_type 字符串(如 "antigravity"、"kiro" 等) - let original_provider_type = c.provider_type.to_string(); - - CredentialInfo { - id: c.uuid.clone(), - provider_type: map_pool_provider_type(&original_provider_type), - original_provider_type: Some(original_provider_type), - supported_models, - is_healthy: c.is_healthy, - current_load: None, - } - }) - .collect(); - - if !cred_infos.is_empty() { - orchestrator.update_credentials(cred_infos).await; - tracing::info!("已从凭证池同步 {} 个凭证到编排器", credentials.len()); - } - - tracing::info!("模型编排器已初始化"); - Ok(()) -} - -/// 从 credential 提取支持的模型列表 -fn extract_supported_models( - credential: &crate::models::provider_pool_model::CredentialData, -) -> Vec { - use crate::models::provider_pool_model::CredentialData; - - match credential { - CredentialData::ClaudeKey { .. } | CredentialData::ClaudeOAuth { .. } => { - vec![ - "claude-opus-4-5-20251101".to_string(), - "claude-opus-4-20250514".to_string(), - "claude-sonnet-4-5-20250929".to_string(), - "claude-sonnet-4-20250514".to_string(), - "claude-haiku-4-5-20251001".to_string(), - "claude-3-7-sonnet-20250219".to_string(), - "claude-3-5-haiku-20241022".to_string(), - ] - } - CredentialData::OpenAIKey { .. } => { - vec![ - "gpt-5.2-codex".to_string(), - "gpt-5.2".to_string(), - "gpt-5.1-codex-max".to_string(), - "gpt-5.1-codex".to_string(), - "gpt-5.1-codex-mini".to_string(), - "gpt-5.1".to_string(), - "gpt-5-codex".to_string(), - "gpt-5-codex-mini".to_string(), - "gpt-5".to_string(), - "gpt-4o".to_string(), - "gpt-4o-mini".to_string(), - ] - } - CredentialData::GeminiOAuth { .. } => { - vec![ - "gemini-3-pro-preview".to_string(), - "gemini-3-flash-preview".to_string(), - "gemini-2.5-pro".to_string(), - "gemini-2.5-flash".to_string(), - "gemini-2.5-flash-lite".to_string(), - ] - } - CredentialData::GeminiApiKey { - excluded_models, .. - } => { - let all_models = vec![ - "gemini-3-pro-preview".to_string(), - "gemini-3-flash-preview".to_string(), - "gemini-2.5-pro".to_string(), - "gemini-2.5-flash".to_string(), - "gemini-2.5-flash-lite".to_string(), - ]; - all_models - .into_iter() - .filter(|m| !excluded_models.contains(m)) - .collect() - } - CredentialData::KiroOAuth { .. } => { - vec![ - "claude-opus-4-5".to_string(), - "claude-opus-4-5-20251101".to_string(), - "claude-haiku-4-5".to_string(), - "claude-sonnet-4-5".to_string(), - "claude-sonnet-4-5-20250929".to_string(), - "claude-sonnet-4-20250514".to_string(), - "claude-3-7-sonnet-20250219".to_string(), - ] - } - CredentialData::CodexOAuth { .. } => { - vec!["codex-mini-latest".to_string()] - } - CredentialData::AntigravityOAuth { .. } => { - vec![ - // Max 等级 - "gemini-3-pro-preview".to_string(), - "gemini-3-pro-image-preview".to_string(), - "gemini-claude-opus-4-5-thinking".to_string(), - // Pro 等级 - "gemini-2.5-flash".to_string(), - "gemini-2.5-computer-use-preview-10-2025".to_string(), - "gemini-claude-sonnet-4-5".to_string(), - "gemini-claude-sonnet-4-5-thinking".to_string(), - // Mini 等级 - "gemini-3-flash-preview".to_string(), - ] - } - _ => vec![], - } -} - -/// 映射 PoolProviderType 到 orchestrator 的 ProviderType -fn map_pool_provider_type(pool_type: &str) -> ProviderType { - match pool_type.to_lowercase().as_str() { - "claude" | "claude_oauth" => ProviderType::Anthropic, - "openai" => ProviderType::OpenAI, - "gemini" | "gemini_api_key" | "gemini_oauth" => ProviderType::Google, - "kiro" => ProviderType::Kiro, - "codex" => ProviderType::OpenAI, - "antigravity" => ProviderType::Antigravity, - _ => ProviderType::Custom, - } -} - -/// 获取编排器配置 -#[tauri::command] -pub async fn get_orchestrator_config() -> Result { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - Ok(orchestrator.get_config().await) -} - -/// 更新编排器配置 -#[tauri::command] -pub async fn update_orchestrator_config(config: OrchestratorConfig) -> Result<(), String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - orchestrator.update_config(config).await; - Ok(()) -} - -// ============================================================================ -// 模型池命令 -// ============================================================================ - -/// 获取模型池统计 -#[tauri::command] -pub async fn get_pool_stats() -> Result { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - Ok(orchestrator.get_pool_stats().await) -} - -/// 获取指定等级的模型列表 -#[tauri::command] -pub async fn get_tier_models(tier: String) -> Result, String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - let service_tier = - ServiceTier::parse_str(&tier).ok_or_else(|| format!("无效的服务等级: {tier}"))?; - - Ok(orchestrator.get_models(service_tier).await) -} - -/// 获取所有可用模型 -#[tauri::command] -pub async fn get_all_models() -> Result, String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - Ok(orchestrator.get_all_models().await) -} - -// ============================================================================ -// 凭证管理命令 -// ============================================================================ - -/// 凭证信息请求 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CredentialInfoRequest { - pub id: String, - pub provider_type: String, - pub supported_models: Vec, - pub is_healthy: bool, - pub current_load: Option, -} - -impl From for CredentialInfo { - fn from(req: CredentialInfoRequest) -> Self { - // 保存原始的 provider_type 字符串 - let original_provider_type = req.provider_type.clone(); - CredentialInfo { - id: req.id, - provider_type: ProviderType::parse_str(&req.provider_type) - .unwrap_or(ProviderType::Custom), - original_provider_type: Some(original_provider_type), - supported_models: req.supported_models, - is_healthy: req.is_healthy, - current_load: req.current_load, - } - } -} - -/// 更新凭证列表 -#[tauri::command] -pub async fn update_orchestrator_credentials( - credentials: Vec, -) -> Result<(), String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - let creds: Vec = credentials.into_iter().map(Into::into).collect(); - orchestrator.update_credentials(creds).await; - - Ok(()) -} - -/// 添加凭证 -#[tauri::command] -pub async fn add_orchestrator_credential(credential: CredentialInfoRequest) -> Result<(), String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - orchestrator.add_credential(credential.into()).await; - Ok(()) -} - -/// 移除凭证 -#[tauri::command] -pub async fn remove_orchestrator_credential(credential_id: String) -> Result<(), String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - orchestrator.remove_credential(&credential_id).await; - Ok(()) -} - -/// 标记凭证为不健康 -#[tauri::command] -pub async fn mark_credential_unhealthy( - model_id: String, - credential_id: String, -) -> Result<(), String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - orchestrator.mark_unhealthy(&model_id, &credential_id).await; - Ok(()) -} - -/// 标记凭证为健康 -#[tauri::command] -pub async fn mark_credential_healthy(credential_id: String) -> Result<(), String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - orchestrator.mark_healthy(&credential_id).await; - Ok(()) -} - -/// 更新凭证负载 -#[tauri::command] -pub async fn update_credential_load(credential_id: String, load: u8) -> Result<(), String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - orchestrator.update_load(&credential_id, load).await; - Ok(()) -} - -// ============================================================================ -// 模型选择命令 -// ============================================================================ - -/// 选择请求 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SelectionRequest { - pub tier: String, - pub task_hint: Option, - pub requires_vision: Option, - pub requires_tools: Option, - pub preferred_provider: Option, - pub excluded_models: Option>, - pub strategy_id: Option, -} - -/// 选择模型 -#[tauri::command] -pub async fn select_model(request: SelectionRequest) -> Result { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - let tier = ServiceTier::parse_str(&request.tier) - .ok_or_else(|| format!("无效的服务等级: {}", request.tier))?; - - let mut ctx = SelectionContext::new(tier); - - if let Some(hint) = &request.task_hint { - ctx.task_hint = match hint.to_lowercase().as_str() { - "coding" => Some(TaskHint::Coding), - "writing" => Some(TaskHint::Writing), - "analysis" => Some(TaskHint::Analysis), - "chat" => Some(TaskHint::Chat), - "translation" => Some(TaskHint::Translation), - "summarization" => Some(TaskHint::Summarization), - "math" => Some(TaskHint::Math), - _ => Some(TaskHint::Other), - }; - } - - if let Some(vision) = request.requires_vision { - ctx.requires_vision = vision; - } - - if let Some(tools) = request.requires_tools { - ctx.requires_tools = tools; - } - - if let Some(provider) = request.preferred_provider { - ctx.preferred_provider = Some(provider); - } - - if let Some(excluded) = request.excluded_models { - ctx.excluded_models = excluded; - } - - let result = if let Some(strategy_id) = &request.strategy_id { - orchestrator.select_with_strategy(strategy_id, &ctx).await - } else { - orchestrator.select(&ctx).await - }; - - result.map_err(|e| e.to_string()) -} - -/// 快速选择模型 -#[tauri::command] -pub async fn quick_select_model() -> Result { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - orchestrator.quick_select().await.map_err(|e| e.to_string()) -} - -/// 为特定任务选择模型 -#[tauri::command] -pub async fn select_model_for_task(tier: String, task: String) -> Result { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - let service_tier = - ServiceTier::parse_str(&tier).ok_or_else(|| format!("无效的服务等级: {tier}"))?; - - let task_hint = match task.to_lowercase().as_str() { - "coding" => TaskHint::Coding, - "writing" => TaskHint::Writing, - "analysis" => TaskHint::Analysis, - "chat" => TaskHint::Chat, - "translation" => TaskHint::Translation, - "summarization" => TaskHint::Summarization, - "math" => TaskHint::Math, - _ => TaskHint::Other, - }; - - orchestrator - .select_for_task(service_tier, task_hint) - .await - .map_err(|e| e.to_string()) -} - -// ============================================================================ -// 策略命令 -// ============================================================================ - -/// 列出所有可用策略 -#[tauri::command] -pub async fn list_strategies() -> Result, String> { - let orchestrator = get_global_orchestrator().ok_or("编排器未初始化")?; - - Ok(orchestrator.list_strategies().await) -} - -/// 获取服务等级列表 -#[tauri::command] -pub fn list_service_tiers() -> Vec { - ServiceTier::all() - .iter() - .map(|t| ServiceTierInfo { - id: format!("{t:?}").to_lowercase(), - display_name: t.display_name().to_string(), - description: t.description().to_string(), - level: t.level(), - }) - .collect() -} - -/// 服务等级信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ServiceTierInfo { - pub id: String, - pub display_name: String, - pub description: String, - pub level: u8, -} - -/// 获取任务类型列表 -#[tauri::command] -pub fn list_task_hints() -> Vec { - vec![ - TaskHintInfo { - id: "coding".to_string(), - display_name: "代码".to_string(), - description: "代码生成、编辑、调试".to_string(), - }, - TaskHintInfo { - id: "writing".to_string(), - display_name: "写作".to_string(), - description: "文章、报告、创意写作".to_string(), - }, - TaskHintInfo { - id: "analysis".to_string(), - display_name: "分析".to_string(), - description: "数据分析、推理、研究".to_string(), - }, - TaskHintInfo { - id: "chat".to_string(), - display_name: "对话".to_string(), - description: "日常对话、问答".to_string(), - }, - TaskHintInfo { - id: "translation".to_string(), - display_name: "翻译".to_string(), - description: "语言翻译".to_string(), - }, - TaskHintInfo { - id: "summarization".to_string(), - display_name: "摘要".to_string(), - description: "文本摘要、总结".to_string(), - }, - TaskHintInfo { - id: "math".to_string(), - display_name: "数学".to_string(), - description: "数学计算、推理".to_string(), - }, - TaskHintInfo { - id: "other".to_string(), - display_name: "其他".to_string(), - description: "其他任务".to_string(), - }, - ] -} - -/// 任务类型信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TaskHintInfo { - pub id: String, - pub display_name: String, - pub description: String, -} diff --git a/src-tauri/src/commands/persona_cmd.rs b/src-tauri/src/commands/persona_cmd.rs index 92647ef83..1b9983b85 100644 --- a/src-tauri/src/commands/persona_cmd.rs +++ b/src-tauri/src/commands/persona_cmd.rs @@ -287,7 +287,7 @@ pub struct GeneratedPersonaCommandResult { /// AI 一键生成人设 /// /// 根据用户提供的简单描述,调用 AI 生成完整的人设配置。 -/// 自动从凭证池选择可用凭证进行调用。 +/// 自动从 API Key Provider 选择可用凭证进行调用。 /// /// # 参数 /// - `prompt`: 用户描述,例如"一个幽默风趣的科技博主" diff --git a/src-tauri/src/commands/provider_pool_cmd.rs b/src-tauri/src/commands/provider_pool_cmd.rs deleted file mode 100644 index f09053d58..000000000 --- a/src-tauri/src/commands/provider_pool_cmd.rs +++ /dev/null @@ -1,3865 +0,0 @@ -//! Provider Pool Tauri 命令 - -#![allow(dead_code)] - -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::database::DbConnection; -use crate::models::provider_pool_model::{ - AddCredentialRequest, CredentialData, CredentialDisplay, HealthCheckResult, OAuthStatus, - PoolProviderType, ProviderCredential, ProviderPoolOverview, UpdateCredentialRequest, -}; -use chrono::Utc; -use lime_credential::CredentialSyncService; -use lime_services::provider_pool_service::ProviderPoolService; -use std::fs; -use std::path::{Path, PathBuf}; -use std::sync::Arc; -use tauri::{Emitter, State}; -use uuid::Uuid; - -pub struct ProviderPoolServiceState(pub Arc); - -/// 凭证同步服务状态封装 -pub struct CredentialSyncServiceState(pub Option>); - -/// 展开路径中的 ~ 为用户主目录 -fn expand_tilde(path: &str) -> String { - if let Some(stripped) = path.strip_prefix("~/") { - if let Some(home) = dirs::home_dir() { - return home.join(stripped).to_string_lossy().to_string(); - } - } - path.to_string() -} - -/// 获取应用凭证存储目录 -fn get_credentials_dir() -> Result { - let app_data_dir = dirs::data_dir() - .ok_or_else(|| "无法获取应用数据目录".to_string())? - .join("lime") - .join("credentials"); - - // 确保目录存在 - if !app_data_dir.exists() { - fs::create_dir_all(&app_data_dir).map_err(|e| format!("创建凭证存储目录失败: {e}"))?; - } - - Ok(app_data_dir) -} - -/// 复制并重命名 OAuth 凭证文件 -/// -/// 对于 Kiro 凭证,会自动合并 clientIdHash 文件中的 client_id/client_secret, -/// 使副本文件完全独立,支持多账号场景。 -fn copy_and_rename_credential_file( - source_path: &str, - provider_type: &str, -) -> Result { - let expanded_source = expand_tilde(source_path); - let source = Path::new(&expanded_source); - - // 验证源文件存在 - if !source.exists() { - return Err(format!("凭证文件不存在: {expanded_source}")); - } - - // 生成新的文件名:{provider_type}_{uuid}_{timestamp}.json - let uuid = Uuid::new_v4().to_string(); - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - let new_filename = format!( - "{}_{}_{}_{}.json", - provider_type, - &uuid[..8], // 使用 UUID 前8位 - timestamp, - provider_type - ); - - // 获取目标目录 - let credentials_dir = get_credentials_dir()?; - let target_path = credentials_dir.join(&new_filename); - - // 对于 Kiro 凭证,需要合并 clientIdHash 文件中的 client_id/client_secret - if provider_type == "kiro" { - let content = fs::read_to_string(source).map_err(|e| format!("读取凭证文件失败: {e}"))?; - let mut creds: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {e}"))?; - - // 检测 refreshToken 是否被截断(仅记录警告,不阻止添加) - // 正常的 refreshToken 长度应该在 500+ 字符,如果小于 100 字符则可能被截断 - // 注意:即使 refreshToken 被截断,也允许添加凭证,在刷新时才会提示错误 - if let Some(refresh_token) = creds.get("refreshToken").and_then(|v| v.as_str()) { - let token_len = refresh_token.len(); - - // 检测常见的截断模式 - let is_truncated = - token_len < 100 || refresh_token.ends_with("...") || refresh_token.contains("..."); - - if is_truncated { - // 安全地截取前 50 个字符(避免 UTF-8 边界 panic) - let preview: String = refresh_token.chars().take(50).collect(); - tracing::warn!( - "[KIRO] 检测到 refreshToken 可能被截断!长度: {}, 内容: {}... (仍允许添加,刷新时会提示)", - token_len, - preview - ); - // 不再阻止添加,只记录警告 - // 在刷新 Token 时会检测并提示用户 - } else { - tracing::info!("[KIRO] refreshToken 长度检查通过: {} 字符", token_len); - } - } else { - tracing::warn!("[KIRO] 凭证文件中没有 refreshToken 字段"); - } - - let aws_sso_cache_dir = dirs::home_dir() - .ok_or_else(|| "无法获取用户主目录".to_string())? - .join(".aws") - .join("sso") - .join("cache"); - - // 尝试从 clientIdHash 文件或扫描目录获取 client_id/client_secret - let mut found_credentials = false; - - // 方式1:如果有 clientIdHash,读取对应文件 - if let Some(hash) = creds.get("clientIdHash").and_then(|v| v.as_str()) { - let hash_file_path = aws_sso_cache_dir.join(format!("{hash}.json")); - - if hash_file_path.exists() { - if let Ok(hash_content) = fs::read_to_string(&hash_file_path) { - if let Ok(hash_json) = serde_json::from_str::(&hash_content) - { - if let Some(client_id) = hash_json.get("clientId") { - creds["clientId"] = client_id.clone(); - } - if let Some(client_secret) = hash_json.get("clientSecret") { - creds["clientSecret"] = client_secret.clone(); - } - if creds.get("clientId").is_some() && creds.get("clientSecret").is_some() { - found_credentials = true; - tracing::info!( - "[KIRO] 已从 clientIdHash 文件合并 client_id/client_secret 到副本" - ); - } - } - } - } - } - - // 方式2:如果没有 clientIdHash 或未找到,扫描目录中的其他 JSON 文件 - if !found_credentials && aws_sso_cache_dir.exists() { - tracing::info!( - "[KIRO] 没有 clientIdHash 或未找到,扫描目录查找 client_id/client_secret" - ); - if let Ok(entries) = fs::read_dir(&aws_sso_cache_dir) { - for entry in entries.flatten() { - let file_path = entry.path(); - // 跳过主凭证文件和备份文件 - if file_path.extension().map(|e| e == "json").unwrap_or(false) { - let file_name = - file_path.file_name().and_then(|n| n.to_str()).unwrap_or(""); - if file_name.starts_with("kiro-auth-token") { - continue; - } - if let Ok(file_content) = fs::read_to_string(&file_path) { - if let Ok(file_json) = - serde_json::from_str::(&file_content) - { - let has_client_id = - file_json.get("clientId").and_then(|v| v.as_str()).is_some(); - let has_client_secret = file_json - .get("clientSecret") - .and_then(|v| v.as_str()) - .is_some(); - if has_client_id && has_client_secret { - creds["clientId"] = file_json["clientId"].clone(); - creds["clientSecret"] = file_json["clientSecret"].clone(); - found_credentials = true; - tracing::info!( - "[KIRO] 已从 {} 合并 client_id/client_secret 到副本", - file_name - ); - break; - } - } - } - } - } - } - } - - if !found_credentials { - // 检查认证方式 - let auth_method = creds - .get("authMethod") - .and_then(|v| v.as_str()) - .unwrap_or("social"); - - if auth_method.to_lowercase() == "idc" { - // IdC 认证必须有 clientId/clientSecret - tracing::error!( - "[KIRO] IdC 认证方式缺少 clientId/clientSecret,无法创建有效的凭证副本" - ); - return Err( - "IdC 认证凭证不完整:缺少 clientId/clientSecret。\n\n💡 解决方案:\n1. 确保 ~/.aws/sso/cache/ 目录下有对应的 clientIdHash 文件\n2. 如果使用 AWS IAM Identity Center,请确保已完成完整的 SSO 登录流程\n3. 或者尝试使用 Social 认证方式的凭证".to_string() - ); - } else { - tracing::warn!("[KIRO] 未找到 client_id/client_secret,将使用 social 认证方式"); - } - } - - // 写入合并后的凭证到副本文件 - let merged_content = - serde_json::to_string_pretty(&creds).map_err(|e| format!("序列化凭证失败: {e}"))?; - fs::write(&target_path, merged_content).map_err(|e| format!("写入凭证文件失败: {e}"))?; - } else { - // 其他类型直接复制 - fs::copy(source, &target_path).map_err(|e| format!("复制凭证文件失败: {e}"))?; - } - - // 返回新的文件路径 - Ok(target_path.to_string_lossy().to_string()) -} - -/// 删除凭证文件(如果在应用存储目录中) -fn cleanup_credential_file(file_path: &str) -> Result<(), String> { - let path = Path::new(file_path); - - // 只删除在应用凭证存储目录中的文件 - if let Ok(credentials_dir) = get_credentials_dir() { - if let Ok(canonical_path) = path.canonicalize() { - if let Ok(canonical_dir) = credentials_dir.canonicalize() { - if canonical_path.starts_with(canonical_dir) { - if let Err(e) = fs::remove_file(&canonical_path) { - // 只记录警告,不中断删除过程 - println!("Warning: Failed to delete credential file: {e}"); - } - } - } - } - } - - Ok(()) -} - -/// 获取凭证池概览 -#[tauri::command] -pub fn get_provider_pool_overview( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, -) -> Result, String> { - pool_service.0.get_overview(&db) -} - -/// 获取指定类型的凭证列表 -#[tauri::command] -pub fn get_provider_pool_credentials( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - provider_type: String, -) -> Result, String> { - pool_service.0.get_by_type(&db, &provider_type) -} - -/// 添加凭证 -/// -/// 添加凭证到数据库,并同步到 YAML 配置文件 -/// Requirements: 1.1, 1.2 -#[tauri::command] -pub fn add_provider_pool_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - sync_service: State<'_, CredentialSyncServiceState>, - request: AddCredentialRequest, -) -> Result { - // 添加到数据库 - let credential = pool_service.0.add_credential( - &db, - &request.provider_type, - request.credential, - request.name, - request.check_health, - request.check_model_name, - )?; - - // 同步到 YAML 配置(如果同步服务可用) - if let Some(ref sync) = sync_service.0 { - if let Err(e) = sync.add_credential(&credential) { - // 记录警告但不中断操作 - tracing::warn!("同步凭证到 YAML 失败: {}", e); - } - } - - Ok(credential) -} - -/// 更新凭证 -/// 更新凭证 -/// -/// 更新数据库中的凭证,并同步到 YAML 配置文件 -/// Requirements: 1.1, 1.2 -#[tauri::command] -pub fn update_provider_pool_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - sync_service: State<'_, CredentialSyncServiceState>, - uuid: String, - request: UpdateCredentialRequest, -) -> Result { - tracing::info!( - "[UPDATE_CREDENTIAL] 收到更新请求: uuid={}, name={:?}, check_model_name={:?}, not_supported_models={:?}", - uuid, - request.name, - request.check_model_name, - request.not_supported_models - ); - // 如果需要重新上传文件,先处理文件上传 - let credential = if let Some(new_file_path) = request.new_creds_file_path { - // 获取当前凭证以确定类型 - let conn = db.lock().map_err(|e| e.to_string())?; - let current_credential = ProviderPoolDao::get_by_uuid(&conn, &uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("凭证不存在: {uuid}"))?; - - // 根据凭证类型复制新文件 - let new_stored_path = match ¤t_credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - // 清理旧文件 - cleanup_credential_file(creds_file_path)?; - copy_and_rename_credential_file(&new_file_path, "kiro")? - } - CredentialData::GeminiOAuth { - creds_file_path, .. - } => { - // 清理旧文件 - cleanup_credential_file(creds_file_path)?; - copy_and_rename_credential_file(&new_file_path, "gemini")? - } - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => { - // 清理旧文件 - cleanup_credential_file(creds_file_path)?; - copy_and_rename_credential_file(&new_file_path, "antigravity")? - } - _ => { - return Err("只有 OAuth 凭证支持重新上传文件".to_string()); - } - }; - - // 更新凭证数据 - let mut updated_cred = current_credential; - - // 更新凭证数据中的文件路径 - match &mut updated_cred.credential { - CredentialData::KiroOAuth { creds_file_path } => { - *creds_file_path = new_stored_path; - } - CredentialData::GeminiOAuth { - creds_file_path, - project_id, - } => { - *creds_file_path = new_stored_path; - if let Some(new_pid) = request.new_project_id { - *project_id = Some(new_pid); - } - } - CredentialData::AntigravityOAuth { - creds_file_path, - project_id, - } => { - *creds_file_path = new_stored_path; - if let Some(new_pid) = request.new_project_id { - *project_id = Some(new_pid); - } - } - _ => {} - } - - // 应用其他更新 - // 处理 name:空字符串表示清除,None 表示不修改 - if let Some(name) = request.name { - updated_cred.name = if name.is_empty() { None } else { Some(name) }; - } - if let Some(is_disabled) = request.is_disabled { - updated_cred.is_disabled = is_disabled; - } - if let Some(check_health) = request.check_health { - updated_cred.check_health = check_health; - } - // 处理 check_model_name:空字符串表示清除,None 表示不修改 - if let Some(check_model_name) = request.check_model_name { - updated_cred.check_model_name = if check_model_name.is_empty() { - None - } else { - Some(check_model_name) - }; - } - if let Some(not_supported_models) = request.not_supported_models { - updated_cred.not_supported_models = not_supported_models; - } - - updated_cred.updated_at = Utc::now(); - - // 保存到数据库 - ProviderPoolDao::update(&conn, &updated_cred).map_err(|e| e.to_string())?; - - updated_cred - } else if request.new_base_url.is_some() || request.new_api_key.is_some() { - // 更新 API Key 凭证的 api_key 和/或 base_url - let conn = db.lock().map_err(|e| e.to_string())?; - let mut current_credential = ProviderPoolDao::get_by_uuid(&conn, &uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("凭证不存在: {uuid}"))?; - - // 更新 api_key 和 base_url - match &mut current_credential.credential { - CredentialData::OpenAIKey { api_key, base_url } => { - if let Some(new_key) = request.new_api_key { - if !new_key.is_empty() { - *api_key = new_key; - } - } - if let Some(new_url) = request.new_base_url { - *base_url = if new_url.is_empty() { - None - } else { - Some(new_url) - }; - } - } - CredentialData::ClaudeKey { api_key, base_url } => { - if let Some(new_key) = request.new_api_key { - if !new_key.is_empty() { - *api_key = new_key; - } - } - if let Some(new_url) = request.new_base_url { - *base_url = if new_url.is_empty() { - None - } else { - Some(new_url) - }; - } - } - _ => { - return Err("只有 API Key 凭证支持修改 API Key 和 Base URL".to_string()); - } - } - - // 应用其他更新 - // 处理 name:空字符串表示清除,None 表示不修改 - if let Some(name) = request.name { - current_credential.name = if name.is_empty() { None } else { Some(name) }; - } - if let Some(is_disabled) = request.is_disabled { - current_credential.is_disabled = is_disabled; - } - if let Some(check_health) = request.check_health { - current_credential.check_health = check_health; - } - // 处理 check_model_name:空字符串表示清除,None 表示不修改 - if let Some(check_model_name) = request.check_model_name { - current_credential.check_model_name = if check_model_name.is_empty() { - None - } else { - Some(check_model_name) - }; - } - if let Some(not_supported_models) = request.not_supported_models { - current_credential.not_supported_models = not_supported_models; - } - - current_credential.updated_at = Utc::now(); - - // 保存到数据库 - ProviderPoolDao::update(&conn, ¤t_credential).map_err(|e| e.to_string())?; - - current_credential - } else { - // 常规更新,不涉及文件 - pool_service.0.update_credential( - &db, - &uuid, - request.name, - request.is_disabled, - request.check_health, - request.check_model_name, - request.not_supported_models, - request.new_proxy_url, - )? - }; - - // 同步到 YAML 配置(如果同步服务可用) - if let Some(ref sync) = sync_service.0 { - if let Err(e) = sync.update_credential(&credential) { - // 记录警告但不中断操作 - tracing::warn!("同步凭证更新到 YAML 失败: {}", e); - } - } - - Ok(credential) -} - -/// 删除凭证 -/// 删除凭证 -/// -/// 从数据库删除凭证,并同步到 YAML 配置文件 -/// Requirements: 1.1, 1.2 -#[tauri::command] -pub fn delete_provider_pool_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - sync_service: State<'_, CredentialSyncServiceState>, - uuid: String, - provider_type: Option, -) -> Result { - // 从数据库删除 - let result = pool_service.0.delete_credential(&db, &uuid)?; - - // 同步到 YAML 配置(如果同步服务可用且提供了 provider_type) - if let Some(ref sync) = sync_service.0 { - if let Some(pt) = provider_type { - if let Ok(pool_type) = pt.parse::() { - if let Err(e) = sync.remove_credential(pool_type, &uuid) { - // 记录警告但不中断操作 - tracing::warn!("从 YAML 删除凭证失败: {}", e); - } - } - } - } - - Ok(result) -} - -/// 切换凭证启用/禁用状态 -#[tauri::command] -pub fn toggle_provider_pool_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - uuid: String, - is_disabled: bool, -) -> Result { - pool_service - .0 - .update_credential(&db, &uuid, None, Some(is_disabled), None, None, None, None) -} - -/// 重置凭证计数器 -#[tauri::command] -pub fn reset_provider_pool_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - uuid: String, -) -> Result<(), String> { - pool_service.0.reset_counters(&db, &uuid) -} - -/// 重置指定类型的所有凭证健康状态 -#[tauri::command] -pub fn reset_provider_pool_health( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - provider_type: String, -) -> Result { - pool_service.0.reset_health_by_type(&db, &provider_type) -} - -/// 执行单个凭证的健康检查 -#[tauri::command] -pub async fn check_provider_pool_credential_health( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - uuid: String, -) -> Result { - tracing::info!("[DEBUG] 开始健康检查 for uuid: {}", uuid); - let result = pool_service.0.check_credential_health(&db, &uuid).await; - match &result { - Ok(health) => tracing::info!( - "[DEBUG] 健康检查完成: success={}, message={:?}", - health.success, - health.message - ), - Err(err) => tracing::error!("[DEBUG] 健康检查失败: {}", err), - } - result -} - -/// 执行指定类型的所有凭证健康检查 -#[tauri::command] -pub async fn check_provider_pool_type_health( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - provider_type: String, -) -> Result, String> { - pool_service.0.check_type_health(&db, &provider_type).await -} - -/// 添加 Kiro OAuth 凭证(通过文件路径) -#[tauri::command] -pub fn add_kiro_oauth_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - creds_file_path: String, - name: Option, -) -> Result { - tracing::info!("[KIRO] 开始添加凭证,文件路径: {}", creds_file_path); - - // 复制并重命名文件到应用存储目录 - let stored_file_path = match copy_and_rename_credential_file(&creds_file_path, "kiro") { - Ok(path) => { - tracing::info!("[KIRO] 凭证文件已复制到: {}", path); - path - } - Err(e) => { - tracing::error!("[KIRO] 复制凭证文件失败: {}", e); - return Err(e); - } - }; - - match pool_service.0.add_credential( - &db, - "kiro", - CredentialData::KiroOAuth { - creds_file_path: stored_file_path, - }, - name, - Some(true), - None, - ) { - Ok(cred) => { - tracing::info!("[KIRO] 凭证添加成功,UUID: {}", cred.uuid); - Ok(cred) - } - Err(e) => { - tracing::error!("[KIRO] 添加凭证到数据库失败: {}", e); - Err(e) - } - } -} - -/// 从 JSON 内容创建 Kiro 凭证文件并添加到凭证池 -/// -/// 直接粘贴 JSON 内容,无需选择文件 -fn create_kiro_credential_from_json(json_content: &str) -> Result { - // 验证 JSON 格式 - let creds: serde_json::Value = - serde_json::from_str(json_content).map_err(|e| format!("JSON 格式无效: {e}"))?; - - // 验证必要字段 - if creds.get("refreshToken").is_none() { - return Err("凭证 JSON 缺少 refreshToken 字段".to_string()); - } - - // 检测 refreshToken 是否被截断 - if let Some(refresh_token) = creds.get("refreshToken").and_then(|v| v.as_str()) { - let token_len = refresh_token.len(); - let is_truncated = - token_len < 100 || refresh_token.ends_with("...") || refresh_token.contains("..."); - - if is_truncated { - tracing::warn!( - "[KIRO] 检测到 refreshToken 可能被截断!长度: {} (仍允许添加,刷新时会提示)", - token_len - ); - } - } - - // 生成新的文件名 - let uuid = Uuid::new_v4().to_string(); - let timestamp = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - let new_filename = format!("kiro_{}_{}_{}.json", &uuid[..8], timestamp, "kiro"); - - // 获取目标目录 - let credentials_dir = get_credentials_dir()?; - let target_path = credentials_dir.join(&new_filename); - - // 尝试合并 clientId/clientSecret(如果凭证中没有) - let mut merged_creds = creds.clone(); - - // 检查是否需要从外部文件获取 clientId/clientSecret - let has_client_id = merged_creds.get("clientId").is_some(); - let has_client_secret = merged_creds.get("clientSecret").is_some(); - - if !has_client_id || !has_client_secret { - let aws_sso_cache_dir = dirs::home_dir() - .ok_or_else(|| "无法获取用户主目录".to_string())? - .join(".aws") - .join("sso") - .join("cache"); - - let mut found_credentials = false; - - // 方式1:如果有 clientIdHash,读取对应文件 - if let Some(hash) = merged_creds.get("clientIdHash").and_then(|v| v.as_str()) { - let hash_file_path = aws_sso_cache_dir.join(format!("{hash}.json")); - - if hash_file_path.exists() { - if let Ok(hash_content) = fs::read_to_string(&hash_file_path) { - if let Ok(hash_json) = serde_json::from_str::(&hash_content) - { - if let Some(client_id) = hash_json.get("clientId") { - merged_creds["clientId"] = client_id.clone(); - } - if let Some(client_secret) = hash_json.get("clientSecret") { - merged_creds["clientSecret"] = client_secret.clone(); - } - if merged_creds.get("clientId").is_some() - && merged_creds.get("clientSecret").is_some() - { - found_credentials = true; - tracing::info!( - "[KIRO] 已从 clientIdHash 文件合并 client_id/client_secret" - ); - } - } - } - } - } - - // 方式2:扫描目录中的其他 JSON 文件 - if !found_credentials && aws_sso_cache_dir.exists() { - tracing::info!("[KIRO] 扫描目录查找 client_id/client_secret"); - if let Ok(entries) = fs::read_dir(&aws_sso_cache_dir) { - for entry in entries.flatten() { - let file_path = entry.path(); - if file_path.extension().map(|e| e == "json").unwrap_or(false) { - let file_name = - file_path.file_name().and_then(|n| n.to_str()).unwrap_or(""); - if file_name.starts_with("kiro-auth-token") { - continue; - } - if let Ok(file_content) = fs::read_to_string(&file_path) { - if let Ok(file_json) = - serde_json::from_str::(&file_content) - { - let has_cid = - file_json.get("clientId").and_then(|v| v.as_str()).is_some(); - let has_csec = file_json - .get("clientSecret") - .and_then(|v| v.as_str()) - .is_some(); - if has_cid && has_csec { - merged_creds["clientId"] = file_json["clientId"].clone(); - merged_creds["clientSecret"] = - file_json["clientSecret"].clone(); - found_credentials = true; - tracing::info!( - "[KIRO] 从 {} 合并 client_id/client_secret", - file_name - ); - break; - } - } - } - } - } - } - } - - if !found_credentials { - let auth_method = merged_creds - .get("authMethod") - .and_then(|v| v.as_str()) - .unwrap_or("social"); - - if auth_method.to_lowercase() == "idc" { - tracing::error!( - "[KIRO] IdC 认证方式缺少 clientId/clientSecret,无法创建有效的凭证" - ); - return Err( - "IdC 认证凭证不完整:缺少 clientId/clientSecret。\n\n💡 解决方案:\n1. 确保 ~/.aws/sso/cache/ 目录下有对应的 clientIdHash 文件\n2. 如果使用 AWS IAM Identity Center,请确保已完成完整的 SSO 登录流程\n3. 或者尝试使用 Social 认证方式的凭证".to_string() - ); - } else { - tracing::warn!("[KIRO] 未找到 client_id/client_secret,将使用 social 认证方式"); - } - } - } - - // 写入凭证文件 - let merged_content = - serde_json::to_string_pretty(&merged_creds).map_err(|e| format!("序列化凭证失败: {e}"))?; - fs::write(&target_path, merged_content).map_err(|e| format!("写入凭证文件失败: {e}"))?; - - tracing::info!("[KIRO] 凭证文件已创建: {:?}", target_path); - - Ok(target_path.to_string_lossy().to_string()) -} - -/// 添加 Kiro OAuth 凭证(通过 JSON 内容) -/// -/// 直接粘贴凭证 JSON 内容,无需选择文件 -#[tauri::command] -pub fn add_kiro_from_json( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - json_content: String, - name: Option, -) -> Result { - tracing::info!( - "[KIRO] 开始从 JSON 添加凭证,内容长度: {}", - json_content.len() - ); - - // 从 JSON 内容创建凭证文件 - let stored_file_path = match create_kiro_credential_from_json(&json_content) { - Ok(path) => { - tracing::info!("[KIRO] 凭证文件已创建: {}", path); - path - } - Err(e) => { - tracing::error!("[KIRO] 创建凭证文件失败: {}", e); - return Err(e); - } - }; - - match pool_service.0.add_credential( - &db, - "kiro", - CredentialData::KiroOAuth { - creds_file_path: stored_file_path, - }, - name, - Some(true), - None, - ) { - Ok(cred) => { - tracing::info!("[KIRO] 凭证添加成功,UUID: {}", cred.uuid); - Ok(cred) - } - Err(e) => { - tracing::error!("[KIRO] 添加凭证到数据库失败: {}", e); - Err(e) - } - } -} - -/// 添加 Gemini OAuth 凭证(通过文件路径) -#[tauri::command] -pub fn add_gemini_oauth_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - creds_file_path: String, - project_id: Option, - name: Option, -) -> Result { - // 复制并重命名文件到应用存储目录 - let stored_file_path = copy_and_rename_credential_file(&creds_file_path, "gemini")?; - - pool_service.0.add_credential( - &db, - "gemini", - CredentialData::GeminiOAuth { - creds_file_path: stored_file_path, - project_id, - }, - name, - Some(true), - None, - ) -} - -/// 添加 Antigravity OAuth 凭证(通过文件路径) -#[tauri::command] -pub fn add_antigravity_oauth_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - creds_file_path: String, - project_id: Option, - name: Option, -) -> Result { - // 复制并重命名文件到应用存储目录 - let stored_file_path = copy_and_rename_credential_file(&creds_file_path, "antigravity")?; - - pool_service.0.add_credential( - &db, - "antigravity", - CredentialData::AntigravityOAuth { - creds_file_path: stored_file_path, - project_id, - }, - name, - Some(true), - None, - ) -} - -/// 添加 OpenAI API Key 凭证 -#[tauri::command] -pub fn add_openai_key_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - api_key: String, - base_url: Option, - name: Option, -) -> Result { - pool_service.0.add_credential( - &db, - "openai", - CredentialData::OpenAIKey { api_key, base_url }, - name, - Some(true), - None, - ) -} - -/// 添加 Claude API Key 凭证 -#[tauri::command] -pub fn add_claude_key_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - api_key: String, - base_url: Option, - name: Option, -) -> Result { - pool_service.0.add_credential( - &db, - "claude", - CredentialData::ClaudeKey { api_key, base_url }, - name, - Some(true), - None, - ) -} - -/// 添加 Gemini API Key 凭证 -#[tauri::command] -pub fn add_gemini_api_key_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - api_key: String, - base_url: Option, - excluded_models: Option>, - name: Option, -) -> Result { - pool_service.0.add_credential( - &db, - "gemini_api_key", - CredentialData::GeminiApiKey { - api_key, - base_url, - excluded_models: excluded_models.unwrap_or_default(), - }, - name, - Some(true), - None, - ) -} - -/// 添加 Codex OAuth 凭证(通过文件路径) -#[tauri::command] -pub fn add_codex_oauth_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - creds_file_path: String, - api_base_url: Option, - name: Option, -) -> Result { - // 复制并重命名文件到应用存储目录 - let stored_file_path = copy_and_rename_credential_file(&creds_file_path, "codex")?; - - pool_service.0.add_credential( - &db, - "codex", - CredentialData::CodexOAuth { - creds_file_path: stored_file_path, - api_base_url, - }, - name, - Some(true), - None, - ) -} - -/// 添加 Claude OAuth 凭证(通过文件路径) -#[tauri::command] -pub fn add_claude_oauth_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - creds_file_path: String, - name: Option, -) -> Result { - // 复制并重命名文件到应用存储目录 - let stored_file_path = copy_and_rename_credential_file(&creds_file_path, "claude_oauth")?; - - pool_service.0.add_credential( - &db, - "claude_oauth", - CredentialData::ClaudeOAuth { - creds_file_path: stored_file_path, - }, - name, - Some(true), - None, - ) -} - -/// 刷新凭证的 OAuth Token -#[tauri::command] -pub async fn refresh_pool_credential_token( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - uuid: String, -) -> Result { - tracing::info!("[DEBUG] 开始刷新 Token for uuid: {}", uuid); - let result = pool_service.0.refresh_credential_token(&db, &uuid).await; - match &result { - Ok(msg) => tracing::info!("[DEBUG] Token 刷新成功: {}", msg), - Err(err) => tracing::error!("[DEBUG] Token 刷新失败: {}", err), - } - result -} - -/// 获取凭证的 OAuth 状态 -#[tauri::command] -pub fn get_pool_credential_oauth_status( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - uuid: String, -) -> Result { - pool_service.0.get_credential_oauth_status(&db, &uuid) -} - -/// 调试 Kiro 凭证加载(从默认路径) -/// P0 安全修复:仅在 debug 构建中可用 -#[cfg(debug_assertions)] -#[tauri::command] -pub async fn debug_kiro_credentials() -> Result { - use crate::providers::kiro::KiroProvider; - - let mut provider = KiroProvider::new(); - - let mut result = String::new(); - result.push_str("🔍 开始 Kiro 凭证调试 (默认路径)...\n\n"); - - match provider.load_credentials().await { - Ok(_) => { - result.push_str("✅ 凭证加载成功!\n"); - result.push_str(&format!( - "📄 认证方式: {:?}\n", - provider.credentials.auth_method - )); - result.push_str(&format!( - "🔑 有 client_id: {}\n", - provider.credentials.client_id.is_some() - )); - result.push_str(&format!( - "🔒 有 client_secret: {}\n", - provider.credentials.client_secret.is_some() - )); - result.push_str(&format!( - "🏷️ 有 clientIdHash: {}\n", - provider.credentials.client_id_hash.is_some() - )); - - // P0 安全修复:不再输出敏感信息(clientIdHash、token 前缀等) - let detected_method = provider.detect_auth_method(); - result.push_str(&format!("🎯 检测到的认证方式: {detected_method}\n")); - - result.push_str("\n🚀 尝试刷新 token...\n"); - match provider.refresh_token().await { - Ok(token) => { - result.push_str(&format!("✅ Token 刷新成功! Token 长度: {}\n", token.len())); - // 不再输出 token 前缀 - } - Err(e) => { - result.push_str(&format!("❌ Token 刷新失败: {e}\n")); - } - } - } - Err(e) => { - result.push_str(&format!("❌ 凭证加载失败: {e}\n")); - } - } - - Ok(result) -} - -/// P0 安全修复:release 构建中禁用 debug 命令 -#[cfg(not(debug_assertions))] -#[tauri::command] -pub async fn debug_kiro_credentials() -> Result { - Err("此调试命令仅在开发构建中可用".to_string()) -} - -/// 测试用户上传的凭证文件 -/// P0 安全修复:仅在 debug 构建中可用,且不输出敏感信息 -#[cfg(debug_assertions)] -#[tauri::command] -pub async fn test_user_credentials() -> Result { - use crate::providers::kiro::KiroProvider; - - let mut result = String::new(); - result.push_str("🧪 测试用户上传的凭证文件...\n\n"); - - // 测试用户上传的凭证文件路径 - let user_creds_path = dirs::home_dir() - .ok_or("无法获取用户主目录".to_string())? - .join("Library/Application Support/lime/credentials/kiro_d8da9d58_1765757992_kiro.json"); - - // P0 安全修复:不输出完整路径,仅显示文件是否存在 - result.push_str("📂 检查用户凭证文件...\n"); - - // 检查文件是否存在 - if !user_creds_path.exists() { - result.push_str("❌ 用户凭证文件不存在!\n"); - result.push_str("💡 请确保文件路径正确,或重新上传凭证文件\n"); - return Ok(result); - } - - result.push_str("✅ 用户凭证文件存在\n\n"); - - // 读取并解析用户凭证文件 - match std::fs::read_to_string(&user_creds_path) { - Ok(content) => { - result.push_str("✅ 成功读取凭证文件\n"); - result.push_str(&format!("📄 文件大小: {} 字节\n", content.len())); - - // 尝试解析 JSON - match serde_json::from_str::(&content) { - Ok(json) => { - result.push_str("✅ JSON 格式有效\n"); - - // 检查关键字段(仅显示是否存在,不显示值) - let has_access_token = - json.get("accessToken").and_then(|v| v.as_str()).is_some(); - let has_refresh_token = - json.get("refreshToken").and_then(|v| v.as_str()).is_some(); - let auth_method = json.get("authMethod").and_then(|v| v.as_str()); - let has_client_id_hash = - json.get("clientIdHash").and_then(|v| v.as_str()).is_some(); - let region = json.get("region").and_then(|v| v.as_str()); - - result.push_str(&format!("🔑 有 accessToken: {has_access_token}\n")); - result.push_str(&format!("🔄 有 refreshToken: {has_refresh_token}\n")); - result.push_str(&format!("📄 authMethod: {auth_method:?}\n")); - // P0 安全修复:不输出 clientIdHash 值 - result.push_str(&format!("🏷️ 有 clientIdHash: {has_client_id_hash}\n")); - result.push_str(&format!("🌍 region: {region:?}\n")); - - // 使用 KiroProvider 测试加载 - result.push_str("\n🔧 使用 KiroProvider 测试加载...\n"); - - let mut provider = KiroProvider::new(); - provider.creds_path = Some(user_creds_path.clone()); - - match provider - .load_credentials_from_path(&user_creds_path.to_string_lossy()) - .await - { - Ok(_) => { - result.push_str("✅ KiroProvider 加载成功!\n"); - result.push_str(&format!( - "📄 最终认证方式: {:?}\n", - provider.credentials.auth_method - )); - result.push_str(&format!( - "🔑 最终有 client_id: {}\n", - provider.credentials.client_id.is_some() - )); - result.push_str(&format!( - "🔒 最终有 client_secret: {}\n", - provider.credentials.client_secret.is_some() - )); - - let detected_method = provider.detect_auth_method(); - result.push_str(&format!("🎯 检测到的认证方式: {detected_method}\n")); - - result.push_str("\n🚀 尝试刷新 token...\n"); - match provider.refresh_token().await { - Ok(token) => { - result.push_str(&format!( - "✅ Token 刷新成功! Token 长度: {}\n", - token.len() - )); - // P0 安全修复:不输出 token 前缀 - } - Err(e) => { - result.push_str(&format!("❌ Token 刷新失败: {e}\n")); - } - } - } - Err(e) => { - result.push_str(&format!("❌ KiroProvider 加载失败: {e}\n")); - } - } - } - Err(e) => { - result.push_str(&format!("❌ JSON 格式无效: {e}\n")); - } - } - } - Err(e) => { - result.push_str(&format!("❌ 无法读取凭证文件: {e}\n")); - } - } - - Ok(result) -} - -/// P0 安全修复:release 构建中禁用 test_user_credentials 命令 -#[cfg(not(debug_assertions))] -#[tauri::command] -pub async fn test_user_credentials() -> Result { - Err("此调试命令仅在开发构建中可用".to_string()) -} - -/// 迁移 Private 配置到凭证池 -/// -/// 从 providers 配置中读取单个凭证配置,迁移到凭证池中并标记为 Private 来源 -/// Requirements: 6.4 -#[tauri::command] -pub fn migrate_private_config_to_pool( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - config: crate::config::Config, -) -> Result { - let result = pool_service.0.migrate_private_config(&db, &config)?; - Ok(MigrationResultResponse { - migrated_count: result.migrated_count, - skipped_count: result.skipped_count, - errors: result.errors, - }) -} - -/// 迁移结果响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct MigrationResultResponse { - /// 成功迁移的凭证数量 - pub migrated_count: usize, - /// 跳过的凭证数量(已存在) - pub skipped_count: usize, - /// 错误信息列表 - pub errors: Vec, -} - -/// 获取 Antigravity OAuth 授权 URL 并等待回调(不自动打开浏览器) -/// -/// 启动服务器后通过事件发送授权 URL,然后等待回调 -/// 成功后返回凭证 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct AntigravityAuthUrlResponse { - pub auth_url: String, -} - -#[tauri::command] -pub async fn get_antigravity_auth_url_and_wait( - app: tauri::AppHandle, - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - name: Option, - skip_project_id_fetch: Option, -) -> Result { - use crate::providers::antigravity; - - tracing::info!("[Antigravity OAuth] 启动服务器并获取授权 URL"); - - // 启动服务器并获取授权 URL - let (auth_url, wait_future) = - antigravity::start_oauth_server_and_get_url(skip_project_id_fetch.unwrap_or(false)) - .await - .map_err(|e| format!("启动 OAuth 服务器失败: {e}"))?; - - tracing::info!("[Antigravity OAuth] 授权 URL: {}", auth_url); - - // 通过事件发送授权 URL 给前端 - let _ = app.emit( - "antigravity-auth-url", - AntigravityAuthUrlResponse { - auth_url: auth_url.clone(), - }, - ); - - // 等待回调 - let result = wait_future.await.map_err(|e| e.to_string())?; - - tracing::info!( - "[Antigravity OAuth] 登录成功,凭证保存到: {}", - result.creds_file_path - ); - - // 从凭证中获取 project_id - let project_id = result.credentials.project_id.clone(); - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "antigravity", - CredentialData::AntigravityOAuth { - creds_file_path: result.creds_file_path, - project_id, - }, - name, - Some(true), - None, - )?; - - tracing::info!( - "[Antigravity OAuth] 凭证已添加到凭证池: {}", - credential.uuid - ); - - Ok(credential) -} - -/// 启动 Antigravity OAuth 登录流程 -/// -/// 打开浏览器让用户登录 Google 账号,获取 Antigravity 凭证 -#[tauri::command] -pub async fn start_antigravity_oauth_login( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - name: Option, - skip_project_id_fetch: Option, -) -> Result { - use crate::providers::antigravity; - - tracing::info!("[Antigravity OAuth] 开始 OAuth 登录流程"); - - // 启动 OAuth 登录 - let result = antigravity::start_oauth_login(skip_project_id_fetch.unwrap_or(false)) - .await - .map_err(|e| format!("Antigravity OAuth 登录失败: {e}"))?; - - tracing::info!( - "[Antigravity OAuth] 登录成功,凭证保存到: {}", - result.creds_file_path - ); - - // 从凭证中获取 project_id - let project_id = result.credentials.project_id.clone(); - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "antigravity", - CredentialData::AntigravityOAuth { - creds_file_path: result.creds_file_path, - project_id, - }, - name, - Some(true), - None, - )?; - - tracing::info!( - "[Antigravity OAuth] 凭证已添加到凭证池: {}", - credential.uuid - ); - - Ok(credential) -} - -/// Codex OAuth 授权 URL 响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct CodexAuthUrlResponse { - pub auth_url: String, -} - -/// 获取 Codex OAuth 授权 URL 并等待回调(不自动打开浏览器) -/// -/// 启动服务器后通过事件发送授权 URL,然后等待回调 -/// 成功后返回凭证 -#[tauri::command] -pub async fn get_codex_auth_url_and_wait( - app: tauri::AppHandle, - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - name: Option, -) -> Result { - use crate::providers::codex; - - tracing::info!("[Codex OAuth] 启动服务器并获取授权 URL"); - - // 启动服务器并获取授权 URL - let (auth_url, wait_future) = codex::start_codex_oauth_server_and_get_url() - .await - .map_err(|e| format!("启动 OAuth 服务器失败: {e}"))?; - - tracing::info!("[Codex OAuth] 授权 URL: {}", auth_url); - - // 通过事件发送授权 URL 给前端 - let _ = app.emit( - "codex-auth-url", - CodexAuthUrlResponse { - auth_url: auth_url.clone(), - }, - ); - - // 等待回调 - let result = wait_future.await.map_err(|e| e.to_string())?; - - tracing::info!( - "[Codex OAuth] 登录成功,凭证保存到: {}", - result.creds_file_path - ); - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "codex", - CredentialData::CodexOAuth { - creds_file_path: result.creds_file_path, - api_base_url: None, - }, - name, - Some(true), - None, - )?; - - tracing::info!("[Codex OAuth] 凭证已添加到凭证池: {}", credential.uuid); - - Ok(credential) -} - -/// 启动 Codex OAuth 登录流程 -/// -/// 打开浏览器让用户登录 OpenAI 账号,获取 Codex 凭证 -#[tauri::command] -pub async fn start_codex_oauth_login( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - name: Option, -) -> Result { - use crate::providers::codex; - - tracing::info!("[Codex OAuth] 开始 OAuth 登录流程"); - - // 启动 OAuth 登录 - let result = codex::start_codex_oauth_login() - .await - .map_err(|e| format!("Codex OAuth 登录失败: {e}"))?; - - tracing::info!( - "[Codex OAuth] 登录成功,凭证保存到: {}", - result.creds_file_path - ); - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "codex", - CredentialData::CodexOAuth { - creds_file_path: result.creds_file_path, - api_base_url: None, - }, - name, - Some(true), - None, - )?; - - tracing::info!("[Codex OAuth] 凭证已添加到凭证池: {}", credential.uuid); - - Ok(credential) -} - -/// Claude OAuth 授权 URL 响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct ClaudeOAuthAuthUrlResponse { - pub auth_url: String, -} - -/// 获取 Claude OAuth 授权 URL 并等待回调(不自动打开浏览器) -/// -/// 启动服务器后通过事件发送授权 URL,然后等待回调 -/// 成功后返回凭证 -/// Claude OAuth 授权 URL 响应(新流程) -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct ClaudeOAuthParamsResponse { - pub auth_url: String, - pub code_verifier: String, - pub state: String, -} - -/// 获取 Claude OAuth 授权 URL(新流程:手动输入授权码) -/// -/// 生成授权 URL 和 PKCE 参数,用户需要: -/// 1. 打开 auth_url 进行授权 -/// 2. 授权后从页面复制授权码 -/// 3. 调用 exchange_claude_oauth_code 交换 token -#[tauri::command] -pub async fn get_claude_oauth_auth_url_and_wait( - app: tauri::AppHandle, - _db: State<'_, DbConnection>, - _pool_service: State<'_, ProviderPoolServiceState>, - _name: Option, -) -> Result { - use crate::providers::claude_oauth; - - tracing::info!("[Claude OAuth] 生成授权 URL(手动授权码流程)"); - - // 生成授权参数 - let params = claude_oauth::generate_claude_oauth_params() - .map_err(|e| format!("生成授权参数失败: {e}"))?; - - tracing::info!("[Claude OAuth] 授权 URL: {}", params.auth_url); - - // 通过事件发送授权 URL 给前端 - let _ = app.emit( - "claude-oauth-auth-url", - ClaudeOAuthAuthUrlResponse { - auth_url: params.auth_url.clone(), - }, - ); - - // 打开浏览器 - if let Err(e) = open::that(¶ms.auth_url) { - tracing::warn!("[Claude OAuth] 无法打开浏览器: {}. 请手动打开 URL.", e); - } - - Ok(ClaudeOAuthParamsResponse { - auth_url: params.auth_url, - code_verifier: params.code_verifier, - state: params.state, - }) -} - -/// 使用授权码交换 Claude OAuth Token -/// -/// 用户在浏览器中授权后,复制授权码,调用此命令交换 token -#[tauri::command] -pub async fn exchange_claude_oauth_code( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - authorization_code: String, - code_verifier: String, - state: String, - name: Option, -) -> Result { - use crate::providers::claude_oauth; - - tracing::info!("[Claude OAuth] 使用授权码交换 Token"); - - // 交换 Token - let result = claude_oauth::exchange_claude_authorization_code( - &authorization_code, - &code_verifier, - &state, - ) - .await - .map_err(|e| format!("Claude OAuth Token 交换失败: {e}"))?; - - tracing::info!( - "[Claude OAuth] 登录成功,凭证保存到: {}", - result.creds_file_path - ); - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "claude_oauth", - CredentialData::ClaudeOAuth { - creds_file_path: result.creds_file_path, - }, - name, - Some(true), - None, - )?; - - tracing::info!("[Claude OAuth] 凭证已添加到凭证池: {}", credential.uuid); - - Ok(credential) -} - -/// 启动 Claude OAuth 登录流程(兼容旧接口,现在返回授权参数) -/// -/// 打开浏览器让用户登录 Claude 账号 -/// 注意:新流程需要用户手动复制授权码,然后调用 exchange_claude_oauth_code -#[tauri::command] -pub async fn start_claude_oauth_login( - _db: State<'_, DbConnection>, - _pool_service: State<'_, ProviderPoolServiceState>, - _name: Option, -) -> Result { - use crate::providers::claude_oauth; - - tracing::info!("[Claude OAuth] 开始 OAuth 登录流程(手动授权码模式)"); - - // 生成授权参数并打开浏览器 - let params = claude_oauth::start_claude_oauth_login() - .await - .map_err(|e| format!("Claude OAuth 登录失败: {e}"))?; - - Ok(ClaudeOAuthParamsResponse { - auth_url: params.auth_url, - code_verifier: params.code_verifier, - state: params.state, - }) -} - -/// Claude Cookie 自动授权响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct ClaudeCookieOAuthResponse { - pub organization_uuid: Option, - pub capabilities: Vec, -} - -/// 使用 Cookie (sessionKey) 自动完成 Claude OAuth 授权 -/// -/// 这是一个更便捷的授权方式,用户只需要提供从浏览器 Cookie 中获取的 sessionKey, -/// 系统会自动完成整个 OAuth 流程,无需手动复制授权码。 -/// -/// # 参数 -/// - `session_key`: 从浏览器 Cookie 中获取的 sessionKey -/// - `is_setup_token`: 是否为 Setup Token 模式(只需要推理权限,无 refresh_token) -/// - `name`: 凭证名称(可选) -#[tauri::command] -pub async fn claude_oauth_with_cookie( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - session_key: String, - is_setup_token: Option, - name: Option, -) -> Result { - use crate::providers::claude_oauth; - - let is_setup = is_setup_token.unwrap_or(false); - tracing::info!( - "[Claude OAuth] 开始 Cookie 自动授权流程,is_setup_token: {}", - is_setup - ); - - // 执行 Cookie 自动授权 - let result = claude_oauth::oauth_with_cookie(&session_key, is_setup) - .await - .map_err(|e| format!("Claude Cookie 授权失败: {e}"))?; - - tracing::info!( - "[Claude OAuth] Cookie 授权成功,凭证保存到: {}", - result.creds_file_path - ); - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "claude_oauth", - CredentialData::ClaudeOAuth { - creds_file_path: result.creds_file_path, - }, - name, - Some(true), - None, - )?; - - tracing::info!( - "[Claude OAuth] 凭证已添加到凭证池: {}, org_uuid: {:?}", - credential.uuid, - result.organization_uuid - ); - - Ok(credential) -} - -/// -/// 获取 Kiro 凭证的 Machine ID 指纹信息 -/// -/// 返回凭证的唯一设备指纹,用于在 UI 中展示 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct KiroFingerprintInfo { - /// Machine ID(SHA256 哈希,64 字符) - pub machine_id: String, - /// Machine ID 的短格式(前 16 字符) - pub machine_id_short: String, - /// 指纹来源(profileArn / clientId / system) - pub source: String, - /// 认证方式 - pub auth_method: String, -} - -#[tauri::command] -pub async fn get_kiro_credential_fingerprint( - db: State<'_, DbConnection>, - uuid: String, -) -> Result { - use crate::database::dao::provider_pool::ProviderPoolDao; - use crate::providers::kiro::{generate_machine_id_from_credentials, KiroProvider}; - - // 获取凭证文件路径(在锁释放前完成) - let creds_file_path = { - let conn = db.lock().map_err(|e| e.to_string())?; - let credential = ProviderPoolDao::get_by_uuid(&conn, &uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("凭证不存在: {uuid}"))?; - - // 检查是否为 Kiro 凭证 - match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(), - _ => return Err("只有 Kiro 凭证支持获取指纹信息".to_string()), - } - }; // conn 在这里释放 - - // 加载凭证文件(异步操作,锁已释放) - let mut provider = KiroProvider::new(); - provider - .load_credentials_from_path(&creds_file_path) - .await - .map_err(|e| format!("加载凭证失败: {e}"))?; - - // 确定指纹来源 - let (source, profile_arn, client_id) = if provider.credentials.profile_arn.is_some() { - ( - "profileArn".to_string(), - provider.credentials.profile_arn.as_deref(), - None, - ) - } else if provider.credentials.client_id.is_some() { - ( - "clientId".to_string(), - None, - provider.credentials.client_id.as_deref(), - ) - } else { - ("system".to_string(), None, None) - }; - - // 生成 Machine ID - let machine_id = generate_machine_id_from_credentials(profile_arn, client_id); - // 安全地截取前 16 个字符(避免越界 panic) - let machine_id_short: String = machine_id.chars().take(16).collect(); - - // 获取认证方式 - let auth_method = provider - .credentials - .auth_method - .clone() - .unwrap_or_else(|| "social".to_string()); - - Ok(KiroFingerprintInfo { - machine_id, - machine_id_short, - source, - auth_method, - }) -} - -/// Gemini OAuth 授权 URL 响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct GeminiAuthUrlResponse { - pub auth_url: String, - pub session_id: String, -} - -use once_cell::sync::Lazy; -/// Gemini OAuth 会话存储(用于存储 code_verifier) -use std::collections::HashMap; -use tokio::sync::RwLock; - -static GEMINI_OAUTH_SESSIONS: Lazy< - RwLock>, -> = Lazy::new(|| RwLock::new(HashMap::new())); - -/// 获取 Gemini OAuth 授权 URL(不等待回调) -/// -/// 生成授权 URL 和 session_id,通过事件发送给前端 -/// 用户需要手动复制授权码回来,然后调用 exchange_gemini_code -#[tauri::command] -pub async fn get_gemini_auth_url_and_wait( - app: tauri::AppHandle, - _db: State<'_, DbConnection>, - _pool_service: State<'_, ProviderPoolServiceState>, - _name: Option, -) -> Result { - use crate::providers::gemini; - - tracing::info!("[Gemini OAuth] 生成授权 URL"); - - // 生成授权 URL 和会话信息 - let (auth_url, session) = gemini::generate_gemini_auth_url_with_session(); - let session_id = session.session_id.clone(); - - tracing::info!("[Gemini OAuth] 授权 URL: {}", auth_url); - tracing::info!("[Gemini OAuth] Session ID: {}", session_id); - - // 存储会话信息(用于后续交换 token) - { - let mut sessions = GEMINI_OAUTH_SESSIONS.write().await; - sessions.insert(session_id.clone(), session); - - // 清理过期的会话(超过 10 分钟) - let now = chrono::Utc::now().timestamp(); - sessions.retain(|_, s| now - s.created_at < 600); - } - - // 通过事件发送授权 URL 给前端 - let _ = app.emit( - "gemini-auth-url", - GeminiAuthUrlResponse { - auth_url: auth_url.clone(), - session_id: session_id.clone(), - }, - ); - - // 返回错误,让前端知道需要用户手动输入授权码 - // 这不是真正的错误,只是流程需要用户交互 - Err(format!("AUTH_URL:{auth_url}")) -} - -/// 用 Gemini 授权码交换 Token 并添加凭证 -#[tauri::command] -pub async fn exchange_gemini_code( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - code: String, - session_id: Option, - name: Option, -) -> Result { - use crate::providers::gemini; - - tracing::info!("[Gemini OAuth] 开始交换授权码"); - - // 获取 code_verifier - let code_verifier = if let Some(ref sid) = session_id { - let sessions = GEMINI_OAUTH_SESSIONS.read().await; - sessions - .get(sid) - .map(|s| s.code_verifier.clone()) - .ok_or_else(|| "会话已过期,请重新获取授权 URL".to_string())? - } else { - // 如果没有 session_id,尝试使用最近的会话 - let sessions = GEMINI_OAUTH_SESSIONS.read().await; - sessions - .values() - .max_by_key(|s| s.created_at) - .map(|s| s.code_verifier.clone()) - .ok_or_else(|| "没有可用的会话,请先获取授权 URL".to_string())? - }; - - // 交换 token 并创建凭证 - let result = gemini::exchange_gemini_code_and_create_credentials(&code, &code_verifier) - .await - .map_err(|e| format!("交换授权码失败: {e}"))?; - - tracing::info!( - "[Gemini OAuth] 登录成功,凭证保存到: {}", - result.creds_file_path - ); - - // 清理使用过的会话 - if let Some(ref sid) = session_id { - let mut sessions = GEMINI_OAUTH_SESSIONS.write().await; - sessions.remove(sid); - } - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "gemini", - CredentialData::GeminiOAuth { - creds_file_path: result.creds_file_path, - project_id: None, // 项目 ID 会在健康检查时自动获取 - }, - name, - Some(true), - None, - )?; - - tracing::info!("[Gemini OAuth] 凭证已添加到凭证池: {}", credential.uuid); - - Ok(credential) -} - -/// 启动 Gemini OAuth 登录流程 -/// -/// 打开浏览器让用户登录 Google 账号,获取 Gemini 凭证 -#[tauri::command] -pub async fn start_gemini_oauth_login( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - name: Option, -) -> Result { - use crate::providers::gemini; - - tracing::info!("[Gemini OAuth] 开始 OAuth 登录流程"); - - // 启动 OAuth 登录 - let result = gemini::start_gemini_oauth_login() - .await - .map_err(|e| format!("Gemini OAuth 登录失败: {e}"))?; - - tracing::info!( - "[Gemini OAuth] 登录成功,凭证保存到: {}", - result.creds_file_path - ); - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "gemini", - CredentialData::GeminiOAuth { - creds_file_path: result.creds_file_path, - project_id: None, - }, - name, - Some(true), - None, - )?; - - tracing::info!("[Gemini OAuth] 凭证已添加到凭证池: {}", credential.uuid); - - Ok(credential) -} - -// ============ Kiro Builder ID 登录相关命令 ============ - -/// Kiro Builder ID 登录状态 -#[derive(Debug, Clone)] -struct KiroBuilderIdLoginState { - /// OIDC 客户端 ID - client_id: String, - /// OIDC 客户端密钥 - client_secret: String, - /// 设备码 - device_code: String, - /// 用户码 - user_code: String, - /// 验证 URI - verification_uri: String, - /// 轮询间隔(秒) - interval: i64, - /// 过期时间戳 - expires_at: i64, - /// 区域 - region: String, -} - -/// 全局 Builder ID 登录状态存储 -static KIRO_BUILDER_ID_LOGIN_STATE: Lazy>> = - Lazy::new(|| RwLock::new(None)); - -/// Kiro Builder ID 登录启动响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct KiroBuilderIdLoginResponse { - /// 是否成功 - pub success: bool, - /// 用户码(用于显示给用户) - #[serde(rename = "userCode")] - pub user_code: Option, - /// 验证 URI(用户需要访问的 URL) - #[serde(rename = "verificationUri")] - pub verification_uri: Option, - /// 过期时间(秒) - #[serde(rename = "expiresIn")] - pub expires_in: Option, - /// 轮询间隔(秒) - pub interval: Option, - /// 错误信息 - pub error: Option, -} - -/// Kiro Builder ID 轮询响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct KiroBuilderIdPollResponse { - /// 是否成功 - pub success: bool, - /// 是否完成授权 - pub completed: bool, - /// 状态(pending / slow_down) - pub status: Option, - /// 错误信息 - pub error: Option, -} - -/// 启动 Kiro Builder ID 登录 -/// -/// 使用 OIDC Device Authorization Flow 进行登录 -#[tauri::command] -pub async fn start_kiro_builder_id_login( - region: Option, -) -> Result { - let region = region.unwrap_or_else(|| "us-east-1".to_string()); - let oidc_base = format!("https://oidc.{region}.amazonaws.com"); - let start_url = "https://view.awsapps.com/start"; - let scopes = vec![ - "codewhisperer:completions", - "codewhisperer:analysis", - "codewhisperer:conversations", - "codewhisperer:transformations", - "codewhisperer:taskassist", - ]; - - tracing::info!("[Kiro Builder ID] 开始登录流程,区域: {}", region); - - // Step 1: 注册 OIDC 客户端 - tracing::info!("[Kiro Builder ID] Step 1: 注册 OIDC 客户端..."); - let client = reqwest::Client::new(); - - let reg_body = serde_json::json!({ - "clientName": "Lime Kiro Manager", - "clientType": "public", - "scopes": scopes, - "grantTypes": ["urn:ietf:params:oauth:grant-type:device_code", "refresh_token"], - "issuerUrl": start_url - }); - - let reg_res = client - .post(format!("{oidc_base}/client/register")) - .header("Content-Type", "application/json") - .json(®_body) - .send() - .await - .map_err(|e| format!("注册客户端请求失败: {e}"))?; - - if !reg_res.status().is_success() { - let err_text = reg_res.text().await.unwrap_or_default(); - return Ok(KiroBuilderIdLoginResponse { - success: false, - user_code: None, - verification_uri: None, - expires_in: None, - interval: None, - error: Some(format!("注册客户端失败: {err_text}")), - }); - } - - let reg_data: serde_json::Value = reg_res - .json() - .await - .map_err(|e| format!("解析注册响应失败: {e}"))?; - - let client_id = reg_data["clientId"] - .as_str() - .ok_or("响应中缺少 clientId")? - .to_string(); - let client_secret = reg_data["clientSecret"] - .as_str() - .ok_or("响应中缺少 clientSecret")? - .to_string(); - - tracing::info!( - "[Kiro Builder ID] 客户端注册成功: {}...", - client_id.chars().take(30).collect::() - ); - - // Step 2: 发起设备授权 - tracing::info!("[Kiro Builder ID] Step 2: 发起设备授权..."); - let auth_body = serde_json::json!({ - "clientId": client_id, - "clientSecret": client_secret, - "startUrl": start_url - }); - - let auth_res = client - .post(format!("{oidc_base}/device_authorization")) - .header("Content-Type", "application/json") - .json(&auth_body) - .send() - .await - .map_err(|e| format!("设备授权请求失败: {e}"))?; - - if !auth_res.status().is_success() { - let err_text = auth_res.text().await.unwrap_or_default(); - return Ok(KiroBuilderIdLoginResponse { - success: false, - user_code: None, - verification_uri: None, - expires_in: None, - interval: None, - error: Some(format!("设备授权失败: {err_text}")), - }); - } - - let auth_data: serde_json::Value = auth_res - .json() - .await - .map_err(|e| format!("解析授权响应失败: {e}"))?; - - let device_code = auth_data["deviceCode"] - .as_str() - .ok_or("响应中缺少 deviceCode")? - .to_string(); - let user_code = auth_data["userCode"] - .as_str() - .ok_or("响应中缺少 userCode")? - .to_string(); - let verification_uri = auth_data["verificationUriComplete"] - .as_str() - .or_else(|| auth_data["verificationUri"].as_str()) - .ok_or("响应中缺少 verificationUri")? - .to_string(); - let interval = auth_data["interval"].as_i64().unwrap_or(5); - let expires_in = auth_data["expiresIn"].as_i64().unwrap_or(600); - - tracing::info!("[Kiro Builder ID] 设备码获取成功,user_code: {}", user_code); - - // 保存登录状态 - let expires_at = chrono::Utc::now().timestamp() + expires_in; - { - let mut state = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - *state = Some(KiroBuilderIdLoginState { - client_id, - client_secret, - device_code, - user_code: user_code.clone(), - verification_uri: verification_uri.clone(), - interval, - expires_at, - region, - }); - } - - Ok(KiroBuilderIdLoginResponse { - success: true, - user_code: Some(user_code), - verification_uri: Some(verification_uri), - expires_in: Some(expires_in), - interval: Some(interval), - error: None, - }) -} - -/// 轮询 Kiro Builder ID 授权状态 -#[tauri::command] -pub async fn poll_kiro_builder_id_auth() -> Result { - let state = { - let state_guard = KIRO_BUILDER_ID_LOGIN_STATE.read().await; - match state_guard.as_ref() { - Some(s) => s.clone(), - None => { - return Ok(KiroBuilderIdPollResponse { - success: false, - completed: false, - status: None, - error: Some("没有进行中的登录".to_string()), - }); - } - } - }; - - // 检查是否过期 - if chrono::Utc::now().timestamp() > state.expires_at { - // 清除状态 - { - let mut state_guard = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - *state_guard = None; - } - return Ok(KiroBuilderIdPollResponse { - success: false, - completed: false, - status: None, - error: Some("授权已过期,请重新开始".to_string()), - }); - } - - let oidc_base = format!("https://oidc.{}.amazonaws.com", state.region); - let client = reqwest::Client::new(); - - let token_body = serde_json::json!({ - "clientId": state.client_id, - "clientSecret": state.client_secret, - "grantType": "urn:ietf:params:oauth:grant-type:device_code", - "deviceCode": state.device_code - }); - - let token_res = client - .post(format!("{oidc_base}/token")) - .header("Content-Type", "application/json") - .json(&token_body) - .send() - .await - .map_err(|e| format!("Token 请求失败: {e}"))?; - - let status = token_res.status(); - - if status.is_success() { - // 授权成功 - let token_data: serde_json::Value = token_res - .json() - .await - .map_err(|e| format!("解析 Token 响应失败: {e}"))?; - - tracing::info!("[Kiro Builder ID] 授权成功!"); - - // 保存凭证到文件 - let access_token = token_data["accessToken"].as_str().unwrap_or("").to_string(); - let refresh_token = token_data["refreshToken"] - .as_str() - .unwrap_or("") - .to_string(); - let expires_in = token_data["expiresIn"].as_i64().unwrap_or(3600); - - // 创建凭证 JSON - let creds_json = serde_json::json!({ - "accessToken": access_token, - "refreshToken": refresh_token, - "clientId": state.client_id, - "clientSecret": state.client_secret, - "region": state.region, - "authMethod": "idc", - "expiresAt": chrono::Utc::now().timestamp() + expires_in - }); - - // 保存到临时状态,等待 add_kiro_from_builder_id_auth 调用 - // 这里我们把凭证 JSON 存储到一个临时位置 - { - let mut sessions = KIRO_BUILDER_ID_CREDENTIALS.write().await; - sessions.insert("pending".to_string(), creds_json); - } - - // 清除登录状态 - { - let mut state_guard = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - *state_guard = None; - } - - Ok(KiroBuilderIdPollResponse { - success: true, - completed: true, - status: None, - error: None, - }) - } else if status.as_u16() == 400 { - let err_data: serde_json::Value = token_res - .json() - .await - .map_err(|e| format!("解析错误响应失败: {e}"))?; - - let error = err_data["error"].as_str().unwrap_or("unknown"); - - match error { - "authorization_pending" => Ok(KiroBuilderIdPollResponse { - success: true, - completed: false, - status: Some("pending".to_string()), - error: None, - }), - "slow_down" => { - // 增加轮询间隔 - { - let mut state_guard = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - if let Some(ref mut s) = *state_guard { - s.interval += 5; - } - } - Ok(KiroBuilderIdPollResponse { - success: true, - completed: false, - status: Some("slow_down".to_string()), - error: None, - }) - } - "expired_token" => { - // 清除状态 - { - let mut state_guard = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - *state_guard = None; - } - Ok(KiroBuilderIdPollResponse { - success: false, - completed: false, - status: None, - error: Some("设备码已过期".to_string()), - }) - } - "access_denied" => { - // 清除状态 - { - let mut state_guard = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - *state_guard = None; - } - Ok(KiroBuilderIdPollResponse { - success: false, - completed: false, - status: None, - error: Some("用户拒绝授权".to_string()), - }) - } - _ => { - // 清除状态 - { - let mut state_guard = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - *state_guard = None; - } - Ok(KiroBuilderIdPollResponse { - success: false, - completed: false, - status: None, - error: Some(format!("授权错误: {error}")), - }) - } - } - } else { - Ok(KiroBuilderIdPollResponse { - success: false, - completed: false, - status: None, - error: Some(format!("未知响应: {status}")), - }) - } -} - -/// 临时存储 Builder ID 登录成功后的凭证 -static KIRO_BUILDER_ID_CREDENTIALS: Lazy>> = - Lazy::new(|| RwLock::new(HashMap::new())); - -/// 取消 Kiro Builder ID 登录 -#[tauri::command] -pub async fn cancel_kiro_builder_id_login() -> Result { - tracing::info!("[Kiro Builder ID] 取消登录"); - { - let mut state = KIRO_BUILDER_ID_LOGIN_STATE.write().await; - *state = None; - } - { - let mut creds = KIRO_BUILDER_ID_CREDENTIALS.write().await; - creds.remove("pending"); - } - Ok(true) -} - -/// 从 Builder ID 授权结果添加 Kiro 凭证 -#[tauri::command] -pub async fn add_kiro_from_builder_id_auth( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - name: Option, -) -> Result { - // 获取待处理的凭证 - let creds_json = { - let mut creds = KIRO_BUILDER_ID_CREDENTIALS.write().await; - creds - .remove("pending") - .ok_or("没有待处理的 Builder ID 凭证")? - }; - - // 将凭证 JSON 转换为字符串 - let json_content = - serde_json::to_string_pretty(&creds_json).map_err(|e| format!("序列化凭证失败: {e}"))?; - - // 使用现有的 create_kiro_credential_from_json 函数创建凭证文件 - let stored_file_path = create_kiro_credential_from_json(&json_content)?; - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "kiro", - CredentialData::KiroOAuth { - creds_file_path: stored_file_path, - }, - name, - Some(true), - None, - )?; - - tracing::info!("[Kiro Builder ID] 凭证已添加到凭证池: {}", credential.uuid); - - Ok(credential) -} - -// ============ Kiro Social Auth 登录相关命令 (Google/GitHub) ============ - -/// Kiro Auth 端点 -const KIRO_AUTH_ENDPOINT: &str = "https://prod.us-east-1.auth.desktop.kiro.dev"; - -/// Kiro Social Auth 登录状态 -#[derive(Debug, Clone)] -struct KiroSocialAuthLoginState { - /// 登录提供商 (Google / Github) - provider: String, - /// PKCE code_verifier - code_verifier: String, - /// PKCE code_challenge - code_challenge: String, - /// OAuth state - oauth_state: String, - /// 过期时间戳 - expires_at: i64, -} - -/// 全局 Social Auth 登录状态存储 -static KIRO_SOCIAL_AUTH_LOGIN_STATE: Lazy>> = - Lazy::new(|| RwLock::new(None)); - -/// Kiro Social Auth 登录启动响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct KiroSocialAuthLoginResponse { - /// 是否成功 - pub success: bool, - /// 登录 URL - #[serde(rename = "loginUrl")] - pub login_url: Option, - /// OAuth state(用于验证回调) - pub state: Option, - /// 错误信息 - pub error: Option, -} - -/// Kiro Social Auth Token 交换响应 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct KiroSocialAuthTokenResponse { - /// 是否成功 - pub success: bool, - /// 错误信息 - pub error: Option, -} - -/// 生成 PKCE code_verifier -fn generate_code_verifier() -> String { - use rand::Rng; - let mut rng = rand::thread_rng(); - let bytes: Vec = (0..64).map(|_| rng.gen()).collect(); - base64_url_encode(&bytes)[..128.min(base64_url_encode(&bytes).len())].to_string() -} - -/// 生成 PKCE code_challenge (SHA256) -fn generate_code_challenge(verifier: &str) -> String { - use sha2::{Digest, Sha256}; - let mut hasher = Sha256::new(); - hasher.update(verifier.as_bytes()); - let result = hasher.finalize(); - base64_url_encode(&result) -} - -/// 生成 OAuth state -fn generate_oauth_state() -> String { - use rand::Rng; - let mut rng = rand::thread_rng(); - let bytes: Vec = (0..32).map(|_| rng.gen()).collect(); - base64_url_encode(&bytes) -} - -/// Base64 URL 编码(无填充) -fn base64_url_encode(data: &[u8]) -> String { - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; - URL_SAFE_NO_PAD.encode(data) -} - -/// 启动 Kiro Social Auth 登录 (Google/GitHub) -/// -/// 使用 PKCE OAuth 流程进行登录 -/// 打开系统默认浏览器进行 OAuth 登录 -#[tauri::command] -pub async fn start_kiro_social_auth_login( - provider: String, -) -> Result { - // 验证 provider - let provider_normalized = match provider.to_lowercase().as_str() { - "google" => "Google", - "github" => "Github", - _ => { - return Ok(KiroSocialAuthLoginResponse { - success: false, - login_url: None, - state: None, - error: Some(format!("不支持的登录提供商: {provider}")), - }); - } - }; - - tracing::info!("[Kiro Social Auth] 开始 {} 登录流程", provider_normalized); - - // 生成 PKCE - let code_verifier = generate_code_verifier(); - let code_challenge = generate_code_challenge(&code_verifier); - let oauth_state = generate_oauth_state(); - - // 构建登录 URL - // 使用本地回调服务器接收授权码 - let redirect_uri = "http://127.0.0.1:19823/kiro-social-callback"; - - let login_url = format!( - "{}/login?idp={}&redirect_uri={}&code_challenge={}&code_challenge_method=S256&state={}", - KIRO_AUTH_ENDPOINT, - provider_normalized, - urlencoding::encode(redirect_uri), - urlencoding::encode(&code_challenge), - urlencoding::encode(&oauth_state) - ); - - tracing::info!("[Kiro Social Auth] 登录 URL: {}", login_url); - - // 保存登录状态(10 分钟过期) - let expires_at = chrono::Utc::now().timestamp() + 600; - { - let mut state = KIRO_SOCIAL_AUTH_LOGIN_STATE.write().await; - *state = Some(KiroSocialAuthLoginState { - provider: provider_normalized.to_string(), - code_verifier, - code_challenge, - oauth_state: oauth_state.clone(), - expires_at, - }); - } - - Ok(KiroSocialAuthLoginResponse { - success: true, - login_url: Some(login_url), - state: Some(oauth_state), - error: None, - }) -} - -/// 交换 Kiro Social Auth Token -/// -/// 用授权码交换 access_token 和 refresh_token -#[tauri::command] -pub async fn exchange_kiro_social_auth_token( - code: String, - state: String, -) -> Result { - tracing::info!("[Kiro Social Auth] 交换 Token..."); - - // 获取并验证登录状态 - let login_state = { - let state_guard = KIRO_SOCIAL_AUTH_LOGIN_STATE.read().await; - match state_guard.as_ref() { - Some(s) => s.clone(), - None => { - return Ok(KiroSocialAuthTokenResponse { - success: false, - error: Some("没有进行中的社交登录".to_string()), - }); - } - } - }; - - // 验证 state - if state != login_state.oauth_state { - // 清除状态 - { - let mut state_guard = KIRO_SOCIAL_AUTH_LOGIN_STATE.write().await; - *state_guard = None; - } - return Ok(KiroSocialAuthTokenResponse { - success: false, - error: Some("状态参数不匹配,可能存在安全风险".to_string()), - }); - } - - // 检查是否过期 - if chrono::Utc::now().timestamp() > login_state.expires_at { - // 清除状态 - { - let mut state_guard = KIRO_SOCIAL_AUTH_LOGIN_STATE.write().await; - *state_guard = None; - } - return Ok(KiroSocialAuthTokenResponse { - success: false, - error: Some("登录已过期,请重新开始".to_string()), - }); - } - - let redirect_uri = "http://127.0.0.1:19823/kiro-social-callback"; - - // 交换 Token - let client = reqwest::Client::new(); - let token_body = serde_json::json!({ - "code": code, - "code_verifier": login_state.code_verifier, - "redirect_uri": redirect_uri - }); - - let token_res = client - .post(format!("{KIRO_AUTH_ENDPOINT}/oauth/token")) - .header("Content-Type", "application/json") - .json(&token_body) - .send() - .await - .map_err(|e| format!("Token 交换请求失败: {e}"))?; - - if !token_res.status().is_success() { - let err_text = token_res.text().await.unwrap_or_default(); - // 清除状态 - { - let mut state_guard = KIRO_SOCIAL_AUTH_LOGIN_STATE.write().await; - *state_guard = None; - } - return Ok(KiroSocialAuthTokenResponse { - success: false, - error: Some(format!("Token 交换失败: {err_text}")), - }); - } - - let token_data: serde_json::Value = token_res - .json() - .await - .map_err(|e| format!("解析 Token 响应失败: {e}"))?; - - tracing::info!("[Kiro Social Auth] Token 交换成功!"); - - // 提取凭证 - let access_token = token_data["accessToken"].as_str().unwrap_or("").to_string(); - let refresh_token = token_data["refreshToken"] - .as_str() - .unwrap_or("") - .to_string(); - let profile_arn = token_data["profileArn"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expiresIn"].as_i64().unwrap_or(3600); - - // 创建凭证 JSON - let creds_json = serde_json::json!({ - "accessToken": access_token, - "refreshToken": refresh_token, - "profileArn": profile_arn, - "authMethod": "social", - "provider": login_state.provider, - "expiresAt": chrono::Utc::now().timestamp() + expires_in - }); - - // 保存到临时状态 - { - let mut creds = KIRO_BUILDER_ID_CREDENTIALS.write().await; - creds.insert("pending".to_string(), creds_json); - } - - // 清除登录状态 - { - let mut state_guard = KIRO_SOCIAL_AUTH_LOGIN_STATE.write().await; - *state_guard = None; - } - - Ok(KiroSocialAuthTokenResponse { - success: true, - error: None, - }) -} - -/// 取消 Kiro Social Auth 登录 -#[tauri::command] -pub async fn cancel_kiro_social_auth_login() -> Result { - tracing::info!("[Kiro Social Auth] 取消登录"); - { - let mut state = KIRO_SOCIAL_AUTH_LOGIN_STATE.write().await; - *state = None; - } - Ok(true) -} - -// ============ Playwright 指纹浏览器登录相关命令 ============ - -/// Playwright 可用性状态 -/// -/// Requirements: 2.1, 2.2 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct PlaywrightStatus { - /// 浏览器是否可用 - pub available: bool, - /// 浏览器可执行文件路径 - pub browser_path: Option, - /// 浏览器来源: "system" 或 "playwright" - pub browser_source: Option, - /// 错误信息 - pub error: Option, -} - -/// 获取系统 Chrome 可执行文件路径 -fn get_system_chrome_path() -> Option { - #[cfg_attr( - not(any(target_os = "macos", target_os = "windows")), - allow(unused_variables) - )] - let home = dirs::home_dir().unwrap_or_else(|| PathBuf::from(".")); - - #[cfg(target_os = "macos")] - { - let paths = [ - PathBuf::from("/Applications/Google Chrome.app/Contents/MacOS/Google Chrome"), - PathBuf::from("/Applications/Chromium.app/Contents/MacOS/Chromium"), - home.join("Applications/Google Chrome.app/Contents/MacOS/Google Chrome"), - ]; - for path in paths { - if path.exists() { - return Some(path.to_string_lossy().to_string()); - } - } - } - - #[cfg(target_os = "windows")] - { - let paths = [ - PathBuf::from("C:\\Program Files\\Google\\Chrome\\Application\\chrome.exe"), - PathBuf::from("C:\\Program Files (x86)\\Google\\Chrome\\Application\\chrome.exe"), - home.join("AppData\\Local\\Google\\Chrome\\Application\\chrome.exe"), - ]; - for path in paths { - if path.exists() { - return Some(path.to_string_lossy().to_string()); - } - } - } - - #[cfg(target_os = "linux")] - { - let paths = [ - PathBuf::from("/usr/bin/google-chrome"), - PathBuf::from("/usr/bin/google-chrome-stable"), - PathBuf::from("/usr/bin/chromium"), - PathBuf::from("/usr/bin/chromium-browser"), - PathBuf::from("/snap/bin/chromium"), - ]; - for path in paths { - if path.exists() { - return Some(path.to_string_lossy().to_string()); - } - } - } - - None -} - -/// 获取 Playwright 浏览器缓存目录 -fn get_playwright_cache_dir() -> PathBuf { - let home = dirs::home_dir().unwrap_or_else(|| PathBuf::from(".")); - - #[cfg(target_os = "macos")] - { - home.join("Library").join("Caches").join("ms-playwright") - } - - #[cfg(target_os = "windows")] - { - home.join("AppData").join("Local").join("ms-playwright") - } - - #[cfg(target_os = "linux")] - { - home.join(".cache").join("ms-playwright") - } - - #[cfg(not(any(target_os = "macos", target_os = "windows", target_os = "linux")))] - { - home.join(".cache").join("ms-playwright") - } -} - -/// 获取 Playwright Chromium 浏览器可执行文件路径 -/// -/// 搜索常见的 Chromium 版本目录 -fn get_playwright_browser_path() -> Option { - let cache_dir = get_playwright_cache_dir(); - - // Playwright 常见的 Chromium 版本目录 - let chromium_versions = [ - "chromium-1140", - "chromium-1134", - "chromium-1124", - "chromium-1117", - "chromium-1112", - "chromium-1108", - "chromium-1105", - "chromium-1097", - "chromium-1091", - "chromium-1084", - "chromium-1080", - "chromium-1076", - "chromium-1067", - "chromium-1060", - "chromium-1055", - "chromium-1048", - "chromium-1045", - "chromium-1041", - "chromium-1033", - "chromium-1028", - "chromium-1024", - "chromium-1020", - "chromium-1015", - "chromium-1012", - "chromium-1008", - "chromium-1005", - "chromium-1000", - "chromium", - ]; - - for version in chromium_versions { - #[cfg(target_os = "macos")] - let exec_path = cache_dir - .join(version) - .join("chrome-mac") - .join("Chromium.app") - .join("Contents") - .join("MacOS") - .join("Chromium"); - - #[cfg(target_os = "windows")] - let exec_path = cache_dir - .join(version) - .join("chrome-win") - .join("chrome.exe"); - - #[cfg(target_os = "linux")] - let exec_path = cache_dir.join(version).join("chrome-linux").join("chrome"); - - #[cfg(not(any(target_os = "macos", target_os = "windows", target_os = "linux")))] - let exec_path = cache_dir.join(version).join("chrome-linux").join("chrome"); - - if exec_path.exists() { - return Some(exec_path.to_string_lossy().to_string()); - } - } - - None -} - -/// 获取可用的浏览器路径(优先系统 Chrome) -fn get_available_browser_path() -> Option<(String, String)> { - // 优先使用系统 Chrome - if let Some(path) = get_system_chrome_path() { - return Some((path, "system".to_string())); - } - - // 其次使用 Playwright Chromium - if let Some(path) = get_playwright_browser_path() { - return Some((path, "playwright".to_string())); - } - - None -} - -/// 检查浏览器是否可用(优先系统 Chrome) -/// -/// 检测系统 Chrome 或 Playwright Chromium 是否存在 -/// Requirements: 2.1, 2.2 -#[tauri::command] -pub async fn check_playwright_available() -> Result { - tracing::info!("[Browser] 检查浏览器可用性..."); - - match get_available_browser_path() { - Some((browser_path, source)) => { - tracing::info!("[Browser] 找到 {} 浏览器: {}", source, browser_path); - Ok(PlaywrightStatus { - available: true, - browser_path: Some(browser_path), - browser_source: Some(source), - error: None, - }) - } - None => { - let error_msg = - "未找到可用的浏览器。请安装 Google Chrome 或运行: npx playwright install chromium" - .to_string(); - tracing::warn!("[Browser] {}", error_msg); - Ok(PlaywrightStatus { - available: false, - browser_path: None, - browser_source: None, - error: Some(error_msg), - }) - } - } -} - -/// Playwright 安装进度事件 -/// -/// 用于向前端发送安装进度信息 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct PlaywrightInstallProgress { - /// 进度消息 - pub message: String, - /// 是否完成 - pub done: bool, - /// 是否成功(仅在 done=true 时有效) - pub success: Option, -} - -/// 安装 Playwright Chromium 浏览器 -/// -/// 执行 npm install playwright && npx playwright install chromium -/// Requirements: 6.1, 6.2 -#[tauri::command] -pub async fn install_playwright(app: tauri::AppHandle) -> Result { - use tokio::io::{AsyncBufReadExt, BufReader}; - use tokio::process::Command; - - tracing::info!("[Playwright] 开始安装 Playwright..."); - - // 发送进度事件 - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: "正在查找 Playwright 脚本目录...".to_string(), - done: false, - success: None, - }, - ); - - // 尝试多个可能的脚本目录路径 - let possible_paths = vec![ - // 开发模式:从 CARGO_MANIFEST_DIR 推导 - PathBuf::from(env!("CARGO_MANIFEST_DIR")) - .parent() - .unwrap_or(&PathBuf::from(".")) - .join("scripts") - .join("playwright-login"), - // 生产模式:应用数据目录 - dirs::data_dir() - .unwrap_or_default() - .join("lime") - .join("scripts") - .join("playwright-login"), - // 当前工作目录 - std::env::current_dir() - .unwrap_or_default() - .join("scripts") - .join("playwright-login"), - ]; - - let mut script_dir: Option = None; - for path in &possible_paths { - tracing::info!("[Playwright] 检查路径: {:?}", path); - if path.join("package.json").exists() { - script_dir = Some(path.clone()); - break; - } - } - - let script_dir = match script_dir { - Some(dir) => dir, - None => { - let error = format!( - "找不到 Playwright 脚本目录。已检查路径:\n{}", - possible_paths - .iter() - .map(|p| format!(" - {p:?}")) - .collect::>() - .join("\n") - ); - tracing::error!("[Playwright] {}", error); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - }; - - tracing::info!("[Playwright] 使用脚本目录: {:?}", script_dir); - - // 步骤 1: 安装 npm 依赖 - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: format!("正在安装 npm 依赖... ({})", script_dir.display()), - done: false, - success: None, - }, - ); - - let npm_install = Command::new("npm") - .arg("install") - .current_dir(&script_dir) - .stdout(std::process::Stdio::piped()) - .stderr(std::process::Stdio::piped()) - .spawn(); - - match npm_install { - Ok(mut child) => { - // 收集 stderr 输出用于错误报告 - let mut stderr_output = String::new(); - if let Some(stderr) = child.stderr.take() { - let mut reader = BufReader::new(stderr).lines(); - while let Ok(Some(line)) = reader.next_line().await { - tracing::debug!("[Playwright npm] {}", line); - stderr_output.push_str(&line); - stderr_output.push('\n'); - } - } - - let status = child.wait().await; - match status { - Ok(s) if s.success() => { - tracing::info!("[Playwright] npm install 成功"); - // 发送成功消息 - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: "npm 依赖安装成功,准备安装 Chromium 浏览器...".to_string(), - done: false, - success: None, - }, - ); - } - Ok(s) => { - let error = if stderr_output.is_empty() { - format!("npm install 失败,退出码: {:?}", s.code()) - } else { - format!("npm install 失败: {}", stderr_output.trim()) - }; - tracing::error!("[Playwright] {}", error); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - Err(e) => { - let error = format!("npm install 执行失败: {e}"); - tracing::error!("[Playwright] {}", error); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - } - } - Err(e) => { - let error = format!("无法启动 npm: {e}。请确保已安装 Node.js"); - tracing::error!("[Playwright] {}", error); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - } - - // 步骤 2: 安装 Chromium 浏览器 - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: "正在安装 Chromium 浏览器 (npx playwright install chromium)...".to_string(), - done: false, - success: None, - }, - ); - - let playwright_install = Command::new("npx") - .args(["playwright", "install", "chromium"]) - .current_dir(&script_dir) - .stdout(std::process::Stdio::piped()) - .stderr(std::process::Stdio::piped()) - .spawn(); - - match playwright_install { - Ok(mut child) => { - // 同时收集 stdout 和 stderr - let mut stdout_output = String::new(); - let mut stderr_output = String::new(); - - // 读取 stdout 并发送进度 - if let Some(stdout) = child.stdout.take() { - let app_clone = app.clone(); - let mut reader = BufReader::new(stdout).lines(); - while let Ok(Some(line)) = reader.next_line().await { - tracing::info!("[Playwright install] {}", line); - stdout_output.push_str(&line); - stdout_output.push('\n'); - // 发送下载进度 - if line.contains("Downloading") - || line.contains("%") - || line.contains("chromium") - { - let _ = app_clone.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: line.clone(), - done: false, - success: None, - }, - ); - } - } - } - - // 读取 stderr - if let Some(stderr) = child.stderr.take() { - let mut reader = BufReader::new(stderr).lines(); - while let Ok(Some(line)) = reader.next_line().await { - tracing::warn!("[Playwright install stderr] {}", line); - stderr_output.push_str(&line); - stderr_output.push('\n'); - } - } - - let status = child.wait().await; - match status { - Ok(s) if s.success() => { - tracing::info!("[Playwright] Chromium 安装成功"); - } - Ok(s) => { - // 优先使用 stderr,如果为空则使用 stdout - let output = if !stderr_output.is_empty() { - stderr_output.trim().to_string() - } else if !stdout_output.is_empty() { - stdout_output.trim().to_string() - } else { - format!("退出码: {:?}", s.code()) - }; - let error = format!("Chromium 安装失败: {output}"); - tracing::error!("[Playwright] {}", error); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - Err(e) => { - let error = format!("Chromium 安装执行失败: {e}"); - tracing::error!("[Playwright] {}", error); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - } - } - Err(e) => { - let error = format!("无法启动 npx: {e}"); - tracing::error!("[Playwright] {}", error); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - } - - // 验证安装结果 - let status = check_playwright_available().await?; - - if status.available { - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: "Playwright 安装成功!".to_string(), - done: true, - success: Some(true), - }, - ); - tracing::info!( - "[Playwright] 安装完成,浏览器路径: {:?}", - status.browser_path - ); - } else { - let error = - "安装完成但未检测到浏览器,请手动运行: npx playwright install chromium".to_string(); - let _ = app.emit( - "playwright-install-progress", - PlaywrightInstallProgress { - message: error.clone(), - done: true, - success: Some(false), - }, - ); - return Err(error); - } - - Ok(status) -} - -/// Playwright 登录进度事件 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct PlaywrightLoginProgress { - pub message: String, -} - -/// Playwright 登录结果 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct PlaywrightLoginResult { - pub success: bool, - pub code: Option, - pub state: Option, - pub error: Option, -} - -/// 全局 Playwright 登录进程状态 -static PLAYWRIGHT_LOGIN_PROCESS: Lazy>> = - Lazy::new(|| RwLock::new(None)); - -/// 获取 Playwright 登录脚本路径 -fn get_playwright_script_path() -> PathBuf { - // 开发模式下使用项目目录中的脚本 - let dev_path = PathBuf::from(env!("CARGO_MANIFEST_DIR")) - .parent() - .unwrap_or(&PathBuf::from(".")) - .join("scripts") - .join("playwright-login") - .join("index.js"); - - if dev_path.exists() { - return dev_path; - } - - // 生产模式下使用打包的资源 - if let Some(data_dir) = dirs::data_dir() { - let prod_path = data_dir - .join("lime") - .join("scripts") - .join("playwright-login") - .join("index.js"); - if prod_path.exists() { - return prod_path; - } - } - - // 回退到开发路径 - dev_path -} - -/// 启动 Kiro Playwright 登录 -/// -/// 使用 Playwright 指纹浏览器进行 OAuth 登录 -/// Requirements: 3.1, 3.4, 3.5, 4.3, 4.4 -#[tauri::command] -pub async fn start_kiro_playwright_login( - app: tauri::AppHandle, - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - provider: String, - name: Option, -) -> Result { - use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; - use tokio::process::Command; - - // 验证 provider - let provider_normalized = match provider.to_lowercase().as_str() { - "google" => "Google", - "github" => "Github", - "builderid" => "BuilderId", - _ => { - return Err(format!("不支持的登录提供商: {provider}")); - } - }; - - tracing::info!("[Playwright Login] 开始 {} 登录流程", provider_normalized); - - // 检查 Playwright 是否可用 - let status = check_playwright_available().await?; - if !status.available { - return Err(status - .error - .unwrap_or_else(|| "Playwright 不可用".to_string())); - } - - // 生成 PKCE - let code_verifier = generate_code_verifier(); - let code_challenge = generate_code_challenge(&code_verifier); - let oauth_state = generate_oauth_state(); - - // 构建 OAuth URL - let redirect_uri = "http://localhost:19824/callback"; - let auth_url = format!( - "{}/login?idp={}&redirect_uri={}&code_challenge={}&code_challenge_method=S256&state={}", - KIRO_AUTH_ENDPOINT, - provider_normalized, - urlencoding::encode(redirect_uri), - urlencoding::encode(&code_challenge), - urlencoding::encode(&oauth_state) - ); - - tracing::info!("[Playwright Login] OAuth URL: {}", auth_url); - - // 获取脚本路径 - let script_path = get_playwright_script_path(); - if !script_path.exists() { - return Err(format!("Playwright 登录脚本不存在: {script_path:?}")); - } - - tracing::info!("[Playwright Login] 脚本路径: {:?}", script_path); - - // 启动 Node.js 进程 - let mut child = Command::new("node") - .arg(&script_path) - .stdin(std::process::Stdio::piped()) - .stdout(std::process::Stdio::piped()) - .stderr(std::process::Stdio::piped()) - .kill_on_drop(true) - .spawn() - .map_err(|e| format!("启动 Playwright 进程失败: {e}"))?; - - let stdin = child.stdin.take().ok_or("无法获取 stdin")?; - let stdout = child.stdout.take().ok_or("无法获取 stdout")?; - - // 保存进程引用 - { - let mut process_guard = PLAYWRIGHT_LOGIN_PROCESS.write().await; - *process_guard = Some(child); - } - - let mut stdin = tokio::io::BufWriter::new(stdin); - let mut reader = BufReader::new(stdout); - - // 等待就绪信号 - let mut line = String::new(); - reader - .read_line(&mut line) - .await - .map_err(|e| format!("读取就绪信号失败: {e}"))?; - - let ready_response: serde_json::Value = - serde_json::from_str(line.trim()).map_err(|e| format!("解析就绪信号失败: {e}"))?; - - if ready_response.get("action").and_then(|v| v.as_str()) != Some("ready") { - return Err("Playwright 脚本未就绪".to_string()); - } - - tracing::info!("[Playwright Login] Sidecar 已就绪"); - - // 发送登录请求 - let login_request = serde_json::json!({ - "action": "login", - "provider": provider_normalized, - "authUrl": auth_url, - "callbackUrl": redirect_uri - }); - - let request_str = - serde_json::to_string(&login_request).map_err(|e| format!("序列化请求失败: {e}"))?; - - stdin - .write_all(request_str.as_bytes()) - .await - .map_err(|e| format!("发送请求失败: {e}"))?; - stdin - .write_all(b"\n") - .await - .map_err(|e| format!("发送换行失败: {e}"))?; - stdin - .flush() - .await - .map_err(|e| format!("刷新 stdin 失败: {e}"))?; - - tracing::info!("[Playwright Login] 已发送登录请求"); - - // 读取响应 - let mut code: Option = None; - let mut state: Option = None; - - loop { - line.clear(); - match reader.read_line(&mut line).await { - Ok(0) => { - // EOF - break; - } - Ok(_) => { - let trimmed = line.trim(); - if trimmed.is_empty() { - continue; - } - - match serde_json::from_str::(trimmed) { - Ok(response) => { - let action = response - .get("action") - .and_then(|v| v.as_str()) - .unwrap_or(""); - let success = response - .get("success") - .and_then(|v| v.as_bool()) - .unwrap_or(false); - - match action { - "progress" => { - if let Some(data) = response.get("data") { - if let Some(message) = - data.get("message").and_then(|v| v.as_str()) - { - tracing::info!("[Playwright Login] 进度: {}", message); - let _ = app.emit( - "playwright-login-progress", - PlaywrightLoginProgress { - message: message.to_string(), - }, - ); - } - } - } - "login" => { - if success { - if let Some(data) = response.get("data") { - code = data - .get("code") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - state = data - .get("state") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - } - } else { - let error = response - .get("data") - .and_then(|d| d.get("error")) - .and_then(|v| v.as_str()) - .unwrap_or("未知错误"); - - // 清理进程 - { - let mut process_guard = - PLAYWRIGHT_LOGIN_PROCESS.write().await; - *process_guard = None; - } - - return Err(format!("Playwright 登录失败: {error}")); - } - break; - } - "error" => { - let error = response - .get("data") - .and_then(|d| d.get("error")) - .and_then(|v| v.as_str()) - .unwrap_or("未知错误"); - - // 清理进程 - { - let mut process_guard = PLAYWRIGHT_LOGIN_PROCESS.write().await; - *process_guard = None; - } - - return Err(format!("Playwright 错误: {error}")); - } - _ => {} - } - } - Err(e) => { - tracing::warn!("[Playwright Login] 解析响应失败: {} - {}", e, trimmed); - } - } - } - Err(e) => { - // 清理进程 - { - let mut process_guard = PLAYWRIGHT_LOGIN_PROCESS.write().await; - *process_guard = None; - } - return Err(format!("读取响应失败: {e}")); - } - } - } - - // 清理进程 - { - let mut process_guard = PLAYWRIGHT_LOGIN_PROCESS.write().await; - *process_guard = None; - } - - // 验证结果 - let auth_code = code.ok_or("未获取到授权码")?; - - // 验证 state - if let Some(returned_state) = &state { - if returned_state != &oauth_state { - return Err("状态参数不匹配,可能存在安全风险".to_string()); - } - } - - tracing::info!("[Playwright Login] 获取到授权码,开始交换 Token"); - - // 交换 Token - let client = reqwest::Client::new(); - let token_body = serde_json::json!({ - "code": auth_code, - "code_verifier": code_verifier, - "redirect_uri": redirect_uri - }); - - let token_res = client - .post(format!("{KIRO_AUTH_ENDPOINT}/oauth/token")) - .header("Content-Type", "application/json") - .json(&token_body) - .send() - .await - .map_err(|e| format!("Token 交换请求失败: {e}"))?; - - if !token_res.status().is_success() { - let err_text = token_res.text().await.unwrap_or_default(); - return Err(format!("Token 交换失败: {err_text}")); - } - - let token_data: serde_json::Value = token_res - .json() - .await - .map_err(|e| format!("解析 Token 响应失败: {e}"))?; - - tracing::info!("[Playwright Login] Token 交换成功!"); - - // 提取凭证 - let access_token = token_data["accessToken"].as_str().unwrap_or("").to_string(); - let refresh_token = token_data["refreshToken"] - .as_str() - .unwrap_or("") - .to_string(); - let profile_arn = token_data["profileArn"].as_str().map(|s| s.to_string()); - let expires_in = token_data["expiresIn"].as_i64().unwrap_or(3600); - - // 创建凭证 JSON - let creds_json = serde_json::json!({ - "accessToken": access_token, - "refreshToken": refresh_token, - "profileArn": profile_arn, - "authMethod": "social", - "provider": provider_normalized, - "loginMethod": "playwright", - "expiresAt": chrono::Utc::now().timestamp() + expires_in - }); - - // 将凭证 JSON 转换为字符串并创建凭证文件 - let json_content = - serde_json::to_string_pretty(&creds_json).map_err(|e| format!("序列化凭证失败: {e}"))?; - - let stored_file_path = create_kiro_credential_from_json(&json_content)?; - - // 添加到凭证池 - let credential = pool_service.0.add_credential( - &db, - "kiro", - CredentialData::KiroOAuth { - creds_file_path: stored_file_path, - }, - name, - Some(true), - None, - )?; - - tracing::info!("[Playwright Login] 凭证已添加到凭证池: {}", credential.uuid); - - Ok(credential) -} - -/// 取消 Kiro Playwright 登录 -/// -/// 终止正在进行的 Playwright 登录进程 -/// Requirements: 5.3 -#[tauri::command] -pub async fn cancel_kiro_playwright_login() -> Result { - tracing::info!("[Playwright Login] 取消登录"); - - let mut process_guard = PLAYWRIGHT_LOGIN_PROCESS.write().await; - - if let Some(mut child) = process_guard.take() { - // 尝试发送取消命令 - if let Some(mut stdin) = child.stdin.take() { - use tokio::io::AsyncWriteExt; - - let cancel_request = serde_json::json!({ - "action": "cancel" - }); - - if let Ok(request_str) = serde_json::to_string(&cancel_request) { - let _ = stdin.write_all(request_str.as_bytes()).await; - let _ = stdin.write_all(b"\n").await; - let _ = stdin.flush().await; - } - } - - // 等待一小段时间让进程优雅退出 - tokio::time::sleep(tokio::time::Duration::from_millis(500)).await; - - // 强制终止进程 - let _ = child.kill().await; - - tracing::info!("[Playwright Login] 登录进程已终止"); - Ok(true) - } else { - tracing::info!("[Playwright Login] 没有正在进行的登录"); - Ok(false) - } -} - -/// 启动 Kiro Social Auth 回调服务器 -/// -/// 启动一个本地 HTTP 服务器来接收 OAuth 回调 -#[tauri::command] -pub async fn start_kiro_social_auth_callback_server(app: tauri::AppHandle) -> Result { - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - use tokio::net::TcpListener; - - tracing::info!("[Kiro Social Auth] 启动回调服务器..."); - - // 尝试绑定端口 - let listener = TcpListener::bind("127.0.0.1:19823") - .await - .map_err(|e| format!("无法启动回调服务器: {e}"))?; - - tracing::info!("[Kiro Social Auth] 回调服务器已启动在 127.0.0.1:19823"); - - // 在后台处理连接 - let app_handle = app.clone(); - tokio::spawn(async move { - // 只处理一个连接 - if let Ok((mut socket, _)) = listener.accept().await { - let mut buffer = [0u8; 4096]; - if let Ok(n) = socket.read(&mut buffer).await { - let request = String::from_utf8_lossy(&buffer[..n]); - - // 解析请求获取 code 和 state - if let Some(path_line) = request.lines().next() { - if let Some(path) = path_line.split_whitespace().nth(1) { - if path.starts_with("/kiro-social-callback") { - // 解析查询参数 - let mut code = None; - let mut state = None; - - if let Some(query_start) = path.find('?') { - let query = &path[query_start + 1..]; - for param in query.split('&') { - let parts: Vec<&str> = param.splitn(2, '=').collect(); - if parts.len() == 2 { - match parts[0] { - "code" => { - code = Some( - urlencoding::decode(parts[1]) - .unwrap_or_default() - .to_string(), - ) - } - "state" => { - state = Some( - urlencoding::decode(parts[1]) - .unwrap_or_default() - .to_string(), - ) - } - _ => {} - } - } - } - } - - // 发送成功响应页面 - let html = r#" - - - - 登录成功 - - - -
-

✓ 登录成功

-

您可以关闭此窗口并返回应用

-
- -"#; - - let response = format!( - "HTTP/1.1 200 OK\r\nContent-Type: text/html; charset=utf-8\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", - html.len(), - html - ); - - let _ = socket.write_all(response.as_bytes()).await; - - // 发送事件到前端 - if let (Some(code), Some(state)) = (code, state) { - let _ = app_handle.emit( - "kiro-social-auth-callback", - serde_json::json!({ - "code": code, - "state": state - }), - ); - } - } - } - } - } - } - }); - - Ok(true) -} - -// ============ Playwright 可用性检测测试 ============ - -#[cfg(test)] -mod playwright_tests { - use super::*; - - /// **Property 1: Playwright 可用性检测正确性** - /// **Validates: Requirements 2.2** - /// - /// *For any* 文件系统状态,Playwright 可用性检测函数应该: - /// - 当 Playwright 浏览器可执行文件存在时返回 `available: true` - /// - 当可执行文件不存在时返回 `available: false` - /// - 返回的 `browserPath` 应该是实际检测到的路径或 `None` - - #[test] - fn test_get_playwright_cache_dir_returns_valid_path() { - // Feature: playwright-fingerprint-login, Property 1: Playwright 可用性检测正确性 - // 测试缓存目录路径生成 - let cache_dir = get_playwright_cache_dir(); - - // 路径应该包含 ms-playwright - assert!( - cache_dir.to_string_lossy().contains("ms-playwright"), - "缓存目录应包含 ms-playwright: {cache_dir:?}" - ); - - // 路径应该是绝对路径或相对于 home 目录 - #[cfg(target_os = "macos")] - assert!( - cache_dir.to_string_lossy().contains("Library/Caches"), - "macOS 缓存目录应在 Library/Caches 下: {cache_dir:?}" - ); - - #[cfg(target_os = "windows")] - assert!( - cache_dir.to_string_lossy().contains("AppData\\Local"), - "Windows 缓存目录应在 AppData\\Local 下: {:?}", - cache_dir - ); - - #[cfg(target_os = "linux")] - assert!( - cache_dir.to_string_lossy().contains(".cache"), - "Linux 缓存目录应在 .cache 下: {:?}", - cache_dir - ); - } - - #[test] - fn test_get_playwright_browser_path_returns_none_when_not_installed() { - // Feature: playwright-fingerprint-login, Property 1: Playwright 可用性检测正确性 - // 当 Playwright 未安装时,应返回 None - // 注意:这个测试在 Playwright 已安装的环境中可能会失败 - // 我们主要测试函数不会 panic - let result = get_playwright_browser_path(); - - // 函数应该正常返回(不 panic) - // 结果可能是 Some 或 None,取决于环境 - match result { - Some(path) => { - // 如果找到了路径,验证路径格式 - assert!(!path.is_empty(), "浏览器路径不应为空"); - assert!( - path.contains("chromium") - || path.contains("Chromium") - || path.contains("chrome"), - "路径应包含 chromium/chrome: {path}" - ); - } - None => { - // 未找到浏览器,这是预期的情况之一 - } - } - } - - #[test] - fn test_playwright_status_serialization() { - // Feature: playwright-fingerprint-login, Property 1: Playwright 可用性检测正确性 - // 测试 PlaywrightStatus 结构体的序列化 - - // 测试可用状态 - let available_status = PlaywrightStatus { - available: true, - browser_path: Some("/path/to/chromium".to_string()), - browser_source: Some("playwright".to_string()), - error: None, - }; - - let json = serde_json::to_string(&available_status).unwrap(); - assert!(json.contains("\"available\":true")); - assert!(json.contains("\"browser_path\":\"/path/to/chromium\"")); - - // 测试不可用状态 - let unavailable_status = PlaywrightStatus { - available: false, - browser_path: None, - browser_source: None, - error: Some("未安装".to_string()), - }; - - let json = serde_json::to_string(&unavailable_status).unwrap(); - assert!(json.contains("\"available\":false")); - assert!(json.contains("\"error\":\"未安装\"")); - } - - #[test] - fn test_playwright_status_deserialization() { - // Feature: playwright-fingerprint-login, Property 1: Playwright 可用性检测正确性 - // 测试 PlaywrightStatus 结构体的反序列化 - - let json = r#"{"available":true,"browser_path":"/test/path","error":null}"#; - let status: PlaywrightStatus = serde_json::from_str(json).unwrap(); - - assert!(status.available); - assert_eq!(status.browser_path, Some("/test/path".to_string())); - assert!(status.error.is_none()); - } - - #[test] - fn test_playwright_status_invariants() { - // Feature: playwright-fingerprint-login, Property 1: Playwright 可用性检测正确性 - // 测试状态不变量: - // - 当 available=true 时,browser_path 应该有值 - // - 当 available=false 时,error 应该有值 - - // 可用状态的不变量 - let available_status = PlaywrightStatus { - available: true, - browser_path: Some("/path".to_string()), - browser_source: Some("system".to_string()), - error: None, - }; - assert!( - available_status.available && available_status.browser_path.is_some(), - "可用状态应有 browser_path" - ); - - // 不可用状态的不变量 - let unavailable_status = PlaywrightStatus { - available: false, - browser_path: None, - browser_source: None, - error: Some("错误".to_string()), - }; - assert!( - !unavailable_status.available && unavailable_status.error.is_some(), - "不可用状态应有 error" - ); - } -} - -/// 获取单个凭证的健康状态 -/// Requirements: 4.4 -#[tauri::command] -pub async fn get_credential_health( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - uuid: String, -) -> Result, String> { - pool_service.0.get_credential_health(&db, &uuid) -} - -/// 获取所有凭证的健康状态 -/// Requirements: 4.4 -#[tauri::command] -pub async fn get_all_credential_health( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, -) -> Result, String> { - pool_service.0.get_all_credential_health(&db) -} diff --git a/src-tauri/src/commands/usage_cmd.rs b/src-tauri/src/commands/usage_cmd.rs deleted file mode 100644 index 9cd663260..000000000 --- a/src-tauri/src/commands/usage_cmd.rs +++ /dev/null @@ -1,351 +0,0 @@ -//! Usage Tauri 命令 -//! -//! 提供 Kiro 用量查询的 Tauri 命令接口。 - -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::database::DbConnection; -use crate::models::provider_pool_model::{CredentialData, PoolProviderType}; -use crate::TokenCacheServiceState; -use lime_services::usage_service::{self, UsageInfo}; -use tauri::State; - -/// 默认 Kiro 版本号 -const DEFAULT_KIRO_VERSION: &str = "1.0.0"; - -/// 获取 Kiro 用量信息 -/// -/// **Validates: Requirements 1.1** -/// -/// # Arguments -/// * `credential_uuid` - 凭证的 UUID -/// * `db` - 数据库连接 -/// * `token_cache` - Token 缓存服务 -/// -/// # Returns -/// * `Ok(UsageInfo)` - 成功时返回用量信息 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_kiro_usage( - credential_uuid: String, - db: State<'_, DbConnection>, - token_cache: State<'_, TokenCacheServiceState>, -) -> Result { - // 1. 获取凭证信息 - let credential = { - let conn = db.lock().map_err(|e| e.to_string())?; - ProviderPoolDao::get_by_uuid(&conn, &credential_uuid) - .map_err(|e| e.to_string())? - .ok_or_else(|| format!("凭证不存在: {credential_uuid}"))? - }; - - // 2. 验证是否为 Kiro 凭证 - if credential.provider_type != PoolProviderType::Kiro { - return Err(format!( - "不支持的凭证类型: {:?},仅支持 Kiro 凭证", - credential.provider_type - )); - } - - // 3. 获取凭证文件路径 - let creds_file_path = match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(), - _ => return Err("凭证数据类型不匹配".to_string()), - }; - - // 4. 获取有效的 access_token - let access_token = token_cache - .0 - .get_valid_token(&db, &credential_uuid) - .await - .map_err(|e| { - // 提供更友好的错误信息 - if e.contains("401") || e.contains("Bad credentials") || e.contains("过期") || e.contains("无效") { - format!("刷新 Kiro Token 失败: OAuth 凭证已过期或无效,需要重新认证。\n💡 解决方案:\n1. 删除当前 OAuth 凭证\n2. 重新添加 OAuth 凭证\n3. 确保使用最新的凭证文件\n\n技术详情:{e}") - } else { - e - } - })?; - - // 5. 从凭证文件读取 auth_method 和 profile_arn - let (auth_method, profile_arn) = read_kiro_credential_info(&creds_file_path)?; - - // 6. 获取 machine_id - let machine_id = get_machine_id()?; - - // 7. 调用 Usage API - let usage_info = usage_service::get_usage_limits_safe( - &access_token, - &auth_method, - profile_arn.as_deref(), - &machine_id, - DEFAULT_KIRO_VERSION, - ) - .await; - - Ok(usage_info) -} - -/// 从 Kiro 凭证文件读取 auth_method 和 profile_arn -fn read_kiro_credential_info(creds_file_path: &str) -> Result<(String, Option), String> { - // 展开 ~ 路径 - let expanded_path = expand_tilde(creds_file_path); - - // 读取文件 - let content = - std::fs::read_to_string(&expanded_path).map_err(|e| format!("读取凭证文件失败: {e}"))?; - - // 解析 JSON - let json: serde_json::Value = - serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {e}"))?; - - // 获取 auth_method,默认为 "social" - let auth_method = json - .get("authMethod") - .and_then(|v| v.as_str()) - .unwrap_or("social") - .to_string(); - - // 获取 profile_arn(可选) - let profile_arn = json - .get("profileArn") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - - Ok((auth_method, profile_arn)) -} - -/// 展开路径中的 ~ 为用户主目录 -fn expand_tilde(path: &str) -> String { - if let Some(stripped) = path.strip_prefix("~/") { - if let Some(home) = dirs::home_dir() { - return home.join(stripped).to_string_lossy().to_string(); - } - } - path.to_string() -} - -/// 获取设备 ID(SHA256 哈希) -fn get_machine_id() -> Result { - // 尝试获取系统 machine-id - let raw_id = get_raw_machine_id()?; - - // 计算 SHA256 哈希 - use sha2::{Digest, Sha256}; - let mut hasher = Sha256::new(); - hasher.update(raw_id.as_bytes()); - let result = hasher.finalize(); - - Ok(format!("{result:x}")) -} - -/// 获取原始设备 ID -fn get_raw_machine_id() -> Result { - #[cfg(target_os = "macos")] - { - // macOS: 使用 IOPlatformUUID - use std::process::Command; - let output = Command::new("ioreg") - .args(["-rd1", "-c", "IOPlatformExpertDevice"]) - .output() - .map_err(|e| format!("执行 ioreg 失败: {e}"))?; - - let stdout = String::from_utf8_lossy(&output.stdout); - for line in stdout.lines() { - if line.contains("IOPlatformUUID") { - if let Some(uuid) = line.split('"').nth(3) { - return Ok(uuid.to_string()); - } - } - } - Err("无法获取 IOPlatformUUID".to_string()) - } - - #[cfg(target_os = "linux")] - { - // Linux: 读取 /etc/machine-id - std::fs::read_to_string("/etc/machine-id") - .map(|s| s.trim().to_string()) - .map_err(|e| format!("读取 /etc/machine-id 失败: {}", e)) - } - - #[cfg(target_os = "windows")] - { - // Windows: 使用注册表中的 MachineGuid - use std::os::windows::process::CommandExt; - use std::process::Command; - let output = Command::new("reg") - .args([ - "query", - "HKEY_LOCAL_MACHINE\\SOFTWARE\\Microsoft\\Cryptography", - "/v", - "MachineGuid", - ]) - .creation_flags(0x08000000) // CREATE_NO_WINDOW - .output() - .map_err(|e| format!("执行 reg query 失败: {}", e))?; - - let stdout = String::from_utf8_lossy(&output.stdout); - for line in stdout.lines() { - if line.contains("MachineGuid") { - if let Some(guid) = line.split_whitespace().last() { - return Ok(guid.to_string()); - } - } - } - Err("无法获取 MachineGuid".to_string()) - } - - #[cfg(not(any(target_os = "macos", target_os = "linux", target_os = "windows")))] - { - Err("不支持的操作系统".to_string()) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_expand_tilde() { - let path = "~/test/path"; - let expanded = expand_tilde(path); - assert!(!expanded.starts_with("~/")); - assert!(expanded.ends_with("test/path")); - } - - #[test] - fn test_expand_tilde_no_tilde() { - let path = "/absolute/path"; - let expanded = expand_tilde(path); - assert_eq!(expanded, path); - } - - #[test] - fn test_get_machine_id() { - // 这个测试在不同平台上行为不同 - let result = get_machine_id(); - // 应该能成功获取 machine_id - assert!(result.is_ok(), "Failed to get machine_id: {result:?}"); - // machine_id 应该是 64 字符的十六进制字符串(SHA256) - let id = result.unwrap(); - assert_eq!(id.len(), 64, "Machine ID should be 64 hex chars"); - assert!( - id.chars().all(|c| c.is_ascii_hexdigit()), - "Machine ID should be hex" - ); - } -} - -// ============================================================================ -// 集成测试 -// ============================================================================ - -#[cfg(test)] -mod integration_tests { - use super::*; - - /// 测试 read_kiro_credential_info 函数 - /// 验证能正确解析 Kiro 凭证文件中的 auth_method 和 profile_arn - #[test] - fn test_read_kiro_credential_info_social() { - // 创建临时文件 - let temp_dir = std::env::temp_dir(); - let temp_file = temp_dir.join("test_kiro_creds_social.json"); - - let creds_json = serde_json::json!({ - "accessToken": "test_access_token", - "refreshToken": "test_refresh_token", - "authMethod": "social", - "profileArn": "arn:aws:iam::123456789:profile/test" - }); - - std::fs::write(&temp_file, serde_json::to_string(&creds_json).unwrap()).unwrap(); - - let result = read_kiro_credential_info(temp_file.to_str().unwrap()); - assert!(result.is_ok()); - - let (auth_method, profile_arn) = result.unwrap(); - assert_eq!(auth_method, "social"); - assert_eq!( - profile_arn, - Some("arn:aws:iam::123456789:profile/test".to_string()) - ); - - // 清理 - let _ = std::fs::remove_file(&temp_file); - } - - /// 测试 read_kiro_credential_info 函数 - IdC 认证 - #[test] - fn test_read_kiro_credential_info_idc() { - let temp_dir = std::env::temp_dir(); - let temp_file = temp_dir.join("test_kiro_creds_idc.json"); - - let creds_json = serde_json::json!({ - "accessToken": "test_access_token", - "refreshToken": "test_refresh_token", - "authMethod": "idc" - }); - - std::fs::write(&temp_file, serde_json::to_string(&creds_json).unwrap()).unwrap(); - - let result = read_kiro_credential_info(temp_file.to_str().unwrap()); - assert!(result.is_ok()); - - let (auth_method, profile_arn) = result.unwrap(); - assert_eq!(auth_method, "idc"); - assert_eq!(profile_arn, None); - - // 清理 - let _ = std::fs::remove_file(&temp_file); - } - - /// 测试 read_kiro_credential_info 函数 - 默认 auth_method - #[test] - fn test_read_kiro_credential_info_default_auth_method() { - let temp_dir = std::env::temp_dir(); - let temp_file = temp_dir.join("test_kiro_creds_default.json"); - - // 没有 authMethod 字段,应该默认为 "social" - let creds_json = serde_json::json!({ - "accessToken": "test_access_token", - "refreshToken": "test_refresh_token" - }); - - std::fs::write(&temp_file, serde_json::to_string(&creds_json).unwrap()).unwrap(); - - let result = read_kiro_credential_info(temp_file.to_str().unwrap()); - assert!(result.is_ok()); - - let (auth_method, profile_arn) = result.unwrap(); - assert_eq!(auth_method, "social"); - assert_eq!(profile_arn, None); - - // 清理 - let _ = std::fs::remove_file(&temp_file); - } - - /// 测试 read_kiro_credential_info 函数 - 文件不存在 - #[test] - fn test_read_kiro_credential_info_file_not_found() { - let result = read_kiro_credential_info("/nonexistent/path/to/creds.json"); - assert!(result.is_err()); - assert!(result.unwrap_err().contains("读取凭证文件失败")); - } - - /// 测试 read_kiro_credential_info 函数 - 无效 JSON - #[test] - fn test_read_kiro_credential_info_invalid_json() { - let temp_dir = std::env::temp_dir(); - let temp_file = temp_dir.join("test_kiro_creds_invalid.json"); - - std::fs::write(&temp_file, "not valid json").unwrap(); - - let result = read_kiro_credential_info(temp_file.to_str().unwrap()); - assert!(result.is_err()); - assert!(result.unwrap_err().contains("解析凭证文件失败")); - - // 清理 - let _ = std::fs::remove_file(&temp_file); - } -} diff --git a/src-tauri/src/config/tests.rs b/src-tauri/src/config/tests.rs index e839a1da0..4f0be8f73 100644 --- a/src-tauri/src/config/tests.rs +++ b/src-tauri/src/config/tests.rs @@ -2461,7 +2461,6 @@ fn arb_valid_provider_type() -> impl Strategy { Just("gemini".to_string()), Just("openai".to_string()), Just("claude".to_string()), - Just("antigravity".to_string()), Just("vertex".to_string()), Just("gemini_api_key".to_string()), Just("codex".to_string()), diff --git a/src-tauri/src/dev_bridge.rs b/src-tauri/src/dev_bridge.rs index bbb32cf30..11b62705b 100644 --- a/src-tauri/src/dev_bridge.rs +++ b/src-tauri/src/dev_bridge.rs @@ -36,7 +36,7 @@ use lime_infra::telemetry::StatsAggregator; #[cfg(debug_assertions)] use lime_services::{ api_key_provider_service::ApiKeyProviderService, model_registry_service::ModelRegistryService, - provider_pool_service::ProviderPoolService, skill_service::SkillService, + skill_service::SkillService, }; #[cfg(debug_assertions)] use tauri::{AppHandle, EventId, Listener}; @@ -69,7 +69,6 @@ pub struct DevBridgeState { pub server: app::AppState, pub logs: app::LogState, pub db: Option, - pub pool_service: Arc, pub api_key_provider_service: Arc, pub connect_state: Arc>>, pub model_registry: Arc>>, @@ -130,7 +129,6 @@ impl DevBridgeServer { server: app::AppState, logs: app::LogState, db: Option, - pool_service: Arc, api_key_provider_service: Arc, connect_state: Arc>>, model_registry: Arc>>, @@ -144,7 +142,6 @@ impl DevBridgeServer { server, logs, db, - pool_service, api_key_provider_service, connect_state, model_registry, @@ -249,7 +246,9 @@ async fn stream_events( let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::(); let listener_event_name = event_name.clone(); - let listener_id = app_handle.listen_any(listener_event_name.clone(), move |event| { + // 只监听 AppHandle 事件目标;listen_any 会同时收到 app/window 目标,浏览器 SSE 会把 + // 同一个 runtime delta 转发两次,导致流式文字逐 token 重复。 + let listener_id = app_handle.listen(listener_event_name.clone(), move |event| { let payload = event.payload(); let payload_value = serde_json::from_str::(payload) .unwrap_or_else(|_| serde_json::Value::String(payload.to_string())); diff --git a/src-tauri/src/dev_bridge/dispatcher.rs b/src-tauri/src/dev_bridge/dispatcher.rs index a07c7a147..dbe61be34 100644 --- a/src-tauri/src/dev_bridge/dispatcher.rs +++ b/src-tauri/src/dev_bridge/dispatcher.rs @@ -218,9 +218,6 @@ mod tests { &config.logging, ))), db: Some(make_test_db()), - pool_service: Arc::new( - lime_services::provider_pool_service::ProviderPoolService::new(), - ), api_key_provider_service: Arc::new( lime_services::api_key_provider_service::ApiKeyProviderService::new(), ), @@ -835,32 +832,6 @@ mod tests { assert!(error.to_string().contains("模型注册服务未初始化")); } - #[tokio::test] - async fn provider_pool_model_commands_are_bridged() { - let state = make_test_state(); - - let models_by_provider = handle_command(&state, "get_all_models_by_provider", None) - .await - .expect("models by provider should route through dev bridge"); - assert!(models_by_provider.is_object()); - - let available_models = handle_command(&state, "get_all_available_models", None) - .await - .expect("available models should route through dev bridge"); - assert!(available_models.is_array()); - - let default_models = handle_command( - &state, - "get_default_models_for_provider", - Some(serde_json::json!({ - "providerType": "openai" - })), - ) - .await - .expect("default models should route through dev bridge"); - assert!(default_models.is_array()); - } - #[tokio::test] async fn test_api_key_provider_connection_is_bridged() { let state = make_test_state(); diff --git a/src-tauri/src/dev_bridge/dispatcher/models.rs b/src-tauri/src/dev_bridge/dispatcher/models.rs index 2bce57444..d54d84793 100644 --- a/src-tauri/src/dev_bridge/dispatcher/models.rs +++ b/src-tauri/src/dev_bridge/dispatcher/models.rs @@ -1,4 +1,4 @@ -use super::{args_or_default, get_db, get_string_arg}; +use super::{args_or_default, get_string_arg}; use crate::dev_bridge::DevBridgeState; use lime_server_utils::load_model_registry_provider_ids_from_resources; use lime_services::model_registry_service::ModelRegistryService; @@ -68,24 +68,6 @@ pub(super) async fn try_handle( "get_model_registry_provider_ids" => { serde_json::to_value(load_model_registry_provider_ids_from_resources()?)? } - "get_all_models_by_provider" => { - let model_service = lime_services::model_service::ModelService::new(); - serde_json::to_value(model_service.get_all_models_by_provider(get_db(state)?)?)? - } - "get_all_available_models" => { - let model_service = lime_services::model_service::ModelService::new(); - serde_json::to_value(model_service.get_all_available_models(get_db(state)?)?)? - } - "get_default_models_for_provider" => { - let args = args_or_default(args); - let provider_type = get_string_arg(&args, "providerType", "provider_type")?; - let parsed_provider_type: crate::models::provider_pool_model::PoolProviderType = - provider_type.parse().map_err(|err: String| err)?; - let model_service = lime_services::model_service::ModelService::new(); - serde_json::to_value( - model_service.get_default_models_for_provider(&parsed_provider_type), - )? - } "fetch_provider_models_auto" => { let args = args_or_default(args); let provider_id = get_string_arg(&args, "providerId", "provider_id")?; diff --git a/src-tauri/src/dev_bridge/dispatcher/providers.rs b/src-tauri/src/dev_bridge/dispatcher/providers.rs index 466cce373..09f665891 100644 --- a/src-tauri/src/dev_bridge/dispatcher/providers.rs +++ b/src-tauri/src/dev_bridge/dispatcher/providers.rs @@ -144,33 +144,6 @@ pub(super) async fn try_handle( .await?, )? } - "aster_agent_configure_from_pool" => { - let app_handle = require_app_handle(state)?; - let aster_state = app_handle.state::(); - let db = app_handle.state::(); - let args = args_or_default(args); - let request = parse_nested_arg::< - crate::commands::aster_agent_cmd::ConfigureFromPoolRequest, - >(&args, "request")?; - let session_id = get_string_arg(&args, "session_id", "sessionId")?; - - serde_json::to_value( - crate::commands::aster_agent_cmd::aster_agent_configure_from_pool( - aster_state, - db, - request, - session_id, - ) - .await?, - )? - } - "get_provider_pool_overview" => { - if let Some(db) = &state.db { - serde_json::to_value(state.pool_service.get_overview(db)?)? - } else { - serde_json::json!([]) - } - } "get_api_key_providers" => { if let Some(db) = &state.db { let providers = state.api_key_provider_service.get_all_providers(db)?; @@ -285,17 +258,6 @@ pub(super) async fn try_handle( .await?, )? } - "get_provider_pool_credentials" => { - if let Some(db) = &state.db { - let conn = db.lock().map_err(|e| e.to_string())?; - let credentials = - crate::database::dao::provider_pool::ProviderPoolDao::get_all(&conn) - .unwrap_or_default(); - serde_json::to_value(credentials)? - } else { - serde_json::json!([]) - } - } "get_provider_ui_state" => { let args = args_or_default(args); let key = get_string_arg(&args, "key", "key")?; diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index e3672059d..88e4fa8e4 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -25,7 +25,6 @@ pub use lime_providers::providers; // 从 core crate 重新导出(保持 crate::xxx 路径兼容) pub use lime_core::connect; pub use lime_core::content; -pub use lime_core::credential; pub use lime_core::database; pub use lime_core::memory; pub use lime_core::session_files; @@ -68,8 +67,7 @@ use lime_core::models; mod tests; // 重新导出核心类型以保持向后兼容 -pub use app::{AppState, LogState, ProviderType, TokenCacheServiceState, TrayManagerState}; -pub use lime_services::provider_pool_service::ProviderPoolService; +pub use app::{AppState, LogState, ProviderType, TrayManagerState}; // 重新导出 run 函数(main.rs 入口) pub use app::run; diff --git a/src-tauri/src/services/README.md b/src-tauri/src/services/README.md index e71d6931a..3a8916486 100644 --- a/src-tauri/src/services/README.md +++ b/src-tauri/src/services/README.md @@ -5,20 +5,17 @@ ## 架构说明 业务服务层,封装核心业务逻辑。 -提供凭证池管理、Token 缓存、MCP 同步等功能。 +提供 API Key Provider、模型注册表、MCP 同步等功能。旧凭证池、Token 缓存与 Kiro 事件服务已退役。 ## 文件索引 - `mod.rs` - 模块入口 - `site_adapter_import_service.rs` - 外部适配器来源导入与 Lime 标准编译层 -- `provider_pool_service.rs` - Provider 凭证池服务(多凭证轮询) -- `token_cache_service.rs` - Token 缓存服务 - `mcp_service.rs` - MCP 服务器管理 - `mcp_sync.rs` - MCP 配置同步 - `prompt_service.rs` - Prompt 管理服务 - `prompt_sync.rs` - Prompt 同步 - `skill_service.rs` - 技能管理服务 -- `usage_service.rs` - 使用量统计服务 - `backup_service.rs` - 备份服务 - `live_sync.rs` - 实时同步服务 - `switch.rs` - 开关服务 @@ -30,7 +27,6 @@ - `browser_runtime_window.rs` - 浏览器运行时调试独立窗口管理 - `general_chat/` - 通用对话服务模块(会话管理、消息存储) - `api_key_provider_service.rs` - API Key Provider 服务 -- `kiro_event_service.rs` - Kiro 事件服务 - `machine_id_service.rs` - 机器 ID 服务 - `model_registry_service.rs` - 模型注册表服务 - `persona_service.rs` - 人设服务(创建、列表、更新、删除、设置默认、模板) diff --git a/src-tauri/src/skills/README.md b/src-tauri/src/skills/README.md index 6ee0132df..0fc585b7c 100644 --- a/src-tauri/src/skills/README.md +++ b/src-tauri/src/skills/README.md @@ -55,12 +55,12 @@ AI 通过 SkillTool 发现新 Skill ### LimeLlmProvider -使用 ProviderPoolService 选择凭证并调用 LLM API。 +使用 API Key Provider 选择凭证并调用 LLM API。 **功能**: -- 通过 ProviderPoolService 选择可用凭证 +- 通过 API Key Provider 选择可用凭证 - 支持指定 provider 类型和 model 参数 -- 智能降级到 API Key Provider +- 不再读取旧凭证池、OAuth token 文件或本地 CLI 凭证目录 ### TauriExecutionCallback @@ -101,7 +101,6 @@ skills/ │ └── 社媒产物 Tool/Artifact 事件补投影 ├── llm_provider.rs (桥接) │ └── crates/skills/src/lime_llm_provider.rs -│ ├── ProviderPoolService (凭证池管理) │ └── ApiKeyProviderService (API Key 服务) └── execution_callback.rs └── tauri::AppHandle (事件发送) diff --git a/src-tauri/src/tests/credential_tests.rs b/src-tauri/src/tests/credential_tests.rs deleted file mode 100644 index 2d74d0879..000000000 --- a/src-tauri/src/tests/credential_tests.rs +++ /dev/null @@ -1,1926 +0,0 @@ -//! 凭证池属性测试 -//! -//! 使用 proptest 进行属性测试 - -#![allow(dead_code)] - -use crate::ProviderType; -use lime_core::credential::{Credential, CredentialData, CredentialPool}; -use lime_credential::{BalanceStrategy, LoadBalancer}; -use proptest::prelude::*; -use std::collections::HashSet; -use std::sync::Arc; - -/// 生成随机的 ProviderType -fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - ] -} - -/// 生成随机的 CredentialData -fn arb_credential_data() -> impl Strategy { - prop_oneof![ - // OAuth 凭证 - ("[a-zA-Z0-9]{10,50}", prop::option::of("[a-zA-Z0-9]{10,50}")).prop_map( - |(access_token, refresh_token)| { - CredentialData::OAuth { - access_token, - refresh_token, - expires_at: None, - } - } - ), - // API Key 凭证 - ( - "[a-zA-Z0-9]{10,50}", - prop::option::of("https?://[a-z]+\\.[a-z]+") - ) - .prop_map(|(key, base_url)| { CredentialData::ApiKey { key, base_url } }), - ] -} - -/// 生成随机的 Credential -fn arb_credential() -> impl Strategy { - ( - "[a-zA-Z0-9_-]{1,32}", // id - arb_provider_type(), - arb_credential_data(), - ) - .prop_map(|(id, provider, data)| Credential::new(id, provider, data)) -} - -/// 生成具有唯一 ID 的凭证列表 -fn arb_unique_credentials(max_count: usize) -> impl Strategy> { - prop::collection::vec(arb_credential(), 1..=max_count).prop_map(|creds| { - // 确保 ID 唯一 - let mut seen = std::collections::HashSet::new(); - creds - .into_iter() - .filter(|c| seen.insert(c.id.clone())) - .collect() - }) -} - -proptest! { - /// **Feature: enhancement-roadmap, Property 1: 凭证池添加不变性** - /// *对于任意* 凭证池和有效凭证,添加凭证后池的大小应增加 1,且池中应包含该凭证 - /// **Validates: Requirements 1.1** - #[test] - fn prop_pool_add_invariant( - provider in arb_provider_type(), - credential in arb_credential() - ) { - let pool = CredentialPool::new(provider); - let initial_size = pool.len(); - let cred_id = credential.id.clone(); - - // 添加凭证 - let result = pool.add(credential); - prop_assert!(result.is_ok(), "添加凭证应该成功"); - - // 验证不变性:大小增加 1 - prop_assert_eq!( - pool.len(), - initial_size + 1, - "添加凭证后池大小应增加 1" - ); - - // 验证不变性:池中包含该凭证 - prop_assert!( - pool.contains(&cred_id), - "池中应包含刚添加的凭证" - ); - } - - /// **Feature: enhancement-roadmap, Property 1: 凭证池添加不变性(批量)** - /// *对于任意* 凭证池和多个有效凭证,添加 N 个凭证后池的大小应增加 N - /// **Validates: Requirements 1.1** - #[test] - fn prop_pool_add_multiple_invariant( - provider in arb_provider_type(), - credentials in arb_unique_credentials(10) - ) { - let pool = CredentialPool::new(provider); - let initial_size = pool.len(); - let cred_count = credentials.len(); - let cred_ids: Vec<_> = credentials.iter().map(|c| c.id.clone()).collect(); - - // 添加所有凭证 - for cred in credentials { - let result = pool.add(cred); - prop_assert!(result.is_ok(), "添加凭证应该成功"); - } - - // 验证不变性:大小增加 N - prop_assert_eq!( - pool.len(), - initial_size + cred_count, - "添加 {} 个凭证后池大小应增加 {}", - cred_count, - cred_count - ); - - // 验证不变性:池中包含所有凭证 - for id in &cred_ids { - prop_assert!( - pool.contains(id), - "池中应包含凭证 {}", - id - ); - } - } -} - -proptest! { - /// **Feature: enhancement-roadmap, Property 2: 凭证移除不变性** - /// *对于任意* 非空凭证池和池中存在的凭证 ID,移除该凭证后其他凭证应保持不变 - /// **Validates: Requirements 1.3** - #[test] - fn prop_pool_remove_invariant( - provider in arb_provider_type(), - credentials in arb_unique_credentials(10), - remove_index in 0usize..10usize - ) { - // 确保有足够的凭证 - prop_assume!(!credentials.is_empty()); - let remove_index = remove_index % credentials.len(); - - let pool = CredentialPool::new(provider); - - // 添加所有凭证 - let cred_ids: Vec<_> = credentials.iter().map(|c| c.id.clone()).collect(); - for cred in credentials { - pool.add(cred).unwrap(); - } - - let initial_size = pool.len(); - let id_to_remove = &cred_ids[remove_index]; - - // 记录其他凭证的 ID - let other_ids: Vec<_> = cred_ids - .iter() - .filter(|id| *id != id_to_remove) - .cloned() - .collect(); - - // 移除凭证 - let result = pool.remove(id_to_remove); - prop_assert!(result.is_ok(), "移除凭证应该成功"); - - // 验证不变性:大小减少 1 - prop_assert_eq!( - pool.len(), - initial_size - 1, - "移除凭证后池大小应减少 1" - ); - - // 验证不变性:被移除的凭证不再存在 - prop_assert!( - !pool.contains(id_to_remove), - "被移除的凭证不应存在于池中" - ); - - // 验证不变性:其他凭证保持不变 - for id in &other_ids { - prop_assert!( - pool.contains(id), - "其他凭证 {} 应保持不变", - id - ); - } - } - - /// **Feature: enhancement-roadmap, Property 2: 凭证移除不变性(连续移除)** - /// *对于任意* 凭证池,连续移除所有凭证后池应为空 - /// **Validates: Requirements 1.3** - #[test] - fn prop_pool_remove_all_invariant( - provider in arb_provider_type(), - credentials in arb_unique_credentials(10) - ) { - prop_assume!(!credentials.is_empty()); - - let pool = CredentialPool::new(provider); - - // 添加所有凭证 - let cred_ids: Vec<_> = credentials.iter().map(|c| c.id.clone()).collect(); - for cred in credentials { - pool.add(cred).unwrap(); - } - - // 逐个移除所有凭证 - for id in &cred_ids { - let result = pool.remove(id); - prop_assert!(result.is_ok(), "移除凭证 {} 应该成功", id); - } - - // 验证不变性:池为空 - prop_assert!(pool.is_empty(), "移除所有凭证后池应为空"); - prop_assert_eq!(pool.len(), 0, "移除所有凭证后池大小应为 0"); - } -} - -/// 生成具有唯一 ID 且属于同一 Provider 的凭证列表 -fn arb_unique_credentials_same_provider( - provider: ProviderType, - min_count: usize, - max_count: usize, -) -> impl Strategy> { - prop::collection::vec(arb_credential_data(), min_count..=max_count).prop_map(move |data_list| { - data_list - .into_iter() - .enumerate() - .map(|(i, data)| Credential::new(format!("cred-{i}"), provider, data)) - .collect() - }) -} - -proptest! { - /// **Feature: enhancement-roadmap, Property 3: 轮询均匀性** - /// *对于任意* 包含 N 个活跃凭证的池,连续 N 次选择应返回 N 个不同的凭证 - /// **Validates: Requirements 1.2 (验收标准 1)** - #[test] - fn prop_round_robin_uniformity( - provider in arb_provider_type(), - cred_count in 2usize..=10usize - ) { - // 创建负载均衡器 - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 添加 N 个凭证 - for i in 0..cred_count { - let cred = Credential::new( - format!("cred-{i}"), - provider, - CredentialData::ApiKey { - key: format!("key-{i}"), - base_url: None, - }, - ); - pool.add(cred).unwrap(); - } - - lb.register_pool(pool); - - // 连续选择 N 次 - let mut selected_ids: Vec = Vec::with_capacity(cred_count); - for _ in 0..cred_count { - let cred = lb.select(provider).unwrap(); - selected_ids.push(cred.id.clone()); - } - - // 验证:N 次选择应返回 N 个不同的凭证 - let unique_ids: HashSet<_> = selected_ids.iter().collect(); - prop_assert_eq!( - unique_ids.len(), - cred_count, - "连续 {} 次选择应返回 {} 个不同的凭证,但只得到 {} 个不同的凭证: {:?}", - cred_count, - cred_count, - unique_ids.len(), - selected_ids - ); - } - - /// **Feature: enhancement-roadmap, Property 3: 轮询均匀性(多轮)** - /// *对于任意* 包含 N 个活跃凭证的池,连续 2N 次选择应每个凭证被选中 2 次 - /// **Validates: Requirements 1.2 (验收标准 1)** - #[test] - fn prop_round_robin_uniformity_multiple_rounds( - provider in arb_provider_type(), - cred_count in 2usize..=5usize, - rounds in 2usize..=4usize - ) { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 添加 N 个凭证 - for i in 0..cred_count { - let cred = Credential::new( - format!("cred-{i}"), - provider, - CredentialData::ApiKey { - key: format!("key-{i}"), - base_url: None, - }, - ); - pool.add(cred).unwrap(); - } - - lb.register_pool(pool); - - // 连续选择 N * rounds 次 - let total_selections = cred_count * rounds; - let mut selection_counts: std::collections::HashMap = - std::collections::HashMap::new(); - - for _ in 0..total_selections { - let cred = lb.select(provider).unwrap(); - *selection_counts.entry(cred.id.clone()).or_insert(0) += 1; - } - - // 验证:每个凭证应被选中 rounds 次 - for (id, count) in &selection_counts { - prop_assert_eq!( - *count, - rounds, - "凭证 {} 应被选中 {} 次,但实际被选中 {} 次", - id, - rounds, - count - ); - } - - // 验证:应该有 N 个不同的凭证被选中 - prop_assert_eq!( - selection_counts.len(), - cred_count, - "应有 {} 个不同的凭证被选中,但实际有 {} 个", - cred_count, - selection_counts.len() - ); - } -} - -proptest! { - /// **Feature: enhancement-roadmap, Property 5: 健康状态转换** - /// *对于任意* 凭证,连续 3 次失败后状态应变为不健康 - /// **Validates: Requirements 1.3 (验收标准 2)** - #[test] - fn prop_health_state_transition( - provider in arb_provider_type(), - failure_threshold in 1u32..=5u32 - ) { - use lime_core::credential::{CredentialStatus, HealthCheckConfig, HealthChecker}; - use std::time::Duration; - - // 创建带自定义阈值的健康检查器 - let config = HealthCheckConfig { - check_interval: Duration::from_secs(60), - failure_threshold, - recovery_threshold: 1, - }; - let checker = HealthChecker::new(config); - let pool = CredentialPool::new(provider); - - // 添加凭证 - let cred = Credential::new( - "test-cred".to_string(), - provider, - CredentialData::ApiKey { - key: "test-key".to_string(), - base_url: None, - }, - ); - pool.add(cred).unwrap(); - - // 记录 (failure_threshold - 1) 次失败,不应标记为不健康 - for i in 0..(failure_threshold - 1) { - let marked = checker.record_failure(&pool, "test-cred").unwrap(); - prop_assert!( - !marked, - "第 {} 次失败不应标记为不健康(阈值: {})", - i + 1, - failure_threshold - ); - - let cred = pool.get("test-cred").unwrap(); - prop_assert!( - matches!(cred.status, CredentialStatus::Active), - "第 {} 次失败后状态应仍为 Active", - i + 1 - ); - } - - // 第 failure_threshold 次失败应标记为不健康 - let marked = checker.record_failure(&pool, "test-cred").unwrap(); - prop_assert!( - marked, - "第 {} 次失败应标记为不健康", - failure_threshold - ); - - let cred = pool.get("test-cred").unwrap(); - prop_assert!( - matches!(cred.status, CredentialStatus::Unhealthy { .. }), - "达到阈值后状态应为 Unhealthy,但实际为 {:?}", - cred.status - ); - - // 验证连续失败次数 - prop_assert_eq!( - cred.stats.consecutive_failures, - failure_threshold, - "连续失败次数应为 {}", - failure_threshold - ); - } - - /// **Feature: enhancement-roadmap, Property 5: 健康状态转换(恢复)** - /// *对于任意* 不健康的凭证,成功后应恢复为健康状态 - /// **Validates: Requirements 1.3 (验收标准 2)** - #[test] - fn prop_health_state_recovery( - provider in arb_provider_type(), - latency_ms in 1u64..1000u64 - ) { - use lime_core::credential::{CredentialStatus, HealthChecker}; - - let checker = HealthChecker::with_defaults(); - let pool = CredentialPool::new(provider); - - // 添加凭证 - let cred = Credential::new( - "test-cred".to_string(), - provider, - CredentialData::ApiKey { - key: "test-key".to_string(), - base_url: None, - }, - ); - pool.add(cred).unwrap(); - - // 标记为不健康 - pool.mark_unhealthy("test-cred", "test reason".to_string()).unwrap(); - - // 验证状态为不健康 - let cred = pool.get("test-cred").unwrap(); - prop_assert!( - matches!(cred.status, CredentialStatus::Unhealthy { .. }), - "凭证应为不健康状态" - ); - - // 记录成功应恢复 - let recovered = checker.record_success(&pool, "test-cred", latency_ms).unwrap(); - prop_assert!( - recovered, - "成功后应恢复为健康状态" - ); - - // 验证状态已恢复 - let cred = pool.get("test-cred").unwrap(); - prop_assert!( - matches!(cred.status, CredentialStatus::Active), - "恢复后状态应为 Active,但实际为 {:?}", - cred.status - ); - - // 验证连续失败次数已重置 - prop_assert_eq!( - cred.stats.consecutive_failures, - 0, - "恢复后连续失败次数应为 0" - ); - } - - /// **Feature: enhancement-roadmap, Property 5: 健康状态转换(成功重置失败计数)** - /// *对于任意* 凭证,成功请求应重置连续失败计数 - /// **Validates: Requirements 1.3 (验收标准 2)** - #[test] - fn prop_success_resets_failure_count( - provider in arb_provider_type(), - failures_before in 1u32..3u32, - latency_ms in 1u64..1000u64 - ) { - use lime_core::credential::HealthChecker; - - let checker = HealthChecker::with_defaults(); - let pool = CredentialPool::new(provider); - - // 添加凭证 - let cred = Credential::new( - "test-cred".to_string(), - provider, - CredentialData::ApiKey { - key: "test-key".to_string(), - base_url: None, - }, - ); - pool.add(cred).unwrap(); - - // 记录一些失败(但不超过阈值) - for _ in 0..failures_before { - checker.record_failure(&pool, "test-cred").unwrap(); - } - - // 验证有连续失败 - let cred = pool.get("test-cred").unwrap(); - prop_assert_eq!( - cred.stats.consecutive_failures, - failures_before, - "应有 {} 次连续失败", - failures_before - ); - - // 记录成功 - checker.record_success(&pool, "test-cred", latency_ms).unwrap(); - - // 验证连续失败次数已重置 - let cred = pool.get("test-cred").unwrap(); - prop_assert_eq!( - cred.stats.consecutive_failures, - 0, - "成功后连续失败次数应重置为 0" - ); - } -} - -proptest! { - /// **Feature: enhancement-roadmap, Property 4: 冷却状态转换** - /// *对于任意* 凭证,当标记为冷却后,在冷却期内不应被选中;冷却期结束后应恢复可选 - /// **Validates: Requirements 1.2 (验收标准 2, 4)** - #[test] - fn prop_cooldown_state_transition( - provider in arb_provider_type(), - cooldown_index in 0usize..5usize - ) { - use chrono::{Duration, Utc}; - use lime_core::credential::CredentialStatus; - - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 添加 5 个凭证 - let cred_count = 5usize; - for i in 0..cred_count { - let cred = Credential::new( - format!("cred-{i}"), - provider, - CredentialData::ApiKey { - key: format!("key-{i}"), - base_url: None, - }, - ); - pool.add(cred).unwrap(); - } - - lb.register_pool(pool.clone()); - - let cooldown_id = format!("cred-{cooldown_index}"); - - // 标记一个凭证为冷却状态(1小时后恢复) - lb.mark_cooldown(provider, &cooldown_id, Duration::hours(1)).unwrap(); - - // 验证:冷却中的凭证不应被选中 - // 连续选择 (N-1) * 2 次,应该不会选中冷却中的凭证 - let selections = (cred_count - 1) * 2; - for _ in 0..selections { - let selected = lb.select(provider).unwrap(); - prop_assert_ne!( - selected.id, - cooldown_id.clone(), - "冷却中的凭证 {} 不应被选中", - &cooldown_id - ); - } - - // 模拟冷却期结束:直接设置状态为过去的时间 - { - let mut entry = pool.credentials.get_mut(&cooldown_id).unwrap(); - entry.status = CredentialStatus::Cooldown { - until: Utc::now() - Duration::seconds(1), - }; - } - - // 验证:冷却期结束后应恢复可选 - // 连续选择 N 次,应该能选中之前冷却的凭证 - let mut found_recovered = false; - for _ in 0..cred_count { - let selected = lb.select(provider).unwrap(); - if selected.id == cooldown_id { - found_recovered = true; - break; - } - } - - prop_assert!( - found_recovered, - "冷却期结束后,凭证 {} 应该能被选中", - cooldown_id - ); - - // 验证:恢复后的凭证状态应为 Active - let cred = pool.get(&cooldown_id).unwrap(); - prop_assert!( - matches!(cred.status, CredentialStatus::Active), - "冷却期结束后,凭证状态应为 Active,但实际为 {:?}", - cred.status - ); - } - - /// **Feature: enhancement-roadmap, Property 4: 冷却状态转换(所有凭证冷却)** - /// *对于任意* 凭证池,当所有凭证都处于冷却状态时,选择应返回错误 - /// **Validates: Requirements 1.2 (验收标准 3)** - #[test] - fn prop_all_cooldown_returns_error( - provider in arb_provider_type(), - cred_count in 1usize..=5usize - ) { - use chrono::Duration; - use lime_core::credential::PoolError; - - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 添加凭证 - for i in 0..cred_count { - let cred = Credential::new( - format!("cred-{i}"), - provider, - CredentialData::ApiKey { - key: format!("key-{i}"), - base_url: None, - }, - ); - pool.add(cred).unwrap(); - } - - lb.register_pool(pool); - - // 将所有凭证标记为冷却 - for i in 0..cred_count { - lb.mark_cooldown(provider, &format!("cred-{i}"), Duration::hours(1)) - .unwrap(); - } - - // 验证:选择应返回 NoAvailableCredential 错误 - let result = lb.select(provider); - prop_assert!( - matches!(result, Err(PoolError::NoAvailableCredential)), - "所有凭证冷却时,选择应返回 NoAvailableCredential 错误,但实际返回 {:?}", - result - ); - - // 验证:应该能获取最早恢复时间 - let recovery = lb.earliest_recovery(provider); - prop_assert!( - recovery.is_some(), - "所有凭证冷却时,应该能获取最早恢复时间" - ); - } -} - -// ============ 凭证同步服务属性测试 ============ - -use crate::models::provider_pool_model::{ - CredentialData as PoolCredentialData, PoolProviderType, ProviderCredential, -}; -use lime_core::config::{Config, ConfigManager}; -use lime_credential::CredentialSyncService; -use std::sync::RwLock; -use tempfile::TempDir; - -/// 创建临时测试环境 -fn create_test_env() -> (TempDir, Arc>) { - let temp_dir = TempDir::new().expect("创建临时目录失败"); - let config_path = temp_dir.path().join("config.yaml"); - - // 创建配置管理器 - let mut config = Config::default(); - config.auth_dir = temp_dir.path().join("auth").to_string_lossy().to_string(); - - let mut manager = ConfigManager::new(config_path); - manager.set_config(config); - manager.save().expect("保存配置失败"); - - (temp_dir, Arc::new(RwLock::new(manager))) -} - -/// 生成随机的 PoolProviderType(仅支持同步的类型) -fn arb_sync_provider_type() -> impl Strategy { - prop_oneof![ - Just(PoolProviderType::Kiro), - Just(PoolProviderType::Gemini), - Just(PoolProviderType::OpenAI), - Just(PoolProviderType::Claude), - ] -} - -/// 生成随机的 API Key 凭证数据 -fn arb_api_key_credential() -> impl Strategy { - prop_oneof![ - ( - "[a-zA-Z0-9]{20,50}", - prop::option::of("https://[a-z]+\\.[a-z]+/v1") - ) - .prop_map(|(api_key, base_url)| { - ( - PoolProviderType::OpenAI, - PoolCredentialData::OpenAIKey { api_key, base_url }, - ) - }), - ( - "[a-zA-Z0-9]{20,50}", - prop::option::of("https://[a-z]+\\.[a-z]+") - ) - .prop_map(|(api_key, base_url)| { - ( - PoolProviderType::Claude, - PoolCredentialData::ClaudeKey { api_key, base_url }, - ) - }), - ] -} - -proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: config-credential-export, Property 1: Credential Sync Round Trip** - /// *For any* credential added to the credential pool, saving to YAML and then loading - /// from YAML should produce an equivalent credential configuration. - /// **Validates: Requirements 1.1, 1.2, 1.5** - #[test] - fn prop_credential_sync_round_trip( - (provider_type, cred_data) in arb_api_key_credential(), - is_disabled in proptest::bool::ANY - ) { - let (_temp_dir, config_manager) = create_test_env(); - let sync_service = CredentialSyncService::new(config_manager.clone()); - - // 创建凭证 - let mut credential = ProviderCredential::new(provider_type, cred_data.clone()); - credential.is_disabled = is_disabled; - let original_uuid = credential.uuid.clone(); - - // 添加凭证 - let add_result = sync_service.add_credential(&credential); - prop_assert!(add_result.is_ok(), "添加凭证应该成功: {:?}", add_result); - - // 从配置加载凭证 - let loaded = sync_service.load_from_config(); - prop_assert!(loaded.is_ok(), "加载凭证应该成功: {:?}", loaded); - - let loaded_creds = loaded.unwrap(); - - // 查找对应的凭证 - let found = loaded_creds.iter().find(|c| c.uuid == original_uuid); - prop_assert!(found.is_some(), "应该能找到添加的凭证"); - - let loaded_cred = found.unwrap(); - - // 验证凭证属性 - prop_assert_eq!( - &loaded_cred.uuid, - &original_uuid, - "UUID 应该一致" - ); - prop_assert_eq!( - loaded_cred.provider_type, - provider_type, - "Provider 类型应该一致" - ); - prop_assert_eq!( - loaded_cred.is_disabled, - is_disabled, - "禁用状态应该一致" - ); - - // 验证凭证数据 - match (&loaded_cred.credential, &cred_data) { - ( - PoolCredentialData::OpenAIKey { api_key: loaded_key, base_url: loaded_url }, - PoolCredentialData::OpenAIKey { api_key: orig_key, base_url: orig_url }, - ) => { - prop_assert_eq!(loaded_key, orig_key, "API Key 应该一致"); - prop_assert_eq!(loaded_url, orig_url, "Base URL 应该一致"); - } - ( - PoolCredentialData::ClaudeKey { api_key: loaded_key, base_url: loaded_url }, - PoolCredentialData::ClaudeKey { api_key: orig_key, base_url: orig_url }, - ) => { - prop_assert_eq!(loaded_key, orig_key, "API Key 应该一致"); - prop_assert_eq!(loaded_url, orig_url, "Base URL 应该一致"); - } - _ => { - prop_assert!(false, "凭证类型不匹配"); - } - } - } - - /// **Feature: config-credential-export, Property 1: Credential Sync Round Trip (Multiple)** - /// *For any* set of credentials, adding them all and then loading should preserve all. - /// **Validates: Requirements 1.1, 1.2, 1.5** - #[test] - fn prop_credential_sync_round_trip_multiple( - cred_count in 1usize..=5usize - ) { - let (_temp_dir, config_manager) = create_test_env(); - let sync_service = CredentialSyncService::new(config_manager.clone()); - - // 创建多个凭证 - let mut original_uuids = Vec::new(); - for i in 0..cred_count { - let cred_data = if i % 2 == 0 { - PoolCredentialData::OpenAIKey { - api_key: format!("sk-test-key-{i}"), - base_url: Some("https://api.openai.com/v1".to_string()), - } - } else { - PoolCredentialData::ClaudeKey { - api_key: format!("sk-ant-test-key-{i}"), - base_url: None, - } - }; - - let provider_type = if i % 2 == 0 { - PoolProviderType::OpenAI - } else { - PoolProviderType::Claude - }; - - let credential = ProviderCredential::new(provider_type, cred_data); - original_uuids.push(credential.uuid.clone()); - - let add_result = sync_service.add_credential(&credential); - prop_assert!(add_result.is_ok(), "添加凭证 {} 应该成功", i); - } - - // 从配置加载凭证 - let loaded = sync_service.load_from_config().unwrap(); - - // 验证所有凭证都被加载 - prop_assert_eq!( - loaded.len(), - cred_count, - "加载的凭证数量应该与添加的一致" - ); - - // 验证每个 UUID 都存在 - for uuid in &original_uuids { - let found = loaded.iter().any(|c| &c.uuid == uuid); - prop_assert!(found, "应该能找到 UUID 为 {} 的凭证", uuid); - } - } -} - -proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: config-credential-export, Property 9: OAuth Token File Handling** - /// *For any* OAuth credential, the token file should be stored in auth-dir on add, - /// included in export bundles, and restored to auth-dir on import. - /// **Validates: Requirements 2.1, 2.4, 3.3, 4.4** - #[test] - fn prop_oauth_token_file_handling( - token_content in "[a-zA-Z0-9]{50,200}", - provider_idx in 0usize..3usize - ) { - let (temp_dir, config_manager) = create_test_env(); - let sync_service = CredentialSyncService::new(config_manager.clone()); - - // 创建源 token 文件 - let source_token_dir = temp_dir.path().join("source_tokens"); - std::fs::create_dir_all(&source_token_dir).expect("创建源目录失败"); - - let source_token_path = source_token_dir.join("token.json"); - let token_json = format!(r#"{{"access_token": "{token_content}", "refresh_token": "refresh-{token_content}", "expires_at": "2025-12-31T23:59:59Z"}}"#); - std::fs::write(&source_token_path, &token_json).expect("写入源 token 文件失败"); - - // 根据索引选择 provider 类型 - let (provider_type, cred_data) = match provider_idx { - 0 => ( - PoolProviderType::Kiro, - PoolCredentialData::KiroOAuth { - creds_file_path: source_token_path.to_string_lossy().to_string(), - }, - ), - 1 => ( - PoolProviderType::Gemini, - PoolCredentialData::GeminiOAuth { - creds_file_path: source_token_path.to_string_lossy().to_string(), - project_id: None, - }, - ), - _ => ( - PoolProviderType::Kiro, - PoolCredentialData::KiroOAuth { - creds_file_path: source_token_path.to_string_lossy().to_string(), - }, - ), - }; - - // 创建凭证 - let credential = ProviderCredential::new(provider_type, cred_data); - let original_uuid = credential.uuid.clone(); - - // 添加凭证(应该复制 token 文件到 auth_dir) - let add_result = sync_service.add_credential(&credential); - prop_assert!(add_result.is_ok(), "添加 OAuth 凭证应该成功: {:?}", add_result); - - // 验证 token 文件已复制到 auth_dir - let auth_dir = sync_service.get_auth_dir().expect("获取 auth_dir 失败"); - let provider_name = match provider_type { - PoolProviderType::Kiro => "kiro", - PoolProviderType::Gemini => "gemini", - _ => "unknown", - }; - let expected_token_path = auth_dir.join(provider_name).join(format!("{original_uuid}.json")); - - prop_assert!( - expected_token_path.exists(), - "Token 文件应该存在于 auth_dir: {:?}", - expected_token_path - ); - - // 验证 token 文件内容一致 - let copied_content = std::fs::read_to_string(&expected_token_path) - .expect("读取复制的 token 文件失败"); - prop_assert_eq!( - copied_content, - token_json, - "Token 文件内容应该一致" - ); - - // 从配置加载凭证 - let loaded = sync_service.load_from_config().expect("加载凭证失败"); - let loaded_cred = loaded.iter().find(|c| c.uuid == original_uuid); - prop_assert!(loaded_cred.is_some(), "应该能找到加载的凭证"); - - // 验证加载的凭证指向正确的 token 文件路径 - let loaded_cred = loaded_cred.unwrap(); - let loaded_path = match &loaded_cred.credential { - PoolCredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(), - PoolCredentialData::GeminiOAuth { creds_file_path, .. } => creds_file_path.clone(), - _ => String::new(), - }; - - prop_assert_eq!( - loaded_path, - expected_token_path.to_string_lossy().to_string(), - "加载的凭证应该指向 auth_dir 中的 token 文件" - ); - - // 删除凭证(应该删除 token 文件) - let remove_result = sync_service.remove_credential(provider_type, &original_uuid); - prop_assert!(remove_result.is_ok(), "删除凭证应该成功: {:?}", remove_result); - - // 验证 token 文件已被删除 - prop_assert!( - !expected_token_path.exists(), - "删除凭证后 token 文件应该被删除" - ); - } - - /// **Feature: config-credential-export, Property 9: OAuth Token File Update** - /// *For any* OAuth credential update, the token file should be updated in auth-dir. - /// **Validates: Requirements 2.1, 2.4** - #[test] - fn prop_oauth_token_file_update( - initial_content in "[a-zA-Z0-9]{50,100}", - updated_content in "[a-zA-Z0-9]{50,100}" - ) { - let (temp_dir, config_manager) = create_test_env(); - let sync_service = CredentialSyncService::new(config_manager.clone()); - - // 创建初始 token 文件 - let source_token_dir = temp_dir.path().join("source_tokens"); - std::fs::create_dir_all(&source_token_dir).expect("创建源目录失败"); - - let source_token_path = source_token_dir.join("token.json"); - let initial_json = format!(r#"{{"access_token": "{initial_content}"}}"#); - std::fs::write(&source_token_path, &initial_json).expect("写入初始 token 文件失败"); - - // 创建凭证 - let credential = ProviderCredential::new( - PoolProviderType::Kiro, - PoolCredentialData::KiroOAuth { - creds_file_path: source_token_path.to_string_lossy().to_string(), - }, - ); - let original_uuid = credential.uuid.clone(); - - // 添加凭证 - sync_service.add_credential(&credential).expect("添加凭证失败"); - - // 更新源 token 文件内容 - let updated_json = format!(r#"{{"access_token": "{updated_content}"}}"#); - std::fs::write(&source_token_path, &updated_json).expect("更新源 token 文件失败"); - - // 更新凭证 - let mut updated_credential = credential.clone(); - updated_credential.credential = PoolCredentialData::KiroOAuth { - creds_file_path: source_token_path.to_string_lossy().to_string(), - }; - - let update_result = sync_service.update_credential(&updated_credential); - prop_assert!(update_result.is_ok(), "更新凭证应该成功: {:?}", update_result); - - // 验证 auth_dir 中的 token 文件已更新 - let auth_dir = sync_service.get_auth_dir().expect("获取 auth_dir 失败"); - let token_path = auth_dir.join("kiro").join(format!("{original_uuid}.json")); - - let stored_content = std::fs::read_to_string(&token_path) - .expect("读取存储的 token 文件失败"); - prop_assert_eq!( - stored_content, - updated_json, - "存储的 token 文件内容应该已更新" - ); - } -} - -// ============ Per-Key Proxy Selection Property Tests ============ - -proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: cliproxyapi-parity, Property 14: Per-Key Proxy Selection** - /// *For any* credential with proxy_url set, requests using that credential - /// SHALL use the per-key proxy; otherwise, the global proxy SHALL be used. - /// **Validates: Requirements 7.1, 7.2** - #[test] - fn prop_credential_per_key_proxy_selection( - provider in arb_provider_type(), - per_key_proxy in "[a-z0-9]{1,10}", - global_proxy in "[a-z0-9]{1,10}" - ) { - let per_key_url = format!("http://{per_key_proxy}:8080"); - let global_url = format!("http://{global_proxy}:8080"); - - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin) - .with_global_proxy(Some(global_url.clone())); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建带 Per-Key 代理的凭证 - let cred_with_proxy = Credential::new( - "cred-with-proxy".to_string(), - provider, - CredentialData::ApiKey { - key: "key-1".to_string(), - base_url: None, - }, - ).with_proxy(Some(per_key_url.clone())); - - // 创建不带 Per-Key 代理的凭证 - let cred_without_proxy = Credential::new( - "cred-without-proxy".to_string(), - provider, - CredentialData::ApiKey { - key: "key-2".to_string(), - base_url: None, - }, - ); - - pool.add(cred_with_proxy).unwrap(); - pool.add(cred_without_proxy).unwrap(); - lb.register_pool(pool.clone()); - - // 验证带 Per-Key 代理的凭证 - let cred = pool.get("cred-with-proxy").unwrap(); - prop_assert_eq!( - cred.proxy_url(), - Some(per_key_url.as_str()), - "带 Per-Key 代理的凭证应该返回 Per-Key 代理 URL" - ); - - // 验证代理选择逻辑 - let selected_proxy = lb.proxy_factory().select_proxy(cred.proxy_url()); - prop_assert_eq!( - selected_proxy, - Some(per_key_url.as_str()), - "Per-Key 代理应该优先于全局代理" - ); - - // 验证不带 Per-Key 代理的凭证 - let cred = pool.get("cred-without-proxy").unwrap(); - prop_assert_eq!( - cred.proxy_url(), - None, - "不带 Per-Key 代理的凭证应该返回 None" - ); - - // 验证回退到全局代理 - let selected_proxy = lb.proxy_factory().select_proxy(cred.proxy_url()); - prop_assert_eq!( - selected_proxy, - Some(global_url.as_str()), - "无 Per-Key 代理时应该使用全局代理" - ); - } - - /// **Feature: cliproxyapi-parity, Property 14: Per-Key Proxy Selection** - /// *For any* credential without proxy_url and no global proxy, - /// no proxy SHALL be used. - /// **Validates: Requirements 7.1, 7.2** - #[test] - fn prop_credential_no_proxy_when_none_configured( - provider in arb_provider_type() - ) { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建不带代理的凭证 - let cred = Credential::new( - "cred-no-proxy".to_string(), - provider, - CredentialData::ApiKey { - key: "key-1".to_string(), - base_url: None, - }, - ); - - pool.add(cred).unwrap(); - lb.register_pool(pool.clone()); - - // 验证凭证没有代理 - let cred = pool.get("cred-no-proxy").unwrap(); - prop_assert_eq!( - cred.proxy_url(), - None, - "凭证应该没有 Per-Key 代理" - ); - - // 验证代理选择返回 None - let selected_proxy = lb.proxy_factory().select_proxy(cred.proxy_url()); - prop_assert_eq!( - selected_proxy, - None, - "无全局代理且无 Per-Key 代理时应该不使用代理" - ); - } - - /// **Feature: cliproxyapi-parity, Property 14: Per-Key Proxy Selection** - /// *For any* credential with proxy_url, select_with_client SHALL create - /// a client configured with that proxy. - /// **Validates: Requirements 7.1, 7.2** - #[test] - fn prop_select_with_client_uses_per_key_proxy( - provider in arb_provider_type(), - // Hostname must start with a letter to be valid - proxy_host in "[a-z][a-z0-9]{0,9}" - ) { - let proxy_url = format!("http://{proxy_host}:8080"); - - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建带代理的凭证 - let cred = Credential::new( - "cred-1".to_string(), - provider, - CredentialData::ApiKey { - key: "key-1".to_string(), - base_url: None, - }, - ).with_proxy(Some(proxy_url.clone())); - - pool.add(cred).unwrap(); - lb.register_pool(pool); - - // 使用 select_with_client 选择凭证 - let selection = lb.select_with_client(provider); - prop_assert!(selection.is_ok(), "select_with_client 应该成功"); - - let selection = selection.unwrap(); - prop_assert_eq!( - selection.credential.proxy_url(), - Some(proxy_url.as_str()), - "选中的凭证应该有正确的代理 URL" - ); - } -} - -// ============ Proxy Failover Property Tests ============ - -proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: cliproxyapi-parity, Property 15: Proxy Failover** - /// *For any* credential where proxy connection fails, the system - /// SHALL attempt the next available credential. - /// **Validates: Requirements 7.4** - #[test] - fn prop_proxy_failover_attempts_next_credential( - provider in arb_provider_type(), - cred_count in 2usize..=5usize - ) { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建多个凭证,第一个有无效代理,其他有有效代理 - for i in 0..cred_count { - let proxy_url = if i == 0 { - // 第一个凭证使用无效代理协议 - Some("ftp://invalid-proxy:21".to_string()) - } else { - // 其他凭证使用有效代理 - Some(format!("http://valid-proxy-{i}:8080")) - }; - - let cred = Credential::new( - format!("cred-{i}"), - provider, - CredentialData::ApiKey { - key: format!("key-{i}"), - base_url: None, - }, - ).with_proxy(proxy_url); - - pool.add(cred).unwrap(); - } - - lb.register_pool(pool); - - // 使用 select_with_failover 应该跳过无效代理的凭证 - let result = lb.select_with_failover(provider, None); - prop_assert!(result.is_ok(), "故障转移应该成功找到有效凭证"); - - let selection = result.unwrap(); - // 选中的凭证不应该是第一个(无效代理的那个) - prop_assert_ne!( - selection.credential.id, - "cred-0", - "应该跳过无效代理的凭证" - ); - } - - /// **Feature: cliproxyapi-parity, Property 15: Proxy Failover** - /// *For any* set of credentials with all valid proxies, select_with_failover - /// SHALL succeed on the first attempt. - /// **Validates: Requirements 7.4** - #[test] - fn prop_proxy_failover_succeeds_with_valid_proxies( - provider in arb_provider_type(), - cred_count in 1usize..=5usize - ) { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建多个凭证,都有有效代理 - for i in 0..cred_count { - let cred = Credential::new( - format!("cred-{i}"), - provider, - CredentialData::ApiKey { - key: format!("key-{i}"), - base_url: None, - }, - ).with_proxy(Some(format!("http://proxy-{i}:8080"))); - - pool.add(cred).unwrap(); - } - - lb.register_pool(pool); - - // 使用 select_with_failover 应该成功 - let result = lb.select_with_failover(provider, None); - prop_assert!(result.is_ok(), "所有代理有效时应该成功"); - } - - /// **Feature: cliproxyapi-parity, Property 15: Proxy Failover** - /// *For any* set of credentials with all invalid proxies, select_with_failover - /// SHALL fail after trying all credentials. - /// **Validates: Requirements 7.4** - #[test] - fn prop_proxy_failover_fails_when_all_invalid( - provider in arb_provider_type(), - cred_count in 1usize..=3usize - ) { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建多个凭证,都有无效代理 - for i in 0..cred_count { - let cred = Credential::new( - format!("cred-{i}"), - provider, - CredentialData::ApiKey { - key: format!("key-{i}"), - base_url: None, - }, - ).with_proxy(Some(format!("ftp://invalid-proxy-{i}:21"))); - - pool.add(cred).unwrap(); - } - - lb.register_pool(pool); - - // 使用 select_with_failover 应该失败 - let result = lb.select_with_failover(provider, None); - prop_assert!(result.is_err(), "所有代理无效时应该失败"); - } - - /// **Feature: cliproxyapi-parity, Property 15: Proxy Failover** - /// *For any* credential without proxy, select_with_failover SHALL succeed - /// using no proxy. - /// **Validates: Requirements 7.4** - #[test] - fn prop_proxy_failover_succeeds_without_proxy( - provider in arb_provider_type() - ) { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建不带代理的凭证 - let cred = Credential::new( - "cred-no-proxy".to_string(), - provider, - CredentialData::ApiKey { - key: "key-1".to_string(), - base_url: None, - }, - ); - - pool.add(cred).unwrap(); - lb.register_pool(pool); - - // 使用 select_with_failover 应该成功 - let result = lb.select_with_failover(provider, None); - prop_assert!(result.is_ok(), "无代理凭证应该成功"); - - let selection = result.unwrap(); - prop_assert_eq!( - selection.credential.proxy_url(), - None, - "选中的凭证应该没有代理" - ); - } - - /// **Feature: cliproxyapi-parity, Property 15: Proxy Failover** - /// *For any* failover_on_proxy_error call, the system SHALL record - /// the failure and attempt to select a new credential. - /// **Validates: Requirements 7.4** - #[test] - fn prop_failover_on_proxy_error_records_failure( - provider in arb_provider_type() - ) { - let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); - let pool = Arc::new(CredentialPool::new(provider)); - - // 创建两个凭证 - let cred1 = Credential::new( - "cred-1".to_string(), - provider, - CredentialData::ApiKey { - key: "key-1".to_string(), - base_url: None, - }, - ).with_proxy(Some("http://proxy1:8080".to_string())); - - let cred2 = Credential::new( - "cred-2".to_string(), - provider, - CredentialData::ApiKey { - key: "key-2".to_string(), - base_url: None, - }, - ).with_proxy(Some("http://proxy2:8080".to_string())); - - pool.add(cred1).unwrap(); - pool.add(cred2).unwrap(); - lb.register_pool(pool.clone()); - - // 调用 failover_on_proxy_error - let result = lb.failover_on_proxy_error(provider, "cred-1"); - prop_assert!(result.is_ok(), "故障转移应该成功"); - - // 验证失败被记录 - let cred1 = pool.get("cred-1").unwrap(); - prop_assert_eq!( - cred1.stats.consecutive_failures, - 1, - "失败应该被记录" - ); - } -} - -// ============ 配额管理器属性测试 ============ - -use lime_core::config::QuotaExceededConfig; -use lime_credential::QuotaManager; - -/// 生成随机的配额超限配置 -fn arb_quota_config() -> impl Strategy { - (proptest::bool::ANY, proptest::bool::ANY, 1u64..=3600u64).prop_map( - |(switch_project, switch_preview_model, cooldown_seconds)| QuotaExceededConfig { - switch_project, - switch_preview_model, - cooldown_seconds, - }, - ) -} - -/// 生成随机的凭证 ID -fn arb_credential_id() -> impl Strategy { - "[a-zA-Z0-9_-]{1,32}".prop_map(|s| s) -} - -/// 生成随机的错误消息 -fn arb_error_message() -> impl Strategy { - prop_oneof![ - // 配额超限相关消息 - Just("Rate limit exceeded".to_string()), - Just("Quota exceeded for this API".to_string()), - Just("Too many requests".to_string()), - Just("Request was throttled".to_string()), - Just("limit exceeded".to_string()), - // 非配额超限消息 - Just("Bad Request".to_string()), - Just("Internal Server Error".to_string()), - Just("Not Found".to_string()), - Just("Unauthorized".to_string()), - Just("Service Unavailable".to_string()), - ] -} - -/// 生成随机的 HTTP 状态码 -fn arb_status_code() -> impl Strategy> { - prop_oneof![ - Just(None), - Just(Some(200u16)), - Just(Some(400u16)), - Just(Some(401u16)), - Just(Some(403u16)), - Just(Some(404u16)), - Just(Some(429u16)), // 配额超限 - Just(Some(500u16)), - Just(Some(502u16)), - Just(Some(503u16)), - Just(Some(504u16)), - ] -} - -/// 生成随机的模型名称 -fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gemini-2.5-pro".to_string()), - Just("gemini-2.5-flash".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - // 已经是预览版本 - Just("gemini-2.5-pro-preview".to_string()), - Just("claude-3-opus-preview-20240101".to_string()), - ] -} - -proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: cliproxyapi-parity, Property 16: Quota Exceeded Detection** - /// *For any* API response indicating quota exceeded (HTTP 429 or specific error codes), - /// the credential SHALL be marked as temporarily unavailable. - /// **Validates: Requirements 8.1** - #[test] - fn prop_quota_exceeded_detection( - config in arb_quota_config(), - credential_id in arb_credential_id(), - status_code in arb_status_code(), - error_message in arb_error_message() - ) { - let manager = QuotaManager::new(config); - - // 检测是否为配额超限错误 - let is_quota_error = QuotaManager::is_quota_exceeded_error(status_code, &error_message); - - // 验证 429 状态码总是被检测为配额超限 - if status_code == Some(429) { - prop_assert!( - is_quota_error, - "HTTP 429 应该被检测为配额超限错误" - ); - } - - // 验证包含配额关键词的消息被检测为配额超限 - let error_lower = error_message.to_lowercase(); - let has_quota_keyword = ["quota", "rate limit", "rate_limit", "too many requests", "exceeded", "limit exceeded", "throttl"] - .iter() - .any(|kw| error_lower.contains(kw)); - - if has_quota_keyword { - prop_assert!( - is_quota_error, - "包含配额关键词的消息应该被检测为配额超限错误: {}", - error_message - ); - } - - // 如果检测到配额超限,标记凭证 - if is_quota_error { - let record = manager.mark_quota_exceeded(&credential_id, &error_message); - - // 验证凭证被标记为不可用 - prop_assert!( - !manager.is_available(&credential_id), - "配额超限后凭证应该不可用" - ); - - // 验证记录包含正确的信息 - prop_assert_eq!( - record.credential_id, - credential_id, - "记录的凭证 ID 应该正确" - ); - prop_assert_eq!( - record.reason, - error_message, - "记录的原因应该正确" - ); - - // 验证冷却结束时间在未来 - prop_assert!( - record.cooldown_until > chrono::Utc::now(), - "冷却结束时间应该在未来" - ); - } - } - - /// **Feature: cliproxyapi-parity, Property 16: Quota Exceeded Detection (Multiple Credentials)** - /// *For any* set of credentials, marking multiple as quota exceeded should track each independently. - /// **Validates: Requirements 8.1** - #[test] - fn prop_quota_exceeded_detection_multiple( - config in arb_quota_config(), - cred_count in 1usize..=10usize - ) { - let manager = QuotaManager::new(config); - - // 标记多个凭证为配额超限 - let mut marked_ids = Vec::new(); - for i in 0..cred_count { - let cred_id = format!("cred-{i}"); - manager.mark_quota_exceeded(&cred_id, "Rate limit exceeded"); - marked_ids.push(cred_id); - } - - // 验证所有凭证都被标记 - prop_assert_eq!( - manager.exceeded_count(), - cred_count, - "超限凭证数量应该正确" - ); - - // 验证每个凭证都不可用 - for id in &marked_ids { - prop_assert!( - !manager.is_available(id), - "凭证 {} 应该不可用", - id - ); - } - - // 验证未标记的凭证仍然可用 - prop_assert!( - manager.is_available("untracked-cred"), - "未标记的凭证应该可用" - ); - } -} - -proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: cliproxyapi-parity, Property 17: Quota Auto-Switch** - /// *For any* quota-exceeded credential when switch_project is enabled, - /// the next request SHALL use a different available credential. - /// **Validates: Requirements 8.2** - #[test] - fn prop_quota_auto_switch( - cred_count in 2usize..=10usize, - failed_index in 0usize..10usize, - model in arb_model_name() - ) { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: false, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - - // 创建凭证 ID 列表 - let available: Vec = (0..cred_count) - .map(|i| format!("cred-{i}")) - .collect(); - - let failed_index = failed_index % cred_count; - let failed_cred = &available[failed_index]; - - // 处理配额超限 - let result = manager.handle_quota_exceeded( - failed_cred, - &model, - &available, - "Rate limit exceeded", - ); - - // 验证:应该切换到不同的凭证 - prop_assert!( - result.switched, - "当 switch_project 启用且有其他可用凭证时,应该切换" - ); - - // 验证:新凭证不是失败的凭证 - let failed_cred_string = failed_cred.to_string(); - prop_assert_ne!( - result.new_credential_id.as_ref(), - Some(&failed_cred_string), - "新凭证不应该是失败的凭证" - ); - - // 验证:新凭证在可用列表中 - prop_assert!( - available.contains(result.new_credential_id.as_ref().unwrap()), - "新凭证应该在可用列表中" - ); - - // 验证:失败的凭证被标记为不可用 - prop_assert!( - !manager.is_available(failed_cred), - "失败的凭证应该被标记为不可用" - ); - } - - /// **Feature: cliproxyapi-parity, Property 17: Quota Auto-Switch (Disabled)** - /// *For any* quota-exceeded credential when switch_project is disabled, - /// the system SHALL NOT automatically switch to another credential. - /// **Validates: Requirements 8.2** - #[test] - fn prop_quota_auto_switch_disabled( - cred_count in 2usize..=10usize, - failed_index in 0usize..10usize, - model in arb_model_name() - ) { - let config = QuotaExceededConfig { - switch_project: false, - switch_preview_model: false, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - - // 创建凭证 ID 列表 - let available: Vec = (0..cred_count) - .map(|i| format!("cred-{i}")) - .collect(); - - let failed_index = failed_index % cred_count; - let failed_cred = &available[failed_index]; - - // 处理配额超限 - let result = manager.handle_quota_exceeded( - failed_cred, - &model, - &available, - "Rate limit exceeded", - ); - - // 验证:不应该切换凭证 - prop_assert!( - !result.switched, - "当 switch_project 禁用时,不应该切换凭证" - ); - - // 验证:失败的凭证仍然被标记为不可用 - prop_assert!( - !manager.is_available(failed_cred), - "失败的凭证应该被标记为不可用" - ); - } - - /// **Feature: cliproxyapi-parity, Property 17: Quota Auto-Switch (All Exhausted)** - /// *For any* set of credentials where all are quota-exceeded, - /// the system SHALL return an appropriate error. - /// **Validates: Requirements 8.2, 8.4** - #[test] - fn prop_quota_auto_switch_all_exhausted( - cred_count in 1usize..=5usize, - model in arb_model_name() - ) { - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: false, - cooldown_seconds: 300, - }; - let manager = QuotaManager::new(config); - - // 创建凭证 ID 列表 - let available: Vec = (0..cred_count) - .map(|i| format!("cred-{i}")) - .collect(); - - // 标记所有凭证为配额超限 - for cred_id in &available { - manager.mark_quota_exceeded(cred_id, "Rate limit exceeded"); - } - - // 处理最后一个凭证的配额超限 - let result = manager.handle_quota_exceeded( - &available[0], - &model, - &available, - "Rate limit exceeded", - ); - - // 验证:不应该切换(没有可用凭证) - prop_assert!( - !result.switched, - "当所有凭证都超限时,不应该切换" - ); - - // 验证:消息应该表明所有凭证都超限 - prop_assert!( - result.message.contains("所有凭证配额超限"), - "消息应该表明所有凭证都超限: {}", - result.message - ); - } -} - -proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: cliproxyapi-parity, Property 18: Quota Cooldown Expiration** - /// *For any* quota-exceeded credential, after the cooldown period expires, - /// the credential SHALL be restored to available status. - /// **Validates: Requirements 8.5** - #[test] - fn prop_quota_cooldown_expiration( - cred_count in 1usize..=10usize - ) { - // 使用 0 秒冷却时间,立即过期 - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 0, // 立即过期 - }; - let manager = QuotaManager::new(config); - - // 标记多个凭证为配额超限 - let cred_ids: Vec = (0..cred_count) - .map(|i| format!("cred-{i}")) - .collect(); - - for cred_id in &cred_ids { - manager.mark_quota_exceeded(cred_id, "Rate limit exceeded"); - } - - // 验证所有凭证都被标记 - prop_assert_eq!( - manager.exceeded_count(), - cred_count, - "所有凭证应该被标记为超限" - ); - - // 等待一小段时间确保过期 - std::thread::sleep(std::time::Duration::from_millis(100)); - - // 清理过期记录 - let cleaned = manager.cleanup_expired(); - - // 验证:所有记录都被清理 - prop_assert_eq!( - cleaned, - cred_count, - "所有过期记录应该被清理" - ); - - // 验证:所有凭证都恢复可用 - for cred_id in &cred_ids { - prop_assert!( - manager.is_available(cred_id), - "凭证 {} 应该恢复可用", - cred_id - ); - } - - // 验证:超限计数为 0 - prop_assert_eq!( - manager.exceeded_count(), - 0, - "超限凭证数量应该为 0" - ); - } - - /// **Feature: cliproxyapi-parity, Property 18: Quota Cooldown Expiration (Not Expired)** - /// *For any* quota-exceeded credential within the cooldown period, - /// the credential SHALL remain unavailable. - /// **Validates: Requirements 8.5** - #[test] - fn prop_quota_cooldown_not_expired( - cred_count in 1usize..=10usize - ) { - // 使用较长的冷却时间 - let config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 3600, // 1 小时 - }; - let manager = QuotaManager::new(config); - - // 标记多个凭证为配额超限 - let cred_ids: Vec = (0..cred_count) - .map(|i| format!("cred-{i}")) - .collect(); - - for cred_id in &cred_ids { - manager.mark_quota_exceeded(cred_id, "Rate limit exceeded"); - } - - // 尝试清理(不应该清理任何记录) - let cleaned = manager.cleanup_expired(); - - // 验证:没有记录被清理 - prop_assert_eq!( - cleaned, - 0, - "未过期的记录不应该被清理" - ); - - // 验证:所有凭证仍然不可用 - for cred_id in &cred_ids { - prop_assert!( - !manager.is_available(cred_id), - "凭证 {} 应该仍然不可用", - cred_id - ); - } - - // 验证:超限计数不变 - prop_assert_eq!( - manager.exceeded_count(), - cred_count, - "超限凭证数量应该不变" - ); - } - - /// **Feature: cliproxyapi-parity, Property 18: Quota Cooldown Expiration (Partial)** - /// *For any* set of credentials with mixed expiration states, - /// only expired credentials SHALL be restored. - /// **Validates: Requirements 8.5** - #[test] - fn prop_quota_cooldown_partial_expiration( - expired_count in 1usize..=5usize, - active_count in 1usize..=5usize - ) { - // 创建两个管理器:一个立即过期,一个长时间冷却 - let _expired_config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 0, // 立即过期 - }; - let active_config = QuotaExceededConfig { - switch_project: true, - switch_preview_model: true, - cooldown_seconds: 3600, // 1 小时 - }; - - // 使用一个管理器,但手动设置不同的过期时间 - let manager = QuotaManager::new(active_config); - - // 标记一些凭证为立即过期 - let expired_ids: Vec = (0..expired_count) - .map(|i| format!("expired-{i}")) - .collect(); - - // 标记一些凭证为长时间冷却 - let active_ids: Vec = (0..active_count) - .map(|i| format!("active-{i}")) - .collect(); - - // 先标记所有凭证 - for cred_id in &expired_ids { - manager.mark_quota_exceeded(cred_id, "Rate limit exceeded"); - } - for cred_id in &active_ids { - manager.mark_quota_exceeded(cred_id, "Rate limit exceeded"); - } - - // 手动将 expired_ids 的冷却时间设置为过去 - for cred_id in &expired_ids { - manager.set_cooldown_until(cred_id, chrono::Utc::now() - chrono::Duration::seconds(1)); - } - - // 清理过期记录 - let cleaned = manager.cleanup_expired(); - - // 验证:只有过期的记录被清理 - prop_assert_eq!( - cleaned, - expired_count, - "只有过期的记录应该被清理" - ); - - // 验证:过期的凭证恢复可用 - for cred_id in &expired_ids { - prop_assert!( - manager.is_available(cred_id), - "过期的凭证 {} 应该恢复可用", - cred_id - ); - } - - // 验证:未过期的凭证仍然不可用 - for cred_id in &active_ids { - prop_assert!( - !manager.is_available(cred_id), - "未过期的凭证 {} 应该仍然不可用", - cred_id - ); - } - } -} diff --git a/src-tauri/src/tests/mod.rs b/src-tauri/src/tests/mod.rs index 1a9d4497f..47a0268fc 100644 --- a/src-tauri/src/tests/mod.rs +++ b/src-tauri/src/tests/mod.rs @@ -1,3 +1,2 @@ -mod credential_tests; mod processor_tests; pub(crate) mod runtime_test_support; diff --git a/src-tauri/src/tests/processor_tests.rs b/src-tauri/src/tests/processor_tests.rs index 33b2be979..dc5b8deb3 100644 --- a/src-tauri/src/tests/processor_tests.rs +++ b/src-tauri/src/tests/processor_tests.rs @@ -2,13 +2,11 @@ use crate::ProviderType; use lime_processor::*; -use lime_services::provider_pool_service::ProviderPoolService; use std::sync::Arc; #[test] fn test_request_processor_new() { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); // 验证所有组件都已初始化 assert!(Arc::strong_count(&processor.router) >= 1); @@ -20,13 +18,11 @@ fn test_request_processor_new() { assert!(Arc::strong_count(&processor.plugins) >= 1); assert!(Arc::strong_count(&processor.stats) >= 1); assert!(Arc::strong_count(&processor.tokens) >= 1); - assert!(Arc::strong_count(&processor.pool_service) >= 1); } #[tokio::test] async fn test_request_processor_components() { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); // 验证路由器可以正常使用 { @@ -65,8 +61,7 @@ async fn test_request_processor_components() { #[tokio::test] async fn test_resolve_model_with_alias() { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); // 添加别名映射 { @@ -89,8 +84,7 @@ async fn test_resolve_model_with_alias() { #[tokio::test] async fn test_resolve_model_for_context() { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); // 添加别名映射 { @@ -115,8 +109,7 @@ async fn test_resolve_model_for_context() { #[tokio::test] async fn test_route_model_returns_default() { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); // 默认路由器为空,所有模型都应返回 None let (provider, is_default) = processor.route_model("gemini-2.5-flash").await; @@ -130,8 +123,7 @@ async fn test_route_model_returns_default() { #[tokio::test] async fn test_route_for_context() { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); // 创建请求上下文 let mut ctx = RequestContext::new("gemini-2.5-flash".to_string()); @@ -146,8 +138,7 @@ async fn test_route_for_context() { #[tokio::test] async fn test_resolve_and_route() { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); // 添加别名映射 { @@ -262,8 +253,7 @@ proptest! { fn prop_request_stats_recorded( log in arb_request_log() ) { - let pool_service = Arc::new(ProviderPoolService::new()); - let processor = RequestProcessor::with_defaults(pool_service); + let processor = RequestProcessor::with_defaults(); let original_id = log.id.clone(); let original_provider = log.provider; diff --git a/src-tauri/tauri.conf.headless.json b/src-tauri/tauri.conf.headless.json index ef50f0ea0..e8a0a6f07 100644 --- a/src-tauri/tauri.conf.headless.json +++ b/src-tauri/tauri.conf.headless.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "Lime", - "version": "1.21.0", + "version": "1.22.0", "identifier": "com.limecloud.lime.headless", "build": { "beforeDevCommand": "npm run dev:web-bridge", diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 3da69b9f2..ab63adb60 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "Lime", - "version": "1.21.0", + "version": "1.22.0", "identifier": "com.limecloud.lime", "build": { "beforeDevCommand": "node scripts/start-tauri-dev-server.mjs", diff --git a/src-tauri/tests/real_web_search_preflight_short_input.rs b/src-tauri/tests/real_web_search_preflight_short_input.rs index 1f82b72e0..128c014a8 100644 --- a/src-tauri/tests/real_web_search_preflight_short_input.rs +++ b/src-tauri/tests/real_web_search_preflight_short_input.rs @@ -92,12 +92,19 @@ async fn test_real_web_search_preflight_short_input_continue() { false, ); let mut tracker = WebSearchExecutionTracker::default(); + let mut metadata = std::collections::HashMap::new(); + metadata.insert("webSearchEnabled".to_string(), serde_json::json!(true)); + let turn_context = aster::session::TurnContextOverride { + metadata, + ..aster::session::TurnContextOverride::default() + }; let execution = execute_web_search_preflight_if_needed( agent, &session_id, "继续", None, None, + Some(turn_context), &policy, &mut tracker, ) diff --git a/src/components/README.md b/src/components/README.md index af48f2b37..5f063f30f 100644 --- a/src/components/README.md +++ b/src/components/README.md @@ -14,7 +14,7 @@ React 组件层,包含 UI 组件和业务组件。 - `flow-monitor/` - LLM 流量监控组件 - `mcp/` - MCP 服务器管理组件(配置管理、运行时控制、工具/提示词/资源浏览与调用) - `plugins/` - 插件管理组件 -- `provider-pool/` - Provider 凭证池管理组件 +- `provider-pool/api-key/` - API Key Provider 设置组件(目录名保留历史路径) - `routing/` - 路由规则配置组件 - `smart-input/` - 截图/语音浮窗共享组件(当前仅保留快捷键设置) - `settings-v2/` - 设置页面组件(当前主实现) diff --git a/src/components/agent/chat/AgentChatWorkspace.tsx b/src/components/agent/chat/AgentChatWorkspace.tsx index edee45498..6a772579f 100644 --- a/src/components/agent/chat/AgentChatWorkspace.tsx +++ b/src/components/agent/chat/AgentChatWorkspace.tsx @@ -486,6 +486,14 @@ export function AgentChatWorkspace({ openBrowserAssistOnMount = false, initialSiteSkillLaunch, }: AgentChatWorkspaceProps) { + // 性能埋点:记录组件渲染开始时间 + const workspaceRenderT0 = useRef(performance.now()); + useEffect(() => { + console.info( + `[PERF] AgentChatWorkspace mounted: ${(performance.now() - workspaceRenderT0.current).toFixed(0)}ms`, + ); + }, []); + const normalizedEntryTheme = normalizeInitialTheme(initialTheme); const shouldAutoCollapseClassicClawSidebar = agentEntry === "claw"; const defaultTopicSidebarVisible = @@ -733,6 +741,7 @@ export function AgentChatWorkspace({ let cancelled = false; const startedAt = Date.now(); + const perfT0 = performance.now(); logAgentDebug("AgentChatPage", "resolveDefaultProjectAlias.start", { externalProjectId: externalProjectId ?? null, }); @@ -828,6 +837,9 @@ export function AgentChatWorkspace({ projectId: defaultProject.id, rootPath: resolvedRootPath, }); + console.info( + `[PERF] resolveDefaultProjectAlias: ${(performance.now() - perfT0).toFixed(0)}ms`, + ); } catch (error) { if (cancelled) { return; @@ -3121,6 +3133,7 @@ export function AgentChatWorkspace({ let disposed = false; setGeneralWorkbenchEntryCheckPending(true); + const perfT0 = performance.now(); void (async () => { try { @@ -3129,6 +3142,10 @@ export function AgentChatWorkspace({ 3, ).catch(() => null); + console.info( + `[PERF] executionRunGetGeneralWorkbenchState: ${(performance.now() - perfT0).toFixed(0)}ms`, + ); + if (disposed) { return; } diff --git a/src/components/agent/chat/agentChatWorkspaceLoader.ts b/src/components/agent/chat/agentChatWorkspaceLoader.ts index 11f139c45..caccc66c7 100644 --- a/src/components/agent/chat/agentChatWorkspaceLoader.ts +++ b/src/components/agent/chat/agentChatWorkspaceLoader.ts @@ -3,7 +3,7 @@ const MODULE_IMPORT_FAILURE_PATTERNS = [ "Failed to fetch dynamically imported module", ] as const; -const DEFAULT_RETRY_DELAYS_MS = [200, 600] as const; +const DEFAULT_RETRY_DELAYS_MS = [120, 320] as const; function isRetryableModuleImportFailure(error: unknown): boolean { if (!(error instanceof Error)) { @@ -51,3 +51,7 @@ export async function loadModuleWithRetry( export function loadAgentChatWorkspaceModule() { return loadModuleWithRetry(() => import("./AgentChatWorkspace")); } + +export function preloadAgentChatWorkspaceModule(): void { + void loadAgentChatWorkspaceModule(); +} diff --git a/src/components/agent/chat/components/ChatModelSelector.integration.test.tsx b/src/components/agent/chat/components/ChatModelSelector.integration.test.tsx index f91ad2ba8..7efbe6b23 100644 --- a/src/components/agent/chat/components/ChatModelSelector.integration.test.tsx +++ b/src/components/agent/chat/components/ChatModelSelector.integration.test.tsx @@ -16,7 +16,6 @@ const { mockLoadProviderModels, mockUseConfiguredProviders, mockUseProviderModels, - mockProviderPoolGetOverview, mockApiKeyProvidersGetProviders, mockEmitProviderDataChanged, mockWechatChannelSetRuntimeModel, @@ -38,7 +37,6 @@ const { mockLoadProviderModels: vi.fn(), mockUseConfiguredProviders: vi.fn(), mockUseProviderModels: vi.fn(), - mockProviderPoolGetOverview: vi.fn(), mockApiKeyProvidersGetProviders: vi.fn(), mockEmitProviderDataChanged: vi.fn(), mockWechatChannelSetRuntimeModel: vi.fn(async () => undefined), @@ -119,12 +117,6 @@ vi.mock("@/hooks/useProviderModels", () => ({ useProviderModels: mockUseProviderModels, })); -vi.mock("@/lib/api/providerPool", () => ({ - providerPoolApi: { - getOverview: mockProviderPoolGetOverview, - }, -})); - vi.mock("@/lib/api/apiKeyProvider", () => ({ apiKeyProviderApi: { getProviders: mockApiKeyProvidersGetProviders, @@ -321,7 +313,6 @@ beforeEach(() => { mockCreateAgentRuntimeSession.mockResolvedValue("created-session"); mockUpdateAgentRuntimeSession.mockResolvedValue(undefined); mockSafeListen.mockResolvedValue(() => {}); - mockProviderPoolGetOverview.mockResolvedValue([]); mockApiKeyProvidersGetProviders.mockResolvedValue([]); mockEmitProviderDataChanged.mockImplementation(() => {}); mockWechatChannelSetRuntimeModel.mockResolvedValue(undefined); diff --git a/src/components/agent/chat/components/MessageList.test.tsx b/src/components/agent/chat/components/MessageList.test.tsx index d41f9a1c0..9436d50bf 100644 --- a/src/components/agent/chat/components/MessageList.test.tsx +++ b/src/components/agent/chat/components/MessageList.test.tsx @@ -271,6 +271,26 @@ describe("MessageList", () => { expect(leadingContent?.textContent).toContain("scene summary heading"); expect(messageColumn?.firstElementChild).toBe(leadingContent); expect(messageColumn?.textContent).toContain("第一条消息"); + expect(messageColumn?.className).toContain("justify-start"); + }); + + it("短对话首帧应贴近输入区底部,避免发送后消息短暂贴顶", () => { + const container = render([ + { + id: "msg-user-first-frame", + role: "user", + content: "你好", + timestamp: new Date("2026-04-25T10:00:00.000Z"), + } as Message, + ]); + + const messageColumn = container.querySelector( + '[data-testid="message-list-column"]', + ); + + expect(messageColumn?.textContent).toContain("你好"); + expect(messageColumn?.className).toContain("min-h-full"); + expect(messageColumn?.className).toContain("justify-end"); }); it("自动恢复生成会话时应展示恢复占位而不是空白引导", () => { @@ -726,7 +746,47 @@ describe("MessageList", () => { expect(container.textContent).not.toContain("正在整理相关信息"); }); - it("首个文本分片到来前,不应渲染空白 assistant 气泡,只保留运行态行", () => { + it("assistant 已有正文且仍在发送时,不应在消息尾部追加处理中状态回复", () => { + const now = new Date(); + const messages: Message[] = [ + { + id: "msg-user-active-status-tail", + role: "user", + content: "hello", + timestamp: now, + }, + { + id: "msg-assistant-active-status-tail", + role: "assistant", + content: "我正在处理你的请求。", + timestamp: new Date(now.getTime() + 1000), + isThinking: true, + runtimeStatus: { + phase: "routing", + title: "处理中", + detail: "正在等待模型输出。", + checkpoints: ["请求已发送"], + }, + }, + ]; + + const container = render(messages, { + isSending: true, + }); + + expect( + container.querySelector('[data-testid="streaming-renderer"]'), + ).not.toBeNull(); + expect(container.textContent).toContain("我正在处理你的请求。"); + expect( + container.querySelector('[data-testid="assistant-message-meta-footer"]'), + ).toBeNull(); + expect( + container.querySelector('[data-testid="inputbar-runtime-status-line"]'), + ).toBeNull(); + }); + + it("首个文本分片到来前,不应把运行态当作 assistant 回复渲染", () => { const now = new Date(); const messages: Message[] = [ { @@ -753,11 +813,11 @@ describe("MessageList", () => { ).toBeNull(); expect( container.querySelector('[data-testid="assistant-message-meta-footer"]'), - ).not.toBeNull(); + ).toBeNull(); expect( container.querySelector('[data-testid="inputbar-runtime-status-line"]'), - ).not.toBeNull(); - expect(container.textContent).toContain("处理中"); + ).toBeNull(); + expect(container.textContent).not.toContain("处理中"); expect(container.textContent).not.toContain(""); }); @@ -794,15 +854,15 @@ describe("MessageList", () => { ).toBeNull(); expect( container.querySelector('[data-testid="assistant-message-meta-footer"]'), - ).not.toBeNull(); + ).toBeNull(); expect( container.querySelector('[data-testid="inputbar-runtime-status-line"]'), - ).not.toBeNull(); - expect(container.textContent).toContain("处理中"); + ).toBeNull(); + expect(container.textContent).not.toContain("处理中"); expect(container.textContent).not.toContain("Built-in Tool"); }); - it("assistant 占位消息只有启动态 runtimeStatus 时,也不应保留空白气泡", () => { + it("assistant 占位消息只有启动态 runtimeStatus 时,也不应保留状态回复", () => { const now = new Date(); const messages: Message[] = [ { @@ -840,11 +900,11 @@ describe("MessageList", () => { ).toBeNull(); expect( container.querySelector('[data-testid="assistant-message-meta-footer"]'), - ).not.toBeNull(); + ).toBeNull(); expect( container.querySelector('[data-testid="inputbar-runtime-status-line"]'), - ).not.toBeNull(); - expect(container.textContent).toContain("处理中"); + ).toBeNull(); + expect(container.textContent).not.toContain("处理中"); expect(container.textContent).not.toContain("正在启动处理流程"); }); @@ -2563,10 +2623,10 @@ describe("MessageList", () => { ).toBeNull(); expect( container.querySelector('[data-testid="assistant-message-meta-footer"]'), - ).not.toBeNull(); + ).toBeNull(); expect( container.querySelector('[data-testid="inputbar-runtime-status-line"]'), - ).not.toBeNull(); + ).toBeNull(); }); it("本地工具批次的阶段结论不应再进入主消息流时间线", () => { @@ -3385,7 +3445,7 @@ describe("MessageList", () => { ).toBeTruthy(); expect( container.querySelector('[data-testid="assistant-message-meta-footer"]'), - ).not.toBeNull(); + ).toBeNull(); expect( container.querySelector('[data-testid="agent-thread-reliability-panel"]'), ).toBeNull(); diff --git a/src/components/agent/chat/components/MessageList.tsx b/src/components/agent/chat/components/MessageList.tsx index d3b75d377..f30cda65f 100644 --- a/src/components/agent/chat/components/MessageList.tsx +++ b/src/components/agent/chat/components/MessageList.tsx @@ -763,6 +763,15 @@ const MessageListInner: React.FC = ({ () => buildMessageTurnGroups(renderedMessages), [renderedMessages], ); + const hasStickyTopContent = + Boolean(leadingContent) || + persistedHiddenHistoryCount > 0 || + hiddenHistoryCount > 0; + const shouldBottomAnchorMessageStack = + messageGroups.length > 0 && + !hasStickyTopContent && + !isRestoringSession && + !isTaskCenterEmptyState; const renderGroups = useMemo( () => messageGroups.map((group) => { @@ -1088,19 +1097,6 @@ const MessageListInner: React.FC = ({ ), ) : []; - const shouldRenderTailRuntimeStatusLine = - msg.role === "assistant" && - msg.id === lastAssistantMessageId && - isConversationTailAssistant && - Boolean(tailRuntimeStatusLine); - const shouldRenderUsageFooter = - isConversationTailAssistant && - !shouldRenderTailRuntimeStatusLine && - !msg.isThinking && - Boolean(msg.usage); - const shouldRenderStatusPill = - !shouldRenderTailRuntimeStatusLine && - shouldRenderRuntimeStatusPill(msg.runtimeStatus); const messageCanvasShortcutTitle = messageSavedSiteContentTarget ? resolveSiteSavedContentTargetDisplayName( messageSavedSiteContentTarget, @@ -1125,6 +1121,23 @@ const MessageListInner: React.FC = ({ !msg.taskPreview; const hasAssistantBodyContent = msg.role !== "assistant" || !shouldCollapseAssistantShell; + const shouldSuppressActiveRuntimeLine = + tailRuntimeStatusLine?.status === "running" || + tailRuntimeStatusLine?.status === "queued"; + const shouldRenderTailRuntimeStatusLine = + msg.role === "assistant" && + msg.id === lastAssistantMessageId && + isConversationTailAssistant && + Boolean(tailRuntimeStatusLine) && + !shouldSuppressActiveRuntimeLine; + const shouldRenderUsageFooter = + isConversationTailAssistant && + !shouldRenderTailRuntimeStatusLine && + !msg.isThinking && + Boolean(msg.usage); + const shouldRenderStatusPill = + !shouldRenderTailRuntimeStatusLine && + shouldRenderRuntimeStatusPill(msg.runtimeStatus); const assistantMetaFooter = msg.role === "assistant" && (shouldRenderTailRuntimeStatusLine || @@ -1525,11 +1538,13 @@ const MessageListInner: React.FC = ({ >
{leadingContent ? (
{leadingContent}
diff --git a/src/components/agent/chat/hooks/agentRuntimeAdapter.test.ts b/src/components/agent/chat/hooks/agentRuntimeAdapter.test.ts index c528f2c25..309d161b3 100644 --- a/src/components/agent/chat/hooks/agentRuntimeAdapter.test.ts +++ b/src/components/agent/chat/hooks/agentRuntimeAdapter.test.ts @@ -122,4 +122,23 @@ describe("defaultAgentRuntimeAdapter", () => { workspaceId: "workspace-9", }); }); + + it("generateSessionTitle 应透传标题预览文本", async () => { + const client = { + ...mockRuntimeClient, + generateAgentRuntimeSessionTitle: vi.fn().mockResolvedValue("新标题"), + }; + const adapter = createAgentRuntimeAdapter({ + client, + }); + + await expect( + adapter.generateSessionTitle?.("session-9", "user:请整理支付异常"), + ).resolves.toBe("新标题"); + + expect(client.generateAgentRuntimeSessionTitle).toHaveBeenCalledWith( + "session-9", + "user:请整理支付异常", + ); + }); }); diff --git a/src/components/agent/chat/hooks/agentRuntimeAdapter.ts b/src/components/agent/chat/hooks/agentRuntimeAdapter.ts index 5b5c39390..62290c814 100644 --- a/src/components/agent/chat/hooks/agentRuntimeAdapter.ts +++ b/src/components/agent/chat/hooks/agentRuntimeAdapter.ts @@ -70,7 +70,10 @@ export interface AgentRuntimeAdapter { providerType: string, model: string, ): Promise; - generateSessionTitle?(sessionId: string): Promise; + generateSessionTitle?( + sessionId: string, + previewText?: string, + ): Promise; submitOp(op: AgentOp): Promise; compactSession(sessionId: string, eventName: string): Promise; interruptTurn(sessionId: string): Promise; @@ -169,8 +172,8 @@ export function createAgentRuntimeAdapter({ model_name: model, }); }, - async generateSessionTitle(sessionId) { - return client.generateAgentRuntimeSessionTitle(sessionId); + async generateSessionTitle(sessionId, previewText) { + return client.generateAgentRuntimeSessionTitle(sessionId, previewText); }, async submitOp(op) { switch (op.type) { diff --git a/src/components/agent/chat/hooks/agentStreamRuntimeHandler.ts b/src/components/agent/chat/hooks/agentStreamRuntimeHandler.ts index 581c60408..a811e60c9 100644 --- a/src/components/agent/chat/hooks/agentStreamRuntimeHandler.ts +++ b/src/components/agent/chat/hooks/agentStreamRuntimeHandler.ts @@ -54,6 +54,20 @@ import { buildToolResultArtifactFromToolResult, } from "../utils/taskPreviewFromToolResult"; +function appendWithOverlapDetection(base: string, chunk: string): string { + if (!base) return chunk; + if (!chunk) return base; + if (chunk.startsWith(base)) return chunk; + if (base.endsWith(chunk)) return base; + const maxOverlap = Math.min(base.length, chunk.length); + for (let overlap = maxOverlap; overlap > 0; overlap -= 1) { + if (base.slice(-overlap) === chunk.slice(0, overlap)) { + return base + chunk.slice(overlap); + } + } + return base + chunk; +} + type MessageParts = NonNullable; interface StreamObserver { @@ -560,7 +574,10 @@ export function handleTurnStreamEvent({ ? { ...msg, isThinking: true, - thinkingContent: (msg.thinkingContent || "") + data.text, + thinkingContent: appendWithOverlapDetection( + msg.thinkingContent || "", + data.text, + ), contentParts: appendThinkingToParts( msg.contentParts || [], data.text, diff --git a/src/components/agent/chat/hooks/useAgentSession.ts b/src/components/agent/chat/hooks/useAgentSession.ts index fc3087cae..c04c4728b 100644 --- a/src/components/agent/chat/hooks/useAgentSession.ts +++ b/src/components/agent/chat/hooks/useAgentSession.ts @@ -1139,8 +1139,9 @@ export function useAgentSession(options: UseAgentSessionOptions) { const detail = resumeSessionStartHooks ? await runtime.getSession(topicId, { resumeSessionStartHooks: true, + historyLimit: 40, }) - : await runtime.getSession(topicId); + : await runtime.getSession(topicId, { historyLimit: 40 }); logAgentDebug("useAgentSession", "switchTopic.fetchDetail.success", { itemsCount: detail.items?.length ?? 0, messagesCount: detail.messages.length, diff --git a/src/components/agent/chat/hooks/useAgentStream.ts b/src/components/agent/chat/hooks/useAgentStream.ts index c03fd20fa..f3752ef05 100644 --- a/src/components/agent/chat/hooks/useAgentStream.ts +++ b/src/components/agent/chat/hooks/useAgentStream.ts @@ -53,17 +53,35 @@ function appendThinkingToParts( const lastPart = nextParts[nextParts.length - 1]; if (lastPart?.type === "thinking") { - nextParts[nextParts.length - 1] = { - type: "thinking", - text: lastPart.text + textDelta, - }; + const base = lastPart.text; + const chunk = textDelta; + let merged: string; + if (!base) { + merged = chunk; + } else if (!chunk) { + merged = base; + } else if (chunk.startsWith(base)) { + merged = chunk; + } else if (base.endsWith(chunk)) { + merged = base; + } else { + const maxOverlap = Math.min(base.length, chunk.length); + let found = false; + merged = base + chunk; + for (let overlap = maxOverlap; overlap > 0; overlap -= 1) { + if (base.slice(-overlap) === chunk.slice(0, overlap)) { + merged = base + chunk.slice(overlap); + found = true; + break; + } + } + void found; + } + nextParts[nextParts.length - 1] = { type: "thinking", text: merged }; return nextParts; } - nextParts.push({ - type: "thinking", - text: textDelta, - }); + nextParts.push({ type: "thinking", text: textDelta }); return nextParts; } diff --git a/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx b/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx index adb1fd4e7..2a07e90d6 100644 --- a/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx +++ b/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx @@ -11,6 +11,7 @@ const { mockListAgentRuntimeSessions, mockGetAgentRuntimeSession, mockGetAgentRuntimeThreadRead, + mockGenerateAgentRuntimeSessionTitle, mockUpdateAgentRuntimeSession, mockDeleteAgentRuntimeSession, mockCompactAgentRuntimeSession, @@ -36,6 +37,7 @@ const { mockListAgentRuntimeSessions: vi.fn(), mockGetAgentRuntimeSession: vi.fn(), mockGetAgentRuntimeThreadRead: vi.fn(), + mockGenerateAgentRuntimeSessionTitle: vi.fn(), mockUpdateAgentRuntimeSession: vi.fn(), mockDeleteAgentRuntimeSession: vi.fn(), mockCompactAgentRuntimeSession: vi.fn(), @@ -74,6 +76,7 @@ vi.mock("@/lib/api/agentRuntime", () => ({ listAgentRuntimeSessions: mockListAgentRuntimeSessions, getAgentRuntimeSession: mockGetAgentRuntimeSession, getAgentRuntimeThreadRead: mockGetAgentRuntimeThreadRead, + generateAgentRuntimeSessionTitle: mockGenerateAgentRuntimeSessionTitle, updateAgentRuntimeSession: mockUpdateAgentRuntimeSession, deleteAgentRuntimeSession: mockDeleteAgentRuntimeSession, compactAgentRuntimeSession: mockCompactAgentRuntimeSession, @@ -90,6 +93,7 @@ vi.mock("@/lib/api/agentRuntime", () => ({ listAgentRuntimeSessions: mockListAgentRuntimeSessions, getAgentRuntimeSession: mockGetAgentRuntimeSession, getAgentRuntimeThreadRead: mockGetAgentRuntimeThreadRead, + generateAgentRuntimeSessionTitle: mockGenerateAgentRuntimeSessionTitle, updateAgentRuntimeSession: mockUpdateAgentRuntimeSession, deleteAgentRuntimeSession: mockDeleteAgentRuntimeSession, compactAgentRuntimeSession: mockCompactAgentRuntimeSession, @@ -299,7 +303,7 @@ function seedSessionSnapshots( ); } - beforeEach(() => { +beforeEach(() => { ( globalThis as typeof globalThis & { IS_REACT_ACT_ENVIRONMENT?: boolean; @@ -312,6 +316,7 @@ function seedSessionSnapshots( mockListAgentRuntimeSessions.mockReset(); mockGetAgentRuntimeSession.mockReset(); mockGetAgentRuntimeThreadRead.mockReset(); + mockGenerateAgentRuntimeSessionTitle.mockReset(); mockUpdateAgentRuntimeSession.mockReset(); mockDeleteAgentRuntimeSession.mockReset(); mockCompactAgentRuntimeSession.mockReset(); @@ -324,14 +329,14 @@ function seedSessionSnapshots( mockSafeListen.mockReset(); mockParseSkillSlashCommand.mockReset(); mockTryExecuteSlashSkillCommand.mockReset(); - mockWechatChannelSetRuntimeModel.mockReset(); - mockGetDefaultProvider.mockReset(); - mockResolveClawWorkspaceProviderSelection.mockReset(); - mockScheduleMinimumDelayIdleTask.mockReset(); - mockScheduleMinimumDelayIdleTask.mockImplementation((task: () => void) => { - task(); - return () => undefined; - }); + mockWechatChannelSetRuntimeModel.mockReset(); + mockGetDefaultProvider.mockReset(); + mockResolveClawWorkspaceProviderSelection.mockReset(); + mockScheduleMinimumDelayIdleTask.mockReset(); + mockScheduleMinimumDelayIdleTask.mockImplementation((task: () => void) => { + task(); + return () => undefined; + }); mockToast.success.mockReset(); mockToast.error.mockReset(); mockToast.info.mockReset(); @@ -348,6 +353,7 @@ function seedSessionSnapshots( messages: [], }); mockGetAgentRuntimeThreadRead.mockResolvedValue(undefined); + mockGenerateAgentRuntimeSessionTitle.mockResolvedValue(""); mockUpdateAgentRuntimeSession.mockResolvedValue(undefined); mockDeleteAgentRuntimeSession.mockResolvedValue(undefined); mockCompactAgentRuntimeSession.mockResolvedValue(undefined); @@ -840,7 +846,7 @@ describe("useAsterAgentChat 任务快照", () => { expect(mockGetAgentRuntimeSession).toHaveBeenCalledWith( sessionId, - undefined, + { historyLimit: 40 }, ); expect(mockInterruptAgentRuntimeTurn).toHaveBeenCalledWith({ session_id: sessionId, @@ -978,7 +984,7 @@ describe("useAsterAgentChat 任务快照", () => { expect(mockGetAgentRuntimeSession).toHaveBeenCalledWith( sessionId, - undefined, + { historyLimit: 40 }, ); expect(harness.getValue().queuedTurns).toEqual([ { @@ -2833,7 +2839,7 @@ describe("useAsterAgentChat runtime routing", () => { expect(mockGetAgentRuntimeSession).toHaveBeenCalledWith( sessionId, - undefined, + { historyLimit: 40 }, ); expect(harness.getValue().currentTurnId).toBe("turn-real-1"); expect(harness.getValue().threadItems).toEqual( @@ -3041,7 +3047,7 @@ describe("useAsterAgentChat runtime routing", () => { .find((msg) => msg.role === "assistant"); expect(mockGetAgentRuntimeSession).toHaveBeenCalledWith( sessionId, - undefined, + { historyLimit: 40 }, ); expect(assistantMessage).toBeTruthy(); expect(assistantMessage?.content).toContain( @@ -5802,7 +5808,7 @@ describe("useAsterAgentChat 偏好持久化", () => { await flushEffects(); expect(mockGetAgentRuntimeSession).toHaveBeenCalledWith( sessionId, - undefined, + { historyLimit: 40 }, ); expect( JSON.parse( @@ -5870,7 +5876,7 @@ describe("useAsterAgentChat 偏好持久化", () => { expect(mockGetAgentRuntimeSession).toHaveBeenCalledWith( sessionId, - undefined, + { historyLimit: 40 }, ); expect(harness.getValue().sessionId).toBeNull(); expect( @@ -6115,8 +6121,8 @@ describe("useAsterAgentChat 偏好持久化", () => { }); await act(async () => { - harness.getValue().setProviderType("antigravity"); - harness.getValue().setModel("gemini-3-pro-image-preview"); + harness.getValue().setProviderType("gemini"); + harness.getValue().setModel("gemini-3-pro-preview"); }); await flushEffects(); @@ -6532,7 +6538,7 @@ describe("useAsterAgentChat 偏好持久化", () => { expect(mockScheduleMinimumDelayIdleTask).not.toHaveBeenCalled(); expect(mockGetAgentRuntimeSession).toHaveBeenCalledWith( "topic-stale", - undefined, + { historyLimit: 40 }, ); await act(async () => { @@ -9423,6 +9429,122 @@ describe("useAsterAgentChat 兼容接口", () => { } }); + it("自动标题生成进行中时,话题状态刷新不应取消导航标题回写", async () => { + const workspaceId = "ws-auto-title"; + const sessionId = "session-auto-title"; + const generatedTitle = "支付页错误定位"; + const deferredTitle = createDeferred(); + mockCreateAgentRuntimeSession.mockResolvedValue(sessionId); + mockGenerateAgentRuntimeSessionTitle.mockReturnValueOnce( + deferredTitle.promise, + ); + + const harness = mountHook(workspaceId); + + try { + await flushEffects(); + + await act(async () => { + await harness.getValue().createFreshSession(); + }); + await flushEffects(); + + await act(async () => { + harness.getValue().setMessages([ + { + id: "msg-user-title", + role: "user", + content: "帮我定位支付页提交时报 500 的问题", + timestamp: new Date(), + }, + { + id: "msg-assistant-title", + role: "assistant", + content: "我会先检查支付页请求链路。", + timestamp: new Date(), + }, + ]); + }); + await flushEffects(); + + expect(mockGenerateAgentRuntimeSessionTitle).toHaveBeenCalledWith( + sessionId, + expect.stringContaining("帮我定位支付页提交时报 500 的问题"), + ); + + await act(async () => { + harness.getValue().updateTopicSnapshot(sessionId, { + lastPreview: "正在检查支付页请求链路。", + messagesCount: 2, + status: "running", + }); + }); + + await act(async () => { + deferredTitle.resolve(generatedTitle); + await Promise.resolve(); + }); + await flushEffects(); + + expect(mockUpdateAgentRuntimeSession).toHaveBeenCalledWith({ + session_id: sessionId, + name: generatedTitle, + }); + expect( + harness.getValue().topics.find((topic) => topic.id === sessionId) + ?.title, + ).toBe(generatedTitle); + } finally { + harness.unmount(); + } + }); + + it("自动标题应覆盖新对话占位标题", async () => { + const workspaceId = "ws-auto-title-new-dialogue"; + const sessionId = "session-new-dialogue-title"; + const generatedTitle = "国际新闻摘要"; + mockCreateAgentRuntimeSession.mockResolvedValue(sessionId); + mockGenerateAgentRuntimeSessionTitle.mockResolvedValue(generatedTitle); + + const harness = mountHook(workspaceId); + + try { + await flushEffects(); + + await act(async () => { + await harness.getValue().createFreshSession("新对话"); + }); + await flushEffects(); + + await act(async () => { + harness.getValue().setMessages([ + { + id: "msg-user-new-dialogue-title", + role: "user", + content: "请用 WebSearch 查询今天国际新闻,简短列出 3 条。", + timestamp: new Date(), + }, + ]); + }); + await flushEffects(); + + expect(mockGenerateAgentRuntimeSessionTitle).toHaveBeenCalledWith( + sessionId, + expect.stringContaining("今天国际新闻"), + ); + expect(mockUpdateAgentRuntimeSession).toHaveBeenCalledWith({ + session_id: sessionId, + name: generatedTitle, + }); + expect( + harness.getValue().topics.find((topic) => topic.id === sessionId) + ?.title, + ).toBe(generatedTitle); + } finally { + harness.unmount(); + } + }); + it("deleteTopic 应调用后端并刷新话题列表", async () => { const createdAt = Math.floor(Date.now() / 1000); let currentSessions = [ diff --git a/src/components/agent/chat/hooks/useAsterAgentChat.ts b/src/components/agent/chat/hooks/useAsterAgentChat.ts index 3dd5a97e7..0293c81d9 100644 --- a/src/components/agent/chat/hooks/useAsterAgentChat.ts +++ b/src/components/agent/chat/hooks/useAsterAgentChat.ts @@ -42,6 +42,43 @@ type UseAsterAgentChatRuntimeOptions = UseAsterAgentChatOptions & { const AUTO_TITLE_DEFERRED_LOAD_MS = 10_000; const AUTO_TITLE_IDLE_TIMEOUT_MS = 2_000; +const AUTO_TITLE_PLACEHOLDER_TITLES = new Set([ + "", + "新任务", + "新话题", + "新对话", +]); + +function isAutoTitlePlaceholder(title: string | null | undefined): boolean { + return AUTO_TITLE_PLACEHOLDER_TITLES.has(title?.trim() ?? ""); +} + +function isPreviewDerivedTitle( + title: string | null | undefined, + messages: Array<{ role: string; content: unknown }>, +): boolean { + const normalizedTitle = title?.trim(); + if (!normalizedTitle) { + return false; + } + + const firstAssistantMessage = messages.find( + (message) => + message.role === "assistant" && + typeof message.content === "string" && + message.content.trim().length > 0, + ); + if (!firstAssistantMessage || typeof firstAssistantMessage.content !== "string") { + return false; + } + + const normalizedMessage = firstAssistantMessage.content.trim(); + const messagePrefix = normalizedMessage.slice(0, Math.max(16, normalizedTitle.length)); + return ( + normalizedMessage.startsWith(normalizedTitle) || + normalizedTitle.startsWith(messagePrefix) + ); +} export function useAsterAgentChat(options: UseAsterAgentChatRuntimeOptions) { const { @@ -276,6 +313,17 @@ export function useAsterAgentChat(options: UseAsterAgentChatRuntimeOptions) { const sessionTopics = session.topics; const sessionSetTopics = session.setTopics; const currentSessionId = session.sessionId; + const activeSessionTitle = useMemo(() => { + const activeSessionId = currentSessionId?.trim(); + if (!activeSessionId) { + return null; + } + + const activeTopic = sessionTopics.find( + (topic) => topic.id === activeSessionId, + ); + return activeTopic?.title?.trim() ?? null; + }, [currentSessionId, sessionTopics]); useEffect(() => { const activeSessionId = currentSessionId?.trim(); @@ -283,17 +331,12 @@ export function useAsterAgentChat(options: UseAsterAgentChatRuntimeOptions) { return; } - const activeTopic = sessionTopics.find( - (topic) => topic.id === activeSessionId, - ); - if (!activeTopic) { + if (activeSessionTitle === null) { return; } - const activeTitle = activeTopic?.title?.trim() || ""; const shouldAutoGenerateTitle = - activeTitle === "" || - activeTitle === "新任务" || - activeTitle === "新话题"; + isAutoTitlePlaceholder(activeSessionTitle) || + isPreviewDerivedTitle(activeSessionTitle, sessionMessages); if (!shouldAutoGenerateTitle) { autoTitleCompletedSessionIdsRef.current.add(activeSessionId); return; @@ -324,14 +367,27 @@ export function useAsterAgentChat(options: UseAsterAgentChatRuntimeOptions) { () => { void (async () => { try { + const conversationText = sessionMessages + .filter( + (msg) => + (msg.role === "user" || msg.role === "assistant") && + typeof msg.content === "string" && + msg.content.trim().length > 0, + ) + .map((msg) => `${msg.role}:${msg.content}`) + .join("\n") + .slice(-1000); + const generatedTitle = ( - await runtime.generateSessionTitle?.(activeSessionId) + await runtime.generateSessionTitle?.( + activeSessionId, + conversationText, + ) )?.trim(); if ( cancelled || !generatedTitle || - generatedTitle === "新任务" || - generatedTitle === "新话题" + isAutoTitlePlaceholder(generatedTitle) ) { return; } @@ -374,11 +430,11 @@ export function useAsterAgentChat(options: UseAsterAgentChatRuntimeOptions) { } }; }, [ + activeSessionTitle, currentSessionId, runtime, sessionMessages, sessionSetTopics, - sessionTopics, stream.isSending, ]); diff --git a/src/components/agent/chat/index.tsx b/src/components/agent/chat/index.tsx index e6a9c0c74..65961c69b 100644 --- a/src/components/agent/chat/index.tsx +++ b/src/components/agent/chat/index.tsx @@ -1,6 +1,9 @@ -import { Suspense, lazy, useEffect } from "react"; +import { Suspense, lazy, useEffect, useRef } from "react"; import type { AgentChatWorkspaceProps } from "./agentChatWorkspaceContract"; -import { loadAgentChatWorkspaceModule } from "./agentChatWorkspaceLoader"; +import { + loadAgentChatWorkspaceModule, + preloadAgentChatWorkspaceModule, +} from "./agentChatWorkspaceLoader"; const WORKSPACE_LOADING_FALLBACK = (
@@ -9,10 +12,17 @@ const WORKSPACE_LOADING_FALLBACK = ( ); const LazyAgentChatWorkspace = lazy(async () => { + const t0 = performance.now(); const module = await loadAgentChatWorkspaceModule(); + console.info( + `[PERF] AgentChatWorkspace module loaded: ${(performance.now() - t0).toFixed(0)}ms`, + ); return { default: module.AgentChatWorkspace }; }); +// 在模块加载时立即预热,避免首次进入聊天页时才触发动态 import +preloadAgentChatWorkspaceModule(); + export type { AgentChatWorkspaceProps, WorkflowProgressSnapshot, @@ -33,6 +43,14 @@ export function AgentChatPage(props: AgentChatWorkspaceProps) { onWorkflowProgressChange, } = props; + // 性能埋点:记录路由进入时间 + const mountT0 = useRef(performance.now()); + useEffect(() => { + console.info( + `[PERF] AgentChatPage mounted: ${(performance.now() - mountT0.current).toFixed(0)}ms`, + ); + }, []); + const hasDirectWorkspaceIntent = Boolean(initialUserPrompt?.trim()) || Boolean(initialUserImages?.length) || @@ -48,6 +66,11 @@ export function AgentChatPage(props: AgentChatWorkspaceProps) { ? true : props.showChatPanel; + // 用首次渲染时的时间戳作为强制重挂载的 key,避免复用旧工作区实例导致旧状态闪烁 + const forcedMountKey = useRef( + shouldForceClawWorkspace ? Date.now() : null, + ); + useEffect(() => { if (!shouldForceClawWorkspace) { return; @@ -67,6 +90,7 @@ export function AgentChatPage(props: AgentChatWorkspaceProps) { diff --git a/src/components/agent/chat/team-workspace-runtime/liveRuntimeProjector.ts b/src/components/agent/chat/team-workspace-runtime/liveRuntimeProjector.ts index 8cacca6ad..cedcdbd11 100644 --- a/src/components/agent/chat/team-workspace-runtime/liveRuntimeProjector.ts +++ b/src/components/agent/chat/team-workspace-runtime/liveRuntimeProjector.ts @@ -98,7 +98,31 @@ function normalizeLiveActivityText( } function appendLiveActivityDraft(previous: string | undefined, chunk: string) { - return normalizeLiveActivityText(`${previous ?? ""}${chunk}`) ?? undefined; + const base = previous ?? ""; + if (!base) { + return normalizeLiveActivityText(chunk) ?? undefined; + } + + if (!chunk) { + return normalizeLiveActivityText(base) ?? undefined; + } + + if (chunk.startsWith(base)) { + return normalizeLiveActivityText(chunk) ?? undefined; + } + + if (base.endsWith(chunk)) { + return normalizeLiveActivityText(base) ?? undefined; + } + + const maxOverlap = Math.min(base.length, chunk.length); + for (let overlap = maxOverlap; overlap > 0; overlap -= 1) { + if (base.slice(-overlap) === chunk.slice(0, overlap)) { + return normalizeLiveActivityText(`${base}${chunk.slice(overlap)}`) ?? undefined; + } + } + + return normalizeLiveActivityText(`${base}${chunk}`) ?? undefined; } function buildActivityEntry(params: { diff --git a/src/components/agent/chat/types.ts b/src/components/agent/chat/types.ts index 15e0e43aa..cb645029b 100644 --- a/src/components/agent/chat/types.ts +++ b/src/components/agent/chat/types.ts @@ -394,14 +394,6 @@ export const PROVIDER_CONFIG: Record< "claude-haiku-4-5-20251001", ], }, - kiro: { - label: "Kiro", - models: [ - "claude-opus-4-5-20251101", - "claude-sonnet-4-5-20250929", - "claude-sonnet-4-20250514", - ], - }, openai: { label: "OpenAI", models: [ @@ -435,27 +427,6 @@ export const PROVIDER_CONFIG: Record< label: "Codex", models: [], // 从后端别名配置动态加载 }, - claude_oauth: { - label: "Claude OAuth", - models: [ - "claude-opus-4-5-20251101", - "claude-sonnet-4-5-20250929", - "claude-sonnet-4-20250514", - ], - }, - antigravity: { - label: "Antigravity", - models: [ - "gemini-3-pro-preview", - "gemini-3-pro-image-preview", - "gemini-3-flash-preview", - "gemini-2.5-flash", - "gemini-2.5-computer-use-preview-10-2025", - "gemini-claude-sonnet-4-5", - "gemini-claude-sonnet-4-5-thinking", - "gemini-claude-opus-4-5-thinking", - ], - }, submodel: { label: "Submodel", models: [ diff --git a/src/components/input-kit/ModelSelector.test.tsx b/src/components/input-kit/ModelSelector.test.tsx index 1430c5fc1..1140b83cc 100644 --- a/src/components/input-kit/ModelSelector.test.tsx +++ b/src/components/input-kit/ModelSelector.test.tsx @@ -311,6 +311,53 @@ describe("ModelSelector", () => { expect(pageText).toContain("无多模态"); }); + it("展开后应把 Lime 云端模型与本地供应商分组显示", () => { + mockUseConfiguredProviders.mockReturnValue({ + providers: [ + { + key: "lime-hub", + label: "Lime 云端", + registryId: "lime-hub", + type: "openai", + providerId: "lime-hub", + apiHost: "https://llm.limeai.run", + }, + { + key: "custom-codex", + label: "Codex Custom", + registryId: "custom-codex", + fallbackRegistryId: "codex", + type: "codex", + providerId: "custom-codex", + apiHost: "https://api.openai.com/v1", + }, + ], + loading: false, + }); + + const { container } = renderModelSelector({ + providerType: "lime-hub", + model: "gpt-5.5", + }); + + const trigger = container.querySelector( + 'button[role="combobox"]', + ) as HTMLButtonElement | null; + if (!trigger) { + throw new Error("未找到模型选择触发器"); + } + + act(() => { + trigger.click(); + }); + + const pageText = document.body.textContent || ""; + expect(pageText).toContain("云端模型"); + expect(pageText).toContain("本地与自定义"); + expect(pageText).toContain("Lime 云端"); + expect(pageText).toContain("Codex Custom"); + }); + it("未知 anthropic-compatible Provider 应在选择器中展示显式缓存提示", () => { mockUseConfiguredProviders.mockReturnValue({ providers: [ diff --git a/src/components/input-kit/ModelSelector.tsx b/src/components/input-kit/ModelSelector.tsx index c3d44489c..abe0ec261 100644 --- a/src/components/input-kit/ModelSelector.tsx +++ b/src/components/input-kit/ModelSelector.tsx @@ -180,6 +180,22 @@ export const ModelSelector: React.FC = ({ : configuredProviders, [configuredProviders, providerFilter], ); + const cloudProviders = useMemo( + () => + visibleProviders.filter( + (provider) => + provider.key === "lime-hub" || provider.providerId === "lime-hub", + ), + [visibleProviders], + ); + const localProviders = useMemo( + () => + visibleProviders.filter( + (provider) => + provider.key !== "lime-hub" && provider.providerId !== "lime-hub", + ), + [visibleProviders], + ); const selectedProvider = useMemo(() => { return findConfiguredProviderBySelection(configuredProviders, providerType); }, [configuredProviders, providerType]); @@ -596,48 +612,67 @@ export const ModelSelector: React.FC = ({ 当前已选供应商暂不可用
) : null} - {visibleProviders.map((provider) => { - const isSelected = selectedProvider?.key === provider.key; - const providerPromptCacheMode = getProviderPromptCacheMode( - provider.type, - provider.promptCacheMode, - provider.apiHost, - ); + {[ + { title: "云端模型", providers: cloudProviders }, + { title: "本地与自定义", providers: localProviders }, + ] + .filter((section) => section.providers.length > 0) + .map((section) => ( +
+
+ {section.title} +
+ {section.providers.map((provider) => { + const isSelected = + selectedProvider?.key === provider.key; + const providerPromptCacheMode = + getProviderPromptCacheMode( + provider.type, + provider.promptCacheMode, + provider.apiHost, + ); - return ( - - ); - })} + return ( + + ); + })} +
+ ))} )}
diff --git a/src/components/provider-pool/AddCredentialModal.tsx b/src/components/provider-pool/AddCredentialModal.tsx deleted file mode 100644 index 6a67eafa3..000000000 --- a/src/components/provider-pool/AddCredentialModal.tsx +++ /dev/null @@ -1,533 +0,0 @@ -/** - * 添加凭证模态框 - * 根据 Provider 类型显示不同的表单 - */ - -import { useState } from "react"; -import { Key, FolderOpen } from "lucide-react"; -import { open } from "@tauri-apps/plugin-dialog"; -import { Modal } from "@/components/Modal"; -import { providerPoolApi, PoolProviderType } from "@/lib/api/providerPool"; -import { AntigravityForm } from "./credential-forms/AntigravityForm"; -import { CodexForm } from "./credential-forms/CodexForm"; -import { ClaudeOAuthForm } from "./credential-forms/ClaudeOAuthForm"; -import { GeminiForm } from "./credential-forms/GeminiForm"; -import { KiroForm } from "./credential-forms/KiroForm"; -import { defaultCredsPath, providerLabels } from "./credential-forms/types"; - -interface AddCredentialModalProps { - providerType: PoolProviderType; - onClose: () => void; - onSuccess: () => void; -} - -export function AddCredentialModal({ - providerType, - onClose, - onSuccess, -}: AddCredentialModalProps) { - const [name, setName] = useState(""); - const [loading, setLoading] = useState(false); - const [error, setError] = useState(null); - - // OAuth 字段 - const [credsFilePath, setCredsFilePath] = useState( - defaultCredsPath[providerType] || "", - ); - const [projectId, setProjectId] = useState(""); - const [apiBaseUrl, setApiBaseUrl] = useState(""); - - // API Key 字段 - const [apiKey, setApiKey] = useState(""); - const [baseUrl, setBaseUrl] = useState(""); - const [excludedModels, setExcludedModels] = useState([]); - - // 判断是否为 OAuth 类型(不包括有特殊表单的 antigravity、codex、claude_oauth、qwen、iflow、gemini、kiro) - const isSimpleOAuth: string[] = []; // Kiro 现在有自己的表单 - const isApiKey = ["openai", "claude", "gemini_api_key"].includes( - providerType, - ); - const primaryActionButtonClassName = - "rounded-lg border border-emerald-200 bg-[linear-gradient(135deg,#0ea5e9_0%,#14b8a6_52%,#10b981_100%)] px-4 py-2 text-sm text-white shadow-sm shadow-emerald-950/15 hover:opacity-95 disabled:opacity-50"; - - const handleSelectFile = async () => { - try { - const selected = await open({ - multiple: false, - filters: [{ name: "JSON", extensions: ["json"] }], - }); - if (selected) { - setCredsFilePath(selected as string); - } - } catch (e) { - console.error("Failed to open file dialog:", e); - } - }; - - // Antigravity 表单 - const antigravityForm = AntigravityForm({ - name, - credsFilePath, - setCredsFilePath, - projectId, - setProjectId, - onSelectFile: handleSelectFile, - loading, - setLoading, - setError, - onSuccess, - }); - - // Codex 表单 - const codexForm = CodexForm({ - name, - credsFilePath, - setCredsFilePath, - apiBaseUrl, - setApiBaseUrl, - onSelectFile: handleSelectFile, - loading, - setLoading, - setError, - onSuccess, - }); - - // Claude OAuth 表单 - const claudeOAuthForm = ClaudeOAuthForm({ - name, - credsFilePath, - setCredsFilePath, - onSelectFile: handleSelectFile, - loading, - setLoading, - setError, - onSuccess, - }); - - // Gemini 表单 - const geminiForm = GeminiForm({ - name, - credsFilePath, - setCredsFilePath, - projectId, - setProjectId, - onSelectFile: handleSelectFile, - loading, - setLoading, - setError, - onSuccess, - }); - - // Kiro 表单 - const kiroForm = KiroForm({ - name, - credsFilePath, - setCredsFilePath, - onSelectFile: handleSelectFile, - loading, - setLoading, - setError, - onSuccess, - }); - - // 简单 OAuth 和 API Key 的提交处理 - const handleSubmit = async () => { - setLoading(true); - setError(null); - - try { - const trimmedName = name.trim() || undefined; - - if (isSimpleOAuth.includes(providerType)) { - if (!credsFilePath) { - setError("请选择凭证文件"); - setLoading(false); - return; - } - - switch (providerType) { - case "gemini": - await providerPoolApi.addGeminiOAuth( - credsFilePath, - projectId.trim() || undefined, - trimmedName, - ); - break; - } - } else if (isApiKey) { - if (!apiKey) { - setError("请输入 API Key"); - setLoading(false); - return; - } - - switch (providerType) { - case "openai": - await providerPoolApi.addOpenAIKey( - apiKey, - baseUrl.trim() || undefined, - trimmedName, - ); - break; - case "claude": - await providerPoolApi.addClaudeKey( - apiKey, - baseUrl.trim() || undefined, - trimmedName, - ); - break; - case "gemini_api_key": - await providerPoolApi.addGeminiApiKey( - apiKey, - baseUrl.trim() || undefined, - excludedModels.length > 0 ? excludedModels : undefined, - trimmedName, - ); - break; - } - } - - onSuccess(); - } catch (e) { - setError(e instanceof Error ? e.message : String(e)); - } finally { - setLoading(false); - } - }; - - // 渲染简单 OAuth 表单 - const renderSimpleOAuthForm = () => ( - <> -
- -
- setCredsFilePath(e.target.value)} - placeholder="输入凭证文件的完整路径..." - className="flex-1 rounded-lg border bg-background px-3 py-2 text-sm" - /> - -
-

- {providerType === "gemini" && "默认路径: ~/.gemini/oauth_creds.json"} -

-
- - {providerType === "gemini" && ( -
- - setProjectId(e.target.value)} - placeholder="Google Cloud Project ID..." - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -
- )} - - ); - - // 渲染 API Key 表单 - const renderApiKeyForm = () => ( - <> -
- -
- - setApiKey(e.target.value)} - placeholder="sk-..." - className="w-full rounded-lg border bg-background pl-10 pr-3 py-2 text-sm" - /> -
-
- -
- - setBaseUrl(e.target.value)} - placeholder={ - providerType === "openai" - ? "https://api.openai.com/v1" - : providerType === "claude" - ? "https://api.anthropic.com/v1" - : "https://generativelanguage.googleapis.com" - } - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -

- 留空使用默认 URL,或输入自定义代理地址 -

-
- - {providerType === "gemini_api_key" && ( -
- - - setExcludedModels( - e.target.value - .split(",") - .map((s) => s.trim()) - .filter((s) => s), - ) - } - placeholder="gemini-1.5-pro, gemini-1.5-flash..." - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -

- 用逗号分隔多个模型名称,支持通配符(如 *-preview) -

-
- )} - - ); - - // 渲染底部按钮 - const renderFooterButton = () => { - // Antigravity 登录模式 - if (providerType === "antigravity" && antigravityForm.mode === "login") { - if (!antigravityForm.authUrl) { - return ( - - ); - } - return null; - } - - // Antigravity 文件模式 - if (providerType === "antigravity" && antigravityForm.mode === "file") { - return ( - - ); - } - - // Codex 登录模式 - if (providerType === "codex" && codexForm.mode === "login") { - if (!codexForm.authUrl) { - return ( - - ); - } - return null; - } - - // Codex 文件模式 - if (providerType === "codex" && codexForm.mode === "file") { - return ( - - ); - } - - // Claude OAuth Cookie 模式 - if (providerType === "claude_oauth" && claudeOAuthForm.mode === "cookie") { - return ( - - ); - } - - // Claude OAuth 登录模式 - if (providerType === "claude_oauth" && claudeOAuthForm.mode === "login") { - if (!claudeOAuthForm.authUrl) { - return ( - - ); - } - return null; - } - - // Claude OAuth 文件模式 - if (providerType === "claude_oauth" && claudeOAuthForm.mode === "file") { - return ( - - ); - } - - // Gemini 登录模式 - if (providerType === "gemini" && geminiForm.mode === "login") { - if (!geminiForm.authUrl) { - return ( - - ); - } - return null; - } - - // Gemini 文件模式 - if (providerType === "gemini" && geminiForm.mode === "file") { - return ( - - ); - } - - // Kiro JSON 模式 - if (providerType === "kiro" && kiroForm.mode === "json") { - return ( - - ); - } - - // Kiro 文件模式 - if (providerType === "kiro" && kiroForm.mode === "file") { - return ( - - ); - } - - // 其他类型 - return ( - - ); - }; - - return ( - - {/* Header */} -
-

- 添加 {providerLabels[providerType]} 凭证 -

-
- - {/* Content */} -
- {/* 名称字段 */} -
- - setName(e.target.value)} - placeholder="给这个凭证起个名字..." - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -
- - {/* 根据类型渲染不同表单 */} - {providerType === "antigravity" && antigravityForm.render()} - {providerType === "codex" && codexForm.render()} - {providerType === "claude_oauth" && claudeOAuthForm.render()} - {providerType === "gemini" && geminiForm.render()} - {providerType === "kiro" && kiroForm.render()} - {isSimpleOAuth.includes(providerType) && renderSimpleOAuthForm()} - {isApiKey && renderApiKeyForm()} - - {/* 错误提示 */} - {error && ( -
- {error} -
- )} -
- - {/* Footer */} -
- - {renderFooterButton()} -
-
- ); -} diff --git a/src/components/provider-pool/AmpConfigSection.tsx b/src/components/provider-pool/AmpConfigSection.tsx deleted file mode 100644 index 4a4fb3479..000000000 --- a/src/components/provider-pool/AmpConfigSection.tsx +++ /dev/null @@ -1,245 +0,0 @@ -import { useState } from "react"; -import { - Plus, - Trash2, - Globe, - ArrowRight, - Terminal, - CheckCircle2, - AlertTriangle, -} from "lucide-react"; -import type { AmpConfig, AmpModelMapping } from "@/lib/api/providerRuntime"; - -interface AmpConfigSectionProps { - config: AmpConfig; - onChange: (config: AmpConfig) => void; - onSave?: () => Promise; -} - -export function AmpConfigSection({ - config, - onChange, - onSave, -}: AmpConfigSectionProps) { - const primaryActionButtonClassName = - "rounded-lg border border-emerald-200 bg-[linear-gradient(135deg,#0ea5e9_0%,#14b8a6_52%,#10b981_100%)] text-sm text-white shadow-sm shadow-emerald-950/15 hover:opacity-95 disabled:opacity-50"; - const [saving, setSaving] = useState(false); - const [message, setMessage] = useState<{ - type: "success" | "error"; - text: string; - } | null>(null); - const [editingMapping, setEditingMapping] = useState(false); - const [mappingFrom, setMappingFrom] = useState(""); - const [mappingTo, setMappingTo] = useState(""); - - // Ensure model_mappings is always an array - const modelMappings = config?.model_mappings ?? []; - - const updateConfig = (updates: Partial) => { - onChange({ ...config, ...updates }); - }; - - const addMapping = () => { - if (!mappingFrom.trim() || !mappingTo.trim()) return; - const newMapping: AmpModelMapping = { - from: mappingFrom.trim(), - to: mappingTo.trim(), - }; - updateConfig({ - model_mappings: [...modelMappings, newMapping], - }); - setMappingFrom(""); - setMappingTo(""); - }; - - const removeMapping = (from: string) => { - updateConfig({ - model_mappings: modelMappings.filter((m) => m.from !== from), - }); - }; - - const handleSave = async () => { - if (!onSave) return; - setSaving(true); - setMessage(null); - try { - await onSave(); - setMessage({ type: "success", text: "Amp CLI 配置已保存" }); - setTimeout(() => setMessage(null), 3000); - } catch (e: unknown) { - const errorMessage = e instanceof Error ? e.message : String(e); - setMessage({ type: "error", text: `保存失败: ${errorMessage}` }); - } - setSaving(false); - }; - - return ( -
-
- -
-

Amp CLI 集成

-

- 配置 Amp CLI 的路由和模型映射 -

-
-
- - {/* 消息提示 */} - {message && ( -
- {message.type === "success" ? ( - - ) : ( - - )} - {message.text} -
- )} - -
- {/* Upstream URL */} -
- - - updateConfig({ upstream_url: e.target.value || null }) - } - placeholder="https://ampcode.com" - className="w-full px-3 py-2 rounded-lg border bg-background text-sm focus:ring-2 focus:ring-primary/20 focus:border-primary outline-none" - /> -

- Amp CLI 管理端点的上游服务器地址 -

-
- - {/* Restrict Management to Localhost */} - - - {/* Model Mappings */} -
- -

- 将不可用的模型请求映射到可用的替代模型 -

- - {/* Existing Mappings */} - {modelMappings.length > 0 && ( -
- {modelMappings.map((mapping) => ( -
- - {mapping.from} - - - - {mapping.to} - - -
- ))} -
- )} - - {/* Add New Mapping */} - {editingMapping ? ( -
-
- setMappingFrom(e.target.value)} - placeholder="源模型 (如 claude-opus-4.5)" - className="flex-1 px-3 py-1.5 rounded border bg-background text-sm" - /> - - setMappingTo(e.target.value)} - placeholder="目标模型 (如 claude-sonnet-4)" - className="flex-1 px-3 py-1.5 rounded border bg-background text-sm" - /> -
-
- - -
-
- ) : ( - - )} -
- - {/* Save Button */} - {onSave && ( - - )} -
-
- ); -} diff --git a/src/components/provider-pool/CodexSection.tsx b/src/components/provider-pool/CodexSection.tsx deleted file mode 100644 index c5ef100d1..000000000 --- a/src/components/provider-pool/CodexSection.tsx +++ /dev/null @@ -1,218 +0,0 @@ -import { useState } from "react"; -import { Plus, Trash2, FolderOpen, LogIn, RefreshCw } from "lucide-react"; -import { open } from "@tauri-apps/plugin-dialog"; -import type { CredentialEntry } from "@/lib/api/providerRuntime"; - -interface CodexSectionProps { - entries: CredentialEntry[]; - onChange: (entries: CredentialEntry[]) => void; - onOAuthLogin?: (id: string) => Promise; - onRefreshToken?: (id: string) => Promise; -} - -export function CodexSection({ - entries, - onChange, - onOAuthLogin, - onRefreshToken, -}: CodexSectionProps) { - const [loginLoading, setLoginLoading] = useState(null); - const [refreshLoading, setRefreshLoading] = useState(null); - - const addEntry = () => { - const newEntry: CredentialEntry = { - id: `codex-${Date.now()}`, - token_file: "~/.codex/oauth.json", - disabled: false, - proxy_url: null, - }; - onChange([...entries, newEntry]); - }; - - const updateEntry = (id: string, updates: Partial) => { - onChange(entries.map((e) => (e.id === id ? { ...e, ...updates } : e))); - }; - - const removeEntry = (id: string) => { - onChange(entries.filter((e) => e.id !== id)); - }; - - const handleSelectFile = async (id: string) => { - try { - const selected = await open({ - multiple: false, - filters: [{ name: "JSON", extensions: ["json"] }], - }); - if (selected) { - updateEntry(id, { token_file: selected as string }); - } - } catch (e) { - console.error("Failed to open file dialog:", e); - } - }; - - const handleOAuthLogin = async (id: string) => { - if (!onOAuthLogin) return; - setLoginLoading(id); - try { - await onOAuthLogin(id); - } catch (e) { - console.error("OAuth login failed:", e); - } finally { - setLoginLoading(null); - } - }; - - const handleRefreshToken = async (id: string) => { - if (!onRefreshToken) return; - setRefreshLoading(id); - try { - await onRefreshToken(id); - } catch (e) { - console.error("Token refresh failed:", e); - } finally { - setRefreshLoading(null); - } - }; - - return ( -
-
-
- -
-

OpenAI Codex OAuth

-

- 通过 OAuth 认证使用 OpenAI Codex 服务 -

-
-
- -
- - {entries.length === 0 ? ( -
-

暂无 Codex OAuth 凭证

-

点击上方"添加"按钮添加凭证

-
- ) : ( -
- {entries.map((entry) => ( -
- {/* Header */} -
- - {entry.id} - -
- - -
-
- - {/* Token File Path */} -
- -
- - updateEntry(entry.id, { token_file: e.target.value }) - } - placeholder="~/.codex/oauth.json" - className="flex-1 px-3 py-1.5 rounded border bg-background text-sm" - /> - -
-
- - {/* Proxy URL */} -
- - - updateEntry(entry.id, { proxy_url: e.target.value || null }) - } - placeholder="socks5://127.0.0.1:1080" - className="w-full px-3 py-1.5 rounded border bg-background text-sm" - /> -
- - {/* OAuth Actions */} -
- {onOAuthLogin && ( - - )} - {onRefreshToken && ( - - )} -
-
- ))} -
- )} -
- ); -} diff --git a/src/components/provider-pool/CredentialCard.test.ts b/src/components/provider-pool/CredentialCard.test.ts deleted file mode 100644 index 9652be275..000000000 --- a/src/components/provider-pool/CredentialCard.test.ts +++ /dev/null @@ -1,400 +0,0 @@ -/** - * @file CredentialCard 属性测试 - * @description 测试 OAuth 凭证卡片信息完整性 - * @module components/provider-pool/CredentialCard.test - * - * **Feature: provider-ui-refactor** - * **Property 3: OAuth 凭证卡片信息完整性** - * **Validates: Requirements 2.2** - */ - -import { describe, expect } from "vitest"; -import { test } from "@fast-check/vitest"; -import * as fc from "fast-check"; -import type { - CredentialDisplay, - PoolProviderType, - CredentialSource, -} from "@/lib/api/providerPool"; - -// ============================================================================ -// 辅助函数(用于测试) -// ============================================================================ - -/** - * 提取 OAuth 凭证卡片显示信息 - * 用于属性测试验证 Requirements 2.2 - * - * @param credential OAuth 凭证数据 - * @returns 卡片显示信息 - */ -export function extractOAuthCardDisplayInfo(credential: CredentialDisplay): { - hasHealthStatus: boolean; - hasUsageCount: boolean; - hasActionButtons: boolean; - healthStatus: "healthy" | "unhealthy" | "disabled"; - usageCount: number; - errorCount: number; -} { - // 健康状态:根据 is_healthy 和 is_disabled 判断 - let healthStatus: "healthy" | "unhealthy" | "disabled"; - if (credential.is_disabled) { - healthStatus = "disabled"; - } else if (credential.is_healthy) { - healthStatus = "healthy"; - } else { - healthStatus = "unhealthy"; - } - - return { - // 健康状态始终存在(通过 is_healthy 和 is_disabled 字段) - hasHealthStatus: - typeof credential.is_healthy === "boolean" && - typeof credential.is_disabled === "boolean", - // 使用次数始终存在(通过 usage_count 字段) - hasUsageCount: typeof credential.usage_count === "number", - // 操作按钮始终存在(卡片组件固定渲染) - hasActionButtons: true, - healthStatus, - usageCount: credential.usage_count, - errorCount: credential.error_count, - }; -} - -/** - * 验证 OAuth 凭证卡片是否包含所有必要信息 - * - * @param credential OAuth 凭证数据 - * @returns 是否包含所有必要信息 - */ -export function isOAuthCardComplete(credential: CredentialDisplay): boolean { - const info = extractOAuthCardDisplayInfo(credential); - return info.hasHealthStatus && info.hasUsageCount && info.hasActionButtons; -} - -/** - * 获取 OAuth 凭证的操作按钮列表 - * 根据凭证类型返回应该显示的操作按钮 - * - * @param credential OAuth 凭证数据 - * @returns 操作按钮列表 - */ -export function getOAuthCardActionButtons( - credential: CredentialDisplay, -): string[] { - const buttons: string[] = [ - "toggle", // 启用/禁用 - "edit", // 编辑 - "checkHealth", // 检测健康 - "reset", // 重置 - "delete", // 删除 - ]; - - // OAuth 类型凭证额外显示刷新 Token 按钮 - if (credential.credential_type.includes("oauth")) { - buttons.push("refreshToken"); - } - - return buttons; -} - -// ============================================================================ -// 测试数据生成器 -// ============================================================================ - -/** - * 生成有效的 ISO 日期字符串 - */ -const validDateArbitrary = fc - .integer({ - min: new Date("2020-01-01").getTime(), - max: new Date("2030-12-31").getTime(), - }) - .map((timestamp) => new Date(timestamp).toISOString()); - -/** - * OAuth Provider 类型 - */ -const oauthProviderTypes: PoolProviderType[] = [ - "kiro", - "gemini", - "antigravity", - "codex", - "claude_oauth", -]; - -/** - * OAuth 凭证类型 - */ -const oauthCredentialTypes = [ - "kiro_oauth", - "gemini_oauth", - "antigravity_oauth", - "codex_oauth", - "claude_oauth", -]; - -/** - * 凭证来源类型 - */ -const credentialSourceArbitrary: fc.Arbitrary = - fc.constantFrom("manual", "imported", "private"); - -/** - * 生成随机 OAuth 凭证显示数据 - */ -const oauthCredentialArbitrary: fc.Arbitrary = fc.record({ - uuid: fc.uuid(), - provider_type: fc.constantFrom(...oauthProviderTypes), - credential_type: fc.constantFrom(...oauthCredentialTypes), - name: fc.option(fc.string({ minLength: 1, maxLength: 100 }), { - nil: undefined, - }), - display_credential: fc.string({ minLength: 1, maxLength: 50 }), - is_healthy: fc.boolean(), - is_disabled: fc.boolean(), - check_health: fc.boolean(), - check_model_name: fc.option(fc.string({ minLength: 1, maxLength: 50 }), { - nil: undefined, - }), - not_supported_models: fc.array(fc.string({ minLength: 1, maxLength: 50 }), { - maxLength: 5, - }), - usage_count: fc.nat({ max: 100000 }), - error_count: fc.nat({ max: 10000 }), - last_used: fc.option(validDateArbitrary, { nil: undefined }), - last_error_time: fc.option(validDateArbitrary, { nil: undefined }), - last_error_message: fc.option(fc.string({ minLength: 1, maxLength: 200 }), { - nil: undefined, - }), - last_health_check_time: fc.option(validDateArbitrary, { nil: undefined }), - last_health_check_model: fc.option( - fc.string({ minLength: 1, maxLength: 50 }), - { nil: undefined }, - ), - oauth_status: fc.option( - fc.record({ - has_access_token: fc.boolean(), - has_refresh_token: fc.boolean(), - is_token_valid: fc.boolean(), - expiry_info: fc.option(fc.string({ minLength: 1, maxLength: 100 }), { - nil: undefined, - }), - creds_path: fc.string({ minLength: 1, maxLength: 200 }), - }), - { nil: undefined }, - ), - token_cache_status: fc.option( - fc.record({ - has_cached_token: fc.boolean(), - is_valid: fc.boolean(), - is_expiring_soon: fc.boolean(), - expiry_time: fc.option(validDateArbitrary, { nil: undefined }), - last_refresh: fc.option(validDateArbitrary, { nil: undefined }), - refresh_error_count: fc.nat({ max: 100 }), - last_refresh_error: fc.option( - fc.string({ minLength: 1, maxLength: 200 }), - { nil: undefined }, - ), - }), - { nil: undefined }, - ), - created_at: validDateArbitrary, - updated_at: validDateArbitrary, - source: credentialSourceArbitrary, - base_url: fc.option(fc.webUrl(), { nil: undefined }), - api_key: fc.option(fc.string({ minLength: 1, maxLength: 100 }), { - nil: undefined, - }), - proxy_url: fc.option(fc.webUrl(), { nil: undefined }), -}); - -// ============================================================================ -// Property 3: OAuth 凭证卡片信息完整性 -// ============================================================================ - -describe("Property 3: OAuth 凭证卡片信息完整性", () => { - /** - * Property 3: OAuth 凭证卡片信息完整性 - * - * *对于任意* OAuth 凭证,渲染后的卡片应包含健康状态、使用次数和操作按钮 - * - * **Validates: Requirements 2.2** - */ - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "每个 OAuth 凭证卡片应包含健康状态、使用次数和操作按钮", - (credential: CredentialDisplay) => { - const displayInfo = extractOAuthCardDisplayInfo(credential); - - // 验证健康状态存在 - expect(displayInfo.hasHealthStatus).toBe(true); - - // 验证使用次数存在 - expect(displayInfo.hasUsageCount).toBe(true); - - // 验证操作按钮存在 - expect(displayInfo.hasActionButtons).toBe(true); - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "健康状态应为 healthy、unhealthy 或 disabled 之一", - (credential: CredentialDisplay) => { - const displayInfo = extractOAuthCardDisplayInfo(credential); - - expect(["healthy", "unhealthy", "disabled"]).toContain( - displayInfo.healthStatus, - ); - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "使用次数应为非负整数", - (credential: CredentialDisplay) => { - const displayInfo = extractOAuthCardDisplayInfo(credential); - - expect(Number.isInteger(displayInfo.usageCount)).toBe(true); - expect(displayInfo.usageCount).toBeGreaterThanOrEqual(0); - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "错误次数应为非负整数", - (credential: CredentialDisplay) => { - const displayInfo = extractOAuthCardDisplayInfo(credential); - - expect(Number.isInteger(displayInfo.errorCount)).toBe(true); - expect(displayInfo.errorCount).toBeGreaterThanOrEqual(0); - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "OAuth 凭证应包含刷新 Token 操作按钮", - (credential: CredentialDisplay) => { - const buttons = getOAuthCardActionButtons(credential); - - // OAuth 凭证应该有刷新 Token 按钮 - expect(buttons).toContain("refreshToken"); - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "所有凭证应包含基本操作按钮", - (credential: CredentialDisplay) => { - const buttons = getOAuthCardActionButtons(credential); - - // 所有凭证都应该有这些基本按钮 - expect(buttons).toContain("toggle"); - expect(buttons).toContain("edit"); - expect(buttons).toContain("checkHealth"); - expect(buttons).toContain("reset"); - expect(buttons).toContain("delete"); - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "禁用状态应正确反映在健康状态中", - (credential: CredentialDisplay) => { - const displayInfo = extractOAuthCardDisplayInfo(credential); - - if (credential.is_disabled) { - expect(displayInfo.healthStatus).toBe("disabled"); - } - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "健康凭证(未禁用)应显示为 healthy", - (credential: CredentialDisplay) => { - const displayInfo = extractOAuthCardDisplayInfo(credential); - - if (!credential.is_disabled && credential.is_healthy) { - expect(displayInfo.healthStatus).toBe("healthy"); - } - }, - ); - - test.prop([oauthCredentialArbitrary], { numRuns: 100 })( - "不健康凭证(未禁用)应显示为 unhealthy", - (credential: CredentialDisplay) => { - const displayInfo = extractOAuthCardDisplayInfo(credential); - - if (!credential.is_disabled && !credential.is_healthy) { - expect(displayInfo.healthStatus).toBe("unhealthy"); - } - }, - ); -}); - -// ============================================================================ -// 边界情况测试 -// ============================================================================ - -describe("OAuth 凭证卡片边界情况", () => { - test("使用次数为 0 的凭证应正确显示", () => { - const credential: CredentialDisplay = { - uuid: "test-uuid", - provider_type: "kiro", - credential_type: "kiro_oauth", - display_credential: "test@example.com", - is_healthy: true, - is_disabled: false, - check_health: true, - not_supported_models: [], - usage_count: 0, - error_count: 0, - created_at: new Date().toISOString(), - updated_at: new Date().toISOString(), - source: "manual", - }; - - const displayInfo = extractOAuthCardDisplayInfo(credential); - expect(displayInfo.usageCount).toBe(0); - expect(displayInfo.hasUsageCount).toBe(true); - }); - - test("高使用次数的凭证应正确显示", () => { - const credential: CredentialDisplay = { - uuid: "test-uuid", - provider_type: "gemini", - credential_type: "gemini_oauth", - display_credential: "test@example.com", - is_healthy: true, - is_disabled: false, - check_health: true, - not_supported_models: [], - usage_count: 999999, - error_count: 100, - created_at: new Date().toISOString(), - updated_at: new Date().toISOString(), - source: "imported", - }; - - const displayInfo = extractOAuthCardDisplayInfo(credential); - expect(displayInfo.usageCount).toBe(999999); - expect(displayInfo.errorCount).toBe(100); - }); - - test("完整的 OAuth 凭证应通过完整性检查", () => { - const credential: CredentialDisplay = { - uuid: "test-uuid", - provider_type: "kiro", - credential_type: "kiro_oauth", - name: "Test Credential", - display_credential: "test@example.com", - is_healthy: true, - is_disabled: false, - check_health: true, - not_supported_models: [], - usage_count: 100, - error_count: 5, - last_used: new Date().toISOString(), - last_health_check_time: new Date().toISOString(), - created_at: new Date().toISOString(), - updated_at: new Date().toISOString(), - source: "manual", - }; - - expect(isOAuthCardComplete(credential)).toBe(true); - }); -}); diff --git a/src/components/provider-pool/CredentialCard.tsx b/src/components/provider-pool/CredentialCard.tsx deleted file mode 100644 index 793199568..000000000 --- a/src/components/provider-pool/CredentialCard.tsx +++ /dev/null @@ -1,1011 +0,0 @@ -import { useState } from "react"; -import { - Heart, - HeartOff, - Trash2, - RotateCcw, - Activity, - Power, - PowerOff, - Clock, - AlertTriangle, - RefreshCw, - Settings, - Upload, - Lock, - User, - Globe, - BarChart3, - ChevronUp, - Fingerprint, - Copy, - Check, - Timer, - MonitorDown, -} from "lucide-react"; -import type { - CredentialDisplay, - CredentialSource, -} from "@/lib/api/providerPool"; -import { - getKiroCredentialFingerprint, - switchKiroToLocal, - type KiroFingerprintInfo, - type SwitchToLocalResult, - kiroCredentialApi, -} from "@/lib/api/providerPool"; -import { usageApi, type UsageInfo } from "@/lib/api/usage"; -import { UsageDisplay } from "./UsageDisplay"; - -interface CredentialCardProps { - credential: CredentialDisplay; - onToggle: () => void; - onDelete: () => void; - onReset: () => void; - onCheckHealth: () => void; - onRefreshToken?: () => void; - onEdit: () => void; - deleting: boolean; - checkingHealth: boolean; - refreshingToken?: boolean; - /** 是否为 Kiro 凭证(支持用量查询) */ - isKiroCredential?: boolean; - /** 是否为当前本地使用的凭证 */ - isLocalActive?: boolean; - /** 切换到本地成功后的回调 */ - onSwitchToLocal?: () => void; -} - -export function CredentialCard({ - credential, - onToggle, - onDelete, - onReset, - onCheckHealth, - onRefreshToken, - onEdit, - deleting, - checkingHealth, - refreshingToken, - isKiroCredential, - isLocalActive, - onSwitchToLocal, -}: CredentialCardProps) { - // 用量查询状态 - const [usageExpanded, setUsageExpanded] = useState(false); - const [usageLoading, setUsageLoading] = useState(false); - const [usageInfo, setUsageInfo] = useState(null); - const [usageError, setUsageError] = useState(null); - - // 指纹信息状态(仅 Kiro 凭证) - const [fingerprintInfo, setFingerprintInfo] = - useState(null); - const [fingerprintLoading, setFingerprintLoading] = useState(false); - const [fingerprintExpanded, setFingerprintExpanded] = useState(false); - const [fingerprintCopied, setFingerprintCopied] = useState(false); - - // Kiro 增强状态管理 - const [kiroHealthScore, setKiroHealthScore] = useState(null); - const [kiroStatusLoading, setKiroStatusLoading] = useState(false); - const [kiroRefreshing, setKiroRefreshing] = useState(false); - const [kiroStatusExpanded, setKiroStatusExpanded] = useState(false); - - // 切换到本地状态 - const [switchingToLocal, setSwitchingToLocal] = useState(false); - const [switchResult, setSwitchResult] = useState( - null, - ); - - // 查询指纹信息 - const handleCheckFingerprint = async () => { - if (fingerprintExpanded && fingerprintInfo) { - // 已展开且有数据,直接折叠 - setFingerprintExpanded(false); - return; - } - - setFingerprintExpanded(true); - setFingerprintLoading(true); - - try { - const info = await getKiroCredentialFingerprint(credential.uuid); - setFingerprintInfo(info); - } catch (e) { - console.error("获取指纹信息失败:", e); - } finally { - setFingerprintLoading(false); - } - }; - - // 复制 Machine ID - const handleCopyMachineId = async () => { - if (!fingerprintInfo) return; - try { - await navigator.clipboard.writeText(fingerprintInfo.machine_id); - setFingerprintCopied(true); - setTimeout(() => setFingerprintCopied(false), 2000); - } catch (e) { - console.error("复制失败:", e); - } - }; - - // 查询用量 - const handleCheckUsage = async () => { - if (usageExpanded && usageInfo) { - // 已展开且有数据,直接折叠 - setUsageExpanded(false); - return; - } - - setUsageExpanded(true); - setUsageLoading(true); - setUsageError(null); - - try { - const info = await usageApi.getKiroUsage(credential.uuid); - setUsageInfo(info); - } catch (e) { - setUsageError(e instanceof Error ? e.message : String(e)); - } finally { - setUsageLoading(false); - } - }; - - // 获取 Kiro 详细状态 - const handleCheckKiroStatus = async () => { - if (kiroStatusExpanded) { - setKiroStatusExpanded(false); - return; - } - - setKiroStatusExpanded(true); - setKiroStatusLoading(true); - - try { - const status = await kiroCredentialApi.getCredentialStatus( - credential.uuid, - ); - setKiroHealthScore(status.health_score || 0); - } catch (e) { - console.error("获取 Kiro 状态失败:", e); - } finally { - setKiroStatusLoading(false); - } - }; - - // 快速刷新 Kiro Token - const handleQuickRefresh = async () => { - setKiroRefreshing(true); - - try { - const result = await kiroCredentialApi.refreshCredential(credential.uuid); - if (result.success) { - // 刷新成功,可以显示成功消息 - console.log("Token 刷新成功:", result.message); - // 可以触发页面数据刷新 - if (onRefreshToken) { - onRefreshToken(); - } - } else { - console.error("Token 刷新失败:", result.error || result.message); - } - } catch (e) { - console.error("Token 刷新异常:", e); - } finally { - setKiroRefreshing(false); - } - }; - - // 切换到本地 - const handleSwitchToLocal = async () => { - setSwitchingToLocal(true); - setSwitchResult(null); - - try { - const result = await switchKiroToLocal(credential.uuid); - setSwitchResult(result); - - if (result.success) { - console.log("切换到本地成功:", result.message); - // 调用回调通知父组件刷新本地活跃凭证 - if (onSwitchToLocal) { - onSwitchToLocal(); - } - } else { - console.error("切换到本地失败:", result.message); - } - - // 3秒后自动清除结果提示 - setTimeout(() => { - setSwitchResult(null); - }, 5000); - } catch (e) { - console.error("切换到本地异常:", e); - setSwitchResult({ - success: false, - message: e instanceof Error ? e.message : String(e), - requires_action: false, - requires_kiro_restart: false, - }); - } finally { - setSwitchingToLocal(false); - } - }; - - const formatDate = (dateStr?: string) => { - if (!dateStr) return "从未"; - const date = new Date(dateStr); - return date.toLocaleString("zh-CN", { - month: "2-digit", - day: "2-digit", - hour: "2-digit", - minute: "2-digit", - }); - }; - - const getCredentialTypeLabel = (type: string) => { - const labels: Record = { - kiro_oauth: "OAuth", - gemini_oauth: "OAuth", - qwen_oauth: "OAuth", - antigravity_oauth: "OAuth", - openai_key: "API Key", - claude_key: "API Key", - codex_oauth: "OAuth", - claude_oauth: "OAuth", - iflow_oauth: "OAuth", - iflow_cookie: "Cookie", - }; - return labels[type] || type; - }; - - const getSourceLabel = (source: CredentialSource) => { - const labels: Record< - CredentialSource, - { text: string; icon: typeof User; color: string } - > = { - manual: { - text: "手动添加", - icon: User, - color: "bg-sky-100 text-sky-700", - }, - imported: { - text: "导入", - icon: Upload, - color: "bg-emerald-100 text-emerald-700", - }, - private: { - text: "私有", - icon: Lock, - color: "bg-amber-100 text-amber-700", - }, - }; - return labels[source] || labels.manual; - }; - - const sourceInfo = getSourceLabel(credential.source || "manual"); - const SourceIcon = sourceInfo.icon; - - const isHealthy = credential.is_healthy && !credential.is_disabled; - const hasError = credential.error_count > 0; - const isOAuth = credential.credential_type.includes("oauth"); - - return ( -
- {/* 第一行:状态图标 + 名称 + 标签 + 操作按钮 */} -
- {/* Status Icon */} -
- {credential.is_disabled ? ( - - ) : isHealthy ? ( - - ) : ( - - )} -
- - {/* Main Info */} -
-

- {credential.name || `凭证 #${credential.uuid.slice(0, 8)}`} -

-
- - {getCredentialTypeLabel(credential.credential_type)} - - - - {sourceInfo.text} - - {credential.proxy_url && ( - - - 代理 - - )} -
-
- - {/* Actions */} -
- - - - - - - {isOAuth && onRefreshToken && ( - - )} - - {/* 指纹信息按钮 - 仅 Kiro 凭证显示 */} - {isKiroCredential && ( - - )} - - {/* 用量查询按钮 - 仅 Kiro 凭证显示 */} - {isKiroCredential && ( - - )} - - {/* Kiro 详细状态按钮 - 仅 Kiro 凭证显示 */} - {isKiroCredential && ( - - )} - - {/* Kiro 快速刷新按钮 - 仅 Kiro 凭证显示 */} - {isKiroCredential && ( - - )} - - {/* Kiro 切换到本地按钮 - 仅 Kiro 凭证显示 */} - {isKiroCredential && ( - - )} - - - - -
-
- - {/* 第二行:统计信息 - 使用网格布局 */} -
-
- {/* 使用次数 */} -
- -
-
使用次数
-
- {credential.usage_count} -
-
-
- - {/* 错误次数 */} -
- -
-
错误次数
-
- {credential.error_count} -
-
-
- - {/* 最后使用 */} -
- -
-
最后使用
-
- {formatDate(credential.last_used)} -
-
-
- - {/* Token 有效期 - OAuth 凭证显示 */} - {isOAuth ? ( -
- -
-
- Token 有效期 -
- {credential.token_cache_status?.expiry_time ? ( -
- {formatDate(credential.token_cache_status.expiry_time)} -
- ) : ( -
--
- )} -
-
- ) : ( -
/* 占位 */ - )} - - {/* 健康检查/健康分数 */} - {isKiroCredential && kiroHealthScore !== null ? ( - // 为 Kiro 凭证显示健康分数 -
-
= 80 - ? "bg-emerald-500" - : kiroHealthScore >= 60 - ? "bg-amber-500" - : kiroHealthScore >= 40 - ? "bg-orange-500" - : "bg-red-500" - }`} - > - ★ -
-
-
健康分数
-
= 80 - ? "text-emerald-600" - : kiroHealthScore >= 60 - ? "text-amber-600" - : kiroHealthScore >= 40 - ? "text-orange-600" - : "text-red-600" - }`} - > - {Math.round(kiroHealthScore)} -
-
-
- ) : credential.last_health_check_time ? ( - // 为其他凭证显示健康检查时间 -
- -
-
健康检查
-
- {formatDate(credential.last_health_check_time)} -
-
-
- ) : ( -
/* 占位 */ - )} -
-
- - {/* 第三行:UUID */} -
-

- {credential.uuid} -

-
- - {/* Mobile Stats - shown on small screens */} -
-
-
- - 使用: - {credential.usage_count} -
-
- - 错误: - {credential.error_count} -
-
- - 最后使用: - {formatDate(credential.last_used)} -
-
-
- - {/* Error Message */} - {credential.last_error_message && ( -
-
- {credential.last_error_message.slice(0, 150)} - {credential.last_error_message.length > 150 && "..."} -
- {/* 重新授权提示 */} - {(credential.last_error_message.includes("invalid_grant") || - credential.last_error_message.includes("重新授权") || - credential.last_error_message.includes("凭证已过期")) && ( -
-
- - 💡 需要重新授权 - - {onRefreshToken && ( - - )} -
-

- 请删除此凭证并重新添加,或尝试刷新 Token -

-
- )} -
- )} - - {/* 切换到本地结果提示 - 仅 Kiro 凭证 */} - {isKiroCredential && switchResult && ( -
-
- {switchResult.success ? ( - - ) : ( - - )} - {switchResult.message} -
- {switchResult.success && switchResult.requires_kiro_restart && ( -
- 请重启 Kiro IDE 使配置生效 -
- )} -
- )} - - {/* 指纹信息展示区域 - 仅 Kiro 凭证 */} - {isKiroCredential && fingerprintExpanded && ( -
-
- - - 设备指纹 - - -
- - {fingerprintLoading ? ( -
-
- 加载中... -
- ) : fingerprintInfo ? ( -
-
- - Machine ID: - - - {fingerprintInfo.machine_id_short}... - - -
-
- - 来源: - - {fingerprintInfo.source} - - - - 认证: - - {fingerprintInfo.auth_method} - - -
-
- ) : ( -
- 无法获取指纹信息 -
- )} -
- )} - - {/* Kiro 详细状态面板 - 仅 Kiro 凭证 */} - {isKiroCredential && kiroStatusExpanded && ( -
-
- - - Kiro 详细状态 - - -
- - {kiroStatusLoading ? ( -
-
- 加载中... -
- ) : kiroHealthScore !== null ? ( -
- {/* 健康分数详情 */} -
-
- - 健康分数 - -
= 80 - ? "bg-emerald-100 text-emerald-700" - : kiroHealthScore >= 60 - ? "bg-amber-100 text-amber-700" - : kiroHealthScore >= 40 - ? "bg-orange-100 text-orange-700" - : "bg-red-100 text-red-700" - }`} - > - {Math.round(kiroHealthScore)} / 100 -
-
- - {/* 健康分数条 */} -
-
= 80 - ? "bg-emerald-500" - : kiroHealthScore >= 60 - ? "bg-amber-500" - : kiroHealthScore >= 40 - ? "bg-orange-500" - : "bg-red-500" - }`} - style={{ - width: `${Math.max(0, Math.min(100, kiroHealthScore))}%`, - }} - >
-
- - {/* 健康状态描述 */} -
- {credential.is_disabled - ? "凭证已被自动禁用,需手动重新启用" - : kiroHealthScore >= 80 - ? "凭证状态良好,可正常使用" - : kiroHealthScore >= 60 - ? "凭证状态一般,建议注意监控" - : kiroHealthScore >= 40 - ? "凭证状态较差,可能有风险" - : "凭证状态异常,需要立即处理"} -
-
- - {/* 状态指标 */} -
-
-
- - - 冷却时间 - -
-
- 根据使用频率计算的建议等待时间 -
-
- -
-
- - - 使用权重 - -
-
- 在轮询池中的权重分配 -
-
-
- - {/* 快速操作 */} -
- {credential.is_disabled ? ( - // 已禁用凭证显示重新启用按钮 - - ) : ( - // 正常凭证显示刷新和检查按钮 - <> - - - - )} -
-
- ) : ( -
- 无法获取状态信息,请重试 -
- )} -
- )} - - {/* 用量信息展示区域 - 仅 Kiro 凭证 */} - {isKiroCredential && usageExpanded && ( -
-
- - - Kiro 用量 - - -
- - {usageError ? ( -
- {usageError} -
- ) : usageInfo ? ( - - ) : ( - - )} -
- )} -
- ); -} diff --git a/src/components/provider-pool/CredentialCard.ui.test.tsx b/src/components/provider-pool/CredentialCard.ui.test.tsx deleted file mode 100644 index 91f8a65f9..000000000 --- a/src/components/provider-pool/CredentialCard.ui.test.tsx +++ /dev/null @@ -1,114 +0,0 @@ -import React from "react"; -import { act } from "react"; -import { createRoot, type Root } from "react-dom/client"; -import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; -import { CredentialCard } from "./CredentialCard"; -import type { CredentialDisplay } from "@/lib/api/providerPool"; - -vi.mock("@/lib/api/providerPool", async () => { - const actual = await vi.importActual( - "@/lib/api/providerPool", - ); - return { - ...actual, - getKiroCredentialFingerprint: vi.fn(), - switchKiroToLocal: vi.fn(), - kiroCredentialApi: {}, - }; -}); - -vi.mock("@/lib/api/usage", () => ({ - usageApi: {}, -})); - -vi.mock("./UsageDisplay", () => ({ - UsageDisplay: () =>
usage
, -})); - -interface MountedRoot { - container: HTMLDivElement; - root: Root; -} - -const mountedRoots: MountedRoot[] = []; - -function createCredential( - overrides: Partial = {}, -): CredentialDisplay { - return { - uuid: "credential-1", - provider_type: "openai", - credential_type: "openai_key", - name: "主账号", - display_credential: "sk-***", - is_healthy: false, - is_disabled: true, - check_health: true, - not_supported_models: [], - usage_count: 12, - error_count: 1, - created_at: "2026-04-18T10:00:00.000Z", - updated_at: "2026-04-18T10:00:00.000Z", - source: "manual", - ...overrides, - }; -} - -function renderCard(credential = createCredential()) { - const container = document.createElement("div"); - document.body.appendChild(container); - const root = createRoot(container); - - act(() => { - root.render( - , - ); - }); - - mountedRoots.push({ container, root }); - return container; -} - -beforeEach(() => { - ( - globalThis as typeof globalThis & { - IS_REACT_ACT_ENVIRONMENT?: boolean; - } - ).IS_REACT_ACT_ENVIRONMENT = true; -}); - -afterEach(() => { - while (mountedRoots.length > 0) { - const mounted = mountedRoots.pop(); - if (!mounted) { - break; - } - act(() => { - mounted.root.unmount(); - }); - mounted.container.remove(); - } -}); - -describe("CredentialCard UI", () => { - it("禁用态卡片应保持浅色主题,不再包含深色背景 fallback", () => { - const container = renderCard(); - const card = container.querySelector(".rounded-xl.border-2"); - const buttons = container.querySelectorAll("button"); - - expect(card).toBeTruthy(); - expect(card?.className).toContain("bg-slate-50/80"); - expect(card?.className).not.toContain("dark:bg-slate-900/60"); - expect(buttons[0]?.className).toContain("bg-emerald-100"); - expect(buttons[0]?.className).not.toContain("dark:bg-emerald-900/30"); - }); -}); diff --git a/src/components/provider-pool/CredentialCardContextMenu.tsx b/src/components/provider-pool/CredentialCardContextMenu.tsx deleted file mode 100644 index fb4af4069..000000000 --- a/src/components/provider-pool/CredentialCardContextMenu.tsx +++ /dev/null @@ -1,149 +0,0 @@ -/** - * 凭证卡片右键菜单组件 - * - * 为凭证池的凭证卡片提供右键菜单功能 - * 支持复制 ID、刷新 Token、查看详情、启用/禁用、删除等操作 - * - * @module components/provider-pool/CredentialCardContextMenu - */ - -import React, { useState } from "react"; -import { Copy, RefreshCw, Info, Power, PowerOff, Trash2 } from "lucide-react"; -import { - ContextMenu, - ContextMenuContent, - ContextMenuItem, - ContextMenuSeparator, - ContextMenuShortcut, - ContextMenuTrigger, -} from "@/components/ui/context-menu"; -import { ConfirmDialog } from "@/components/ConfirmDialog"; -import { toast } from "sonner"; -import type { CredentialDisplay } from "@/lib/api/providerPool"; - -interface CredentialCardContextMenuProps { - /** 凭证数据 */ - credential: CredentialDisplay; - /** 子元素 */ - children: React.ReactNode; - /** 刷新 Token 回调 */ - onRefreshToken?: () => void; - /** 切换启用状态回调 */ - onToggle: () => void; - /** 删除回调 */ - onDelete: () => void; - /** 是否为 OAuth 凭证 */ - isOAuth?: boolean; -} - -export function CredentialCardContextMenu({ - credential, - children, - onRefreshToken, - onToggle, - onDelete, - isOAuth = false, -}: CredentialCardContextMenuProps) { - const [showDeleteDialog, setShowDeleteDialog] = useState(false); - - // 复制凭证 ID - const handleCopyId = async () => { - try { - await navigator.clipboard.writeText(credential.uuid); - toast.success("已复制凭证 ID"); - } catch (error) { - console.error("复制失败:", error); - toast.error("复制失败"); - } - }; - - // 刷新 Token - const handleRefreshToken = () => { - if (onRefreshToken) { - onRefreshToken(); - toast.info("正在刷新 Token..."); - } - }; - - // 查看详情(展开卡片详情) - const handleViewDetail = () => { - // 触发卡片展开,这里通过复制 ID 并提示用户点击卡片查看 - toast.info("请点击卡片查看详细信息"); - }; - - // 确认删除 - const handleConfirmDelete = () => { - onDelete(); - setShowDeleteDialog(false); - }; - - return ( - <> - - {children} - - {/* 复制凭证 ID */} - - - 复制凭证 ID - C - - - {/* 刷新 Token - 仅 OAuth 凭证显示 */} - {isOAuth && onRefreshToken && ( - - - 刷新 Token - R - - )} - - {/* 查看详情 */} - - - 查看详情 - I - - - - - {/* 启用/禁用 */} - - {credential.is_disabled ? ( - <> - - 启用凭证 - - ) : ( - <> - - 禁用凭证 - - )} - E - - - {/* 删除 */} - setShowDeleteDialog(true)} - className="text-red-600 focus:text-red-600" - > - - 删除凭证 - ⌫ - - - - - {/* 删除确认对话框 */} - setShowDeleteDialog(false)} - /> - - ); -} diff --git a/src/components/provider-pool/EditCredentialModal.tsx b/src/components/provider-pool/EditCredentialModal.tsx deleted file mode 100644 index e67d549d5..000000000 --- a/src/components/provider-pool/EditCredentialModal.tsx +++ /dev/null @@ -1,541 +0,0 @@ -import { useState, useEffect } from "react"; -import { - Eye, - EyeOff, - Settings, - Upload, - CheckCircle, - Ban, - Globe, -} from "lucide-react"; -import { open } from "@tauri-apps/plugin-dialog"; -import { Modal } from "@/components/Modal"; -import { - CredentialDisplay, - UpdateCredentialRequest, - PoolProviderType, -} from "@/lib/api/providerPool"; -import { validateProxyUrl } from "@/lib/utils"; - -interface EditCredentialModalProps { - credential: CredentialDisplay | null; - isOpen: boolean; - onClose: () => void; - onEdit: (uuid: string, request: UpdateCredentialRequest) => Promise; -} - -// 各 Provider 支持的模型列表 (参考 AIClient-2-API/src/provider-models.js) -const providerModels: Record = { - kiro: [ - "claude-opus-4-5", - "claude-opus-4-5-20251101", - "claude-haiku-4-5", - "claude-sonnet-4-5", - "claude-sonnet-4-5-20250929", - "claude-sonnet-4-20250514", - "claude-3-7-sonnet-20250219", - ], - gemini: [ - "gemini-2.5-flash", - "gemini-2.5-flash-lite", - "gemini-2.5-pro", - "gemini-2.5-pro-preview-06-05", - "gemini-2.5-flash-preview-09-2025", - "gemini-3-pro-preview", - ], - antigravity: [ - "gemini-3-pro-preview", - "gemini-3-pro-image-preview", - "gemini-2.5-computer-use-preview-10-2025", - "gemini-claude-sonnet-4-5", - "gemini-claude-sonnet-4-5-thinking", - ], - openai: [], // 自定义 API,无预设模型 - claude: [], // 自定义 API,无预设模型 - codex: ["gpt-4o", "gpt-4o-mini", "o1", "o1-mini", "o3-mini"], // Codex(OAuth / API Key) - claude_oauth: [ - "claude-3-5-sonnet-latest", - "claude-3-5-haiku-latest", - "claude-sonnet-4-20250514", - ], // Claude OAuth - gemini_api_key: [ - "gemini-2.5-flash", - "gemini-2.5-flash-lite", - "gemini-2.5-pro", - "gemini-2.5-pro-preview-06-05", - "gemini-2.5-flash-preview-09-2025", - "gemini-3-pro-preview", - ], // Gemini API Key -}; - -export function EditCredentialModal({ - credential, - isOpen, - onClose, - onEdit, -}: EditCredentialModalProps) { - const [name, setName] = useState(""); - const [checkHealth, setCheckHealth] = useState(true); - const [checkModelName, setCheckModelName] = useState(""); - const [notSupportedModels, setNotSupportedModels] = useState([]); - const [loading, setLoading] = useState(false); - const [error, setError] = useState(null); - const [showCredentialDetails, setShowCredentialDetails] = useState(false); - - // 重新上传文件相关状态 - const [newCredFilePath, setNewCredFilePath] = useState(""); - const [newProjectId, setNewProjectId] = useState(""); - // API Key 相关状态 - const [newBaseUrl, setNewBaseUrl] = useState(""); - const [newApiKey, setNewApiKey] = useState(""); - const [showApiKey, setShowApiKey] = useState(false); - - // 代理 URL 相关状态 - const [proxyUrl, setProxyUrl] = useState(""); - const [proxyError, setProxyError] = useState(null); - - // 初始化表单数据 - useEffect(() => { - if (credential) { - console.log("[EditCredentialModal] 初始化表单数据:", { - uuid: credential.uuid, - name: credential.name, - check_model_name: credential.check_model_name, - not_supported_models: credential.not_supported_models, - base_url: credential.base_url, - api_key: credential.api_key ? "***" : undefined, - }); - setName(credential.name || ""); - setCheckHealth(credential.check_health); - setCheckModelName(credential.check_model_name || ""); - setNotSupportedModels(credential.not_supported_models || []); - setNewCredFilePath(""); - setNewProjectId(""); - // 初始化 base_url 为已保存的值 - setNewBaseUrl(credential.base_url || ""); - // 初始化 api_key 为已保存的值 - setNewApiKey(credential.api_key || ""); - setShowApiKey(false); - // 初始化代理 URL 为已保存的值 - setProxyUrl(credential.proxy_url || ""); - setProxyError(null); - setError(null); - } - }, [credential]); - - if (!isOpen || !credential) { - return null; - } - - const isOAuth = credential.credential_type.includes("oauth"); - const isApiKey = credential.credential_type.includes("key"); - - // 获取当前 provider 类型 - const getProviderType = (): PoolProviderType => { - if (credential.credential_type.includes("kiro")) return "kiro"; - if (credential.credential_type.includes("gemini")) return "gemini"; - if (credential.credential_type.includes("codex")) return "codex"; - if (credential.credential_type === "claude_oauth") return "claude_oauth"; - if (credential.credential_type.includes("openai")) return "openai"; - if (credential.credential_type.includes("claude")) return "claude"; - return "kiro"; - }; - - const currentProviderModels = providerModels[getProviderType()] || []; - - const handleSelectNewFile = async () => { - try { - const selected = await open({ - multiple: false, - filters: [{ name: "JSON", extensions: ["json"] }], - }); - if (selected) { - setNewCredFilePath(selected as string); - } - } catch (e) { - console.error("Failed to open file dialog:", e); - } - }; - - const getMaskedCredentialInfo = () => { - if (isOAuth) { - const path = credential.display_credential; - const parts = path.split("/"); - if (parts.length > 1) { - const fileName = parts[parts.length - 1]; - const dirPath = parts.slice(0, -1).join("/"); - return `${dirPath}/***${fileName.slice(-8)}`; - } - return `***${path.slice(-12)}`; - } else { - return credential.display_credential; - } - }; - - const toggleModelSupport = (model: string) => { - setNotSupportedModels((prev) => - prev.includes(model) ? prev.filter((m) => m !== model) : [...prev, model], - ); - }; - - const handleProxyUrlChange = (value: string) => { - setProxyUrl(value); - if (value && !validateProxyUrl(value)) { - setProxyError( - "代理 URL 格式无效,请使用 http://、https:// 或 socks5:// 开头的地址", - ); - } else { - setProxyError(null); - } - }; - - const handleSubmit = async () => { - // 验证代理 URL 格式 - if (proxyUrl && !validateProxyUrl(proxyUrl)) { - setProxyError( - "代理 URL 格式无效,请使用 http://、https:// 或 socks5:// 开头的地址", - ); - return; - } - - setLoading(true); - setError(null); - - try { - const updateRequest: UpdateCredentialRequest = { - // 始终传递 name,空字符串表示清除名称 - name: name.trim(), - check_health: checkHealth, - // 始终传递 check_model_name,空字符串表示清除 - check_model_name: checkModelName.trim(), - // 始终传递 not_supported_models,即使为空数组(用于清除选择) - not_supported_models: notSupportedModels, - new_creds_file_path: newCredFilePath.trim() || undefined, - new_project_id: newProjectId.trim() || undefined, - // API Key 的 base_url(始终传递当前值,空字符串表示使用默认 URL) - new_base_url: isApiKey ? newBaseUrl.trim() : undefined, - // API Key 的 api_key(始终传递当前值) - new_api_key: isApiKey ? newApiKey.trim() : undefined, - // 代理 URL:始终传递当前值,空字符串表示清除代理 - new_proxy_url: proxyUrl.trim(), - }; - - console.log("[EditCredentialModal] 提交更新请求:", updateRequest); - await onEdit(credential.uuid, updateRequest); - onClose(); - } catch (e) { - setError(e instanceof Error ? e.message : String(e)); - } finally { - setLoading(false); - } - }; - - return ( - - {/* Header */} -
-

- - 编辑凭证 -

-
- - {/* Content - Scrollable */} -
-
- {/* 名称 + 健康检查 */} -
-
- - setName(e.target.value)} - placeholder="给这个凭证起个名字..." - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -
-
- - -
-
- - {/* 检查模型名称 */} -
- - setCheckModelName(e.target.value)} - placeholder="用于健康检查的模型名称..." - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -
- - {/* OAuth凭据文件路径 */} - {isOAuth && ( -
- -
- - - -
- {newCredFilePath && ( -
- - 新文件已选择: {newCredFilePath.split("/").pop()} -
- )} -
- )} - - {/* Gemini Project ID */} - {credential.credential_type === "gemini_oauth" && newCredFilePath && ( -
- - setNewProjectId(e.target.value)} - placeholder="留空保持当前项目ID..." - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -
- )} - - {/* API Key 编辑 */} - {isApiKey && ( - <> -
- -
- setNewApiKey(e.target.value)} - placeholder="留空保持当前 Key,或输入新的 API Key..." - className="flex-1 rounded-lg border bg-background px-3 py-2 text-sm" - /> - -
-

- 当前: {credential.display_credential} -

-
-
- - setNewBaseUrl(e.target.value)} - placeholder={ - credential.credential_type === "openai_key" - ? "https://api.openai.com" - : "https://api.anthropic.com" - } - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -

- 留空使用默认 URL,或输入自定义代理地址(不要包含 /v1 后缀) -

-
- - )} - - {/* 不支持的模型 - Checkbox Grid */} -
-
- - - - 选择此提供商不支持的模型,系统会自动排除这些模型 - -
-
- {currentProviderModels.map((model) => ( - - ))} -
-
- - {/* 高级选项:代理设置 */} -
-
- - - - (高级选项) - -
-
-
- - handleProxyUrlChange(e.target.value)} - placeholder="例如: http://127.0.0.1:7890 或 socks5://127.0.0.1:1080" - className={`w-full rounded-lg border bg-background px-3 py-2 text-sm ${ - proxyError ? "border-red-500" : "" - }`} - /> - {proxyError ? ( -

{proxyError}

- ) : ( -

- 留空则使用全局代理设置。支持 http://、https://、socks5:// - 协议 -

- )} -
-
-

- 代理优先级说明: -

-
    -
  • 此凭证代理优先于全局代理
  • -
  • 留空时使用全局代理设置
  • -
  • 全局代理可在「设置 → 通用」中配置
  • -
-
-
-
- - {/* 使用统计(只读) */} -
- -
-
- - 使用次数 - - {credential.usage_count} -
-
- - 错误次数 - - {credential.error_count} -
-
- - 最后使用 - - - {credential.last_used || "从未"} - -
-
-
- - {/* Error */} - {error && ( -
- {error} -
- )} -
-
- - {/* Footer */} -
- - -
-
- ); -} diff --git a/src/components/provider-pool/ErrorDisplay.test.tsx b/src/components/provider-pool/ErrorDisplay.test.tsx deleted file mode 100644 index 3e4e0f1d7..000000000 --- a/src/components/provider-pool/ErrorDisplay.test.tsx +++ /dev/null @@ -1,70 +0,0 @@ -import React from "react"; -import { act } from "react"; -import { createRoot, type Root } from "react-dom/client"; -import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; -import { ErrorDisplay, type ErrorInfo } from "./ErrorDisplay"; - -interface MountedRoot { - container: HTMLDivElement; - root: Root; -} - -const mountedRoots: MountedRoot[] = []; - -function renderErrorDisplay(errors: ErrorInfo[]) { - const container = document.createElement("div"); - document.body.appendChild(container); - const root = createRoot(container); - - act(() => { - root.render( - , - ); - }); - - mountedRoots.push({ container, root }); - return container; -} - -beforeEach(() => { - ( - globalThis as typeof globalThis & { - IS_REACT_ACT_ENVIRONMENT?: boolean; - } - ).IS_REACT_ACT_ENVIRONMENT = true; -}); - -afterEach(() => { - while (mountedRoots.length > 0) { - const mounted = mountedRoots.pop(); - if (!mounted) { - break; - } - act(() => { - mounted.root.unmount(); - }); - mounted.container.remove(); - } -}); - -describe("ErrorDisplay", () => { - it("错误通知应保持浅色主题按钮和表面", () => { - const container = renderErrorDisplay([ - { - id: "error-1", - message: "检测失败,请稍后重试", - type: "general", - }, - ]); - - const notice = container.querySelector(".rounded-lg.border"); - expect(notice).toBeTruthy(); - expect(notice?.className).toContain("bg-slate-50"); - expect(notice?.className).not.toContain("dark:bg-slate-950/30"); - - const buttons = container.querySelectorAll("button"); - expect(buttons).toHaveLength(2); - expect(buttons[0]?.className).toContain("bg-white"); - expect(buttons[0]?.className).not.toContain("dark:bg-slate-900/70"); - }); -}); diff --git a/src/components/provider-pool/ErrorDisplay.tsx b/src/components/provider-pool/ErrorDisplay.tsx deleted file mode 100644 index 9240b2f31..000000000 --- a/src/components/provider-pool/ErrorDisplay.tsx +++ /dev/null @@ -1,244 +0,0 @@ -import { useState, useEffect } from "react"; -import { - AlertTriangle, - X, - RotateCcw, - Trash2, - Settings, - CheckCircle2, - KeyRound, -} from "lucide-react"; - -export interface ErrorInfo { - id: string; - message: string; - type: - | "delete" - | "toggle" - | "reset" - | "health_check" - | "refresh_token" - | "migrate" - | "config" - | "general" - | "success" - | "reauth"; // 需要重新授权 - uuid?: string; // 相关凭证的UUID(如果有的话) -} - -interface ErrorDisplayProps { - errors: ErrorInfo[]; - onDismiss: (id: string) => void; - onRetry?: (error: ErrorInfo) => void; -} - -const ErrorTypeConfig = { - delete: { - icon: Trash2, - color: "text-red-600", - bgColor: "bg-red-50", - borderColor: "border-red-200", - }, - toggle: { - icon: Settings, - color: "text-emerald-600", - bgColor: "bg-emerald-50", - borderColor: "border-emerald-200", - }, - reset: { - icon: RotateCcw, - color: "text-amber-600", - bgColor: "bg-amber-50", - borderColor: "border-amber-200", - }, - health_check: { - icon: AlertTriangle, - color: "text-amber-600", - bgColor: "bg-amber-50", - borderColor: "border-amber-200", - }, - refresh_token: { - icon: RotateCcw, - color: "text-sky-600", - bgColor: "bg-sky-50", - borderColor: "border-sky-200", - }, - migrate: { - icon: AlertTriangle, - color: "text-emerald-600", - bgColor: "bg-emerald-50", - borderColor: "border-emerald-200", - }, - config: { - icon: Settings, - color: "text-sky-600", - bgColor: "bg-sky-50", - borderColor: "border-sky-200", - }, - general: { - icon: AlertTriangle, - color: "text-slate-600", - bgColor: "bg-slate-50", - borderColor: "border-slate-200", - }, - success: { - icon: CheckCircle2, - color: "text-emerald-600", - bgColor: "bg-emerald-50", - borderColor: "border-emerald-200", - }, - reauth: { - icon: KeyRound, - color: "text-amber-600", - bgColor: "bg-amber-50", - borderColor: "border-amber-200", - }, -}; - -function ErrorItem({ - error, - onDismiss, - onRetry, -}: { - error: ErrorInfo; - onDismiss: (id: string) => void; - onRetry?: (error: ErrorInfo) => void; -}) { - const config = ErrorTypeConfig[error.type]; - const IconComponent = config.icon; - - return ( -
-
- -
-
- {error.message} -
-
- {onRetry && ( - - )} - -
-
-
-
- ); -} - -export function ErrorDisplay({ - errors, - onDismiss, - onRetry, -}: ErrorDisplayProps) { - // 自动关闭通知 - useEffect(() => { - const timers: ReturnType[] = []; - - errors.forEach((error) => { - // 成功消息 3 秒后自动关闭,其他类型 15 秒后自动关闭 - if (error.type === "success") { - const timer = setTimeout(() => { - onDismiss(error.id); - }, 3000); // 3秒后自动关闭 - timers.push(timer); - } else if (error.type === "general" || error.message.includes("💡")) { - const timer = setTimeout(() => { - onDismiss(error.id); - }, 15000); // 15秒后自动关闭 - timers.push(timer); - } - }); - - return () => { - timers.forEach((timer) => clearTimeout(timer)); - }; - }, [errors, onDismiss]); - - if (errors.length === 0) { - return null; - } - - return ( -
-
- {errors.map((error) => ( - - ))} -
-
- ); -} - -// Hook for managing errors and success messages -// eslint-disable-next-line react-refresh/only-export-components -export function useErrorDisplay() { - const [errors, setErrors] = useState([]); - - const showError = ( - message: string, - type: ErrorInfo["type"] = "general", - uuid?: string, - ) => { - // 检查是否已经存在相同的错误消息(基于 message, type, uuid 的组合) - setErrors((prev) => { - const isDuplicate = prev.some( - (existing) => - existing.message === message && - existing.type === type && - existing.uuid === uuid, - ); - - if (isDuplicate) { - return prev; // 如果重复,不添加新的错误 - } - - const id = - Date.now().toString() + Math.random().toString(36).substr(2, 9); - const error: ErrorInfo = { id, message, type, uuid }; - return [...prev, error]; - }); - }; - - const showSuccess = (message: string, uuid?: string) => { - const id = Date.now().toString() + Math.random().toString(36).substr(2, 9); - const info: ErrorInfo = { id, message, type: "success", uuid }; - setErrors((prev) => [...prev, info]); - }; - - const dismissError = (id: string) => { - setErrors((prev) => prev.filter((error) => error.id !== id)); - }; - - const clearErrors = () => { - setErrors([]); - }; - - return { - errors, - showError, - showSuccess, - dismissError, - clearErrors, - }; -} diff --git a/src/components/provider-pool/GeminiApiKeySection.tsx b/src/components/provider-pool/GeminiApiKeySection.tsx deleted file mode 100644 index fda1c835b..000000000 --- a/src/components/provider-pool/GeminiApiKeySection.tsx +++ /dev/null @@ -1,284 +0,0 @@ -import { useState } from "react"; -import { Plus, Trash2, Key, Globe, Ban, Eye, EyeOff } from "lucide-react"; -import type { GeminiApiKeyEntry } from "@/lib/api/providerRuntime"; - -interface GeminiApiKeySectionProps { - entries: GeminiApiKeyEntry[] | undefined; - onChange: (entries: GeminiApiKeyEntry[]) => void; -} - -export function GeminiApiKeySection({ - entries = [], - onChange, -}: GeminiApiKeySectionProps) { - const primaryActionButtonClassName = - "rounded-lg border border-emerald-200 bg-[linear-gradient(135deg,#0ea5e9_0%,#14b8a6_52%,#10b981_100%)] text-sm text-white shadow-sm shadow-emerald-950/15 hover:opacity-95"; - // 确保 entries 是数组 - const safeEntries = Array.isArray(entries) ? entries : []; - const [showKeys, setShowKeys] = useState>(new Set()); - const [editingExclusions, setEditingExclusions] = useState( - null, - ); - const [exclusionInput, setExclusionInput] = useState(""); - - const toggleShowKey = (id: string) => { - const newSet = new Set(showKeys); - if (newSet.has(id)) { - newSet.delete(id); - } else { - newSet.add(id); - } - setShowKeys(newSet); - }; - - const addEntry = () => { - const newEntry: GeminiApiKeyEntry = { - id: `gemini-api-${Date.now()}`, - api_key: "", - base_url: null, - proxy_url: null, - excluded_models: [], - disabled: false, - }; - onChange([...safeEntries, newEntry]); - }; - - const updateEntry = (id: string, updates: Partial) => { - onChange(safeEntries.map((e) => (e.id === id ? { ...e, ...updates } : e))); - }; - - const removeEntry = (id: string) => { - onChange(safeEntries.filter((e) => e.id !== id)); - }; - - const addExclusion = (id: string) => { - if (!exclusionInput.trim()) return; - const entry = safeEntries.find((e) => e.id === id); - if (entry) { - const currentExcluded = Array.isArray(entry.excluded_models) - ? entry.excluded_models - : []; - updateEntry(id, { - excluded_models: [...currentExcluded, exclusionInput.trim()], - }); - setExclusionInput(""); - } - }; - - const removeExclusion = (id: string, model: string) => { - const entry = safeEntries.find((e) => e.id === id); - if (entry) { - const currentExcluded = Array.isArray(entry.excluded_models) - ? entry.excluded_models - : []; - updateEntry(id, { - excluded_models: currentExcluded.filter((m) => m !== model), - }); - } - }; - - return ( -
-
-
- -
-

Gemini API Key 多账号

-

- 配置多个 Gemini API Key 实现负载均衡 -

-
-
- -
- - {safeEntries.length === 0 ? ( -
-

暂无 Gemini API Key

-

点击上方"添加"按钮添加 API Key

-
- ) : ( -
- {safeEntries.map((entry) => { - // 确保 excluded_models 是数组 - const excludedModels = Array.isArray(entry.excluded_models) - ? entry.excluded_models - : []; - return ( -
- {/* Header */} -
- - {entry.id} - -
- - -
-
- - {/* API Key */} -
- -
- - updateEntry(entry.id, { api_key: e.target.value }) - } - placeholder="AIzaSy..." - className="w-full px-3 py-1.5 pr-10 rounded border bg-background text-sm font-mono" - /> - -
-
- - {/* Base URL */} -
- - - updateEntry(entry.id, { - base_url: e.target.value || null, - }) - } - placeholder="https://generativelanguage.googleapis.com" - className="w-full px-3 py-1.5 rounded border bg-background text-sm" - /> -
- - {/* Proxy URL */} -
- - - updateEntry(entry.id, { - proxy_url: e.target.value || null, - }) - } - placeholder="socks5://127.0.0.1:1080" - className="w-full px-3 py-1.5 rounded border bg-background text-sm" - /> -
- - {/* Excluded Models */} -
- -
- {excludedModels.map((model) => ( - - {model} - - - ))} -
- {editingExclusions === entry.id ? ( -
- setExclusionInput(e.target.value)} - placeholder="gemini-2.5-pro 或 *-preview" - className="flex-1 px-2 py-1 rounded border bg-background text-sm" - onKeyDown={(e) => { - if (e.key === "Enter") { - addExclusion(entry.id); - } - }} - /> - - -
- ) : ( - - )} -

- 支持通配符,如 *-preview 匹配所有预览模型 -

-
-
- ); - })} -
- )} -
- ); -} diff --git a/src/components/provider-pool/ProviderPoolPage.tsx b/src/components/provider-pool/ProviderPoolPage.tsx deleted file mode 100644 index efc81dff0..000000000 --- a/src/components/provider-pool/ProviderPoolPage.tsx +++ /dev/null @@ -1,702 +0,0 @@ -/** - * @file ProviderPoolPage 组件 - * @description 凭证池管理页面,支持 OAuth 凭证卡片布局和 API Key 左右分栏布局 - * @module components/provider-pool/ProviderPoolPage - * - * **Feature: provider-ui-refactor** - * **Validates: Requirements 1.1, 2.1, 2.2, 2.3** - */ - -import { - useState, - useEffect, - forwardRef, - useImperativeHandle, - useRef, -} from "react"; -import { - RefreshCw, - Plus, - Heart, - HeartOff, - RotateCcw, - Activity, -} from "lucide-react"; -import { useProviderPool } from "@/hooks/useProviderPool"; -import { useApiKeyProvider } from "@/hooks/useApiKeyProvider"; -import { CredentialCard } from "./CredentialCard"; -import { CredentialCardContextMenu } from "./CredentialCardContextMenu"; -import { AddCredentialModal } from "./AddCredentialModal"; -import { EditCredentialModal } from "./EditCredentialModal"; -import { ErrorDisplay, useErrorDisplay } from "./ErrorDisplay"; -import { ConfirmDialog } from "@/components/ConfirmDialog"; -import { ProviderIcon } from "@/icons/providers"; -import { ApiKeyProviderSection } from "./api-key"; -import type { ApiKeyProviderSectionRef } from "./api-key"; -import { RelayProvidersSection } from "./RelayProvidersSection"; -import { AsrProviderSection } from "@/components/voice"; -import { - getLocalKiroCredentialUuid, - type PoolProviderType, - type CredentialDisplay, - type UpdateCredentialRequest, -} from "@/lib/api/providerPool"; - -export interface ProviderPoolPageRef { - refresh: () => void; -} - -interface ProviderPoolPageProps { - hideHeader?: boolean; -} - -// OAuth 类型凭证(需要上传凭证文件或登录授权) -const oauthProviderTypes: PoolProviderType[] = [ - "kiro", - "gemini", - "antigravity", - "codex", - "claude_oauth", -]; - -// 配置类型 tab(非凭证池) -type ConfigTabType = "connect"; - -// 所有 tab 类型 -type TabType = PoolProviderType | ConfigTabType; - -const providerLabels: Record = { - kiro: "Kiro (AWS)", - gemini: "Gemini (Google)", - antigravity: "Antigravity (Gemini 3 Pro)", - openai: "OpenAI", - claude: "Claude (Anthropic)", - codex: "Codex (OpenAI)", - claude_oauth: "Claude OAuth", - gemini_api_key: "Gemini", -}; - -// 判断是否为配置类型 tab -const isConfigTab = (tab: TabType): tab is ConfigTabType => { - return tab === "connect"; -}; - -// 分类类型 -type CategoryType = "oauth" | "apikey" | "connect" | "voice"; - -export const ProviderPoolPage = forwardRef< - ProviderPoolPageRef, - ProviderPoolPageProps ->(({ hideHeader = false }, ref) => { - const [addModalOpen, setAddModalOpen] = useState(false); - const [editModalOpen, setEditModalOpen] = useState(false); - const [editingCredential, setEditingCredential] = - useState(null); - const [activeCategory, setActiveCategory] = useState("apikey"); - const [activeTab, setActiveTab] = useState("kiro"); - const [deletingCredentials, setDeletingCredentials] = useState>( - new Set(), - ); - const [deleteConfirm, setDeleteConfirm] = useState(null); - const { errors, showError, showSuccess, dismissError } = useErrorDisplay(); - - // ApiKeyProviderSection 的 ref - const apiKeyProviderSectionRef = useRef(null); - - const { - overview, - loading, - error, - checkingHealth, - refreshingToken, - refresh, - deleteCredential, - toggleCredential, - resetCredential, - resetHealth, - checkCredentialHealth, - checkTypeHealth, - refreshCredentialToken, - updateCredential, - } = useProviderPool(); - - // API Key Provider Hook - const { refresh: refreshApiKeyProviders } = useApiKeyProvider(); - - const [_migrating, setMigrating] = useState(false); - - // Kiro 本地活跃凭证 UUID - const [localActiveUuid, setLocalActiveUuid] = useState(null); - - // 获取本地活跃的 Kiro 凭证 UUID - const fetchLocalActiveUuid = async () => { - try { - const uuid = await getLocalKiroCredentialUuid(); - console.log("[ProviderPoolPage] Local active Kiro UUID:", uuid); - setLocalActiveUuid(uuid); - } catch (e) { - console.error("Failed to get local Kiro credential:", e); - setLocalActiveUuid(null); - } - }; - - useEffect(() => { - // 只在 Kiro tab 时检测本地活跃凭证 - if (activeTab === "kiro") { - fetchLocalActiveUuid(); - } - }, [activeTab, overview]); - - useImperativeHandle(ref, () => ({ - refresh: () => { - refresh(); - refreshApiKeyProviders(); - }, - })); - - const handleDeleteClick = (uuid: string) => { - setDeleteConfirm(uuid); - }; - - const handleDeleteConfirm = async () => { - if (!deleteConfirm) return; - const uuid = deleteConfirm; - setDeleteConfirm(null); - setDeletingCredentials((prev) => new Set(prev).add(uuid)); - try { - const providerType = !isConfigTab(activeTab) - ? (activeTab as PoolProviderType) - : undefined; - await deleteCredential(uuid, providerType); - } catch (e) { - showError(e instanceof Error ? e.message : String(e), "delete", uuid); - } finally { - setDeletingCredentials((prev) => { - const next = new Set(prev); - next.delete(uuid); - return next; - }); - } - }; - - const handleToggle = async (credential: CredentialDisplay) => { - try { - await toggleCredential(credential.uuid, !credential.is_disabled); - } catch (e) { - showError( - e instanceof Error ? e.message : String(e), - "toggle", - credential.uuid, - ); - } - }; - - const handleReset = async (uuid: string) => { - try { - await resetCredential(uuid); - } catch (e) { - showError(e instanceof Error ? e.message : String(e), "reset", uuid); - } - }; - - const handleCheckHealth = async (uuid: string) => { - try { - const result = await checkCredentialHealth(uuid); - if (result.success) { - showSuccess("健康检查通过!", uuid); - } else { - showError(result.message || "健康检查未通过", "health_check", uuid); - } - } catch (e) { - showError( - e instanceof Error ? e.message : String(e), - "health_check", - uuid, - ); - } - }; - - const handleCheckTypeHealth = async (providerType: PoolProviderType) => { - try { - await checkTypeHealth(providerType); - } catch (e) { - showError(e instanceof Error ? e.message : String(e), "health_check"); - } - }; - - const handleResetTypeHealth = async (providerType: PoolProviderType) => { - try { - await resetHealth(providerType); - } catch (e) { - showError(e instanceof Error ? e.message : String(e), "reset"); - } - }; - - // 迁移 Private 配置到凭证池 - const _handleMigratePrivateConfig = async () => { - setMigrating(true); - try { - // const config = await getConfig(); - // const result = await migratePrivateConfig(config); - const result = { - migrated_count: 0, - skipped_count: 0, - errors: [] as string[], - }; - if (result.migrated_count > 0) { - showSuccess( - `成功迁移 ${result.migrated_count} 个凭证${result.skipped_count > 0 ? `,跳过 ${result.skipped_count} 个已存在的凭证` : ""}`, - ); - } else if (result.skipped_count > 0) { - showSuccess(`所有凭证已存在,跳过 ${result.skipped_count} 个`); - } else { - showSuccess("没有需要迁移的凭证"); - } - if (result.errors.length > 0) { - showError(`部分迁移失败: ${result.errors.join(", ")}`, "migrate"); - } - } catch (e) { - showError(e instanceof Error ? e.message : String(e), "migrate"); - } finally { - setMigrating(false); - } - }; - - const handleRefreshToken = async (uuid: string) => { - try { - await refreshCredentialToken(uuid); - showSuccess("Token 刷新成功!", uuid); - } catch (e) { - showError( - e instanceof Error ? e.message : String(e), - "refresh_token", - uuid, - ); - } - }; - - const handleEdit = (credential: CredentialDisplay) => { - setEditingCredential(credential); - setEditModalOpen(true); - }; - - const handleEditSubmit = async ( - uuid: string, - request: UpdateCredentialRequest, - ) => { - try { - await updateCredential(uuid, request); - } catch (e) { - throw new Error( - `编辑失败: ${e instanceof Error ? e.message : String(e)}`, - ); - } - }; - - const closeEditModal = () => { - setEditModalOpen(false); - setEditingCredential(null); - }; - - const openAddModal = () => { - setAddModalOpen(true); - }; - - const getProviderOverview = (providerType: PoolProviderType) => { - return overview.find((p) => p.provider_type === providerType); - }; - - const getCredentialCount = (providerType: PoolProviderType) => { - const pool = getProviderOverview(providerType); - return pool?.credentials?.length || 0; - }; - - // Current tab data (仅用于 OAuth 凭证 tab) - const currentPool = - !isConfigTab(activeTab) && activeCategory === "oauth" - ? getProviderOverview(activeTab as PoolProviderType) - : null; - const currentStats = currentPool?.stats; - const currentCredentials = currentPool?.credentials || []; - - if (hideHeader) { - return ( -
- -
- ); - } - - return ( -
- {!hideHeader && ( -
-
-

- Provider 与凭证 -

-

- 默认先管理 API Key - Provider。OAuth、语音和中转服务保留在同一入口,但不再和日常 - Provider 配置抢同一视觉焦点。 -

-
-
- )} - - {error && ( -
- {error} -
- )} - - {/* Category Tabs - 第一行:分类选择 */} -
- - - - -
- - {/* OAuth 凭证分类 - Provider 选择图标网格 */} - {activeCategory === "oauth" && ( -
- {oauthProviderTypes.map((providerType) => { - const count = getCredentialCount(providerType); - const isActive = activeTab === providerType; - return ( - - ); - })} -
- )} - - {/* Connect 分类 - 中转商列表 */} - {activeCategory === "connect" && ( -
- -
- )} - - {/* API Key 分类 - 左右分栏布局 */} - {activeCategory === "apikey" && ( -
- -
- )} - - {/* 语音服务分类 */} - {activeCategory === "voice" && ( -
- -
- )} - - {/* OAuth 凭证内容 - 卡片布局 */} - {activeCategory === "oauth" && - !isConfigTab(activeTab) && - (loading ? ( -
- -
- ) : ( -
- {/* Stats and Actions Bar */} -
-
- {currentStats && currentStats.total > 0 && ( -
- - - 健康: {currentStats.healthy} - - - - 不健康: {currentStats.unhealthy} - - - 总计: {currentStats.total} - -
- )} -
-
- {currentCredentials.length > 0 && ( - <> - - - - - )} -
-
- - {/* Credentials List */} - {currentCredentials.length === 0 ? ( -
-

- 暂无 {providerLabels[activeTab as PoolProviderType]} 凭证 -

-

点击上方"添加凭证"按钮添加

- -
- ) : ( -
- {currentCredentials.map((credential) => { - // 判断是否为 OAuth 类型(需要刷新 Token 功能) - const isOAuthType = - credential.credential_type.includes("oauth"); - // 判断是否为 Kiro 凭证(支持用量查询) - const isKiroCredential = activeTab === "kiro"; - const isLocalActive = - isKiroCredential && credential.uuid === localActiveUuid; - - if (isKiroCredential) { - console.log( - `[ProviderPoolPage] Credential ${credential.uuid.substring(0, 8)}: isLocalActive=${isLocalActive}, localActiveUuid=${localActiveUuid?.substring(0, 8)}`, - ); - } - - return ( - handleRefreshToken(credential.uuid) - : undefined - } - onToggle={() => handleToggle(credential)} - onDelete={() => handleDeleteClick(credential.uuid)} - isOAuth={isOAuthType} - > -
- handleToggle(credential)} - onDelete={() => handleDeleteClick(credential.uuid)} - onReset={() => handleReset(credential.uuid)} - onCheckHealth={() => - handleCheckHealth(credential.uuid) - } - onRefreshToken={ - isOAuthType - ? () => handleRefreshToken(credential.uuid) - : undefined - } - onEdit={() => handleEdit(credential)} - deleting={deletingCredentials.has(credential.uuid)} - checkingHealth={checkingHealth === credential.uuid} - refreshingToken={refreshingToken === credential.uuid} - isKiroCredential={isKiroCredential} - isLocalActive={isLocalActive} - onSwitchToLocal={ - isKiroCredential ? fetchLocalActiveUuid : undefined - } - /> -
-
- ); - })} -
- )} -
- ))} - - {/* Add Credential Modal (仅 OAuth 凭证 tab) */} - {addModalOpen && - activeCategory === "oauth" && - !isConfigTab(activeTab) && ( - { - setAddModalOpen(false); - }} - onSuccess={() => { - setAddModalOpen(false); - refresh(); - }} - /> - )} - - {/* Edit Credential Modal */} - - - {/* Error Display */} - { - switch (error.type) { - case "health_check": - if (error.uuid) { - handleCheckHealth(error.uuid); - } - break; - case "refresh_token": - if (error.uuid) { - handleRefreshToken(error.uuid); - } - break; - case "reset": - if (error.uuid) { - handleReset(error.uuid); - } - break; - } - dismissError(error.id); - }} - /> - - setDeleteConfirm(null)} - /> -
- ); -}); - -ProviderPoolPage.displayName = "ProviderPoolPage"; diff --git a/src/components/provider-pool/README.md b/src/components/provider-pool/README.md deleted file mode 100644 index c8695f777..000000000 --- a/src/components/provider-pool/README.md +++ /dev/null @@ -1,79 +0,0 @@ -# Provider Pool 组件 - -本目录包含凭证池管理界面的所有组件。 - -## 组件列表 - -| 文件 | 描述 | -| ------------------------------- | ---------------------------------------------------------------- | -| `ProviderPoolPage.tsx` | 凭证池管理主页面,支持 OAuth 凭证卡片布局和 API Key 左右分栏布局 | -| `CredentialCard.tsx` | OAuth 凭证卡片组件,显示健康状态、使用统计和操作按钮 | -| `CredentialCardContextMenu.tsx` | 凭证卡片右键菜单组件 | -| `AddCredentialModal.tsx` | 添加凭证模态框组件 | -| `EditCredentialModal.tsx` | 编辑凭证模态框组件 | -| `ErrorDisplay.tsx` | 错误显示组件 | -| `UsageDisplay.tsx` | 用量显示组件 | -| `RelayProvidersSection.tsx` | Connect 中转商列表组件,展示已验证的中转服务商 | -| `VertexAISection.tsx` | Vertex AI 配置区域组件 | -| `AmpConfigSection.tsx` | Amp CLI 配置区域组件 | -| `GeminiApiKeySection.tsx` | Gemini API Key 配置区域组件 | -| `CodexSection.tsx` | Codex 配置区域组件 | -| `IFlowSection.tsx` | iFlow 配置区域组件 | -| `OAuthPluginTab.tsx` | OAuth 插件标签页组件 | -| `index.ts` | 组件导出入口 | - -## 子目录 - -| 目录 | 描述 | -| ------------------- | ----------------------------------------- | -| `api-key/` | API Key Provider 管理组件(左右分栏布局) | -| `credential-forms/` | 各类凭证表单组件 | - -## 测试文件 - -| 文件 | 描述 | -| ------------------------ | --------------------------------------------- | -| `CredentialCard.test.ts` | Property 3 属性测试:OAuth 凭证卡片信息完整性 | - -## 使用示例 - -```tsx -import { ProviderPoolPage } from "@/components/provider-pool"; - -function App() { - return ; -} -``` - -## 相关需求 - -- Requirements 1.1: API Key Provider 左右分栏布局 -- Requirements 2.1, 2.2, 2.3: OAuth 凭证保持卡片布局 -- Requirements 3.1-3.6: 完整支持 System Provider 类型 -- Connect: 中转商浏览和一键添加功能 - -## 架构说明 - -ProviderPoolPage 支持四种分类: - -1. **OAuth 凭证** - 使用卡片式布局显示 OAuth 类型凭证 -2. **API Key** - 使用左右分栏布局(ApiKeyProviderSection) -3. **Connect** - 中转商列表,支持浏览和一键获取 API Key -4. **语音服务** - 语音 Provider 管理入口 - -## Prompt Cache 认知边界 - -Provider Pool 页面当前已把 Prompt Cache 能力前置到 Provider UI,而不是等到对话发出后才暴露: - -- `anthropic`:按官方 Anthropic 能力链展示 -- `anthropic-compatible`:先按已知官方 Anthropic 兼容端点识别自动缓存,未知端点才回退为“仅显式缓存” -- 其它 Provider:默认不展示 Prompt Cache 标签或 notice - -当前页面上的主要提示落点: - -- 左侧 Provider 列表:`显式缓存` badge -- 右侧 Provider 详情头部:`显式缓存` badge -- 新增自定义 Provider:amber Prompt Cache 提示 -- 编辑 Provider 配置:amber Prompt Cache 提示 - -这条语义的当前事实源位于 `src/lib/model/providerPromptCacheSupport.ts`。如果以后要新增新的缓存能力提示,优先扩这份 helper,不要在单个组件里自行判断。 diff --git a/src/components/provider-pool/RelayProvidersSection.tsx b/src/components/provider-pool/RelayProvidersSection.tsx deleted file mode 100644 index 74344f976..000000000 --- a/src/components/provider-pool/RelayProvidersSection.tsx +++ /dev/null @@ -1,259 +0,0 @@ -/** - * @file RelayProvidersSection 组件 - * @description 中转商列表展示组件,支持浏览和一键跳转获取 API Key - * @module components/provider-pool/RelayProvidersSection - * - * _Requirements: Connect 中转商浏览功能_ - */ - -import { useState } from "react"; -import { - RefreshCw, - ExternalLink, - Globe, - Mail, - MessageCircle, - Shield, - ShieldCheck, - Zap, - Clock, - AlertCircle, -} from "lucide-react"; -import { useRelayRegistry } from "@/hooks/useRelayRegistry"; -import type { RelayInfo } from "@/hooks/useDeepLink"; -import { open } from "@tauri-apps/plugin-shell"; - -/** - * 中转商卡片组件 - */ -function RelayProviderCard({ provider }: { provider: RelayInfo }) { - const [imageError, setImageError] = useState(false); - const primaryActionButtonClassName = - "rounded-lg border border-emerald-200 bg-[linear-gradient(135deg,#0ea5e9_0%,#14b8a6_52%,#10b981_100%)] text-sm font-medium text-white shadow-sm shadow-emerald-950/15 hover:opacity-95 transition-colors"; - - // 打开外部链接 - const handleOpenLink = async (url: string) => { - try { - await open(url); - } catch (e) { - console.error("打开链接失败:", e); - // 回退到 window.open - window.open(url, "_blank"); - } - }; - - // 获取 API Key 的链接(优先使用 dashboard,其次 website) - const getApiKeyLink = () => { - return provider.links.dashboard || provider.links.website; - }; - - return ( -
- {/* 头部:Logo + 名称 */} -
- {/* Logo */} -
- {provider.branding.logo && !imageError ? ( - {provider.name} setImageError(true)} - /> - ) : ( - - )} -
- - {/* 名称和描述 */} -
-
-

- {provider.name} -

- {provider.features.verified && ( - - )} -
-

- {provider.description} -

-
-
- - {/* API 信息 */} -
- - - {provider.api.protocol.toUpperCase()} - - {provider.features.streaming && ( - - 流式响应 - - )} - {provider.features.models && provider.features.models.length > 0 && ( - - {provider.features.models.length} 个模型 - - )} -
- - {/* 功能特性 */} - {provider.features.models && provider.features.models.length > 0 && ( -
-

支持模型:

-
- {provider.features.models.slice(0, 5).map((model) => ( - - {model} - - ))} - {provider.features.models.length > 5 && ( - - +{provider.features.models.length - 5} - - )} -
-
- )} - - {/* 操作按钮 */} -
- {getApiKeyLink() && ( - - )} - - {provider.links.docs && ( - - )} - - {provider.contact.email && ( - - )} - - {provider.contact.discord && ( - - )} -
-
- ); -} - -/** - * 中转商列表组件 - */ -export function RelayProvidersSection() { - const { providers, isLoading, error, refresh } = useRelayRegistry(); - - return ( -
- {/* 头部说明 */} -
-
-
- -

Lime Connect

-
-

- 浏览已验证的 AI API 中转服务商,获取 API Key - 后可通过链接一键添加到凭证池 -

-
- -
- - {/* 错误提示 */} - {error && ( -
- - {error.message} - -
- )} - - {/* 加载状态 */} - {isLoading && providers.length === 0 && ( -
- -
- )} - - {/* 空状态 */} - {!isLoading && providers.length === 0 && !error && ( -
- -

暂无中转商

-

点击刷新按钮加载中转商列表

- -
- )} - - {/* 中转商列表 */} - {providers.length > 0 && ( -
- {providers.map((provider) => ( - - ))} -
- )} - - {/* 底部说明 */} - {providers.length > 0 && ( -
-

- 获取 API Key 后,中转商会提供一个 lime://{" "} - 链接,点击即可一键添加到凭证池 -

-
- )} -
- ); -} - -export default RelayProvidersSection; diff --git a/src/components/provider-pool/UsageDisplay.tsx b/src/components/provider-pool/UsageDisplay.tsx deleted file mode 100644 index 370ebaa59..000000000 --- a/src/components/provider-pool/UsageDisplay.tsx +++ /dev/null @@ -1,140 +0,0 @@ -import { AlertTriangle, TrendingUp, Zap, Wallet } from "lucide-react"; -import type { UsageInfo } from "@/lib/api/usage"; - -interface UsageDisplayProps { - usage: UsageInfo; - loading?: boolean; -} - -/** - * 用量显示组件 - * - * 显示订阅类型、总额度、已使用、余额 - * 低余额时显示警告样式 - * - * _Requirements: 3.3, 3.4_ - */ -export function UsageDisplay({ usage, loading }: UsageDisplayProps) { - if (loading) { - return ( -
-
-
-
-
-
-
-
- ); - } - - // 计算使用百分比 - const usagePercent = - usage.usageLimit > 0 - ? Math.round((usage.currentUsage / usage.usageLimit) * 100) - : 0; - - // 格式化数字 - const formatNumber = (num: number) => { - if (num >= 1000000) { - return `${(num / 1000000).toFixed(1)}M`; - } - if (num >= 1000) { - return `${(num / 1000).toFixed(1)}K`; - } - return num.toFixed(1); - }; - - return ( -
- {/* 标题和警告 */} -
-
- - - {usage.subscriptionTitle || "用量信息"} - -
- {usage.isLowBalance && ( -
- - 余额不足 -
- )} -
- - {/* 进度条 */} -
-
-
-
-
- 已使用 {usagePercent}% - 剩余 {100 - usagePercent}% -
-
- - {/* 数据统计 */} -
-
-
- - 总额度 -
-
- {formatNumber(usage.usageLimit)} -
-
- -
-
- - 已使用 -
-
- {formatNumber(usage.currentUsage)} -
-
- -
-
- - 余额 -
-
- {formatNumber(usage.balance)} -
-
-
-
- ); -} diff --git a/src/components/provider-pool/VertexAISection.tsx b/src/components/provider-pool/VertexAISection.tsx deleted file mode 100644 index e631f0c27..000000000 --- a/src/components/provider-pool/VertexAISection.tsx +++ /dev/null @@ -1,304 +0,0 @@ -import { useState } from "react"; -import { - Plus, - Trash2, - Key, - Globe, - Eye, - EyeOff, - ArrowRight, -} from "lucide-react"; -import type { - VertexApiKeyEntry, - VertexModelAlias, -} from "@/lib/api/providerRuntime"; - -interface VertexAISectionProps { - entries: VertexApiKeyEntry[] | undefined; - onChange: (entries: VertexApiKeyEntry[]) => void; -} - -export function VertexAISection({ - entries = [], - onChange, -}: VertexAISectionProps) { - const primaryActionButtonClassName = - "rounded-lg border border-emerald-200 bg-[linear-gradient(135deg,#0ea5e9_0%,#14b8a6_52%,#10b981_100%)] text-sm text-white shadow-sm shadow-emerald-950/15 hover:opacity-95"; - // 确保 entries 是数组 - const safeEntries = Array.isArray(entries) ? entries : []; - const [showKeys, setShowKeys] = useState>(new Set()); - const [editingAliases, setEditingAliases] = useState(null); - const [aliasName, setAliasName] = useState(""); - const [aliasAlias, setAliasAlias] = useState(""); - - const toggleShowKey = (id: string) => { - const newSet = new Set(showKeys); - if (newSet.has(id)) { - newSet.delete(id); - } else { - newSet.add(id); - } - setShowKeys(newSet); - }; - - const addEntry = () => { - const newEntry: VertexApiKeyEntry = { - id: `vertex-${Date.now()}`, - api_key: "", - base_url: null, - models: [], - proxy_url: null, - disabled: false, - }; - onChange([...safeEntries, newEntry]); - }; - - const updateEntry = (id: string, updates: Partial) => { - onChange(safeEntries.map((e) => (e.id === id ? { ...e, ...updates } : e))); - }; - - const removeEntry = (id: string) => { - onChange(safeEntries.filter((e) => e.id !== id)); - }; - - const addAlias = (id: string) => { - if (!aliasName.trim() || !aliasAlias.trim()) return; - const entry = safeEntries.find((e) => e.id === id); - if (entry) { - const newAlias: VertexModelAlias = { - name: aliasName.trim(), - alias: aliasAlias.trim(), - }; - const currentModels = Array.isArray(entry.models) ? entry.models : []; - updateEntry(id, { - models: [...currentModels, newAlias], - }); - setAliasName(""); - setAliasAlias(""); - } - }; - - const removeAlias = (id: string, aliasToRemove: string) => { - const entry = safeEntries.find((e) => e.id === id); - if (entry) { - const currentModels = Array.isArray(entry.models) ? entry.models : []; - updateEntry(id, { - models: currentModels.filter((m) => m.alias !== aliasToRemove), - }); - } - }; - - return ( -
-
-
- -
-

Vertex AI

-

- 配置 Google Vertex AI API Key 和模型别名 -

-
-
- -
- - {safeEntries.length === 0 ? ( -
-

暂无 Vertex AI 凭证

-

点击上方"添加"按钮添加凭证

-
- ) : ( -
- {safeEntries.map((entry) => { - // 确保 models 是数组 - const models = Array.isArray(entry.models) ? entry.models : []; - return ( -
- {/* Header */} -
- - {entry.id} - -
- - -
-
- - {/* API Key */} -
- -
- - updateEntry(entry.id, { api_key: e.target.value }) - } - placeholder="vk-..." - className="w-full px-3 py-1.5 pr-10 rounded border bg-background text-sm font-mono" - /> - -
-
- - {/* Base URL */} -
- - - updateEntry(entry.id, { - base_url: e.target.value || null, - }) - } - placeholder="https://example.com/api" - className="w-full px-3 py-1.5 rounded border bg-background text-sm" - /> -
- - {/* Proxy URL */} -
- - - updateEntry(entry.id, { - proxy_url: e.target.value || null, - }) - } - placeholder="socks5://127.0.0.1:1080" - className="w-full px-3 py-1.5 rounded border bg-background text-sm" - /> -
- - {/* Model Aliases */} -
- -
- {models.map((model) => ( -
- {model.alias} - - - {model.name} - - -
- ))} -
- {editingAliases === entry.id ? ( -
-
- setAliasAlias(e.target.value)} - placeholder="客户端别名" - className="flex-1 px-2 py-1 rounded border bg-background text-sm" - /> - - setAliasName(e.target.value)} - placeholder="上游模型名" - className="flex-1 px-2 py-1 rounded border bg-background text-sm" - /> -
-
- - -
-
- ) : ( - - )} -

- 将客户端请求的模型名映射到上游实际模型名 -

-
-
- ); - })} -
- )} -
- ); -} diff --git a/src/components/provider-pool/api-key/README.md b/src/components/provider-pool/api-key/README.md index b96a76f65..eeb3f5b84 100644 --- a/src/components/provider-pool/api-key/README.md +++ b/src/components/provider-pool/api-key/README.md @@ -46,7 +46,7 @@ ```tsx import { ApiKeyProviderSection } from "@/components/provider-pool/api-key"; -function ProviderPoolPage() { +function ProviderSettingsPage() { return (
diff --git a/src/components/provider-pool/credential-forms/AntigravityForm.tsx b/src/components/provider-pool/credential-forms/AntigravityForm.tsx deleted file mode 100644 index fb584bda1..000000000 --- a/src/components/provider-pool/credential-forms/AntigravityForm.tsx +++ /dev/null @@ -1,152 +0,0 @@ -/** - * Antigravity 凭证添加表单 - * 支持 Google OAuth 登录和文件导入两种模式 - */ - -import { useState, useEffect } from "react"; -import { onAntigravityAuthUrl } from "@/lib/api/providerAuthEvents"; -import { providerPoolApi } from "@/lib/api/providerPool"; -import { ModeSelector } from "./ModeSelector"; -import { FileImportForm } from "./FileImportForm"; -import { OAuthUrlDisplay } from "./OAuthUrlDisplay"; - -interface AntigravityFormProps { - name: string; - credsFilePath: string; - setCredsFilePath: (path: string) => void; - projectId: string; - setProjectId: (id: string) => void; - onSelectFile: () => void; - loading: boolean; - setLoading: (loading: boolean) => void; - setError: (error: string | null) => void; - onSuccess: () => void; -} - -export function AntigravityForm({ - name, - credsFilePath, - setCredsFilePath, - projectId, - setProjectId, - onSelectFile, - loading: _loading, - setLoading, - setError, - onSuccess, -}: AntigravityFormProps) { - const [mode, setMode] = useState<"login" | "file">("login"); - const [authUrl, setAuthUrl] = useState(null); - const [waitingForCallback, setWaitingForCallback] = useState(false); - - // 监听后端发送的授权 URL 事件 - useEffect(() => { - let unlisten: (() => void) | undefined; - - const setupListener = async () => { - unlisten = await onAntigravityAuthUrl((payload) => { - setAuthUrl(payload.auth_url); - }); - }; - - setupListener(); - - return () => { - if (unlisten) unlisten(); - }; - }, []); - - // 获取授权 URL 并启动服务器等待回调 - const handleGetAuthUrl = async () => { - setLoading(true); - setError(null); - setAuthUrl(null); - setWaitingForCallback(true); - - try { - const trimmedName = name.trim() || undefined; - await providerPoolApi.getAntigravityAuthUrlAndWait(trimmedName, false); - onSuccess(); - } catch (e) { - const errorMsg = e instanceof Error ? e.message : String(e); - setError(errorMsg); - setWaitingForCallback(false); - } finally { - setLoading(false); - } - }; - - // 文件导入提交 - const handleFileSubmit = async () => { - if (!credsFilePath) { - setError("请选择凭证文件"); - return; - } - - setLoading(true); - setError(null); - - try { - const trimmedName = name.trim() || undefined; - await providerPoolApi.addAntigravityOAuth( - credsFilePath, - projectId.trim() || undefined, - trimmedName, - ); - onSuccess(); - } catch (e) { - setError(e instanceof Error ? e.message : String(e)); - } finally { - setLoading(false); - } - }; - - return { - mode, - authUrl, - waitingForCallback, - handleGetAuthUrl, - handleFileSubmit, - render: () => ( - <> - - - {mode === "login" ? ( -
-
-

- 点击下方按钮获取授权 - URL,然后复制到浏览器(支持指纹浏览器)完成登录。 -

-

- 授权成功后,凭证将自动保存并添加到凭证池。 -

-
- - -
- ) : ( - - )} - - ), - }; -} diff --git a/src/components/provider-pool/credential-forms/ClaudeOAuthForm.tsx b/src/components/provider-pool/credential-forms/ClaudeOAuthForm.tsx deleted file mode 100644 index a4fda1ef0..000000000 --- a/src/components/provider-pool/credential-forms/ClaudeOAuthForm.tsx +++ /dev/null @@ -1,260 +0,0 @@ -/** - * Claude OAuth 凭证添加表单 - * 支持三种模式: - * 1. OAuth 登录 - 通过授权 URL 手动复制授权码 - * 2. Cookie 授权 - 使用 sessionKey 自动完成 OAuth 流程 - * 3. 文件导入 - 导入已有的凭证文件 - */ - -import { useState, useEffect } from "react"; -import { onClaudeOAuthAuthUrl } from "@/lib/api/providerAuthEvents"; -import { Cookie, Key, FileJson } from "lucide-react"; -import { providerPoolApi } from "@/lib/api/providerPool"; -import { FileImportForm } from "./FileImportForm"; -import { OAuthUrlDisplay } from "./OAuthUrlDisplay"; - -interface ClaudeOAuthFormProps { - name: string; - credsFilePath: string; - setCredsFilePath: (path: string) => void; - onSelectFile: () => void; - loading: boolean; - setLoading: (loading: boolean) => void; - setError: (error: string | null) => void; - onSuccess: () => void; -} - -type AuthMode = "login" | "cookie" | "file"; - -export function ClaudeOAuthForm({ - name, - credsFilePath, - setCredsFilePath, - onSelectFile, - loading: _loading, - setLoading, - setError, - onSuccess, -}: ClaudeOAuthFormProps) { - const [mode, setMode] = useState("cookie"); - const [authUrl, setAuthUrl] = useState(null); - const [waitingForCallback, setWaitingForCallback] = useState(false); - const [sessionKey, setSessionKey] = useState(""); - const [isSetupToken, setIsSetupToken] = useState(false); - - // 监听后端发送的授权 URL 事件 - useEffect(() => { - let unlisten: (() => void) | undefined; - - const setupListener = async () => { - unlisten = await onClaudeOAuthAuthUrl((payload) => { - setAuthUrl(payload.auth_url); - }); - }; - - setupListener(); - - return () => { - if (unlisten) unlisten(); - }; - }, []); - - // 获取授权 URL 并启动服务器等待回调 - const handleGetAuthUrl = async () => { - setLoading(true); - setError(null); - setAuthUrl(null); - setWaitingForCallback(true); - - try { - const trimmedName = name.trim() || undefined; - await providerPoolApi.getClaudeOAuthAuthUrlAndWait(trimmedName); - onSuccess(); - } catch (e) { - const errorMsg = e instanceof Error ? e.message : String(e); - setError(errorMsg); - setWaitingForCallback(false); - } finally { - setLoading(false); - } - }; - - // Cookie 自动授权 - const handleCookieSubmit = async () => { - if (!sessionKey.trim()) { - setError("请输入 sessionKey"); - return; - } - - setLoading(true); - setError(null); - - try { - const trimmedName = name.trim() || undefined; - await providerPoolApi.claudeOAuthWithCookie( - sessionKey.trim(), - isSetupToken, - trimmedName, - ); - onSuccess(); - } catch (e) { - setError(e instanceof Error ? e.message : String(e)); - } finally { - setLoading(false); - } - }; - - // 文件导入提交 - const handleFileSubmit = async () => { - if (!credsFilePath) { - setError("请选择凭证文件"); - return; - } - - setLoading(true); - setError(null); - - try { - const trimmedName = name.trim() || undefined; - await providerPoolApi.addClaudeOAuth(credsFilePath, trimmedName); - onSuccess(); - } catch (e) { - setError(e instanceof Error ? e.message : String(e)); - } finally { - setLoading(false); - } - }; - - // 模式选择器 - const renderModeSelector = () => ( -
- - - -
- ); - - // Cookie 授权表单 - const renderCookieForm = () => ( -
-
-

- 使用浏览器 Cookie 中的 sessionKey 自动完成 OAuth - 授权,无需手动复制授权码。 -

-

- 获取方式:在 claude.ai 登录后,打开开发者工具 → Application → Cookies - → 复制 sessionKey 的值 -

-
- -
- -