From 9d71e355da48e03b44ffa4f27d08da6d48229579 Mon Sep 17 00:00:00 2001 From: SaladDay Date: Mon, 3 Aug 2026 09:21:23 +0000 Subject: [PATCH] feat(pi): add first-class Coding Agent support --- docs/pi-extensions-design-appendix-zh.md | 85 + scripts/generate-pi-native-oracle.mjs | 1354 ++++++++ scripts/pi-transport-capture.mjs | 776 +++++ src-tauri/Cargo.lock | 13 +- src-tauri/Cargo.toml | 7 +- src-tauri/src/app_config.rs | 65 +- src-tauri/src/architecture_tests.rs | 1001 ++++++ src-tauri/src/commands/config.rs | 14 + src-tauri/src/commands/failover.rs | 173 + src-tauri/src/commands/import_export.rs | 104 +- src-tauri/src/commands/misc.rs | 41 +- src-tauri/src/commands/mod.rs | 2 + src-tauri/src/commands/pi.rs | 76 + src-tauri/src/commands/prompt.rs | 64 +- src-tauri/src/commands/proxy.rs | 58 + src-tauri/src/commands/s3_sync.rs | 36 +- src-tauri/src/commands/settings.rs | 61 +- src-tauri/src/commands/skill.rs | 9 + src-tauri/src/commands/sync_support.rs | 12 +- src-tauri/src/commands/webdav_sync.rs | 36 +- src-tauri/src/config.rs | 130 +- src-tauri/src/database/dao/mod.rs | 3 + src-tauri/src/database/dao/pi_catalog.rs | 276 ++ src-tauri/src/database/dao/pi_projections.rs | 204 ++ src-tauri/src/database/dao/prompts.rs | 204 +- src-tauri/src/database/dao/provider_write.rs | 89 +- src-tauri/src/database/dao/providers.rs | 28 + .../src/database/dao/skill_deployments.rs | 290 ++ src-tauri/src/database/dao/skills.rs | 97 +- src-tauri/src/database/mod.rs | 2 + src-tauri/src/deeplink/mod.rs | 4 + src-tauri/src/deeplink/parser.rs | 31 +- src-tauri/src/deeplink/provider.rs | 40 + src-tauri/src/deeplink/tests.rs | 65 + src-tauri/src/lib.rs | 76 +- src-tauri/src/pi_config/composer.rs | 979 ++++++ src-tauri/src/pi_config/document.rs | 1250 +++++++ src-tauri/src/pi_config/gateway.rs | 1458 ++++++++ src-tauri/src/pi_config/mod.rs | 218 ++ src-tauri/src/pi_config/model.rs | 1144 +++++++ src-tauri/src/pi_config/native.rs | 1107 ++++++ .../native_inspection_certification.rs | 834 +++++ src-tauri/src/pi_config/native_settings.rs | 324 ++ src-tauri/src/pi_config/raw_schema.rs | 1042 ++++++ src-tauri/src/pi_config/shared_file.rs | 1770 ++++++++++ src-tauri/src/prompt.rs | 2 +- src-tauri/src/prompt_files.rs | 2 + src-tauri/src/provider.rs | 4 + src-tauri/src/proxy/handler_config.rs | 7 + src-tauri/src/proxy/handlers.rs | 12 +- src-tauri/src/proxy/mod.rs | 2 + src-tauri/src/proxy/pi_handler.rs | 1739 ++++++++++ src-tauri/src/proxy/pi_runtime.rs | 1402 ++++++++ src-tauri/src/proxy/provider_router.rs | 16 +- src-tauri/src/proxy/providers/mod.rs | 12 +- src-tauri/src/proxy/response_processor.rs | 21 +- src-tauri/src/proxy/server.rs | 22 + src-tauri/src/proxy/types.rs | 11 + src-tauri/src/proxy/usage/calculator.rs | 52 +- src-tauri/src/proxy/usage/logger.rs | 31 +- src-tauri/src/proxy/usage/mod.rs | 3 + src-tauri/src/proxy/usage/semantics.rs | 32 + src-tauri/src/services/config.rs | 4 + src-tauri/src/services/mcp.rs | 48 +- src-tauri/src/services/mod.rs | 3 + src-tauri/src/services/pi_catalog.rs | 2118 ++++++++++++ src-tauri/src/services/pi_prompt_files.rs | 433 +++ src-tauri/src/services/prompt.rs | 1178 +++++++ src-tauri/src/services/provider/live.rs | 26 + src-tauri/src/services/provider/mod.rs | 341 +- src-tauri/src/services/proxy.rs | 2156 +++++++++++- src-tauri/src/services/skill.rs | 720 +++- src-tauri/src/services/skill_deployment.rs | 1353 ++++++++ src-tauri/src/services/sql_helpers.rs | 37 +- src-tauri/src/services/stream_check.rs | 56 + src-tauri/src/services/usage_stats.rs | 86 +- src-tauri/src/session_manager/mod.rs | 10 +- .../src/session_manager/providers/mod.rs | 1 + src-tauri/src/session_manager/providers/pi.rs | 731 ++++ src-tauri/src/session_manager/terminal/mod.rs | 2 +- src-tauri/src/settings.rs | 283 +- src-tauri/src/store.rs | 1 + src/App.tsx | 115 +- src/components/AppSwitcher.tsx | 15 +- src/components/DeepLinkImportDialog.tsx | 10 + src/components/UsageScriptModal.tsx | 10 + src/components/common/AppToggleGroup.tsx | 32 +- .../prompts/PiNativePromptResources.tsx | 411 +++ src/components/prompts/PromptFormPanel.tsx | 1 + src/components/prompts/PromptPanel.tsx | 54 +- .../providers/AddProviderDialog.tsx | 9 +- .../providers/EditProviderDialog.tsx | 10 +- .../providers/PiNativeCatalogPanel.tsx | 335 ++ src/components/providers/ProviderCard.tsx | 22 +- src/components/providers/ProviderList.tsx | 2 +- .../providers/forms/EndpointSpeedTest.tsx | 1 + .../providers/forms/PiProviderForm.tsx | 589 ++++ .../providers/forms/ProviderForm.tsx | 4 + src/components/proxy/FailoverToggle.tsx | 10 +- src/components/proxy/ProxyPanel.tsx | 133 +- src/components/proxy/ProxyToggle.tsx | 60 +- .../sessions/SessionManagerPage.tsx | 55 +- src/components/settings/AboutSection.tsx | 11 +- .../settings/AppVisibilitySettings.tsx | 13 +- src/components/settings/DirectorySettings.tsx | 13 + .../settings/ImportExportSection.tsx | 5 +- src/components/settings/ProxyTabContent.tsx | 13 +- src/components/settings/SettingsPage.tsx | 1 + src/components/skills/UnifiedSkillsPanel.tsx | 82 +- src/components/usage/PricingConfigPanel.tsx | 6 +- src/components/usage/UsageDashboard.tsx | 1 + src/components/usage/UsageHero.tsx | 26 +- src/config/appConfig.tsx | 55 +- src/hooks/useDirectorySettings.ts | 15 +- src/hooks/useImportExport.ts | 5 + src/hooks/usePromptActions.ts | 27 +- src/hooks/useProxyStatus.ts | 29 +- src/hooks/useSettings.ts | 19 +- src/hooks/useSettingsForm.ts | 4 + src/hooks/useSkills.ts | 30 + src/i18n/locales/en.json | 152 +- src/i18n/locales/ja.json | 152 +- src/i18n/locales/zh-TW.json | 152 +- src/i18n/locales/zh.json | 152 +- src/icons/extracted/index.ts | 1 + src/lib/api/deeplink.ts | 4 +- src/lib/api/index.ts | 8 +- src/lib/api/pi.ts | 84 + src/lib/api/prompts.ts | 83 + src/lib/api/settings.ts | 1 + src/lib/api/skills.ts | 27 +- src/lib/api/types.ts | 3 +- src/lib/query/failover.ts | 10 +- src/lib/query/mutations.ts | 32 +- src/types.ts | 23 + src/types/proxy.ts | 2 + src/types/usage.ts | 47 +- tests/components/ImportExportSection.test.tsx | 17 + .../PiNativePromptResources.test.tsx | 144 + tests/components/PiProviderForm.test.tsx | 162 + .../PromptPanel.piReconciliation.test.tsx | 67 + tests/components/ProxyTabContent.apps.test.ts | 3 +- tests/components/SessionManagerPage.test.tsx | 18 + tests/components/UnifiedSkillsPanel.test.tsx | 74 +- tests/config/localeCoverage.test.ts | 100 + tests/fixtures/pi/module-boundaries-v1.json | 13 + .../pi/native-oracle/composer-oracle-v1.json | 2186 ++++++++++++ .../pi/native-oracle/field-coverage-v1.json | 3000 +++++++++++++++++ .../pi/native-oracle/provenance-v1.json | 96 + .../provider-schema.snapshot.json | 1708 ++++++++++ .../pi/native-oracle/raw-oracle-v1.json | 723 ++++ .../pi/native-oracle/transport-oracle-v1.json | 82 + tests/hooks/useDirectorySettings.test.tsx | 2 + tests/hooks/useImportSkillsFromApps.test.tsx | 1 + tests/hooks/useProxyStatus.test.tsx | 70 + tests/hooks/useSettings.test.tsx | 29 +- tests/hooks/useSettingsForm.test.tsx | 4 + tests/msw/state.ts | 4 + tests/types/usage.test.ts | 17 + 159 files changed, 39663 insertions(+), 632 deletions(-) create mode 100644 docs/pi-extensions-design-appendix-zh.md create mode 100644 scripts/generate-pi-native-oracle.mjs create mode 100644 scripts/pi-transport-capture.mjs create mode 100644 src-tauri/src/architecture_tests.rs create mode 100644 src-tauri/src/commands/pi.rs create mode 100644 src-tauri/src/database/dao/pi_catalog.rs create mode 100644 src-tauri/src/database/dao/pi_projections.rs create mode 100644 src-tauri/src/database/dao/skill_deployments.rs create mode 100644 src-tauri/src/pi_config/composer.rs create mode 100644 src-tauri/src/pi_config/document.rs create mode 100644 src-tauri/src/pi_config/gateway.rs create mode 100644 src-tauri/src/pi_config/mod.rs create mode 100644 src-tauri/src/pi_config/model.rs create mode 100644 src-tauri/src/pi_config/native.rs create mode 100644 src-tauri/src/pi_config/native_inspection_certification.rs create mode 100644 src-tauri/src/pi_config/native_settings.rs create mode 100644 src-tauri/src/pi_config/raw_schema.rs create mode 100644 src-tauri/src/pi_config/shared_file.rs create mode 100644 src-tauri/src/proxy/pi_handler.rs create mode 100644 src-tauri/src/proxy/pi_runtime.rs create mode 100644 src-tauri/src/proxy/usage/semantics.rs create mode 100644 src-tauri/src/services/pi_catalog.rs create mode 100644 src-tauri/src/services/pi_prompt_files.rs create mode 100644 src-tauri/src/services/skill_deployment.rs create mode 100644 src-tauri/src/session_manager/providers/pi.rs create mode 100644 src/components/prompts/PiNativePromptResources.tsx create mode 100644 src/components/providers/PiNativeCatalogPanel.tsx create mode 100644 src/components/providers/forms/PiProviderForm.tsx create mode 100644 src/lib/api/pi.ts create mode 100644 tests/components/PiNativePromptResources.test.tsx create mode 100644 tests/components/PiProviderForm.test.tsx create mode 100644 tests/components/PromptPanel.piReconciliation.test.tsx create mode 100644 tests/config/localeCoverage.test.ts create mode 100644 tests/fixtures/pi/module-boundaries-v1.json create mode 100644 tests/fixtures/pi/native-oracle/composer-oracle-v1.json create mode 100644 tests/fixtures/pi/native-oracle/field-coverage-v1.json create mode 100644 tests/fixtures/pi/native-oracle/provenance-v1.json create mode 100644 tests/fixtures/pi/native-oracle/provider-schema.snapshot.json create mode 100644 tests/fixtures/pi/native-oracle/raw-oracle-v1.json create mode 100644 tests/fixtures/pi/native-oracle/transport-oracle-v1.json create mode 100644 tests/types/usage.test.ts diff --git a/docs/pi-extensions-design-appendix-zh.md b/docs/pi-extensions-design-appendix-zh.md new file mode 100644 index 000000000..58fc8ea6b --- /dev/null +++ b/docs/pi-extensions-design-appendix-zh.md @@ -0,0 +1,85 @@ +# Pi Extensions 一等支持设计附录 + +> 状态:仅设计,不含实现。主工程不实现 extensions;themes 明确不在范围内。 +> authority:任何上游路径、加载顺序、启用规则或 tool 组合语义,在实施前都必须 +> 由 pin `ab366ebe94cacd419d986be454f12b1b9913aaca` 的 oracle 或 +> `scripts/pi-transport-capture.mjs` 实际执行确认,本文不以源码阅读代替证据。 + +## 目标 + +未来让用户在 cc-switch 中观察、导入、启用和停用 Pi extensions,同时满足: + +1. 原生文件/目录是真相,`exists = active`,不制造与 Pi 分叉的 enabled 影子状态; +2. 只接管 cc-switch 明确拥有或用户显式、可验证采用的内容; +3. extension 提供的 tools 与 pinned core tools 分开展示,不伪装为 MCP server; +4. 所有写操作可并发检测、可补偿,portable import 不覆盖目标设备的未知资产; +5. 复用前置 C inspection 与当前 Skill/Prompt 的共享文件、fingerprint、ownership + 和协调器原语。 + +## 模块边界 + +```text +PiExtensionInspector (只读、oracle 驱动) + │ + ├── NativeExtensionObservation + │ path / fingerprint / manifest / contributed tools + │ validity / reasons / ownership + │ + └── PiExtensionCoordinator (唯一写入口) + ├── exact-content adoption + ├── CAS + atomic replace + ├── ownership ledger transaction + ├── catalog epoch / UI invalidation + └── compensation on partial failure +``` + +- `PiExtensionInspector` 只负责原生观察和结构化诊断;不得写数据库或文件。 +- `PiExtensionCoordinator` 是唯一写入口;命令、deeplink、portable reconcile 与 + UI mutation 都调用它,禁止各自复制目录。 +- 通用 `shared_file`、Skill tree fingerprint 与 ownership ledger 可复用; + extension-specific manifest/加载规则必须先新增捕获向量,不能借 Skill 规则猜测。 +- gateway 只消费 coordinator 发布后的 immutable runtime snapshot;extension + 不能在请求中途直接改 candidate/header 计划。 + +## 状态模型 + +建议公开三个正交维度: + +- `discovery`: `absent | active | invalid | unknown`,只来自 native observation; +- `ownership`: `external | adoptable_exact | managed | conflict`; +- `capability`: `inspectable | manageable | unsupported | unknown`。 + +不得增加独立 `enabled` 布尔值。用户点击“停用”时,语义是对受管原生资产执行可逆 +移除;外部资产只能显式采用后再管理。内容变化导致 fingerprint 不匹配时进入 +`conflict`,不得覆盖。 + +## Tools 与 MCP + +- capture 已确认 pinned core tool inventory 为 + `bash/edit/find/grep/ls/read/write`;未来 capture 应分别记录每个 extension + 注入前后的 tool inventory 与来源。 +- UI 将 tools 按 `core` / `extension:` 分组,并展示冲突与覆盖次序的实测结果。 +- MCP 页面仍不为 Pi 建虚假 registry。即使某个 extension 通过自身机制连接外部 + tool,也属于 extension capability,除非未来 pinned Pi 真正提供 MCP registry + 且由新证据和契约明确升级。 + +## Portable 与冲突策略 + +- 备份只携带 cc-switch 拥有的 extension 描述、内容 hash 和期望状态,不携带绝对 + 目录、设备 token 或未知外部目录。 +- 导入先观察目标设备;missing 可部署,exact 可采用,different 必须 conflict, + 绝不“最后写入者获胜”。 +- 多 extension 贡献同名 tool、command 或资源时 fail-closed;只有 oracle/capture + 证明 Pi 的确定 precedence 且产品明确展示该覆盖时,才允许自动解析。 + +## 实施前验收 + +1. 扩展 transport capture:发现路径、空/损坏 manifest、启停、重复 ID、资源覆盖、 + tool inventory、相对路径和 symlink 负例。 +2. 冻结 schema/oracle provenance,建立 lossless raw observation;未知形状为 + `unknown`,不得整目录连坐隐藏合法兄弟。 +3. 服务级测试覆盖显式采用、并发 CAS loser、写后补偿、portable reconcile、 + 外部冲突、目录越界与 UI 的 `exists = active`。 +4. 中英文 UI 与可访问性完成后,再进入独立实现与盲审。 + +Themes 不与 extensions 共用该项目:其资源语义、预览与安全面另行立项。 diff --git a/scripts/generate-pi-native-oracle.mjs b/scripts/generate-pi-native-oracle.mjs new file mode 100644 index 000000000..a577f8f17 --- /dev/null +++ b/scripts/generate-pi-native-oracle.mjs @@ -0,0 +1,1354 @@ +#!/usr/bin/env node + +/** + * Generate the vendored Pi raw-schema and composer oracles. + * + * The composer expectations are not a local reimplementation. This harness + * bundles and executes the pinned upstream `composeModelProvider` function. + * Two inert transport shims satisfy imports that are unreachable while + * composing credential-blind custom catalog models; the exact shim bytes and + * harness bytes are recorded in provenance. + */ + +import { createHash } from "node:crypto"; +import { execFileSync } from "node:child_process"; +import { + mkdirSync, + mkdtempSync, + readFileSync, + rmSync, + writeFileSync, +} from "node:fs"; +import { createRequire } from "node:module"; +import { dirname, join, resolve } from "node:path"; +import { fileURLToPath, pathToFileURL } from "node:url"; + +const [piRootArgument, outputArgument] = process.argv.slice(2); +if (!piRootArgument || !outputArgument) { + throw new Error( + "usage: generate-pi-native-oracle.mjs ", + ); +} + +const piRoot = resolve(piRootArgument); +const outputDirectory = resolve(outputArgument); +const repositoryRoot = resolve(dirname(fileURLToPath(import.meta.url)), ".."); +const generatorRelativePath = "scripts/generate-pi-native-oracle.mjs"; +const generatorPath = join(repositoryRoot, generatorRelativePath); +const requireFromPi = createRequire(join(piRoot, "package.json")); +const { Type } = requireFromPi("typebox"); +const { Check } = requireFromPi("typebox/value"); +const { buildSync, version: esbuildVersion } = requireFromPi("esbuild"); +const typeboxVersion = JSON.parse( + readFileSync( + join(dirname(dirname(requireFromPi.resolve("typebox"))), "package.json"), + "utf8", + ), +).version; + +const expectedPiCommit = "ab366ebe94cacd419d986be454f12b1b9913aaca"; +const piRepository = "https://github.com/earendil-works/pi.git"; +const modelConfigRelativePath = + "packages/coding-agent/src/core/model-config.ts"; +const composerRelativePath = + "packages/coding-agent/src/core/provider-composer.ts"; +const resolverRelativePath = + "packages/coding-agent/src/core/resolve-config-value.ts"; + +const modelConfigPath = join(piRoot, modelConfigRelativePath); +const composerPath = join(piRoot, composerRelativePath); +const resolverPath = join(piRoot, resolverRelativePath); +const modelConfigSource = readFileSync(modelConfigPath, "utf8"); +const composerSource = readFileSync(composerPath, "utf8"); +const resolverSource = readFileSync(resolverPath, "utf8"); +const generatorSource = readFileSync(generatorPath, "utf8"); +const actualPiCommit = execFileSync( + "git", + ["-C", piRoot, "rev-parse", "HEAD"], + { encoding: "utf8" }, +).trim(); + +if (actualPiCommit !== expectedPiCommit) { + throw new Error( + `Pi checkout is ${actualPiCommit}, expected ${expectedPiCommit}`, + ); +} + +function sha256(content) { + return createHash("sha256").update(content).digest("hex"); +} + +function writeJson(filename, value) { + const content = `${JSON.stringify(value, null, 2)}\n`; + writeFileSync(join(outputDirectory, filename), content); + return { content, sha256: sha256(content) }; +} + +function extractModelsSchema(source) { + const start = source.indexOf("const PercentileCutoffsSchema"); + const end = source.indexOf("const validateModelsConfig"); + if (start < 0 || end <= start) { + throw new Error("pinned model-config.ts schema block was not found"); + } + const schemaProgram = source.slice(start, end); + return Function( + "Type", + `${schemaProgram}; return ModelsConfigSchema;`, + )(Type); +} + +const modelsSchema = extractModelsSchema(modelConfigSource); +const providerSchema = + modelsSchema.properties.providers.patternProperties["^.*$"]; + +const allThinkingLevels = { + off: null, + minimal: "minimal-effort", + low: "low-effort", + medium: "medium-effort", + high: "high-effort", + xhigh: "xhigh-effort", + max: "max-effort", +}; + +const allCostFields = { + input: 0.11, + output: 0.22, + cacheRead: 0.033, + cacheWrite: 0.044, + tiers: [ + { + inputTokensAbove: 1000.5, + input: 0.55, + output: 0.66, + cacheRead: 0.077, + cacheWrite: 0.088, + }, + ], +}; + +const allCompatFields = { + supportsStore: true, + supportsDeveloperRole: true, + supportsReasoningEffort: true, + supportsUsageInStreaming: true, + maxTokensField: "max_completion_tokens", + requiresToolResultName: true, + requiresAssistantAfterToolResult: true, + requiresThinkingAsText: true, + requiresReasoningContentOnAssistantMessages: true, + thinkingFormat: "chat-template", + chatTemplateKwargs: { + stringValue: "literal", + numberValue: 1.25, + booleanValue: true, + nullValue: null, + variableValue: { + $var: "thinking.effort", + omitWhenOff: true, + }, + }, + cacheControlFormat: "anthropic", + openRouterRouting: { + allow_fallbacks: true, + require_parameters: true, + data_collection: "deny", + zdr: true, + enforce_distillable_text: true, + order: ["provider-a", "provider-b"], + only: ["provider-a"], + ignore: ["provider-z"], + quantizations: ["fp8"], + sort: { + by: "price", + partition: null, + }, + max_price: { + prompt: 1.1, + completion: "2.2", + image: 3.3, + audio: "4.4", + request: 5.5, + }, + preferred_min_throughput: { + p50: 10.5, + p75: 9.5, + p90: 8.5, + p99: 7.5, + }, + preferred_max_latency: { + p50: 100.5, + p75: 200.5, + p90: 300.5, + p99: 400.5, + }, + }, + vercelGatewayRouting: { + only: ["provider-a"], + order: ["provider-a", "provider-b"], + }, + supportsOpenAIGrammarTools: true, + supportsStrictMode: true, + sendSessionAffinityHeaders: true, + deferredToolsMode: "kimi", + sessionAffinityFormat: "openrouter", + supportsLongCacheRetention: true, + supportsToolSearch: true, + supportsEagerToolInputStreaming: true, + supportsCacheControlOnTools: true, + supportsTemperature: true, + forceAdaptiveThinking: true, + allowEmptySignature: true, + supportsStrictTools: true, + supportsToolReferences: true, +}; + +const rawInputs = [ + { + id: "all-schema-fields-valid", + input: { + name: "All Fields Provider", + baseUrl: "https://all-fields.example/v1", + apiKey: "literal-all-fields-key", + api: "openai-responses", + oauth: "radius", + headers: { + "x-provider-field": "provider-value", + }, + compat: allCompatFields, + authHeader: true, + models: [ + { + id: "all-fields-model", + name: "All Fields Model", + api: "anthropic-messages", + baseUrl: "https://all-fields-model.example/v1", + reasoning: true, + thinkingLevelMap: allThinkingLevels, + input: ["text", "image"], + cost: allCostFields, + contextWindow: 128000.5, + maxTokens: 16384.25, + headers: { + "x-model-field": "model-value", + }, + compat: allCompatFields, + }, + ], + modelOverrides: { + "all-fields-model": { + name: "All Fields Override", + reasoning: false, + thinkingLevelMap: allThinkingLevels, + input: ["image", "text"], + cost: allCostFields, + contextWindow: 256000.75, + maxTokens: 32768.5, + headers: { + "x-override-field": "override-value", + }, + compat: allCompatFields, + }, + }, + }, + }, + { id: "empty-provider-object", input: {} }, + { id: "null-provider", input: null }, + { + id: "additional-provider-property", + input: { futureProviderField: { nested: true } }, + }, + { id: "empty-present-base-url", input: { baseUrl: "" } }, + { id: "non-url-base-url-is-raw-string", input: { baseUrl: "not a URL" } }, + { id: "radius-oauth", input: { oauth: "radius" } }, + { id: "unknown-oauth-literal", input: { oauth: "other" } }, + { id: "models-null", input: { models: null } }, + { + id: "model-missing-id", + input: { models: [{ api: "openai-responses" }] }, + }, + { id: "model-empty-id", input: { models: [{ id: "" }] } }, + { + id: "integer-model-numbers", + input: { + models: [{ id: "integer", contextWindow: 128000, maxTokens: 16384 }], + }, + }, + { + id: "fractional-model-numbers", + input: { + models: [ + { + id: "fractional", + contextWindow: 128000.5, + maxTokens: 16384.25, + }, + ], + }, + }, + { + id: "negative-number-is-still-typebox-number", + input: { models: [{ id: "negative", maxTokens: -1.5 }] }, + }, + { + id: "string-is-not-number", + input: { models: [{ id: "string-limit", maxTokens: "16384" }] }, + }, + { + id: "complete-cost-with-fractions", + input: { + models: [ + { + id: "priced", + cost: { + input: 0.1, + output: 0.2, + cacheRead: 0.03, + cacheWrite: 0.04, + tiers: [ + { + inputTokensAbove: 1000.5, + input: 0.5, + output: 0.6, + cacheRead: 0.07, + cacheWrite: 0.08, + }, + ], + }, + }, + ], + }, + }, + { + id: "incomplete-model-cost", + input: { + models: [ + { + id: "priced", + cost: { input: 0.1, output: 0.2, cacheRead: 0.03 }, + }, + ], + }, + }, + { + id: "compat-openrouter-recursive-valid", + input: { + compat: { + thinkingFormat: "chat-template", + chatTemplateKwargs: { + temperature: 0.25, + thinking: { + $var: "thinking.effort", + omitWhenOff: true, + }, + nullable: null, + }, + openRouterRouting: { + data_collection: "deny", + sort: { by: "price", partition: null }, + max_price: { prompt: 1.5, completion: "2.0" }, + preferred_min_throughput: { p50: 10.5, p99: 2 }, + }, + }, + }, + }, + { + id: "compat-nested-union-invalid-in-every-branch", + input: { + compat: { + openRouterRouting: { + preferred_min_throughput: { p50: "fast" }, + }, + supportsToolSearch: "yes", + supportsTemperature: 1, + }, + }, + }, + { + id: "compat-union-additional-field", + input: { compat: { futureCompatField: [1, 2, 3] } }, + }, + { id: "compat-null", input: { compat: null } }, + { + id: "thinking-map-unknown-key-and-string-value", + input: { + models: [ + { + id: "thinking", + thinkingLevelMap: { + low: null, + max: "future-effort", + future: { nested: true }, + }, + }, + ], + }, + }, + { + id: "thinking-map-known-key-invalid-value", + input: { + models: [{ id: "thinking", thinkingLevelMap: { low: 2 } }], + }, + }, + { + id: "compat-sort-string-union-branch", + input: { compat: { openRouterRouting: { sort: "price" } } }, + }, + { + id: "compat-sort-object-union-branch", + input: { + compat: { + openRouterRouting: { sort: { by: "price", partition: null } }, + }, + }, + }, + { + id: "compat-sort-invalid-union-boundary", + input: { compat: { openRouterRouting: { sort: false } } }, + }, + { + id: "headers-record-valid", + input: { headers: { Authorization: "Bearer literal", "x-count": "2" } }, + }, + { + id: "headers-record-invalid-value", + input: { headers: { "x-count": 2 } }, + }, + { + id: "model-override-fractional-and-extra", + input: { + modelOverrides: { + "model/a": { + contextWindow: 10.5, + maxTokens: 2.25, + futureOverrideField: true, + }, + }, + }, + }, + { + id: "input-union-invalid", + input: { models: [{ id: "input", input: ["text", "audio"] }] }, + }, +]; + +const rawOracle = { + version: 1, + execution: { + engine: "typebox Value.Check", + typeboxVersion, + schemaTarget: + "ModelsConfigSchema.properties.providers.patternProperties['^.*$']", + }, + cases: rawInputs.map(({ id, input }) => ({ + id, + input, + expectedValid: Check(providerSchema, input), + })), +}; + +const composerCases = [ + { + id: "combined-all-fields-precedence", + providerId: "custom-all-fields", + input: { + name: "All Fields Provider", + baseUrl: "https://all-fields.example/v1", + apiKey: "literal-all-fields-key", + api: "openai-responses", + oauth: "radius", + headers: { + "x-provider-field": "provider-value", + }, + compat: allCompatFields, + authHeader: true, + models: [ + { + id: "all-fields-model", + name: "All Fields Model", + api: "anthropic-messages", + baseUrl: "https://all-fields-model.example/v1", + reasoning: true, + thinkingLevelMap: allThinkingLevels, + input: ["text", "image"], + cost: allCostFields, + contextWindow: 128000.5, + maxTokens: 16384.25, + headers: { + "x-model-field": "model-value", + }, + compat: allCompatFields, + }, + ], + modelOverrides: { + "all-fields-model": { + name: "All Fields Override", + reasoning: false, + thinkingLevelMap: allThinkingLevels, + input: ["image", "text"], + cost: allCostFields, + contextWindow: 256000.75, + maxTokens: 32768.5, + headers: { + "x-override-field": "override-value", + }, + compat: allCompatFields, + }, + }, + }, + }, + { + id: "provider-fields-inherited", + providerId: "custom-provider-fields", + input: { + name: "Provider Layer Name", + baseUrl: "https://provider-fields.example/v1", + apiKey: "provider-layer-key", + api: "openai-responses", + oauth: "radius", + headers: { + "x-provider-field": "provider-layer-value", + }, + compat: allCompatFields, + authHeader: true, + models: [{ id: "provider-inherited-model" }], + }, + }, + { + id: "model-fields-executed", + providerId: "custom-model-fields", + input: { + apiKey: "model-layer-key", + models: [ + { + id: "model-layer-id", + name: "Model Layer Name", + api: "anthropic-messages", + baseUrl: "https://model-fields.example/v1", + reasoning: true, + thinkingLevelMap: allThinkingLevels, + input: ["text", "image"], + cost: allCostFields, + contextWindow: 64000.5, + maxTokens: 8192.25, + headers: { + "x-model-field": "model-layer-value", + }, + compat: allCompatFields, + }, + ], + }, + }, + { + id: "override-fields-executed", + providerId: "custom-override-fields", + input: { + apiKey: "override-layer-key", + api: "openai-completions", + baseUrl: "https://override-provider.example/v1", + models: [ + { + id: "override-target", + name: "Definition Before Override", + reasoning: true, + }, + ], + modelOverrides: { + "override-target": { + name: "Override Layer Name", + reasoning: false, + thinkingLevelMap: allThinkingLevels, + input: ["image", "text"], + cost: allCostFields, + contextWindow: 96000.75, + maxTokens: 12288.5, + headers: { + "x-override-field": "override-layer-value", + }, + compat: allCompatFields, + }, + }, + }, + }, + { + id: "radius-oauth-requires-provider-base-url", + providerId: "custom-radius-missing-base", + input: { + apiKey: "radius-key", + oauth: "radius", + models: [ + { + id: "radius-model", + api: "openai-responses", + baseUrl: "https://model-only.example/v1", + }, + ], + }, + }, + { + id: "official-defaults", + providerId: "custom-defaults", + input: { + api: "openai-responses", + baseUrl: "https://default.example/v1", + apiKey: "literal-secret", + models: [{ id: "default-model" }], + }, + }, + { + id: "later-model-first-model-fallback", + providerId: "custom-fallback", + input: { + apiKey: "literal-secret", + models: [ + { + id: "first", + api: "anthropic-messages", + baseUrl: "https://first.example/v1", + }, + { id: "second" }, + ], + }, + }, + { + id: "duplicate-model-replaces-existing-slot", + providerId: "custom-duplicate", + input: { + apiKey: "literal-secret", + models: [ + { + id: "same", + name: "First", + api: "openai-completions", + baseUrl: "https://first.example/v1", + }, + { + id: "same", + name: "Second", + api: "openai-responses", + baseUrl: "https://second.example/v1", + }, + ], + }, + }, + { + id: "provider-model-override-precedence", + providerId: "custom-precedence", + input: { + api: "openai-completions", + baseUrl: "https://provider.example/v1", + apiKey: "literal-secret", + headers: { layer: "provider", providerOnly: "yes" }, + compat: { + supportsDeveloperRole: true, + openRouterRouting: { zdr: true }, + }, + models: [ + { + id: "precedence", + name: "Definition", + api: "openai-responses", + baseUrl: "https://model.example/v1", + reasoning: false, + thinkingLevelMap: { high: "model-high", future: "model-future" }, + contextWindow: 1000.5, + maxTokens: 100.25, + headers: { layer: "definition", definitionOnly: "yes" }, + compat: { + supportsStore: true, + openRouterRouting: { only: ["definition"], zdr: false }, + }, + cost: { + input: 1, + output: 2, + cacheRead: 0.1, + cacheWrite: 0.2, + }, + }, + ], + modelOverrides: { + precedence: { + name: "Override", + reasoning: true, + thinkingLevelMap: { + high: "override-high", + futureOverride: "preserve", + }, + contextWindow: 2000.75, + headers: { layer: "override", overrideOnly: "yes" }, + compat: { + supportsStore: false, + openRouterRouting: { order: ["override"] }, + }, + cost: { output: 3.5 }, + }, + }, + }, + }, + { + id: "fractional-cost-and-limits", + providerId: "custom-fractional", + input: { + api: "google-generative-ai", + baseUrl: "https://fractional.example/v1", + apiKey: "literal-secret", + models: [ + { + id: "fractional", + contextWindow: 128000.5, + maxTokens: 16384.25, + cost: { + input: 0.125, + output: 0.375, + cacheRead: 0.0625, + cacheWrite: 0.1875, + }, + }, + ], + }, + }, + { + id: "unmatched-override-is-ignored-by-composer", + providerId: "custom-unmatched", + input: { + api: "google-generative-ai", + baseUrl: "https://google.example/v1", + apiKey: "literal-secret", + models: [{ id: "known" }], + modelOverrides: { + missing: { + maxTokens: 7.5, + headers: { ignored: "yes" }, + }, + }, + }, + }, + { + id: "invalid-nonempty-url-is-composed", + providerId: "custom-invalid-url", + input: { + api: "openai-responses", + baseUrl: "not a URL", + apiKey: "literal-secret", + models: [{ id: "model" }], + }, + }, + { + id: "unknown-api-is-composed", + providerId: "custom-future-api", + input: { + api: "future-wire-v9", + baseUrl: "https://future.example/v9", + apiKey: "literal-secret", + models: [{ id: "future-model" }], + }, + }, + { + id: "unknown-thinking-fields-are-preserved", + providerId: "custom-thinking", + input: { + api: "anthropic-messages", + baseUrl: "https://thinking.example/v1", + apiKey: "literal-secret", + models: [ + { + id: "thinking", + thinkingLevelMap: { + high: "future-high", + future: { nested: ["opaque"] }, + }, + }, + ], + }, + }, + { + id: "empty-url-fails-upstream-composition", + providerId: "custom-empty-url", + input: { + api: "openai-responses", + baseUrl: "", + apiKey: "literal-secret", + models: [{ id: "model" }], + }, + }, +]; + +function joinFieldPath(parent, token) { + const escaped = token.replaceAll("~", "~0").replaceAll("/", "~1"); + return `${parent}/${escaped}`; +} + +function collectSchemaFieldPaths(schema, pointer = "", output = new Set()) { + if (!schema || typeof schema !== "object" || Array.isArray(schema)) { + return output; + } + for (const branch of schema.anyOf ?? []) { + collectSchemaFieldPaths(branch, pointer, output); + } + for (const [name, child] of Object.entries(schema.properties ?? {})) { + const childPointer = joinFieldPath(pointer, name); + output.add(childPointer); + collectSchemaFieldPaths(child, childPointer, output); + } + for (const child of Object.values(schema.patternProperties ?? {})) { + const childPointer = joinFieldPath(pointer, "*"); + output.add(childPointer); + collectSchemaFieldPaths(child, childPointer, output); + } + if (schema.items) { + collectSchemaFieldPaths( + schema.items, + joinFieldPath(pointer, "*"), + output, + ); + } + return output; +} + +function isPatternContainer(pointer) { + return ( + pointer.endsWith("/headers") || + pointer.endsWith("/modelOverrides") || + pointer.endsWith("/chatTemplateKwargs") + ); +} + +function collectInputFieldPaths(value, pointer = "", output = new Set()) { + if (Array.isArray(value)) { + for (const child of value) { + collectInputFieldPaths(child, joinFieldPath(pointer, "*"), output); + } + return output; + } + if (!value || typeof value !== "object") { + return output; + } + for (const [name, child] of Object.entries(value)) { + const childPointer = joinFieldPath( + pointer, + isPatternContainer(pointer) ? "*" : name, + ); + output.add(childPointer); + collectInputFieldPaths(child, childPointer, output); + } + return output; +} + +const schemaFieldPaths = [...collectSchemaFieldPaths(providerSchema)].sort(); +const rawCaseFieldPaths = new Map( + rawInputs.map((entry) => [entry.id, collectInputFieldPaths(entry.input)]), +); +const composerCaseFieldPaths = new Map( + composerCases.map((entry) => [entry.id, collectInputFieldPaths(entry.input)]), +); +function requiredComposerBehaviorCase(fieldPath) { + if (fieldPath === "/models" || fieldPath.startsWith("/models/")) { + return "model-fields-executed"; + } + if ( + fieldPath === "/modelOverrides" || + fieldPath.startsWith("/modelOverrides/") + ) { + return "override-fields-executed"; + } + return "provider-fields-inherited"; +} + +const fieldCoverageEntries = schemaFieldPaths.map((fieldPath) => { + const rawOracleCases = [...rawCaseFieldPaths] + .filter(([, paths]) => paths.has(fieldPath)) + .map(([id]) => id); + const composerOracleCases = [...composerCaseFieldPaths] + .filter(([, paths]) => paths.has(fieldPath)) + .map(([id]) => id); + const composerBehaviorCase = requiredComposerBehaviorCase(fieldPath); + if (!composerOracleCases.includes(composerBehaviorCase)) { + throw new Error( + `Pi field ${fieldPath} is not exercised at its own composition layer by ${composerBehaviorCase}`, + ); + } + return { + fieldPath, + rawOracleCases, + composerOracleCases, + composerBehaviorCase, + }; +}); +const uncoveredRawFields = fieldCoverageEntries + .filter((entry) => entry.rawOracleCases.length === 0) + .map((entry) => entry.fieldPath); +const uncoveredComposerFields = fieldCoverageEntries + .filter((entry) => entry.composerOracleCases.length === 0) + .map((entry) => entry.fieldPath); +if (uncoveredRawFields.length > 0 || uncoveredComposerFields.length > 0) { + throw new Error( + `Pi field coverage is incomplete; raw=${JSON.stringify( + uncoveredRawFields, + )}, composer=${JSON.stringify(uncoveredComposerFields)}`, + ); +} + +const aiShimSource = ` +export function lazyStream() { + throw new Error("transport shim must not execute during composer oracle generation"); +} +`; +const compatShimSource = ` +export function getApiProvider() { + throw new Error("transport shim must not execute during composer oracle generation"); +} +`; + +function projectUpstreamModel( + model, + providerConfig, + resolveCompatibilityRequestConfig, +) { + const requestConfig = resolveCompatibilityRequestConfig( + model, + providerConfig, + undefined, + ); + return { + id: model.id, + name: model.name, + api: model.api, + provider: model.provider, + baseUrl: model.baseUrl, + reasoning: model.reasoning, + ...(model.thinkingLevelMap === undefined + ? {} + : { thinkingLevelMap: model.thinkingLevelMap }), + input: model.input, + cost: model.cost, + contextWindow: model.contextWindow, + maxTokens: model.maxTokens, + ...(model.compat === undefined ? {} : { compat: model.compat }), + ...(requestConfig.headers === undefined + ? {} + : { headers: requestConfig.headers }), + authHeader: requestConfig.authHeader, + }; +} + +async function runPinnedComposer() { + const harnessDirectory = mkdtempSync( + join(piRoot, ".cc-switch-composer-oracle-"), + ); + try { + const aiShimPath = join(harnessDirectory, "pi-ai-shim.mjs"); + const compatShimPath = join(harnessDirectory, "pi-ai-compat-shim.mjs"); + const bundlePath = join(harnessDirectory, "provider-composer.mjs"); + writeFileSync(aiShimPath, aiShimSource); + writeFileSync(compatShimPath, compatShimSource); + buildSync({ + entryPoints: [composerPath], + bundle: true, + platform: "node", + format: "esm", + target: "node22", + outfile: bundlePath, + packages: "external", + alias: { + "@earendil-works/pi-ai": aiShimPath, + "@earendil-works/pi-ai/compat": compatShimPath, + }, + logLevel: "silent", + }); + const bundledComposer = await import(pathToFileURL(bundlePath).href); + const { + composeModelProvider, + resolveCompatibilityRequestConfig, + } = bundledComposer; + const cases = await Promise.all(composerCases.map(async ({ id, providerId, input }) => { + try { + const modelConfig = { + getProvider(candidateId) { + return candidateId === providerId ? input : undefined; + }, + }; + const provider = composeModelProvider( + providerId, + undefined, + modelConfig, + undefined, + ); + const models = provider.getModels(); + const modelIds = new Set(models.map((model) => model.id)); + let authExecution; + if (provider.auth.apiKey) { + try { + const result = await provider.auth.apiKey.resolve({ + ctx: { + env: async () => undefined, + }, + }); + authExecution = { + status: "success", + entryFunction: "Provider.auth.apiKey.resolve", + result: result ?? null, + }; + } catch (error) { + authExecution = { + status: "error", + entryFunction: "Provider.auth.apiKey.resolve", + error: error instanceof Error ? error.message : String(error), + }; + } + } else { + authExecution = { + status: "unavailable", + entryFunction: "Provider.auth.apiKey.resolve", + reason: "pinned composer exposed no API-key auth method", + }; + } + return { + id, + providerId, + input, + execution: { + status: "success", + entryFunctions: [ + "composeModelProvider", + "Provider.getModels", + "resolveCompatibilityRequestConfig", + "Provider.auth.apiKey.resolve", + ], + }, + authExecution, + expected: { + provider: { + id: provider.id, + name: provider.name, + ...(provider.baseUrl === undefined + ? {} + : { baseUrl: provider.baseUrl }), + }, + models: models.map((model) => + projectUpstreamModel( + model, + input, + resolveCompatibilityRequestConfig, + ), + ), + ignoredOverrideKeys: Object.keys(input.modelOverrides ?? {}).filter( + (modelId) => !modelIds.has(modelId), + ), + }, + }; + } catch (error) { + return { + id, + providerId, + input, + execution: { + status: "error", + entryFunctions: ["composeModelProvider"], + }, + expectedError: + error instanceof Error ? error.message : String(error), + }; + } + })); + return { + cases, + harness: { + bundler: `esbuild@${esbuildVersion}`, + upstreamEntry: composerRelativePath, + entryFunctions: [ + "composeModelProvider", + "Provider.getModels", + "resolveCompatibilityRequestConfig", + "Provider.auth.apiKey.resolve", + ], + transportShims: { + "pi-ai": sha256(aiShimSource), + "pi-ai/compat": sha256(compatShimSource), + }, + }, + }; + } finally { + rmSync(harnessDirectory, { recursive: true, force: true }); + } +} + +async function runPinnedTransportResolver() { + const harnessDirectory = mkdtempSync( + join(piRoot, ".cc-switch-transport-oracle-"), + ); + try { + const bundlePath = join(harnessDirectory, "resolve-config-value.mjs"); + buildSync({ + entryPoints: [resolverPath], + bundle: true, + platform: "node", + format: "esm", + target: "node22", + outfile: bundlePath, + packages: "external", + logLevel: "silent", + }); + const resolver = await import(pathToFileURL(bundlePath).href); + const cases = [ + { + id: "literal-value", + input: "literal-secret", + environment: {}, + }, + { + id: "environment-template", + input: "prefix-${PI_ORACLE_VALUE}-suffix", + environment: { PI_ORACLE_VALUE: "environment-secret" }, + }, + { + id: "escaped-dollar-and-bang", + input: "$$literal-$!bang", + environment: {}, + }, + { + id: "shell-command", + input: "!printf pi-command-value", + environment: {}, + }, + { + id: "missing-environment", + input: "${PI_ORACLE_MISSING}", + environment: {}, + }, + ].map((entry) => { + try { + const value = resolver.resolveConfigValueOrThrow( + entry.input, + `oracle ${entry.id}`, + entry.environment, + ); + return { + ...entry, + execution: { + status: "success", + entryFunction: "resolveConfigValueOrThrow", + }, + expected: value, + }; + } catch (error) { + return { + ...entry, + execution: { + status: "error", + entryFunction: "resolveConfigValueOrThrow", + }, + expectedError: + error instanceof Error ? error.message : String(error), + }; + } + }); + const headerInput = { + "x-literal": "literal-header", + "x-environment": "${PI_ORACLE_HEADER}", + "x-command": "!printf pi-command-header", + }; + const headerEnvironment = { PI_ORACLE_HEADER: "environment-header" }; + const expectedHeaders = resolver.resolveHeadersOrThrow( + headerInput, + "oracle headers", + headerEnvironment, + ); + return { + engine: "pinned upstream TypeScript", + bundler: `esbuild@${esbuildVersion}`, + upstreamEntry: resolverRelativePath, + platform: process.platform, + cases, + headerCase: { + id: "provider-header-materialization", + input: headerInput, + environment: headerEnvironment, + execution: { + status: "success", + entryFunction: "resolveHeadersOrThrow", + }, + expected: expectedHeaders, + }, + }; + } finally { + rmSync(harnessDirectory, { recursive: true, force: true }); + } +} + +const composerExecution = await runPinnedComposer(); +const transportExecution = await runPinnedTransportResolver(); +if ( + !rawOracle.cases.find((entry) => entry.id === "all-schema-fields-valid") + ?.expectedValid +) { + throw new Error("the all-fields raw vector was not accepted by pinned TypeBox"); +} +if ( + composerExecution.cases.find( + (entry) => entry.id === "combined-all-fields-precedence", + )?.execution.status !== "success" +) { + throw new Error("the all-fields vector did not execute in the pinned composer"); +} +const composerOracle = { + version: 1, + execution: { + engine: "pinned upstream TypeScript", + piCommit: actualPiCommit, + ...composerExecution.harness, + }, + cases: composerExecution.cases, + failClosedCases: [ + { + id: "builtin-overlay-without-pinned-base-catalog", + kind: "built_in_overlay", + unavailableContext: "pinned built-in Provider instance and model catalog", + rustExpectedStatus: "unknown", + reasonCode: "catalog_required", + }, + { + id: "extension-overlay-without-extension-registration", + kind: "extension_overlay", + unavailableContext: "extension ProviderConfigInput and registered model catalog", + rustExpectedStatus: "unknown", + reasonCode: "catalog_required", + }, + ], +}; +const transportOracle = { + version: 1, + piCommit: actualPiCommit, + ...transportExecution, +}; + +function collectSchemaOperators(value, operators = new Set()) { + if (Array.isArray(value)) { + for (const child of value) collectSchemaOperators(child, operators); + return operators; + } + if (!value || typeof value !== "object") return operators; + for (const [key, child] of Object.entries(value)) { + if ( + [ + "additionalProperties", + "anyOf", + "const", + "items", + "minLength", + "patternProperties", + "properties", + "required", + "type", + ].includes(key) + ) { + operators.add(key); + } + collectSchemaOperators(child, operators); + } + return operators; +} + +mkdirSync(outputDirectory, { recursive: true }); +const schemaArtifact = writeJson( + "provider-schema.snapshot.json", + providerSchema, +); +const rawArtifact = writeJson("raw-oracle-v1.json", rawOracle); +const composerArtifact = writeJson( + "composer-oracle-v1.json", + composerOracle, +); +const transportArtifact = writeJson( + "transport-oracle-v1.json", + transportOracle, +); +const fieldCoverageArtifact = writeJson("field-coverage-v1.json", { + version: 1, + piCommit: actualPiCommit, + schemaTarget: + "ModelsConfigSchema.properties.providers.patternProperties['^.*$']", + assertion: + "every schema field is present in an actual TypeBox Value.Check vector and in a pinned composer execution dedicated to its provider, model, or override layer", + fields: fieldCoverageEntries, +}); + +const provenance = { + version: 1, + pi: { + repository: piRepository, + commit: actualPiCommit, + }, + typeboxVersion, + sources: { + modelConfig: { + path: modelConfigRelativePath, + sha256: sha256(modelConfigSource), + }, + providerComposer: { + path: composerRelativePath, + sha256: sha256(composerSource), + }, + resolveConfigValue: { + path: resolverRelativePath, + sha256: sha256(resolverSource), + }, + }, + artifacts: { + "provider-schema.snapshot.json": schemaArtifact.sha256, + "raw-oracle-v1.json": rawArtifact.sha256, + "composer-oracle-v1.json": composerArtifact.sha256, + "transport-oracle-v1.json": transportArtifact.sha256, + "field-coverage-v1.json": fieldCoverageArtifact.sha256, + }, + harness: { + path: generatorRelativePath, + sha256: sha256(generatorSource), + upstreamEntry: composerRelativePath, + entryFunctions: composerExecution.harness.entryFunctions, + bundler: composerExecution.harness.bundler, + transportShims: composerExecution.harness.transportShims, + assertion: + "expected composer outputs were captured by executing the pinned upstream entry functions; transport shims are unreachable during credential-blind model composition", + transportResolver: { + upstreamEntry: resolverRelativePath, + entryFunctions: [ + "resolveConfigValueOrThrow", + "resolveHeadersOrThrow", + ], + bundler: transportExecution.bundler, + platform: transportExecution.platform, + }, + }, + uncoveredSemantics: [ + ...composerOracle.failClosedCases.map( + ({ id, unavailableContext, rustExpectedStatus, reasonCode }) => ({ + id, + unavailableContext, + rustExpectedStatus, + reasonCode, + }), + ), + { + id: "radius-oauth-credential-lifecycle", + unavailableContext: + "interactive Radius login, refresh credentials, and network exchange", + rustExpectedStatus: "direct_only", + reasonCode: "missing_gateway_credential", + }, + ], + schemaOperatorInventory: [ + ...collectSchemaOperators(providerSchema), + ].sort(), + evaluatorOperatorAllowlist: [ + "additionalProperties", + "anyOf", + "const", + "items", + "minLength", + "patternProperties", + "properties", + "required", + "type", + ], +}; +writeJson("provenance-v1.json", provenance); diff --git a/scripts/pi-transport-capture.mjs b/scripts/pi-transport-capture.mjs new file mode 100644 index 000000000..3ab5d5806 --- /dev/null +++ b/scripts/pi-transport-capture.mjs @@ -0,0 +1,776 @@ +#!/usr/bin/env node +/** + * Pi transport request-capture harness. + * + * 现有的 native-oracle 只执行 pinned Pi 的 schema evaluator / composer / + * config-value resolver,**不执行** adapter 与厂商 SDK 的头合并,因此 + * "Pi 实际发出什么认证头" 一直只能靠读源码推断。本脚本补上这一层: + * 起一个本地 HTTP 抓包端点当 baseUrl,用 pinned Pi 的 adapter 真发一次 + * 请求,记录实际发出的 header。 + * + * 用法: + * PI_CHECKOUT=/path/to/pinned/pi node scripts/pi-transport-capture.mjs + * + * 不含任何密钥:测试用的 apiKey 是本地抓包用的假值;若要打真实端点, + * 通过环境变量传入(PI_CAPTURE_BASE_URL / PI_CAPTURE_API_KEY),不要写进文件。 + * + * 输出为 JSON,可作为 transport 断言的出处依据。若要升级为受冻结的 + * oracle 夹具,请比照 scripts/generate-pi-native-oracle.mjs 补 provenance + * (pinned commit、源码哈希、bundler 版本)。 + */ + +import { createServer } from "node:http"; +import { execFileSync } from "node:child_process"; +import { createHash } from "node:crypto"; +import { createRequire } from "node:module"; +import { + mkdirSync, + mkdtempSync, + readFileSync, + rmSync, + writeFileSync, +} from "node:fs"; +import { join, resolve } from "node:path"; +import { pathToFileURL } from "node:url"; + +const PI = process.env.PI_CHECKOUT; +const EXPECTED_PI_COMMIT = "ab366ebe94cacd419d986be454f12b1b9913aaca"; +if (!PI) { + console.error( + "PI_CHECKOUT must point at a pinned Pi checkout (with node_modules).", + ); + process.exit(2); +} +const piCommit = execFileSync("git", ["-C", PI, "rev-parse", "HEAD"], { + encoding: "utf8", +}).trim(); +if (piCommit !== EXPECTED_PI_COMMIT) { + throw new Error( + `Pi checkout pin mismatch: expected ${EXPECTED_PI_COMMIT}, got ${piCommit}`, + ); +} + +const requireFromPi = createRequire(join(PI, "package.json")); +const { buildSync, version: esbuildVersion } = requireFromPi("esbuild"); +const codingAgentPackagePath = join(PI, "packages/coding-agent/package.json"); +const codingAgentPackageBytes = readFileSync(codingAgentPackagePath); +const codingAgentPackage = JSON.parse(codingAgentPackageBytes.toString("utf8")); +const distributionMetadata = { + source: "packages/coding-agent/package.json", + sha256: createHash("sha256").update(codingAgentPackageBytes).digest("hex"), + name: codingAgentPackage.name, + version: codingAgentPackage.version, + bin: codingAgentPackage.bin, + piConfig: codingAgentPackage.piConfig, +}; + +const ANTHROPIC_SSE = + 'event: message_start\ndata: {"type":"message_start","message":{"id":"m","type":"message","role":"assistant","model":"m","content":[],"stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":1,"output_tokens":1}}}\n\n' + + 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":1}}\n\n' + + 'event: message_stop\ndata: {"type":"message_stop"}\n\n'; +const OPENAI_SSE = + 'data: {"type":"response.completed","response":{"id":"r","status":"completed","output":[],"usage":{"input_tokens":1,"output_tokens":1}}}\n\n' + + "data: [DONE]\n\n"; + +const captured = []; +const server = createServer((request, response) => { + const chunks = []; + request.on("data", (chunk) => chunks.push(chunk)); + request.on("end", () => { + captured.push({ url: request.url, headers: { ...request.headers } }); + response.writeHead(200, { "content-type": "text/event-stream" }); + response.end(request.url.includes("messages") ? ANTHROPIC_SSE : OPENAI_SSE); + }); +}); +await new Promise((resolve) => server.listen(0, "127.0.0.1", resolve)); +const baseUrl = + process.env.PI_CAPTURE_BASE_URL ?? + `http://127.0.0.1:${server.address().port}`; + +const harnessDirectory = mkdtempSync(join(PI, ".cc-switch-transport-capture-")); +process.on("exit", () => + rmSync(harnessDirectory, { recursive: true, force: true }), +); +const aiShimPath = join(harnessDirectory, "pi-ai-shim.mjs"); +const compatShimPath = join(harnessDirectory, "pi-ai-compat-shim.mjs"); +writeFileSync( + aiShimPath, + [ + 'export function lazyStream() { throw new Error("transport shim must not execute during compat capture"); }', + "let uuidSequence = 0;", + 'export function uuidv7() { uuidSequence += 1; return `00000000-0000-7000-8000-${String(uuidSequence).padStart(12, "0")}`; }', + "export class EventStream {}", + "export class ModelsError extends Error {}", + "export function validateToolArguments() { return undefined; }", + 'export function contentText(value) { return typeof value === "string" ? value : ""; }', + 'export function retryAssistantCall() { throw new Error("resource capture must not call AI"); }', + "export function parseStreamingJson() { return undefined; }", + "export function modelsAreEqual(left, right) { return left === right; }", + "export function createModels() { return {}; }", + "export function getBuiltinModelDataGeneratedAt() { return undefined; }", + "export function builtinProviders() { return []; }", + "export function radiusProvider() { return undefined; }", + "", + ].join("\n"), +); +writeFileSync( + compatShimPath, + [ + 'export function getApiProvider() { throw new Error("transport shim must not execute during compat capture"); }', + "export function clampThinkingLevel(value) { return value; }", + "export async function cleanupSessionResources() {}", + "export function getSupportedThinkingLevels() { return []; }", + "export function isContextOverflow() { return false; }", + "export function isRetryableAssistantError() { return false; }", + "export function modelsAreEqual(left, right) { return left === right; }", + "export function resetApiProviders() {}", + 'export function streamSimple() { throw new Error("resource capture must not call AI"); }', + 'export function stream() { throw new Error("resource capture must not call AI"); }', + 'export function completeSimple() { throw new Error("resource capture must not call AI"); }', + "", + ].join("\n"), +); +const entryPoint = join(harnessDirectory, "entry.mjs"); +writeFileSync( + entryPoint, + [ + `export { streamSimple as anthropicMessages } from "${PI}/packages/ai/src/api/anthropic-messages.ts";`, + `export { streamSimple as openaiResponses } from "${PI}/packages/ai/src/api/openai-responses.ts";`, + `export { streamSimple as openaiCompletions } from "${PI}/packages/ai/src/api/openai-completions.ts";`, + `export { streamSimple as googleGenerativeAi } from "${PI}/packages/ai/src/api/google-generative-ai.ts";`, + `export { composeModelProvider } from "${PI}/packages/coding-agent/src/core/provider-composer.ts";`, + `export { resolveConfigValueOrThrow } from "${PI}/packages/coding-agent/src/core/resolve-config-value.ts";`, + `export { loadSkills } from "${PI}/packages/coding-agent/src/core/skills.ts";`, + `export { loadPromptTemplates } from "${PI}/packages/coding-agent/src/core/prompt-templates.ts";`, + `export { SessionManager } from "${PI}/packages/coding-agent/src/core/session-manager.ts";`, + `export { parseArgs } from "${PI}/packages/coding-agent/src/cli/args.ts";`, + `export { createAllToolDefinitions } from "${PI}/packages/coding-agent/src/core/tools/index.ts";`, + ].join("\n"), +); +const bundlePath = join(harnessDirectory, "bundle.mjs"); +buildSync({ + entryPoints: [entryPoint], + bundle: true, + platform: "node", + format: "esm", + outfile: bundlePath, + external: ["node:*"], + packages: "external", + alias: { + "@earendil-works/pi-ai": aiShimPath, + "@earendil-works/pi-ai/compat": compatShimPath, + }, + logLevel: "silent", +}); +const adapters = await import(pathToFileURL(bundlePath).href); + +// ResourceLoader pulls the complete coding-agent resource graph. Bundle the +// real pinned ResourceLoader separately with broad AI stubs; the captured +// instruction behavior remains real while unrelated generated model data is +// kept outside this resource-only probe. +const resourceEntryPoint = join(harnessDirectory, "resource-entry.mjs"); +const resourceBundlePath = join(harnessDirectory, "resource-bundle.mjs"); +writeFileSync( + resourceEntryPoint, + `export { DefaultResourceLoader } from "${PI}/packages/coding-agent/src/core/resource-loader.ts";\n`, +); +buildSync({ + entryPoints: [resourceEntryPoint], + bundle: true, + platform: "node", + format: "esm", + outfile: resourceBundlePath, + external: ["node:*"], + packages: "external", + alias: { + "@earendil-works/pi-ai/providers/all": aiShimPath, + "@earendil-works/pi-ai/oauth": aiShimPath, + "@earendil-works/pi-ai": aiShimPath, + "@earendil-works/pi-ai/compat": compatShimPath, + }, + logLevel: "silent", +}); +const resourceAdapters = await import(pathToFileURL(resourceBundlePath).href); + +const API_BY_ADAPTER = { + anthropicMessages: "anthropic-messages", + openaiResponses: "openai-responses", + openaiCompletions: "openai-completions", + googleGenerativeAi: "google-generative-ai", +}; + +/** 每个用例只改变凭证与显式 header,其余保持最小合法模型。 */ +const CASES = [ + ["anthropicMessages", "plain-key", "sk-ant-api03-plain", {}], + ["anthropicMessages", "oauth-token", "sk-ant-oat01-token", {}], + [ + "anthropicMessages", + "oauth-with-explicit-x-api-key", + "sk-ant-oat01-token", + { "x-api-key": "explicit-secret" }, + ], + [ + "anthropicMessages", + "oauth-with-explicit-authorization", + "sk-ant-oat01-token", + { authorization: "Bearer configured" }, + ], + [ + "anthropicMessages", + "explicit-x-api-key", + "synthesized-secret", + { "x-api-key": "explicit-secret" }, + ], + [ + "anthropicMessages", + "explicit-authorization", + "synthesized-secret", + { authorization: "Bearer configured" }, + ], + ["openaiResponses", "plain-key", "sk-plain", {}], + ["openaiResponses", "oauth-shaped-token", "sk-ant-oat01-not-anthropic", {}], + [ + "openaiResponses", + "explicit-authorization", + "synthesized-secret", + { authorization: "Bearer configured" }, + ], + [ + "openaiCompletions", + "explicit-authorization", + "synthesized-secret", + { authorization: "Bearer configured" }, + ], + [ + "openaiCompletions", + "explicit-x-api-key", + "synthesized-secret", + { "x-api-key": "explicit-secret" }, + ], + ["googleGenerativeAi", "plain-key", "google-plain", {}], + [ + "googleGenerativeAi", + "explicit-x-goog-api-key", + "google-synthesized", + { "x-goog-api-key": "google-explicit" }, + ], +]; + +const results = []; +for (const [adapter, label, apiKey, headers] of CASES) { + const model = { + id: "m", + name: "m", + api: API_BY_ADAPTER[adapter], + provider: "candidate", + baseUrl, + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 1000, + maxTokens: 100, + }; + const before = captured.length; + let error; + try { + const stream = adapters[adapter]( + model, + { messages: [{ role: "user", content: "hi" }] }, + { + apiKey: process.env.PI_CAPTURE_API_KEY ?? apiKey, + headers, + maxTokens: 16, + }, + ); + for await (const _event of stream) { + // drain + } + } catch (caught) { + error = String(caught); + } + const request = + captured.length > before ? captured[captured.length - 1] : undefined; + results.push({ + adapter: API_BY_ADAPTER[adapter], + case: label, + requestSent: Boolean(request), + requestUrl: request?.url, + error: request ? undefined : error, + authHeaders: request + ? Object.fromEntries( + Object.entries(request.headers).filter(([name]) => + [ + "authorization", + "x-api-key", + "x-goog-api-key", + "anthropic-beta", + "anthropic-version", + "openai-beta", + ].includes(name), + ), + ) + : undefined, + }); +} + +const compatInput = { + api: "openai-responses", + baseUrl: "https://compat.example/v1", + apiKey: "literal", + compat: { + openRouterRouting: ["first", "second"], + chatTemplateKwargs: "ab", + baseOnly: true, + }, + models: [{ id: "m", compat: { supportsStore: true } }], + modelOverrides: { + m: { + compat: { + openRouterRouting: null, + chatTemplateKwargs: { named: true }, + overlayOnly: true, + }, + }, + }, +}; +const compatProvider = adapters.composeModelProvider( + "compat-spread", + undefined, + { + getProvider(providerId) { + return providerId === "compat-spread" ? compatInput : undefined; + }, + }, + undefined, +); +const compatSpread = compatProvider.getModels()[0].compat; + +const minimalProvider = adapters.composeModelProvider( + "minimal-provider", + undefined, + { + getProvider(providerId) { + return providerId === "minimal-provider" + ? { + name: "Minimal provider", + api: "openai-responses", + baseUrl: "https://minimal.example/v1", + apiKey: "literal", + models: [{ id: "minimal-model", name: "Minimal model" }], + } + : undefined; + }, + }, + undefined, +); +const minimalModel = minimalProvider.getModels()[0]; +const minimalProviderComposition = { + id: minimalModel.id, + name: minimalModel.name, + provider: minimalModel.provider, + api: minimalModel.api, + baseUrl: minimalModel.baseUrl, + reasoning: minimalModel.reasoning, + input: minimalModel.input, + cost: minimalModel.cost, + contextWindow: minimalModel.contextWindow, + maxTokens: minimalModel.maxTokens, +}; + +function jsonSafeJavaScriptValue(value) { + if (typeof value === "string") { + const codeUnits = Array.from({ length: value.length }, (_, index) => + value.charCodeAt(index), + ); + const hasLoneSurrogate = codeUnits.some((unit, index) => { + if (unit >= 0xd800 && unit <= 0xdbff) { + return !( + index + 1 < codeUnits.length && + codeUnits[index + 1] >= 0xdc00 && + codeUnits[index + 1] <= 0xdfff + ); + } + if (unit >= 0xdc00 && unit <= 0xdfff) { + return !( + index > 0 && + codeUnits[index - 1] >= 0xd800 && + codeUnits[index - 1] <= 0xdbff + ); + } + return false; + }); + return hasLoneSurrogate + ? { + $javascriptStringUtf16: codeUnits.map((unit) => + unit.toString(16).padStart(4, "0"), + ), + } + : value; + } + if (Array.isArray(value)) { + return value.map(jsonSafeJavaScriptValue); + } + if (value && typeof value === "object") { + return Object.fromEntries( + Object.entries(value).map(([key, child]) => [ + key, + jsonSafeJavaScriptValue(child), + ]), + ); + } + return value; +} + +function captureCompatSpread(label, baseValue, overlayValue) { + const providerId = `compat-${label}`; + const provider = adapters.composeModelProvider( + providerId, + undefined, + { + getProvider(candidate) { + if (candidate !== providerId) return undefined; + return { + api: "openai-responses", + baseUrl: "https://compat.example/v1", + apiKey: "literal", + compat: { chatTemplateKwargs: baseValue }, + models: [{ id: "m" }], + modelOverrides: { + m: { compat: { chatTemplateKwargs: overlayValue } }, + }, + }; + }, + }, + undefined, + ); + return { + label, + baseValue, + overlayValue, + result: jsonSafeJavaScriptValue( + provider.getModels()[0].compat.chatTemplateKwargs, + ), + }; +} + +const compatEdgeCases = [ + captureCompatSpread("ascii-string-to-string", "ab", "cd"), + captureCompatSpread("astral-string-to-object", "😀", { named: true }), + captureCompatSpread("astral-string-fully-overridden", "😀", { + 0: "repaired-high", + 1: "repaired-low", + named: true, + }), + captureCompatSpread("string-to-array", "ab", ["first", "second"]), +]; + +const resolverInputs = [ + "literal-secret", + "cash$money", + "café$literal", + "$$literal-$!bang", + "prefix-${PI_CAPTURE_MISSING}-suffix", +]; +const resolverCases = resolverInputs.map((input) => { + try { + return { + input, + status: "success", + result: adapters.resolveConfigValueOrThrow( + input, + "transport capture", + {}, + ), + }; + } catch (error) { + return { + input, + status: "error", + error: error instanceof Error ? error.message : String(error), + }; + } +}); + +// Execute Pi's real discovery entry point. The directories are deliberately +// created in reverse order: a deterministic winner must come from Pi's +// discovery ordering, not filesystem insertion order. +const skillProbeRoot = join(harnessDirectory, "skill-discovery"); +const skillAgentDir = join(skillProbeRoot, "agent"); +const skillProjectDir = join(skillProbeRoot, "project"); +for (const directory of [ + skillProjectDir, + join(skillAgentDir, "skills", "b-second"), + join(skillAgentDir, "skills", "a-first"), +]) { + mkdirSync(directory, { recursive: true }); +} +writeFileSync( + join(skillAgentDir, "skills", "b-second", "SKILL.md"), + "---\nname: duplicate\ndescription: second\n---\nsecond\n", +); +writeFileSync( + join(skillAgentDir, "skills", "a-first", "SKILL.md"), + "---\nname: duplicate\ndescription: first\n---\nfirst\n", +); +const skillDiscovery = jsonSafeJavaScriptValue( + await adapters.loadSkills({ + cwd: skillProjectDir, + agentDir: skillAgentDir, + skillPaths: [], + includeDefaults: true, + }), +); + +const promptAgentDir = join(harnessDirectory, "prompt-agent"); +const promptProjectDir = join(harnessDirectory, "prompt-project"); +mkdirSync(join(promptAgentDir, "prompts", "nested"), { recursive: true }); +mkdirSync(promptProjectDir, { recursive: true }); +writeFileSync( + join(promptAgentDir, "prompts", "review.md"), + "---\ndescription: Review captured changes\nargument-hint: \n---\nReview $1\n", +); +writeFileSync(join(promptAgentDir, "prompts", "empty.md"), ""); +writeFileSync( + join(promptAgentDir, "prompts", "nested", "ignored.md"), + "nested", +); +const promptTemplateDiscovery = jsonSafeJavaScriptValue( + adapters + .loadPromptTemplates({ + cwd: promptProjectDir, + agentDir: promptAgentDir, + promptPaths: [], + includeDefaults: true, + }) + .map((template) => ({ + name: template.name, + description: template.description, + argumentHint: template.argumentHint, + content: template.content, + source: template.sourceInfo?.source, + scope: template.sourceInfo?.scope, + relativeFile: + template.filePath === + join(promptAgentDir, "prompts", `${template.name}.md`) + ? `prompts/${template.name}.md` + : template.filePath, + })), +); + +// File presence, including a zero-byte file, is the native activation state +// for Pi's global instruction resources. Execute the real resource loader so +// cc-switch does not infer that rule from a parser implementation. +const instructionAgentDir = join(harnessDirectory, "instruction-agent"); +const instructionProjectDir = join(harnessDirectory, "instruction-project"); +mkdirSync(instructionAgentDir, { recursive: true }); +mkdirSync(instructionProjectDir, { recursive: true }); +for (const filename of ["AGENTS.md", "SYSTEM.md", "APPEND_SYSTEM.md"]) { + writeFileSync(join(instructionAgentDir, filename), ""); +} +const instructionLoader = new resourceAdapters.DefaultResourceLoader({ + cwd: instructionProjectDir, + agentDir: instructionAgentDir, + noExtensions: true, + noSkills: true, + noPromptTemplates: true, + noThemes: true, +}); +await instructionLoader.reload(); +const emptyInstructionFiles = { + agentsFiles: instructionLoader.getAgentsFiles().agentsFiles.map((entry) => ({ + relativeFile: entry.path.startsWith(instructionAgentDir) + ? entry.path.slice(instructionAgentDir.length + 1) + : entry.path, + content: entry.content, + })), + systemPrompt: instructionLoader.getSystemPrompt(), + systemPromptSource: instructionLoader.getSystemPromptSource()?.path, + appendSystemPrompt: instructionLoader.getAppendSystemPrompt(), + appendSystemPromptSources: instructionLoader + .getAppendSystemPromptSources() + .map((entry) => entry.path), +}; + +// Exercise Pi's real SessionManager instead of inferring sessionDir or JSONL +// shape from its TypeScript source. Relative sessionDir is resolved against the +// launching process cwd, while the header keeps the explicit project cwd. +const sessionProjectDir = join(harnessDirectory, "session project"); +mkdirSync(sessionProjectDir, { recursive: true }); +const originalCwd = process.cwd(); +process.chdir(sessionProjectDir); +const capturedSession = adapters.SessionManager.create( + sessionProjectDir, + ".pi/sessions", + { id: "cc-switch-capture-session" }, +); +capturedSession.appendSessionInfo("Captured session"); +capturedSession.appendMessage({ + role: "user", + content: [{ type: "text", text: "captured question" }], + timestamp: 1_700_000_000_000, +}); +capturedSession.appendMessage({ + role: "assistant", + content: [{ type: "text", text: "captured answer" }], + api: "openai-responses", + provider: "capture", + model: "capture-model", + usage: { + input: 1, + output: 1, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 2, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: 1_700_000_001_000, +}); +const capturedSessionFile = capturedSession.getSessionFile(); +if (!capturedSessionFile) { + throw new Error("pinned SessionManager did not persist the capture session"); +} +const capturedSessionPath = resolve(capturedSessionFile); +const parsedSessionArgs = adapters.parseArgs([ + "--session", + capturedSessionPath, +]); +const parsedVersionArgs = adapters.parseArgs(["--version"]); +const sessionCliSemantics = { + argv: ["--session", capturedSessionPath], + parsedSession: parsedSessionArgs.session, + diagnostics: parsedSessionArgs.diagnostics, + versionArgv: ["--version"], + parsedVersion: parsedVersionArgs.version, + versionDiagnostics: parsedVersionArgs.diagnostics, +}; +const capturedSessionLines = capturedSessionFile + ? readFileSync(capturedSessionFile, "utf8") + .trim() + .split("\n") + .map((line) => JSON.parse(line)) + : []; +const sessionDirectorySemantics = { + processCwd: sessionProjectDir, + suppliedProjectCwd: sessionProjectDir, + suppliedSessionDir: ".pi/sessions", + resolvedSessionDir: capturedSession.getSessionDir(), + headerKeys: Object.keys(capturedSession.getHeader() ?? {}).sort(), + entryShapes: capturedSession.getEntries().map((entry) => ({ + type: entry.type, + keys: Object.keys(entry).sort(), + messageRole: entry.type === "message" ? entry.message.role : undefined, + messageKeys: + entry.type === "message" ? Object.keys(entry.message).sort() : undefined, + })), + persistedLineTypes: capturedSessionLines.map((entry) => entry.type), + listAllCount: ( + await adapters.SessionManager.listAll(capturedSession.getSessionDir()) + ).length, +}; + +const branchedSession = adapters.SessionManager.create( + sessionProjectDir, + ".pi/sessions", + { id: "cc-switch-capture-branch" }, +); +const branchRootId = branchedSession.appendMessage({ + role: "user", + content: [{ type: "text", text: "branch root" }], + timestamp: 1_700_000_002_000, +}); +branchedSession.appendSessionInfo("Abandoned branch name"); +branchedSession.appendMessage({ + role: "assistant", + content: [{ type: "text", text: "abandoned answer" }], + api: "openai-responses", + provider: "capture", + model: "capture-model", + usage: { + input: 1, + output: 1, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 2, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: 1_700_000_003_000, +}); +branchedSession.branch(branchRootId); +branchedSession.appendMessage({ + role: "user", + content: [{ type: "text", text: "active branch" }], + timestamp: 1_700_000_004_000, +}); +const sessionBranchSemantics = { + sessionName: branchedSession.getSessionName(), + activeEntryTypes: branchedSession.getBranch().map((entry) => entry.type), +}; + +const malformedSessionFile = join( + capturedSession.getSessionDir(), + "cc-switch-capture-malformed.jsonl", +); +writeFileSync( + malformedSessionFile, + [ + JSON.stringify(capturedSession.getHeader()), + "{not valid json", + ...capturedSession + .getEntries() + .map((entry) => JSON.stringify(entry)), + "", + ].join("\n"), +); +let malformedSessionSemantics; +try { + const malformedSession = adapters.SessionManager.open(malformedSessionFile); + malformedSessionSemantics = { + status: "accepted", + entryTypes: malformedSession.getEntries().map((entry) => entry.type), + }; +} catch (error) { + malformedSessionSemantics = { + status: "rejected", + error: error instanceof Error ? error.message : String(error), + }; +} +process.chdir(originalCwd); + +// The executed built-in tool factory is the authoritative core inventory. +// Extension-provided tools remain possible, but Pi exposes no native MCP +// registry for cc-switch to mirror as an app toggle. +const nativeToolInventory = Object.values( + adapters.createAllToolDefinitions(sessionProjectDir), +) + .map((tool) => tool.name) + .sort(); + +server.close(); +console.log( + JSON.stringify( + { + bundler: `esbuild@${esbuildVersion}`, + piCheckout: PI, + piCommit, + distributionMetadata, + baseUrl, + results, + minimalProviderComposition, + compatSpread, + compatEdgeCases, + resolverCases, + skillDiscovery, + promptTemplateDiscovery, + emptyInstructionFiles, + sessionDirectorySemantics, + sessionCliSemantics, + sessionBranchSemantics, + malformedSessionSemantics, + nativeToolInventory, + }, + null, + 2, + ), +); diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index c484796f2..5d0411bf2 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -783,6 +783,8 @@ dependencies = [ "indexmap 2.13.0", "json-five", "json5", + "jsonc-parser", + "libc", "log", "objc2 0.5.2", "objc2-app-kit 0.2.2", @@ -2799,6 +2801,15 @@ dependencies = [ "serde", ] +[[package]] +name = "jsonc-parser" +version = "0.33.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a0560e3f9a9a03ea6b6e90b41138c5db9e21526c99eb192c1a26c68176593285" +dependencies = [ + "serde_json", +] + [[package]] name = "jsonptr" version = "0.6.3" @@ -4687,7 +4698,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys 0.4.15", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 56e223c91..e5bd10ee7 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -23,7 +23,8 @@ test-hooks = [] tauri-build = { version = "2.4.0", features = [] } [dependencies] -serde_json = { version = "1.0", features = ["preserve_order"] } +serde_json = { version = "1.0", features = ["arbitrary_precision", "preserve_order"] } +jsonc-parser = { version = "0.33", features = ["cst", "serde_json"] } serde = { version = "1.0", features = ["derive"] } log = "0.4" chrono = { version = "0.4", features = ["serde"] } @@ -78,6 +79,7 @@ indexmap = { version = "2", features = ["serde"] } rust_decimal = "1.33" uuid = { version = "1.11", features = ["v4"] } sha2 = "0.10" +libc = "0.2" hmac = "0.12" json5 = "0.4" json-five = "0.3.1" @@ -94,6 +96,9 @@ winreg = "0.52" windows-sys = { version = "0.61", features = [ "Win32_Globalization", "Win32_Storage_FileSystem", + "Win32_System_Diagnostics_ToolHelp", + "Win32_System_JobObjects", + "Win32_System_Threading", "Win32_UI_Shell", ] } diff --git a/src-tauri/src/app_config.rs b/src-tauri/src/app_config.rs index 74eac3f84..ae3e6fb2f 100644 --- a/src-tauri/src/app_config.rs +++ b/src-tauri/src/app_config.rs @@ -32,6 +32,7 @@ impl McpApps { AppType::OpenCode => self.opencode, AppType::OpenClaw => false, // OpenClaw doesn't support MCP AppType::Hermes => self.hermes, + AppType::Pi => false, // Pi core has no native MCP registry. AppType::ClaudeDesktop => false, } } @@ -46,6 +47,7 @@ impl McpApps { AppType::OpenCode => self.opencode = enabled, AppType::OpenClaw => {} // OpenClaw doesn't support MCP, ignore AppType::Hermes => self.hermes = enabled, + AppType::Pi => {} // Pi core has no native MCP registry. AppType::ClaudeDesktop => {} // Claude Desktop 3P provider config doesn't support MCP here } } @@ -100,6 +102,8 @@ pub struct SkillApps { pub opencode: bool, #[serde(default)] pub hermes: bool, + #[serde(default)] + pub pi: bool, } impl SkillApps { @@ -112,6 +116,7 @@ impl SkillApps { AppType::GrokBuild => self.grokbuild, AppType::OpenCode => self.opencode, AppType::Hermes => self.hermes, + AppType::Pi => self.pi, AppType::OpenClaw => false, // OpenClaw doesn't support Skills AppType::ClaudeDesktop => false, } @@ -126,6 +131,7 @@ impl SkillApps { AppType::GrokBuild => self.grokbuild = enabled, AppType::OpenCode => self.opencode = enabled, AppType::Hermes => self.hermes = enabled, + AppType::Pi => self.pi = enabled, AppType::OpenClaw => {} // OpenClaw doesn't support Skills, ignore AppType::ClaudeDesktop => {} // Claude Desktop 3P profiles don't use CC Switch skill sync } @@ -152,6 +158,9 @@ impl SkillApps { if self.hermes { apps.push(AppType::Hermes); } + if self.pi { + apps.push(AppType::Pi); + } apps } @@ -163,6 +172,7 @@ impl SkillApps { && !self.grokbuild && !self.opencode && !self.hermes + && !self.pi } /// 仅启用指定应用(其他应用设为禁用) @@ -357,6 +367,8 @@ pub struct PromptRoot { pub openclaw: PromptConfig, #[serde(default)] pub hermes: PromptConfig, + #[serde(default)] + pub pi: PromptConfig, } use crate::config::{copy_file, get_app_config_dir, get_app_config_path, write_json_file}; @@ -381,6 +393,7 @@ pub enum AppType { OpenCode, OpenClaw, Hermes, + Pi, } impl AppType { @@ -394,6 +407,7 @@ impl AppType { AppType::OpenCode => "opencode", AppType::OpenClaw => "openclaw", AppType::Hermes => "hermes", + AppType::Pi => "pi", } } @@ -419,6 +433,7 @@ impl AppType { AppType::OpenCode, AppType::OpenClaw, AppType::Hermes, + AppType::Pi, ] .into_iter() } @@ -438,10 +453,11 @@ impl FromStr for AppType { "opencode" => Ok(AppType::OpenCode), "openclaw" => Ok(AppType::OpenClaw), "hermes" => Ok(AppType::Hermes), + "pi" => Ok(AppType::Pi), other => Err(AppError::localized( "unsupported_app", - format!("不支持的应用标识: '{other}'。可选值: claude, claude-desktop, codex, gemini, grokbuild, opencode, openclaw, hermes。"), - format!("Unsupported app id: '{other}'. Allowed: claude, claude-desktop, codex, gemini, grokbuild, opencode, openclaw, hermes."), + format!("不支持的应用标识: '{other}'。可选值: claude, claude-desktop, codex, gemini, grokbuild, opencode, openclaw, hermes, pi。"), + format!("Unsupported app id: '{other}'. Allowed: claude, claude-desktop, codex, gemini, grokbuild, opencode, openclaw, hermes, pi."), )), } } @@ -467,6 +483,9 @@ pub struct CommonConfigSnippets { #[serde(default, skip_serializing_if = "Option::is_none")] pub hermes: Option, + + #[serde(default, skip_serializing_if = "Option::is_none")] + pub pi: Option, } impl CommonConfigSnippets { @@ -481,6 +500,7 @@ impl CommonConfigSnippets { AppType::OpenCode => self.opencode.as_ref(), AppType::OpenClaw => self.openclaw.as_ref(), AppType::Hermes => self.hermes.as_ref(), + AppType::Pi => self.pi.as_ref(), } } @@ -495,6 +515,7 @@ impl CommonConfigSnippets { AppType::OpenCode => self.opencode = snippet, AppType::OpenClaw => self.openclaw = snippet, AppType::Hermes => self.hermes = snippet, + AppType::Pi => self.pi = snippet, } } } @@ -539,6 +560,7 @@ impl Default for MultiAppConfig { apps.insert("opencode".to_string(), ProviderManager::default()); apps.insert("openclaw".to_string(), ProviderManager::default()); apps.insert("hermes".to_string(), ProviderManager::default()); + apps.insert("pi".to_string(), ProviderManager::default()); Self { version: 2, @@ -626,6 +648,12 @@ impl MultiAppConfig { .insert("gemini".to_string(), ProviderManager::default()); updated = true; } + if !config.apps.contains_key("pi") { + config + .apps + .insert("pi".to_string(), ProviderManager::default()); + updated = true; + } // 执行 MCP 迁移(v3.6.x → v3.7.0) let migrated = config.migrate_mcp_to_unified()?; @@ -691,34 +719,6 @@ impl MultiAppConfig { } } - /// 获取指定客户端的 MCP 配置(不可变引用) - pub fn mcp_for(&self, app: &AppType) -> &McpConfig { - match app { - AppType::Claude => &self.mcp.claude, - AppType::ClaudeDesktop => &self.mcp.claude_desktop, - AppType::Codex => &self.mcp.codex, - AppType::Gemini => &self.mcp.gemini, - AppType::GrokBuild => &self.mcp.grokbuild, - AppType::OpenCode => &self.mcp.opencode, - AppType::OpenClaw => &self.mcp.openclaw, - AppType::Hermes => &self.mcp.hermes, - } - } - - /// 获取指定客户端的 MCP 配置(可变引用) - pub fn mcp_for_mut(&mut self, app: &AppType) -> &mut McpConfig { - match app { - AppType::Claude => &mut self.mcp.claude, - AppType::ClaudeDesktop => &mut self.mcp.claude_desktop, - AppType::Codex => &mut self.mcp.codex, - AppType::Gemini => &mut self.mcp.gemini, - AppType::GrokBuild => &mut self.mcp.grokbuild, - AppType::OpenCode => &mut self.mcp.opencode, - AppType::OpenClaw => &mut self.mcp.openclaw, - AppType::Hermes => &mut self.mcp.hermes, - } - } - /// 创建默认配置并自动导入已存在的提示词文件 fn default_with_auto_import() -> Result { log::info!("首次启动,创建默认配置并检测提示词文件"); @@ -733,6 +733,7 @@ impl MultiAppConfig { Self::auto_import_prompt_if_exists(&mut config, AppType::OpenCode)?; Self::auto_import_prompt_if_exists(&mut config, AppType::OpenClaw)?; Self::auto_import_prompt_if_exists(&mut config, AppType::Hermes)?; + Self::auto_import_prompt_if_exists(&mut config, AppType::Pi)?; Ok(config) } @@ -757,6 +758,7 @@ impl MultiAppConfig { || !self.prompts.opencode.prompts.is_empty() || !self.prompts.openclaw.prompts.is_empty() || !self.prompts.hermes.prompts.is_empty() + || !self.prompts.pi.prompts.is_empty() { return Ok(false); } @@ -772,6 +774,7 @@ impl MultiAppConfig { AppType::OpenCode, AppType::OpenClaw, AppType::Hermes, + AppType::Pi, ] { // 复用已有的单应用导入逻辑 if Self::auto_import_prompt_if_exists(self, app)? { @@ -846,6 +849,7 @@ impl MultiAppConfig { AppType::OpenCode => &mut config.prompts.opencode.prompts, AppType::OpenClaw => &mut config.prompts.openclaw.prompts, AppType::Hermes => &mut config.prompts.hermes.prompts, + AppType::Pi => &mut config.prompts.pi.prompts, }; prompts.insert(id, prompt); @@ -889,6 +893,7 @@ impl MultiAppConfig { AppType::OpenCode => &self.mcp.opencode.servers, AppType::OpenClaw => continue, // OpenClaw MCP is still in development, skip AppType::Hermes => continue, // Hermes didn't exist in v3.6.x, skip + AppType::Pi => continue, // Pi didn't exist in v3.6.x, skip }; for (id, entry) in old_servers { diff --git a/src-tauri/src/architecture_tests.rs b/src-tauri/src/architecture_tests.rs new file mode 100644 index 000000000..b6d06cf82 --- /dev/null +++ b/src-tauri/src/architecture_tests.rs @@ -0,0 +1,1001 @@ +#![cfg(test)] + +use regex::Regex; +use serde_json::json; +use std::collections::{BTreeMap, BTreeSet}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::sync::LazyLock; +use syn::visit::{self, Visit}; +use syn::{ + Attribute, ExprLit, ImplItem, Item, ItemImpl, ItemStruct, ItemUse, Lit, Meta, Path as SynPath, + Token, Type, UseTree, +}; + +static PROVIDER_DML: LazyLock = LazyLock::new(|| { + Regex::new( + r#"(?is)\b(?:INSERT(?:\s+OR\s+(?:ABORT|FAIL|IGNORE|REPLACE|ROLLBACK))?\s+INTO|UPDATE(?:\s+OR\s+(?:ABORT|FAIL|IGNORE|REPLACE|ROLLBACK))?|DELETE\s+FROM)\s+["`\[]?(providers|provider_endpoints)\b"#, + ) + .expect("compile provider DML classifier") +}); +static RESTORE_CREATE: LazyLock = LazyLock::new(|| { + Regex::new(r"(?is)\bCREATE\s+(?:TEMP(?:ORARY)?\s+)?(?:TABLE|INDEX|TRIGGER|VIEW)\b") + .expect("compile restore DDL classifier") +}); + +#[derive(Debug, Clone, PartialEq, Eq)] +struct Violation { + kind: &'static str, + path: String, + detail: String, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum CfgTruth { + True, + False, + Unknown, +} + +impl CfgTruth { + fn not(self) -> Self { + match self { + Self::True => Self::False, + Self::False => Self::True, + Self::Unknown => Self::Unknown, + } + } + + fn and(self, other: Self) -> Self { + match (self, other) { + (Self::False, _) | (_, Self::False) => Self::False, + (Self::True, Self::True) => Self::True, + _ => Self::Unknown, + } + } + + fn or(self, other: Self) -> Self { + match (self, other) { + (Self::True, _) | (_, Self::True) => Self::True, + (Self::False, Self::False) => Self::False, + _ => Self::Unknown, + } + } +} + +/// Evaluate a cfg predicate for a production build where `test = false`. +/// Other target/features are deliberately unknown: the scanner may exclude an +/// item only when the predicate is definitely false in every production +/// configuration. +fn production_cfg_truth(meta: &Meta) -> CfgTruth { + match meta { + Meta::Path(path) if path.is_ident("test") => CfgTruth::False, + Meta::Path(_) | Meta::NameValue(_) => CfgTruth::Unknown, + Meta::List(list) if list.path.is_ident("not") => list + .parse_args::() + .map(|inner| production_cfg_truth(&inner).not()) + .unwrap_or(CfgTruth::Unknown), + Meta::List(list) if list.path.is_ident("all") || list.path.is_ident("any") => { + let Ok(items) = list + .parse_args_with(syn::punctuated::Punctuated::::parse_terminated) + else { + return CfgTruth::Unknown; + }; + if list.path.is_ident("all") { + items.iter().fold(CfgTruth::True, |truth, item| { + truth.and(production_cfg_truth(item)) + }) + } else { + items.iter().fold(CfgTruth::False, |truth, item| { + truth.or(production_cfg_truth(item)) + }) + } + } + Meta::List(_) => CfgTruth::Unknown, + } +} + +fn is_cfg_test(attrs: &[Attribute]) -> bool { + attrs.iter().any(|attribute| { + let Meta::List(list) = &attribute.meta else { + return false; + }; + if !list.path.is_ident("cfg") { + return false; + } + list.parse_args::() + .map(|predicate| production_cfg_truth(&predicate) == CfgTruth::False) + .unwrap_or(false) + }) +} + +fn flatten_use(tree: &UseTree, prefix: &mut Vec, output: &mut Vec>) { + match tree { + UseTree::Path(path) => { + prefix.push(path.ident.to_string()); + flatten_use(&path.tree, prefix, output); + prefix.pop(); + } + UseTree::Name(name) => { + let mut path = prefix.clone(); + path.push(name.ident.to_string()); + output.push(path); + } + UseTree::Rename(rename) => { + let mut path = prefix.clone(); + path.push(rename.ident.to_string()); + output.push(path); + } + UseTree::Group(group) => { + for item in &group.items { + flatten_use(item, prefix, output); + } + } + UseTree::Glob(_) => output.push(prefix.clone()), + } +} + +fn path_segments(path: &SynPath) -> Vec { + path.segments + .iter() + .map(|segment| segment.ident.to_string()) + .collect() +} + +struct ArchitectureVisitor<'a> { + path: &'a str, + source_module: Option<&'static str>, + violations: Vec, + internal_edges: BTreeSet, +} + +fn pi_config_source_module(path: &str) -> Option<&'static str> { + ["raw_schema", "composer", "gateway", "native", "model"] + .into_iter() + .find(|module| { + let root = format!("pi_config/{module}"); + path.ends_with(&format!("{root}.rs")) || path.contains(&format!("{root}/")) + }) +} + +impl ArchitectureVisitor<'_> { + fn record_dependency(&mut self, segments: &[String]) { + let Some(source_module) = self.source_module else { + return; + }; + let qualified_internal_path = segments.len() > 1 + && segments.iter().any(|segment| { + matches!(segment.as_str(), "super" | "self" | "crate" | "pi_config") + }); + if qualified_internal_path { + for module in [ + "raw_schema", + "composer", + "gateway", + "native", + "model", + "document", + ] { + if module != source_module && segments.iter().any(|segment| segment == module) { + self.internal_edges + .insert(format!("{source_module}->{module}")); + } + } + } + + let imports_gateway = segments + .iter() + .any(|segment| segment == "gateway" || segment.starts_with("PiGateway")); + let imports_model_module = segments.len() > 1 + && segments.iter().any(|segment| segment == "model") + && segments.iter().any(|segment| { + matches!(segment.as_str(), "super" | "self" | "crate" | "pi_config") + }); + let imports_managed = imports_model_module + || segments.iter().any(|segment| { + segment == "PiApiFamily" + || segment.starts_with("PiManaged") + || segment == "PiEffectiveModel" + }); + if matches!(source_module, "raw_schema" | "composer") + && (imports_gateway || imports_managed) + { + self.violations.push(Violation { + kind: "cross_layer_import", + path: self.path.to_string(), + detail: format!( + "{source_module} imports managed/gateway path {}", + segments.join("::") + ), + }); + } + + let imports_raw_valid = segments + .iter() + .any(|segment| segment == "raw_schema" || segment == "PiRawValidProvider"); + if source_module == "gateway" && imports_raw_valid { + self.violations.push(Violation { + kind: "cross_layer_import", + path: self.path.to_string(), + detail: format!( + "gateway constructs or imports raw-valid path {}", + segments.join("::") + ), + }); + } + } + + fn inspect_string(&mut self, literal: &str) { + if PROVIDER_DML.is_match(literal) && !provider_dml_path_allowed(self.path) { + self.violations.push(Violation { + kind: "provider_dml", + path: self.path.to_string(), + detail: "provider table DML exists outside typed row/state, endpoint, migration, or canonical-copy authority".to_string(), + }); + } + if self.path.ends_with("database/backup.rs") && RESTORE_CREATE.is_match(literal) { + self.violations.push(Violation { + kind: "restore_create", + path: self.path.to_string(), + detail: "restore production code contains CREATE TABLE/INDEX/TRIGGER/VIEW" + .to_string(), + }); + } + } +} + +impl<'ast> Visit<'ast> for ArchitectureVisitor<'_> { + fn visit_item(&mut self, item: &'ast Item) { + let attrs = match item { + Item::Const(item) => &item.attrs, + Item::Enum(item) => &item.attrs, + Item::ExternCrate(item) => &item.attrs, + Item::Fn(item) => &item.attrs, + Item::ForeignMod(item) => &item.attrs, + Item::Impl(item) => &item.attrs, + Item::Macro(item) => &item.attrs, + Item::Mod(item) => &item.attrs, + Item::Static(item) => &item.attrs, + Item::Struct(item) => &item.attrs, + Item::Trait(item) => &item.attrs, + Item::TraitAlias(item) => &item.attrs, + Item::Type(item) => &item.attrs, + Item::Union(item) => &item.attrs, + Item::Use(item) => &item.attrs, + Item::Verbatim(_) => { + visit::visit_item(self, item); + return; + } + _ => { + visit::visit_item(self, item); + return; + } + }; + if !is_cfg_test(attrs) { + visit::visit_item(self, item); + } + } + + fn visit_impl_item(&mut self, item: &'ast ImplItem) { + let attrs = match item { + ImplItem::Const(item) => &item.attrs, + ImplItem::Fn(item) => &item.attrs, + ImplItem::Type(item) => &item.attrs, + ImplItem::Macro(item) => &item.attrs, + ImplItem::Verbatim(_) => { + visit::visit_impl_item(self, item); + return; + } + _ => { + visit::visit_impl_item(self, item); + return; + } + }; + if !is_cfg_test(attrs) { + visit::visit_impl_item(self, item); + } + } + + fn visit_ident(&mut self, identifier: &'ast syn::Ident) { + let identifier = identifier.to_string(); + if matches!( + identifier.as_str(), + "save_provider" | "save_provider_row_on_tx" | "save_provider_aggregate" + ) || identifier.starts_with("upsert_provider") + { + self.violations.push(Violation { + kind: "forbidden_provider_symbol", + path: self.path.to_string(), + detail: format!("forbidden provider write symbol '{identifier}'"), + }); + } + } + + fn visit_expr_lit(&mut self, expression: &'ast ExprLit) { + if let Lit::Str(literal) = &expression.lit { + self.inspect_string(&literal.value()); + } + visit::visit_expr_lit(self, expression); + } + + fn visit_item_use(&mut self, item: &'ast ItemUse) { + let mut paths = Vec::new(); + flatten_use(&item.tree, &mut Vec::new(), &mut paths); + for path in paths { + self.record_dependency(&path); + } + visit::visit_item_use(self, item); + } + + fn visit_path(&mut self, path: &'ast SynPath) { + self.record_dependency(&path_segments(path)); + visit::visit_path(self, path); + } +} + +fn provider_dml_path_allowed(path: &str) -> bool { + [ + "database/dao/provider_write.rs", + "database/dao/providers.rs", + "database/dao/failover.rs", + "database/schema.rs", + "database/migration.rs", + "database/backup.rs", + ] + .iter() + .any(|allowed| path.ends_with(allowed)) +} + +fn scan_source(path: &str, source: &str) -> (Vec, BTreeSet) { + scan_source_as_module(path, source, pi_config_source_module(path)) +} + +fn scan_source_as_module( + path: &str, + source: &str, + source_module: Option<&'static str>, +) -> (Vec, BTreeSet) { + let syntax = match syn::parse_file(source) { + Ok(syntax) => syntax, + Err(error) => { + return ( + vec![Violation { + kind: "parse_error", + path: path.to_string(), + detail: error.to_string(), + }], + BTreeSet::new(), + ) + } + }; + // 文件级 #![cfg(test)] 的文件(认证套件等)不进入任何构建的生产目标, + // 不参与生产扫描;该属性的存在性由认证套件的注册元测试强制。 + if is_cfg_test(&syntax.attrs) { + return (Vec::new(), BTreeSet::new()); + } + let mut visitor = ArchitectureVisitor { + path, + source_module, + violations: Vec::new(), + internal_edges: BTreeSet::new(), + }; + visitor.visit_file(&syntax); + (visitor.violations, visitor.internal_edges) +} + +fn rust_sources(root: &Path) -> Vec { + fn walk(directory: &Path, output: &mut Vec) { + let mut entries = fs::read_dir(directory) + .unwrap_or_else(|error| panic!("read {}: {error}", directory.display())) + .collect::, _>>() + .unwrap_or_else(|error| panic!("read entry in {}: {error}", directory.display())); + entries.sort_by_key(|entry| entry.path()); + for entry in entries { + let path = entry.path(); + if path.is_dir() { + walk(&path, output); + } else if path.extension().is_some_and(|extension| extension == "rs") { + output.push(path); + } + } + } + let mut output = Vec::new(); + walk(root, &mut output); + output +} + +#[derive(Debug, Default)] +struct ModulePathOptions { + explicit: Vec, + definitely_explicit: bool, +} + +impl ModulePathOptions { + fn default_path_is_possible(&self) -> bool { + !self.definitely_explicit + } +} + +fn collect_conditional_module_paths( + meta: &Meta, + activation: CfgTruth, + output: &mut ModulePathOptions, +) { + match meta { + Meta::NameValue(value) if value.path.is_ident("path") => { + let syn::Expr::Lit(expression) = &value.value else { + return; + }; + let Lit::Str(path) = &expression.lit else { + return; + }; + if activation != CfgTruth::False { + output.explicit.push(PathBuf::from(path.value())); + output.definitely_explicit |= activation == CfgTruth::True; + } + } + Meta::List(list) if list.path.is_ident("cfg_attr") => { + let Ok(items) = list + .parse_args_with(syn::punctuated::Punctuated::::parse_terminated) + else { + return; + }; + let mut items = items.iter(); + let Some(predicate) = items.next() else { + return; + }; + let activation = activation.and(production_cfg_truth(predicate)); + for attribute in items { + collect_conditional_module_paths(attribute, activation, output); + } + } + _ => {} + } +} + +fn module_path_options(attrs: &[Attribute]) -> ModulePathOptions { + let mut output = ModulePathOptions::default(); + for attribute in attrs { + collect_conditional_module_paths(&attribute.meta, CfgTruth::True, &mut output); + } + output.explicit.sort(); + output.explicit.dedup(); + output +} + +fn default_submodule_directory(source_path: &Path) -> PathBuf { + let parent = source_path + .parent() + .expect("Rust source has a parent directory"); + match source_path.file_stem().and_then(|stem| stem.to_str()) { + Some("lib" | "main" | "mod") | None => parent.to_path_buf(), + Some(stem) => parent.join(stem), + } +} + +fn collect_module_targets( + items: &[Item], + source_path: &Path, + inline_modules: &[String], + output: &mut Vec, +) { + for item in items { + let Item::Mod(module) = item else { + continue; + }; + if is_cfg_test(&module.attrs) { + continue; + } + if let Some((_, nested)) = &module.content { + let mut nested_modules = inline_modules.to_vec(); + nested_modules.push(module.ident.to_string()); + collect_module_targets(nested, source_path, &nested_modules, output); + continue; + } + + let options = module_path_options(&module.attrs); + let module_directory = default_submodule_directory(source_path); + let inline_directory = inline_modules + .iter() + .fold(module_directory, |directory, module| directory.join(module)); + let explicit_base = if inline_modules.is_empty() { + source_path + .parent() + .expect("Rust source has a parent directory") + .to_path_buf() + } else { + inline_directory.clone() + }; + let default_path_is_possible = options.default_path_is_possible(); + output.extend( + options + .explicit + .into_iter() + .map(|relative| explicit_base.join(relative)), + ); + if default_path_is_possible { + let module_name = module.ident.to_string(); + output.push(inline_directory.join(format!("{module_name}.rs"))); + output.push(inline_directory.join(module_name).join("mod.rs")); + } + } +} + +fn inherited_module_owners( + sources: &[(PathBuf, String, String)], +) -> BTreeMap> { + let mut owners = BTreeMap::>::new(); + for (path, relative, _) in sources { + if let Some(owner) = pi_config_source_module(relative) { + owners + .entry(fs::canonicalize(path).expect("canonicalize Rust source")) + .or_default() + .insert(owner); + } + } + + loop { + let mut changed = false; + for (path, _, source) in sources { + let canonical = fs::canonicalize(path).expect("canonicalize Rust source"); + let source_owners = owners.get(&canonical).cloned().unwrap_or_default(); + if source_owners.is_empty() { + continue; + } + let Ok(syntax) = syn::parse_file(source) else { + continue; + }; + if is_cfg_test(&syntax.attrs) { + continue; + } + let mut targets = Vec::new(); + collect_module_targets(&syntax.items, path, &[], &mut targets); + for target in targets { + let Ok(target) = fs::canonicalize(target) else { + continue; + }; + let target_owners = owners.entry(target).or_default(); + let before = target_owners.len(); + target_owners.extend(source_owners.iter().copied()); + changed |= target_owners.len() != before; + } + } + if !changed { + return owners; + } + } +} + +fn scan_production_tree( + manifest_dir: &Path, + source_root: &Path, +) -> (Vec, BTreeSet) { + let mut violations = Vec::new(); + let mut edges = BTreeSet::new(); + let sources = rust_sources(source_root) + .into_iter() + .map(|path| { + let relative = path + .strip_prefix(manifest_dir) + .expect("source is under manifest") + .to_string_lossy() + .replace('\\', "/"); + let source = fs::read_to_string(&path) + .unwrap_or_else(|error| panic!("read {}: {error}", path.display())); + (path, relative, source) + }) + .collect::>(); + let owners = inherited_module_owners(&sources); + let scanned_paths = sources + .iter() + .map(|(path, _, _)| fs::canonicalize(path).expect("canonicalize Rust source")) + .collect::>(); + for (target, target_owners) in &owners { + if !scanned_paths.contains(target) { + violations.push(Violation { + kind: "module_path_outside_scan", + path: target.display().to_string(), + detail: format!( + "a #[path] module owned by {} escapes the scanned source tree", + target_owners.iter().copied().collect::>().join(",") + ), + }); + } + } + for (path, relative, source) in sources { + let canonical = fs::canonicalize(&path).expect("canonicalize Rust source"); + let inherited = owners.get(&canonical).cloned().unwrap_or_default(); + if inherited.is_empty() { + let (mut source_violations, source_edges) = scan_source(&relative, &source); + violations.append(&mut source_violations); + edges.extend(source_edges); + } else { + for owner in inherited { + let (mut source_violations, source_edges) = + scan_source_as_module(&relative, &source, Some(owner)); + violations.append(&mut source_violations); + edges.extend(source_edges); + } + } + } + (violations, edges) +} + +fn type_leaf(ty: &Type) -> String { + match ty { + Type::Path(path) => path + .path + .segments + .last() + .map(|segment| segment.ident.to_string()) + .unwrap_or_else(|| "unknown".to_string()), + Type::Reference(reference) => type_leaf(&reference.elem), + _ => "other".to_string(), + } +} + +fn provider_write_api_snapshot(source: &str) -> serde_json::Value { + let syntax = syn::parse_file(source).expect("parse provider write authority"); + let type_names = [ + "ProviderKey", + "ProviderRowCreate", + "ProviderRowUpdate", + "NewEndpoint", + "NewProviderAggregate", + "RenameProvider", + ]; + let method_names = [ + "create_provider", + "update_provider", + "rename_db_only_additive_provider", + "add_provider_endpoint", + "remove_provider_endpoint", + "touch_provider_endpoint", + ]; + let mut types = BTreeMap::>::new(); + let mut methods = BTreeMap::>::new(); + for item in syntax.items { + if let Item::Struct(ItemStruct { ident, fields, .. }) = &item { + if type_names.contains(&ident.to_string().as_str()) { + types.insert( + ident.to_string(), + fields + .iter() + .filter_map(|field| field.ident.as_ref().map(ToString::to_string)) + .collect(), + ); + } + } + if let Item::Impl(ItemImpl { self_ty, items, .. }) = item { + if type_leaf(&self_ty) != "Database" { + continue; + } + for item in items { + let ImplItem::Fn(function) = item else { + continue; + }; + let name = function.sig.ident.to_string(); + if !method_names.contains(&name.as_str()) { + continue; + } + let inputs = function + .sig + .inputs + .iter() + .filter_map(|argument| match argument { + syn::FnArg::Receiver(_) => None, + syn::FnArg::Typed(argument) => Some(type_leaf(&argument.ty)), + }) + .collect(); + methods.insert(name, inputs); + } + } + } + json!({ + "manifestVersion": 1, + "codeAuthority": "src-tauri/src/database/dao/provider_write.rs", + "types": types, + "databaseMethods": methods, + }) +} + +#[test] +fn architecture_scanner_accepts_production_tree_and_matches_machine_snapshots() { + let manifest_dir = Path::new(env!("CARGO_MANIFEST_DIR")); + let source_root = manifest_dir.join("src"); + let (violations, edges) = scan_production_tree(manifest_dir, &source_root); + assert!( + violations.is_empty(), + "architecture violations:\n{}", + violations + .iter() + .map(|violation| format!( + "{} {}: {}", + violation.kind, violation.path, violation.detail + )) + .collect::>() + .join("\n") + ); + + // Integration tests compile as external consumers and used to retain + // calls to the deleted generic writer even after `cargo test --lib` + // passed. They may contain deliberate SQL fixtures, so this pass checks + // only that the forbidden API surface is absent from actual Rust syntax. + let mut forbidden_test_symbols = Vec::new(); + for path in rust_sources(&manifest_dir.join("tests")) { + let relative = path + .strip_prefix(manifest_dir) + .expect("integration test is under manifest") + .to_string_lossy() + .replace('\\', "/"); + let source = fs::read_to_string(&path) + .unwrap_or_else(|error| panic!("read {}: {error}", path.display())); + let (source_violations, _) = scan_source(&relative, &source); + forbidden_test_symbols.extend( + source_violations + .into_iter() + .filter(|violation| violation.kind == "forbidden_provider_symbol"), + ); + } + assert!( + forbidden_test_symbols.is_empty(), + "forbidden provider write symbols remain in integration tests:\n{}", + forbidden_test_symbols + .iter() + .map(|violation| format!("{}: {}", violation.path, violation.detail)) + .collect::>() + .join("\n") + ); + + let module_snapshot: serde_json::Value = serde_json::from_str(include_str!( + "../../tests/fixtures/pi/module-boundaries-v1.json" + )) + .expect("parse module boundary snapshot"); + assert_eq!( + json!({ + "manifestVersion": 1, + "codeAuthority": "src-tauri/src/architecture_tests.rs", + "edges": edges, + }), + module_snapshot + ); + + let provider_source = fs::read_to_string(source_root.join("database/dao/provider_write.rs")) + .expect("read provider write authority"); + let provider_snapshot: serde_json::Value = serde_json::from_str(include_str!( + "../../tests/fixtures/pi/provider-write-api-v1.json" + )) + .expect("parse provider write API snapshot"); + assert_eq!( + provider_write_api_snapshot(&provider_source), + provider_snapshot + ); +} + +#[test] +fn architecture_scanner_negative_fixtures_prove_each_guard_fires() { + let cases = [ + ( + "src/services/forbidden.rs", + "fn call() { save_provider(); }", + "forbidden_provider_symbol", + ), + ( + "src/services/rogue.rs", + r#"fn write(conn: &Db) { conn.execute("UPDATE OR REPLACE providers SET name = 'x'", []); }"#, + "provider_dml", + ), + ( + "src/database/backup.rs", + r#"fn stage(conn: &Db) { conn.execute("CREATE TABLE leaked (id INTEGER)", []); }"#, + "restore_create", + ), + ( + "src/pi_config/composer.rs", + "use super::gateway::PiGatewayApiFamily;", + "cross_layer_import", + ), + ( + "src/pi_config/gateway.rs", + "use super::raw_schema::PiRawValidProvider;", + "cross_layer_import", + ), + ]; + for (path, source, expected_kind) in cases { + let (violations, _) = scan_source(path, source); + assert!( + violations + .iter() + .any(|violation| violation.kind == expected_kind), + "negative fixture {path} did not trigger {expected_kind}: {violations:?}" + ); + } + + let ignored_test_fixture = r#" + #[cfg(test)] + mod tests { + fn legacy() { + save_provider(); + let _ = "CREATE TABLE ignored (id INTEGER)"; + let _ = "DELETE FROM providers"; + } + } + "#; + let (violations, _) = scan_source("src/services/ignored.rs", ignored_test_fixture); + assert!( + violations.is_empty(), + "#[cfg(test)] content must be excluded: {violations:?}" + ); + + let production_not_test_fixture = r#" + #[cfg(not(test))] + fn production_escape_attempt() { + save_provider(); + } + "#; + let (violations, _) = scan_source("src/services/production.rs", production_not_test_fixture); + assert!( + violations + .iter() + .any(|violation| violation.kind == "forbidden_provider_symbol"), + "#[cfg(not(test))] production content must remain visible: {violations:?}" + ); + + let test_only_conjunction = r#" + #[cfg(all(test, unix))] + fn test_only() { + save_provider(); + } + "#; + let (violations, _) = scan_source("src/services/test_only.rs", test_only_conjunction); + assert!( + violations.is_empty(), + "a predicate requiring test must stay excluded: {violations:?}" + ); + + let temp = tempfile::tempdir().expect("temp architecture tree"); + let nested = temp.path().join("src/pi_config/composer/tests/mod.rs"); + fs::create_dir_all(nested.parent().expect("nested parent")).expect("create nested tree"); + fs::write( + &nested, + r#" + use super::super::gateway::PiGatewayApiFamily; + fn production_escape_attempt() { + save_provider(); + } + "#, + ) + .expect("write production nested-module fixture"); + let (violations, _) = scan_production_tree(temp.path(), &temp.path().join("src")); + for expected_kind in ["cross_layer_import", "forbidden_provider_symbol"] { + assert!( + violations + .iter() + .any(|violation| violation.kind == expected_kind), + "production code under tests/ must inherit composer ownership and trigger \ + {expected_kind}: {violations:?}" + ); + } + + let custom_path_root = tempfile::tempdir().expect("temp custom-path architecture tree"); + let composer = custom_path_root.path().join("src/pi_config/composer.rs"); + let shared = custom_path_root.path().join("src/shared/helper.rs"); + fs::create_dir_all(composer.parent().expect("composer parent")).expect("create pi_config"); + fs::create_dir_all(shared.parent().expect("shared parent")).expect("create shared"); + fs::write( + &composer, + r#" + #[path = "../shared/helper.rs"] + mod helper; + "#, + ) + .expect("write custom-path parent"); + fs::write( + &shared, + r#" + use crate::pi_config::gateway::PiGatewayApiFamily; + fn production_escape_attempt() { + save_provider(); + } + "#, + ) + .expect("write custom-path child"); + let (violations, _) = scan_production_tree( + custom_path_root.path(), + &custom_path_root.path().join("src"), + ); + for expected_kind in ["cross_layer_import", "forbidden_provider_symbol"] { + assert!( + violations.iter().any(|violation| { + violation.kind == expected_kind && violation.path.ends_with("src/shared/helper.rs") + }), + "#[path] modules must inherit composer ownership and trigger \ + {expected_kind}: {violations:?}" + ); + } + + let conditional_path_root = + tempfile::tempdir().expect("temp conditional-path architecture tree"); + let composer = conditional_path_root + .path() + .join("src/pi_config/composer.rs"); + let shared = conditional_path_root.path().join("src/shared/helper.rs"); + fs::create_dir_all(composer.parent().expect("composer parent")).expect("create pi_config"); + fs::create_dir_all(shared.parent().expect("shared parent")).expect("create shared"); + fs::write( + &composer, + r#" + #[cfg_attr(not(test), path = "../shared/helper.rs")] + mod helper; + "#, + ) + .expect("write conditional-path parent"); + fs::write( + &shared, + r#" + use crate::pi_config::gateway::PiGatewayApiFamily; + fn production_escape_attempt() { + save_provider(); + } + "#, + ) + .expect("write conditional-path child"); + let (violations, _) = scan_production_tree( + conditional_path_root.path(), + &conditional_path_root.path().join("src"), + ); + for expected_kind in ["cross_layer_import", "forbidden_provider_symbol"] { + assert!( + violations.iter().any(|violation| { + violation.kind == expected_kind && violation.path.ends_with("src/shared/helper.rs") + }), + "production cfg_attr(path) modules must inherit composer ownership and trigger \ + {expected_kind}: {violations:?}" + ); + } + + let transitive_path_root = tempfile::tempdir().expect("temp transitive-path architecture tree"); + let composer = transitive_path_root + .path() + .join("src/pi_config/composer.rs"); + let helper = transitive_path_root.path().join("src/shared/helper.rs"); + let leaf = transitive_path_root + .path() + .join("src/shared/helper/leaf.rs"); + fs::create_dir_all(composer.parent().expect("composer parent")).expect("create pi_config"); + fs::create_dir_all(helper.parent().expect("helper parent")).expect("create shared"); + fs::create_dir_all(leaf.parent().expect("leaf parent")).expect("create helper module"); + fs::write( + &composer, + r#" + #[path = "../shared/helper.rs"] + mod helper; + "#, + ) + .expect("write transitive parent"); + fs::write(&helper, "mod leaf;").expect("write transitive child"); + fs::write( + &leaf, + r#" + use crate::pi_config::gateway::PiGatewayApiFamily; + fn production_escape_attempt() { + save_provider(); + } + "#, + ) + .expect("write transitive leaf"); + let (violations, _) = scan_production_tree( + transitive_path_root.path(), + &transitive_path_root.path().join("src"), + ); + for expected_kind in ["cross_layer_import", "forbidden_provider_symbol"] { + assert!( + violations.iter().any(|violation| { + violation.kind == expected_kind + && violation.path.ends_with("src/shared/helper/leaf.rs") + }), + "ordinary descendants of #[path] modules must inherit composer ownership and \ + trigger {expected_kind}: {violations:?}" + ); + } +} diff --git a/src-tauri/src/commands/config.rs b/src-tauri/src/commands/config.rs index d09377162..6cbf0e128 100644 --- a/src-tauri/src/commands/config.rs +++ b/src-tauri/src/commands/config.rs @@ -135,6 +135,18 @@ pub async fn get_config_status( Ok(ConfigStatus { exists, path }) } + AppType::Pi => { + let config_path = + crate::pi_config::native::get_pi_models_path().map_err(|e| e.to_string())?; + let path = crate::pi_config::native::get_pi_agent_dir() + .map_err(|e| e.to_string())? + .to_string_lossy() + .to_string(); + Ok(ConfigStatus { + exists: config_path.exists(), + path, + }) + } } } @@ -156,6 +168,7 @@ pub async fn get_config_dir(app: String) -> Result { AppType::OpenCode => crate::opencode_config::get_opencode_dir(), AppType::OpenClaw => crate::openclaw_config::get_openclaw_dir(), AppType::Hermes => crate::hermes_config::get_hermes_dir(), + AppType::Pi => crate::pi_config::native::get_pi_agent_dir().map_err(|e| e.to_string())?, }; Ok(dir.to_string_lossy().to_string()) @@ -174,6 +187,7 @@ pub async fn open_config_folder(handle: AppHandle, app: String) -> Result crate::opencode_config::get_opencode_dir(), AppType::OpenClaw => crate::openclaw_config::get_openclaw_dir(), AppType::Hermes => crate::hermes_config::get_hermes_dir(), + AppType::Pi => crate::pi_config::native::get_pi_agent_dir().map_err(|e| e.to_string())?, }; if !config_dir.exists() { diff --git a/src-tauri/src/commands/failover.rs b/src-tauri/src/commands/failover.rs index a0c40ce66..dab762cfe 100644 --- a/src-tauri/src/commands/failover.rs +++ b/src-tauri/src/commands/failover.rs @@ -2,6 +2,7 @@ //! //! 管理代理模式下的故障转移队列(基于 providers 表的 in_failover_queue 字段) +use crate::app_config::AppType; use crate::database::FailoverQueueItem; use crate::provider::Provider; use crate::store::AppState; @@ -39,6 +40,50 @@ pub async fn add_to_failover_queue( app_type: String, provider_id: String, ) -> Result<(), String> { + if app_type == "pi" { + let _guard = state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()) + .await; + if state + .db + .get_provider_aggregate("pi", &provider_id) + .map_err(|error| error.to_string())? + .is_none() + { + return Err(format!("Pi provider does not exist: {provider_id}")); + } + let was_member = state + .db + .is_in_failover_queue("pi", &provider_id) + .map_err(|error| error.to_string())?; + let epoch = state.proxy_service.begin_pi_catalog_mutation().await; + if let Err(error) = state.db.add_to_failover_queue("pi", &provider_id) { + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await; + return Err(error.to_string()); + } + if let Err(error) = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await + { + if !was_member { + let _ = state.db.remove_from_failover_queue("pi", &provider_id); + } + let rollback_epoch = state.proxy_service.begin_pi_catalog_mutation().await; + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(rollback_epoch) + .await; + return Err(format!( + "Pi failover queue changed but runtime publication failed: {error}" + )); + } + return Ok(()); + } state .db .add_to_failover_queue(&app_type, &provider_id) @@ -52,6 +97,42 @@ pub async fn remove_from_failover_queue( app_type: String, provider_id: String, ) -> Result<(), String> { + if app_type == "pi" { + let _guard = state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()) + .await; + let was_member = state + .db + .is_in_failover_queue("pi", &provider_id) + .map_err(|error| error.to_string())?; + let epoch = state.proxy_service.begin_pi_catalog_mutation().await; + if let Err(error) = state.db.remove_from_failover_queue("pi", &provider_id) { + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await; + return Err(error.to_string()); + } + if let Err(error) = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await + { + if was_member { + let _ = state.db.add_to_failover_queue("pi", &provider_id); + } + let rollback_epoch = state.proxy_service.begin_pi_catalog_mutation().await; + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(rollback_epoch) + .await; + return Err(format!( + "Pi failover queue changed but runtime publication failed: {error}" + )); + } + return Ok(()); + } state .db .remove_from_failover_queue(&app_type, &provider_id) @@ -64,6 +145,9 @@ pub async fn get_auto_failover_enabled( state: tauri::State<'_, AppState>, app_type: String, ) -> Result { + if app_type == "pi" { + return Ok(crate::settings::get_pi_proxy_settings().auto_failover_enabled); + } state .db .get_proxy_config_for_app(&app_type) @@ -86,6 +170,10 @@ pub async fn set_auto_failover_enabled( "[Failover] Setting auto_failover_enabled: app_type='{app_type}', enabled={enabled}" ); + if app_type == "pi" { + return set_pi_auto_failover_enabled(&app, state.inner(), enabled).await; + } + // 读取当前配置 let mut config = state .db @@ -180,3 +268,88 @@ pub async fn set_auto_failover_enabled( Ok(()) } + +async fn set_pi_auto_failover_enabled( + app: &tauri::AppHandle, + state: &AppState, + enabled: bool, +) -> Result<(), String> { + let _guard = state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()) + .await; + let previous_config = crate::settings::get_pi_proxy_settings(); + if enabled && !crate::settings::pi_takeover_enabled() { + return Err("Pi gateway takeover must be enabled before failover".to_string()); + } + + let mut auto_added = None; + if enabled + && state + .db + .get_failover_queue("pi") + .map_err(|error| error.to_string())? + .is_empty() + { + let current = + crate::services::pi_catalog::PiCatalogCoordinator::current_native_provider(state) + .map_err(|error| error.to_string())? + .ok_or_else(|| { + "Pi failover queue is empty and no current provider is selected".to_string() + })?; + state + .db + .add_to_failover_queue("pi", ¤t) + .map_err(|error| error.to_string())?; + auto_added = Some(current); + } + + let mut next = previous_config.clone(); + next.auto_failover_enabled = enabled; + let epoch = state.proxy_service.begin_pi_catalog_mutation().await; + if let Err(error) = crate::settings::update_pi_proxy_settings(next) { + if let Some(provider_id) = auto_added { + let _ = state.db.remove_from_failover_queue("pi", &provider_id); + } + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await; + return Err(error.to_string()); + } + if let Err(error) = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await + { + let _ = crate::settings::update_pi_proxy_settings(previous_config); + if let Some(provider_id) = auto_added { + let _ = state.db.remove_from_failover_queue("pi", &provider_id); + } + let rollback_epoch = state.proxy_service.begin_pi_catalog_mutation().await; + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(rollback_epoch) + .await; + return Err(format!( + "Pi failover preference changed but runtime publication failed: {error}" + )); + } + + let _ = app.emit( + "provider-switched", + serde_json::json!({ + "appType": "pi", + "providerId": + crate::services::pi_catalog::PiCatalogCoordinator::current_native_provider(state) + .map_err(|error| error.to_string())?, + "source": "failoverPreferenceChanged" + }), + ); + if let Ok(new_menu) = crate::tray::create_tray_menu(app, state) { + if let Some(tray) = app.tray_by_id(crate::tray::TRAY_ID) { + let _ = tray.set_menu(Some(new_menu)); + } + } + Ok(()) +} diff --git a/src-tauri/src/commands/import_export.rs b/src-tauri/src/commands/import_export.rs index 048935b90..7734d6a84 100644 --- a/src-tauri/src/commands/import_export.rs +++ b/src-tauri/src/commands/import_export.rs @@ -44,26 +44,56 @@ pub async fn import_config_from_file( state: State<'_, AppState>, ) -> Result { let db = state.db.clone(); - let db_for_sync = db.clone(); - tauri::async_runtime::spawn_blocking(move || { - let path_buf = PathBuf::from(&filePath); - let backup_id = db.import_sql(&path_buf)?; - let warning = post_sync_warning_from_result(Ok(run_post_import_sync(db_for_sync))); - if let Some(msg) = warning.as_ref() { - log::warn!("[Import] post-import sync warning: {msg}"); + let app_state = state.inner().clone(); + let pi_guard = app_state + .proxy_service + .lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) + .await; + app_state + .proxy_service + .prepare_pi_portable_import_under_lock(&pi_guard) + .await + .map_err(|error| format!("导入前恢复 Pi 直连投影失败: {error}"))?; + + let import_path = filePath.clone(); + let import_result = + tauri::async_runtime::spawn_blocking(move || db.import_sql(&PathBuf::from(import_path))) + .await + .map_err(|error| AppError::Message(format!("SQL import task failed: {error}"))) + .and_then(|result| result); + let backup_id = match import_result { + Ok(backup_id) => backup_id, + Err(error) => { + let recovery = app_state + .proxy_service + .recover_pi_after_aborted_portable_import_under_lock(&pi_guard) + .await; + return Err(match recovery { + Ok(()) => error.to_string(), + Err(recovery) => { + format!("{error}; Pi gateway recovery after aborted import failed: {recovery}") + } + }); } - Ok::<_, AppError>(success_payload_with_warning(backup_id, warning)) - }) - .await - .map_err(|e| format!("导入配置失败: {e}"))? - .map_err(|e: AppError| e.to_string()) + }; + drop(pi_guard); + + let sync_state = app_state.clone(); + let warning = post_sync_warning_from_result( + tauri::async_runtime::spawn_blocking(move || run_post_import_sync(&sync_state)) + .await + .map_err(|error| error.to_string()), + ); + if let Some(msg) = warning.as_ref() { + log::warn!("[Import] post-import sync warning: {msg}"); + } + Ok(success_payload_with_warning(backup_id, warning)) } #[tauri::command] pub async fn sync_current_providers_live(state: State<'_, AppState>) -> Result { - let db = state.db.clone(); + let app_state = state.inner().clone(); tauri::async_runtime::spawn_blocking(move || { - let app_state = AppState::new(db); ProviderService::sync_current_to_live(&app_state)?; Ok::<_, AppError>(json!({ "success": true, @@ -154,10 +184,50 @@ pub async fn restore_db_backup( filename: String, ) -> Result { let db = state.db.clone(); - tauri::async_runtime::spawn_blocking(move || db.restore_from_backup(&filename)) + let app_state = state.inner().clone(); + let pi_guard = app_state + .proxy_service + .lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) + .await; + app_state + .proxy_service + .prepare_pi_portable_import_under_lock(&pi_guard) .await - .map_err(|e| format!("Restore failed: {e}"))? - .map_err(|e: AppError| e.to_string()) + .map_err(|error| format!("Restore preparation failed: {error}"))?; + + let restore_result = + tauri::async_runtime::spawn_blocking(move || db.restore_from_backup(&filename)) + .await + .map_err(|error| AppError::Message(format!("Restore task failed: {error}"))) + .and_then(|result| result); + let restored = match restore_result { + Ok(restored) => restored, + Err(error) => { + let recovery = app_state + .proxy_service + .recover_pi_after_aborted_portable_import_under_lock(&pi_guard) + .await; + return Err(match recovery { + Ok(()) => error.to_string(), + Err(recovery) => { + format!("{error}; Pi gateway recovery after aborted restore failed: {recovery}") + } + }); + } + }; + drop(pi_guard); + + let sync_state = app_state.clone(); + match tauri::async_runtime::spawn_blocking(move || run_post_import_sync(&sync_state)).await { + Ok(Ok(())) => {} + Ok(Err(error)) => { + log::warn!("[Restore] database restored but post-restore sync failed: {error}"); + } + Err(error) => { + log::warn!("[Restore] database restored but post-restore sync task failed: {error}"); + } + } + Ok(restored) } /// Rename a database backup file diff --git a/src-tauri/src/commands/misc.rs b/src-tauri/src/commands/misc.rs index 0cf46e426..f5e8e8a29 100644 --- a/src-tauri/src/commands/misc.rs +++ b/src-tauri/src/commands/misc.rs @@ -111,8 +111,8 @@ pub struct ToolVersion { wsl_distro: Option, } -const VALID_TOOLS: [&str; 7] = [ - "claude", "codex", "gemini", "grok", "opencode", "openclaw", "hermes", +const VALID_TOOLS: [&str; 8] = [ + "claude", "codex", "gemini", "grok", "opencode", "openclaw", "hermes", "pi", ]; #[derive(Debug, Clone, serde::Deserialize)] @@ -433,6 +433,7 @@ fn tool_display_name(tool: &str) -> &'static str { "opencode" => "OpenCode", "openclaw" => "OpenClaw", "hermes" => "Hermes", + "pi" => "Pi", _ => "Unknown", } } @@ -513,6 +514,7 @@ fn npm_install_command_for(tool: &str) -> Option<&'static str> { "grok" => Some("npm i -g @xai-official/grok@latest"), "opencode" => Some("npm i -g opencode-ai@latest"), "openclaw" => Some("npm i -g openclaw@latest"), + "pi" => Some("npm i -g @earendil-works/pi-coding-agent@latest"), _ => None, } } @@ -807,6 +809,9 @@ async fn get_single_tool_version_impl( } "openclaw" => fetch_npm_latest_for_tool(&client, "openclaw", tool, local).await, "hermes" => fetch_pypi_latest_version(&client, "hermes-agent").await, + "pi" => { + fetch_npm_latest_for_tool(&client, "@earendil-works/pi-coding-agent", tool, local).await + } _ => None, }; @@ -2071,6 +2076,7 @@ fn npm_package_for(tool: &str) -> Option<&'static str> { "grok" => Some("@xai-official/grok"), "opencode" => Some("opencode-ai"), "openclaw" => Some("openclaw"), + "pi" => Some("@earendil-works/pi-coding-agent"), _ => None, } } @@ -2789,6 +2795,7 @@ fn wsl_distro_for_tool(tool: &str) -> Option { "opencode" => crate::settings::get_opencode_override_dir(), "openclaw" => crate::settings::get_openclaw_override_dir(), "hermes" => crate::settings::get_hermes_override_dir(), + "pi" => crate::settings::get_pi_override_dir(), _ => None, }?; @@ -3926,6 +3933,24 @@ mod tests { ); } + #[test] + fn pi_lifecycle_metadata_matches_pinned_distribution() { + let requested = vec!["unsupported".to_string(), "pi".to_string()]; + assert_eq!(normalize_requested_tools(&requested), vec!["pi"]); + assert_eq!(tool_display_name("pi"), "Pi"); + assert_eq!( + npm_package_for("pi"), + Some("@earendil-works/pi-coding-agent") + ); + assert_eq!( + npm_install_command_for("pi"), + Some("npm i -g @earendil-works/pi-coding-agent@latest") + ); + // The verified distribution exposes `pi --version`, but no updater + // contract is assumed; upgrades stay on the package-manager path. + assert_eq!(official_update_args("pi"), None); + } + #[test] fn test_compare_semver() { use std::cmp::Ordering; @@ -5331,6 +5356,13 @@ mod tests { assert_eq!(cmd, "npm i -g openclaw@latest"); } + #[test] + fn pi_install_uses_the_verified_pinned_package() { + let cmd = install_command_for("pi"); + assert_eq!(cmd, "npm i -g @earendil-works/pi-coding-agent@latest"); + assert!(!cmd.contains("||")); + } + #[test] fn update_fallbacks_use_official_cli_only_when_supported() { assert_eq!( @@ -5360,6 +5392,11 @@ mod tests { static_fallback_command("openclaw"), "openclaw update --yes || npm i -g openclaw@latest" ); + assert_eq!( + static_fallback_command("pi"), + "npm i -g @earendil-works/pi-coding-agent@latest" + ); + assert!(!static_fallback_command("pi").contains("pi update")); } #[test] diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 08e270919..b289bf7c7 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -17,6 +17,7 @@ mod misc; mod model_fetch; mod omo; mod openclaw; +mod pi; mod plugin; mod profile; mod prompt; @@ -53,6 +54,7 @@ pub use misc::*; pub use model_fetch::*; pub use omo::*; pub use openclaw::*; +pub(crate) use pi::*; pub use plugin::*; pub use profile::*; pub use prompt::*; diff --git a/src-tauri/src/commands/pi.rs b/src-tauri/src/commands/pi.rs new file mode 100644 index 000000000..4dadc7059 --- /dev/null +++ b/src-tauri/src/commands/pi.rs @@ -0,0 +1,76 @@ +use crate::pi_config::model::PiNativeDiagnostic; +use crate::pi_config::native_settings::{read_pi_native_defaults, PiNativeDefaults}; +use crate::services::pi_catalog::{PiCatalogCoordinator, PiCatalogMutation}; +use crate::session_manager::providers::pi::PiSessionDiscovery; +use crate::store::AppState; +use tauri::State; + +/// Read-only diagnostics come exclusively from the Pre-C certified inspection +/// service. This command does not infer manageability or gateway status. +#[tauri::command] +pub(crate) fn get_pi_native_catalog( + state: State<'_, AppState>, +) -> Result, String> { + PiCatalogCoordinator::inspect_native(state.inner()).map_err(|error| error.to_string()) +} + +#[tauri::command] +pub(crate) fn import_pi_native_provider( + state: State<'_, AppState>, + #[allow(non_snake_case)] providerKey: String, + #[allow(non_snake_case)] expectedFingerprint: String, +) -> Result { + let result = PiCatalogCoordinator::apply( + state.inner(), + PiCatalogMutation::ImportNative { + provider_key: providerKey, + expected_fingerprint: expectedFingerprint, + }, + ) + .map_err(|error| error.to_string())?; + result + .provider_id + .ok_or_else(|| "Pi import did not return a provider id".to_string()) +} + +#[tauri::command] +pub(crate) fn set_pi_default_model( + state: State<'_, AppState>, + #[allow(non_snake_case)] providerId: String, + #[allow(non_snake_case)] modelId: String, +) -> Result { + PiCatalogCoordinator::apply( + state.inner(), + PiCatalogMutation::SetDefault { + provider_id: providerId, + model_id: modelId, + }, + ) + .map(|_| true) + .map_err(|error| error.to_string()) +} + +#[tauri::command] +pub(crate) fn get_pi_native_defaults() -> Result { + read_pi_native_defaults().map_err(|error| error.to_string()) +} + +#[tauri::command] +pub(crate) fn get_pi_session_discovery() -> PiSessionDiscovery { + crate::session_manager::providers::pi::session_discovery() +} + +/// Explicitly rotate the device-local gateway bearer and republish every +/// managed Pi projection. Existing Pi processes must restart because they +/// retain the previous projected credential in memory. +#[tauri::command] +pub(crate) async fn reset_pi_gateway_credential( + state: State<'_, AppState>, +) -> Result { + state + .proxy_service + .rotate_pi_gateway_token() + .await + .map(|()| true) + .map_err(|error| error.to_string()) +} diff --git a/src-tauri/src/commands/prompt.rs b/src-tauri/src/commands/prompt.rs index 20bd9f2a3..8f0ebcdc9 100644 --- a/src-tauri/src/commands/prompt.rs +++ b/src-tauri/src/commands/prompt.rs @@ -5,7 +5,11 @@ use tauri::State; use crate::app_config::AppType; use crate::prompt::Prompt; -use crate::services::PromptService; +use crate::services::pi_prompt_files::{ + PiPromptFileKind, PiPromptFileService, PiPromptFileSnapshot, PiPromptTemplate, + PiPromptTemplateService, +}; +use crate::services::prompt::{PiPromptLibraryStatus, PromptService}; use crate::store::AppState; #[tauri::command] @@ -62,3 +66,61 @@ pub async fn get_current_prompt_file_content(app: String) -> Result, +) -> Result { + PromptService::get_pi_library_status(&state).map_err(|error| error.to_string()) +} + +#[tauri::command] +pub async fn reconcile_pi_prompt_library(state: State<'_, AppState>) -> Result<(), String> { + PromptService::reconcile_pi_library(&state).map_err(|error| error.to_string()) +} + +#[tauri::command] +pub async fn get_pi_prompt_file(kind: PiPromptFileKind) -> Result { + PiPromptFileService::read(kind).map_err(|error| error.to_string()) +} + +#[tauri::command] +pub async fn replace_pi_prompt_file( + kind: PiPromptFileKind, + #[allow(non_snake_case)] expectedRevision: String, + content: String, +) -> Result { + PiPromptFileService::replace(kind, &expectedRevision, &content) + .map_err(|error| error.to_string()) +} + +#[tauri::command] +pub async fn delete_pi_prompt_file( + kind: PiPromptFileKind, + #[allow(non_snake_case)] expectedRevision: String, +) -> Result { + PiPromptFileService::delete(kind, &expectedRevision).map_err(|error| error.to_string()) +} + +#[tauri::command] +pub async fn list_pi_prompt_templates() -> Result, String> { + PiPromptTemplateService::list().map_err(|error| error.to_string()) +} + +#[tauri::command] +pub async fn upsert_pi_prompt_template( + slug: String, + #[allow(non_snake_case)] expectedRevision: String, + content: String, +) -> Result { + PiPromptTemplateService::upsert(&slug, &expectedRevision, &content) + .map_err(|error| error.to_string()) +} + +#[tauri::command] +pub async fn delete_pi_prompt_template( + slug: String, + #[allow(non_snake_case)] expectedRevision: String, +) -> Result { + PiPromptTemplateService::delete(&slug, &expectedRevision).map_err(|error| error.to_string()) +} diff --git a/src-tauri/src/commands/proxy.rs b/src-tauri/src/commands/proxy.rs index 05bc08ae4..73e0cb63b 100644 --- a/src-tauri/src/commands/proxy.rs +++ b/src-tauri/src/commands/proxy.rs @@ -26,6 +26,7 @@ pub async fn stop_proxy_server(state: tauri::State<'_, AppState>) -> Result<(), || takeover.grokbuild || takeover.opencode || takeover.openclaw + || takeover.pi { return Err( "仍有应用处于代理接管状态,请先在设置中关闭对应应用接管后再停止本地路由。".to_string(), @@ -120,6 +121,9 @@ pub async fn get_proxy_config_for_app( state: tauri::State<'_, AppState>, app_type: String, ) -> Result { + if app_type == "pi" { + return Ok(crate::settings::get_pi_app_proxy_config()); + } let db = &state.db; db.get_proxy_config_for_app(&app_type) .await @@ -138,6 +142,60 @@ pub async fn update_proxy_config_for_app( let app_type = config.app_type.clone(); let circuit_config = CircuitBreakerConfig::from(&config); + if app_type == "pi" { + let _guard = state + .proxy_service + .lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) + .await; + let previous = crate::settings::get_pi_proxy_settings(); + if config.enabled != crate::settings::pi_takeover_enabled() { + return Err( + "Pi enabled state is owned by set_proxy_takeover_for_app, not proxy config" + .to_string(), + ); + } + let next = crate::settings::PiProxySettings { + auto_failover_enabled: config.auto_failover_enabled, + max_retries: config.max_retries, + streaming_first_byte_timeout: config.streaming_first_byte_timeout, + streaming_idle_timeout: config.streaming_idle_timeout, + non_streaming_timeout: config.non_streaming_timeout, + circuit_failure_threshold: config.circuit_failure_threshold, + circuit_success_threshold: config.circuit_success_threshold, + circuit_timeout_seconds: config.circuit_timeout_seconds, + circuit_error_rate_threshold: config.circuit_error_rate_threshold, + circuit_min_requests: config.circuit_min_requests, + }; + let epoch = state.proxy_service.begin_pi_catalog_mutation().await; + if let Err(error) = crate::settings::update_pi_proxy_settings(next) { + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await; + return Err(error.to_string()); + } + if let Err(error) = state + .proxy_service + .reconcile_pi_runtime_at_epoch(epoch) + .await + { + let _ = crate::settings::update_pi_proxy_settings(previous); + let rollback_epoch = state.proxy_service.begin_pi_catalog_mutation().await; + let _ = state + .proxy_service + .reconcile_pi_runtime_at_epoch(rollback_epoch) + .await; + return Err(format!( + "Pi proxy config changed but runtime publication failed: {error}" + )); + } + state + .proxy_service + .update_circuit_breaker_config_for_app(&app_type, circuit_config) + .await?; + return Ok(()); + } + db.update_proxy_config_for_app(config) .await .map_err(|e| e.to_string())?; diff --git a/src-tauri/src/commands/s3_sync.rs b/src-tauri/src/commands/s3_sync.rs index 462b20439..738e4e326 100644 --- a/src-tauri/src/commands/s3_sync.rs +++ b/src-tauri/src/commands/s3_sync.rs @@ -107,18 +107,44 @@ pub async fn s3_sync_upload(state: State<'_, AppState>) -> Result #[tauri::command] pub async fn s3_sync_download(state: State<'_, AppState>) -> Result { let db = state.db.clone(); - let db_for_sync = db.clone(); + let app_state = state.inner().clone(); let mut settings = require_enabled_s3_settings()?; let _auto_sync_suppression = crate::services::s3_auto_sync::AutoSyncSuppressionGuard::new(); + let pi_guard = app_state + .proxy_service + .lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) + .await; + app_state + .proxy_service + .prepare_pi_portable_import_under_lock(&pi_guard) + .await + .map_err(|error| format!("S3 下载前恢复 Pi 直连投影失败: {error}"))?; let sync_result = run_with_s3_lock(s3_sync_service::download(&db, &mut settings)).await; - let mut result = map_sync_result(sync_result, |error| { - persist_sync_error(&mut settings, error, "manual") - })?; + let mut result = match sync_result { + Ok(result) => result, + Err(error) => { + persist_sync_error(&mut settings, &error, "manual"); + let recovery = app_state + .proxy_service + .recover_pi_after_aborted_portable_import_under_lock(&pi_guard) + .await; + return Err(match recovery { + Ok(()) => error.to_string(), + Err(recovery) => { + format!( + "{error}; Pi gateway recovery after aborted S3 download failed: {recovery}" + ) + } + }); + } + }; + drop(pi_guard); // Post-download sync is best-effort: snapshot restore has already succeeded. + let sync_state = app_state.clone(); let warning = post_sync_warning_from_result( - tauri::async_runtime::spawn_blocking(move || run_post_import_sync(db_for_sync)) + tauri::async_runtime::spawn_blocking(move || run_post_import_sync(&sync_state)) .await .map_err(|e| e.to_string()), ); diff --git a/src-tauri/src/commands/settings.rs b/src-tauri/src/commands/settings.rs index f6cdb8111..855848db0 100644 --- a/src-tauri/src/commands/settings.rs +++ b/src-tauri/src/commands/settings.rs @@ -48,6 +48,13 @@ fn merge_settings_for_save( // 开关)后、前端 query 缓存刷新前的一次全量保存会把旧 marker 重放回来, // 重新开启时被"复活"的标记挡住而漏迁。 incoming.local_migrations = existing.local_migrations.clone(); + // Pi gateway credential is an installation secret. Settings IPC can + // neither observe it (frontend projection clears it) nor mutate it. + incoming.pi_gateway_token = existing.pi_gateway_token.clone(); + // Pi proxy behavior is committed through the proxy commands so a generic + // settings round-trip cannot bypass the switch/epoch publication boundary. + incoming.pi_takeover_enabled = existing.pi_takeover_enabled; + incoming.pi_proxy = existing.pi_proxy.clone(); incoming } @@ -63,12 +70,29 @@ pub async fn save_settings( state: tauri::State<'_, crate::store::AppState>, settings: crate::settings::AppSettings, ) -> Result { + // The frontend settings projection intentionally cannot mutate Pi's + // takeover bit or gateway secret. Serialize the read/merge/write with Pi + // catalog mutations so a concurrent toggle cannot be overwritten by a + // stale full-settings payload. + let pi_guard = state + .proxy_service + .lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) + .await; let existing = crate::settings::get_settings(); let merged = merge_settings_for_save(settings, &existing); let unify_codex_changed = merged.unify_codex_session_history != existing.unify_codex_session_history; let unify_codex_enabled = merged.unify_codex_session_history; - crate::settings::update_settings(merged).map_err(|e| e.to_string())?; + state + .proxy_service + .replace_settings_with_pi_directory_boundary_under_lock( + &pi_guard, + &existing, + merged.clone(), + ) + .await + .map_err(|e| e.to_string())?; + drop(pi_guard); // 统一会话开关变更时立即重写当前官方 Codex 供应商的 live 配置, // 不必等下一次切换才生效。 @@ -82,7 +106,18 @@ pub async fn save_settings( crate::services::provider::reapply_current_codex_official_live(state.inner()) { log::warn!("统一 Codex 会话历史开关变更后重写 live 配置失败,回滚设置: {err}"); - if let Err(rollback_err) = crate::settings::update_settings(existing) { + let pi_guard = state + .proxy_service + .lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) + .await; + let current = crate::settings::get_settings(); + if let Err(rollback_err) = state + .proxy_service + .replace_settings_with_pi_directory_boundary_under_lock( + &pi_guard, ¤t, existing, + ) + .await + { log::error!("回滚统一会话开关设置失败: {rollback_err}"); } return Err(format!( @@ -618,6 +653,28 @@ mod tests { assert!(merged.local_migrations.is_none()); } + + #[test] + fn save_settings_cannot_bypass_pi_gateway_publication_ownership() { + let existing = AppSettings { + pi_takeover_enabled: true, + pi_proxy: crate::settings::PiProxySettings { + max_retries: 7, + ..crate::settings::PiProxySettings::default() + }, + ..AppSettings::default() + }; + let incoming = AppSettings { + pi_takeover_enabled: false, + pi_proxy: crate::settings::PiProxySettings::default(), + ..AppSettings::default() + }; + + let merged = merge_settings_for_save(incoming, &existing); + + assert!(merged.pi_takeover_enabled); + assert_eq!(merged.pi_proxy.max_retries, 7); + } } /// 获取开机自启状态 diff --git a/src-tauri/src/commands/skill.rs b/src-tauri/src/commands/skill.rs index 6ac90ac54..44f41f55b 100644 --- a/src-tauri/src/commands/skill.rs +++ b/src-tauri/src/commands/skill.rs @@ -11,7 +11,9 @@ use crate::services::skill::{ SkillService, SkillStorageLocation, SkillUninstallResult, SkillUpdateInfo, SkillsShSearchResult, }; +use crate::services::skill_deployment::{PiSkillDeploymentService, SkillAppStatus}; use crate::store::AppState; +use std::collections::BTreeMap; use std::str::FromStr; use std::sync::Arc; use tauri::State; @@ -32,6 +34,13 @@ pub fn get_installed_skills(app_state: State<'_, AppState>) -> Result, +) -> Result, String> { + PiSkillDeploymentService::inspect_all(&app_state.db).map_err(|error| error.to_string()) +} + #[tauri::command] pub fn get_skill_backups() -> Result, String> { SkillService::list_backups().map_err(|e| e.to_string()) diff --git a/src-tauri/src/commands/sync_support.rs b/src-tauri/src/commands/sync_support.rs index 00793b3b8..4c685e194 100644 --- a/src-tauri/src/commands/sync_support.rs +++ b/src-tauri/src/commands/sync_support.rs @@ -1,15 +1,13 @@ -use serde_json::{json, Value}; -use std::sync::Arc; - -use crate::database::Database; use crate::error::AppError; use crate::services::provider::ProviderService; +use crate::services::PromptService; use crate::settings; use crate::store::AppState; +use serde_json::{json, Value}; -pub(crate) fn run_post_import_sync(db: Arc) -> Result<(), AppError> { - let app_state = AppState::new(db); - ProviderService::sync_current_to_live(&app_state)?; +pub(crate) fn run_post_import_sync(app_state: &AppState) -> Result<(), AppError> { + PromptService::reconcile_pi_portable_import(app_state)?; + ProviderService::sync_current_to_live(app_state)?; settings::reload_settings()?; Ok(()) } diff --git a/src-tauri/src/commands/webdav_sync.rs b/src-tauri/src/commands/webdav_sync.rs index 31879d13b..553d89827 100644 --- a/src-tauri/src/commands/webdav_sync.rs +++ b/src-tauri/src/commands/webdav_sync.rs @@ -115,18 +115,44 @@ pub async fn webdav_sync_upload(state: State<'_, AppState>) -> Result) -> Result { let db = state.db.clone(); - let db_for_sync = db.clone(); + let app_state = state.inner().clone(); let mut settings = require_enabled_webdav_settings()?; let _auto_sync_suppression = crate::services::webdav_auto_sync::AutoSyncSuppressionGuard::new(); + let pi_guard = app_state + .proxy_service + .lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) + .await; + app_state + .proxy_service + .prepare_pi_portable_import_under_lock(&pi_guard) + .await + .map_err(|error| format!("WebDAV 下载前恢复 Pi 直连投影失败: {error}"))?; let sync_result = run_with_webdav_lock(webdav_sync_service::download(&db, &mut settings)).await; - let mut result = map_sync_result(sync_result, |error| { - persist_sync_error(&mut settings, error, "manual") - })?; + let mut result = match sync_result { + Ok(result) => result, + Err(error) => { + persist_sync_error(&mut settings, &error, "manual"); + let recovery = app_state + .proxy_service + .recover_pi_after_aborted_portable_import_under_lock(&pi_guard) + .await; + return Err(match recovery { + Ok(()) => error.to_string(), + Err(recovery) => { + format!( + "{error}; Pi gateway recovery after aborted WebDAV download failed: {recovery}" + ) + } + }); + } + }; + drop(pi_guard); // Post-download sync is best-effort: snapshot restore has already succeeded. + let sync_state = app_state.clone(); let warning = post_sync_warning_from_result( - tauri::async_runtime::spawn_blocking(move || run_post_import_sync(db_for_sync)) + tauri::async_runtime::spawn_blocking(move || run_post_import_sync(&sync_state)) .await .map_err(|e| e.to_string()), ); diff --git a/src-tauri/src/config.rs b/src-tauri/src/config.rs index ed4686522..ac33dc66c 100644 --- a/src-tauri/src/config.rs +++ b/src-tauri/src/config.rs @@ -295,6 +295,22 @@ pub fn write_text_file(path: &Path, data: &str) -> Result<(), AppError> { /// 原子写入:写入临时文件后 rename 替换,避免半写状态 pub fn atomic_write(path: &Path, data: &[u8]) -> Result<(), AppError> { + atomic_write_durable(path, data, None) +} + +/// Durable same-directory atomic replacement. +/// +/// Existing permissions are preserved. `new_file_mode` controls only a newly +/// created Unix file (settings and other local secrets pass `0o600`). The +/// temporary file is created exclusively, synced before replacement, and the +/// containing directory is synced afterwards on Unix. +pub(crate) fn atomic_write_durable( + path: &Path, + data: &[u8], + new_file_mode: Option, +) -> Result<(), AppError> { + #[cfg(not(unix))] + let _ = new_file_mode; if let Some(parent) = path.parent() { fs::create_dir_all(parent).map_err(|e| AppError::io(parent, e))?; } @@ -302,51 +318,95 @@ pub fn atomic_write(path: &Path, data: &[u8]) -> Result<(), AppError> { let parent = path .parent() .ok_or_else(|| AppError::Config("无效的路径".to_string()))?; - let mut tmp = parent.to_path_buf(); let file_name = path .file_name() .ok_or_else(|| AppError::Config("无效的文件名".to_string()))? .to_string_lossy() .to_string(); - let ts = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_nanos(); - tmp.push(format!("{file_name}.tmp.{ts}")); + let tmp = parent.join(format!( + ".{file_name}.{}.tmp", + uuid::Uuid::new_v4().simple() + )); - { - let mut f = fs::File::create(&tmp).map_err(|e| AppError::io(&tmp, e))?; - f.write_all(data).map_err(|e| AppError::io(&tmp, e))?; - f.flush().map_err(|e| AppError::io(&tmp, e))?; - } - - #[cfg(unix)] - { - use std::os::unix::fs::PermissionsExt; - if let Ok(meta) = fs::metadata(path) { - let perm = meta.permissions().mode(); - let _ = fs::set_permissions(&tmp, fs::Permissions::from_mode(perm)); + let result = (|| -> Result<(), AppError> { + let mut options = fs::OpenOptions::new(); + options.create_new(true).write(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options.mode(new_file_mode.unwrap_or(0o666)); } - } + let mut file = options + .open(&tmp) + .map_err(|error| AppError::io(&tmp, error))?; + file.write_all(data) + .map_err(|error| AppError::io(&tmp, error))?; + file.flush().map_err(|error| AppError::io(&tmp, error))?; + file.sync_all().map_err(|error| AppError::io(&tmp, error))?; + drop(file); - #[cfg(windows)] - { - // Windows 上 rename 目标存在会失败,先移除再重命名(尽量接近原子性) - if path.exists() { - let _ = fs::remove_file(path); + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + let mode = fs::metadata(path) + .map(|metadata| metadata.permissions().mode()) + .unwrap_or_else(|_| new_file_mode.unwrap_or(0o666)); + fs::set_permissions(&tmp, fs::Permissions::from_mode(mode)) + .map_err(|error| AppError::io(&tmp, error))?; } - fs::rename(&tmp, path).map_err(|e| AppError::IoContext { - context: format!("原子替换失败: {} -> {}", tmp.display(), path.display()), - source: e, - })?; - } - #[cfg(not(windows))] - { - fs::rename(&tmp, path).map_err(|e| AppError::IoContext { - context: format!("原子替换失败: {} -> {}", tmp.display(), path.display()), - source: e, - })?; + replace_file_atomically(&tmp, path)?; + #[cfg(unix)] + fs::File::open(parent) + .and_then(|directory| directory.sync_all()) + .map_err(|error| AppError::io(parent, error))?; + Ok(()) + })(); + if result.is_err() { + let _ = fs::remove_file(&tmp); + } + result +} + +#[cfg(not(windows))] +fn replace_file_atomically(temp_path: &Path, path: &Path) -> Result<(), AppError> { + fs::rename(temp_path, path).map_err(|source| AppError::IoContext { + context: format!( + "原子替换失败: {} -> {}", + temp_path.display(), + path.display() + ), + source, + }) +} + +#[cfg(windows)] +fn replace_file_atomically(temp_path: &Path, path: &Path) -> Result<(), AppError> { + use std::os::windows::ffi::OsStrExt; + use windows_sys::Win32::Storage::FileSystem::{ + MoveFileExW, MOVEFILE_REPLACE_EXISTING, MOVEFILE_WRITE_THROUGH, + }; + + let source: Vec = temp_path.as_os_str().encode_wide().chain(Some(0)).collect(); + let destination: Vec = path.as_os_str().encode_wide().chain(Some(0)).collect(); + // SAFETY: both buffers are NUL-terminated and remain alive for the + // duration of this synchronous Win32 call. + let moved = unsafe { + MoveFileExW( + source.as_ptr(), + destination.as_ptr(), + MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH, + ) + }; + if moved == 0 { + return Err(AppError::IoContext { + context: format!( + "原子替换失败: {} -> {}", + temp_path.display(), + path.display() + ), + source: std::io::Error::last_os_error(), + }); } Ok(()) } diff --git a/src-tauri/src/database/dao/mod.rs b/src-tauri/src/database/dao/mod.rs index f2b7ed1b6..15d174143 100644 --- a/src-tauri/src/database/dao/mod.rs +++ b/src-tauri/src/database/dao/mod.rs @@ -4,6 +4,8 @@ pub mod failover; pub mod mcp; +pub(crate) mod pi_catalog; +pub mod pi_projections; pub mod profiles; pub mod prompts; pub mod provider_write; @@ -13,6 +15,7 @@ pub mod providers; pub mod providers_seed; pub mod proxy; pub mod settings; +pub mod skill_deployments; pub mod skills; pub mod stream_check; pub mod universal_providers; diff --git a/src-tauri/src/database/dao/pi_catalog.rs b/src-tauri/src/database/dao/pi_catalog.rs new file mode 100644 index 000000000..6ce69a006 --- /dev/null +++ b/src-tauri/src/database/dao/pi_catalog.rs @@ -0,0 +1,276 @@ +//! Transactional database half of Pi catalog coordination. +//! +//! Provider row/endpoint SQL remains owned by the certified provider-write +//! primitives. This module only composes those primitives with Pi's exact-key +//! ownership ledger in one SQLite transaction. + +use super::pi_projections::PiProviderProjection; +use super::provider_write::{ + insert_endpoint, insert_row, restore_provider_aggregate_on_tx, NewEndpoint, + NewProviderAggregate, ProviderKey, ProviderRowUpdate, +}; +use super::providers::delete_provider_on_tx; +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use crate::provider::{ProviderAggregate, ProviderMutationInput}; +use indexmap::IndexMap; +use rusqlite::params; + +impl Database { + pub(crate) fn restore_pi_catalog_snapshot( + &self, + aggregates: &IndexMap, + projections: &[PiProviderProjection], + current_provider: Option<&str>, + ) -> Result<(), AppError> { + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + tx.execute("DELETE FROM pi_provider_projections", []) + .map_err(|error| AppError::Database(error.to_string()))?; + let current_ids = { + let mut statement = tx + .prepare("SELECT id FROM providers WHERE app_type = 'pi'") + .map_err(|error| AppError::Database(error.to_string()))?; + let ids = statement + .query_map([], |row| row.get::<_, String>(0)) + .map_err(|error| AppError::Database(error.to_string()))? + .collect::, _>>() + .map_err(|error| AppError::Database(error.to_string()))?; + ids + }; + for provider_id in current_ids + .iter() + .filter(|provider_id| !aggregates.contains_key(provider_id.as_str())) + { + // Only rows created after the snapshot are removed. Updating + // providers which existed in the snapshot preserves dependent + // provider_health history instead of triggering ON DELETE CASCADE. + delete_provider_on_tx(&tx, "pi", provider_id)?; + } + + for aggregate in aggregates.values() { + let key = ProviderKey::new("pi", aggregate.provider.id.clone())?; + let mut input = provider_mutation_input(aggregate); + if let Some(meta) = input.meta.as_mut() { + meta.custom_endpoints.clear(); + } + let row = ProviderRowUpdate::from_input(&input)?; + let endpoints = aggregate + .endpoints + .values() + .cloned() + .map(NewEndpoint::try_from) + .collect::, _>>()?; + restore_provider_aggregate_on_tx( + &tx, + &key, + &row, + aggregate.provider.created_at, + aggregate.provider.sort_index, + current_provider == Some(key.id()), + aggregate.provider.in_failover_queue, + &endpoints, + )?; + } + for projection in projections { + tx.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4)", + params![ + projection.provider_id, + projection.provider_key, + projection.created_at, + projection.updated_at + ], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn create_pi_catalog_provider( + &self, + input: NewProviderAggregate, + provider_key: &str, + ) -> Result { + if input.key.app_type() != "pi" || provider_key.trim().is_empty() { + return Err(AppError::InvalidInput( + "Pi catalog create requires app_type=pi and a non-empty native key".to_string(), + )); + } + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + insert_row( + &tx, + &input.key, + &input.row.content, + input.row.created_at, + input.sort_index, + false, + input.in_failover_queue, + )?; + for endpoint in &input.initial_endpoints { + insert_endpoint(&tx, &input.key, endpoint)?; + } + let now = chrono::Utc::now().timestamp_millis(); + tx.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES (?1, ?2, ?3, ?3)", + params![input.key.id(), provider_key, now], + ) + .map_err(|error| match &error { + rusqlite::Error::SqliteFailure(code, _) + if matches!( + code.extended_code, + rusqlite::ffi::SQLITE_CONSTRAINT_PRIMARYKEY + | rusqlite::ffi::SQLITE_CONSTRAINT_UNIQUE + ) => + { + AppError::Conflict(format!( + "Pi native provider key '{provider_key}' is already claimed" + )) + } + _ => AppError::Database(error.to_string()), + })?; + tx.commit() + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(PiProviderProjection { + provider_id: input.key.id().to_string(), + provider_key: provider_key.to_string(), + created_at: now, + updated_at: now, + }) + } + + pub(crate) fn update_pi_catalog_provider( + &self, + key: &ProviderKey, + row: &ProviderRowUpdate, + ) -> Result<(), AppError> { + if key.app_type() != "pi" { + return Err(AppError::InvalidInput( + "Pi catalog update requires app_type=pi".to_string(), + )); + } + self.update_provider(key, row) + } + + pub(crate) fn delete_pi_catalog_provider( + &self, + provider_id: &str, + ) -> Result, AppError> { + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + let projection = tx + .query_row( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections + WHERE provider_id = ?1", + [provider_id], + |row| { + Ok(PiProviderProjection { + provider_id: row.get(0)?, + provider_key: row.get(1)?, + created_at: row.get(2)?, + updated_at: row.get(3)?, + }) + }, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))?; + delete_provider_on_tx(&tx, "pi", provider_id)?; + tx.execute( + "DELETE FROM pi_provider_projections WHERE provider_id = ?1", + [provider_id], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + tx.commit() + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(projection) + } + + pub(crate) fn restore_pi_catalog_provider( + &self, + aggregate: &ProviderAggregate, + was_current: bool, + projection: Option<&PiProviderProjection>, + ) -> Result<(), AppError> { + let key = ProviderKey::new("pi", aggregate.provider.id.clone())?; + let mut input = provider_mutation_input(aggregate); + if let Some(meta) = input.meta.as_mut() { + meta.custom_endpoints.clear(); + } + let row = ProviderRowUpdate::from_input(&input)?; + let endpoints = aggregate + .endpoints + .values() + .cloned() + .map(NewEndpoint::try_from) + .collect::, _>>()?; + + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + restore_provider_aggregate_on_tx( + &tx, + &key, + &row, + aggregate.provider.created_at, + aggregate.provider.sort_index, + was_current, + aggregate.provider.in_failover_queue, + &endpoints, + )?; + tx.execute( + "DELETE FROM pi_provider_projections WHERE provider_id = ?1", + [key.id()], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if let Some(projection) = projection { + tx.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4)", + params![ + projection.provider_id, + projection.provider_key, + projection.created_at, + projection.updated_at + ], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) + } +} + +fn provider_mutation_input(aggregate: &ProviderAggregate) -> ProviderMutationInput { + let provider = &aggregate.provider; + ProviderMutationInput { + id: provider.id.clone(), + name: provider.name.clone(), + settings_config: provider.settings_config.clone(), + website_url: provider.website_url.clone(), + category: provider.category.clone(), + created_at: provider.created_at, + sort_index: provider.sort_index, + notes: provider.notes.clone(), + meta: provider.meta.clone(), + icon: provider.icon.clone(), + icon_color: provider.icon_color.clone(), + in_failover_queue: provider.in_failover_queue, + } +} + +use rusqlite::OptionalExtension; diff --git a/src-tauri/src/database/dao/pi_projections.rs b/src-tauri/src/database/dao/pi_projections.rs new file mode 100644 index 000000000..59007f286 --- /dev/null +++ b/src-tauri/src/database/dao/pi_projections.rs @@ -0,0 +1,204 @@ +//! Device-local ownership ledger for exact keys in Pi's shared models.json. + +// The projection writer is introduced in a later contract-ordered commit. +#![allow(dead_code)] + +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use indexmap::IndexMap; +use rusqlite::{params, OptionalExtension}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiProviderProjection { + pub provider_id: String, + pub provider_key: String, + pub created_at: i64, + pub updated_at: i64, +} + +fn decode_projection(row: &rusqlite::Row<'_>) -> rusqlite::Result { + Ok(PiProviderProjection { + provider_id: row.get(0)?, + provider_key: row.get(1)?, + created_at: row.get(2)?, + updated_at: row.get(3)?, + }) +} + +impl Database { + pub(crate) fn get_pi_projection( + &self, + provider_id: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections WHERE provider_id = ?1", + [provider_id], + decode_projection, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn get_pi_projection_for_key( + &self, + provider_key: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections WHERE provider_key = ?1", + [provider_key], + decode_projection, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn get_pi_projection_manifest( + &self, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + let mut stmt = conn + .prepare( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections ORDER BY provider_id", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let rows = stmt + .query_map([], decode_projection) + .map_err(|error| AppError::Database(error.to_string()))?; + let mut manifest = IndexMap::new(); + for row in rows { + let projection = row.map_err(|error| AppError::Database(error.to_string()))?; + manifest.insert(projection.provider_id.clone(), projection); + } + Ok(manifest) + } + + /// Claim an exact key. Existing exact claims are idempotent; either-side + /// collisions fail and are never rewritten. + pub(crate) fn claim_pi_projection_key( + &self, + provider_id: &str, + provider_key: &str, + ) -> Result { + if provider_id.trim().is_empty() || provider_key.trim().is_empty() { + return Err(AppError::Config( + "Pi projection provider id and key must be non-empty".to_string(), + )); + } + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + let by_provider = tx + .query_row( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections WHERE provider_id = ?1", + [provider_id], + decode_projection, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))?; + if let Some(existing) = by_provider { + if existing.provider_key != provider_key { + return Err(AppError::Config(format!( + "Pi provider '{provider_id}' already owns key '{}', not '{provider_key}'", + existing.provider_key + ))); + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string()))?; + return Ok(existing); + } + if let Some(existing_owner) = tx + .query_row( + "SELECT provider_id FROM pi_provider_projections WHERE provider_key = ?1", + [provider_key], + |row| row.get::<_, String>(0), + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))? + { + return Err(AppError::Config(format!( + "Pi key '{provider_key}' is already owned by provider '{existing_owner}'" + ))); + } + let now = chrono::Utc::now().timestamp_millis(); + tx.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES (?1, ?2, ?3, ?3)", + params![provider_id, provider_key, now], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + tx.commit() + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(PiProviderProjection { + provider_id: provider_id.to_string(), + provider_key: provider_key.to_string(), + created_at: now, + updated_at: now, + }) + } + + pub(crate) fn delete_pi_projection_key( + &self, + provider_id: &str, + expected_key: &str, + ) -> Result { + let conn = lock_conn!(self.conn); + let removed = conn + .execute( + "DELETE FROM pi_provider_projections + WHERE provider_id = ?1 AND provider_key = ?2", + params![provider_id, expected_key], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if removed == 0 + && conn + .query_row( + "SELECT 1 FROM pi_provider_projections WHERE provider_id = ?1", + [provider_id], + |_| Ok(()), + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))? + .is_some() + { + return Err(AppError::Config(format!( + "refusing to delete Pi projection '{provider_id}': expected key changed" + ))); + } + Ok(removed == 1) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn projection_claims_are_exact_idempotent_and_collision_safe() -> Result<(), AppError> { + let db = Database::memory()?; + let first = db.claim_pi_projection_key("provider-a", "native-a")?; + let repeated = db.claim_pi_projection_key("provider-a", "native-a")?; + assert_eq!(first, repeated); + assert!(db + .claim_pi_projection_key("provider-a", "native-b") + .is_err()); + assert!(db + .claim_pi_projection_key("provider-b", "native-a") + .is_err()); + assert_eq!(db.get_pi_projection_manifest()?.len(), 1); + assert!(db.delete_pi_projection_key("provider-a", "wrong").is_err()); + assert!(db.get_pi_projection("provider-a")?.is_some()); + assert!(db.delete_pi_projection_key("provider-a", "native-a")?); + assert!(db.get_pi_projection_for_key("native-a")?.is_none()); + Ok(()) + } +} diff --git a/src-tauri/src/database/dao/prompts.rs b/src-tauri/src/database/dao/prompts.rs index 1c274504f..d0f0010fa 100644 --- a/src-tauri/src/database/dao/prompts.rs +++ b/src-tauri/src/database/dao/prompts.rs @@ -6,51 +6,106 @@ use crate::database::{lock_conn, Database}; use crate::error::AppError; use crate::prompt::Prompt; use indexmap::IndexMap; -use rusqlite::params; +use rusqlite::{params, Connection, Transaction}; + +fn query_prompts(conn: &Connection, app_type: &str) -> Result, AppError> { + let mut stmt = conn + .prepare( + "SELECT id, name, content, description, enabled, created_at, updated_at + FROM prompts WHERE app_type = ?1 + ORDER BY created_at ASC, id ASC", + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + let prompt_iter = stmt + .query_map(params![app_type], |row| { + let id: String = row.get(0)?; + let name: String = row.get(1)?; + let content: String = row.get(2)?; + let description: Option = row.get(3)?; + let enabled: bool = row.get(4)?; + let created_at: Option = row.get(5)?; + let updated_at: Option = row.get(6)?; + + Ok(( + id.clone(), + Prompt { + id, + name, + content, + description, + enabled, + created_at, + updated_at, + }, + )) + }) + .map_err(|e| AppError::Database(e.to_string()))?; + + let mut prompts = IndexMap::new(); + for prompt_res in prompt_iter { + let (id, prompt) = prompt_res.map_err(|e| AppError::Database(e.to_string()))?; + prompts.insert(id, prompt); + } + Ok(prompts) +} + +fn validate_prompt_selection(prompts: &IndexMap) -> Result<(), AppError> { + if prompts.values().filter(|prompt| prompt.enabled).count() > 1 { + return Err(AppError::InvalidInput( + "at most one prompt may be enabled for an app".to_string(), + )); + } + Ok(()) +} + +fn replace_prompt_rows( + transaction: &Transaction<'_>, + app_type: &str, + prompts: &IndexMap, +) -> Result<(), AppError> { + transaction + .execute("DELETE FROM prompts WHERE app_type = ?1", [app_type]) + .map_err(|error| AppError::Database(error.to_string()))?; + let mut statement = transaction + .prepare( + "INSERT OR REPLACE INTO prompts ( + id, app_type, name, content, description, enabled, created_at, updated_at + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + for prompt in prompts.values() { + statement + .execute(params![ + prompt.id, + app_type, + prompt.name, + prompt.content, + prompt.description, + prompt.enabled, + prompt.created_at, + prompt.updated_at, + ]) + .map_err(|error| AppError::Database(error.to_string()))?; + } + Ok(()) +} + +fn prompt_libraries_equal( + left: &IndexMap, + right: &IndexMap, +) -> bool { + left.len() == right.len() + && left + .iter() + .all(|(id, prompt)| right.get(id) == Some(prompt)) +} impl Database { /// 获取指定应用类型的所有提示词 pub fn get_prompts(&self, app_type: &str) -> Result, AppError> { let conn = lock_conn!(self.conn); - let mut stmt = conn - .prepare( - "SELECT id, name, content, description, enabled, created_at, updated_at - FROM prompts WHERE app_type = ?1 - ORDER BY created_at ASC, id ASC", - ) - .map_err(|e| AppError::Database(e.to_string()))?; - - let prompt_iter = stmt - .query_map(params![app_type], |row| { - let id: String = row.get(0)?; - let name: String = row.get(1)?; - let content: String = row.get(2)?; - let description: Option = row.get(3)?; - let enabled: bool = row.get(4)?; - let created_at: Option = row.get(5)?; - let updated_at: Option = row.get(6)?; - - Ok(( - id.clone(), - Prompt { - id, - name, - content, - description, - enabled, - created_at, - updated_at, - }, - )) - }) - .map_err(|e| AppError::Database(e.to_string()))?; - - let mut prompts = IndexMap::new(); - for prompt_res in prompt_iter { - let (id, prompt) = prompt_res.map_err(|e| AppError::Database(e.to_string()))?; - prompts.insert(id, prompt); - } - Ok(prompts) + query_prompts(&conn, app_type) } /// 保存提示词 @@ -75,6 +130,75 @@ impl Database { Ok(()) } + /// Persist a complete prompt-library selection atomically. + /// + /// Pi projects the single enabled row into AGENTS.md. A sequence of + /// individual `save_prompt` calls can expose two enabled rows (or none) to + /// concurrent readers, so selection changes use one SQLite transaction. + pub(crate) fn save_prompt_selection( + &self, + app_type: &str, + prompts: &IndexMap, + ) -> Result<(), AppError> { + validate_prompt_selection(prompts)?; + let mut conn = lock_conn!(self.conn); + let transaction = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + replace_prompt_rows(&transaction, app_type, prompts)?; + transaction + .commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + /// Atomically publish a complete prompt library only while its full + /// before-image still matches. This is the database half of Pi's + /// native-file/portable-library compare-and-swap boundary. + pub(crate) fn compare_exchange_prompt_selection( + &self, + app_type: &str, + expected: &IndexMap, + replacement: &IndexMap, + ) -> Result<(), AppError> { + validate_prompt_selection(replacement)?; + self.compare_exchange_prompt_selection_unchecked(app_type, expected, replacement) + } + + /// Restore a captured before-image only if the database still contains the + /// exact attempted projection. The before-image may predate the current + /// single-selection invariant, so compensation must preserve it byte for + /// byte instead of refusing to restore legacy rows. + pub(crate) fn restore_prompt_selection_if_attempted( + &self, + app_type: &str, + attempted: &IndexMap, + before: &IndexMap, + ) -> Result<(), AppError> { + self.compare_exchange_prompt_selection_unchecked(app_type, attempted, before) + } + + fn compare_exchange_prompt_selection_unchecked( + &self, + app_type: &str, + expected: &IndexMap, + replacement: &IndexMap, + ) -> Result<(), AppError> { + let mut conn = lock_conn!(self.conn); + let transaction = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + let observed = query_prompts(&transaction, app_type)?; + if !prompt_libraries_equal(&observed, expected) { + return Err(AppError::Conflict(format!( + "{app_type} prompt library changed since it was read" + ))); + } + replace_prompt_rows(&transaction, app_type, replacement)?; + transaction + .commit() + .map_err(|error| AppError::Database(error.to_string())) + } + /// 删除提示词 pub fn delete_prompt(&self, app_type: &str, id: &str) -> Result<(), AppError> { let conn = lock_conn!(self.conn); diff --git a/src-tauri/src/database/dao/provider_write.rs b/src-tauri/src/database/dao/provider_write.rs index 3ddab2ac2..e42454ec2 100644 --- a/src-tauri/src/database/dao/provider_write.rs +++ b/src-tauri/src/database/dao/provider_write.rs @@ -37,14 +37,14 @@ impl ProviderKey { #[derive(Debug, Clone)] pub struct ProviderRowUpdate { - name: String, - settings_config: Value, - website_url: Option, - category: Option, - notes: Option, - meta: ProviderMeta, - icon: Option, - icon_color: Option, + pub(super) name: String, + pub(super) settings_config: Value, + pub(super) website_url: Option, + pub(super) category: Option, + pub(super) notes: Option, + pub(super) meta: ProviderMeta, + pub(super) icon: Option, + pub(super) icon_color: Option, } impl ProviderRowUpdate { @@ -71,15 +71,15 @@ impl ProviderRowUpdate { #[derive(Debug, Clone)] pub struct ProviderRowCreate { - content: ProviderRowUpdate, - created_at: Option, + pub(super) content: ProviderRowUpdate, + pub(super) created_at: Option, } #[derive(Debug, Clone)] pub struct NewEndpoint { - url: String, - added_at: Option, - last_used: Option, + pub(super) url: String, + pub(super) added_at: Option, + pub(super) last_used: Option, } impl NewEndpoint { @@ -116,11 +116,11 @@ impl TryFrom for NewEndpoint { #[derive(Debug, Clone)] pub struct NewProviderAggregate { - key: ProviderKey, - row: ProviderRowCreate, - sort_index: Option, - in_failover_queue: bool, - initial_endpoints: Vec, + pub(super) key: ProviderKey, + pub(super) row: ProviderRowCreate, + pub(super) sort_index: Option, + pub(super) in_failover_queue: bool, + pub(super) initial_endpoints: Vec, } impl NewProviderAggregate { @@ -219,7 +219,7 @@ fn encode_row(row: &ProviderRowUpdate) -> Result<(String, String), AppError> { Ok((settings_config, meta)) } -fn insert_row( +pub(super) fn insert_row( tx: &Transaction<'_>, key: &ProviderKey, row: &ProviderRowUpdate, @@ -272,7 +272,7 @@ fn insert_row( Ok(()) } -fn insert_endpoint( +pub(super) fn insert_endpoint( tx: &Transaction<'_>, key: &ProviderKey, endpoint: &NewEndpoint, @@ -357,7 +357,7 @@ pub(super) fn restore_provider_aggregate_on_tx( Ok(()) } -fn update_row( +pub(super) fn update_row( tx: &Transaction<'_>, key: &ProviderKey, row: &ProviderRowUpdate, @@ -618,23 +618,38 @@ impl Database { pub(crate) fn update_provider_sort_index( &self, - key: &ProviderKey, - sort_index: usize, + updates: &[(ProviderKey, usize)], ) -> Result<(), AppError> { - let conn = lock_conn!(self.conn); - if conn - .execute( - "UPDATE providers SET sort_index = ?1 WHERE id = ?2 AND app_type = ?3", - params![sort_index, key.id, key.app_type], - ) - .map_err(|error| AppError::Database(error.to_string()))? - != 1 - { - return Err(AppError::NotFound(format!( - "provider '{}/{}'", - key.app_type, key.id - ))); + let mut seen = std::collections::HashSet::with_capacity(updates.len()); + for (key, _) in updates { + if !seen.insert((key.app_type().to_string(), key.id().to_string())) { + return Err(AppError::InvalidInput(format!( + "duplicate provider sort update for '{}/{}'", + key.app_type(), + key.id() + ))); + } } - Ok(()) + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + for (key, sort_index) in updates { + if tx + .execute( + "UPDATE providers SET sort_index = ?1 WHERE id = ?2 AND app_type = ?3", + params![sort_index, key.id, key.app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))? + != 1 + { + return Err(AppError::NotFound(format!( + "provider '{}/{}'", + key.app_type, key.id + ))); + } + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) } } diff --git a/src-tauri/src/database/dao/providers.rs b/src-tauri/src/database/dao/providers.rs index 1d094348b..b44726e3a 100644 --- a/src-tauri/src/database/dao/providers.rs +++ b/src-tauri/src/database/dao/providers.rs @@ -6,6 +6,24 @@ use indexmap::IndexMap; use rusqlite::{params, OptionalExtension, Row}; use std::collections::{HashMap, HashSet}; +pub(super) fn delete_provider_on_tx( + tx: &rusqlite::Transaction<'_>, + app_type: &str, + id: &str, +) -> Result<(), AppError> { + if tx + .execute( + "DELETE FROM providers WHERE id = ?1 AND app_type = ?2", + params![id, app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))? + != 1 + { + return Err(AppError::NotFound(format!("provider '{app_type}/{id}'"))); + } + Ok(()) +} + pub(super) struct StoredProviderRow { id: String, name: String, @@ -295,6 +313,16 @@ impl Database { Ok(()) } + pub(crate) fn clear_current_provider_for_app(&self, app_type: &str) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + conn.execute( + "UPDATE providers SET is_current = 0 WHERE app_type = ?1", + params![app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(()) + } + pub fn set_omo_provider_current( &self, app_type: &str, diff --git a/src-tauri/src/database/dao/skill_deployments.rs b/src-tauri/src/database/dao/skill_deployments.rs new file mode 100644 index 000000000..7f1133ee6 --- /dev/null +++ b/src-tauri/src/database/dao/skill_deployments.rs @@ -0,0 +1,290 @@ +//! Device-local evidence for Pi Skill deployments. + +// Pi skill reconciliation consumes this ledger in a later contract-ordered commit. +#![allow(dead_code)] + +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use rusqlite::{params, OptionalExtension}; +use serde::{Deserialize, Serialize}; +use std::str::FromStr; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum SkillDeploymentMethod { + Symlink, + Copy, +} + +impl SkillDeploymentMethod { + pub(crate) const fn as_str(self) -> &'static str { + match self { + Self::Symlink => "symlink", + Self::Copy => "copy", + } + } +} + +impl FromStr for SkillDeploymentMethod { + type Err = AppError; + + fn from_str(value: &str) -> Result { + match value { + "symlink" => Ok(Self::Symlink), + "copy" => Ok(Self::Copy), + _ => Err(AppError::Database(format!( + "unknown Pi Skill deployment method '{value}'" + ))), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct SkillDeployment { + pub skill_id: String, + pub destination: String, + pub destination_key: String, + pub method: SkillDeploymentMethod, + pub source_identity: String, + pub deployed_digest: Option, + pub created_at: i64, + pub updated_at: i64, +} + +fn decode_deployment(row: &rusqlite::Row<'_>) -> rusqlite::Result { + let method: String = row.get(3)?; + let method = method.parse().map_err(|error: AppError| { + rusqlite::Error::FromSqlConversionFailure(3, rusqlite::types::Type::Text, Box::new(error)) + })?; + Ok(SkillDeployment { + skill_id: row.get(0)?, + destination: row.get(1)?, + destination_key: row.get(2)?, + method, + source_identity: row.get(4)?, + deployed_digest: row.get(5)?, + created_at: row.get(6)?, + updated_at: row.get(7)?, + }) +} + +impl Database { + pub(crate) fn set_pi_skill_desired( + &self, + skill_id: &str, + desired_enabled: bool, + ) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + let changed = conn + .execute( + "UPDATE skills SET enabled_pi = ?1 WHERE id = ?2", + params![desired_enabled, skill_id], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if changed != 1 { + return Err(AppError::Conflict(format!( + "Pi Skill '{skill_id}' disappeared before desired state was saved" + ))); + } + Ok(()) + } + + pub(crate) fn get_pi_skill_deployment( + &self, + skill_id: &str, + destination_key: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT skill_id, destination, destination_key, method, + source_identity, deployed_digest, created_at, updated_at + FROM skill_deployments + WHERE app_type = 'pi' AND skill_id = ?1 AND destination_key = ?2", + params![skill_id, destination_key], + decode_deployment, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn get_pi_skill_deployments( + &self, + skill_id: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + let mut stmt = conn + .prepare( + "SELECT skill_id, destination, destination_key, method, + source_identity, deployed_digest, created_at, updated_at + FROM skill_deployments + WHERE app_type = 'pi' AND skill_id = ?1 + ORDER BY created_at, destination_key", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let rows = stmt + .query_map([skill_id], decode_deployment) + .map_err(|error| AppError::Database(error.to_string()))?; + rows.map(|row| row.map_err(|error| AppError::Database(error.to_string()))) + .collect() + } + + pub(crate) fn save_pi_skill_deployment( + &self, + deployment: &SkillDeployment, + ) -> Result<(), AppError> { + self.save_pi_skill_deployment_with_desired(deployment, None) + } + + /// Commit ledger evidence and, for a user toggle, the desired Pi bit in + /// the same SQLite transaction. Filesystem publication happens before + /// this point; a failed transaction is therefore safe to compensate by + /// restoring the staged destination without exposing split DB authority. + pub(crate) fn save_pi_skill_deployment_with_desired( + &self, + deployment: &SkillDeployment, + desired_enabled: Option, + ) -> Result<(), AppError> { + if deployment.skill_id.trim().is_empty() + || deployment.destination.trim().is_empty() + || deployment.destination_key.trim().is_empty() + || deployment.source_identity.trim().is_empty() + { + return Err(AppError::Config( + "Pi Skill deployment identity fields must be non-empty".to_string(), + )); + } + let mut conn = lock_conn!(self.conn); + let transaction = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + transaction + .execute( + "INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, deployed_digest, created_at, updated_at + ) VALUES ('pi', ?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) + ON CONFLICT(app_type, skill_id, destination_key) DO UPDATE SET + destination = excluded.destination, + method = excluded.method, + source_identity = excluded.source_identity, + deployed_digest = excluded.deployed_digest, + updated_at = excluded.updated_at", + params![ + deployment.skill_id, + deployment.destination, + deployment.destination_key, + deployment.method.as_str(), + deployment.source_identity, + deployment.deployed_digest, + deployment.created_at, + deployment.updated_at, + ], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if let Some(desired_enabled) = desired_enabled { + let changed = transaction + .execute( + "UPDATE skills SET enabled_pi = ?1 WHERE id = ?2", + params![desired_enabled, deployment.skill_id], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if changed != 1 { + return Err(AppError::Conflict(format!( + "Pi Skill '{}' disappeared before deployment commit", + deployment.skill_id + ))); + } + } + transaction + .commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn delete_pi_skill_deployment( + &self, + skill_id: &str, + destination_key: &str, + ) -> Result { + self.delete_pi_skill_deployment_with_desired(skill_id, destination_key, None) + } + + pub(crate) fn delete_pi_skill_deployment_with_desired( + &self, + skill_id: &str, + destination_key: &str, + desired_enabled: Option, + ) -> Result { + let mut conn = lock_conn!(self.conn); + let transaction = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + if let Some(desired_enabled) = desired_enabled { + let changed = transaction + .execute( + "UPDATE skills SET enabled_pi = ?1 WHERE id = ?2", + params![desired_enabled, skill_id], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if changed != 1 { + return Err(AppError::Conflict(format!( + "Pi Skill '{skill_id}' disappeared before deployment cleanup" + ))); + } + } + let removed = transaction + .execute( + "DELETE FROM skill_deployments + WHERE app_type = 'pi' AND skill_id = ?1 AND destination_key = ?2", + params![skill_id, destination_key], + ) + .map_err(|error| AppError::Database(error.to_string()))? + == 1; + transaction + .commit() + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(removed) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn deployment(skill_id: &str, destination_key: &str) -> SkillDeployment { + SkillDeployment { + skill_id: skill_id.into(), + destination: format!("/tmp/{destination_key}"), + destination_key: destination_key.into(), + method: SkillDeploymentMethod::Copy, + source_identity: format!("source:{skill_id}"), + deployed_digest: Some("sha256:initial".into()), + created_at: 10, + updated_at: 10, + } + } + + #[test] + fn skill_ledger_preserves_created_at_and_rejects_destination_collision() -> Result<(), AppError> + { + let db = Database::memory()?; + db.save_pi_skill_deployment(&deployment("one", "destination"))?; + let mut updated = deployment("one", "destination"); + updated.updated_at = 20; + updated.deployed_digest = Some("sha256:updated".into()); + db.save_pi_skill_deployment(&updated)?; + let saved = db + .get_pi_skill_deployment("one", "destination")? + .expect("deployment"); + assert_eq!(saved.created_at, 10); + assert_eq!(saved.updated_at, 20); + assert_eq!(saved.deployed_digest.as_deref(), Some("sha256:updated")); + + assert!(db + .save_pi_skill_deployment(&deployment("two", "destination")) + .is_err()); + assert_eq!(db.get_pi_skill_deployments("one")?.len(), 1); + assert!(db.delete_pi_skill_deployment("one", "destination")?); + Ok(()) + } +} diff --git a/src-tauri/src/database/dao/skills.rs b/src-tauri/src/database/dao/skills.rs index 488fde29d..1f075336f 100644 --- a/src-tauri/src/database/dao/skills.rs +++ b/src-tauri/src/database/dao/skills.rs @@ -23,7 +23,8 @@ impl Database { .prepare( "SELECT id, name, description, directory, repo_owner, repo_name, repo_branch, readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, - enabled_opencode, enabled_hermes, installed_at, content_hash, updated_at + enabled_opencode, enabled_hermes, enabled_pi, + installed_at, content_hash, updated_at FROM skills ORDER BY name ASC", ) .map_err(|e| AppError::Database(e.to_string()))?; @@ -46,10 +47,11 @@ impl Database { grokbuild: row.get(11)?, opencode: row.get(12)?, hermes: row.get(13)?, + pi: row.get(14)?, }, - installed_at: row.get(14)?, - content_hash: row.get(15)?, - updated_at: row.get::<_, i64>(16).unwrap_or(0), + installed_at: row.get(15)?, + content_hash: row.get(16)?, + updated_at: row.get::<_, i64>(17).unwrap_or(0), }) }) .map_err(|e| AppError::Database(e.to_string()))?; @@ -69,7 +71,8 @@ impl Database { .prepare( "SELECT id, name, description, directory, repo_owner, repo_name, repo_branch, readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, - enabled_opencode, enabled_hermes, installed_at, content_hash, updated_at + enabled_opencode, enabled_hermes, enabled_pi, + installed_at, content_hash, updated_at FROM skills WHERE id = ?1", ) .map_err(|e| AppError::Database(e.to_string()))?; @@ -91,10 +94,11 @@ impl Database { grokbuild: row.get(11)?, opencode: row.get(12)?, hermes: row.get(13)?, + pi: row.get(14)?, }, - installed_at: row.get(14)?, - content_hash: row.get(15)?, - updated_at: row.get::<_, i64>(16).unwrap_or(0), + installed_at: row.get(15)?, + content_hash: row.get(16)?, + updated_at: row.get::<_, i64>(17).unwrap_or(0), }) }); @@ -109,11 +113,28 @@ impl Database { pub fn save_skill(&self, skill: &InstalledSkill) -> Result<(), AppError> { let conn = lock_conn!(self.conn); conn.execute( - "INSERT OR REPLACE INTO skills + "INSERT INTO skills (id, name, description, directory, repo_owner, repo_name, repo_branch, readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, enabled_opencode, enabled_hermes, - installed_at, content_hash, updated_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17)", + enabled_pi, installed_at, content_hash, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + description = excluded.description, + directory = excluded.directory, + repo_owner = excluded.repo_owner, + repo_name = excluded.repo_name, + repo_branch = excluded.repo_branch, + readme_url = excluded.readme_url, + enabled_claude = excluded.enabled_claude, + enabled_codex = excluded.enabled_codex, + enabled_gemini = excluded.enabled_gemini, + enabled_grokbuild = excluded.enabled_grokbuild, + enabled_opencode = excluded.enabled_opencode, + enabled_hermes = excluded.enabled_hermes, + installed_at = excluded.installed_at, + content_hash = excluded.content_hash, + updated_at = excluded.updated_at", params![ skill.id, skill.name, @@ -129,6 +150,7 @@ impl Database { skill.apps.grokbuild, skill.apps.opencode, skill.apps.hermes, + skill.apps.pi, skill.installed_at, skill.content_hash, skill.updated_at, @@ -160,8 +182,8 @@ impl Database { let conn = lock_conn!(self.conn); let affected = conn .execute( - "UPDATE skills SET enabled_claude = ?1, enabled_codex = ?2, enabled_gemini = ?3, enabled_grokbuild = ?4, enabled_opencode = ?5, enabled_hermes = ?6 WHERE id = ?7", - params![apps.claude, apps.codex, apps.gemini, apps.grokbuild, apps.opencode, apps.hermes, id], + "UPDATE skills SET enabled_claude = ?1, enabled_codex = ?2, enabled_gemini = ?3, enabled_grokbuild = ?4, enabled_opencode = ?5, enabled_hermes = ?6, enabled_pi = ?7 WHERE id = ?8", + params![apps.claude, apps.codex, apps.gemini, apps.grokbuild, apps.opencode, apps.hermes, apps.pi, id], ) .map_err(|e| AppError::Database(e.to_string()))?; Ok(affected > 0) @@ -262,3 +284,52 @@ impl Database { Ok(count) } } + +#[cfg(test)] +mod tests { + use super::*; + + fn installed_skill() -> InstalledSkill { + InstalledSkill { + id: "owner/repo:skill".into(), + name: "Skill".into(), + description: Some("before".into()), + directory: "skill".into(), + repo_owner: Some("owner".into()), + repo_name: Some("repo".into()), + repo_branch: Some("main".into()), + readme_url: None, + apps: SkillApps::default(), + installed_at: 10, + content_hash: Some("sha256:before".into()), + updated_at: 11, + } + } + + #[test] + fn legacy_skill_save_preserves_pi_desired_state() -> Result<(), AppError> { + let db = Database::memory()?; + let mut skill = installed_skill(); + db.save_skill(&skill)?; + { + let conn = lock_conn!(db.conn); + conn.execute( + "UPDATE skills SET enabled_pi = 1 WHERE id = ?1", + [&skill.id], + )?; + } + + skill.name = "Updated".into(); + skill.content_hash = Some("sha256:after".into()); + db.save_skill(&skill)?; + + let conn = lock_conn!(db.conn); + let saved: (String, String, bool) = conn.query_row( + "SELECT name, content_hash, enabled_pi FROM skills WHERE id = ?1", + [&skill.id], + |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)), + )?; + assert_eq!(saved, ("Updated".into(), "sha256:after".into(), true)); + Ok(()) + } +} diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index 23a05284a..69514f192 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -32,6 +32,7 @@ mod schema; mod tests; // DAO 类型导出供外部使用 +pub(crate) use dao::pi_projections::PiProviderProjection; pub use dao::provider_write::{ NewEndpoint, NewProviderAggregate, ProviderKey, ProviderRowUpdate, RenameProvider, }; @@ -43,6 +44,7 @@ pub(crate) use dao::proxy::{ validate_cost_multiplier, validate_pricing_source, PRICING_SOURCE_REQUEST, PRICING_SOURCE_RESPONSE, }; +pub(crate) use dao::skill_deployments::{SkillDeployment, SkillDeploymentMethod}; pub use dao::FailoverQueueItem; pub use dao::Profile; diff --git a/src-tauri/src/deeplink/mod.rs b/src-tauri/src/deeplink/mod.rs index 6bee98878..a89c39208 100644 --- a/src-tauri/src/deeplink/mod.rs +++ b/src-tauri/src/deeplink/mod.rs @@ -66,6 +66,10 @@ pub struct DeepLinkImportRequest { /// Optional model name #[serde(skip_serializing_if = "Option::is_none")] pub model: Option, + /// Native API identifier. Pi provider links require this explicitly; + /// cc-switch never infers a protocol from a URL or model name. + #[serde(skip_serializing_if = "Option::is_none")] + pub api: Option, /// Optional notes/description #[serde(skip_serializing_if = "Option::is_none")] pub notes: Option, diff --git a/src-tauri/src/deeplink/parser.rs b/src-tauri/src/deeplink/parser.rs index ee854ff4f..d6c8ea02b 100644 --- a/src-tauri/src/deeplink/parser.rs +++ b/src-tauri/src/deeplink/parser.rs @@ -81,10 +81,10 @@ fn parse_provider_deeplink( // Validate app type if !matches!( app.as_str(), - "claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes" + "claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes" | "pi" ) { return Err(AppError::InvalidInput(format!( - "Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', or 'hermes', got '{app}'" + "Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', 'hermes', or 'pi', got '{app}'" ))); } @@ -116,6 +116,7 @@ fn parse_provider_deeplink( // Extract optional fields let model = params.get("model").cloned(); + let api = params.get("api").cloned(); let notes = params.get("notes").cloned(); let haiku_model = params.get("haikuModel").cloned(); let sonnet_model = params.get("sonnetModel").cloned(); @@ -127,6 +128,24 @@ fn parse_provider_deeplink( let config = params.get("config").cloned(); let config_format = params.get("configFormat").cloned(); let config_url = params.get("configUrl").cloned(); + if app == "pi" { + if model.as_deref().is_none_or(|value| value.trim().is_empty()) { + return Err(AppError::InvalidInput( + "Pi provider deep links require a non-empty 'model' parameter".to_string(), + )); + } + if api.as_deref().is_none_or(|value| value.trim().is_empty()) { + return Err(AppError::InvalidInput( + "Pi provider deep links require an explicit non-empty 'api' parameter".to_string(), + )); + } + if config.is_some() || config_url.is_some() { + return Err(AppError::InvalidInput( + "Pi provider deep links use explicit endpoint/api/model fields; embedded or remote config payloads are not supported" + .to_string(), + )); + } + } let enabled = params.get("enabled").and_then(|v| v.parse::().ok()); // Extract usage script fields (v3.9+) @@ -153,6 +172,7 @@ fn parse_provider_deeplink( api_key, icon, model, + api, notes, haiku_model, sonnet_model, @@ -190,10 +210,10 @@ fn parse_prompt_deeplink( // Validate app type if !matches!( app.as_str(), - "claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes" + "claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes" | "pi" ) { return Err(AppError::InvalidInput(format!( - "Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', or 'hermes', got '{app}'" + "Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', 'hermes', or 'pi', got '{app}'" ))); } @@ -225,6 +245,7 @@ fn parse_prompt_deeplink( endpoint: None, api_key: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -298,6 +319,7 @@ fn parse_mcp_deeplink( endpoint: None, api_key: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -353,6 +375,7 @@ fn parse_skill_deeplink( endpoint: None, api_key: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, diff --git a/src-tauri/src/deeplink/provider.rs b/src-tauri/src/deeplink/provider.rs index 4ab1fd13e..2b13f2646 100644 --- a/src-tauri/src/deeplink/provider.rs +++ b/src-tauri/src/deeplink/provider.rs @@ -160,6 +160,7 @@ pub(crate) fn build_provider_from_request( AppType::OpenCode => build_opencode_settings(request), AppType::OpenClaw => build_additive_app_settings(request), AppType::Hermes => build_hermes_settings(request), + AppType::Pi => build_pi_settings(request)?, }; // Build usage script configuration if provided @@ -591,6 +592,45 @@ fn build_hermes_settings(request: &DeepLinkImportRequest) -> serde_json::Value { json!(config) } +/// Pi deep links intentionally carry one explicit model, endpoint and native +/// API identifier. Map only that closed subset; richer Pi catalogs use native +/// inspection/import or the Pi editor. No URL/model heuristic may invent the +/// protocol or model identity. +fn build_pi_settings(request: &DeepLinkImportRequest) -> Result { + let endpoint = get_primary_endpoint(request); + let model = request + .model + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .ok_or_else(|| { + AppError::InvalidInput( + "Pi provider deep links require a non-empty model identifier".to_string(), + ) + })?; + let api = request + .api + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + .ok_or_else(|| { + AppError::InvalidInput( + "Pi provider deep links require an explicit API identifier".to_string(), + ) + })?; + + Ok(json!({ + "name": request.name, + "baseUrl": endpoint, + "apiKey": request.api_key, + "api": api, + "models": [{ + "id": model, + "name": model + }] + })) +} + // ============================================================================= // Config Merge Logic // ============================================================================= diff --git a/src-tauri/src/deeplink/tests.rs b/src-tauri/src/deeplink/tests.rs index 49fcea1d1..62eced364 100644 --- a/src-tauri/src/deeplink/tests.rs +++ b/src-tauri/src/deeplink/tests.rs @@ -89,6 +89,61 @@ fn test_parse_deeplink_with_notes() { assert_eq!(request.notes, Some("Test notes".to_string())); } +#[test] +fn pi_provider_deeplink_requires_and_preserves_explicit_native_identity() { + use super::provider::build_provider_from_request; + + let request = parse_deeplink_url( + "ccswitch://v1/import?resource=provider&app=pi&name=Pi%20Provider&homepage=https%3A%2F%2Fexample.com&endpoint=https%3A%2F%2Fapi.example.com%2Fv1&apiKey=sk-test&model=opaque-model&api=future-native-api", + ) + .expect("parse explicit Pi provider link"); + assert_eq!(request.app.as_deref(), Some("pi")); + assert_eq!(request.api.as_deref(), Some("future-native-api")); + assert_eq!(request.model.as_deref(), Some("opaque-model")); + + let provider = build_provider_from_request(&AppType::Pi, &request).expect("build Pi provider"); + assert_eq!( + provider.settings_config, + serde_json::json!({ + "name": "Pi Provider", + "baseUrl": "https://api.example.com/v1", + "apiKey": "sk-test", + "api": "future-native-api", + "models": [{ + "id": "opaque-model", + "name": "opaque-model" + }] + }), + "deeplinks must not invent a model, protocol, capability, pricing, or limit field" + ); +} + +#[test] +fn pi_provider_deeplink_rejects_implicit_model_or_protocol() { + let missing_api = "ccswitch://v1/import?resource=provider&app=pi&name=Pi&endpoint=https%3A%2F%2Fapi.example.com&apiKey=sk-test&model=opaque-model"; + assert!(parse_deeplink_url(missing_api) + .expect_err("Pi api must be explicit") + .to_string() + .contains("'api'")); + + let missing_model = "ccswitch://v1/import?resource=provider&app=pi&name=Pi&endpoint=https%3A%2F%2Fapi.example.com&apiKey=sk-test&api=openai-responses"; + assert!(parse_deeplink_url(missing_model) + .expect_err("Pi model must be explicit") + .to_string() + .contains("'model'")); +} + +#[test] +fn pi_prompt_deeplink_is_accepted_by_the_shared_prompt_path() { + let content = BASE64_STANDARD.encode("Pinned Pi AGENTS content"); + let url = format!( + "ccswitch://v1/import?resource=prompt&app=pi&name=Pi%20AGENTS&content={content}&enabled=false" + ); + let request = parse_deeplink_url(&url).expect("parse Pi prompt deeplink"); + assert_eq!(request.app.as_deref(), Some("pi")); + assert_eq!(request.content.as_deref(), Some(content.as_str())); +} + #[test] fn test_parse_grokbuild_provider() { use super::provider::build_provider_from_request; @@ -210,6 +265,7 @@ fn test_build_gemini_provider_with_model() { api_key: Some("test-api-key".to_string()), icon: None, model: Some("gemini-2.0-flash".to_string()), + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -263,6 +319,7 @@ fn test_build_gemini_provider_without_model() { api_key: Some("test-api-key".to_string()), icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -309,6 +366,7 @@ fn test_deeplink_usage_script_does_not_copy_provider_credentials() { api_key: Some("sk-main".to_string()), icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -356,6 +414,7 @@ fn usage_script_request(code: &str, usage_enabled: Option) -> DeepLinkImpo api_key: Some("sk-main".to_string()), icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -439,6 +498,7 @@ fn test_deeplink_usage_script_omits_explicit_credentials_that_match_provider() { api_key: Some("sk-main".to_string()), icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -487,6 +547,7 @@ fn test_deeplink_usage_script_preserves_distinct_usage_credentials() { api_key: Some("sk-main".to_string()), icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -540,6 +601,7 @@ fn test_parse_and_merge_config_claude() { api_key: None, icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -663,6 +725,7 @@ fn test_parse_and_merge_config_url_override() { api_key: Some("sk-new".to_string()), // URL param should override icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, @@ -726,6 +789,7 @@ fn test_build_claude_provider_preserves_custom_env_fields() { icon: None, // URL param: must win over the same key in config (haiku-from-config) model: Some("main-model".to_string()), + api: None, notes: None, haiku_model: Some("haiku-from-url".to_string()), sonnet_model: None, @@ -781,6 +845,7 @@ fn test_build_claude_provider_without_config_unchanged() { api_key: Some("sk".to_string()), icon: None, model: None, + api: None, notes: None, haiku_model: None, sonnet_model: None, diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 2e0844105..8d9d33ac6 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -25,6 +25,7 @@ mod model_capabilities; mod openclaw_config; mod opencode_config; mod panic_hook; +mod pi_config; mod prompt; mod prompt_files; mod provider; @@ -38,6 +39,9 @@ mod tray; mod usage_events; mod usage_script; +#[cfg(test)] +mod architecture_tests; + pub use app_config::{AppType, InstalledSkill, McpApps, McpServer, MultiAppConfig, SkillApps}; pub use codex_config::{ get_codex_auth_path, get_codex_config_path, read_codex_live_settings, write_codex_live_atomic, @@ -949,6 +953,7 @@ pub fn run() { crate::app_config::AppType::OpenCode, crate::app_config::AppType::OpenClaw, crate::app_config::AppType::Hermes, + crate::app_config::AppType::Pi, ] { match crate::services::prompt::PromptService::import_from_file_on_first_launch( &app_state, @@ -1327,6 +1332,12 @@ pub fn run() { commands::remove_provider_from_live_config, commands::switch_provider, commands::import_default_config, + commands::get_pi_native_catalog, + commands::import_pi_native_provider, + commands::set_pi_default_model, + commands::get_pi_native_defaults, + commands::get_pi_session_discovery, + commands::reset_pi_gateway_credential, commands::get_claude_desktop_status, commands::get_claude_desktop_default_routes, commands::import_claude_desktop_providers_from_claude, @@ -1411,6 +1422,14 @@ pub fn run() { commands::enable_prompt, commands::import_prompt_from_file, commands::get_current_prompt_file_content, + commands::get_pi_prompt_library_status, + commands::reconcile_pi_prompt_library, + commands::get_pi_prompt_file, + commands::replace_pi_prompt_file, + commands::delete_pi_prompt_file, + commands::list_pi_prompt_templates, + commands::upsert_pi_prompt_template, + commands::delete_pi_prompt_template, // Profile management (项目配置方案) commands::list_profiles, commands::create_profile, @@ -1465,6 +1484,7 @@ pub fn run() { commands::restore_env_backup, // Skill management (v3.10.0+ unified) commands::get_installed_skills, + commands::get_pi_skill_statuses, commands::get_skill_backups, commands::delete_skill_backup, commands::install_skill_unified, @@ -1830,7 +1850,11 @@ pub async fn cleanup_before_exit(app_handle: &tauri::AppHandle) { } }; let live_taken_over = proxy_service.detect_takeover_in_live_configs(); - let needs_restore = has_backups || live_taken_over; + let needs_restore = cleanup_before_exit_needed( + has_backups, + live_taken_over, + crate::settings::pi_takeover_enabled(), + ); if needs_restore { log::info!("检测到接管残留,开始恢复 Live 配置(保留代理状态)..."); @@ -1854,6 +1878,14 @@ pub async fn cleanup_before_exit(app_handle: &tauri::AppHandle) { } } +fn cleanup_before_exit_needed( + has_live_backups: bool, + legacy_live_taken_over: bool, + pi_takeover_enabled: bool, +) -> bool { + has_live_backups || legacy_live_taken_over || pi_takeover_enabled +} + /// 主动从系统托盘移除托盘图标。 /// /// `std::process::exit` 会绕过 Tauri 运行时,触发不了 `TrayIcon::drop()`, @@ -1884,7 +1916,10 @@ pub(crate) fn remove_tray_icon_before_exit(app_handle: &tauri::AppHandle) { /// 则自动启动代理服务并接管对应应用的 Live 配置。 const PROXY_STARTUP_APP_TYPES: [&str; 4] = ["claude", "codex", "gemini", "grokbuild"]; -async fn enabled_proxy_apps_on_startup(db: &database::Database) -> Vec<&'static str> { +async fn enabled_proxy_apps_on_startup( + db: &database::Database, + pi_takeover_enabled: bool, +) -> Vec<&'static str> { let mut apps = Vec::new(); for app_type in PROXY_STARTUP_APP_TYPES { if db @@ -1895,12 +1930,16 @@ async fn enabled_proxy_apps_on_startup(db: &database::Database) -> Vec<&'static apps.push(app_type); } } + if pi_takeover_enabled { + apps.push("pi"); + } apps } async fn restore_proxy_state_on_startup(state: &store::AppState) { // 收集需要恢复接管的应用列表(从 proxy_config.enabled 读取) - let apps_to_restore = enabled_proxy_apps_on_startup(&state.db).await; + let apps_to_restore = + enabled_proxy_apps_on_startup(&state.db, crate::settings::pi_takeover_enabled()).await; if apps_to_restore.is_empty() { log::debug!("启动时无需恢复代理状态"); @@ -1921,7 +1960,15 @@ async fn restore_proxy_state_on_startup(state: &store::AppState) { } Err(e) => { log::error!("✗ 恢复 {app_type} 的代理接管状态失败: {e}"); - // 失败时清除该应用的状态,避免下次启动再次尝试 + // Pi desired state is device-local user intent. Keep it + // pending/degraded so a transient bind or projection failure + // is retried on the next startup. + if app_type == "pi" { + continue; + } + // Legacy live-config apps retain their historical cleanup + // behavior because their enabled bit also describes a live + // file takeover, not an independent desired/operational pair. if let Err(clear_err) = state .proxy_service .set_takeover_for_app(app_type, false) @@ -2223,9 +2270,9 @@ pub fn restart_process(app_handle: &tauri::AppHandle) -> ! { #[cfg(test)] mod tests { use super::{ - classify_exit_request, enabled_proxy_apps_on_startup, redact_url_for_log, - redact_url_for_log_with_secrets, redact_url_origin_for_log, runtime_log_level_allows, - ExitRequestAction, + classify_exit_request, cleanup_before_exit_needed, enabled_proxy_apps_on_startup, + redact_url_for_log, redact_url_for_log_with_secrets, redact_url_origin_for_log, + runtime_log_level_allows, ExitRequestAction, }; use crate::database::Database; @@ -2347,8 +2394,21 @@ mod tests { .await .expect("enable Grok Build proxy config"); - let apps = enabled_proxy_apps_on_startup(&db).await; + let apps = enabled_proxy_apps_on_startup(&db, false).await; assert_eq!(apps, vec!["grokbuild"]); } + + #[tokio::test] + async fn startup_restore_republishes_persisted_pi_takeover() { + let db = Database::memory().expect("initialize database"); + let apps = enabled_proxy_apps_on_startup(&db, true).await; + assert_eq!(apps, vec!["pi"]); + } + + #[test] + fn process_exit_cleanup_includes_pi_takeover_without_legacy_live_backups() { + assert!(cleanup_before_exit_needed(false, false, true)); + assert!(!cleanup_before_exit_needed(false, false, false)); + } } diff --git a/src-tauri/src/pi_config/composer.rs b/src-tauri/src/pi_config/composer.rs new file mode 100644 index 000000000..a38eece7c --- /dev/null +++ b/src-tauri/src/pi_config/composer.rs @@ -0,0 +1,979 @@ +//! Credential-blind Pi native model composition. +//! +//! The only Pi-layer input is [`PiRawValidProvider`]. This module does not +//! import managed DTOs or gateway families, and it never resolves credentials, +//! environment variables, commands, files, or network resources. + +#![allow(dead_code)] + +use super::{ + merge_pi_compat, + raw_schema::{PiRawApiId, PiRawValidProvider}, +}; +use serde_json::{json, Map, Value}; +use std::collections::{BTreeMap, HashSet}; + +const PROVIDER_FIELDS: &[&str] = &[ + "name", + "baseUrl", + "apiKey", + "api", + "oauth", + "headers", + "compat", + "authHeader", + "models", + "modelOverrides", +]; +const MODEL_FIELDS: &[&str] = &[ + "id", + "name", + "baseUrl", + "api", + "reasoning", + "thinkingLevelMap", + "input", + "cost", + "contextWindow", + "maxTokens", + "headers", + "compat", +]; +const OVERRIDE_FIELDS: &[&str] = &[ + "name", + "reasoning", + "thinkingLevelMap", + "input", + "cost", + "contextWindow", + "maxTokens", + "headers", + "compat", +]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum PiComposerStatus { + Composed, + Failed, + Unknown, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum PiComposerReasonCode { + CatalogRequired, + MissingExplicitModels, + MissingEffectiveApi, + MissingEffectiveEndpoint, + NonPositiveModelLimit, + UnrepresentableCompat, + CompositionFailed, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PiComposerReason { + pub code: PiComposerReasonCode, + pub json_pointer: String, +} + +/// One configured header together with the source pointer that Pi resolves. +/// +/// `headers` remains the pinned composer's flattened observable result, while +/// these entries retain the provider-vs-model boundary needed to reproduce +/// the later `ModelRuntime` merge on the wire. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PiComposedHeader { + pub name: String, + pub value: String, + pub json_pointer: String, +} + +/// The lossless native result of pinned Pi composition. +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct PiComposedNativeModel { + pub id: String, + pub name: String, + pub api: PiRawApiId, + pub provider: String, + pub base_url: String, + pub reasoning: bool, + pub thinking_level_map: Option, + pub input: Value, + pub cost: Value, + pub context_window: Value, + pub max_tokens: Value, + pub headers: BTreeMap, + pub provider_headers: Vec, + pub model_headers: Vec, + pub compat: Option, + pub api_key: Option, + pub oauth: Option, + pub auth_header: bool, + pub provider_extra: BTreeMap, + pub model_extra: BTreeMap, + pub override_extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct PiNativeComposition { + pub status: PiComposerStatus, + pub provider_id: Option, + pub provider_name: Option, + pub provider_base_url: Option, + pub models: Vec, + pub ignored_override_keys: Vec, + pub reasons: Vec, +} + +impl PiNativeComposition { + pub(super) fn unavailable_without_valid_raw() -> Self { + Self { + status: PiComposerStatus::Unknown, + provider_id: None, + provider_name: None, + provider_base_url: None, + models: Vec::new(), + ignored_override_keys: Vec::new(), + reasons: Vec::new(), + } + } + + pub(super) fn catalog_required(pointer: &str) -> Self { + Self { + status: PiComposerStatus::Unknown, + provider_id: None, + provider_name: None, + provider_base_url: None, + models: Vec::new(), + ignored_override_keys: Vec::new(), + reasons: vec![PiComposerReason { + code: PiComposerReasonCode::CatalogRequired, + json_pointer: pointer.to_string(), + }], + } + } + + fn failed(code: PiComposerReasonCode, pointer: impl Into) -> Self { + Self { + status: PiComposerStatus::Failed, + provider_id: None, + provider_name: None, + provider_base_url: None, + models: Vec::new(), + ignored_override_keys: Vec::new(), + reasons: vec![PiComposerReason { + code, + json_pointer: pointer.into(), + }], + } + } + + fn unknown(code: PiComposerReasonCode, pointer: impl Into) -> Self { + Self { + status: PiComposerStatus::Unknown, + provider_id: None, + provider_name: None, + provider_base_url: None, + models: Vec::new(), + ignored_override_keys: Vec::new(), + reasons: vec![PiComposerReason { + code, + json_pointer: pointer.into(), + }], + } + } +} + +pub(super) fn compose_explicit_custom_catalog( + provider_id: &str, + provider: &PiRawValidProvider, +) -> PiNativeComposition { + let Some(provider_object) = provider.raw().as_object() else { + return PiNativeComposition::failed(PiComposerReasonCode::CompositionFailed, ""); + }; + let Some(definitions) = provider_object + .get("models") + .and_then(Value::as_array) + .filter(|models| !models.is_empty()) + else { + return PiNativeComposition::failed(PiComposerReasonCode::MissingExplicitModels, "/models"); + }; + + let provider_api = provider_object.get("api").and_then(Value::as_str); + let provider_base_url = provider_object.get("baseUrl").and_then(Value::as_str); + if provider_object.get("oauth").and_then(Value::as_str) == Some("radius") + && provider_base_url.is_none() + { + return PiNativeComposition::failed( + PiComposerReasonCode::MissingEffectiveEndpoint, + "/baseUrl", + ); + } + let provider_compat = provider_object.get("compat").cloned(); + let provider_header_entries = header_entries(provider_object.get("headers"), "/headers"); + let provider_headers = provider_header_entries + .iter() + .map(|entry| (entry.name.clone(), entry.value.clone())) + .collect::>(); + let provider_extra = unknown_fields(provider_object, PROVIDER_FIELDS); + let api_key = provider_object + .get("apiKey") + .and_then(Value::as_str) + .map(ToOwned::to_owned); + let oauth = provider_object.get("oauth").cloned(); + let auth_header = provider_object + .get("authHeader") + .and_then(Value::as_bool) + .unwrap_or(false); + let overrides = provider_object + .get("modelOverrides") + .and_then(Value::as_object); + + let mut models: Vec = Vec::with_capacity(definitions.len()); + for (index, definition_value) in definitions.iter().enumerate() { + let Some(definition) = definition_value.as_object() else { + return PiNativeComposition::failed( + PiComposerReasonCode::CompositionFailed, + format!("/models/{index}"), + ); + }; + let Some(id) = definition.get("id").and_then(Value::as_str) else { + return PiNativeComposition::failed( + PiComposerReasonCode::CompositionFailed, + format!("/models/{index}/id"), + ); + }; + let existing_index = models.iter().position(|model| model.id == id); + let defaults = existing_index + .and_then(|position| models.get(position)) + .or_else(|| models.first()); + + let api_value = definition + .get("api") + .and_then(Value::as_str) + .or(provider_api) + .or_else(|| defaults.map(|model| model.api.as_str())); + let Some(api_value) = api_value else { + return PiNativeComposition::failed( + PiComposerReasonCode::MissingEffectiveApi, + format!("/models/{index}/api"), + ); + }; + let Some(api) = PiRawApiId::new(api_value) else { + return PiNativeComposition::failed( + PiComposerReasonCode::MissingEffectiveApi, + format!("/models/{index}/api"), + ); + }; + + let base_url = definition + .get("baseUrl") + .and_then(Value::as_str) + .or(provider_base_url) + .or_else(|| defaults.map(|model| model.base_url.as_str())); + let Some(base_url) = base_url.filter(|value| !value.is_empty()) else { + return PiNativeComposition::failed( + PiComposerReasonCode::MissingEffectiveEndpoint, + format!("/models/{index}/baseUrl"), + ); + }; + + for (field, code) in [ + ("contextWindow", PiComposerReasonCode::NonPositiveModelLimit), + ("maxTokens", PiComposerReasonCode::NonPositiveModelLimit), + ] { + if definition + .get(field) + .and_then(Value::as_f64) + .is_some_and(|value| value <= 0.0) + { + return PiNativeComposition::failed(code, format!("/models/{index}/{field}")); + } + } + + let compat = + match merge_pi_compat(provider_compat.clone(), definition.get("compat").cloned()) { + Ok(compat) => compat, + Err(_) => { + return PiNativeComposition::unknown( + PiComposerReasonCode::UnrepresentableCompat, + format!("/models/{index}/compat"), + ) + } + }; + let model = PiComposedNativeModel { + id: id.to_string(), + name: definition + .get("name") + .and_then(Value::as_str) + .unwrap_or(id) + .to_string(), + api, + provider: provider_id.to_string(), + base_url: base_url.to_string(), + reasoning: definition + .get("reasoning") + .and_then(Value::as_bool) + .unwrap_or(false), + thinking_level_map: definition.get("thinkingLevelMap").cloned(), + input: definition + .get("input") + .cloned() + .unwrap_or_else(|| json!(["text"])), + cost: definition.get("cost").cloned().unwrap_or_else(default_cost), + context_window: definition + .get("contextWindow") + .cloned() + .unwrap_or_else(|| json!(128000)), + max_tokens: definition + .get("maxTokens") + .cloned() + .unwrap_or_else(|| json!(16384)), + headers: BTreeMap::new(), + provider_headers: provider_header_entries.clone(), + model_headers: Vec::new(), + compat, + api_key: api_key.clone(), + oauth: oauth.clone(), + auth_header, + provider_extra: provider_extra.clone(), + model_extra: unknown_fields(definition, MODEL_FIELDS), + override_extra: BTreeMap::new(), + }; + if let Some(existing_index) = existing_index { + models[existing_index] = model; + } else { + models.push(model); + } + } + + for model in &mut models { + // Pinned Pi's rawModelHeaders uses Array.find, so duplicate model + // definitions obtain request headers from the first definition even + // though the later definition replaces the composed model slot. + let (definition_index, definition) = definitions + .iter() + .enumerate() + .find_map(|(index, definition)| { + definition + .as_object() + .filter(|definition| { + definition.get("id").and_then(Value::as_str) == Some(model.id.as_str()) + }) + .map(|definition| (index, definition)) + }) + .expect("raw-valid composed model has a source definition"); + let model_override = + overrides.and_then(|overrides| overrides.get(&model.id).and_then(Value::as_object)); + + // rawModelHeaders constructs one case-sensitive JavaScript object from + // override headers followed by the first matching model definition. + // Exact-name replacement keeps its insertion slot; differently-cased + // names remain distinct until ModelRuntime performs its later + // case-insensitive HTTP merge. + let mut model_headers = Vec::new(); + if let Some(model_override) = model_override { + overlay_header_entries( + &mut model_headers, + header_entries( + model_override.get("headers"), + &format!("/modelOverrides/{}/headers", escape_json_pointer(&model.id)), + ), + ); + } + overlay_header_entries( + &mut model_headers, + header_entries( + definition.get("headers"), + &format!("/models/{definition_index}/headers"), + ), + ); + + let mut headers = provider_headers.clone(); + for entry in &model_headers { + headers.insert(entry.name.clone(), entry.value.clone()); + } + model.headers = headers; + model.model_headers = model_headers; + + if let Some(model_override) = model_override { + if let Some(name) = model_override.get("name").and_then(Value::as_str) { + model.name = name.to_string(); + } + if let Some(reasoning) = model_override.get("reasoning").and_then(Value::as_bool) { + model.reasoning = reasoning; + } + if let Some(override_map) = model_override + .get("thinkingLevelMap") + .and_then(Value::as_object) + { + let mut merged = model + .thinking_level_map + .take() + .and_then(|value| value.as_object().cloned()) + .unwrap_or_default(); + merged.extend(override_map.clone()); + model.thinking_level_map = Some(Value::Object(merged)); + } + if let Some(input) = model_override.get("input") { + model.input = input.clone(); + } + if let Some(cost) = model_override.get("cost").and_then(Value::as_object) { + model.cost = merge_cost(&model.cost, cost); + } + if let Some(context_window) = model_override.get("contextWindow") { + model.context_window = context_window.clone(); + } + if let Some(max_tokens) = model_override.get("maxTokens") { + model.max_tokens = max_tokens.clone(); + } + model.compat = match merge_pi_compat( + model.compat.clone(), + model_override.get("compat").cloned(), + ) { + Ok(compat) => compat, + Err(_) => { + return PiNativeComposition::unknown( + PiComposerReasonCode::UnrepresentableCompat, + format!("/modelOverrides/{}/compat", escape_json_pointer(&model.id)), + ) + } + }; + model.override_extra = unknown_fields(model_override, OVERRIDE_FIELDS); + } + } + + let model_ids = models + .iter() + .map(|model| model.id.as_str()) + .collect::>(); + let ignored_override_keys = overrides + .into_iter() + .flat_map(|overrides| overrides.keys()) + .filter(|model_id| !model_ids.contains(model_id.as_str())) + .cloned() + .collect(); + + PiNativeComposition { + status: PiComposerStatus::Composed, + provider_id: Some(provider_id.to_string()), + provider_name: Some( + provider_object + .get("name") + .and_then(Value::as_str) + .unwrap_or(provider_id) + .to_string(), + ), + provider_base_url: provider_base_url.map(ToOwned::to_owned), + models, + ignored_override_keys, + reasons: Vec::new(), + } +} + +fn default_cost() -> Value { + json!({ + "input": 0, + "output": 0, + "cacheRead": 0, + "cacheWrite": 0 + }) +} + +fn merge_cost(base: &Value, overlay: &Map) -> Value { + let base = base.as_object(); + let mut merged = Map::new(); + for key in ["input", "output", "cacheRead", "cacheWrite", "tiers"] { + if let Some(value) = overlay + .get(key) + .or_else(|| base.and_then(|base| base.get(key))) + { + merged.insert(key.to_string(), value.clone()); + } + } + Value::Object(merged) +} + +fn header_entries(value: Option<&Value>, base_pointer: &str) -> Vec { + value + .and_then(Value::as_object) + .into_iter() + .flat_map(|object| object.iter()) + .filter_map(|(name, value)| { + value.as_str().map(|value| PiComposedHeader { + name: name.clone(), + value: value.to_string(), + json_pointer: format!("{base_pointer}/{}", escape_json_pointer(name)), + }) + }) + .collect() +} + +fn overlay_header_entries( + base: &mut Vec, + overlay: impl IntoIterator, +) { + for entry in overlay { + if let Some(existing) = base.iter_mut().find(|existing| existing.name == entry.name) { + *existing = entry; + } else { + base.push(entry); + } + } +} + +fn escape_json_pointer(segment: &str) -> String { + segment.replace('~', "~0").replace('/', "~1") +} + +fn unknown_fields(object: &Map, recognized: &[&str]) -> BTreeMap { + object + .iter() + .filter(|(key, _)| !recognized.contains(&key.as_str())) + .map(|(key, value)| (key.clone(), value.clone())) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::pi_config::raw_schema::{evaluate_provider_value, PiRawValidity}; + use serde::Deserialize; + + const COMPOSER_ORACLE_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/composer-oracle-v1.json"); + + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct ComposerOracle { + cases: Vec, + fail_closed_cases: Vec, + } + + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct ComposerOracleCase { + id: String, + provider_id: String, + input: Value, + execution: Execution, + #[serde(default)] + auth_execution: Option, + #[serde(default)] + expected: Option, + #[serde(default)] + expected_error: Option, + } + + #[derive(Deserialize)] + struct Execution { + status: String, + } + + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct FailClosedCase { + id: String, + rust_expected_status: String, + reason_code: String, + } + + fn model_as_oracle_value(model: &PiComposedNativeModel) -> Value { + let mut object = Map::new(); + object.insert("id".into(), json!(model.id)); + object.insert("name".into(), json!(model.name)); + object.insert("api".into(), json!(model.api.as_str())); + object.insert("provider".into(), json!(model.provider)); + object.insert("baseUrl".into(), json!(model.base_url)); + object.insert("reasoning".into(), json!(model.reasoning)); + if let Some(thinking) = &model.thinking_level_map { + object.insert("thinkingLevelMap".into(), thinking.clone()); + } + object.insert("input".into(), model.input.clone()); + object.insert("cost".into(), model.cost.clone()); + object.insert("contextWindow".into(), model.context_window.clone()); + object.insert("maxTokens".into(), model.max_tokens.clone()); + object.insert("authHeader".into(), json!(model.auth_header)); + if let Some(compat) = &model.compat { + object.insert("compat".into(), compat.clone()); + } + if !model.headers.is_empty() { + object.insert("headers".into(), json!(model.headers)); + } + Value::Object(object) + } + + fn provider_as_oracle_value(composition: &PiNativeComposition) -> Value { + let mut object = Map::new(); + object.insert( + "id".into(), + json!(composition + .provider_id + .as_ref() + .expect("composed provider id")), + ); + object.insert( + "name".into(), + json!(composition + .provider_name + .as_ref() + .expect("composed provider name")), + ); + if let Some(base_url) = &composition.provider_base_url { + object.insert("baseUrl".into(), json!(base_url)); + } + Value::Object(object) + } + + fn json_numbers_equal(left: &Value, right: &Value) -> bool { + match (left, right) { + (Value::Number(left), Value::Number(right)) => left.as_f64() == right.as_f64(), + (Value::Array(left), Value::Array(right)) => { + left.len() == right.len() + && left + .iter() + .zip(right) + .all(|(left, right)| json_numbers_equal(left, right)) + } + (Value::Object(left), Value::Object(right)) => { + left.len() == right.len() + && left.iter().all(|(key, left)| { + right + .get(key) + .is_some_and(|right| json_numbers_equal(left, right)) + }) + } + _ => left == right, + } + } + + #[test] + fn rust_composer_matches_actual_pinned_upstream_execution() { + let oracle: ComposerOracle = + serde_json::from_str(COMPOSER_ORACLE_SOURCE).expect("parse composer oracle"); + for case in oracle.cases { + let raw = evaluate_provider_value(&case.input); + if case.execution.status == "error" { + assert!( + case.expected_error.is_some(), + "upstream error vector '{}' records its actual error", + case.id + ); + match raw.validity { + PiRawValidity::Invalid => {} + PiRawValidity::Valid => { + let result = compose_explicit_custom_catalog( + &case.provider_id, + raw.valid_provider.as_ref().expect("raw-valid provider"), + ); + assert_eq!( + result.status, + PiComposerStatus::Failed, + "raw-valid upstream error case '{}'", + case.id + ); + } + PiRawValidity::Unknown => { + panic!("oracle case '{}' unexpectedly became Unknown", case.id) + } + } + continue; + } + assert_eq!(raw.validity, PiRawValidity::Valid, "case '{}'", case.id); + let result = compose_explicit_custom_catalog( + &case.provider_id, + raw.valid_provider.as_ref().expect("raw-valid provider"), + ); + assert_eq!( + result.status, + PiComposerStatus::Composed, + "case '{}'", + case.id + ); + let auth_execution = case + .auth_execution + .as_ref() + .expect("successful composer case records actual auth execution"); + assert_eq!( + auth_execution.pointer("/status").and_then(Value::as_str), + Some("success"), + "case '{}'", + case.id + ); + let actual_resolved_key = auth_execution + .pointer("/result/auth/apiKey") + .and_then(Value::as_str) + .expect("successful literal composer vector resolves an API key"); + assert!( + result + .models + .iter() + .all(|model| { model.api_key.as_deref() == Some(actual_resolved_key) }), + "case '{}' preserves the same literal key that pinned Pi resolved", + case.id + ); + if case + .input + .get("authHeader") + .and_then(Value::as_bool) + .unwrap_or(false) + { + let expected_bearer = format!("Bearer {actual_resolved_key}"); + assert_eq!( + auth_execution + .pointer("/result/auth/headers/Authorization") + .and_then(Value::as_str), + Some(expected_bearer.as_str()), + "case '{}' uses pinned Pi authHeader behavior", + case.id + ); + } + let actual = json!({ + "provider": provider_as_oracle_value(&result), + "models": result + .models + .iter() + .map(model_as_oracle_value) + .collect::>(), + "ignoredOverrideKeys": result.ignored_override_keys, + }); + let expected = case.expected.expect("successful upstream expected output"); + assert!( + json_numbers_equal(&actual, &expected), + "oracle case '{}'\nactual: {actual:#}\nexpected: {expected:#}", + case.id + ); + } + } + + #[test] + fn unavailable_upstream_catalog_semantics_are_explicitly_unknown() { + let oracle: ComposerOracle = + serde_json::from_str(COMPOSER_ORACLE_SOURCE).expect("parse composer oracle"); + assert_eq!(oracle.fail_closed_cases.len(), 2); + for case in oracle.fail_closed_cases { + assert_eq!(case.rust_expected_status, "unknown", "case '{}'", case.id); + assert_eq!(case.reason_code, "catalog_required", "case '{}'", case.id); + let result = PiNativeComposition::catalog_required("/models"); + assert_eq!(result.status, PiComposerStatus::Unknown); + assert_eq!( + result.reasons[0].code, + PiComposerReasonCode::CatalogRequired + ); + } + } + + #[test] + fn credential_expressions_are_preserved_without_execution() { + let value = json!({ + "api": "openai-responses", + "baseUrl": "https://example.test/v1", + "apiKey": "!read-secret", + "oauth": "radius", + "authHeader": true, + "headers": {"x-tenant": "${TENANT}"}, + "models": [{"id": "m"}] + }); + let raw = evaluate_provider_value(&value); + let composed = compose_explicit_custom_catalog( + "deferred", + raw.valid_provider.as_ref().expect("raw-valid"), + ); + assert_eq!(composed.status, PiComposerStatus::Composed); + assert_eq!(composed.models[0].api_key.as_deref(), Some("!read-secret")); + assert_eq!(composed.models[0].oauth, Some(json!("radius"))); + assert!(composed.models[0].auth_header); + assert_eq!(composed.models[0].headers["x-tenant"], "${TENANT}"); + } + + #[test] + fn pinned_cost_override_reconstructs_only_known_cost_members() { + let value = json!({ + "api": "anthropic-messages", + "baseUrl": "https://cost.example", + "apiKey": "literal", + "models": [{ + "id": "m", + "cost": { + "input": 1, + "output": 2, + "cacheRead": 0.1, + "cacheWrite": 0.2, + "futureRate": 9 + } + }], + "modelOverrides": { + "m": {"cost": {"output": 3}} + } + }); + let raw = evaluate_provider_value(&value); + let composed = compose_explicit_custom_catalog( + "cost-shape", + raw.valid_provider.as_ref().expect("raw-valid"), + ); + assert_eq!( + composed.models[0].cost, + json!({ + "input": 1, + "output": 3, + "cacheRead": 0.1, + "cacheWrite": 0.2 + }), + "pinned applyModelOverride drops unknown base cost keys when an override exists" + ); + } + + #[test] + fn compat_spread_matches_pinned_composer_request_capture() { + let value = json!({ + "api": "openai-responses", + "baseUrl": "https://compat.example/v1", + "apiKey": "literal", + "compat": { + "openRouterRouting": ["first", "second"], + "chatTemplateKwargs": "ab", + "baseOnly": true + }, + "models": [{ + "id": "m", + "compat": {"supportsStore": true} + }], + "modelOverrides": { + "m": { + "compat": { + "openRouterRouting": null, + "chatTemplateKwargs": {"named": true}, + "overlayOnly": true + } + } + } + }); + let raw = evaluate_provider_value(&value); + let composed = compose_explicit_custom_catalog( + "compat-spread", + raw.valid_provider.as_ref().expect("raw-valid"), + ); + + assert_eq!( + composed.models[0].compat, + Some(json!({ + "openRouterRouting": {"0": "first", "1": "second"}, + "chatTemplateKwargs": {"0": "a", "1": "b", "named": true}, + "baseOnly": true, + "supportsStore": true, + "overlayOnly": true + })), + "captured by scripts/pi-transport-capture.mjs at the pinned Pi commit" + ); + } + + #[test] + fn compat_spread_fails_closed_when_pinned_output_requires_lone_surrogates() { + let value = json!({ + "api": "openai-responses", + "baseUrl": "https://compat.example/v1", + "apiKey": "literal", + "compat": {"chatTemplateKwargs": "😀"}, + "models": [{"id": "m"}], + "modelOverrides": { + "m": {"compat": {"chatTemplateKwargs": {"named": true}}} + } + }); + let raw = evaluate_provider_value(&value); + let composition = compose_explicit_custom_catalog( + "compat-surrogate", + raw.valid_provider.as_ref().expect("raw-valid"), + ); + + assert_eq!(composition.status, PiComposerStatus::Unknown); + assert_eq!( + composition.reasons, + vec![PiComposerReason { + code: PiComposerReasonCode::UnrepresentableCompat, + json_pointer: "/modelOverrides/m/compat".to_string(), + }], + "capture records UTF-16 d83d/de00 as two lone-surrogate values, which \ + serde_json::Value cannot represent" + ); + } + + #[test] + fn header_layers_retain_runtime_precedence_and_source_pointers() { + let value = json!({ + "api": "anthropic-messages", + "baseUrl": "https://headers.example", + "apiKey": "literal", + "headers": {"authorization": "Bearer provider"}, + "models": [{ + "id": "m", + "headers": {"Authorization": "Bearer model"} + }], + "modelOverrides": { + "m": {"headers": {"x-layer": "override"}} + } + }); + let raw = evaluate_provider_value(&value); + let composed = compose_explicit_custom_catalog( + "header-layers", + raw.valid_provider.as_ref().expect("raw-valid"), + ); + let model = &composed.models[0]; + assert_eq!( + model.provider_headers[0].json_pointer, + "/headers/authorization" + ); + assert_eq!( + model + .model_headers + .iter() + .map(|entry| (entry.name.as_str(), entry.value.as_str())) + .collect::>(), + vec![("x-layer", "override"), ("Authorization", "Bearer model")] + ); + } + + #[test] + fn unknown_provider_model_and_override_fields_are_retained_losslessly() { + let value = json!({ + "api": "future-wire-v9", + "baseUrl": "https://example.test/v9", + "apiKey": "literal", + "futureProviderShape": { + "nested": [1, {"flag": true}] + }, + "models": [{ + "id": "m", + "futureModelShape": { + "mode": "novel", + "threshold": 0.125 + } + }], + "modelOverrides": { + "m": { + "futureOverrideShape": [ + null, + {"preserve": "exactly"} + ] + } + } + }); + let raw = evaluate_provider_value(&value); + let composed = compose_explicit_custom_catalog( + "lossless", + raw.valid_provider.as_ref().expect("raw-valid"), + ); + assert_eq!(composed.status, PiComposerStatus::Composed); + let model = &composed.models[0]; + assert_eq!( + model.provider_extra["futureProviderShape"], + json!({"nested": [1, {"flag": true}]}) + ); + assert_eq!( + model.model_extra["futureModelShape"], + json!({"mode": "novel", "threshold": 0.125}) + ); + assert_eq!( + model.override_extra["futureOverrideShape"], + json!([null, {"preserve": "exactly"}]) + ); + } +} diff --git a/src-tauri/src/pi_config/document.rs b/src-tauri/src/pi_config/document.rs new file mode 100644 index 000000000..b5a359ef5 --- /dev/null +++ b/src-tauri/src/pi_config/document.rs @@ -0,0 +1,1250 @@ +//! Read-only access to Pi's shared `models.json` document. +//! +//! The semantic parse mirrors Pi's pinned `stripJsonComments()` behavior. +//! A CST parse is additionally required so callers can fingerprint one exact +//! provider value without making unrelated entries part of the revision. + +use crate::error::AppError; +use indexmap::IndexMap; +use jsonc_parser::cst::{ + CstArray, CstContainerNode, CstInputValue, CstLeafNode, CstNode, CstObject, CstObjectProp, + CstRootNode, +}; +use jsonc_parser::ParseOptions; +use regex::Regex; +use serde_json::Value; +use sha2::{Digest, Sha256}; +use std::collections::HashMap; +use std::fs; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, LazyLock, Mutex, MutexGuard}; + +use super::shared_file::{ + compare_exchange_shared_file_bytes, delete_shared_file, read_shared_file, + sync_shared_file_parent, +}; + +const MAX_PI_MODELS_BYTES: u64 = 8 * 1024 * 1024; +const EMPTY_MODELS_DOCUMENT: &str = "{\"providers\":{}}"; +const MAX_MUTATION_ATTEMPTS: usize = 3; + +static PI_JSON_LINE_COMMENTS: LazyLock = LazyLock::new(|| { + Regex::new(r#""(?:\\.|[^"\\])*"|//[^\n]*"#).expect("Pi JSON line-comment regex must compile") +}); +static PI_JSON_TRAILING_COMMAS: LazyLock = LazyLock::new(|| { + Regex::new(r#""(?:\\.|[^"\\])*"|,(\s*[}\]])"#) + .expect("Pi JSON trailing-comma regex must compile") +}); +static PATH_LOCKS: LazyLock>>>> = + LazyLock::new(|| Mutex::new(HashMap::new())); +#[cfg(test)] +static BEFORE_PROVIDER_VERIFY: LazyLock>>> = + LazyLock::new(|| Mutex::new(HashMap::new())); + +fn path_lock(path: &Path) -> Result>, AppError> { + let mut locks = PATH_LOCKS + .lock() + .map_err(|error| AppError::Config(format!("Pi path-lock registry is poisoned: {error}")))?; + Ok(locks + .entry(path.to_path_buf()) + .or_insert_with(|| Arc::new(Mutex::new(()))) + .clone()) +} + +fn lock_path(lock: &Mutex<()>) -> Result, AppError> { + lock.lock() + .map_err(|error| AppError::Config(format!("Pi config path lock is poisoned: {error}"))) +} + +#[derive(Debug, Clone)] +pub(super) struct PiRawProviderEntry { + pub value: Value, + pub raw_source: String, +} + +#[derive(Debug, Clone)] +pub(super) struct PiModelsDocument { + providers: IndexMap, +} + +impl PiModelsDocument { + pub fn providers(&self) -> &IndexMap { + &self.providers + } +} + +fn pi_models_parse_options() -> ParseOptions { + // Pinned Pi accepts standard double-quoted JSON with `//` comments and + // trailing commas. jsonc-parser is broader by default, so all other + // extensions stay disabled. + ParseOptions { + allow_comments: true, + allow_loose_object_property_names: false, + allow_trailing_commas: true, + allow_missing_commas: false, + allow_single_quoted_strings: false, + allow_hexadecimal_numbers: false, + allow_unary_plus_numbers: false, + } +} + +fn jsonc_error(path: &Path, message: impl std::fmt::Display) -> AppError { + AppError::Config(format!( + "JSON parse error in Pi models file {}: {message}", + path.display() + )) +} + +/// Mirrors Pi commit `ab366ebe94cacd419d986be454f12b1b9913aaca` +/// (`packages/coding-agent/src/utils/json.ts`). +fn strip_pi_json_comments(input: &str) -> String { + let without_comments = + PI_JSON_LINE_COMMENTS.replace_all(input, |captures: ®ex::Captures<'_>| { + let matched = captures + .get(0) + .expect("the full regex match is always present") + .as_str(); + if matched.starts_with('"') { + matched.to_string() + } else { + String::new() + } + }); + PI_JSON_TRAILING_COMMAS + .replace_all(&without_comments, |captures: ®ex::Captures<'_>| { + captures + .get(1) + .or_else(|| captures.get(0)) + .expect("the full regex match is always present") + .as_str() + .to_string() + }) + .into_owned() +} + +fn cst_property_name(property: &CstObjectProp) -> Option { + property.name()?.decoded_value().ok() +} + +fn last_cst_property(object: &CstObject, name: &str) -> Option { + object + .properties() + .into_iter() + .rev() + .find(|property| cst_property_name(property).as_deref() == Some(name)) +} + +fn cst_object(node: CstNode, path: &Path, label: &str) -> Result { + match node { + CstNode::Container(CstContainerNode::Object(object)) => Ok(object), + _ => Err(jsonc_error(path, format!("{label} must be an object"))), + } +} + +fn cst_input(value: &Value) -> CstInputValue { + match value { + Value::Null => CstInputValue::Null, + Value::Bool(value) => CstInputValue::Bool(*value), + Value::Number(value) => CstInputValue::Number(value.to_string()), + Value::String(value) => CstInputValue::String(value.clone()), + Value::Array(values) => CstInputValue::Array(values.iter().map(cst_input).collect()), + Value::Object(values) => CstInputValue::Object( + values + .iter() + .map(|(key, value)| (key.clone(), cst_input(value))) + .collect(), + ), + } +} + +fn replace_cst_node(node: CstNode, replacement: &Value) -> Result<(), AppError> { + let replacement = cst_input(replacement); + let replaced = match node { + CstNode::Container(CstContainerNode::Array(node)) => node.replace_with(replacement), + CstNode::Container(CstContainerNode::Object(node)) => node.replace_with(replacement), + CstNode::Leaf(CstLeafNode::BooleanLit(node)) => node.replace_with(replacement), + CstNode::Leaf(CstLeafNode::NullKeyword(node)) => node.replace_with(replacement), + CstNode::Leaf(CstLeafNode::NumberLit(node)) => node.replace_with(replacement), + CstNode::Leaf(CstLeafNode::StringLit(node)) => node.replace_with(replacement), + CstNode::Leaf(CstLeafNode::WordLit(node)) => node.replace_with(replacement), + CstNode::Container(CstContainerNode::Root(_)) + | CstNode::Container(CstContainerNode::ObjectProp(_)) + | CstNode::Leaf(CstLeafNode::Token(_)) + | CstNode::Leaf(CstLeafNode::Whitespace(_)) + | CstNode::Leaf(CstLeafNode::Newline(_)) + | CstNode::Leaf(CstLeafNode::Comment(_)) => None, + }; + replaced.map(|_| ()).ok_or_else(|| { + AppError::Config("Pi models.json CST became disconnected during update".to_string()) + }) +} + +fn patch_cst_object( + object: &CstObject, + before: &serde_json::Map, + after: &serde_json::Map, +) -> Result<(), AppError> { + for key in before.keys().filter(|key| !after.contains_key(*key)) { + let matching = object + .properties() + .into_iter() + .filter(|property| cst_property_name(property).as_deref() == Some(key.as_str())) + .collect::>(); + for property in matching.into_iter().rev() { + property.remove(); + } + } + + for (key, after_value) in after { + if let Some(before_value) = before.get(key) { + let property = last_cst_property(object, key).ok_or_else(|| { + AppError::Config(format!( + "Pi models.json CST is missing existing property '{key}'" + )) + })?; + let value = property.value().ok_or_else(|| { + AppError::Config(format!("Pi models.json CST property '{key}' has no value")) + })?; + patch_cst_node(value, before_value, after_value)?; + } else { + object.append(key, cst_input(after_value)); + } + } + Ok(()) +} + +fn patch_cst_array(array: &CstArray, before: &[Value], after: &[Value]) -> Result<(), AppError> { + let elements = array.elements(); + if elements.len() != before.len() { + return Err(AppError::Config( + "Pi models.json CST array does not match its parsed value".to_string(), + )); + } + + for (index, (before_value, after_value)) in before.iter().zip(after).enumerate() { + patch_cst_node(elements[index].clone(), before_value, after_value)?; + } + for element in elements.into_iter().skip(after.len()).rev() { + element.remove(); + } + for value in after.iter().skip(before.len()) { + array.append(cst_input(value)); + } + Ok(()) +} + +fn patch_cst_node(node: CstNode, before: &Value, after: &Value) -> Result<(), AppError> { + if before == after { + return Ok(()); + } + match (&node, before, after) { + ( + CstNode::Container(CstContainerNode::Object(object)), + Value::Object(before), + Value::Object(after), + ) => patch_cst_object(object, before, after), + ( + CstNode::Container(CstContainerNode::Array(array)), + Value::Array(before), + Value::Array(after), + ) => patch_cst_array(array, before, after), + _ => replace_cst_node(node, after), + } +} + +fn parse_models_source(path: &Path, source: &str) -> Result { + let document: Value = serde_json::from_str(&strip_pi_json_comments(source)) + .map_err(|error| AppError::json(path, error))?; + let root = CstRootNode::parse(source, &pi_models_parse_options()) + .map_err(|error| jsonc_error(path, error))?; + if root.to_serde_value().as_ref() != Some(&document) { + return Err(jsonc_error(path, "CST does not match Pi's parsed document")); + } + + let semantic_providers = document + .as_object() + .and_then(|root| root.get("providers")) + .and_then(Value::as_object) + .ok_or_else(|| jsonc_error(path, "root must contain a providers object"))?; + + let root_object = cst_object( + root.value() + .ok_or_else(|| jsonc_error(path, "document must contain a JSON value"))?, + path, + "root", + )?; + let providers_node = last_cst_property(&root_object, "providers") + .and_then(|property| property.value()) + .ok_or_else(|| jsonc_error(path, "CST is missing the providers value"))?; + let providers_object = cst_object(providers_node, path, "providers")?; + + let mut providers = IndexMap::with_capacity(semantic_providers.len()); + for (provider_key, value) in semantic_providers { + let raw_source = last_cst_property(&providers_object, provider_key) + .and_then(|property| property.value()) + .ok_or_else(|| { + jsonc_error( + path, + format!("CST is missing provider entry '{provider_key}'"), + ) + })? + .to_string(); + providers.insert( + provider_key.clone(), + PiRawProviderEntry { + value: value.clone(), + raw_source, + }, + ); + } + Ok(PiModelsDocument { providers }) +} + +fn read_models_bytes(path: &Path) -> Result>, AppError> { + Ok(read_shared_file(path, MAX_PI_MODELS_BYTES, "Pi models file")?.bytes) +} + +fn parse_pi_models_document( + path: &Path, + bytes: Option<&[u8]>, +) -> Result { + let bytes = bytes.unwrap_or_else(|| EMPTY_MODELS_DOCUMENT.as_bytes()); + let source = std::str::from_utf8(bytes) + .map_err(|error| jsonc_error(path, format!("file is not UTF-8: {error}")))?; + parse_models_source(path, source) +} + +pub(super) fn read_pi_models_document(path: &Path) -> Result { + let bytes = read_models_bytes(path)?; + parse_pi_models_document(path, bytes.as_deref()) +} + +pub(super) fn pi_raw_provider_fingerprint(raw_source: &str) -> String { + format!("sha256:{:x}", Sha256::digest(raw_source.as_bytes())) +} + +fn serialize_models_mutation( + path: &Path, + before: Option<&[u8]>, + mutator: &impl Fn(&mut Value) -> Result<(), AppError>, +) -> Result, AppError> { + if let Some(bytes) = before { + let source = std::str::from_utf8(bytes) + .map_err(|error| jsonc_error(path, format!("file is not UTF-8: {error}")))?; + let mut document: Value = serde_json::from_str(&strip_pi_json_comments(source)) + .map_err(|error| AppError::json(path, error))?; + let root = CstRootNode::parse(source, &pi_models_parse_options()) + .map_err(|error| jsonc_error(path, error))?; + if root.to_serde_value().as_ref() != Some(&document) { + return Err(jsonc_error(path, "CST does not match Pi's parsed document")); + } + let original = document.clone(); + mutator(&mut document)?; + let root_value = root + .value() + .ok_or_else(|| jsonc_error(path, "document must contain a JSON value"))?; + patch_cst_node(root_value, &original, &document)?; + if root.to_serde_value().as_ref() != Some(&document) { + return Err(AppError::Config( + "Pi models.json CST update did not produce the requested document".to_string(), + )); + } + return Ok(root.to_string().into_bytes()); + } + + let mut document: Value = serde_json::from_str(EMPTY_MODELS_DOCUMENT) + .expect("empty Pi models document is valid JSON"); + mutator(&mut document)?; + let mut serialized = serde_json::to_vec_pretty(&document) + .map_err(|source| AppError::JsonSerialize { source })?; + serialized.push(b'\n'); + Ok(serialized) +} + +/// Patch only the explicitly named provider keys in Pi's shared models.json. +/// +/// Unknown root fields, unowned provider entries, comments, and formatting are +/// preserved by the CST patch. An optimistic fingerprint check prevents a +/// Pi/user write observed before replacement from being silently overwritten. +#[cfg(test)] +pub(crate) fn apply_pi_provider_patch( + path: &Path, + patch: &IndexMap>, +) -> Result<(), AppError> { + apply_pi_provider_patch_checked(path, None, None, patch).map(|_| ()) +} + +fn apply_pi_provider_patch_checked( + path: &Path, + expected: Option<&IndexMap>>, + expected_fingerprints: Option<&IndexMap>, + patch: &IndexMap>, +) -> Result { + let lock = path_lock(path)?; + let _guard = lock_path(&lock)?; + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).map_err(|error| AppError::io(parent, error))?; + } + + for _ in 0..MAX_MUTATION_ATTEMPTS { + let before = read_models_bytes(path)?; + if let Some(expected) = expected { + ensure_provider_values_match(path, before.as_deref(), expected)?; + } + if let Some(expected_fingerprints) = expected_fingerprints { + ensure_provider_fingerprints_match(path, before.as_deref(), expected_fingerprints)?; + } + if before.is_none() && patch.values().all(Option::is_none) { + return Ok(PiDocumentCommit { bytes: None }); + } + let serialized = serialize_models_mutation(path, before.as_deref(), &|document| { + let providers = document + .as_object_mut() + .and_then(|root| root.get_mut("providers")) + .and_then(Value::as_object_mut) + .ok_or_else(|| jsonc_error(path, "root must contain a providers object"))?; + for (provider_key, replacement) in patch { + match replacement { + Some(value) => { + providers.insert(provider_key.clone(), value.clone()); + } + None => { + providers.remove(provider_key); + } + } + } + Ok(()) + })?; + if before.as_deref() == Some(serialized.as_slice()) { + return Ok(PiDocumentCommit { bytes: before }); + } + + match compare_exchange_shared_file_bytes( + path, + before.as_deref(), + &serialized, + MAX_PI_MODELS_BYTES, + None, + "Pi models file", + ) { + Ok(_) => { + return Ok(PiDocumentCommit { + bytes: Some(serialized), + }) + } + Err(AppError::Conflict(_)) => continue, + Err(error) => return Err(error), + } + } + + Err(AppError::Conflict(format!( + "Pi models file changed concurrently too many times: {}", + path.display() + ))) +} + +struct PiDocumentCommit { + bytes: Option>, +} + +fn provider_values_from_bytes<'a>( + path: &Path, + bytes: Option<&[u8]>, + provider_keys: impl IntoIterator, +) -> Result>, AppError> { + let document = parse_pi_models_document(path, bytes)?; + Ok(provider_keys + .into_iter() + .map(|key| { + let value = document + .providers() + .get(key) + .map(|entry| entry.value.clone()); + (key.clone(), value) + }) + .collect()) +} + +fn ensure_provider_values_match( + path: &Path, + bytes: Option<&[u8]>, + expected: &IndexMap>, +) -> Result<(), AppError> { + let observed = provider_values_from_bytes(path, bytes, expected.keys())?; + if let Some((provider_key, expected_value)) = expected + .iter() + .find(|(provider_key, expected_value)| observed.get(*provider_key) != Some(*expected_value)) + { + return Err(AppError::Conflict(format!( + "Pi provider key '{provider_key}' changed since directory/catalog preflight \ + (expected {}, observed {})", + provider_value_label(expected_value), + provider_value_label( + observed + .get(provider_key) + .expect("every requested provider key is observed") + ) + ))); + } + Ok(()) +} + +fn ensure_provider_fingerprints_match( + path: &Path, + bytes: Option<&[u8]>, + expected: &IndexMap, +) -> Result<(), AppError> { + let document = parse_pi_models_document(path, bytes)?; + for (provider_key, expected_fingerprint) in expected { + let observed = document + .providers() + .get(provider_key) + .map(|entry| pi_raw_provider_fingerprint(&entry.raw_source)); + if observed.as_deref() != Some(expected_fingerprint) { + return Err(AppError::Conflict(format!( + "Pi native provider '{provider_key}' changed since inspection \ + (expected raw fingerprint {expected_fingerprint}, observed {})", + observed.as_deref().unwrap_or("missing") + ))); + } + } + Ok(()) +} + +fn provider_fingerprints_from_bytes<'a>( + path: &Path, + bytes: Option<&[u8]>, + provider_keys: impl IntoIterator, +) -> Result, AppError> { + let document = parse_pi_models_document(path, bytes)?; + Ok(provider_keys + .into_iter() + .filter_map(|provider_key| { + document.providers().get(provider_key).map(|entry| { + ( + provider_key.clone(), + pi_raw_provider_fingerprint(&entry.raw_source), + ) + }) + }) + .collect()) +} + +fn provider_value_label(value: &Option) -> &'static str { + if value.is_some() { + "present" + } else { + "absent" + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PiProviderValuesSnapshot { + pub file_existed: bool, + pub values: IndexMap>, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PiProviderPatchReceipt { + path: PathBuf, + before: PiProviderValuesSnapshot, + attempted: IndexMap>, + attempted_fingerprints: IndexMap, + attempted_file_existed: bool, +} + +impl PiProviderPatchReceipt { + pub(crate) fn attempted_values(&self) -> &IndexMap> { + &self.attempted + } + + pub(crate) fn attempted_snapshot(&self) -> PiProviderValuesSnapshot { + PiProviderValuesSnapshot { + file_existed: self.attempted_file_existed, + values: self.attempted.clone(), + } + } + + /// Restore only provider keys which still contain this operation's exact + /// attempted values. Unrelated root/provider edits are preserved, while a + /// concurrent edit of an owned key turns compensation into an explicit + /// conflict instead of being overwritten. + pub(crate) fn rollback(&self) -> Result<(), AppError> { + if let Err(error) = apply_pi_provider_patch_checked( + &self.path, + Some(&self.attempted), + Some(&self.attempted_fingerprints), + &self.before.values, + ) { + let observed = + snapshot_pi_provider_values(&self.path, self.before.values.keys().cloned())?; + if observed.values != self.before.values { + return Err(error); + } + // A namespace mutation may have committed before its durability + // barrier reported failure. Re-observe the exact semantic state + // and retry the parent sync before declaring compensation done. + sync_shared_file_parent(&self.path).map_err(|sync_error| { + AppError::Config(format!( + "Pi provider rollback reached the previous exact-key state but could not \ + confirm directory durability ({error}; retry={sync_error})" + )) + })?; + } + remove_new_empty_document(&self.path, &self.before) + } +} + +pub(crate) fn snapshot_pi_provider_values( + path: &Path, + provider_keys: impl IntoIterator, +) -> Result { + let bytes = read_models_bytes(path)?; + let file_existed = bytes.is_some(); + Ok(PiProviderValuesSnapshot { + file_existed, + values: provider_values_from_bytes( + path, + bytes.as_deref(), + provider_keys.into_iter().collect::>().iter(), + )?, + }) +} + +/// Revalidate a previously captured exact-key set without mutating the +/// document. This is the ownership-claim barrier used when runtime takeover is +/// disabled and therefore has no gateway projection write of its own. +pub(crate) fn verify_pi_provider_values( + path: &Path, + expected: &IndexMap>, +) -> Result<(), AppError> { + let lock = path_lock(path)?; + let _guard = lock_path(&lock)?; + #[cfg(test)] + run_before_provider_verify_hook(path)?; + let bytes = read_models_bytes(path)?; + ensure_provider_values_match(path, bytes.as_deref(), expected) +} + +/// Revalidate semantic exact-key values and raw entry fingerprints under one +/// path lock. Native import uses the stronger raw barrier because its ownership +/// token is the fingerprint returned by public inspection. +pub(crate) fn verify_pi_provider_preconditions( + path: &Path, + expected_values: &IndexMap>, + expected_fingerprints: &IndexMap, +) -> Result<(), AppError> { + let lock = path_lock(path)?; + let _guard = lock_path(&lock)?; + #[cfg(test)] + run_before_provider_verify_hook(path)?; + let bytes = read_models_bytes(path)?; + ensure_provider_values_match(path, bytes.as_deref(), expected_values)?; + ensure_provider_fingerprints_match(path, bytes.as_deref(), expected_fingerprints) +} + +#[cfg(test)] +fn run_before_provider_verify_hook(path: &Path) -> Result<(), AppError> { + if let Some(replacement) = BEFORE_PROVIDER_VERIFY + .lock() + .map_err(|error| AppError::Lock(error.to_string()))? + .remove(path) + { + fs::write(path, replacement).map_err(|error| AppError::io(path, error))?; + } + Ok(()) +} + +#[cfg(test)] +pub(crate) fn replace_before_next_pi_provider_verify(path: &Path, bytes: &[u8]) { + BEFORE_PROVIDER_VERIFY + .lock() + .expect("Pi provider verify hook lock") + .insert(path.to_path_buf(), bytes.to_vec()); +} + +fn attempted_provider_values( + before: &PiProviderValuesSnapshot, + patch: &IndexMap>, +) -> IndexMap> { + before + .values + .iter() + .map(|(provider_key, previous)| { + ( + provider_key.clone(), + patch + .get(provider_key) + .cloned() + .unwrap_or_else(|| previous.clone()), + ) + }) + .collect() +} + +/// Publish an exact-key patch only if every preflighted provider value still +/// matches. Whole-file CAS retries preserve unrelated edits, but re-check the +/// provider precondition before every retry. +pub(crate) fn apply_pi_provider_patch_with_receipt( + path: &Path, + before: &PiProviderValuesSnapshot, + patch: &IndexMap>, +) -> Result { + apply_pi_provider_patch_with_receipt_and_fingerprints(path, before, None, patch) +} + +/// Publish an exact-key patch while atomically binding selected entries to the +/// raw inspection fingerprints which authorized an ownership claim. +pub(crate) fn apply_pi_provider_patch_with_receipt_and_fingerprints( + path: &Path, + before: &PiProviderValuesSnapshot, + expected_fingerprints: Option<&IndexMap>, + patch: &IndexMap>, +) -> Result { + if patch + .keys() + .any(|provider_key| !before.values.contains_key(provider_key)) + { + return Err(AppError::InvalidInput( + "Pi provider patch contains a key which was not preflighted".to_string(), + )); + } + let attempted = attempted_provider_values(before, patch); + let commit = + apply_pi_provider_patch_checked(path, Some(&before.values), expected_fingerprints, patch)?; + let attempted_fingerprints = + provider_fingerprints_from_bytes(path, commit.bytes.as_deref(), attempted.keys())?; + Ok(PiProviderPatchReceipt { + path: path.to_path_buf(), + before: before.clone(), + attempted, + attempted_fingerprints, + attempted_file_existed: commit.bytes.is_some(), + }) +} + +/// Remove a file which this operation created only when its bytes are still +/// the canonical empty document. If Pi or the user added any other content, +/// the file is retained. +fn remove_new_empty_document( + path: &Path, + before: &PiProviderValuesSnapshot, +) -> Result<(), AppError> { + if before.file_existed { + return Ok(()); + } + + let canonical_empty = serialize_models_mutation(path, None, &|_| Ok(()))?; + let mut last_error = None; + for _ in 0..MAX_MUTATION_ATTEMPTS { + let current = read_shared_file(path, MAX_PI_MODELS_BYTES, "Pi models file")?; + match current.bytes.as_deref() { + None => { + if last_error.is_some() { + sync_shared_file_parent(path)?; + } + return Ok(()); + } + Some(bytes) if bytes != canonical_empty.as_slice() => { + // A concurrent writer added unrelated content after our key + // rollback. File ownership therefore belongs to that writer. + return Ok(()); + } + Some(_) => {} + } + match delete_shared_file( + path, + ¤t.revision, + MAX_PI_MODELS_BYTES, + "Pi models file", + ) { + Ok(_) => return Ok(()), + Err(error) => last_error = Some(error), + } + } + Err(last_error.unwrap_or_else(|| { + AppError::Config("Pi empty models document cleanup did not make progress".to_string()) + })) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs::File; + use std::io::Write; + + #[test] + fn reads_exact_pi_jsonc_dialect_and_retains_entry_cst() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{ + // root + "providers": { + "custom": { + // nested comment + "baseUrl": "https://example.test/v1", + "models": [{"id": "model//literal",},], + }, + }, +} +"#, + ) + .expect("write"); + + let document = read_pi_models_document(&path).expect("read"); + let entry = &document.providers()["custom"]; + assert_eq!(entry.value["models"][0]["id"], "model//literal"); + assert!(entry.raw_source.contains("// nested comment")); + assert!(entry.raw_source.contains("\"model//literal\"")); + } + + #[test] + fn exact_key_patch_preserves_unowned_entries_comments_and_root_fields() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{ + "theme": "native", + "providers": { + // user-owned + "native": {"models": [{"id": "native"}]}, + "managed": {"models": [{"id": "old"}]} + } +} +"#, + ) + .expect("write"); + let patch = IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"models": [{"id": "new"}]})), + )]); + + apply_pi_provider_patch(&path, &patch).expect("patch"); + + let saved = fs::read_to_string(&path).expect("read"); + assert!(saved.contains("// user-owned")); + assert!(saved.contains("\"theme\": \"native\"")); + let document = read_pi_models_document(&path).expect("parse"); + assert_eq!( + document.providers()["native"].value["models"][0]["id"], + "native" + ); + assert_eq!( + document.providers()["managed"].value["models"][0]["id"], + "new" + ); + } + + #[test] + fn missing_models_snapshot_restores_absence_without_deleting_external_content() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("snapshot absence"); + let receipt = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([("managed".to_string(), Some(serde_json::json!({"api": "x"})))]), + ) + .expect("publish"); + receipt.rollback().expect("restore absence"); + assert!( + !path.exists(), + "rollback must not leave an empty shadow file" + ); + + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("snapshot absence"); + let receipt = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([("managed".to_string(), Some(serde_json::json!({"api": "x"})))]), + ) + .expect("publish"); + let mut external = fs::read_to_string(&path).expect("published document"); + let root_end = external.rfind("\n}").expect("root closing brace"); + external.insert_str(root_end, ",\n \"external\": true"); + fs::write(&path, external).expect("external root-only update"); + receipt.rollback().expect("restore managed key"); + let restored: Value = + serde_json::from_slice(&fs::read(&path).expect("external file retained")) + .expect("parse restored"); + assert_eq!(restored.get("external"), Some(&Value::Bool(true))); + assert!(restored + .get("providers") + .and_then(Value::as_object) + .is_some_and(serde_json::Map::is_empty)); + } + + #[test] + fn checked_provider_publish_rejects_a_key_changed_after_preflight() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write(&path, r#"{"providers":{}}"#).expect("seed"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + crate::pi_config::shared_file::replace_before_next_compare_exchange( + &path, + br#"{"providers":{"managed":{"api":"external"}}}"#, + ); + let error = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"api": "attempted"})), + )]), + ) + .expect_err("external key must win"); + assert!(matches!(error, AppError::Conflict(_))); + let document = read_pi_models_document(&path).expect("external document"); + assert_eq!(document.providers()["managed"].value["api"], "external"); + } + + #[test] + fn provider_receipt_rollback_rejects_a_newer_exact_key_but_preserves_it() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"root":"preserved","providers":{"managed":{"api":"before"}}}"#, + ) + .expect("seed"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + let receipt = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"api": "attempted"})), + )]), + ) + .expect("publish"); + crate::pi_config::shared_file::replace_before_next_compare_exchange( + &path, + br#"{"root":"external","providers":{"managed":{"api":"external"}}}"#, + ); + + let error = receipt + .rollback() + .expect_err("rollback must not overwrite the newer key"); + assert!(matches!(error, AppError::Conflict(_))); + let document: Value = + serde_json::from_slice(&fs::read(&path).expect("read external")).expect("parse"); + assert_eq!(document["root"], "external"); + assert_eq!(document["providers"]["managed"]["api"], "external"); + } + + #[cfg(unix)] + #[test] + fn create_parent_sync_failure_conditionally_removes_the_committed_provider() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + crate::pi_config::shared_file::fail_next_parent_sync_for_test(&path); + + let error = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"api": "attempted"})), + )]), + ) + .expect_err("durability failure must remain visible"); + assert!(error.to_string().contains("injected")); + assert!( + !path.exists(), + "a failed create must not leave a committed provider or empty shadow file" + ); + } + + #[cfg(unix)] + #[test] + fn replace_parent_sync_failure_conditionally_restores_the_previous_provider() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"root":"preserved","providers":{"managed":{"api":"before"}}}"#, + ) + .expect("seed"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + crate::pi_config::shared_file::fail_next_parent_sync_for_test(&path); + + apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"api": "attempted"})), + )]), + ) + .expect_err("durability failure must remain visible"); + let restored: Value = + serde_json::from_slice(&fs::read(&path).expect("restored document")).expect("parse"); + assert_eq!(restored["root"], "preserved"); + assert_eq!(restored["providers"]["managed"]["api"], "before"); + } + + #[cfg(unix)] + #[test] + fn remove_key_parent_sync_failure_conditionally_restores_the_deleted_provider() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{"managed":{"api":"before"},"external":{"api":"keep"}}}"#, + ) + .expect("seed"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + crate::pi_config::shared_file::fail_next_parent_sync_for_test(&path); + + apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([("managed".to_string(), None)]), + ) + .expect_err("durability failure must remain visible"); + let restored: Value = + serde_json::from_slice(&fs::read(&path).expect("restored document")).expect("parse"); + assert_eq!(restored["providers"]["managed"]["api"], "before"); + assert_eq!(restored["providers"]["external"]["api"], "keep"); + } + + #[test] + fn transient_conflict_recovery_failure_still_restores_external_provider_state() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write(&path, br#"{"providers":{"managed":{"api":"before"}}}"#).expect("seed"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + let external = br#"{"providers":{"managed":{"api":"external"}}}"#; + crate::pi_config::shared_file::replace_before_next_compare_exchange(&path, external); + crate::pi_config::shared_file::fail_next_rollback_restore_for_test(&path); + + let error = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"api": "attempted"})), + )]), + ) + .expect_err("an uncertain conflict must remain an error"); + assert!(matches!(error, AppError::Conflict(_))); + let restored = read_pi_models_document(&path).expect("restored document"); + assert_eq!(restored.providers()["managed"].value["api"], "external"); + } + + #[test] + fn transient_delete_recovery_failure_still_restores_external_provider_state() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write(&path, br#"{"providers":{"managed":{"api":"before"}}}"#).expect("seed"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + let external = br#"{"providers":{"managed":{"api":"external"}}}"#; + crate::pi_config::shared_file::replace_before_next_compare_exchange(&path, external); + crate::pi_config::shared_file::fail_next_rollback_restore_for_test(&path); + + let error = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([("managed".to_string(), None)]), + ) + .expect_err("an uncertain delete conflict must remain an error"); + assert!(matches!(error, AppError::Conflict(_))); + let restored = read_pi_models_document(&path).expect("restored document"); + assert_eq!(restored.providers()["managed"].value["api"], "external"); + } + + #[test] + fn precondition_conflict_never_compensates_an_external_same_value_writer() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + let external = br#"{ + "providers": { + // external ownership and formatting must survive + "managed": {"api": "attempted"} + } +}"#; + crate::pi_config::shared_file::replace_before_next_compare_exchange(&path, external); + + let error = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"api": "attempted"})), + )]), + ) + .expect_err("the external create must own the conflict"); + assert!(matches!(error, AppError::Conflict(_))); + assert_eq!( + fs::read(&path).expect("external document preserved"), + external, + "semantic equality must never be used as writer identity" + ); + } + + #[test] + fn receipt_rollback_rejects_same_value_entry_with_new_raw_ownership() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write(&path, br#"{"providers":{"managed":{"api":"before"}}}"#).expect("seed"); + let before = + snapshot_pi_provider_values(&path, ["managed".to_string()]).expect("preflight"); + let receipt = apply_pi_provider_patch_with_receipt( + &path, + &before, + &IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"api": "attempted"})), + )]), + ) + .expect("publish"); + let external = br#"{ + "providers": { + // same value, independently published raw entry + "managed": { "api": "attempted" } + } +}"#; + fs::write(&path, external).expect("external same-value rewrite"); + + let error = receipt + .rollback() + .expect_err("raw ownership change must stop compensation"); + assert!(matches!(error, AppError::Conflict(_))); + assert_eq!(fs::read(&path).expect("external retained"), external); + } + + #[test] + fn external_rename_during_patch_is_reparsed_before_owned_fields_change() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{"managed":{"models":[{"id":"old"}]}}}"#, + ) + .expect("seed"); + crate::pi_config::shared_file::replace_before_next_compare_exchange( + &path, + br#"{ + "externalRevision": 7, + "providers": { + "native": {"models": [{"id": "external"}]}, + "managed": {"models": [{"id": "old"}]} + } +}"#, + ); + let patch = IndexMap::from([( + "managed".to_string(), + Some(serde_json::json!({"models": [{"id": "new"}]})), + )]); + + apply_pi_provider_patch(&path, &patch).expect("retry patch"); + + let saved: Value = serde_json::from_slice(&fs::read(&path).expect("read")).expect("parse"); + assert_eq!(saved["externalRevision"], 7); + assert_eq!(saved["providers"]["native"]["models"][0]["id"], "external"); + assert_eq!(saved["providers"]["managed"]["models"][0]["id"], "new"); + } + + #[test] + fn exact_key_delete_does_not_delete_same_content_sibling() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{ + "managed":{"models":[{"id":"same"}]}, + "native":{"models":[{"id":"same"}]} +}}"#, + ) + .expect("write"); + let patch = IndexMap::from([("managed".to_string(), None)]); + + apply_pi_provider_patch(&path, &patch).expect("delete"); + + let document = read_pi_models_document(&path).expect("parse"); + assert!(!document.providers().contains_key("managed")); + assert!(document.providers().contains_key("native")); + } + + #[test] + fn parses_javascript_overflow_without_hiding_sibling_entries() { + let document = parse_models_source( + Path::new("models.json"), + r#"{ + "providers": { + "healthy": {"models": [{"id": "healthy"}]}, + "overflow": {"models": [{"id": "m", "contextWindow": 1e400}]} + } +}"#, + ) + .expect("JSON.parse accepts the numeric token before TypeBox validation"); + + assert!(document.providers().contains_key("healthy")); + let overflow = &document.providers()["overflow"].value["models"][0]["contextWindow"]; + assert!( + overflow + .as_number() + .is_some_and(|number| number.as_f64().is_none()), + "the raw evaluator must still be able to distinguish non-finite JavaScript Number" + ); + } + + #[test] + fn rejects_json_extensions_that_pinned_pi_rejects() { + let temp = tempfile::tempdir().expect("tempdir"); + for (case, source) in [ + ("single-quoted", "{'providers': {}}"), + ("bare-key", "{providers: {}}"), + ("block-comment", "{\"providers\": {/* nope */}}"), + ("missing-comma", "{\"providers\": {} \"other\": 1}"), + ] { + let path = temp.path().join(format!("{case}.json")); + fs::write(&path, source).expect("write"); + assert!(read_pi_models_document(&path).is_err(), "{case} must fail"); + } + } + + #[test] + fn missing_file_is_an_empty_catalog_and_oversized_file_fails() { + let temp = tempfile::tempdir().expect("tempdir"); + let missing = temp.path().join("missing-models.json"); + assert!(read_pi_models_document(&missing) + .expect("missing") + .providers() + .is_empty()); + + let oversized = temp.path().join("oversized-models.json"); + let mut file = File::create(&oversized).expect("create"); + file.write_all(b"{").expect("seed"); + file.set_len(MAX_PI_MODELS_BYTES + 1).expect("extend"); + assert!(read_pi_models_document(&oversized).is_err()); + } + + #[cfg(unix)] + #[test] + fn symlink_is_never_followed() { + use std::os::unix::fs::symlink; + + let temp = tempfile::tempdir().expect("tempdir"); + let target = temp.path().join("target.json"); + let link = temp.path().join("models.json"); + fs::write(&target, EMPTY_MODELS_DOCUMENT).expect("target"); + symlink(&target, &link).expect("symlink"); + assert!(read_pi_models_document(&link).is_err()); + } +} diff --git a/src-tauri/src/pi_config/gateway.rs b/src-tauri/src/pi_config/gateway.rs new file mode 100644 index 000000000..8e96e05f5 --- /dev/null +++ b/src-tauri/src/pi_config/gateway.rs @@ -0,0 +1,1458 @@ +//! Gateway-only Pi API families and candidate header planning. +//! +//! The four-family enum is intentionally confined to this module. Raw and +//! managed layers retain opaque API identifiers and cannot accidentally reject +//! a future Pi family merely because the gateway has not implemented it. + +#![allow(dead_code)] + +use super::composer::{ + PiComposedHeader, PiComposedNativeModel, PiComposerStatus, PiNativeComposition, +}; +use http::{HeaderMap, HeaderName, HeaderValue}; +use std::fmt; +use url::Url; + +/// Headers owned by HTTP framing, the proxy hop, or gateway-generated request +/// identity. Configured providers may not override them. +const GATEWAY_OWNED_HEADERS: &[&str] = &[ + "connection", + "content-length", + "forwarded", + "host", + "keep-alive", + "te", + "trailer", + "transfer-encoding", + "upgrade", + "x-forwarded-for", + "x-forwarded-host", + "x-forwarded-port", + "x-forwarded-proto", + "x-real-ip", + "cf-connecting-ip", + "cf-ipcountry", + "cf-ray", + "cf-visitor", + "true-client-ip", + "fastly-client-ip", + "x-azure-clientip", + "x-azure-fdid", + "x-azure-ref", + "akamai-origin-hop", + "x-akamai-config-log-detail", + "x-request-id", + "x-correlation-id", + "x-trace-id", + "x-amzn-trace-id", + "x-b3-traceid", + "x-b3-spanid", + "x-b3-parentspanid", + "x-b3-sampled", + "traceparent", + "tracestate", +]; +const CANDIDATE_AUTH_HEADERS: &[&str] = &["authorization", "x-api-key", "x-goog-api-key"]; +const PROTOCOL_HEADERS: &[&str] = &[ + "anthropic-version", + "anthropic-beta", + "openai-beta", + "openai-version", +]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ConfiguredHeaderClass { + GatewayOwned, + Protocol, + CandidateAuth, + CandidateLocal, +} + +fn configured_header_class(name: &HeaderName) -> ConfiguredHeaderClass { + let name = name.as_str(); + if name.starts_with("proxy-") || GATEWAY_OWNED_HEADERS.contains(&name) { + ConfiguredHeaderClass::GatewayOwned + } else if PROTOCOL_HEADERS.contains(&name) { + ConfiguredHeaderClass::Protocol + } else if CANDIDATE_AUTH_HEADERS.contains(&name) { + ConfiguredHeaderClass::CandidateAuth + } else { + ConfiguredHeaderClass::CandidateLocal + } +} + +/// Whether an inbound client header must be replaced by candidate-local or +/// gateway-owned transport state before forwarding to an upstream. +/// +/// Keep this predicate beside configured-header classification so request +/// filtering and provider validation cannot drift into two deny lists. +pub(crate) fn gateway_replaces_incoming_header(name: &HeaderName) -> bool { + configured_header_class(name) != ConfiguredHeaderClass::CandidateLocal +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub(crate) enum PiGatewayApiFamily { + AnthropicMessages, + OpenAiCompletions, + OpenAiResponses, + GoogleGenerativeAi, +} + +impl PiGatewayApiFamily { + pub(super) const ALL: [Self; 4] = [ + Self::AnthropicMessages, + Self::OpenAiCompletions, + Self::OpenAiResponses, + Self::GoogleGenerativeAi, + ]; + + pub(crate) const fn as_str(self) -> &'static str { + match self { + Self::AnthropicMessages => "anthropic-messages", + Self::OpenAiCompletions => "openai-completions", + Self::OpenAiResponses => "openai-responses", + Self::GoogleGenerativeAi => "google-generative-ai", + } + } + + fn parse(value: &str) -> Option { + Self::ALL + .into_iter() + .find(|family| family.as_str() == value) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum PiGatewayCapability { + Proxyable, + DirectOnly, + Unknown, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum PiGatewayReasonCode { + UnsupportedFamily, + UnsupportedCredentialKind, + InvalidEndpoint, + MissingCredential, + InvalidHeaderName, + InvalidHeaderValue, + ProtectedHeader, + DeferredValueUnavailable, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PiGatewayReason { + pub code: PiGatewayReasonCode, + pub json_pointer: String, +} + +#[derive(Debug, Clone)] +pub(crate) struct PiGatewayAssessment { + pub capability: PiGatewayCapability, + pub reasons: Vec, + pub plans: Vec, +} + +#[derive(Clone, PartialEq, Eq)] +struct DeferredHeaderValue { + raw: String, +} + +impl fmt::Debug for DeferredHeaderValue { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("DeferredHeaderValue()") + } +} + +impl DeferredHeaderValue { + fn new(raw: impl Into) -> Self { + Self { raw: raw.into() } + } + + fn materialize( + &self, + resolver: &impl DeferredValueResolver, + pointer: &str, + ) -> Result { + let materialized = if is_deferred(&self.raw) { + resolver.resolve(&self.raw).ok_or_else(|| PiGatewayReason { + code: PiGatewayReasonCode::DeferredValueUnavailable, + json_pointer: pointer.to_string(), + })? + } else { + self.raw.clone() + }; + parse_transport_header_value(&materialized).ok_or_else(|| PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: pointer.to_string(), + }) + } +} + +#[derive(Clone)] +pub(crate) struct CandidateHeaderPlan { + family: PiGatewayApiFamily, + endpoint: Url, + credential: DeferredHeaderValue, + auth_header: bool, + provider_headers: Vec, + model_headers: Vec, + protocol_identity_predictable: bool, +} + +impl fmt::Debug for CandidateHeaderPlan { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("CandidateHeaderPlan") + .field("family", &self.family) + .field("endpoint", &"") + .field("credential", &"") + .field("auth_header", &self.auth_header) + .field("provider_headers", &self.provider_headers) + .field("model_headers", &self.model_headers) + .field( + "protocol_identity_predictable", + &self.protocol_identity_predictable, + ) + .finish() + } +} + +#[derive(Debug, Clone)] +struct PlannedHeader { + name: HeaderName, + json_pointer: String, + value: DeferredHeaderValue, + class: ConfiguredHeaderClass, +} + +#[derive(Clone)] +pub(crate) struct MaterializedCandidate { + pub endpoint: Url, + pub headers: HeaderMap, + family: PiGatewayApiFamily, + protocol_headers: HeaderMap, + protocol_identity_predictable: bool, +} + +impl fmt::Debug for MaterializedCandidate { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let header_names = self.headers.keys().collect::>(); + let protocol_header_names = self.protocol_headers.keys().collect::>(); + formatter + .debug_struct("MaterializedCandidate") + .field("endpoint", &"") + .field("header_names", &header_names) + .field("family", &self.family) + .field("protocol_header_names", &protocol_header_names) + .field( + "protocol_identity_predictable", + &self.protocol_identity_predictable, + ) + .finish() + } +} + +pub(crate) trait DeferredValueResolver { + fn resolve(&self, expression: &str) -> Option; +} + +impl DeferredValueResolver for F +where + F: Fn(&str) -> Option, +{ + fn resolve(&self, expression: &str) -> Option { + self(expression) + } +} + +impl CandidateHeaderPlan { + fn build( + model: &PiComposedNativeModel, + model_index: usize, + allow_anthropic_oauth: bool, + ) -> Result> { + let mut reasons = Vec::new(); + let Some(family) = PiGatewayApiFamily::parse(model.api.as_str()) else { + reasons.push(PiGatewayReason { + code: PiGatewayReasonCode::UnsupportedFamily, + json_pointer: format!("/models/{model_index}/api"), + }); + return Err(reasons); + }; + let endpoint = match parse_pi_gateway_endpoint( + &model.base_url, + format!("/models/{model_index}/baseUrl"), + ) { + Ok(endpoint) => endpoint, + Err(reason) => { + reasons.push(reason); + return Err(reasons); + } + }; + let credential = model.api_key.as_ref(); + match credential { + None => reasons.push(PiGatewayReason { + code: PiGatewayReasonCode::MissingCredential, + json_pointer: "/apiKey".to_string(), + }), + Some(credential) if !is_deferred(credential) => { + if parse_transport_header_value(credential).is_none() { + reasons.push(PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + }); + } else if family == PiGatewayApiFamily::AnthropicMessages + && is_anthropic_oauth_credential(credential) + && !allow_anthropic_oauth + { + reasons.push(PiGatewayReason { + code: PiGatewayReasonCode::UnsupportedCredentialKind, + json_pointer: "/apiKey".to_string(), + }); + } + } + Some(_) => {} + }; + + // Anthropic's credential kind changes the pinned wire protocol: + // OAuth adds a different beta profile. A command-valued credential + // therefore cannot be pre-classified without executing the command, + // and commands are deliberately executed only for the one real + // network attempt. + let mut protocol_identity_predictable = !matches!( + (family, credential), + (PiGatewayApiFamily::AnthropicMessages, Some(value)) if value.starts_with('!') + ); + let provider_headers = plan_configured_headers( + &model.provider_headers, + &mut reasons, + &mut protocol_identity_predictable, + ); + let model_headers = plan_configured_headers( + &model.model_headers, + &mut reasons, + &mut protocol_identity_predictable, + ); + if !reasons.is_empty() { + return Err(reasons); + } + let Some(credential) = credential.cloned() else { + // Missing credentials are accumulated above so configured-header + // diagnostics can be returned in the same assessment. + return Err(vec![PiGatewayReason { + code: PiGatewayReasonCode::MissingCredential, + json_pointer: "/apiKey".to_string(), + }]); + }; + + Ok(Self { + family, + endpoint, + credential: DeferredHeaderValue::new(credential), + auth_header: model.auth_header, + provider_headers, + model_headers, + protocol_identity_predictable, + }) + } + + pub(super) fn materialize( + &self, + resolver: &impl DeferredValueResolver, + ) -> Result { + self.materialize_with_policy(resolver, false) + } + + pub(crate) fn materialize_for_runtime( + &self, + resolver: &impl DeferredValueResolver, + ) -> Result { + self.materialize_with_policy(resolver, true) + } + + pub(crate) fn with_endpoint(&self, endpoint: &str) -> Result { + let endpoint = parse_pi_gateway_endpoint(endpoint, "/customEndpoints")?; + let mut candidate = self.clone(); + candidate.endpoint = endpoint; + Ok(candidate) + } + + /// Resolve only the fields which define failover protocol identity. + /// + /// This deliberately does not touch tenant/custom headers and does not + /// synthesize outbound authentication. The handler uses it before circuit + /// admission so a skipped primary can still constrain compatible + /// failovers without executing unrelated credential commands. + pub(crate) fn materialize_protocol_identity( + &self, + resolver: &impl DeferredValueResolver, + ) -> Result, PiGatewayReason> { + if !self.protocol_identity_predictable { + return Ok(None); + } + + let mut protocol_headers = HeaderMap::new(); + if self.family == PiGatewayApiFamily::AnthropicMessages { + // OAuth changes the pinned Anthropic beta contract, so credential + // kind is the sole auth detail needed by protocol identity. + let credential = self.credential.materialize(resolver, "/apiKey")?; + let anthropic_oauth = credential.to_str().is_ok_and(is_anthropic_oauth_credential); + protocol_headers.insert( + HeaderName::from_static("anthropic-version"), + HeaderValue::from_static("2023-06-01"), + ); + protocol_headers.insert( + HeaderName::from_static("anthropic-beta"), + HeaderValue::from_static(if anthropic_oauth { + "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14" + } else { + "interleaved-thinking-2025-05-14" + }), + ); + } + apply_protocol_headers(&self.provider_headers, resolver, &mut protocol_headers)?; + apply_protocol_headers(&self.model_headers, resolver, &mut protocol_headers)?; + Ok(Some((self.family, protocol_headers))) + } + + fn materialize_with_policy( + &self, + resolver: &impl DeferredValueResolver, + allow_anthropic_oauth: bool, + ) -> Result { + // A new map is allocated for every candidate. No value from a prior + // candidate can survive failover. + let mut headers = HeaderMap::new(); + let mut protocol_headers = HeaderMap::new(); + let credential = self.credential.materialize(resolver, "/apiKey")?; + if self.family == PiGatewayApiFamily::AnthropicMessages + && credential.to_str().is_ok_and(is_anthropic_oauth_credential) + && !allow_anthropic_oauth + { + return Err(PiGatewayReason { + code: PiGatewayReasonCode::UnsupportedCredentialKind, + json_pointer: "/apiKey".to_string(), + }); + } + let bearer_credential = credential.clone(); + let anthropic_oauth = self.family == PiGatewayApiFamily::AnthropicMessages + && credential.to_str().is_ok_and(is_anthropic_oauth_credential); + let (auth_name, auth_value) = match self.family { + PiGatewayApiFamily::AnthropicMessages if anthropic_oauth => { + let credential = credential.to_str().map_err(|_| PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + })?; + let bearer = + HeaderValue::from_str(&format!("Bearer {credential}")).map_err(|_| { + PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + } + })?; + (HeaderName::from_static("authorization"), bearer) + } + PiGatewayApiFamily::AnthropicMessages => { + (HeaderName::from_static("x-api-key"), credential) + } + PiGatewayApiFamily::GoogleGenerativeAi => { + (HeaderName::from_static("x-goog-api-key"), credential) + } + PiGatewayApiFamily::OpenAiCompletions | PiGatewayApiFamily::OpenAiResponses => { + let credential = credential.to_str().map_err(|_| PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + })?; + let bearer = + HeaderValue::from_str(&format!("Bearer {credential}")).map_err(|_| { + PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + } + })?; + (HeaderName::from_static("authorization"), bearer) + } + }; + headers.insert(auth_name, auth_value); + + // The pinned Anthropic SDK contributes this protocol default before + // configured headers. Identity must therefore compare the final value + // even when the user omitted it. + if self.family == PiGatewayApiFamily::AnthropicMessages { + let name = HeaderName::from_static("anthropic-version"); + let value = HeaderValue::from_static("2023-06-01"); + protocol_headers.insert(name.clone(), value.clone()); + headers.insert(name, value); + let beta = if anthropic_oauth { + "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14" + } else { + "interleaved-thinking-2025-05-14" + }; + let name = HeaderName::from_static("anthropic-beta"); + let value = HeaderValue::from_static(beta); + protocol_headers.insert(name.clone(), value.clone()); + headers.insert(name, value); + } + + // Provider headers are part of provider auth resolution. Pinned SDKs + // merge them after synthesized family auth. + apply_planned_headers( + &self.provider_headers, + resolver, + &mut headers, + &mut protocol_headers, + )?; + + // Provider-level authHeader is applied after provider headers. + if self.auth_header { + let credential = bearer_credential.to_str().map_err(|_| PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + })?; + let bearer = HeaderValue::from_str(&format!("Bearer {credential}")).map_err(|_| { + PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + } + })?; + headers.insert(HeaderName::from_static("authorization"), bearer); + } + + // ModelRuntime then performs a case-insensitive model-header overlay. + // Keeping this as a separate phase is essential: flattening the two + // layers before authHeader can send a different credential than Pi. + apply_planned_headers( + &self.model_headers, + resolver, + &mut headers, + &mut protocol_headers, + )?; + + // Main-project OAuth transport is a completed policy boundary, not a + // partial SDK-header overlay. A configured auth header may override + // synthesized auth for ordinary credentials (matching pinned Pi), but + // an Anthropic OAuth credential is always transported as that exact + // Bearer and never alongside x-api-key. + if anthropic_oauth { + headers.remove(HeaderName::from_static("x-api-key")); + let credential = bearer_credential.to_str().map_err(|_| PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + })?; + let bearer = HeaderValue::from_str(&format!("Bearer {credential}")).map_err(|_| { + PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: "/apiKey".to_string(), + } + })?; + headers.insert(HeaderName::from_static("authorization"), bearer); + let beta = HeaderValue::from_static( + "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14", + ); + headers.insert(HeaderName::from_static("anthropic-beta"), beta.clone()); + protocol_headers.insert(HeaderName::from_static("anthropic-beta"), beta); + } + + let host = authority_header(&self.endpoint).ok_or_else(|| PiGatewayReason { + code: PiGatewayReasonCode::InvalidEndpoint, + json_pointer: "/baseUrl".to_string(), + })?; + headers.insert(HeaderName::from_static("host"), host); + Ok(MaterializedCandidate { + endpoint: self.endpoint.clone(), + headers, + family: self.family, + protocol_headers, + protocol_identity_predictable: self.protocol_identity_predictable, + }) + } +} + +fn plan_configured_headers( + entries: &[PiComposedHeader], + reasons: &mut Vec, + protocol_identity_predictable: &mut bool, +) -> Vec { + let mut planned = Vec::with_capacity(entries.len()); + for entry in entries { + let Ok(name) = HeaderName::from_bytes(entry.name.as_bytes()) else { + reasons.push(PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderName, + json_pointer: entry.json_pointer.clone(), + }); + continue; + }; + let class = configured_header_class(&name); + if class == ConfiguredHeaderClass::GatewayOwned { + reasons.push(PiGatewayReason { + code: PiGatewayReasonCode::ProtectedHeader, + json_pointer: entry.json_pointer.clone(), + }); + continue; + } + if !is_deferred(&entry.value) && parse_transport_header_value(&entry.value).is_none() { + reasons.push(PiGatewayReason { + code: PiGatewayReasonCode::InvalidHeaderValue, + json_pointer: entry.json_pointer.clone(), + }); + continue; + } + if class == ConfiguredHeaderClass::Protocol { + *protocol_identity_predictable &= !entry.value.starts_with('!'); + } + planned.push(PlannedHeader { + name, + json_pointer: entry.json_pointer.clone(), + value: DeferredHeaderValue::new(entry.value.clone()), + class, + }); + } + planned +} + +fn apply_planned_headers( + planned: &[PlannedHeader], + resolver: &impl DeferredValueResolver, + headers: &mut HeaderMap, + protocol_headers: &mut HeaderMap, +) -> Result<(), PiGatewayReason> { + for entry in planned { + let value = entry.value.materialize(resolver, &entry.json_pointer)?; + if entry.class == ConfiguredHeaderClass::Protocol { + protocol_headers.insert(entry.name.clone(), value.clone()); + } + headers.insert(entry.name.clone(), value); + } + Ok(()) +} + +fn apply_protocol_headers( + planned: &[PlannedHeader], + resolver: &impl DeferredValueResolver, + protocol_headers: &mut HeaderMap, +) -> Result<(), PiGatewayReason> { + for entry in planned + .iter() + .filter(|entry| entry.class == ConfiguredHeaderClass::Protocol) + { + protocol_headers.insert( + entry.name.clone(), + entry.value.materialize(resolver, &entry.json_pointer)?, + ); + } + Ok(()) +} + +impl MaterializedCandidate { + pub(crate) fn family(&self) -> PiGatewayApiFamily { + self.family + } + + pub(crate) fn failover_protocol_identity(&self) -> Option<(PiGatewayApiFamily, &HeaderMap)> { + // Auth, tenant and arbitrary custom headers are deliberately excluded. + self.protocol_identity_predictable + .then_some((self.family, &self.protocol_headers)) + } + + pub(crate) fn family_name(&self) -> &'static str { + self.family.as_str() + } +} + +impl CandidateHeaderPlan { + pub(crate) fn family(&self) -> PiGatewayApiFamily { + self.family + } + + pub(crate) fn protocol_identity_is_predictable(&self) -> bool { + self.protocol_identity_predictable + } + + pub(crate) fn endpoint(&self) -> &Url { + &self.endpoint + } +} + +pub(super) fn assess_composition(composition: &PiNativeComposition) -> PiGatewayAssessment { + if composition.status != PiComposerStatus::Composed { + return PiGatewayAssessment { + capability: PiGatewayCapability::Unknown, + reasons: Vec::new(), + plans: Vec::new(), + }; + } + let mut plans = Vec::with_capacity(composition.models.len()); + let mut reasons = Vec::new(); + for (index, model) in composition.models.iter().enumerate() { + match CandidateHeaderPlan::build(model, index, false) { + Ok(plan) => plans.push(plan), + Err(mut model_reasons) => reasons.append(&mut model_reasons), + } + } + if reasons.is_empty() && plans.len() == composition.models.len() { + PiGatewayAssessment { + capability: PiGatewayCapability::Proxyable, + reasons, + plans, + } + } else { + PiGatewayAssessment { + capability: PiGatewayCapability::DirectOnly, + reasons, + plans: Vec::new(), + } + } +} + +/// Main-project data plane assessment. The certified Pre-C assessment remains +/// unchanged and honestly reports Anthropic OAuth as DirectOnly; this entry +/// point becomes reachable only with the complete OAuth transport policy. +pub(crate) fn assess_composition_for_runtime( + composition: &PiNativeComposition, +) -> PiGatewayAssessment { + if composition.status != PiComposerStatus::Composed { + return PiGatewayAssessment { + capability: PiGatewayCapability::Unknown, + reasons: Vec::new(), + plans: Vec::new(), + }; + } + let mut plans = Vec::with_capacity(composition.models.len()); + let mut reasons = Vec::new(); + for (index, model) in composition.models.iter().enumerate() { + match CandidateHeaderPlan::build(model, index, true) { + Ok(plan) => plans.push(plan), + Err(mut model_reasons) => reasons.append(&mut model_reasons), + } + } + if reasons.is_empty() && plans.len() == composition.models.len() { + PiGatewayAssessment { + capability: PiGatewayCapability::Proxyable, + reasons, + plans, + } + } else { + PiGatewayAssessment { + capability: PiGatewayCapability::DirectOnly, + reasons, + plans: Vec::new(), + } + } +} + +fn authority_header(url: &Url) -> Option { + let host = match url.host()? { + url::Host::Domain(value) => value.to_string(), + url::Host::Ipv4(value) => value.to_string(), + url::Host::Ipv6(value) => format!("[{value}]"), + }; + let authority = match url.port() { + Some(port) => format!("{host}:{port}"), + None => host, + }; + HeaderValue::from_str(&authority).ok() +} + +fn is_deferred(value: &str) -> bool { + value.starts_with('!') || value.contains('$') +} + +fn is_anthropic_oauth_credential(value: &str) -> bool { + // Pinned Pi's Anthropic adapter uses `includes`, not a prefix test. + value.contains("sk-ant-oat") +} + +fn parse_transport_header_value(value: &str) -> Option { + if !value.bytes().all(|byte| matches!(byte, 0x20..=0x7e)) { + return None; + } + HeaderValue::from_str(value).ok() +} + +/// Parse the one endpoint domain accepted by Pi's gateway. +/// +/// Control-plane endpoint mutations and runtime candidate construction both +/// use this function so a value cannot be accepted for storage and rejected +/// only after a request starts. Pi-native base URLs may remain visible as +/// direct-only diagnostics; this validator governs gateway-owned routes. +pub(crate) fn parse_pi_gateway_endpoint( + value: &str, + json_pointer: impl Into, +) -> Result { + let json_pointer = json_pointer.into(); + let endpoint = Url::parse(value).map_err(|_| PiGatewayReason { + code: PiGatewayReasonCode::InvalidEndpoint, + json_pointer: json_pointer.clone(), + })?; + if matches!(endpoint.scheme(), "http" | "https") + && endpoint.host().is_some() + && endpoint.username().is_empty() + && endpoint.password().is_none() + { + Ok(endpoint) + } else { + Err(PiGatewayReason { + code: PiGatewayReasonCode::InvalidEndpoint, + json_pointer, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::pi_config::composer::compose_explicit_custom_catalog; + use crate::pi_config::raw_schema::evaluate_provider_value; + use serde_json::{json, Value}; + use std::collections::BTreeMap; + + const TRANSPORT_ORACLE_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/transport-oracle-v1.json"); + + fn composed(input: serde_json::Value) -> PiNativeComposition { + let raw = evaluate_provider_value(&input); + compose_explicit_custom_catalog( + "candidate", + raw.valid_provider.as_ref().expect("raw-valid input"), + ) + } + + #[test] + fn only_gateway_module_closes_the_four_family_set() { + assert_eq!( + PiGatewayApiFamily::ALL.map(PiGatewayApiFamily::as_str), + [ + "anthropic-messages", + "openai-completions", + "openai-responses", + "google-generative-ai", + ] + ); + } + + #[test] + fn sensitive_gateway_debug_output_redacts_credentials_and_header_values() { + let credential = "sk-debug-credential-never-log"; + let header_secret = "debug-header-value-never-log"; + let query_secret = "debug-query-never-log"; + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": format!("https://example.test/v1?token={query_secret}"), + "apiKey": credential, + "headers": {"x-private": header_secret}, + "models": [{"id": "m"}] + })); + let assessment = assess_composition(&composition); + assert_eq!(assessment.capability, PiGatewayCapability::Proxyable); + let assessment_debug = format!("{assessment:?}"); + assert!(!assessment_debug.contains(credential)); + assert!(!assessment_debug.contains(header_secret)); + assert!(!assessment_debug.contains(query_secret)); + + let materialized = assessment.plans[0] + .materialize(&|_expression: &str| None) + .expect("literal plan materializes"); + let materialized_debug = format!("{materialized:?}"); + assert!(!materialized_debug.contains(credential)); + assert!(!materialized_debug.contains(header_secret)); + assert!(!materialized_debug.contains(query_secret)); + assert!(materialized_debug.contains("authorization")); + assert!(materialized_debug.contains("x-private")); + } + + #[test] + fn gateway_rejects_endpoint_userinfo_and_never_debugs_it() { + let userinfo_secret = "userinfo-secret-never-log"; + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": format!("https://user:{userinfo_secret}@example.test/v1"), + "apiKey": "credential", + "models": [{"id": "m"}] + })); + let assessment = assess_composition(&composition); + assert_eq!(assessment.capability, PiGatewayCapability::DirectOnly); + assert_eq!( + assessment.reasons[0].code, + PiGatewayReasonCode::InvalidEndpoint + ); + assert!(!format!("{assessment:?}").contains(userinfo_secret)); + + let safe = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://example.test/v1", + "apiKey": "credential", + "models": [{"id": "m"}] + })); + let plan = assess_composition(&safe).plans.remove(0); + let error = plan + .with_endpoint(&format!( + "https://user:{userinfo_secret}@failover.example/v1" + )) + .expect_err("custom endpoint userinfo must be rejected"); + assert_eq!(error.code, PiGatewayReasonCode::InvalidEndpoint); + } + + #[test] + fn unknown_api_is_composed_but_direct_only() { + let composition = composed(json!({ + "api": "future-wire-v9", + "baseUrl": "https://future.example/v9", + "apiKey": "literal", + "models": [{"id": "future"}] + })); + assert_eq!(composition.status, PiComposerStatus::Composed); + let gateway = assess_composition(&composition); + assert_eq!(gateway.capability, PiGatewayCapability::DirectOnly); + assert_eq!( + gateway.reasons[0].code, + PiGatewayReasonCode::UnsupportedFamily + ); + } + + #[test] + fn capability_plan_rejects_invalid_protected_and_hop_headers() { + for (name, expected) in [ + ("bad header", PiGatewayReasonCode::InvalidHeaderName), + ("Host", PiGatewayReasonCode::ProtectedHeader), + ("Content-Length", PiGatewayReasonCode::ProtectedHeader), + ("Connection", PiGatewayReasonCode::ProtectedHeader), + ] { + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://example.test/v1", + "apiKey": "literal", + "headers": {name: "value"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&composition); + assert_eq!( + gateway.capability, + PiGatewayCapability::DirectOnly, + "{name}" + ); + assert_eq!(gateway.reasons[0].code, expected, "{name}"); + } + + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://example.test/v1", + "apiKey": "literal", + "headers": {"x-inject": "ok\r\nbad: value"}, + "models": [{"id": "m"}] + })); + // TypeBox rejects non-header string syntax only at the gateway layer; + // the raw schema intentionally accepts arbitrary strings. + let gateway = assess_composition(&composition); + assert_eq!( + gateway.reasons[0].code, + PiGatewayReasonCode::InvalidHeaderValue + ); + } + + #[test] + fn deferred_materialization_is_per_candidate_and_precedes_network_io() { + let first = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://first.example:8443/v1", + "apiKey": "${FIRST_KEY}", + "headers": {"x-tenant": "${FIRST_TENANT}"}, + "models": [{"id": "m"}] + })); + let second = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://second.example/v1", + "apiKey": "${SECOND_KEY}", + "headers": {"x-tenant": "${SECOND_TENANT}"}, + "models": [{"id": "m"}] + })); + let first_plan = assess_composition(&first).plans.remove(0); + let second_plan = assess_composition(&second).plans.remove(0); + + let first_values = BTreeMap::from([ + ("${FIRST_KEY}".to_string(), "first-secret".to_string()), + ("${FIRST_TENANT}".to_string(), "tenant-a".to_string()), + ]); + let first_materialized = first_plan + .materialize(&|expression: &str| first_values.get(expression).cloned()) + .expect("first materialization"); + assert_eq!( + first_materialized.headers[&HeaderName::from_static("host")], + "first.example:8443" + ); + assert_eq!( + first_materialized.headers[&HeaderName::from_static("authorization")], + "Bearer first-secret" + ); + + let missing = second_plan.materialize(&|_expression: &str| None); + assert_eq!( + missing.expect_err("missing deferred value").code, + PiGatewayReasonCode::DeferredValueUnavailable + ); + + let second_values = BTreeMap::from([ + ("${SECOND_KEY}".to_string(), "second-secret".to_string()), + ("${SECOND_TENANT}".to_string(), "tenant-b".to_string()), + ]); + let second_materialized = second_plan + .materialize(&|expression: &str| second_values.get(expression).cloned()) + .expect("second materialization"); + assert_eq!( + second_materialized.headers[&HeaderName::from_static("host")], + "second.example" + ); + assert_eq!( + second_materialized.headers[&HeaderName::from_static("authorization")], + "Bearer second-secret" + ); + assert_eq!( + second_materialized.headers[&HeaderName::from_static("x-tenant")], + "tenant-b" + ); + assert_ne!( + first_materialized.headers[&HeaderName::from_static("authorization")], + second_materialized.headers[&HeaderName::from_static("authorization")] + ); + assert_eq!( + first_materialized.failover_protocol_identity(), + second_materialized.failover_protocol_identity() + ); + } + + #[test] + fn protocol_identity_uses_final_values_but_excludes_auth_tenant_and_custom_headers() { + let first = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://first.example/v1", + "apiKey": "first-secret", + "headers": { + "anthropic-version": "${VERSION_A}", + "x-tenant": "tenant-a", + "x-private": "private-a" + }, + "models": [{"id": "m"}] + })); + let second = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://second.example/v1", + "apiKey": "second-secret", + "headers": { + "anthropic-version": "${VERSION_B}", + "x-tenant": "tenant-b", + "x-private": "private-b" + }, + "models": [{"id": "m"}] + })); + let first = assess_composition(&first) + .plans + .remove(0) + .materialize(&|expression: &str| { + (expression == "${VERSION_A}").then(|| "2023-06-01".to_string()) + }) + .expect("first protocol materialization"); + let same_protocol = assess_composition(&second) + .plans + .remove(0) + .materialize(&|expression: &str| { + (expression == "${VERSION_B}").then(|| "2023-06-01".to_string()) + }) + .expect("second protocol materialization"); + assert_eq!( + first.failover_protocol_identity(), + same_protocol.failover_protocol_identity(), + "auth, origin, tenant and arbitrary custom headers do not affect wire identity" + ); + + let changed = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://third.example/v1", + "apiKey": "third-secret", + "headers": {"anthropic-version": "2024-01-01"}, + "models": [{"id": "m"}] + })); + let changed = assess_composition(&changed) + .plans + .remove(0) + .materialize(&|_expression: &str| None) + .expect("changed protocol materialization"); + assert_ne!( + first.failover_protocol_identity(), + changed.failover_protocol_identity() + ); + } + + #[test] + fn deferred_protocol_headers_fail_before_candidate_use_and_commands_are_ineligible() { + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://candidate.example/v1", + "apiKey": "literal", + "headers": {"openai-version": "${OPENAI_VERSION}"}, + "models": [{"id": "m"}] + })); + let plan = assess_composition(&composition).plans.remove(0); + assert_eq!( + plan.materialize(&|_expression: &str| None) + .expect_err("unresolved protocol material") + .code, + PiGatewayReasonCode::DeferredValueUnavailable + ); + + let command = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://candidate.example/v1", + "apiKey": "literal", + "headers": {"openai-version": "!resolve-version"}, + "models": [{"id": "m"}] + })); + let command = assess_composition(&command) + .plans + .remove(0) + .materialize(&|expression: &str| { + (expression == "!resolve-version").then(|| "2024-01-01".to_string()) + }) + .expect("command materializes for direct candidate use"); + assert!( + command.failover_protocol_identity().is_none(), + "unpredictable protocol commands are failover-ineligible" + ); + } + + #[test] + fn deferred_header_values_use_one_post_resolution_validator() { + let expression = "!echo café"; + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://candidate.example/v1", + "apiKey": "literal", + "headers": {"x-tenant": expression}, + "models": [{"id": "m"}] + })); + let plan = assess_composition(&composition).plans.remove(0); + let resolved = plan + .materialize(&|value: &str| { + (value == expression).then(|| "resolved-secret".to_string()) + }) + .expect("the resolved visible-ASCII value is valid"); + assert_eq!( + resolved.headers[&HeaderName::from_static("x-tenant")], + "resolved-secret" + ); + assert_eq!( + plan.materialize(&|value: &str| (value == expression).then(|| "café".to_string())) + .expect_err("the resolved value still passes through transport validation") + .code, + PiGatewayReasonCode::InvalidHeaderValue + ); + + let literal = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://candidate.example/v1", + "apiKey": "literal", + "headers": {"x-tenant": "café"}, + "models": [{"id": "m"}] + })); + assert_eq!( + assess_composition(&literal).reasons[0].code, + PiGatewayReasonCode::InvalidHeaderValue + ); + } + + #[test] + fn auth_header_adds_candidate_local_bearer_without_reusing_another_candidate() { + let composition = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://candidate.example/v1", + "apiKey": "${KEY}", + "authHeader": true, + "models": [{"id": "m"}] + })); + let materialized = assess_composition(&composition) + .plans + .remove(0) + .materialize(&|expression: &str| { + (expression == "${KEY}").then(|| "candidate-secret".to_string()) + }) + .expect("authHeader materialization"); + assert_eq!( + materialized.headers[&HeaderName::from_static("x-api-key")], + "candidate-secret" + ); + assert_eq!( + materialized.headers[&HeaderName::from_static("authorization")], + "Bearer candidate-secret" + ); + } + + #[test] + fn model_headers_overlay_provider_auth_case_insensitively() { + for (provider_name, model_name) in [ + ("authorization", "Authorization"), + ("Authorization", "authorization"), + ] { + let composition = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://candidate.example/v1", + "apiKey": "candidate-secret", + "authHeader": true, + "headers": {provider_name: "Bearer provider-token"}, + "models": [{ + "id": "m", + "headers": {model_name: "Bearer model-token"} + }] + })); + let materialized = assess_composition(&composition) + .plans + .remove(0) + .materialize(&|_expression: &str| None) + .expect("layered header materialization"); + assert_eq!( + materialized.headers[&HeaderName::from_static("authorization")], + "Bearer model-token", + "pinned ModelRuntime applies model headers after provider authHeader" + ); + } + } + + #[test] + fn anthropic_protocol_identity_includes_the_sdk_default() { + let omitted = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://first.example/v1", + "apiKey": "first", + "models": [{"id": "m"}] + })); + let explicit = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://second.example/v1", + "apiKey": "second", + "headers": {"anthropic-version": "2023-06-01"}, + "models": [{"id": "m"}] + })); + let omitted = assess_composition(&omitted) + .plans + .remove(0) + .materialize(&|_expression: &str| None) + .expect("omitted SDK default"); + let explicit = assess_composition(&explicit) + .plans + .remove(0) + .materialize(&|_expression: &str| None) + .expect("explicit SDK default"); + assert_eq!( + omitted.headers[&HeaderName::from_static("anthropic-version")], + "2023-06-01" + ); + assert_eq!( + omitted.failover_protocol_identity(), + explicit.failover_protocol_identity(), + "wire-equivalent omitted and explicit SDK defaults must remain failover-compatible" + ); + } + + #[test] + fn four_families_materialize_their_own_auth_headers() { + for (family, auth_name, expected_value) in [ + ("anthropic-messages", "x-api-key", "secret"), + ("openai-completions", "authorization", "Bearer secret"), + ("openai-responses", "authorization", "Bearer secret"), + ("google-generative-ai", "x-goog-api-key", "secret"), + ] { + let composition = composed(json!({ + "api": family, + "baseUrl": "https://candidate.example/v1", + "apiKey": "secret", + "models": [{"id": "m"}] + })); + let mut assessment = assess_composition(&composition); + assert_eq!(assessment.capability, PiGatewayCapability::Proxyable); + let materialized = assessment + .plans + .remove(0) + .materialize(&|_expression: &str| None) + .expect("literal materialization"); + assert_eq!( + materialized.headers[&HeaderName::from_bytes(auth_name.as_bytes()).unwrap()], + expected_value + ); + } + } + + #[test] + fn candidate_host_header_preserves_ipv6_authority_brackets() { + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": "http://[::1]:8443/v1", + "apiKey": "secret", + "models": [{"id": "m"}] + })); + let materialized = assess_composition(&composition) + .plans + .remove(0) + .materialize(&|_expression: &str| None) + .expect("IPv6 endpoint"); + assert_eq!( + materialized.headers[&HeaderName::from_static("host")], + "[::1]:8443" + ); + } + + #[test] + fn completed_runtime_oauth_policy_forces_bearer_beta_and_no_x_api_key() { + let composition = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://candidate.example", + "apiKey": "prefix-sk-ant-oat01-token-suffix", + "headers": { + "authorization": "Bearer configured", + "x-api-key": "configured-api-key", + "anthropic-beta": "configured-beta" + }, + "models": [{"id": "m"}] + })); + let certified = assess_composition(&composition); + assert_eq!(certified.capability, PiGatewayCapability::DirectOnly); + assert_eq!( + certified.reasons[0].code, + PiGatewayReasonCode::UnsupportedCredentialKind + ); + + let mut runtime = assess_composition_for_runtime(&composition); + assert_eq!(runtime.capability, PiGatewayCapability::Proxyable); + let materialized = runtime + .plans + .remove(0) + .materialize_for_runtime(&|_expression: &str| None) + .expect("the complete main-project OAuth policy is proxyable"); + assert_eq!( + materialized.headers[&HeaderName::from_static("authorization")], + "Bearer prefix-sk-ant-oat01-token-suffix" + ); + assert!(materialized.headers.get("x-api-key").is_none()); + assert_eq!( + materialized.headers[&HeaderName::from_static("anthropic-beta")], + "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14" + ); + } + + #[test] + fn anthropic_command_credentials_are_single_direct_attempts_after_materialization() { + for (resolved, expected_auth, oauth) in [ + ("ordinary-secret", "ordinary-secret", false), + ( + "prefix-sk-ant-oat01-token-suffix", + "Bearer prefix-sk-ant-oat01-token-suffix", + true, + ), + ] { + let composition = composed(json!({ + "api": "anthropic-messages", + "baseUrl": "https://candidate.example", + "apiKey": "!credential-command", + "models": [{"id": "m"}] + })); + let mut assessment = assess_composition_for_runtime(&composition); + assert_eq!(assessment.capability, PiGatewayCapability::Proxyable); + let plan = assessment.plans.remove(0); + assert!( + !plan.protocol_identity_is_predictable(), + "credential commands affecting Anthropic beta must not be pre-executed" + ); + let materialized = plan + .materialize_for_runtime(&|expression: &str| { + (expression == "!credential-command").then(|| resolved.to_string()) + }) + .expect("materialize command result once"); + let auth_name = if oauth { "authorization" } else { "x-api-key" }; + assert_eq!( + materialized.headers[&HeaderName::from_static(auth_name)], + expected_auth + ); + assert_eq!( + materialized.headers.get("x-api-key").is_none(), + oauth, + "OAuth must never be proxied through x-api-key" + ); + } + } + + #[test] + fn deferred_values_replay_actual_pinned_pi_transport_results() { + let oracle: Value = + serde_json::from_str(TRANSPORT_ORACLE_SOURCE).expect("parse transport oracle"); + for case in oracle["cases"].as_array().expect("transport cases") { + let input = case["input"].as_str().expect("transport input"); + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://candidate.example/v1", + "apiKey": input, + "models": [{"id": "m"}] + })); + let plan = assess_composition(&composition).plans.remove(0); + match case.pointer("/execution/status").and_then(Value::as_str) { + Some("success") => { + let expected = case["expected"].as_str().expect("actual Pi result"); + let materialized = plan + .materialize(&|expression: &str| { + (expression == input).then(|| expected.to_string()) + }) + .expect("replay actual Pi resolver result"); + let expected_bearer = format!("Bearer {expected}"); + assert_eq!( + materialized.headers[&HeaderName::from_static("authorization")], + expected_bearer, + "transport case '{}'", + case["id"] + ); + } + Some("error") => { + assert!(case["expectedError"].is_string()); + assert_eq!( + plan.materialize(&|_expression: &str| None) + .expect_err("actual Pi resolver failure must discard candidate") + .code, + PiGatewayReasonCode::DeferredValueUnavailable, + "transport case '{}'", + case["id"] + ); + } + status => panic!("unexpected transport status {status:?}"), + } + } + + let header_case = &oracle["headerCase"]; + let input_headers = header_case["input"] + .as_object() + .expect("transport header input"); + let expected_headers = header_case["expected"] + .as_object() + .expect("actual Pi header output"); + let composition = composed(json!({ + "api": "openai-responses", + "baseUrl": "https://candidate.example/v1", + "apiKey": "literal-key", + "headers": input_headers, + "models": [{"id": "m"}] + })); + let materialized = assess_composition(&composition) + .plans + .remove(0) + .materialize(&|expression: &str| { + input_headers.iter().find_map(|(name, configured)| { + (configured.as_str() == Some(expression)) + .then(|| expected_headers[name].as_str().map(ToOwned::to_owned)) + .flatten() + }) + }) + .expect("replay actual Pi header resolver results"); + for (name, expected) in expected_headers { + let name = HeaderName::from_bytes(name.as_bytes()).expect("oracle header name"); + assert_eq!( + materialized.headers[&name], + expected.as_str().expect("oracle header value") + ); + } + } +} diff --git a/src-tauri/src/pi_config/mod.rs b/src-tauri/src/pi_config/mod.rs new file mode 100644 index 000000000..9d7ec41f7 --- /dev/null +++ b/src-tauri/src/pi_config/mod.rs @@ -0,0 +1,218 @@ +//! Pi Coding Agent integration boundaries. +//! +//! This module deliberately separates the managed control-plane model from +//! Pi's shared files and from the proxy data plane. Callers must use the +//! typed model resolver rather than reimplementing provider/model inheritance. + +use indexmap::IndexMap; +use serde_json::{Map, Value}; + +pub(crate) mod composer; +pub(crate) mod document; +pub(crate) mod gateway; +pub(crate) mod model; +pub(crate) mod native; +#[cfg(test)] +mod native_inspection_certification; +pub(crate) mod native_settings; +pub(crate) mod raw_schema; +pub(crate) mod shared_file; + +const PI_COMPAT_NESTED_SPREAD_KEYS: [&str; 3] = [ + "openRouterRouting", + "vercelGatewayRouting", + "chatTemplateKwargs", +]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct PiCompatMergeError; + +#[derive(Debug, Clone, PartialEq)] +enum JavaScriptSpreadValue { + Json(Value), + LoneSurrogate, +} + +type JavaScriptSpreadMap = IndexMap; + +/// Mirror pinned Pi's `mergeCompat` JavaScript object-spread semantics. +/// +/// Arrays expose numeric enumerable properties, strings expose character +/// properties, objects expose their own fields, and the remaining JSON +/// primitives expose none. Existing key positions are retained when an +/// overlay replaces their values, matching object spread. +/// +/// A JavaScript string is indexed by UTF-16 code unit. Spreading an astral +/// character therefore creates lone-surrogate string values, which cannot be +/// represented by Rust `String` or `serde_json::Value`. That shape is rejected +/// explicitly so callers can fail closed instead of emitting a different +/// composed model. +fn merge_pi_compat( + base: Option, + overlay: Option, +) -> Result, PiCompatMergeError> { + let Some(overlay) = overlay else { + return Ok(base); + }; + if !javascript_truthy(&overlay) { + return Ok(base); + } + + let mut merged = javascript_object_spread(base.as_ref()); + merged.extend(javascript_object_spread(Some(&overlay))); + + for key in PI_COMPAT_NESTED_SPREAD_KEYS { + let base_value = javascript_property(base.as_ref(), key); + let overlay_value = javascript_property(Some(&overlay), key); + if base_value.is_some_and(javascript_is_object) + || overlay_value.is_some_and(javascript_is_object) + { + let mut nested = javascript_object_spread(base_value); + nested.extend(javascript_object_spread(overlay_value)); + merged.insert( + key.to_string(), + JavaScriptSpreadValue::Json(Value::Object(finish_javascript_object_spread( + nested, + )?)), + ); + } + } + Ok(Some(Value::Object(finish_javascript_object_spread( + merged, + )?))) +} + +fn javascript_truthy(value: &Value) -> bool { + match value { + Value::Null | Value::Bool(false) => false, + Value::Number(value) => value.as_f64().is_none_or(|value| value != 0.0), + Value::String(value) => !value.is_empty(), + Value::Bool(true) | Value::Array(_) | Value::Object(_) => true, + } +} + +fn javascript_is_object(value: &Value) -> bool { + matches!(value, Value::Array(_) | Value::Object(_)) +} + +fn javascript_property<'a>(value: Option<&'a Value>, key: &str) -> Option<&'a Value> { + value.and_then(Value::as_object)?.get(key) +} + +fn javascript_object_spread(value: Option<&Value>) -> JavaScriptSpreadMap { + match value { + Some(Value::Object(object)) => object + .iter() + .map(|(key, value)| (key.clone(), JavaScriptSpreadValue::Json(value.clone()))) + .collect(), + Some(Value::Array(values)) => values + .iter() + .enumerate() + .map(|(index, value)| { + ( + index.to_string(), + JavaScriptSpreadValue::Json(value.clone()), + ) + }) + .collect(), + Some(Value::String(value)) => value + .encode_utf16() + .enumerate() + .map(|(index, unit)| { + let value = char::from_u32(u32::from(unit)) + .map(|character| { + JavaScriptSpreadValue::Json(Value::String(character.to_string())) + }) + .unwrap_or(JavaScriptSpreadValue::LoneSurrogate); + (index.to_string(), value) + }) + .collect(), + Some(Value::Null | Value::Bool(_) | Value::Number(_)) | None => IndexMap::new(), + } +} + +fn finish_javascript_object_spread( + spread: JavaScriptSpreadMap, +) -> Result, PiCompatMergeError> { + spread + .into_iter() + .map(|(key, value)| match value { + JavaScriptSpreadValue::Json(value) => Ok((key, value)), + JavaScriptSpreadValue::LoneSurrogate => Err(PiCompatMergeError), + }) + .collect() +} + +#[cfg(test)] +mod compat_spread_tests { + use super::*; + use serde_json::json; + + #[test] + fn compat_nested_values_follow_javascript_object_spread() { + let merged = merge_pi_compat( + Some(json!({ + "openRouterRouting": ["first", "second"], + "chatTemplateKwargs": "ab", + "baseOnly": true + })), + Some(json!({ + "openRouterRouting": null, + "chatTemplateKwargs": {"named": true}, + "overlayOnly": true + })), + ) + .expect("representable compat spread") + .expect("truthy overlay produces an object"); + + assert_eq!( + merged, + json!({ + "openRouterRouting": {"0": "first", "1": "second"}, + "chatTemplateKwargs": {"0": "a", "1": "b", "named": true}, + "baseOnly": true, + "overlayOnly": true + }) + ); + } + + #[test] + fn compat_falsy_overlay_returns_base_without_spreading() { + let base = Some(json!({"openRouterRouting": ["kept"]})); + assert_eq!(merge_pi_compat(base.clone(), Some(Value::Null)), Ok(base)); + } + + #[test] + fn compat_spread_rejects_unrepresentable_javascript_surrogates() { + assert_eq!( + merge_pi_compat( + Some(json!({"chatTemplateKwargs": "😀"})), + Some(json!({"chatTemplateKwargs": {"named": true}})), + ), + Err(PiCompatMergeError) + ); + } + + #[test] + fn compat_spread_checks_surrogates_after_later_properties_override_them() { + assert_eq!( + merge_pi_compat( + Some(json!({"chatTemplateKwargs": "😀"})), + Some(json!({ + "chatTemplateKwargs": { + "0": "repaired-high", + "1": "repaired-low", + "named": true + } + })), + ), + Ok(Some(json!({ + "chatTemplateKwargs": { + "0": "repaired-high", + "1": "repaired-low", + "named": true + } + }))) + ); + } +} diff --git a/src-tauri/src/pi_config/model.rs b/src-tauri/src/pi_config/model.rs new file mode 100644 index 000000000..2d0441a78 --- /dev/null +++ b/src-tauri/src/pi_config/model.rs @@ -0,0 +1,1144 @@ +//! Managed Pi control-plane types. +//! +//! API identifiers are intentionally opaque here. The closed set of API +//! families that the gateway can proxy belongs to `gateway.rs`; admitting a +//! new Pi API identifier into the managed control plane must not require a +//! gateway enum update. + +#![allow(dead_code)] + +use super::merge_pi_compat; +use indexmap::IndexMap; +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::collections::{BTreeMap, HashSet}; +use thiserror::Error; +use url::Url; + +pub(crate) type PiHeaderMap = IndexMap; +pub(crate) type PiThinkingLevelMap = BTreeMap; + +/// Pi uses JavaScript/TypeBox `Number`, not `Integer`, for model limits. +#[derive(Debug, Clone, Copy, PartialEq, PartialOrd, Serialize)] +#[serde(transparent)] +pub(crate) struct PiNumber(f64); + +impl PiNumber { + pub(crate) const DEFAULT_CONTEXT_WINDOW: Self = Self(128_000.0); + pub(crate) const DEFAULT_MAX_TOKENS: Self = Self(16_384.0); + + pub(crate) fn new(value: f64) -> Result { + if value.is_finite() { + Ok(Self(value)) + } else { + Err(PiNumberError) + } + } + + pub(crate) const fn get(self) -> f64 { + self.0 + } +} + +impl<'de> Deserialize<'de> for PiNumber { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + let value = f64::deserialize(deserializer)?; + Self::new(value).map_err(serde::de::Error::custom) + } +} + +impl TryFrom for PiNumber { + type Error = PiNumberError; + + fn try_from(value: f64) -> Result { + Self::new(value) + } +} + +impl From for f64 { + fn from(value: PiNumber) -> Self { + value.get() + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)] +#[error("Pi Number must be finite")] +pub(crate) struct PiNumberError; + +/// Opaque managed API identifier. +/// +/// This is not the gateway support enum. Any non-empty Pi-valid identifier +/// round-trips through this type, including identifiers introduced after the +/// pinned gateway implementation. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize)] +#[serde(transparent)] +pub(crate) struct PiManagedApiId(String); + +impl PiManagedApiId { + pub(crate) fn new(value: impl Into) -> Result { + let value = value.into(); + if value.is_empty() { + return Err(PiConfigError::EmptyApiId); + } + Ok(Self(value)) + } + + pub(crate) fn as_str(&self) -> &str { + &self.0 + } +} + +impl<'de> Deserialize<'de> for PiManagedApiId { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + let value = String::deserialize(deserializer)?; + Self::new(value).map_err(serde::de::Error::custom) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub(crate) enum PiModelInput { + Text, + Image, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiModelCostRates { + pub input: f64, + pub output: f64, + pub cache_read: f64, + pub cache_write: f64, +} + +impl Default for PiModelCostRates { + fn default() -> Self { + Self { + input: 0.0, + output: 0.0, + cache_read: 0.0, + cache_write: 0.0, + } + } +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiModelCostTier { + pub input_tokens_above: f64, + pub input: f64, + pub output: f64, + pub cache_read: f64, + pub cache_write: f64, + #[serde(flatten)] + pub extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiModelCost { + #[serde(flatten)] + pub rates: PiModelCostRates, + #[serde(skip_serializing_if = "Option::is_none")] + pub tiers: Option>, + #[serde(flatten)] + pub extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiModelCostOverride { + #[serde(skip_serializing_if = "Option::is_none")] + pub input: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub output: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_write: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tiers: Option>, + #[serde(flatten)] + pub extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiManagedModelOverride { + #[serde(skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub reasoning: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub thinking_level_map: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub input: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub cost: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub context_window: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(default, skip_serializing_if = "IndexMap::is_empty")] + pub headers: PiHeaderMap, + #[serde(skip_serializing_if = "Option::is_none")] + pub compat: Option, + #[serde(flatten)] + pub extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiManagedModel { + pub id: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub base_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub api: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub reasoning: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub thinking_level_map: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub input: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub cost: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub context_window: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(default, skip_serializing_if = "IndexMap::is_empty")] + pub headers: PiHeaderMap, + #[serde(skip_serializing_if = "Option::is_none")] + pub compat: Option, + #[serde(flatten)] + pub extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiManagedProviderConfig { + #[serde(skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub base_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub api: Option, + /// Literal or deferred transport material. The managed layer never + /// executes env/command/file/network resolution. + #[serde(skip_serializing_if = "Option::is_none")] + pub api_key: Option, + #[serde(default, skip_serializing_if = "IndexMap::is_empty")] + pub headers: PiHeaderMap, + #[serde(skip_serializing_if = "Option::is_none")] + pub auth_header: Option, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub models: Vec, + #[serde(default, skip_serializing_if = "BTreeMap::is_empty")] + pub model_overrides: BTreeMap, + #[serde(skip_serializing_if = "Option::is_none")] + pub compat: Option, + #[serde(flatten)] + pub extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiEffectiveModel { + pub id: String, + pub name: String, + pub api: PiManagedApiId, + pub base_url: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub api_key: Option, + pub auth_header: bool, + pub reasoning: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub thinking_level_map: Option, + pub input: Vec, + pub cost: PiModelCost, + pub context_window: PiNumber, + pub max_tokens: PiNumber, + pub headers: PiHeaderMap, + /// Provider auth headers and the later model overlay are kept distinct so + /// a gateway can reproduce pinned Pi's case-insensitive runtime merge + /// without changing the serialized effective projection. + #[serde(skip)] + pub header_layers: PiEffectiveHeaderLayers, + #[serde(skip_serializing_if = "Option::is_none")] + pub compat: Option, + pub provider_extra: BTreeMap, + pub model_extra: BTreeMap, + pub override_extra: BTreeMap, +} + +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub(crate) struct PiEffectiveHeaderLayers { + pub provider: PiHeaderMap, + pub model: PiHeaderMap, +} + +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub(crate) enum PiConfigError { + #[error("Pi provider must declare at least one managed model")] + ProviderHasNoModels, + #[error("Pi API id cannot be empty")] + EmptyApiId, + #[error("Pi model id cannot be empty")] + EmptyModelId, + #[error("Pi provider contains duplicate model id '{0}'")] + DuplicateModelId(String), + #[error("Pi model '{0}' does not exist in this provider")] + ModelNotFound(String), + #[error("Pi model '{model_id}' has no effective API id")] + MissingEffectiveApi { model_id: String }, + #[error("Pi model '{model_id}' has no effective endpoint")] + MissingEffectiveEndpoint { model_id: String }, + #[error("Pi endpoint at '{json_pointer}' must be an absolute HTTP(S) URL: {reason}")] + InvalidEndpoint { + json_pointer: String, + reason: String, + }, + #[error("Pi model override '{0}' does not match an explicit model id")] + UnknownModelOverride(String), + #[error("Pi compat at '{json_pointer}' must be an object")] + InvalidCompat { json_pointer: String }, + #[error( + "Pi compat at '{json_pointer}' requires JavaScript UTF-16 values that cannot be represented" + )] + UnrepresentableCompat { json_pointer: String }, + #[error("Pi {field} cannot be empty when present")] + EmptyOptionalField { field: &'static str }, + #[error("Pi model '{model_id}' {field} must be greater than zero")] + NonPositiveModelLimit { + model_id: String, + field: &'static str, + }, + #[error("Pi model '{model_id}' contains an invalid value for thinking level '{level}'")] + InvalidThinkingLevelValue { model_id: String, level: String }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiNativeEntryKind { + BuiltInOverlay, + CustomCatalog, + ExtensionOverlay, + UnknownShape, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiRawNativeValidity { + Valid, + Invalid, + Unknown, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiManagedAssessment { + Manageable, + Unsupported, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiCompositionStatus { + Composed, + Failed, + Unknown, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "status", rename_all = "snake_case")] +pub(crate) enum PiManagementStatus { + Importable, + Managed { + #[serde(rename = "providerId")] + provider_id: String, + }, + Unsupported, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiGatewayStatus { + Proxyable, + DirectOnly, + Unknown, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiDiagnosticLayer { + RawSchema, + Managed, + Composition, + Gateway, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiReasonCode { + RawSchemaMismatch, + RawSchemaUnsupportedOperator, + RawSchemaPinDrift, + RawSchemaAmbiguous, + CatalogRequired, + ModelOverridesOnly, + MissingExplicitModels, + ManagedTypeConversionFailed, + EmptyOptionalField, + EmptyModelId, + DuplicateModelId, + UnknownModelOverride, + MissingEffectiveApi, + MissingEffectiveEndpoint, + InvalidEndpoint, + InvalidCompat, + UnrepresentableCompat, + NonPositiveModelLimit, + InvalidThinkingLevel, + CompositionFailed, + GatewayCredentialUnavailable, + UnsupportedCredentialKind, + UnsupportedGatewayFamily, + InvalidHeaderName, + InvalidHeaderValue, + ProtectedHeader, + DeferredValueUnavailable, +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiDiagnosticReason { + pub layer: PiDiagnosticLayer, + pub code: PiReasonCode, + #[serde(skip_serializing_if = "Option::is_none")] + pub json_pointer: Option, +} + +impl PiDiagnosticReason { + pub(crate) fn new( + layer: PiDiagnosticLayer, + code: PiReasonCode, + json_pointer: Option, + ) -> Self { + Self { + layer, + code, + json_pointer, + } + } +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiNativeDiagnostic { + pub provider_key: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub display_name: Option, + pub fingerprint: String, + pub kind: PiNativeEntryKind, + pub raw_validity: PiRawNativeValidity, + pub managed_assessment: PiManagedAssessment, + pub composition_status: PiCompositionStatus, + pub management_status: PiManagementStatus, + pub gateway_status: PiGatewayStatus, + pub reasons: Vec, +} + +const THINKING_LEVELS: [&str; 7] = ["off", "minimal", "low", "medium", "high", "xhigh", "max"]; + +pub(crate) fn validate_pi_managed_provider( + provider: &PiManagedProviderConfig, +) -> Result<(), PiConfigError> { + if provider.models.is_empty() { + return Err(PiConfigError::ProviderHasNoModels); + } + validate_optional_text(provider.name.as_deref(), "provider name")?; + validate_optional_text(provider.api_key.as_deref(), "provider apiKey")?; + validate_present_endpoint(provider.base_url.as_deref(), "/baseUrl")?; + validate_compat(provider.compat.as_ref(), "/compat")?; + + let mut model_ids = HashSet::with_capacity(provider.models.len()); + for (index, model) in provider.models.iter().enumerate() { + if model.id.is_empty() { + return Err(PiConfigError::EmptyModelId); + } + if !model_ids.insert(model.id.as_str()) { + return Err(PiConfigError::DuplicateModelId(model.id.clone())); + } + validate_optional_text(model.name.as_deref(), "model name")?; + validate_present_endpoint( + model.base_url.as_deref(), + &format!("/models/{index}/baseUrl"), + )?; + validate_compat(model.compat.as_ref(), &format!("/models/{index}/compat"))?; + validate_model_limit(model, model.context_window, "contextWindow")?; + validate_model_limit(model, model.max_tokens, "maxTokens")?; + validate_thinking_levels(&model.id, model.thinking_level_map.as_ref())?; + let _ = effective_pi_model_unchecked(provider, model)?; + } + + for (model_id, model_override) in &provider.model_overrides { + if !model_ids.contains(model_id.as_str()) { + return Err(PiConfigError::UnknownModelOverride(model_id.clone())); + } + validate_optional_text(model_override.name.as_deref(), "model override name")?; + validate_optional_limit(model_id, model_override.context_window, "contextWindow")?; + validate_optional_limit(model_id, model_override.max_tokens, "maxTokens")?; + validate_thinking_levels(model_id, model_override.thinking_level_map.as_ref())?; + validate_compat( + model_override.compat.as_ref(), + &format!("/modelOverrides/{}/compat", escape_json_pointer(model_id)), + )?; + } + Ok(()) +} + +pub(crate) fn effective_pi_model( + provider: &PiManagedProviderConfig, + model_id: &str, +) -> Result { + let model = provider + .models + .iter() + .find(|model| model.id == model_id) + .ok_or_else(|| PiConfigError::ModelNotFound(model_id.to_string()))?; + validate_pi_managed_provider(provider)?; + effective_pi_model_unchecked(provider, model) +} + +fn effective_pi_model_unchecked( + provider: &PiManagedProviderConfig, + model: &PiManagedModel, +) -> Result { + let api = model + .api + .clone() + .or_else(|| provider.api.clone()) + .ok_or_else(|| PiConfigError::MissingEffectiveApi { + model_id: model.id.clone(), + })?; + let base_url = model + .base_url + .as_ref() + .or(provider.base_url.as_ref()) + .ok_or_else(|| PiConfigError::MissingEffectiveEndpoint { + model_id: model.id.clone(), + })?; + let model_override = provider.model_overrides.get(&model.id); + let mut model_headers = PiHeaderMap::new(); + if let Some(model_override) = model_override { + model_headers.extend(model_override.headers.clone()); + } + model_headers.extend(model.headers.clone()); + let mut headers = provider.headers.clone(); + headers.extend(model_headers.clone()); + + let thinking_level_map = merge_thinking_level_maps( + model.thinking_level_map.as_ref(), + model_override.and_then(|entry| entry.thinking_level_map.as_ref()), + ); + + let base_cost = model.cost.clone().unwrap_or_default(); + let cost = model_override + .and_then(|entry| entry.cost.as_ref()) + .map(|entry| apply_cost_override(base_cost.clone(), entry)) + .unwrap_or(base_cost); + + let compat = merge_pi_compat(provider.compat.clone(), model.compat.clone()).map_err(|_| { + PiConfigError::UnrepresentableCompat { + json_pointer: "/compat".to_string(), + } + })?; + let compat = merge_pi_compat( + compat, + model_override.and_then(|entry| entry.compat.clone()), + ) + .map_err(|_| PiConfigError::UnrepresentableCompat { + json_pointer: format!("/modelOverrides/{}/compat", escape_json_pointer(&model.id)), + })?; + + Ok(PiEffectiveModel { + id: model.id.clone(), + name: model_override + .and_then(|entry| entry.name.clone()) + .or_else(|| model.name.clone()) + .unwrap_or_else(|| model.id.clone()), + api, + base_url: base_url.clone(), + api_key: provider.api_key.clone(), + auth_header: provider.auth_header.unwrap_or(false), + reasoning: model_override + .and_then(|entry| entry.reasoning) + .or(model.reasoning) + .unwrap_or(false), + thinking_level_map, + input: model_override + .and_then(|entry| entry.input.clone()) + .or_else(|| model.input.clone()) + .unwrap_or_else(|| vec![PiModelInput::Text]), + cost, + context_window: model_override + .and_then(|entry| entry.context_window) + .or(model.context_window) + .unwrap_or(PiNumber::DEFAULT_CONTEXT_WINDOW), + max_tokens: model_override + .and_then(|entry| entry.max_tokens) + .or(model.max_tokens) + .unwrap_or(PiNumber::DEFAULT_MAX_TOKENS), + headers, + header_layers: PiEffectiveHeaderLayers { + provider: provider.headers.clone(), + model: model_headers, + }, + compat, + provider_extra: provider.extra.clone(), + model_extra: model.extra.clone(), + override_extra: model_override + .map(|entry| entry.extra.clone()) + .unwrap_or_default(), + }) +} + +fn merge_thinking_level_maps( + base: Option<&PiThinkingLevelMap>, + overlay: Option<&PiThinkingLevelMap>, +) -> Option { + match (base, overlay) { + (None, None) => None, + (Some(value), None) | (None, Some(value)) => Some(value.clone()), + (Some(base), Some(overlay)) => { + let mut merged = base.clone(); + merged.extend(overlay.clone()); + Some(merged) + } + } +} + +fn apply_cost_override(mut base: PiModelCost, model_override: &PiModelCostOverride) -> PiModelCost { + base.rates = PiModelCostRates { + input: model_override.input.unwrap_or(base.rates.input), + output: model_override.output.unwrap_or(base.rates.output), + cache_read: model_override.cache_read.unwrap_or(base.rates.cache_read), + cache_write: model_override.cache_write.unwrap_or(base.rates.cache_write), + }; + if let Some(tiers) = &model_override.tiers { + base.tiers = Some(tiers.clone()); + } + base.extra.extend(model_override.extra.clone()); + base +} + +fn validate_optional_text(value: Option<&str>, field: &'static str) -> Result<(), PiConfigError> { + if value.is_some_and(str::is_empty) { + return Err(PiConfigError::EmptyOptionalField { field }); + } + Ok(()) +} + +fn validate_model_limit( + model: &PiManagedModel, + value: Option, + field: &'static str, +) -> Result<(), PiConfigError> { + validate_optional_limit(&model.id, value, field) +} + +fn validate_optional_limit( + model_id: &str, + value: Option, + field: &'static str, +) -> Result<(), PiConfigError> { + if value.is_some_and(|value| value.get() <= 0.0) { + return Err(PiConfigError::NonPositiveModelLimit { + model_id: model_id.to_string(), + field, + }); + } + Ok(()) +} + +fn validate_present_endpoint( + endpoint: Option<&str>, + json_pointer: &str, +) -> Result<(), PiConfigError> { + let Some(endpoint) = endpoint else { + return Ok(()); + }; + if endpoint.trim().is_empty() { + return Err(PiConfigError::InvalidEndpoint { + json_pointer: json_pointer.to_string(), + reason: "endpoint cannot be empty".to_string(), + }); + } + let parsed = Url::parse(endpoint).map_err(|error| PiConfigError::InvalidEndpoint { + json_pointer: json_pointer.to_string(), + reason: error.to_string(), + })?; + if !matches!(parsed.scheme(), "http" | "https") || parsed.host().is_none() { + return Err(PiConfigError::InvalidEndpoint { + json_pointer: json_pointer.to_string(), + reason: "endpoint must be an absolute HTTP(S) URL".to_string(), + }); + } + Ok(()) +} + +fn validate_compat(compat: Option<&Value>, json_pointer: &str) -> Result<(), PiConfigError> { + if compat.is_some_and(|value| !value.is_object()) { + return Err(PiConfigError::InvalidCompat { + json_pointer: json_pointer.to_string(), + }); + } + Ok(()) +} + +fn escape_json_pointer(value: &str) -> String { + value.replace('~', "~0").replace('/', "~1") +} + +fn validate_thinking_levels( + model_id: &str, + levels: Option<&PiThinkingLevelMap>, +) -> Result<(), PiConfigError> { + let Some(levels) = levels else { + return Ok(()); + }; + for (level, value) in levels { + if THINKING_LEVELS.contains(&level.as_str()) && !(value.is_string() || value.is_null()) { + return Err(PiConfigError::InvalidThinkingLevelValue { + model_id: model_id.to_string(), + level: level.clone(), + }); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn api(value: &str) -> PiManagedApiId { + PiManagedApiId::new(value).expect("non-empty API id") + } + + fn number(value: f64) -> PiNumber { + PiNumber::new(value).expect("finite Pi Number") + } + + fn model(id: &str) -> PiManagedModel { + PiManagedModel { + id: id.to_string(), + name: None, + base_url: None, + api: None, + reasoning: None, + thinking_level_map: None, + input: None, + cost: None, + context_window: None, + max_tokens: None, + headers: PiHeaderMap::new(), + compat: None, + extra: BTreeMap::new(), + } + } + + fn provider(models: Vec) -> PiManagedProviderConfig { + PiManagedProviderConfig { + name: Some("Example".into()), + base_url: Some("https://example.com/api".into()), + api: Some(api("anthropic-messages")), + api_key: Some("$PI_KEY".into()), + headers: PiHeaderMap::new(), + auth_header: None, + models, + model_overrides: BTreeMap::new(), + compat: None, + extra: BTreeMap::new(), + } + } + + #[test] + fn opaque_api_ids_round_trip_without_gateway_narrowing() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "baseUrl": "https://future.example/v9", + "api": "future-wire-v9", + "models": [{"id": "future"}] + })) + .expect("opaque API id"); + validate_pi_managed_provider(&config).expect("future API is manageable"); + let effective = effective_pi_model(&config, "future").expect("effective future model"); + assert_eq!(effective.api.as_str(), "future-wire-v9"); + assert_eq!( + serde_json::to_value(&config).expect("serialize")["api"], + "future-wire-v9" + ); + } + + #[test] + fn provider_and_model_level_inheritance_is_exactly_two_levels() { + let inherited = effective_pi_model(&provider(vec![model("claude")]), "claude") + .expect("provider defaults"); + assert_eq!(inherited.api.as_str(), "anthropic-messages"); + assert_eq!(inherited.base_url, "https://example.com/api"); + + let mut self_contained = model("future"); + self_contained.api = Some(api("future-api")); + self_contained.base_url = Some("https://model.example/v2".into()); + let mut config = provider(vec![self_contained]); + config.api = None; + config.base_url = None; + let effective = effective_pi_model(&config, "future").expect("model defaults"); + assert_eq!(effective.api.as_str(), "future-api"); + assert_eq!(effective.base_url, "https://model.example/v2"); + } + + #[test] + fn schema_valid_whitespace_strings_remain_manageable() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "name": " ", + "api": "anthropic-messages", + "baseUrl": "https://example.com", + "apiKey": " ", + "models": [{"id": " ", "name": " "}] + })) + .expect("pinned schema-valid whitespace fields"); + validate_pi_managed_provider(&config) + .expect("managed validation must use pinned minLength semantics without trimming"); + assert_eq!( + effective_pi_model(&config, " ") + .expect("whitespace model id") + .id, + " " + ); + } + + #[test] + fn override_precedence_and_nested_compat_are_stable() { + let mut base_model = model("m"); + base_model.reasoning = Some(false); + base_model.headers = PiHeaderMap::from([ + ("layer".into(), "model".into()), + ("model".into(), "yes".into()), + ]); + base_model.compat = Some(json!({ + "supportsStore": true, + "openRouterRouting": {"only": ["model"], "zdr": true} + })); + let mut config = provider(vec![base_model]); + config.headers = PiHeaderMap::from([ + ("layer".into(), "provider".into()), + ("provider".into(), "yes".into()), + ]); + config.compat = Some(json!({"supportsDeveloperRole": true})); + config.model_overrides.insert( + "m".into(), + PiManagedModelOverride { + reasoning: Some(true), + headers: PiHeaderMap::from([ + ("layer".into(), "override".into()), + ("override".into(), "yes".into()), + ]), + compat: Some(json!({ + "supportsStore": false, + "openRouterRouting": {"order": ["override"]} + })), + ..Default::default() + }, + ); + + let effective = effective_pi_model(&config, "m").expect("effective"); + assert!(effective.reasoning); + assert_eq!(effective.headers["layer"], "model"); + assert_eq!(effective.header_layers.provider["layer"], "provider"); + assert_eq!(effective.header_layers.model["layer"], "model"); + assert_eq!( + effective.compat, + Some(json!({ + "supportsDeveloperRole": true, + "supportsStore": false, + "openRouterRouting": { + "only": ["model"], + "zdr": true, + "order": ["override"] + } + })) + ); + } + + #[test] + fn effective_compat_uses_the_pinned_composer_spread_result() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "api": "openai-responses", + "baseUrl": "https://compat.example/v1", + "compat": { + "openRouterRouting": ["first", "second"], + "chatTemplateKwargs": "ab", + "baseOnly": true + }, + "models": [{ + "id": "m", + "compat": {"supportsStore": true} + }], + "modelOverrides": { + "m": { + "compat": { + "openRouterRouting": null, + "chatTemplateKwargs": {"named": true}, + "overlayOnly": true + } + } + } + })) + .expect("deserialize pinned compat vector"); + + assert_eq!( + effective_pi_model(&config, "m") + .expect("effective compat") + .compat, + Some(json!({ + "openRouterRouting": {"0": "first", "1": "second"}, + "chatTemplateKwargs": {"0": "a", "1": "b", "named": true}, + "baseOnly": true, + "supportsStore": true, + "overlayOnly": true + })) + ); + } + + #[test] + fn effective_compat_rejects_unrepresentable_javascript_surrogate_spread() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "api": "openai-responses", + "baseUrl": "https://compat.example/v1", + "compat": {"chatTemplateKwargs": "😀"}, + "models": [{"id": "m"}], + "modelOverrides": { + "m": {"compat": {"chatTemplateKwargs": {"named": true}}} + } + })) + .expect("deserialize pinned compat vector"); + + assert_eq!( + effective_pi_model(&config, "m"), + Err(PiConfigError::UnrepresentableCompat { + json_pointer: "/modelOverrides/m/compat".to_string(), + }) + ); + } + + #[test] + fn fractional_pi_numbers_survive_managed_round_trip() { + let mut fractional = model("fractional"); + fractional.context_window = Some(number(128000.5)); + fractional.max_tokens = Some(number(16384.25)); + let config = provider(vec![fractional]); + let encoded = serde_json::to_value(&config).expect("serialize"); + let decoded: PiManagedProviderConfig = + serde_json::from_value(encoded).expect("deserialize"); + let effective = effective_pi_model(&decoded, "fractional").expect("effective"); + assert_eq!(effective.context_window.get(), 128000.5); + assert_eq!(effective.max_tokens.get(), 16384.25); + } + + #[test] + fn managed_validation_rejects_duplicates_but_preserves_schema_valid_thinking_maps() { + let duplicate = provider(vec![model("same"), model("same")]); + assert_eq!( + validate_pi_managed_provider(&duplicate), + Err(PiConfigError::DuplicateModelId("same".into())) + ); + + let mut future_thinking = model("thinking"); + future_thinking.thinking_level_map = Some(BTreeMap::from([ + ("high".into(), json!("native-high")), + ("future".into(), json!({"opaque": true})), + ])); + let future_config = provider(vec![future_thinking]); + validate_pi_managed_provider(&future_config) + .expect("pinned-schema additional keys remain manageable"); + assert_eq!( + serde_json::to_value(&future_config) + .expect("serialize") + .pointer("/models/0/thinkingLevelMap/future"), + Some(&json!({"opaque": true})) + ); + + let mut invalid_known = model("invalid-known"); + invalid_known.thinking_level_map = Some(BTreeMap::from([("low".into(), json!(2))])); + assert_eq!( + validate_pi_managed_provider(&provider(vec![invalid_known])), + Err(PiConfigError::InvalidThinkingLevelValue { + model_id: "invalid-known".into(), + level: "low".into() + }) + ); + } + + #[test] + fn future_cost_members_survive_model_override_and_effective_projection() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "api": "anthropic-messages", + "baseUrl": "https://cost.example", + "models": [{ + "id": "m", + "cost": { + "input": 1.0, + "output": 2.0, + "cacheRead": 0.5, + "cacheWrite": 0.25, + "futureRate": {"opaque": true}, + "tiers": [{ + "inputTokensAbove": 100.0, + "input": 1.0, + "output": 2.0, + "cacheRead": 0.5, + "cacheWrite": 0.25, + "futureTierField": ["preserved"] + }] + } + }], + "modelOverrides": { + "m": { + "cost": { + "output": 3.0, + "futureOverrideRate": "preserved" + } + } + } + })) + .expect("deserialize future cost members"); + validate_pi_managed_provider(&config).expect("future cost members are manageable"); + + let round_trip = serde_json::to_value(&config).expect("serialize managed config"); + assert_eq!( + round_trip.pointer("/models/0/cost/futureRate"), + Some(&json!({"opaque": true})) + ); + assert_eq!( + round_trip.pointer("/models/0/cost/tiers/0/futureTierField"), + Some(&json!(["preserved"])) + ); + let effective = + serde_json::to_value(effective_pi_model(&config, "m").expect("effective model")) + .expect("serialize effective model"); + assert_eq!(effective.pointer("/cost/output"), Some(&json!(3.0))); + assert_eq!( + effective.pointer("/cost/futureRate"), + Some(&json!({"opaque": true})) + ); + assert_eq!( + effective.pointer("/cost/futureOverrideRate"), + Some(&json!("preserved")) + ); + assert_eq!( + effective.pointer("/cost/tiers/0/futureTierField"), + Some(&json!(["preserved"])) + ); + } + + #[test] + fn empty_cost_tiers_remain_distinct_from_absent_at_managed_and_effective_boundaries() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "api": "anthropic-messages", + "baseUrl": "https://cost.example", + "models": [ + { + "id": "empty", + "cost": { + "input": 1.0, + "output": 2.0, + "cacheRead": 0.5, + "cacheWrite": 0.25, + "tiers": [] + } + }, + { + "id": "absent", + "cost": { + "input": 1.0, + "output": 2.0, + "cacheRead": 0.5, + "cacheWrite": 0.25 + } + } + ] + })) + .expect("deserialize explicit and absent tiers"); + + let managed = serde_json::to_value(&config).expect("serialize managed config"); + assert_eq!(managed.pointer("/models/0/cost/tiers"), Some(&json!([]))); + assert_eq!(managed.pointer("/models/1/cost/tiers"), None); + + let empty_effective = serde_json::to_value( + effective_pi_model(&config, "empty").expect("effective empty tiers"), + ) + .expect("serialize effective empty tiers"); + let absent_effective = serde_json::to_value( + effective_pi_model(&config, "absent").expect("effective absent tiers"), + ) + .expect("serialize effective absent tiers"); + assert_eq!(empty_effective.pointer("/cost/tiers"), Some(&json!([]))); + assert_eq!(absent_effective.pointer("/cost/tiers"), None); + } + + #[test] + fn diagnostic_reason_serialization_is_structured() { + let reason = PiDiagnosticReason::new( + PiDiagnosticLayer::Gateway, + PiReasonCode::UnsupportedGatewayFamily, + Some("/models/0/api".into()), + ); + assert_eq!( + serde_json::to_value(reason).expect("serialize"), + json!({ + "layer": "gateway", + "code": "unsupported_gateway_family", + "jsonPointer": "/models/0/api" + }) + ); + + let credential_reason = PiDiagnosticReason::new( + PiDiagnosticLayer::Gateway, + PiReasonCode::UnsupportedCredentialKind, + Some("/apiKey".into()), + ); + assert_eq!( + serde_json::to_value(credential_reason).expect("serialize"), + json!({ + "layer": "gateway", + "code": "unsupported_credential_kind", + "jsonPointer": "/apiKey" + }) + ); + + let compat_reason = PiDiagnosticReason::new( + PiDiagnosticLayer::Composition, + PiReasonCode::UnrepresentableCompat, + Some("/modelOverrides/m/compat".into()), + ); + assert_eq!( + serde_json::to_value(compat_reason).expect("serialize"), + json!({ + "layer": "composition", + "code": "unrepresentable_compat", + "jsonPointer": "/modelOverrides/m/compat" + }) + ); + } +} diff --git a/src-tauri/src/pi_config/native.rs b/src-tauri/src/pi_config/native.rs new file mode 100644 index 000000000..7d0df8d34 --- /dev/null +++ b/src-tauri/src/pi_config/native.rs @@ -0,0 +1,1107 @@ +//! Public, side-effect-free inspection service for Pi's native catalog. +//! +//! This module orchestrates independent raw, managed, composer and gateway +//! assessments. No assessment is allowed to gate execution of a sibling layer. + +#![allow(dead_code)] + +use super::composer::{ + compose_explicit_custom_catalog, PiComposerReasonCode, PiComposerStatus, PiNativeComposition, +}; +use super::document::{pi_raw_provider_fingerprint, read_pi_models_document, PiRawProviderEntry}; +use super::gateway::{ + assess_composition, PiGatewayAssessment, PiGatewayCapability, PiGatewayReasonCode, +}; +use super::model::{ + validate_pi_managed_provider, PiCompositionStatus, PiConfigError, PiDiagnosticLayer, + PiDiagnosticReason, PiGatewayStatus, PiManagedAssessment, PiManagedProviderConfig, + PiManagementStatus, PiNativeDiagnostic, PiNativeEntryKind, PiRawNativeValidity, PiReasonCode, +}; +use super::raw_schema::{ + evaluate_provider_value, PiRawReasonCode, PiRawSchemaEvaluation, PiRawValidity, +}; +use crate::config::get_home_dir; +use crate::error::AppError; +use serde_json::Value; +use std::collections::{BTreeMap, HashSet}; +use std::path::{Path, PathBuf}; +use url::Url; + +/// Built-in provider IDs at Pi commit +/// `ab366ebe94cacd419d986be454f12b1b9913aaca`. +const PI_BUILTIN_PROVIDER_KEYS: &[&str] = &[ + "amazon-bedrock", + "ant-ling", + "anthropic", + "azure-openai-responses", + "cerebras", + "cloudflare-ai-gateway", + "cloudflare-workers-ai", + "deepseek", + "fireworks", + "github-copilot", + "google", + "google-vertex", + "groq", + "huggingface", + "kimi-coding", + "minimax", + "minimax-cn", + "mistral", + "moonshotai", + "moonshotai-cn", + "nvidia", + "openai", + "openai-codex", + "opencode", + "opencode-go", + "openrouter", + "qwen-token-plan", + "qwen-token-plan-cn", + "radius", + "together", + "vercel-ai-gateway", + "xai", + "xiaomi", + "xiaomi-token-plan-ams", + "xiaomi-token-plan-cn", + "xiaomi-token-plan-sgp", + "zai", + "zai-coding-cn", +]; + +const RECOGNIZED_PROVIDER_FIELDS: &[&str] = &[ + "name", + "baseUrl", + "apiKey", + "api", + "oauth", + "headers", + "compat", + "authHeader", + "models", + "modelOverrides", +]; + +#[derive(Debug, Clone)] +pub(crate) struct PiNativeEntryInspection { + pub diagnostic: PiNativeDiagnostic, + pub managed_config: Option, + pub composition: PiNativeComposition, +} + +#[derive(Debug)] +struct ManagedResult { + assessment: PiManagedAssessment, + config: Option, + reasons: Vec, +} + +/// The public read-only service entry used by commands and certification tests. +pub(crate) struct PiNativeInspectionService; + +impl PiNativeInspectionService { + pub(crate) fn inspect_current( + managed_claims: &BTreeMap, + ) -> Result, AppError> { + Self::inspect_catalog(&get_pi_models_path()?, managed_claims) + } + + pub(crate) fn inspect_catalog( + path: &Path, + managed_claims: &BTreeMap, + ) -> Result, AppError> { + let document = read_pi_models_document(path)?; + Ok(document + .providers() + .iter() + .map(|(provider_key, entry)| { + analyze_native_entry(provider_key, entry, managed_claims).diagnostic + }) + .collect()) + } + + pub(crate) fn inspect_entry( + path: &Path, + provider_key: &str, + managed_claims: &BTreeMap, + ) -> Result, AppError> { + let document = read_pi_models_document(path)?; + Ok(document + .providers() + .get(provider_key) + .map(|entry| analyze_native_entry(provider_key, entry, managed_claims))) + } +} + +pub(crate) fn inspect_current_pi_native_catalog( + managed_claims: &BTreeMap, +) -> Result, AppError> { + PiNativeInspectionService::inspect_current(managed_claims) +} + +pub(crate) fn inspect_pi_native_catalog( + path: &Path, + managed_claims: &BTreeMap, +) -> Result, AppError> { + PiNativeInspectionService::inspect_catalog(path, managed_claims) +} + +pub(crate) fn inspect_pi_native_entry( + path: &Path, + provider_key: &str, + managed_claims: &BTreeMap, +) -> Result, AppError> { + PiNativeInspectionService::inspect_entry(path, provider_key, managed_claims) +} + +/// Compose a database-authoritative managed provider through the same raw and +/// composer layers used by native inspection. Runtime construction must not +/// reimplement inheritance or field semantics. +pub(crate) fn compose_managed_pi_provider( + provider_key: &str, + config: &PiManagedProviderConfig, +) -> Result { + validate_pi_managed_provider(config) + .map_err(|error| AppError::InvalidInput(error.to_string()))?; + let value = + serde_json::to_value(config).map_err(|source| AppError::JsonSerialize { source })?; + let raw = evaluate_provider_value(&value); + let provider = raw.valid_provider.as_ref().ok_or_else(|| { + AppError::Config(format!( + "managed Pi provider '{provider_key}' did not pass the pinned raw schema" + )) + })?; + Ok(compose_explicit_custom_catalog(provider_key, provider)) +} + +fn normalize_pi_agent_dir(value: &str, home: &Path) -> Result { + if value == "~" { + return Ok(home.to_path_buf()); + } + if let Some(suffix) = value.strip_prefix("~/") { + return Ok(home.join(suffix)); + } + #[cfg(windows)] + if let Some(suffix) = value.strip_prefix("~\\") { + return Ok(home.join(suffix)); + } + if value.starts_with("file://") { + let url = Url::parse(value).map_err(|error| { + AppError::Config(format!("invalid Pi agent directory URL: {error}")) + })?; + return url.to_file_path().map_err(|_| { + AppError::Config(format!( + "Pi agent directory URL is not a local file path: {value}" + )) + }); + } + Ok(PathBuf::from(value)) +} + +pub(crate) fn get_pi_agent_dir_for_override( + override_dir: Option<&str>, +) -> Result { + if let Some(override_dir) = override_dir + .map(str::trim) + .filter(|value| !value.is_empty()) + { + return Ok(crate::settings::resolve_override_path(override_dir)); + } + let Some(raw) = std::env::var_os("PI_CODING_AGENT_DIR") else { + return Ok(get_home_dir().join(".pi").join("agent")); + }; + if raw.is_empty() { + return Ok(get_home_dir().join(".pi").join("agent")); + } + normalize_pi_agent_dir(&raw.to_string_lossy(), &get_home_dir()) +} + +pub(crate) fn get_pi_agent_dir() -> Result { + if let Some(override_dir) = crate::settings::get_pi_override_dir() { + return Ok(override_dir); + } + get_pi_agent_dir_for_override(None) +} + +pub(crate) fn get_pi_models_path_for_override( + override_dir: Option<&str>, +) -> Result { + Ok(get_pi_agent_dir_for_override(override_dir)?.join("models.json")) +} + +pub(crate) fn get_pi_models_path() -> Result { + Ok(get_pi_agent_dir()?.join("models.json")) +} + +fn analyze_native_entry( + provider_key: &str, + entry: &PiRawProviderEntry, + managed_claims: &BTreeMap, +) -> PiNativeEntryInspection { + // Raw, composition and managed conversion are deliberately invoked from + // the same immutable JSON value. Neither result controls whether a sibling + // assessment is attempted. + let raw = evaluate_provider_value(&entry.value); + let raw_validity = map_raw_validity(raw.validity); + let kind = classify_kind(provider_key, &entry.value, raw.validity); + let composition = match (raw.valid_provider.as_ref(), kind) { + (Some(provider), PiNativeEntryKind::CustomCatalog) => { + compose_explicit_custom_catalog(provider_key, provider) + } + (Some(_), _) => PiNativeComposition::catalog_required("/models"), + (None, _) => PiNativeComposition::unavailable_without_valid_raw(), + }; + let managed = assess_managed(raw.validity, kind, &entry.value); + let gateway = if raw.validity == PiRawValidity::Valid { + assess_composition(&composition) + } else { + PiGatewayAssessment { + capability: PiGatewayCapability::Unknown, + reasons: Vec::new(), + plans: Vec::new(), + } + }; + + let management_status = managed_claims + .get(provider_key) + .map(|provider_id| PiManagementStatus::Managed { + provider_id: provider_id.clone(), + }) + .unwrap_or_else(|| { + if raw.validity == PiRawValidity::Valid + && managed.assessment == PiManagedAssessment::Manageable + { + PiManagementStatus::Importable + } else { + PiManagementStatus::Unsupported + } + }); + + let mut reasons = map_raw_reasons(&raw); + extend_reasons(&mut reasons, managed.reasons.clone()); + extend_reasons(&mut reasons, map_composer_reasons(&composition)); + extend_reasons(&mut reasons, map_gateway_reasons(&gateway)); + + PiNativeEntryInspection { + diagnostic: PiNativeDiagnostic { + provider_key: provider_key.to_string(), + display_name: entry + .value + .get("name") + .and_then(Value::as_str) + .map(ToOwned::to_owned), + fingerprint: pi_raw_provider_fingerprint(&entry.raw_source), + kind, + raw_validity, + managed_assessment: managed.assessment, + composition_status: map_composition_status(composition.status), + management_status, + gateway_status: map_gateway_status(gateway.capability), + reasons, + }, + managed_config: managed.config, + composition, + } +} + +fn assess_managed( + raw_validity: PiRawValidity, + kind: PiNativeEntryKind, + value: &Value, +) -> ManagedResult { + if raw_validity != PiRawValidity::Valid { + return ManagedResult { + assessment: PiManagedAssessment::Unsupported, + config: None, + reasons: Vec::new(), + }; + } + if kind != PiNativeEntryKind::CustomCatalog { + let mut reasons = vec![diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::CatalogRequired, + "/models", + )]; + if value + .get("modelOverrides") + .and_then(Value::as_object) + .is_some_and(|overrides| !overrides.is_empty()) + { + reasons.push(diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::ModelOverridesOnly, + "/modelOverrides", + )); + } + return ManagedResult { + assessment: PiManagedAssessment::Unsupported, + config: None, + reasons, + }; + } + + let config = match serde_json::from_value::(value.clone()) { + Ok(config) => config, + Err(_) => { + return ManagedResult { + assessment: PiManagedAssessment::Unsupported, + config: None, + reasons: vec![diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::ManagedTypeConversionFailed, + "", + )], + }; + } + }; + let mut reasons = collect_managed_reasons(&config); + if reasons.is_empty() { + if let Err(error) = validate_pi_managed_provider(&config) { + reasons.push(managed_validation_reason(error)); + } + } + if !reasons.is_empty() { + return ManagedResult { + assessment: PiManagedAssessment::Unsupported, + config: None, + reasons, + }; + } + ManagedResult { + assessment: PiManagedAssessment::Manageable, + config: Some(config), + reasons, + } +} + +fn collect_managed_reasons(config: &PiManagedProviderConfig) -> Vec { + let mut reasons = Vec::new(); + let mut ids = HashSet::with_capacity(config.models.len()); + if config.models.is_empty() { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::MissingExplicitModels, + "/models", + ), + ); + } + if config + .base_url + .as_deref() + .is_some_and(|value| !valid_http_endpoint(value)) + { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::InvalidEndpoint, + "/baseUrl", + ), + ); + } + for (index, model) in config.models.iter().enumerate() { + let pointer = format!("/models/{index}"); + if model.id.is_empty() { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::EmptyModelId, + &format!("{pointer}/id"), + ), + ); + } else if !ids.insert(model.id.as_str()) { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::DuplicateModelId, + &format!("{pointer}/id"), + ), + ); + } + if model + .base_url + .as_deref() + .is_some_and(|value| !valid_http_endpoint(value)) + { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::InvalidEndpoint, + &format!("{pointer}/baseUrl"), + ), + ); + } + if model.api.is_none() && config.api.is_none() { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::MissingEffectiveApi, + &format!("{pointer}/api"), + ), + ); + } + if model.base_url.is_none() && config.base_url.is_none() { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::MissingEffectiveEndpoint, + &format!("{pointer}/baseUrl"), + ), + ); + } + for (field, value) in [ + ("contextWindow", model.context_window), + ("maxTokens", model.max_tokens), + ] { + if value.is_some_and(|value| value.get() <= 0.0) { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::NonPositiveModelLimit, + &format!("{pointer}/{field}"), + ), + ); + } + } + } + + for (model_id, model_override) in &config.model_overrides { + let pointer = format!("/modelOverrides/{}", escape_json_pointer(model_id)); + if !ids.contains(model_id.as_str()) { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::UnknownModelOverride, + &pointer, + ), + ); + } + for (field, value) in [ + ("contextWindow", model_override.context_window), + ("maxTokens", model_override.max_tokens), + ] { + if value.is_some_and(|value| value.get() <= 0.0) { + add_reason( + &mut reasons, + diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::NonPositiveModelLimit, + &format!("{pointer}/{field}"), + ), + ); + } + } + } + reasons +} + +fn managed_validation_reason(error: PiConfigError) -> PiDiagnosticReason { + if let PiConfigError::UnrepresentableCompat { json_pointer } = &error { + return diagnostic_reason( + PiDiagnosticLayer::Managed, + PiReasonCode::UnrepresentableCompat, + json_pointer, + ); + } + let (code, pointer) = match error { + PiConfigError::ProviderHasNoModels => (PiReasonCode::MissingExplicitModels, "/models"), + PiConfigError::EmptyApiId => (PiReasonCode::ManagedTypeConversionFailed, "/api"), + PiConfigError::EmptyModelId => (PiReasonCode::EmptyModelId, "/models"), + PiConfigError::DuplicateModelId(_) => (PiReasonCode::DuplicateModelId, "/models"), + PiConfigError::ModelNotFound(_) => (PiReasonCode::ManagedTypeConversionFailed, "/models"), + PiConfigError::MissingEffectiveApi { .. } => (PiReasonCode::MissingEffectiveApi, "/models"), + PiConfigError::MissingEffectiveEndpoint { .. } => { + (PiReasonCode::MissingEffectiveEndpoint, "/models") + } + PiConfigError::InvalidEndpoint { .. } => (PiReasonCode::InvalidEndpoint, "/baseUrl"), + PiConfigError::UnknownModelOverride(_) => { + (PiReasonCode::UnknownModelOverride, "/modelOverrides") + } + PiConfigError::InvalidCompat { .. } => (PiReasonCode::InvalidCompat, "/compat"), + PiConfigError::UnrepresentableCompat { .. } => { + unreachable!("handled before the exhaustive mapping") + } + PiConfigError::EmptyOptionalField { .. } => (PiReasonCode::EmptyOptionalField, ""), + PiConfigError::NonPositiveModelLimit { .. } => { + (PiReasonCode::NonPositiveModelLimit, "/models") + } + PiConfigError::InvalidThinkingLevelValue { .. } => { + (PiReasonCode::InvalidThinkingLevel, "/models") + } + }; + diagnostic_reason(PiDiagnosticLayer::Managed, code, pointer) +} + +fn classify_kind( + provider_key: &str, + value: &Value, + raw_validity: PiRawValidity, +) -> PiNativeEntryKind { + if PI_BUILTIN_PROVIDER_KEYS.contains(&provider_key) { + return PiNativeEntryKind::BuiltInOverlay; + } + if raw_validity != PiRawValidity::Valid { + return PiNativeEntryKind::UnknownShape; + } + let Some(object) = value.as_object() else { + return PiNativeEntryKind::UnknownShape; + }; + if object + .get("models") + .and_then(Value::as_array) + .is_some_and(|models| !models.is_empty()) + { + PiNativeEntryKind::CustomCatalog + } else if object + .keys() + .any(|key| RECOGNIZED_PROVIDER_FIELDS.contains(&key.as_str())) + { + PiNativeEntryKind::ExtensionOverlay + } else { + PiNativeEntryKind::UnknownShape + } +} + +fn map_raw_validity(validity: PiRawValidity) -> PiRawNativeValidity { + match validity { + PiRawValidity::Valid => PiRawNativeValidity::Valid, + PiRawValidity::Invalid => PiRawNativeValidity::Invalid, + PiRawValidity::Unknown => PiRawNativeValidity::Unknown, + } +} + +fn map_composition_status(status: PiComposerStatus) -> PiCompositionStatus { + match status { + PiComposerStatus::Composed => PiCompositionStatus::Composed, + PiComposerStatus::Failed => PiCompositionStatus::Failed, + PiComposerStatus::Unknown => PiCompositionStatus::Unknown, + } +} + +fn map_gateway_status(capability: PiGatewayCapability) -> PiGatewayStatus { + match capability { + PiGatewayCapability::Proxyable => PiGatewayStatus::Proxyable, + PiGatewayCapability::DirectOnly => PiGatewayStatus::DirectOnly, + PiGatewayCapability::Unknown => PiGatewayStatus::Unknown, + } +} + +fn map_raw_reasons(raw: &PiRawSchemaEvaluation) -> Vec { + raw.reasons + .iter() + .map(|reason| { + diagnostic_reason( + PiDiagnosticLayer::RawSchema, + match reason.code { + PiRawReasonCode::SchemaMismatch => PiReasonCode::RawSchemaMismatch, + PiRawReasonCode::UnsupportedOperator => { + PiReasonCode::RawSchemaUnsupportedOperator + } + PiRawReasonCode::PinDrift => PiReasonCode::RawSchemaPinDrift, + PiRawReasonCode::AmbiguousSchema => PiReasonCode::RawSchemaAmbiguous, + }, + &reason.json_pointer, + ) + }) + .collect() +} + +fn map_composer_reasons(composition: &PiNativeComposition) -> Vec { + composition + .reasons + .iter() + .map(|reason| { + diagnostic_reason( + PiDiagnosticLayer::Composition, + match reason.code { + PiComposerReasonCode::CatalogRequired => PiReasonCode::CatalogRequired, + PiComposerReasonCode::MissingExplicitModels => { + PiReasonCode::MissingExplicitModels + } + PiComposerReasonCode::MissingEffectiveApi => PiReasonCode::MissingEffectiveApi, + PiComposerReasonCode::MissingEffectiveEndpoint => { + PiReasonCode::MissingEffectiveEndpoint + } + PiComposerReasonCode::NonPositiveModelLimit => { + PiReasonCode::NonPositiveModelLimit + } + PiComposerReasonCode::UnrepresentableCompat => { + PiReasonCode::UnrepresentableCompat + } + PiComposerReasonCode::CompositionFailed => PiReasonCode::CompositionFailed, + }, + &reason.json_pointer, + ) + }) + .collect() +} + +fn map_gateway_reasons(gateway: &PiGatewayAssessment) -> Vec { + gateway + .reasons + .iter() + .map(|reason| { + diagnostic_reason( + PiDiagnosticLayer::Gateway, + match reason.code { + PiGatewayReasonCode::UnsupportedFamily => { + PiReasonCode::UnsupportedGatewayFamily + } + PiGatewayReasonCode::UnsupportedCredentialKind => { + PiReasonCode::UnsupportedCredentialKind + } + PiGatewayReasonCode::InvalidEndpoint => PiReasonCode::InvalidEndpoint, + PiGatewayReasonCode::MissingCredential => { + PiReasonCode::GatewayCredentialUnavailable + } + PiGatewayReasonCode::InvalidHeaderName => PiReasonCode::InvalidHeaderName, + PiGatewayReasonCode::InvalidHeaderValue => PiReasonCode::InvalidHeaderValue, + PiGatewayReasonCode::ProtectedHeader => PiReasonCode::ProtectedHeader, + PiGatewayReasonCode::DeferredValueUnavailable => { + PiReasonCode::DeferredValueUnavailable + } + }, + &reason.json_pointer, + ) + }) + .collect() +} + +fn diagnostic_reason( + layer: PiDiagnosticLayer, + code: PiReasonCode, + pointer: &str, +) -> PiDiagnosticReason { + PiDiagnosticReason::new(layer, code, Some(pointer.to_string())) +} + +fn add_reason(reasons: &mut Vec, reason: PiDiagnosticReason) { + if !reasons.contains(&reason) { + reasons.push(reason); + } +} + +fn extend_reasons( + reasons: &mut Vec, + candidates: impl IntoIterator, +) { + for reason in candidates { + add_reason(reasons, reason); + } +} + +fn valid_http_endpoint(value: &str) -> bool { + Url::parse(value) + .ok() + .is_some_and(|url| matches!(url.scheme(), "http" | "https") && url.host().is_some()) +} + +fn escape_json_pointer(value: &str) -> String { + value.replace('~', "~0").replace('/', "~1") +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::pi_config::model::effective_pi_model; + use serde_json::json; + use std::fs; + + fn by_key<'a>(diagnostics: &'a [PiNativeDiagnostic], key: &str) -> &'a PiNativeDiagnostic { + diagnostics + .iter() + .find(|diagnostic| diagnostic.provider_key == key) + .expect("diagnostic") + } + + fn has_reason( + diagnostic: &PiNativeDiagnostic, + layer: PiDiagnosticLayer, + code: PiReasonCode, + pointer: &str, + ) -> bool { + diagnostic.reasons.iter().any(|reason| { + reason.layer == layer + && reason.code == code + && reason.json_pointer.as_deref() == Some(pointer) + }) + } + + #[test] + fn public_inspection_service_certifies_the_native_state_matrix() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{ + "providers": { + "anthropic": {"baseUrl": "https://builtin.example"}, + "extension": {"api": "openai-responses"}, + "future-custom": { + "api": "future-wire-v9", + "baseUrl": "https://future.example/v9", + "apiKey": "$FUTURE_KEY", + "models": [{"id": "future"}] + }, + "known-custom": { + "api": "openai-responses", + "baseUrl": "https://known.example/v1", + "apiKey": "!read-secret", + "headers": {"x-tenant": "${TENANT}"}, + "models": [{"id": "known"}] + }, + "malformed": {"models": "not-an-array"} + } +}"#, + ) + .expect("write fixture"); + let bytes_before = fs::read(&path).expect("before"); + let claims = BTreeMap::new(); + let diagnostics = + PiNativeInspectionService::inspect_catalog(&path, &claims).expect("inspect service"); + assert_eq!(fs::read(&path).expect("after"), bytes_before); + assert_eq!(diagnostics.len(), 5); + + for key in ["anthropic", "extension"] { + let diagnostic = by_key(&diagnostics, key); + assert_eq!(diagnostic.raw_validity, PiRawNativeValidity::Valid); + assert_eq!(diagnostic.composition_status, PiCompositionStatus::Unknown); + assert_eq!(diagnostic.gateway_status, PiGatewayStatus::Unknown); + assert!(has_reason( + diagnostic, + PiDiagnosticLayer::Composition, + PiReasonCode::CatalogRequired, + "/models" + )); + } + + let future = by_key(&diagnostics, "future-custom"); + assert_eq!(future.raw_validity, PiRawNativeValidity::Valid); + assert_eq!(future.composition_status, PiCompositionStatus::Composed); + assert_eq!(future.managed_assessment, PiManagedAssessment::Manageable); + assert_eq!(future.management_status, PiManagementStatus::Importable); + assert_eq!(future.gateway_status, PiGatewayStatus::DirectOnly); + assert!(has_reason( + future, + PiDiagnosticLayer::Gateway, + PiReasonCode::UnsupportedGatewayFamily, + "/models/0/api" + )); + + let known = by_key(&diagnostics, "known-custom"); + assert_eq!(known.composition_status, PiCompositionStatus::Composed); + assert_eq!(known.management_status, PiManagementStatus::Importable); + assert_eq!(known.gateway_status, PiGatewayStatus::Proxyable); + + let malformed = by_key(&diagnostics, "malformed"); + assert_eq!(malformed.raw_validity, PiRawNativeValidity::Invalid); + assert_eq!(malformed.composition_status, PiCompositionStatus::Unknown); + assert_eq!(malformed.gateway_status, PiGatewayStatus::Unknown); + } + + #[test] + fn public_inspection_fails_closed_for_unrepresentable_compat_spread() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{"surrogate":{ + "api":"openai-responses", + "baseUrl":"https://compat.example/v1", + "apiKey":"literal", + "compat":{"chatTemplateKwargs":"😀"}, + "models":[{"id":"m"}], + "modelOverrides":{"m":{"compat":{"chatTemplateKwargs":{"named":true}}}} +}}}"#, + ) + .expect("write"); + + let diagnostic = + &PiNativeInspectionService::inspect_catalog(&path, &BTreeMap::new()).unwrap()[0]; + assert_eq!(diagnostic.raw_validity, PiRawNativeValidity::Valid); + assert_eq!( + diagnostic.managed_assessment, + PiManagedAssessment::Unsupported + ); + assert_eq!(diagnostic.composition_status, PiCompositionStatus::Unknown); + assert_eq!(diagnostic.gateway_status, PiGatewayStatus::Unknown); + assert!(has_reason( + diagnostic, + PiDiagnosticLayer::Managed, + PiReasonCode::UnrepresentableCompat, + "/modelOverrides/m/compat" + )); + assert!(has_reason( + diagnostic, + PiDiagnosticLayer::Composition, + PiReasonCode::UnrepresentableCompat, + "/modelOverrides/m/compat" + )); + } + + #[test] + fn public_inspection_accepts_surrogates_overridden_before_the_final_result() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{"surrogate":{ + "api":"openai-responses", + "baseUrl":"https://compat.example/v1", + "apiKey":"literal", + "compat":{"chatTemplateKwargs":"😀"}, + "models":[{"id":"m"}], + "modelOverrides":{"m":{"compat":{"chatTemplateKwargs":{ + "0":"repaired-high", + "1":"repaired-low", + "named":true + }}}} +}}}"#, + ) + .expect("write"); + + let inspection = + PiNativeInspectionService::inspect_entry(&path, "surrogate", &BTreeMap::new()) + .unwrap() + .expect("provider"); + assert_eq!( + inspection.diagnostic.managed_assessment, + PiManagedAssessment::Manageable + ); + assert_eq!( + inspection.diagnostic.composition_status, + PiCompositionStatus::Composed + ); + assert_eq!( + inspection.diagnostic.gateway_status, + PiGatewayStatus::Proxyable + ); + assert!(!inspection + .diagnostic + .reasons + .iter() + .any(|reason| reason.code == PiReasonCode::UnrepresentableCompat)); + assert_eq!( + inspection.composition.models[0].compat, + Some(json!({ + "chatTemplateKwargs": { + "0": "repaired-high", + "1": "repaired-low", + "named": true + } + })) + ); + } + + #[test] + fn managed_rejection_does_not_control_raw_composition() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{"duplicate":{ + "apiKey":"literal", + "models":[ + {"id":"same","api":"openai-responses","baseUrl":"https://one.example"}, + {"id":"same","api":"future-wire","baseUrl":"https://two.example"} + ], + "modelOverrides":{"missing":{"maxTokens":7.5}} +}}}"#, + ) + .expect("write"); + let diagnostic = + &PiNativeInspectionService::inspect_catalog(&path, &BTreeMap::new()).unwrap()[0]; + assert_eq!(diagnostic.raw_validity, PiRawNativeValidity::Valid); + assert_eq!( + diagnostic.managed_assessment, + PiManagedAssessment::Unsupported + ); + assert_eq!(diagnostic.composition_status, PiCompositionStatus::Composed); + assert!(has_reason( + diagnostic, + PiDiagnosticLayer::Managed, + PiReasonCode::DuplicateModelId, + "/models/1/id" + )); + assert!(has_reason( + diagnostic, + PiDiagnosticLayer::Managed, + PiReasonCode::UnknownModelOverride, + "/modelOverrides/missing" + )); + } + + #[test] + fn unknown_thinking_shape_is_lossless_through_managed_and_effective_boundaries() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{"thinking":{ + "api":"anthropic-messages", + "baseUrl":"https://thinking.example", + "apiKey":"literal", + "models":[{ + "id":"m", + "thinkingLevelMap":{"high":"future-high","future":{"opaque":true}} + }] +}}}"#, + ) + .expect("write"); + let inspection = + PiNativeInspectionService::inspect_entry(&path, "thinking", &BTreeMap::new()) + .expect("inspect") + .expect("entry"); + assert_eq!( + inspection.diagnostic.raw_validity, + PiRawNativeValidity::Valid + ); + assert_eq!( + inspection.diagnostic.composition_status, + PiCompositionStatus::Composed + ); + assert_eq!( + inspection.composition.models[0] + .thinking_level_map + .as_ref() + .expect("opaque thinking")["future"], + json!({"opaque": true}) + ); + assert_eq!( + inspection.diagnostic.managed_assessment, + PiManagedAssessment::Manageable + ); + let managed = serde_json::to_value( + inspection + .managed_config + .as_ref() + .expect("schema-valid managed config"), + ) + .expect("serialize managed config"); + assert_eq!( + managed.pointer("/models/0/thinkingLevelMap/future"), + Some(&json!({"opaque": true})) + ); + let effective = serde_json::to_value( + effective_pi_model( + inspection.managed_config.as_ref().expect("managed config"), + "m", + ) + .expect("effective model"), + ) + .expect("serialize effective model"); + assert_eq!( + effective.pointer("/thinkingLevelMap/future"), + Some(&json!({"opaque": true})) + ); + } + + #[test] + fn public_inspection_accepts_schema_valid_whitespace_strings() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write( + &path, + r#"{"providers":{"whitespace":{ + "name":" ", + "api":"anthropic-messages", + "baseUrl":"https://whitespace.example", + "apiKey":" ", + "models":[{"id":" ","name":" "}] +}}}"#, + ) + .expect("write"); + let inspection = + PiNativeInspectionService::inspect_entry(&path, "whitespace", &BTreeMap::new()) + .expect("inspect") + .expect("entry"); + assert_eq!( + inspection.diagnostic.raw_validity, + PiRawNativeValidity::Valid + ); + assert_eq!( + inspection.diagnostic.managed_assessment, + PiManagedAssessment::Manageable + ); + assert_eq!( + inspection.diagnostic.management_status, + PiManagementStatus::Importable + ); + assert_eq!( + inspection + .managed_config + .as_ref() + .expect("managed config") + .models[0] + .id, + " " + ); + } + + #[test] + fn exact_entry_fingerprint_changes_only_when_that_raw_entry_changes() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + let claims = BTreeMap::new(); + fs::write( + &path, + r#"{"providers":{ + "target":{"api":"openai-responses","baseUrl":"https://target","apiKey":"x","models":[{"id":"m"}]}, + "sibling":{"api":"openai-responses","baseUrl":"https://sibling","apiKey":"x","models":[{"id":"m"}]} +}}"#, + ) + .expect("write"); + let first = PiNativeInspectionService::inspect_entry(&path, "target", &claims) + .unwrap() + .unwrap() + .diagnostic + .fingerprint; + fs::write( + &path, + r#"{"providers":{ + "target":{"api":"openai-responses","baseUrl":"https://target","apiKey":"x","models":[{"id":"m"}]}, + "sibling":{"api":"openai-responses","baseUrl":"https://changed","apiKey":"x","models":[{"id":"m"}]} +}}"#, + ) + .expect("write sibling"); + let after_sibling = PiNativeInspectionService::inspect_entry(&path, "target", &claims) + .unwrap() + .unwrap() + .diagnostic + .fingerprint; + assert_eq!(first, after_sibling); + } + + #[test] + fn agent_dir_normalization_matches_pi_path_semantics() { + let temp = tempfile::tempdir().expect("tempdir"); + let file_url = Url::from_file_path(temp.path()) + .expect("absolute temp path") + .to_string(); + assert_eq!( + normalize_pi_agent_dir(&file_url, Path::new("/unused")).expect("file URL"), + temp.path() + ); + let spaced = " relative agent dir "; + assert_eq!( + normalize_pi_agent_dir(spaced, Path::new("/unused")).expect("spaced path"), + PathBuf::from(spaced) + ); + } + + #[test] + fn pi_config_error_stays_managed_only() { + let error = PiConfigError::EmptyApiId; + assert_eq!(error.to_string(), "Pi API id cannot be empty"); + } +} diff --git a/src-tauri/src/pi_config/native_inspection_certification.rs b/src-tauri/src/pi_config/native_inspection_certification.rs new file mode 100644 index 000000000..dd379177f --- /dev/null +++ b/src-tauri/src/pi_config/native_inspection_certification.rs @@ -0,0 +1,834 @@ +#![cfg(test)] +//! 只读 native inspection 契约测试。 +//! +//! ## 目标 +//! **pinned Pi 决定什么是合法**。本仓库的 DTO 形状、网关支持范围、头部策略 +//! 都不得成为"合法性"的来源:schema 接受的,managed 不得拒绝也不得丢值; +//! Pi 会发出的,网关不得降级;Pi 不接受的形态,我们也不假装支持。 +//! +//! ## 六条裁决及其上游证据 +//! C1【无损性】pinned schema 对 `thinkingLevelMap` 只约束 7 个标准键 +//! (string|null;oracle 实证 `low: 2` 非法),额外键无约束(oracle 实证 +//! `future: {nested:true}` 合法);`cost`/tier 同样接受未来键。managed 与 +//! **effective 边界**(`effective_pi_model` 是 projection/routing/failover +//! 的共同入口)都必须无损,**空容器与缺席必须保持可区分**(`{}` 之于 +//! thinkingLevelMap、`[]` 之于 cost.tiers 同理)。 +//! 据此取代两个既有测试中把收窄固化为断言的部分: +//! `managed_narrowing_rejects_duplicates_or_unknown_thinking_keys` 与 +//! `unknown_thinking_shape_is_lossless_for_composer_and_narrowed_separately` +//! (授权改写、可改名;DuplicateModelId 与 composer 无损两个语义由本套件 +//! 直接接管)。若 `InvalidThinkingLevel` 变体因此不再可构造,授权移除。 +//! C2【认证头】authorization / x-api-key / x-goog-api-key 是候选认证头, +//! 不是 protected。取值次序据 pinned SDK 与 composer 源码:authHeader 未 +//! 设时显式头优先于 apiKey 合成值(Anthropic/OpenAI SDK 按"合成 auth → +//! 显式 headers"合并,后项覆盖);authHeader:true 时合成 Bearer 反过来 +//! 优先(pinned provider-composer 在自定义头之后写入,且只写 +//! Authorization、不动 x-api-key)。**header-only 凭证对四族都不是 Pi 原生 +//! 可请求形态**:pinned `ModelRuntime.prepareRequest()` 先解析 auth,得不到 +//! AuthResult 即抛 "Provider is not configured",在合并 headers 之前返回, +//! 而 headers 本身永不产生 AuthResult(Google adapter 更是无条件要 apiKey)。 +//! 故无 apiKey 时维持 MissingCredential 降级,但认证头本身仍不得被报为 +//! ProtectedHeader。 +//! C3【传输层】放宽认证头不得连带放宽传输层:逐跳头完整覆盖并以 `proxy-` +//! **前缀**拒绝;契约 header 六分类中的 Gateway/HTTP owned(proxy trace / +//! CDN 客户端身份 / 分布式追踪)同样拒绝,清单与生产 forwarder 无条件 +//! 剥离的集合对齐。 +//! C4【deferred 值的校验时机】pinned `resolveConfigValueOrThrow()` 先执行 +//! `!command` / 展开 `${ENV}`,再使用结果;**从不按 HTTP 头规则校验原始 +//! 表达式**(命令输出 trim,环境模板不 trim,解析结果亦不做头合法性校验)。 +//! 因此原始表达式含头非法字符、而解析结果合法的配置必须被接受;头合法性 +//! 校验只能发生在物化之后(这是网关自身的传输约束,保留)。**字面量值仍在 +//! 判定期校验,且该规则对 credential 与 header 一视同仁**——判定期说 +//! "可代理"而每次物化必然失败,是判定层与执行层自相矛盾。 +//! +//! C5【凭证种类,2026-08-02 新增,**已 request-capture 实证**】pinned +//! Anthropic 传输层以 `apiKey.includes("sk-ant-oat")`(子串,非前缀)判定 +//! OAuth,命中则以 `Authorization: Bearer` 发送、**不发 x-api-key**,并附 +//! `anthropic-beta: claude-code-20250219,oauth-2025-04-20,...`;**models.json +//! 里的字面量 apiKey 同样会走该分支**;该判定**只在 Anthropic 族**,同形 +//! token 在 OpenAI 族仍按普通 Bearer 发送。因此网关不得把这类凭证当普通 +//! x-api-key 代理:字面量命中即判定期 DirectOnly 并给结构化理由(不得是 +//! MissingCredential);deferred 凭证判定期不可知,则**物化期解析出命中值 +//! 时必须失败**,绝不发出错误的认证形态。 +//! **完整 OAuth 传输(Bearer + oauth beta 值)不在前置 C 范围**——按 +//! 项目范围划分,gateway 数据面属主工程,且需要 +//! 先补 request-capture oracle。本工程只保证判定诚实、不发错凭证。 +//! C6【entry 隔离,2026-08-02 新增】pinned Pi 逐 entry 做 TypeBox 判定, +//! 单个 entry 的取值错误(如 `contextWindow: 1e400`)只令该 entry 非法; +//! 整文件解析失败会让合法的兄弟 entry 被连坐隐藏,违反四层判定"每个 +//! entry 独立"的核心设计。 +//! +//! ## 实现方义务(不在本文件断言,交盲审核查) +//! O1 `compat` 需复现 JavaScript object-spread 对嵌套值(尤其数组)的语义; +//! O2 架构扫描器:cfg 布尔语义(`cfg(not(test))` 的生产代码必须被扫描)、 +//! 不得按 `tests/` 路径整体跳过文件、嵌套模块须继承父层归属。 +//! +//! ## 上游实证(request-capture,2026-08-02) +//! `scripts/pi-transport-capture.mjs` 以本地抓包端点作 baseUrl,用 pinned Pi +//! 的 adapter 真发请求,实测矩阵(据此 C2/C5 不再是"读源码推断"): +//! - anthropic 普通 key → `x-api-key: `; +//! - anthropic `sk-ant-oat...` → `authorization: Bearer ` + +//! `anthropic-beta: claude-code-20250219,oauth-2025-04-20,...`,**无 x-api-key**; +//! - anthropic apiKey + 显式 `x-api-key` → 发**显式值**(显式覆盖合成); +//! - anthropic apiKey + 显式 `authorization` → 两者**并存** +//! (`authorization` 取显式值,`x-api-key` 取合成值); +//! - openai responses/completions + 显式 `authorization` → 发**显式值**; +//! - openai + `sk-ant-oat` 形状 token → 仍是普通 `Bearer`,无 OAuth 特殊处理; +//! - openai completions + 显式 `x-api-key` → 与合成 `authorization` **并存**。 +//! +//! ## 残余 +//! Google 族两值并存的优先级、头名大小写变体未实测;命令输出 trim 与环境模板 +//! 不 trim 的差异属数据面语义,本只读面不断言;完整 OAuth 传输实现按范围表 +//! 归主工程(harness 已就位,可直接扩为受冻结的 transport oracle); +//! 其余按盲审 finding 处理。 + +use super::composer::compose_explicit_custom_catalog; +use super::gateway::{assess_composition, PiGatewayCapability, PiGatewayReasonCode}; +use super::model::{ + effective_pi_model, validate_pi_managed_provider, PiConfigError, PiManagedAssessment, + PiManagedProviderConfig, PiManagementStatus, PiRawNativeValidity, +}; +use super::native::{inspect_pi_native_catalog, inspect_pi_native_entry}; +use super::raw_schema::evaluate_provider_value; +use serde_json::{json, Value}; +use sha2::{Digest, Sha256}; +use std::collections::BTreeMap; +use std::fs; +use std::path::{Path, PathBuf}; + +fn repo_root() -> PathBuf { + Path::new(env!("CARGO_MANIFEST_DIR")) + .parent() + .expect("workspace root") + .to_path_buf() +} + +fn write_catalog(value: &Value) -> (tempfile::TempDir, PathBuf) { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write(&path, serde_json::to_string_pretty(value).expect("encode")).expect("write"); + (temp, path) +} + +fn composed_catalog(value: Value) -> super::composer::PiNativeComposition { + let raw = evaluate_provider_value(&value); + compose_explicit_custom_catalog( + "candidate", + raw.valid_provider.as_ref().expect("raw-valid input"), + ) +} + +fn has_gateway_reason( + gateway: &super::gateway::PiGatewayAssessment, + code: PiGatewayReasonCode, +) -> bool { + gateway.reasons.iter().any(|reason| reason.code == code) +} + +// --------------------------------------------------------------------------- +// pinned 夹具冻结——oracle 是上游出处工件,不得为过测试再生成 +// --------------------------------------------------------------------------- + +const PINNED_FIXTURES: &[(&str, &str)] = &[ + ( + "tests/fixtures/pi/native-oracle/composer-oracle-v1.json", + "f7e54bb84e5fd6d50e5762dc304834410fa73ef608c2f9c42475c5983f8e0cf5", + ), + ( + "tests/fixtures/pi/native-oracle/field-coverage-v1.json", + "b8b85e611cf1dbef86c611df185ba8ac2d64160087d0c6e47747f838a0fafe42", + ), + ( + "tests/fixtures/pi/native-oracle/provenance-v1.json", + "6b2f9570ecc58d54ebe3da094530fee1c8c0d4a8265fa9c3199218582cb8dbcb", + ), + ( + "tests/fixtures/pi/native-oracle/provider-schema.snapshot.json", + "e498c9f1b344eee1bd3c3ba74d1b648dcb835378cfad92800ec80078b825745c", + ), + ( + "tests/fixtures/pi/native-oracle/raw-oracle-v1.json", + "5aaa37160f96a0fe50867d900ca38c73f13aba769e156a883324368d9dbeeb9a", + ), + ( + "tests/fixtures/pi/native-oracle/transport-oracle-v1.json", + "b2c816e53b60da5cd6352d2c23934939e9f6dd0077971488fe9dd36fa723e855", + ), + ( + "tests/fixtures/pi/module-boundaries-v1.json", + "a69ab84fc0db323d5eb8ddc63555a9c69613dda865962f077cd8691639951b4d", + ), +]; + +#[test] +fn certify_pinned_fixtures_are_frozen() { + for (relative, expected) in PINNED_FIXTURES { + let bytes = fs::read(repo_root().join(relative)) + .unwrap_or_else(|e| panic!("read fixture {relative}: {e}")); + assert_eq!( + &format!("{:x}", Sha256::digest(bytes)), + expected, + "pinned fixture '{relative}' drifted; fixtures are upstream provenance \ + artifacts and may only change under adjudication" + ); + } +} + +// --------------------------------------------------------------------------- +// 被取代测试中必须保留的语义,由本套件直接接管 +// --------------------------------------------------------------------------- + +#[test] +fn certify_duplicate_model_id_rejection_is_preserved() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "api": "anthropic-messages", + "baseUrl": "https://dup.example", + "apiKey": "literal", + "models": [{"id": "same"}, {"id": "same"}] + })) + .expect("deserialize managed provider"); + assert_eq!( + validate_pi_managed_provider(&config), + Err(PiConfigError::DuplicateModelId("same".into())), + "duplicate model ids must keep being rejected" + ); +} + +#[test] +fn certify_composer_thinking_losslessness_guard() { + let odd_map = json!({"high": "h", "future": {"opaque": true}}); + let composition = composed_catalog(json!({ + "api": "anthropic-messages", + "baseUrl": "https://thinking.example", + "apiKey": "literal", + "models": [{"id": "m", "thinkingLevelMap": odd_map}] + })); + assert_eq!( + composition.models[0].thinking_level_map.as_ref(), + Some(&odd_map), + "composer keeps the raw thinkingLevelMap value verbatim" + ); +} + +// --------------------------------------------------------------------------- +// C1:schema 合法值必须无损直到 effective 边界 +// --------------------------------------------------------------------------- + +#[test] +fn certify_managed_losslessness_through_effective_boundary() { + // 标准键只取 schema 允许的 string|null;额外键覆盖全部 JSON 类型。 + let model_map = json!({ + "high": "native-high", + "medium": null, + "future-level": "textual", + "vendor": {"opaque": {"nested": true}}, + "budget": 42, + "enabled": true + }); + // 与 model_map 共有 "high",用于绑定 override 的覆盖方向。 + let override_map = json!({ + "high": "override-high", + "low": "override-low", + "another-future": [1, "two", null] + }); + let cost = json!({ + "input": 1.5, + "output": 2.5, + "cacheRead": 0.5, + "cacheWrite": 0.25, + "futureRate": 9.0, + "tiers": [{ + "inputTokensAbove": 100.0, + "input": 1.0, + "output": 2.0, + "cacheRead": 0.5, + "cacheWrite": 0.25, + "futureTierField": "opaque" + }] + }); + let catalog = json!({ + "providers": { + "thinking": { + "api": "anthropic-messages", + "baseUrl": "https://thinking.example", + "apiKey": "literal", + "models": [ + {"id": "m", "thinkingLevelMap": model_map.clone(), "cost": cost.clone()}, + {"id": "empty-map", "thinkingLevelMap": {}}, + {"id": "absent-map"} + ], + "modelOverrides": {"m": {"thinkingLevelMap": override_map.clone()}} + } + } + }); + let (_temp, path) = write_catalog(&catalog); + let inspection = inspect_pi_native_entry(&path, "thinking", &BTreeMap::new()) + .expect("inspect") + .expect("entry present"); + + assert_eq!( + inspection.diagnostic.raw_validity, + PiRawNativeValidity::Valid, + "the pinned schema accepts additional thinkingLevelMap and cost members" + ); + assert_eq!( + inspection.diagnostic.managed_assessment, + PiManagedAssessment::Manageable, + "managed must not reject what the executed pin accepts" + ); + assert_eq!( + inspection.diagnostic.management_status, + PiManagementStatus::Importable + ); + // 以序列化后的字符串码断言,便于 InvalidThinkingLevel 变体被整体移除。 + let reasons = serde_json::to_value(&inspection.diagnostic.reasons).expect("serialize reasons"); + assert!( + !reasons + .as_array() + .expect("reasons array") + .iter() + .any(|reason| reason["code"] == "invalid_thinking_level"), + "no invalid_thinking_level reason may fire for schema-valid input" + ); + + let managed = inspection.managed_config.expect("managed config"); + let round_trip = serde_json::to_value(&managed).expect("serialize managed config"); + assert_eq!( + round_trip.pointer("/models/0/thinkingLevelMap"), + Some(&model_map), + "model thinkingLevelMap must round-trip losslessly" + ); + assert_eq!( + round_trip.pointer("/modelOverrides/m/thinkingLevelMap"), + Some(&override_map), + "override thinkingLevelMap must round-trip losslessly" + ); + assert_eq!( + round_trip.pointer("/models/0/cost"), + Some(&cost), + "cost and tier members must round-trip losslessly, including future keys" + ); + // 空对象与缺席是两种原生形态,序列化必须保持可区分。 + assert_eq!( + round_trip.pointer("/models/1/thinkingLevelMap"), + Some(&json!({})), + "an explicitly empty thinkingLevelMap must survive as an empty object" + ); + assert_eq!( + round_trip.pointer("/models/2/thinkingLevelMap"), + None, + "an absent thinkingLevelMap must stay absent" + ); + + // effective 是 projection / runtime / routing / failover 的共同入口: + // DTO 修好后在这里二次收窄同样是丢值。 + let effective = effective_pi_model(&managed, "m").expect("effective model"); + let effective_value = serde_json::to_value(&effective).expect("serialize effective model"); + let mut merged = model_map.as_object().expect("model map").clone(); + for (key, value) in override_map.as_object().expect("override map") { + merged.insert(key.clone(), value.clone()); + } + assert_eq!( + effective_value.pointer("/thinkingLevelMap"), + Some(&Value::Object(merged)), + "the effective model must carry the merged map losslessly, with override \ + entries winning on shared keys" + ); + assert_eq!( + effective_value.pointer("/cost"), + Some(&cost), + "the effective model must not drop cost members either" + ); +} + +// --------------------------------------------------------------------------- +// C2:候选认证头不是 protected +// --------------------------------------------------------------------------- + +#[test] +fn certify_auth_candidate_headers_are_not_protected() { + // (a) Anthropic:显式 x-api-key 不得被拒,取值优先于 apiKey 合成值。 + let explicit = composed_catalog(json!({ + "api": "anthropic-messages", + "baseUrl": "https://anthropic.example", + "apiKey": "synthesized-secret", + "headers": {"x-api-key": "explicit-secret"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&explicit); + assert!( + !has_gateway_reason(&gateway, PiGatewayReasonCode::ProtectedHeader), + "x-api-key is candidate-auth, not protected" + ); + assert_eq!(gateway.capability, PiGatewayCapability::Proxyable); + let materialized = gateway.plans[0] + .materialize(&|_: &str| None) + .expect("materialize literal candidate"); + assert_eq!( + materialized.headers[&http::HeaderName::from_static("x-api-key")], + http::HeaderValue::from_static("explicit-secret"), + "explicit config header value takes precedence over synthesized family auth" + ); + // 认证头永远不进 failover 协议身份。 + if let Some((_, protocol_headers)) = materialized.failover_protocol_identity() { + assert!( + !protocol_headers.contains_key(http::HeaderName::from_static("x-api-key")), + "auth headers must stay out of the failover protocol identity" + ); + } + + // (b) OpenAI-Responses:显式 authorization 同理。 + let bearer = composed_catalog(json!({ + "api": "openai-responses", + "baseUrl": "https://openai.example/v1", + "apiKey": "synthesized-secret", + "headers": {"authorization": "Bearer configured-token"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&bearer); + assert!( + !has_gateway_reason(&gateway, PiGatewayReasonCode::ProtectedHeader), + "authorization is candidate-auth, not protected" + ); + assert_eq!(gateway.capability, PiGatewayCapability::Proxyable); + assert_eq!( + gateway.plans[0] + .materialize(&|_: &str| None) + .expect("materialize") + .headers[&http::HeaderName::from_static("authorization")], + http::HeaderValue::from_static("Bearer configured-token") + ); + + // (c) Google:显式认证头与 apiKey 并存,不得拒绝、不得降级 + // (取值优先级不断言——Google SDK 顺序无上游证据)。 + let google = composed_catalog(json!({ + "api": "google-generative-ai", + "baseUrl": "https://gemini.example", + "apiKey": "literal", + "headers": {"x-goog-api-key": "explicit-secret"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&google); + assert!( + !has_gateway_reason(&gateway, PiGatewayReasonCode::ProtectedHeader), + "x-goog-api-key is candidate-auth, not protected" + ); + assert_eq!(gateway.capability, PiGatewayCapability::Proxyable); +} + +// --------------------------------------------------------------------------- +// C2:authHeader:true 时合成 Bearer 覆盖显式 Authorization +// --------------------------------------------------------------------------- + +#[test] +fn certify_auth_header_bearer_overrides_explicit_authorization() { + let composition = composed_catalog(json!({ + "api": "anthropic-messages", + "baseUrl": "https://anthropic.example", + "apiKey": "synthesized-secret", + "authHeader": true, + "headers": {"authorization": "Bearer explicit-token"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&composition); + assert_eq!( + gateway.capability, + PiGatewayCapability::Proxyable, + "an explicit authorization header must not downgrade an authHeader model" + ); + let materialized = gateway.plans[0] + .materialize(&|_: &str| None) + .expect("materialize literal candidate"); + assert_eq!( + materialized.headers[&http::HeaderName::from_static("authorization")], + http::HeaderValue::from_static("Bearer synthesized-secret"), + "with authHeader:true the synthesized Bearer wins (pinned composer writes it \ + after the explicit headers)" + ); + assert_eq!( + materialized.headers[&http::HeaderName::from_static("x-api-key")], + http::HeaderValue::from_static("synthesized-secret"), + "the Bearer step only rewrites Authorization; family auth stays synthesized" + ); +} + +// --------------------------------------------------------------------------- +// C3:传输层与网关自有身份头 +// --------------------------------------------------------------------------- + +/// 逐跳/传输头。末四项是合成名字:精确枚举无法覆盖,必须按 `proxy-` 前缀拒绝。 +const HOP_BY_HOP_HEADERS: &[&str] = &[ + "host", + "connection", + "content-length", + "transfer-encoding", + "te", + "trailer", + "upgrade", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "proxy-connection", + "proxy-future-extension", + "proxy-tenant-routing", + "proxy-x9", +]; + +/// Gateway/HTTP owned:proxy trace / CDN 客户端身份 / 分布式追踪。 +/// 与生产 forwarder 无条件剥离的集合对齐,两侧同进退。 +const GATEWAY_OWNED_HEADERS: &[&str] = &[ + "forwarded", + "x-forwarded-for", + "x-forwarded-host", + "x-forwarded-port", + "x-forwarded-proto", + "x-real-ip", + "cf-connecting-ip", + "cf-ipcountry", + "cf-ray", + "cf-visitor", + "true-client-ip", + "fastly-client-ip", + "x-azure-clientip", + "x-azure-fdid", + "x-azure-ref", + "akamai-origin-hop", + "x-akamai-config-log-detail", + "x-request-id", + "x-correlation-id", + "x-trace-id", + "x-amzn-trace-id", + "x-b3-traceid", + "x-b3-spanid", + "x-b3-parentspanid", + "x-b3-sampled", + "traceparent", + "tracestate", +]; + +#[test] +fn certify_transport_owned_headers_stay_protected() { + let cases = HOP_BY_HOP_HEADERS + .iter() + .map(|header| ("hop-by-hop", *header)) + .chain( + GATEWAY_OWNED_HEADERS + .iter() + .map(|header| ("gateway-owned", *header)), + ); + for (class, header) in cases { + let composition = composed_catalog(json!({ + "api": "openai-responses", + "baseUrl": "https://openai.example/v1", + "apiKey": "literal", + "headers": {header: "value"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&composition); + assert!( + has_gateway_reason(&gateway, PiGatewayReasonCode::ProtectedHeader), + "{class} header '{header}' must be reported as ProtectedHeader" + ); + assert_eq!( + gateway.capability, + PiGatewayCapability::DirectOnly, + "{class} header '{header}' must keep the model DirectOnly" + ); + } +} + +// --------------------------------------------------------------------------- +// C2:header-only 凭证四族皆非 Pi 原生可请求形态 +// --------------------------------------------------------------------------- + +#[test] +fn certify_header_only_credentials_stay_direct_only() { + // pinned ModelRuntime.prepareRequest() 先解析 auth,得不到 AuthResult 即抛 + // "Provider is not configured",在合并 headers 之前返回;headers 永不产生 + // AuthResult。因此"只有认证头、无 apiKey"必须降级——但认证头本身依然是 + // candidate-auth,不得被报为 ProtectedHeader。 + for (api, header) in [ + ("anthropic-messages", "x-api-key"), + ("openai-completions", "authorization"), + ("openai-responses", "authorization"), + ("google-generative-ai", "x-goog-api-key"), + ] { + let composition = composed_catalog(json!({ + "api": api, + "baseUrl": "https://example.test/v1", + "headers": {header: "header-secret"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&composition); + assert_eq!( + gateway.capability, + PiGatewayCapability::DirectOnly, + "{api}: header-only credentials are not a requestable pinned Pi form" + ); + assert!( + has_gateway_reason(&gateway, PiGatewayReasonCode::MissingCredential), + "{api}: a missing apiKey must be reported as MissingCredential" + ); + assert!( + !has_gateway_reason(&gateway, PiGatewayReasonCode::ProtectedHeader), + "{api}: the auth header itself must not be reported as protected" + ); + } +} + +// --------------------------------------------------------------------------- +// C4:deferred 值只能在物化之后校验 +// --------------------------------------------------------------------------- + +#[test] +fn certify_deferred_header_values_are_validated_after_resolution() { + // 原始表达式含头非法字符(非可见 ASCII),解析结果合法。pinned Pi 先执行 + // 再用结果,从不校验原始表达式,故这类配置必须被接受。 + let expression = "!echo café"; + let deferred = composed_catalog(json!({ + "api": "openai-responses", + "baseUrl": "https://openai.example/v1", + "apiKey": "literal", + "headers": {"x-tenant": expression}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&deferred); + assert!( + !has_gateway_reason(&gateway, PiGatewayReasonCode::InvalidHeaderValue), + "a deferred expression must not be validated as an HTTP header value before \ + it is resolved" + ); + assert_eq!(gateway.capability, PiGatewayCapability::Proxyable); + let materialized = gateway.plans[0] + .materialize(&|value: &str| (value == expression).then(|| "resolved-secret".to_string())) + .expect("materialize resolved candidate"); + assert_eq!( + materialized.headers[&http::HeaderName::from_static("x-tenant")], + http::HeaderValue::from_static("resolved-secret"), + "the resolved value is what reaches the candidate" + ); + + // 防过度放宽:字面量(非 deferred)含头非法字符仍必须当场拒绝。 + let literal = composed_catalog(json!({ + "api": "openai-responses", + "baseUrl": "https://openai.example/v1", + "apiKey": "literal", + "headers": {"x-tenant": "café"}, + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&literal); + assert!( + has_gateway_reason(&gateway, PiGatewayReasonCode::InvalidHeaderValue), + "a literal header value outside visible ASCII must still be rejected" + ); +} + +// --------------------------------------------------------------------------- +// C5:OAuth 凭证绝不能按 x-api-key 代理 +// --------------------------------------------------------------------------- + +#[test] +fn certify_oauth_credentials_are_never_proxied_as_api_key() { + // 字面量命中:判定期即可知,必须 DirectOnly 并给出结构化理由—— + // 而不是宣称可代理再发出错误的认证形态。 + let literal = composed_catalog(json!({ + "api": "anthropic-messages", + "baseUrl": "https://anthropic.example", + "apiKey": "sk-ant-oat01-example-token", + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&literal); + assert_eq!( + gateway.capability, + PiGatewayCapability::DirectOnly, + "pinned Pi sends an sk-ant-oat credential as an OAuth Bearer with oauth beta \ + headers; proxying it as x-api-key would send the wrong auth form" + ); + assert!( + !gateway.reasons.is_empty(), + "the downgrade must carry a structured reason" + ); + assert!( + !has_gateway_reason(&gateway, PiGatewayReasonCode::MissingCredential), + "the credential is present; MissingCredential would misreport the cause" + ); + + // deferred 凭证:判定期不可知,允许 Proxyable;但物化解析出命中值时必须 + // 失败,绝不发出错误的认证形态。 + let deferred = composed_catalog(json!({ + "api": "anthropic-messages", + "baseUrl": "https://anthropic.example", + "apiKey": "!load-token", + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&deferred); + assert_eq!( + gateway.capability, + PiGatewayCapability::Proxyable, + "a deferred credential's kind is unknowable at plan time" + ); + assert!( + gateway.plans[0] + .materialize(&|_: &str| Some("sk-ant-oat01-resolved".to_string())) + .is_err(), + "materialising a resolved OAuth credential must fail rather than send it as \ + a plain api key" + ); + + // 防过度收窄:普通 Anthropic key 不受影响;非 Anthropic 族不适用该判定 + // (pinned 的 includes 检查只在 Anthropic 传输层)。 + for (api, key) in [ + ("anthropic-messages", "sk-ant-api03-plain"), + ("openai-responses", "sk-ant-oat01-not-anthropic"), + ] { + let plain = composed_catalog(json!({ + "api": api, + "baseUrl": "https://plain.example/v1", + "apiKey": key, + "models": [{"id": "m"}] + })); + assert_eq!( + assess_composition(&plain).capability, + PiGatewayCapability::Proxyable, + "{api}: the OAuth rule must not over-reach" + ); + } +} + +// --------------------------------------------------------------------------- +// C4 扩展:字面量凭证必须在判定期校验 +// --------------------------------------------------------------------------- + +#[test] +fn certify_literal_credentials_are_validated_at_plan_time() { + // 判定期宣称"可代理"、而每次物化必然失败,是判定层与执行层自相矛盾。 + let illegal = composed_catalog(json!({ + "api": "openai-responses", + "baseUrl": "https://openai.example/v1", + "apiKey": "café", + "models": [{"id": "m"}] + })); + let gateway = assess_composition(&illegal); + assert!( + gateway.plans.is_empty() || gateway.plans[0].materialize(&|_: &str| None).is_err(), + "sanity: this literal credential can never materialise" + ); + assert_eq!( + gateway.capability, + PiGatewayCapability::DirectOnly, + "a literal credential that can never materialise must not be judged proxyable" + ); + assert!( + !gateway.reasons.is_empty(), + "the downgrade must carry a structured reason" + ); + + // 对称约束:deferred 凭证仍不得因原始表达式在判定期被拒(C4)。 + let deferred = composed_catalog(json!({ + "api": "openai-responses", + "baseUrl": "https://openai.example/v1", + "apiKey": "!echo café", + "models": [{"id": "m"}] + })); + assert_eq!( + assess_composition(&deferred).capability, + PiGatewayCapability::Proxyable, + "a deferred credential must not be validated as a header value before it is \ + resolved" + ); +} + +// --------------------------------------------------------------------------- +// C1 扩展:空容器与缺席必须可区分 +// --------------------------------------------------------------------------- + +#[test] +fn certify_empty_containers_stay_distinct_from_absent() { + let base_rates = json!({ + "input": 1.0, "output": 2.0, "cacheRead": 0.5, "cacheWrite": 0.25 + }); + let mut with_empty = base_rates.as_object().expect("rates").clone(); + with_empty.insert("tiers".into(), json!([])); + let catalog = json!({ + "providers": { + "tiers": { + "api": "anthropic-messages", + "baseUrl": "https://tiers.example", + "apiKey": "literal", + "models": [ + {"id": "empty-tiers", "cost": Value::Object(with_empty)}, + {"id": "absent-tiers", "cost": base_rates.clone()} + ] + } + } + }); + let (_temp, path) = write_catalog(&catalog); + let managed = inspect_pi_native_entry(&path, "tiers", &BTreeMap::new()) + .expect("inspect") + .expect("entry present") + .managed_config + .expect("managed config"); + let round_trip = serde_json::to_value(&managed).expect("serialize managed config"); + assert_eq!( + round_trip.pointer("/models/0/cost/tiers"), + Some(&json!([])), + "an explicitly empty tiers list must survive as an empty list" + ); + assert_eq!( + round_trip.pointer("/models/1/cost/tiers"), + None, + "an absent tiers list must stay absent" + ); +} + +// --------------------------------------------------------------------------- +// C6:单个 entry 的错误不得连坐兄弟 entry +// --------------------------------------------------------------------------- + +#[test] +fn certify_one_bad_entry_does_not_hide_its_siblings() { + // pinned Pi 逐 entry 判定:`contextWindow: 1e400` 只令该 entry 非法。 + // 整文件解析失败会让合法条目一并消失,破坏"每个 entry 独立"的判定设计。 + let source = r#"{ + "providers": { + "healthy": { + "api": "anthropic-messages", + "baseUrl": "https://healthy.example", + "apiKey": "literal", + "models": [{"id": "m"}] + }, + "overflow": { + "api": "anthropic-messages", + "baseUrl": "https://overflow.example", + "apiKey": "literal", + "models": [{"id": "m", "contextWindow": 1e400}] + } + } +}"#; + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + fs::write(&path, source).expect("write"); + + let diagnostics = inspect_pi_native_catalog(&path, &BTreeMap::new()) + .expect("one malformed entry must not fail the whole catalog"); + assert_eq!(diagnostics.len(), 2, "both entries must still be reported"); + let healthy = diagnostics + .iter() + .find(|diagnostic| diagnostic.provider_key == "healthy") + .expect("healthy entry present"); + assert_eq!( + healthy.raw_validity, + PiRawNativeValidity::Valid, + "a legal sibling must not be hidden by a malformed entry" + ); + assert_eq!(healthy.management_status, PiManagementStatus::Importable); + let overflow = diagnostics + .iter() + .find(|diagnostic| diagnostic.provider_key == "overflow") + .expect("overflow entry present"); + assert_ne!( + overflow.raw_validity, + PiRawNativeValidity::Valid, + "the out-of-range contextWindow entry itself must not be judged valid" + ); +} diff --git a/src-tauri/src/pi_config/native_settings.rs b/src-tauri/src/pi_config/native_settings.rs new file mode 100644 index 000000000..91b9a4644 --- /dev/null +++ b/src-tauri/src/pi_config/native_settings.rs @@ -0,0 +1,324 @@ +//! Exact-field access to Pi's shared `settings.json`. +//! +//! cc-switch owns only `defaultProvider` and `defaultModel`. Every other field +//! remains Pi/user-owned and survives each mutation unchanged. + +use crate::error::AppError; +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::sync::{LazyLock, Mutex}; + +use super::shared_file::{ + compare_exchange_shared_file_bytes, delete_shared_file, read_shared_file, replace_shared_file, + SharedFileSnapshot, +}; + +const MAX_PI_SETTINGS_BYTES: u64 = 1024 * 1024; +const MAX_WRITE_ATTEMPTS: usize = 3; +static SETTINGS_WRITE_LOCK: LazyLock> = LazyLock::new(|| Mutex::new(())); + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiNativeDefaults { + #[serde(skip_serializing_if = "Option::is_none")] + pub default_provider: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub default_model: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub session_dir: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PiNativeDefaultsReceipt { + path: PathBuf, + before: SharedFileSnapshot, + after: SharedFileSnapshot, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum PiNativeDefaultsRollback { + Restored, + Superseded, +} + +impl PiNativeDefaultsReceipt { + /// Restore the exact file revision replaced by this write. A newer Pi/user + /// edit wins and is reported as Superseded rather than being overwritten. + pub(crate) fn rollback(&self) -> Result { + let result = match self.before.bytes.as_deref() { + Some(bytes) => replace_shared_file( + &self.path, + &self.after.revision, + bytes, + MAX_PI_SETTINGS_BYTES, + None, + "Pi settings rollback", + ) + .map(|_| ()), + None => delete_shared_file( + &self.path, + &self.after.revision, + MAX_PI_SETTINGS_BYTES, + "Pi settings rollback", + ) + .map(|_| ()), + }; + match result { + Ok(()) => Ok(PiNativeDefaultsRollback::Restored), + Err(AppError::Conflict(_)) => Ok(PiNativeDefaultsRollback::Superseded), + Err(error) => Err(error), + } + } +} + +pub(crate) fn get_pi_settings_path() -> Result { + Ok(super::native::get_pi_agent_dir()?.join("settings.json")) +} + +pub(crate) fn read_pi_native_defaults() -> Result { + read_pi_native_defaults_at(&get_pi_settings_path()?) +} + +pub(crate) fn read_pi_native_defaults_at(path: &Path) -> Result { + let document = read_settings_document(path)?; + let root = document.as_object().ok_or_else(|| { + AppError::Config(format!( + "Pi settings root must be an object: {}", + path.display() + )) + })?; + Ok(PiNativeDefaults { + default_provider: optional_string(root, "defaultProvider", path)?, + default_model: optional_string(root, "defaultModel", path)?, + session_dir: optional_string(root, "sessionDir", path)?, + }) +} + +pub(crate) fn set_pi_native_default_with_receipt( + provider_key: &str, + model_id: &str, +) -> Result { + if provider_key.trim().is_empty() || model_id.trim().is_empty() { + return Err(AppError::InvalidInput( + "Pi default provider and model must be non-empty".to_string(), + )); + } + mutate_settings_document(&get_pi_settings_path()?, |root| { + root.insert( + "defaultProvider".to_string(), + Value::String(provider_key.to_string()), + ); + root.insert( + "defaultModel".to_string(), + Value::String(model_id.to_string()), + ); + Ok(()) + }) +} + +fn optional_string( + root: &Map, + key: &str, + path: &Path, +) -> Result, AppError> { + match root.get(key) { + None | Some(Value::Null) => Ok(None), + Some(Value::String(value)) => Ok(Some(value.clone())), + Some(_) => Err(AppError::Config(format!( + "Pi settings field '{key}' must be a string: {}", + path.display() + ))), + } +} + +fn mutate_settings_document( + path: &Path, + mut mutator: impl FnMut(&mut Map) -> Result<(), AppError>, +) -> Result { + let _guard = SETTINGS_WRITE_LOCK + .lock() + .map_err(|error| AppError::Config(format!("Pi settings lock is poisoned: {error}")))?; + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).map_err(|error| AppError::io(parent, error))?; + } + + for _ in 0..MAX_WRITE_ATTEMPTS { + let before = read_shared_file(path, MAX_PI_SETTINGS_BYTES, "Pi settings")?; + let mut document = match before.bytes.as_deref() { + Some(bytes) => { + serde_json::from_slice(bytes).map_err(|error| AppError::json(path, error))? + } + None => Value::Object(Map::new()), + }; + let root = document.as_object_mut().ok_or_else(|| { + AppError::Config(format!( + "Pi settings root must be an object: {}", + path.display() + )) + })?; + mutator(root)?; + let mut serialized = serde_json::to_vec_pretty(&document) + .map_err(|source| AppError::JsonSerialize { source })?; + serialized.push(b'\n'); + + match compare_exchange_shared_file_bytes( + path, + before.bytes.as_deref(), + &serialized, + MAX_PI_SETTINGS_BYTES, + None, + "Pi settings", + ) { + Ok(after) => { + return Ok(PiNativeDefaultsReceipt { + path: path.to_path_buf(), + before, + after, + }) + } + Err(AppError::Conflict(_)) => continue, + Err(error) => return Err(error), + } + } + + Err(AppError::Conflict(format!( + "Pi settings changed concurrently too many times: {}", + path.display() + ))) +} + +fn read_settings_document(path: &Path) -> Result { + match read_shared_file(path, MAX_PI_SETTINGS_BYTES, "Pi settings")?.bytes { + Some(bytes) => serde_json::from_slice(&bytes).map_err(|error| AppError::json(path, error)), + None => Ok(Value::Object(Map::new())), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn default_patch_preserves_every_unowned_field() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("settings.json"); + fs::write( + &path, + serde_json::to_vec_pretty(&json!({ + "theme": "custom", + "packages": ["npm:foreign"], + "sessionDir": "/tmp/pi-sessions", + "defaultProvider": "old", + "defaultModel": "old-model" + })) + .expect("serialize"), + ) + .expect("write"); + + mutate_settings_document(&path, |root| { + root.insert("defaultProvider".into(), json!("managed")); + root.insert("defaultModel".into(), json!("model")); + Ok(()) + }) + .expect("mutate"); + + let saved: Value = serde_json::from_slice(&fs::read(&path).expect("read")).expect("parse"); + assert_eq!(saved["theme"], "custom"); + assert_eq!(saved["packages"], json!(["npm:foreign"])); + assert_eq!(saved["sessionDir"], "/tmp/pi-sessions"); + assert_eq!(saved["defaultProvider"], "managed"); + assert_eq!(saved["defaultModel"], "model"); + } + + #[test] + fn external_rename_during_settings_patch_is_reparsed_before_retry() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("settings.json"); + fs::write( + &path, + br#"{"theme":"before","defaultProvider":"old","defaultModel":"old"}"#, + ) + .expect("seed"); + crate::pi_config::shared_file::replace_before_next_compare_exchange( + &path, + br#"{"theme":"external","packages":["foreign"],"defaultProvider":"old","defaultModel":"old"}"#, + ); + + mutate_settings_document(&path, |root| { + root.insert("defaultProvider".into(), json!("managed")); + root.insert("defaultModel".into(), json!("model")); + Ok(()) + }) + .expect("retry mutation"); + + let saved: Value = serde_json::from_slice(&fs::read(&path).expect("read")).expect("parse"); + assert_eq!(saved["theme"], "external"); + assert_eq!(saved["packages"], json!(["foreign"])); + assert_eq!(saved["defaultProvider"], "managed"); + assert_eq!(saved["defaultModel"], "model"); + } + + #[cfg(unix)] + #[test] + fn settings_symlink_is_rejected() { + use std::os::unix::fs::symlink; + let temp = tempfile::tempdir().expect("tempdir"); + let target = temp.path().join("target.json"); + let path = temp.path().join("settings.json"); + fs::write(&target, "{}").expect("target"); + symlink(&target, &path).expect("symlink"); + assert!(read_pi_native_defaults_at(&path).is_err()); + } + + #[test] + fn rollback_receipt_never_overwrites_a_newer_external_default() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("settings.json"); + fs::write( + &path, + serde_json::to_vec_pretty(&json!({ + "theme": "before", + "defaultProvider": "old", + "defaultModel": "old-model" + })) + .expect("serialize"), + ) + .expect("write"); + + let receipt = mutate_settings_document(&path, |root| { + root.insert("defaultProvider".into(), json!("attempted")); + root.insert("defaultModel".into(), json!("attempted-model")); + Ok(()) + }) + .expect("write attempted defaults"); + fs::write( + &path, + serde_json::to_vec_pretty(&json!({ + "theme": "external", + "defaultProvider": "external", + "defaultModel": "external-model" + })) + .expect("serialize external"), + ) + .expect("external write"); + + assert_eq!( + receipt.rollback().expect("rollback decision"), + PiNativeDefaultsRollback::Superseded + ); + assert_eq!( + read_pi_native_defaults_at(&path) + .expect("live defaults") + .default_provider + .as_deref(), + Some("external") + ); + assert_eq!( + read_settings_document(&path).expect("live document")["theme"], + "external" + ); + } +} diff --git a/src-tauri/src/pi_config/raw_schema.rs b/src-tauri/src/pi_config/raw_schema.rs new file mode 100644 index 000000000..cd70a7d23 --- /dev/null +++ b/src-tauri/src/pi_config/raw_schema.rs @@ -0,0 +1,1042 @@ +//! Lossless raw Pi catalog types and pinned TypeBox schema evaluation. +//! +//! This module deliberately imports neither managed control-plane types nor +//! gateway types. A raw-valid provider remains its original JSON value until +//! the independent managed assessor or composer explicitly consumes it. + +#![allow(dead_code)] + +use regex::Regex; +use serde_json::{Map, Value}; +use sha2::{Digest, Sha256}; +use std::collections::BTreeSet; +use std::sync::LazyLock; + +const PI_REPOSITORY: &str = "https://github.com/earendil-works/pi.git"; +const PI_COMMIT: &str = "ab366ebe94cacd419d986be454f12b1b9913aaca"; +const TYPEBOX_VERSION: &str = "1.3.7"; +const MODEL_CONFIG_PATH: &str = "packages/coding-agent/src/core/model-config.ts"; +const MODEL_CONFIG_SHA256: &str = + "62141770d675ad6357a72e07354355f0eda29281c0e5be1b48d2360f341c7360"; +const PROVIDER_COMPOSER_PATH: &str = "packages/coding-agent/src/core/provider-composer.ts"; +const PROVIDER_COMPOSER_SHA256: &str = + "17308a4179b330526eabf6c917fa13e9dbd9ece90d1555b870e87d39b5b60d9d"; +const RESOLVE_CONFIG_VALUE_PATH: &str = "packages/coding-agent/src/core/resolve-config-value.ts"; +const RESOLVE_CONFIG_VALUE_SHA256: &str = + "0f53dad47fe5d5d8837c022b7951ccd3bd5a9b577bd662f0986272110e83bcc7"; +const SCHEMA_SHA256: &str = "e498c9f1b344eee1bd3c3ba74d1b648dcb835378cfad92800ec80078b825745c"; +const RAW_ORACLE_SHA256: &str = "5aaa37160f96a0fe50867d900ca38c73f13aba769e156a883324368d9dbeeb9a"; +const COMPOSER_ORACLE_SHA256: &str = + "f7e54bb84e5fd6d50e5762dc304834410fa73ef608c2f9c42475c5983f8e0cf5"; +const TRANSPORT_ORACLE_SHA256: &str = + "b2c816e53b60da5cd6352d2c23934939e9f6dd0077971488fe9dd36fa723e855"; +const FIELD_COVERAGE_SHA256: &str = + "b8b85e611cf1dbef86c611df185ba8ac2d64160087d0c6e47747f838a0fafe42"; +const HARNESS_PATH: &str = "scripts/generate-pi-native-oracle.mjs"; +const HARNESS_SHA256: &str = "f7a138831284b48ef655ef500a63313f0fc89cf08d319094895027ba4777cc20"; +const EVALUATOR_OPERATOR_ALLOWLIST: &[&str] = &[ + "additionalProperties", + "anyOf", + "const", + "items", + "minLength", + "patternProperties", + "properties", + "required", + "type", +]; + +const SCHEMA_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/provider-schema.snapshot.json"); +const RAW_ORACLE_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/raw-oracle-v1.json"); +const COMPOSER_ORACLE_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/composer-oracle-v1.json"); +const TRANSPORT_ORACLE_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/transport-oracle-v1.json"); +const FIELD_COVERAGE_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/field-coverage-v1.json"); +const PROVENANCE_SOURCE: &str = + include_str!("../../../tests/fixtures/pi/native-oracle/provenance-v1.json"); +const GENERATOR_SOURCE: &str = include_str!("../../../scripts/generate-pi-native-oracle.mjs"); + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PiRawApiId(String); + +impl PiRawApiId { + pub(crate) fn new(value: impl Into) -> Option { + let value = value.into(); + (!value.is_empty()).then_some(Self(value)) + } + + pub(crate) fn as_str(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq)] +pub(super) struct PiRawValidProvider { + raw: Value, +} + +impl PiRawValidProvider { + fn new(raw: Value) -> Self { + Self { raw } + } + + pub(super) fn raw(&self) -> &Value { + &self.raw + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum PiRawValidity { + Valid, + Invalid, + Unknown, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum PiRawReasonCode { + SchemaMismatch, + UnsupportedOperator, + PinDrift, + AmbiguousSchema, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) struct PiRawReason { + pub code: PiRawReasonCode, + pub json_pointer: String, +} + +#[derive(Debug, Clone)] +pub(super) struct PiRawSchemaEvaluation { + pub validity: PiRawValidity, + pub valid_provider: Option, + pub reasons: Vec, +} + +#[derive(Debug)] +struct OracleBundle { + provider_schema: Value, +} + +#[derive(Debug, Clone, Copy)] +struct OracleSources<'a> { + schema: &'a str, + raw_oracle: &'a str, + composer_oracle: &'a str, + transport_oracle: &'a str, + field_coverage: &'a str, + provenance: &'a str, + generator: &'a str, +} + +static ORACLE_BUNDLE: LazyLock> = + LazyLock::new(load_and_verify_oracle_bundle); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum UnknownKind { + UnsupportedOperator, + AmbiguousSchema, +} + +#[derive(Debug, PartialEq, Eq)] +enum SchemaOutcome { + Valid, + Invalid(String), + Unknown { + kind: UnknownKind, + instance_pointer: String, + }, +} + +pub(super) fn evaluate_provider_value(value: &Value) -> PiRawSchemaEvaluation { + evaluate_provider_value_against(value, &ORACLE_BUNDLE) +} + +fn evaluate_provider_value_against( + value: &Value, + oracle_bundle: &Result, +) -> PiRawSchemaEvaluation { + let bundle = match oracle_bundle { + Ok(bundle) => bundle, + Err(_) => { + return PiRawSchemaEvaluation { + validity: PiRawValidity::Unknown, + valid_provider: None, + reasons: vec![PiRawReason { + code: PiRawReasonCode::PinDrift, + json_pointer: String::new(), + }], + }; + } + }; + outcome_to_evaluation(evaluate_schema(&bundle.provider_schema, value, ""), value) +} + +fn outcome_to_evaluation(outcome: SchemaOutcome, value: &Value) -> PiRawSchemaEvaluation { + match outcome { + SchemaOutcome::Valid => PiRawSchemaEvaluation { + validity: PiRawValidity::Valid, + valid_provider: Some(PiRawValidProvider::new(value.clone())), + reasons: Vec::new(), + }, + SchemaOutcome::Invalid(pointer) => PiRawSchemaEvaluation { + validity: PiRawValidity::Invalid, + valid_provider: None, + reasons: vec![PiRawReason { + code: PiRawReasonCode::SchemaMismatch, + json_pointer: pointer, + }], + }, + SchemaOutcome::Unknown { + kind, + instance_pointer, + } => PiRawSchemaEvaluation { + validity: PiRawValidity::Unknown, + valid_provider: None, + reasons: vec![PiRawReason { + code: match kind { + UnknownKind::UnsupportedOperator => PiRawReasonCode::UnsupportedOperator, + UnknownKind::AmbiguousSchema => PiRawReasonCode::AmbiguousSchema, + }, + json_pointer: instance_pointer, + }], + }, + } +} + +fn load_and_verify_oracle_bundle() -> Result { + load_and_verify_oracle_bundle_from(OracleSources { + schema: SCHEMA_SOURCE, + raw_oracle: RAW_ORACLE_SOURCE, + composer_oracle: COMPOSER_ORACLE_SOURCE, + transport_oracle: TRANSPORT_ORACLE_SOURCE, + field_coverage: FIELD_COVERAGE_SOURCE, + provenance: PROVENANCE_SOURCE, + generator: GENERATOR_SOURCE, + }) +} + +fn load_and_verify_oracle_bundle_from(sources: OracleSources<'_>) -> Result { + let provenance: Value = + serde_json::from_str(sources.provenance).map_err(|error| error.to_string())?; + if provenance.pointer("/version").and_then(Value::as_u64) != Some(1) { + return Err("provenance version does not match the pinned value".to_string()); + } + verify_string( + "/pi/repository", + provenance.pointer("/pi/repository"), + PI_REPOSITORY, + )?; + verify_string("/pi/commit", provenance.pointer("/pi/commit"), PI_COMMIT)?; + verify_string( + "/typeboxVersion", + provenance.pointer("/typeboxVersion"), + TYPEBOX_VERSION, + )?; + verify_string( + "/sources/modelConfig/path", + provenance.pointer("/sources/modelConfig/path"), + MODEL_CONFIG_PATH, + )?; + verify_string( + "/sources/modelConfig/sha256", + provenance.pointer("/sources/modelConfig/sha256"), + MODEL_CONFIG_SHA256, + )?; + verify_string( + "/sources/providerComposer/path", + provenance.pointer("/sources/providerComposer/path"), + PROVIDER_COMPOSER_PATH, + )?; + verify_string( + "/sources/providerComposer/sha256", + provenance.pointer("/sources/providerComposer/sha256"), + PROVIDER_COMPOSER_SHA256, + )?; + verify_string( + "/sources/resolveConfigValue/path", + provenance.pointer("/sources/resolveConfigValue/path"), + RESOLVE_CONFIG_VALUE_PATH, + )?; + verify_string( + "/sources/resolveConfigValue/sha256", + provenance.pointer("/sources/resolveConfigValue/sha256"), + RESOLVE_CONFIG_VALUE_SHA256, + )?; + for (filename, source, expected_hash) in [ + ( + "provider-schema.snapshot.json", + sources.schema, + SCHEMA_SHA256, + ), + ("raw-oracle-v1.json", sources.raw_oracle, RAW_ORACLE_SHA256), + ( + "composer-oracle-v1.json", + sources.composer_oracle, + COMPOSER_ORACLE_SHA256, + ), + ( + "transport-oracle-v1.json", + sources.transport_oracle, + TRANSPORT_ORACLE_SHA256, + ), + ( + "field-coverage-v1.json", + sources.field_coverage, + FIELD_COVERAGE_SHA256, + ), + ] { + if sha256_hex(source.as_bytes()) != expected_hash { + return Err(format!("{filename} does not match the code pin")); + } + let pointer = format!("/artifacts/{}", escape_json_pointer(filename)); + verify_string(&pointer, provenance.pointer(&pointer), expected_hash)?; + } + verify_string( + "/harness/path", + provenance.pointer("/harness/path"), + HARNESS_PATH, + )?; + if sha256_hex(sources.generator.as_bytes()) != HARNESS_SHA256 { + return Err("oracle execution harness does not match the code pin".to_string()); + } + verify_string( + "/harness/sha256", + provenance.pointer("/harness/sha256"), + HARNESS_SHA256, + )?; + verify_string( + "/harness/upstreamEntry", + provenance.pointer("/harness/upstreamEntry"), + PROVIDER_COMPOSER_PATH, + )?; + let entry_functions = provenance + .pointer("/harness/entryFunctions") + .and_then(Value::as_array) + .ok_or_else(|| "composer harness entry functions are missing".to_string())?; + for expected in [ + "composeModelProvider", + "Provider.getModels", + "resolveCompatibilityRequestConfig", + "Provider.auth.apiKey.resolve", + ] { + if !entry_functions + .iter() + .any(|value| value.as_str() == Some(expected)) + { + return Err(format!( + "composer harness does not record upstream entry function '{expected}'" + )); + } + } + let transport_functions = provenance + .pointer("/harness/transportResolver/entryFunctions") + .and_then(Value::as_array) + .ok_or_else(|| "transport resolver entry functions are missing".to_string())?; + for expected in ["resolveConfigValueOrThrow", "resolveHeadersOrThrow"] { + if !transport_functions + .iter() + .any(|value| value.as_str() == Some(expected)) + { + return Err(format!( + "transport harness does not record upstream entry function '{expected}'" + )); + } + } + + let allowlist = provenance + .pointer("/evaluatorOperatorAllowlist") + .and_then(Value::as_array) + .ok_or_else(|| "provenance evaluator allowlist is missing".to_string())? + .iter() + .map(|value| { + value + .as_str() + .map(ToOwned::to_owned) + .ok_or_else(|| "provenance evaluator allowlist is malformed".to_string()) + }) + .collect::, _>>()?; + if allowlist + != EVALUATOR_OPERATOR_ALLOWLIST + .iter() + .map(|value| (*value).to_string()) + .collect::>() + { + return Err("provenance evaluator allowlist drifted".to_string()); + } + + let schema: Value = serde_json::from_str(sources.schema).map_err(|error| error.to_string())?; + // The vendored snapshot is the provider schema itself. The harness records + // the extraction target used in the pinned ModelsConfigSchema. + let provider_schema = schema.clone(); + let mut inventory = BTreeSet::new(); + collect_schema_operators(&schema, &mut inventory)?; + let expected_inventory = provenance + .pointer("/schemaOperatorInventory") + .and_then(Value::as_array) + .ok_or_else(|| "provenance schema operator inventory is missing".to_string())? + .iter() + .map(|value| { + value + .as_str() + .map(ToOwned::to_owned) + .ok_or_else(|| "provenance schema operator inventory is malformed".to_string()) + }) + .collect::, _>>()?; + if inventory != expected_inventory { + return Err("schema operator inventory drifted".to_string()); + } + verify_field_coverage( + &schema, + sources.field_coverage, + sources.raw_oracle, + sources.composer_oracle, + )?; + Ok(OracleBundle { provider_schema }) +} + +fn verify_field_coverage( + schema: &Value, + source: &str, + raw_oracle_source: &str, + composer_oracle_source: &str, +) -> Result<(), String> { + let coverage: Value = serde_json::from_str(source).map_err(|error| error.to_string())?; + let raw_oracle: Value = + serde_json::from_str(raw_oracle_source).map_err(|error| error.to_string())?; + let composer_oracle: Value = + serde_json::from_str(composer_oracle_source).map_err(|error| error.to_string())?; + verify_string("/piCommit", coverage.pointer("/piCommit"), PI_COMMIT)?; + let entries = coverage + .pointer("/fields") + .and_then(Value::as_array) + .ok_or_else(|| "field coverage entries are missing".to_string())?; + let mut covered = BTreeSet::new(); + for entry in entries { + let field_path = entry + .get("fieldPath") + .and_then(Value::as_str) + .ok_or_else(|| "field coverage path is malformed".to_string())?; + if !covered.insert(field_path.to_string()) { + return Err(format!("field coverage duplicates '{field_path}'")); + } + let raw_cases = entry + .get("rawOracleCases") + .and_then(Value::as_array) + .filter(|cases| !cases.is_empty()) + .ok_or_else(|| { + format!("field coverage '{field_path}' has no rawOracleCases evidence") + })?; + let mut has_successful_raw_execution = false; + for case_id in raw_cases { + let case_id = case_id.as_str().ok_or_else(|| { + format!("field coverage '{field_path}' has a malformed raw case id") + })?; + let case = oracle_case(&raw_oracle, case_id).ok_or_else(|| { + format!("field coverage '{field_path}' cites unknown raw case '{case_id}'") + })?; + if !input_has_field_path(&case["input"], field_path) { + return Err(format!( + "raw case '{case_id}' does not contain covered field '{field_path}'" + )); + } + has_successful_raw_execution |= + case.get("expectedValid").and_then(Value::as_bool) == Some(true); + } + if !has_successful_raw_execution { + return Err(format!( + "field coverage '{field_path}' has no successful pinned TypeBox execution" + )); + } + + let composer_cases = entry + .get("composerOracleCases") + .and_then(Value::as_array) + .filter(|cases| !cases.is_empty()) + .ok_or_else(|| { + format!("field coverage '{field_path}' has no composerOracleCases evidence") + })?; + let expected_behavior_case = if field_path == "/models" + || field_path.starts_with("/models/") + { + "model-fields-executed" + } else if field_path == "/modelOverrides" || field_path.starts_with("/modelOverrides/") { + "override-fields-executed" + } else { + "provider-fields-inherited" + }; + verify_string( + &format!("{field_path}/composerBehaviorCase"), + entry.get("composerBehaviorCase"), + expected_behavior_case, + )?; + let mut has_own_layer_execution = false; + for case_id in composer_cases { + let case_id = case_id.as_str().ok_or_else(|| { + format!("field coverage '{field_path}' has a malformed composer case id") + })?; + let case = oracle_case(&composer_oracle, case_id).ok_or_else(|| { + format!("field coverage '{field_path}' cites unknown composer case '{case_id}'") + })?; + if !input_has_field_path(&case["input"], field_path) { + return Err(format!( + "composer case '{case_id}' does not contain covered field '{field_path}'" + )); + } + if case_id == expected_behavior_case + && case.pointer("/execution/status").and_then(Value::as_str) == Some("success") + && case.get("expected").is_some_and(|value| !value.is_null()) + { + has_own_layer_execution = true; + } + } + if !has_own_layer_execution { + return Err(format!( + "field coverage '{field_path}' has no successful pinned Pi execution at its own layer" + )); + } + } + let mut expected = BTreeSet::new(); + collect_schema_field_paths(schema, "", &mut expected)?; + if covered != expected { + let missing = expected.difference(&covered).cloned().collect::>(); + let stale = covered.difference(&expected).cloned().collect::>(); + return Err(format!( + "field coverage differs from pinned schema; missing={missing:?}, stale={stale:?}" + )); + } + Ok(()) +} + +fn oracle_case<'a>(oracle: &'a Value, id: &str) -> Option<&'a Value> { + oracle + .get("cases") + .and_then(Value::as_array)? + .iter() + .find(|case| case.get("id").and_then(Value::as_str) == Some(id)) +} + +fn input_has_field_path(value: &Value, field_path: &str) -> bool { + fn descend(value: &Value, segments: &[&str]) -> bool { + let Some((head, tail)) = segments.split_first() else { + return true; + }; + if *head == "*" { + return match value { + Value::Array(values) => values.iter().any(|value| descend(value, tail)), + Value::Object(values) => values.values().any(|value| descend(value, tail)), + _ => false, + }; + } + let decoded = head.replace("~1", "/").replace("~0", "~"); + value + .get(decoded.as_str()) + .is_some_and(|value| descend(value, tail)) + } + + let segments = field_path + .strip_prefix('/') + .unwrap_or(field_path) + .split('/') + .collect::>(); + descend(value, &segments) +} + +fn collect_schema_field_paths( + schema: &Value, + pointer: &str, + output: &mut BTreeSet, +) -> Result<(), String> { + let object = schema + .as_object() + .ok_or_else(|| format!("schema field node at '{pointer}' is not an object"))?; + if let Some(branches) = object.get("anyOf") { + for branch in branches + .as_array() + .ok_or_else(|| format!("schema anyOf at '{pointer}' is not an array"))? + { + collect_schema_field_paths(branch, pointer, output)?; + } + } + if let Some(properties) = object.get("properties") { + for (name, child) in properties + .as_object() + .ok_or_else(|| format!("schema properties at '{pointer}' are malformed"))? + { + let child_pointer = join_json_pointer(pointer, name); + output.insert(child_pointer.clone()); + collect_schema_field_paths(child, &child_pointer, output)?; + } + } + if let Some(patterns) = object.get("patternProperties") { + for child in patterns + .as_object() + .ok_or_else(|| format!("schema patterns at '{pointer}' are malformed"))? + .values() + { + let child_pointer = join_json_pointer(pointer, "*"); + output.insert(child_pointer.clone()); + collect_schema_field_paths(child, &child_pointer, output)?; + } + } + if let Some(items) = object.get("items") { + collect_schema_field_paths(items, &join_json_pointer(pointer, "*"), output)?; + } + Ok(()) +} + +fn verify_string(label: &str, actual: Option<&Value>, expected: &str) -> Result<(), String> { + if actual.and_then(Value::as_str) == Some(expected) { + Ok(()) + } else { + Err(format!("{label} does not match the pinned value")) + } +} + +fn sha256_hex(bytes: &[u8]) -> String { + format!("{:x}", Sha256::digest(bytes)) +} + +fn collect_schema_operators( + schema: &Value, + inventory: &mut BTreeSet, +) -> Result<(), String> { + let object = schema + .as_object() + .ok_or_else(|| "schema node is not an object".to_string())?; + for key in object.keys() { + inventory.insert(key.clone()); + } + for map_name in ["properties", "patternProperties"] { + if let Some(children) = object.get(map_name) { + let children = children + .as_object() + .ok_or_else(|| format!("{map_name} is not an object"))?; + for child in children.values() { + collect_schema_operators(child, inventory)?; + } + } + } + if let Some(branches) = object.get("anyOf") { + for branch in branches + .as_array() + .ok_or_else(|| "anyOf is not an array".to_string())? + { + collect_schema_operators(branch, inventory)?; + } + } + if let Some(items) = object.get("items") { + collect_schema_operators(items, inventory)?; + } + if let Some(Value::Object(_)) = object.get("additionalProperties") { + collect_schema_operators( + object + .get("additionalProperties") + .expect("checked additionalProperties"), + inventory, + )?; + } + Ok(()) +} + +fn evaluate_schema(schema: &Value, instance: &Value, pointer: &str) -> SchemaOutcome { + let Some(schema) = schema.as_object() else { + return ambiguous(pointer); + }; + if schema + .keys() + .any(|key| !EVALUATOR_OPERATOR_ALLOWLIST.contains(&key.as_str())) + { + return unsupported(pointer); + } + + if let Some(expected_type) = schema.get("type") { + let Some(expected_type) = expected_type.as_str() else { + return ambiguous(pointer); + }; + let matches = match expected_type { + "object" => instance.is_object(), + "array" => instance.is_array(), + "string" => instance.is_string(), + // TypeBox Number follows JavaScript Number semantics and rejects + // NaN/Infinity. With serde_json arbitrary precision, an overflow + // token remains a Number but has no finite f64 representation. + "number" => instance + .as_number() + .and_then(serde_json::Number::as_f64) + .is_some(), + "boolean" => instance.is_boolean(), + "null" => instance.is_null(), + _ => return unsupported(pointer), + }; + if !matches { + return SchemaOutcome::Invalid(pointer.to_string()); + } + } + + if let Some(constant) = schema.get("const") { + if !json_value_equal(instance, constant) { + return SchemaOutcome::Invalid(pointer.to_string()); + } + } + + if let Some(min_length) = schema.get("minLength") { + let Some(min_length) = min_length.as_u64() else { + return ambiguous(pointer); + }; + let Some(value) = instance.as_str() else { + return SchemaOutcome::Invalid(pointer.to_string()); + }; + if value.encode_utf16().count() < min_length as usize { + return SchemaOutcome::Invalid(pointer.to_string()); + } + } + + if let Some(branches) = schema.get("anyOf") { + let Some(branches) = branches.as_array() else { + return ambiguous(pointer); + }; + if branches.is_empty() { + return SchemaOutcome::Invalid(pointer.to_string()); + } + let mut first_unknown = None; + for branch in branches { + match evaluate_schema(branch, instance, pointer) { + SchemaOutcome::Valid => return SchemaOutcome::Valid, + SchemaOutcome::Invalid(_) => {} + unknown @ SchemaOutcome::Unknown { .. } => { + first_unknown.get_or_insert(unknown); + } + } + } + return first_unknown.unwrap_or_else(|| SchemaOutcome::Invalid(pointer.to_string())); + } + + if let Some(items) = schema.get("items") { + let Some(values) = instance.as_array() else { + return SchemaOutcome::Invalid(pointer.to_string()); + }; + for (index, value) in values.iter().enumerate() { + let child_pointer = join_json_pointer(pointer, &index.to_string()); + match evaluate_schema(items, value, &child_pointer) { + SchemaOutcome::Valid => {} + outcome => return outcome, + } + } + } + + if schema.contains_key("required") + || schema.contains_key("properties") + || schema.contains_key("patternProperties") + || schema.contains_key("additionalProperties") + { + let Some(object) = instance.as_object() else { + return SchemaOutcome::Invalid(pointer.to_string()); + }; + match evaluate_object(schema, object, pointer) { + SchemaOutcome::Valid => {} + outcome => return outcome, + } + } + + SchemaOutcome::Valid +} + +fn evaluate_object( + schema: &Map, + instance: &Map, + pointer: &str, +) -> SchemaOutcome { + if let Some(required) = schema.get("required") { + let Some(required) = required.as_array() else { + return ambiguous(pointer); + }; + for name in required { + let Some(name) = name.as_str() else { + return ambiguous(pointer); + }; + if !instance.contains_key(name) { + return SchemaOutcome::Invalid(join_json_pointer(pointer, name)); + } + } + } + + let properties = match schema.get("properties") { + None => None, + Some(Value::Object(properties)) => Some(properties), + Some(_) => return ambiguous(pointer), + }; + let pattern_properties = match schema.get("patternProperties") { + None => Vec::new(), + Some(Value::Object(patterns)) => { + let mut compiled = Vec::with_capacity(patterns.len()); + for (pattern, child_schema) in patterns { + let Ok(regex) = Regex::new(pattern) else { + return ambiguous(pointer); + }; + compiled.push((regex, child_schema)); + } + compiled + } + Some(_) => return ambiguous(pointer), + }; + + for (name, value) in instance { + let child_pointer = join_json_pointer(pointer, name); + let mut covered = false; + if let Some(child_schema) = properties.and_then(|properties| properties.get(name)) { + covered = true; + match evaluate_schema(child_schema, value, &child_pointer) { + SchemaOutcome::Valid => {} + outcome => return outcome, + } + } + for (pattern, child_schema) in &pattern_properties { + if pattern.is_match(name) { + covered = true; + match evaluate_schema(child_schema, value, &child_pointer) { + SchemaOutcome::Valid => {} + outcome => return outcome, + } + } + } + if !covered { + match schema.get("additionalProperties") { + None | Some(Value::Bool(true)) => {} + Some(Value::Bool(false)) => { + return SchemaOutcome::Invalid(child_pointer); + } + Some(child_schema @ Value::Object(_)) => { + match evaluate_schema(child_schema, value, &child_pointer) { + SchemaOutcome::Valid => {} + outcome => return outcome, + } + } + Some(_) => return ambiguous(pointer), + } + } + } + SchemaOutcome::Valid +} + +fn json_value_equal(left: &Value, right: &Value) -> bool { + match (left, right) { + (Value::Number(left), Value::Number(right)) => left.as_f64() == right.as_f64(), + _ => left == right, + } +} + +fn join_json_pointer(parent: &str, token: &str) -> String { + format!("{parent}/{}", escape_json_pointer(token)) +} + +fn escape_json_pointer(value: &str) -> String { + value.replace('~', "~0").replace('/', "~1") +} + +fn unsupported(pointer: &str) -> SchemaOutcome { + SchemaOutcome::Unknown { + kind: UnknownKind::UnsupportedOperator, + instance_pointer: pointer.to_string(), + } +} + +fn ambiguous(pointer: &str) -> SchemaOutcome { + SchemaOutcome::Unknown { + kind: UnknownKind::AmbiguousSchema, + instance_pointer: pointer.to_string(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde::Deserialize; + use serde_json::json; + + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct RawOracle { + cases: Vec, + } + + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct RawOracleCase { + id: String, + input: Value, + expected_valid: bool, + } + + #[test] + fn rust_evaluator_matches_typebox_1_3_7_oracle() { + assert!( + ORACLE_BUNDLE.is_ok(), + "vendored oracle bundle must verify: {:?}", + ORACLE_BUNDLE.as_ref().err() + ); + let oracle: RawOracle = serde_json::from_str(RAW_ORACLE_SOURCE).expect("parse raw oracle"); + for case in oracle.cases { + let actual = evaluate_provider_value(&case.input).validity; + let expected = if case.expected_valid { + PiRawValidity::Valid + } else { + PiRawValidity::Invalid + }; + assert_eq!(actual, expected, "oracle case '{}'", case.id); + let evaluated = evaluate_provider_value(&case.input); + assert_eq!( + evaluated + .valid_provider + .as_ref() + .map(PiRawValidProvider::raw), + case.expected_valid.then_some(&case.input), + "raw-valid type barrier case '{}'", + case.id + ); + } + } + + #[test] + fn every_schema_field_is_bound_to_real_successful_pi_executions() { + let schema: Value = serde_json::from_str(SCHEMA_SOURCE).expect("parse schema"); + verify_field_coverage( + &schema, + FIELD_COVERAGE_SOURCE, + RAW_ORACLE_SOURCE, + COMPOSER_ORACLE_SOURCE, + ) + .expect("all fields cite actual pinned Pi executions"); + + let invented_raw_case = + FIELD_COVERAGE_SOURCE.replacen("all-schema-fields-valid", "invented-raw-case", 1); + assert!( + verify_field_coverage( + &schema, + &invented_raw_case, + RAW_ORACLE_SOURCE, + COMPOSER_ORACLE_SOURCE, + ) + .is_err(), + "a field cannot be certified by a made-up execution id" + ); + + let invented_composer_case = FIELD_COVERAGE_SOURCE.replacen( + "combined-all-fields-precedence", + "invented-composer-case", + 1, + ); + assert!( + verify_field_coverage( + &schema, + &invented_composer_case, + RAW_ORACLE_SOURCE, + COMPOSER_ORACLE_SOURCE, + ) + .is_err(), + "a field cannot cite a composer execution that did not occur" + ); + } + + #[test] + fn unsupported_operator_is_unknown_only_when_input_traverses_it() { + let schema = json!({ + "type": "object", + "properties": { + "optional": { + "type": "string", + "unevaluatedProperties": false + } + } + }); + assert_eq!( + evaluate_schema(&schema, &json!({}), ""), + SchemaOutcome::Valid + ); + assert!(matches!( + evaluate_schema(&schema, &json!({"optional": "value"}), ""), + SchemaOutcome::Unknown { + kind: UnknownKind::UnsupportedOperator, + instance_pointer + } if instance_pointer == "/optional" + )); + } + + #[test] + fn additional_properties_operator_is_explicitly_supported() { + let schema = json!({ + "type": "object", + "properties": {"known": {"type": "string"}}, + "additionalProperties": false + }); + assert_eq!( + evaluate_schema(&schema, &json!({"known": "yes"}), ""), + SchemaOutcome::Valid + ); + assert_eq!( + evaluate_schema(&schema, &json!({"unknown": true}), ""), + SchemaOutcome::Invalid("/unknown".into()) + ); + } + + #[test] + fn artifact_or_provenance_drift_fails_closed_as_raw_unknown() { + let tampered_schema = SCHEMA_SOURCE.replacen("\"baseUrl\"", "\"baseUrl-tampered\"", 1); + let schema_bundle = load_and_verify_oracle_bundle_from(OracleSources { + schema: &tampered_schema, + raw_oracle: RAW_ORACLE_SOURCE, + composer_oracle: COMPOSER_ORACLE_SOURCE, + transport_oracle: TRANSPORT_ORACLE_SOURCE, + field_coverage: FIELD_COVERAGE_SOURCE, + provenance: PROVENANCE_SOURCE, + generator: GENERATOR_SOURCE, + }); + assert!(schema_bundle.is_err()); + let schema_result = evaluate_provider_value_against(&json!({}), &schema_bundle); + assert_eq!(schema_result.validity, PiRawValidity::Unknown); + assert_eq!( + schema_result.reasons, + vec![PiRawReason { + code: PiRawReasonCode::PinDrift, + json_pointer: String::new(), + }] + ); + + let tampered_provenance = + PROVENANCE_SOURCE.replacen(PI_COMMIT, "0000000000000000000000000000000000000000", 1); + let provenance_bundle = load_and_verify_oracle_bundle_from(OracleSources { + schema: SCHEMA_SOURCE, + raw_oracle: RAW_ORACLE_SOURCE, + composer_oracle: COMPOSER_ORACLE_SOURCE, + transport_oracle: TRANSPORT_ORACLE_SOURCE, + field_coverage: FIELD_COVERAGE_SOURCE, + provenance: &tampered_provenance, + generator: GENERATOR_SOURCE, + }); + assert!(provenance_bundle.is_err()); + assert_eq!( + evaluate_provider_value_against(&json!({}), &provenance_bundle).validity, + PiRawValidity::Unknown + ); + + let tampered_coverage = FIELD_COVERAGE_SOURCE.replacen("\"/api\"", "\"/api-tampered\"", 1); + let coverage_bundle = load_and_verify_oracle_bundle_from(OracleSources { + schema: SCHEMA_SOURCE, + raw_oracle: RAW_ORACLE_SOURCE, + composer_oracle: COMPOSER_ORACLE_SOURCE, + transport_oracle: TRANSPORT_ORACLE_SOURCE, + field_coverage: &tampered_coverage, + provenance: PROVENANCE_SOURCE, + generator: GENERATOR_SOURCE, + }); + assert!(coverage_bundle.is_err()); + assert_eq!( + evaluate_provider_value_against(&json!({}), &coverage_bundle).validity, + PiRawValidity::Unknown + ); + } +} diff --git a/src-tauri/src/pi_config/shared_file.rs b/src-tauri/src/pi_config/shared_file.rs new file mode 100644 index 000000000..362f404ce --- /dev/null +++ b/src-tauri/src/pi_config/shared_file.rs @@ -0,0 +1,1770 @@ +//! Reusable safety boundary for exact Pi-owned/shared files. +//! +//! Callers choose an exact path and size limit. This layer supplies bounded +//! regular-file reads, symlink rejection, optimistic revisions, per-path +//! process locking, and OS-backed compare/exchange replacement. The latter is +//! deliberately stronger than "read, compare, rename": the displaced path is +//! inspected after one atomic namespace operation, so an external Pi/user +//! rename in the commit window is restored instead of overwritten. + +use crate::error::AppError; +use sha2::{Digest, Sha256}; +use std::collections::HashMap; +use std::fs::{self, File, Metadata, OpenOptions}; +use std::io::{Read, Take, Write}; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, LazyLock, Mutex}; + +static FILE_LOCKS: LazyLock>>>> = + LazyLock::new(|| Mutex::new(HashMap::new())); + +#[cfg(test)] +type CompareExchangeHooks = HashMap>; +#[cfg(test)] +static BEFORE_COMPARE_EXCHANGE: LazyLock> = + LazyLock::new(|| Mutex::new(HashMap::new())); +#[cfg(test)] +static BEFORE_ROLLBACK_EXCHANGE: LazyLock> = + LazyLock::new(|| Mutex::new(HashMap::new())); +#[cfg(test)] +static FAIL_NEXT_ROLLBACK_RESTORE: LazyLock>> = + LazyLock::new(|| Mutex::new(HashMap::new())); +#[cfg(test)] +static FAIL_AFTER_ATOMIC_ROLLBACK_SWAP: LazyLock>> = + LazyLock::new(|| Mutex::new(HashMap::new())); +#[cfg(all(test, unix))] +static FAIL_NEXT_PARENT_SYNC: LazyLock>> = + LazyLock::new(|| Mutex::new(HashMap::new())); + +#[derive(Debug)] +struct StagedReplacement { + path: PathBuf, + identity: FileIdentity, +} + +/// Stable-enough namespace identity for conditional cleanup. +/// +/// Installed-file classification pairs it with exact bytes; cleanup at a +/// private random path may use identity alone. Windows obtains the value from +/// an open handle instead of unstable `MetadataExt` APIs, keeping the pinned +/// Rust toolchain buildable without weakening identity to timestamps or +/// content alone. +#[derive(Debug, Clone, PartialEq, Eq)] +struct FileIdentity { + #[cfg(unix)] + device: u64, + #[cfg(unix)] + inode: u64, + #[cfg(windows)] + volume: u32, + #[cfg(windows)] + index: u64, + #[cfg(not(any(unix, windows)))] + len: u64, + #[cfg(not(any(unix, windows)))] + modified: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct SharedFileSnapshot { + pub revision: String, + pub bytes: Option>, +} + +impl SharedFileSnapshot { + pub(crate) fn exists(&self) -> bool { + self.bytes.is_some() + } +} + +pub(crate) fn read_shared_file( + path: &Path, + max_bytes: u64, + label: &str, +) -> Result { + let bytes = read_regular_bytes(path, max_bytes, label)?; + Ok(SharedFileSnapshot { + revision: revision(bytes.as_deref()), + bytes, + }) +} + +pub(crate) fn replace_shared_file( + path: &Path, + expected_revision: &str, + bytes: &[u8], + max_bytes: u64, + new_file_mode: Option, + label: &str, +) -> Result { + if bytes.len() as u64 > max_bytes { + return Err(AppError::InvalidInput(format!( + "{label} exceeds the {max_bytes}-byte limit" + ))); + } + let lock = path_lock(path)?; + let _guard = lock + .lock() + .map_err(|error| AppError::Config(format!("Pi file lock is poisoned: {error}")))?; + let current = read_shared_file(path, max_bytes, label)?; + ensure_revision(path, expected_revision, ¤t.revision)?; + compare_exchange_under_lock( + path, + current.bytes.as_deref(), + Some(bytes), + max_bytes, + new_file_mode, + label, + ) +} + +pub(crate) fn delete_shared_file( + path: &Path, + expected_revision: &str, + max_bytes: u64, + label: &str, +) -> Result { + let lock = path_lock(path)?; + let _guard = lock + .lock() + .map_err(|error| AppError::Config(format!("Pi file lock is poisoned: {error}")))?; + let current = read_shared_file(path, max_bytes, label)?; + ensure_revision(path, expected_revision, ¤t.revision)?; + if !current.exists() { + return Ok(false); + } + compare_exchange_under_lock(path, current.bytes.as_deref(), None, max_bytes, None, label)?; + Ok(true) +} + +/// Atomically replace the exact bytes a caller parsed. +/// +/// This is the common commit primitive for shared Pi documents. Callers may +/// retry `Conflict` after reparsing, but other failures are fail-closed. The +/// target never gets replaced merely because a pre-rename fingerprint happened +/// to match. +pub(crate) fn compare_exchange_shared_file_bytes( + path: &Path, + expected: Option<&[u8]>, + replacement: &[u8], + max_bytes: u64, + new_file_mode: Option, + label: &str, +) -> Result { + if replacement.len() as u64 > max_bytes { + return Err(AppError::InvalidInput(format!( + "{label} exceeds the {max_bytes}-byte limit" + ))); + } + let lock = path_lock(path)?; + let _guard = lock + .lock() + .map_err(|error| AppError::Config(format!("Pi file lock is poisoned: {error}")))?; + compare_exchange_under_lock( + path, + expected, + Some(replacement), + max_bytes, + new_file_mode, + label, + ) +} + +/// Retry durability for a namespace state which was observed after a +/// compare/exchange returned an error. The caller is still responsible for +/// verifying the exact live value before treating the mutation as committed +/// or compensated. +pub(crate) fn sync_shared_file_parent(path: &Path) -> Result<(), AppError> { + let parent = path.parent().ok_or_else(|| { + AppError::InvalidInput(format!( + "Pi shared-file path has no parent: {}", + path.display() + )) + })?; + sync_parent(parent) +} + +fn compare_exchange_under_lock( + path: &Path, + expected: Option<&[u8]>, + replacement: Option<&[u8]>, + max_bytes: u64, + new_file_mode: Option, + label: &str, +) -> Result { + if replacement.is_some_and(|bytes| bytes.len() as u64 > max_bytes) { + return Err(AppError::InvalidInput(format!( + "{label} exceeds the {max_bytes}-byte limit" + ))); + } + let parent = path.parent().ok_or_else(|| { + AppError::InvalidInput(format!("{label} path has no parent: {}", path.display())) + })?; + fs::create_dir_all(parent).map_err(|error| AppError::io(parent, error))?; + + // This preflight rejects symlinks/non-regular files and avoids a namespace + // operation when the conflict is already visible. Correctness still rests + // on inspecting the displaced file after the atomic operation below. + let before = read_regular_bytes(path, max_bytes, label)?; + if before.as_deref() != expected { + return Err(concurrent_change(path, label)); + } + + let staged = replacement + .map(|bytes| stage_replacement(path, bytes, before.is_some(), new_file_mode)) + .transpose()?; + run_before_compare_exchange_hook(path)?; + + let result = match (expected, replacement, staged.as_ref()) { + (None, Some(bytes), Some(staged)) => match rename_noreplace(&staged.path, path) { + Ok(()) => { + match sync_parent(parent) { + Ok(()) => Ok(snapshot(Some(bytes))), + Err(publish_error) => { + match rollback_installed_file( + path, + &staged.identity, + bytes, + None, + parent, + max_bytes, + label, + ) { + Ok(InstalledRollback::Restored) => Err(publish_error), + Ok(InstalledRollback::Superseded) => Err(AppError::Config(format!( + "{label} create lost its durability barrier ({publish_error}); \ + a concurrent external state won and was preserved" + ))), + Err(rollback_error) => { + if path_is_installed( + path, + &staged.identity, + bytes, + max_bytes, + label, + ) && sync_parent(parent).is_ok() + { + // A rollback can fail for reasons unrelated to + // the now-visible canonical file. A successful + // second durability barrier makes the only + // honest outcome a committed success. + Ok(snapshot(Some(bytes))) + } else { + Err(ambiguous_publication( + label, + &publish_error, + &rollback_error, + )) + } + } + } + } + } + } + Err(error) if is_destination_exists(&error) => Err(concurrent_change(path, label)), + Err(error) => Err(rename_error("create", &staged.path, path, error)), + }, + (Some(expected), Some(replacement), Some(staged)) => { + replace_existing_if_equal(path, staged, expected, replacement, max_bytes, label) + } + (Some(expected), None, None) => delete_existing_if_equal(path, expected, max_bytes, label), + (None, None, None) => Ok(snapshot(None)), + _ => Err(AppError::Config( + "invalid Pi shared-file compare/exchange state".to_string(), + )), + }; + + if let Some(staged) = staged { + // Never use byte equality as writer identity: an external writer may + // independently publish the same bytes with different ownership. + // Staging cleanup is safe only while the original staged inode/file-id + // remains at this private path. + let _ = remove_file_if_identity(&staged.path, &staged.identity); + } + result +} + +fn replace_existing_if_equal( + path: &Path, + staged: &StagedReplacement, + expected: &[u8], + replacement: &[u8], + max_bytes: u64, + label: &str, +) -> Result { + let parent = path + .parent() + .expect("a path accepted by compare/exchange has a parent"); + let displaced = match install_over_existing(&staged.path, path) { + Ok(displaced) => displaced, + Err(error) + if matches!( + error.kind(), + std::io::ErrorKind::NotFound | std::io::ErrorKind::AlreadyExists + ) => + { + return Err(concurrent_change(path, label)); + } + Err(error) => return Err(rename_error("exchange", &staged.path, path, error)), + }; + let displaced_bytes = read_regular_bytes(&displaced, max_bytes, label); + if matches!( + displaced_bytes, + Ok(ref bytes) if bytes.as_deref() == Some(expected) + ) { + return match sync_parent(parent) { + Ok(()) => { + finalize_recovery_artifact(&displaced, parent, label); + Ok(snapshot(Some(replacement))) + } + Err(publish_error) => match rollback_installed_file( + path, + &staged.identity, + replacement, + Some(&displaced), + parent, + max_bytes, + label, + ) { + Ok(InstalledRollback::Restored) => Err(publish_error), + Ok(InstalledRollback::Superseded) => Err(AppError::Config(format!( + "{label} replacement lost its durability barrier ({publish_error}); \ + a concurrent external state won and was preserved" + ))), + Err(rollback_error) => { + if path_is_installed(path, &staged.identity, replacement, max_bytes, label) + && sync_parent(parent).is_ok() + { + finalize_recovery_artifact(&displaced, parent, label); + Ok(snapshot(Some(replacement))) + } else { + Err(ambiguous_publication( + label, + &publish_error, + &rollback_error, + )) + } + } + }, + }; + } + + match rollback_installed_file( + path, + &staged.identity, + replacement, + Some(&displaced), + parent, + max_bytes, + label, + ) { + Ok(InstalledRollback::Restored) => match displaced_bytes { + Ok(_) => Err(concurrent_change(path, label)), + Err(error) => Err(AppError::Conflict(format!( + "{label} became unsafe during atomic replacement and was restored: {error}" + ))), + }, + Ok(InstalledRollback::Superseded) => Err(AppError::Config(format!( + "{label} changed again during rollback; all external bytes were preserved \ + and require explicit recovery/reconciliation" + ))), + Err(error) => Err(AppError::Config(format!( + "{label} changed during atomic replacement and could not be restored safely; \ + the displaced bytes remain at {}: {error}", + displaced.display() + ))), + } +} + +fn delete_existing_if_equal( + path: &Path, + expected: &[u8], + max_bytes: u64, + label: &str, +) -> Result { + let quarantine = sibling_temp_path(path, "delete"); + rename_noreplace(path, &quarantine) + .map_err(|error| rename_error("quarantine", path, &quarantine, error))?; + let quarantined = read_regular_bytes(&quarantine, max_bytes, label); + if matches!( + quarantined, + Ok(ref bytes) if bytes.as_deref() == Some(expected) + ) { + let parent = path + .parent() + .expect("a path accepted by compare/exchange has a parent"); + return match sync_parent(parent) { + Ok(()) => { + finalize_recovery_artifact(&quarantine, parent, label); + Ok(snapshot(None)) + } + Err(publish_error) => { + match restore_quarantined_file(&quarantine, path, parent, label) { + Ok(InstalledRollback::Restored) => Err(publish_error), + Ok(InstalledRollback::Superseded) => Err(AppError::Config(format!( + "{label} delete lost its durability barrier ({publish_error}); \ + a concurrent external state won and was preserved" + ))), + Err(rollback_error) => { + if path_is_missing(path) && sync_parent(parent).is_ok() { + finalize_recovery_artifact(&quarantine, parent, label); + Ok(snapshot(None)) + } else { + Err(ambiguous_publication( + label, + &publish_error, + &rollback_error, + )) + } + } + } + } + }; + } + + let parent = path + .parent() + .expect("a path accepted by compare/exchange has a parent"); + match restore_quarantined_file(&quarantine, path, parent, label) { + Ok(InstalledRollback::Restored) => match quarantined { + Ok(_) => Err(concurrent_change(path, label)), + Err(error) => Err(AppError::Conflict(format!( + "{label} became unsafe during delete and was restored: {error}" + ))), + }, + Ok(InstalledRollback::Superseded) => Err(AppError::Config(format!( + "{label} changed again while a delete was being restored; the newer \ + canonical state and displaced bytes at {} were both preserved", + quarantine.display() + ))), + Err(error) => Err(AppError::Config(format!( + "{label} changed during delete and the displaced bytes remain at {}: {error}", + quarantine.display() + ))), + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum InstalledRollback { + Restored, + Superseded, +} + +/// Remove an installed replacement without ever identifying it by content alone. +/// +/// The canonical path is first moved to a private recovery name and its +/// inode/file-id and exact bytes are compared with the staged witness. If an +/// external writer won between the failed durability barrier and this rollback, +/// that writer is restored (or retained at a recovery path) instead of being +/// overwritten. +fn rollback_installed_file( + path: &Path, + installed_identity: &FileIdentity, + installed_bytes: &[u8], + displaced: Option<&Path>, + parent: &Path, + max_bytes: u64, + label: &str, +) -> Result { + if let Err(error) = run_before_rollback_restore_hook(path) { + // The hook models a transient namespace rollback failure. Retrying the + // actual no-replace operation is part of the production guarantee. + log::warn!("retrying {label} rollback after transient failure: {error}"); + } + if let Some(displaced) = displaced { + return rollback_replacement_atomically( + path, + installed_identity, + installed_bytes, + displaced, + parent, + max_bytes, + label, + ); + } + + // A failed create has no before-image to exchange back into place. Isolate + // the canonical entry and identify it by file-id plus exact bytes before + // deleting it. If another writer won, restore that writer instead. + let rejected = sibling_temp_path(path, "rejected"); + match rename_noreplace(path, &rejected) { + Ok(()) => {} + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + sync_parent(parent)?; + return Ok(InstalledRollback::Superseded); + } + Err(error) => { + return Err(rename_error( + "isolate failed publication", + path, + &rejected, + error, + )); + } + } + + let rejected_identity = + file_identity_from_path(&rejected).map_err(|error| AppError::io(&rejected, error))?; + let rejected_is_installed = installed_identity == &rejected_identity + && read_regular_bytes(&rejected, max_bytes, label) + .ok() + .flatten() + .as_deref() + == Some(installed_bytes); + if !rejected_is_installed { + let restored = match rename_noreplace(&rejected, path) { + Ok(()) => InstalledRollback::Superseded, + Err(error) if is_destination_exists(&error) => InstalledRollback::Superseded, + Err(error) => { + return Err(rename_error( + "restore external winner", + &rejected, + path, + error, + )); + } + }; + sync_parent(parent)?; + return Ok(restored); + } + + sync_parent(parent)?; + remove_file_if_identity(&rejected, installed_identity)?; + sync_parent(parent)?; + Ok(InstalledRollback::Restored) +} + +/// Restore a replacement with one canonical-preserving namespace operation. +/// +/// The displaced before-image becomes canonical atomically and the value which +/// occupied the canonical path moves to a private recovery path. If that value +/// is not our staged file, a second external writer won; atomically put it back +/// and retain the older external value as a recovery artifact. A crash at any +/// point leaves a real file at the canonical path rather than a gap between two +/// renames. +fn rollback_replacement_atomically( + path: &Path, + installed_identity: &FileIdentity, + installed_bytes: &[u8], + displaced: &Path, + parent: &Path, + max_bytes: u64, + label: &str, +) -> Result { + if path_is_missing(path) { + // An external delete superseded the failed publication. Do not + // resurrect the older displaced value; retain it for reconciliation. + sync_parent(parent)?; + return Ok(InstalledRollback::Superseded); + } + if !path_is_installed(path, installed_identity, installed_bytes, max_bytes, label) { + // A newer external value is already canonical. Leave it there; the + // older displaced value is already a recovery artifact. The identity + // check cannot be made atomic with an uncooperative writer, but it + // removes the broad two-exchange window from the normal supersession + // path. + sync_parent(parent)?; + return Ok(InstalledRollback::Superseded); + } + run_before_rollback_exchange_hook(path)?; + // Namespace identity, rather than readability or content type, is the + // rollback witness. A symlink/directory which appeared in the commit + // window is unsafe to consume, but it still belongs back at the canonical + // path instead of being overwritten by our staged regular file. + let displaced_identity = + file_identity_from_path(displaced).map_err(|error| AppError::io(displaced, error))?; + + let swapped_out = match restore_displaced_atomically(displaced, path) { + Ok(swapped_out) => swapped_out, + Err(error) + if error.kind() == std::io::ErrorKind::NotFound + && path_is_missing(path) + && !path_is_missing(displaced) => + { + sync_parent(parent)?; + return Ok(InstalledRollback::Superseded); + } + Err(error) => { + return Err(rename_error( + "atomically restore displaced file", + displaced, + path, + error, + )); + } + }; + // Make the canonical-preserving recovery durable before inspecting or + // cleaning either recovery name. A crash after this barrier may require + // reconciliation, but it cannot replay a canonical-path gap. + sync_parent(parent)?; + + #[cfg(test)] + fail_after_atomic_rollback_swap_for_test(path)?; + + let swapped_is_installed = path_is_installed( + &swapped_out, + installed_identity, + installed_bytes, + max_bytes, + label, + ); + if swapped_is_installed { + remove_file_if_identity(&swapped_out, installed_identity)?; + sync_parent(parent)?; + return Ok(if path_has_identity(path, &displaced_identity) { + InstalledRollback::Restored + } else { + // The rollback itself succeeded, then an external writer changed + // or deleted the canonical value. That newer state is authority. + InstalledRollback::Superseded + }); + } + + // The canonical path changed again after our original exchange. The first + // atomic restore placed the older displaced value at the canonical path + // and preserved the newer value at `swapped_out`. Put that newer value + // back with another atomic operation; even a crash between the two swaps + // leaves a valid external value at the canonical path. + if !path_has_identity(path, &displaced_identity) { + sync_parent(parent)?; + return Ok(InstalledRollback::Superseded); + } + let older_recovery = match restore_displaced_atomically(&swapped_out, path) { + Ok(older_recovery) => older_recovery, + Err(error) + if error.kind() == std::io::ErrorKind::NotFound + && path_is_missing(path) + && !path_is_missing(&swapped_out) => + { + sync_parent(parent)?; + return Ok(InstalledRollback::Superseded); + } + Err(error) => { + let _ = sync_parent(parent); + return Err(AppError::Config(format!( + "{label} preserved a newer external value at {} but could not atomically \ + restore it to {}: {error}", + swapped_out.display(), + path.display() + ))); + } + }; + sync_parent(parent)?; + log::warn!( + "{label} changed again during rollback; an external value remains canonical and another \ + external value is preserved at {}", + older_recovery.display() + ); + Ok(InstalledRollback::Superseded) +} + +fn restore_quarantined_file( + quarantine: &Path, + path: &Path, + parent: &Path, + label: &str, +) -> Result { + if let Err(error) = run_before_rollback_restore_hook(path) { + log::warn!("retrying {label} delete rollback after transient failure: {error}"); + } + match rename_noreplace(quarantine, path) { + Ok(()) => { + sync_parent(parent)?; + Ok(InstalledRollback::Restored) + } + Err(error) if is_destination_exists(&error) => { + sync_parent(parent)?; + Ok(InstalledRollback::Superseded) + } + Err(error) => Err(rename_error( + "restore quarantined file", + quarantine, + path, + error, + )), + } +} + +fn path_is_installed( + path: &Path, + expected_identity: &FileIdentity, + expected_bytes: &[u8], + max_bytes: u64, + label: &str, +) -> bool { + file_identity_from_path(path) + .ok() + .is_some_and(|actual| &actual == expected_identity) + && read_regular_bytes(path, max_bytes, label) + .ok() + .flatten() + .as_deref() + == Some(expected_bytes) +} + +fn path_is_missing(path: &Path) -> bool { + matches!( + fs::symlink_metadata(path), + Err(error) if error.kind() == std::io::ErrorKind::NotFound + ) +} + +fn path_has_identity(path: &Path, expected: &FileIdentity) -> bool { + file_identity_from_path(path) + .ok() + .is_some_and(|actual| &actual == expected) +} + +fn remove_file_if_identity(path: &Path, expected: &FileIdentity) -> Result { + let actual = match fs::symlink_metadata(path) { + Ok(actual) => actual, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false), + Err(error) => return Err(AppError::io(path, error)), + }; + if !actual.file_type().is_file() + || file_identity_from_path(path).map_err(|error| AppError::io(path, error))? != *expected + { + return Ok(false); + } + fs::remove_file(path).map_err(|error| AppError::io(path, error))?; + Ok(true) +} + +fn finalize_recovery_artifact(path: &Path, parent: &Path, label: &str) { + match fs::remove_file(path) { + Ok(()) => { + if sync_parent(parent).is_err() { + if let Err(error) = sync_parent(parent) { + // The canonical mutation already passed its own barrier. + // This cleanup barrier only governs whether a private + // recovery name can reappear after a crash, so it must not + // turn a committed operation into a false failure. + log::warn!( + "{label} committed, but recovery-artifact cleanup durability is uncertain at {}: {error}", + path.display() + ); + } + } + } + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => { + log::warn!( + "{label} committed, but its private recovery artifact remains at {}: {error}", + path.display() + ); + } + } +} + +fn ambiguous_publication( + label: &str, + publish_error: &AppError, + rollback_error: &AppError, +) -> AppError { + AppError::Config(format!( + "{label} lost its durability barrier ({publish_error}) and neither conditional rollback \ + nor a second durability barrier established a safe outcome: {rollback_error}" + )) +} + +fn stage_replacement( + path: &Path, + bytes: &[u8], + preserve_mode: bool, + new_file_mode: Option, +) -> Result { + #[cfg(not(unix))] + let _ = (preserve_mode, new_file_mode); + let staged = sibling_temp_path(path, "cas"); + let mut options = OpenOptions::new(); + options.create_new(true).write(true); + #[cfg(unix)] + let requested_mode = { + use std::os::unix::fs::{OpenOptionsExt, PermissionsExt}; + let mode = if preserve_mode { + Some( + fs::metadata(path) + .map(|metadata| metadata.permissions().mode()) + .unwrap_or_else(|_| new_file_mode.unwrap_or(0o666)), + ) + } else { + new_file_mode + }; + options.mode(mode.unwrap_or(0o666)); + mode.map(|mode| mode & 0o7777) + }; + let mut file = options + .open(&staged) + .map_err(|error| AppError::io(&staged, error))?; + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + // open(2) always applies umask, including when OpenOptionsExt::mode is + // used. Restore the exact requested bits before the durable commit. + if let Some(requested_mode) = requested_mode { + file.set_permissions(fs::Permissions::from_mode(requested_mode)) + .map_err(|error| AppError::io(&staged, error))?; + } + } + file.write_all(bytes) + .map_err(|error| AppError::io(&staged, error))?; + file.flush() + .and_then(|_| file.sync_all()) + .map_err(|error| AppError::io(&staged, error))?; + let identity = file_identity_from_file(&file).map_err(|error| AppError::io(&staged, error))?; + drop(file); + Ok(StagedReplacement { + path: staged, + identity, + }) +} + +fn sibling_temp_path(path: &Path, purpose: &str) -> PathBuf { + let parent = path + .parent() + .expect("a path accepted by compare/exchange has a parent"); + let name = path + .file_name() + .expect("a path accepted by compare/exchange has a file name") + .to_string_lossy(); + parent.join(format!( + ".{name}.{purpose}.{}", + uuid::Uuid::new_v4().simple() + )) +} + +fn snapshot(bytes: Option<&[u8]>) -> SharedFileSnapshot { + SharedFileSnapshot { + revision: revision(bytes), + bytes: bytes.map(ToOwned::to_owned), + } +} + +fn concurrent_change(path: &Path, label: &str) -> AppError { + AppError::Conflict(format!( + "{label} changed since it was read: {}", + path.display() + )) +} + +fn rename_error( + operation: &str, + source: &Path, + destination: &Path, + error: std::io::Error, +) -> AppError { + AppError::IoContext { + context: format!( + "Pi shared-file {operation} failed: {} -> {}", + source.display(), + destination.display() + ), + source: error, + } +} + +#[cfg(unix)] +fn sync_parent(parent: &Path) -> Result<(), AppError> { + #[cfg(test)] + if fail_parent_sync_for_test(parent)? { + return Err(AppError::io( + parent, + std::io::Error::other("injected Pi parent-directory sync failure"), + )); + } + File::open(parent) + .and_then(|directory| directory.sync_all()) + .map_err(|error| AppError::io(parent, error)) +} + +#[cfg(not(unix))] +fn sync_parent(_parent: &Path) -> Result<(), AppError> { + Ok(()) +} + +#[cfg(any(target_os = "linux", target_os = "macos"))] +fn c_path(path: &Path) -> std::io::Result { + use std::os::unix::ffi::OsStrExt; + std::ffi::CString::new(path.as_os_str().as_bytes()) + .map_err(|_| std::io::Error::new(std::io::ErrorKind::InvalidInput, "path contains NUL")) +} + +#[cfg(target_os = "linux")] +fn rename_noreplace(source: &Path, destination: &Path) -> std::io::Result<()> { + let source = c_path(source)?; + let destination = c_path(destination)?; + // SAFETY: both C strings remain alive and renameat2 performs one + // synchronous namespace operation. + let result = unsafe { + libc::renameat2( + libc::AT_FDCWD, + source.as_ptr(), + libc::AT_FDCWD, + destination.as_ptr(), + libc::RENAME_NOREPLACE, + ) + }; + if result == 0 { + Ok(()) + } else { + Err(std::io::Error::last_os_error()) + } +} + +#[cfg(target_os = "linux")] +fn exchange_paths(left: &Path, right: &Path) -> std::io::Result<()> { + let left = c_path(left)?; + let right = c_path(right)?; + // SAFETY: both C strings remain alive and renameat2 performs one + // synchronous namespace operation. + let result = unsafe { + libc::renameat2( + libc::AT_FDCWD, + left.as_ptr(), + libc::AT_FDCWD, + right.as_ptr(), + libc::RENAME_EXCHANGE, + ) + }; + if result == 0 { + Ok(()) + } else { + Err(std::io::Error::last_os_error()) + } +} + +#[cfg(target_os = "macos")] +fn rename_with_flags(source: &Path, destination: &Path, flags: u32) -> std::io::Result<()> { + let source = c_path(source)?; + let destination = c_path(destination)?; + // SAFETY: both C strings remain alive and renamex_np performs one + // synchronous namespace operation. + let result = unsafe { libc::renamex_np(source.as_ptr(), destination.as_ptr(), flags) }; + if result == 0 { + Ok(()) + } else { + Err(std::io::Error::last_os_error()) + } +} + +#[cfg(target_os = "macos")] +fn rename_noreplace(source: &Path, destination: &Path) -> std::io::Result<()> { + rename_with_flags(source, destination, libc::RENAME_EXCL) +} + +#[cfg(target_os = "macos")] +fn exchange_paths(left: &Path, right: &Path) -> std::io::Result<()> { + rename_with_flags(left, right, libc::RENAME_SWAP) +} + +#[cfg(windows)] +fn wide_path(path: &Path) -> Vec { + use std::os::windows::ffi::OsStrExt; + path.as_os_str().encode_wide().chain(Some(0)).collect() +} + +#[cfg(windows)] +fn rename_noreplace(source: &Path, destination: &Path) -> std::io::Result<()> { + use windows_sys::Win32::Storage::FileSystem::{MoveFileExW, MOVEFILE_WRITE_THROUGH}; + let source = wide_path(source); + let destination = wide_path(destination); + // SAFETY: both buffers are NUL-terminated and remain alive during the + // synchronous Win32 call. No REPLACE_EXISTING flag is supplied. + let result = unsafe { + MoveFileExW( + source.as_ptr(), + destination.as_ptr(), + MOVEFILE_WRITE_THROUGH, + ) + }; + if result != 0 { + Ok(()) + } else { + Err(std::io::Error::last_os_error()) + } +} + +#[cfg(windows)] +fn replace_file_with_backup(path: &Path, replacement: &Path, backup: &Path) -> std::io::Result<()> { + use windows_sys::Win32::Storage::FileSystem::{ReplaceFileW, REPLACEFILE_WRITE_THROUGH}; + let path_wide = wide_path(path); + let replacement_wide = wide_path(replacement); + let backup_wide = wide_path(backup); + // SAFETY: all buffers are NUL-terminated and remain alive during the + // synchronous Win32 call. + let result = unsafe { + ReplaceFileW( + path_wide.as_ptr(), + replacement_wide.as_ptr(), + backup_wide.as_ptr(), + REPLACEFILE_WRITE_THROUGH, + std::ptr::null(), + std::ptr::null(), + ) + }; + if result != 0 { + Ok(()) + } else { + let error = std::io::Error::last_os_error(); + if error.raw_os_error() == Some(1177) { + // ERROR_UNABLE_TO_MOVE_REPLACEMENT_2 is a documented partial + // success: `path` moved to `backup`, `replacement` kept its name, + // and the canonical name may be absent. Restore the backup without + // overwriting a concurrent writer before surfacing the failure. + return Err(recover_partial_replace_backup(path, backup, error)); + } + Err(error) + } +} + +#[cfg(windows)] +fn recover_partial_replace_backup( + path: &Path, + backup: &Path, + replace_error: std::io::Error, +) -> std::io::Error { + match rename_noreplace(backup, path) { + Ok(()) => std::io::Error::new( + replace_error.kind(), + format!("{replace_error}; restored the partial backup to the canonical path"), + ), + Err(recovery_error) if is_destination_exists(&recovery_error) => std::io::Error::new( + replace_error.kind(), + format!( + "{replace_error}; a concurrent canonical value won and the partial backup remains at {}", + backup.display() + ), + ), + Err(recovery_error) => std::io::Error::new( + replace_error.kind(), + format!( + "{replace_error}; the partial backup remains at {} and could not be restored: {recovery_error}", + backup.display() + ), + ), + } +} + +#[cfg(any(target_os = "linux", target_os = "macos"))] +fn install_over_existing(staged: &Path, path: &Path) -> std::io::Result { + exchange_paths(staged, path)?; + Ok(staged.to_path_buf()) +} + +#[cfg(any(target_os = "linux", target_os = "macos"))] +fn restore_displaced_atomically(displaced: &Path, path: &Path) -> std::io::Result { + exchange_paths(displaced, path)?; + Ok(displaced.to_path_buf()) +} + +#[cfg(windows)] +fn install_over_existing(staged: &Path, path: &Path) -> std::io::Result { + let backup = sibling_temp_path(path, "displaced"); + replace_file_with_backup(path, staged, &backup)?; + Ok(backup) +} + +#[cfg(windows)] +fn restore_displaced_atomically(displaced: &Path, path: &Path) -> std::io::Result { + let backup = sibling_temp_path(path, "rejected"); + replace_file_with_backup(path, displaced, &backup)?; + Ok(backup) +} + +#[cfg(not(any(target_os = "linux", target_os = "macos", windows)))] +fn rename_noreplace(_source: &Path, _destination: &Path) -> std::io::Result<()> { + Err(std::io::Error::new( + std::io::ErrorKind::Unsupported, + "atomic no-replace rename is unsupported on this platform", + )) +} + +#[cfg(not(any(target_os = "linux", target_os = "macos", windows)))] +fn install_over_existing(_staged: &Path, _path: &Path) -> std::io::Result { + Err(std::io::Error::new( + std::io::ErrorKind::Unsupported, + "atomic file exchange is unsupported on this platform", + )) +} + +#[cfg(not(any(target_os = "linux", target_os = "macos", windows)))] +fn restore_displaced_atomically(_displaced: &Path, _path: &Path) -> std::io::Result { + Err(std::io::Error::new( + std::io::ErrorKind::Unsupported, + "atomic file recovery is unsupported on this platform", + )) +} + +fn is_destination_exists(error: &std::io::Error) -> bool { + error.kind() == std::io::ErrorKind::AlreadyExists + || matches!(error.raw_os_error(), Some(libc::EEXIST)) +} + +/// Publish a staged file or directory without replacing a path created by a +/// concurrent writer. +pub(crate) fn publish_path_noreplace(source: &Path, destination: &Path) -> std::io::Result<()> { + rename_noreplace(source, destination) +} + +#[cfg(test)] +pub(crate) fn replace_before_next_compare_exchange(path: &Path, bytes: &[u8]) { + BEFORE_COMPARE_EXCHANGE + .lock() + .expect("compare/exchange test hook lock") + .insert(path.to_path_buf(), bytes.to_vec()); +} + +#[cfg(all(test, unix))] +pub(crate) fn fail_next_parent_sync_for_test(path: &Path) { + let parent = path + .parent() + .expect("a shared-file test path must have a parent") + .to_path_buf(); + let mut failures = FAIL_NEXT_PARENT_SYNC + .lock() + .expect("parent sync test hook lock"); + failures + .entry(parent) + .and_modify(|remaining| *remaining = remaining.saturating_add(1)) + .or_insert(1); +} + +#[cfg(all(test, unix))] +fn fail_parent_sync_for_test(parent: &Path) -> Result { + let mut failures = FAIL_NEXT_PARENT_SYNC + .lock() + .map_err(|error| AppError::Config(format!("Pi sync test hook is poisoned: {error}")))?; + let Some(remaining) = failures.get_mut(parent) else { + return Ok(false); + }; + *remaining = remaining.saturating_sub(1); + if *remaining == 0 { + failures.remove(parent); + } + Ok(true) +} + +#[cfg(test)] +fn run_file_replacement_hook( + hooks: &Mutex, + path: &Path, +) -> Result<(), AppError> { + let replacement = { + let mut hook = hooks + .lock() + .map_err(|error| AppError::Config(format!("Pi CAS test hook is poisoned: {error}")))?; + hook.remove(path) + }; + if let Some(bytes) = replacement { + crate::config::atomic_write_durable(path, &bytes, None)?; + } + Ok(()) +} + +#[cfg(test)] +fn run_before_compare_exchange_hook(path: &Path) -> Result<(), AppError> { + run_file_replacement_hook(&BEFORE_COMPARE_EXCHANGE, path) +} + +#[cfg(not(test))] +fn run_before_compare_exchange_hook(_path: &Path) -> Result<(), AppError> { + Ok(()) +} + +#[cfg(test)] +fn replace_before_next_rollback_exchange(path: &Path, bytes: &[u8]) { + BEFORE_ROLLBACK_EXCHANGE + .lock() + .expect("rollback exchange test hook lock") + .insert(path.to_path_buf(), bytes.to_vec()); +} + +#[cfg(test)] +pub(crate) fn fail_next_rollback_restore_for_test(path: &Path) { + let mut failures = FAIL_NEXT_ROLLBACK_RESTORE + .lock() + .expect("rollback restore test hook lock"); + failures + .entry(path.to_path_buf()) + .and_modify(|remaining| *remaining = remaining.saturating_add(1)) + .or_insert(1); +} + +#[cfg(test)] +fn fail_next_after_atomic_rollback_swap_for_test(path: &Path) { + let mut failures = FAIL_AFTER_ATOMIC_ROLLBACK_SWAP + .lock() + .expect("atomic rollback-swap test hook lock"); + failures + .entry(path.to_path_buf()) + .and_modify(|remaining| *remaining = remaining.saturating_add(1)) + .or_insert(1); +} + +#[cfg(test)] +fn fail_after_atomic_rollback_swap_for_test(path: &Path) -> Result<(), AppError> { + let mut failures = FAIL_AFTER_ATOMIC_ROLLBACK_SWAP.lock().map_err(|error| { + AppError::Config(format!("Pi atomic rollback-swap hook is poisoned: {error}")) + })?; + let Some(remaining) = failures.get_mut(path) else { + return Ok(()); + }; + *remaining = remaining.saturating_sub(1); + if *remaining == 0 { + failures.remove(path); + } + Err(AppError::Config( + "injected stop after canonical-preserving rollback swap".to_string(), + )) +} + +#[cfg(test)] +fn run_before_rollback_restore_hook(path: &Path) -> std::io::Result<()> { + let mut failures = FAIL_NEXT_ROLLBACK_RESTORE + .lock() + .map_err(|error| std::io::Error::other(error.to_string()))?; + let Some(remaining) = failures.get_mut(path) else { + return Ok(()); + }; + *remaining = remaining.saturating_sub(1); + if *remaining == 0 { + failures.remove(path); + } + Err(std::io::Error::other( + "injected Pi rollback-restore failure", + )) +} + +#[cfg(not(test))] +fn run_before_rollback_restore_hook(_path: &Path) -> std::io::Result<()> { + Ok(()) +} + +#[cfg(test)] +fn run_before_rollback_exchange_hook(path: &Path) -> Result<(), AppError> { + run_file_replacement_hook(&BEFORE_ROLLBACK_EXCHANGE, path) +} + +#[cfg(not(test))] +fn run_before_rollback_exchange_hook(_path: &Path) -> Result<(), AppError> { + Ok(()) +} + +fn ensure_revision(path: &Path, expected: &str, actual: &str) -> Result<(), AppError> { + if expected == actual { + Ok(()) + } else { + Err(AppError::Conflict(format!( + "Pi file changed since it was read: {}", + path.display() + ))) + } +} + +fn path_lock(path: &Path) -> Result>, AppError> { + let mut locks = FILE_LOCKS + .lock() + .map_err(|error| AppError::Config(format!("Pi file-lock registry is poisoned: {error}")))?; + Ok(locks + .entry(path.to_path_buf()) + .or_insert_with(|| Arc::new(Mutex::new(()))) + .clone()) +} + +fn revision(bytes: Option<&[u8]>) -> String { + bytes.map_or_else( + || "missing".to_string(), + |bytes| format!("sha256:{:x}", Sha256::digest(bytes)), + ) +} + +#[cfg(unix)] +fn open_read_only(path: &Path) -> std::io::Result { + use std::os::unix::fs::OpenOptionsExt; + OpenOptions::new() + .read(true) + .custom_flags(libc::O_NOFOLLOW) + .open(path) +} + +#[cfg(windows)] +fn open_read_only(path: &Path) -> std::io::Result { + use std::os::windows::fs::OpenOptionsExt; + use windows_sys::Win32::Storage::FileSystem::FILE_FLAG_OPEN_REPARSE_POINT; + OpenOptions::new() + .read(true) + .custom_flags(FILE_FLAG_OPEN_REPARSE_POINT) + .open(path) +} + +#[cfg(not(any(unix, windows)))] +fn open_read_only(path: &Path) -> std::io::Result { + OpenOptions::new().read(true).open(path) +} + +#[cfg(unix)] +fn file_identity_from_metadata(metadata: &Metadata) -> FileIdentity { + use std::os::unix::fs::MetadataExt; + FileIdentity { + device: metadata.dev(), + inode: metadata.ino(), + } +} + +#[cfg(unix)] +fn file_identity_from_file(file: &File) -> std::io::Result { + file.metadata() + .map(|metadata| file_identity_from_metadata(&metadata)) +} + +#[cfg(unix)] +fn file_identity_from_path(path: &Path) -> std::io::Result { + fs::symlink_metadata(path).map(|metadata| file_identity_from_metadata(&metadata)) +} + +#[cfg(windows)] +fn file_identity_from_file(file: &File) -> std::io::Result { + use std::mem::MaybeUninit; + use std::os::windows::io::AsRawHandle; + use windows_sys::Win32::Storage::FileSystem::{ + GetFileInformationByHandle, BY_HANDLE_FILE_INFORMATION, + }; + + let mut information = MaybeUninit::::zeroed(); + // SAFETY: `file` owns a valid handle for the duration of the synchronous + // call and `information` points to writable, correctly sized storage. + let result = + unsafe { GetFileInformationByHandle(file.as_raw_handle(), information.as_mut_ptr()) }; + if result == 0 { + return Err(std::io::Error::last_os_error()); + } + // SAFETY: a successful GetFileInformationByHandle initializes the entire + // BY_HANDLE_FILE_INFORMATION structure. + let information = unsafe { information.assume_init() }; + Ok(FileIdentity { + volume: information.dwVolumeSerialNumber, + index: ((information.nFileIndexHigh as u64) << 32) | information.nFileIndexLow as u64, + }) +} + +#[cfg(windows)] +fn file_identity_from_path(path: &Path) -> std::io::Result { + use std::os::windows::fs::OpenOptionsExt; + use windows_sys::Win32::Storage::FileSystem::{ + FILE_FLAG_BACKUP_SEMANTICS, FILE_FLAG_OPEN_REPARSE_POINT, + }; + + let file = OpenOptions::new() + .read(true) + .custom_flags(FILE_FLAG_BACKUP_SEMANTICS | FILE_FLAG_OPEN_REPARSE_POINT) + .open(path)?; + file_identity_from_file(&file) +} + +#[cfg(not(any(unix, windows)))] +fn file_identity_from_metadata(metadata: &Metadata) -> FileIdentity { + FileIdentity { + len: metadata.len(), + modified: metadata.modified().ok(), + } +} + +#[cfg(not(any(unix, windows)))] +fn file_identity_from_file(file: &File) -> std::io::Result { + file.metadata() + .map(|metadata| file_identity_from_metadata(&metadata)) +} + +#[cfg(not(any(unix, windows)))] +fn file_identity_from_path(path: &Path) -> std::io::Result { + fs::symlink_metadata(path).map(|metadata| file_identity_from_metadata(&metadata)) +} + +#[cfg(windows)] +fn is_reparse_point(metadata: &Metadata) -> bool { + use std::os::windows::fs::MetadataExt; + use windows_sys::Win32::Storage::FileSystem::FILE_ATTRIBUTE_REPARSE_POINT; + metadata.file_attributes() & FILE_ATTRIBUTE_REPARSE_POINT != 0 +} + +#[cfg(not(windows))] +fn is_reparse_point(_metadata: &Metadata) -> bool { + false +} + +fn read_limited( + mut reader: Take<&mut File>, + path: &Path, + max_bytes: u64, +) -> Result, AppError> { + let mut bytes = Vec::new(); + reader + .read_to_end(&mut bytes) + .map_err(|error| AppError::io(path, error))?; + if bytes.len() as u64 > max_bytes { + return Err(AppError::InvalidInput(format!( + "Pi file exceeds the {max_bytes}-byte limit: {}", + path.display() + ))); + } + Ok(bytes) +} + +fn read_regular_bytes( + path: &Path, + max_bytes: u64, + label: &str, +) -> Result>, AppError> { + let initial = match fs::symlink_metadata(path) { + Ok(metadata) => metadata, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(error) => return Err(AppError::io(path, error)), + }; + if !initial.file_type().is_file() || is_reparse_point(&initial) || initial.len() > max_bytes { + return Err(AppError::InvalidInput(format!( + "{label} must be a bounded regular file: {}", + path.display() + ))); + } + let mut file = open_read_only(path).map_err(|error| AppError::io(path, error))?; + let opened = file.metadata().map_err(|error| AppError::io(path, error))?; + let opened_identity = + file_identity_from_file(&file).map_err(|error| AppError::io(path, error))?; + let bytes = read_limited(Read::by_ref(&mut file).take(max_bytes + 1), path, max_bytes)?; + let completed = file.metadata().map_err(|error| AppError::io(path, error))?; + let completed_identity = + file_identity_from_file(&file).map_err(|error| AppError::io(path, error))?; + let current = fs::symlink_metadata(path).map_err(|error| AppError::io(path, error))?; + if !current.file_type().is_file() || is_reparse_point(¤t) { + return Err(AppError::Conflict(format!( + "{label} changed during read: {}", + path.display() + ))); + } + let current_identity = + file_identity_from_path(path).map_err(|error| AppError::io(path, error))?; + if !opened.file_type().is_file() + || is_reparse_point(&opened) + || !completed.file_type().is_file() + || is_reparse_point(&completed) + || opened_identity != completed_identity + || completed_identity != current_identity + || opened.len() != bytes.len() as u64 + || completed.len() != bytes.len() as u64 + || current.len() != bytes.len() as u64 + || opened.modified().ok() != completed.modified().ok() + { + return Err(AppError::Conflict(format!( + "{label} changed during read: {}", + path.display() + ))); + } + Ok(Some(bytes)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn compare_and_replace_distinguishes_missing_and_content_revisions() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + let missing = read_shared_file(&path, 1024, "test").expect("missing"); + assert_eq!(missing.revision, "missing"); + let written = replace_shared_file(&path, "missing", b"one", 1024, Some(0o600), "test") + .expect("create"); + assert!(written.revision.starts_with("sha256:")); + assert!(replace_shared_file(&path, "missing", b"two", 1024, None, "test").is_err()); + let replaced = replace_shared_file(&path, &written.revision, b"two", 1024, None, "test") + .expect("replace"); + assert_eq!(replaced.bytes.as_deref(), Some(b"two".as_slice())); + assert!(delete_shared_file(&path, &written.revision, 1024, "test").is_err()); + assert!(delete_shared_file(&path, &replaced.revision, 1024, "test").expect("delete")); + } + + #[test] + fn external_rename_in_the_commit_window_is_restored_without_data_loss() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + fs::write(&path, b"observed").expect("seed"); + let observed = read_shared_file(&path, 1024, "test").expect("snapshot"); + + replace_before_next_compare_exchange(&path, b"external replacement"); + let replace = replace_shared_file(&path, &observed.revision, b"ours", 1024, None, "test"); + assert!(matches!(replace, Err(AppError::Conflict(_)))); + assert_eq!( + fs::read(&path).expect("external bytes restored"), + b"external replacement" + ); + + let observed = read_shared_file(&path, 1024, "test").expect("snapshot"); + replace_before_next_compare_exchange(&path, b"external before delete"); + let delete = delete_shared_file(&path, &observed.revision, 1024, "test"); + assert!(matches!(delete, Err(AppError::Conflict(_)))); + assert_eq!( + fs::read(&path).expect("external bytes restored"), + b"external before delete" + ); + } + + #[test] + fn second_external_rename_surfaces_the_recovery_artifact() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + fs::write(&path, b"observed").expect("seed"); + let observed = read_shared_file(&path, 1024, "test").expect("snapshot"); + + replace_before_next_compare_exchange(&path, b"external-a"); + replace_before_next_rollback_exchange(&path, b"external-b"); + let error = replace_shared_file(&path, &observed.revision, b"ours", 1024, None, "test") + .expect_err("the second race needs explicit recovery"); + let message = error.to_string(); + assert!( + matches!(error, AppError::Config(_)), + "recovery conflicts must not be auto-retried: {message}" + ); + assert!(message.contains("explicit recovery")); + assert_eq!( + fs::read(&path).expect("newest external version restored"), + b"external-b" + ); + assert!( + fs::read_dir(temp.path()) + .expect("recovery directory") + .filter_map(Result::ok) + .any(|entry| fs::read(entry.path()).ok().as_deref() == Some(b"external-a")), + "the displaced external bytes must remain in a named recovery artifact" + ); + } + + #[test] + fn atomic_rollback_stop_keeps_a_canonical_external_file() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + fs::write(&path, b"observed").expect("seed"); + let observed = read_shared_file(&path, 1024, "test").expect("snapshot"); + + replace_before_next_compare_exchange(&path, b"external-a"); + replace_before_next_rollback_exchange(&path, b"external-b"); + fail_next_after_atomic_rollback_swap_for_test(&path); + replace_shared_file(&path, &observed.revision, b"ours", 1024, None, "test") + .expect_err("the injected stop interrupts rollback cleanup"); + + assert_eq!( + fs::read(&path).expect("canonical path remains present"), + b"external-a", + "the first atomic recovery step must never create a canonical-path gap" + ); + assert!( + fs::read_dir(temp.path()) + .expect("recovery directory") + .filter_map(Result::ok) + .any(|entry| fs::read(entry.path()).ok().as_deref() == Some(b"external-b")), + "the newer external value must remain recoverable if cleanup never runs" + ); + } + + #[test] + fn external_delete_wins_over_replacement_rollback() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + let staged = temp.path().join("staged"); + let displaced = temp.path().join("displaced"); + fs::write(&staged, b"ours").expect("staged witness"); + fs::write(&displaced, b"external-before").expect("displaced value"); + let identity = file_identity_from_path(&staged).expect("staged identity"); + + let outcome = rollback_replacement_atomically( + &path, + &identity, + b"ours", + &displaced, + temp.path(), + 1024, + "test", + ) + .expect("external deletion is a safe superseding state"); + + assert_eq!(outcome, InstalledRollback::Superseded); + assert!( + !path.exists(), + "rollback must not resurrect the deleted path" + ); + assert_eq!( + fs::read(&displaced).expect("older value retained"), + b"external-before" + ); + } + + #[test] + fn visible_external_winner_is_not_exchanged_during_rollback() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + let installed = temp.path().join("installed-witness"); + let displaced = temp.path().join("displaced"); + fs::write(&path, b"external-newer").expect("canonical external winner"); + fs::write(&installed, b"ours").expect("installed witness"); + fs::write(&displaced, b"external-older").expect("older displaced value"); + let identity = file_identity_from_path(&installed).expect("installed identity"); + + let outcome = rollback_replacement_atomically( + &path, + &identity, + b"ours", + &displaced, + temp.path(), + 1024, + "test", + ) + .expect("visible external authority is a safe superseding state"); + + assert_eq!(outcome, InstalledRollback::Superseded); + assert_eq!( + fs::read(&path).expect("newer value stays canonical"), + b"external-newer" + ); + assert_eq!( + fs::read(&displaced).expect("older value stays recoverable"), + b"external-older" + ); + } + + #[cfg(unix)] + #[test] + fn unsafe_displaced_entry_is_restored_by_identity_without_being_followed() { + use std::os::unix::fs::symlink; + + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + let displaced = temp.path().join("displaced"); + let target = temp.path().join("target"); + fs::write(&path, b"ours").expect("installed value"); + fs::write(&target, b"external-target").expect("external target"); + symlink(&target, &displaced).expect("unsafe displaced entry"); + let identity = file_identity_from_path(&path).expect("installed identity"); + + let outcome = rollback_replacement_atomically( + &path, + &identity, + b"ours", + &displaced, + temp.path(), + 1024, + "test", + ) + .expect("namespace identity is sufficient to restore an unsafe entry"); + + assert_eq!(outcome, InstalledRollback::Restored); + assert!( + fs::symlink_metadata(&path) + .expect("canonical entry") + .file_type() + .is_symlink(), + "the external namespace entry must be restored instead of consumed" + ); + assert_eq!(fs::read_link(&path).expect("symlink target"), target); + assert_eq!( + fs::read(&target).expect("target remains untouched"), + b"external-target" + ); + assert!(!displaced.exists(), "our rejected file must be removed"); + } + + #[cfg(windows)] + #[test] + fn windows_partial_replace_restores_backup_to_missing_canonical() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + let backup = temp.path().join("backup"); + fs::write(&backup, b"external-before").expect("partial backup"); + + let error = + recover_partial_replace_backup(&path, &backup, std::io::Error::from_raw_os_error(1177)); + + assert!(error.to_string().contains("restored")); + assert_eq!( + fs::read(&path).expect("canonical restored"), + b"external-before" + ); + assert!(!backup.exists()); + } + + #[cfg(unix)] + #[test] + fn create_sync_failure_returns_error_only_after_removing_its_file_identity() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + fail_next_parent_sync_for_test(&path); + + let error = replace_shared_file(&path, "missing", b"created", 1024, None, "test") + .expect_err("the injected durability failure remains visible"); + assert!(error.to_string().contains("injected")); + assert!( + !path.exists(), + "Err must not leave the attempted create live" + ); + } + + #[cfg(unix)] + #[test] + fn replace_sync_failure_returns_error_only_after_restoring_before_image() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + fs::write(&path, b"before").expect("seed"); + let before = read_shared_file(&path, 1024, "test").expect("snapshot"); + fail_next_parent_sync_for_test(&path); + + let error = replace_shared_file(&path, &before.revision, b"after", 1024, None, "test") + .expect_err("the injected durability failure remains visible"); + assert!(error.to_string().contains("injected")); + assert_eq!(fs::read(&path).expect("before restored"), b"before"); + } + + #[cfg(unix)] + #[test] + fn delete_sync_failure_returns_error_only_after_restoring_before_image() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + fs::write(&path, b"before").expect("seed"); + let before = read_shared_file(&path, 1024, "test").expect("snapshot"); + fail_next_parent_sync_for_test(&path); + + let error = delete_shared_file(&path, &before.revision, 1024, "test") + .expect_err("the injected durability failure remains visible"); + assert!(error.to_string().contains("injected")); + assert_eq!(fs::read(&path).expect("before restored"), b"before"); + } + + #[cfg(unix)] + #[test] + fn replacement_preserves_exact_existing_permissions_despite_umask() { + use std::os::unix::fs::PermissionsExt; + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("shared.md"); + fs::write(&path, b"before").expect("seed"); + fs::set_permissions(&path, fs::Permissions::from_mode(0o764)).expect("set mode"); + let before = read_shared_file(&path, 1024, "test").expect("snapshot"); + + replace_shared_file(&path, &before.revision, b"after", 1024, None, "test") + .expect("replace"); + assert_eq!( + fs::metadata(&path).expect("metadata").permissions().mode() & 0o7777, + 0o764 + ); + } + + #[cfg(unix)] + #[test] + fn symlink_targets_fail_closed() { + use std::os::unix::fs::symlink; + let temp = tempfile::tempdir().expect("tempdir"); + let target = temp.path().join("target"); + let path = temp.path().join("shared"); + fs::write(&target, b"secret").expect("target"); + symlink(&target, &path).expect("symlink"); + assert!(read_shared_file(&path, 1024, "test").is_err()); + assert!(replace_shared_file(&path, "missing", b"overwrite", 1024, None, "test").is_err()); + assert_eq!(fs::read(&target).expect("target remains"), b"secret"); + } +} diff --git a/src-tauri/src/prompt.rs b/src-tauri/src/prompt.rs index ce9b8b474..2bf080fea 100644 --- a/src-tauri/src/prompt.rs +++ b/src-tauri/src/prompt.rs @@ -1,6 +1,6 @@ use serde::{Deserialize, Serialize}; -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct Prompt { pub id: String, pub name: String, diff --git a/src-tauri/src/prompt_files.rs b/src-tauri/src/prompt_files.rs index c3cf6968d..565cce52d 100644 --- a/src-tauri/src/prompt_files.rs +++ b/src-tauri/src/prompt_files.rs @@ -26,6 +26,7 @@ pub fn prompt_file_path(app: &AppType) -> Result { AppType::OpenCode => get_opencode_dir(), AppType::OpenClaw => get_openclaw_dir(), AppType::Hermes => crate::hermes_config::get_hermes_dir(), + AppType::Pi => crate::pi_config::native::get_pi_agent_dir()?, AppType::ClaudeDesktop => unreachable!("handled above"), }; @@ -35,6 +36,7 @@ pub fn prompt_file_path(app: &AppType) -> Result { AppType::Gemini => "GEMINI.md", AppType::GrokBuild | AppType::OpenCode | AppType::OpenClaw => "AGENTS.md", AppType::Hermes => "SOUL.md", + AppType::Pi => "AGENTS.md", AppType::ClaudeDesktop => unreachable!("handled above"), }; diff --git a/src-tauri/src/provider.rs b/src-tauri/src/provider.rs index 5c9dd1bc6..fa764d533 100644 --- a/src-tauri/src/provider.rs +++ b/src-tauri/src/provider.rs @@ -346,6 +346,10 @@ impl Provider { str_at(settings.get("base_url")), str_at(settings.get("api_key")), ), + AppType::Pi => ( + str_at(settings.get("baseUrl")), + str_at(settings.get("apiKey")), + ), // OpenClaw (openclaw.json) flattens credentials at the top level, camelCase. AppType::OpenClaw => ( str_at(settings.get("baseUrl")), diff --git a/src-tauri/src/proxy/handler_config.rs b/src-tauri/src/proxy/handler_config.rs index 2e3855a67..b5d2c0ad5 100644 --- a/src-tauri/src/proxy/handler_config.rs +++ b/src-tauri/src/proxy/handler_config.rs @@ -4,6 +4,7 @@ use crate::app_config::AppType; use crate::proxy::usage::parser::TokenUsage; +use crate::proxy::usage::InputTokenSemantics; use serde_json::Value; /// 使用量解析器类型别名 @@ -31,6 +32,8 @@ pub struct UsageParserConfig { pub model_extractor: StreamModelExtractor, /// 流式 usage 事件预过滤器 pub stream_event_filter: Option, + /// Semantics of `TokenUsage::input_tokens` produced by these parsers. + pub input_token_semantics: InputTokenSemantics, /// 应用类型字符串(用于日志记录) pub app_type_str: &'static str, } @@ -141,6 +144,7 @@ pub const CLAUDE_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { response_parser: TokenUsage::from_claude_response, model_extractor: claude_model_extractor, stream_event_filter: Some(claude_stream_usage_event_filter), + input_token_semantics: InputTokenSemantics::FreshExcludesCache, app_type_str: "claude", }; @@ -150,6 +154,7 @@ pub const OPENAI_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { response_parser: TokenUsage::from_openai_response, model_extractor: openai_model_extractor, stream_event_filter: Some(openai_stream_usage_event_filter), + input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets, app_type_str: "codex", }; @@ -159,6 +164,7 @@ pub const CODEX_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { response_parser: TokenUsage::from_codex_response_auto, model_extractor: codex_auto_model_extractor, stream_event_filter: Some(codex_stream_usage_event_filter), + input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets, app_type_str: "codex", }; @@ -168,6 +174,7 @@ pub const GEMINI_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { response_parser: TokenUsage::from_gemini_response, model_extractor: gemini_model_extractor, stream_event_filter: Some(gemini_stream_usage_event_filter), + input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets, app_type_str: "gemini", }; diff --git a/src-tauri/src/proxy/handlers.rs b/src-tauri/src/proxy/handlers.rs index bb2606486..8981994c7 100644 --- a/src-tauri/src/proxy/handlers.rs +++ b/src-tauri/src/proxy/handlers.rs @@ -39,7 +39,7 @@ use super::{ server::ProxyState, sse::{strip_sse_field, take_sse_block}, types::*, - usage::parser::TokenUsage, + usage::{parser::TokenUsage, InputTokenSemantics}, ProxyError, }; use crate::app_config::AppType; @@ -338,6 +338,7 @@ async fn write_claude_usage_log(state: &ProxyState, log: ClaudeUsageLog) { &log.model, &log.request_model, &log.outbound_model, + InputTokenSemantics::FreshExcludesCache, log.usage, log.latency_ms, None, @@ -465,6 +466,7 @@ async fn handle_claude_transform( &model, &request_model, &outbound_model, + InputTokenSemantics::FreshExcludesCache, usage, latency_ms, first_token_ms, @@ -1133,6 +1135,7 @@ async fn handle_codex_responses_namespace_restore( &model, &request_model, &outbound_model, + InputTokenSemantics::TotalIncludesCacheBuckets, usage, latency_ms, None, @@ -1245,6 +1248,7 @@ async fn handle_codex_chat_to_responses_transform( &model, &request_model, &outbound_model, + InputTokenSemantics::TotalIncludesCacheBuckets, usage, latency_ms, first_token_ms, @@ -1366,6 +1370,7 @@ async fn handle_codex_chat_to_responses_transform( &model, &request_model, &outbound_model, + InputTokenSemantics::TotalIncludesCacheBuckets, usage, latency_ms, None, @@ -1531,6 +1536,7 @@ async fn handle_codex_anthropic_to_responses_transform( &model, &request_model, &outbound_model, + InputTokenSemantics::TotalIncludesCacheBuckets, usage, latency_ms, None, @@ -1618,6 +1624,7 @@ fn build_codex_anthropic_sse_response( &model, &request_model, &outbound_model, + InputTokenSemantics::TotalIncludesCacheBuckets, usage, latency_ms, first_token_ms, @@ -2590,6 +2597,7 @@ fn log_forward_error( is_streaming, Some(ctx.session_id.clone()), None, + InputTokenSemantics::FreshExcludesCache, ) { log::warn!("记录失败请求日志失败: {e}"); } @@ -2607,6 +2615,7 @@ async fn log_usage( model: &str, request_model: &str, outbound_model: &str, + input_token_semantics: InputTokenSemantics, usage: TokenUsage, latency_ms: u64, first_token_ms: Option, @@ -2640,6 +2649,7 @@ async fn log_usage( model.to_string(), request_model.to_string(), pricing_model.to_string(), + input_token_semantics, usage, multiplier, latency_ms, diff --git a/src-tauri/src/proxy/mod.rs b/src-tauri/src/proxy/mod.rs index d1dc85808..1e4de3033 100644 --- a/src-tauri/src/proxy/mod.rs +++ b/src-tauri/src/proxy/mod.rs @@ -21,6 +21,8 @@ pub(crate) mod json_canonical; pub mod log_codes; pub mod media_sanitizer; pub mod model_mapper; +pub(crate) mod pi_handler; +pub(crate) mod pi_runtime; pub mod provider_router; pub mod providers; pub mod response_processor; diff --git a/src-tauri/src/proxy/pi_handler.rs b/src-tauri/src/proxy/pi_handler.rs new file mode 100644 index 000000000..f153aa169 --- /dev/null +++ b/src-tauri/src/proxy/pi_handler.rs @@ -0,0 +1,1739 @@ +//! Native Pi gateway transport. +//! +//! Pi's SDK has already serialized the request before it reaches this route. +//! The handler therefore preserves method, path, query, and body bytes; it only +//! replaces gateway/client transport headers with candidate-local material +//! and selects a wire-compatible failover target from one immutable lease. + +use super::pi_runtime::{ + infer_family, PiMaterializedAttempt, PiRequestCandidate, PiRuntimeSnapshot, +}; +use super::server::ProxyState; +use super::usage::{InputTokenSemantics, TokenUsage, UsageLogger}; +use super::ProxyError; +use crate::database::PRICING_SOURCE_REQUEST; +use crate::pi_config::gateway::gateway_replaces_incoming_header; +use axum::body::Body; +use axum::extract::{Path, State}; +use axum::response::Response; +use bytes::Bytes; +use futures::{stream::BoxStream, StreamExt}; +use http::header::{ + AUTHORIZATION, CONNECTION, CONTENT_LENGTH, HOST, PROXY_AUTHENTICATE, PROXY_AUTHORIZATION, TE, + TRAILER, TRANSFER_ENCODING, UPGRADE, +}; +use http::{HeaderMap, HeaderName, StatusCode}; +use http_body_util::BodyExt; +use serde_json::Value; +use std::time::{Duration, Instant}; + +const USAGE_CAPTURE_LIMIT: usize = 4 * 1024 * 1024; +const SSE_PREFLIGHT_LIMIT: usize = 1024 * 1024; + +pub(crate) async fn handle_pi_native( + State(state): State, + Path((route_token, wildcard_path)): Path<(String, String)>, + request: axum::extract::Request, +) -> Result { + let forwarded_path = format!("/{}", wildcard_path.trim_start_matches('/')); + let family = infer_family(&forwarded_path).ok_or_else(|| { + ProxyError::InvalidRequest(format!( + "unsupported Pi native gateway path: {forwarded_path}" + )) + })?; + let snapshot = state + .pi_runtime + .lease(state.pi_server_generation) + .ok_or(ProxyError::NoAvailableProvider)?; + authenticate_gateway(&snapshot, family, request.headers())?; + + let (parts, body) = request.into_parts(); + let method = parts.method; + let uri = parts.uri; + let path_and_query = uri.query().map_or_else( + || forwarded_path.clone(), + |query| format!("{forwarded_path}?{query}"), + ); + let incoming_headers = parts.headers; + let body = body + .collect() + .await + .map_err(|error| ProxyError::InvalidRequest(format!("failed to read Pi request: {error}")))? + .to_bytes(); + let request_json = if body.is_empty() { + Value::Null + } else { + serde_json::from_slice::(&body).map_err(|error| { + ProxyError::InvalidRequest(format!("Pi request body is not valid JSON: {error}")) + })? + }; + let model_id = request_model(family, &forwarded_path, &request_json)?; + let route = snapshot + .route(&route_token, family, &model_id) + .map_err(|error| ProxyError::ConfigError(error.to_string()))?; + let admission = state + .pi_runtime + .admission_guard(state.pi_server_generation, &snapshot) + .await + .ok_or(ProxyError::NoAvailableProvider)?; + drop(admission); + + let is_streaming = request_is_streaming(&uri, &incoming_headers, &request_json); + // Retry policy counts actual upstream sends. Circuit-open candidates, + // protocol-ineligible failovers, and materialization failures must not + // consume the budget or hide a later eligible candidate. + let mut network_budget = + NetworkAttemptBudget::new((route.app_config.max_retries as usize).saturating_add(1)); + let attempts = route.candidates; + let request_headers = filtered_incoming_headers(&incoming_headers); + let started = Instant::now(); + let session_id = + crate::proxy::extract_session_id(&incoming_headers, &request_json, "pi").session_id; + let mut protocol_anchor = ProtocolAnchor::for_primary(attempts.first()); + let mut last_error = None; + let mut pending_retryable: Option = None; + record_request_start(&state).await; + + let mut index = 0; + while index < attempts.len() && network_budget.has_remaining() { + let provider_id = attempts[index].provider_id.clone(); + let provider_end = attempts[index..] + .iter() + .position(|candidate| candidate.provider_id != provider_id) + .map_or(attempts.len(), |offset| index + offset); + let permit = state + .provider_router + .allow_provider_request(&provider_id, "pi") + .await; + if !permit.allowed { + last_error = Some(format!("Pi provider '{provider_id}' circuit is open")); + index = provider_end; + continue; + } + + let mut provider_health_failure = None; + for candidate in attempts[index..provider_end].iter().cloned() { + if !network_budget.has_remaining() { + break; + } + let Some(single_direct_attempt) = + begin_protocol_materialization(&mut protocol_anchor, candidate.is_failover) + else { + continue; + }; + let materialized = match materialize_candidate(candidate, path_and_query.clone()).await + { + Ok(candidate) => candidate, + Err(error) => { + last_error = Some(error.to_string()); + // Deferred materialization is provider-level state, not an + // endpoint health result. Do not execute the same command + // again for every endpoint; move to a compatible provider. + break; + } + }; + let protocol_identity = materialized + .transport + .failover_protocol_identity() + .map(|(family, headers)| (family.as_str().to_string(), headers.clone())); + if !single_direct_attempt + && !protocol_identity_allows_attempt(&mut protocol_anchor, protocol_identity) + { + continue; + } + + let outgoing_headers = + merge_candidate_headers(&request_headers, &materialized.transport.headers); + let timeout_seconds = if is_streaming { + route.app_config.streaming_first_byte_timeout + } else { + route.app_config.non_streaming_timeout + }; + if !network_budget.begin_send() { + break; + } + if let Some(pending) = pending_retryable.take() { + if pending.provider_id != provider_id { + settle_provider_health( + &state, + route.catalog_epoch, + &pending.provider_id, + pending.used_half_open_permit, + pending.provider_health.clone(), + ) + .await; + } else { + debug_assert_eq!( + pending.used_half_open_permit, permit.used_half_open_permit, + "one provider group must retain one circuit-breaker permit" + ); + } + // A real later send has now begun, so the earlier fallback + // response is no longer client-visible. + drop(pending); + } + let send = crate::proxy::http_client::get() + .request(method.clone(), materialized.url.clone()) + .headers(outgoing_headers) + .body(body.clone()) + .send(); + let response = match if timeout_seconds > 0 { + tokio::time::timeout(Duration::from_secs(u64::from(timeout_seconds)), send) + .await + .map_err(|_| ()) + } else { + Ok(send.await) + } { + Ok(Ok(response)) => response, + Ok(Err(error)) => { + let error = if error.is_timeout() { + "Pi upstream request timed out".to_string() + } else { + "Pi upstream request failed before response".to_string() + }; + provider_health_failure = Some(error.clone()); + last_error = Some(error); + continue; + } + Err(()) => { + let error = "Pi upstream response-header timeout".to_string(); + provider_health_failure = Some(error.clone()); + last_error = Some(error); + continue; + } + }; + + let status = response.status(); + let status_disposition = upstream_status_disposition(status); + if status_disposition.is_retryable() && network_budget.has_remaining() { + let error = format!("Pi upstream returned retryable status {status}"); + let status_health = status_disposition.provider_health(); + if status_health == ProviderHealthDisposition::Unhealthy { + provider_health_failure = Some(error.clone()); + } + last_error = Some(error); + let selected_is_failover = materialized.is_failover; + pending_retryable = Some(PendingRetryableResponse { + response, + materialized, + provider_id: provider_id.clone(), + used_half_open_permit: permit.used_half_open_permit, + selected_is_failover, + provider_health: ProviderHealthOutcome::from_status( + status, + provider_health_failure.as_deref(), + ), + }); + match status_disposition { + UpstreamStatusDisposition::RetryEndpoint => continue, + // Every endpoint in one provider group is cloned from the + // same credential plan. Preserve this response as a + // fallback, but reserve the remaining network budget for + // a provider that can own a different credential. + UpstreamStatusDisposition::RetryProvider => break, + UpstreamStatusDisposition::ReturnResponse => { + unreachable!("a non-retryable status cannot enter the retry branch") + } + } + } + let selected_is_failover = materialized.is_failover; + let provider_health = + ProviderHealthOutcome::from_status(status, provider_health_failure.as_deref()); + match prepare_response( + state.clone(), + response, + materialized, + route.catalog_epoch, + model_id.clone(), + session_id.clone(), + started, + is_streaming, + route.app_config.streaming_first_byte_timeout, + route.app_config.streaming_idle_timeout, + route.app_config.non_streaming_timeout, + permit.used_half_open_permit, + selected_is_failover, + provider_health.clone(), + ) + .await + { + Ok(prepared) => { + if !prepared.finalization_deferred { + settle_provider_health( + &state, + route.catalog_epoch, + &provider_id, + permit.used_half_open_permit, + provider_health, + ) + .await; + record_request_finish( + &state, + status.is_success(), + selected_is_failover, + (!status.is_success()) + .then(|| format!("Pi upstream returned {status}")), + ) + .await; + } + return Ok(prepared.response); + } + Err(ProxyError::ForwardFailed(error)) | Err(ProxyError::Timeout(error)) => { + provider_health_failure = Some(error.clone()); + last_error = Some(error); + continue; + } + Err(error) => { + release_or_record_provider( + &state, + route.catalog_epoch, + &provider_id, + permit.used_half_open_permit, + provider_health_failure.clone(), + ) + .await; + record_request_finish(&state, false, false, Some(error.to_string())).await; + return Err(error); + } + } + } + + if pending_retryable + .as_ref() + .is_none_or(|pending| pending.provider_id != provider_id) + { + release_or_record_provider( + &state, + route.catalog_epoch, + &provider_id, + permit.used_half_open_permit, + provider_health_failure, + ) + .await; + } + index = provider_end; + } + + if let Some(pending) = pending_retryable { + let status = pending.response.status(); + let provider_id = pending.provider_id.clone(); + let used_half_open_permit = pending.used_half_open_permit; + let selected_is_failover = pending.selected_is_failover; + let provider_health = pending.provider_health; + match prepare_response( + state.clone(), + pending.response, + pending.materialized, + route.catalog_epoch, + model_id.clone(), + session_id.clone(), + started, + is_streaming, + route.app_config.streaming_first_byte_timeout, + route.app_config.streaming_idle_timeout, + route.app_config.non_streaming_timeout, + used_half_open_permit, + selected_is_failover, + provider_health.clone(), + ) + .await + { + Ok(prepared) => { + if !prepared.finalization_deferred { + settle_provider_health( + &state, + route.catalog_epoch, + &provider_id, + used_half_open_permit, + provider_health, + ) + .await; + record_request_finish( + &state, + status.is_success(), + selected_is_failover, + (!status.is_success()).then(|| format!("Pi upstream returned {status}")), + ) + .await; + } + return Ok(prepared.response); + } + Err(error) => { + record_provider_result( + &state, + route.catalog_epoch, + &provider_id, + used_half_open_permit, + false, + Some(error.to_string()), + ) + .await; + record_request_finish(&state, false, selected_is_failover, Some(error.to_string())) + .await; + return Err(error); + } + } + } + + let error = + last_error.unwrap_or_else(|| "no wire-compatible Pi candidate was available".to_string()); + record_request_finish(&state, false, false, Some(error.clone())).await; + Err(ProxyError::ForwardFailed(error)) +} + +struct PendingRetryableResponse { + response: reqwest::Response, + materialized: PiMaterializedAttempt, + provider_id: String, + used_half_open_permit: bool, + selected_is_failover: bool, + provider_health: ProviderHealthOutcome, +} + +#[derive(Debug)] +struct NetworkAttemptBudget { + remaining: usize, +} + +impl NetworkAttemptBudget { + fn new(max_attempts: usize) -> Self { + Self { + remaining: max_attempts, + } + } + + fn has_remaining(&self) -> bool { + self.remaining > 0 + } + + /// Consume budget only immediately before an actual upstream send. + fn begin_send(&mut self) -> bool { + if self.remaining == 0 { + return false; + } + self.remaining -= 1; + true + } +} + +/// Return whether this attempt consumes the one direct-only grant, or `None` +/// when the candidate must not even be materialized. +fn begin_protocol_materialization(anchor: &mut ProtocolAnchor, is_failover: bool) -> Option { + match anchor { + ProtocolAnchor::DirectOnlyPending if !is_failover => { + // Consume before materialization so an error in a later + // credential/custom header cannot run the protocol command again. + *anchor = ProtocolAnchor::Ineligible; + Some(true) + } + ProtocolAnchor::DirectOnlyPending | ProtocolAnchor::Ineligible => None, + ProtocolAnchor::Unset | ProtocolAnchor::Predictable(_) => Some(false), + } +} + +#[derive(Debug)] +enum ProtocolAnchor { + Unset, + Predictable((String, HeaderMap)), + DirectOnlyPending, + Ineligible, +} + +impl ProtocolAnchor { + fn for_primary(primary: Option<&PiRequestCandidate>) -> Self { + let Some(primary) = primary else { + return Self::Unset; + }; + if !primary.protocol_identity_is_predictable() { + return Self::DirectOnlyPending; + } + match primary.planned_protocol_identity() { + Ok(Some(identity)) => Self::Predictable(identity), + Ok(None) => Self::DirectOnlyPending, + // If the primary's protocol identity cannot be established, a + // backup must not self-declare compatibility. Give the primary + // exactly one direct materialization; failure remains fail-closed. + Err(_) => Self::DirectOnlyPending, + } + } +} + +fn protocol_identity_allows_attempt( + anchor: &mut ProtocolAnchor, + candidate: Option<(String, HeaderMap)>, +) -> bool { + match (&*anchor, candidate) { + (ProtocolAnchor::Unset, Some(candidate)) => { + *anchor = ProtocolAnchor::Predictable(candidate); + true + } + (ProtocolAnchor::Unset, None) => false, + (ProtocolAnchor::Predictable(primary), Some(candidate)) => primary == &candidate, + (ProtocolAnchor::Predictable(_), None) + | (ProtocolAnchor::DirectOnlyPending, _) + | (ProtocolAnchor::Ineligible, _) => false, + } +} + +async fn materialize_candidate( + candidate: PiRequestCandidate, + path_and_query: String, +) -> Result { + tokio::task::spawn_blocking(move || candidate.materialize(&path_and_query)) + .await + .map_err(|error| { + ProxyError::Internal(format!("Pi candidate materialization task failed: {error}")) + })? + .map_err(|error| ProxyError::ConfigError(error.to_string())) +} + +async fn record_provider_result( + state: &ProxyState, + catalog_epoch: u64, + provider_id: &str, + used_half_open_permit: bool, + success: bool, + error: Option, +) { + let Some(_guard) = state.pi_runtime.writeback_guard(catalog_epoch).await else { + state + .provider_router + .release_permit_neutral(provider_id, "pi", used_half_open_permit) + .await; + return; + }; + if let Err(record_error) = state + .provider_router + .record_result(provider_id, "pi", used_half_open_permit, success, error) + .await + { + log::warn!("failed to update Pi provider health: {record_error}"); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ProviderHealthDisposition { + Healthy, + Unhealthy, + Neutral, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum UpstreamStatusDisposition { + ReturnResponse, + RetryEndpoint, + RetryProvider, +} + +impl UpstreamStatusDisposition { + const fn is_retryable(self) -> bool { + !matches!(self, Self::ReturnResponse) + } + + const fn provider_health(self) -> ProviderHealthDisposition { + match self { + Self::ReturnResponse => ProviderHealthDisposition::Healthy, + Self::RetryEndpoint => ProviderHealthDisposition::Unhealthy, + Self::RetryProvider => ProviderHealthDisposition::Neutral, + } + } +} + +fn upstream_status_disposition(status: StatusCode) -> UpstreamStatusDisposition { + if matches!(status, StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) { + // Pi's provider owns authentication (API key or OAuth), while its + // custom endpoints only replace the URL. Authentication rejection is + // therefore neutral for endpoint health and can only benefit from a + // distinct provider credential. + return UpstreamStatusDisposition::RetryProvider; + } + if (!status.is_client_error() && !status.is_server_error()) + || matches!( + status, + StatusCode::BAD_REQUEST + | StatusCode::METHOD_NOT_ALLOWED + | StatusCode::NOT_ACCEPTABLE + | StatusCode::PAYLOAD_TOO_LARGE + | StatusCode::URI_TOO_LONG + | StatusCode::UNSUPPORTED_MEDIA_TYPE + | StatusCode::UNPROCESSABLE_ENTITY + | StatusCode::NOT_IMPLEMENTED + ) + { + UpstreamStatusDisposition::ReturnResponse + } else { + UpstreamStatusDisposition::RetryEndpoint + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct ProviderHealthOutcome { + disposition: ProviderHealthDisposition, + error: Option, +} + +impl ProviderHealthOutcome { + fn from_status(status: StatusCode, prior_failure: Option<&str>) -> Self { + match upstream_status_disposition(status).provider_health() { + ProviderHealthDisposition::Healthy => Self { + disposition: ProviderHealthDisposition::Healthy, + error: None, + }, + ProviderHealthDisposition::Unhealthy => Self { + disposition: ProviderHealthDisposition::Unhealthy, + error: Some(format!("Pi upstream returned {status}")), + }, + ProviderHealthDisposition::Neutral => prior_failure.map_or_else( + || Self { + disposition: ProviderHealthDisposition::Neutral, + error: None, + }, + |error| Self { + // A credential rejection is neutral by itself, but it + // cannot erase a real failure from an earlier endpoint + // covered by the same provider-level permit. + disposition: ProviderHealthDisposition::Unhealthy, + error: Some(error.to_string()), + }, + ), + } + } +} + +async fn settle_provider_health( + state: &ProxyState, + catalog_epoch: u64, + provider_id: &str, + used_half_open_permit: bool, + outcome: ProviderHealthOutcome, +) { + match outcome.disposition { + ProviderHealthDisposition::Healthy => { + record_provider_result( + state, + catalog_epoch, + provider_id, + used_half_open_permit, + true, + None, + ) + .await; + } + ProviderHealthDisposition::Unhealthy => { + record_provider_result( + state, + catalog_epoch, + provider_id, + used_half_open_permit, + false, + outcome.error, + ) + .await; + } + ProviderHealthDisposition::Neutral => { + state + .provider_router + .release_permit_neutral(provider_id, "pi", used_half_open_permit) + .await; + } + } +} + +async fn release_or_record_provider( + state: &ProxyState, + catalog_epoch: u64, + provider_id: &str, + used_half_open_permit: bool, + provider_health_failure: Option, +) { + if provider_health_failure.is_some() { + record_provider_result( + state, + catalog_epoch, + provider_id, + used_half_open_permit, + false, + provider_health_failure, + ) + .await; + } else { + state + .provider_router + .release_permit_neutral(provider_id, "pi", used_half_open_permit) + .await; + } +} + +async fn record_request_start(state: &ProxyState) { + let mut status = state.status.write().await; + status.total_requests = status.total_requests.saturating_add(1); + status.last_request_at = Some(chrono::Utc::now().to_rfc3339()); +} + +async fn record_request_finish( + state: &ProxyState, + success: bool, + used_failover: bool, + error: Option, +) { + let mut status = state.status.write().await; + if success { + status.success_requests = status.success_requests.saturating_add(1); + status.last_error = None; + } else { + status.failed_requests = status.failed_requests.saturating_add(1); + status.last_error = error; + } + if used_failover { + status.failover_count = status.failover_count.saturating_add(1); + } + if status.total_requests > 0 { + status.success_rate = + (status.success_requests as f32 / status.total_requests as f32) * 100.0; + } +} + +fn authenticate_gateway( + snapshot: &PiRuntimeSnapshot, + family: crate::pi_config::gateway::PiGatewayApiFamily, + headers: &HeaderMap, +) -> Result<(), ProxyError> { + let value = match family { + crate::pi_config::gateway::PiGatewayApiFamily::AnthropicMessages => headers + .get("x-api-key") + .and_then(|value| value.to_str().ok()), + crate::pi_config::gateway::PiGatewayApiFamily::GoogleGenerativeAi => headers + .get("x-goog-api-key") + .and_then(|value| value.to_str().ok()), + crate::pi_config::gateway::PiGatewayApiFamily::OpenAiCompletions + | crate::pi_config::gateway::PiGatewayApiFamily::OpenAiResponses => headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.strip_prefix("Bearer ")), + } + .ok_or_else(|| ProxyError::AuthError("missing Pi gateway credential".to_string()))?; + if !snapshot.token_matches(value) { + return Err(ProxyError::AuthError( + "invalid Pi gateway credential".to_string(), + )); + } + Ok(()) +} + +fn request_model( + family: crate::pi_config::gateway::PiGatewayApiFamily, + path: &str, + body: &Value, +) -> Result { + if family == crate::pi_config::gateway::PiGatewayApiFamily::GoogleGenerativeAi { + let encoded = path + .strip_prefix("/models/") + .and_then(|rest| rest.split(':').next()) + .filter(|model| !model.is_empty()) + .ok_or_else(|| { + ProxyError::InvalidRequest("Pi Google request has no model in its path".to_string()) + })?; + return percent_decode(encoded).ok_or_else(|| { + ProxyError::InvalidRequest("Pi Google model path has invalid escaping".to_string()) + }); + } + body.get("model") + .and_then(Value::as_str) + .filter(|model| !model.is_empty()) + .map(str::to_string) + .ok_or_else(|| { + ProxyError::InvalidRequest("Pi request body has no model identifier".to_string()) + }) +} + +fn percent_decode(value: &str) -> Option { + let bytes = value.as_bytes(); + let mut decoded = Vec::with_capacity(bytes.len()); + let mut index = 0; + while index < bytes.len() { + if bytes[index] != b'%' { + decoded.push(bytes[index]); + index += 1; + continue; + } + let high = *bytes.get(index + 1)?; + let low = *bytes.get(index + 2)?; + decoded.push(hex(high)? << 4 | hex(low)?); + index += 3; + } + String::from_utf8(decoded).ok() +} + +fn hex(value: u8) -> Option { + match value { + b'0'..=b'9' => Some(value - b'0'), + b'a'..=b'f' => Some(value - b'a' + 10), + b'A'..=b'F' => Some(value - b'A' + 10), + _ => None, + } +} + +fn request_is_streaming(uri: &http::Uri, headers: &HeaderMap, body: &Value) -> bool { + body.get("stream").and_then(Value::as_bool).unwrap_or(false) + || uri + .query() + .is_some_and(|query| query.split('&').any(|part| part == "alt=sse")) + || headers + .get(http::header::ACCEPT) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| value.contains("text/event-stream")) +} + +fn filtered_incoming_headers(headers: &HeaderMap) -> HeaderMap { + let connection_named = connection_named_headers(headers); + let mut filtered = HeaderMap::new(); + for (name, value) in headers { + if matches!( + *name, + HOST | CONTENT_LENGTH + | CONNECTION + | TRANSFER_ENCODING + | TE + | TRAILER + | UPGRADE + | AUTHORIZATION + | PROXY_AUTHENTICATE + | PROXY_AUTHORIZATION + ) || gateway_replaces_incoming_header(name) + || connection_named.contains(name) + { + continue; + } + filtered.append(name.clone(), value.clone()); + } + filtered +} + +fn merge_candidate_headers(incoming: &HeaderMap, candidate: &HeaderMap) -> HeaderMap { + let mut merged = incoming.clone(); + for (name, value) in candidate { + merged.insert(name.clone(), value.clone()); + } + merged.remove(CONTENT_LENGTH); + merged.remove(TRANSFER_ENCODING); + merged +} + +struct PreparedPiResponse { + response: Response, + finalization_deferred: bool, +} + +struct PiStreamFinalization { + state: ProxyState, + candidate: PiMaterializedAttempt, + catalog_epoch: u64, + request_model: String, + session_id: String, + started: Instant, + is_streaming: bool, + status: StatusCode, + content_is_sse: bool, + used_half_open_permit: bool, + selected_is_failover: bool, + complete_provider_health: ProviderHealthOutcome, +} + +enum PiStreamTermination { + Complete { captured: Option> }, + UpstreamFailure { message: String }, + DownstreamDropped, +} + +struct PiStreamDisposition { + provider_health: ProviderHealthDisposition, + provider_error: Option, + request_success: bool, + request_error: Option, +} + +fn pi_stream_disposition( + status: StatusCode, + complete_provider_health: &ProviderHealthOutcome, + termination: &PiStreamTermination, +) -> PiStreamDisposition { + match termination { + PiStreamTermination::Complete { .. } => PiStreamDisposition { + provider_health: complete_provider_health.disposition, + provider_error: complete_provider_health.error.clone(), + request_success: status.is_success(), + request_error: (!status.is_success()).then(|| format!("Pi upstream returned {status}")), + }, + PiStreamTermination::UpstreamFailure { message } => PiStreamDisposition { + provider_health: ProviderHealthDisposition::Unhealthy, + provider_error: Some(message.clone()), + request_success: false, + request_error: Some(message.clone()), + }, + PiStreamTermination::DownstreamDropped => PiStreamDisposition { + // A downstream cancellation says nothing about upstream health. + provider_health: ProviderHealthDisposition::Neutral, + provider_error: None, + request_success: false, + request_error: Some( + "Pi downstream client closed before the upstream stream completed".to_string(), + ), + }, + } +} + +struct PiStreamFinalizer { + pending: Option, +} + +impl PiStreamFinalizer { + fn new(finalization: PiStreamFinalization) -> Self { + Self { + pending: Some(finalization), + } + } + + fn finish(mut self, termination: PiStreamTermination) { + if let Some(finalization) = self.pending.take() { + finalization.spawn(termination); + } + } +} + +impl Drop for PiStreamFinalizer { + fn drop(&mut self) { + if let Some(finalization) = self.pending.take() { + finalization.spawn(PiStreamTermination::DownstreamDropped); + } + } +} + +impl PiStreamFinalization { + fn spawn(self, termination: PiStreamTermination) { + let Ok(runtime) = tokio::runtime::Handle::try_current() else { + log::error!("Pi stream finalization lost because no Tokio runtime is available"); + return; + }; + runtime.spawn(async move { + self.apply(termination).await; + }); + } + + async fn apply(self, termination: PiStreamTermination) { + let disposition = + pi_stream_disposition(self.status, &self.complete_provider_health, &termination); + settle_provider_health( + &self.state, + self.catalog_epoch, + &self.candidate.provider_id, + self.used_half_open_permit, + ProviderHealthOutcome { + disposition: disposition.provider_health, + error: disposition.provider_error, + }, + ) + .await; + record_request_finish( + &self.state, + disposition.request_success, + self.selected_is_failover, + disposition.request_error.clone(), + ) + .await; + + match termination { + PiStreamTermination::Complete { captured } => { + if let Some(_guard) = self + .state + .pi_runtime + .writeback_guard(self.catalog_epoch) + .await + { + log_pi_usage( + &self.state, + &self.candidate, + &self.request_model, + &self.session_id, + self.started, + self.is_streaming, + self.status, + self.content_is_sse, + captured.as_deref(), + ) + .await; + } + } + PiStreamTermination::UpstreamFailure { message } => { + if let Some(_guard) = self + .state + .pi_runtime + .writeback_guard(self.catalog_epoch) + .await + { + log_pi_stream_error(&self, StatusCode::BAD_GATEWAY.as_u16(), &message); + } + } + PiStreamTermination::DownstreamDropped => { + if let Some(_guard) = self + .state + .pi_runtime + .writeback_guard(self.catalog_epoch) + .await + { + if let Some(message) = disposition.request_error.as_deref() { + log_pi_stream_error(&self, 499, message); + } + } + } + } + } +} + +fn log_pi_stream_error(finalization: &PiStreamFinalization, status_code: u16, message: &str) { + let logging_enabled = finalization + .state + .config + .try_read() + .map(|config| config.enable_logging) + .unwrap_or(true); + if !logging_enabled { + return; + } + let logger = UsageLogger::new(&finalization.state.db); + if let Err(error) = logger.log_error_with_context( + uuid::Uuid::new_v4().to_string(), + finalization.candidate.provider_id.clone(), + "pi".to_string(), + finalization.request_model.clone(), + status_code, + message.to_string(), + finalization.started.elapsed().as_millis() as u64, + finalization.is_streaming, + (!finalization.session_id.is_empty()).then(|| finalization.session_id.clone()), + Some(finalization.candidate.transport.family_name().to_string()), + InputTokenSemantics::for_pi_family(finalization.candidate.transport.family()), + ) { + log::warn!("failed to record Pi gateway stream error: {error}"); + } +} + +#[allow(clippy::too_many_arguments)] +async fn prepare_response( + state: ProxyState, + response: reqwest::Response, + candidate: PiMaterializedAttempt, + catalog_epoch: u64, + request_model: String, + session_id: String, + started: Instant, + is_streaming: bool, + first_semantic_timeout_seconds: u32, + streaming_idle_timeout_seconds: u32, + non_streaming_timeout_seconds: u32, + used_half_open_permit: bool, + selected_is_failover: bool, + complete_provider_health: ProviderHealthOutcome, +) -> Result { + let status = response.status(); + let headers = filtered_response_headers(response.headers()); + let content_is_sse = response + .headers() + .get(http::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| value.starts_with("text/event-stream")); + if !is_streaming && !content_is_sse { + let read = response.bytes(); + let bytes = if non_streaming_timeout_seconds > 0 { + tokio::time::timeout( + Duration::from_secs(u64::from(non_streaming_timeout_seconds)), + read, + ) + .await + .map_err(|_| { + ProxyError::Timeout( + "Pi non-streaming response body exceeded its timeout".to_string(), + ) + })? + .map_err(|error| { + ProxyError::ForwardFailed(format!("Pi upstream response body failed: {error}")) + })? + } else { + read.await.map_err(|error| { + ProxyError::ForwardFailed(format!("Pi upstream response body failed: {error}")) + })? + }; + if let Some(_guard) = state.pi_runtime.writeback_guard(catalog_epoch).await { + state.current_providers.write().await.insert( + "pi".to_string(), + ( + candidate.provider_id.clone(), + candidate.provider_name.clone(), + ), + ); + log_pi_usage( + &state, + &candidate, + &request_model, + &session_id, + started, + false, + status, + false, + Some(&bytes), + ) + .await; + } + let mut builder = Response::builder().status(status); + *builder.headers_mut().ok_or_else(|| { + ProxyError::Internal("failed to build Pi response headers".to_string()) + })? = headers; + let response = builder.body(Body::from(bytes)).map_err(|error| { + ProxyError::Internal(format!("failed to build Pi response: {error}")) + })?; + return Ok(PreparedPiResponse { + response, + finalization_deferred: false, + }); + } + + let mut stream = response.bytes_stream().boxed(); + let prefix = if content_is_sse || (is_streaming && status.is_success()) { + preflight_sse( + &mut stream, + first_semantic_timeout_seconds, + SSE_PREFLIGHT_LIMIT, + ) + .await? + } else { + Vec::new() + }; + if let Some(_guard) = state.pi_runtime.writeback_guard(catalog_epoch).await { + state.current_providers.write().await.insert( + "pi".to_string(), + ( + candidate.provider_id.clone(), + candidate.provider_name.clone(), + ), + ); + } + let body_stream = logged_body_stream( + state, + stream, + prefix, + candidate, + catalog_epoch, + request_model, + session_id, + started, + is_streaming || content_is_sse, + status, + content_is_sse, + streaming_idle_timeout_seconds, + used_half_open_permit, + selected_is_failover, + complete_provider_health, + ); + let mut builder = Response::builder().status(status); + *builder + .headers_mut() + .ok_or_else(|| ProxyError::Internal("failed to build Pi response headers".to_string()))? = + headers; + let response = builder + .body(Body::from_stream(body_stream)) + .map_err(|error| ProxyError::Internal(format!("failed to build Pi response: {error}")))?; + Ok(PreparedPiResponse { + response, + finalization_deferred: true, + }) +} + +fn filtered_response_headers(headers: &HeaderMap) -> HeaderMap { + let connection_named = connection_named_headers(headers); + let mut filtered = HeaderMap::new(); + for (name, value) in headers { + if matches!( + *name, + CONNECTION + | CONTENT_LENGTH + | TRANSFER_ENCODING + | TE + | TRAILER + | UPGRADE + | PROXY_AUTHENTICATE + | PROXY_AUTHORIZATION + ) || connection_named.contains(name) + { + continue; + } + filtered.append(name.clone(), value.clone()); + } + filtered +} + +fn connection_named_headers(headers: &HeaderMap) -> std::collections::HashSet { + headers + .get_all(CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()) + .flat_map(|value| value.split(',')) + .filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok()) + .collect() +} + +async fn preflight_sse( + stream: &mut BoxStream<'static, Result>, + timeout_seconds: u32, + byte_limit: usize, +) -> Result, ProxyError> { + let deadline = (timeout_seconds > 0) + .then(|| tokio::time::Instant::now() + Duration::from_secs(u64::from(timeout_seconds))); + let mut chunks = Vec::new(); + let mut buffer = Vec::new(); + loop { + let next = match deadline { + Some(deadline) => tokio::time::timeout_at(deadline, stream.next()) + .await + .map_err(|_| { + ProxyError::Timeout( + "Pi SSE produced no semantic event before the first-event timeout" + .to_string(), + ) + })?, + None => stream.next().await, + }; + let chunk = next + .ok_or_else(|| { + ProxyError::ForwardFailed( + "Pi SSE ended before its first semantic event".to_string(), + ) + })? + .map_err(|_| { + ProxyError::ForwardFailed( + "Pi SSE failed before its first semantic event".to_string(), + ) + })?; + if buffer.len().saturating_add(chunk.len()) > byte_limit { + return Err(ProxyError::ForwardFailed( + "Pi SSE prelude exceeded the bounded commit fence".to_string(), + )); + } + buffer.extend_from_slice(&chunk); + chunks.push(chunk); + if contains_semantic_sse_event(&buffer) { + return Ok(chunks); + } + } +} + +fn contains_semantic_sse_event(bytes: &[u8]) -> bool { + let text = String::from_utf8_lossy(bytes).replace("\r\n", "\n"); + text.split("\n\n").any(|block| { + block.lines().any(|line| { + line.strip_prefix("data:") + .is_some_and(|data| !data.trim().is_empty()) + }) + }) +} + +#[allow(clippy::too_many_arguments)] +fn logged_body_stream( + state: ProxyState, + mut stream: BoxStream<'static, Result>, + prefix: Vec, + candidate: PiMaterializedAttempt, + catalog_epoch: u64, + request_model: String, + session_id: String, + started: Instant, + is_streaming: bool, + status: StatusCode, + content_is_sse: bool, + streaming_idle_timeout_seconds: u32, + used_half_open_permit: bool, + selected_is_failover: bool, + complete_provider_health: ProviderHealthOutcome, +) -> impl futures::Stream> + Send + 'static { + // Construct the guard before the generator is polled. Axum may drop a + // response body without ever polling it when the client disconnects after + // headers, and the HalfOpen permit must still be released in that case. + let finalizer = PiStreamFinalizer::new(PiStreamFinalization { + state, + candidate, + catalog_epoch, + request_model, + session_id, + started, + is_streaming, + status, + content_is_sse, + used_half_open_permit, + selected_is_failover, + complete_provider_health, + }); + async_stream::stream! { + let mut captured = Vec::new(); + let mut capture_open = true; + for chunk in prefix { + capture_usage_bytes(&mut captured, &mut capture_open, &chunk); + yield Ok(chunk); + } + loop { + let next = if streaming_idle_timeout_seconds > 0 { + match tokio::time::timeout( + Duration::from_secs(u64::from(streaming_idle_timeout_seconds)), + stream.next(), + ) + .await + { + Ok(next) => next, + Err(_) => { + let message = "Pi upstream stream exceeded its idle timeout".to_string(); + finalizer.finish(PiStreamTermination::UpstreamFailure { + message: message.clone(), + }); + yield Err(std::io::Error::new( + std::io::ErrorKind::TimedOut, + message, + )); + return; + } + } + } else { + stream.next().await + }; + let Some(result) = next else { + break; + }; + match result { + Ok(chunk) => { + capture_usage_bytes(&mut captured, &mut capture_open, &chunk); + yield Ok(chunk); + } + Err(error) => { + let message = format!("Pi upstream response failed: {error}"); + finalizer.finish(PiStreamTermination::UpstreamFailure { + message: message.clone(), + }); + yield Err(std::io::Error::other(message)); + return; + } + } + } + finalizer.finish(PiStreamTermination::Complete { + captured: capture_open.then_some(captured), + }); + } +} + +fn capture_usage_bytes(captured: &mut Vec, capture_open: &mut bool, chunk: &[u8]) { + if !*capture_open { + return; + } + if captured.len().saturating_add(chunk.len()) > USAGE_CAPTURE_LIMIT { + captured.clear(); + *capture_open = false; + return; + } + captured.extend_from_slice(chunk); +} + +#[allow(clippy::too_many_arguments)] +async fn log_pi_usage( + state: &ProxyState, + candidate: &PiMaterializedAttempt, + request_model: &str, + session_id: &str, + started: Instant, + is_streaming: bool, + status: StatusCode, + content_is_sse: bool, + captured: Option<&[u8]>, +) { + let logging_enabled = state + .config + .try_read() + .map(|config| config.enable_logging) + .unwrap_or(true); + if !logging_enabled { + return; + } + let usage = captured + .and_then(|bytes| parse_usage(candidate.transport.family_name(), bytes, content_is_sse)) + .unwrap_or_default(); + let response_model = usage + .model + .clone() + .unwrap_or_else(|| request_model.to_string()); + let logger = UsageLogger::new(&state.db); + let input_token_semantics = InputTokenSemantics::for_pi_family(candidate.transport.family()); + if !status.is_success() { + let _ = logger.log_error( + uuid::Uuid::new_v4().to_string(), + candidate.provider_id.clone(), + "pi".to_string(), + response_model, + status.as_u16(), + format!("Pi upstream returned {status}"), + started.elapsed().as_millis() as u64, + input_token_semantics, + ); + return; + } + let (multiplier, pricing_source) = logger + .resolve_pricing_config(&candidate.provider_id, "pi") + .await; + let pricing_model = if pricing_source == PRICING_SOURCE_REQUEST { + request_model.to_string() + } else { + response_model.clone() + }; + let request_id = usage.dedup_request_id(Some(("pi", candidate.provider_id.as_str()))); + if let Err(error) = logger.log_with_calculation( + request_id, + candidate.provider_id.clone(), + "pi".to_string(), + response_model, + request_model.to_string(), + pricing_model, + input_token_semantics, + usage, + multiplier, + started.elapsed().as_millis() as u64, + None, + status.as_u16(), + (!session_id.is_empty()).then(|| session_id.to_string()), + Some(candidate.transport.family_name().to_string()), + is_streaming, + ) { + log::warn!("failed to record Pi gateway usage: {error}"); + } +} + +fn parse_usage(family: &str, bytes: &[u8], is_sse: bool) -> Option { + if is_sse { + let events = sse_json_events(bytes); + return match family { + "anthropic-messages" => TokenUsage::from_claude_stream_events(&events), + "openai-completions" => TokenUsage::from_openai_stream_events(&events), + "openai-responses" => TokenUsage::from_codex_stream_events_auto(&events), + "google-generative-ai" => TokenUsage::from_gemini_stream_chunks(&events), + _ => None, + }; + } + let body = serde_json::from_slice::(bytes).ok()?; + match family { + "anthropic-messages" => TokenUsage::from_claude_response(&body), + "openai-completions" => TokenUsage::from_openai_response(&body), + "openai-responses" => TokenUsage::from_codex_response_auto(&body), + "google-generative-ai" => TokenUsage::from_gemini_response(&body), + _ => None, + } +} + +fn sse_json_events(bytes: &[u8]) -> Vec { + let text = String::from_utf8_lossy(bytes).replace("\r\n", "\n"); + text.split("\n\n") + .flat_map(str::lines) + .filter_map(|line| line.strip_prefix("data:")) + .map(str::trim) + .filter(|data| !data.is_empty() && *data != "[DONE]") + .filter_map(|data| serde_json::from_str(data).ok()) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn retry_policy_matches_pi_contract_matrix() { + for status in [StatusCode::UNAUTHORIZED, StatusCode::FORBIDDEN] { + assert_eq!( + upstream_status_disposition(status), + UpstreamStatusDisposition::RetryProvider, + "{status}" + ); + } + for status in [ + StatusCode::NOT_FOUND, + StatusCode::REQUEST_TIMEOUT, + StatusCode::CONFLICT, + StatusCode::TOO_MANY_REQUESTS, + StatusCode::IM_A_TEAPOT, + StatusCode::BAD_GATEWAY, + ] { + assert_eq!( + upstream_status_disposition(status), + UpstreamStatusDisposition::RetryEndpoint, + "{status}" + ); + } + for status in [ + StatusCode::BAD_REQUEST, + StatusCode::METHOD_NOT_ALLOWED, + StatusCode::NOT_ACCEPTABLE, + StatusCode::PAYLOAD_TOO_LARGE, + StatusCode::URI_TOO_LONG, + StatusCode::UNSUPPORTED_MEDIA_TYPE, + StatusCode::UNPROCESSABLE_ENTITY, + StatusCode::NOT_IMPLEMENTED, + ] { + assert_eq!( + upstream_status_disposition(status), + UpstreamStatusDisposition::ReturnResponse, + "{status}" + ); + } + } + + #[tokio::test] + async fn credential_rejections_remain_retryable_but_health_neutral() { + let app = axum::Router::new() + .route( + "/unauthorized", + axum::routing::get(|| async { StatusCode::UNAUTHORIZED }), + ) + .route( + "/forbidden", + axum::routing::get(|| async { StatusCode::FORBIDDEN }), + ); + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .expect("bind local Pi capture endpoint"); + let address = listener.local_addr().expect("local capture address"); + let server = tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("serve local Pi capture endpoint"); + }); + + for (path, expected) in [ + ("unauthorized", StatusCode::UNAUTHORIZED), + ("forbidden", StatusCode::FORBIDDEN), + ] { + let response = reqwest::get(format!("http://{address}/{path}")) + .await + .expect("request local Pi capture endpoint"); + assert_eq!(response.status(), expected); + assert_eq!( + upstream_status_disposition(response.status()), + UpstreamStatusDisposition::RetryProvider + ); + assert_eq!( + ProviderHealthOutcome::from_status(response.status(), None), + ProviderHealthOutcome { + disposition: ProviderHealthDisposition::Neutral, + error: None, + } + ); + } + assert_eq!( + ProviderHealthOutcome::from_status( + StatusCode::UNAUTHORIZED, + Some("earlier endpoint failed"), + ), + ProviderHealthOutcome { + disposition: ProviderHealthDisposition::Unhealthy, + error: Some("earlier endpoint failed".to_string()), + } + ); + server.abort(); + } + + #[test] + fn sse_comments_do_not_cross_the_commit_fence() { + assert!(!contains_semantic_sse_event( + b": keep-alive\n\n: another\n\n" + )); + assert!(contains_semantic_sse_event( + b": keep-alive\n\nevent: message_start\ndata: {\"type\":\"message_start\"}\n\n" + )); + } + + #[test] + fn gateway_header_filter_removes_client_auth_and_protocol_identity() { + let mut incoming = HeaderMap::new(); + incoming.insert( + AUTHORIZATION, + http::HeaderValue::from_static("Bearer gateway"), + ); + incoming.insert("x-api-key", http::HeaderValue::from_static("gateway")); + incoming.insert( + "anthropic-version", + http::HeaderValue::from_static("client-version"), + ); + incoming.insert("x-request-local", http::HeaderValue::from_static("kept")); + let filtered = filtered_incoming_headers(&incoming); + assert!(filtered.get(AUTHORIZATION).is_none()); + assert!(filtered.get("x-api-key").is_none()); + assert!(filtered.get("anthropic-version").is_none()); + assert_eq!(filtered["x-request-local"], "kept"); + } + + #[test] + fn percent_decoding_is_strict_utf8() { + assert_eq!( + percent_decode("gemini%2D2.5").as_deref(), + Some("gemini-2.5") + ); + assert!(percent_decode("%ZZ").is_none()); + assert!(percent_decode("%ff").is_none()); + } + + #[test] + fn unavailable_primary_identity_is_direct_only_and_cannot_self_anchor_from_failover() { + let mut anchor = ProtocolAnchor::DirectOnlyPending; + assert_eq!( + begin_protocol_materialization(&mut anchor, false), + Some(true) + ); + assert_eq!(begin_protocol_materialization(&mut anchor, true), None); + assert!(matches!(anchor, ProtocolAnchor::Ineligible)); + } + + #[test] + fn skipped_candidates_do_not_reduce_the_network_retry_budget() { + let mut budget = NetworkAttemptBudget::new(2); + + // Circuit, protocol, and materialization skips never call begin_send. + for _ in 0..4 { + assert!(budget.has_remaining()); + } + assert!(budget.begin_send()); + assert!(budget.has_remaining()); + assert!(budget.begin_send()); + assert!(!budget.has_remaining()); + assert!(!budget.begin_send()); + } + + #[test] + fn protocol_anchor_blocks_replay_when_identity_is_unpredictable_or_changes() { + let mut unpredictable = ProtocolAnchor::DirectOnlyPending; + assert_eq!( + begin_protocol_materialization(&mut unpredictable, false), + Some(true) + ); + assert_eq!( + begin_protocol_materialization(&mut unpredictable, false), + None + ); + assert_eq!( + begin_protocol_materialization(&mut unpredictable, true), + None + ); + + let identity = Some(("openai-responses".to_string(), HeaderMap::new())); + let mut predictable = ProtocolAnchor::Unset; + assert!(protocol_identity_allows_attempt( + &mut predictable, + identity.clone() + )); + assert!(protocol_identity_allows_attempt(&mut predictable, identity)); + let mut changed_headers = HeaderMap::new(); + changed_headers.insert("openai-version", http::HeaderValue::from_static("changed")); + assert!(!protocol_identity_allows_attempt( + &mut predictable, + Some(("openai-responses".to_string(), changed_headers)) + )); + } + + #[test] + fn gateway_header_filters_share_protected_and_dynamic_hop_by_hop_rules() { + let mut incoming = HeaderMap::new(); + incoming.insert( + CONNECTION, + http::HeaderValue::from_static("x-private-hop, x-another-hop"), + ); + incoming.insert( + "x-private-hop", + http::HeaderValue::from_static("must-not-forward"), + ); + incoming.insert( + "x-another-hop", + http::HeaderValue::from_static("must-not-forward"), + ); + incoming.insert( + "cf-connecting-ip", + http::HeaderValue::from_static("203.0.113.5"), + ); + incoming.insert("traceparent", http::HeaderValue::from_static("00-spoofed")); + incoming.insert( + "x-candidate-local", + http::HeaderValue::from_static("preserved"), + ); + + let request = filtered_incoming_headers(&incoming); + assert!(request.get("x-private-hop").is_none()); + assert!(request.get("x-another-hop").is_none()); + assert!(request.get("cf-connecting-ip").is_none()); + assert!(request.get("traceparent").is_none()); + assert_eq!(request["x-candidate-local"], "preserved"); + + let response = filtered_response_headers(&incoming); + assert!(response.get("x-private-hop").is_none()); + assert!(response.get("x-another-hop").is_none()); + assert_eq!(response["x-candidate-local"], "preserved"); + } + + #[test] + fn streaming_health_waits_for_the_terminal_outcome() { + let complete = pi_stream_disposition( + StatusCode::OK, + &ProviderHealthOutcome::from_status(StatusCode::OK, None), + &PiStreamTermination::Complete { + captured: Some(Vec::new()), + }, + ); + assert_eq!(complete.provider_health, ProviderHealthDisposition::Healthy); + assert!(complete.request_success); + + let truncated = pi_stream_disposition( + StatusCode::OK, + &ProviderHealthOutcome::from_status(StatusCode::OK, None), + &PiStreamTermination::UpstreamFailure { + message: "truncated".to_string(), + }, + ); + assert_eq!( + truncated.provider_health, + ProviderHealthDisposition::Unhealthy + ); + assert!(!truncated.request_success); + assert_eq!(truncated.provider_error.as_deref(), Some("truncated")); + + let client_drop = pi_stream_disposition( + StatusCode::OK, + &ProviderHealthOutcome::from_status(StatusCode::OK, None), + &PiStreamTermination::DownstreamDropped, + ); + assert_eq!( + client_drop.provider_health, + ProviderHealthDisposition::Neutral + ); + assert!(!client_drop.request_success); + + let non_retryable_error = pi_stream_disposition( + StatusCode::BAD_REQUEST, + &ProviderHealthOutcome::from_status(StatusCode::BAD_REQUEST, None), + &PiStreamTermination::Complete { captured: None }, + ); + assert_eq!( + non_retryable_error.provider_health, + ProviderHealthDisposition::Healthy + ); + assert!(!non_retryable_error.request_success); + + for status in [StatusCode::UNAUTHORIZED, StatusCode::FORBIDDEN] { + let credential_rejection = pi_stream_disposition( + status, + &ProviderHealthOutcome::from_status(status, None), + &PiStreamTermination::Complete { captured: None }, + ); + assert_eq!( + credential_rejection.provider_health, + ProviderHealthDisposition::Neutral + ); + assert!(!credential_rejection.request_success); + } + } +} diff --git a/src-tauri/src/proxy/pi_runtime.rs b/src-tauri/src/proxy/pi_runtime.rs new file mode 100644 index 000000000..3558c3c84 --- /dev/null +++ b/src-tauri/src/proxy/pi_runtime.rs @@ -0,0 +1,1402 @@ +//! Immutable Pi gateway catalog and native projection planning. +//! +//! The database remains the managed provider authority. A snapshot is built +//! from complete provider aggregates, exact-key ownership claims, and a stable +//! device token; only after the matching `models.json` patch succeeds is that +//! snapshot published for request admission. + +use crate::database::Database; +use crate::error::AppError; +use crate::pi_config::composer::PiComposedNativeModel; +use crate::pi_config::gateway::{ + assess_composition_for_runtime, CandidateHeaderPlan, MaterializedCandidate, PiGatewayApiFamily, + PiGatewayCapability, PiGatewayReason, +}; +use crate::pi_config::model::PiManagedProviderConfig; +use crate::pi_config::native::compose_managed_pi_provider; +use crate::provider::ProviderAggregate; +use crate::proxy::types::AppProxyConfig; +use crate::settings::GatewayToken; +use indexmap::IndexMap; +use serde_json::{Map, Value}; +use sha2::{Digest, Sha256}; +use std::collections::{BTreeMap, HashMap}; +use std::fmt; +use std::io::Read; +use std::process::{Command, Stdio}; +use std::sync::{mpsc, Arc, RwLock}; +use std::time::{Duration, Instant}; +use tokio::sync::{OwnedRwLockReadGuard, RwLock as AsyncRwLock}; +use url::Url; + +const PI_APP: &str = "pi"; +const COMMAND_TIMEOUT: Duration = Duration::from_secs(10); +const COMMAND_OUTPUT_LIMIT: u64 = 1024 * 1024; + +#[derive(Debug, Clone)] +struct PiRuntimeModel { + provider_id: String, + provider_name: String, + family: PiGatewayApiFamily, + wire_profile: Vec, + plan: CandidateHeaderPlan, + custom_endpoint_plans: Vec, +} + +#[derive(Debug, Clone)] +struct PiRuntimeProvider { + models: HashMap, +} + +#[derive(Debug, Clone)] +struct PiRouteBinding { + provider_id: String, +} + +#[derive(Clone)] +struct PiNativeProjectionWitness(IndexMap>); + +impl fmt::Debug for PiNativeProjectionWitness { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + let keys = self + .0 + .iter() + .map(|(key, value)| (key, if value.is_some() { "present" } else { "absent" })) + .collect::>(); + formatter + .debug_tuple("PiNativeProjectionWitness") + .field(&keys) + .finish() + } +} + +/// One immutable catalog matching a successfully published native projection. +#[derive(Clone)] +pub(crate) struct PiRuntimeSnapshot { + pub(crate) server_generation: u64, + pub(crate) catalog_epoch: u64, + gateway_token: GatewayToken, + /// Exact native values which made this immutable runtime reachable. + /// + /// A fenced runtime may only be re-published while these values still + /// match `models.json`; retaining the projection alongside the routes + /// avoids reconstructing ownership from mutable database/settings state. + native_projection: PiNativeProjectionWitness, + app_config: AppProxyConfig, + providers: HashMap, + failover_ids: Vec, + routes: HashMap, +} + +impl fmt::Debug for PiRuntimeSnapshot { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("PiRuntimeSnapshot") + .field("server_generation", &self.server_generation) + .field("catalog_epoch", &self.catalog_epoch) + .field("gateway_token", &self.gateway_token) + .field("native_projection", &self.native_projection) + .field("provider_count", &self.providers.len()) + .field("failover_ids", &self.failover_ids) + .field("route_count", &self.routes.len()) + .finish_non_exhaustive() + } +} + +#[derive(Debug, Clone)] +pub(crate) struct PiRequestCandidate { + pub(crate) provider_id: String, + pub(crate) provider_name: String, + pub(crate) family: PiGatewayApiFamily, + pub(crate) plan: CandidateHeaderPlan, + pub(crate) is_failover: bool, +} + +#[derive(Debug, Clone)] +pub(crate) struct PiRequestRoute { + pub(crate) catalog_epoch: u64, + pub(crate) app_config: AppProxyConfig, + pub(crate) candidates: Vec, +} + +#[derive(Debug)] +pub(crate) struct PiMaterializedAttempt { + pub(crate) provider_id: String, + pub(crate) provider_name: String, + pub(crate) is_failover: bool, + pub(crate) transport: MaterializedCandidate, + pub(crate) url: Url, +} + +pub(crate) struct PiRuntimeBuild { + pub(crate) snapshot: Arc, + /// Exact keys only. A direct-only provider deliberately keeps its original + /// database projection while proxyable siblings point at the gateway. + pub(crate) projection_patch: IndexMap>, + pub(crate) direct_only_provider_ids: Vec, +} + +impl PiRuntimeSnapshot { + pub(crate) fn token_matches(&self, candidate: &str) -> bool { + self.gateway_token.constant_time_eq(candidate) + } + + pub(crate) fn route( + &self, + route_token: &str, + family: PiGatewayApiFamily, + model_id: &str, + ) -> Result { + let binding = self + .routes + .get(route_token) + .ok_or_else(|| AppError::NotFound("unknown Pi gateway provider route".to_string()))?; + let primary = self + .providers + .get(&binding.provider_id) + .and_then(|provider| provider.models.get(model_id)) + .filter(|model| model.family == family) + .ok_or_else(|| { + AppError::InvalidInput(format!( + "Pi provider '{}' does not expose model '{model_id}' for {}", + binding.provider_id, + family.as_str() + )) + })?; + + let mut candidates = expand_model_attempts(primary, false); + if self.app_config.auto_failover_enabled { + for provider_id in &self.failover_ids { + if provider_id == &binding.provider_id { + continue; + } + let Some(candidate) = self + .providers + .get(provider_id) + .and_then(|provider| provider.models.get(model_id)) + else { + continue; + }; + if candidate.family != family + || candidate.wire_profile != primary.wire_profile + || !candidate.plan.protocol_identity_is_predictable() + { + continue; + } + candidates.extend(expand_model_attempts(candidate, true)); + } + } + + Ok(PiRequestRoute { + catalog_epoch: self.catalog_epoch, + app_config: self.app_config.clone(), + candidates, + }) + } +} + +fn expand_model_attempts(model: &PiRuntimeModel, is_failover: bool) -> Vec { + let mut plans = Vec::with_capacity(model.custom_endpoint_plans.len().saturating_add(1)); + plans.push(model.plan.clone()); + plans.extend(model.custom_endpoint_plans.iter().cloned()); + plans + .into_iter() + .map(|plan| PiRequestCandidate { + provider_id: model.provider_id.clone(), + provider_name: model.provider_name.clone(), + family: model.family, + plan, + is_failover, + }) + .collect() +} + +impl PiRequestCandidate { + pub(crate) fn materialize( + self, + forwarded_path_and_query: &str, + ) -> Result { + let resolver_failure = std::cell::Cell::new(false); + let transport = self + .plan + .materialize_for_runtime(&|expression: &str| { + let resolved = resolve_pi_config_value(expression); + resolver_failure.set(resolver_failure.get() || resolved.is_err()); + resolved.ok() + }) + .map_err(|reason| { + if resolver_failure.get() { + AppError::Config( + "failed to resolve a deferred Pi gateway credential or header".to_string(), + ) + } else { + gateway_reason(reason) + } + })?; + let url = build_family_url(self.family, &transport.endpoint, forwarded_path_and_query)?; + Ok(PiMaterializedAttempt { + provider_id: self.provider_id, + provider_name: self.provider_name, + is_failover: self.is_failover, + transport, + url, + }) + } + + pub(crate) fn protocol_identity_is_predictable(&self) -> bool { + self.plan.protocol_identity_is_predictable() + } + + pub(crate) fn planned_protocol_identity( + &self, + ) -> Result, AppError> { + let resolver_failure = std::cell::Cell::new(false); + let identity = self + .plan + .materialize_protocol_identity(&|expression: &str| { + // Protocol !commands are marked unpredictable before this + // method is called. An auth command must not be run merely to + // decide whether a circuit-skipped primary permits failover. + if expression.starts_with('!') { + resolver_failure.set(true); + return None; + } + let resolved = resolve_pi_config_value(expression); + resolver_failure.set(resolver_failure.get() || resolved.is_err()); + resolved.ok() + }) + .map_err(|reason| { + if resolver_failure.get() { + AppError::Config( + "failed to pre-resolve Pi primary protocol identity".to_string(), + ) + } else { + gateway_reason(reason) + } + })?; + Ok(identity.map(|(family, headers)| (family.as_str().to_string(), headers))) + } +} + +fn gateway_reason(reason: PiGatewayReason) -> AppError { + AppError::Config(format!( + "Pi gateway candidate rejected at {}: {:?}", + reason.json_pointer, reason.code + )) +} + +/// Build the immutable runtime and its exact native projection in one pass. +pub(crate) fn build_pi_runtime( + db: &Database, + server_generation: u64, + catalog_epoch: u64, + gateway_origin: &Url, + gateway_token: GatewayToken, + app_config: AppProxyConfig, +) -> Result { + if catalog_epoch % 2 != 0 { + return Err(AppError::Config( + "Pi runtime publication requires an even catalog epoch".to_string(), + )); + } + let aggregates = db.get_all_provider_aggregates(PI_APP)?; + let manifest = db.get_pi_projection_manifest()?; + if aggregates.len() != manifest.len() + || aggregates + .keys() + .any(|provider_id| !manifest.contains_key(provider_id)) + { + return Err(AppError::Conflict( + "Pi provider aggregates and exact-key ownership claims diverged".to_string(), + )); + } + + let mut providers = HashMap::new(); + let mut routes = HashMap::new(); + let mut projection_patch = IndexMap::new(); + let mut direct_only_provider_ids = Vec::new(); + for (provider_id, aggregate) in aggregates { + let projection = manifest.get(&provider_id).ok_or_else(|| { + AppError::Conflict(format!( + "Pi provider '{provider_id}' has no exact-key claim" + )) + })?; + let config = decode_managed_config(&aggregate)?; + let composition = compose_managed_pi_provider(&projection.provider_key, &config)?; + let assessment = assess_composition_for_runtime(&composition); + if assessment.capability != PiGatewayCapability::Proxyable + || assessment.plans.len() != composition.models.len() + { + projection_patch.insert( + projection.provider_key.clone(), + Some(serde_json::to_value(&config).map_err(|source| { + AppError::Config(format!( + "failed to serialize direct-only Pi provider: {source}" + )) + })?), + ); + direct_only_provider_ids.push(provider_id); + continue; + } + + let token = pi_route_token(&provider_id, &projection.provider_key); + let local_base = gateway_origin + .join(&format!("pi/{token}")) + .map_err(|error| AppError::Config(format!("invalid Pi gateway origin: {error}")))?; + let endpoints = aggregate.endpoints.keys().cloned().collect::>(); + let mut models = HashMap::new(); + for ((model, plan), expected) in composition + .models + .iter() + .zip(assessment.plans) + .zip(config.models.iter()) + { + if model.id != expected.id { + return Err(AppError::Config(format!( + "Pi composer changed managed model order for '{provider_id}'" + ))); + } + let family = plan.family(); + let runtime_model = runtime_model(&aggregate, model, family, plan, endpoints.clone())?; + if models.insert(model.id.clone(), runtime_model).is_some() { + return Err(AppError::Conflict(format!( + "duplicate Pi model '{}' in provider '{provider_id}'", + model.id + ))); + } + } + if routes + .insert( + token, + PiRouteBinding { + provider_id: provider_id.clone(), + }, + ) + .is_some() + { + return Err(AppError::Conflict( + "Pi gateway route digest collision".to_string(), + )); + } + let projected = project_config_for_gateway(&config, &local_base, &gateway_token)?; + projection_patch.insert(projection.provider_key.clone(), Some(projected)); + providers.insert(provider_id, PiRuntimeProvider { models }); + } + + let failover_ids = db + .get_failover_queue(PI_APP)? + .into_iter() + .map(|item| item.provider_id) + .collect(); + Ok(PiRuntimeBuild { + snapshot: Arc::new(PiRuntimeSnapshot { + server_generation, + catalog_epoch, + gateway_token, + native_projection: PiNativeProjectionWitness(projection_patch.clone()), + app_config, + providers, + failover_ids, + routes, + }), + projection_patch, + direct_only_provider_ids, + }) +} + +pub(crate) fn direct_pi_projection_patch( + db: &Database, +) -> Result>, AppError> { + let aggregates = db.get_all_provider_aggregates(PI_APP)?; + let manifest = db.get_pi_projection_manifest()?; + if aggregates.len() != manifest.len() { + return Err(AppError::Conflict( + "Pi provider aggregates and exact-key claims diverged".to_string(), + )); + } + let mut patch = IndexMap::new(); + for (provider_id, aggregate) in aggregates { + let projection = manifest.get(&provider_id).ok_or_else(|| { + AppError::Conflict(format!( + "Pi provider '{provider_id}' has no exact-key claim" + )) + })?; + let config = decode_managed_config(&aggregate)?; + patch.insert( + projection.provider_key.clone(), + Some(serde_json::to_value(config).map_err(|source| { + AppError::Config(format!("failed to serialize Pi provider: {source}")) + })?), + ); + } + Ok(patch) +} + +/// Render one managed provider for an in-progress catalog mutation. This is +/// the same planning boundary used by the full runtime build, so the +/// coordinator never carries a second notion of "proxyable". +pub(crate) fn project_managed_pi_config( + provider_id: &str, + provider_key: &str, + config: &PiManagedProviderConfig, + gateway_origin: &Url, + gateway_token: &GatewayToken, +) -> Result { + let composition = compose_managed_pi_provider(provider_key, config)?; + let assessment = assess_composition_for_runtime(&composition); + if assessment.capability != PiGatewayCapability::Proxyable + || assessment.plans.len() != composition.models.len() + { + return serde_json::to_value(config).map_err(|source| AppError::JsonSerialize { source }); + } + let token = pi_route_token(provider_id, provider_key); + let local_base = gateway_origin + .join(&format!("pi/{token}")) + .map_err(|error| AppError::Config(format!("invalid Pi gateway origin: {error}")))?; + project_config_for_gateway(config, &local_base, gateway_token) +} + +fn decode_managed_config( + aggregate: &ProviderAggregate, +) -> Result { + serde_json::from_value(aggregate.provider.settings_config.clone()).map_err(|error| { + AppError::Config(format!( + "managed Pi provider '{}' is invalid: {error}", + aggregate.provider.id + )) + }) +} + +fn runtime_model( + aggregate: &ProviderAggregate, + model: &PiComposedNativeModel, + family: PiGatewayApiFamily, + plan: CandidateHeaderPlan, + endpoints: Vec, +) -> Result { + let custom_endpoint_plans = + build_custom_endpoint_plans(&plan, endpoints, &aggregate.provider.id); + Ok(PiRuntimeModel { + provider_id: aggregate.provider.id.clone(), + provider_name: aggregate.provider.name.clone(), + family, + wire_profile: canonical_wire_profile(model)?, + plan, + custom_endpoint_plans, + }) +} + +fn build_custom_endpoint_plans( + primary: &CandidateHeaderPlan, + endpoints: Vec, + provider_id: &str, +) -> Vec { + let mut plans = Vec::::with_capacity(endpoints.len()); + for endpoint in endpoints { + match primary.with_endpoint(&endpoint) { + Ok(candidate) + if candidate.endpoint() != primary.endpoint() + && !plans + .iter() + .any(|existing| existing.endpoint() == candidate.endpoint()) => + { + plans.push(candidate); + } + Ok(_) => {} + Err(reason) => { + // The write boundary rejects these values. Keeping this + // defensive compatibility path prevents an old/corrupt + // auxiliary endpoint from disabling the valid primary route. + // Never log the URL: it may contain the very userinfo which + // caused rejection. + log::warn!( + "ignoring invalid persisted Pi custom endpoint for provider \ + '{provider_id}': {:?}", + reason.code + ); + } + } + } + plans +} + +fn canonical_wire_profile(model: &PiComposedNativeModel) -> Result, AppError> { + let mut profile = Map::new(); + profile.insert("reasoning".to_string(), Value::Bool(model.reasoning)); + profile.insert( + "thinkingLevelMap".to_string(), + model.thinking_level_map.clone().unwrap_or(Value::Null), + ); + profile.insert("input".to_string(), model.input.clone()); + profile.insert("contextWindow".to_string(), model.context_window.clone()); + profile.insert("maxTokens".to_string(), model.max_tokens.clone()); + profile.insert( + "compat".to_string(), + model.compat.clone().unwrap_or(Value::Null), + ); + profile.insert( + "providerExtra".to_string(), + serde_json::to_value(&model.provider_extra) + .map_err(|source| AppError::JsonSerialize { source })?, + ); + profile.insert( + "modelExtra".to_string(), + serde_json::to_value(&model.model_extra) + .map_err(|source| AppError::JsonSerialize { source })?, + ); + profile.insert( + "overrideExtra".to_string(), + serde_json::to_value(&model.override_extra) + .map_err(|source| AppError::JsonSerialize { source })?, + ); + serde_json::to_vec(&canonical_json(&Value::Object(profile))) + .map_err(|source| AppError::JsonSerialize { source }) +} + +fn canonical_json(value: &Value) -> Value { + match value { + Value::Array(values) => Value::Array(values.iter().map(canonical_json).collect()), + Value::Object(values) => { + let sorted = values + .iter() + .map(|(key, value)| (key.clone(), canonical_json(value))) + .collect::>(); + Value::Object(sorted.into_iter().collect()) + } + scalar => scalar.clone(), + } +} + +fn project_config_for_gateway( + config: &PiManagedProviderConfig, + local_base: &Url, + gateway_token: &GatewayToken, +) -> Result { + let mut value = + serde_json::to_value(config).map_err(|source| AppError::JsonSerialize { source })?; + let root = value + .as_object_mut() + .ok_or_else(|| AppError::Config("Pi provider projection is not an object".to_string()))?; + root.insert( + "apiKey".to_string(), + Value::String(gateway_token.expose().to_string()), + ); + root.remove("headers"); + root.remove("authHeader"); + root.remove("oauth"); + let provider_has_base = root.contains_key("baseUrl"); + if provider_has_base { + root.insert( + "baseUrl".to_string(), + Value::String(local_base.as_str().trim_end_matches('/').to_string()), + ); + } + let models = root + .get_mut("models") + .and_then(Value::as_array_mut) + .ok_or_else(|| AppError::Config("Pi provider projection has no models".to_string()))?; + for model in models { + let object = model + .as_object_mut() + .ok_or_else(|| AppError::Config("Pi model projection is not an object".to_string()))?; + object.remove("headers"); + if object.contains_key("baseUrl") || !provider_has_base { + object.insert( + "baseUrl".to_string(), + Value::String(local_base.as_str().trim_end_matches('/').to_string()), + ); + } + } + if let Some(overrides) = root + .get_mut("modelOverrides") + .and_then(Value::as_object_mut) + { + for model_override in overrides.values_mut() { + if let Some(object) = model_override.as_object_mut() { + object.remove("headers"); + } + } + } + Ok(value) +} + +fn pi_route_token(provider_id: &str, provider_key: &str) -> String { + let mut digest = Sha256::new(); + digest.update(b"cc-switch:pi-route:v2\0"); + digest.update(provider_id.as_bytes()); + digest.update([0]); + digest.update(provider_key.as_bytes()); + digest + .finalize() + .iter() + .map(|byte| format!("{byte:02x}")) + .collect() +} + +/// Process-local publication point. Odd epochs close admission; an even +/// snapshot is leased by `Arc`, so requests already admitted keep a coherent +/// catalog while a replacement is prepared. +#[derive(Debug, Default)] +pub(crate) struct PiRuntimeStore { + publication: RwLock, + epoch_gate: Arc>, +} + +#[derive(Debug, Default)] +struct PiRuntimePublication { + current: Option>, + catalog_epoch: u64, +} + +impl PiRuntimeStore { + pub(crate) async fn begin_mutation(&self) -> u64 { + let _guard = self.epoch_gate.write().await; + let mut publication = self + .publication + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let current = publication.catalog_epoch; + let odd = if current % 2 == 0 { + current.saturating_add(1) + } else { + current + }; + publication.catalog_epoch = odd; + odd.saturating_add(1) + } + + pub(crate) fn next_even_epoch(&self) -> Result { + let current = self + .publication + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .catalog_epoch; + if current % 2 != 0 { + return Err(AppError::Conflict( + "cannot publish a sorted Pi runtime while catalog admission is fenced".to_string(), + )); + } + let next = current.saturating_add(2); + if next % 2 != 0 { + return Err(AppError::Config( + "Pi catalog epoch overflowed its even publication sequence".to_string(), + )); + } + Ok(next) + } + + pub(crate) async fn publish(&self, snapshot: Arc) -> Result<(), AppError> { + if snapshot.catalog_epoch % 2 != 0 { + return Err(AppError::Config( + "cannot publish an odd Pi catalog epoch".to_string(), + )); + } + let _guard = self.epoch_gate.write().await; + let mut publication = self + .publication + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + publication.catalog_epoch = snapshot.catalog_epoch; + publication.current = Some(snapshot); + Ok(()) + } + + pub(crate) async fn close(&self, even_epoch: u64) -> Result<(), AppError> { + if even_epoch % 2 != 0 { + return Err(AppError::Config( + "Pi admission close requires an even terminal epoch".to_string(), + )); + } + let _guard = self.epoch_gate.write().await; + let mut publication = self + .publication + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + publication.current = None; + publication.catalog_epoch = even_epoch; + Ok(()) + } + + pub(crate) async fn republish_current(&self, even_epoch: u64) -> Result { + if even_epoch % 2 != 0 { + return Err(AppError::Config( + "Pi runtime re-publication requires an even epoch".to_string(), + )); + } + let _guard = self.epoch_gate.write().await; + let mut publication = self + .publication + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let current = publication.current.as_ref().cloned(); + let Some(current) = current else { + publication.catalog_epoch = even_epoch; + return Ok(false); + }; + let mut next = (*current).clone(); + next.catalog_epoch = even_epoch; + publication.current = Some(Arc::new(next)); + publication.catalog_epoch = even_epoch; + Ok(true) + } + + pub(crate) fn lease(&self, server_generation: u64) -> Option> { + let publication = self + .publication + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if publication.catalog_epoch % 2 != 0 { + return None; + } + publication + .current + .as_ref() + .filter(|snapshot| { + snapshot.server_generation == server_generation + && snapshot.catalog_epoch == publication.catalog_epoch + }) + .cloned() + } + + pub(crate) fn is_admitting(&self, server_generation: u64) -> bool { + self.lease(server_generation).is_some() + } + + /// Retain the credential of the last published generation for exact + /// native-projection compensation even while an odd epoch fences new + /// admission. This does not mint or rotate credentials and never crosses + /// IPC; it is only an ownership witness for restoring `models.json`. + pub(crate) fn retained_gateway_token( + &self, + server_generation: u64, + ) -> Option { + self.publication + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .current + .as_ref() + .filter(|snapshot| snapshot.server_generation == server_generation) + .map(|snapshot| snapshot.gateway_token.clone()) + } + + /// Return the exact native projection paired with the fenced runtime. + /// This remains available while admission is at an odd epoch so recovery + /// can prove that re-publication would still describe Pi's live file. + pub(crate) fn retained_native_projection(&self) -> Option>> { + self.publication + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .current + .as_ref() + .map(|snapshot| snapshot.native_projection.0.clone()) + } + + pub(crate) async fn admission_guard( + self: &Arc, + server_generation: u64, + snapshot: &Arc, + ) -> Option> { + let guard = self.epoch_gate.clone().read_owned().await; + let publication = self + .publication + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let current = publication.current.as_ref().is_some_and(|current| { + snapshot.catalog_epoch % 2 == 0 + && publication.catalog_epoch == snapshot.catalog_epoch + && current.server_generation == server_generation + && Arc::ptr_eq(current, snapshot) + }); + current.then_some(guard) + } + + pub(crate) async fn writeback_guard( + self: &Arc, + expected_epoch: u64, + ) -> Option> { + let guard = self.epoch_gate.clone().read_owned().await; + let current = self + .publication + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .catalog_epoch; + (expected_epoch % 2 == 0 && current == expected_epoch).then_some(guard) + } +} + +pub(crate) fn infer_family(path: &str) -> Option { + if path == "/v1/messages" { + Some(PiGatewayApiFamily::AnthropicMessages) + } else if path == "/chat/completions" { + Some(PiGatewayApiFamily::OpenAiCompletions) + } else if matches!(path, "/responses" | "/responses/compact") { + Some(PiGatewayApiFamily::OpenAiResponses) + } else if path.starts_with("/models/") || path == "/models" { + Some(PiGatewayApiFamily::GoogleGenerativeAi) + } else { + None + } +} + +fn build_family_url( + family: PiGatewayApiFamily, + base: &Url, + path_and_query: &str, +) -> Result { + let (path, query) = path_and_query + .split_once('?') + .map_or((path_and_query, None), |(path, query)| (path, Some(query))); + if infer_family(path) != Some(family) { + return Err(AppError::InvalidInput(format!( + "Pi gateway path '{path}' does not match {}", + family.as_str() + ))); + } + let mut url = base.clone(); + let base_path = base.path().trim_end_matches('/'); + let suffix = path.trim_start_matches('/'); + let combined = if base_path.is_empty() || base_path == "/" { + format!("/{suffix}") + } else { + format!("{base_path}/{suffix}") + }; + url.set_path(&combined); + url.set_query(query); + url.set_fragment(None); + Ok(url) +} + +fn resolve_pi_config_value(expression: &str) -> Result { + if let Some(command) = expression.strip_prefix('!') { + return execute_config_command(command); + } + expand_environment(expression) +} + +fn expand_environment(input: &str) -> Result { + const ESCAPED_DOLLAR: char = '\u{e000}'; + const ESCAPED_BANG: char = '\u{e001}'; + let chars = input.chars().collect::>(); + let mut output = String::new(); + let mut index = 0; + while index < chars.len() { + if chars[index] != '$' { + output.push(chars[index]); + index += 1; + continue; + } + if chars.get(index + 1) == Some(&'$') { + output.push(ESCAPED_DOLLAR); + index += 2; + continue; + } + if chars.get(index + 1) == Some(&'!') { + output.push(ESCAPED_BANG); + index += 2; + continue; + } + let (name, next) = if chars.get(index + 1) == Some(&'{') { + let Some(end) = chars[index + 2..].iter().position(|value| *value == '}') else { + return Err("unterminated Pi environment expression".to_string()); + }; + let end = index + 2 + end; + (chars[index + 2..end].iter().collect::(), end + 1) + } else { + let mut end = index + 1; + while end < chars.len() && (chars[end] == '_' || chars[end].is_ascii_alphanumeric()) { + end += 1; + } + if end == index + 1 { + output.push('$'); + index += 1; + continue; + } + (chars[index + 1..end].iter().collect::(), end) + }; + if name.is_empty() + || !name + .chars() + .next() + .is_some_and(|value| value == '_' || value.is_ascii_alphabetic()) + { + return Err("invalid Pi environment variable name".to_string()); + } + let value = std::env::var(&name) + .map_err(|_| format!("Pi environment variable '{name}' is unavailable"))?; + output.push_str(&value); + index = next; + } + Ok(output + .replace(ESCAPED_DOLLAR, "$") + .replace(ESCAPED_BANG, "!")) +} + +fn execute_config_command(script: &str) -> Result { + if script.trim().is_empty() { + return Err("empty Pi config command".to_string()); + } + let mut command = if cfg!(windows) { + let mut command = Command::new("cmd"); + command.args(["/D", "/S", "/C", script]); + command + } else { + let mut command = Command::new("/bin/sh"); + command.args(["-c", script]); + #[cfg(unix)] + { + use std::os::unix::process::CommandExt; + command.process_group(0); + } + command + }; + #[cfg(windows)] + { + use std::os::windows::process::CommandExt; + use windows_sys::Win32::System::Threading::CREATE_SUSPENDED; + // The shell must not execute user code before it belongs to the + // kill-on-close Job Object. Its primary thread is resumed only after + // CommandTree::attach succeeds. + command.creation_flags(CREATE_SUSPENDED); + } + let mut child = command + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .map_err(|error| format!("failed to start Pi config command: {error}"))?; + let command_tree = CommandTree::attach(&mut child)?; + #[cfg(windows)] + if let Err(error) = resume_suspended_process(child.id()) { + command_tree.terminate(&mut child); + let _ = child.wait(); + return Err(error); + } + let stdout = child + .stdout + .take() + .ok_or_else(|| "failed to capture Pi config command stdout".to_string())?; + let stderr = child + .stderr + .take() + .ok_or_else(|| "failed to capture Pi config command stderr".to_string())?; + let (stdout_sender, stdout_reader) = mpsc::sync_channel(1); + let (stderr_sender, stderr_reader) = mpsc::sync_channel(1); + std::thread::spawn(move || { + let _ = stdout_sender.send(read_bounded(stdout)); + }); + std::thread::spawn(move || { + let _ = stderr_sender.send(read_bounded(stderr)); + }); + let started = Instant::now(); + let status = loop { + match child.try_wait() { + Ok(Some(status)) => break status, + Ok(None) if started.elapsed() < COMMAND_TIMEOUT => { + std::thread::sleep(Duration::from_millis(10)); + } + Ok(None) => { + command_tree.terminate(&mut child); + let _ = child.wait(); + return Err("Pi config command timed out".to_string()); + } + Err(error) => { + command_tree.terminate(&mut child); + let _ = child.wait(); + return Err(format!("failed to wait for Pi config command: {error}")); + } + } + }; + // A successful shell may leave descendants holding inherited pipe handles. + // Terminate the whole tree before draining, and keep the original deadline + // over both process wait and output collection. + command_tree.terminate(&mut child); + let deadline = started + COMMAND_TIMEOUT; + let stdout = receive_command_output( + &stdout_reader, + deadline, + "Pi config command stdout did not close before timeout", + )??; + let stderr = receive_command_output( + &stderr_reader, + deadline, + "Pi config command stderr did not close before timeout", + )??; + if !status.success() { + return Err(format!( + "Pi config command exited unsuccessfully: {}", + String::from_utf8_lossy(&stderr).trim() + )); + } + String::from_utf8(stdout) + .map(|value| value.trim().to_string()) + .map_err(|_| "Pi config command output is not UTF-8".to_string()) +} + +fn read_bounded(reader: impl Read) -> Result, String> { + let mut output = Vec::new(); + reader + .take(COMMAND_OUTPUT_LIMIT + 1) + .read_to_end(&mut output) + .map_err(|error| format!("failed to read Pi config command output: {error}"))?; + if output.len() as u64 > COMMAND_OUTPUT_LIMIT { + return Err("Pi config command output exceeded 1 MiB".to_string()); + } + Ok(output) +} + +fn receive_command_output( + receiver: &mpsc::Receiver, String>>, + deadline: Instant, + timeout_message: &str, +) -> Result, String>, String> { + let remaining = deadline.saturating_duration_since(Instant::now()); + receiver + .recv_timeout(remaining) + .map_err(|error| match error { + mpsc::RecvTimeoutError::Timeout => timeout_message.to_string(), + mpsc::RecvTimeoutError::Disconnected => { + "Pi config command output reader stopped unexpectedly".to_string() + } + }) +} + +struct CommandTree { + #[cfg(windows)] + job: windows_sys::Win32::Foundation::HANDLE, +} + +impl CommandTree { + fn attach(child: &mut std::process::Child) -> Result { + #[cfg(windows)] + { + use std::mem::size_of; + use std::os::windows::io::AsRawHandle; + use windows_sys::Win32::Foundation::CloseHandle; + use windows_sys::Win32::System::JobObjects::{ + AssignProcessToJobObject, CreateJobObjectW, JobObjectExtendedLimitInformation, + SetInformationJobObject, JOBOBJECT_EXTENDED_LIMIT_INFORMATION, + JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE, + }; + + // SAFETY: all pointers are either null or point to initialized + // values for the duration of their synchronous Win32 calls. + unsafe { + let job = CreateJobObjectW(std::ptr::null(), std::ptr::null()); + if job.is_null() { + let _ = child.kill(); + let _ = child.wait(); + return Err(format!( + "failed to create Pi config command job: {}", + std::io::Error::last_os_error() + )); + } + let mut limits = JOBOBJECT_EXTENDED_LIMIT_INFORMATION::default(); + limits.BasicLimitInformation.LimitFlags = JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE; + if SetInformationJobObject( + job, + JobObjectExtendedLimitInformation, + (&raw const limits).cast(), + size_of::() as u32, + ) == 0 + { + let error = std::io::Error::last_os_error(); + CloseHandle(job); + let _ = child.kill(); + let _ = child.wait(); + return Err(format!( + "failed to configure Pi config command job: {error}" + )); + } + if AssignProcessToJobObject(job, child.as_raw_handle() as _) == 0 { + let error = std::io::Error::last_os_error(); + CloseHandle(job); + let _ = child.kill(); + let _ = child.wait(); + return Err(format!( + "failed to assign Pi config command to its job: {error}" + )); + } + Ok(Self { job }) + } + } + #[cfg(not(windows))] + { + let _ = child; + Ok(Self {}) + } + } + + fn terminate(&self, child: &mut std::process::Child) { + #[cfg(unix)] + unsafe { + let _ = libc::kill(-(child.id() as i32), libc::SIGKILL); + } + #[cfg(windows)] + unsafe { + let _ = child; + let _ = windows_sys::Win32::System::JobObjects::TerminateJobObject(self.job, 1); + } + #[cfg(not(any(unix, windows)))] + { + let _ = child.kill(); + } + } +} + +#[cfg(windows)] +fn resume_suspended_process(process_id: u32) -> Result<(), String> { + use std::mem::size_of; + use windows_sys::Win32::Foundation::{CloseHandle, INVALID_HANDLE_VALUE}; + use windows_sys::Win32::System::Diagnostics::ToolHelp::{ + CreateToolhelp32Snapshot, Thread32First, Thread32Next, TH32CS_SNAPTHREAD, THREADENTRY32, + }; + use windows_sys::Win32::System::Threading::{OpenThread, ResumeThread, THREAD_SUSPEND_RESUME}; + + // SAFETY: the snapshot and thread handles are checked before use and + // closed on every path. CREATE_SUSPENDED prevents the target from adding + // threads while this enumeration runs. + unsafe { + let snapshot = CreateToolhelp32Snapshot(TH32CS_SNAPTHREAD, 0); + if snapshot == INVALID_HANDLE_VALUE { + return Err(format!( + "failed to enumerate the suspended Pi config command: {}", + std::io::Error::last_os_error() + )); + } + let mut entry = THREADENTRY32 { + dwSize: size_of::() as u32, + ..THREADENTRY32::default() + }; + let mut has_entry = Thread32First(snapshot, &mut entry); + while has_entry != 0 { + if entry.th32OwnerProcessID == process_id { + let thread = OpenThread(THREAD_SUSPEND_RESUME, 0, entry.th32ThreadID); + if thread.is_null() { + let error = std::io::Error::last_os_error(); + CloseHandle(snapshot); + return Err(format!( + "failed to open the suspended Pi config command thread: {error}" + )); + } + let result = ResumeThread(thread); + let error = (result == u32::MAX).then(std::io::Error::last_os_error); + CloseHandle(thread); + CloseHandle(snapshot); + return match error { + Some(error) => Err(format!( + "failed to resume the suspended Pi config command: {error}" + )), + None => Ok(()), + }; + } + has_entry = Thread32Next(snapshot, &mut entry); + } + CloseHandle(snapshot); + } + Err("the suspended Pi config command had no primary thread".to_string()) +} + +#[cfg(windows)] +impl Drop for CommandTree { + fn drop(&mut self) { + // JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE is the final safety net if an + // early return occurs before explicit termination. + unsafe { + let _ = windows_sys::Win32::Foundation::CloseHandle(self.job); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn native_projection_debug_exposes_only_key_presence() { + let secret = "runtime-native-secret-never-log"; + let witness = PiNativeProjectionWitness(IndexMap::from([ + ( + "managed".to_string(), + Some(json!({"apiKey": secret, "headers": {"x-private": secret}})), + ), + ("removed".to_string(), None), + ])); + let debug = format!("{witness:?}"); + assert!(!debug.contains(secret)); + assert!(debug.contains("managed")); + assert!(debug.contains("present")); + assert!(debug.contains("removed")); + assert!(debug.contains("absent")); + } + + #[test] + fn invalid_persisted_custom_endpoint_is_quarantined_without_losing_primary_route() { + let config: PiManagedProviderConfig = serde_json::from_value(json!({ + "name": "Provider", + "api": "openai-responses", + "baseUrl": "https://primary.example/v1", + "apiKey": "literal", + "models": [{"id": "model-a"}] + })) + .expect("managed config"); + let composition = + compose_managed_pi_provider("provider", &config).expect("compose provider"); + let primary = assess_composition_for_runtime(&composition) + .plans + .into_iter() + .next() + .expect("primary plan"); + let custom_endpoint_plans = build_custom_endpoint_plans( + &primary, + vec![ + "https://user:secret@invalid.example/v1".to_string(), + "https://mirror.example/v1".to_string(), + "https://mirror.example/v1".to_string(), + "https://primary.example/v1".to_string(), + ], + "provider", + ); + assert_eq!( + custom_endpoint_plans + .iter() + .map(|plan| plan.endpoint().as_str()) + .collect::>(), + ["https://mirror.example/v1"] + ); + + let model = PiRuntimeModel { + provider_id: "provider".to_string(), + provider_name: "Provider".to_string(), + family: PiGatewayApiFamily::OpenAiResponses, + wire_profile: Vec::new(), + plan: primary, + custom_endpoint_plans, + }; + let attempts = expand_model_attempts(&model, false); + assert_eq!(attempts.len(), 2); + assert_eq!( + attempts[0].plan.endpoint().as_str(), + "https://primary.example/v1" + ); + assert_eq!( + attempts[1].plan.endpoint().as_str(), + "https://mirror.example/v1" + ); + } + + #[test] + fn environment_resolution_matches_vendored_transport_oracle() { + std::env::set_var("PI_RUNTIME_TEST_VALUE", "environment-secret"); + assert_eq!( + resolve_pi_config_value("prefix-${PI_RUNTIME_TEST_VALUE}-suffix").unwrap(), + "prefix-environment-secret-suffix" + ); + assert_eq!( + resolve_pi_config_value("$$literal-$!bang").unwrap(), + "$literal-!bang" + ); + std::env::remove_var("PI_RUNTIME_TEST_VALUE"); + } + + #[cfg(unix)] + #[test] + fn command_deadline_covers_descendants_holding_output_pipes() { + let started = Instant::now(); + let output = + execute_config_command("sleep 30 & printf inherited-pipe").expect("command output"); + assert_eq!(output, "inherited-pipe"); + assert!( + started.elapsed() < Duration::from_secs(2), + "background descendants must be terminated before output drain" + ); + } + + #[test] + fn command_resolution_matches_vendored_transport_oracle() { + #[cfg(unix)] + assert_eq!( + resolve_pi_config_value("!printf pi-command-value").unwrap(), + "pi-command-value" + ); + } + + #[cfg(windows)] + #[test] + fn command_job_assignment_precedes_descendant_execution() { + let started = Instant::now(); + for _ in 0..8 { + let output = execute_config_command( + "start \"\" /b cmd /D /S /C \"ping 127.0.0.1 -n 30 >nul\" & PiComposedNativeModel { + PiComposedNativeModel { + id: "m".to_string(), + name: "M".to_string(), + api: crate::pi_config::raw_schema::PiRawApiId::new("openai-responses".to_string()) + .unwrap(), + provider: "p".to_string(), + base_url: "https://example.test/v1".to_string(), + reasoning: false, + thinking_level_map: None, + input, + cost: json!({"input": 1}), + context_window: json!(1000), + max_tokens: json!(100), + headers: BTreeMap::new(), + provider_headers: Vec::new(), + model_headers: Vec::new(), + compat: None, + api_key: Some("secret".to_string()), + oauth: None, + auth_header: false, + provider_extra: serde_json::from_value(extra).unwrap(), + model_extra: BTreeMap::new(), + override_extra: BTreeMap::new(), + } + } + let first = model(json!({"z": 1, "a": {"b": 2, "a": 1}}), json!(["text"])); + let reordered = model(json!({"a": {"a": 1, "b": 2}, "z": 1}), json!(["text"])); + assert_eq!( + canonical_wire_profile(&first).unwrap(), + canonical_wire_profile(&reordered).unwrap() + ); + let changed = model( + json!({"z": 1, "a": {"b": 2, "a": 1}}), + json!(["image", "text"]), + ); + assert_ne!( + canonical_wire_profile(&first).unwrap(), + canonical_wire_profile(&changed).unwrap() + ); + } +} diff --git a/src-tauri/src/proxy/provider_router.rs b/src-tauri/src/proxy/provider_router.rs index 2baa11fa6..b552ec229 100644 --- a/src-tauri/src/proxy/provider_router.rs +++ b/src-tauri/src/proxy/provider_router.rs @@ -29,6 +29,16 @@ impl ProviderRouter { } } + async fn app_proxy_config( + &self, + app_type: &str, + ) -> Result { + if app_type == AppType::Pi.as_str() { + return Ok(crate::settings::get_pi_app_proxy_config()); + } + self.db.get_proxy_config_for_app(app_type).await + } + /// 选择可用的供应商(支持故障转移) /// /// 返回按优先级排序的可用供应商列表: @@ -40,7 +50,7 @@ impl ProviderRouter { let mut circuit_open_count = 0usize; // 检查该应用的自动故障转移开关是否开启(从 proxy_config 表读取) - let auto_failover_enabled = match self.db.get_proxy_config_for_app(app_type).await { + let auto_failover_enabled = match self.app_proxy_config(app_type).await { Ok(config) => config.auto_failover_enabled, Err(e) => { log::error!("[{app_type}] 读取 proxy_config 失败: {e},默认禁用故障转移"); @@ -132,7 +142,7 @@ impl ProviderRouter { error_msg: Option, ) -> Result<(), AppError> { // 1. 按应用独立获取熔断器配置 - let failure_threshold = match self.db.get_proxy_config_for_app(app_type).await { + let failure_threshold = match self.app_proxy_config(app_type).await { Ok(app_config) => app_config.circuit_failure_threshold, Err(_) => 5, // 默认值 }; @@ -251,7 +261,7 @@ impl ProviderRouter { let app_type = key.split(':').next().unwrap_or("claude"); // 按应用独立读取熔断器配置 - let config = match self.db.get_proxy_config_for_app(app_type).await { + let config = match self.app_proxy_config(app_type).await { Ok(app_config) => crate::proxy::circuit_breaker::CircuitBreakerConfig { failure_threshold: app_config.circuit_failure_threshold, success_threshold: app_config.circuit_success_threshold, diff --git a/src-tauri/src/proxy/providers/mod.rs b/src-tauri/src/proxy/providers/mod.rs index 625e1d0bb..ab0f4b35a 100644 --- a/src-tauri/src/proxy/providers/mod.rs +++ b/src-tauri/src/proxy/providers/mod.rs @@ -205,7 +205,11 @@ impl ProviderType { ProviderType::Gemini } AppType::GrokBuild => ProviderType::Codex, - AppType::OpenCode | AppType::OpenClaw | AppType::Hermes => ProviderType::Codex, + AppType::OpenCode | AppType::OpenClaw | AppType::Hermes | AppType::Pi => { + // Generic callers cannot infer Pi's wire family from AppType; + // the dedicated Pi runtime routes by effective model API. + ProviderType::Codex + } } } @@ -259,7 +263,11 @@ pub fn get_adapter(app_type: &AppType) -> Box { AppType::Codex => Box::new(CodexAdapter::new()), AppType::Gemini => Box::new(GeminiAdapter::new()), AppType::GrokBuild => Box::new(CodexAdapter::new()), - AppType::OpenCode | AppType::OpenClaw | AppType::Hermes => Box::new(CodexAdapter::new()), + AppType::OpenCode | AppType::OpenClaw | AppType::Hermes | AppType::Pi => { + // Pi requests use the dedicated per-model adapter path. Keep the + // generic fallback deterministic for non-routing utilities. + Box::new(CodexAdapter::new()) + } } } diff --git a/src-tauri/src/proxy/response_processor.rs b/src-tauri/src/proxy/response_processor.rs index 5402a43f7..bf40d5495 100644 --- a/src-tauri/src/proxy/response_processor.rs +++ b/src-tauri/src/proxy/response_processor.rs @@ -254,11 +254,11 @@ pub async fn handle_non_streaming( spawn_log_usage( state, ctx, + parser_config.input_token_semantics, usage, &model, &ctx.request_model, status.as_u16(), - false, ); } else { let model = json_value @@ -271,11 +271,11 @@ pub async fn handle_non_streaming( spawn_log_usage( state, ctx, + parser_config.input_token_semantics, TokenUsage::default(), &model, &ctx.request_model, status.as_u16(), - false, ); log::debug!( "[{}] 未能解析 usage 信息,跳过记录", @@ -291,11 +291,11 @@ pub async fn handle_non_streaming( spawn_log_usage( state, ctx, + parser_config.input_token_semantics, TokenUsage::default(), ctx.outbound_model.as_deref().unwrap_or(&ctx.request_model), &ctx.request_model, status.as_u16(), - false, ); } } else { @@ -488,6 +488,7 @@ pub(crate) fn create_usage_collector( let start_time = ctx.start_time; let stream_parser = parser_config.stream_parser; let model_extractor = parser_config.model_extractor; + let input_token_semantics = parser_config.input_token_semantics; let session_id = ctx.session_id.clone(); Some(SseUsageCollector::new( @@ -512,6 +513,7 @@ pub(crate) fn create_usage_collector( &model, &request_model, &outbound_model, + input_token_semantics, usage, latency_ms, first_token_ms, @@ -538,6 +540,7 @@ pub(crate) fn create_usage_collector( &model, &request_model, &outbound_model, + input_token_semantics, TokenUsage::default(), latency_ms, first_token_ms, @@ -557,11 +560,11 @@ pub(crate) fn create_usage_collector( fn spawn_log_usage( state: &ProxyState, ctx: &RequestContext, + input_token_semantics: super::usage::InputTokenSemantics, usage: TokenUsage, model: &str, request_model: &str, status_code: u16, - is_streaming: bool, ) { // Check enable_logging before spawning the log task if let Ok(config) = state.config.try_read() { @@ -591,10 +594,11 @@ fn spawn_log_usage( &model, &request_model, &outbound_model, + input_token_semantics, usage, latency_ms, None, - is_streaming, + false, status_code, Some(session_id), ) @@ -624,6 +628,7 @@ async fn log_usage_internal( model: &str, request_model: &str, outbound_model: &str, + input_token_semantics: super::usage::InputTokenSemantics, usage: TokenUsage, latency_ms: u64, first_token_ms: Option, @@ -661,6 +666,7 @@ async fn log_usage_internal( model.to_string(), request_model.to_string(), pricing_model.to_string(), + input_token_semantics, usage, multiplier, latency_ms, @@ -1001,6 +1007,8 @@ mod tests { codex_chat_history: Arc::new(CodexChatHistoryStore::default()), app_handle: None, failover_manager: Arc::new(FailoverSwitchManager::new(db)), + pi_runtime: Arc::new(crate::proxy::pi_runtime::PiRuntimeStore::default()), + pi_server_generation: 0, } } @@ -1072,6 +1080,7 @@ mod tests { "resp-model", "req-model", "req-model", + crate::proxy::usage::InputTokenSemantics::FreshExcludesCache, usage, 10, None, @@ -1142,6 +1151,7 @@ mod tests { "resp-model", "req-model", "outbound-model", + crate::proxy::usage::InputTokenSemantics::FreshExcludesCache, usage, 10, None, @@ -1222,6 +1232,7 @@ mod tests { "resp-model", "req-model", "req-model", + crate::proxy::usage::InputTokenSemantics::FreshExcludesCache, usage, 10, None, diff --git a/src-tauri/src/proxy/server.rs b/src-tauri/src/proxy/server.rs index 1d9bdc37c..2cc8baa71 100644 --- a/src-tauri/src/proxy/server.rs +++ b/src-tauri/src/proxy/server.rs @@ -12,6 +12,7 @@ use super::{ failover_switch::FailoverSwitchManager, handlers, log_codes::srv as log_srv, + pi_runtime::PiRuntimeStore, provider_router::ProviderRouter, providers::{codex_chat_history::CodexChatHistoryStore, gemini_shadow::GeminiShadowStore}, types::*, @@ -48,6 +49,11 @@ pub struct ProxyState { pub app_handle: Option, /// 故障转移切换管理器 pub failover_manager: Arc, + /// Immutable Pi catalog publication point shared with `ProxyService`. + pub pi_runtime: Arc, + /// Listener instance identity. A runtime built for an older listener can + /// never admit requests through this state. + pub pi_server_generation: u64, } /// 代理HTTP服务器 @@ -57,6 +63,7 @@ pub struct ProxyServer { shutdown_tx: Arc>>>, /// 服务器任务句柄,用于等待服务器实际关闭 server_handle: Arc>>>, + pi_server_generation: u64, } impl ProxyServer { @@ -64,6 +71,8 @@ impl ProxyServer { config: ProxyConfig, db: Arc, app_handle: Option, + pi_runtime: Arc, + pi_server_generation: u64, ) -> Self { // 创建共享的 ProviderRouter(熔断器状态将跨所有请求保持) let provider_router = Arc::new(ProviderRouter::new(db.clone())); @@ -81,6 +90,8 @@ impl ProxyServer { codex_chat_history: Arc::new(CodexChatHistoryStore::default()), app_handle, failover_manager, + pi_runtime, + pi_server_generation, }; Self { @@ -88,9 +99,14 @@ impl ProxyServer { state, shutdown_tx: Arc::new(RwLock::new(None)), server_handle: Arc::new(RwLock::new(None)), + pi_server_generation, } } + pub(crate) fn pi_server_generation(&self) -> u64 { + self.pi_server_generation + } + pub async fn start(&self) -> Result { // 检查是否已在运行 if self.shutdown_tx.read().await.is_some() { @@ -364,6 +380,12 @@ impl ProxyServer { .route("/gemini/v1beta/*path", any(handlers::handle_gemini)) // Gemini 的 GA 版本也叫 /v1,给原 SDK 留一条出口 .route("/gemini/v1/*path", any(handlers::handle_gemini)) + // Pi native SDK requests retain their family-specific path below + // the opaque provider route token. + .route( + "/pi/:route_token/*path", + any(super::pi_handler::handle_pi_native), + ) // 提高默认请求体大小限制(避免 413 Payload Too Large) .layer(DefaultBodyLimit::max(200 * 1024 * 1024)) .with_state(self.state.clone()) diff --git a/src-tauri/src/proxy/types.rs b/src-tauri/src/proxy/types.rs index 358e68822..e070d6116 100644 --- a/src-tauri/src/proxy/types.rs +++ b/src-tauri/src/proxy/types.rs @@ -116,6 +116,17 @@ pub struct ProxyTakeoverStatus { pub grokbuild: bool, pub opencode: bool, pub openclaw: bool, + pub pi: bool, + pub pi_operational_state: PiTakeoverOperationalState, +} + +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum PiTakeoverOperationalState { + #[default] + Disabled, + Active, + Degraded, } /// Provider健康状态 diff --git a/src-tauri/src/proxy/usage/calculator.rs b/src-tauri/src/proxy/usage/calculator.rs index 0defaeda1..6615d0030 100644 --- a/src-tauri/src/proxy/usage/calculator.rs +++ b/src-tauri/src/proxy/usage/calculator.rs @@ -3,6 +3,7 @@ //! 使用高精度 Decimal 类型避免浮点数精度问题 use super::parser::TokenUsage; +use super::semantics::InputTokenSemantics; use rust_decimal::Decimal; use std::str::FromStr; @@ -46,13 +47,17 @@ impl CostCalculator { pricing: &ModelPricing, cost_multiplier: Decimal, ) -> CostBreakdown { - Self::calculate_with_cache_semantics(usage, pricing, cost_multiplier, false) + Self::calculate_with_input_semantics( + InputTokenSemantics::FreshExcludesCache, + usage, + pricing, + cost_multiplier, + ) } - /// 按 app_type 选择输入 token 语义后计算成本。 - /// - /// Codex/OpenAI Responses 与 Gemini 的输入 token 字段包含 cache read 部分; - /// Claude/Anthropic 的 input_tokens 已经是 fresh input。 + /// Compatibility helper for existing callers. Live request paths use + /// [`Self::calculate_with_input_semantics`] so product app ownership never + /// stands in for the actual response parser/wire family. pub fn calculate_for_app( app_type: &str, usage: &TokenUsage, @@ -61,32 +66,37 @@ impl CostCalculator { ) -> CostBreakdown { let input_includes_cache_read = crate::services::sql_helpers::is_cache_inclusive_app(app_type); - Self::calculate_with_cache_semantics( + Self::calculate_with_input_semantics( + if input_includes_cache_read { + InputTokenSemantics::TotalIncludesCacheBuckets + } else { + InputTokenSemantics::FreshExcludesCache + }, usage, pricing, cost_multiplier, - input_includes_cache_read, ) } - fn calculate_with_cache_semantics( + pub fn calculate_with_input_semantics( + input_semantics: InputTokenSemantics, usage: &TokenUsage, pricing: &ModelPricing, cost_multiplier: Decimal, - input_includes_cache_read: bool, ) -> CostBreakdown { let million = Decimal::from(1_000_000); // OpenAI/Gemini 风格的 input_tokens 包含缓存读取和写入,需要扣除后再按输入价计费; // Claude/Anthropic 风格的 input_tokens 已经是 fresh input,不能再次扣减。 - let billable_input_tokens = if input_includes_cache_read { - usage - .input_tokens - .saturating_sub(usage.cache_read_tokens) - .saturating_sub(usage.cache_creation_tokens) - } else { - usage.input_tokens - }; + let billable_input_tokens = + if input_semantics == InputTokenSemantics::TotalIncludesCacheBuckets { + usage + .input_tokens + .saturating_sub(usage.cache_read_tokens) + .saturating_sub(usage.cache_creation_tokens) + } else { + usage.input_tokens + }; // 各项基础成本(不含倍率) let input_cost = @@ -112,13 +122,15 @@ impl CostCalculator { } } - pub fn try_calculate_for_app( - app_type: &str, + pub fn try_calculate_with_input_semantics( + input_semantics: InputTokenSemantics, usage: &TokenUsage, pricing: Option<&ModelPricing>, cost_multiplier: Decimal, ) -> Option { - pricing.map(|p| Self::calculate_for_app(app_type, usage, p, cost_multiplier)) + pricing.map(|pricing| { + Self::calculate_with_input_semantics(input_semantics, usage, pricing, cost_multiplier) + }) } } diff --git a/src-tauri/src/proxy/usage/logger.rs b/src-tauri/src/proxy/usage/logger.rs index 7b6362ff3..9f5670ffc 100644 --- a/src-tauri/src/proxy/usage/logger.rs +++ b/src-tauri/src/proxy/usage/logger.rs @@ -2,9 +2,9 @@ use super::calculator::{CostBreakdown, CostCalculator, ModelPricing}; use super::parser::TokenUsage; +use super::semantics::InputTokenSemantics; use crate::database::{Database, PRICING_SOURCE_REQUEST, PRICING_SOURCE_RESPONSE}; use crate::error::AppError; -use crate::services::sql_helpers::{INPUT_TOKEN_SEMANTICS_FRESH, INPUT_TOKEN_SEMANTICS_TOTAL}; use crate::services::usage_stats::{find_model_pricing_row, is_placeholder_pricing_model}; use rusqlite::OptionalExtension; use rust_decimal::Decimal; @@ -72,6 +72,9 @@ pub struct RequestLog { /// 用 model/request_model 猜——路由接管下三者可能各不相同。 /// 错误行(未计价)为空字符串。 pub pricing_model: String, + /// Copied from the response parser/wire family at request admission. + /// Product app ownership is intentionally not consulted at write time. + pub input_token_semantics: InputTokenSemantics, pub usage: TokenUsage, pub cost: Option, pub latency_ms: u64, @@ -121,12 +124,7 @@ impl<'a> UsageLogger<'a> { }; let created_at = chrono::Utc::now().timestamp(); - let input_token_semantics = - if crate::services::sql_helpers::is_cache_inclusive_app(log.app_type.as_str()) { - INPUT_TOKEN_SEMANTICS_TOTAL - } else { - INPUT_TOKEN_SEMANTICS_FRESH - }; + let input_token_semantics = log.input_token_semantics.stored_value(); let semantic = UsageSemantic::from_log(log, input_token_semantics); let existing = Self::load_existing_semantic(&conn, &log.request_id)?; @@ -266,6 +264,7 @@ impl<'a> UsageLogger<'a> { status_code: u16, error_message: String, latency_ms: u64, + input_token_semantics: InputTokenSemantics, ) -> Result<(), AppError> { let request_model = model.clone(); let log = RequestLog { @@ -276,6 +275,7 @@ impl<'a> UsageLogger<'a> { request_model, // 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行) pricing_model: String::new(), + input_token_semantics, usage: TokenUsage::default(), cost: None, latency_ms, @@ -307,6 +307,7 @@ impl<'a> UsageLogger<'a> { is_streaming: bool, session_id: Option, provider_type: Option, + input_token_semantics: InputTokenSemantics, ) -> Result<(), AppError> { let request_model = model.clone(); let log = RequestLog { @@ -317,6 +318,7 @@ impl<'a> UsageLogger<'a> { request_model, // 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行) pricing_model: String::new(), + input_token_semantics, usage: TokenUsage::default(), cost: None, latency_ms, @@ -451,6 +453,7 @@ impl<'a> UsageLogger<'a> { model: String, request_model: String, pricing_model: String, + input_token_semantics: InputTokenSemantics, usage: TokenUsage, cost_multiplier: Decimal, latency_ms: u64, @@ -471,8 +474,8 @@ impl<'a> UsageLogger<'a> { log::warn!("[USG-002] 模型定价未找到,成本将记录为 0: {pricing_model}"); } - let cost = CostCalculator::try_calculate_for_app( - &app_type, + let cost = CostCalculator::try_calculate_with_input_semantics( + input_token_semantics, &usage, pricing.as_ref(), cost_multiplier, @@ -485,6 +488,7 @@ impl<'a> UsageLogger<'a> { model, request_model, pricing_model, + input_token_semantics, usage, cost, latency_ms, @@ -513,6 +517,7 @@ mod tests { model: "gpt-5.6".to_string(), request_model: "gpt-5.6".to_string(), pricing_model: "gpt-5.6".to_string(), + input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets, usage: TokenUsage { input_tokens, output_tokens: 5, @@ -566,6 +571,7 @@ mod tests { "test-model".to_string(), "req-model".to_string(), "test-model".to_string(), + InputTokenSemantics::FreshExcludesCache, usage, Decimal::from(1), 100, @@ -751,6 +757,7 @@ mod tests { 500, "Internal Server Error".to_string(), 50, + InputTokenSemantics::FreshExcludesCache, )?; // 验证错误记录已插入 @@ -778,6 +785,7 @@ mod tests { model: "grok-4.5".to_string(), request_model: "grok-4.5".to_string(), pricing_model: String::new(), + input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets, usage: TokenUsage::default(), cost: None, latency_ms: 1, @@ -798,7 +806,10 @@ mod tests { [], |row| row.get(0), )?; - assert_eq!(semantics, INPUT_TOKEN_SEMANTICS_TOTAL); + assert_eq!( + semantics, + InputTokenSemantics::TotalIncludesCacheBuckets.stored_value() + ); Ok(()) } } diff --git a/src-tauri/src/proxy/usage/mod.rs b/src-tauri/src/proxy/usage/mod.rs index 08ef51018..9ea2628f1 100644 --- a/src-tauri/src/proxy/usage/mod.rs +++ b/src-tauri/src/proxy/usage/mod.rs @@ -5,6 +5,7 @@ pub mod calculator; pub mod logger; pub mod parser; +pub mod semantics; // 仅导出内部使用的类型,避免未使用警告 #[allow(unused_imports)] @@ -13,3 +14,5 @@ pub use calculator::{CostBreakdown, CostCalculator, ModelPricing}; pub use logger::{RequestLog, UsageLogger}; #[allow(unused_imports)] pub use parser::TokenUsage; +#[allow(unused_imports)] +pub use semantics::InputTokenSemantics; diff --git a/src-tauri/src/proxy/usage/semantics.rs b/src-tauri/src/proxy/usage/semantics.rs new file mode 100644 index 000000000..8190edfa8 --- /dev/null +++ b/src-tauri/src/proxy/usage/semantics.rs @@ -0,0 +1,32 @@ +//! Input-token semantics carried with every live proxy request. +//! +//! `app_type` is a product/UI ownership dimension. It must not decide whether +//! an upstream's input count already contains cache buckets: Pi can route the +//! same logical app through four different wire families. + +use crate::pi_config::gateway::PiGatewayApiFamily; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(i64)] +pub enum InputTokenSemantics { + /// OpenAI Responses/Completions and Google usage totals include cached + /// input. Fresh input is total minus the reported cache buckets. + TotalIncludesCacheBuckets = 1, + /// Anthropic reports fresh input separately from cache reads/creation. + FreshExcludesCache = 2, +} + +impl InputTokenSemantics { + pub const fn stored_value(self) -> i64 { + self as i64 + } + + pub const fn for_pi_family(family: PiGatewayApiFamily) -> Self { + match family { + PiGatewayApiFamily::AnthropicMessages => Self::FreshExcludesCache, + PiGatewayApiFamily::OpenAiCompletions + | PiGatewayApiFamily::OpenAiResponses + | PiGatewayApiFamily::GoogleGenerativeAi => Self::TotalIncludesCacheBuckets, + } + } +} diff --git a/src-tauri/src/services/config.rs b/src-tauri/src/services/config.rs index 7d465810a..813630297 100644 --- a/src-tauri/src/services/config.rs +++ b/src-tauri/src/services/config.rs @@ -138,6 +138,10 @@ impl ConfigService { AppType::Hermes => { // Hermes uses additive mode, no live sync needed } + AppType::Pi => { + // Pi's shared models/settings documents are owned by the + // catalog coordinator, never by this legacy live-sync path. + } } Ok(()) diff --git a/src-tauri/src/services/mcp.rs b/src-tauri/src/services/mcp.rs index b57591b82..5160225cd 100644 --- a/src-tauri/src/services/mcp.rs +++ b/src-tauri/src/services/mcp.rs @@ -147,6 +147,13 @@ impl McpService { AppType::Hermes => { mcp::sync_single_server_to_hermes(&Default::default(), &server.id, &server.server)?; } + AppType::Pi => { + return Err(AppError::localized( + "mcp.pi.unsupported", + "固定版本的 Pi 核心没有原生 MCP 注册表", + "The pinned Pi core has no native MCP registry", + )); + } } Ok(()) } @@ -183,6 +190,13 @@ impl McpService { AppType::Hermes => { mcp::remove_server_from_hermes(id)?; } + AppType::Pi => { + return Err(AppError::localized( + "mcp.pi.unsupported", + "固定版本的 Pi 核心没有原生 MCP 注册表", + "The pinned Pi core has no native MCP registry", + )); + } } Ok(()) } @@ -227,7 +241,10 @@ impl McpService { servers: &IndexMap, app: &AppType, ) -> Result<(), AppError> { - if matches!(app, AppType::OpenClaw | AppType::ClaudeDesktop) { + if matches!( + app, + AppType::OpenClaw | AppType::ClaudeDesktop | AppType::Pi + ) { return Ok(()); } @@ -544,3 +561,32 @@ impl McpService { } } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::database::Database; + use std::sync::Arc; + + #[test] + fn global_mcp_projection_treats_pi_as_explicitly_not_applicable() { + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + let mut servers = IndexMap::new(); + servers.insert( + "example".to_string(), + McpServer { + id: "example".to_string(), + name: "Example".to_string(), + server: serde_json::json!({"command": "example"}), + apps: Default::default(), + description: None, + homepage: None, + docs: None, + tags: Vec::new(), + }, + ); + + McpService::project_servers_to_app(&state, &servers, &AppType::Pi) + .expect("Pi is intentionally outside the pinned core MCP registry"); + } +} diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 2fce3bec7..fb5a36946 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -8,6 +8,8 @@ pub mod mcp; pub mod model_fetch; pub mod model_pricing; pub mod omo; +pub(crate) mod pi_catalog; +pub mod pi_prompt_files; pub mod profile; pub mod prompt; pub mod provider; @@ -21,6 +23,7 @@ pub mod session_usage_gemini; pub mod session_usage_grokbuild; pub mod session_usage_opencode; pub mod skill; +pub(crate) mod skill_deployment; pub mod speedtest; pub mod sql_helpers; pub mod stream_check; diff --git a/src-tauri/src/services/pi_catalog.rs b/src-tauri/src/services/pi_catalog.rs new file mode 100644 index 000000000..c5d43883d --- /dev/null +++ b/src-tauri/src/services/pi_catalog.rs @@ -0,0 +1,2118 @@ +//! Ordered mutations for Pi's managed provider catalog. +//! +//! SQLite is the managed aggregate authority, `pi_provider_projections` owns +//! exact keys in Pi's shared `models.json`, and Pi's `settings.json` owns the +//! native default. Every public mutation acquires the same Pi switch boundary; +//! callers must not compose the database and native-file primitives directly. + +use crate::app_config::AppType; +use crate::database::{ + NewEndpoint, NewProviderAggregate, PiProviderProjection, ProviderKey, ProviderRowUpdate, +}; +use crate::error::AppError; +use crate::pi_config::document::{ + apply_pi_provider_patch_with_receipt, snapshot_pi_provider_values, PiProviderPatchReceipt, + PiProviderValuesSnapshot, +}; +use crate::pi_config::gateway::parse_pi_gateway_endpoint; +use crate::pi_config::model::{ + effective_pi_model, validate_pi_managed_provider, PiManagedProviderConfig, PiManagementStatus, +}; +use crate::pi_config::native::{ + get_pi_models_path, inspect_pi_native_entry, PiNativeInspectionService, +}; +use crate::pi_config::native_settings::{ + read_pi_native_defaults, set_pi_native_default_with_receipt, PiNativeDefaults, + PiNativeDefaultsReceipt, PiNativeDefaultsRollback, +}; +use crate::provider::{ProviderAggregate, ProviderMutationInput}; +use crate::settings; +use crate::store::AppState; +use indexmap::IndexMap; +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::collections::BTreeMap; + +const PI_APP: &str = "pi"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum PiCatalogAuthority { + Published, + PreviousRestored, + MutatedDatabaseAuthoritative, + ProjectionPending, +} + +impl PiCatalogAuthority { + fn as_str(self) -> &'static str { + match self { + Self::Published => "published", + Self::PreviousRestored => "previous_restored", + Self::MutatedDatabaseAuthoritative => "mutated_database_authoritative", + Self::ProjectionPending => "projection_pending", + } + } +} + +#[derive(Debug)] +pub(crate) enum PiCatalogMutation { + CreateProvider { + input: ProviderMutationInput, + provider_key: String, + activate_if_first: bool, + }, + UpdateProvider { + input: ProviderMutationInput, + }, + DeleteProvider { + provider_id: String, + }, + AddEndpoint { + provider_id: String, + url: String, + }, + RemoveEndpoint { + provider_id: String, + url: String, + }, + ImportNative { + provider_key: String, + expected_fingerprint: String, + }, + SetDefault { + provider_id: String, + model_id: String, + }, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiCatalogMutationResult { + pub authority: PiCatalogAuthority, + #[serde(skip_serializing_if = "Option::is_none")] + pub provider_id: Option, + #[serde(skip)] + native_defaults_receipt: Option, + #[serde(skip)] + native_patch_receipt: Option, + #[serde(skip)] + native_fingerprint_preconditions: IndexMap, +} + +pub(crate) struct PiCatalogCoordinator; + +struct PiCatalogSnapshot { + aggregates: IndexMap, + projections: Vec, + db_current: Option, + // Capture validates the complete shared projection before any DB write. + // Exact compensation itself uses per-operation before/attempted receipts. + native: PiProviderValuesSnapshot, +} + +impl PiCatalogCoordinator { + pub(crate) fn update_route_order( + state: &AppState, + updates: Vec<(ProviderKey, usize)>, + ) -> Result { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + let aggregates = state.db.get_all_provider_aggregates(PI_APP)?; + for (key, _) in &updates { + if !aggregates.contains_key(key.id()) { + return Err(AppError::NotFound(format!( + "Pi provider '{}' cannot be sorted because it does not exist", + key.id() + ))); + } + } + let projections = state + .db + .get_pi_projection_manifest()? + .into_values() + .collect::>(); + let db_current = state.db.get_current_provider(PI_APP)?; + + state.db.update_provider_sort_index(&updates)?; + if let Err(error) = + futures::executor::block_on(state.proxy_service.publish_pi_runtime_order()) + { + let rollback = state.db.restore_pi_catalog_snapshot( + &aggregates, + &projections, + db_current.as_deref(), + ); + return Err(AppError::Config(format!( + "failed to publish sorted Pi runtime: {error}; DB rollback={}", + rollback + .err() + .map_or_else(|| "ok".to_string(), |value| value.to_string()) + ))); + } + Ok(true) + } + + pub(crate) fn apply( + state: &AppState, + mutation: PiCatalogMutation, + ) -> Result { + let additional_native_key = match &mutation { + PiCatalogMutation::CreateProvider { provider_key, .. } + | PiCatalogMutation::ImportNative { provider_key, .. } => Some(provider_key.clone()), + _ => None, + }; + Self::run_with_runtime_reconcile(state, additional_native_key.as_deref(), || match mutation + { + PiCatalogMutation::CreateProvider { + input, + provider_key, + activate_if_first, + } => Self::create(state, input, provider_key, activate_if_first), + PiCatalogMutation::UpdateProvider { input } => Self::update(state, input), + PiCatalogMutation::DeleteProvider { provider_id } => Self::delete(state, &provider_id), + PiCatalogMutation::AddEndpoint { provider_id, url } => { + Self::add_endpoint(state, &provider_id, &url) + } + PiCatalogMutation::RemoveEndpoint { provider_id, url } => { + Self::remove_endpoint(state, &provider_id, &url) + } + PiCatalogMutation::ImportNative { + provider_key, + expected_fingerprint, + } => Self::import_native(state, &provider_key, &expected_fingerprint), + PiCatalogMutation::SetDefault { + provider_id, + model_id, + } => Self::set_default(state, &provider_id, &model_id), + }) + } + + /// Reconcile portable provider rows with this device's exact-key ledger. + /// + /// SQL/WebDAV/S3 exports intentionally omit device-local projection rows. + /// Imported Pi providers use their immutable provider id as their native + /// key. A missing native key can therefore be claimed and published + /// without guessing; an existing unclaimed key is never overwritten. + pub(crate) fn reconcile_portable_import(state: &AppState) -> Result<(), AppError> { + Self::run_with_runtime_reconcile(state, None, || { + let models_path = get_pi_models_path()?; + let (native_defaults_receipt, native_patch_receipt) = + Self::reconcile_portable_catalog_at( + state, + &models_path, + |provider_id, provider_key, config| { + state.proxy_service.project_pi_provider_value( + provider_id, + provider_key, + config, + ) + }, + )?; + Ok(success(None) + .with_native_defaults_receipt(native_defaults_receipt) + .with_native_patch_receipt(native_patch_receipt)) + }) + .map(|_| ()) + } + + fn run_with_runtime_reconcile( + state: &AppState, + additional_native_key: Option<&str>, + operation: impl FnOnce() -> Result, + ) -> Result { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + Self::reconcile_current_indexes_from_native(state)?; + let snapshot = PiCatalogSnapshot::capture(state, additional_native_key)?; + let catalog_epoch = + futures::executor::block_on(state.proxy_service.begin_pi_catalog_mutation()); + let result = operation(); + let native_defaults_receipt = result + .as_ref() + .ok() + .and_then(|result| result.native_defaults_receipt.as_ref()) + .cloned(); + let native_patch_receipt = result + .as_ref() + .ok() + .and_then(|result| result.native_patch_receipt.as_ref()) + .cloned(); + let native_fingerprint_preconditions = result + .as_ref() + .ok() + .map(|result| result.native_fingerprint_preconditions.clone()) + .unwrap_or_default(); + let mut expected_native = snapshot.native.clone(); + if let Some(receipt) = native_patch_receipt.as_ref() { + for (provider_key, attempted) in receipt.attempted_values() { + expected_native + .values + .insert(provider_key.clone(), attempted.clone()); + } + } + let reconcile = futures::executor::block_on( + state + .proxy_service + .reconcile_pi_runtime_at_epoch_with_native_claim_precondition( + catalog_epoch, + Some(&expected_native), + (!native_fingerprint_preconditions.is_empty()) + .then_some(&native_fingerprint_preconditions), + ), + ); + match (result, reconcile) { + (Ok(result), Ok(_)) => Ok(result), + (Ok(_), Err(error)) => { + if let Err(rollback_error) = snapshot.restore( + state, + native_defaults_receipt.as_ref(), + native_patch_receipt.as_ref(), + ) { + let _ = futures::executor::block_on( + state.proxy_service.close_pi_runtime_at_epoch(catalog_epoch), + ); + return Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!( + "Pi catalog runtime publication failed ({error}); snapshot rollback failed ({rollback_error})" + ), + )); + } + match futures::executor::block_on( + state + .proxy_service + .reconcile_pi_runtime_at_epoch_with_native_precondition( + catalog_epoch, + Some(&snapshot.native), + ), + ) { + Ok(_) => Err(authority_error( + PiCatalogAuthority::PreviousRestored, + format!( + "Pi catalog runtime publication failed and the previous catalog was restored: {error}" + ), + )), + Err(rollback_error) => { + let _ = futures::executor::block_on( + state.proxy_service.close_pi_runtime_at_epoch(catalog_epoch), + ); + Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!( + "Pi catalog runtime publication failed ({error}); the database snapshot was restored but runtime recovery failed ({rollback_error})" + ), + )) + } + } + } + (Err(error), Ok(_)) => Err(error), + (Err(error), Err(reconcile_error)) => { + let _ = futures::executor::block_on( + state.proxy_service.close_pi_runtime_at_epoch(catalog_epoch), + ); + Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!( + "{error}; additionally failed to reconcile Pi admission: {reconcile_error}" + ), + )) + } + } + } + + fn reconcile_portable_catalog_at( + state: &AppState, + models_path: &std::path::Path, + mut project: impl FnMut(&str, &str, &PiManagedProviderConfig) -> Result, + ) -> Result< + ( + Option, + Option, + ), + AppError, + > { + let providers = state.db.get_all_providers(PI_APP)?; + let manifest = state.db.get_pi_projection_manifest()?; + let claimed_keys = manifest + .values() + .map(|projection| { + ( + projection.provider_key.clone(), + projection.provider_id.clone(), + ) + }) + .collect::>(); + + struct PlannedProjection { + provider_id: String, + provider_key: String, + config: PiManagedProviderConfig, + projected: Value, + needs_claim: bool, + } + + let mut plans = Vec::with_capacity(providers.len()); + let mut planned_key_owners = BTreeMap::::new(); + for provider in providers.values() { + let config: PiManagedProviderConfig = + serde_json::from_value(provider.settings_config.clone()).map_err(|error| { + AppError::InvalidInput(format!( + "imported Pi provider '{}' cannot be decoded: {error}", + provider.id + )) + })?; + validate_pi_managed_provider(&config).map_err(|error| { + AppError::InvalidInput(format!( + "imported Pi provider '{}' is invalid: {error}", + provider.id + )) + })?; + let (provider_key, needs_claim) = match manifest.get(&provider.id) { + Some(projection) => (projection.provider_key.clone(), false), + None => (non_empty_native_key(&provider.id)?.to_string(), true), + }; + if let Some(owner) = + planned_key_owners.insert(provider_key.clone(), provider.id.clone()) + { + if owner != provider.id { + return Err(AppError::Conflict(format!( + "imported Pi providers '{owner}' and '{}' normalize to the same native key '{provider_key}'", + provider.id + ))); + } + } + if let Some(owner) = claimed_keys.get(&provider_key) { + if owner != &provider.id { + return Err(AppError::Conflict(format!( + "cannot project imported Pi provider '{}': native key '{}' is owned by '{}'", + provider.id, provider_key, owner + ))); + } + } + let projected = project(&provider.id, &provider_key, &config)?; + plans.push(PlannedProjection { + provider_id: provider.id.clone(), + provider_key, + config, + projected, + needs_claim, + }); + } + + let before_file = snapshot_pi_provider_values( + models_path, + plans.iter().map(|plan| plan.provider_key.clone()), + )?; + for plan in plans.iter().filter(|plan| plan.needs_claim) { + if before_file + .values + .get(&plan.provider_key) + .and_then(Option::as_ref) + .is_some() + { + return Err(AppError::Conflict(format!( + "cannot project imported Pi provider '{}': unclaimed native key '{}' already exists", + plan.provider_id, plan.provider_key + ))); + } + } + // Native defaults are another authority input to the reconciliation + // plan. Read and validate them before changing either the exact-key + // ownership ledger or models.json, so an unreadable settings.json is a + // fail-before-write error rather than a partially compensated publish. + let previous_defaults = read_pi_native_defaults()?; + + // Portable restore replaces provider rows but deliberately preserves + // the device-local exact-key ledger. Claims whose providers disappeared + // must be released before publishing a runtime, otherwise the manifest + // and provider aggregate sets can never converge. Their native values + // remain untouched and become user-owned; without a retained provider + // aggregate there is no safe expected value with which to delete them. + let orphaned_claims = manifest + .values() + .filter(|projection| !providers.contains_key(&projection.provider_id)) + .map(|projection| { + ( + projection.provider_id.clone(), + projection.provider_key.clone(), + ) + }) + .collect::>(); + let mut released_claims = Vec::new(); + for (provider_id, provider_key) in &orphaned_claims { + match state.db.delete_pi_projection_key(provider_id, provider_key) { + Ok(true) => released_claims.push((provider_id.clone(), provider_key.clone())), + Ok(false) => { + return Err(compensate_portable_projection_ledger( + state, + &[], + &released_claims, + AppError::Conflict(format!( + "Pi projection claim '{provider_id}' disappeared during import reconciliation" + )), + )); + } + Err(error) => { + return Err(compensate_portable_projection_ledger( + state, + &[], + &released_claims, + error, + )); + } + } + } + + let mut newly_claimed = Vec::new(); + for plan in plans.iter().filter(|plan| plan.needs_claim) { + match state + .db + .claim_pi_projection_key(&plan.provider_id, &plan.provider_key) + { + Ok(_) => newly_claimed.push((plan.provider_id.clone(), plan.provider_key.clone())), + Err(error) => { + return Err(compensate_portable_projection_ledger( + state, + &newly_claimed, + &released_claims, + error, + )); + } + } + } + + let patch = plans + .iter() + .map(|plan| (plan.provider_key.clone(), Some(plan.projected.clone()))) + .collect::>(); + let native_patch_receipt = if patch.is_empty() { + None + } else { + match apply_pi_provider_patch_with_receipt(models_path, &before_file, &patch) { + Ok(receipt) => Some(receipt), + Err(error) => { + return Err(compensate_portable_projection_ledger( + state, + &newly_claimed, + &released_claims, + error, + )); + } + } + }; + + let native_selection = match ( + previous_defaults.default_provider.as_deref(), + previous_defaults.default_model.as_deref(), + ) { + (Some(provider_key), Some(model_id)) => plans + .iter() + .find(|plan| { + plan.provider_key == provider_key + && effective_pi_model(&plan.config, model_id).is_ok() + }) + .map(|plan| (plan.provider_id.clone(), model_id.to_string())), + _ => None, + }; + // Only use a DB/local fallback when Pi has no native selection at all. + // A partial, invalid, or unowned native default is still real state and + // must be surfaced rather than silently overwritten during import. + let current_selection = if native_selection.is_some() { + native_selection + } else if previous_defaults.default_provider.is_none() + && previous_defaults.default_model.is_none() + { + crate::settings::get_effective_current_provider(&state.db, &AppType::Pi)? + .and_then(|provider_id| plans.iter().find(|plan| plan.provider_id == provider_id)) + .map(|plan| { + ( + plan.provider_id.clone(), + plan.config + .models + .first() + .expect("validated Pi provider has at least one model") + .id + .clone(), + ) + }) + } else { + None + }; + + let mut native_defaults_receipt = None; + if let Some((current_provider, model_id)) = current_selection { + match Self::set_default(state, ¤t_provider, &model_id) { + Ok(result) => native_defaults_receipt = result.native_defaults_receipt, + Err(error) => { + let file_restored = native_patch_receipt + .as_ref() + .map_or(Ok(()), PiProviderPatchReceipt::rollback); + let claims_restored = compensate_portable_projection_ledger( + state, + &newly_claimed, + &released_claims, + error, + ); + if file_restored.is_err() { + return Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!( + "{claims_restored}; imported Pi default compensation was incomplete" + ), + )); + } + return Err(claims_restored); + } + } + } + Ok((native_defaults_receipt, native_patch_receipt)) + } + + pub(crate) fn inspect_native( + state: &AppState, + ) -> Result, AppError> { + PiNativeInspectionService::inspect_current(&managed_claims(state)?) + } + + /// Resolve the provider that Pi itself will start with. + /// + /// `settings.json` is the live authority. The device-local and SQLite + /// current markers are compensation/indexing aids and must never make the + /// UI claim that a different provider is active after an external Pi edit. + pub(crate) fn current_native_provider(state: &AppState) -> Result, AppError> { + let defaults = read_pi_native_defaults()?; + resolve_native_current_provider(state, &defaults) + } + + fn reconcile_current_indexes_from_native(state: &AppState) -> Result<(), AppError> { + let defaults = read_pi_native_defaults()?; + let native_current = resolve_native_current_provider(state, &defaults)?; + let previous_local = settings::get_current_provider(&AppType::Pi); + let previous_db = state.db.get_current_provider(PI_APP)?; + if previous_local == native_current && previous_db == native_current { + return Ok(()); + } + + settings::set_current_provider(&AppType::Pi, native_current.as_deref())?; + if let Err(error) = restore_db_current(state, native_current.as_deref()) { + let local_restored = + settings::set_current_provider(&AppType::Pi, previous_local.as_deref()).is_ok(); + let db_restored = restore_db_current(state, previous_db.as_deref()).is_ok(); + return Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!( + "failed to align Pi current indexes with native settings: {error}; rollback: local={local_restored}, db={db_restored}" + ), + )); + } + Ok(()) + } + + fn create( + state: &AppState, + input: ProviderMutationInput, + provider_key: String, + activate_if_first: bool, + ) -> Result { + let catalog_was_empty = state.db.get_all_providers(PI_APP)?.is_empty(); + let provider_key = non_empty_native_key(&provider_key)?; + let config = managed_config(&input)?; + validate_pi_initial_endpoints(&input)?; + let provider_id = input.id.clone(); + if state.db.get_pi_projection_for_key(provider_key)?.is_some() { + return Err(AppError::Conflict(format!( + "Pi native provider key '{provider_key}' is already managed" + ))); + } + + let models_path = get_pi_models_path()?; + let before_file = snapshot_pi_provider_values(&models_path, [provider_key.to_string()])?; + if before_file + .values + .get(provider_key) + .and_then(Option::as_ref) + .is_some() + { + return Err(AppError::Conflict(format!( + "Pi native provider key '{provider_key}' already exists; import it instead" + ))); + } + + let previous_defaults = read_pi_native_defaults()?; + let previous_local = settings::get_current_provider(&AppType::Pi); + let previous_db = state.db.get_current_provider(PI_APP)?; + let projected = + state + .proxy_service + .project_pi_provider_value(&provider_id, provider_key, &config)?; + state.db.create_pi_catalog_provider( + NewProviderAggregate::from_input(PI_APP, input)?, + provider_key, + )?; + + let projection = IndexMap::from([(provider_key.to_string(), Some(projected))]); + let native_patch_receipt = + match apply_pi_provider_patch_with_receipt(&models_path, &before_file, &projection) { + Ok(receipt) => receipt, + Err(projection_error) => { + return Err(compensate_created_provider( + state, + &provider_id, + None, + projection_error, + )); + } + }; + + let no_selected_provider = previous_local.is_none() && previous_db.is_none(); + let native_defaults_empty = previous_defaults.default_provider.is_none() + && previous_defaults.default_model.is_none(); + let mut result = success(Some(provider_id.clone())); + if activate_if_first && catalog_was_empty && no_selected_provider && native_defaults_empty { + let first_model = config + .models + .first() + .expect("validated Pi config has at least one model") + .id + .clone(); + match Self::set_default(state, &provider_id, &first_model) { + Ok(activation) => { + result.native_defaults_receipt = activation.native_defaults_receipt; + } + Err(error) => { + return Err(compensate_created_provider( + state, + &provider_id, + Some(&native_patch_receipt), + error, + )); + } + } + } + + Ok(result.with_native_patch_receipt(Some(native_patch_receipt))) + } + + fn update( + state: &AppState, + input: ProviderMutationInput, + ) -> Result { + let config = managed_config(&input)?; + let provider_id = input.id.clone(); + let projection = state.db.get_pi_projection(&provider_id)?.ok_or_else(|| { + AppError::Conflict(format!( + "Pi provider '{provider_id}' has no exact-key ownership claim" + )) + })?; + let previous = state + .db + .get_provider_aggregate(PI_APP, &provider_id)? + .ok_or_else(|| AppError::NotFound(format!("Pi provider '{provider_id}'")))?; + let previous_defaults = read_pi_native_defaults()?; + let previous_db = state.db.get_current_provider(PI_APP)?; + let was_current = previous_db.as_deref() == Some(&provider_id); + let models_path = get_pi_models_path()?; + let before_file = + snapshot_pi_provider_values(&models_path, [projection.provider_key.clone()])?; + let projected = state.proxy_service.project_pi_provider_value( + &provider_id, + &projection.provider_key, + &config, + )?; + + let key = ProviderKey::new(PI_APP, provider_id.clone())?; + state + .db + .update_pi_catalog_provider(&key, &ProviderRowUpdate::from_input(&input)?)?; + let patch = IndexMap::from([(projection.provider_key.clone(), Some(projected))]); + let native_patch_receipt = + match apply_pi_provider_patch_with_receipt(&models_path, &before_file, &patch) { + Ok(receipt) => receipt, + Err(error) => { + return Err(compensate_existing_provider( + state, + &previous, + was_current, + &projection, + None, + error, + )); + } + }; + let mut result = success(Some(provider_id.clone())); + if previous_defaults.default_provider.as_deref() == Some(&projection.provider_key) { + if let Some(default_model) = previous_defaults.default_model.as_deref() { + if effective_pi_model(&config, default_model).is_err() { + let replacement = config + .models + .first() + .expect("validated Pi config has at least one model") + .id + .clone(); + match Self::set_default(state, &provider_id, &replacement) { + Ok(default_update) => { + result.native_defaults_receipt = default_update.native_defaults_receipt; + } + Err(error) => { + return Err(compensate_existing_provider( + state, + &previous, + was_current, + &projection, + Some(&native_patch_receipt), + error, + )); + } + } + } + } + } + Ok(result.with_native_patch_receipt(Some(native_patch_receipt))) + } + + fn delete(state: &AppState, provider_id: &str) -> Result { + let projection = state + .db + .get_pi_projection(provider_id)? + .ok_or_else(|| AppError::NotFound(format!("Pi provider '{provider_id}'")))?; + let previous = state + .db + .get_provider_aggregate(PI_APP, provider_id)? + .ok_or_else(|| AppError::NotFound(format!("Pi provider '{provider_id}'")))?; + let db_current = state.db.get_current_provider(PI_APP)?; + let native_defaults = read_pi_native_defaults()?; + if native_defaults.default_provider.as_deref() == Some(&projection.provider_key) { + return Err(AppError::Conflict( + "the active Pi provider cannot be deleted".to_string(), + )); + } + + let models_path = get_pi_models_path()?; + let before_file = + snapshot_pi_provider_values(&models_path, [projection.provider_key.clone()])?; + let was_current = db_current.as_deref() == Some(provider_id); + state.db.delete_pi_catalog_provider(provider_id)?; + let patch = IndexMap::from([(projection.provider_key.clone(), None)]); + let native_patch_receipt = + match apply_pi_provider_patch_with_receipt(&models_path, &before_file, &patch) { + Ok(receipt) => receipt, + Err(error) => { + return Err(compensate_existing_provider( + state, + &previous, + was_current, + &projection, + None, + error, + )); + } + }; + Ok(success(Some(provider_id.to_string())) + .with_native_patch_receipt(Some(native_patch_receipt))) + } + + fn add_endpoint( + state: &AppState, + provider_id: &str, + url: &str, + ) -> Result { + ensure_managed_provider(state, provider_id)?; + let normalized = normalize_gateway_endpoint_for_write(url)?; + state.db.add_provider_endpoint( + &ProviderKey::new(PI_APP, provider_id)?, + NewEndpoint::now(normalized)?, + )?; + Ok(success(Some(provider_id.to_string()))) + } + + fn remove_endpoint( + state: &AppState, + provider_id: &str, + url: &str, + ) -> Result { + ensure_managed_provider(state, provider_id)?; + let normalized = normalize_endpoint_key(url)?; + state + .db + .remove_provider_endpoint(&ProviderKey::new(PI_APP, provider_id)?, &normalized)?; + Ok(success(Some(provider_id.to_string()))) + } + + fn import_native( + state: &AppState, + provider_key: &str, + expected_fingerprint: &str, + ) -> Result { + let provider_key = non_empty_native_key(provider_key)?; + let claims = managed_claims(state)?; + let inspection = inspect_pi_native_entry(&get_pi_models_path()?, provider_key, &claims)? + .ok_or_else(|| AppError::NotFound(format!("Pi native provider '{provider_key}'")))?; + if inspection.diagnostic.fingerprint != expected_fingerprint { + return Err(AppError::Conflict(format!( + "Pi native provider '{provider_key}' changed since inspection" + ))); + } + if inspection.diagnostic.management_status != PiManagementStatus::Importable { + return Err(AppError::Conflict(format!( + "Pi native provider '{provider_key}' is not importable" + ))); + } + let config = inspection.managed_config.ok_or_else(|| { + AppError::Conflict(format!( + "Pi native provider '{provider_key}' is not representable by the managed catalog" + )) + })?; + validate_pi_managed_provider(&config) + .map_err(|error| AppError::InvalidInput(error.to_string()))?; + let display_name = config + .name + .clone() + .unwrap_or_else(|| provider_key.to_string()); + let input = ProviderMutationInput { + id: provider_key.to_string(), + name: display_name, + settings_config: serde_json::to_value(config).map_err(|source| { + AppError::Config(format!( + "failed to serialize imported Pi provider: {source}" + )) + })?, + website_url: None, + category: Some("imported".to_string()), + created_at: Some(chrono::Utc::now().timestamp_millis()), + sort_index: state + .db + .get_all_providers(PI_APP)? + .values() + .filter_map(|provider| provider.sort_index) + .max() + .map_or(Some(0), |index| index.checked_add(1)), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }; + state.db.create_pi_catalog_provider( + NewProviderAggregate::from_input(PI_APP, input)?, + provider_key, + )?; + Ok( + success(Some(provider_key.to_string())).with_native_fingerprint_precondition( + provider_key.to_string(), + inspection.diagnostic.fingerprint, + ), + ) + } + + fn set_default( + state: &AppState, + provider_id: &str, + model_id: &str, + ) -> Result { + let aggregate = ensure_managed_provider(state, provider_id)?; + let config: PiManagedProviderConfig = serde_json::from_value( + aggregate.provider.settings_config.clone(), + ) + .map_err(|error| { + AppError::Config(format!( + "managed Pi provider '{provider_id}' is invalid: {error}" + )) + })?; + effective_pi_model(&config, model_id) + .map_err(|error| AppError::InvalidInput(error.to_string()))?; + let projection = state.db.get_pi_projection(provider_id)?.ok_or_else(|| { + AppError::Conflict(format!( + "Pi provider '{provider_id}' has no exact-key ownership claim" + )) + })?; + + let previous_local = settings::get_current_provider(&AppType::Pi); + let previous_db = state.db.get_current_provider(PI_APP)?; + let native_defaults_receipt = + set_pi_native_default_with_receipt(&projection.provider_key, model_id)?; + if let Err(error) = settings::set_current_provider(&AppType::Pi, Some(provider_id)) { + return match native_defaults_receipt.rollback() { + Ok(PiNativeDefaultsRollback::Restored) => Err(authority_error( + PiCatalogAuthority::PreviousRestored, + error.to_string(), + )), + Ok(PiNativeDefaultsRollback::Superseded) => { + let indexes = Self::reconcile_current_indexes_from_native(state).map_or_else( + |index_error| format!("; current-index reconcile failed: {index_error}"), + |()| "; external native defaults were preserved".to_string(), + ); + Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!("{error}{indexes}"), + )) + } + Err(rollback_error) => Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!("{error}; native-default rollback failed: {rollback_error}"), + )), + }; + } + if let Err(error) = state.db.set_current_provider(PI_APP, provider_id) { + return match native_defaults_receipt.rollback() { + Ok(PiNativeDefaultsRollback::Restored) => { + let local_restored = + settings::set_current_provider(&AppType::Pi, previous_local.as_deref()); + let db_restored = restore_db_current(state, previous_db.as_deref()); + let authority = if local_restored.is_ok() && db_restored.is_ok() { + PiCatalogAuthority::PreviousRestored + } else { + PiCatalogAuthority::ProjectionPending + }; + Err(authority_error(authority, error.to_string())) + } + Ok(PiNativeDefaultsRollback::Superseded) => { + let indexes = Self::reconcile_current_indexes_from_native(state).map_or_else( + |index_error| format!("; current-index reconcile failed: {index_error}"), + |()| "; external native defaults were preserved".to_string(), + ); + Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!("{error}{indexes}"), + )) + } + Err(rollback_error) => { + let indexes = Self::reconcile_current_indexes_from_native(state).map_or_else( + |index_error| format!("; current-index reconcile failed: {index_error}"), + |()| String::new(), + ); + Err(authority_error( + PiCatalogAuthority::ProjectionPending, + format!( + "{error}; native-default rollback failed: {rollback_error}{indexes}" + ), + )) + } + }; + } + Ok(success(Some(provider_id.to_string())) + .with_native_defaults_receipt(Some(native_defaults_receipt))) + } +} + +fn managed_config(input: &ProviderMutationInput) -> Result { + let config: PiManagedProviderConfig = serde_json::from_value(input.settings_config.clone()) + .map_err(|error| { + AppError::InvalidInput(format!("invalid managed Pi provider config: {error}")) + })?; + validate_pi_managed_provider(&config) + .map_err(|error| AppError::InvalidInput(error.to_string()))?; + Ok(config) +} + +fn ensure_managed_provider( + state: &AppState, + provider_id: &str, +) -> Result { + if state.db.get_pi_projection(provider_id)?.is_none() { + return Err(AppError::Conflict(format!( + "Pi provider '{provider_id}' has no exact-key ownership claim" + ))); + } + state + .db + .get_provider_aggregate(PI_APP, provider_id)? + .ok_or_else(|| AppError::NotFound(format!("Pi provider '{provider_id}'"))) +} + +fn managed_claims(state: &AppState) -> Result, AppError> { + Ok(state + .db + .get_pi_projection_manifest()? + .into_values() + .map(|projection| (projection.provider_key, projection.provider_id)) + .collect()) +} + +fn non_empty_native_key(value: &str) -> Result<&str, AppError> { + let value = value.trim(); + if value.is_empty() { + Err(AppError::InvalidInput( + "Pi native provider key cannot be empty".to_string(), + )) + } else { + Ok(value) + } +} + +fn normalize_endpoint_key(value: &str) -> Result { + let normalized = value.trim().trim_end_matches('/').to_string(); + if normalized.is_empty() { + return Err(AppError::InvalidInput( + "Pi endpoint URL cannot be empty".to_string(), + )); + } + Ok(normalized) +} + +fn normalize_gateway_endpoint_for_write(value: &str) -> Result { + let normalized = normalize_endpoint_key(value)?; + parse_pi_gateway_endpoint(&normalized, "/customEndpoints").map_err(|_| { + AppError::InvalidInput( + "Pi endpoint must be an absolute HTTP(S) URL without embedded credentials".to_string(), + ) + })?; + Ok(normalized) +} + +fn validate_pi_initial_endpoints(input: &ProviderMutationInput) -> Result<(), AppError> { + if let Some(meta) = input.meta.as_ref() { + for endpoint in meta.custom_endpoints.values() { + normalize_gateway_endpoint_for_write(&endpoint.url)?; + } + } + Ok(()) +} + +fn compensate_created_provider( + state: &AppState, + provider_id: &str, + native_patch_receipt: Option<&PiProviderPatchReceipt>, + cause: AppError, +) -> AppError { + let db_restored = state.db.delete_pi_catalog_provider(provider_id); + let file_restored = native_patch_receipt.map_or(Ok(()), PiProviderPatchReceipt::rollback); + if db_restored.is_ok() && file_restored.is_ok() { + authority_error(PiCatalogAuthority::PreviousRestored, cause.to_string()) + } else { + authority_error( + if db_restored.is_err() { + PiCatalogAuthority::MutatedDatabaseAuthoritative + } else { + PiCatalogAuthority::ProjectionPending + }, + format!("{}; create compensation was incomplete", cause), + ) + } +} + +fn compensate_portable_projection_ledger( + state: &AppState, + new_claims: &[(String, String)], + released_claims: &[(String, String)], + cause: AppError, +) -> AppError { + let mut restored = true; + for (provider_id, provider_key) in new_claims.iter().rev() { + if state + .db + .delete_pi_projection_key(provider_id, provider_key) + .is_err() + { + // Continue compensating the remaining claims even after one row + // fails. A short-circuit here would strand every earlier claim. + restored = false; + } + } + for (provider_id, provider_key) in released_claims { + if state + .db + .claim_pi_projection_key(provider_id, provider_key) + .is_err() + { + restored = false; + } + } + if restored { + authority_error(PiCatalogAuthority::PreviousRestored, cause) + } else { + authority_error( + PiCatalogAuthority::ProjectionPending, + format!("{cause}; imported projection-ledger compensation was incomplete"), + ) + } +} + +fn compensate_existing_provider( + state: &AppState, + previous: &ProviderAggregate, + was_current: bool, + projection: &crate::database::PiProviderProjection, + native_patch_receipt: Option<&PiProviderPatchReceipt>, + cause: AppError, +) -> AppError { + let db_restored = state + .db + .restore_pi_catalog_provider(previous, was_current, Some(projection)); + let file_restored = native_patch_receipt.map_or(Ok(()), PiProviderPatchReceipt::rollback); + if db_restored.is_ok() && file_restored.is_ok() { + authority_error(PiCatalogAuthority::PreviousRestored, cause.to_string()) + } else { + authority_error( + if db_restored.is_err() { + PiCatalogAuthority::MutatedDatabaseAuthoritative + } else { + PiCatalogAuthority::ProjectionPending + }, + format!("{}; catalog compensation was incomplete", cause), + ) + } +} + +fn restore_db_current(state: &AppState, previous: Option<&str>) -> Result<(), AppError> { + match previous { + Some(provider_id) => state.db.set_current_provider(PI_APP, provider_id), + None => state.db.clear_current_provider_for_app(PI_APP), + } +} + +impl PiCatalogSnapshot { + fn capture(state: &AppState, additional_native_key: Option<&str>) -> Result { + let aggregates = state.db.get_all_provider_aggregates(PI_APP)?; + let manifest = state.db.get_pi_projection_manifest()?; + let projections = manifest.values().cloned().collect::>(); + let mut native_keys = BTreeMap::::new(); + for projection in manifest.values() { + native_keys.insert(projection.provider_key.clone(), ()); + } + // Portable providers have no device-local claim yet and reconcile + // under their normalized id. Include every possible new key in the + // snapshot so a later runtime failure can remove it again. + for provider_id in aggregates.keys() { + native_keys.insert(non_empty_native_key(provider_id)?.to_string(), ()); + } + if let Some(key) = additional_native_key { + native_keys.insert(non_empty_native_key(key)?.to_string(), ()); + } + let models_path = get_pi_models_path()?; + let native = snapshot_pi_provider_values(&models_path, native_keys.into_keys())?; + Ok(Self { + aggregates, + projections, + db_current: state.db.get_current_provider(PI_APP)?, + native, + }) + } + + fn restore( + &self, + state: &AppState, + native_defaults_receipt: Option<&PiNativeDefaultsReceipt>, + native_patch_receipt: Option<&PiProviderPatchReceipt>, + ) -> Result<(), AppError> { + let mut failures = Vec::new(); + if let Err(error) = state.db.restore_pi_catalog_snapshot( + &self.aggregates, + &self.projections, + self.db_current.as_deref(), + ) { + failures.push(format!("database={error}")); + } + if let Some(receipt) = native_patch_receipt { + if let Err(error) = receipt.rollback() { + failures.push(format!("models={error}")); + } + } + if let Some(receipt) = native_defaults_receipt { + if let Err(error) = receipt.rollback() { + failures.push(format!("native_defaults={error}")); + } + } + // Current markers are derived indexes. Re-read the native authority so + // an external selection which superseded our receipt is preserved. + if let Err(error) = PiCatalogCoordinator::reconcile_current_indexes_from_native(state) { + failures.push(format!("current_indexes={error}")); + } + if failures.is_empty() { + Ok(()) + } else { + Err(AppError::Config(failures.join(", "))) + } + } +} + +fn resolve_native_current_provider( + state: &AppState, + defaults: &PiNativeDefaults, +) -> Result, AppError> { + let Some(provider_key) = defaults.default_provider.as_deref() else { + return Ok(None); + }; + let Some(projection) = state.db.get_pi_projection_for_key(provider_key)? else { + return Ok(None); + }; + Ok(state + .db + .get_provider_aggregate(PI_APP, &projection.provider_id)? + .is_some() + .then_some(projection.provider_id)) +} + +fn success(provider_id: Option) -> PiCatalogMutationResult { + PiCatalogMutationResult { + authority: PiCatalogAuthority::Published, + provider_id, + native_defaults_receipt: None, + native_patch_receipt: None, + native_fingerprint_preconditions: IndexMap::new(), + } +} + +impl PiCatalogMutationResult { + fn with_native_defaults_receipt(mut self, receipt: Option) -> Self { + self.native_defaults_receipt = receipt; + self + } + + fn with_native_patch_receipt(mut self, receipt: Option) -> Self { + self.native_patch_receipt = receipt; + self + } + + fn with_native_fingerprint_precondition( + mut self, + provider_key: String, + fingerprint: String, + ) -> Self { + self.native_fingerprint_preconditions + .insert(provider_key, fingerprint); + self + } +} + +fn authority_error(authority: PiCatalogAuthority, message: impl std::fmt::Display) -> AppError { + AppError::Message(format!( + "{message} [authoritative_state={}]", + authority.as_str() + )) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::database::Database; + use serde_json::json; + use std::sync::Arc; + + struct TestHome(Option); + + impl TestHome { + fn install(path: &std::path::Path) -> Result { + let previous = std::env::var_os("CC_SWITCH_TEST_HOME"); + std::env::set_var("CC_SWITCH_TEST_HOME", path); + crate::settings::reload_settings()?; + Ok(Self(previous)) + } + } + + impl Drop for TestHome { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("CC_SWITCH_TEST_HOME", value), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + let _ = crate::settings::reload_settings(); + } + } + + fn managed_input(id: &str, base_url: &str) -> ProviderMutationInput { + ProviderMutationInput { + id: id.to_string(), + name: id.to_string(), + settings_config: json!({ + "name": id, + "api": "openai-responses", + "baseUrl": base_url, + "apiKey": "literal-key", + "models": [{"id": "model-a", "name": "Model A"}] + }), + website_url: None, + category: None, + created_at: Some(1), + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + } + } + + fn insert_portable_pi_provider(db: &Database, id: &str) -> Result { + let config = json!({ + "name": "Portable Pi", + "api": "openai-responses", + "baseUrl": "https://portable.example/v1", + "apiKey": "literal-key", + "headers": {"x-capture": "present"}, + "models": [{ + "id": "portable-model", + "name": "Portable model", + "compat": {"supportsStore": true} + }] + }); + let conn = crate::database::lock_conn!(db.conn); + conn.execute( + "INSERT INTO providers + (id, app_type, name, settings_config, meta, is_current) + VALUES (?1, 'pi', 'Portable Pi', ?2, '{}', 0)", + rusqlite::params![id, config.to_string()], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(config) + } + + fn configure_pi_directory(path: &std::path::Path) -> Result<(), AppError> { + let mut app_settings = crate::settings::get_settings(); + app_settings.pi_config_dir = Some(path.to_string_lossy().into_owned()); + app_settings.pi_takeover_enabled = false; + crate::settings::update_settings(app_settings) + } + + #[test] + #[serial_test::serial] + #[cfg(unix)] + fn create_durability_failure_compensates_database_and_native_projection() -> Result<(), AppError> + { + let temp = tempfile::tempdir().expect("tempdir"); + let _home = TestHome::install(temp.path())?; + let pi_dir = temp.path().join("pi-agent"); + configure_pi_directory(&pi_dir)?; + let state = AppState::new(Arc::new(Database::memory()?)); + let models_path = pi_dir.join("models.json"); + crate::pi_config::shared_file::fail_next_parent_sync_for_test(&models_path); + + let error = PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: managed_input("ambiguous-create", "https://create.example/v1"), + provider_key: "ambiguous-native".to_string(), + activate_if_first: false, + }, + ) + .expect_err("directory durability failure must fail the service operation"); + assert!(error.to_string().contains("injected")); + assert!(state + .db + .get_provider_aggregate(PI_APP, "ambiguous-create")? + .is_none()); + assert!(state.db.get_pi_projection("ambiguous-create")?.is_none()); + assert!( + !models_path.exists(), + "the failed create must not leave a provider or empty shadow file" + ); + Ok(()) + } + + #[test] + #[serial_test::serial] + fn pi_endpoint_writes_reuse_gateway_validation_without_partial_state() -> Result<(), AppError> { + let temp = tempfile::tempdir().expect("tempdir"); + let _home = TestHome::install(temp.path())?; + let pi_dir = temp.path().join("pi-agent"); + configure_pi_directory(&pi_dir)?; + let state = AppState::new(Arc::new(Database::memory()?)); + let bad_url = "https://user:secret@mirror.example/v1"; + + let mut invalid_initial = managed_input("invalid-initial", "https://primary.example/v1"); + invalid_initial.meta = Some(crate::provider::ProviderMeta { + custom_endpoints: std::collections::HashMap::from([( + bad_url.to_string(), + crate::settings::CustomEndpoint { + url: bad_url.to_string(), + added_at: Some(1), + last_used: None, + }, + )]), + ..Default::default() + }); + let error = PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: invalid_initial, + provider_key: "invalid-initial".to_string(), + activate_if_first: false, + }, + ) + .expect_err("initial gateway endpoints must be validated before create"); + assert!(matches!(error, AppError::InvalidInput(_))); + assert!(!error.to_string().contains("secret")); + assert!(state + .db + .get_provider_aggregate(PI_APP, "invalid-initial")? + .is_none()); + + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: managed_input("endpoint-owner", "https://primary.example/v1"), + provider_key: "endpoint-owner".to_string(), + activate_if_first: false, + }, + )?; + let error = PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::AddEndpoint { + provider_id: "endpoint-owner".to_string(), + url: bad_url.to_string(), + }, + ) + .expect_err("an unusable endpoint must not be persisted"); + assert!(matches!(error, AppError::InvalidInput(_))); + assert!(!error.to_string().contains("secret")); + assert!(state + .db + .get_provider_aggregate(PI_APP, "endpoint-owner")? + .expect("provider remains") + .endpoints + .is_empty()); + + // A database restored from an older build may already contain this + // value. Runtime quarantines it, while deletion remains total over the + // persisted endpoint domain so the user can repair the row. + let key = ProviderKey::new(PI_APP, "endpoint-owner")?; + state + .db + .add_provider_endpoint(&key, NewEndpoint::new(bad_url, Some(2), None)?)?; + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::RemoveEndpoint { + provider_id: "endpoint-owner".to_string(), + url: bad_url.to_string(), + }, + )?; + assert!(state + .db + .get_provider_aggregate(PI_APP, "endpoint-owner")? + .expect("provider remains") + .endpoints + .is_empty()); + Ok(()) + } + + #[test] + #[serial_test::serial] + fn native_import_revalidates_exact_values_before_claiming_success() -> Result<(), AppError> { + let temp = tempfile::tempdir().expect("tempdir"); + let _home = TestHome::install(temp.path())?; + let pi_dir = temp.path().join("pi-agent"); + std::fs::create_dir_all(&pi_dir).expect("Pi directory"); + configure_pi_directory(&pi_dir)?; + let state = AppState::new(Arc::new(Database::memory()?)); + let models_path = pi_dir.join("models.json"); + let original = json!({ + "name": "Native import", + "api": "openai-responses", + "baseUrl": "https://original.example/v1", + "apiKey": "literal-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + std::fs::write( + &models_path, + serde_json::to_vec_pretty(&json!({"providers": {"native": original}})) + .expect("serialize original"), + ) + .expect("write original"); + let fingerprint = inspect_pi_native_entry(&models_path, "native", &BTreeMap::new())? + .expect("native entry") + .diagnostic + .fingerprint; + let external = serde_json::to_vec_pretty(&json!({ + "providers": { + "native": { + "name": "External replacement", + "api": "openai-responses", + "baseUrl": "https://external.example/v1", + "apiKey": "external-key", + "models": [{"id": "external-model"}] + } + } + })) + .expect("serialize external"); + crate::pi_config::document::replace_before_next_pi_provider_verify(&models_path, &external); + + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::ImportNative { + provider_key: "native".to_string(), + expected_fingerprint: fingerprint, + }, + ) + .expect_err("an external edit before the final barrier must reject import"); + assert!(state.db.get_provider_aggregate(PI_APP, "native")?.is_none()); + assert!(state.db.get_pi_projection("native")?.is_none()); + assert_eq!( + std::fs::read(&models_path).expect("external native file"), + external, + "the external writer remains authoritative" + ); + Ok(()) + } + + #[test] + #[serial_test::serial] + fn native_import_revalidates_raw_fingerprint_after_semantically_equivalent_edit( + ) -> Result<(), AppError> { + let temp = tempfile::tempdir().expect("tempdir"); + let _home = TestHome::install(temp.path())?; + let pi_dir = temp.path().join("pi-agent"); + std::fs::create_dir_all(&pi_dir).expect("Pi directory"); + configure_pi_directory(&pi_dir)?; + let state = AppState::new(Arc::new(Database::memory()?)); + let models_path = pi_dir.join("models.json"); + let original = br#"{ + "providers": { + "native": { + "name": "Native import", + "api": "openai-responses", + "baseUrl": "https://original.example/v1", + "apiKey": "literal-key", + "models": [{"id": "model-a", "name": "Model A"}] + } + } +}"#; + std::fs::write(&models_path, original).expect("write original"); + let fingerprint = inspect_pi_native_entry(&models_path, "native", &BTreeMap::new())? + .expect("native entry") + .diagnostic + .fingerprint; + let external = br#"{ + "providers": { + "native": { + // raw ownership changed while parsed values stayed identical + "name": "Native import", + "api": "openai-responses", + "baseUrl": "https://original.example/v1", + "apiKey": "literal-key", + "models": [ + {"id": "model-a", "name": "Model A"}, + ], + }, + }, +}"#; + crate::pi_config::document::replace_before_next_pi_provider_verify(&models_path, external); + + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::ImportNative { + provider_key: "native".to_string(), + expected_fingerprint: fingerprint, + }, + ) + .expect_err("a raw-only external edit must reject the ownership claim"); + assert!(state.db.get_provider_aggregate(PI_APP, "native")?.is_none()); + assert!(state.db.get_pi_projection("native")?.is_none()); + assert_eq!( + std::fs::read(&models_path).expect("external native file"), + external, + "raw external ownership must remain intact" + ); + Ok(()) + } + + #[test] + fn portable_sql_import_claims_only_an_absent_identity_key_and_is_idempotent( + ) -> Result<(), AppError> { + let source = Database::memory()?; + let expected = insert_portable_pi_provider(&source, "portable-pi")?; + let exported = source.export_sql_string()?; + assert!(!exported.contains("INSERT INTO \"pi_provider_projections\"")); + + let target = Arc::new(Database::memory()?); + target.import_sql_string(&exported)?; + assert!(target.get_pi_projection("portable-pi")?.is_none()); + let state = AppState::new(target.clone()); + let temp = tempfile::tempdir().expect("tempdir"); + let models_path = temp.path().join("models.json"); + + let project_direct = + |_: &str, _: &str, config: &PiManagedProviderConfig| -> Result { + serde_json::to_value(config).map_err(|source| AppError::JsonSerialize { source }) + }; + PiCatalogCoordinator::reconcile_portable_catalog_at(&state, &models_path, project_direct)?; + PiCatalogCoordinator::reconcile_portable_catalog_at(&state, &models_path, project_direct)?; + + let projection = target + .get_pi_projection("portable-pi")? + .expect("projection"); + assert_eq!(projection.provider_key, "portable-pi"); + let document: Value = + serde_json::from_slice(&std::fs::read(&models_path).expect("read models")) + .expect("parse models"); + assert_eq!( + document.pointer("/providers/portable-pi"), + Some(&expected), + "all portable provider fields must survive SQL import and native publication" + ); + Ok(()) + } + + #[test] + fn portable_import_never_claims_or_overwrites_an_unowned_native_key() -> Result<(), AppError> { + let db = Arc::new(Database::memory()?); + insert_portable_pi_provider(&db, "occupied")?; + let state = AppState::new(db.clone()); + let temp = tempfile::tempdir().expect("tempdir"); + let models_path = temp.path().join("models.json"); + let original = json!({ + "providers": { + "occupied": { + "api": "anthropic-messages", + "baseUrl": "https://native.example", + "apiKey": "native-secret", + "models": [{"id": "native-model"}] + } + } + }); + std::fs::write( + &models_path, + serde_json::to_vec_pretty(&original).expect("serialize"), + ) + .expect("write models"); + + let error = PiCatalogCoordinator::reconcile_portable_catalog_at( + &state, + &models_path, + |_, _, config| { + serde_json::to_value(config).map_err(|source| AppError::JsonSerialize { source }) + }, + ) + .expect_err("unowned native value must block import publication"); + assert!(error.to_string().contains("unclaimed native key")); + assert!(db.get_pi_projection("occupied")?.is_none()); + let after: Value = + serde_json::from_slice(&std::fs::read(&models_path).expect("read models")) + .expect("parse models"); + assert_eq!(after, original); + Ok(()) + } + + #[test] + fn portable_import_releases_orphan_claim_but_preserves_its_native_value() -> Result<(), AppError> + { + let source = Database::memory()?; + let remote = insert_portable_pi_provider(&source, "remote")?; + let exported = source.export_sql_string()?; + + let target = Arc::new(Database::memory()?); + insert_portable_pi_provider(&target, "local")?; + target.claim_pi_projection_key("local", "local-native")?; + target.import_sql_string(&exported)?; + assert!( + target.get_provider_aggregate(PI_APP, "local")?.is_none(), + "portable provider rows must be replaced" + ); + assert!( + target.get_pi_projection("local")?.is_some(), + "the device-local claim is intentionally preserved by restore" + ); + + let state = AppState::new(target.clone()); + let temp = tempfile::tempdir().expect("tempdir"); + let models_path = temp.path().join("models.json"); + let local_native = json!({ + "api": "anthropic-messages", + "baseUrl": "https://local.example", + "apiKey": "local-secret", + "models": [{"id": "local-model"}] + }); + std::fs::write( + &models_path, + serde_json::to_vec_pretty(&json!({ + "providers": {"local-native": local_native.clone()} + })) + .expect("serialize"), + ) + .expect("write models"); + + PiCatalogCoordinator::reconcile_portable_catalog_at( + &state, + &models_path, + |_, _, config| { + serde_json::to_value(config).map_err(|source| AppError::JsonSerialize { source }) + }, + )?; + + let manifest = target.get_pi_projection_manifest()?; + assert_eq!(manifest.len(), 1); + assert_eq!(manifest["remote"].provider_key, "remote"); + let document: Value = + serde_json::from_slice(&std::fs::read(&models_path).expect("read models")) + .expect("parse models"); + assert_eq!( + document.pointer("/providers/local-native"), + Some(&local_native), + "an orphaned exact-key value has no safe deletion proof and must become user-owned" + ); + assert_eq!(document.pointer("/providers/remote"), Some(&remote)); + Ok(()) + } + + #[test] + fn portable_import_of_empty_pi_catalog_releases_all_device_claims() -> Result<(), AppError> { + let empty_source = Database::memory()?; + let exported = empty_source.export_sql_string()?; + let target = Arc::new(Database::memory()?); + insert_portable_pi_provider(&target, "local")?; + target.claim_pi_projection_key("local", "local-native")?; + target.import_sql_string(&exported)?; + + let state = AppState::new(target.clone()); + let temp = tempfile::tempdir().expect("tempdir"); + let models_path = temp.path().join("models.json"); + let source = json!({"providers": {"local-native": {"models": [{"id": "m"}]}}}); + std::fs::write( + &models_path, + serde_json::to_vec_pretty(&source).expect("serialize"), + ) + .expect("models"); + + PiCatalogCoordinator::reconcile_portable_catalog_at( + &state, + &models_path, + |_, _, config| { + serde_json::to_value(config).map_err(|source| AppError::JsonSerialize { source }) + }, + )?; + + assert!(target.get_pi_projection_manifest()?.is_empty()); + assert_eq!( + serde_json::from_slice::(&std::fs::read(&models_path).expect("read models")) + .expect("parse"), + source, + "empty portable catalogs must not create, rewrite, or delete native values" + ); + Ok(()) + } + + #[test] + #[serial_test::serial] + fn catalog_projection_runtime_failures_restore_new_normalized_keys_and_file_absence( + ) -> Result<(), AppError> { + struct HomeGuard(Option); + impl Drop for HomeGuard { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("CC_SWITCH_TEST_HOME", value), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + let _ = crate::settings::reload_settings(); + } + } + + let temp = tempfile::tempdir().expect("tempdir"); + let _home = HomeGuard(std::env::var_os("CC_SWITCH_TEST_HOME")); + std::env::set_var("CC_SWITCH_TEST_HOME", temp.path()); + crate::settings::reload_settings()?; + let pi_dir = temp.path().join("pi-agent"); + let mut app_settings = crate::settings::get_settings(); + app_settings.pi_config_dir = Some(pi_dir.to_string_lossy().into_owned()); + app_settings.pi_takeover_enabled = false; + crate::settings::update_settings(app_settings)?; + + let db = Arc::new(Database::memory()?); + insert_portable_pi_provider(&db, "portable-pi")?; + let state = AppState::new(db.clone()); + state.proxy_service.fail_next_pi_reconcile_for_test(); + let error = PiCatalogCoordinator::reconcile_portable_import(&state) + .expect_err("injected runtime publication failure"); + assert!(error.to_string().contains("previous catalog was restored")); + assert!( + db.get_pi_projection("portable-pi")?.is_none(), + "a failed portable projection must release its new exact-key claim" + ); + assert!( + !pi_dir.join("models.json").exists(), + "a failed portable projection must restore a previously absent native file" + ); + assert!( + db.get_provider_aggregate(PI_APP, "portable-pi")?.is_some(), + "portable provider rows predate reconciliation and remain authoritative" + ); + + state.proxy_service.fail_next_pi_reconcile_for_test(); + let error = PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: managed_input("new-provider", "https://new.example/v1"), + provider_key: " new-native ".to_string(), + activate_if_first: false, + }, + ) + .expect_err("injected create publication failure"); + assert!(error.to_string().contains("previous catalog was restored")); + assert!(db.get_provider_aggregate(PI_APP, "new-provider")?.is_none()); + assert!(db.get_pi_projection("new-provider")?.is_none()); + assert!( + !pi_dir.join("models.json").exists(), + "the normalized additional create key must be part of the rollback snapshot" + ); + Ok(()) + } + + #[test] + #[serial_test::serial] + fn external_native_switch_repairs_indexes_before_deleting_the_inactive_provider( + ) -> Result<(), AppError> { + struct HomeGuard(Option); + impl Drop for HomeGuard { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("CC_SWITCH_TEST_HOME", value), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + let _ = crate::settings::reload_settings(); + } + } + + let temp = tempfile::tempdir().expect("tempdir"); + let _home = HomeGuard(std::env::var_os("CC_SWITCH_TEST_HOME")); + std::env::set_var("CC_SWITCH_TEST_HOME", temp.path()); + crate::settings::reload_settings()?; + let pi_dir = temp.path().join("pi-agent"); + let mut app_settings = crate::settings::get_settings(); + app_settings.pi_config_dir = Some(pi_dir.to_string_lossy().into_owned()); + app_settings.pi_takeover_enabled = false; + crate::settings::update_settings(app_settings)?; + + let db = Arc::new(Database::memory()?); + let state = AppState::new(db.clone()); + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: managed_input("provider-a", "https://a.example/v1"), + provider_key: "native-a".to_string(), + activate_if_first: true, + }, + )?; + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: managed_input("provider-b", "https://b.example/v1"), + provider_key: "native-b".to_string(), + activate_if_first: true, + }, + )?; + assert_eq!( + settings::get_current_provider(&AppType::Pi).as_deref(), + Some("provider-a") + ); + assert_eq!( + db.get_current_provider(PI_APP)?.as_deref(), + Some("provider-a") + ); + + crate::pi_config::native_settings::set_pi_native_default_with_receipt( + "native-b", "model-a", + )?; + assert_eq!( + PiCatalogCoordinator::current_native_provider(&state)?.as_deref(), + Some("provider-b"), + "native settings must immediately drive displayed current state" + ); + + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::DeleteProvider { + provider_id: "provider-a".to_string(), + }, + )?; + assert!(db.get_provider_aggregate(PI_APP, "provider-a")?.is_none()); + assert!(db.get_provider_aggregate(PI_APP, "provider-b")?.is_some()); + assert_eq!( + settings::get_current_provider(&AppType::Pi).as_deref(), + Some("provider-b") + ); + assert_eq!( + db.get_current_provider(PI_APP)?.as_deref(), + Some("provider-b") + ); + Ok(()) + } + + #[test] + #[serial_test::serial] + fn runtime_publication_failure_restores_endpoint_database_and_native_snapshot( + ) -> Result<(), AppError> { + struct HomeGuard(Option); + impl Drop for HomeGuard { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("CC_SWITCH_TEST_HOME", value), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + let _ = crate::settings::reload_settings(); + } + } + + let temp = tempfile::tempdir().expect("tempdir"); + let _home = HomeGuard(std::env::var_os("CC_SWITCH_TEST_HOME")); + std::env::set_var("CC_SWITCH_TEST_HOME", temp.path()); + crate::settings::reload_settings()?; + let pi_dir = temp.path().join("pi-agent"); + let mut app_settings = crate::settings::get_settings(); + app_settings.pi_config_dir = Some(pi_dir.to_string_lossy().into_owned()); + app_settings.pi_takeover_enabled = false; + crate::settings::update_settings(app_settings)?; + + let db = Arc::new(Database::memory()?); + let state = AppState::new(db.clone()); + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: managed_input("provider-a", "https://a.example/v1"), + provider_key: "native-a".to_string(), + activate_if_first: true, + }, + )?; + let before_aggregate = db + .get_provider_aggregate(PI_APP, "provider-a")? + .expect("provider"); + futures::executor::block_on(db.update_provider_health( + "provider-a", + PI_APP, + false, + Some("captured health".to_string()), + ))?; + let before_health = + futures::executor::block_on(db.get_provider_health("provider-a", PI_APP))?; + let models_path = get_pi_models_path()?; + let before_models = std::fs::read(&models_path).expect("models"); + + state.proxy_service.fail_next_pi_reconcile_for_test(); + let error = PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::AddEndpoint { + provider_id: "provider-a".to_string(), + url: "https://endpoint.example/v1".to_string(), + }, + ) + .expect_err("injected publication failure"); + assert!(error.to_string().contains("previous catalog was restored")); + + let after_aggregate = db + .get_provider_aggregate(PI_APP, "provider-a")? + .expect("provider remains"); + assert_eq!( + serde_json::to_value(after_aggregate).expect("after aggregate"), + serde_json::to_value(before_aggregate).expect("before aggregate") + ); + assert_eq!( + std::fs::read(models_path).expect("models after"), + before_models + ); + let after_health = + futures::executor::block_on(db.get_provider_health("provider-a", PI_APP))?; + assert_eq!( + serde_json::to_value(after_health).expect("after health"), + serde_json::to_value(before_health).expect("before health"), + "catalog compensation must not cascade-delete provider health history" + ); + Ok(()) + } + + #[tokio::test] + #[serial_test::serial] + async fn route_sort_publishes_an_even_runtime_without_rewriting_native_models() { + let result: Result<(), AppError> = async { + struct HomeGuard(Option); + impl Drop for HomeGuard { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("CC_SWITCH_TEST_HOME", value), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + let _ = crate::settings::reload_settings(); + } + } + + let temp = tempfile::tempdir().expect("tempdir"); + let _home = HomeGuard(std::env::var_os("CC_SWITCH_TEST_HOME")); + std::env::set_var("CC_SWITCH_TEST_HOME", temp.path()); + crate::settings::reload_settings()?; + let mut app_settings = crate::settings::get_settings(); + app_settings.pi_config_dir = + Some(temp.path().join("pi-agent").to_string_lossy().into_owned()); + app_settings.pi_takeover_enabled = false; + crate::settings::update_settings(app_settings)?; + + let db = Arc::new(Database::memory()?); + let state = AppState::new(db.clone()); + for (id, key) in [("provider-a", "native-a"), ("provider-b", "native-b")] { + PiCatalogCoordinator::apply( + &state, + PiCatalogMutation::CreateProvider { + input: managed_input(id, &format!("https://{id}.example/v1")), + provider_key: key.to_string(), + activate_if_first: true, + }, + )?; + } + state + .proxy_service + .set_takeover_for_app(PI_APP, true) + .await + .map_err(AppError::Message)?; + let models_path = get_pi_models_path()?; + let before_models = std::fs::read(&models_path).expect("models"); + + PiCatalogCoordinator::update_route_order( + &state, + vec![ + (ProviderKey::new(PI_APP, "provider-b")?, 0), + (ProviderKey::new(PI_APP, "provider-a")?, 1), + ], + )?; + + assert_eq!( + std::fs::read(&models_path).expect("models after sort"), + before_models, + "sorting is DB/runtime-only and must not rewrite Pi's shared file" + ); + let aggregates = db.get_all_provider_aggregates(PI_APP)?; + assert_eq!(aggregates["provider-b"].provider.sort_index, Some(0)); + assert_eq!(aggregates["provider-a"].provider.sort_index, Some(1)); + assert_eq!( + state + .proxy_service + .get_takeover_status() + .await + .map_err(AppError::Message)? + .pi_operational_state, + crate::proxy::types::PiTakeoverOperationalState::Active, + "sort publication must never leave the runtime at an odd admission epoch" + ); + + let mut corrupted_settings = crate::settings::get_settings(); + corrupted_settings.pi_gateway_token = None; + crate::settings::update_settings(corrupted_settings)?; + let missing_token = PiCatalogCoordinator::update_route_order( + &state, + vec![ + (ProviderKey::new(PI_APP, "provider-a")?, 0), + (ProviderKey::new(PI_APP, "provider-b")?, 1), + ], + ) + .expect_err("pure sorting must not silently rotate the gateway credential"); + assert!(missing_token + .to_string() + .contains("credential is unavailable")); + let aggregates = db.get_all_provider_aggregates(PI_APP)?; + assert_eq!(aggregates["provider-b"].provider.sort_index, Some(0)); + assert_eq!(aggregates["provider-a"].provider.sort_index, Some(1)); + assert_eq!( + std::fs::read(&models_path).expect("models after rejected sort"), + before_models + ); + state + .proxy_service + .set_takeover_for_app(PI_APP, false) + .await + .map_err(AppError::Message)?; + Ok(()) + } + .await; + result.expect("route sort"); + } +} diff --git a/src-tauri/src/services/pi_prompt_files.rs b/src-tauri/src/services/pi_prompt_files.rs new file mode 100644 index 000000000..5b1520011 --- /dev/null +++ b/src-tauri/src/services/pi_prompt_files.rs @@ -0,0 +1,433 @@ +//! Pi native instruction files and prompt templates. +//! +//! AGENTS.md is also the Prompt-library projection. SYSTEM.md and +//! APPEND_SYSTEM.md are direct native resources: file presence is activation +//! and there is no shadow enabled flag. + +use crate::error::AppError; +use crate::pi_config::native::get_pi_agent_dir; +use crate::pi_config::shared_file::{delete_shared_file, read_shared_file, replace_shared_file}; +use serde::{Deserialize, Serialize}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, LazyLock}; +use tokio::sync::{Mutex, OwnedMutexGuard}; + +const MAX_PROMPT_FILE_BYTES: u64 = 1024 * 1024; +const MAX_TEMPLATE_SLUG_BYTES: usize = 128; +static INSTRUCTION_FILE_LOCK: LazyLock>> = LazyLock::new(|| Arc::new(Mutex::new(()))); + +pub(crate) type PiInstructionFileGuard = OwnedMutexGuard<()>; + +pub(crate) fn lock_instruction_files() -> Result { + Ok(futures::executor::block_on( + INSTRUCTION_FILE_LOCK.clone().lock_owned(), + )) +} + +pub(crate) async fn lock_instruction_files_async() -> PiInstructionFileGuard { + INSTRUCTION_FILE_LOCK.clone().lock_owned().await +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum PiPromptFileKind { + GlobalContext, + SystemOverride, + SystemAppend, +} + +impl PiPromptFileKind { + fn filename(self) -> &'static str { + match self { + Self::GlobalContext => "AGENTS.md", + Self::SystemOverride => "SYSTEM.md", + Self::SystemAppend => "APPEND_SYSTEM.md", + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct PiPromptFileSnapshot { + pub kind: PiPromptFileKind, + pub path: String, + pub exists: bool, + pub revision: String, + pub content: String, +} + +pub struct PiPromptFileService; + +impl PiPromptFileService { + pub fn read(kind: PiPromptFileKind) -> Result { + let guard = lock_instruction_files()?; + Self::read_under_guard(&guard, kind) + } + + pub fn replace( + kind: PiPromptFileKind, + expected_revision: &str, + content: &str, + ) -> Result { + if kind == PiPromptFileKind::GlobalContext { + return Err(AppError::InvalidInput( + "Pi AGENTS.md is managed through the Prompt library".to_string(), + )); + } + validate_direct_instruction_content(content)?; + let guard = lock_instruction_files()?; + Self::replace_under_guard(&guard, kind, expected_revision, content) + } + + pub fn delete(kind: PiPromptFileKind, expected_revision: &str) -> Result { + if kind == PiPromptFileKind::GlobalContext { + return Err(AppError::InvalidInput( + "Pi AGENTS.md is managed through the Prompt library".to_string(), + )); + } + let guard = lock_instruction_files()?; + Self::delete_under_guard(&guard, kind, expected_revision) + } + + pub(crate) fn read_under_guard( + _guard: &PiInstructionFileGuard, + kind: PiPromptFileKind, + ) -> Result { + Self::read_at(&get_pi_agent_dir()?, kind) + } + + pub(crate) fn replace_under_guard( + _guard: &PiInstructionFileGuard, + kind: PiPromptFileKind, + expected_revision: &str, + content: &str, + ) -> Result { + Self::replace_at(&get_pi_agent_dir()?, kind, expected_revision, content) + } + + pub(crate) fn delete_under_guard( + _guard: &PiInstructionFileGuard, + kind: PiPromptFileKind, + expected_revision: &str, + ) -> Result { + Self::delete_at(&get_pi_agent_dir()?, kind, expected_revision) + } + + fn read_at(root: &Path, kind: PiPromptFileKind) -> Result { + let path = root.join(kind.filename()); + let snapshot = read_shared_file(&path, MAX_PROMPT_FILE_BYTES, "Pi prompt file")?; + let exists = snapshot.exists(); + let content = match snapshot.bytes { + Some(bytes) => String::from_utf8(bytes).map_err(|error| { + AppError::InvalidInput(format!( + "Pi prompt file must be UTF-8 ({}): {error}", + path.display() + )) + })?, + None => String::new(), + }; + Ok(PiPromptFileSnapshot { + kind, + path: path.to_string_lossy().into_owned(), + exists, + revision: snapshot.revision, + content, + }) + } + + fn replace_at( + root: &Path, + kind: PiPromptFileKind, + expected_revision: &str, + content: &str, + ) -> Result { + fs::create_dir_all(root).map_err(|error| AppError::io(root, error))?; + let path = root.join(kind.filename()); + replace_shared_file( + &path, + expected_revision, + content.as_bytes(), + MAX_PROMPT_FILE_BYTES, + Some(0o600), + "Pi prompt file", + )?; + Self::read_at(root, kind) + } + + fn delete_at( + root: &Path, + kind: PiPromptFileKind, + expected_revision: &str, + ) -> Result { + delete_shared_file( + &root.join(kind.filename()), + expected_revision, + MAX_PROMPT_FILE_BYTES, + "Pi prompt file", + ) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct PiPromptTemplate { + pub slug: String, + pub content: String, + pub revision: String, +} + +pub struct PiPromptTemplateService; + +impl PiPromptTemplateService { + pub fn list() -> Result, AppError> { + Self::list_at(&get_pi_agent_dir()?.join("prompts")) + } + + pub fn upsert( + slug: &str, + expected_revision: &str, + content: &str, + ) -> Result { + Self::upsert_at( + &get_pi_agent_dir()?.join("prompts"), + slug, + expected_revision, + content, + ) + } + + pub fn delete(slug: &str, expected_revision: &str) -> Result { + validate_template_slug(slug)?; + delete_shared_file( + &template_path(&get_pi_agent_dir()?.join("prompts"), slug), + expected_revision, + MAX_PROMPT_FILE_BYTES, + "Pi prompt template", + ) + } + + fn list_at(dir: &Path) -> Result, AppError> { + let entries = match fs::read_dir(dir) { + Ok(entries) => entries, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(Vec::new()), + Err(error) => return Err(AppError::io(dir, error)), + }; + let mut templates = Vec::new(); + for entry in entries { + let entry = entry.map_err(|error| AppError::io(dir, error))?; + let path = entry.path(); + let Some(slug) = path.file_stem().and_then(|value| value.to_str()) else { + continue; + }; + if path.extension().and_then(|value| value.to_str()) != Some("md") + || validate_template_slug(slug).is_err() + { + continue; + } + let snapshot = read_shared_file(&path, MAX_PROMPT_FILE_BYTES, "Pi prompt template")?; + let Some(bytes) = snapshot.bytes else { + continue; + }; + let content = String::from_utf8(bytes).map_err(|error| { + AppError::InvalidInput(format!( + "Pi prompt template must be UTF-8 ({}): {error}", + path.display() + )) + })?; + templates.push(PiPromptTemplate { + slug: slug.to_string(), + content, + revision: snapshot.revision, + }); + } + templates.sort_by(|left, right| left.slug.cmp(&right.slug)); + Ok(templates) + } + + fn upsert_at( + dir: &Path, + slug: &str, + expected_revision: &str, + content: &str, + ) -> Result { + validate_template_slug(slug)?; + fs::create_dir_all(dir).map_err(|error| AppError::io(dir, error))?; + let snapshot = replace_shared_file( + &template_path(dir, slug), + expected_revision, + content.as_bytes(), + MAX_PROMPT_FILE_BYTES, + Some(0o600), + "Pi prompt template", + )?; + Ok(PiPromptTemplate { + slug: slug.to_string(), + content: content.to_string(), + revision: snapshot.revision, + }) + } +} + +fn template_path(dir: &Path, slug: &str) -> PathBuf { + dir.join(format!("{slug}.md")) +} + +fn validate_direct_instruction_content(content: &str) -> Result<(), AppError> { + if content.trim().is_empty() { + Err(AppError::InvalidInput( + "Pi SYSTEM.md and APPEND_SYSTEM.md content cannot be blank; delete the file to deactivate it" + .to_string(), + )) + } else { + Ok(()) + } +} + +fn validate_template_slug(slug: &str) -> Result<(), AppError> { + let valid = !slug.is_empty() + && slug.len() <= MAX_TEMPLATE_SLUG_BYTES + && slug != "." + && slug != ".." + && slug.trim() == slug + && !slug.starts_with('.') + && !slug.ends_with('.') + && !slug + .chars() + .any(|character| character.is_control() || matches!(character, '/' | '\\')); + if valid { + Ok(()) + } else { + Err(AppError::InvalidInput( + "Pi prompt-template slug must be one visible filename (1-128 UTF-8 bytes)".to_string(), + )) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn all_instruction_files_use_presence_and_revision_as_native_state() { + let temp = tempfile::tempdir().expect("tempdir"); + for kind in [ + PiPromptFileKind::GlobalContext, + PiPromptFileKind::SystemOverride, + PiPromptFileKind::SystemAppend, + ] { + let missing = PiPromptFileService::read_at(temp.path(), kind).expect("missing"); + assert!(!missing.exists); + // scripts/pi-transport-capture.mjs executes pinned Pi's + // DefaultResourceLoader at + // ab366ebe94cacd419d986be454f12b1b9913aaca and records all three + // zero-byte files as present resources. + let empty = PiPromptFileService::replace_at(temp.path(), kind, "missing", "") + .expect("create empty instruction file"); + assert!(empty.exists); + assert_eq!(empty.content, ""); + let saved = + PiPromptFileService::replace_at(temp.path(), kind, &empty.revision, "content") + .expect("replace"); + assert!(saved.exists); + assert_eq!(saved.content, "content"); + assert!(PiPromptFileService::delete_at(temp.path(), kind, "missing").is_err()); + assert!( + PiPromptFileService::delete_at(temp.path(), kind, &saved.revision).expect("delete") + ); + } + } + + #[test] + fn direct_instruction_save_rejects_blank_content_without_redefining_native_presence() { + for content in ["", " \n\t"] { + assert!(validate_direct_instruction_content(content).is_err()); + } + assert!(validate_direct_instruction_content("# Explicit override").is_ok()); + } + + #[cfg(unix)] + #[test] + fn direct_instruction_entry_never_reports_failure_with_its_attempt_live() { + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp + .path() + .join(PiPromptFileKind::SystemOverride.filename()); + + crate::pi_config::shared_file::fail_next_parent_sync_for_test(&path); + PiPromptFileService::replace_at( + temp.path(), + PiPromptFileKind::SystemOverride, + "missing", + "created", + ) + .expect_err("failed create must be compensated"); + assert!(!path.exists()); + + let before = PiPromptFileService::replace_at( + temp.path(), + PiPromptFileKind::SystemOverride, + "missing", + "before", + ) + .expect("seed"); + crate::pi_config::shared_file::fail_next_parent_sync_for_test(&path); + PiPromptFileService::replace_at( + temp.path(), + PiPromptFileKind::SystemOverride, + &before.revision, + "after", + ) + .expect_err("failed replace must restore its before-image"); + assert_eq!( + fs::read_to_string(&path).expect("before restored"), + "before" + ); + + let before = PiPromptFileService::read_at(temp.path(), PiPromptFileKind::SystemOverride) + .expect("snapshot"); + crate::pi_config::shared_file::fail_next_parent_sync_for_test(&path); + PiPromptFileService::delete_at( + temp.path(), + PiPromptFileKind::SystemOverride, + &before.revision, + ) + .expect_err("failed delete must restore its before-image"); + assert_eq!( + fs::read_to_string(&path).expect("before restored"), + "before" + ); + } + + #[test] + fn templates_reject_ambiguous_or_traversing_slugs() { + for slug in [ + "", + ".", + "..", + ".hidden", + "trailing.", + " padded", + "a/b", + r"a\b", + ] { + assert!(validate_template_slug(slug).is_err(), "{slug:?}"); + } + for slug in ["review-pr", "release.v2", "评审", "SYSTEM"] { + assert!(validate_template_slug(slug).is_ok(), "{slug:?}"); + } + } + + #[test] + fn empty_template_is_present_and_round_trips_like_pinned_pi() { + // scripts/pi-transport-capture.mjs executes pinned Pi + // ab366ebe94cacd419d986be454f12b1b9913aaca and confirms that an empty + // prompts/empty.md is discovered as an active template. + let temp = tempfile::tempdir().expect("tempdir"); + let created = PiPromptTemplateService::upsert_at(temp.path(), "empty", "missing", "") + .expect("create empty template"); + assert_eq!(created.content, ""); + let listed = PiPromptTemplateService::list_at(temp.path()).expect("list templates"); + assert_eq!(listed, vec![created]); + } +} diff --git a/src-tauri/src/services/prompt.rs b/src-tauri/src/services/prompt.rs index d0dde040d..726c01f31 100644 --- a/src-tauri/src/services/prompt.rs +++ b/src-tauri/src/services/prompt.rs @@ -2,10 +2,17 @@ use indexmap::IndexMap; use crate::app_config::AppType; use crate::config::write_text_file; +use crate::database::Database; use crate::error::AppError; use crate::prompt::Prompt; use crate::prompt_files::prompt_file_path; +use crate::services::pi_prompt_files::{ + lock_instruction_files, PiInstructionFileGuard, PiPromptFileKind, PiPromptFileService, + PiPromptFileSnapshot, +}; use crate::store::AppState; +use serde::Serialize; +use sha2::{Digest, Sha256}; /// 安全地获取当前 Unix 时间戳 fn get_unix_timestamp() -> Result { @@ -17,20 +24,50 @@ fn get_unix_timestamp() -> Result { pub struct PromptService; +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct PiPromptLibraryStatus { + pub native_exists: bool, + pub native_revision: String, + pub matched_prompt_id: Option, + pub needs_reconciliation: bool, +} + impl PromptService { pub fn get_prompts( state: &AppState, app: AppType, ) -> Result, AppError> { + if matches!(app, AppType::Pi) { + let guard = lock_instruction_files()?; + return Self::inspect_pi_library_under_guard(state.db.as_ref(), &guard) + .map(|(prompts, _)| prompts); + } state.db.get_prompts(app.as_str()) } + /// Inspect Pi's live AGENTS.md without adopting it into the portable + /// library. File presence and exact bytes are the effective active truth; + /// persisted `enabled` flags are only a projection that explicit + /// reconciliation may repair. + pub fn get_pi_library_status(state: &AppState) -> Result { + let guard = lock_instruction_files()?; + Self::inspect_pi_library_under_guard(state.db.as_ref(), &guard).map(|(_, status)| status) + } + + pub fn reconcile_pi_library(state: &AppState) -> Result<(), AppError> { + Self::reconcile_pi_portable_import(state) + } + pub fn upsert_prompt( state: &AppState, app: AppType, _id: &str, prompt: Prompt, ) -> Result<(), AppError> { + if matches!(app, AppType::Pi) { + return Self::upsert_pi_prompt(state, prompt); + } // 检查是否为已启用的提示词 let is_enabled = prompt.enabled; @@ -58,6 +95,32 @@ impl PromptService { } pub fn delete_prompt(state: &AppState, app: AppType, id: &str) -> Result<(), AppError> { + if matches!(app, AppType::Pi) { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + let guard = lock_instruction_files()?; + let prompts = state.db.get_prompts(AppType::Pi.as_str())?; + let snapshot = + PiPromptFileService::read_under_guard(&guard, PiPromptFileKind::GlobalContext)?; + reject_pi_prompt_mutation_during_drift(&prompts, &snapshot, id, None)?; + if prompts.get(id).is_some_and(|prompt| prompt.enabled) + || preferred_pi_prompt_match(&prompts, &snapshot) == Some(id) + { + return Err(AppError::InvalidInput( + "无法删除 Pi 当前生效的提示词".to_string(), + )); + } + let mut after = prompts.clone(); + after.shift_remove(id); + return state.db.compare_exchange_prompt_selection( + AppType::Pi.as_str(), + &prompts, + &after, + ); + } let prompts = state.db.get_prompts(app.as_str())?; if let Some(prompt) = prompts.get(id) { @@ -71,6 +134,9 @@ impl PromptService { } pub fn enable_prompt(state: &AppState, app: AppType, id: &str) -> Result<(), AppError> { + if matches!(app, AppType::Pi) { + return Self::enable_pi_prompt(state, id); + } // 回填当前 live 文件内容到已启用的提示词,或创建备份 let target_path = prompt_file_path(&app)?; if target_path.exists() { @@ -144,6 +210,36 @@ impl PromptService { } pub fn import_from_file(state: &AppState, app: AppType) -> Result { + if matches!(app, AppType::Pi) { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + let guard = lock_instruction_files()?; + let snapshot = + PiPromptFileService::read_under_guard(&guard, PiPromptFileKind::GlobalContext)?; + if !snapshot.exists { + return Err(AppError::Message("Pi AGENTS.md does not exist".to_string())); + } + let before = state.db.get_prompts(AppType::Pi.as_str())?; + let timestamp = get_unix_timestamp()?; + let id = format!("imported-{timestamp}"); + let prompt = Prompt { + id: id.clone(), + name: format!( + "导入的提示词 {}", + chrono::Local::now().format("%Y-%m-%d %H:%M") + ), + content: snapshot.content.clone(), + description: Some("从 Pi AGENTS.md 导入".to_string()), + enabled: true, + created_at: Some(timestamp), + updated_at: Some(timestamp), + }; + Self::import_pi_snapshot(state, &guard, &snapshot, &before, prompt)?; + return Ok(id); + } let file_path = prompt_file_path(&app)?; if !file_path.exists() { @@ -173,6 +269,12 @@ impl PromptService { } pub fn get_current_file_content(app: AppType) -> Result, AppError> { + if matches!(app, AppType::Pi) { + let guard = lock_instruction_files()?; + let snapshot = + PiPromptFileService::read_under_guard(&guard, PiPromptFileKind::GlobalContext)?; + return Ok(snapshot.exists.then_some(snapshot.content)); + } let file_path = prompt_file_path(&app)?; if !file_path.exists() { return Ok(None); @@ -188,6 +290,40 @@ impl PromptService { state: &AppState, app: AppType, ) -> Result { + if matches!(app, AppType::Pi) { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + let guard = lock_instruction_files()?; + let existing = state.db.get_prompts(app.as_str())?; + if !existing.is_empty() { + return Ok(0); + } + let snapshot = + PiPromptFileService::read_under_guard(&guard, PiPromptFileKind::GlobalContext)?; + if !snapshot.exists { + return Ok(0); + } + let timestamp = get_unix_timestamp()?; + let id = format!("auto-imported-{timestamp}"); + let prompt = Prompt { + id: id.clone(), + name: format!( + "Auto-imported Prompt {}", + chrono::Local::now().format("%Y-%m-%d %H:%M") + ), + content: snapshot.content.clone(), + description: Some("Automatically imported on first launch".to_string()), + enabled: true, + created_at: Some(timestamp), + updated_at: Some(timestamp), + }; + Self::import_pi_snapshot(state, &guard, &snapshot, &existing, prompt)?; + return Ok(1); + } + // 幂等性保护:该应用已有提示词则跳过 let existing = state.db.get_prompts(app.as_str())?; if !existing.is_empty() { @@ -239,4 +375,1046 @@ impl PromptService { log::info!("自动导入完成: {}", app.as_str()); Ok(1) } + + /// Reconcile portable Pi prompt rows to this device's native AGENTS.md. + /// + /// Prompt content is portable, but native instruction files are not. The + /// live file therefore decides which library row is active after an + /// import: exact content adopts an existing row, otherwise a local + /// counterpart is added. A missing file disables every row. The file is + /// never created, replaced, or deleted by portable reconciliation. + pub(crate) fn reconcile_pi_portable_import(state: &AppState) -> Result<(), AppError> { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + let guard = lock_instruction_files()?; + Self::reconcile_pi_native_under_guard(state.db.as_ref(), &guard) + } + + pub(crate) fn reconcile_pi_native_under_guard( + db: &Database, + guard: &PiInstructionFileGuard, + ) -> Result<(), AppError> { + const MAX_EXTERNAL_RETRIES: usize = 3; + let original = db.get_prompts(AppType::Pi.as_str())?; + + for _ in 0..MAX_EXTERNAL_RETRIES { + let snapshot = + PiPromptFileService::read_under_guard(guard, PiPromptFileKind::GlobalContext)?; + let prompts = build_pi_reconciled_library(&original, &snapshot); + let precommit = + PiPromptFileService::read_under_guard(guard, PiPromptFileKind::GlobalContext)?; + if precommit.revision != snapshot.revision { + continue; + } + + db.compare_exchange_prompt_selection(AppType::Pi.as_str(), &original, &prompts)?; + #[cfg(test)] + apply_pi_native_binding_after_save_hooks_for_test(db, &snapshot.path)?; + let verified = + match PiPromptFileService::read_under_guard(guard, PiPromptFileKind::GlobalContext) + { + Ok(verified) => verified, + Err(error) => { + restore_pi_library_after_failed_native_binding( + db, &original, &prompts, &error, + )?; + return Err(error); + } + }; + if verified.revision == snapshot.revision { + return Ok(()); + } + restore_pi_library_after_failed_native_binding( + db, + &original, + &prompts, + &AppError::Conflict( + "Pi AGENTS.md changed after prompt-library publication".to_string(), + ), + )?; + } + + Err(AppError::Conflict( + "Pi AGENTS.md kept changing during portable prompt reconciliation".to_string(), + )) + } + + fn inspect_pi_library_under_guard( + db: &Database, + guard: &PiInstructionFileGuard, + ) -> Result<(IndexMap, PiPromptLibraryStatus), AppError> { + let snapshot = + PiPromptFileService::read_under_guard(guard, PiPromptFileKind::GlobalContext)?; + let mut prompts = db.get_prompts(AppType::Pi.as_str())?; + let persisted_enabled = prompts + .iter() + .filter_map(|(id, prompt)| prompt.enabled.then_some(id.clone())) + .collect::>(); + let matched_prompt_id = + preferred_pi_prompt_match(&prompts, &snapshot).map(ToOwned::to_owned); + let expected_enabled = matched_prompt_id.iter().cloned().collect::>(); + let needs_reconciliation = persisted_enabled != expected_enabled + || (snapshot.exists && matched_prompt_id.is_none()); + + for (id, prompt) in &mut prompts { + prompt.enabled = matched_prompt_id.as_deref() == Some(id.as_str()); + } + + Ok(( + prompts, + PiPromptLibraryStatus { + native_exists: snapshot.exists, + native_revision: snapshot.revision, + matched_prompt_id, + needs_reconciliation, + }, + )) + } + + fn upsert_pi_prompt(state: &AppState, prompt: Prompt) -> Result<(), AppError> { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + let guard = lock_instruction_files()?; + let before = state.db.get_prompts(AppType::Pi.as_str())?; + let snapshot = + PiPromptFileService::read_under_guard(&guard, PiPromptFileKind::GlobalContext)?; + reject_pi_prompt_mutation_during_drift(&before, &snapshot, &prompt.id, Some(&prompt))?; + let mut prompts = before.clone(); + let previous = prompts.insert(prompt.id.clone(), prompt.clone()); + let current_enabled = previous + .as_ref() + .filter(|candidate| candidate.enabled) + .or_else(|| { + before + .values() + .find(|candidate| candidate.id != prompt.id && candidate.enabled) + }); + + if prompt.enabled { + ensure_pi_library_projection_matches(&snapshot, current_enabled)?; + for candidate in prompts.values_mut() { + candidate.enabled = candidate.id == prompt.id; + } + let published = PiPromptFileService::replace_under_guard( + &guard, + PiPromptFileKind::GlobalContext, + &snapshot.revision, + &prompt.content, + )?; + if let Err(error) = + state + .db + .compare_exchange_prompt_selection(AppType::Pi.as_str(), &before, &prompts) + { + restore_pi_prompt_file(&guard, &published, &snapshot)?; + return Err(error); + } + return Ok(()); + } + + if previous.as_ref().is_some_and(|value| value.enabled) { + ensure_pi_library_projection_matches(&snapshot, previous.as_ref())?; + let removed = PiPromptFileService::delete_under_guard( + &guard, + PiPromptFileKind::GlobalContext, + &snapshot.revision, + )?; + if let Err(error) = + state + .db + .compare_exchange_prompt_selection(AppType::Pi.as_str(), &before, &prompts) + { + if removed { + let missing = PiPromptFileService::read_under_guard( + &guard, + PiPromptFileKind::GlobalContext, + )?; + PiPromptFileService::replace_under_guard( + &guard, + PiPromptFileKind::GlobalContext, + &missing.revision, + &snapshot.content, + )?; + } + return Err(error); + } + return Ok(()); + } + + state + .db + .compare_exchange_prompt_selection(AppType::Pi.as_str(), &before, &prompts) + } + + fn enable_pi_prompt(state: &AppState, id: &str) -> Result<(), AppError> { + let _switch_guard = futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ); + let guard = lock_instruction_files()?; + let before = state.db.get_prompts(AppType::Pi.as_str())?; + let target = before + .get(id) + .cloned() + .ok_or_else(|| AppError::InvalidInput(format!("提示词 {id} 不存在")))?; + let snapshot = + PiPromptFileService::read_under_guard(&guard, PiPromptFileKind::GlobalContext)?; + ensure_pi_library_projection_matches( + &snapshot, + before.values().find(|candidate| candidate.enabled), + )?; + let published = PiPromptFileService::replace_under_guard( + &guard, + PiPromptFileKind::GlobalContext, + &snapshot.revision, + &target.content, + )?; + + let mut after = before.clone(); + for prompt in after.values_mut() { + prompt.enabled = prompt.id == id; + } + if let Err(error) = + state + .db + .compare_exchange_prompt_selection(AppType::Pi.as_str(), &before, &after) + { + restore_pi_prompt_file(&guard, &published, &snapshot)?; + return Err(error); + } + Ok(()) + } + + fn import_pi_snapshot( + state: &AppState, + guard: &PiInstructionFileGuard, + snapshot: &PiPromptFileSnapshot, + before: &IndexMap, + prompt: Prompt, + ) -> Result<(), AppError> { + let mut prompts = before.clone(); + for prompt in prompts.values_mut() { + prompt.enabled = false; + } + // Import is an explicit reconciliation action. The native file is + // already active by presence, so its exact DB counterpart becomes the + // sole enabled library entry without rewriting the user-owned file. + prompts.insert(prompt.id.clone(), prompt); + + let precommit = + PiPromptFileService::read_under_guard(guard, PiPromptFileKind::GlobalContext)?; + if precommit.revision != snapshot.revision { + return Err(AppError::Conflict( + "Pi AGENTS.md changed before prompt import publication".to_string(), + )); + } + state + .db + .compare_exchange_prompt_selection(AppType::Pi.as_str(), before, &prompts)?; + #[cfg(test)] + apply_pi_native_binding_after_save_hooks_for_test(state.db.as_ref(), &snapshot.path)?; + let verified = + match PiPromptFileService::read_under_guard(guard, PiPromptFileKind::GlobalContext) { + Ok(verified) => verified, + Err(error) => { + restore_pi_library_after_failed_native_binding( + state.db.as_ref(), + before, + &prompts, + &error, + )?; + return Err(error); + } + }; + if verified.revision != snapshot.revision { + let error = AppError::Conflict("Pi AGENTS.md changed during prompt import".to_string()); + restore_pi_library_after_failed_native_binding( + state.db.as_ref(), + before, + &prompts, + &error, + )?; + return Err(error); + } + Ok(()) + } +} + +fn preferred_pi_prompt_match<'a>( + prompts: &'a IndexMap, + snapshot: &PiPromptFileSnapshot, +) -> Option<&'a str> { + if !snapshot.exists { + return None; + } + prompts + .iter() + .find_map(|(id, prompt)| { + (prompt.enabled && prompt.content == snapshot.content).then_some(id.as_str()) + }) + .or_else(|| { + prompts.iter().find_map(|(id, prompt)| { + (prompt.content == snapshot.content).then_some(id.as_str()) + }) + }) +} + +fn pi_library_needs_reconciliation( + prompts: &IndexMap, + snapshot: &PiPromptFileSnapshot, +) -> bool { + let persisted_enabled = prompts + .iter() + .filter_map(|(id, prompt)| prompt.enabled.then_some(id.as_str())) + .collect::>(); + let matched = preferred_pi_prompt_match(prompts, snapshot); + persisted_enabled != matched.into_iter().collect::>() + || (snapshot.exists && matched.is_none()) +} + +fn reject_pi_prompt_mutation_during_drift( + prompts: &IndexMap, + snapshot: &PiPromptFileSnapshot, + target_id: &str, + replacement: Option<&Prompt>, +) -> Result<(), AppError> { + if !pi_library_needs_reconciliation(prompts, snapshot) { + return Ok(()); + } + let target_is_persisted_active = prompts.get(target_id).is_some_and(|prompt| prompt.enabled); + let target_is_live_match = preferred_pi_prompt_match(prompts, snapshot) == Some(target_id); + let replacement_claims_live = replacement.is_some_and(|prompt| { + prompt.enabled || (snapshot.exists && prompt.content == snapshot.content) + }); + if target_is_persisted_active || target_is_live_match || replacement_claims_live { + return Err(AppError::Conflict( + "Pi AGENTS.md and the prompt library disagree; reconcile native truth before changing \ + an active prompt" + .to_string(), + )); + } + Ok(()) +} + +fn build_pi_reconciled_library( + original: &IndexMap, + snapshot: &PiPromptFileSnapshot, +) -> IndexMap { + let mut prompts = original.clone(); + for prompt in prompts.values_mut() { + prompt.enabled = false; + } + if !snapshot.exists { + return prompts; + } + + let active_id = preferred_pi_prompt_match(original, snapshot) + .map(ToOwned::to_owned) + .unwrap_or_else(|| { + let digest = format!("{:x}", Sha256::digest(snapshot.content.as_bytes())); + let base = format!("native-{digest}"); + let mut id = base.clone(); + let mut suffix = 1_u32; + while prompts.contains_key(&id) { + id = format!("{base}-{suffix}"); + suffix += 1; + } + let timestamp = chrono::Utc::now().timestamp(); + prompts.insert( + id.clone(), + Prompt { + id: id.clone(), + name: "Imported from Pi AGENTS.md".to_string(), + content: snapshot.content.clone(), + description: Some( + "Device-local native state preserved during portable import".to_string(), + ), + enabled: false, + created_at: Some(timestamp), + updated_at: Some(timestamp), + }, + ); + id + }); + prompts + .get_mut(&active_id) + .expect("selected Pi prompt is present") + .enabled = true; + prompts +} + +fn restore_pi_library_after_failed_native_binding( + db: &Database, + original: &IndexMap, + attempted: &IndexMap, + cause: &AppError, +) -> Result<(), AppError> { + db.restore_prompt_selection_if_attempted(AppType::Pi.as_str(), attempted, original) + .map_err(|restore_error| { + AppError::Config(format!( + "Pi prompt-library publication lost its native revision ({cause}) and failed to \ + restore the previous portable library without overwriting a newer database \ + revision: {restore_error}" + )) + }) +} + +#[cfg(test)] +static PI_NATIVE_BINDING_AFTER_SAVE_REPLACEMENTS: std::sync::LazyLock< + std::sync::Mutex>, +> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::VecDeque::new())); + +#[cfg(test)] +fn apply_pi_native_binding_after_save_hooks_for_test( + db: &Database, + path: &str, +) -> Result<(), AppError> { + if let Some(replacement) = PI_NATIVE_BINDING_AFTER_SAVE_DB_REPLACEMENTS + .lock() + .map_err(|error| AppError::Lock(error.to_string()))? + .pop_front() + { + db.save_prompt_selection(AppType::Pi.as_str(), &replacement)?; + } + let replacement = PI_NATIVE_BINDING_AFTER_SAVE_REPLACEMENTS + .lock() + .map_err(|error| AppError::Lock(error.to_string()))? + .pop_front(); + if let Some(replacement) = replacement { + std::fs::write(path, replacement).map_err(|error| AppError::io(path, error))?; + } + Ok(()) +} + +#[cfg(test)] +static PI_NATIVE_BINDING_AFTER_SAVE_DB_REPLACEMENTS: std::sync::LazyLock< + std::sync::Mutex>>, +> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::VecDeque::new())); + +#[cfg(test)] +fn replace_pi_library_after_next_reconcile_save_for_test(replacement: IndexMap) { + PI_NATIVE_BINDING_AFTER_SAVE_DB_REPLACEMENTS + .lock() + .expect("Pi reconcile DB hook lock") + .push_back(replacement); +} + +#[cfg(test)] +fn replace_pi_agents_after_each_native_binding_save_for_test( + replacements: impl IntoIterator, +) { + *PI_NATIVE_BINDING_AFTER_SAVE_REPLACEMENTS + .lock() + .expect("Pi reconcile hook lock") = + replacements.into_iter().map(ToOwned::to_owned).collect(); +} + +fn ensure_pi_library_projection_matches( + snapshot: &PiPromptFileSnapshot, + current_enabled: Option<&Prompt>, +) -> Result<(), AppError> { + match current_enabled { + Some(prompt) if snapshot.exists && snapshot.content == prompt.content => Ok(()), + Some(_) => Err(AppError::Conflict( + "Pi AGENTS.md changed outside CC Switch; import or reconcile it before switching prompts" + .to_string(), + )), + None if !snapshot.exists => Ok(()), + None => Err(AppError::Conflict( + "Pi AGENTS.md is user-owned; import it before enabling a library prompt".to_string(), + )), + } +} + +fn restore_pi_prompt_file( + guard: &PiInstructionFileGuard, + published: &PiPromptFileSnapshot, + previous: &PiPromptFileSnapshot, +) -> Result<(), AppError> { + if previous.exists { + PiPromptFileService::replace_under_guard( + guard, + PiPromptFileKind::GlobalContext, + &published.revision, + &previous.content, + )?; + } else { + PiPromptFileService::delete_under_guard( + guard, + PiPromptFileKind::GlobalContext, + &published.revision, + )?; + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::database::Database; + use serial_test::serial; + use std::ffi::OsString; + use std::sync::Arc; + + struct EnvRestore { + key: &'static str, + previous: Option, + } + + impl EnvRestore { + fn set(key: &'static str, value: &std::path::Path) -> Self { + let previous = std::env::var_os(key); + std::env::set_var(key, value); + Self { key, previous } + } + } + + impl Drop for EnvRestore { + fn drop(&mut self) { + match self.previous.take() { + Some(value) => std::env::set_var(self.key, value), + None => std::env::remove_var(self.key), + } + } + } + + fn prompt(id: &str, content: &str, enabled: bool, timestamp: i64) -> Prompt { + Prompt { + id: id.to_string(), + name: id.to_string(), + content: content.to_string(), + description: None, + enabled, + created_at: Some(timestamp), + updated_at: Some(timestamp), + } + } + + #[test] + #[serial] + fn public_pi_import_reconciles_an_active_empty_agents_file() { + // Pinned-Pi provenance: scripts/pi-transport-capture.mjs executes + // DefaultResourceLoader at ab366ebe94cacd419d986be454f12b1b9913aaca + // and records an existing zero-byte AGENTS.md as an active resource. + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + + let imported = + PromptService::import_from_file(&state, AppType::Pi).expect("import empty AGENTS.md"); + let prompts = PromptService::get_prompts(&state, AppType::Pi).expect("read prompts"); + let active = prompts.get(&imported).expect("imported prompt"); + assert!(active.enabled); + assert_eq!(active.content, ""); + assert_eq!( + prompts.values().filter(|prompt| prompt.enabled).count(), + 1, + "the native active file must have exactly one active DB owner" + ); + assert_eq!( + PromptService::get_current_file_content(AppType::Pi).expect("read live"), + Some(String::new()) + ); + } + + #[test] + #[serial] + fn public_pi_import_rolls_back_portable_state_if_native_revision_changes() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "native-before").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + replace_pi_agents_after_each_native_binding_save_for_test(["native-after"]); + + let error = PromptService::import_from_file(&state, AppType::Pi) + .expect_err("a stale native snapshot must not become active portable state"); + assert!(matches!(error, AppError::Conflict(_))); + assert!( + state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("restored portable library") + .is_empty(), + "failed import must conditionally restore its exact DB before-image" + ); + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("external native edit"), + "native-after" + ); + } + + #[test] + #[serial] + fn first_launch_pi_import_rolls_back_if_native_revision_changes() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "native-before").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + replace_pi_agents_after_each_native_binding_save_for_test(["native-after"]); + + let error = PromptService::import_from_file_on_first_launch(&state, AppType::Pi) + .expect_err("first-launch import must bind its DB row to one native revision"); + assert!(matches!(error, AppError::Conflict(_))); + assert!(state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("restored portable library") + .is_empty()); + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("external native edit"), + "native-after" + ); + } + + #[test] + fn unowned_empty_agents_file_is_not_treated_as_absent() { + let snapshot = PiPromptFileSnapshot { + kind: PiPromptFileKind::GlobalContext, + path: "AGENTS.md".to_string(), + exists: true, + revision: "present-empty".to_string(), + content: String::new(), + }; + assert!(matches!( + ensure_pi_library_projection_matches(&snapshot, None), + Err(AppError::Conflict(_)) + )); + } + + #[test] + #[serial] + fn portable_prompt_reconciliation_preserves_native_bytes_and_rebuilds_active_truth() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + let native_content = "device-local AGENTS"; + std::fs::write(temp.path().join("AGENTS.md"), native_content).expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &Prompt { + id: "portable-active".to_string(), + name: "Portable active".to_string(), + content: "incoming portable content".to_string(), + description: None, + enabled: true, + created_at: Some(1), + updated_at: Some(1), + }, + ) + .expect("seed portable prompt"); + + PromptService::reconcile_pi_portable_import(&state).expect("reconcile"); + + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("read native"), + native_content, + "portable import must not overwrite the device-local native file" + ); + let prompts = state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("read prompts"); + let active = prompts + .values() + .filter(|prompt| prompt.enabled) + .collect::>(); + assert_eq!(active.len(), 1); + assert_eq!(active[0].content, native_content); + assert!(!prompts["portable-active"].enabled); + } + + #[test] + #[serial] + fn portable_prompt_reconciliation_disables_shadow_state_when_agents_is_absent() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &Prompt { + id: "portable-active".to_string(), + name: "Portable active".to_string(), + content: "incoming portable content".to_string(), + description: None, + enabled: true, + created_at: Some(1), + updated_at: Some(1), + }, + ) + .expect("seed portable prompt"); + + PromptService::reconcile_pi_portable_import(&state).expect("reconcile"); + + let prompts = state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("read prompts"); + assert!(prompts.values().all(|prompt| !prompt.enabled)); + assert!(!temp.path().join("AGENTS.md").exists()); + } + + #[test] + #[serial] + fn reading_pi_prompts_does_not_adopt_external_agents_drift() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "managed-before").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &Prompt { + id: "managed".to_string(), + name: "Managed".to_string(), + content: "managed-before".to_string(), + description: None, + enabled: true, + created_at: Some(1), + updated_at: Some(1), + }, + ) + .expect("seed prompt"); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &Prompt { + id: "other".to_string(), + name: "Other".to_string(), + content: "other-content".to_string(), + description: None, + enabled: false, + created_at: Some(2), + updated_at: Some(2), + }, + ) + .expect("seed alternate prompt"); + + std::fs::write(temp.path().join("AGENTS.md"), "external-after").expect("external edit"); + let prompts = PromptService::get_prompts(&state, AppType::Pi).expect("read prompt list"); + assert_eq!(prompts["managed"].content, "managed-before"); + assert!( + prompts.values().all(|prompt| !prompt.enabled), + "read-only inspection must report native truth instead of the stale DB projection" + ); + assert!( + state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("read persisted projection")["managed"] + .enabled, + "inspection must not silently adopt or rewrite the portable library" + ); + assert!( + PromptService::enable_prompt(&state, AppType::Pi, "other").is_err(), + "the write boundary must still report the external drift conflict" + ); + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("live AGENTS.md"), + "external-after" + ); + } + + #[test] + #[serial] + fn explicit_pi_library_reconciliation_adopts_external_native_truth() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "external-after").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &Prompt { + id: "stale".to_string(), + name: "Stale".to_string(), + content: "managed-before".to_string(), + description: None, + enabled: true, + created_at: Some(1), + updated_at: Some(1), + }, + ) + .expect("seed stale projection"); + + let status = PromptService::get_pi_library_status(&state).expect("inspect"); + assert!(status.native_exists); + assert!(status.matched_prompt_id.is_none()); + assert!(status.needs_reconciliation); + + PromptService::reconcile_pi_library(&state).expect("explicit reconcile"); + let prompts = PromptService::get_prompts(&state, AppType::Pi).expect("read reconciled"); + let active = prompts + .values() + .filter(|prompt| prompt.enabled) + .collect::>(); + assert_eq!(active.len(), 1); + assert_eq!(active[0].content, "external-after"); + assert!( + !PromptService::get_pi_library_status(&state) + .expect("reinspect") + .needs_reconciliation + ); + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("native remains"), + "external-after" + ); + } + + #[test] + #[serial] + fn missing_agents_file_is_effectively_inactive_until_explicit_reconciliation() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &Prompt { + id: "shadow".to_string(), + name: "Shadow".to_string(), + content: "not live".to_string(), + description: None, + enabled: true, + created_at: Some(1), + updated_at: Some(1), + }, + ) + .expect("seed shadow"); + + assert!(PromptService::get_prompts(&state, AppType::Pi) + .expect("read effective") + .values() + .all(|prompt| !prompt.enabled)); + let status = PromptService::get_pi_library_status(&state).expect("inspect"); + assert!(!status.native_exists); + assert!(status.needs_reconciliation); + + PromptService::reconcile_pi_library(&state).expect("explicit reconcile"); + assert!(state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("read persisted") + .values() + .all(|prompt| !prompt.enabled)); + assert!(!temp.path().join("AGENTS.md").exists()); + } + + #[test] + #[serial] + fn duplicate_prompt_content_prefers_the_persisted_enabled_identity() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "same-content").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &prompt("first", "same-content", false, 1), + ) + .expect("first duplicate"); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &prompt("persisted-active", "same-content", true, 2), + ) + .expect("active duplicate"); + + let effective = PromptService::get_prompts(&state, AppType::Pi).expect("effective prompts"); + assert!(!effective["first"].enabled); + assert!(effective["persisted-active"].enabled); + let status = PromptService::get_pi_library_status(&state).expect("status"); + assert_eq!( + status.matched_prompt_id.as_deref(), + Some("persisted-active") + ); + assert!(!status.needs_reconciliation); + } + + #[test] + #[serial] + fn external_prompt_drift_blocks_deleting_or_disabling_the_live_match() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "external-live").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &prompt("stale-active", "stale-content", true, 1), + ) + .expect("stale prompt"); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &prompt("live-match", "external-live", false, 2), + ) + .expect("live match"); + let before = serde_json::to_value( + state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("before prompts"), + ) + .expect("serialize before"); + + assert!(matches!( + PromptService::delete_prompt(&state, AppType::Pi, "live-match"), + Err(AppError::Conflict(_)) + )); + assert!(matches!( + PromptService::upsert_prompt( + &state, + AppType::Pi, + "live-match", + prompt("live-match", "external-live", false, 3), + ), + Err(AppError::Conflict(_)) + )); + assert_eq!( + serde_json::to_value( + state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("after prompts") + ) + .expect("serialize after"), + before + ); + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("live remains"), + "external-live" + ); + } + + #[test] + #[serial] + fn repeated_external_reconcile_races_restore_the_original_library() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "native-0").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &prompt("portable", "portable-content", true, 1), + ) + .expect("portable prompt"); + let before = serde_json::to_value( + state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("before prompts"), + ) + .expect("serialize before"); + replace_pi_agents_after_each_native_binding_save_for_test([ + "native-1", "native-2", "native-3", + ]); + + assert!(matches!( + PromptService::reconcile_pi_library(&state), + Err(AppError::Conflict(_)) + )); + assert_eq!( + serde_json::to_value( + state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("restored prompts") + ) + .expect("serialize restored"), + before, + "a failed reconcile must leave no imported rows or selection changes" + ); + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("latest native"), + "native-3" + ); + } + + #[test] + #[serial] + fn reconcile_compensation_preserves_a_concurrent_portable_library() { + let temp = tempfile::tempdir().expect("tempdir"); + let _restore = EnvRestore::set("PI_CODING_AGENT_DIR", temp.path()); + std::fs::write(temp.path().join("AGENTS.md"), "native-0").expect("seed AGENTS.md"); + let state = AppState::new(Arc::new(Database::memory().expect("database"))); + state + .db + .save_prompt( + AppType::Pi.as_str(), + &prompt("before", "portable-before", true, 1), + ) + .expect("seed original library"); + + let imported = IndexMap::from([( + "restored-import".to_string(), + prompt("restored-import", "portable-restored", true, 2), + )]); + replace_pi_library_after_next_reconcile_save_for_test(imported.clone()); + replace_pi_agents_after_each_native_binding_save_for_test(["native-1"]); + + let error = PromptService::reconcile_pi_library(&state) + .expect_err("stale compensation must not overwrite a concurrent import"); + assert!( + error.to_string().contains("newer database revision"), + "the conflict must explain why compensation stopped: {error}" + ); + assert_eq!( + state + .db + .get_prompts(AppType::Pi.as_str()) + .expect("preserved imported library"), + imported + ); + assert_eq!( + std::fs::read_to_string(temp.path().join("AGENTS.md")).expect("latest native"), + "native-1" + ); + } + + #[test] + fn prompt_selection_compensation_can_restore_an_exact_legacy_before_image() { + let db = Database::memory().expect("database"); + db.save_prompt(AppType::Pi.as_str(), &prompt("legacy-a", "a", true, 1)) + .expect("first legacy row"); + db.save_prompt(AppType::Pi.as_str(), &prompt("legacy-b", "b", true, 2)) + .expect("second legacy row"); + let before = db + .get_prompts(AppType::Pi.as_str()) + .expect("legacy before-image"); + let mut attempted = before.clone(); + attempted.get_mut("legacy-a").expect("first row").enabled = false; + + db.compare_exchange_prompt_selection(AppType::Pi.as_str(), &before, &attempted) + .expect("publish valid projection"); + db.restore_prompt_selection_if_attempted(AppType::Pi.as_str(), &attempted, &before) + .expect("restore exact legacy image"); + + assert_eq!( + db.get_prompts(AppType::Pi.as_str()) + .expect("restored library"), + before + ); + } } diff --git a/src-tauri/src/services/provider/live.rs b/src-tauri/src/services/provider/live.rs index 555d48a21..67dc866b5 100644 --- a/src-tauri/src/services/provider/live.rs +++ b/src-tauri/src/services/provider/live.rs @@ -530,6 +530,7 @@ fn settings_contain_common_config(app_type: &AppType, settings: &Value, snippet: | AppType::OpenCode | AppType::OpenClaw | AppType::Hermes + | AppType::Pi | AppType::ClaudeDesktop => false, } } @@ -604,6 +605,7 @@ pub(crate) fn remove_common_config_from_settings( | AppType::OpenCode | AppType::OpenClaw | AppType::Hermes + | AppType::Pi | AppType::ClaudeDesktop => Ok(settings.clone()), } } @@ -663,6 +665,7 @@ fn apply_common_config_to_settings( | AppType::OpenCode | AppType::OpenClaw | AppType::Hermes + | AppType::Pi | AppType::ClaudeDesktop => Ok(settings.clone()), } } @@ -1165,6 +1168,13 @@ pub(crate) fn write_live_snapshot(app_type: &AppType, provider: &Provider) -> Re crate::hermes_config::set_provider(&provider.id, provider.settings_config.clone())?; log::debug!("Hermes provider '{}' written to live config", provider.id); } + AppType::Pi => { + return Err(AppError::localized( + "pi.live.requires_catalog_coordinator", + "Pi 的共享 models.json 必须通过 Pi 目录协调器写入", + "Pi's shared models.json must be written through the Pi catalog coordinator", + )); + } } Ok(()) } @@ -1281,6 +1291,10 @@ fn sync_current_provider_for_app_respecting_takeover( pub fn sync_current_to_live(state: &AppState) -> Result<(), AppError> { // Sync providers based on mode for app_type in AppType::all() { + if matches!(app_type, AppType::Pi) { + crate::services::pi_catalog::PiCatalogCoordinator::reconcile_portable_import(state)?; + continue; + } if app_type.is_additive_mode() { // Provider rename and every additive live mutation share this // per-app lock. Acquire it before reading the catalog so a key @@ -1426,6 +1440,11 @@ pub fn read_live_settings(app_type: AppType) -> Result { let config = crate::hermes_config::yaml_to_json(&yaml_config)?; Ok(config) } + AppType::Pi => Err(AppError::localized( + "pi.live.requires_catalog_inspection", + "Pi 的共享 models.json 必须通过 Pi 原生目录检查服务读取", + "Pi's shared models.json must be read through the Pi native catalog inspection service", + )), } } @@ -1534,6 +1553,13 @@ pub fn import_default_config(state: &AppState, app_type: AppType) -> Result { + return Err(AppError::localized( + "pi.import.requires_catalog_coordinator", + "Pi 原生供应商必须通过 Pi 目录导入流程导入", + "Native Pi providers must be imported through the Pi catalog import flow", + )); + } // OpenCode, OpenClaw and Hermes use additive mode and are handled by early return above AppType::OpenCode | AppType::OpenClaw | AppType::Hermes => { unreachable!("additive mode apps are handled by early return") diff --git a/src-tauri/src/services/provider/mod.rs b/src-tauri/src/services/provider/mod.rs index 2d5d55df8..9e285dfc0 100644 --- a/src-tauri/src/services/provider/mod.rs +++ b/src-tauri/src/services/provider/mod.rs @@ -20,6 +20,7 @@ use crate::database::{ use crate::error::AppError; use crate::provider::{Provider, ProviderMutationInput, UsageResult}; use crate::services::mcp::McpService; +use crate::services::pi_catalog::{PiCatalogCoordinator, PiCatalogMutation}; use crate::settings::CustomEndpoint; use crate::store::AppState; @@ -478,6 +479,32 @@ mod tests { } } + fn pi_provider(id: &str) -> Provider { + Provider { + id: id.to_string(), + name: format!("Pi Provider {id}"), + settings_config: json!({ + "name": format!("Pi Provider {id}"), + "api": "openai-responses", + "baseUrl": "https://pi.example/v1", + "apiKey": "test-key", + "models": [ + {"id": "model-a", "name": "Model A"}, + {"id": "model-b", "name": "Model B"} + ] + }), + website_url: None, + category: Some("custom".to_string()), + created_at: Some(1), + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + } + } + fn opencode_provider(id: &str) -> Provider { Provider { id: id.to_string(), @@ -632,6 +659,174 @@ mod tests { }); } + #[test] + #[serial] + fn pi_provider_service_create_hydrates_all_endpoints_and_publishes_native_default() { + with_test_home(|state, home| { + let original_settings = crate::settings::get_settings(); + let mut isolated_settings = original_settings.clone(); + isolated_settings.pi_config_dir = Some( + home.join(".pi") + .join("agent") + .to_string_lossy() + .into_owned(), + ); + isolated_settings.current_provider_pi = None; + crate::settings::update_settings(isolated_settings) + .expect("install isolated Pi settings"); + + let outcome = (|| -> Result<_, AppError> { + let expected_endpoints = HashMap::from([ + ( + "https://one.pi.example".to_string(), + endpoint("https://one.pi.example", None, Some(11)), + ), + ( + "https://two.pi.example".to_string(), + endpoint("https://two.pi.example", Some(20), None), + ), + ]); + let mut provider = pi_provider("managed-pi"); + provider.meta = Some(ProviderMeta { + custom_endpoints: expected_endpoints.clone(), + ..Default::default() + }); + ProviderService::add( + state, + AppType::Pi, + provider_to_mutation_input(provider), + false, + )?; + + let aggregate = state + .db + .get_provider_aggregate("pi", "managed-pi")? + .ok_or_else(|| AppError::NotFound("managed Pi aggregate".to_string()))?; + let hydrated = aggregate.endpoints.into_iter().collect::>(); + let models: Value = serde_json::from_slice( + &fs::read(home.join(".pi/agent/models.json")) + .map_err(|error| AppError::io(home, error))?, + ) + .map_err(|error| AppError::json(home, error))?; + let defaults = crate::pi_config::native_settings::read_pi_native_defaults()?; + Ok((expected_endpoints, hydrated, models, defaults)) + })(); + + crate::settings::update_settings(original_settings).expect("restore process settings"); + let (expected, hydrated, models, defaults) = outcome.expect("Pi service create"); + assert_eq!(hydrated, expected); + assert_eq!( + models.pointer("/providers/managed-pi/models/0/id"), + Some(&json!("model-a")) + ); + assert_eq!(defaults.default_provider.as_deref(), Some("managed-pi")); + assert_eq!(defaults.default_model.as_deref(), Some("model-a")); + }); + } + + #[test] + #[serial] + fn pi_provider_service_update_rehomes_a_removed_active_model() { + with_test_home(|state, home| { + let original_settings = crate::settings::get_settings(); + let mut isolated_settings = original_settings.clone(); + isolated_settings.pi_config_dir = Some( + home.join(".pi") + .join("agent") + .to_string_lossy() + .into_owned(), + ); + isolated_settings.current_provider_pi = None; + crate::settings::update_settings(isolated_settings) + .expect("install isolated Pi settings"); + + let outcome = (|| -> Result<_, AppError> { + let provider = pi_provider("managed-pi-update"); + ProviderService::add( + state, + AppType::Pi, + provider_to_mutation_input(provider.clone()), + false, + )?; + + let mut updated = provider; + updated.settings_config["models"] = json!([{"id": "model-b", "name": "Model B"}]); + ProviderService::update( + state, + AppType::Pi, + Some("managed-pi-update"), + provider_to_mutation_input(updated), + )?; + + let defaults = crate::pi_config::native_settings::read_pi_native_defaults()?; + let models: Value = serde_json::from_slice( + &fs::read(home.join(".pi/agent/models.json")) + .map_err(|error| AppError::io(home, error))?, + ) + .map_err(|error| AppError::json(home, error))?; + Ok((defaults, models)) + })(); + + crate::settings::update_settings(original_settings).expect("restore process settings"); + let (defaults, models) = outcome.expect("Pi service update"); + assert_eq!( + defaults.default_provider.as_deref(), + Some("managed-pi-update") + ); + assert_eq!(defaults.default_model.as_deref(), Some("model-b")); + assert_eq!( + models.pointer("/providers/managed-pi-update/models"), + Some(&json!([{"id": "model-b", "name": "Model B"}])) + ); + }); + } + + #[test] + #[serial] + fn pi_provider_service_current_follows_the_native_default_after_external_edit() { + with_test_home(|state, home| { + let original_settings = crate::settings::get_settings(); + let mut isolated_settings = original_settings.clone(); + isolated_settings.pi_config_dir = Some( + home.join(".pi") + .join("agent") + .to_string_lossy() + .into_owned(), + ); + isolated_settings.current_provider_pi = None; + crate::settings::update_settings(isolated_settings) + .expect("install isolated Pi settings"); + + let outcome = (|| -> Result<_, AppError> { + for provider_id in ["native-first", "native-second"] { + ProviderService::add( + state, + AppType::Pi, + provider_to_mutation_input(pi_provider(provider_id)), + false, + )?; + } + assert_eq!( + state.db.get_current_provider("pi")?.as_deref(), + Some("native-first") + ); + crate::pi_config::native_settings::set_pi_native_default_with_receipt( + "native-second", + "model-b", + )?; + + ProviderService::current(state, AppType::Pi) + })(); + + crate::settings::update_settings(original_settings).expect("restore process settings"); + assert_eq!( + outcome.expect("resolve current Pi provider"), + "native-second", + "the UI current marker must follow Pi's live settings, not a stale DB marker" + ); + }); + } + #[test] #[serial] fn provider_service_create_canonicalizes_initial_endpoint_identity() { @@ -2567,8 +2762,13 @@ requires_openai_auth = true ..Default::default() }); - ProviderService::update(&state, AppType::ClaudeDesktop, None, updated.clone()) - .expect("update current provider"); + ProviderService::update( + &state, + AppType::ClaudeDesktop, + None, + provider_to_mutation_input(updated.clone()), + ) + .expect("update current provider"); let backup = db .get_live_backup("claude-desktop") @@ -3404,6 +3604,10 @@ impl ProviderService { if app_type.is_additive_mode() { return Ok(String::new()); } + if matches!(app_type, AppType::Pi) { + return PiCatalogCoordinator::current_native_provider(state) + .map(|provider| provider.unwrap_or_default()); + } crate::settings::get_effective_current_provider(&state.db, &app_type) .map(|opt| opt.unwrap_or_default()) } @@ -3415,6 +3619,18 @@ impl ProviderService { input: ProviderMutationInput, add_to_live: bool, ) -> Result { + if matches!(app_type, AppType::Pi) { + let provider_key = input.id.clone(); + PiCatalogCoordinator::apply( + state, + PiCatalogMutation::CreateProvider { + input, + provider_key, + activate_if_first: true, + }, + )?; + return Ok(true); + } let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); let mut provider: Provider = input.into(); // Normalize Claude model keys @@ -3473,6 +3689,16 @@ impl ProviderService { // Reject endpoint-bearing edit payloads before any live or DB side // effect. Endpoints have their own typed mutation API. ProviderRowUpdate::from_input(&input)?; + if matches!(app_type, AppType::Pi) { + let original_id = original_id.unwrap_or(input.id.as_str()); + if original_id != input.id { + return Err(AppError::InvalidInput( + "Pi provider identity and native projection key cannot be renamed".to_string(), + )); + } + PiCatalogCoordinator::apply(state, PiCatalogMutation::UpdateProvider { input })?; + return Ok(true); + } let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); let mut provider: Provider = input.into(); let original_id = original_id.unwrap_or(provider.id.as_str()).to_string(); @@ -3713,6 +3939,15 @@ impl ProviderService { /// 同时检查本地 settings 和数据库的当前供应商,防止删除任一端正在使用的供应商。 /// 对于累加模式应用(OpenCode, OpenClaw),可以随时删除任意供应商,同时从 live 配置中移除。 pub fn delete(state: &AppState, app_type: AppType, id: &str) -> Result<(), AppError> { + if matches!(app_type, AppType::Pi) { + PiCatalogCoordinator::apply( + state, + PiCatalogMutation::DeleteProvider { + provider_id: id.to_string(), + }, + )?; + return Ok(()); + } let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); // Additive mode apps - no current provider concept if app_type.is_additive_mode() { @@ -3853,6 +4088,32 @@ impl ProviderService { /// d. Write target provider config to live files /// e. Sync MCP configuration pub fn switch(state: &AppState, app_type: AppType, id: &str) -> Result { + if matches!(app_type, AppType::Pi) { + let provider = state + .db + .get_provider_aggregate(AppType::Pi.as_str(), id)? + .ok_or_else(|| AppError::NotFound(format!("Pi provider '{id}'")))?; + let config: crate::pi_config::model::PiManagedProviderConfig = + serde_json::from_value(provider.provider.settings_config).map_err(|error| { + AppError::Config(format!("managed Pi provider '{id}' is invalid: {error}")) + })?; + let model_id = config + .models + .first() + .ok_or_else(|| { + AppError::InvalidInput(format!("Pi provider '{id}' has no selectable models")) + })? + .id + .clone(); + PiCatalogCoordinator::apply( + state, + PiCatalogMutation::SetDefault { + provider_id: id.to_string(), + model_id, + }, + )?; + return Ok(SwitchResult::default()); + } // The same per-app lock also guards additive provider key changes and // bulk live sync. Acquire it before observing the provider map so a // queued rename cannot leave this switch holding a stale source key. @@ -4393,6 +4654,7 @@ impl ProviderService { AppType::OpenCode => Self::extract_opencode_common_config(&provider.settings_config), AppType::OpenClaw => Self::extract_openclaw_common_config(&provider.settings_config), AppType::Hermes => Ok(String::new()), // Hermes doesn't use common config snippets + AppType::Pi => Ok(String::new()), // Pi owns a shared exact-key catalog, not snippets } } @@ -4410,6 +4672,7 @@ impl ProviderService { AppType::OpenCode => Self::extract_opencode_common_config(settings_config), AppType::OpenClaw => Self::extract_openclaw_common_config(settings_config), AppType::Hermes => Ok(String::new()), // Hermes doesn't use common config snippets + AppType::Pi => Ok(String::new()), } } @@ -4979,6 +5242,16 @@ impl ProviderService { provider_id: &str, url: String, ) -> Result<(), AppError> { + if matches!(app_type, AppType::Pi) { + PiCatalogCoordinator::apply( + state, + PiCatalogMutation::AddEndpoint { + provider_id: provider_id.to_string(), + url, + }, + )?; + return Ok(()); + } endpoints::add_custom_endpoint(state, app_type, provider_id, url) } @@ -4989,6 +5262,16 @@ impl ProviderService { provider_id: &str, url: String, ) -> Result<(), AppError> { + if matches!(app_type, AppType::Pi) { + PiCatalogCoordinator::apply( + state, + PiCatalogMutation::RemoveEndpoint { + provider_id: provider_id.to_string(), + url, + }, + )?; + return Ok(()); + } endpoints::remove_custom_endpoint(state, app_type, provider_id, url) } @@ -4999,6 +5282,13 @@ impl ProviderService { provider_id: &str, url: String, ) -> Result<(), AppError> { + let _pi_switch_guard = matches!(app_type, AppType::Pi).then(|| { + futures::executor::block_on( + state + .proxy_service + .lock_switch_for_app(AppType::Pi.as_str()), + ) + }); endpoints::update_endpoint_last_used(state, app_type, provider_id, url) } @@ -5008,12 +5298,18 @@ impl ProviderService { app_type: AppType, updates: Vec, ) -> Result { - for update in updates { - let key = ProviderKey::new(app_type.as_str(), update.id)?; - state - .db - .update_provider_sort_index(&key, update.sort_index)?; + // Validate the whole payload before entering the app-specific + // ordering boundary. + let updates = updates + .into_iter() + .map(|update| { + ProviderKey::new(app_type.as_str(), update.id).map(|key| (key, update.sort_index)) + }) + .collect::, _>>()?; + if matches!(app_type, AppType::Pi) { + return PiCatalogCoordinator::update_route_order(state, updates); } + state.db.update_provider_sort_index(&updates)?; Ok(true) } @@ -5176,6 +5472,25 @@ impl ProviderService { )); } } + AppType::Pi => { + let config: crate::pi_config::model::PiManagedProviderConfig = + serde_json::from_value(provider.settings_config.clone()).map_err(|error| { + AppError::localized( + "provider.pi.settings.invalid", + format!("Pi 配置无法解析: {error}"), + format!("Pi configuration cannot be decoded: {error}"), + ) + })?; + crate::pi_config::model::validate_pi_managed_provider(&config).map_err( + |error| { + AppError::localized( + "provider.pi.settings.invalid", + format!("Pi 配置无效: {error}"), + format!("Invalid Pi configuration: {error}"), + ) + }, + )?; + } } // Validate and clean UsageScript configuration (common for all app types) @@ -5404,6 +5719,18 @@ impl ProviderService { Ok((api_key, base_url)) } + AppType::Pi => { + let config: crate::pi_config::model::PiManagedProviderConfig = + serde_json::from_value(provider.settings_config.clone()).map_err(|error| { + AppError::Config(format!("invalid Pi provider configuration: {error}")) + })?; + let model = config.models.first().ok_or_else(|| { + AppError::Config("Pi provider has no configured model".to_string()) + })?; + let effective = crate::pi_config::model::effective_pi_model(&config, &model.id) + .map_err(|error| AppError::Config(format!("invalid Pi model: {error}")))?; + Ok((effective.api_key.unwrap_or_default(), effective.base_url)) + } } } } diff --git a/src-tauri/src/services/proxy.rs b/src-tauri/src/services/proxy.rs index 9dcaaa392..2a80ea5b5 100644 --- a/src-tauri/src/services/proxy.rs +++ b/src-tauri/src/services/proxy.rs @@ -5,7 +5,11 @@ use crate::app_config::AppType; use crate::config::{get_claude_settings_path, read_json_file, write_json_file}; use crate::database::Database; +use crate::error::AppError; use crate::provider::Provider; +use crate::proxy::pi_runtime::{ + build_pi_runtime, direct_pi_projection_patch, project_managed_pi_config, PiRuntimeStore, +}; use crate::proxy::server::ProxyServer; use crate::proxy::switch_lock::SwitchLockManager; use crate::proxy::types::*; @@ -14,12 +18,19 @@ use crate::services::provider::{ provider_to_mutation_input, reconcile_provider_record_with_precondition, write_live_with_common_config, ReconcilePrecondition, }; +use crate::services::{pi_prompt_files::lock_instruction_files_async, prompt::PromptService}; use serde_json::{json, Map, Value}; use std::str::FromStr; -use std::sync::Arc; +use std::sync::{ + atomic::{AtomicU64, Ordering}, + Arc, RwLock as StdRwLock, +}; use tauri::Emitter; use tokio::sync::RwLock; +#[cfg(test)] +use std::sync::atomic::AtomicBool; + /// 用于接管 Live 配置时的占位符(避免客户端提示缺少 key,同时不泄露真实 Token) const PROXY_TOKEN_PLACEHOLDER: &str = "PROXY_MANAGED"; @@ -63,6 +74,40 @@ pub struct ProxyService { /// AppHandle,用于传递给 ProxyServer 以支持故障转移时的 UI 更新 app_handle: Arc>>, switch_locks: SwitchLockManager, + pi_runtime: Arc, + pi_server_sequence: Arc, + pi_listener: Arc>>, + #[cfg(test)] + fail_next_pi_reconcile: Arc, +} + +#[derive(Debug, Clone)] +struct PiListenerIdentity { + server_generation: u64, + gateway_origin: url::Url, +} + +fn pi_loopback_origin(address: &str, port: u16) -> Result { + let host = if address.eq_ignore_ascii_case("localhost") { + "127.0.0.1".to_string() + } else { + let ip = address.parse::().map_err(|_| { + AppError::Config(format!( + "Pi gateway requires an explicit loopback listener, got '{address}'" + )) + })?; + if !ip.is_loopback() { + return Err(AppError::Config(format!( + "Pi gateway refuses non-loopback listener '{address}'" + ))); + } + match ip { + std::net::IpAddr::V4(value) => value.to_string(), + std::net::IpAddr::V6(value) => format!("[{value}]"), + } + }; + url::Url::parse(&format!("http://{host}:{port}/")) + .map_err(|error| AppError::Config(format!("invalid Pi listener origin: {error}"))) } #[derive(Debug, Clone, Copy, Default)] @@ -77,6 +122,11 @@ impl ProxyService { server: Arc::new(RwLock::new(None)), app_handle: Arc::new(RwLock::new(None)), switch_locks: SwitchLockManager::new(), + pi_runtime: Arc::new(PiRuntimeStore::default()), + pi_server_sequence: Arc::new(AtomicU64::new(0)), + pi_listener: Arc::new(StdRwLock::new(None)), + #[cfg(test)] + fail_next_pi_reconcile: Arc::new(AtomicBool::new(false)), } } @@ -526,8 +576,736 @@ impl ProxyService { self.switch_locks.lock_for_app(app_type).await } + pub(crate) async fn begin_pi_catalog_mutation(&self) -> u64 { + self.pi_runtime.begin_mutation().await + } + + pub(crate) async fn close_pi_runtime_at_epoch( + &self, + catalog_epoch: u64, + ) -> Result<(), AppError> { + self.pi_runtime.close(catalog_epoch).await + } + + pub(crate) fn project_pi_provider_value( + &self, + provider_id: &str, + provider_key: &str, + config: &crate::pi_config::model::PiManagedProviderConfig, + ) -> Result { + if !crate::settings::pi_takeover_enabled() { + return serde_json::to_value(config) + .map_err(|source| AppError::JsonSerialize { source }); + } + let listener = self + .pi_listener + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + .ok_or_else(|| { + AppError::Conflict( + "Pi gateway projection is desired but no loopback listener is active" + .to_string(), + ) + })?; + let token = crate::settings::get_or_create_pi_gateway_token()?; + project_managed_pi_config( + provider_id, + provider_key, + config, + &listener.gateway_origin, + &token, + ) + } + + pub(crate) async fn reconcile_pi_runtime_at_epoch( + &self, + catalog_epoch: u64, + ) -> Result, AppError> { + self.reconcile_pi_runtime_at_epoch_with_native_precondition(catalog_epoch, None) + .await + } + + pub(crate) async fn reconcile_pi_runtime_at_epoch_with_native_precondition( + &self, + catalog_epoch: u64, + expected_native: Option<&crate::pi_config::document::PiProviderValuesSnapshot>, + ) -> Result, AppError> { + self.reconcile_pi_runtime_at_epoch_with_native_claim_precondition( + catalog_epoch, + expected_native, + None, + ) + .await + } + + pub(crate) async fn reconcile_pi_runtime_at_epoch_with_native_claim_precondition( + &self, + catalog_epoch: u64, + expected_native: Option<&crate::pi_config::document::PiProviderValuesSnapshot>, + expected_fingerprints: Option<&indexmap::IndexMap>, + ) -> Result, AppError> { + #[cfg(test)] + if self.fail_next_pi_reconcile.swap(false, Ordering::AcqRel) { + return Err(AppError::Config( + "injected Pi runtime reconciliation failure".to_string(), + )); + } + if !crate::settings::pi_takeover_enabled() { + self.pi_runtime.close(catalog_epoch).await?; + if let Some(expected_native) = expected_native { + let models_path = crate::pi_config::native::get_pi_models_path()?; + match expected_fingerprints { + Some(expected_fingerprints) => { + crate::pi_config::document::verify_pi_provider_preconditions( + &models_path, + &expected_native.values, + expected_fingerprints, + )?; + } + None => crate::pi_config::document::verify_pi_provider_values( + &models_path, + &expected_native.values, + )?, + } + } + return Ok(Vec::new()); + } + let listener = self + .pi_listener + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + .ok_or_else(|| { + AppError::Conflict( + "Pi takeover is desired but the loopback listener is unavailable".to_string(), + ) + })?; + let token = crate::settings::get_or_create_pi_gateway_token()?; + let app_config = crate::settings::get_pi_app_proxy_config(); + let build = build_pi_runtime( + self.db.as_ref(), + listener.server_generation, + catalog_epoch, + &listener.gateway_origin, + token, + app_config, + )?; + let models_path = crate::pi_config::native::get_pi_models_path()?; + let native_receipt = if build.projection_patch.is_empty() { + if let Some(expected_native) = expected_native { + match expected_fingerprints { + Some(expected_fingerprints) => { + crate::pi_config::document::verify_pi_provider_preconditions( + &models_path, + &expected_native.values, + expected_fingerprints, + )?; + } + None => crate::pi_config::document::verify_pi_provider_values( + &models_path, + &expected_native.values, + )?, + } + } + None + } else { + let expected = match expected_native { + Some(expected) => { + crate::pi_config::document::PiProviderValuesSnapshot { + file_existed: expected.file_existed, + values: build + .projection_patch + .keys() + .map(|provider_key| { + expected + .values + .get(provider_key) + .cloned() + .map(|value| (provider_key.clone(), value)) + .ok_or_else(|| { + AppError::Config(format!( + "Pi runtime provider key '{provider_key}' was not included in the catalog preflight" + )) + }) + }) + .collect::>()?, + } + } + None => self.preflight_pi_owned_projection_at( + &models_path, + &build.projection_patch, + )?, + }; + Some( + crate::pi_config::document::apply_pi_provider_patch_with_receipt_and_fingerprints( + &models_path, + &expected, + expected_fingerprints, + &build.projection_patch, + )?, + ) + }; + if let Err(error) = self.pi_runtime.publish(build.snapshot).await { + let rollback = native_receipt.as_ref().map_or( + Ok(()), + crate::pi_config::document::PiProviderPatchReceipt::rollback, + ); + return Err(AppError::Config(format!( + "failed to publish Pi runtime after native projection: {error}; native rollback={rollback:?}" + ))); + } + Ok(build.direct_only_provider_ids) + } + + /// Re-open a fenced runtime only while Pi's live exact-key projection + /// still matches the immutable snapshot which produced that runtime. + /// + /// Rebuilding this witness from the mutable database or settings would + /// turn a failed compensation into stale admission, so the projection is + /// retained inside `PiRuntimeSnapshot`. + async fn republish_current_pi_runtime_if_native_matches( + &self, + catalog_epoch: u64, + ) -> Result { + let Some(expected_projection) = self.pi_runtime.retained_native_projection() else { + self.pi_runtime.close(catalog_epoch).await?; + return Ok(false); + }; + let verification = crate::pi_config::native::get_pi_models_path().and_then(|models_path| { + crate::pi_config::document::verify_pi_provider_values( + &models_path, + &expected_projection, + ) + }); + if let Err(error) = verification { + self.pi_runtime.close(catalog_epoch).await?; + return Err(AppError::Conflict(format!( + "Pi native projection changed while runtime admission was fenced; admission \ + remains closed: {error}" + ))); + } + self.pi_runtime.republish_current(catalog_epoch).await + } + + #[cfg(test)] + pub(crate) fn fail_next_pi_reconcile_for_test(&self) { + self.fail_next_pi_reconcile.store(true, Ordering::Release); + } + + pub(crate) async fn reconcile_pi_runtime(&self) -> Result, AppError> { + let epoch = self.pi_runtime.begin_mutation().await; + self.reconcile_pi_runtime_at_epoch(epoch).await + } + + /// Publish a DB-only catalog ordering without closing admission or + /// touching the native projection. The caller holds Pi's switch lock and + /// must restore the DB order if this preparation fails. + pub(crate) async fn publish_pi_runtime_order(&self) -> Result<(), AppError> { + if !crate::settings::pi_takeover_enabled() { + return Ok(()); + } + let listener = self + .pi_listener + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + .ok_or_else(|| { + AppError::Conflict( + "Pi takeover is desired but the loopback listener is unavailable".to_string(), + ) + })?; + let epoch = self.pi_runtime.next_even_epoch()?; + let build = build_pi_runtime( + self.db.as_ref(), + listener.server_generation, + epoch, + &listener.gateway_origin, + crate::settings::get_pi_gateway_token()?, + crate::settings::get_pi_app_proxy_config(), + )?; + self.pi_runtime.publish(build.snapshot).await + } + + fn restore_pi_direct_projection_at( + &self, + models_path: &std::path::Path, + ) -> Result { + if self + .pi_listener + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .is_none() + { + return self.confirm_pi_projection_is_already_direct_at(models_path); + } + let gateway = self.current_pi_gateway_projection()?; + self.restore_pi_direct_projection_at_with_expected_gateway(models_path, &gateway) + } + + fn confirm_pi_projection_is_already_direct_at( + &self, + models_path: &std::path::Path, + ) -> Result { + let direct = direct_pi_projection_patch(self.db.as_ref())?; + let before = crate::pi_config::document::snapshot_pi_provider_values( + models_path, + direct.keys().cloned(), + )?; + if let Some(provider_key) = direct.iter().find_map(|(provider_key, expected)| { + (before.values.get(provider_key) != Some(expected)).then_some(provider_key) + }) { + return Err(AppError::Conflict(format!( + "Pi listener is unavailable and provider key '{provider_key}' is not already in its direct projection" + ))); + } + crate::pi_config::document::apply_pi_provider_patch_with_receipt( + models_path, + &before, + &direct, + ) + } + + fn current_pi_gateway_projection( + &self, + ) -> Result>, AppError> { + let listener = self + .pi_listener + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + .ok_or_else(|| { + AppError::Conflict( + "cannot restore Pi direct projection without its active listener identity" + .to_string(), + ) + })?; + let gateway_token = crate::settings::get_pi_gateway_token().or_else(|settings_error| { + self.pi_runtime + .retained_gateway_token(listener.server_generation) + .ok_or(settings_error) + })?; + let gateway = build_pi_runtime( + self.db.as_ref(), + listener.server_generation, + 0, + &listener.gateway_origin, + gateway_token, + crate::settings::get_pi_app_proxy_config(), + )? + .projection_patch; + Ok(gateway) + } + + fn preflight_pi_owned_projection_at( + &self, + models_path: &std::path::Path, + expected_gateway: &indexmap::IndexMap>, + ) -> Result { + let direct = direct_pi_projection_patch(self.db.as_ref())?; + let direct_keys = direct.keys().collect::>(); + let gateway_keys = expected_gateway + .keys() + .collect::>(); + if direct_keys != gateway_keys { + return Err(AppError::Conflict( + "Pi direct and gateway projections cover different exact-key ownership".to_string(), + )); + } + let before = crate::pi_config::document::snapshot_pi_provider_values( + models_path, + direct.keys().cloned(), + )?; + for (provider_key, direct_value) in &direct { + let observed = before + .values + .get(provider_key) + .expect("every owned key was preflighted"); + let gateway_value = expected_gateway + .get(provider_key) + .expect("gateway ownership keys were checked"); + if observed != direct_value && observed != gateway_value { + return Err(AppError::Conflict(format!( + "Pi provider key '{provider_key}' changed outside CC Switch" + ))); + } + } + Ok(before) + } + + fn restore_pi_direct_projection_at_with_expected_gateway( + &self, + models_path: &std::path::Path, + expected_gateway: &indexmap::IndexMap>, + ) -> Result { + let direct = direct_pi_projection_patch(self.db.as_ref())?; + let before = self.preflight_pi_owned_projection_at(models_path, expected_gateway)?; + crate::pi_config::document::apply_pi_provider_patch_with_receipt( + models_path, + &before, + &direct, + ) + } + + fn restore_pi_direct_projection( + &self, + ) -> Result { + let models_path = crate::pi_config::native::get_pi_models_path()?; + self.restore_pi_direct_projection_at(&models_path) + } + + /// Replace settings while preserving the ownership boundary when the Pi + /// native directory changes. The caller must hold Pi's switch lock from + /// the settings snapshot through this call. + pub(crate) async fn replace_settings_with_pi_directory_boundary_under_lock( + &self, + _guard: &tokio::sync::OwnedMutexGuard<()>, + existing: &crate::settings::AppSettings, + next: crate::settings::AppSettings, + ) -> Result<(), AppError> { + let old_models_path = crate::pi_config::native::get_pi_models_path_for_override( + existing.pi_config_dir.as_deref(), + )?; + let new_models_path = crate::pi_config::native::get_pi_models_path_for_override( + next.pi_config_dir.as_deref(), + )?; + + if old_models_path == new_models_path { + return crate::settings::update_settings(next); + } + + let direct_patch = direct_pi_projection_patch(self.db.as_ref())?; + let new_native_before = crate::pi_config::document::snapshot_pi_provider_values( + &new_models_path, + direct_patch.keys().cloned(), + )?; + for (provider_key, expected) in &direct_patch { + if let Some(existing_value) = new_native_before + .values + .get(provider_key) + .and_then(Option::as_ref) + { + if Some(existing_value) != expected.as_ref() { + return Err(AppError::Conflict(format!( + "Pi directory change would overwrite unowned provider key '{provider_key}' in {}", + new_models_path.display() + ))); + } + } + } + + // Prompt operations and directory ownership share one sendable mutex. + // Holding it across runtime publication prevents a prompt write from + // committing against the new root while a failed directory move is + // rolling settings and the DB selection back to the old root. + let prompt_guard = lock_instruction_files_async().await; + let previous_prompts = self.db.get_prompts(AppType::Pi.as_str())?; + + if !existing.pi_takeover_enabled { + crate::settings::update_settings(next)?; + let native_receipt = + match crate::pi_config::document::apply_pi_provider_patch_with_receipt( + &new_models_path, + &new_native_before, + &direct_patch, + ) { + Ok(receipt) => receipt, + Err(error) => { + let settings_restored = + crate::settings::update_settings(existing.clone()).is_ok(); + return Err(AppError::Config(format!( + "failed to publish managed Pi providers in the new directory: {error}; rollback: settings={settings_restored}" + ))); + } + }; + if let Err(error) = + PromptService::reconcile_pi_native_under_guard(self.db.as_ref(), &prompt_guard) + { + let settings_restored = crate::settings::update_settings(existing.clone()).is_ok(); + let prompts_restored = self + .db + .save_prompt_selection(AppType::Pi.as_str(), &previous_prompts) + .is_ok(); + let new_native_restored = native_receipt.rollback().is_ok(); + return Err(AppError::Config(format!( + "failed to reconcile Pi prompts in the new directory: {error}; rollback: settings={settings_restored}, prompts={prompts_restored}, native={new_native_restored}" + ))); + } + if let Err(error) = + crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all(&self.db) + { + let settings_restored = crate::settings::update_settings(existing.clone()).is_ok(); + let prompts_restored = self + .db + .save_prompt_selection(AppType::Pi.as_str(), &previous_prompts) + .is_ok(); + let new_native_restored = native_receipt.rollback().is_ok(); + let skills_restored = settings_restored + && crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all( + &self.db, + ) + .is_ok(); + return Err(AppError::Config(format!( + "failed to reconcile Pi Skills in the new directory: {error}; rollback: settings={settings_restored}, prompts={prompts_restored}, native={new_native_restored}, skills={skills_restored}" + ))); + } + return Ok(()); + } + + // Admission closes before the old native file is restored. A Pi + // process that already loaded the old gateway projection therefore + // cannot enter a catalog whose directory ownership is in flight. + let epoch = self.pi_runtime.begin_mutation().await; + let old_direct_receipt = match self.restore_pi_direct_projection_at(&old_models_path) { + Ok(receipt) => receipt, + Err(error) => { + let _ = self + .republish_current_pi_runtime_if_native_matches(epoch) + .await; + return Err(error); + } + }; + + if let Err(error) = crate::settings::update_settings(next) { + let expected_old_direct = old_direct_receipt.attempted_snapshot(); + let runtime_restored = self + .reconcile_pi_runtime_at_epoch_with_native_precondition( + epoch, + Some(&expected_old_direct), + ) + .await; + return Err(AppError::Config(if runtime_restored.is_ok() { + format!( + "failed to save the new Pi directory; the previous gateway projection was restored: {error}" + ) + } else { + format!( + "failed to save the new Pi directory and restore the previous gateway projection: {error}" + ) + })); + } + + if let Err(error) = + PromptService::reconcile_pi_native_under_guard(self.db.as_ref(), &prompt_guard) + { + let settings_restored = crate::settings::update_settings(existing.clone()).is_ok(); + let prompts_restored = self + .db + .save_prompt_selection(AppType::Pi.as_str(), &previous_prompts) + .is_ok(); + let skills_restored = settings_restored + && crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all( + &self.db, + ) + .is_ok(); + let runtime_restored = if settings_restored { + let expected_old_direct = old_direct_receipt.attempted_snapshot(); + self.reconcile_pi_runtime_at_epoch_with_native_precondition( + epoch, + Some(&expected_old_direct), + ) + .await + .is_ok() + } else { + self.republish_current_pi_runtime_if_native_matches(epoch) + .await + .unwrap_or(false) + }; + return Err(AppError::Config(format!( + "failed to reconcile Pi prompts in the new directory: {error}; rollback: settings={settings_restored}, prompts={prompts_restored}, skills={skills_restored}, gateway={runtime_restored}" + ))); + } + + if let Err(error) = + crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all(&self.db) + { + let settings_restored = crate::settings::update_settings(existing.clone()).is_ok(); + let prompts_restored = self + .db + .save_prompt_selection(AppType::Pi.as_str(), &previous_prompts) + .is_ok(); + let skills_restored = settings_restored + && crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all( + &self.db, + ) + .is_ok(); + let runtime_restored = if settings_restored { + let expected_old_direct = old_direct_receipt.attempted_snapshot(); + self.reconcile_pi_runtime_at_epoch_with_native_precondition( + epoch, + Some(&expected_old_direct), + ) + .await + .is_ok() + } else { + self.republish_current_pi_runtime_if_native_matches(epoch) + .await + .unwrap_or(false) + }; + return Err(AppError::Config(format!( + "failed to reconcile Pi Skills in the new directory: {error}; rollback: settings={settings_restored}, prompts={prompts_restored}, skills={skills_restored}, gateway={runtime_restored}" + ))); + } + + if let Err(error) = self + .reconcile_pi_runtime_at_epoch_with_native_precondition(epoch, Some(&new_native_before)) + .await + { + let settings_restored = crate::settings::update_settings(existing.clone()).is_ok(); + let prompts_restored = self + .db + .save_prompt_selection(AppType::Pi.as_str(), &previous_prompts) + .is_ok(); + let skills_restored = settings_restored + && crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all( + &self.db, + ) + .is_ok(); + let old_gateway_restored = if settings_restored { + let expected_old_direct = old_direct_receipt.attempted_snapshot(); + self.reconcile_pi_runtime_at_epoch_with_native_precondition( + epoch, + Some(&expected_old_direct), + ) + .await + .is_ok() + } else { + self.republish_current_pi_runtime_if_native_matches(epoch) + .await + .unwrap_or(false) + }; + return Err(AppError::Config(format!( + "failed to publish Pi in the new native directory: {error}; rollback: settings={settings_restored}, prompts={prompts_restored}, skills={skills_restored}, old_gateway={old_gateway_restored}" + ))); + } + + Ok(()) + } + + async fn suspend_pi_takeover_projection(&self) -> Result<(), AppError> { + if !crate::settings::pi_takeover_enabled() { + return Ok(()); + } + let epoch = self.pi_runtime.begin_mutation().await; + let direct_receipt = match self.restore_pi_direct_projection() { + Ok(receipt) => receipt, + Err(error) => { + let _ = self + .republish_current_pi_runtime_if_native_matches(epoch) + .await; + return Err(error); + } + }; + if let Err(error) = self.pi_runtime.close(epoch).await { + let expected_direct = direct_receipt.attempted_snapshot(); + let rollback = self + .reconcile_pi_runtime_at_epoch_with_native_precondition( + epoch, + Some(&expected_direct), + ) + .await; + return Err(AppError::Config(format!( + "failed to close Pi admission after direct projection: {error}; gateway rollback={rollback:?}" + ))); + } + Ok(()) + } + + /// Enter the portable-import boundary while the caller holds Pi's switch + /// lock. The old database still exists at this point, so this is the last + /// safe moment to restore every managed native key before SQL/binary + /// replacement can remove its direct configuration. + pub(crate) async fn prepare_pi_portable_import_under_lock( + &self, + _guard: &tokio::sync::OwnedMutexGuard<()>, + ) -> Result<(), AppError> { + self.suspend_pi_takeover_projection().await + } + + /// Re-publish the old catalog after an import operation aborts. Callers + /// retain Pi's switch lock from prepare through this compensation. + pub(crate) async fn recover_pi_after_aborted_portable_import_under_lock( + &self, + _guard: &tokio::sync::OwnedMutexGuard<()>, + ) -> Result<(), AppError> { + if crate::settings::pi_takeover_enabled() { + self.reconcile_pi_runtime().await.map(|_| ()) + } else { + Ok(()) + } + } + + async fn suspend_pi_takeover_for_process_exit(&self) -> Result<(), String> { + if !crate::settings::pi_takeover_enabled() { + return Ok(()); + } + let _guard = self.switch_locks.lock_for_app(AppType::Pi.as_str()).await; + self.suspend_pi_takeover_projection() + .await + .map_err(|error| { + format!( + "failed to restore Pi's direct native projection before process exit: {error}" + ) + }) + } + + pub(crate) async fn rotate_pi_gateway_token(&self) -> Result<(), AppError> { + let _guard = self.switch_locks.lock_for_app(AppType::Pi.as_str()).await; + let previous = crate::settings::get_or_create_pi_gateway_token()?; + let takeover_enabled = crate::settings::pi_takeover_enabled(); + let native_before = if takeover_enabled { + let models_path = crate::pi_config::native::get_pi_models_path()?; + let gateway = self.current_pi_gateway_projection()?; + Some(self.preflight_pi_owned_projection_at(&models_path, &gateway)?) + } else { + None + }; + crate::settings::reset_pi_gateway_token()?; + if !takeover_enabled { + return Ok(()); + } + let native_before = native_before.expect("enabled takeover has a native preflight"); + let epoch = self.pi_runtime.begin_mutation().await; + if let Err(error) = self + .reconcile_pi_runtime_at_epoch_with_native_precondition(epoch, Some(&native_before)) + .await + { + let token_restored = crate::settings::replace_pi_gateway_token(previous); + let runtime_restored = if token_restored.is_ok() { + let rollback_epoch = self.pi_runtime.begin_mutation().await; + self.reconcile_pi_runtime_at_epoch_with_native_precondition( + rollback_epoch, + Some(&native_before), + ) + .await + } else { + Err(AppError::Config( + "failed to restore the previous Pi gateway credential".to_string(), + )) + }; + let rollback = if token_restored.is_ok() && runtime_restored.is_ok() { + "the previous credential and runtime were restored" + } else { + "credential rotation rollback was incomplete" + }; + return Err(AppError::Config(format!( + "Pi gateway credential rotation failed: {error}; {rollback}" + ))); + } + Ok(()) + } + /// 启动代理服务器 pub async fn start(&self) -> Result { + // Listener identity and the Pi runtime/native projection form one + // publication boundary. A direct start command must serialize with + // provider edits, takeover toggles, and listener reconfiguration just + // like every other operation that can replace that identity. + let _pi_guard = self.switch_locks.lock_for_app(AppType::Pi.as_str()).await; + self.start_with_pi_lock_held().await + } + + async fn start_with_pi_lock_held(&self) -> Result { // 1. 启动时自动设置 proxy_enabled = true let mut global_config = self .db @@ -563,7 +1341,17 @@ impl ProxyService { // 4. 创建并启动服务器 let app_handle = self.app_handle.read().await.clone(); - let server = ProxyServer::new(config.clone(), self.db.clone(), app_handle); + let pi_server_generation = self + .pi_server_sequence + .fetch_add(1, Ordering::AcqRel) + .saturating_add(1); + let server = ProxyServer::new( + config.clone(), + self.db.clone(), + app_handle, + self.pi_runtime.clone(), + pi_server_generation, + ); let info = server .start() .await @@ -577,8 +1365,41 @@ impl ProxyService { } // 5. 保存服务器实例 + match pi_loopback_origin(&info.address, info.port) { + Ok(gateway_origin) => { + *self + .pi_listener + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = + Some(PiListenerIdentity { + server_generation: server.pi_server_generation(), + gateway_origin, + }); + } + Err(error) => { + *self + .pi_listener + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = None; + log::warn!("Pi gateway admission unavailable: {error}"); + } + } *self.server.write().await = Some(server); + if crate::settings::pi_takeover_enabled() { + if let Err(error) = self.reconcile_pi_runtime().await { + let rollback = self.suspend_pi_takeover_projection().await; + if let Err(rollback_error) = rollback { + return Err(format!( + "Pi desired takeover could not be reconciled after listener start: {error}; direct-mode rollback failed, so the live listener was retained: {rollback_error}" + )); + } + return Err(format!( + "Pi desired takeover could not be reconciled after listener start; direct mode was restored and the shared listener was retained: {error}" + )); + } + } + log::info!("代理服务器已启动: {}:{}", info.address, info.port); Ok(info) } @@ -600,6 +1421,57 @@ impl ProxyService { .map_err(|e| format!("保存动态代理端口失败: {e}")) } + async fn restore_proxy_config_snapshot(&self, config: &ProxyConfig) -> Result<(), String> { + self.db + .update_proxy_config(config.clone()) + .await + .map_err(|error| format!("恢复原代理配置失败: {error}")) + } + + fn replace_pi_listener_identity(&self, server: &ProxyServer, info: &ProxyServerInfo) { + *self + .pi_listener + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = + pi_loopback_origin(&info.address, info.port) + .ok() + .map(|gateway_origin| PiListenerIdentity { + server_generation: server.pi_server_generation(), + gateway_origin, + }); + } + + async fn recover_previous_listener( + &self, + previous_server: Option, + previous_config: &ProxyConfig, + ) -> Result<(), String> { + let server = previous_server.ok_or_else(|| "原代理监听器实例不可用".to_string())?; + let info = server + .start() + .await + .map_err(|error| format!("恢复原代理监听器失败: {error}"))?; + + let mut recovered_config = previous_config.clone(); + recovered_config.listen_port = info.port; + if let Err(error) = self.restore_proxy_config_snapshot(&recovered_config).await { + let _ = server.stop().await; + return Err(error); + } + + self.replace_pi_listener_identity(&server, &info); + *self.server.write().await = Some(server); + if crate::settings::pi_takeover_enabled() { + if let Err(error) = self.reconcile_pi_runtime().await { + let direct_rollback = self.suspend_pi_takeover_projection().await; + return Err(format!( + "原监听器已恢复,但 Pi 网关重建失败: {error}; 直连回滚: {direct_rollback:?}" + )); + } + } + Ok(()) + } + async fn start_before_takeover_if_ephemeral_port(&self) -> Result { let config = self .db @@ -724,6 +1596,24 @@ impl ProxyService { // OpenCode and OpenClaw don't support proxy features, always return false let opencode_enabled = false; let openclaw_enabled = false; + let pi_enabled = crate::settings::pi_takeover_enabled(); + let pi_operational_state = if !pi_enabled { + PiTakeoverOperationalState::Disabled + } else { + let listener = self + .pi_listener + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone(); + if listener + .as_ref() + .is_some_and(|listener| self.pi_runtime.is_admitting(listener.server_generation)) + { + PiTakeoverOperationalState::Active + } else { + PiTakeoverOperationalState::Degraded + } + }; Ok(ProxyTakeoverStatus { claude: claude_enabled, @@ -732,6 +1622,8 @@ impl ProxyService { grokbuild: grokbuild_enabled, opencode: opencode_enabled, openclaw: openclaw_enabled, + pi: pi_enabled, + pi_operational_state, }) } @@ -743,6 +1635,9 @@ impl ProxyService { let app = AppType::from_str(app_type).map_err(|e| format!("无效的应用类型: {e}"))?; let app_type_str = app.as_str(); let _guard = self.switch_locks.lock_for_app(app_type_str).await; + if app == AppType::Pi { + return self.set_pi_takeover_locked(enabled).await; + } if enabled { // 1) 代理服务未运行则自动启动 @@ -925,6 +1820,136 @@ impl ProxyService { Ok(()) } + async fn set_pi_takeover_locked(&self, enabled: bool) -> Result<(), String> { + if enabled { + let was_enabled = crate::settings::pi_takeover_enabled(); + if !was_enabled { + // Desired intent is durable before bind/token/projection work. + // Operational failures remain retryable on the next startup. + crate::settings::set_pi_takeover_enabled(true) + .map_err(|error| error.to_string())?; + } + if !self.is_running().await { + self.start_with_pi_lock_held().await?; + } + if self + .pi_listener + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .is_none() + { + let epoch = self.pi_runtime.begin_mutation().await; + let _ = self.pi_runtime.close(epoch).await; + return Err( + "Pi gateway requires the proxy listener to bind an explicit loopback address" + .to_string(), + ); + } + // Token creation is stable and occurs only after a listener exists. + if let Err(error) = crate::settings::get_or_create_pi_gateway_token() { + if !was_enabled { + let epoch = self.pi_runtime.begin_mutation().await; + let _ = self.pi_runtime.close(epoch).await; + } + return Err(error.to_string()); + } + + if let Err(error) = self.reconcile_pi_runtime().await { + if was_enabled { + let epoch = self.pi_runtime.begin_mutation().await; + let restored = self + .republish_current_pi_runtime_if_native_matches(epoch) + .await; + return Err(if matches!(restored, Ok(true)) { + format!( + "failed to refresh Pi gateway catalog; the previous runtime remains active: {error}" + ) + } else { + format!( + "failed to refresh Pi gateway catalog and restore admission: {error}" + ) + }); + } + + let direct_projection = self.restore_pi_direct_projection(); + match direct_projection { + Ok(_) => { + let epoch = self.pi_runtime.begin_mutation().await; + let admission_closed = self.pi_runtime.close(epoch).await; + return Err(if admission_closed.is_ok() { + format!( + "failed to publish Pi gateway catalog; direct mode was restored and desired takeover remains pending: {error}" + ) + } else { + format!( + "failed to publish Pi gateway catalog and close Pi admission: {error}" + ) + }); + } + Err(restore_error) => { + // Keep the desired bit true when restoring the native + // direct projection itself failed. Startup + // reconciliation can retry, and the error remains + // explicit instead of claiming a successful rollback + // while models.json still points local. + let epoch = self.pi_runtime.begin_mutation().await; + let _ = self.pi_runtime.close(epoch).await; + return Err(format!( + "failed to publish Pi gateway catalog: {error}; failed to restore the native direct projection: {restore_error}" + )); + } + } + } + self.refresh_active_target_from_current_provider(&AppType::Pi) + .await; + return Ok(()); + } + + if !crate::settings::pi_takeover_enabled() { + return Ok(()); + } + let epoch = self.pi_runtime.begin_mutation().await; + let direct_receipt = match self.restore_pi_direct_projection() { + Ok(receipt) => receipt, + Err(error) => { + let _ = self + .republish_current_pi_runtime_if_native_matches(epoch) + .await; + return Err(format!( + "failed to restore Pi's direct native projection: {error}" + )); + } + }; + if let Err(error) = crate::settings::set_pi_takeover_enabled(false) { + // Desired state did not change; rebuild the gateway projection and + // restore admission so the native file cannot be left lying. + let expected_direct = direct_receipt.attempted_snapshot(); + let _ = self + .reconcile_pi_runtime_at_epoch_with_native_precondition( + epoch, + Some(&expected_direct), + ) + .await; + return Err(format!("failed to persist Pi takeover disable: {error}")); + } + + self.pi_runtime + .close(epoch) + .await + .map_err(|error| error.to_string())?; + + let status = self.get_takeover_status().await?; + if !status.claude + && !status.codex + && !status.gemini + && !status.grokbuild + && self.is_running().await + { + let _ = self.stop().await; + } + Ok(()) + } + /// 同步关闭指定应用的 Live 接管(恢复配置并清标志,不停止代理服务)。 /// /// 用于 `ProfileService::apply` 等 sync 路径:调用者所在线程可能没有 Tokio @@ -1293,7 +2318,36 @@ impl ProxyService { /// 停止代理服务器 pub async fn stop(&self) -> Result<(), String> { + // `stop()` has several internal callers (profile switching included) + // that do not pass through the public takeover command. Never leave + // models.json pointing at a listener that this call is about to stop. + // The desired bit remains true so a later start can republish a fresh + // listener identity. + let _pi_guard = if crate::settings::pi_takeover_enabled() { + Some(self.switch_locks.lock_for_app(AppType::Pi.as_str()).await) + } else { + None + }; + if crate::settings::pi_takeover_enabled() { + self.suspend_pi_takeover_projection() + .await + .map_err(|error| { + format!( + "refusing to stop the proxy while Pi's direct projection cannot be restored: {error}" + ) + })?; + } + self.stop_listener_only().await + } + + async fn stop_listener_only(&self) -> Result<(), String> { if let Some(server) = self.server.write().await.take() { + let terminal_epoch = self.pi_runtime.begin_mutation().await; + let _ = self.pi_runtime.close(terminal_epoch).await; + *self + .pi_listener + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = None; server .stop() .await @@ -1324,8 +2378,12 @@ impl ProxyService { /// /// 会清除 settings 表中的代理状态,下次启动不会自动恢复。 pub async fn stop_with_restore(&self) -> Result<(), String> { + if crate::settings::pi_takeover_enabled() { + self.set_takeover_for_app(AppType::Pi.as_str(), false) + .await?; + } // 1. 停止代理服务器(即使未运行也继续执行恢复逻辑) - if let Err(e) = self.stop().await { + if let Err(e) = self.stop_listener_only().await { log::warn!("停止代理服务器失败(将继续恢复 Live 配置): {e}"); } @@ -1371,8 +2429,14 @@ impl ProxyService { /// /// 用于程序正常退出时,保留代理状态以便下次启动时自动恢复 pub async fn stop_with_restore_keep_state(&self) -> Result<(), String> { + // Pi's gateway projection lives in models.json rather than the legacy + // Live-backup table. Restore its direct exact-key projection while the + // listener is still alive, but retain the desired bit so startup can + // publish a fresh gateway projection. + self.suspend_pi_takeover_for_process_exit().await?; + // 1. 停止代理服务器(即使未运行也继续执行恢复逻辑) - if let Err(e) = self.stop().await { + if let Err(e) = self.stop_listener_only().await { log::warn!("停止代理服务器失败(将继续恢复 Live 配置): {e}"); } @@ -2206,7 +3270,7 @@ impl ProxyService { /// 检查是否处于 Live 接管模式 pub async fn is_takeover_active(&self) -> Result { let status = self.get_takeover_status().await?; - Ok(status.claude || status.codex || status.gemini || status.grokbuild) + Ok(status.claude || status.codex || status.gemini || status.grokbuild || status.pi) } /// 从异常退出中恢复(启动时调用) @@ -3099,6 +4163,11 @@ impl ProxyService { /// 更新代理配置 pub async fn update_config(&self, config: &ProxyConfig) -> Result<(), String> { + // Listener identity is part of Pi's native projection contract. + // Serialize address changes with every Pi catalog/takeover mutation so + // models.json can never race from one listener generation to another. + let _pi_guard = self.switch_locks.lock_for_app(AppType::Pi.as_str()).await; + // 记录旧配置用于判定是否需要重启 let previous = self .db @@ -3126,28 +4195,84 @@ impl ProxyService { || new_config.listen_port != previous.listen_port; if require_restart { - if let Some(server) = server_guard.take() { - server - .stop() - .await - .map_err(|e| format!("重启前停止代理服务器失败: {e}"))?; + // A listener restart first returns Pi to its direct, pinned-native + // representation while the old listener is still reachable. The + // desired bit is retained and a fresh gateway projection is only + // published after the new listener identity is known. + if let Err(error) = self.suspend_pi_takeover_projection().await { + let config_rollback = self.restore_proxy_config_snapshot(&previous).await; + return Err(format!( + "重启代理前恢复 Pi 直连投影失败: {error}; 配置回滚: {config_rollback:?}" + )); } + let previous_server = server_guard.take(); + if let Some(server) = previous_server.as_ref() { + if let Err(error) = server.stop().await { + *self + .pi_listener + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = None; + drop(server_guard); + let recovery = self + .recover_previous_listener(previous_server, &previous) + .await; + return Err(format!( + "重启前停止代理服务器失败: {error}; Pi 已恢复直连,原监听器恢复: {recovery:?}" + )); + } + } + *self + .pi_listener + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = None; + let app_handle = self.app_handle.read().await.clone(); - let new_server = ProxyServer::new(new_config.clone(), self.db.clone(), app_handle); - let info = new_server - .start() - .await - .map_err(|e| format!("重启代理服务器失败: {e}"))?; + let pi_server_generation = self + .pi_server_sequence + .fetch_add(1, Ordering::AcqRel) + .saturating_add(1); + let new_server = ProxyServer::new( + new_config.clone(), + self.db.clone(), + app_handle, + self.pi_runtime.clone(), + pi_server_generation, + ); + let info = match new_server.start().await { + Ok(info) => info, + Err(error) => { + drop(server_guard); + let recovery = self + .recover_previous_listener(previous_server, &previous) + .await; + return Err(format!( + "重启代理服务器失败: {error}; 原监听器恢复: {recovery:?}" + )); + } + }; if let Err(e) = self .persist_ephemeral_listen_port_if_needed(&new_config, info.port) .await { let _ = new_server.stop().await; - return Err(e); + drop(server_guard); + let recovery = self + .recover_previous_listener(previous_server, &previous) + .await; + return Err(format!("{e}; 原监听器恢复: {recovery:?}")); } + self.replace_pi_listener_identity(&new_server, &info); *server_guard = Some(new_server); + if crate::settings::pi_takeover_enabled() { + if let Err(error) = self.reconcile_pi_runtime().await { + let direct_rollback = self.suspend_pi_takeover_projection().await; + return Err(format!( + "重建 Pi 网关运行时失败: {error}; 新监听器保持运行,Pi 直连回滚: {direct_rollback:?}" + )); + } + } log::info!("代理配置已更新,服务器已自动重启应用最新配置"); // 如果当前存在任意 app 的 Live 接管,需要同步更新 Live 中的代理地址(否则客户端仍指向旧端口) @@ -3249,7 +4374,9 @@ impl ProxyService { #[cfg(test)] mod tests { use super::*; - use crate::provider::ProviderMeta; + use crate::database::NewProviderAggregate; + use crate::provider::{ProviderMeta, ProviderMutationInput}; + use crate::AppState; use serial_test::serial; use std::env; use tempfile::TempDir; @@ -3298,6 +4425,10 @@ mod tests { Some(value) => env::set_var("CC_SWITCH_TEST_HOME", value), None => env::remove_var("CC_SWITCH_TEST_HOME"), } + // Tests mutate the process-global settings cache after redirecting + // the home directory. Restore the cache after the environment so + // later serial tests do not inherit a dead temporary Pi override. + let _ = crate::settings::reload_settings(); } } @@ -3318,6 +4449,997 @@ mod tests { format!("http://127.0.0.1:{}/v1", status.port) } + #[test] + fn process_exit_projection_restores_claimed_pi_values_and_preserves_unowned_values() { + let db = Arc::new(Database::memory().expect("in-memory database")); + let config = serde_json::json!({ + "name": "Managed Pi", + "api": "openai-responses", + "baseUrl": "https://managed.example/v1", + "apiKey": "managed-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + let input = ProviderMutationInput { + id: "managed-pi".to_string(), + name: "Managed Pi".to_string(), + settings_config: config.clone(), + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }; + db.create_pi_catalog_provider( + NewProviderAggregate::from_input("pi", input).expect("build aggregate"), + "managed-pi", + ) + .expect("seed managed Pi provider"); + let state = AppState::new(db); + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("models.json"); + let unowned = serde_json::json!({ + "api": "anthropic-messages", + "baseUrl": "https://native.example", + "apiKey": "native-key", + "models": [{"id": "native-model"}] + }); + let expected_gateway = serde_json::json!({ + "api": "openai-responses", + "baseUrl": "http://127.0.0.1:15721/pi/route", + "apiKey": "gateway-token", + "models": [{"id": "model-a"}] + }); + std::fs::write( + &path, + serde_json::to_vec_pretty(&serde_json::json!({ + "providers": { + "managed-pi": expected_gateway.clone(), + "native": unowned.clone() + } + })) + .expect("serialize gateway projection"), + ) + .expect("write gateway projection"); + + state + .proxy_service + .restore_pi_direct_projection_at_with_expected_gateway( + &path, + &indexmap::IndexMap::from([("managed-pi".to_string(), Some(expected_gateway))]), + ) + .expect("restore direct Pi projection"); + + let restored: Value = + serde_json::from_slice(&std::fs::read(&path).expect("read restored Pi catalog")) + .expect("parse restored Pi catalog"); + assert_eq!(restored.pointer("/providers/managed-pi"), Some(&config)); + assert_eq!(restored.pointer("/providers/native"), Some(&unowned)); + } + + #[tokio::test] + #[serial] + async fn changing_pi_directory_moves_gateway_ownership_and_restores_the_old_file() { + let home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let old_dir = home.dir.path().join("old-pi"); + let new_dir = home.dir.path().join("new-pi"); + std::fs::create_dir_all(&old_dir).expect("old Pi directory"); + std::fs::create_dir_all(&new_dir).expect("new Pi directory"); + std::fs::write(old_dir.join("AGENTS.md"), "old-root-agents").expect("old Pi AGENTS.md"); + std::fs::write(new_dir.join("AGENTS.md"), "new-root-agents").expect("new Pi AGENTS.md"); + + let direct = json!({ + "name": "Managed Pi", + "api": "openai-responses", + "baseUrl": "https://managed.example/v1", + "apiKey": "managed-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + for directory in [&old_dir, &new_dir] { + std::fs::write( + directory.join("models.json"), + serde_json::to_vec_pretty(&json!({ + "providers": {"managed-pi": direct.clone()} + })) + .expect("serialize Pi models"), + ) + .expect("seed Pi models"); + } + + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some(old_dir.to_string_lossy().into_owned()); + crate::settings::update_settings(settings).expect("set old Pi directory"); + + let db = Arc::new(Database::memory().expect("in-memory database")); + let skill_source = crate::services::skill::SkillService::get_ssot_dir() + .expect("SSOT") + .join("directory-move"); + std::fs::create_dir_all(&skill_source).expect("skill source"); + std::fs::write( + skill_source.join("SKILL.md"), + "---\nname: directory-move\ndescription: directory move\n---\n", + ) + .expect("skill manifest"); + db.save_skill(&crate::app_config::InstalledSkill { + id: "local:directory-move".to_string(), + name: "Directory move".to_string(), + description: Some("directory move".to_string()), + directory: "directory-move".to_string(), + repo_owner: None, + repo_name: None, + repo_branch: None, + readme_url: None, + apps: crate::app_config::SkillApps::only(&AppType::Pi), + installed_at: 1, + content_hash: None, + updated_at: 1, + }) + .expect("save skill"); + crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all(&db) + .expect("deploy skill in old root"); + assert!(old_dir.join("skills").join("directory-move").exists()); + db.save_prompt( + AppType::Pi.as_str(), + &crate::prompt::Prompt { + id: "old-root".to_string(), + name: "Old root".to_string(), + content: "old-root-agents".to_string(), + description: None, + enabled: true, + created_at: Some(1), + updated_at: Some(1), + }, + ) + .expect("old prompt"); + db.save_prompt( + AppType::Pi.as_str(), + &crate::prompt::Prompt { + id: "new-root".to_string(), + name: "New root".to_string(), + content: "new-root-agents".to_string(), + description: None, + enabled: false, + created_at: Some(2), + updated_at: Some(2), + }, + ) + .expect("new prompt"); + use_ephemeral_proxy_port(&db).await; + let input = ProviderMutationInput { + id: "managed-pi".to_string(), + name: "Managed Pi".to_string(), + settings_config: direct.clone(), + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }; + db.create_pi_catalog_provider( + NewProviderAggregate::from_input("pi", input).expect("aggregate"), + "managed-pi", + ) + .expect("seed Pi catalog"); + + let service = ProxyService::new(db.clone()); + service + .set_takeover_for_app("pi", true) + .await + .expect("enable Pi takeover"); + + let guard = service.lock_switch_for_app("pi").await; + let existing = crate::settings::get_settings(); + let mut next = existing.clone(); + next.pi_config_dir = Some(new_dir.to_string_lossy().into_owned()); + service + .replace_settings_with_pi_directory_boundary_under_lock(&guard, &existing, next) + .await + .expect("move Pi directory ownership"); + drop(guard); + + let old_document: Value = serde_json::from_slice( + &std::fs::read(old_dir.join("models.json")).expect("read old Pi models"), + ) + .expect("parse old Pi models"); + assert_eq!( + old_document.pointer("/providers/managed-pi"), + Some(&direct), + "the old native directory must be restored to direct Pi semantics" + ); + + let new_document: Value = serde_json::from_slice( + &std::fs::read(new_dir.join("models.json")).expect("read new Pi models"), + ) + .expect("parse new Pi models"); + let new_base_url = new_document + .pointer("/providers/managed-pi/baseUrl") + .and_then(Value::as_str) + .expect("new gateway base URL"); + let new_base_url = url::Url::parse(new_base_url).expect("parse new gateway base URL"); + assert_eq!(new_base_url.host_str(), Some("127.0.0.1")); + assert!(new_base_url.path().starts_with("/pi/")); + assert_eq!( + crate::settings::get_settings().pi_config_dir.as_deref(), + Some(new_dir.to_string_lossy().as_ref()) + ); + let prompts = db.get_prompts(AppType::Pi.as_str()).expect("Pi prompts"); + assert!(!prompts["old-root"].enabled); + assert!(prompts["new-root"].enabled); + assert_eq!( + std::fs::read_to_string(old_dir.join("AGENTS.md")).expect("old AGENTS"), + "old-root-agents", + "directory changes must not migrate or clean the old native file" + ); + assert_eq!( + std::fs::read_to_string(new_dir.join("AGENTS.md")).expect("new AGENTS"), + "new-root-agents" + ); + assert!( + !old_dir.join("skills").join("directory-move").exists(), + "verified stale Skill ownership must be cleaned after new deployment" + ); + assert!(new_dir.join("skills").join("directory-move").exists()); + let deployments = db + .get_pi_skill_deployments("local:directory-move") + .expect("skill deployments"); + assert_eq!(deployments.len(), 1); + assert!(deployments[0] + .destination + .contains(new_dir.to_string_lossy().as_ref())); + + service + .set_takeover_for_app("pi", false) + .await + .expect("disable Pi takeover"); + } + + #[tokio::test] + #[serial] + async fn changing_pi_directory_rejects_an_unowned_managed_key_before_side_effects() { + let home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let old_dir = home.dir.path().join("old-pi"); + let new_dir = home.dir.path().join("new-pi"); + std::fs::create_dir_all(&old_dir).expect("old Pi directory"); + std::fs::create_dir_all(&new_dir).expect("new Pi directory"); + + let managed = json!({ + "name": "Managed Pi", + "api": "openai-responses", + "baseUrl": "https://managed.example/v1", + "apiKey": "managed-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + let foreign = json!({ + "api": "anthropic-messages", + "baseUrl": "https://foreign.example", + "apiKey": "foreign-secret", + "models": [{"id": "foreign-model"}] + }); + let foreign_document = serde_json::to_vec_pretty(&json!({ + "providers": {"managed-pi": foreign} + })) + .expect("foreign models"); + std::fs::write(new_dir.join("models.json"), &foreign_document).expect("seed foreign key"); + + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some(old_dir.to_string_lossy().into_owned()); + settings.pi_takeover_enabled = false; + crate::settings::update_settings(settings).expect("old Pi settings"); + let db = Arc::new(Database::memory().expect("database")); + db.create_pi_catalog_provider( + NewProviderAggregate::from_input( + AppType::Pi.as_str(), + ProviderMutationInput { + id: "managed-pi".to_string(), + name: "Managed Pi".to_string(), + settings_config: managed, + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }, + ) + .expect("aggregate"), + "managed-pi", + ) + .expect("managed provider"); + + let service = ProxyService::new(db); + let switch_guard = service.lock_switch_for_app(AppType::Pi.as_str()).await; + let existing = crate::settings::get_settings(); + let mut next = existing.clone(); + next.pi_config_dir = Some(new_dir.to_string_lossy().into_owned()); + let error = service + .replace_settings_with_pi_directory_boundary_under_lock(&switch_guard, &existing, next) + .await + .expect_err("unowned target key must reject the directory move"); + assert!(error.to_string().contains("overwrite unowned provider key")); + assert_eq!( + crate::settings::get_settings().pi_config_dir, + existing.pi_config_dir, + "preflight rejection must not switch the native authority" + ); + assert_eq!( + std::fs::read(new_dir.join("models.json")).expect("foreign models remain"), + foreign_document + ); + } + + #[tokio::test] + #[serial] + async fn changing_pi_directory_does_not_overwrite_an_exact_key_changed_after_preflight() { + let home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let old_dir = home.dir.path().join("old-pi"); + let new_dir = home.dir.path().join("new-pi"); + std::fs::create_dir_all(&old_dir).expect("old Pi directory"); + std::fs::create_dir_all(&new_dir).expect("new Pi directory"); + + let direct = json!({ + "name": "Managed Pi", + "api": "openai-responses", + "baseUrl": "https://managed.example/v1", + "apiKey": "managed-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + let external = json!({ + "name": "External Pi", + "api": "openai-responses", + "baseUrl": "https://external.example/v1", + "apiKey": "external-key", + "models": [{"id": "external-model", "name": "External"}] + }); + let new_models_path = new_dir.join("models.json"); + std::fs::write( + &new_models_path, + serde_json::to_vec(&json!({"providers": {}})).expect("serialize direct document"), + ) + .expect("seed direct document"); + + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some(old_dir.to_string_lossy().into_owned()); + settings.pi_takeover_enabled = false; + crate::settings::update_settings(settings).expect("old Pi settings"); + let db = Arc::new(Database::memory().expect("database")); + db.create_pi_catalog_provider( + NewProviderAggregate::from_input( + AppType::Pi.as_str(), + ProviderMutationInput { + id: "managed-pi".to_string(), + name: "Managed Pi".to_string(), + settings_config: direct, + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }, + ) + .expect("aggregate"), + "managed-pi", + ) + .expect("managed provider"); + crate::pi_config::shared_file::replace_before_next_compare_exchange( + &new_models_path, + &serde_json::to_vec(&json!({"providers": {"managed-pi": external.clone()}})) + .expect("serialize external document"), + ); + + let service = ProxyService::new(db); + let switch_guard = service.lock_switch_for_app(AppType::Pi.as_str()).await; + let existing = crate::settings::get_settings(); + let mut next = existing.clone(); + next.pi_config_dir = Some(new_dir.to_string_lossy().into_owned()); + let error = service + .replace_settings_with_pi_directory_boundary_under_lock(&switch_guard, &existing, next) + .await + .expect_err("the post-preflight native edit must win"); + assert!(error.to_string().contains("changed since")); + assert_eq!( + crate::settings::get_settings().pi_config_dir, + existing.pi_config_dir + ); + let live: Value = serde_json::from_slice( + &std::fs::read(&new_models_path).expect("read external native document"), + ) + .expect("parse external native document"); + assert_eq!(live.pointer("/providers/managed-pi"), Some(&external)); + } + + #[tokio::test] + #[serial] + async fn changing_pi_directory_without_takeover_reconciles_missing_agents_truth() { + let home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let old_dir = home.dir.path().join("old-direct-pi"); + let new_dir = home.dir.path().join("new-direct-pi"); + std::fs::create_dir_all(&old_dir).expect("old Pi directory"); + std::fs::create_dir_all(&new_dir).expect("new Pi directory"); + std::fs::write(old_dir.join("AGENTS.md"), "old-only").expect("old AGENTS"); + + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some(old_dir.to_string_lossy().into_owned()); + settings.pi_takeover_enabled = false; + crate::settings::update_settings(settings).expect("old directory settings"); + + let db = Arc::new(Database::memory().expect("database")); + let direct = json!({ + "name": "Managed direct Pi", + "api": "openai-responses", + "baseUrl": "https://direct.example/v1", + "apiKey": "direct-key", + "models": [{"id": "direct-model", "name": "Direct model"}] + }); + db.create_pi_catalog_provider( + NewProviderAggregate::from_input( + AppType::Pi.as_str(), + ProviderMutationInput { + id: "managed-direct".to_string(), + name: "Managed direct Pi".to_string(), + settings_config: direct.clone(), + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }, + ) + .expect("aggregate"), + "managed-direct", + ) + .expect("managed direct provider"); + db.save_prompt( + AppType::Pi.as_str(), + &crate::prompt::Prompt { + id: "old-only".to_string(), + name: "Old only".to_string(), + content: "old-only".to_string(), + description: None, + enabled: true, + created_at: Some(1), + updated_at: Some(1), + }, + ) + .expect("prompt"); + let service = ProxyService::new(db.clone()); + let switch_guard = service.lock_switch_for_app(AppType::Pi.as_str()).await; + let existing = crate::settings::get_settings(); + let mut next = existing.clone(); + next.pi_config_dir = Some(new_dir.to_string_lossy().into_owned()); + service + .replace_settings_with_pi_directory_boundary_under_lock(&switch_guard, &existing, next) + .await + .expect("change direct Pi directory"); + + assert!( + db.get_prompts(AppType::Pi.as_str()) + .expect("prompts") + .values() + .all(|prompt| !prompt.enabled), + "a missing AGENTS.md in the new root is the inactive authority" + ); + assert_eq!( + std::fs::read_to_string(old_dir.join("AGENTS.md")).expect("old AGENTS survives"), + "old-only" + ); + assert!(!new_dir.join("AGENTS.md").exists()); + let new_models: Value = serde_json::from_slice( + &std::fs::read(new_dir.join("models.json")).expect("new direct models"), + ) + .expect("parse new direct models"); + assert_eq!( + new_models.pointer("/providers/managed-direct"), + Some(&direct), + "a direct-mode directory change must publish every managed provider in the new root" + ); + } + + #[tokio::test] + #[serial] + async fn pi_gateway_returns_the_last_real_retryable_upstream_response_when_no_later_send_occurs( + ) { + let home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let pi_dir = home.dir.path().join("pi"); + std::fs::create_dir_all(&pi_dir).expect("Pi directory"); + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some(pi_dir.to_string_lossy().into_owned()); + settings.pi_takeover_enabled = false; + crate::settings::update_settings(settings).expect("Pi settings"); + + let upstream_hits = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let hits = upstream_hits.clone(); + let upstream = axum::Router::new().fallback(axum::routing::any(move || { + let hits = hits.clone(); + async move { + hits.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + let mut response = axum::response::Response::builder() + .status(http::StatusCode::TOO_MANY_REQUESTS) + .header("x-upstream-marker", "real-429") + .header(http::header::CONTENT_TYPE, "application/json") + .body(axum::body::Body::from( + r#"{"error":{"message":"upstream rate limit"}}"#, + )) + .expect("upstream response"); + response + .headers_mut() + .insert("retry-after", http::HeaderValue::from_static("17")); + response + } + })); + let upstream_listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .expect("upstream listener"); + let upstream_address = upstream_listener.local_addr().expect("upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(upstream_listener, upstream) + .await + .expect("upstream server"); + }); + + let db = Arc::new(Database::memory().expect("database")); + use_ephemeral_proxy_port(&db).await; + let direct = json!({ + "name": "Retryable upstream", + "api": "openai-responses", + "baseUrl": format!("http://{upstream_address}/v1"), + "apiKey": "upstream-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + db.create_pi_catalog_provider( + NewProviderAggregate::from_input( + AppType::Pi.as_str(), + ProviderMutationInput { + id: "retryable-upstream".to_string(), + name: "Retryable upstream".to_string(), + settings_config: direct.clone(), + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }, + ) + .expect("aggregate"), + "retryable-upstream", + ) + .expect("managed provider"); + std::fs::write( + pi_dir.join("models.json"), + serde_json::to_vec(&json!({ + "providers": {"retryable-upstream": direct} + })) + .expect("serialize native catalog"), + ) + .expect("native catalog"); + + let service = ProxyService::new(db); + service + .set_takeover_for_app(AppType::Pi.as_str(), true) + .await + .expect("enable Pi gateway"); + let projected: Value = serde_json::from_slice( + &std::fs::read(pi_dir.join("models.json")).expect("projected catalog"), + ) + .expect("parse projected catalog"); + let base_url = projected + .pointer("/providers/retryable-upstream/baseUrl") + .and_then(Value::as_str) + .expect("gateway base URL"); + let gateway_token = projected + .pointer("/providers/retryable-upstream/apiKey") + .and_then(Value::as_str) + .expect("gateway token"); + + let response = reqwest::Client::new() + .post(format!("{base_url}/responses")) + .bearer_auth(gateway_token) + .json(&json!({"model": "model-a", "input": "hello"})) + .send() + .await + .expect("gateway response"); + let status = response.status(); + let marker = response + .headers() + .get("x-upstream-marker") + .and_then(|value| value.to_str().ok()) + .map(str::to_string); + let retry_after = response + .headers() + .get("retry-after") + .and_then(|value| value.to_str().ok()) + .map(str::to_string); + let body = response.text().await.expect("upstream body"); + upstream_task.abort(); + + assert_eq!(status, reqwest::StatusCode::TOO_MANY_REQUESTS); + assert_eq!(marker.as_deref(), Some("real-429")); + assert_eq!(retry_after.as_deref(), Some("17")); + assert_eq!(body, r#"{"error":{"message":"upstream rate limit"}}"#); + assert_eq!( + upstream_hits.load(std::sync::atomic::Ordering::SeqCst), + 1, + "local candidate exhaustion must not invent a phantom retry" + ); + service + .set_takeover_for_app(AppType::Pi.as_str(), false) + .await + .expect("disable Pi gateway"); + } + + #[tokio::test] + #[serial] + async fn failed_listener_rebind_recovers_pi_on_the_previous_listener() { + let _home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some( + crate::config::get_home_dir() + .join(".pi/agent") + .to_string_lossy() + .into_owned(), + ); + crate::settings::update_settings(settings).expect("set Pi directory"); + + let db = Arc::new(Database::memory().expect("in-memory database")); + use_ephemeral_proxy_port(&db).await; + let config = json!({ + "name": "Managed Pi", + "api": "openai-responses", + "baseUrl": "https://managed.example/v1", + "apiKey": "managed-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + let input = ProviderMutationInput { + id: "managed-pi".to_string(), + name: "Managed Pi".to_string(), + settings_config: config.clone(), + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }; + db.create_pi_catalog_provider( + NewProviderAggregate::from_input("pi", input).expect("aggregate"), + "managed-pi", + ) + .expect("seed Pi catalog"); + let models_path = crate::pi_config::native::get_pi_models_path().expect("models path"); + std::fs::create_dir_all(models_path.parent().expect("Pi directory")).expect("Pi directory"); + std::fs::write( + &models_path, + serde_json::to_vec_pretty(&json!({"providers": {"managed-pi": config}})) + .expect("serialize models"), + ) + .expect("write models"); + + let service = ProxyService::new(db.clone()); + service + .set_takeover_for_app("pi", true) + .await + .expect("enable Pi takeover"); + let occupied = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .expect("reserve conflicting port"); + let mut rejected = db.get_proxy_config().await.expect("proxy config"); + rejected.listen_port = occupied.local_addr().expect("reserved address").port(); + assert!(service.update_config(&rejected).await.is_err()); + + let status = service.get_status().await.expect("recovered status"); + assert!(status.running); + let stored = db.get_proxy_config().await.expect("recovered config"); + assert_eq!(stored.listen_port, status.port); + let projected: Value = + serde_json::from_slice(&std::fs::read(&models_path).expect("read models")) + .expect("parse models"); + let projected_base = projected + .pointer("/providers/managed-pi/baseUrl") + .and_then(Value::as_str) + .expect("gateway base url"); + assert!(projected_base.starts_with(&format!("http://127.0.0.1:{}/pi/", status.port))); + assert!(!projected_base.contains(&rejected.listen_port.to_string())); + + let import_guard = service.lock_switch_for_app(AppType::Pi.as_str()).await; + service + .prepare_pi_portable_import_under_lock(&import_guard) + .await + .expect("portable import boundary restores the old direct catalog"); + let import_safe: Value = + serde_json::from_slice(&std::fs::read(&models_path).expect("read import-safe models")) + .expect("parse import-safe models"); + assert_eq!(import_safe.pointer("/providers/managed-pi"), Some(&config)); + service + .recover_pi_after_aborted_portable_import_under_lock(&import_guard) + .await + .expect("aborted portable import restores the gateway"); + drop(import_guard); + let republished: Value = + serde_json::from_slice(&std::fs::read(&models_path).expect("read republished models")) + .expect("parse republished models"); + assert!( + republished + .pointer("/providers/managed-pi/baseUrl") + .and_then(Value::as_str) + .is_some_and( + |base| base.starts_with(&format!("http://127.0.0.1:{}/pi/", status.port)) + ) + ); + drop(occupied); + + service + .stop() + .await + .expect("a direct stop restores Pi before closing the listener"); + let direct: Value = + serde_json::from_slice(&std::fs::read(&models_path).expect("read direct models")) + .expect("parse direct models"); + assert_eq!(direct.pointer("/providers/managed-pi"), Some(&config)); + assert!( + crate::settings::pi_takeover_enabled(), + "safe listener stop preserves desired takeover for the next start" + ); + service + .set_takeover_for_app("pi", false) + .await + .expect("disable Pi takeover"); + } + + #[tokio::test] + #[serial] + async fn initial_pi_bind_failure_keeps_desired_state_and_reports_degraded() { + let _home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some( + crate::config::get_home_dir() + .join(".pi/agent") + .to_string_lossy() + .into_owned(), + ); + settings.pi_takeover_enabled = false; + crate::settings::update_settings(settings).expect("Pi settings"); + + let occupied = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .expect("reserve port"); + let db = Arc::new(Database::memory().expect("database")); + let mut proxy_config = db.get_proxy_config().await.expect("proxy config"); + proxy_config.listen_address = "127.0.0.1".to_string(); + proxy_config.listen_port = occupied.local_addr().expect("address").port(); + db.update_proxy_config(proxy_config) + .await + .expect("fixed occupied port"); + let service = ProxyService::new(db); + + service + .set_takeover_for_app("pi", true) + .await + .expect_err("occupied port must fail"); + assert!( + crate::settings::pi_takeover_enabled(), + "bind failure must not erase user intent" + ); + let status = service.get_takeover_status().await.expect("status"); + assert!(status.pi); + assert_eq!( + status.pi_operational_state, + PiTakeoverOperationalState::Degraded + ); + assert!(!service.is_running().await); + + drop(occupied); + service + .set_takeover_for_app("pi", false) + .await + .expect("explicit disable clears desired state"); + } + + #[tokio::test] + #[serial] + async fn listener_absence_does_not_authorize_overwriting_a_non_direct_native_key() { + let home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let pi_dir = home.dir.path().join("pi"); + std::fs::create_dir_all(&pi_dir).expect("Pi directory"); + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some(pi_dir.to_string_lossy().into_owned()); + settings.pi_takeover_enabled = true; + crate::settings::update_settings(settings).expect("Pi settings"); + + let db = Arc::new(Database::memory().expect("database")); + let direct = json!({ + "name": "Managed Pi", + "api": "openai-responses", + "baseUrl": "https://managed.example/v1", + "apiKey": "managed-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + db.create_pi_catalog_provider( + NewProviderAggregate::from_input( + AppType::Pi.as_str(), + ProviderMutationInput { + id: "managed-pi".to_string(), + name: "Managed Pi".to_string(), + settings_config: direct, + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }, + ) + .expect("aggregate"), + "managed-pi", + ) + .expect("managed provider"); + let external = json!({ + "name": "External edit", + "api": "openai-responses", + "baseUrl": "https://external.example/v1", + "apiKey": "external-key", + "models": [{"id": "external-model"}] + }); + let native_bytes = serde_json::to_vec(&json!({ + "providers": {"managed-pi": external} + })) + .expect("serialize external catalog"); + let models_path = pi_dir.join("models.json"); + std::fs::write(&models_path, &native_bytes).expect("external native catalog"); + + let service = ProxyService::new(db); + let error = service + .set_takeover_for_app(AppType::Pi.as_str(), false) + .await + .expect_err("no listener is not ownership evidence"); + assert!(error.contains("not already in its direct projection")); + assert!(crate::settings::pi_takeover_enabled()); + assert_eq!( + std::fs::read(&models_path).expect("external catalog remains"), + native_bytes + ); + } + + #[tokio::test] + #[serial] + async fn fenced_pi_runtime_is_not_republished_after_external_native_drift() { + let home = TempHome::new(); + crate::settings::reload_settings().expect("reload isolated settings"); + let pi_dir = home.dir.path().join("pi"); + std::fs::create_dir_all(&pi_dir).expect("Pi directory"); + let mut settings = crate::settings::get_settings(); + settings.pi_config_dir = Some(pi_dir.to_string_lossy().into_owned()); + settings.pi_takeover_enabled = true; + crate::settings::update_settings(settings).expect("Pi settings"); + + let db = Arc::new(Database::memory().expect("database")); + let direct = json!({ + "name": "Managed Pi", + "api": "openai-responses", + "baseUrl": "https://managed.example/v1", + "apiKey": "managed-key", + "models": [{"id": "model-a", "name": "Model A"}] + }); + db.create_pi_catalog_provider( + NewProviderAggregate::from_input( + AppType::Pi.as_str(), + ProviderMutationInput { + id: "managed-pi".to_string(), + name: "Managed Pi".to_string(), + settings_config: direct.clone(), + website_url: None, + category: None, + created_at: None, + sort_index: Some(0), + notes: None, + meta: None, + icon: Some("pi".to_string()), + icon_color: None, + in_failover_queue: false, + }, + ) + .expect("aggregate"), + "managed-pi", + ) + .expect("managed provider"); + let models_path = pi_dir.join("models.json"); + std::fs::write( + &models_path, + serde_json::to_vec_pretty(&json!({"providers": {"managed-pi": direct}})) + .expect("direct models"), + ) + .expect("seed direct models"); + + let service = ProxyService::new(db); + *service + .pi_listener + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(PiListenerIdentity { + server_generation: 41, + gateway_origin: url::Url::parse("http://127.0.0.1:15721/").expect("gateway origin"), + }); + service + .reconcile_pi_runtime() + .await + .expect("publish initial runtime"); + assert_eq!( + service + .get_takeover_status() + .await + .expect("active status") + .pi_operational_state, + PiTakeoverOperationalState::Active + ); + + let epoch = service.begin_pi_catalog_mutation().await; + let external = serde_json::to_vec_pretty(&json!({ + "providers": { + "managed-pi": { + "name": "External", + "api": "openai-responses", + "baseUrl": "https://external.example/v1", + "apiKey": "external-key", + "models": [{"id": "external-model"}] + } + } + })) + .expect("external models"); + std::fs::write(&models_path, &external).expect("external edit"); + + assert!(matches!( + service + .republish_current_pi_runtime_if_native_matches(epoch) + .await, + Err(AppError::Conflict(_)) + )); + assert_eq!( + service + .get_takeover_status() + .await + .expect("degraded status") + .pi_operational_state, + PiTakeoverOperationalState::Degraded + ); + assert_eq!( + std::fs::read(&models_path).expect("external file remains"), + external + ); + } + fn seed_codex_model_template() { let codex_dir = crate::codex_config::get_codex_config_dir(); std::fs::create_dir_all(&codex_dir).expect("create codex dir"); diff --git a/src-tauri/src/services/skill.rs b/src-tauri/src/services/skill.rs index 84f02e9eb..fc507cd59 100644 --- a/src-tauri/src/services/skill.rs +++ b/src-tauri/src/services/skill.rs @@ -565,6 +565,9 @@ impl SkillService { return Ok(custom.join("skills")); } } + AppType::Pi => { + return Ok(crate::pi_config::native::get_pi_agent_dir()?.join("skills")); + } } // 默认路径:回退到用户主目录下的标准位置。 @@ -581,6 +584,7 @@ impl SkillService { AppType::OpenCode => home.join(".config").join("opencode").join("skills"), AppType::OpenClaw => home.join(".openclaw").join("skills"), AppType::Hermes => crate::hermes_config::get_hermes_dir().join("skills"), + AppType::Pi => crate::pi_config::native::get_pi_agent_dir()?.join("skills"), }) } @@ -637,8 +641,19 @@ impl SkillService { // 同一仓库的同名 skill,返回现有记录(可能需要更新启用状态) let mut updated = existing.clone(); updated.apps.set_enabled_for(current_app, true); + if matches!(current_app, AppType::Pi) { + let guard = crate::services::skill_deployment::PiSkillDeploymentService::operation_guard(); + crate::services::skill_deployment::PiSkillDeploymentService::toggle_under_guard( + &guard, + db, + &mut updated, + true, + ) + .map_err(|error| anyhow!(error.to_string()))?; + return Ok(updated); + } db.save_skill(&updated)?; - Self::sync_to_app_dir(&updated.directory, current_app)?; + Self::sync_installed_skill_to_app(db, &updated, current_app)?; log::info!( "Skill {} 已存在,更新 {:?} 启用状态", updated.name, @@ -671,6 +686,7 @@ impl SkillService { } let dest = ssot_dir.join(&install_name); + let destination_preexisted = dest.exists(); let mut repo_branch = skill.repo_branch.clone(); @@ -788,11 +804,39 @@ impl SkillService { updated_at: 0, }; - // 保存到数据库 - db.save_skill(&installed_skill)?; - - // 同步到当前应用目录 - Self::sync_to_app_dir(&install_name, current_app)?; + let installed_skill = if matches!(current_app, AppType::Pi) { + let guard = + crate::services::skill_deployment::PiSkillDeploymentService::operation_guard(); + let mut persisted = installed_skill.clone(); + // Let the deployment coordinator commit the desired bit and ledger + // evidence together. Until then the row is deliberately disabled. + persisted.apps.pi = false; + if let Err(error) = db.save_skill(&persisted) { + if !destination_preexisted { + let _ = fs::remove_dir_all(&dest); + } + return Err(error.into()); + } + if let Err(error) = + crate::services::skill_deployment::PiSkillDeploymentService::toggle_under_guard( + &guard, + db, + &mut persisted, + true, + ) + { + let _ = db.delete_skill(&persisted.id); + if !destination_preexisted { + let _ = fs::remove_dir_all(&dest); + } + return Err(anyhow!(error.to_string())); + } + persisted + } else { + db.save_skill(&installed_skill)?; + Self::sync_installed_skill_to_app(db, &installed_skill, current_app)?; + installed_skill + }; log::info!( "Skill {} 安装成功,已启用 {:?}", @@ -810,6 +854,8 @@ impl SkillService { /// 2. 从 SSOT 删除 /// 3. 从数据库删除 pub fn uninstall(db: &Arc, id: &str) -> Result { + let deployment_guard = + crate::services::skill_deployment::PiSkillDeploymentService::operation_guard(); // 获取 skill 信息 let skill = db .get_installed_skill(id)? @@ -828,8 +874,15 @@ impl SkillService { let backup_path = Self::create_uninstall_backup(&skill)? .map(|path| path.to_string_lossy().to_string()); + crate::services::skill_deployment::PiSkillDeploymentService::remove_before_uninstall_under_guard( + &deployment_guard, + db, + &skill, + ) + .map_err(|error| anyhow!(error.to_string()))?; + // 从所有应用目录删除 - for app in AppType::all() { + for app in AppType::all().filter(|app| !matches!(app, AppType::Pi)) { let _ = Self::remove_from_app(&directory, &app); } @@ -1113,15 +1166,40 @@ impl SkillService { )) })?; + // All Pi deployment mutations, SSOT replacement, and ledger + // reconciliation share one process boundary. Downloading remains + // outside the lock so a slow network cannot block toggles. + let deployment_guard = + crate::services::skill_deployment::PiSkillDeploymentService::operation_guard(); + // 备份旧文件 let _ = Self::create_uninstall_backup(&skill); - // 删除旧 SSOT 目录并复制新文件 + // Stage the exact old SSOT tree in the same directory. Reconstructing + // it from the remote source is not a rollback: local files may differ. let dest = ssot_dir.join(&skill.directory); - if dest.exists() { - fs::remove_dir_all(&dest)?; + let staged_previous = if fs::symlink_metadata(&dest).is_ok() { + let staged = ssot_dir.join(format!( + ".{}.cc-switch-update-{}", + skill.directory, + uuid::Uuid::new_v4().simple() + )); + fs::rename(&dest, &staged)?; + Some(staged) + } else { + None + }; + if let Err(error) = Self::copy_dir_recursive(&source, &dest) { + if let Some(staged) = staged_previous.as_deref() { + fs::rename(staged, &dest).with_context(|| { + format!( + "Skill update copy failed ({error}); restoring {} also failed", + dest.display() + ) + })?; + } + return Err(error); } - Self::copy_dir_recursive(&source, &dest)?; // 计算新哈希 + 解析新元数据 let new_hash = Self::compute_dir_hash(&dest).ok(); @@ -1151,14 +1229,71 @@ impl SkillService { updated_at: chrono::Utc::now().timestamp(), }; - db.save_skill(&updated_skill)?; + if let Err(error) = db.save_skill(&updated_skill) { + let _ = Self::remove_path(&dest); + if let Some(staged) = staged_previous.as_deref() { + fs::rename(staged, &dest).with_context(|| { + format!( + "Skill metadata update failed ({error}); restoring {} also failed", + dest.display() + ) + })?; + } + return Err(error.into()); + } - // 同步到所有已启用的应用目录 - for app in updated_skill.apps.enabled_apps() { - if let Err(e) = Self::sync_to_app_dir(&updated_skill.directory, &app) { + // Pi is the consistency-critical consumer: update its owned + // deployment before best-effort legacy app copies. + if updated_skill.apps.pi { + if let Err(error) = + crate::services::skill_deployment::PiSkillDeploymentService::reconcile_skill_under_guard( + &deployment_guard, + db, + &updated_skill, + ) + { + let db_rollback = db.save_skill(&skill); + let file_rollback = Self::remove_path(&dest).and_then(|_| { + if let Some(staged) = staged_previous.as_deref() { + fs::rename(staged, &dest).map_err(anyhow::Error::from) + } else { + Ok(()) + } + }); + return match (db_rollback, file_rollback) { + (Ok(()), Ok(())) => Err(anyhow!(error.to_string())), + (db_result, file_result) => Err(anyhow!( + "Pi Skill update failed ({error}); DB rollback: {}; file rollback: {}", + db_result + .err() + .map_or_else(|| "ok".to_string(), |value| value.to_string()), + file_result + .err() + .map_or_else(|| "ok".to_string(), |value| value.to_string()) + )), + }; + } + } + + // 同步到所有已启用的其他应用目录 + for app in updated_skill + .apps + .enabled_apps() + .into_iter() + .filter(|app| !matches!(app, AppType::Pi)) + { + if let Err(e) = Self::sync_installed_skill_to_app(db, &updated_skill, &app) { log::warn!("同步更新后的 skill 到 {:?} 失败: {e}", app); } } + if let Some(staged) = staged_previous { + if let Err(error) = Self::remove_path(&staged) { + log::warn!( + "Failed to remove committed Skill update rollback staging '{}': {error}", + staged.display() + ); + } + } log::info!("Skill {} 更新成功", updated_skill.name); Ok(updated_skill) @@ -1201,7 +1336,10 @@ impl SkillService { /// 迁移 Skill 存储位置(在两个 SSOT 目录间移动文件) /// - /// 安全策略:先移文件,后改设置。中途崩溃时设置仍指向旧目录。 + /// Safety strategy: copy first while the old SSOT remains live, switch the + /// setting, reconcile every app, then delete the old trees. Keeping both + /// roots during reconciliation lets the Pi ownership ledger verify its old + /// symlink/copy before atomically replacing it. pub fn migrate_storage( db: &Arc, target: SkillStorageLocation, @@ -1215,6 +1353,9 @@ impl SkillService { }); } + let deployment_guard = + crate::services::skill_deployment::PiSkillDeploymentService::operation_guard(); + // 1. 解析旧目录和新目录(不改设置) let old_dir = Self::get_ssot_dir()?; let new_dir = match target { @@ -1225,18 +1366,18 @@ impl SkillService { }; fs::create_dir_all(&new_dir)?; - // 2. 逐个移动 skill 目录 + // 2. Copy every valid tree. Do not rename/delete the old root before + // Pi has verified the ownership identity recorded in its ledger. let skills = db.get_all_installed_skills()?; let mut result = MigrationResult { migrated_count: 0, skipped_count: 0, errors: vec![], }; + let mut copied = Vec::<(PathBuf, PathBuf)>::new(); for skill in skills.values() { - // 下面是 rename 与 remove_dir_all,脏 directory 可把任意目录搬走或删掉。 - // 软失败:本函数已有 errors 收集通道,记一条继续处理其余 skill, - // 不要整体中断——用户只是在切换存储位置。 + // Invalid DB rows are reported but never joined to either root. let directory = match Self::require_valid_directory(&skill.directory) { Ok(directory) => directory, Err(err) => { @@ -1253,32 +1394,90 @@ impl SkillService { result.skipped_count += 1; continue; } - if dst.exists() { - result.skipped_count += 1; - continue; + if fs::symlink_metadata(&dst).is_ok() { + for (_, copied_destination) in copied.iter().rev() { + let _ = Self::remove_path(copied_destination); + } + return Err(anyhow!( + "Skill storage target already contains an unowned entry: {}", + dst.display() + )); } - - // 优先 rename(同文件系统原子操作),失败则 copy+delete - match fs::rename(&src, &dst) { - Ok(()) => result.migrated_count += 1, - Err(_) => match Self::copy_dir_recursive(&src, &dst) { - Ok(()) => { - let _ = fs::remove_dir_all(&src); - result.migrated_count += 1; - } - Err(e) => { - result.errors.push(format!("{}: {e}", skill.directory)); - } - }, + if let Err(error) = Self::copy_dir_recursive(&src, &dst) { + let _ = Self::remove_path(&dst); + for (_, copied_destination) in copied.iter().rev() { + let _ = Self::remove_path(copied_destination); + } + return Err(error); } + copied.push((src, dst)); + result.migrated_count += 1; } - // 3. 文件移动完成后才持久化设置 - crate::settings::set_skill_storage_location(target)?; + // 3. Switch authority only after every new tree is complete. + if let Err(error) = crate::settings::set_skill_storage_location(target) { + for (_, copied_destination) in copied.iter().rev() { + let _ = Self::remove_path(copied_destination); + } + return Err(error.into()); + } - // 4. 刷新所有应用目录的 symlink(指向新 SSOT) - for app in AppType::all() { - let _ = Self::sync_to_app(db, &app); + // 4. Reconcile Pi under the same mutex, then all legacy app views. + let reconcile_result = + crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all_under_guard( + &deployment_guard, + db, + ) + .map_err(|error| anyhow!(error.to_string())) + .and_then(|()| { + for app in AppType::all().filter(|app| !matches!(app, AppType::Pi)) { + Self::sync_to_app(db, &app)?; + } + Ok(()) + }); + if let Err(error) = reconcile_result { + let mut rollback_errors = Vec::new(); + if let Err(rollback) = crate::settings::set_skill_storage_location(current) { + rollback_errors.push(format!("settings: {rollback}")); + } else { + if let Err(rollback) = + crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all_under_guard( + &deployment_guard, + db, + ) + { + rollback_errors.push(format!("Pi deployment: {rollback}")); + } + for app in AppType::all().filter(|app| !matches!(app, AppType::Pi)) { + if let Err(rollback) = Self::sync_to_app(db, &app) { + rollback_errors.push(format!("{app:?}: {rollback}")); + } + } + } + for (_, copied_destination) in copied.iter().rev() { + if let Err(rollback) = Self::remove_path(copied_destination) { + rollback_errors.push(format!("{}: {rollback}", copied_destination.display())); + } + } + return if rollback_errors.is_empty() { + Err(error) + } else { + Err(anyhow!( + "Skill storage migration failed ({error}); rollback failures: {}", + rollback_errors.join("; ") + )) + }; + } + + // 5. Only after every consumer points at the new root may the old + // sources be removed. Cleanup errors are visible but do not roll back + // an already-consistent authority switch. + for (old_source, _) in &copied { + if let Err(error) = Self::remove_path(old_source) { + result + .errors + .push(format!("{}: {error}", old_source.display())); + } } log::info!( @@ -1403,7 +1602,7 @@ impl SkillService { } if !restored_skill.apps.is_empty() { - if let Err(err) = Self::sync_to_app_dir(&restored_skill.directory, current_app) { + if let Err(err) = Self::sync_installed_skill_to_app(db, &restored_skill, current_app) { let _ = db.delete_skill(&restored_skill.id); let _ = fs::remove_dir_all(&restore_path); return Err(err); @@ -1429,6 +1628,14 @@ impl SkillService { .get_installed_skill(id)? .ok_or_else(|| anyhow!("Skill not found: {id}"))?; + if matches!(app, AppType::Pi) { + crate::services::skill_deployment::PiSkillDeploymentService::toggle( + db, &mut skill, enabled, + ) + .map_err(|error| anyhow!(error.to_string()))?; + return Ok(()); + } + // 更新状态 skill.apps.set_enabled_for(app, enabled); @@ -1517,6 +1724,11 @@ impl SkillService { db: &Arc, imports: Vec, ) -> Result> { + // Import can explicitly acquire or release Pi filesystem ownership. + // Serialize the source scan, SSOT establishment, ownership decision and + // desired-state transaction with every other Pi deployment operation. + let deployment_guard = + crate::services::skill_deployment::PiSkillDeploymentService::operation_guard(); let ssot_dir = Self::get_ssot_dir()?; let agents_lock = parse_agents_lock(); let mut imported = Vec::new(); @@ -1579,8 +1791,12 @@ impl SkillService { // 复制到 SSOT let dest = ssot_dir.join(&dir_name); - if !dest.exists() { - Self::copy_dir_recursive(&source, &dest)?; + let created_ssot = !dest.exists(); + if created_ssot { + if let Err(error) = Self::copy_dir_recursive(&source, &dest) { + let _ = Self::remove_path(&dest); + return Err(error); + } } // 解析元数据 @@ -1588,7 +1804,7 @@ impl SkillService { let (name, description) = Self::read_skill_name_desc(&skill_md, &dir_name); // 启用状态仅信任用户本次显式选择,不再根据“在哪些位置找到”自动推断。 - let apps = selection.apps; + let requested_apps = selection.apps; // 从 lock 文件提取仓库信息 let (id, repo_owner, repo_name, repo_branch, readme_url) = @@ -1599,7 +1815,8 @@ impl SkillService { let content_hash = Self::compute_dir_hash(&ssot_skill_dir).ok(); // 创建记录 - let skill = InstalledSkill { + let previous = db.get_installed_skill(&id)?; + let mut skill = InstalledSkill { id, name, description, @@ -1608,14 +1825,54 @@ impl SkillService { repo_name, repo_branch, readme_url, - apps, + // save_skill intentionally preserves an existing Pi desired bit. + // For a new row keep it disabled until the deployment ledger and + // desired bit can commit in one transaction below. + apps: SkillApps { + pi: previous.as_ref().is_some_and(|installed| installed.apps.pi), + ..requested_apps.clone() + }, installed_at: chrono::Utc::now().timestamp(), content_hash, updated_at: 0, }; // 保存到数据库 - db.save_skill(&skill)?; + if let Err(error) = db.save_skill(&skill) { + if created_ssot { + let _ = Self::remove_path(&dest); + } + return Err(error.into()); + } + if let Err(error) = crate::services::skill_deployment::PiSkillDeploymentService::import_desired_state_under_guard( + &deployment_guard, + db, + &mut skill, + requested_apps.pi, + ) { + let db_rollback = if let Some(previous) = previous.as_ref() { + db.save_skill(previous).map(|_| ()) + } else { + db.delete_skill(&skill.id).map(|_| ()) + }; + let file_rollback = if created_ssot { + Self::remove_path(&dest) + } else { + Ok(()) + }; + return match (db_rollback, file_rollback) { + (Ok(()), Ok(())) => Err(anyhow!(error.to_string())), + (db_result, file_result) => Err(anyhow!( + "Pi Skill import failed ({error}); DB rollback: {}; SSOT rollback: {}", + db_result + .err() + .map_or_else(|| "ok".to_string(), |value| value.to_string()), + file_result + .err() + .map_or_else(|| "ok".to_string(), |value| value.to_string()) + )), + }; + } imported.push(skill); } @@ -1655,6 +1912,19 @@ impl SkillService { crate::settings::get_skill_sync_method() } + fn sync_installed_skill_to_app( + db: &Arc, + skill: &InstalledSkill, + app: &AppType, + ) -> Result<()> { + if matches!(app, AppType::Pi) { + crate::services::skill_deployment::PiSkillDeploymentService::reconcile_skill(db, skill) + .map_err(|error| anyhow!(error.to_string())) + } else { + Self::sync_to_app_dir(&skill.directory, app) + } + } + /// 同步 Skill 到应用目录(使用 symlink 或 copy) /// /// 根据配置和平台选择最佳同步方式: @@ -1665,6 +1935,11 @@ impl SkillService { if matches!(app, AppType::ClaudeDesktop) { return Ok(()); } + if matches!(app, AppType::Pi) { + return Err(anyhow!( + "Pi Skill deployment requires the ownership ledger; use the database-aware reconciler" + )); + } // directory 可能来自被污染的 DB 行(如同步导入的远端快照),join 前必须校验。 let directory = Self::require_valid_directory(directory)?; @@ -1861,6 +2136,10 @@ impl SkillService { if matches!(app, AppType::ClaudeDesktop) { return Ok(()); } + if matches!(app, AppType::Pi) { + return crate::services::skill_deployment::PiSkillDeploymentService::reconcile_all(db) + .map_err(|error| anyhow!(error.to_string())); + } let skills = db.get_all_installed_skills()?; let ssot_dir = Self::get_ssot_dir()?; @@ -2113,7 +2392,7 @@ impl SkillService { } /// 静态方法:解析技能元数据 - fn parse_skill_metadata_static(path: &Path) -> Result { + pub(crate) fn parse_skill_metadata_static(path: &Path) -> Result { let content = fs::read_to_string(path)?; let content = content.trim_start_matches('\u{feff}'); @@ -2733,6 +3012,35 @@ impl SkillService { Ok(()) } + /// Copy into a unique sibling and publish with an OS no-replace rename. + /// A concurrent installer can win, but its directory is never overwritten. + fn copy_dir_noreplace(src: &Path, dest: &Path) -> Result<()> { + let parent = dest + .parent() + .ok_or_else(|| anyhow!("Skill destination has no parent: {}", dest.display()))?; + fs::create_dir_all(parent)?; + let name = dest + .file_name() + .and_then(|value| value.to_str()) + .ok_or_else(|| anyhow!("Skill destination has an invalid name: {}", dest.display()))?; + let staged = parent.join(format!( + ".{name}.cc-switch-install-{}", + uuid::Uuid::new_v4().simple() + )); + if let Err(error) = Self::copy_dir_recursive(src, &staged) { + let _ = fs::remove_dir_all(&staged); + return Err(error); + } + if let Err(error) = crate::pi_config::shared_file::publish_path_noreplace(&staged, dest) { + let _ = fs::remove_dir_all(&staged); + return Err(anyhow!( + "Skill destination was created concurrently ({}): {error}", + dest.display() + )); + } + Ok(()) + } + fn resolve_uninstall_backup_source(skill: &InstalledSkill) -> Result> { // 返回值会被整目录复制进 ~/.cc-switch/skill-backups/ 并由 get_skill_backups // 在界面上列出——脏 directory 在这里等于任意文件读取 + 外泄通道。 @@ -2996,6 +3304,10 @@ impl SkillService { let ssot_dir = Self::get_ssot_dir()?; let mut installed = Vec::new(); let existing_skills = db.get_all_installed_skills()?; + let mut claimed_directories = existing_skills + .values() + .map(|skill| skill.directory.to_ascii_lowercase()) + .collect::>(); let zip_stem = zip_path .file_stem() .and_then(|s| s.to_str()) @@ -3062,6 +3374,32 @@ impl SkillService { ); continue; } + if claimed_directories.contains(&install_name.to_ascii_lowercase()) { + log::warn!( + "Skill directory '{}' appears more than once in the archive, skipping", + install_name + ); + continue; + } + + if matches!(current_app, AppType::Pi) + && meta.as_ref().is_none_or(|metadata| { + metadata + .name + .as_deref() + .is_none_or(|name| name.trim().is_empty()) + || metadata + .description + .as_deref() + .is_none_or(|description| description.trim().is_empty()) + }) + { + return Err(anyhow!(format_skill_error( + "INVALID_SKILL_DIRECTORY", + &[("directory", &install_name)], + Some("checkSkillManifest"), + ))); + } let (name, description) = match meta { Some(m) => ( @@ -3071,18 +3409,35 @@ impl SkillService { None => (install_name.clone(), None), }; + let deployment_guard = matches!(current_app, AppType::Pi) + .then(crate::services::skill_deployment::PiSkillDeploymentService::operation_guard); + let pi_source_digest = if deployment_guard.is_some() { + Some( + crate::services::skill_deployment::PiSkillDeploymentService::source_digest( + &skill_dir, + ) + .map_err(|error| anyhow!(error.to_string()))?, + ) + } else { + None + }; + // 复制到 SSOT let dest = ssot_dir.join(&install_name); - if dest.exists() { - let _ = fs::remove_dir_all(&dest); + if fs::symlink_metadata(&dest).is_ok() { + return Err(anyhow!(format_skill_error( + "SKILL_DIRECTORY_CONFLICT", + &[("directory", &install_name)], + Some("uninstallFirst"), + ))); } - Self::copy_dir_recursive(&skill_dir, &dest)?; + Self::copy_dir_noreplace(&skill_dir, &dest)?; // 计算内容哈希 let content_hash = Self::compute_dir_hash(&dest).ok(); // 创建 InstalledSkill 记录 - let skill = InstalledSkill { + let mut skill = InstalledSkill { id: format!("local:{install_name}"), name, description, @@ -3097,17 +3452,75 @@ impl SkillService { updated_at: 0, }; - // 保存到数据库 - db.save_skill(&skill)?; - - // 同步到当前应用目录 - Self::sync_to_app_dir(&install_name, current_app)?; + if let Some(guard) = deployment_guard.as_ref() { + // The coordinator commits Pi desired state and ownership + // evidence together. Until then the portable row is inert. + skill.apps.pi = false; + if let Err(error) = db.save_skill(&skill) { + let cleanup = + crate::services::skill_deployment::PiSkillDeploymentService::remove_source_if_unchanged( + &dest, + pi_source_digest + .as_deref() + .expect("Pi ZIP publication has a source digest"), + ); + return match cleanup { + Ok(()) => Err(error.into()), + Err(cleanup) => Err(anyhow!( + "Pi Skill ZIP database write failed ({error}); SSOT rollback failed ({cleanup})" + )), + }; + } + if let Err(error) = + crate::services::skill_deployment::PiSkillDeploymentService::toggle_under_guard( + guard, db, &mut skill, true, + ) + { + if let Err(deployment_rollback) = + crate::services::skill_deployment::PiSkillDeploymentService::remove_before_uninstall_under_guard( + guard, db, &skill, + ) + { + return Err(anyhow!( + "Pi Skill ZIP install failed ({error}); native ownership rollback failed ({deployment_rollback}); the DB row and SSOT were retained as recovery evidence" + )); + } + let db_rollback = db.delete_skill(&skill.id); + if !matches!(&db_rollback, Ok(true)) { + return Err(anyhow!( + "Pi Skill ZIP install failed ({error}); DB rollback failed ({}); SSOT was retained", + match db_rollback { + Ok(false) => "row missing".to_string(), + Err(value) => value.to_string(), + Ok(true) => unreachable!(), + } + )); + } + let file_rollback = + crate::services::skill_deployment::PiSkillDeploymentService::remove_source_if_unchanged( + &dest, + pi_source_digest + .as_deref() + .expect("Pi ZIP publication has a source digest"), + ); + return match file_rollback { + Ok(()) => Err(anyhow!(error.to_string())), + Err(file_error) => Err(anyhow!( + "Pi Skill ZIP install failed ({error}); SSOT rollback failed ({file_error})" + )), + }; + } + } else { + db.save_skill(&skill)?; + Self::sync_installed_skill_to_app(db, &skill, current_app)?; + } log::info!( "Skill {} installed from ZIP, enabled for {:?}", skill.name, current_app ); + claimed_directories.insert(install_name.to_ascii_lowercase()); installed.push(skill); } @@ -4013,6 +4426,34 @@ mod tests { ); } + #[test] + fn no_replace_directory_publish_preserves_an_existing_destination() { + let temp = tempdir().expect("tempdir"); + let source = temp.path().join("source"); + let destination = temp.path().join("destination"); + fs::create_dir(&source).expect("source"); + fs::create_dir(&destination).expect("destination"); + fs::write(source.join("SKILL.md"), "managed").expect("source manifest"); + fs::write(destination.join("SKILL.md"), "external").expect("external manifest"); + + SkillService::copy_dir_noreplace(&source, &destination) + .expect_err("an existing destination must win atomically"); + assert_eq!( + fs::read_to_string(destination.join("SKILL.md")).expect("external destination"), + "external" + ); + assert!( + fs::read_dir(temp.path()) + .expect("temp root") + .all(|entry| !entry + .expect("entry") + .file_name() + .to_string_lossy() + .contains("cc-switch-install")), + "a rejected staged publication must be cleaned" + ); + } + #[test] fn extract_local_zip_hands_back_a_guard_that_owns_the_tree() { use std::io::Write; @@ -4142,6 +4583,169 @@ mod tests { } } + #[test] + #[serial_test::serial] + fn importing_a_native_pi_skill_adopts_exact_content_and_can_disable_it() { + struct PiDirGuard(Option); + impl Drop for PiDirGuard { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("PI_CODING_AGENT_DIR", value), + None => std::env::remove_var("PI_CODING_AGENT_DIR"), + } + } + } + struct StorageLocationGuard(SkillStorageLocation); + impl Drop for StorageLocationGuard { + fn drop(&mut self) { + let _ = crate::settings::set_skill_storage_location(self.0); + } + } + + let temp = tempdir().expect("tempdir"); + let _home_guard = TestHomeGuard::set(temp.path()); + let pi_agent_dir = temp.path().join("pi-agent"); + let _pi_dir_guard = PiDirGuard(std::env::var_os("PI_CODING_AGENT_DIR")); + std::env::set_var("PI_CODING_AGENT_DIR", &pi_agent_dir); + let _storage_guard = StorageLocationGuard(crate::settings::get_skill_storage_location()); + crate::settings::set_skill_storage_location(SkillStorageLocation::CcSwitch) + .expect("select isolated SSOT"); + + let native = pi_agent_dir.join("skills").join("native-skill"); + write_skill(&native, "Native Skill"); + fs::write(native.join("details.txt"), "pinned native bytes").expect("native detail"); + let db = Arc::new(Database::memory().expect("memory db")); + + let imported = SkillService::import_from_apps( + &db, + vec![ImportSkillSelection { + directory: "native-skill".to_string(), + apps: SkillApps::only(&AppType::Pi), + }], + ) + .expect("explicit Pi import should adopt the exact native tree"); + assert_eq!(imported.len(), 1); + assert!(imported[0].apps.pi); + + let statuses = + crate::services::skill_deployment::PiSkillDeploymentService::inspect_all(&db) + .expect("inspect Pi deployment"); + let status = statuses + .get(&imported[0].id) + .expect("imported status must exist"); + assert_eq!( + status.ownership, + crate::services::skill_deployment::PiSkillOwnership::Owned + ); + assert_eq!( + status.discovery, + crate::services::skill_deployment::PiSkillDiscovery::Active + ); + assert!(status.effectively_discovered); + + SkillService::toggle_app(&db, &imported[0].id, &AppType::Pi, false) + .expect("owned imported Pi Skill can be disabled"); + assert!( + !native.exists(), + "disabling an explicitly adopted native tree removes that owned deployment" + ); + assert!( + SkillService::get_ssot_dir() + .expect("SSOT") + .join("native-skill") + .join("SKILL.md") + .is_file(), + "disabling Pi must preserve the managed SSOT" + ); + assert!( + !db.get_installed_skill(&imported[0].id) + .expect("read imported row") + .expect("row") + .apps + .pi + ); + } + + #[test] + #[serial_test::serial] + fn pi_zip_collision_rolls_back_database_and_ssot_without_touching_native_skill() { + use std::io::Write; + use zip::write::SimpleFileOptions; + + struct PiDirGuard(Option); + impl Drop for PiDirGuard { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("PI_CODING_AGENT_DIR", value), + None => std::env::remove_var("PI_CODING_AGENT_DIR"), + } + } + } + struct StorageLocationGuard(SkillStorageLocation); + impl Drop for StorageLocationGuard { + fn drop(&mut self) { + let _ = crate::settings::set_skill_storage_location(self.0); + } + } + + let temp = tempdir().expect("tempdir"); + let _home_guard = TestHomeGuard::set(temp.path()); + let pi_agent_dir = temp.path().join("pi-agent"); + let _pi_dir_guard = PiDirGuard(std::env::var_os("PI_CODING_AGENT_DIR")); + std::env::set_var("PI_CODING_AGENT_DIR", &pi_agent_dir); + let _storage_guard = StorageLocationGuard(crate::settings::get_skill_storage_location()); + crate::settings::set_skill_storage_location(SkillStorageLocation::CcSwitch) + .expect("isolated SSOT"); + + let native = pi_agent_dir.join("skills").join("collision"); + write_skill(&native, "Native collision"); + fs::write(native.join("native.txt"), "must survive").expect("native bytes"); + + let mut archive = Vec::new(); + { + let mut zip = zip::ZipWriter::new(std::io::Cursor::new(&mut archive)); + let options = SimpleFileOptions::default(); + zip.start_file("collision/SKILL.md", options) + .expect("manifest entry"); + zip.write_all(b"---\nname: Imported\ndescription: Imported collision\n---\n") + .expect("manifest bytes"); + zip.start_file("collision/imported.txt", options) + .expect("payload entry"); + zip.write_all(b"must not remain").expect("payload bytes"); + zip.finish().expect("finish zip"); + } + let zip_path = temp.path().join("collision.zip"); + fs::write(&zip_path, archive).expect("write zip"); + let db = Arc::new(Database::memory().expect("database")); + + SkillService::install_from_zip(&db, &zip_path, &AppType::Pi) + .expect_err("unowned native collision must fail"); + + assert_eq!( + fs::read_to_string(native.join("native.txt")).expect("native survives"), + "must survive" + ); + assert!( + db.get_installed_skill("local:collision") + .expect("read skill") + .is_none(), + "failed ZIP install must not leave desired state" + ); + assert!( + db.get_pi_skill_deployments("local:collision") + .expect("read ledger") + .is_empty(), + "failed ZIP install must not create ownership evidence" + ); + assert!( + !SkillService::get_ssot_dir() + .expect("SSOT") + .join("collision") + .exists(), + "failed ZIP install must compensate its SSOT copy" + ); + } + fn poisoned_skill(id: &str, directory: &str) -> InstalledSkill { InstalledSkill { id: id.to_string(), diff --git a/src-tauri/src/services/skill_deployment.rs b/src-tauri/src/services/skill_deployment.rs new file mode 100644 index 000000000..d669e7c04 --- /dev/null +++ b/src-tauri/src/services/skill_deployment.rs @@ -0,0 +1,1353 @@ +//! Ownership-safe Pi Skill deployment reconciliation. +//! +//! Desired state lives on the installed Skill row. Filesystem presence alone +//! is never ownership evidence; only the device-local deployment ledger may +//! authorize replacement or deletion. + +use crate::app_config::{AppType, InstalledSkill}; +use crate::database::{Database, SkillDeployment, SkillDeploymentMethod}; +use crate::error::AppError; +use crate::services::skill::{SkillService, SyncMethod}; +use serde::Serialize; +use sha2::{Digest, Sha256}; +use std::collections::{BTreeMap, HashMap}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, Mutex, MutexGuard, OnceLock}; + +pub(crate) struct PiSkillDeploymentService; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum PiSkillOwnership { + Absent, + Owned, + Foreign, + Stale, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum PiSkillDiscovery { + Absent, + Active, + Shadowed, + Invalid, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct SkillAppStatus { + pub desired_enabled: bool, + pub owned_deployment: bool, + pub effectively_discovered: bool, + pub ownership: PiSkillOwnership, + pub discovery: PiSkillDiscovery, + pub issue: Option, +} + +impl PiSkillDeploymentService { + pub(crate) fn reconcile_skill( + db: &Arc, + skill: &InstalledSkill, + ) -> Result<(), AppError> { + let guard = Self::operation_guard(); + Self::reconcile_skill_under_guard(&guard, db, skill) + } + + pub(crate) fn toggle( + db: &Arc, + skill: &mut InstalledSkill, + enabled: bool, + ) -> Result<(), AppError> { + let guard = Self::operation_guard(); + Self::toggle_under_guard(&guard, db, skill, enabled) + } + + pub(crate) fn toggle_under_guard( + _guard: &MutexGuard<'static, ()>, + db: &Arc, + skill: &mut InstalledSkill, + enabled: bool, + ) -> Result<(), AppError> { + if enabled { + let destination = skill_destination(skill)?; + let destination_key = destination_key(&destination); + let existing = db.get_pi_skill_deployment(&skill.id, &destination_key)?; + let source = skill_source(skill)?; + deploy( + db, + skill, + &source, + &destination, + &destination_key, + existing, + Some(true), + )?; + cleanup_stale_deployments_after_commit(db, skill, &destination_key); + } else { + remove_all_recorded_deployments(db, skill, Some(false))?; + } + skill.apps.pi = enabled; + Ok(()) + } + + /// Apply the Pi desired state selected by the user while importing an + /// existing Skill from application directories. + /// + /// Import is the one operation where an unowned Pi destination may become + /// managed: the user explicitly selected that exact native tree. Adoption + /// is allowed only when every byte in the native destination matches the + /// newly established SSOT source. A mere directory/name match is never + /// ownership evidence. + pub(crate) fn import_desired_state_under_guard( + _guard: &MutexGuard<'static, ()>, + db: &Arc, + skill: &mut InstalledSkill, + enabled: bool, + ) -> Result<(), AppError> { + let destination = skill_destination(skill)?; + let destination_key = destination_key(&destination); + let existing = db.get_pi_skill_deployment(&skill.id, &destination_key)?; + if enabled { + let source = skill_source(skill)?; + if existing.is_none() && fs::symlink_metadata(&destination).is_ok() { + adopt_exact_import(db, skill, &source, &destination, &destination_key)?; + } else { + deploy( + db, + skill, + &source, + &destination, + &destination_key, + existing, + Some(true), + )?; + } + cleanup_stale_deployments_after_commit(db, skill, &destination_key); + } else { + remove_all_recorded_deployments(db, skill, Some(false))?; + } + skill.apps.pi = enabled; + Ok(()) + } + + pub(crate) fn reconcile_all(db: &Arc) -> Result<(), AppError> { + let guard = Self::operation_guard(); + Self::reconcile_all_under_guard(&guard, db) + } + + pub(crate) fn operation_guard() -> MutexGuard<'static, ()> { + deployment_lock() + } + + pub(crate) fn reconcile_skill_under_guard( + _guard: &MutexGuard<'static, ()>, + db: &Arc, + skill: &InstalledSkill, + ) -> Result<(), AppError> { + reconcile_skill_unlocked(db, skill) + } + + pub(crate) fn reconcile_all_under_guard( + _guard: &MutexGuard<'static, ()>, + db: &Arc, + ) -> Result<(), AppError> { + for skill in db.get_all_installed_skills()?.values() { + // Portable sync and old databases may contain a poisoned directory + // name. Reject it before any path join, but do not let that inert + // row hide every valid Pi Skill during startup/storage migration. + // Only this syntactic row corruption is skippable: source errors, + // foreign destinations and stale ownership still fail closed. + if let Err(error) = validate_directory_name(&skill.directory) { + log::warn!( + "skipping invalid Pi Skill row '{}' during reconciliation: {error}", + skill.id + ); + continue; + } + reconcile_skill_unlocked(db, skill)?; + } + Ok(()) + } + + pub(crate) fn remove_before_uninstall_under_guard( + _guard: &MutexGuard<'static, ()>, + db: &Arc, + skill: &InstalledSkill, + ) -> Result<(), AppError> { + remove_all_recorded_deployments(db, skill, None) + } + + pub(crate) fn inspect_all( + db: &Arc, + ) -> Result, AppError> { + let _guard = deployment_lock(); + let discovery = scan_pi_discovery()?; + db.get_all_installed_skills()? + .into_iter() + .map(|(id, skill)| { + inspect_skill_status(db, &skill, &discovery).map(|status| (id, status)) + }) + .collect() + } + + pub(crate) fn source_digest(path: &Path) -> Result { + tree_digest(path) + } + + /// Remove a just-published SSOT tree only if the exact bytes installed by + /// this operation still own the path. The namespace move precedes digest + /// validation so an external replacement is restored, never recursively + /// deleted after a check/use race. + pub(crate) fn remove_source_if_unchanged( + path: &Path, + expected_digest: &str, + ) -> Result<(), AppError> { + let staged = stage_destination(path)?; + let observed = tree_digest(&staged); + if !matches!(observed, Ok(ref digest) if digest == expected_digest) { + restore_staged_destination(&staged, path).map_err(|rollback| { + AppError::Conflict(format!( + "Pi Skill SSOT ownership changed and rollback failed ({rollback}); recovery tree: {}", + staged.display() + )) + })?; + return Err(AppError::Conflict(format!( + "Pi Skill SSOT changed before rollback: {}", + path.display() + ))); + } + remove_path(&staged) + } +} + +#[derive(Debug)] +struct PiDiscoveryScan { + by_manifest: HashMap)>, +} + +fn scan_pi_discovery() -> Result { + const MAX_SKILL_MANIFEST_BYTES: u64 = 1024 * 1024; + const MAX_SKILL_DIRECTORIES: usize = 10_000; + + let root = SkillService::get_app_skills_dir(&AppType::Pi) + .map_err(|error| AppError::Config(error.to_string()))?; + let mut entries = match fs::read_dir(&root) { + Ok(entries) => entries + .collect::, _>>() + .map_err(|error| AppError::io(&root, error))?, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + return Ok(PiDiscoveryScan { + by_manifest: HashMap::new(), + }); + } + Err(error) => return Err(AppError::io(&root, error)), + }; + if entries.len() > MAX_SKILL_DIRECTORIES { + return Err(AppError::InvalidInput(format!( + "Pi Skill discovery exceeds {MAX_SKILL_DIRECTORIES} top-level entries" + ))); + } + entries.sort_by_key(std::fs::DirEntry::file_name); + + let mut winner_by_name = HashMap::::new(); + let mut by_manifest = HashMap::new(); + for entry in entries { + let directory = entry.path(); + let metadata = fs::metadata(&directory).map_err(|error| AppError::io(&directory, error))?; + if !metadata.is_dir() { + continue; + } + let manifest = directory.join("SKILL.md"); + let metadata = match fs::symlink_metadata(&manifest) { + Ok(metadata) if metadata.file_type().is_file() => metadata, + Ok(_) => { + by_manifest.insert( + manifest, + ( + PiSkillDiscovery::Invalid, + Some("SKILL.md is not a regular file".to_string()), + ), + ); + continue; + } + Err(error) if error.kind() == std::io::ErrorKind::NotFound => continue, + Err(error) => return Err(AppError::io(&manifest, error)), + }; + if metadata.len() > MAX_SKILL_MANIFEST_BYTES { + by_manifest.insert( + manifest, + ( + PiSkillDiscovery::Invalid, + Some("SKILL.md exceeds the 1 MiB inspection limit".to_string()), + ), + ); + continue; + } + let parsed = SkillService::parse_skill_metadata_static(&manifest) + .map_err(|error| AppError::Config(error.to_string()))?; + let Some(name) = parsed.name.filter(|name| !name.trim().is_empty()) else { + by_manifest.insert( + manifest, + ( + PiSkillDiscovery::Invalid, + Some("SKILL.md has no non-empty frontmatter name".to_string()), + ), + ); + continue; + }; + if parsed + .description + .as_deref() + .is_none_or(|description| description.trim().is_empty()) + { + by_manifest.insert( + manifest, + ( + PiSkillDiscovery::Invalid, + Some("SKILL.md has no non-empty frontmatter description".to_string()), + ), + ); + continue; + } + if let Some(winner) = winner_by_name.get(&name) { + by_manifest.insert( + manifest, + ( + PiSkillDiscovery::Shadowed, + Some(format!( + "skill name '{name}' is shadowed by {}", + winner.display() + )), + ), + ); + } else { + winner_by_name.insert(name, manifest.clone()); + by_manifest.insert(manifest, (PiSkillDiscovery::Active, None)); + } + } + Ok(PiDiscoveryScan { by_manifest }) +} + +fn inspect_skill_status( + db: &Arc, + skill: &InstalledSkill, + discovery: &PiDiscoveryScan, +) -> Result { + let destination = skill_destination(skill)?; + let destination_key = destination_key(&destination); + let deployments = db.get_pi_skill_deployments(&skill.id)?; + let deployment = deployments + .iter() + .find(|deployment| deployment.destination_key == destination_key); + let has_stale_destination = deployments + .iter() + .any(|deployment| deployment.destination_key != destination_key); + let manifest = destination.join("SKILL.md"); + let discovered = discovery.by_manifest.get(&manifest); + let destination_exists = fs::symlink_metadata(&destination).is_ok(); + let owned_deployment = deployment + .is_some_and(|deployment| verify_owned_destination(deployment, &destination).is_ok()); + let ownership = if has_stale_destination { + PiSkillOwnership::Stale + } else { + match (deployment.is_some(), destination_exists, owned_deployment) { + (_, _, true) => PiSkillOwnership::Owned, + (true, _, false) => PiSkillOwnership::Stale, + (false, true, false) => PiSkillOwnership::Foreign, + (false, false, false) => PiSkillOwnership::Absent, + } + }; + let (discovery_status, discovery_issue) = discovered.cloned().unwrap_or_else(|| { + ( + PiSkillDiscovery::Absent, + destination_exists.then(|| "Pi did not discover this destination".to_string()), + ) + }); + let effectively_discovered = discovery_status == PiSkillDiscovery::Active; + let issue = if has_stale_destination { + Some("recorded Pi deployment remains at a previous agent root".to_string()) + } else { + match ownership { + PiSkillOwnership::Stale => { + Some("recorded Pi deployment no longer matches the live filesystem".to_string()) + } + PiSkillOwnership::Foreign if skill.apps.pi => { + Some("desired Pi Skill collides with an unowned live destination".to_string()) + } + _ => discovery_issue, + } + }; + Ok(SkillAppStatus { + desired_enabled: skill.apps.pi, + owned_deployment, + effectively_discovered, + ownership, + discovery: discovery_status, + issue, + }) +} + +fn deployment_lock() -> std::sync::MutexGuard<'static, ()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())) + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) +} + +fn reconcile_skill_unlocked(db: &Arc, skill: &InstalledSkill) -> Result<(), AppError> { + if skill.apps.pi { + let destination = skill_destination(skill)?; + let destination_key = destination_key(&destination); + let existing = db.get_pi_skill_deployment(&skill.id, &destination_key)?; + let source = skill_source(skill)?; + deploy( + db, + skill, + &source, + &destination, + &destination_key, + existing, + None, + )?; + cleanup_stale_deployments_after_commit(db, skill, &destination_key); + Ok(()) + } else { + remove_all_recorded_deployments(db, skill, None) + } +} + +fn cleanup_stale_deployments( + db: &Arc, + skill: &InstalledSkill, + current_destination_key: &str, +) -> Result<(), AppError> { + for deployment in db.get_pi_skill_deployments(&skill.id)? { + if deployment.destination_key != current_destination_key { + remove_recorded_deployment(db, skill, deployment)?; + } + } + Ok(()) +} + +/// `deploy` atomically commits the new native destination, its ownership +/// ledger, and (when requested) desired state before old-root cleanup begins. +/// A drifted old root must remain visible as stale ownership evidence, but it +/// must not turn that committed publication into an error: several callers +/// compensate a returned error by deleting the Skill's SSOT/database row. +fn cleanup_stale_deployments_after_commit( + db: &Arc, + skill: &InstalledSkill, + current_destination_key: &str, +) { + if let Err(error) = cleanup_stale_deployments(db, skill, current_destination_key) { + log::warn!( + "Pi Skill '{}' was published at its current agent root, but stale deployment cleanup \ + remains pending: {error}", + skill.id + ); + } +} + +fn remove_all_recorded_deployments( + db: &Arc, + skill: &InstalledSkill, + desired_enabled: Option, +) -> Result<(), AppError> { + if let Some(desired_enabled) = desired_enabled { + // Desired state is one row-level authority, independent of how many + // old agent roots still have device-local ownership evidence. + db.set_pi_skill_desired(&skill.id, desired_enabled)?; + } + for deployment in db.get_pi_skill_deployments(&skill.id)? { + remove_recorded_deployment(db, skill, deployment)?; + } + Ok(()) +} + +fn remove_recorded_deployment( + db: &Arc, + skill: &InstalledSkill, + deployment: SkillDeployment, +) -> Result<(), AppError> { + let destination = PathBuf::from(&deployment.destination); + if destination_key(&destination) != deployment.destination_key { + return Err(AppError::Conflict(format!( + "Pi Skill '{}' has inconsistent recorded destination identity", + skill.id + ))); + } + match fs::symlink_metadata(&destination) { + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + db.delete_pi_skill_deployment(&skill.id, &deployment.destination_key)?; + Ok(()) + } + Err(error) => Err(AppError::io(&destination, error)), + Ok(_) => { + let key = deployment.destination_key.clone(); + remove_owned(db, skill, &destination, &key, Some(deployment), None) + } + } +} + +fn skill_source(skill: &InstalledSkill) -> Result { + validate_directory_name(&skill.directory)?; + let source = SkillService::get_ssot_dir() + .map_err(|error| AppError::Config(error.to_string()))? + .join(&skill.directory); + validate_source_tree(&source)?; + Ok(source) +} + +fn skill_destination(skill: &InstalledSkill) -> Result { + validate_directory_name(&skill.directory)?; + Ok(SkillService::get_app_skills_dir(&AppType::Pi) + .map_err(|error| AppError::Config(error.to_string()))? + .join(&skill.directory)) +} + +fn validate_directory_name(value: &str) -> Result<(), AppError> { + let path = Path::new(value); + if value.is_empty() + || path.components().count() != 1 + || matches!( + path.components().next(), + Some(std::path::Component::CurDir | std::path::Component::ParentDir) + ) + || value.starts_with('.') + { + return Err(AppError::InvalidInput(format!( + "invalid Pi Skill directory '{value}'" + ))); + } + Ok(()) +} + +fn destination_key(destination: &Path) -> String { + #[cfg(windows)] + { + destination + .to_string_lossy() + .replace('\\', "/") + .to_lowercase() + } + #[cfg(not(windows))] + { + destination.to_string_lossy().into_owned() + } +} + +fn source_identity(source: &Path) -> Result<(String, String), AppError> { + let canonical = source + .canonicalize() + .map_err(|error| AppError::io(source, error))?; + let digest = tree_digest(source)?; + Ok(( + format!("path:{};digest:{digest}", canonical.display()), + digest, + )) +} + +fn adopt_exact_import( + db: &Arc, + skill: &InstalledSkill, + source: &Path, + destination: &Path, + destination_key: &str, +) -> Result<(), AppError> { + validate_source_tree(source)?; + validate_source_tree(destination)?; + let (source_identity, source_digest) = source_identity(source)?; + let destination_digest = tree_digest(destination)?; + if source_digest != destination_digest { + return Err(AppError::Conflict(format!( + "cannot adopt Pi Skill '{}': native destination differs from the imported SSOT", + skill.directory + ))); + } + + let previous_desired = db + .get_installed_skill(&skill.id)? + .ok_or_else(|| { + AppError::Conflict(format!( + "Pi Skill '{}' disappeared before import adoption", + skill.id + )) + })? + .apps + .pi; + let now = chrono::Utc::now().timestamp_millis(); + let deployment = SkillDeployment { + skill_id: skill.id.clone(), + destination: destination.to_string_lossy().into_owned(), + destination_key: destination_key.to_string(), + method: SkillDeploymentMethod::Copy, + source_identity, + deployed_digest: Some(destination_digest), + created_at: now, + updated_at: now, + }; + db.save_pi_skill_deployment_with_desired(&deployment, Some(true))?; + + let final_verification = verify_owned_destination(&deployment, destination).and_then(|_| { + let current_source_digest = tree_digest(source)?; + if current_source_digest == source_digest { + Ok(()) + } else { + Err(AppError::Conflict(format!( + "cannot adopt Pi Skill '{}': SSOT changed during import", + skill.directory + ))) + } + }); + if let Err(error) = final_verification { + let rollback = db.delete_pi_skill_deployment_with_desired( + &skill.id, + destination_key, + Some(previous_desired), + ); + return match rollback { + Ok(true) => Err(error), + Ok(false) => Err(AppError::Conflict(format!( + "Pi Skill import adoption failed ({error}); ownership rollback found no ledger row" + ))), + Err(rollback) => Err(AppError::Conflict(format!( + "Pi Skill import adoption failed ({error}); ownership rollback failed ({rollback})" + ))), + }; + } + Ok(()) +} + +fn deploy( + db: &Arc, + skill: &InstalledSkill, + source: &Path, + destination: &Path, + destination_key: &str, + existing: Option, + desired_enabled: Option, +) -> Result<(), AppError> { + let (source_identity, source_digest) = source_identity(source)?; + if let Some(existing) = existing.as_ref() { + verify_owned_destination(existing, destination)?; + } else if fs::symlink_metadata(destination).is_ok() { + return Err(AppError::Conflict(format!( + "Pi Skill destination already exists without ownership evidence: {}", + destination.display() + ))); + } + + let requested_method = choose_method(); + let previous = existing.clone(); + let staged_previous = if let Some(previous_deployment) = previous.as_ref() { + let staged = stage_destination(destination)?; + if let Err(error) = verify_deployment_identity(previous_deployment, &staged) { + restore_staged_destination(&staged, destination).map_err(|rollback| { + AppError::Conflict(format!( + "Pi Skill changed while it was staged ({error}); restoring it failed ({rollback})" + )) + })?; + return Err(error); + } + Some(staged) + } else { + None + }; + let method = match replace_destination(source, destination, requested_method) { + Ok(method) => method, + Err(error) => { + if let Some(staged) = staged_previous.as_deref() { + restore_staged_destination(staged, destination).map_err(|rollback| { + AppError::Conflict(format!( + "Pi Skill deployment failed ({error}); previous deployment rollback failed ({rollback})" + )) + })?; + } + return Err(error); + } + }; + let now = chrono::Utc::now().timestamp_millis(); + let deployment = SkillDeployment { + skill_id: skill.id.clone(), + destination: destination.to_string_lossy().into_owned(), + destination_key: destination_key.to_string(), + method, + source_identity, + deployed_digest: (method == SkillDeploymentMethod::Copy).then_some(source_digest), + created_at: previous.as_ref().map_or(now, |value| value.created_at), + updated_at: now, + }; + if let Err(error) = verify_owned_destination(&deployment, destination) { + rollback_verified_replacement(&deployment, destination, staged_previous.as_deref()) + .map_err(|rollback| { + AppError::Conflict(format!( + "Pi Skill deployment identity check failed ({error}); rollback failed ({rollback})" + )) + })?; + return Err(error); + } + if let Err(error) = db.save_pi_skill_deployment_with_desired(&deployment, desired_enabled) { + rollback_verified_replacement(&deployment, destination, staged_previous.as_deref()) + .map_err(|rollback| { + AppError::Conflict(format!( + "Pi Skill ledger write failed ({error}); deployment rollback failed ({rollback})" + )) + })?; + return Err(error); + } + if let Some(staged) = staged_previous { + let previous_deployment = previous.as_ref().ok_or_else(|| { + AppError::Config( + "Pi Skill replacement staging lost its previous ownership record".to_string(), + ) + })?; + verify_deployment_identity(previous_deployment, &staged)?; + if let Err(error) = remove_path(&staged) { + log::warn!( + "failed to remove committed Pi Skill rollback staging '{}': {error}", + staged.display() + ); + } + } + Ok(()) +} + +fn remove_owned( + db: &Arc, + skill: &InstalledSkill, + destination: &Path, + destination_key: &str, + existing: Option, + desired_enabled: Option, +) -> Result<(), AppError> { + if let Some(desired_enabled) = desired_enabled { + // Persist user intent before filesystem validation or cleanup. Drift + // therefore reports a conflict while the toggle remains off and the + // ledger is retained as deletion evidence. + db.set_pi_skill_desired(&skill.id, desired_enabled)?; + } + let Some(existing) = existing else { + // Foreign/native discovered directories are preserved. + return Ok(()); + }; + verify_owned_destination(&existing, destination)?; + let staged = stage_destination(destination)?; + if let Err(error) = verify_deployment_identity(&existing, &staged) { + restore_staged_destination(&staged, destination).map_err(|rollback| { + AppError::Conflict(format!( + "Pi Skill changed while it was staged for removal ({error}); restoring it failed ({rollback})" + )) + })?; + return Err(error); + } + if let Err(error) = db.delete_pi_skill_deployment(&skill.id, destination_key) { + restore_staged_destination(&staged, destination).map_err(|rollback| { + AppError::Conflict(format!( + "Pi Skill ledger cleanup failed ({error}); file rollback failed ({rollback})" + )) + })?; + return Err(error); + } + verify_deployment_identity(&existing, &staged)?; + if let Err(error) = remove_path(&staged) { + log::warn!( + "failed to remove disabled Pi Skill rollback staging '{}': {error}", + staged.display() + ); + } + if fs::symlink_metadata(destination).is_ok() { + return Err(AppError::Conflict(format!( + "Pi Skill destination was recreated concurrently after ownership removal: {}", + destination.display() + ))); + } + Ok(()) +} + +fn choose_method() -> SkillDeploymentMethod { + match crate::settings::get_skill_sync_method() { + SyncMethod::Copy => SkillDeploymentMethod::Copy, + SyncMethod::Symlink | SyncMethod::Auto => SkillDeploymentMethod::Symlink, + } +} + +fn verify_owned_destination( + deployment: &SkillDeployment, + destination: &Path, +) -> Result<(), AppError> { + if Path::new(&deployment.destination) != destination { + return Err(AppError::Conflict( + "Pi Skill deployment destination changed since it was recorded".to_string(), + )); + } + verify_deployment_identity(deployment, destination) +} + +fn verify_deployment_identity( + deployment: &SkillDeployment, + destination: &Path, +) -> Result<(), AppError> { + match deployment.method { + SkillDeploymentMethod::Symlink => { + let metadata = fs::symlink_metadata(destination) + .map_err(|error| AppError::io(destination, error))?; + if !metadata.file_type().is_symlink() { + return Err(AppError::Conflict(format!( + "owned Pi Skill symlink was replaced externally: {}", + destination.display() + ))); + } + let target = + fs::read_link(destination).map_err(|error| AppError::io(destination, error))?; + let resolved = if target.is_absolute() { + target + } else { + destination + .parent() + .unwrap_or_else(|| Path::new(".")) + .join(target) + }; + let canonical = resolved + .canonicalize() + .map_err(|error| AppError::io(&resolved, error))?; + if !deployment + .source_identity + .starts_with(&format!("path:{};", canonical.display())) + { + return Err(AppError::Conflict(format!( + "owned Pi Skill symlink target changed externally: {}", + destination.display() + ))); + } + } + SkillDeploymentMethod::Copy => { + let expected = deployment.deployed_digest.as_deref().ok_or_else(|| { + AppError::Conflict("copied Pi Skill lacks a recorded digest".to_string()) + })?; + if tree_digest(destination)? != expected { + return Err(AppError::Conflict(format!( + "owned Pi Skill copy was modified externally: {}", + destination.display() + ))); + } + } + } + Ok(()) +} + +fn replace_destination( + source: &Path, + destination: &Path, + method: SkillDeploymentMethod, +) -> Result { + if fs::symlink_metadata(destination).is_ok() { + return Err(AppError::Conflict(format!( + "Pi Skill replacement destination is not empty: {}", + destination.display() + ))); + } + let parent = destination + .parent() + .ok_or_else(|| AppError::InvalidInput("Pi Skill destination has no parent".to_string()))?; + fs::create_dir_all(parent).map_err(|error| AppError::io(parent, error))?; + match method { + SkillDeploymentMethod::Symlink => match create_directory_symlink(source, destination) { + Ok(()) => Ok(SkillDeploymentMethod::Symlink), + Err(_) if matches!(crate::settings::get_skill_sync_method(), SyncMethod::Auto) => { + copy_tree_atomic(source, destination)?; + Ok(SkillDeploymentMethod::Copy) + } + Err(error) => Err(error), + }, + SkillDeploymentMethod::Copy => { + copy_tree_atomic(source, destination)?; + Ok(SkillDeploymentMethod::Copy) + } + } +} + +fn stage_destination(destination: &Path) -> Result { + let parent = destination + .parent() + .ok_or_else(|| AppError::InvalidInput("Pi Skill destination has no parent".to_string()))?; + let file_name = destination + .file_name() + .and_then(|name| name.to_str()) + .ok_or_else(|| { + AppError::InvalidInput("Pi Skill destination name is invalid".to_string()) + })?; + let staged = parent.join(format!( + ".{file_name}.cc-switch-rollback-{}", + uuid::Uuid::new_v4().simple() + )); + fs::rename(destination, &staged).map_err(|error| AppError::io(destination, error))?; + Ok(staged) +} + +fn restore_staged_destination(staged: &Path, destination: &Path) -> Result<(), AppError> { + if fs::symlink_metadata(destination).is_ok() { + return Err(AppError::Conflict(format!( + "refusing to overwrite a concurrently created Pi Skill destination: {}", + destination.display() + ))); + } + fs::rename(staged, destination).map_err(|error| AppError::io(destination, error)) +} + +fn rollback_verified_replacement( + deployment: &SkillDeployment, + destination: &Path, + staged_previous: Option<&Path>, +) -> Result<(), AppError> { + let staged_replacement = stage_destination(destination)?; + if let Err(error) = verify_deployment_identity(deployment, &staged_replacement) { + restore_staged_destination(&staged_replacement, destination).map_err(|rollback| { + AppError::Conflict(format!( + "replacement ownership was lost before rollback ({error}); preserving it also failed ({rollback})" + )) + })?; + return Err(error); + } + remove_path(&staged_replacement)?; + if let Some(staged) = staged_previous { + restore_staged_destination(staged, destination)?; + } + Ok(()) +} + +#[cfg(unix)] +fn create_directory_symlink(source: &Path, destination: &Path) -> Result<(), AppError> { + std::os::unix::fs::symlink(source, destination) + .map_err(|error| AppError::io(destination, error)) +} + +#[cfg(windows)] +fn create_directory_symlink(source: &Path, destination: &Path) -> Result<(), AppError> { + std::os::windows::fs::symlink_dir(source, destination) + .map_err(|error| AppError::io(destination, error)) +} + +fn copy_tree_atomic(source: &Path, destination: &Path) -> Result<(), AppError> { + let parent = destination + .parent() + .ok_or_else(|| AppError::InvalidInput("Pi Skill destination has no parent".to_string()))?; + let temp = parent.join(format!(".pi-skill-{}.tmp", uuid::Uuid::new_v4().simple())); + let result = copy_tree(source, &temp).and_then(|_| { + fs::rename(&temp, destination).map_err(|error| AppError::io(destination, error)) + }); + if result.is_err() { + let _ = fs::remove_dir_all(&temp); + } + result +} + +fn copy_tree(source: &Path, destination: &Path) -> Result<(), AppError> { + validate_source_tree(source)?; + fs::create_dir(destination).map_err(|error| AppError::io(destination, error))?; + for entry in fs::read_dir(source).map_err(|error| AppError::io(source, error))? { + let entry = entry.map_err(|error| AppError::io(source, error))?; + let file_type = entry + .file_type() + .map_err(|error| AppError::io(entry.path(), error))?; + let target = destination.join(entry.file_name()); + if file_type.is_dir() { + copy_tree(&entry.path(), &target)?; + } else if file_type.is_file() { + fs::copy(entry.path(), &target).map_err(|error| AppError::io(&target, error))?; + } else { + return Err(AppError::InvalidInput(format!( + "Pi Skill source contains a non-regular entry: {}", + entry.path().display() + ))); + } + } + Ok(()) +} + +fn validate_source_tree(source: &Path) -> Result<(), AppError> { + let metadata = fs::symlink_metadata(source).map_err(|error| AppError::io(source, error))?; + if !metadata.file_type().is_dir() || !source.join("SKILL.md").is_file() { + return Err(AppError::InvalidInput(format!( + "Pi Skill source must be a directory containing SKILL.md: {}", + source.display() + ))); + } + Ok(()) +} + +fn tree_digest(root: &Path) -> Result { + let mut entries = Vec::new(); + collect_digest_entries(root, root, &mut entries)?; + entries.sort_by(|left, right| left.0.cmp(&right.0)); + let mut hasher = Sha256::new(); + for (relative, bytes) in entries { + hasher.update((relative.len() as u64).to_le_bytes()); + hasher.update(relative.as_bytes()); + hasher.update((bytes.len() as u64).to_le_bytes()); + hasher.update(bytes); + } + Ok(format!("sha256:{:x}", hasher.finalize())) +} + +fn collect_digest_entries( + root: &Path, + current: &Path, + output: &mut Vec<(String, Vec)>, +) -> Result<(), AppError> { + for entry in fs::read_dir(current).map_err(|error| AppError::io(current, error))? { + let entry = entry.map_err(|error| AppError::io(current, error))?; + let path = entry.path(); + let kind = entry + .file_type() + .map_err(|error| AppError::io(&path, error))?; + if kind.is_dir() { + collect_digest_entries(root, &path, output)?; + } else if kind.is_file() { + let relative = path + .strip_prefix(root) + .map_err(|_| AppError::Config("Pi Skill path escaped its root".to_string()))? + .to_string_lossy() + .replace('\\', "/"); + output.push(( + relative, + fs::read(&path).map_err(|error| AppError::io(&path, error))?, + )); + } else { + return Err(AppError::InvalidInput(format!( + "Pi Skill tree contains a symlink or special entry: {}", + path.display() + ))); + } + } + Ok(()) +} + +fn remove_path(path: &Path) -> Result<(), AppError> { + let metadata = match fs::symlink_metadata(path) { + Ok(metadata) => metadata, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(()), + Err(error) => return Err(AppError::io(path, error)), + }; + if metadata.file_type().is_symlink() || metadata.file_type().is_file() { + fs::remove_file(path).map_err(|error| AppError::io(path, error)) + } else if metadata.file_type().is_dir() { + fs::remove_dir_all(path).map_err(|error| AppError::io(path, error)) + } else { + Err(AppError::InvalidInput(format!( + "refusing to remove special Pi Skill entry: {}", + path.display() + ))) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::app_config::SkillApps; + + #[test] + fn digest_includes_hidden_files_and_rejects_symlinks() { + let temp = tempfile::tempdir().expect("tempdir"); + fs::write(temp.path().join("SKILL.md"), "skill").expect("manifest"); + fs::write(temp.path().join(".hidden"), "one").expect("hidden"); + let first = tree_digest(temp.path()).expect("digest"); + fs::write(temp.path().join(".hidden"), "two").expect("hidden update"); + assert_ne!(tree_digest(temp.path()).expect("digest"), first); + } + + #[test] + fn rollback_preserves_a_destination_without_matching_ownership_evidence() { + let temp = tempfile::tempdir().expect("tempdir"); + let destination = temp.path().join("skill"); + fs::create_dir(&destination).expect("foreign destination"); + fs::write(destination.join("SKILL.md"), "foreign").expect("foreign manifest"); + let deployment = SkillDeployment { + skill_id: "skill".to_string(), + destination: destination.to_string_lossy().into_owned(), + destination_key: destination_key(&destination), + method: SkillDeploymentMethod::Copy, + source_identity: "path:/managed;digest:sha256:managed".to_string(), + deployed_digest: Some("sha256:not-the-foreign-tree".to_string()), + created_at: 1, + updated_at: 1, + }; + + assert!(rollback_verified_replacement(&deployment, &destination, None).is_err()); + assert_eq!( + fs::read_to_string(destination.join("SKILL.md")).expect("foreign content survives"), + "foreign" + ); + } + + #[test] + fn disabling_a_drifted_owned_skill_persists_intent_and_keeps_ledger() { + let temp = tempfile::tempdir().expect("tempdir"); + let destination = temp.path().join("skill"); + fs::create_dir(&destination).expect("destination"); + fs::write( + destination.join("SKILL.md"), + "---\nname: skill\ndescription: before\n---\n", + ) + .expect("manifest"); + let original_digest = tree_digest(&destination).expect("original digest"); + let mut skill = InstalledSkill { + id: "local:skill".to_string(), + name: "Skill".to_string(), + description: Some("before".to_string()), + directory: "skill".to_string(), + repo_owner: None, + repo_name: None, + repo_branch: None, + readme_url: None, + apps: SkillApps::only(&AppType::Pi), + installed_at: 1, + content_hash: Some(original_digest.clone()), + updated_at: 1, + }; + let db = Arc::new(Database::memory().expect("database")); + db.save_skill(&skill).expect("save skill"); + let key = destination_key(&destination); + db.save_pi_skill_deployment(&SkillDeployment { + skill_id: skill.id.clone(), + destination: destination.to_string_lossy().into_owned(), + destination_key: key.clone(), + method: SkillDeploymentMethod::Copy, + source_identity: format!("path:{};digest:{original_digest}", destination.display()), + deployed_digest: Some(original_digest), + created_at: 1, + updated_at: 1, + }) + .expect("save ledger"); + + fs::write(destination.join("changed.txt"), "external drift").expect("drift"); + let error = remove_owned( + &db, + &skill, + &destination, + &key, + db.get_pi_skill_deployment(&skill.id, &key) + .expect("read ledger"), + Some(false), + ) + .expect_err("drift must block deletion"); + assert!(matches!(error, AppError::Conflict(_))); + + skill = db + .get_installed_skill(&skill.id) + .expect("read skill") + .expect("skill remains"); + assert!(!skill.apps.pi, "desired state must remain disabled"); + assert!( + db.get_pi_skill_deployment(&skill.id, &key) + .expect("read ledger") + .is_some(), + "drift evidence must remain for explicit resolution" + ); + assert!(destination.join("changed.txt").is_file()); + } + + #[test] + #[serial_test::serial] + fn discovery_rejects_a_manifest_without_pinned_required_description() { + struct EnvGuard(Option); + impl Drop for EnvGuard { + fn drop(&mut self) { + match self.0.take() { + Some(value) => std::env::set_var("PI_CODING_AGENT_DIR", value), + None => std::env::remove_var("PI_CODING_AGENT_DIR"), + } + } + } + + let temp = tempfile::tempdir().expect("tempdir"); + let _guard = EnvGuard(std::env::var_os("PI_CODING_AGENT_DIR")); + std::env::set_var("PI_CODING_AGENT_DIR", temp.path()); + let manifest = temp.path().join("skills").join("invalid").join("SKILL.md"); + fs::create_dir_all(manifest.parent().expect("manifest parent")).expect("skills dir"); + fs::write(&manifest, "---\nname: invalid\n---\n").expect("manifest"); + + let discovery = scan_pi_discovery().expect("scan"); + assert_eq!( + discovery.by_manifest[&manifest].0, + PiSkillDiscovery::Invalid + ); + assert!(discovery.by_manifest[&manifest] + .1 + .as_deref() + .is_some_and(|issue| issue.contains("description"))); + } + + #[test] + #[serial_test::serial] + fn agent_root_relocation_deploys_new_destination_before_cleaning_old_ownership() { + struct EnvGuard { + key: &'static str, + previous: Option, + } + impl EnvGuard { + fn set(key: &'static str, value: &Path) -> Self { + let previous = std::env::var_os(key); + std::env::set_var(key, value); + Self { key, previous } + } + } + impl Drop for EnvGuard { + fn drop(&mut self) { + match self.previous.take() { + Some(value) => std::env::set_var(self.key, value), + None => std::env::remove_var(self.key), + } + if self.key == "CC_SWITCH_TEST_HOME" { + let _ = crate::settings::reload_settings(); + } + } + } + + let temp = tempfile::tempdir().expect("tempdir"); + let _home = EnvGuard::set("CC_SWITCH_TEST_HOME", temp.path()); + crate::settings::reload_settings().expect("reload settings"); + let old_root = temp.path().join("old-pi"); + let new_root = temp.path().join("new-pi"); + let _pi_root = EnvGuard::set("PI_CODING_AGENT_DIR", &old_root); + + let source = SkillService::get_ssot_dir() + .expect("SSOT") + .join("relocated"); + fs::create_dir_all(&source).expect("source"); + fs::write( + source.join("SKILL.md"), + "---\nname: relocated\ndescription: relocation test\n---\n", + ) + .expect("manifest"); + let skill = InstalledSkill { + id: "local:relocated".to_string(), + name: "Relocated".to_string(), + description: Some("relocation test".to_string()), + directory: "relocated".to_string(), + repo_owner: None, + repo_name: None, + repo_branch: None, + readme_url: None, + apps: SkillApps::only(&AppType::Pi), + installed_at: 1, + content_hash: None, + updated_at: 1, + }; + let db = Arc::new(Database::memory().expect("database")); + db.save_skill(&skill).expect("save skill"); + PiSkillDeploymentService::reconcile_skill(&db, &skill).expect("old deployment"); + let old_destination = old_root.join("skills").join("relocated"); + assert!(fs::symlink_metadata(&old_destination).is_ok()); + + std::env::set_var("PI_CODING_AGENT_DIR", &new_root); + PiSkillDeploymentService::reconcile_skill(&db, &skill).expect("relocate deployment"); + let new_destination = new_root.join("skills").join("relocated"); + assert!(fs::symlink_metadata(&new_destination).is_ok()); + assert!( + fs::symlink_metadata(&old_destination).is_err(), + "the verified old deployment must be cleaned only after new publication" + ); + let deployments = db + .get_pi_skill_deployments(&skill.id) + .expect("deployment ledger"); + assert_eq!(deployments.len(), 1); + assert_eq!( + deployments[0].destination_key, + destination_key(&new_destination) + ); + } + + #[test] + #[serial_test::serial] + fn drifted_old_root_is_post_commit_status_not_a_failed_new_publication() { + struct EnvGuard { + key: &'static str, + previous: Option, + } + impl EnvGuard { + fn set(key: &'static str, value: &Path) -> Self { + let previous = std::env::var_os(key); + std::env::set_var(key, value); + Self { key, previous } + } + } + impl Drop for EnvGuard { + fn drop(&mut self) { + match self.previous.take() { + Some(value) => std::env::set_var(self.key, value), + None => std::env::remove_var(self.key), + } + if self.key == "CC_SWITCH_TEST_HOME" { + let _ = crate::settings::reload_settings(); + } + } + } + + let temp = tempfile::tempdir().expect("tempdir"); + let _home = EnvGuard::set("CC_SWITCH_TEST_HOME", temp.path()); + crate::settings::reload_settings().expect("reload settings"); + let old_root = temp.path().join("old-pi"); + let new_root = temp.path().join("new-pi"); + let _pi_root = EnvGuard::set("PI_CODING_AGENT_DIR", &old_root); + let source = SkillService::get_ssot_dir() + .expect("SSOT") + .join("drifted-relocation"); + fs::create_dir_all(&source).expect("source"); + fs::write( + source.join("SKILL.md"), + "---\nname: drifted-relocation\ndescription: relocation test\n---\n", + ) + .expect("manifest"); + let mut skill = InstalledSkill { + id: "local:drifted-relocation".to_string(), + name: "Drifted relocation".to_string(), + description: Some("relocation test".to_string()), + directory: "drifted-relocation".to_string(), + repo_owner: None, + repo_name: None, + repo_branch: None, + readme_url: None, + apps: SkillApps::only(&AppType::Pi), + installed_at: 1, + content_hash: None, + updated_at: 1, + }; + let db = Arc::new(Database::memory().expect("database")); + db.save_skill(&skill).expect("save skill"); + PiSkillDeploymentService::reconcile_skill(&db, &skill).expect("old deployment"); + let old_destination = old_root.join("skills").join(&skill.directory); + remove_path(&old_destination).expect("replace old deployment"); + fs::create_dir_all(&old_destination).expect("foreign old destination"); + fs::write(old_destination.join("external.txt"), "external").expect("external drift"); + + std::env::set_var("PI_CODING_AGENT_DIR", &new_root); + PiSkillDeploymentService::toggle(&db, &mut skill, true) + .expect("the new publication is already committed"); + + let new_destination = new_root.join("skills").join(&skill.directory); + assert!(fs::symlink_metadata(&new_destination).is_ok()); + assert_eq!( + fs::read_to_string(old_destination.join("external.txt")).expect("external survives"), + "external" + ); + let stored = db + .get_installed_skill(&skill.id) + .expect("read desired state") + .expect("skill remains"); + assert!(stored.apps.pi); + assert_eq!( + db.get_pi_skill_deployments(&skill.id) + .expect("ledger") + .len(), + 2, + "old drift evidence and the committed new deployment must both remain" + ); + let status = + PiSkillDeploymentService::inspect_all(&db).expect("inspect")[&skill.id].clone(); + assert_eq!(status.ownership, PiSkillOwnership::Stale); + assert!(status.effectively_discovered); + assert!(status.issue.is_some_and(|issue| issue.contains("previous"))); + } +} diff --git a/src-tauri/src/services/sql_helpers.rs b/src-tauri/src/services/sql_helpers.rs index 6a582f90f..e6de52935 100644 --- a/src-tauri/src/services/sql_helpers.rs +++ b/src-tauri/src/services/sql_helpers.rs @@ -54,8 +54,7 @@ pub fn fresh_input_sql(alias: &str) -> String { format!( "CASE \ WHEN {prefix}input_token_semantics = {INPUT_TOKEN_SEMANTICS_FRESH} THEN {prefix}input_tokens \ - WHEN {prefix}app_type IN ({app_type_list}) \ - AND {prefix}input_token_semantics = {INPUT_TOKEN_SEMANTICS_TOTAL} \ + WHEN {prefix}input_token_semantics = {INPUT_TOKEN_SEMANTICS_TOTAL} \ AND {prefix}input_tokens >= ({prefix}cache_read_tokens + {prefix}cache_creation_tokens) \ THEN ({prefix}input_tokens - {prefix}cache_read_tokens - {prefix}cache_creation_tokens) \ WHEN {prefix}app_type IN ({app_type_list}) \ @@ -144,6 +143,40 @@ mod tests { assert_eq!(total, 400 + 500 + 450 + 200); } + #[test] + fn stored_wire_semantics_override_logical_pi_app_type() { + let conn = setup_conn(); + conn.execute( + "INSERT INTO proxy_request_logs ( + request_id, app_type, input_tokens, cache_read_tokens, + cache_creation_tokens, input_token_semantics + ) VALUES + ('pi-openai', 'pi', 1000, 700, 100, 1), + ('pi-anthropic', 'pi', 1000, 700, 100, 2)", + [], + ) + .unwrap(); + + let sql = format!( + "SELECT request_id, {} FROM proxy_request_logs ORDER BY request_id", + fresh_input_sql("") + ); + let values: Vec<(String, i64)> = conn + .prepare(&sql) + .unwrap() + .query_map([], |row| Ok((row.get(0)?, row.get(1)?))) + .unwrap() + .collect::>() + .unwrap(); + assert_eq!( + values, + vec![ + ("pi-anthropic".to_string(), 1000), + ("pi-openai".to_string(), 200), + ] + ); + } + #[test] fn fresh_input_handles_codex_with_cache_exceeding_input() { // Defensive: if a malformed Codex row somehow has cache > input, diff --git a/src-tauri/src/services/stream_check.rs b/src-tauri/src/services/stream_check.rs index 5db28335b..b928efbd6 100644 --- a/src-tauri/src/services/stream_check.rs +++ b/src-tauri/src/services/stream_check.rs @@ -181,6 +181,7 @@ impl StreamCheckService { } AppType::OpenClaw => Self::extract_openclaw_base_url(provider), AppType::Hermes => Self::extract_hermes_base_url(provider), + AppType::Pi => Self::extract_pi_base_url(provider), AppType::ClaudeDesktop => ClaudeAdapter::new() .extract_base_url(provider) .map_err(|e| AppError::Message(format!("Failed to extract base_url: {e}"))), @@ -323,6 +324,31 @@ impl StreamCheckService { }) } + /// Pi endpoint inheritance is owned by the pinned composer. Reachability + /// checks deliberately consume its first effective model instead of + /// reimplementing provider/model fallback rules. + fn extract_pi_base_url(provider: &Provider) -> Result { + let config: crate::pi_config::model::PiManagedProviderConfig = + serde_json::from_value(provider.settings_config.clone()).map_err(|error| { + AppError::InvalidInput(format!( + "Pi provider '{}' is not a managed native configuration: {error}", + provider.id + )) + })?; + let composition = + crate::pi_config::native::compose_managed_pi_provider(&provider.id, &config)?; + composition + .models + .first() + .map(|model| model.base_url.clone()) + .ok_or_else(|| { + AppError::InvalidInput(format!( + "Pi provider '{}' has no effective models", + provider.id + )) + }) + } + /// OpenCode: `{ npm, options: { baseURL, apiKey }, ... }` /// /// 用户未显式填 `options.baseURL` 时,按 `npm`(AI SDK 包)回退到包自带默认端点。 @@ -501,6 +527,36 @@ mod tests { ); } + #[test] + fn pi_reachability_uses_composer_effective_model_endpoint() { + let provider = make_provider(serde_json::json!({ + "name": "Pi", + "api": "openai-responses", + "baseUrl": "https://provider.example/v1", + "apiKey": "literal", + "models": [{ + "id": "model", + "name": "Model", + "baseUrl": "https://model.example/custom", + "reasoning": false, + "input": ["text"], + "cost": { + "input": 0, + "output": 0, + "cacheRead": 0, + "cacheWrite": 0 + }, + "contextWindow": 128000, + "maxTokens": 8192 + }] + })); + + assert_eq!( + StreamCheckService::resolve_base_url(&AppType::Pi, &provider).unwrap(), + "https://model.example/custom" + ); + } + #[test] fn test_resolve_base_url_uses_explicit_url_or_errors_when_missing() { // 有显式 base_url → 直接用 diff --git a/src-tauri/src/services/usage_stats.rs b/src-tauri/src/services/usage_stats.rs index 326cd2837..9db6b10a7 100644 --- a/src-tauri/src/services/usage_stats.rs +++ b/src-tauri/src/services/usage_stats.rs @@ -137,8 +137,9 @@ pub struct RequestLogDetail { pub output_tokens: u32, pub cache_read_tokens: u32, pub cache_creation_tokens: u32, - /// Internal storage semantics; omitted from the UI/API payload. - #[serde(skip)] + /// Persisted request-level semantics used by both pricing and UI cache + /// normalization. This must cross IPC; app-type inference is only a legacy + /// fallback for rows written before the semantics column existed. pub input_token_semantics: i64, pub input_cost_usd: String, pub output_cost_usd: String, @@ -1653,10 +1654,10 @@ impl Database { let detail_sql = format!( "SELECT l.request_id, l.provider_id, {detail_pname} as provider_name, l.app_type, l.model, l.request_model, l.cost_multiplier, - input_tokens, output_tokens, cache_read_tokens, cache_creation_tokens, - input_cost_usd, output_cost_usd, cache_read_cost_usd, cache_creation_cost_usd, total_cost_usd, - is_streaming, latency_ms, first_token_ms, duration_ms, - status_code, error_message, created_at, l.data_source, l.pricing_model, + l.input_tokens, l.output_tokens, l.cache_read_tokens, l.cache_creation_tokens, + l.input_cost_usd, l.output_cost_usd, l.cache_read_cost_usd, l.cache_creation_cost_usd, l.total_cost_usd, + l.is_streaming, l.latency_ms, l.first_token_ms, l.duration_ms, + l.status_code, l.error_message, l.created_at, l.data_source, l.pricing_model, l.input_token_semantics FROM proxy_request_logs l LEFT JOIN providers p ON l.provider_id = p.id AND l.app_type = p.app_type @@ -1897,19 +1898,18 @@ impl Database { // 1. 历史 cache-inclusive 行只包含 cache read;新 total 行还包含 cache write。 // 2. Claude/Anthropic 的 input_tokens 已经是 fresh input,不能再次扣减 // 3. 各项成本是基础成本(不含倍率),倍率只作用于最终总价 - let cache_inclusive_app = - crate::services::sql_helpers::is_cache_inclusive_app(log.app_type.as_str()); - let billable_input_tokens = - if !cache_inclusive_app || log.input_token_semantics == INPUT_TOKEN_SEMANTICS_FRESH { - log.input_tokens as u64 - } else if log.input_token_semantics == INPUT_TOKEN_SEMANTICS_TOTAL { - (log.input_tokens as u64) - .saturating_sub(log.cache_read_tokens as u64) - .saturating_sub(log.cache_creation_tokens as u64) - } else { - // v12 and earlier: input included cache reads but excluded cache writes. - (log.input_tokens as u64).saturating_sub(log.cache_read_tokens as u64) - }; + let billable_input_tokens = if log.input_token_semantics == INPUT_TOKEN_SEMANTICS_FRESH { + log.input_tokens as u64 + } else if log.input_token_semantics == INPUT_TOKEN_SEMANTICS_TOTAL { + (log.input_tokens as u64) + .saturating_sub(log.cache_read_tokens as u64) + .saturating_sub(log.cache_creation_tokens as u64) + } else if crate::services::sql_helpers::is_cache_inclusive_app(log.app_type.as_str()) { + // v12 and earlier: input included cache reads but excluded cache writes. + (log.input_tokens as u64).saturating_sub(log.cache_read_tokens as u64) + } else { + log.input_tokens as u64 + }; let input_cost = rust_decimal::Decimal::from(billable_input_tokens) * pricing.input / million; let output_cost = @@ -2407,6 +2407,54 @@ mod tests { Ok(()) } + #[test] + fn paginated_and_detail_ipc_serialize_persisted_input_semantics() -> Result<(), AppError> { + let db = Database::memory()?; + { + let conn = lock_conn!(db.conn); + insert_usage_log( + &conn, + "pi-semantics-ipc", + "pi", + "pi-provider", + "gpt-test", + "request", + 1, + 1_000, + 5, + 800, + 0, + 200, + "0", + )?; + conn.execute( + "UPDATE proxy_request_logs + SET input_token_semantics = ?1 + WHERE request_id = 'pi-semantics-ipc'", + [INPUT_TOKEN_SEMANTICS_TOTAL], + )?; + } + + let page = db.get_request_logs(&LogFilters::default(), 0, 10)?; + let page_json = + serde_json::to_value(&page).map_err(|error| AppError::Database(error.to_string()))?; + assert_eq!( + page_json["data"][0]["inputTokenSemantics"], + INPUT_TOKEN_SEMANTICS_TOTAL + ); + + let detail = db + .get_request_detail("pi-semantics-ipc")? + .expect("request detail"); + let detail_json = + serde_json::to_value(detail).map_err(|error| AppError::Database(error.to_string()))?; + assert_eq!( + detail_json["inputTokenSemantics"], + INPUT_TOKEN_SEMANTICS_TOTAL + ); + Ok(()) + } + fn create_legacy_nullable_logs_table(conn: &Connection) -> Result<(), AppError> { conn.execute( "CREATE TABLE proxy_request_logs ( diff --git a/src-tauri/src/session_manager/mod.rs b/src-tauri/src/session_manager/mod.rs index 839e94fb0..188c4b055 100644 --- a/src-tauri/src/session_manager/mod.rs +++ b/src-tauri/src/session_manager/mod.rs @@ -4,7 +4,7 @@ pub mod terminal; use serde::{Deserialize, Serialize}; use std::path::{Path, PathBuf}; -use providers::{claude, codex, gemini, grokbuild, hermes, openclaw, opencode}; +use providers::{claude, codex, gemini, grokbuild, hermes, openclaw, opencode, pi}; #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] @@ -56,7 +56,7 @@ pub struct DeleteSessionOutcome { } pub fn scan_sessions() -> Vec { - let (r1, r2, r3, r4, r5, r6, r7) = std::thread::scope(|s| { + let (r1, r2, r3, r4, r5, r6, r7, r8) = std::thread::scope(|s| { let h1 = s.spawn(codex::scan_sessions); let h2 = s.spawn(claude::scan_sessions); let h3 = s.spawn(opencode::scan_sessions); @@ -64,6 +64,7 @@ pub fn scan_sessions() -> Vec { let h5 = s.spawn(gemini::scan_sessions); let h6 = s.spawn(hermes::scan_sessions); let h7 = s.spawn(grokbuild::scan_sessions); + let h8 = s.spawn(pi::scan_sessions); ( h1.join().unwrap_or_default(), h2.join().unwrap_or_default(), @@ -72,6 +73,7 @@ pub fn scan_sessions() -> Vec { h5.join().unwrap_or_default(), h6.join().unwrap_or_default(), h7.join().unwrap_or_default(), + h8.join().unwrap_or_default(), ) }); @@ -83,6 +85,7 @@ pub fn scan_sessions() -> Vec { sessions.extend(r5); sessions.extend(r6); sessions.extend(r7); + sessions.extend(r8); sessions.sort_by(|a, b| { let a_ts = a.last_active_at.or(a.created_at).unwrap_or(0); @@ -111,6 +114,7 @@ pub fn load_messages(provider_id: &str, source_path: &str) -> Result gemini::load_messages(path), "grokbuild" => grokbuild::load_messages(path), "hermes" => hermes::load_messages(path), + "pi" => pi::load_messages(path), _ => Err(format!("Unsupported provider: {provider_id}")), } } @@ -173,6 +177,7 @@ fn delete_session_with_roots( grokbuild::delete_session(&validated_root, &validated_source, session_id) } "hermes" => hermes::delete_session(&validated_root, &validated_source, session_id), + "pi" => pi::delete_session(&validated_root, &validated_source, session_id), _ => Err(format!("Unsupported provider: {provider_id}")), }; } @@ -203,6 +208,7 @@ fn provider_roots(provider_id: &str) -> Result, String> { "gemini" => vec![crate::gemini_config::get_gemini_dir().join("tmp")], "grokbuild" => grokbuild::session_roots(), "hermes" => vec![crate::hermes_config::get_hermes_dir().join("sessions")], + "pi" => pi::session_roots(), _ => return Err(format!("Unsupported provider: {provider_id}")), }; diff --git a/src-tauri/src/session_manager/providers/mod.rs b/src-tauri/src/session_manager/providers/mod.rs index 6a16fb5a8..1cd3fd1b5 100644 --- a/src-tauri/src/session_manager/providers/mod.rs +++ b/src-tauri/src/session_manager/providers/mod.rs @@ -5,4 +5,5 @@ pub mod grokbuild; pub mod hermes; pub mod openclaw; pub mod opencode; +pub mod pi; mod utils; diff --git a/src-tauri/src/session_manager/providers/pi.rs b/src-tauri/src/session_manager/providers/pi.rs new file mode 100644 index 000000000..5b19715c9 --- /dev/null +++ b/src-tauri/src/session_manager/providers/pi.rs @@ -0,0 +1,731 @@ +use std::collections::{HashMap, HashSet}; +use std::fs::{self, File}; +use std::io::{BufRead, BufReader}; +use std::path::{Path, PathBuf}; + +use serde::Serialize; +use serde_json::Value; + +use crate::session_manager::{SessionMessage, SessionMeta}; + +use super::utils::{ + extract_text, parse_timestamp_to_ms, path_basename, truncate_summary, TITLE_MAX_CHARS, +}; + +const PROVIDER_ID: &str = "pi"; +const MAX_TREE_ENTRIES: usize = 500_000; +const MAX_TREE_ID_BYTES: usize = 256; +const MAX_SCAN_DEPTH: usize = 8; +const MAX_SESSION_BYTES: u64 = 128 * 1024 * 1024; + +#[derive(Debug, PartialEq, Eq)] +enum SessionRootResolution { + Available { + root: PathBuf, + source: &'static str, + }, + RequiresProjectContext { + configured_path: String, + source: &'static str, + }, + Unavailable { + reason: String, + }, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +#[serde(tag = "status", rename_all = "snake_case")] +pub enum PiSessionDiscovery { + Available { + root: String, + source: &'static str, + }, + RequiresProjectContext { + #[serde(rename = "configuredPath")] + configured_path: String, + source: &'static str, + }, + Unavailable { + reason: String, + }, +} + +#[derive(Debug)] +struct SessionHeader { + id: String, + cwd: String, + timestamp: Option, + version: u64, +} + +#[derive(Debug)] +struct SessionTree { + header: SessionHeader, + active_ids: HashSet, +} + +#[derive(Default)] +struct ActiveSessionData { + messages: Vec, + first_user_message: Option, + last_message: Option, + explicit_name: Option>, + last_active_at: Option, +} + +/// Pi keeps a relative `sessionDir` relative through SessionManager creation; +/// its file operations therefore depend on the launching process cwd. A global +/// session browser has no authoritative launch cwd, so relative values are +/// deliberately non-enumerable and never fall back to another root. +pub fn session_roots() -> Vec { + match resolve_session_root() { + SessionRootResolution::Available { root, .. } => vec![root], + SessionRootResolution::RequiresProjectContext { .. } + | SessionRootResolution::Unavailable { .. } => Vec::new(), + } +} + +pub fn session_discovery() -> PiSessionDiscovery { + match resolve_session_root() { + SessionRootResolution::Available { root, source } => PiSessionDiscovery::Available { + root: root.to_string_lossy().into_owned(), + source, + }, + SessionRootResolution::RequiresProjectContext { + configured_path, + source, + } => PiSessionDiscovery::RequiresProjectContext { + configured_path, + source, + }, + SessionRootResolution::Unavailable { reason } => PiSessionDiscovery::Unavailable { reason }, + } +} + +fn resolve_session_root() -> SessionRootResolution { + let home = crate::config::get_home_dir(); + if let Some(raw) = std::env::var_os("PI_CODING_AGENT_SESSION_DIR") { + if !raw.is_empty() { + return classify_configured_session_dir( + raw.to_string_lossy().as_ref(), + &home, + "environment", + ); + } + } + + match crate::pi_config::native_settings::read_pi_native_defaults() { + Ok(defaults) => { + if let Some(value) = defaults.session_dir.filter(|value| !value.is_empty()) { + return classify_configured_session_dir(&value, &home, "settings"); + } + } + Err(error) => { + return SessionRootResolution::Unavailable { + reason: error.to_string(), + }; + } + } + + match crate::pi_config::native::get_pi_agent_dir() { + Ok(agent_dir) => SessionRootResolution::Available { + root: agent_dir.join("sessions"), + source: "default", + }, + Err(error) => SessionRootResolution::Unavailable { + reason: error.to_string(), + }, + } +} + +fn classify_configured_session_dir( + value: &str, + home: &Path, + source: &'static str, +) -> SessionRootResolution { + match resolve_global_session_dir(value, home) { + Some(root) => SessionRootResolution::Available { root, source }, + None => SessionRootResolution::RequiresProjectContext { + configured_path: value.to_string(), + source, + }, + } +} + +fn resolve_global_session_dir(value: &str, home: &Path) -> Option { + let path = if value == "~" { + home.to_path_buf() + } else if let Some(suffix) = value + .strip_prefix("~/") + .or_else(|| value.strip_prefix("~\\")) + { + home.join(suffix) + } else { + PathBuf::from(value) + }; + path.is_absolute().then_some(path) +} + +pub fn scan_sessions() -> Vec { + let Some(root) = session_roots().into_iter().next() else { + match session_discovery() { + PiSessionDiscovery::RequiresProjectContext { + configured_path, .. + } => log::warn!( + "Pi sessionDir '{configured_path}' requires a project cwd and cannot be globally enumerated" + ), + PiSessionDiscovery::Unavailable { reason } => { + log::warn!("Pi session discovery unavailable: {reason}") + } + PiSessionDiscovery::Available { .. } => {} + } + return Vec::new(); + }; + scan_sessions_in_root(&root) +} + +fn scan_sessions_in_root(root: &Path) -> Vec { + let mut files = Vec::new(); + collect_jsonl_files(root, 0, &mut files); + files + .into_iter() + .filter_map(|path| match parse_session(&path) { + Ok(session) => Some(session), + Err(error) => { + log::debug!("Skipping invalid Pi session {}: {error}", path.display()); + None + } + }) + .collect() +} + +pub fn load_messages(path: &Path) -> Result, String> { + let root = session_roots() + .into_iter() + .next() + .ok_or_else(|| "Relative Pi sessionDir cannot be globally resolved".to_string())?; + load_messages_with_root(&root, path) +} + +fn load_messages_with_root(root: &Path, path: &Path) -> Result, String> { + let (_, source) = validate_source_under_root(root, path)?; + let tree = read_tree(&source)?; + Ok(read_active_data(&source, &tree)?.messages) +} + +pub fn delete_session(root: &Path, path: &Path, session_id: &str) -> Result { + if !is_valid_tree_id(session_id) { + return Err("Invalid Pi session ID".to_string()); + } + let (_, source) = validate_source_under_root(root, path)?; + let tree = read_tree(&source)?; + if tree.header.id != session_id { + return Err(format!( + "Pi session ID mismatch: expected {session_id}, found {}", + tree.header.id + )); + } + fs::remove_file(&source) + .map_err(|error| format!("Failed to delete Pi session {}: {error}", source.display()))?; + Ok(true) +} + +fn parse_session(path: &Path) -> Result { + let source = path + .canonicalize() + .map_err(|error| format!("Failed to resolve Pi session {}: {error}", path.display()))?; + let source_path = source + .to_str() + .ok_or_else(|| "Pi session path is not valid UTF-8".to_string())? + .to_string(); + let tree = read_tree(&source)?; + let data = read_active_data(&source, &tree)?; + let title = data.explicit_name.flatten().or_else(|| { + data.first_user_message + .as_deref() + .map(|message| truncate_summary(message, TITLE_MAX_CHARS)) + .filter(|message| !message.is_empty()) + .or_else(|| path_basename(&tree.header.cwd)) + }); + let summary = data + .last_message + .as_deref() + .map(|message| truncate_summary(message, 160)) + .filter(|message| !message.is_empty()); + Ok(SessionMeta { + provider_id: PROVIDER_ID.to_string(), + session_id: tree.header.id.clone(), + title, + summary, + project_dir: (!tree.header.cwd.trim().is_empty()).then(|| tree.header.cwd.clone()), + created_at: tree.header.timestamp, + last_active_at: data.last_active_at.or(tree.header.timestamp), + source_path: Some(source_path.clone()), + resume_command: Some(format!( + "pi --session {}", + crate::session_manager::terminal::shell_escape(&source_path) + )), + }) +} + +fn read_tree(path: &Path) -> Result { + validate_file_size(path)?; + let reader = BufReader::new( + File::open(path).map_err(|error| format!("Failed to open Pi session: {error}"))?, + ); + let mut header = None; + let mut parents = HashMap::>::new(); + let mut latest_id = None; + let mut legacy_previous_id = None; + let mut entry_index = 0usize; + for line in reader.lines() { + let line = line.map_err(|error| format!("Failed to read Pi session: {error}"))?; + if line.trim().is_empty() { + continue; + } + let Ok(value) = serde_json::from_str::(&line) else { + continue; + }; + if header.is_none() { + header = Some(parse_header(&value)?); + continue; + } + entry_index += 1; + if entry_index > MAX_TREE_ENTRIES { + return Err(format!( + "Pi session exceeds the {MAX_TREE_ENTRIES}-entry safety limit" + )); + } + let version = header + .as_ref() + .map_or(1, |item: &SessionHeader| item.version); + let Some((id, parent_id)) = + entry_identity(&value, version, entry_index, legacy_previous_id.as_deref()) + else { + continue; + }; + if parents.insert(id.clone(), parent_id).is_some() { + return Err(format!("Pi session contains duplicate entry ID: {id}")); + } + latest_id = Some(id.clone()); + legacy_previous_id = Some(id); + } + let header = header.ok_or_else(|| "Pi session has no valid header".to_string())?; + let mut active_ids = HashSet::new(); + let mut current = latest_id; + while let Some(id) = current { + if !active_ids.insert(id.clone()) { + return Err(format!("Pi session tree contains a cycle at entry {id}")); + } + current = parents + .get(&id) + .ok_or_else(|| format!("Pi session entry references missing parent: {id}"))? + .clone(); + } + Ok(SessionTree { header, active_ids }) +} + +fn read_active_data(path: &Path, tree: &SessionTree) -> Result { + validate_file_size(path)?; + let reader = BufReader::new( + File::open(path).map_err(|error| format!("Failed to open Pi session: {error}"))?, + ); + let mut data = ActiveSessionData::default(); + let mut saw_header = false; + let mut entry_index = 0usize; + let mut legacy_previous_id = None; + for line in reader.lines() { + let line = line.map_err(|error| format!("Failed to read Pi session: {error}"))?; + if line.trim().is_empty() { + continue; + } + let Ok(value) = serde_json::from_str::(&line) else { + continue; + }; + if !saw_header { + if value.get("type").and_then(Value::as_str) == Some("session") { + saw_header = true; + } + continue; + } + entry_index += 1; + if entry_index > MAX_TREE_ENTRIES { + return Err(format!( + "Pi session exceeds the {MAX_TREE_ENTRIES}-entry safety limit" + )); + } + let Some((id, _)) = entry_identity( + &value, + tree.header.version, + entry_index, + legacy_previous_id.as_deref(), + ) else { + continue; + }; + legacy_previous_id = Some(id.clone()); + if value.get("type").and_then(Value::as_str) == Some("session_info") { + data.explicit_name = Some( + value + .get("name") + .and_then(Value::as_str) + .map(str::trim) + .filter(|name| !name.is_empty()) + .map(str::to_string), + ); + } + if !tree.active_ids.contains(&id) { + continue; + } + let entry_timestamp = value.get("timestamp").and_then(parse_timestamp_to_ms); + if let Some(timestamp) = entry_timestamp { + data.last_active_at = Some(timestamp); + } + match value.get("type").and_then(Value::as_str) { + Some("session_info") => {} + Some("message") => { + let Some((role, content)) = value.get("message").and_then(parse_message) else { + continue; + }; + let timestamp = value + .get("message") + .and_then(|message| message.get("timestamp")) + .and_then(parse_timestamp_to_ms) + .or(entry_timestamp); + if role == "user" && data.first_user_message.is_none() { + data.first_user_message = Some(content.clone()); + } + if matches!(role.as_str(), "user" | "assistant") { + data.last_message = Some(content.clone()); + } + data.messages.push(SessionMessage { + role, + content, + ts: timestamp, + }); + } + Some("compaction") | Some("branch_summary") => { + push_system( + &mut data.messages, + value + .get("summary") + .and_then(Value::as_str) + .unwrap_or_default(), + entry_timestamp, + ); + } + Some("custom_message") + if value.get("display").and_then(Value::as_bool) != Some(false) => + { + push_system( + &mut data.messages, + &value.get("content").map(extract_text).unwrap_or_default(), + entry_timestamp, + ); + } + _ => {} + } + } + Ok(data) +} + +fn push_system(messages: &mut Vec, content: &str, ts: Option) { + if !content.trim().is_empty() { + messages.push(SessionMessage { + role: "system".to_string(), + content: content.to_string(), + ts, + }); + } +} + +fn parse_header(value: &Value) -> Result { + if value.get("type").and_then(Value::as_str) != Some("session") { + return Err("Pi session header must be the first valid JSON entry".to_string()); + } + let id = value + .get("id") + .and_then(Value::as_str) + .filter(|id| is_valid_tree_id(id)) + .ok_or_else(|| "Pi session header has an invalid ID".to_string())? + .to_string(); + let version = value.get("version").and_then(Value::as_u64).unwrap_or(1); + if !(1..=3).contains(&version) { + return Err(format!("Unsupported Pi session version: {version}")); + } + Ok(SessionHeader { + id, + cwd: value + .get("cwd") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(), + timestamp: value.get("timestamp").and_then(parse_timestamp_to_ms), + version, + }) +} + +fn entry_identity( + value: &Value, + version: u64, + entry_index: usize, + legacy_previous_id: Option<&str>, +) -> Option<(String, Option)> { + if version < 2 { + return Some(( + format!("legacy-{entry_index}"), + legacy_previous_id.map(str::to_string), + )); + } + let id = value + .get("id") + .and_then(Value::as_str) + .filter(|id| is_valid_tree_id(id))? + .to_string(); + let parent_id = match value.get("parentId") { + None | Some(Value::Null) => None, + Some(Value::String(parent)) if is_valid_tree_id(parent) => Some(parent.clone()), + _ => return None, + }; + Some((id, parent_id)) +} + +fn parse_message(message: &Value) -> Option<(String, String)> { + let role = message.get("role").and_then(Value::as_str)?; + let (display_role, content) = match role { + "user" | "assistant" => ( + role.to_string(), + message.get("content").map(extract_text).unwrap_or_default(), + ), + "toolResult" => ( + "tool".to_string(), + message.get("content").map(extract_text).unwrap_or_default(), + ), + "bashExecution" => ( + "tool".to_string(), + format!( + "$ {}\n{}", + message + .get("command") + .and_then(Value::as_str) + .unwrap_or_default(), + message + .get("output") + .and_then(Value::as_str) + .unwrap_or_default() + ), + ), + "branchSummary" | "compactionSummary" => ( + "system".to_string(), + message + .get("summary") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(), + ), + _ => return None, + }; + (!content.trim().is_empty()).then_some((display_role, content)) +} + +fn validate_source_under_root(root: &Path, path: &Path) -> Result<(PathBuf, PathBuf), String> { + let root = root.canonicalize().map_err(|error| { + format!( + "Failed to resolve Pi session root {}: {error}", + root.display() + ) + })?; + let source = path + .canonicalize() + .map_err(|error| format!("Failed to resolve Pi session {}: {error}", path.display()))?; + if !source.starts_with(&root) { + return Err(format!( + "Pi session source is outside the session root: {}", + path.display() + )); + } + let metadata = fs::symlink_metadata(&source) + .map_err(|error| format!("Failed to inspect Pi session {}: {error}", source.display()))?; + if !metadata.file_type().is_file() + || source.extension().and_then(|value| value.to_str()) != Some("jsonl") + || metadata.len() > MAX_SESSION_BYTES + { + return Err(format!("Invalid Pi session file: {}", source.display())); + } + Ok((root, source)) +} + +fn validate_file_size(path: &Path) -> Result<(), String> { + let metadata = + fs::metadata(path).map_err(|error| format!("Failed to inspect Pi session: {error}"))?; + if metadata.len() > MAX_SESSION_BYTES { + Err(format!( + "Pi session exceeds the {MAX_SESSION_BYTES}-byte safety limit" + )) + } else { + Ok(()) + } +} + +fn is_valid_tree_id(id: &str) -> bool { + let bytes = id.as_bytes(); + !bytes.is_empty() + && bytes.len() <= MAX_TREE_ID_BYTES + && bytes.first().is_some_and(u8::is_ascii_alphanumeric) + && bytes.last().is_some_and(u8::is_ascii_alphanumeric) + && bytes + .iter() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) +} + +fn collect_jsonl_files(root: &Path, depth: usize, output: &mut Vec) { + if depth > MAX_SCAN_DEPTH { + return; + } + let Ok(entries) = fs::read_dir(root) else { + return; + }; + for entry in entries.flatten() { + let Ok(file_type) = entry.file_type() else { + continue; + }; + let path = entry.path(); + if file_type.is_dir() { + collect_jsonl_files(&path, depth + 1, output); + } else if file_type.is_file() + && path.extension().and_then(|value| value.to_str()) == Some("jsonl") + && entry + .metadata() + .is_ok_and(|metadata| metadata.len() <= MAX_SESSION_BYTES) + { + output.push(path); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn latest_leaf_defines_the_active_branch() { + let temp = tempfile::tempdir().expect("tempdir"); + let root = temp.path().join("sessions"); + fs::create_dir_all(&root).expect("root"); + let path = root.join("tree.jsonl"); + fs::write( + &path, + "{\"type\":\"session\",\"version\":3,\"id\":\"session-1\",\"cwd\":\"/work\"}\n\ + {\"type\":\"message\",\"id\":\"root\",\"parentId\":null,\"message\":{\"role\":\"user\",\"content\":\"question\"}}\n\ + {\"type\":\"message\",\"id\":\"dead\",\"parentId\":\"root\",\"message\":{\"role\":\"assistant\",\"content\":\"abandoned\"}}\n\ + {\"type\":\"message\",\"id\":\"live\",\"parentId\":\"root\",\"message\":{\"role\":\"assistant\",\"content\":\"active\"}}\n", + ) + .expect("session"); + let messages = load_messages_with_root(&root, &path).expect("messages"); + assert_eq!( + messages + .into_iter() + .map(|message| message.content) + .collect::>(), + vec!["question", "active"] + ); + } + + #[test] + fn capture_matches_global_name_and_malformed_line_semantics() { + // Executed by scripts/pi-transport-capture.mjs against pinned Pi: + // getSessionName() keeps the latest global session_info even when its + // branch is inactive, and SessionManager.open() skips a malformed line. + let temp = tempfile::tempdir().expect("tempdir"); + let path = temp.path().join("captured.jsonl"); + fs::write( + &path, + "{\"type\":\"session\",\"version\":3,\"id\":\"session-1\",\"cwd\":\"/work\"}\n\ + {\"type\":\"message\",\"id\":\"root\",\"parentId\":null,\"message\":{\"role\":\"user\",\"content\":\"root\"}}\n\ + {not valid json\n\ + {\"type\":\"session_info\",\"id\":\"dead-name\",\"parentId\":\"root\",\"name\":\"Abandoned branch name\"}\n\ + {\"type\":\"message\",\"id\":\"dead\",\"parentId\":\"dead-name\",\"message\":{\"role\":\"assistant\",\"content\":\"abandoned\"}}\n\ + {\"type\":\"message\",\"id\":\"live\",\"parentId\":\"root\",\"message\":{\"role\":\"user\",\"content\":\"active branch\"}}\n", + ) + .expect("captured session"); + + let session = parse_session(&path).expect("parse capture semantics"); + assert_eq!(session.title.as_deref(), Some("Abandoned branch name")); + assert_eq!(session.summary.as_deref(), Some("active branch")); + } + + #[test] + fn relative_root_is_explicitly_non_enumerable() { + assert_eq!( + resolve_global_session_dir(".pi/sessions", Path::new("/home/pi")), + None + ); + assert_eq!( + classify_configured_session_dir(".pi/sessions", Path::new("/home/pi"), "settings"), + SessionRootResolution::RequiresProjectContext { + configured_path: ".pi/sessions".to_string(), + source: "settings", + } + ); + } + + #[test] + fn capture_generated_v3_shape_round_trips_all_consumed_fields() { + // Generated by scripts/pi-transport-capture.mjs against pinned Pi + // ab366ebe94cacd419d986be454f12b1b9913aaca using SessionManager APIs. + let temp = tempfile::tempdir().expect("tempdir"); + let root = temp.path().join("sessions"); + fs::create_dir_all(&root).expect("root"); + let path = root.join("captured.jsonl"); + fs::write( + &path, + "{\"type\":\"session\",\"version\":3,\"id\":\"cc-switch-capture-session\",\"timestamp\":\"2023-11-14T22:13:20.000Z\",\"cwd\":\"/work/captured\",\"parentSession\":null}\n\ + {\"type\":\"session_info\",\"id\":\"00000000-0000-7000-8000-000000000001\",\"parentId\":null,\"timestamp\":\"2023-11-14T22:13:20.100Z\",\"name\":\"Captured session\"}\n\ + {\"type\":\"message\",\"id\":\"00000000-0000-7000-8000-000000000002\",\"parentId\":\"00000000-0000-7000-8000-000000000001\",\"timestamp\":\"2023-11-14T22:13:20.200Z\",\"message\":{\"role\":\"user\",\"content\":[{\"type\":\"text\",\"text\":\"captured question\"}],\"timestamp\":1700000000000}}\n\ + {\"type\":\"message\",\"id\":\"00000000-0000-7000-8000-000000000003\",\"parentId\":\"00000000-0000-7000-8000-000000000002\",\"timestamp\":\"2023-11-14T22:13:21.200Z\",\"message\":{\"role\":\"assistant\",\"content\":[{\"type\":\"text\",\"text\":\"captured answer\"}],\"api\":\"openai-responses\",\"provider\":\"capture\",\"model\":\"capture-model\",\"usage\":{\"input\":1,\"output\":1,\"cacheRead\":0,\"cacheWrite\":0,\"totalTokens\":2,\"cost\":{\"input\":0,\"output\":0,\"cacheRead\":0,\"cacheWrite\":0,\"total\":0}},\"stopReason\":\"stop\",\"timestamp\":1700000001000}}\n", + ) + .expect("captured session"); + + let session = parse_session(&path).expect("parse capture-generated session"); + assert_eq!(session.session_id, "cc-switch-capture-session"); + assert_eq!(session.title.as_deref(), Some("Captured session")); + assert_eq!(session.summary.as_deref(), Some("captured answer")); + assert_eq!(session.project_dir.as_deref(), Some("/work/captured")); + assert_eq!(session.created_at, Some(1_700_000_000_000)); + assert_eq!(session.last_active_at, Some(1_700_000_001_200)); + // scripts/pi-transport-capture.mjs executes pinned Pi's parseArgs with + // ["--session", ] and records that the + // exact path is returned in Args.session. + assert!(session + .resume_command + .as_deref() + .is_some_and(|command| command.starts_with("pi --session "))); + + let messages = load_messages_with_root(&root, &path).expect("load messages"); + assert_eq!( + messages + .iter() + .map(|message| (message.role.as_str(), message.content.as_str(), message.ts)) + .collect::>(), + vec![ + ("user", "captured question", Some(1_700_000_000_000)), + ("assistant", "captured answer", Some(1_700_000_001_000)), + ] + ); + } + + #[test] + fn deletion_requires_containment_and_matching_header_id() { + let temp = tempfile::tempdir().expect("tempdir"); + let root = temp.path().join("sessions"); + fs::create_dir_all(&root).expect("root"); + let path = root.join("session.jsonl"); + fs::write( + &path, + "{\"type\":\"session\",\"version\":3,\"id\":\"session-1\",\"cwd\":\"/work\"}\n", + ) + .expect("session"); + assert!(delete_session(&root, &path, "other").is_err()); + assert!(path.exists()); + assert!(delete_session(&root, &path, "session-1").expect("delete")); + } +} diff --git a/src-tauri/src/session_manager/terminal/mod.rs b/src-tauri/src/session_manager/terminal/mod.rs index 2124e6c17..aaacb1a48 100644 --- a/src-tauri/src/session_manager/terminal/mod.rs +++ b/src-tauri/src/session_manager/terminal/mod.rs @@ -333,7 +333,7 @@ fn build_shell_command(command: &str, cwd: Option<&str>) -> String { /// /// 单引号内不做任何展开,唯一的特例是 `'` 自身无法被表示:用「闭合-转义-重开」 /// 的 `'\''` 序列绕过。 -fn shell_escape(value: &str) -> String { +pub(crate) fn shell_escape(value: &str) -> String { format!("'{}'", value.replace('\'', r"'\''")) } diff --git a/src-tauri/src/settings.rs b/src-tauri/src/settings.rs index 30010ff85..f7d090a7d 100644 --- a/src-tauri/src/settings.rs +++ b/src-tauri/src/settings.rs @@ -17,10 +17,156 @@ pub struct CustomEndpoint { pub last_used: Option, } +/// Device-local Pi gateway behavior. The frozen `proxy_config` schema has a +/// closed four-app domain, so Pi must not manufacture an out-of-contract row. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct PiProxySettings { + #[serde(default)] + pub auto_failover_enabled: bool, + #[serde(default = "default_pi_max_retries")] + pub max_retries: u32, + #[serde(default = "default_pi_first_byte_timeout")] + pub streaming_first_byte_timeout: u32, + #[serde(default = "default_pi_idle_timeout")] + pub streaming_idle_timeout: u32, + #[serde(default = "default_pi_request_timeout")] + pub non_streaming_timeout: u32, + #[serde(default = "default_pi_circuit_failure_threshold")] + pub circuit_failure_threshold: u32, + #[serde(default = "default_pi_circuit_success_threshold")] + pub circuit_success_threshold: u32, + #[serde(default = "default_pi_circuit_timeout")] + pub circuit_timeout_seconds: u32, + #[serde(default = "default_pi_circuit_error_rate")] + pub circuit_error_rate_threshold: f64, + #[serde(default = "default_pi_circuit_min_requests")] + pub circuit_min_requests: u32, +} + +const fn default_pi_max_retries() -> u32 { + 3 +} +const fn default_pi_first_byte_timeout() -> u32 { + 60 +} +const fn default_pi_idle_timeout() -> u32 { + 120 +} +const fn default_pi_request_timeout() -> u32 { + 600 +} +const fn default_pi_circuit_failure_threshold() -> u32 { + 4 +} +const fn default_pi_circuit_success_threshold() -> u32 { + 2 +} +const fn default_pi_circuit_timeout() -> u32 { + 60 +} +const fn default_pi_circuit_min_requests() -> u32 { + 10 +} +fn default_pi_circuit_error_rate() -> f64 { + 0.6 +} + +impl Default for PiProxySettings { + fn default() -> Self { + Self { + auto_failover_enabled: false, + max_retries: default_pi_max_retries(), + streaming_first_byte_timeout: default_pi_first_byte_timeout(), + streaming_idle_timeout: default_pi_idle_timeout(), + non_streaming_timeout: default_pi_request_timeout(), + circuit_failure_threshold: default_pi_circuit_failure_threshold(), + circuit_success_threshold: default_pi_circuit_success_threshold(), + circuit_timeout_seconds: default_pi_circuit_timeout(), + circuit_error_rate_threshold: default_pi_circuit_error_rate(), + circuit_min_requests: default_pi_circuit_min_requests(), + } + } +} + +impl PiProxySettings { + pub(crate) fn validate(&self) -> Result<(), AppError> { + if !self.circuit_error_rate_threshold.is_finite() + || !(0.0..=1.0).contains(&self.circuit_error_rate_threshold) + { + return Err(AppError::InvalidInput( + "Pi circuit error-rate threshold must be finite and within [0, 1]".to_string(), + )); + } + if self.max_retries > 32 { + return Err(AppError::InvalidInput( + "Pi max retries cannot exceed 32".to_string(), + )); + } + Ok(()) + } + + pub(crate) fn app_config(&self, enabled: bool) -> crate::proxy::types::AppProxyConfig { + crate::proxy::types::AppProxyConfig { + app_type: "pi".to_string(), + enabled, + auto_failover_enabled: self.auto_failover_enabled, + max_retries: self.max_retries, + streaming_first_byte_timeout: self.streaming_first_byte_timeout, + streaming_idle_timeout: self.streaming_idle_timeout, + non_streaming_timeout: self.non_streaming_timeout, + circuit_failure_threshold: self.circuit_failure_threshold, + circuit_success_threshold: self.circuit_success_threshold, + circuit_timeout_seconds: self.circuit_timeout_seconds, + circuit_error_rate_threshold: self.circuit_error_rate_threshold, + circuit_min_requests: self.circuit_min_requests, + } + } +} + fn default_true() -> bool { true } +/// Device-local bearer used only between Pi and cc-switch's loopback gateway. +/// +/// Deliberately redacts `Debug`; the value must never be returned by settings +/// IPC, copied into SQLite, or emitted to logs. +#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(transparent)] +pub struct GatewayToken(String); + +impl GatewayToken { + fn generate() -> Self { + Self(format!( + "ccs_pi_{}{}", + uuid::Uuid::new_v4().simple(), + uuid::Uuid::new_v4().simple() + )) + } + + pub(crate) fn expose(&self) -> &str { + &self.0 + } + + pub(crate) fn constant_time_eq(&self, candidate: &str) -> bool { + let expected = self.0.as_bytes(); + let candidate = candidate.as_bytes(); + let mut difference = expected.len() ^ candidate.len(); + let shared = expected.len().min(candidate.len()); + for index in 0..shared { + difference |= usize::from(expected[index] ^ candidate[index]); + } + difference == 0 + } +} + +impl std::fmt::Debug for GatewayToken { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("GatewayToken()") + } +} + /// 主页面显示的应用配置 #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] @@ -46,6 +192,8 @@ pub struct VisibleApps { pub openclaw: bool, #[serde(default)] pub hermes: bool, + #[serde(default = "default_true")] + pub pi: bool, } impl Default for VisibleApps { @@ -59,6 +207,7 @@ impl Default for VisibleApps { opencode: true, openclaw: true, hermes: false, // 默认不显示,需用户手动启用 + pi: true, } } } @@ -75,6 +224,7 @@ impl VisibleApps { AppType::OpenCode => self.opencode, AppType::OpenClaw => self.openclaw, AppType::Hermes => self.hermes, + AppType::Pi => self.pi, } } } @@ -422,6 +572,8 @@ pub struct AppSettings { pub openclaw_config_dir: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub hermes_config_dir: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub pi_config_dir: Option, // ===== 当前供应商 ID(设备级)===== /// 当前 Claude 供应商 ID(本地存储,优先于数据库 is_current) @@ -448,6 +600,29 @@ pub struct AppSettings { /// 当前 Hermes 供应商 ID(本地存储,保持结构一致) #[serde(default, skip_serializing_if = "Option::is_none")] pub current_provider_hermes: Option, + /// 当前 Pi 供应商 ID(本地存储) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub current_provider_pi: Option, + + /// Device-local desired state for Pi's native `models.json` gateway + /// projection. Unlike the shared proxy-config row, this survives database + /// replacement and is reconciled against the live listener on startup. + #[serde(default)] + pub pi_takeover_enabled: bool, + #[serde(default)] + pub pi_proxy: PiProxySettings, + + /// Stable device-installation secret for Pi's loopback gateway. + /// + /// This field is serialized only to the device settings file. The frontend + /// projection clears it and settings-save merge restores the existing + /// value, so ordinary IPC cannot read, replace, or clear it. + #[serde(default, skip_serializing_if = "Option::is_none")] + // Public so integration tests and downstream Rust callers can continue to + // use struct-update syntax with `AppSettings`. IPC still cannot observe or + // mutate the value: the settings command clears it on reads and restores + // the persisted value on writes. + pub pi_gateway_token: Option, // ===== Skill 同步设置 ===== /// Skill 同步方式:auto(默认,优先 symlink)、symlink、copy @@ -533,6 +708,7 @@ impl Default for AppSettings { opencode_config_dir: None, openclaw_config_dir: None, hermes_config_dir: None, + pi_config_dir: None, current_provider_claude: None, current_provider_claude_desktop: None, current_provider_codex: None, @@ -541,6 +717,10 @@ impl Default for AppSettings { current_provider_opencode: None, current_provider_openclaw: None, current_provider_hermes: None, + current_provider_pi: None, + pi_takeover_enabled: false, + pi_proxy: PiProxySettings::default(), + pi_gateway_token: None, skill_sync_method: SyncMethod::default(), skill_storage_location: SkillStorageLocation::default(), webdav_sync: None, @@ -614,6 +794,13 @@ impl AppSettings { .filter(|s| !s.is_empty()) .map(|s| s.to_string()); + self.pi_config_dir = self + .pi_config_dir + .as_ref() + .map(|s| s.trim()) + .filter(|s| !s.is_empty()) + .map(|s| s.to_string()); + self.language = self .language .as_ref() @@ -672,31 +859,9 @@ fn save_settings_file(settings: &AppSettings) -> Result<(), AppError> { fs::create_dir_all(parent).map_err(|e| AppError::io(parent, e))?; } - let json = serde_json::to_string_pretty(&normalized) + let json = serde_json::to_vec_pretty(&normalized) .map_err(|e| AppError::JsonSerialize { source: e })?; - #[cfg(unix)] - { - use std::fs::OpenOptions; - use std::io::Write; - use std::os::unix::fs::OpenOptionsExt; - - let mut file = OpenOptions::new() - .create(true) - .write(true) - .truncate(true) - .mode(0o600) - .open(&path) - .map_err(|e| AppError::io(&path, e))?; - file.write_all(json.as_bytes()) - .map_err(|e| AppError::io(&path, e))?; - } - - #[cfg(not(unix))] - { - fs::write(&path, json).map_err(|e| AppError::io(&path, e))?; - } - - Ok(()) + crate::config::atomic_write_durable(&path, &json, Some(0o600)) } static SETTINGS_STORE: OnceLock> = OnceLock::new(); @@ -705,7 +870,7 @@ fn settings_store() -> &'static RwLock { SETTINGS_STORE.get_or_init(|| RwLock::new(AppSettings::load_from_file())) } -fn resolve_override_path(raw: &str) -> PathBuf { +pub(crate) fn resolve_override_path(raw: &str) -> PathBuf { if raw == "~" { if let Some(home) = dirs::home_dir() { return home; @@ -742,17 +907,17 @@ pub fn get_settings_for_frontend() -> AppSettings { s3.secret_access_key.clear(); } settings.webdav_backup = None; + settings.pi_gateway_token = None; settings } pub fn update_settings(mut new_settings: AppSettings) -> Result<(), AppError> { new_settings.normalize_paths(); - save_settings_file(&new_settings)?; - let mut guard = settings_store().write().unwrap_or_else(|e| { log::warn!("设置锁已毒化,使用恢复值: {e}"); e.into_inner() }); + save_settings_file(&new_settings)?; *guard = new_settings; Ok(()) } @@ -933,6 +1098,64 @@ pub fn get_hermes_override_dir() -> Option { .map(|p| resolve_override_path(p)) } +pub fn get_pi_override_dir() -> Option { + let settings = settings_store().read().ok()?; + settings + .pi_config_dir + .as_ref() + .map(|p| resolve_override_path(p)) +} + +pub(crate) fn get_or_create_pi_gateway_token() -> Result { + let mut token = None; + mutate_settings(|settings| { + let stored = settings + .pi_gateway_token + .get_or_insert_with(GatewayToken::generate); + token = Some(stored.clone()); + })?; + token.ok_or_else(|| AppError::Config("无法创建 Pi 网关凭据".to_string())) +} + +pub(crate) fn get_pi_gateway_token() -> Result { + get_settings().pi_gateway_token.ok_or_else(|| { + AppError::Conflict( + "Pi takeover is active but its gateway credential is unavailable".to_string(), + ) + }) +} + +pub(crate) fn reset_pi_gateway_token() -> Result { + let generated = GatewayToken::generate(); + mutate_settings(|settings| settings.pi_gateway_token = Some(generated.clone()))?; + Ok(generated) +} + +pub(crate) fn replace_pi_gateway_token(token: GatewayToken) -> Result<(), AppError> { + mutate_settings(|settings| settings.pi_gateway_token = Some(token)) +} + +pub(crate) fn pi_takeover_enabled() -> bool { + get_settings().pi_takeover_enabled +} + +pub(crate) fn set_pi_takeover_enabled(enabled: bool) -> Result<(), AppError> { + mutate_settings(|settings| settings.pi_takeover_enabled = enabled) +} + +pub(crate) fn get_pi_proxy_settings() -> PiProxySettings { + get_settings().pi_proxy +} + +pub(crate) fn update_pi_proxy_settings(settings: PiProxySettings) -> Result<(), AppError> { + settings.validate()?; + mutate_settings(|current| current.pi_proxy = settings) +} + +pub(crate) fn get_pi_app_proxy_config() -> crate::proxy::types::AppProxyConfig { + get_pi_proxy_settings().app_config(pi_takeover_enabled()) +} + pub fn preserve_codex_official_auth_on_switch() -> bool { settings_store() .read() @@ -970,6 +1193,7 @@ pub fn get_current_provider(app_type: &AppType) -> Option { AppType::OpenCode => settings.current_provider_opencode.clone(), AppType::OpenClaw => settings.current_provider_openclaw.clone(), AppType::Hermes => settings.current_provider_hermes.clone(), + AppType::Pi => settings.current_provider_pi.clone(), } } @@ -988,6 +1212,7 @@ pub fn set_current_provider(app_type: &AppType, id: Option<&str>) -> Result<(), AppType::OpenCode => settings.current_provider_opencode = id_owned.clone(), AppType::OpenClaw => settings.current_provider_openclaw = id_owned.clone(), AppType::Hermes => settings.current_provider_hermes = id_owned.clone(), + AppType::Pi => settings.current_provider_pi = id_owned.clone(), }) } @@ -1161,6 +1386,10 @@ mod tests { .expect("visible apps"); assert!(visible.is_visible(&AppType::ClaudeDesktop)); + assert!( + visible.is_visible(&AppType::Pi), + "Pi is a first-class app and must be visible when older settings omit its field" + ); } #[test] diff --git a/src-tauri/src/store.rs b/src-tauri/src/store.rs index f34e1de9a..81b83b6b4 100644 --- a/src-tauri/src/store.rs +++ b/src-tauri/src/store.rs @@ -3,6 +3,7 @@ use crate::services::{ProxyService, UsageCache}; use std::sync::Arc; /// 全局应用状态 +#[derive(Clone)] pub struct AppState { pub db: Arc, pub proxy_service: ProxyService, diff --git a/src/App.tsx b/src/App.tsx index e794ab131..8bd26ce78 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -31,6 +31,7 @@ import type { Provider, VisibleApps } from "@/types"; import type { EnvConflict } from "@/types/env"; import { proxyKeys, useProvidersQuery, useSettingsQuery } from "@/lib/query"; import { + piApi, providersApi, settingsApi, type AppId, @@ -59,6 +60,7 @@ import { import { AppSwitcher } from "@/components/AppSwitcher"; import { ProfileSwitcher } from "@/components/profiles/ProfileSwitcher"; import { ProviderList } from "@/components/providers/ProviderList"; +import { PiNativeCatalogPanel } from "@/components/providers/PiNativeCatalogPanel"; import { AddProviderDialog } from "@/components/providers/AddProviderDialog"; import { EditProviderDialog } from "@/components/providers/EditProviderDialog"; import { ConfirmDialog } from "@/components/ConfirmDialog"; @@ -94,6 +96,7 @@ import ToolsPanel from "@/components/openclaw/ToolsPanel"; import AgentsDefaultsPanel from "@/components/openclaw/AgentsDefaultsPanel"; import OpenClawHealthBanner from "@/components/openclaw/OpenClawHealthBanner"; import HermesMemoryPanel from "@/components/hermes/HermesMemoryPanel"; +import { APP_IDS, DEFAULT_VISIBLE_APPS } from "@/config/appConfig"; type View = | "providers" @@ -121,20 +124,9 @@ const DEFAULT_DRAG_BAR_HEIGHT = isWindows() || isLinux() ? 0 : 28; // px const HEADER_HEIGHT = 64; // px const STORAGE_KEY = "cc-switch-last-app"; -const VALID_APPS: AppId[] = [ - "claude", - "claude-desktop", - "codex", - "gemini", - "grokbuild", - "opencode", - "openclaw", - "hermes", -]; - const getInitialApp = (): AppId => { const saved = localStorage.getItem(STORAGE_KEY) as AppId | null; - if (saved && VALID_APPS.includes(saved)) { + if (saved && APP_IDS.includes(saved)) { return saved; } return "claude"; @@ -189,27 +181,16 @@ function App() { isLinux() && (settingsData?.useAppWindowControls ?? false); const dragBarHeight = useAppWindowControls ? 32 : DEFAULT_DRAG_BAR_HEIGHT; const contentTopOffset = dragBarHeight + HEADER_HEIGHT; - const visibleApps: VisibleApps = settingsData?.visibleApps ?? { - claude: true, - "claude-desktop": true, - codex: true, - gemini: true, - grokbuild: true, - opencode: true, - openclaw: true, - hermes: true, - }; + const visibleApps = useMemo( + () => ({ + ...DEFAULT_VISIBLE_APPS, + ...settingsData?.visibleApps, + }), + [settingsData?.visibleApps], + ); const getFirstVisibleApp = (): AppId => { - if (visibleApps.claude) return "claude"; - if (visibleApps["claude-desktop"]) return "claude-desktop"; - if (visibleApps.codex) return "codex"; - if (visibleApps.gemini) return "gemini"; - if (visibleApps.grokbuild) return "grokbuild"; - if (visibleApps.opencode) return "opencode"; - if (visibleApps.openclaw) return "openclaw"; - if (visibleApps.hermes) return "hermes"; - return "claude"; // fallback + return APP_IDS.find((app) => visibleApps[app]) ?? "claude"; }; useEffect(() => { @@ -220,6 +201,10 @@ function App() { // Fallback from sessions view when switching to an app without session support useEffect(() => { + if (currentView === "mcp" && sharedFeatureApp === "pi") { + setCurrentView("providers"); + return; + } if ( currentView === "sessions" && sharedFeatureApp !== "claude" && @@ -228,7 +213,8 @@ function App() { sharedFeatureApp !== "opencode" && sharedFeatureApp !== "openclaw" && sharedFeatureApp !== "gemini" && - sharedFeatureApp !== "hermes" + sharedFeatureApp !== "hermes" && + sharedFeatureApp !== "pi" ) { setCurrentView("providers"); } @@ -295,7 +281,9 @@ function App() { sharedFeatureApp === "opencode" || sharedFeatureApp === "openclaw" || sharedFeatureApp === "gemini" || - sharedFeatureApp === "hermes"; + sharedFeatureApp === "hermes" || + sharedFeatureApp === "pi"; + const hasMcpSupport = sharedFeatureApp !== "pi"; const { addProvider, @@ -726,7 +714,8 @@ function App() { if ( activeApp === "opencode" || activeApp === "openclaw" || - activeApp === "hermes" + activeApp === "hermes" || + activeApp === "pi" ) { let liveProviderIds: string[] = []; try { @@ -741,10 +730,17 @@ function App() { queryKey: openclawKeys.liveProviderIds, queryFn: () => providersApi.getOpenClawLiveProviderIds(), }) - : await queryClient.ensureQueryData({ - queryKey: hermesKeys.liveProviderIds, - queryFn: () => providersApi.getHermesLiveProviderIds(), - }); + : activeApp === "hermes" + ? await queryClient.ensureQueryData({ + queryKey: hermesKeys.liveProviderIds, + queryFn: () => providersApi.getHermesLiveProviderIds(), + }) + : ( + await queryClient.ensureQueryData({ + queryKey: ["pi", "nativeCatalog"], + queryFn: () => piApi.getNativeCatalog(), + }) + ).map((entry) => entry.providerKey); } catch (error) { console.error( "[App] Failed to load live provider IDs for duplication", @@ -978,6 +974,9 @@ function App() { transition={{ duration: 0.15 }} className="space-y-4" > + {activeApp === "pi" && ( + + )} - + {hasMcpSupport && ( + + )} ) : activeApp === "openclaw" ? ( <> @@ -1542,15 +1543,17 @@ function App() { > - + {hasMcpSupport && ( + + )} )} diff --git a/src/components/AppSwitcher.tsx b/src/components/AppSwitcher.tsx index bff3a2054..13d3163b6 100644 --- a/src/components/AppSwitcher.tsx +++ b/src/components/AppSwitcher.tsx @@ -3,6 +3,7 @@ import type { VisibleApps } from "@/types"; import { ProviderIcon } from "@/components/ProviderIcon"; import { cn } from "@/lib/utils"; import { Monitor, Terminal } from "lucide-react"; +import { APP_IDS } from "@/config/appConfig"; const APP_BADGE_ICON: Partial< Record @@ -17,16 +18,6 @@ interface AppSwitcherProps { visibleApps?: VisibleApps; } -const ALL_APPS: AppId[] = [ - "claude", - "claude-desktop", - "codex", - "gemini", - "grokbuild", - "opencode", - "openclaw", - "hermes", -]; const STORAGE_KEY = "cc-switch-last-app"; export function AppSwitcher({ @@ -49,6 +40,7 @@ export function AppSwitcher({ opencode: "opencode", openclaw: "openclaw", hermes: "hermes", + pi: "pi", }; const appDisplayName: Record = { claude: "Claude Code", @@ -59,10 +51,11 @@ export function AppSwitcher({ opencode: "OpenCode", openclaw: "OpenClaw", hermes: "Hermes", + pi: "Pi", }; // Filter apps based on visibility settings (default all visible) - const appsToShow = ALL_APPS.filter((app) => { + const appsToShow = APP_IDS.filter((app) => { if (!visibleApps) return true; return visibleApps[app]; }); diff --git a/src/components/DeepLinkImportDialog.tsx b/src/components/DeepLinkImportDialog.tsx index f2d4b669c..a967a76e0 100644 --- a/src/components/DeepLinkImportDialog.tsx +++ b/src/components/DeepLinkImportDialog.tsx @@ -405,6 +405,16 @@ export function DeepLinkImportDialog() { {/* Model Fields - 根据应用类型显示不同的模型字段 */} + {request.app === "pi" && request.api && ( +
+
+ {t("deeplink.api")} +
+
+ {request.api} +
+
+ )} {request.app === "claude" ? ( <> {/* Claude 四种模型字段 */} diff --git a/src/components/UsageScriptModal.tsx b/src/components/UsageScriptModal.tsx index a2960d7af..442f15f03 100644 --- a/src/components/UsageScriptModal.tsx +++ b/src/components/UsageScriptModal.tsx @@ -284,6 +284,16 @@ const UsageScriptModal: React.FC = ({ apiKey: (config as any).api_key, baseUrl: (config as any).base_url, }; + } else if (appId === "pi") { + // Pi: provider values are camelCase; a model may override baseUrl. + const root = config as any; + const firstModel = Array.isArray(root.models) + ? root.models[0] + : undefined; + return { + apiKey: root.apiKey, + baseUrl: firstModel?.baseUrl || root.baseUrl, + }; } else if (appId === "openclaw") { // OpenClaw: settingsConfig 顶层扁平(camelCase,对应 openclaw.json) return { diff --git a/src/components/common/AppToggleGroup.tsx b/src/components/common/AppToggleGroup.tsx index c41eb1a12..e082a43dd 100644 --- a/src/components/common/AppToggleGroup.tsx +++ b/src/components/common/AppToggleGroup.tsx @@ -11,27 +11,44 @@ interface AppToggleGroupProps { apps: Partial>; onToggle: (app: AppId, enabled: boolean) => void; appIds?: AppId[]; + stateByApp?: Partial>; +} + +export interface AppToggleVisualState { + /** 应用实际是否发现/启用了资源。 */ + active: boolean; + /** 用户期望状态;点击切换时以此取反。 */ + desired: boolean; + statusLabel?: string; + warning?: boolean; } export const AppToggleGroup: React.FC = ({ apps, onToggle, appIds = APP_IDS, + stateByApp, }) => { return (
{appIds.map((app) => { const { label, icon, activeClass } = APP_ICON_MAP[app]; - const enabled = apps[app]; + const visualState = stateByApp?.[app]; + const desired = visualState?.desired ?? Boolean(apps[app]); + const active = visualState?.active ?? desired; + const warning = + visualState?.warning ?? (visualState ? active !== desired : false); return ( @@ -39,8 +56,13 @@ export const AppToggleGroup: React.FC = ({

{label} - {enabled ? " ✓" : ""} + {active ? " ✓" : ""}

+ {visualState?.statusLabel && ( +

+ {visualState.statusLabel} +

+ )}
); diff --git a/src/components/prompts/PiNativePromptResources.tsx b/src/components/prompts/PiNativePromptResources.tsx new file mode 100644 index 000000000..bdaadd40d --- /dev/null +++ b/src/components/prompts/PiNativePromptResources.tsx @@ -0,0 +1,411 @@ +import { useEffect, useState } from "react"; +import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; +import { FilePlus2, Loader2, RefreshCw, Trash2 } from "lucide-react"; +import { useTranslation } from "react-i18next"; +import { toast } from "sonner"; +import { ConfirmDialog } from "@/components/ConfirmDialog"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Textarea } from "@/components/ui/textarea"; +import { + promptsApi, + type PiPromptFileKind, + type PiPromptFileSnapshot, + type PiPromptTemplate, +} from "@/lib/api/prompts"; +import { extractErrorMessage } from "@/utils/errorUtils"; + +const EDITABLE_FILES: Array<{ + kind: Exclude; + filename: string; + titleKey: string; + descriptionKey: string; +}> = [ + { + kind: "system_override", + filename: "SYSTEM.md", + titleKey: "pi.prompts.systemOverride", + descriptionKey: "pi.prompts.systemOverrideDescription", + }, + { + kind: "system_append", + filename: "APPEND_SYSTEM.md", + titleKey: "pi.prompts.systemAppend", + descriptionKey: "pi.prompts.systemAppendDescription", + }, +]; + +function mutationError(error: unknown, fallback: string) { + toast.error(extractErrorMessage(error) || fallback); +} + +function PiInstructionFileEditor({ + kind, + filename, + titleKey, + descriptionKey, +}: (typeof EDITABLE_FILES)[number]) { + const { t } = useTranslation(); + const queryClient = useQueryClient(); + const [draft, setDraft] = useState(""); + const [confirmCreate, setConfirmCreate] = useState(false); + const [confirmDelete, setConfirmDelete] = useState(false); + const queryKey = ["pi", "promptFile", kind] as const; + const query = useQuery({ + queryKey, + queryFn: () => promptsApi.getPiPromptFile(kind), + }); + + useEffect(() => { + if (query.data) setDraft(query.data.content); + }, [query.data?.revision]); + + const save = useMutation({ + mutationFn: () => { + const snapshot = query.data; + if (!snapshot) throw new Error(t("pi.prompts.loadFirst")); + return promptsApi.replacePiPromptFile(kind, snapshot.revision, draft); + }, + onSuccess: (snapshot) => { + queryClient.setQueryData(queryKey, snapshot); + setConfirmCreate(false); + toast.success(t("pi.prompts.fileSaved", { filename })); + }, + onError: (error) => { + mutationError(error, t("pi.prompts.saveFailed")); + void query.refetch(); + }, + }); + const remove = useMutation({ + mutationFn: async () => { + const snapshot = query.data; + if (!snapshot) throw new Error(t("pi.prompts.loadFirst")); + await promptsApi.deletePiPromptFile(kind, snapshot.revision); + return promptsApi.getPiPromptFile(kind); + }, + onSuccess: (snapshot) => { + queryClient.setQueryData(queryKey, snapshot); + setDraft(""); + setConfirmDelete(false); + toast.success(t("pi.prompts.fileDeactivated", { filename })); + }, + onError: (error) => { + mutationError(error, t("pi.prompts.deleteFailed")); + void query.refetch(); + }, + }); + + const busy = save.isPending || remove.isPending; + const changed = Boolean(query.data && draft !== query.data.content); + const blank = !draft.trim(); + const requestSave = () => { + if (kind === "system_override" && query.data && !query.data.exists) { + setConfirmCreate(true); + return; + } + save.mutate(); + }; + + return ( +
+
+
+
+

{t(titleKey)}

+ + {query.data?.exists + ? t("pi.prompts.active") + : t("pi.prompts.inactive")} + +
+

+ {t(descriptionKey)} +

+ {query.data?.path && ( + + {query.data.path} + + )} +
+ +
+ + {query.isLoading ? ( +
+ + {t("common.loading")} +
+ ) : ( + <> +