Compare commits

..

7 Commits

Author SHA1 Message Date
Jason e8e4cae41b Merge origin/main into feat/codex-oauth-account-usage 2026-08-03 21:15:48 +08:00
makoMakoGo eb356e15bd fix(skills): resolve source dir by SKILL.md anchor instead of name (#4153)
* fix(skills): resolve source dir by SKILL.md anchor instead of name

resolve_skill_source_dir previously guessed the source dir via root.join(name).is_dir() without verifying SKILL.md, misjudging same-name non-skill dirs (e.g. the ast-grep plugin wrapper dir in ast-grep/agent-skill) and causing install failure #4141.

Now anchors on SKILL.md: direct + SKILL.md check -> root manifest explicit skills[] -> fallback by name -> root fallback. Adds 5 layout tests.

Closes #4141

* fix(skills): drop speculative manifest resolver path

resolve_via_manifest (parsing root .claude-plugin/marketplace.json &
plugin.json explicit skills[]) is inert for the actual #4141 case: the
real ast-grep/agent-skill marketplace.json declares no skills[] array,
so the manifest branch never produces a candidate. The #4141 fix is
delivered entirely by resolve_skill_source_dir step 1's SKILL.md anchor
plus the pre-existing find_skill_dir_by_name DFS.

Keeping the manifest path would pull npx-skills package-parity semantics
(pluginRoot / source / remote-object source / skills[] / "./"-validation
/ ...) into a bug hotfix, with no real manifest proving it is not dead
code. Drop it to keep this PR a focused #4141 hotfix.

- remove SkillMarketplaceMetadata / SkillManifestPlugin /
  SkillMarketplaceManifest, resolve_via_manifest, sanitize_manifest_path
- narrow resolve_skill_source_dir to 3 steps
  (direct+SKILL.md -> by-name DFS+SKILL.md -> root+SKILL.md -> None)
- replace the two synthetic manifest tests with a negative case:
  same-name wrapper dir without SKILL.md and no inner skill -> None

cargo test --lib resolve_skill_source_dir: 7 passed
cargo clippy --lib: clean
2026-08-03 19:05:51 +08:00
mhy1227 f38722a440 feat(pricing): seed Qwen3.8 Max built-in model pricing (#6053)
* feat(pricing): seed Qwen3.8 Max built-in model pricing

Add insert-if-absent row for qwen3.8-max at 2/6 USD per Mtok input/output with 0.20 cache read.

* fix(pricing): set qwen3.8-max cache write to 2.50

Align cache_write with official explicit context-cache rate (125 percent of input). cache_read stays 0.20 (10 percent hit).

* fix(pricing): correct qwen3.8-max cache read price

---------

Co-authored-by: Jason <farion1231@gmail.com>
2026-08-03 17:57:24 +08:00
saladday bc180a3d9d fix(codex-oauth): scope account quota to auth center 2026-07-14 02:57:35 -04:00
saladday c7a2bff78b Merge remote-tracking branch 'origin/main' into pr-4887
# Conflicts:
#	src/lib/query/subscription.ts
2026-07-14 02:56:31 -04:00
SaladDay d52ab6c5f4 refactor(codex-oauth): stable async loading placeholder for account usage
The account header (login + badges + actions) already renders independently
of the usage query — the quota is fetched async via Tauri invoke + React
Query, so the account never waits on it. Make that visually obvious and
jump-free: while the usage loads, show a spinner inside a placeholder shaped
like the final quota card (same rounded-xl / border / bg-card), so the card
morphs smoothly into the data instead of popping in from an empty gap.
2026-07-01 17:33:35 +00:00
SaladDay 0f3991efc3 feat(codex-oauth): show per-account usage in Auth Center
Each ChatGPT (Codex OAuth) account under Settings → 认证 now displays its
own subscription usage — reset countdowns and per-window progress bars —
directly in the account list, instead of usage only being visible on the
active provider card.

- Add useCodexOauthQuotaByAccountId(accountId) and refactor
  useCodexOauthQuota to delegate to it (shared query key → cache reuse)
- Add CodexOauthAccountQuota, a thin per-account wrapper that reuses the
  existing SubscriptionQuotaView expanded layout (same look and 5-state
  handling as provider cards), with a light spinner on first load
- Render it under each account row in CodexOAuthSection; fetch once when
  the Auth Center opens, manual refresh available (no polling)

Copilot is intentionally left out — same as before, this is Codex-only.
2026-07-01 17:07:04 +00:00
188 changed files with 1842 additions and 47454 deletions
-1
View File
@@ -15,7 +15,6 @@
*.ts text eol=lf *.ts text eol=lf
*.tsx text eol=lf *.tsx text eol=lf
*.js text eol=lf *.js text eol=lf
*.mjs text eol=lf
*.jsx text eol=lf *.jsx text eol=lf
# HTML/CSS files # HTML/CSS files
-85
View File
@@ -1,85 +0,0 @@
# 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 snapshotextension
不能在请求中途直接改 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:<id>` 分组,并展示冲突与覆盖次序的实测结果。
- 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 共用该项目:其资源语义、预览与安全面另行立项。
File diff suppressed because it is too large Load Diff
-789
View File
@@ -1,789 +0,0 @@
#!/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, expandPromptTemplate } 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: <range>\n---\nReview $1\n",
);
writeFileSync(
join(promptAgentDir, "prompts", "release notes.md"),
"This spaced filename cannot be addressed as one slash-command token.\n",
);
writeFileSync(join(promptAgentDir, "prompts", "release.v2.md"), "Release $1\n");
writeFileSync(join(promptAgentDir, "prompts", "评审.md"), "评审 $1\n");
writeFileSync(join(promptAgentDir, "prompts", "empty.md"), "");
writeFileSync(
join(promptAgentDir, "prompts", "nested", "ignored.md"),
"nested",
);
const loadedPromptTemplates = adapters.loadPromptTemplates({
cwd: promptProjectDir,
agentDir: promptAgentDir,
promptPaths: [],
includeDefaults: true,
});
const promptTemplateDiscovery = jsonSafeJavaScriptValue(
loadedPromptTemplates.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,
})),
);
const promptTemplateExpansion = [
"/review captured-range",
"/release notes",
"/release.v2 captured-range",
"/评审 变更",
].map((input) => ({
input,
result: adapters.expandPromptTemplate(input, loadedPromptTemplates),
}));
// 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,
promptTemplateExpansion,
emptyInstructionFiles,
sessionDirectorySemantics,
sessionCliSemantics,
sessionBranchSemantics,
malformedSessionSemantics,
nativeToolInventory,
},
null,
2,
),
);
+1 -13
View File
@@ -783,8 +783,6 @@ dependencies = [
"indexmap 2.13.0", "indexmap 2.13.0",
"json-five", "json-five",
"json5", "json5",
"jsonc-parser",
"libc",
"log", "log",
"objc2 0.5.2", "objc2 0.5.2",
"objc2-app-kit 0.2.2", "objc2-app-kit 0.2.2",
@@ -801,7 +799,6 @@ dependencies = [
"serde_yaml", "serde_yaml",
"serial_test", "serial_test",
"sha2", "sha2",
"syn 2.0.117",
"sys-locale", "sys-locale",
"tauri", "tauri",
"tauri-build", "tauri-build",
@@ -2801,15 +2798,6 @@ dependencies = [
"serde", "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]] [[package]]
name = "jsonptr" name = "jsonptr"
version = "0.6.3" version = "0.6.3"
@@ -4698,7 +4686,7 @@ dependencies = [
"errno", "errno",
"libc", "libc",
"linux-raw-sys 0.4.15", "linux-raw-sys 0.4.15",
"windows-sys 0.59.0", "windows-sys 0.52.0",
] ]
[[package]] [[package]]
+1 -7
View File
@@ -23,8 +23,7 @@ test-hooks = []
tauri-build = { version = "2.4.0", features = [] } tauri-build = { version = "2.4.0", features = [] }
[dependencies] [dependencies]
serde_json = { version = "1.0", features = ["arbitrary_precision", "preserve_order"] } serde_json = { version = "1.0", features = ["preserve_order"] }
jsonc-parser = { version = "0.33", features = ["cst", "serde_json"] }
serde = { version = "1.0", features = ["derive"] } serde = { version = "1.0", features = ["derive"] }
log = "0.4" log = "0.4"
chrono = { version = "0.4", features = ["serde"] } chrono = { version = "0.4", features = ["serde"] }
@@ -79,7 +78,6 @@ indexmap = { version = "2", features = ["serde"] }
rust_decimal = "1.33" rust_decimal = "1.33"
uuid = { version = "1.11", features = ["v4"] } uuid = { version = "1.11", features = ["v4"] }
sha2 = "0.10" sha2 = "0.10"
libc = "0.2"
hmac = "0.12" hmac = "0.12"
json5 = "0.4" json5 = "0.4"
json-five = "0.3.1" json-five = "0.3.1"
@@ -96,9 +94,6 @@ winreg = "0.52"
windows-sys = { version = "0.61", features = [ windows-sys = { version = "0.61", features = [
"Win32_Globalization", "Win32_Globalization",
"Win32_Storage_FileSystem", "Win32_Storage_FileSystem",
"Win32_System_Diagnostics_ToolHelp",
"Win32_System_JobObjects",
"Win32_System_Threading",
"Win32_UI_Shell", "Win32_UI_Shell",
] } ] }
@@ -121,4 +116,3 @@ strip = "symbols"
[dev-dependencies] [dev-dependencies]
serial_test = "3" serial_test = "3"
tempfile = "3" tempfile = "3"
syn = { version = "2", features = ["full", "visit"] }
+30 -35
View File
@@ -32,7 +32,6 @@ impl McpApps {
AppType::OpenCode => self.opencode, AppType::OpenCode => self.opencode,
AppType::OpenClaw => false, // OpenClaw doesn't support MCP AppType::OpenClaw => false, // OpenClaw doesn't support MCP
AppType::Hermes => self.hermes, AppType::Hermes => self.hermes,
AppType::Pi => false, // Pi core has no native MCP registry.
AppType::ClaudeDesktop => false, AppType::ClaudeDesktop => false,
} }
} }
@@ -47,7 +46,6 @@ impl McpApps {
AppType::OpenCode => self.opencode = enabled, AppType::OpenCode => self.opencode = enabled,
AppType::OpenClaw => {} // OpenClaw doesn't support MCP, ignore AppType::OpenClaw => {} // OpenClaw doesn't support MCP, ignore
AppType::Hermes => self.hermes = enabled, 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 AppType::ClaudeDesktop => {} // Claude Desktop 3P provider config doesn't support MCP here
} }
} }
@@ -102,8 +100,6 @@ pub struct SkillApps {
pub opencode: bool, pub opencode: bool,
#[serde(default)] #[serde(default)]
pub hermes: bool, pub hermes: bool,
#[serde(default)]
pub pi: bool,
} }
impl SkillApps { impl SkillApps {
@@ -116,7 +112,6 @@ impl SkillApps {
AppType::GrokBuild => self.grokbuild, AppType::GrokBuild => self.grokbuild,
AppType::OpenCode => self.opencode, AppType::OpenCode => self.opencode,
AppType::Hermes => self.hermes, AppType::Hermes => self.hermes,
AppType::Pi => self.pi,
AppType::OpenClaw => false, // OpenClaw doesn't support Skills AppType::OpenClaw => false, // OpenClaw doesn't support Skills
AppType::ClaudeDesktop => false, AppType::ClaudeDesktop => false,
} }
@@ -131,7 +126,6 @@ impl SkillApps {
AppType::GrokBuild => self.grokbuild = enabled, AppType::GrokBuild => self.grokbuild = enabled,
AppType::OpenCode => self.opencode = enabled, AppType::OpenCode => self.opencode = enabled,
AppType::Hermes => self.hermes = enabled, AppType::Hermes => self.hermes = enabled,
AppType::Pi => self.pi = enabled,
AppType::OpenClaw => {} // OpenClaw doesn't support Skills, ignore AppType::OpenClaw => {} // OpenClaw doesn't support Skills, ignore
AppType::ClaudeDesktop => {} // Claude Desktop 3P profiles don't use CC Switch skill sync AppType::ClaudeDesktop => {} // Claude Desktop 3P profiles don't use CC Switch skill sync
} }
@@ -158,9 +152,6 @@ impl SkillApps {
if self.hermes { if self.hermes {
apps.push(AppType::Hermes); apps.push(AppType::Hermes);
} }
if self.pi {
apps.push(AppType::Pi);
}
apps apps
} }
@@ -172,7 +163,6 @@ impl SkillApps {
&& !self.grokbuild && !self.grokbuild
&& !self.opencode && !self.opencode
&& !self.hermes && !self.hermes
&& !self.pi
} }
/// 仅启用指定应用(其他应用设为禁用) /// 仅启用指定应用(其他应用设为禁用)
@@ -367,8 +357,6 @@ pub struct PromptRoot {
pub openclaw: PromptConfig, pub openclaw: PromptConfig,
#[serde(default)] #[serde(default)]
pub hermes: PromptConfig, pub hermes: PromptConfig,
#[serde(default)]
pub pi: PromptConfig,
} }
use crate::config::{copy_file, get_app_config_dir, get_app_config_path, write_json_file}; use crate::config::{copy_file, get_app_config_dir, get_app_config_path, write_json_file};
@@ -393,7 +381,6 @@ pub enum AppType {
OpenCode, OpenCode,
OpenClaw, OpenClaw,
Hermes, Hermes,
Pi,
} }
impl AppType { impl AppType {
@@ -407,7 +394,6 @@ impl AppType {
AppType::OpenCode => "opencode", AppType::OpenCode => "opencode",
AppType::OpenClaw => "openclaw", AppType::OpenClaw => "openclaw",
AppType::Hermes => "hermes", AppType::Hermes => "hermes",
AppType::Pi => "pi",
} }
} }
@@ -433,7 +419,6 @@ impl AppType {
AppType::OpenCode, AppType::OpenCode,
AppType::OpenClaw, AppType::OpenClaw,
AppType::Hermes, AppType::Hermes,
AppType::Pi,
] ]
.into_iter() .into_iter()
} }
@@ -453,11 +438,10 @@ impl FromStr for AppType {
"opencode" => Ok(AppType::OpenCode), "opencode" => Ok(AppType::OpenCode),
"openclaw" => Ok(AppType::OpenClaw), "openclaw" => Ok(AppType::OpenClaw),
"hermes" => Ok(AppType::Hermes), "hermes" => Ok(AppType::Hermes),
"pi" => Ok(AppType::Pi),
other => Err(AppError::localized( other => Err(AppError::localized(
"unsupported_app", "unsupported_app",
format!("不支持的应用标识: '{other}'。可选值: claude, claude-desktop, codex, gemini, grokbuild, opencode, openclaw, hermes, pi"), 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, pi."), format!("Unsupported app id: '{other}'. Allowed: claude, claude-desktop, codex, gemini, grokbuild, opencode, openclaw, hermes."),
)), )),
} }
} }
@@ -483,9 +467,6 @@ pub struct CommonConfigSnippets {
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
pub hermes: Option<String>, pub hermes: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub pi: Option<String>,
} }
impl CommonConfigSnippets { impl CommonConfigSnippets {
@@ -500,7 +481,6 @@ impl CommonConfigSnippets {
AppType::OpenCode => self.opencode.as_ref(), AppType::OpenCode => self.opencode.as_ref(),
AppType::OpenClaw => self.openclaw.as_ref(), AppType::OpenClaw => self.openclaw.as_ref(),
AppType::Hermes => self.hermes.as_ref(), AppType::Hermes => self.hermes.as_ref(),
AppType::Pi => self.pi.as_ref(),
} }
} }
@@ -515,7 +495,6 @@ impl CommonConfigSnippets {
AppType::OpenCode => self.opencode = snippet, AppType::OpenCode => self.opencode = snippet,
AppType::OpenClaw => self.openclaw = snippet, AppType::OpenClaw => self.openclaw = snippet,
AppType::Hermes => self.hermes = snippet, AppType::Hermes => self.hermes = snippet,
AppType::Pi => self.pi = snippet,
} }
} }
} }
@@ -560,7 +539,6 @@ impl Default for MultiAppConfig {
apps.insert("opencode".to_string(), ProviderManager::default()); apps.insert("opencode".to_string(), ProviderManager::default());
apps.insert("openclaw".to_string(), ProviderManager::default()); apps.insert("openclaw".to_string(), ProviderManager::default());
apps.insert("hermes".to_string(), ProviderManager::default()); apps.insert("hermes".to_string(), ProviderManager::default());
apps.insert("pi".to_string(), ProviderManager::default());
Self { Self {
version: 2, version: 2,
@@ -648,12 +626,6 @@ impl MultiAppConfig {
.insert("gemini".to_string(), ProviderManager::default()); .insert("gemini".to_string(), ProviderManager::default());
updated = true; 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 // 执行 MCP 迁移(v3.6.x → v3.7.0
let migrated = config.migrate_mcp_to_unified()?; let migrated = config.migrate_mcp_to_unified()?;
@@ -719,6 +691,34 @@ 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<Self, AppError> { fn default_with_auto_import() -> Result<Self, AppError> {
log::info!("首次启动,创建默认配置并检测提示词文件"); log::info!("首次启动,创建默认配置并检测提示词文件");
@@ -733,7 +733,6 @@ impl MultiAppConfig {
Self::auto_import_prompt_if_exists(&mut config, AppType::OpenCode)?; 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::OpenClaw)?;
Self::auto_import_prompt_if_exists(&mut config, AppType::Hermes)?; Self::auto_import_prompt_if_exists(&mut config, AppType::Hermes)?;
Self::auto_import_prompt_if_exists(&mut config, AppType::Pi)?;
Ok(config) Ok(config)
} }
@@ -758,7 +757,6 @@ impl MultiAppConfig {
|| !self.prompts.opencode.prompts.is_empty() || !self.prompts.opencode.prompts.is_empty()
|| !self.prompts.openclaw.prompts.is_empty() || !self.prompts.openclaw.prompts.is_empty()
|| !self.prompts.hermes.prompts.is_empty() || !self.prompts.hermes.prompts.is_empty()
|| !self.prompts.pi.prompts.is_empty()
{ {
return Ok(false); return Ok(false);
} }
@@ -774,7 +772,6 @@ impl MultiAppConfig {
AppType::OpenCode, AppType::OpenCode,
AppType::OpenClaw, AppType::OpenClaw,
AppType::Hermes, AppType::Hermes,
AppType::Pi,
] { ] {
// 复用已有的单应用导入逻辑 // 复用已有的单应用导入逻辑
if Self::auto_import_prompt_if_exists(self, app)? { if Self::auto_import_prompt_if_exists(self, app)? {
@@ -849,7 +846,6 @@ impl MultiAppConfig {
AppType::OpenCode => &mut config.prompts.opencode.prompts, AppType::OpenCode => &mut config.prompts.opencode.prompts,
AppType::OpenClaw => &mut config.prompts.openclaw.prompts, AppType::OpenClaw => &mut config.prompts.openclaw.prompts,
AppType::Hermes => &mut config.prompts.hermes.prompts, AppType::Hermes => &mut config.prompts.hermes.prompts,
AppType::Pi => &mut config.prompts.pi.prompts,
}; };
prompts.insert(id, prompt); prompts.insert(id, prompt);
@@ -893,7 +889,6 @@ impl MultiAppConfig {
AppType::OpenCode => &self.mcp.opencode.servers, AppType::OpenCode => &self.mcp.opencode.servers,
AppType::OpenClaw => continue, // OpenClaw MCP is still in development, skip AppType::OpenClaw => continue, // OpenClaw MCP is still in development, skip
AppType::Hermes => continue, // Hermes didn't exist in v3.6.x, 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 { for (id, entry) in old_servers {
File diff suppressed because it is too large Load Diff
+19 -52
View File
@@ -10,10 +10,6 @@ use crate::codex_state_db::codex_state_db_paths;
use crate::config::{atomic_write, copy_file, get_app_config_dir}; use crate::config::{atomic_write, copy_file, get_app_config_dir};
use crate::database::{is_official_seed_id, Database}; use crate::database::{is_official_seed_id, Database};
use crate::error::AppError; use crate::error::AppError;
use crate::services::provider::{
provider_row_fingerprint, provider_to_mutation_input,
reconcile_provider_record_with_precondition, ReconcilePrecondition,
};
use crate::settings::{ use crate::settings::{
CodexOfficialHistoryUnifyMigration, CodexProviderTemplateMigration, CodexOfficialHistoryUnifyMigration, CodexProviderTemplateMigration,
CodexThirdPartyHistoryProviderBucketMigration, CodexThirdPartyHistoryProviderBucketMigration,
@@ -667,8 +663,7 @@ fn migrate_codex_provider_templates_to_custom(
let providers = db.get_all_providers("codex")?; let providers = db.get_all_providers("codex")?;
let mut migrated_provider_ids = Vec::new(); let mut migrated_provider_ids = Vec::new();
for (_, mut provider) in providers { for (_, provider) in providers {
let observed_fingerprint = provider_row_fingerprint(&provider);
if provider.category.as_deref() == Some("official") if provider.category.as_deref() == Some("official")
|| is_official_seed_id(&provider.id) || is_official_seed_id(&provider.id)
|| provider.is_codex_oauth() || provider.is_codex_oauth()
@@ -699,21 +694,8 @@ fn migrate_codex_provider_templates_to_custom(
}; };
backup_provider_settings_config(&provider.id, &provider.settings_config, backup_root)?; backup_provider_settings_config(&provider.id, &provider.settings_config, backup_root)?;
obj.insert("config".to_string(), Value::String(migrated_config_text)); obj.insert("config".to_string(), Value::String(migrated_config_text));
let provider_id = provider.id.clone(); db.update_provider_settings_config("codex", &provider.id, &settings)?;
provider.settings_config = settings; migrated_provider_ids.push(provider.id);
if let Some(meta) = provider.meta.as_mut() {
meta.custom_endpoints.clear();
}
let input = provider_to_mutation_input(provider);
reconcile_provider_record_with_precondition(
db,
"codex",
input,
ReconcilePrecondition::ExpectPresent {
fingerprint: observed_fingerprint,
},
)?;
migrated_provider_ids.push(provider_id);
} }
Ok(CodexProviderTemplateBucketMigrationOutcome { Ok(CodexProviderTemplateBucketMigrationOutcome {
@@ -1457,8 +1439,7 @@ base_url = "https://proxy.example/v1"
), ),
]; ];
for provider in providers { for provider in providers {
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
} }
let mut official = Provider::with_id( let mut official = Provider::with_id(
@@ -1468,8 +1449,7 @@ base_url = "https://proxy.example/v1"
None, None,
); );
official.category = Some("official".to_string()); official.category = Some("official".to_string());
db.reconcile_provider_fixture("codex", &official) db.save_provider("codex", &official).expect("save official");
.expect("save official");
let source_provider_ids = collect_source_model_provider_ids(&db).expect("collect ids"); let source_provider_ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert_eq!( assert_eq!(
@@ -2191,10 +2171,9 @@ base_url = "https://proxy.example/v1"
); );
official.category = Some("official".to_string()); official.category = Some("official".to_string());
db.reconcile_provider_fixture("codex", &third_party) db.save_provider("codex", &third_party)
.expect("save third-party"); .expect("save third-party");
db.reconcile_provider_fixture("codex", &official) db.save_provider("codex", &official).expect("save official");
.expect("save official");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(ids.contains("rightcode")); assert!(ids.contains("rightcode"));
@@ -2217,8 +2196,7 @@ base_url = "https://proxy.example/v1"
); );
provider.category = Some("aggregator".to_string()); provider.category = Some("aggregator".to_string());
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(!ids.contains("my-private-relay")); assert!(!ids.contains("my-private-relay"));
@@ -2238,8 +2216,7 @@ base_url = "https://proxy.example/v1"
); );
provider.category = Some("aggregator".to_string()); provider.category = Some("aggregator".to_string());
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(!ids.contains("my-private-relay")); assert!(!ids.contains("my-private-relay"));
@@ -2267,8 +2244,7 @@ model_provider = "my-private-relay"
); );
provider.category = Some("aggregator".to_string()); provider.category = Some("aggregator".to_string());
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(!ids.contains("my-private-relay")); assert!(!ids.contains("my-private-relay"));
@@ -2288,8 +2264,7 @@ model_provider = "my-private-relay"
); );
provider.category = Some("aggregator".to_string()); provider.category = Some("aggregator".to_string());
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(ids.contains("aihubmix")); assert!(ids.contains("aihubmix"));
@@ -2310,8 +2285,7 @@ model_provider = "my-private-relay"
); );
provider.category = Some("aggregator".to_string()); provider.category = Some("aggregator".to_string());
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(ids.contains("ccswitch")); assert!(ids.contains("ccswitch"));
@@ -2343,8 +2317,7 @@ model = "gpt-5.4"
}), }),
None, None,
); );
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let (outcome, backup_dir) = migrate_provider_templates_for_test(&db); let (outcome, backup_dir) = migrate_provider_templates_for_test(&db);
assert_eq!(outcome.migrated_provider_ids, vec!["legacy".to_string()]); assert_eq!(outcome.migrated_provider_ids, vec!["legacy".to_string()]);
@@ -2417,8 +2390,7 @@ base_url = "https://aihubmix.example/v1"
}), }),
None, None,
); );
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db);
assert_eq!( assert_eq!(
@@ -2474,8 +2446,7 @@ base_url = "http://localhost:8080/v1"
}), }),
None, None,
); );
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db);
assert!(outcome.migrated_provider_ids.is_empty()); assert!(outcome.migrated_provider_ids.is_empty());
@@ -2524,8 +2495,7 @@ base_url = "https://proxy.example/v1"
}), }),
None, None,
); );
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db);
assert!(outcome.migrated_provider_ids.is_empty()); assert!(outcome.migrated_provider_ids.is_empty());
@@ -2582,8 +2552,7 @@ model_provider = "aihubmix"
}), }),
None, None,
); );
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db);
assert_eq!(outcome.migrated_provider_ids, vec!["profiled".to_string()]); assert_eq!(outcome.migrated_provider_ids, vec!["profiled".to_string()]);
@@ -2632,8 +2601,7 @@ model_provider = "aihubmix"
provider.category = Some("custom".to_string()); provider.category = Some("custom".to_string());
provider.created_at = Some(1); provider.created_at = Some(1);
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(!ids.contains("my-private-relay")); assert!(!ids.contains("my-private-relay"));
@@ -2654,8 +2622,7 @@ model_provider = "aihubmix"
); );
provider.category = Some("custom".to_string()); provider.category = Some("custom".to_string());
db.reconcile_provider_fixture("codex", &provider) db.save_provider("codex", &provider).expect("save provider");
.expect("save provider");
let ids = collect_source_model_provider_ids(&db).expect("collect ids"); let ids = collect_source_model_provider_ids(&db).expect("collect ids");
assert!(!ids.contains("my-local-relay")); assert!(!ids.contains("my-local-relay"));
-14
View File
@@ -135,18 +135,6 @@ pub async fn get_config_status(
Ok(ConfigStatus { exists, path }) 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,
})
}
} }
} }
@@ -168,7 +156,6 @@ pub async fn get_config_dir(app: String) -> Result<String, String> {
AppType::OpenCode => crate::opencode_config::get_opencode_dir(), AppType::OpenCode => crate::opencode_config::get_opencode_dir(),
AppType::OpenClaw => crate::openclaw_config::get_openclaw_dir(), AppType::OpenClaw => crate::openclaw_config::get_openclaw_dir(),
AppType::Hermes => crate::hermes_config::get_hermes_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()) Ok(dir.to_string_lossy().to_string())
@@ -187,7 +174,6 @@ pub async fn open_config_folder(handle: AppHandle, app: String) -> Result<bool,
AppType::OpenCode => crate::opencode_config::get_opencode_dir(), AppType::OpenCode => crate::opencode_config::get_opencode_dir(),
AppType::OpenClaw => crate::openclaw_config::get_openclaw_dir(), AppType::OpenClaw => crate::openclaw_config::get_openclaw_dir(),
AppType::Hermes => crate::hermes_config::get_hermes_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() { if !config_dir.exists() {
-379
View File
@@ -2,7 +2,6 @@
//! //!
//! 管理代理模式下的故障转移队列(基于 providers 表的 in_failover_queue 字段) //! 管理代理模式下的故障转移队列(基于 providers 表的 in_failover_queue 字段)
use crate::app_config::AppType;
use crate::database::FailoverQueueItem; use crate::database::FailoverQueueItem;
use crate::provider::Provider; use crate::provider::Provider;
use crate::store::AppState; use crate::store::AppState;
@@ -40,50 +39,6 @@ pub async fn add_to_failover_queue(
app_type: String, app_type: String,
provider_id: String, provider_id: String,
) -> Result<(), 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 state
.db .db
.add_to_failover_queue(&app_type, &provider_id) .add_to_failover_queue(&app_type, &provider_id)
@@ -97,42 +52,6 @@ pub async fn remove_from_failover_queue(
app_type: String, app_type: String,
provider_id: String, provider_id: String,
) -> Result<(), 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 state
.db .db
.remove_from_failover_queue(&app_type, &provider_id) .remove_from_failover_queue(&app_type, &provider_id)
@@ -145,9 +64,6 @@ pub async fn get_auto_failover_enabled(
state: tauri::State<'_, AppState>, state: tauri::State<'_, AppState>,
app_type: String, app_type: String,
) -> Result<bool, String> { ) -> Result<bool, String> {
if app_type == "pi" {
return Ok(crate::settings::get_pi_proxy_settings().auto_failover_enabled);
}
state state
.db .db
.get_proxy_config_for_app(&app_type) .get_proxy_config_for_app(&app_type)
@@ -170,10 +86,6 @@ pub async fn set_auto_failover_enabled(
"[Failover] Setting auto_failover_enabled: app_type='{app_type}', enabled={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 let mut config = state
.db .db
@@ -268,294 +180,3 @@ pub async fn set_auto_failover_enabled(
Ok(()) Ok(())
} }
async fn set_pi_auto_failover_enabled(
app: &tauri::AppHandle,
state: &AppState,
enabled: bool,
) -> Result<(), String> {
let selected_provider = set_pi_auto_failover_enabled_inner(state, enabled).await?;
let _ = app.emit(
"provider-switched",
serde_json::json!({
"appType": "pi",
"providerId": selected_provider,
"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(())
}
async fn set_pi_auto_failover_enabled_inner(
state: &AppState,
enabled: bool,
) -> Result<Option<String>, 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 previous_provider =
crate::services::pi_catalog::PiCatalogCoordinator::current_native_provider(state)
.map_err(|error| error.to_string())?;
let mut auto_added = None;
let mut switched_primary = false;
let selected_provider = if enabled {
let previous_provider = previous_provider.clone().ok_or_else(|| {
"Pi has no current provider, so failover cannot select queue P1".to_string()
})?;
let mut queue = state
.db
.get_failover_queue("pi")
.map_err(|error| error.to_string())?;
if queue.is_empty() {
state
.db
.add_to_failover_queue("pi", &previous_provider)
.map_err(|error| error.to_string())?;
auto_added = Some(previous_provider.clone());
queue = state
.db
.get_failover_queue("pi")
.map_err(|error| error.to_string())?;
}
let primary = queue
.first()
.map(|item| item.provider_id.clone())
.ok_or_else(|| "Pi failover queue is empty".to_string())?;
if primary != previous_provider {
if let Err(error) = set_pi_default_under_switch_guard(state, &guard, &primary) {
if let Some(provider_id) = auto_added.take() {
let _ = state.db.remove_from_failover_queue("pi", &provider_id);
}
return Err(error);
}
switched_primary = true;
}
Some(primary)
} else {
previous_provider.clone()
};
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.take() {
let _ = state.db.remove_from_failover_queue("pi", &provider_id);
}
let rollback = rollback_pi_failover_primary(
state,
&guard,
previous_provider.as_deref(),
switched_primary,
Some(epoch),
)
.await;
return Err(with_pi_failover_rollback(error.to_string(), rollback));
}
if let Err(error) = state
.proxy_service
.reconcile_pi_runtime_at_epoch(epoch)
.await
{
let settings_rollback = crate::settings::update_pi_proxy_settings(previous_config)
.err()
.map(|rollback_error| rollback_error.to_string());
if let Some(provider_id) = auto_added.take() {
let _ = state.db.remove_from_failover_queue("pi", &provider_id);
}
let primary_rollback = rollback_pi_failover_primary(
state,
&guard,
previous_provider.as_deref(),
switched_primary,
None,
)
.await;
let rollback = settings_rollback.or(primary_rollback);
return Err(with_pi_failover_rollback(
format!("Pi failover preference changed but runtime publication failed: {error}"),
rollback,
));
}
Ok(selected_provider)
}
fn set_pi_default_under_switch_guard(
state: &AppState,
guard: &tokio::sync::OwnedMutexGuard<()>,
provider_id: &str,
) -> Result<(), String> {
let aggregate = state
.db
.get_provider_aggregate(AppType::Pi.as_str(), provider_id)
.map_err(|error| error.to_string())?
.ok_or_else(|| format!("Pi provider does not exist: {provider_id}"))?;
let config: crate::pi_config::model::PiManagedProviderConfig =
serde_json::from_value(aggregate.provider.settings_config)
.map_err(|error| format!("managed Pi provider '{provider_id}' is invalid: {error}"))?;
let model_id = config
.models
.first()
.map(|model| model.id.clone())
.ok_or_else(|| format!("Pi provider '{provider_id}' has no selectable models"))?;
crate::services::pi_catalog::PiCatalogCoordinator::apply_under_switch_guard(
state,
guard,
crate::services::pi_catalog::PiCatalogMutation::SetDefault {
provider_id: provider_id.to_string(),
model_id,
},
)
.map(|_| ())
.map_err(|error| error.to_string())
}
async fn rollback_pi_failover_primary(
state: &AppState,
guard: &tokio::sync::OwnedMutexGuard<()>,
previous_provider: Option<&str>,
switched_primary: bool,
pending_epoch: Option<u64>,
) -> Option<String> {
if switched_primary {
let previous_provider =
previous_provider.expect("switching P1 requires a previous Pi provider");
return set_pi_default_under_switch_guard(state, guard, previous_provider)
.err()
.map(|error| format!("primary rollback failed: {error}"));
}
let epoch = match pending_epoch {
Some(epoch) => epoch,
None => state.proxy_service.begin_pi_catalog_mutation().await,
};
state
.proxy_service
.reconcile_pi_runtime_at_epoch(epoch)
.await
.err()
.map(|error| format!("runtime rollback failed: {error}"))
}
fn with_pi_failover_rollback(error: String, rollback: Option<String>) -> String {
match rollback {
Some(rollback) => format!("{error}; {rollback}"),
None => error,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::database::Database;
use crate::provider::ProviderMutationInput;
use crate::services::pi_catalog::{PiCatalogCoordinator, PiCatalogMutation};
use serde_json::json;
use std::sync::Arc;
struct TestHome(Option<std::ffi::OsString>);
impl TestHome {
fn install(path: &std::path::Path) -> Result<Self, crate::error::AppError> {
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) -> ProviderMutationInput {
ProviderMutationInput {
id: id.to_string(),
name: id.to_string(),
settings_config: json!({
"name": id,
"api": "openai-responses",
"baseUrl": format!("https://{id}.example/v1"),
"apiKey": "literal-key",
"models": [{"id": format!("{id}-model"), "name": id}]
}),
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,
}
}
#[tokio::test]
#[serial_test::serial]
async fn enabling_pi_failover_selects_queue_p1() -> Result<(), crate::error::AppError> {
let temp = tempfile::tempdir().expect("tempdir");
let _home = TestHome::install(temp.path())?;
let mut settings = crate::settings::get_settings();
settings.pi_config_dir = Some(temp.path().join("pi-agent").to_string_lossy().into_owned());
settings.pi_takeover_enabled = false;
crate::settings::update_settings(settings)?;
let state = AppState::new(Arc::new(Database::memory()?));
for (provider_id, activate_if_first) in [("provider-a", true), ("provider-b", false)] {
PiCatalogCoordinator::apply(
&state,
PiCatalogMutation::CreateProvider {
input: managed_input(provider_id),
provider_key: provider_id.to_string(),
activate_if_first,
},
)?;
}
state.db.add_to_failover_queue("pi", "provider-b")?;
let mut proxy_config = state.db.get_global_proxy_config().await?;
proxy_config.listen_port = 0;
state.db.update_global_proxy_config(proxy_config).await?;
state
.proxy_service
.set_takeover_for_app("pi", true)
.await
.map_err(crate::error::AppError::Message)?;
let selected = set_pi_auto_failover_enabled_inner(&state, true)
.await
.map_err(crate::error::AppError::Message)?;
assert_eq!(selected.as_deref(), Some("provider-b"));
assert_eq!(
PiCatalogCoordinator::current_native_provider(&state)?.as_deref(),
Some("provider-b")
);
assert!(crate::settings::get_pi_proxy_settings().auto_failover_enabled);
state
.proxy_service
.set_takeover_for_app("pi", false)
.await
.map_err(crate::error::AppError::Message)?;
Ok(())
}
}
+17 -90
View File
@@ -12,7 +12,6 @@ use crate::database::backup::BackupEntry;
use crate::database::Database; use crate::database::Database;
use crate::error::AppError; use crate::error::AppError;
use crate::services::provider::ProviderService; use crate::services::provider::ProviderService;
use crate::services::skill_deployment::PiSkillDeploymentService;
use crate::store::AppState; use crate::store::AppState;
// ─── File import/export ────────────────────────────────────── // ─── File import/export ──────────────────────────────────────
@@ -26,7 +25,7 @@ pub async fn export_config_to_file(
let db = state.db.clone(); let db = state.db.clone();
tauri::async_runtime::spawn_blocking(move || { tauri::async_runtime::spawn_blocking(move || {
let target_path = PathBuf::from(&filePath); let target_path = PathBuf::from(&filePath);
db.export_portable_sql(&target_path)?; db.export_sql(&target_path)?;
Ok::<_, AppError>(json!({ Ok::<_, AppError>(json!({
"success": true, "success": true,
"message": "SQL exported successfully", "message": "SQL exported successfully",
@@ -45,57 +44,26 @@ pub async fn import_config_from_file(
state: State<'_, AppState>, state: State<'_, AppState>,
) -> Result<Value, String> { ) -> Result<Value, String> {
let db = state.db.clone(); let db = state.db.clone();
let app_state = state.inner().clone(); let db_for_sync = db.clone();
let pi_guard = app_state tauri::async_runtime::spawn_blocking(move || {
.proxy_service let path_buf = PathBuf::from(&filePath);
.lock_switch_for_app(crate::app_config::AppType::Pi.as_str()) let backup_id = db.import_sql(&path_buf)?;
.await; let warning = post_sync_warning_from_result(Ok(run_post_import_sync(db_for_sync)));
app_state if let Some(msg) = warning.as_ref() {
.proxy_service log::warn!("[Import] post-import sync warning: {msg}");
.prepare_pi_portable_import_under_lock(&pi_guard) }
.await Ok::<_, AppError>(success_payload_with_warning(backup_id, warning))
.map_err(|error| format!("导入前恢复 Pi 直连投影失败: {error}"))?;
let import_path = filePath.clone();
let import_result = tauri::async_runtime::spawn_blocking(move || {
PiSkillDeploymentService::import_portable_sql(&db, &PathBuf::from(import_path))
}) })
.await .await
.map_err(|error| AppError::Message(format!("SQL import task failed: {error}"))) .map_err(|e| format!("导入配置失败: {e}"))?
.and_then(|result| result); .map_err(|e: AppError| e.to_string())
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}")
}
});
}
};
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] #[tauri::command]
pub async fn sync_current_providers_live(state: State<'_, AppState>) -> Result<Value, String> { pub async fn sync_current_providers_live(state: State<'_, AppState>) -> Result<Value, String> {
let app_state = state.inner().clone(); let db = state.db.clone();
tauri::async_runtime::spawn_blocking(move || { tauri::async_runtime::spawn_blocking(move || {
let app_state = AppState::new(db);
ProviderService::sync_current_to_live(&app_state)?; ProviderService::sync_current_to_live(&app_state)?;
Ok::<_, AppError>(json!({ Ok::<_, AppError>(json!({
"success": true, "success": true,
@@ -186,51 +154,10 @@ pub async fn restore_db_backup(
filename: String, filename: String,
) -> Result<String, String> { ) -> Result<String, String> {
let db = state.db.clone(); let db = state.db.clone();
let app_state = state.inner().clone(); tauri::async_runtime::spawn_blocking(move || db.restore_from_backup(&filename))
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 .await
.map_err(|error| format!("Restore preparation failed: {error}"))?; .map_err(|e| format!("Restore failed: {e}"))?
.map_err(|e: AppError| e.to_string())
let restore_result = tauri::async_runtime::spawn_blocking(move || {
PiSkillDeploymentService::restore_binary_backup_without_pi_ownership(&db, &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 /// Rename a database backup file
+2 -39
View File
@@ -111,8 +111,8 @@ pub struct ToolVersion {
wsl_distro: Option<String>, wsl_distro: Option<String>,
} }
const VALID_TOOLS: [&str; 8] = [ const VALID_TOOLS: [&str; 7] = [
"claude", "codex", "gemini", "grok", "opencode", "openclaw", "hermes", "pi", "claude", "codex", "gemini", "grok", "opencode", "openclaw", "hermes",
]; ];
#[derive(Debug, Clone, serde::Deserialize)] #[derive(Debug, Clone, serde::Deserialize)]
@@ -433,7 +433,6 @@ fn tool_display_name(tool: &str) -> &'static str {
"opencode" => "OpenCode", "opencode" => "OpenCode",
"openclaw" => "OpenClaw", "openclaw" => "OpenClaw",
"hermes" => "Hermes", "hermes" => "Hermes",
"pi" => "Pi",
_ => "Unknown", _ => "Unknown",
} }
} }
@@ -514,7 +513,6 @@ fn npm_install_command_for(tool: &str) -> Option<&'static str> {
"grok" => Some("npm i -g @xai-official/grok@latest"), "grok" => Some("npm i -g @xai-official/grok@latest"),
"opencode" => Some("npm i -g opencode-ai@latest"), "opencode" => Some("npm i -g opencode-ai@latest"),
"openclaw" => Some("npm i -g openclaw@latest"), "openclaw" => Some("npm i -g openclaw@latest"),
"pi" => Some("npm i -g @earendil-works/pi-coding-agent@latest"),
_ => None, _ => None,
} }
} }
@@ -809,9 +807,6 @@ async fn get_single_tool_version_impl(
} }
"openclaw" => fetch_npm_latest_for_tool(&client, "openclaw", tool, local).await, "openclaw" => fetch_npm_latest_for_tool(&client, "openclaw", tool, local).await,
"hermes" => fetch_pypi_latest_version(&client, "hermes-agent").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, _ => None,
}; };
@@ -2076,7 +2071,6 @@ fn npm_package_for(tool: &str) -> Option<&'static str> {
"grok" => Some("@xai-official/grok"), "grok" => Some("@xai-official/grok"),
"opencode" => Some("opencode-ai"), "opencode" => Some("opencode-ai"),
"openclaw" => Some("openclaw"), "openclaw" => Some("openclaw"),
"pi" => Some("@earendil-works/pi-coding-agent"),
_ => None, _ => None,
} }
} }
@@ -2795,7 +2789,6 @@ fn wsl_distro_for_tool(tool: &str) -> Option<String> {
"opencode" => crate::settings::get_opencode_override_dir(), "opencode" => crate::settings::get_opencode_override_dir(),
"openclaw" => crate::settings::get_openclaw_override_dir(), "openclaw" => crate::settings::get_openclaw_override_dir(),
"hermes" => crate::settings::get_hermes_override_dir(), "hermes" => crate::settings::get_hermes_override_dir(),
"pi" => crate::settings::get_pi_override_dir(),
_ => None, _ => None,
}?; }?;
@@ -3933,24 +3926,6 @@ 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] #[test]
fn test_compare_semver() { fn test_compare_semver() {
use std::cmp::Ordering; use std::cmp::Ordering;
@@ -5356,13 +5331,6 @@ mod tests {
assert_eq!(cmd, "npm i -g openclaw@latest"); 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] #[test]
fn update_fallbacks_use_official_cli_only_when_supported() { fn update_fallbacks_use_official_cli_only_when_supported() {
assert_eq!( assert_eq!(
@@ -5392,11 +5360,6 @@ mod tests {
static_fallback_command("openclaw"), static_fallback_command("openclaw"),
"openclaw update --yes || npm i -g openclaw@latest" "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] #[test]
-2
View File
@@ -17,7 +17,6 @@ mod misc;
mod model_fetch; mod model_fetch;
mod omo; mod omo;
mod openclaw; mod openclaw;
mod pi;
mod plugin; mod plugin;
mod profile; mod profile;
mod prompt; mod prompt;
@@ -54,7 +53,6 @@ pub use misc::*;
pub use model_fetch::*; pub use model_fetch::*;
pub use omo::*; pub use omo::*;
pub use openclaw::*; pub use openclaw::*;
pub(crate) use pi::*;
pub use plugin::*; pub use plugin::*;
pub use profile::*; pub use profile::*;
pub use prompt::*; pub use prompt::*;
-76
View File
@@ -1,76 +0,0 @@
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<Vec<PiNativeDiagnostic>, 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<String, String> {
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<bool, String> {
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<PiNativeDefaults, String> {
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<bool, String> {
state
.proxy_service
.rotate_pi_gateway_token()
.await
.map(|()| true)
.map_err(|error| error.to_string())
}
+1 -63
View File
@@ -5,11 +5,7 @@ use tauri::State;
use crate::app_config::AppType; use crate::app_config::AppType;
use crate::prompt::Prompt; use crate::prompt::Prompt;
use crate::services::pi_prompt_files::{ use crate::services::PromptService;
PiPromptFileKind, PiPromptFileService, PiPromptFileSnapshot, PiPromptTemplate,
PiPromptTemplateService,
};
use crate::services::prompt::{PiPromptLibraryStatus, PromptService};
use crate::store::AppState; use crate::store::AppState;
#[tauri::command] #[tauri::command]
@@ -66,61 +62,3 @@ pub async fn get_current_prompt_file_content(app: String) -> Result<Option<Strin
let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?; let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?;
PromptService::get_current_file_content(app_type).map_err(|e| e.to_string()) PromptService::get_current_file_content(app_type).map_err(|e| e.to_string())
} }
#[tauri::command]
pub async fn get_pi_prompt_library_status(
state: State<'_, AppState>,
) -> Result<PiPromptLibraryStatus, String> {
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<PiPromptFileSnapshot, String> {
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<PiPromptFileSnapshot, String> {
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<bool, String> {
PiPromptFileService::delete(kind, &expectedRevision).map_err(|error| error.to_string())
}
#[tauri::command]
pub async fn list_pi_prompt_templates() -> Result<Vec<PiPromptTemplate>, 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<PiPromptTemplate, String> {
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<bool, String> {
PiPromptTemplateService::delete(&slug, &expectedRevision).map_err(|error| error.to_string())
}
+4 -11
View File
@@ -4,9 +4,8 @@ use tauri::{Emitter, Manager, State};
use crate::app_config::AppType; use crate::app_config::AppType;
use crate::commands::copilot::CopilotAuthState; use crate::commands::copilot::CopilotAuthState;
use crate::commands::xai_oauth::XaiOAuthState; use crate::commands::xai_oauth::XaiOAuthState;
use crate::database::NewProviderAggregate;
use crate::error::AppError; use crate::error::AppError;
use crate::provider::{ClaudeDesktopMode, Provider, ProviderMutationInput}; use crate::provider::{ClaudeDesktopMode, Provider};
use crate::services::{ use crate::services::{
EndpointLatency, ProviderService, ProviderSortUpdate, SpeedtestService, SwitchResult, EndpointLatency, ProviderService, ProviderSortUpdate, SpeedtestService, SwitchResult,
}; };
@@ -40,7 +39,7 @@ pub fn get_current_provider(state: State<'_, AppState>, app: String) -> Result<S
pub fn add_provider( pub fn add_provider(
state: State<'_, AppState>, state: State<'_, AppState>,
app: String, app: String,
provider: ProviderMutationInput, provider: Provider,
#[allow(non_snake_case)] addToLive: Option<bool>, #[allow(non_snake_case)] addToLive: Option<bool>,
) -> Result<bool, String> { ) -> Result<bool, String> {
let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?; let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?;
@@ -52,7 +51,7 @@ pub fn add_provider(
pub fn update_provider( pub fn update_provider(
state: State<'_, AppState>, state: State<'_, AppState>,
app: String, app: String,
provider: ProviderMutationInput, provider: Provider,
#[allow(non_snake_case)] originalId: Option<String>, #[allow(non_snake_case)] originalId: Option<String>,
) -> Result<bool, String> { ) -> Result<bool, String> {
let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?; let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?;
@@ -251,13 +250,7 @@ pub fn import_claude_desktop_providers_from_claude(
state state
.db .db
.create_provider( .save_provider(AppType::ClaudeDesktop.as_str(), &desktop_provider)
NewProviderAggregate::from_input(
AppType::ClaudeDesktop.as_str(),
crate::services::provider::provider_to_mutation_input(desktop_provider),
)
.map_err(|e| e.to_string())?,
)
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
imported += 1; imported += 1;
} }
-129
View File
@@ -26,7 +26,6 @@ pub async fn stop_proxy_server(state: tauri::State<'_, AppState>) -> Result<(),
|| takeover.grokbuild || takeover.grokbuild
|| takeover.opencode || takeover.opencode
|| takeover.openclaw || takeover.openclaw
|| takeover.pi
{ {
return Err( return Err(
"仍有应用处于代理接管状态,请先在设置中关闭对应应用接管后再停止本地路由。".to_string(), "仍有应用处于代理接管状态,请先在设置中关闭对应应用接管后再停止本地路由。".to_string(),
@@ -121,9 +120,6 @@ pub async fn get_proxy_config_for_app(
state: tauri::State<'_, AppState>, state: tauri::State<'_, AppState>,
app_type: String, app_type: String,
) -> Result<AppProxyConfig, String> { ) -> Result<AppProxyConfig, String> {
if app_type == "pi" {
return Ok(crate::settings::get_pi_app_proxy_config());
}
let db = &state.db; let db = &state.db;
db.get_proxy_config_for_app(&app_type) db.get_proxy_config_for_app(&app_type)
.await .await
@@ -142,61 +138,6 @@ pub async fn update_proxy_config_for_app(
let app_type = config.app_type.clone(); let app_type = config.app_type.clone();
let circuit_config = CircuitBreakerConfig::from(&config); 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,
..previous.clone()
};
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) db.update_proxy_config_for_app(config)
.await .await
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
@@ -211,9 +152,6 @@ async fn get_default_cost_multiplier_internal(
state: &AppState, state: &AppState,
app_type: &str, app_type: &str,
) -> Result<String, AppError> { ) -> Result<String, AppError> {
if app_type == "pi" {
return Ok(crate::settings::get_pi_default_cost_multiplier());
}
let db = &state.db; let db = &state.db;
db.get_default_cost_multiplier(app_type).await db.get_default_cost_multiplier(app_type).await
} }
@@ -242,9 +180,6 @@ async fn set_default_cost_multiplier_internal(
app_type: &str, app_type: &str,
value: &str, value: &str,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
if app_type == "pi" {
return crate::settings::set_pi_default_cost_multiplier(value);
}
let db = &state.db; let db = &state.db;
db.set_default_cost_multiplier(app_type, value).await db.set_default_cost_multiplier(app_type, value).await
} }
@@ -274,9 +209,6 @@ async fn get_pricing_model_source_internal(
state: &AppState, state: &AppState,
app_type: &str, app_type: &str,
) -> Result<String, AppError> { ) -> Result<String, AppError> {
if app_type == "pi" {
return Ok(crate::settings::get_pi_pricing_model_source());
}
let db = &state.db; let db = &state.db;
db.get_pricing_model_source(app_type).await db.get_pricing_model_source(app_type).await
} }
@@ -305,9 +237,6 @@ async fn set_pricing_model_source_internal(
app_type: &str, app_type: &str,
value: &str, value: &str,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
if app_type == "pi" {
return crate::settings::set_pi_pricing_model_source(value);
}
let db = &state.db; let db = &state.db;
db.set_pricing_model_source(app_type, value).await db.set_pricing_model_source(app_type, value).await
} }
@@ -524,61 +453,3 @@ pub async fn get_circuit_breaker_stats(
let _ = (state, provider_id, app_type); let _ = (state, provider_id, app_type);
Ok(None) Ok(None)
} }
#[cfg(test)]
mod tests {
use super::*;
use crate::database::{lock_conn, Database};
use std::sync::Arc;
struct TestHome(Option<std::ffi::OsString>);
impl TestHome {
fn install(path: &std::path::Path) -> Result<Self, AppError> {
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();
}
}
#[tokio::test]
#[serial_test::serial]
async fn pi_pricing_round_trips_without_out_of_schema_proxy_row() -> Result<(), AppError> {
let temp = tempfile::tempdir().expect("tempdir");
let _home = TestHome::install(temp.path())?;
let state = AppState::new(Arc::new(Database::memory()?));
set_default_cost_multiplier_test_hook(&state, "pi", "1.25").await?;
set_pricing_model_source_test_hook(&state, "pi", "request").await?;
assert_eq!(
get_default_cost_multiplier_test_hook(&state, "pi").await?,
"1.25"
);
assert_eq!(
get_pricing_model_source_test_hook(&state, "pi").await?,
"request"
);
let conn = lock_conn!(state.db.conn);
let pi_rows: i64 = conn
.query_row(
"SELECT COUNT(*) FROM proxy_config WHERE app_type = 'pi'",
[],
|row| row.get(0),
)
.map_err(|error| AppError::Database(error.to_string()))?;
assert_eq!(pi_rows, 0);
Ok(())
}
}
+5 -31
View File
@@ -107,44 +107,18 @@ pub async fn s3_sync_upload(state: State<'_, AppState>) -> Result<Value, String>
#[tauri::command] #[tauri::command]
pub async fn s3_sync_download(state: State<'_, AppState>) -> Result<Value, String> { pub async fn s3_sync_download(state: State<'_, AppState>) -> Result<Value, String> {
let db = state.db.clone(); let db = state.db.clone();
let app_state = state.inner().clone(); let db_for_sync = db.clone();
let mut settings = require_enabled_s3_settings()?; let mut settings = require_enabled_s3_settings()?;
let _auto_sync_suppression = crate::services::s3_auto_sync::AutoSyncSuppressionGuard::new(); 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 sync_result = run_with_s3_lock(s3_sync_service::download(&db, &mut settings)).await;
let mut result = match sync_result { let mut result = map_sync_result(sync_result, |error| {
Ok(result) => result, persist_sync_error(&mut settings, error, "manual")
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. // Post-download sync is best-effort: snapshot restore has already succeeded.
let sync_state = app_state.clone();
let warning = post_sync_warning_from_result( let warning = post_sync_warning_from_result(
tauri::async_runtime::spawn_blocking(move || run_post_import_sync(&sync_state)) tauri::async_runtime::spawn_blocking(move || run_post_import_sync(db_for_sync))
.await .await
.map_err(|e| e.to_string()), .map_err(|e| e.to_string()),
); );
+2 -59
View File
@@ -48,13 +48,6 @@ fn merge_settings_for_save(
// 开关)后、前端 query 缓存刷新前的一次全量保存会把旧 marker 重放回来, // 开关)后、前端 query 缓存刷新前的一次全量保存会把旧 marker 重放回来,
// 重新开启时被"复活"的标记挡住而漏迁。 // 重新开启时被"复活"的标记挡住而漏迁。
incoming.local_migrations = existing.local_migrations.clone(); 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 incoming
} }
@@ -70,29 +63,12 @@ pub async fn save_settings(
state: tauri::State<'_, crate::store::AppState>, state: tauri::State<'_, crate::store::AppState>,
settings: crate::settings::AppSettings, settings: crate::settings::AppSettings,
) -> Result<bool, String> { ) -> Result<bool, String> {
// 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 existing = crate::settings::get_settings();
let merged = merge_settings_for_save(settings, &existing); let merged = merge_settings_for_save(settings, &existing);
let unify_codex_changed = let unify_codex_changed =
merged.unify_codex_session_history != existing.unify_codex_session_history; merged.unify_codex_session_history != existing.unify_codex_session_history;
let unify_codex_enabled = merged.unify_codex_session_history; let unify_codex_enabled = merged.unify_codex_session_history;
state crate::settings::update_settings(merged).map_err(|e| e.to_string())?;
.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 配置, // 统一会话开关变更时立即重写当前官方 Codex 供应商的 live 配置,
// 不必等下一次切换才生效。 // 不必等下一次切换才生效。
@@ -106,18 +82,7 @@ pub async fn save_settings(
crate::services::provider::reapply_current_codex_official_live(state.inner()) crate::services::provider::reapply_current_codex_official_live(state.inner())
{ {
log::warn!("统一 Codex 会话历史开关变更后重写 live 配置失败,回滚设置: {err}"); log::warn!("统一 Codex 会话历史开关变更后重写 live 配置失败,回滚设置: {err}");
let pi_guard = state if let Err(rollback_err) = crate::settings::update_settings(existing) {
.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, &current, existing,
)
.await
{
log::error!("回滚统一会话开关设置失败: {rollback_err}"); log::error!("回滚统一会话开关设置失败: {rollback_err}");
} }
return Err(format!( return Err(format!(
@@ -653,28 +618,6 @@ mod tests {
assert!(merged.local_migrations.is_none()); 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);
}
} }
/// 获取开机自启状态 /// 获取开机自启状态
-9
View File
@@ -11,9 +11,7 @@ use crate::services::skill::{
SkillService, SkillStorageLocation, SkillUninstallResult, SkillUpdateInfo, SkillService, SkillStorageLocation, SkillUninstallResult, SkillUpdateInfo,
SkillsShSearchResult, SkillsShSearchResult,
}; };
use crate::services::skill_deployment::{PiSkillDeploymentService, SkillAppStatus};
use crate::store::AppState; use crate::store::AppState;
use std::collections::BTreeMap;
use std::str::FromStr; use std::str::FromStr;
use std::sync::Arc; use std::sync::Arc;
use tauri::State; use tauri::State;
@@ -34,13 +32,6 @@ pub fn get_installed_skills(app_state: State<'_, AppState>) -> Result<Vec<Instal
SkillService::get_all_installed(&app_state.db).map_err(|e| e.to_string()) SkillService::get_all_installed(&app_state.db).map_err(|e| e.to_string())
} }
#[tauri::command]
pub fn get_pi_skill_statuses(
app_state: State<'_, AppState>,
) -> Result<BTreeMap<String, SkillAppStatus>, String> {
PiSkillDeploymentService::inspect_all(&app_state.db).map_err(|error| error.to_string())
}
#[tauri::command] #[tauri::command]
pub fn get_skill_backups() -> Result<Vec<SkillBackupEntry>, String> { pub fn get_skill_backups() -> Result<Vec<SkillBackupEntry>, String> {
SkillService::list_backups().map_err(|e| e.to_string()) SkillService::list_backups().map_err(|e| e.to_string())
+10 -63
View File
@@ -1,44 +1,17 @@
use serde_json::{json, Value};
use std::sync::Arc;
use crate::database::Database;
use crate::error::AppError; use crate::error::AppError;
use crate::services::provider::ProviderService; use crate::services::provider::ProviderService;
use crate::services::PromptService;
use crate::settings; use crate::settings;
use crate::store::AppState; use crate::store::AppState;
use serde_json::{json, Value};
pub(crate) fn run_post_import_sync(app_state: &AppState) -> Result<(), AppError> { pub(crate) fn run_post_import_sync(db: Arc<Database>) -> Result<(), AppError> {
// Provider synchronization reopens/reconciles Pi's runtime admission after let app_state = AppState::new(db);
// the pre-import boundary closed it. Run that recovery first, then execute ProviderService::sync_current_to_live(&app_state)?;
// every remaining independent projection even if one of them fails. settings::reload_settings()?;
run_post_import_steps( Ok(())
|| ProviderService::sync_current_to_live(app_state),
|| PromptService::reconcile_pi_portable_import(app_state),
settings::reload_settings,
)
}
fn run_post_import_steps(
live_sync: impl FnOnce() -> Result<(), AppError>,
prompt_sync: impl FnOnce() -> Result<(), AppError>,
settings_reload: impl FnOnce() -> Result<(), AppError>,
) -> Result<(), AppError> {
let mut failures = Vec::new();
for (stage, result) in [
("live", live_sync()),
("pi_prompt", prompt_sync()),
("settings", settings_reload()),
] {
if let Err(error) = result {
failures.push(format!("{stage}={error}"));
}
}
if failures.is_empty() {
Ok(())
} else {
Err(AppError::Config(format!(
"post-import reconciliation incomplete: {}",
failures.join("; ")
)))
}
} }
fn post_sync_warning<E: std::fmt::Display>(err: E) -> String { fn post_sync_warning<E: std::fmt::Display>(err: E) -> String {
@@ -82,9 +55,8 @@ pub(crate) fn success_payload_with_warning(backup_id: String, warning: Option<St
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{attach_warning, post_sync_warning_from_result, run_post_import_steps}; use super::{attach_warning, post_sync_warning_from_result};
use serde_json::json; use serde_json::json;
use std::cell::RefCell;
#[test] #[test]
fn post_sync_warning_from_result_returns_none_on_success() { fn post_sync_warning_from_result_returns_none_on_success() {
@@ -122,29 +94,4 @@ mod tests {
Some("post sync warning") Some("post sync warning")
); );
} }
#[test]
fn post_import_steps_recover_live_first_and_do_not_short_circuit() {
let calls = RefCell::new(Vec::new());
let error = run_post_import_steps(
|| {
calls.borrow_mut().push("live");
Ok(())
},
|| {
calls.borrow_mut().push("prompt");
Err(crate::error::AppError::Config("invalid AGENTS.md".into()))
},
|| {
calls.borrow_mut().push("settings");
Err(crate::error::AppError::Config("reload failed".into()))
},
)
.expect_err("independent failures must be reported");
assert_eq!(*calls.borrow(), ["live", "prompt", "settings"]);
let message = error.to_string();
assert!(message.contains("pi_prompt="));
assert!(message.contains("settings="));
}
} }
+5 -31
View File
@@ -115,44 +115,18 @@ pub async fn webdav_sync_upload(state: State<'_, AppState>) -> Result<Value, Str
#[tauri::command] #[tauri::command]
pub async fn webdav_sync_download(state: State<'_, AppState>) -> Result<Value, String> { pub async fn webdav_sync_download(state: State<'_, AppState>) -> Result<Value, String> {
let db = state.db.clone(); let db = state.db.clone();
let app_state = state.inner().clone(); let db_for_sync = db.clone();
let mut settings = require_enabled_webdav_settings()?; let mut settings = require_enabled_webdav_settings()?;
let _auto_sync_suppression = crate::services::webdav_auto_sync::AutoSyncSuppressionGuard::new(); 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 sync_result = run_with_webdav_lock(webdav_sync_service::download(&db, &mut settings)).await;
let mut result = match sync_result { let mut result = map_sync_result(sync_result, |error| {
Ok(result) => result, persist_sync_error(&mut settings, error, "manual")
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. // Post-download sync is best-effort: snapshot restore has already succeeded.
let sync_state = app_state.clone();
let warning = post_sync_warning_from_result( let warning = post_sync_warning_from_result(
tauri::async_runtime::spawn_blocking(move || run_post_import_sync(&sync_state)) tauri::async_runtime::spawn_blocking(move || run_post_import_sync(db_for_sync))
.await .await
.map_err(|e| e.to_string()), .map_err(|e| e.to_string()),
); );
+35 -122
View File
@@ -295,23 +295,6 @@ pub fn write_text_file(path: &Path, data: &str) -> Result<(), AppError> {
/// 原子写入:写入临时文件后 rename 替换,避免半写状态 /// 原子写入:写入临时文件后 rename 替换,避免半写状态
pub fn atomic_write(path: &Path, data: &[u8]) -> Result<(), AppError> { 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 when `required_file_mode` is `None`.
/// Sensitive callers pass an explicit mode (for example `0o600`), which is
/// enforced for both new and existing Unix files before the replacement
/// becomes visible. 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],
required_file_mode: Option<u32>,
) -> Result<(), AppError> {
#[cfg(not(unix))]
let _ = required_file_mode;
if let Some(parent) = path.parent() { if let Some(parent) = path.parent() {
fs::create_dir_all(parent).map_err(|e| AppError::io(parent, e))?; fs::create_dir_all(parent).map_err(|e| AppError::io(parent, e))?;
} }
@@ -319,97 +302,51 @@ pub(crate) fn atomic_write_durable(
let parent = path let parent = path
.parent() .parent()
.ok_or_else(|| AppError::Config("无效的路径".to_string()))?; .ok_or_else(|| AppError::Config("无效的路径".to_string()))?;
let mut tmp = parent.to_path_buf();
let file_name = path let file_name = path
.file_name() .file_name()
.ok_or_else(|| AppError::Config("无效的文件名".to_string()))? .ok_or_else(|| AppError::Config("无效的文件名".to_string()))?
.to_string_lossy() .to_string_lossy()
.to_string(); .to_string();
let tmp = parent.join(format!( let ts = std::time::SystemTime::now()
".{file_name}.{}.tmp", .duration_since(std::time::UNIX_EPOCH)
uuid::Uuid::new_v4().simple() .unwrap_or_default()
)); .as_nanos();
tmp.push(format!("{file_name}.tmp.{ts}"));
let result = (|| -> Result<(), AppError> { {
let mut options = fs::OpenOptions::new(); let mut f = fs::File::create(&tmp).map_err(|e| AppError::io(&tmp, e))?;
options.create_new(true).write(true); f.write_all(data).map_err(|e| AppError::io(&tmp, e))?;
#[cfg(unix)] f.flush().map_err(|e| AppError::io(&tmp, e))?;
{
use std::os::unix::fs::OpenOptionsExt;
options.mode(required_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))?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let mode = required_file_mode.unwrap_or_else(|| {
fs::metadata(path)
.map(|metadata| metadata.permissions().mode())
.unwrap_or(0o666)
});
file.set_permissions(fs::Permissions::from_mode(mode))
.map_err(|error| AppError::io(&tmp, error))?;
}
file.sync_all().map_err(|error| AppError::io(&tmp, error))?;
drop(file);
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))] #[cfg(unix)]
fn replace_file_atomically(temp_path: &Path, path: &Path) -> Result<(), AppError> { {
fs::rename(temp_path, path).map_err(|source| AppError::IoContext { use std::os::unix::fs::PermissionsExt;
context: format!( if let Ok(meta) = fs::metadata(path) {
"原子替换失败: {} -> {}", let perm = meta.permissions().mode();
temp_path.display(), let _ = fs::set_permissions(&tmp, fs::Permissions::from_mode(perm));
path.display() }
), }
source,
})
}
#[cfg(windows)] #[cfg(windows)]
fn replace_file_atomically(temp_path: &Path, path: &Path) -> Result<(), AppError> { {
use std::os::windows::ffi::OsStrExt; // Windows 上 rename 目标存在会失败,先移除再重命名(尽量接近原子性)
use windows_sys::Win32::Storage::FileSystem::{ if path.exists() {
MoveFileExW, MOVEFILE_REPLACE_EXISTING, MOVEFILE_WRITE_THROUGH, let _ = fs::remove_file(path);
}; }
fs::rename(&tmp, path).map_err(|e| AppError::IoContext {
context: format!("原子替换失败: {} -> {}", tmp.display(), path.display()),
source: e,
})?;
}
let source: Vec<u16> = temp_path.as_os_str().encode_wide().chain(Some(0)).collect(); #[cfg(not(windows))]
let destination: Vec<u16> = path.as_os_str().encode_wide().chain(Some(0)).collect(); {
// SAFETY: both buffers are NUL-terminated and remain alive for the fs::rename(&tmp, path).map_err(|e| AppError::IoContext {
// duration of this synchronous Win32 call. context: format!("原子替换失败: {} -> {}", tmp.display(), path.display()),
let moved = unsafe { source: e,
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(()) Ok(())
} }
@@ -563,30 +500,6 @@ mod tests {
); );
} }
#[cfg(unix)]
#[test]
fn sensitive_atomic_write_tightens_an_existing_file_before_publish() {
use std::os::unix::fs::PermissionsExt;
let temp = tempfile::tempdir().expect("tempdir");
let path = temp.path().join("settings.json");
fs::write(&path, b"old").expect("seed settings");
fs::set_permissions(&path, fs::Permissions::from_mode(0o644))
.expect("make legacy settings permissive");
atomic_write_durable(&path, b"new-secret", Some(0o600)).expect("replace settings");
assert_eq!(fs::read(&path).expect("read settings"), b"new-secret");
assert_eq!(
fs::metadata(&path)
.expect("settings metadata")
.permissions()
.mode()
& 0o777,
0o600
);
}
#[test] #[test]
fn sort_json_keys_produces_identical_output_for_different_insertion_orders() { fn sort_json_keys_produces_identical_output_for_different_insertion_orders() {
// 核心保证:同一逻辑配置无论键的插入顺序如何,写出的字节序列必须一致。 // 核心保证:同一逻辑配置无论键的插入顺序如何,写出的字节序列必须一致。
-7
View File
@@ -4,19 +4,12 @@
pub mod failover; pub mod failover;
pub mod mcp; pub mod mcp;
pub(crate) mod pi_catalog;
pub(crate) mod pi_portable_state;
pub mod pi_projections;
pub mod profiles; pub mod profiles;
pub mod prompts; pub mod prompts;
pub mod provider_write;
#[cfg(test)]
mod provider_write_certification;
pub mod providers; pub mod providers;
pub mod providers_seed; pub mod providers_seed;
pub mod proxy; pub mod proxy;
pub mod settings; pub mod settings;
pub mod skill_deployments;
pub mod skills; pub mod skills;
pub mod stream_check; pub mod stream_check;
pub mod universal_providers; pub mod universal_providers;
-276
View File
@@ -1,276 +0,0 @@
//! 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<String, ProviderAggregate>,
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::<Result<Vec<_>, _>>()
.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::<Result<Vec<_>, _>>()?;
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<PiProviderProjection, AppError> {
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<Option<PiProviderProjection>, 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::<Result<Vec<_>, _>>()?;
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;
@@ -1,702 +0,0 @@
//! Portable SQL boundary for Pi's device-local ownership evidence.
//!
//! Provider and Skill configuration is portable. Exact `models.json` claims and
//! native Skill deployment receipts are not: they describe files on one device
//! and must never become ownership proof on another. This adapter keeps that
//! policy outside the frozen generic backup/restore implementation while giving
//! manual SQL, WebDAV, and S3 one shared boundary.
#[cfg(test)]
use crate::database::SkillDeploymentMethod;
use crate::database::{lock_conn, Database, PiProviderProjection, SkillDeployment};
use crate::error::AppError;
use rusqlite::{Connection, OpenFlags, OptionalExtension};
use std::fmt::Write as _;
use std::fs;
use std::path::Path;
const DEVICE_LOCAL_INSERT_PREFIXES: &[&str] = &[
"INSERT INTO \"pi_provider_projections\"",
"INSERT INTO \"skill_deployments\"",
];
#[derive(Debug, Clone, PartialEq, Eq)]
struct PiDeviceLocalState {
projections: Vec<PiProviderProjection>,
skill_deployments: Vec<SkillDeployment>,
}
impl Database {
/// Export a user-portable SQL backup without device-local ownership rows.
pub(crate) fn export_portable_sql_string(&self) -> Result<String, AppError> {
strip_device_local_insert_statements(&self.export_sql_string()?)
}
/// Export a cloud-sync SQL snapshot without device-local ownership rows.
pub(crate) fn export_portable_sql_string_for_sync(&self) -> Result<String, AppError> {
strip_device_local_insert_statements(&self.export_sql_string_for_sync()?)
}
pub(crate) fn export_portable_sql(&self, target_path: &Path) -> Result<(), AppError> {
let dump = self.export_portable_sql_string()?;
if let Some(parent) = target_path.parent() {
fs::create_dir_all(parent).map_err(|error| AppError::io(parent, error))?;
}
crate::config::atomic_write(target_path, dump.as_bytes())
}
/// Import a user-portable SQL backup while retaining this device's evidence.
pub(crate) fn import_portable_sql(&self, source_path: &Path) -> Result<String, AppError> {
let local = self.capture_pi_device_local_state()?;
let sql =
fs::read_to_string(source_path).map_err(|error| AppError::io(source_path, error))?;
self.import_sql_string(&append_pi_device_local_state(&sql, &local)?)
}
/// Import a cloud-sync snapshot while retaining this device's evidence.
pub(crate) fn import_portable_sql_string_for_sync(
&self,
sql: &str,
) -> Result<String, AppError> {
let local = self.capture_pi_device_local_state()?;
self.import_sql_string_for_sync(&append_pi_device_local_state(sql, &local)?)
}
/// Fail closed before the legacy whole-database restore can import
/// device-local Pi ownership evidence.
///
/// Canonical binary restore hardening is a separate project. Until that
/// boundary can retain the receiving device's ledgers atomically, binary
/// restore is supported only when neither side carries ownership rows.
pub(crate) fn ensure_binary_restore_has_no_pi_ownership(
&self,
filename: &str,
) -> Result<(), AppError> {
if filename.contains("..")
|| filename.contains('/')
|| filename.contains('\\')
|| !filename.ends_with(".db")
{
return Err(AppError::InvalidInput(
"Invalid backup filename".to_string(),
));
}
{
let conn = lock_conn!(self.conn);
if pi_device_local_rows_exist(&conn)? {
return Err(binary_restore_ownership_error(
"当前数据库",
"current database",
));
}
}
let backup_path = crate::config::get_app_config_dir()
.join("backups")
.join(filename);
if !backup_path.exists() {
return Err(AppError::InvalidInput(format!(
"Backup file not found: {filename}"
)));
}
let source = Connection::open_with_flags(
&backup_path,
OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
)
.map_err(|error| AppError::Database(format!("无法只读检查备份: {error}")))?;
if pi_device_local_rows_exist(&source)? {
return Err(binary_restore_ownership_error(
"所选备份",
"selected backup",
));
}
Ok(())
}
fn capture_pi_device_local_state(&self) -> Result<PiDeviceLocalState, AppError> {
let conn = lock_conn!(self.conn);
let projections = {
let mut statement = 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 = statement
.query_map([], |row| {
Ok(PiProviderProjection {
provider_id: row.get(0)?,
provider_key: row.get(1)?,
created_at: row.get(2)?,
updated_at: row.get(3)?,
})
})
.map_err(|error| AppError::Database(error.to_string()))?;
rows.collect::<Result<Vec<_>, _>>()
.map_err(|error| AppError::Database(error.to_string()))?
};
let skill_deployments = {
let mut statement = conn
.prepare(
"SELECT skill_id, destination, destination_key, method,
source_identity, deployed_digest, created_at, updated_at
FROM skill_deployments
WHERE app_type = 'pi'
ORDER BY skill_id, destination_key",
)
.map_err(|error| AppError::Database(error.to_string()))?;
let rows = statement
.query_map([], |row| {
let method = row.get::<_, String>(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)?,
})
})
.map_err(|error| AppError::Database(error.to_string()))?;
rows.collect::<Result<Vec<_>, _>>()
.map_err(|error| AppError::Database(error.to_string()))?
};
Ok(PiDeviceLocalState {
projections,
skill_deployments,
})
}
}
fn pi_device_local_rows_exist(conn: &Connection) -> Result<bool, AppError> {
let schema_object_exists = |name: &str| -> Result<bool, AppError> {
conn.query_row(
"SELECT 1 FROM sqlite_master
WHERE name = ?1 COLLATE NOCASE",
[name],
|_| Ok(()),
)
.optional()
.map(|row| row.is_some())
.map_err(|error| AppError::Database(error.to_string()))
};
let projections_exist = if schema_object_exists("pi_provider_projections")? {
conn.query_row(
"SELECT EXISTS(SELECT 1 FROM pi_provider_projections LIMIT 1)",
[],
|row| row.get::<_, bool>(0),
)
.map_err(|error| AppError::Database(error.to_string()))?
} else {
false
};
if projections_exist {
return Ok(true);
}
if schema_object_exists("skill_deployments")? {
return conn
.query_row(
"SELECT EXISTS(
SELECT 1 FROM skill_deployments
WHERE app_type = 'pi'
LIMIT 1
)",
[],
|row| row.get::<_, bool>(0),
)
.map_err(|error| AppError::Database(error.to_string()));
}
Ok(false)
}
/// Merge receiving-device evidence into the same temporary database that the
/// generic importer validates and publishes. No live database replacement can
/// occur before these statements have succeeded.
fn append_pi_device_local_state(sql: &str, local: &PiDeviceLocalState) -> Result<String, AppError> {
ensure_sql_append_boundary(sql)?;
let mut merged = String::with_capacity(sql.len() + 1024);
merged.push_str(sql);
// The leading newline closes a trailing `--` comment; the standalone
// semicolon terminates a final statement that omitted its delimiter.
merged.push_str(
"\n;\n\
DROP TABLE IF EXISTS pi_provider_projections;\n\
CREATE TABLE pi_provider_projections (\n\
provider_id TEXT PRIMARY KEY,\n\
provider_key TEXT NOT NULL UNIQUE,\n\
created_at INTEGER NOT NULL,\n\
updated_at INTEGER NOT NULL\n\
);\n\
DROP TABLE IF EXISTS skill_deployments;\n\
CREATE TABLE skill_deployments (\n\
app_type TEXT NOT NULL CHECK (app_type = 'pi'),\n\
skill_id TEXT NOT NULL,\n\
destination TEXT NOT NULL,\n\
destination_key TEXT NOT NULL,\n\
method TEXT NOT NULL CHECK (method IN ('symlink', 'copy')),\n\
source_identity TEXT NOT NULL,\n\
deployed_digest TEXT,\n\
created_at INTEGER NOT NULL,\n\
updated_at INTEGER NOT NULL,\n\
PRIMARY KEY (app_type, skill_id, destination_key),\n\
UNIQUE (app_type, destination_key)\n\
);\n",
);
for projection in &local.projections {
writeln!(
merged,
"INSERT INTO pi_provider_projections \
(provider_id, provider_key, created_at, updated_at) \
VALUES ({}, {}, {}, {});",
sql_text(&projection.provider_id)?,
sql_text(&projection.provider_key)?,
projection.created_at,
projection.updated_at,
)
.expect("writing to String cannot fail");
}
for deployment in &local.skill_deployments {
writeln!(
merged,
"INSERT INTO skill_deployments \
(app_type, skill_id, destination, destination_key, method, \
source_identity, deployed_digest, created_at, updated_at) \
VALUES ('pi', {}, {}, {}, {}, {}, {}, {}, {});",
sql_text(&deployment.skill_id)?,
sql_text(&deployment.destination)?,
sql_text(&deployment.destination_key)?,
sql_text(deployment.method.as_str())?,
sql_text(&deployment.source_identity)?,
sql_optional_text(deployment.deployed_digest.as_deref())?,
deployment.created_at,
deployment.updated_at,
)
.expect("writing to String cannot fail");
}
Ok(merged)
}
fn sql_optional_text(value: Option<&str>) -> Result<String, AppError> {
value.map_or_else(|| Ok("NULL".to_string()), sql_text)
}
fn sql_text(value: &str) -> Result<String, AppError> {
if value.contains('\0') {
return Err(AppError::InvalidInput(
"device-local Pi ownership text cannot contain NUL".to_string(),
));
}
Ok(format!("'{}'", value.replace('\'', "''")))
}
/// SQLite accepts an unterminated block comment at EOF. Reject such input so
/// it cannot swallow the receiving-device statements appended above. Other
/// unterminated quoted forms are rejected here as a clearer pre-publish error.
fn ensure_sql_append_boundary(sql: &str) -> Result<(), AppError> {
#[derive(Clone, Copy, PartialEq, Eq)]
enum State {
Normal,
SingleQuote,
DoubleQuote,
Backtick,
Bracket,
LineComment,
BlockComment,
}
let bytes = sql.as_bytes();
let mut state = State::Normal;
let mut cursor = 0;
while cursor < bytes.len() {
let current = bytes[cursor];
let next = bytes.get(cursor + 1).copied();
match state {
State::Normal => match (current, next) {
(b'\'', _) => state = State::SingleQuote,
(b'"', _) => state = State::DoubleQuote,
(b'`', _) => state = State::Backtick,
(b'[', _) => state = State::Bracket,
(b'-', Some(b'-')) => {
state = State::LineComment;
cursor += 1;
}
(b'/', Some(b'*')) => {
state = State::BlockComment;
cursor += 1;
}
_ => {}
},
State::SingleQuote if current == b'\'' => {
if next == Some(b'\'') {
cursor += 1;
} else {
state = State::Normal;
}
}
State::DoubleQuote if current == b'"' => {
if next == Some(b'"') {
cursor += 1;
} else {
state = State::Normal;
}
}
State::Backtick if current == b'`' => {
if next == Some(b'`') {
cursor += 1;
} else {
state = State::Normal;
}
}
State::Bracket if current == b']' => state = State::Normal,
State::LineComment if current == b'\n' => state = State::Normal,
State::BlockComment if current == b'*' && next == Some(b'/') => {
state = State::Normal;
cursor += 1;
}
_ => {}
}
cursor += 1;
}
if matches!(state, State::Normal | State::LineComment) {
Ok(())
} else {
Err(AppError::InvalidInput(
"portable SQL ends inside an unterminated quoted value or comment".to_string(),
))
}
}
fn binary_restore_ownership_error(source_zh: &str, source_en: &str) -> AppError {
AppError::localized(
"pi.binary_restore_device_ownership_unsupported",
format!(
"为防止历史设备所有权记录覆盖当前 Pi 原生文件,{source_zh}含 Pi 所有权状态时不能使用数据库备份恢复;请改用可移植 SQL 导入"
),
format!(
"Binary database restore is unavailable because the {source_en} contains device-local Pi ownership state; use portable SQL import instead"
),
)
}
/// Remove complete INSERT statements for the two device-local tables from SQL
/// generated by `Database::dump_sql`. Values may contain quotes, semicolons, or
/// newlines, so line filtering is insufficient; statement boundaries are found
/// only outside SQLite single-quoted literals. Unknown output fails closed.
fn strip_device_local_insert_statements(sql: &str) -> Result<String, AppError> {
let bytes = sql.as_bytes();
let mut output = String::with_capacity(sql.len());
let mut statement_start = 0;
let mut cursor = 0;
let mut in_string = false;
while cursor < bytes.len() {
match bytes[cursor] {
b'\'' if in_string && bytes.get(cursor + 1) == Some(&b'\'') => {
cursor += 2;
continue;
}
b'\'' => in_string = !in_string,
b';' if !in_string => {
let statement_end = cursor + 1;
let statement = &sql[statement_start..statement_end];
if !DEVICE_LOCAL_INSERT_PREFIXES
.iter()
.any(|prefix| statement.trim_start().starts_with(prefix))
{
output.push_str(statement);
}
statement_start = statement_end;
}
_ => {}
}
cursor += 1;
}
if in_string {
return Err(AppError::Config(
"portable SQL export ended inside a quoted value".to_string(),
));
}
output.push_str(&sql[statement_start..]);
if DEVICE_LOCAL_INSERT_PREFIXES
.iter()
.any(|prefix| output.contains(prefix))
{
return Err(AppError::Config(
"portable SQL export retained device-local Pi ownership rows".to_string(),
));
}
Ok(output)
}
#[cfg(test)]
mod tests {
use super::*;
fn deployment(skill_id: &str, destination_key: &str) -> SkillDeployment {
SkillDeployment {
skill_id: skill_id.to_string(),
destination: format!("/device/{destination_key}"),
destination_key: destination_key.to_string(),
method: SkillDeploymentMethod::Copy,
source_identity: format!("path:/source/{skill_id};digest:sha256:one"),
deployed_digest: Some("sha256:one".to_string()),
created_at: 10,
updated_at: 20,
}
}
fn seed_portable_provider(db: &Database) -> Result<(), AppError> {
let conn = lock_conn!(db.conn);
conn.execute(
"INSERT INTO providers
(id, app_type, name, settings_config, meta, is_current)
VALUES ('portable-sentinel', 'codex', 'Portable sentinel', '{}', '{}', 0)",
[],
)
.map_err(|error| AppError::Database(error.to_string()))?;
Ok(())
}
fn add_restore_poison_trigger(sql: &str) -> String {
let insertion = sql
.rfind("COMMIT;")
.expect("CC Switch dump must contain a final COMMIT");
let mut poisoned = sql.to_string();
poisoned.insert_str(
insertion,
"CREATE TRIGGER poison_pi_restore \
BEFORE INSERT ON pi_provider_projections \
BEGIN SELECT RAISE(ABORT, 'poisoned local restore'); END;\n",
);
poisoned
}
#[test]
fn portable_sql_scrubber_handles_multiline_quoted_values() -> Result<(), AppError> {
let sql = concat!(
"-- CC Switch SQLite 导出\n",
"CREATE TABLE \"pi_provider_projections\" (value TEXT);\n",
"INSERT INTO \"pi_provider_projections\" (value) VALUES ('one;\n",
"two ''quoted''');\n",
"CREATE TABLE \"providers\" (value TEXT);\n",
"INSERT INTO \"providers\" (value) VALUES ('portable;\nvalue');\n",
"COMMIT;\n",
);
let scrubbed = strip_device_local_insert_statements(sql)?;
assert!(!scrubbed.contains("INSERT INTO \"pi_provider_projections\""));
assert!(scrubbed.contains("CREATE TABLE \"pi_provider_projections\""));
assert!(scrubbed.contains("INSERT INTO \"providers\""));
assert!(scrubbed.contains("'portable;\nvalue'"));
Ok(())
}
#[test]
fn portable_exports_keep_schema_but_omit_device_evidence() -> Result<(), AppError> {
let db = Database::memory()?;
db.claim_pi_projection_key("local-provider", "local-key")?;
db.save_pi_skill_deployment(&deployment("local-skill", "local-destination"))?;
for exported in [
db.export_portable_sql_string()?,
db.export_portable_sql_string_for_sync()?,
] {
assert!(exported.contains("CREATE TABLE pi_provider_projections"));
assert!(exported.contains("CREATE TABLE skill_deployments"));
assert!(!exported.contains("INSERT INTO \"pi_provider_projections\""));
assert!(!exported.contains("INSERT INTO \"skill_deployments\""));
}
Ok(())
}
#[test]
fn sync_import_discards_remote_evidence_and_restores_local_evidence() -> Result<(), AppError> {
let remote = Database::memory()?;
seed_portable_provider(&remote)?;
remote.claim_pi_projection_key("remote-provider", "shared-key")?;
remote.save_pi_skill_deployment(&deployment("remote-skill", "remote-destination"))?;
// Model an older remote snapshot created before portable row scrubbing.
let remote_sql = add_restore_poison_trigger(&remote.export_sql_string()?);
let local = Database::memory()?;
local.claim_pi_projection_key("local-provider", "shared-key")?;
local.save_pi_skill_deployment(&deployment("local-skill", "local-destination"))?;
local.import_portable_sql_string_for_sync(&remote_sql)?;
assert_eq!(
local
.get_pi_projection_for_key("shared-key")?
.map(|projection| projection.provider_id),
Some("local-provider".to_string())
);
assert!(local.get_pi_projection("remote-provider")?.is_none());
assert_eq!(
local.get_pi_skill_deployments("local-skill")?,
vec![deployment("local-skill", "local-destination")]
);
assert!(local.get_pi_skill_deployments("remote-skill")?.is_empty());
Ok(())
}
#[test]
fn manual_sql_import_preserves_local_evidence() -> Result<(), AppError> {
let remote = Database::memory()?;
seed_portable_provider(&remote)?;
remote.claim_pi_projection_key("remote-provider", "remote-key")?;
let temp = tempfile::tempdir().expect("tempdir");
let path = temp.path().join("portable.sql");
fs::write(
&path,
add_restore_poison_trigger(&remote.export_sql_string()?),
)
.expect("write SQL backup");
let local = Database::memory()?;
local.claim_pi_projection_key("local-provider", "local-key")?;
local.import_portable_sql(&path)?;
assert!(local.get_pi_projection("remote-provider")?.is_none());
assert_eq!(
local
.get_pi_projection("local-provider")?
.map(|projection| projection.provider_key),
Some("local-key".to_string())
);
Ok(())
}
#[test]
fn binary_restore_guard_rejects_live_ownership_before_opening_the_backup(
) -> Result<(), AppError> {
let db = Database::memory()?;
db.claim_pi_projection_key("local-provider", "local-key")?;
let error = db
.ensure_binary_restore_has_no_pi_ownership("missing.db")
.expect_err("live ownership must reject before inspecting a source");
assert!(error.to_string().contains("portable SQL"));
Ok(())
}
#[test]
fn binary_restore_guard_rejects_backup_ownership_but_allows_empty_tables(
) -> Result<(), AppError> {
let temp = tempfile::tempdir().expect("tempdir");
let empty_path = temp.path().join("empty.db");
let owned_path = temp.path().join("owned.db");
for path in [&empty_path, &owned_path] {
let conn =
Connection::open(path).map_err(|error| AppError::Database(error.to_string()))?;
conn.execute_batch(
"CREATE TABLE PI_PROVIDER_PROJECTIONS (
provider_id TEXT PRIMARY KEY,
provider_key TEXT NOT NULL
);
CREATE TABLE SKILL_DEPLOYMENTS (
app_type TEXT NOT NULL
);",
)
.map_err(|error| AppError::Database(error.to_string()))?;
}
let owned =
Connection::open(&owned_path).map_err(|error| AppError::Database(error.to_string()))?;
owned
.execute(
"INSERT INTO pi_provider_projections (provider_id, provider_key)
VALUES ('historical-provider', 'native-key')",
[],
)
.map_err(|error| AppError::Database(error.to_string()))?;
let empty = Connection::open_with_flags(
&empty_path,
OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
)
.map_err(|error| AppError::Database(error.to_string()))?;
assert!(!pi_device_local_rows_exist(&empty)?);
let owned = Connection::open_with_flags(
&owned_path,
OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
)
.map_err(|error| AppError::Database(error.to_string()))?;
assert!(pi_device_local_rows_exist(&owned)?);
drop(owned);
let owned =
Connection::open(&owned_path).map_err(|error| AppError::Database(error.to_string()))?;
owned
.execute("DELETE FROM pi_provider_projections", [])
.map_err(|error| AppError::Database(error.to_string()))?;
owned
.execute("INSERT INTO skill_deployments (app_type) VALUES ('pi')", [])
.map_err(|error| AppError::Database(error.to_string()))?;
drop(owned);
let skill_only = Connection::open_with_flags(
&owned_path,
OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
)
.map_err(|error| AppError::Database(error.to_string()))?;
assert!(pi_device_local_rows_exist(&skill_only)?);
let view_backed =
Connection::open_in_memory().map_err(|error| AppError::Database(error.to_string()))?;
view_backed
.execute_batch(
"CREATE VIEW PI_PROVIDER_PROJECTIONS AS
SELECT 'view-provider' AS provider_id, 'view-key' AS provider_key;",
)
.map_err(|error| AppError::Database(error.to_string()))?;
assert!(
pi_device_local_rows_exist(&view_backed)?,
"a reserved-name view must not bypass binary ownership rejection"
);
Ok(())
}
#[test]
fn portable_import_rejects_an_unclosed_comment_before_live_replacement() -> Result<(), AppError>
{
let remote = Database::memory()?;
seed_portable_provider(&remote)?;
let malicious = format!("{}/*", remote.export_sql_string()?);
let local = Database::memory()?;
local.claim_pi_projection_key("local-provider", "local-key")?;
let error = local
.import_portable_sql_string_for_sync(&malicious)
.expect_err("the appended ownership program must not be swallowed");
assert!(error.to_string().contains("unterminated"));
assert_eq!(
local
.get_pi_projection("local-provider")?
.map(|projection| projection.provider_key),
Some("local-key".to_string())
);
assert!(
local
.get_provider_by_id("portable-sentinel", "codex")?
.is_none(),
"the remote database must not be published"
);
Ok(())
}
}
@@ -1,204 +0,0 @@
//! 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<PiProviderProjection> {
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<Option<PiProviderProjection>, 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<Option<PiProviderProjection>, 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<IndexMap<String, PiProviderProjection>, 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<PiProviderProjection, AppError> {
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<bool, AppError> {
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(())
}
}
+40 -164
View File
@@ -6,106 +6,51 @@ use crate::database::{lock_conn, Database};
use crate::error::AppError; use crate::error::AppError;
use crate::prompt::Prompt; use crate::prompt::Prompt;
use indexmap::IndexMap; use indexmap::IndexMap;
use rusqlite::{params, Connection, Transaction}; use rusqlite::params;
fn query_prompts(conn: &Connection, app_type: &str) -> Result<IndexMap<String, Prompt>, 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<String> = row.get(3)?;
let enabled: bool = row.get(4)?;
let created_at: Option<i64> = row.get(5)?;
let updated_at: Option<i64> = 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<String, Prompt>) -> 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<String, Prompt>,
) -> 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<String, Prompt>,
right: &IndexMap<String, Prompt>,
) -> bool {
left.len() == right.len()
&& left
.iter()
.all(|(id, prompt)| right.get(id) == Some(prompt))
}
impl Database { impl Database {
/// 获取指定应用类型的所有提示词 /// 获取指定应用类型的所有提示词
pub fn get_prompts(&self, app_type: &str) -> Result<IndexMap<String, Prompt>, AppError> { pub fn get_prompts(&self, app_type: &str) -> Result<IndexMap<String, Prompt>, AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
query_prompts(&conn, app_type) 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<String> = row.get(3)?;
let enabled: bool = row.get(4)?;
let created_at: Option<i64> = row.get(5)?;
let updated_at: Option<i64> = 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)
} }
/// 保存提示词 /// 保存提示词
@@ -130,75 +75,6 @@ impl Database {
Ok(()) 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<String, Prompt>,
) -> 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<String, Prompt>,
replacement: &IndexMap<String, Prompt>,
) -> 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<String, Prompt>,
before: &IndexMap<String, Prompt>,
) -> Result<(), AppError> {
self.compare_exchange_prompt_selection_unchecked(app_type, attempted, before)
}
fn compare_exchange_prompt_selection_unchecked(
&self,
app_type: &str,
expected: &IndexMap<String, Prompt>,
replacement: &IndexMap<String, Prompt>,
) -> 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> { pub fn delete_prompt(&self, app_type: &str, id: &str) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
@@ -1,655 +0,0 @@
use crate::database::{lock_conn, Database};
use crate::error::AppError;
use crate::provider::{ProviderMeta, ProviderMutationInput};
use crate::settings::CustomEndpoint;
use rusqlite::{params, OptionalExtension, Transaction};
use serde_json::Value;
use std::collections::HashSet;
use super::providers::{StoredProviderRow, PROVIDER_SELECT};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ProviderKey {
app_type: String,
id: String,
}
impl ProviderKey {
pub fn new(app_type: impl Into<String>, id: impl Into<String>) -> Result<Self, AppError> {
let app_type = app_type.into();
let id = id.into();
if app_type.trim().is_empty() || id.trim().is_empty() {
return Err(AppError::InvalidInput(
"provider app type and id must be non-empty".to_string(),
));
}
Ok(Self { app_type, id })
}
pub fn app_type(&self) -> &str {
&self.app_type
}
pub fn id(&self) -> &str {
&self.id
}
}
#[derive(Debug, Clone)]
pub struct ProviderRowUpdate {
pub(super) name: String,
pub(super) settings_config: Value,
pub(super) website_url: Option<String>,
pub(super) category: Option<String>,
pub(super) notes: Option<String>,
pub(super) meta: ProviderMeta,
pub(super) icon: Option<String>,
pub(super) icon_color: Option<String>,
}
impl ProviderRowUpdate {
pub fn from_input(input: &ProviderMutationInput) -> Result<Self, AppError> {
let meta = input.meta.clone().unwrap_or_default();
if !meta.custom_endpoints.is_empty() {
return Err(AppError::InvalidInput(
"provider update must not contain customEndpoints; use endpoint operations"
.to_string(),
));
}
Ok(Self {
name: input.name.clone(),
settings_config: input.settings_config.clone(),
website_url: input.website_url.clone(),
category: input.category.clone(),
notes: input.notes.clone(),
meta,
icon: input.icon.clone(),
icon_color: input.icon_color.clone(),
})
}
}
#[derive(Debug, Clone)]
pub struct ProviderRowCreate {
pub(super) content: ProviderRowUpdate,
pub(super) created_at: Option<i64>,
}
#[derive(Debug, Clone)]
pub struct NewEndpoint {
pub(super) url: String,
pub(super) added_at: Option<i64>,
pub(super) last_used: Option<i64>,
}
impl NewEndpoint {
pub fn new(
url: impl Into<String>,
added_at: Option<i64>,
last_used: Option<i64>,
) -> Result<Self, AppError> {
let url = url.into();
if url.trim().is_empty() {
return Err(AppError::InvalidInput(
"provider endpoint URL cannot be empty".to_string(),
));
}
Ok(Self {
url,
added_at,
last_used,
})
}
pub fn now(url: impl Into<String>) -> Result<Self, AppError> {
Self::new(url, Some(chrono::Utc::now().timestamp_millis()), None)
}
}
impl TryFrom<CustomEndpoint> for NewEndpoint {
type Error = AppError;
fn try_from(endpoint: CustomEndpoint) -> Result<Self, Self::Error> {
Self::new(endpoint.url, endpoint.added_at, endpoint.last_used)
}
}
#[derive(Debug, Clone)]
pub struct NewProviderAggregate {
pub(super) key: ProviderKey,
pub(super) row: ProviderRowCreate,
pub(super) sort_index: Option<usize>,
pub(super) in_failover_queue: bool,
pub(super) initial_endpoints: Vec<NewEndpoint>,
}
impl NewProviderAggregate {
pub fn from_input(app_type: &str, mut input: ProviderMutationInput) -> Result<Self, AppError> {
let endpoints = input
.meta
.as_mut()
.map(|meta| std::mem::take(&mut meta.custom_endpoints))
.unwrap_or_default();
let mut seen = HashSet::with_capacity(endpoints.len());
let mut initial_endpoints = Vec::with_capacity(endpoints.len());
for (key, endpoint) in endpoints {
let normalized_key = key.trim().trim_end_matches('/').to_string();
let normalized_url = endpoint.url.trim().trim_end_matches('/').to_string();
if normalized_key != normalized_url {
return Err(AppError::InvalidInput(format!(
"provider endpoint key '{key}' must match endpoint URL '{}'",
endpoint.url
)));
}
if !seen.insert(normalized_url.clone()) {
return Err(AppError::InvalidInput(format!(
"duplicate initial provider endpoint '{}'",
endpoint.url
)));
}
initial_endpoints.push(NewEndpoint::new(
normalized_url,
endpoint.added_at,
endpoint.last_used,
)?);
}
let key = ProviderKey::new(app_type, input.id.clone())?;
let row = ProviderRowCreate {
content: ProviderRowUpdate::from_input(&input)?,
created_at: input.created_at,
};
Ok(Self {
key,
row,
sort_index: input.sort_index,
in_failover_queue: input.in_failover_queue,
initial_endpoints,
})
}
}
#[derive(Debug, Clone)]
pub struct RenameProvider {
source: ProviderKey,
target_id: String,
row: ProviderRowUpdate,
}
impl RenameProvider {
pub fn from_input(
source: ProviderKey,
input: &ProviderMutationInput,
) -> Result<Self, AppError> {
if !matches!(source.app_type(), "opencode" | "openclaw") {
return Err(AppError::InvalidInput(
"provider key changes are restricted to additive OpenCode/OpenClaw providers"
.to_string(),
));
}
if source.id() == input.id {
return Err(AppError::InvalidInput(
"provider rename requires a different target id".to_string(),
));
}
if input.id.trim().is_empty() {
return Err(AppError::InvalidInput(
"provider target id must be non-empty".to_string(),
));
}
let mut row = ProviderRowUpdate::from_input(input)?;
// A successful key change always remains DB-only. The service owns
// the corresponding live-file absence check, while the DAO persists
// the durable half of that invariant.
row.meta.live_config_managed = Some(false);
Ok(Self {
source,
target_id: input.id.clone(),
row,
})
}
}
fn encode_row(row: &ProviderRowUpdate) -> Result<(String, String), AppError> {
let settings_config = serde_json::to_string(&row.settings_config).map_err(|error| {
AppError::Database(format!("failed to serialize settings_config: {error}"))
})?;
let meta = serde_json::to_string(&row.meta).map_err(|error| {
AppError::Database(format!("failed to serialize provider meta: {error}"))
})?;
Ok((settings_config, meta))
}
pub(super) fn insert_row(
tx: &Transaction<'_>,
key: &ProviderKey,
row: &ProviderRowUpdate,
created_at: Option<i64>,
sort_index: Option<usize>,
is_current: bool,
in_failover_queue: bool,
) -> Result<(), AppError> {
let (settings_config, meta) = encode_row(row)?;
tx.execute(
"INSERT INTO providers (
id, app_type, name, settings_config, website_url, category,
created_at, sort_index, notes, icon, icon_color, meta,
is_current, in_failover_queue
) VALUES (
?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14
)",
params![
key.id,
key.app_type,
row.name,
settings_config,
row.website_url,
row.category,
created_at,
sort_index,
row.notes,
row.icon,
row.icon_color,
meta,
is_current,
in_failover_queue,
],
)
.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!(
"provider '{}/{}' already exists",
key.app_type, key.id
))
}
_ => AppError::Database(error.to_string()),
})?;
Ok(())
}
pub(super) fn insert_endpoint(
tx: &Transaction<'_>,
key: &ProviderKey,
endpoint: &NewEndpoint,
) -> Result<(), AppError> {
tx.execute(
"INSERT INTO provider_endpoints
(provider_id, app_type, url, added_at, last_used)
VALUES (?1, ?2, ?3, ?4, ?5)",
params![
key.id,
key.app_type,
endpoint.url,
endpoint.added_at,
endpoint.last_used
],
)
.map_err(|error| AppError::Database(error.to_string()))?;
Ok(())
}
/// Exact aggregate replacement is sealed inside the DAO parent module. The
/// catalog compensation coordinator introduced with the ordered mutation
/// pipeline is the only intended caller.
#[allow(dead_code)]
// The certification contract keeps immutable creation time separate from the
// mutable row DTO and calls this sealed helper directly with the full snapshot.
#[allow(clippy::too_many_arguments)]
pub(super) fn restore_provider_aggregate_on_tx(
tx: &Transaction<'_>,
key: &ProviderKey,
row: &ProviderRowUpdate,
created_at: Option<i64>,
sort_index: Option<usize>,
is_current: bool,
in_failover_queue: bool,
endpoints: &[NewEndpoint],
) -> Result<(), AppError> {
let updated = update_row(tx, key, row)?;
if updated == 0 {
insert_row(
tx,
key,
row,
created_at,
sort_index,
is_current,
in_failover_queue,
)?;
} else {
// Exact compensation is the only path allowed to restore immutable
// creation time after a prior aggregate mutation.
tx.execute(
"UPDATE providers SET created_at = ?1 WHERE id = ?2 AND app_type = ?3",
params![created_at, key.id, key.app_type],
)
.map_err(|error| AppError::Database(error.to_string()))?;
}
tx.execute(
"DELETE FROM provider_endpoints WHERE provider_id = ?1 AND app_type = ?2",
params![key.id, key.app_type],
)
.map_err(|error| AppError::Database(error.to_string()))?;
for endpoint in endpoints {
insert_endpoint(tx, key, endpoint)?;
}
// State and order are maintained by their dedicated authorities. Exact
// compensation may restore their captured values without exposing them in
// ProviderRowUpdate.
tx.execute(
"UPDATE providers
SET sort_index = ?1, is_current = ?2, in_failover_queue = ?3
WHERE id = ?4 AND app_type = ?5",
params![
sort_index,
is_current,
in_failover_queue,
key.id,
key.app_type
],
)
.map_err(|error| AppError::Database(error.to_string()))?;
Ok(())
}
pub(super) fn update_row(
tx: &Transaction<'_>,
key: &ProviderKey,
row: &ProviderRowUpdate,
) -> Result<usize, AppError> {
let (settings_config, meta) = encode_row(row)?;
tx.execute(
"UPDATE providers SET
name = ?1,
settings_config = ?2,
website_url = ?3,
category = ?4,
notes = ?5,
icon = ?6,
icon_color = ?7,
meta = ?8
WHERE id = ?9 AND app_type = ?10",
params![
row.name,
settings_config,
row.website_url,
row.category,
row.notes,
row.icon,
row.icon_color,
meta,
key.id,
key.app_type,
],
)
.map_err(|error| AppError::Database(error.to_string()))
}
impl Database {
pub fn create_provider(&self, input: NewProviderAggregate) -> Result<(), AppError> {
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)?;
}
tx.commit()
.map_err(|error| AppError::Database(error.to_string()))
}
pub fn update_provider(
&self,
key: &ProviderKey,
row: &ProviderRowUpdate,
) -> Result<(), AppError> {
let mut conn = lock_conn!(self.conn);
let tx = conn
.transaction()
.map_err(|error| AppError::Database(error.to_string()))?;
if update_row(&tx, key, row)? != 1 {
return Err(AppError::NotFound(format!(
"provider '{}/{}'",
key.app_type, key.id
)));
}
tx.commit()
.map_err(|error| AppError::Database(error.to_string()))
}
pub(crate) fn update_provider_if_content_fingerprint(
&self,
key: &ProviderKey,
expected_fingerprint: &str,
row: &ProviderRowUpdate,
) -> Result<(), AppError> {
let mut conn = lock_conn!(self.conn);
let tx = conn
.transaction()
.map_err(|error| AppError::Database(error.to_string()))?;
let current = tx
.query_row(
&format!("{PROVIDER_SELECT} WHERE id = ?1 AND app_type = ?2"),
params![key.id, key.app_type],
StoredProviderRow::from_row,
)
.optional()
.map_err(|error| AppError::Database(error.to_string()))?
.ok_or_else(|| AppError::NotFound(format!("provider '{}/{}'", key.app_type, key.id)))?
.decode(key.app_type())?;
if current.row_content_fingerprint() != expected_fingerprint {
return Err(AppError::Conflict(format!(
"provider '{}/{}' changed since it was read",
key.app_type, key.id
)));
}
if update_row(&tx, key, row)? != 1 {
return Err(AppError::NotFound(format!(
"provider '{}/{}'",
key.app_type, key.id
)));
}
tx.commit()
.map_err(|error| AppError::Database(error.to_string()))
}
pub(crate) fn rename_db_only_additive_provider(
&self,
input: RenameProvider,
) -> Result<(), AppError> {
let mut conn = lock_conn!(self.conn);
let tx = conn
.transaction()
.map_err(|error| AppError::Database(error.to_string()))?;
let source_state = tx
.query_row(
"SELECT sort_index, is_current, in_failover_queue, category, created_at, meta
FROM providers
WHERE id = ?1 AND app_type = ?2",
params![input.source.id, input.source.app_type],
|row| {
Ok((
row.get::<_, Option<usize>>(0)?,
row.get::<_, bool>(1)?,
row.get::<_, bool>(2)?,
row.get::<_, Option<String>>(3)?,
row.get::<_, Option<i64>>(4)?,
row.get::<_, String>(5)?,
))
},
)
.optional()
.map_err(|error| AppError::Database(error.to_string()))?
.ok_or_else(|| {
AppError::NotFound(format!(
"provider '{}/{}'",
input.source.app_type, input.source.id
))
})?;
if matches!(source_state.3.as_deref(), Some("omo" | "omo-slim")) {
return Err(AppError::InvalidInput(
"OMO/OMO Slim providers cannot be renamed".to_string(),
));
}
let source_meta: ProviderMeta = if source_state.5.trim().is_empty() {
ProviderMeta::default()
} else {
serde_json::from_str(&source_state.5).map_err(|error| {
AppError::Database(format!(
"invalid meta for provider '{}/{}': {error}",
input.source.app_type, input.source.id
))
})?
};
if source_meta.live_config_managed == Some(true) {
return Err(AppError::Conflict(format!(
"provider '{}/{}' became live-managed before rename",
input.source.app_type, input.source.id
)));
}
let target = ProviderKey::new(&input.source.app_type, &input.target_id)?;
insert_row(
&tx,
&target,
&input.row,
source_state.4,
source_state.0,
source_state.1,
source_state.2,
)?;
tx.execute(
"INSERT INTO provider_endpoints
(provider_id, app_type, url, added_at, last_used)
SELECT ?1, app_type, url, added_at, last_used
FROM provider_endpoints
WHERE provider_id = ?2 AND app_type = ?3
ORDER BY id",
params![target.id, input.source.id, input.source.app_type],
)
.map_err(|error| AppError::Database(error.to_string()))?;
if tx
.execute(
"DELETE FROM providers WHERE id = ?1 AND app_type = ?2",
params![input.source.id, input.source.app_type],
)
.map_err(|error| AppError::Database(error.to_string()))?
!= 1
{
return Err(AppError::NotFound(format!(
"provider '{}/{}'",
input.source.app_type, input.source.id
)));
}
tx.commit()
.map_err(|error| AppError::Database(error.to_string()))
}
pub fn add_provider_endpoint(
&self,
key: &ProviderKey,
endpoint: NewEndpoint,
) -> Result<(), AppError> {
let mut conn = lock_conn!(self.conn);
let tx = conn
.transaction()
.map_err(|error| AppError::Database(error.to_string()))?;
insert_endpoint(&tx, key, &endpoint)?;
tx.commit()
.map_err(|error| AppError::Database(error.to_string()))
}
pub fn remove_provider_endpoint(&self, key: &ProviderKey, url: &str) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
if conn
.execute(
"DELETE FROM provider_endpoints
WHERE provider_id = ?1 AND app_type = ?2 AND url = ?3",
params![key.id, key.app_type, url],
)
.map_err(|error| AppError::Database(error.to_string()))?
!= 1
{
return Err(AppError::NotFound(format!(
"provider endpoint '{}/{}/{}'",
key.app_type, key.id, url
)));
}
Ok(())
}
pub fn touch_provider_endpoint(
&self,
key: &ProviderKey,
url: &str,
at: i64,
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
if conn
.execute(
"UPDATE provider_endpoints
SET last_used = ?1
WHERE provider_id = ?2 AND app_type = ?3 AND url = ?4",
params![at, key.id, key.app_type, url],
)
.map_err(|error| AppError::Database(error.to_string()))?
!= 1
{
return Err(AppError::NotFound(format!(
"provider endpoint '{}/{}/{}'",
key.app_type, key.id, url
)));
}
Ok(())
}
pub(crate) fn update_provider_sort_index(
&self,
updates: &[(ProviderKey, usize)],
) -> Result<(), AppError> {
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()
)));
}
}
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()))
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,308 +0,0 @@
//! 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<Self, Self::Err> {
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<String>,
pub created_at: i64,
pub updated_at: i64,
}
fn decode_deployment(row: &rusqlite::Row<'_>) -> rusqlite::Result<SkillDeployment> {
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<Option<SkillDeployment>, 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<Vec<SkillDeployment>, 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 get_all_pi_skill_deployments(&self) -> Result<Vec<SkillDeployment>, 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'
ORDER BY skill_id, created_at, destination_key",
)
.map_err(|error| AppError::Database(error.to_string()))?;
let rows = stmt
.query_map([], 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<bool>,
) -> 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<bool, AppError> {
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<bool>,
) -> Result<bool, AppError> {
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(())
}
}
+13 -84
View File
@@ -23,8 +23,7 @@ impl Database {
.prepare( .prepare(
"SELECT id, name, description, directory, repo_owner, repo_name, repo_branch, "SELECT id, name, description, directory, repo_owner, repo_name, repo_branch,
readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild,
enabled_opencode, enabled_hermes, enabled_pi, enabled_opencode, enabled_hermes, installed_at, content_hash, updated_at
installed_at, content_hash, updated_at
FROM skills ORDER BY name ASC", FROM skills ORDER BY name ASC",
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
@@ -47,11 +46,10 @@ impl Database {
grokbuild: row.get(11)?, grokbuild: row.get(11)?,
opencode: row.get(12)?, opencode: row.get(12)?,
hermes: row.get(13)?, hermes: row.get(13)?,
pi: row.get(14)?,
}, },
installed_at: row.get(15)?, installed_at: row.get(14)?,
content_hash: row.get(16)?, content_hash: row.get(15)?,
updated_at: row.get::<_, i64>(17).unwrap_or(0), updated_at: row.get::<_, i64>(16).unwrap_or(0),
}) })
}) })
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
@@ -71,8 +69,7 @@ impl Database {
.prepare( .prepare(
"SELECT id, name, description, directory, repo_owner, repo_name, repo_branch, "SELECT id, name, description, directory, repo_owner, repo_name, repo_branch,
readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild,
enabled_opencode, enabled_hermes, enabled_pi, enabled_opencode, enabled_hermes, installed_at, content_hash, updated_at
installed_at, content_hash, updated_at
FROM skills WHERE id = ?1", FROM skills WHERE id = ?1",
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
@@ -94,11 +91,10 @@ impl Database {
grokbuild: row.get(11)?, grokbuild: row.get(11)?,
opencode: row.get(12)?, opencode: row.get(12)?,
hermes: row.get(13)?, hermes: row.get(13)?,
pi: row.get(14)?,
}, },
installed_at: row.get(15)?, installed_at: row.get(14)?,
content_hash: row.get(16)?, content_hash: row.get(15)?,
updated_at: row.get::<_, i64>(17).unwrap_or(0), updated_at: row.get::<_, i64>(16).unwrap_or(0),
}) })
}); });
@@ -113,28 +109,11 @@ impl Database {
pub fn save_skill(&self, skill: &InstalledSkill) -> Result<(), AppError> { pub fn save_skill(&self, skill: &InstalledSkill) -> Result<(), AppError> {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
conn.execute( conn.execute(
"INSERT INTO skills "INSERT OR REPLACE INTO skills
(id, name, description, directory, repo_owner, repo_name, repo_branch, (id, name, description, directory, repo_owner, repo_name, repo_branch,
readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, enabled_opencode, enabled_hermes, readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, enabled_opencode, enabled_hermes,
enabled_pi, installed_at, content_hash, updated_at) 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) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17)",
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![ params![
skill.id, skill.id,
skill.name, skill.name,
@@ -150,7 +129,6 @@ impl Database {
skill.apps.grokbuild, skill.apps.grokbuild,
skill.apps.opencode, skill.apps.opencode,
skill.apps.hermes, skill.apps.hermes,
skill.apps.pi,
skill.installed_at, skill.installed_at,
skill.content_hash, skill.content_hash,
skill.updated_at, skill.updated_at,
@@ -182,8 +160,8 @@ impl Database {
let conn = lock_conn!(self.conn); let conn = lock_conn!(self.conn);
let affected = conn let affected = conn
.execute( .execute(
"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", "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, apps.pi, id], params![apps.claude, apps.codex, apps.gemini, apps.grokbuild, apps.opencode, apps.hermes, id],
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
Ok(affected > 0) Ok(affected > 0)
@@ -284,52 +262,3 @@ impl Database {
Ok(count) 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(())
}
}
+1 -47
View File
@@ -32,10 +32,6 @@ mod schema;
mod tests; mod tests;
// DAO 类型导出供外部使用 // DAO 类型导出供外部使用
pub(crate) use dao::pi_projections::PiProviderProjection;
pub use dao::provider_write::{
NewEndpoint, NewProviderAggregate, ProviderKey, ProviderRowUpdate, RenameProvider,
};
pub(crate) use dao::providers_seed::{ pub(crate) use dao::providers_seed::{
is_official_seed_id, CLAUDE_DESKTOP_OFFICIAL_PROVIDER_ID, CODEX_OFFICIAL_PROVIDER_ID, is_official_seed_id, CLAUDE_DESKTOP_OFFICIAL_PROVIDER_ID, CODEX_OFFICIAL_PROVIDER_ID,
GROKBUILD_OFFICIAL_PROVIDER_ID, GROKBUILD_OFFICIAL_PROVIDER_ID,
@@ -44,7 +40,6 @@ pub(crate) use dao::proxy::{
validate_cost_multiplier, validate_pricing_source, PRICING_SOURCE_REQUEST, validate_cost_multiplier, validate_pricing_source, PRICING_SOURCE_REQUEST,
PRICING_SOURCE_RESPONSE, PRICING_SOURCE_RESPONSE,
}; };
pub(crate) use dao::skill_deployments::{SkillDeployment, SkillDeploymentMethod};
pub use dao::FailoverQueueItem; pub use dao::FailoverQueueItem;
pub use dao::Profile; pub use dao::Profile;
@@ -58,7 +53,7 @@ use std::sync::Mutex;
/// 当前 Schema 版本号 /// 当前 Schema 版本号
/// 每次修改表结构时递增,并在 schema.rs 中添加相应的迁移逻辑 /// 每次修改表结构时递增,并在 schema.rs 中添加相应的迁移逻辑
pub(crate) const SCHEMA_VERSION: i32 = 17; pub(crate) const SCHEMA_VERSION: i32 = 16;
/// 安全地序列化 JSON,避免 unwrap panic /// 安全地序列化 JSON,避免 unwrap panic
pub(crate) fn to_json_string<T: Serialize>(value: &T) -> Result<String, AppError> { pub(crate) fn to_json_string<T: Serialize>(value: &T) -> Result<String, AppError> {
@@ -202,11 +197,6 @@ impl Database {
conn: Mutex::new(conn), conn: Mutex::new(conn),
}; };
db.create_tables()?; db.create_tables()?;
// Keep the test database structurally identical to a fresh production
// database. Marking the base DDL as current without running the
// migration chain creates a false-current schema and makes restore
// tests certify columns that do not actually exist.
db.apply_schema_migrations()?;
db.ensure_model_pricing_seeded()?; db.ensure_model_pricing_seeded()?;
Ok(db) Ok(db)
@@ -303,39 +293,3 @@ impl Database {
Ok(count == 0) Ok(count == 0)
} }
} }
#[cfg(test)]
impl Database {
/// Test-fixture reconciliation helper. Production code cannot call this:
/// provider writes there must choose a typed create or update operation.
pub(crate) fn reconcile_provider_fixture(
&self,
app_type: &str,
provider: &crate::provider::Provider,
) -> Result<(), AppError> {
let mut input = crate::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,
};
if self.get_provider_aggregate(app_type, &input.id)?.is_some() {
if let Some(meta) = input.meta.as_mut() {
meta.custom_endpoints.clear();
}
let key = ProviderKey::new(app_type, input.id.clone())?;
let row = ProviderRowUpdate::from_input(&input)?;
self.update_provider(&key, &row)
} else {
self.create_provider(NewProviderAggregate::from_input(app_type, input)?)
}
}
}
+3 -212
View File
@@ -53,10 +53,7 @@ impl Database {
app_type TEXT NOT NULL, app_type TEXT NOT NULL,
url TEXT NOT NULL, url TEXT NOT NULL,
added_at INTEGER, added_at INTEGER,
last_used INTEGER, FOREIGN KEY (provider_id, app_type) REFERENCES providers(id, app_type) ON DELETE CASCADE
FOREIGN KEY (provider_id, app_type)
REFERENCES providers(id, app_type) ON DELETE CASCADE,
UNIQUE (provider_id, app_type, url)
)", )",
[], [],
) )
@@ -100,7 +97,6 @@ impl Database {
enabled_grokbuild BOOLEAN NOT NULL DEFAULT 0, enabled_grokbuild BOOLEAN NOT NULL DEFAULT 0,
enabled_opencode BOOLEAN NOT NULL DEFAULT 0, enabled_opencode BOOLEAN NOT NULL DEFAULT 0,
enabled_hermes BOOLEAN NOT NULL DEFAULT 0, enabled_hermes BOOLEAN NOT NULL DEFAULT 0,
enabled_pi BOOLEAN NOT NULL DEFAULT 0,
installed_at INTEGER NOT NULL DEFAULT 0, installed_at INTEGER NOT NULL DEFAULT 0,
content_hash TEXT, content_hash TEXT,
updated_at INTEGER NOT NULL DEFAULT 0 updated_at INTEGER NOT NULL DEFAULT 0
@@ -109,36 +105,6 @@ impl Database {
) )
.map_err(|e| AppError::Database(e.to_string()))?; .map_err(|e| AppError::Database(e.to_string()))?;
// Reserve the v17 device-local ledgers here so later stacked features
// never mutate the semantics of an already-published migration.
conn.execute(
"CREATE TABLE IF NOT EXISTS pi_provider_projections (
provider_id TEXT PRIMARY KEY,
provider_key TEXT NOT NULL UNIQUE,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
conn.execute(
"CREATE TABLE IF NOT EXISTS skill_deployments (
app_type TEXT NOT NULL CHECK (app_type = 'pi'),
skill_id TEXT NOT NULL,
destination TEXT NOT NULL,
destination_key TEXT NOT NULL,
method TEXT NOT NULL CHECK (method IN ('symlink', 'copy')),
source_identity TEXT NOT NULL,
deployed_digest TEXT,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
PRIMARY KEY (app_type, skill_id, destination_key),
UNIQUE (app_type, destination_key)
)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
// 6. Skill Repos 表 // 6. Skill Repos 表
conn.execute( conn.execute(
"CREATE TABLE IF NOT EXISTS skill_repos ( "CREATE TABLE IF NOT EXISTS skill_repos (
@@ -545,13 +511,6 @@ impl Database {
Self::migrate_v15_to_v16(conn)?; Self::migrate_v15_to_v16(conn)?;
Self::set_user_version(conn, 16)?; Self::set_user_version(conn, 16)?;
} }
16 => {
log::info!(
"迁移数据库从 v16 到 v17(规范化 provider endpoint 并预留设备本地 ledger"
);
Self::migrate_v16_to_v17(conn)?;
Self::set_user_version(conn, 17)?;
}
_ => { _ => {
return Err(AppError::Database(format!( return Err(AppError::Database(format!(
"未知的数据库版本 {version},无法迁移到 {SCHEMA_VERSION}" "未知的数据库版本 {version},无法迁移到 {SCHEMA_VERSION}"
@@ -1564,112 +1523,6 @@ impl Database {
crate::services::session_usage_codex::reset_codex_usage_on_conn(conn, &codex_dir) crate::services::session_usage_codex::reset_codex_usage_on_conn(conn, &codex_dir)
} }
/// v16 -> v17: make endpoint rows a lossless, uniquely owned child
/// collection. The device-local ledger DDL is reserved in the same
/// migration because later stacked PRs must not rewrite a released
/// user_version step.
fn migrate_v16_to_v17(conn: &Connection) -> Result<(), AppError> {
if Self::table_exists(conn, "provider_endpoints")? {
Self::add_column_if_missing(conn, "provider_endpoints", "last_used", "INTEGER")?;
conn.execute_batch(
"UPDATE provider_endpoints AS kept
SET added_at = (
SELECT MIN(other.added_at)
FROM provider_endpoints AS other
WHERE other.provider_id = kept.provider_id
AND other.app_type = kept.app_type
AND other.url = kept.url
),
last_used = (
SELECT MAX(other.last_used)
FROM provider_endpoints AS other
WHERE other.provider_id = kept.provider_id
AND other.app_type = kept.app_type
AND other.url = kept.url
)
WHERE kept.id = (
SELECT MIN(other.id)
FROM provider_endpoints AS other
WHERE other.provider_id = kept.provider_id
AND other.app_type = kept.app_type
AND other.url = kept.url
);
DELETE FROM provider_endpoints
WHERE id NOT IN (
SELECT MIN(id)
FROM provider_endpoints
GROUP BY provider_id, app_type, url
);
DROP TABLE IF EXISTS provider_endpoints_v17_canonical;
CREATE TABLE provider_endpoints_v17_canonical (
id INTEGER PRIMARY KEY AUTOINCREMENT,
provider_id TEXT NOT NULL,
app_type TEXT NOT NULL,
url TEXT NOT NULL,
added_at INTEGER,
last_used INTEGER,
FOREIGN KEY (provider_id, app_type)
REFERENCES providers(id, app_type) ON DELETE CASCADE,
UNIQUE (provider_id, app_type, url)
);
INSERT INTO provider_endpoints_v17_canonical
(id, provider_id, app_type, url, added_at, last_used)
SELECT id, provider_id, app_type, url, added_at, last_used
FROM provider_endpoints;
DROP TABLE provider_endpoints;
ALTER TABLE provider_endpoints_v17_canonical
RENAME TO provider_endpoints;",
)
.map_err(|error| AppError::Database(error.to_string()))?;
} else {
conn.execute(
"CREATE TABLE provider_endpoints (
id INTEGER PRIMARY KEY AUTOINCREMENT,
provider_id TEXT NOT NULL,
app_type TEXT NOT NULL,
url TEXT NOT NULL,
added_at INTEGER,
last_used INTEGER,
FOREIGN KEY (provider_id, app_type)
REFERENCES providers(id, app_type) ON DELETE CASCADE,
UNIQUE (provider_id, app_type, url)
)",
[],
)
.map_err(|error| AppError::Database(error.to_string()))?;
}
if Self::table_exists(conn, "skills")? {
Self::add_column_if_missing(
conn,
"skills",
"enabled_pi",
"BOOLEAN NOT NULL DEFAULT 0",
)?;
}
conn.execute_batch(
"CREATE TABLE IF NOT EXISTS pi_provider_projections (
provider_id TEXT PRIMARY KEY,
provider_key TEXT NOT NULL UNIQUE,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS skill_deployments (
app_type TEXT NOT NULL CHECK (app_type = 'pi'),
skill_id TEXT NOT NULL,
destination TEXT NOT NULL,
destination_key TEXT NOT NULL,
method TEXT NOT NULL CHECK (method IN ('symlink', 'copy')),
source_identity TEXT NOT NULL,
deployed_digest TEXT,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
PRIMARY KEY (app_type, skill_id, destination_key),
UNIQUE (app_type, destination_key)
);",
)
.map_err(|error| AppError::Database(error.to_string()))
}
/// 插入默认模型定价数据 /// 插入默认模型定价数据
/// 格式: (model_id, display_name, input, output, cache_read, cache_creation) /// 格式: (model_id, display_name, input, output, cache_read, cache_creation)
/// 注意: model_id 使用短横线格式(如 claude-haiku-4-5),与 API 返回的模型名称标准化后一致 /// 注意: model_id 使用短横线格式(如 claude-haiku-4-5),与 API 返回的模型名称标准化后一致
@@ -2384,6 +2237,7 @@ impl Database {
"0", "0",
), ),
// Qwen 系列 (阿里巴巴) // Qwen 系列 (阿里巴巴)
("qwen3.8-max", "Qwen3.8 Max", "2", "6", "0.25", "2.50"),
("qwen3.7-max", "Qwen3.7 Max", "2.50", "7.50", "0.25", "0"), ("qwen3.7-max", "Qwen3.7 Max", "2.50", "7.50", "0.25", "0"),
("qwen3.7-plus", "Qwen3.7 Plus", "0.40", "1.60", "0.08", "0"), ("qwen3.7-plus", "Qwen3.7 Plus", "0.40", "1.60", "0.08", "0"),
( (
@@ -3369,7 +3223,7 @@ mod tests {
Database::apply_schema_migrations_on_conn(&conn)?; Database::apply_schema_migrations_on_conn(&conn)?;
assert_eq!(Database::get_user_version(&conn)?, SCHEMA_VERSION); assert_eq!(Database::get_user_version(&conn)?, 16);
let counts: (i64, i64, i64, i64) = conn.query_row( let counts: (i64, i64, i64, i64) = conn.query_row(
"SELECT "SELECT
(SELECT COUNT(*) FROM proxy_request_logs WHERE data_source = 'codex_session'), (SELECT COUNT(*) FROM proxy_request_logs WHERE data_source = 'codex_session'),
@@ -3382,67 +3236,4 @@ mod tests {
assert_eq!(counts, (0, 1, 0, 1)); assert_eq!(counts, (0, 1, 0, 1));
Ok(()) Ok(())
} }
#[test]
fn migrate_v16_to_v17_preserves_endpoint_metadata_and_starts_ledgers_empty(
) -> Result<(), AppError> {
let conn = Connection::open_in_memory()?;
conn.execute_batch(
"CREATE TABLE providers (
id TEXT NOT NULL,
app_type TEXT NOT NULL,
PRIMARY KEY (id, app_type)
);
CREATE TABLE provider_endpoints (
id INTEGER PRIMARY KEY,
provider_id TEXT NOT NULL,
app_type TEXT NOT NULL,
url TEXT NOT NULL,
added_at INTEGER
);
CREATE TABLE skills (
id TEXT PRIMARY KEY,
enabled_codex BOOLEAN NOT NULL DEFAULT 0
);
INSERT INTO providers (id, app_type) VALUES ('provider', 'pi');
INSERT INTO provider_endpoints
(id, provider_id, app_type, url, added_at)
VALUES
(1, 'provider', 'pi', 'https://duplicate.test', 20),
(2, 'provider', 'pi', 'https://duplicate.test', 10);
INSERT INTO skills (id, enabled_codex) VALUES ('existing', 1);",
)?;
Database::set_user_version(&conn, 16)?;
Database::apply_schema_migrations_on_conn(&conn)?;
assert_eq!(Database::get_user_version(&conn)?, SCHEMA_VERSION);
assert!(Database::has_column(
&conn,
"provider_endpoints",
"last_used"
)?);
assert!(Database::has_column(&conn, "skills", "enabled_pi")?);
assert!(Database::table_exists(&conn, "pi_provider_projections")?);
assert!(Database::table_exists(&conn, "skill_deployments")?);
let endpoint: (i64, Option<i64>) = conn.query_row(
"SELECT COUNT(*), MIN(added_at)
FROM provider_endpoints
WHERE provider_id = 'provider'
AND app_type = 'pi'
AND url = 'https://duplicate.test'",
[],
|row| Ok((row.get(0)?, row.get(1)?)),
)?;
assert_eq!(endpoint, (1, Some(10)));
let ledgers: (i64, i64) = conn.query_row(
"SELECT
(SELECT COUNT(*) FROM pi_provider_projections),
(SELECT COUNT(*) FROM skill_deployments)",
[],
|row| Ok((row.get(0)?, row.get(1)?)),
)?;
assert_eq!(ledgers, (0, 0));
Ok(())
}
} }
-2
View File
@@ -1,5 +1,3 @@
#![cfg(test)]
//! 数据库模块测试 //! 数据库模块测试
//! //!
//! 包含 Schema 迁移和基本功能的测试。 //! 包含 Schema 迁移和基本功能的测试。
-4
View File
@@ -66,10 +66,6 @@ pub struct DeepLinkImportRequest {
/// Optional model name /// Optional model name
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>, pub model: Option<String>,
/// 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<String>,
/// Optional notes/description /// Optional notes/description
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub notes: Option<String>, pub notes: Option<String>,
+4 -27
View File
@@ -81,10 +81,10 @@ fn parse_provider_deeplink(
// Validate app type // Validate app type
if !matches!( if !matches!(
app.as_str(), app.as_str(),
"claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes" | "pi" "claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes"
) { ) {
return Err(AppError::InvalidInput(format!( return Err(AppError::InvalidInput(format!(
"Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', 'hermes', or 'pi', got '{app}'" "Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', or 'hermes', got '{app}'"
))); )));
} }
@@ -116,7 +116,6 @@ fn parse_provider_deeplink(
// Extract optional fields // Extract optional fields
let model = params.get("model").cloned(); let model = params.get("model").cloned();
let api = params.get("api").cloned();
let notes = params.get("notes").cloned(); let notes = params.get("notes").cloned();
let haiku_model = params.get("haikuModel").cloned(); let haiku_model = params.get("haikuModel").cloned();
let sonnet_model = params.get("sonnetModel").cloned(); let sonnet_model = params.get("sonnetModel").cloned();
@@ -128,24 +127,6 @@ fn parse_provider_deeplink(
let config = params.get("config").cloned(); let config = params.get("config").cloned();
let config_format = params.get("configFormat").cloned(); let config_format = params.get("configFormat").cloned();
let config_url = params.get("configUrl").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::<bool>().ok()); let enabled = params.get("enabled").and_then(|v| v.parse::<bool>().ok());
// Extract usage script fields (v3.9+) // Extract usage script fields (v3.9+)
@@ -172,7 +153,6 @@ fn parse_provider_deeplink(
api_key, api_key,
icon, icon,
model, model,
api,
notes, notes,
haiku_model, haiku_model,
sonnet_model, sonnet_model,
@@ -210,10 +190,10 @@ fn parse_prompt_deeplink(
// Validate app type // Validate app type
if !matches!( if !matches!(
app.as_str(), app.as_str(),
"claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes" | "pi" "claude" | "codex" | "gemini" | "grokbuild" | "opencode" | "openclaw" | "hermes"
) { ) {
return Err(AppError::InvalidInput(format!( return Err(AppError::InvalidInput(format!(
"Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', 'hermes', or 'pi', got '{app}'" "Invalid app type: must be 'claude', 'codex', 'gemini', 'grokbuild', 'opencode', 'openclaw', or 'hermes', got '{app}'"
))); )));
} }
@@ -245,7 +225,6 @@ fn parse_prompt_deeplink(
endpoint: None, endpoint: None,
api_key: None, api_key: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -319,7 +298,6 @@ fn parse_mcp_deeplink(
endpoint: None, endpoint: None,
api_key: None, api_key: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -375,7 +353,6 @@ fn parse_skill_deeplink(
endpoint: None, endpoint: None,
api_key: None, api_key: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
+16 -64
View File
@@ -109,35 +109,27 @@ pub fn import_provider_from_deeplink(
let provider_id = provider.id.clone(); let provider_id = provider.id.clone();
// All endpoints supplied by one import request belong to the same create // Use ProviderService to add the provider
// intent. Put the non-primary endpoints into the initial aggregate so the ProviderService::add(state, app_type.clone(), provider, true)?;
// provider row and its complete endpoint set commit atomically.
let initial_endpoints = &mut provider // Add extra endpoints as custom endpoints (skip first one as it's the primary)
.meta for ep in all_endpoints.iter().skip(1) {
.get_or_insert_with(ProviderMeta::default) let normalized = ep.trim().trim_end_matches('/').to_string();
.custom_endpoints;
for endpoint in all_endpoints.iter().skip(1) {
let normalized = endpoint.trim().trim_end_matches('/').to_string();
if !normalized.is_empty() { if !normalized.is_empty() {
initial_endpoints.insert( if let Err(e) = ProviderService::add_custom_endpoint(
state,
app_type.clone(),
&provider_id,
normalized.clone(), normalized.clone(),
crate::settings::CustomEndpoint { ) {
url: normalized, log::warn!(
added_at: Some(timestamp), "Failed to add custom endpoint '{}': {e}",
last_used: None, crate::url_for_log(&normalized)
}, );
); }
} }
} }
// ProviderService owns the strict aggregate create.
ProviderService::add(
state,
app_type.clone(),
crate::services::provider::provider_to_mutation_input(provider),
true,
)?;
// If enabled=true, set as current provider // If enabled=true, set as current provider
if merged_request.enabled.unwrap_or(false) { if merged_request.enabled.unwrap_or(false) {
ProviderService::switch(state, app_type.clone(), &provider_id)?; ProviderService::switch(state, app_type.clone(), &provider_id)?;
@@ -160,7 +152,6 @@ pub(crate) fn build_provider_from_request(
AppType::OpenCode => build_opencode_settings(request), AppType::OpenCode => build_opencode_settings(request),
AppType::OpenClaw => build_additive_app_settings(request), AppType::OpenClaw => build_additive_app_settings(request),
AppType::Hermes => build_hermes_settings(request), AppType::Hermes => build_hermes_settings(request),
AppType::Pi => build_pi_settings(request)?,
}; };
// Build usage script configuration if provided // Build usage script configuration if provided
@@ -592,45 +583,6 @@ fn build_hermes_settings(request: &DeepLinkImportRequest) -> serde_json::Value {
json!(config) 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<serde_json::Value, AppError> {
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 // Config Merge Logic
// ============================================================================= // =============================================================================
+1 -101
View File
@@ -1,11 +1,9 @@
#![cfg(test)]
//! Deep link module tests //! Deep link module tests
use super::mcp::parse_mcp_apps; use super::mcp::parse_mcp_apps;
use super::parser::parse_deeplink_url; use super::parser::parse_deeplink_url;
use super::prompt::import_prompt_from_deeplink; use super::prompt::import_prompt_from_deeplink;
use super::provider::{import_provider_from_deeplink, parse_and_merge_config}; use super::provider::parse_and_merge_config;
use super::utils::{infer_homepage_from_endpoint, validate_url}; use super::utils::{infer_homepage_from_endpoint, validate_url};
use super::DeepLinkImportRequest; use super::DeepLinkImportRequest;
use crate::AppType; use crate::AppType;
@@ -89,61 +87,6 @@ fn test_parse_deeplink_with_notes() {
assert_eq!(request.notes, Some("Test notes".to_string())); 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] #[test]
fn test_parse_grokbuild_provider() { fn test_parse_grokbuild_provider() {
use super::provider::build_provider_from_request; use super::provider::build_provider_from_request;
@@ -265,7 +208,6 @@ fn test_build_gemini_provider_with_model() {
api_key: Some("test-api-key".to_string()), api_key: Some("test-api-key".to_string()),
icon: None, icon: None,
model: Some("gemini-2.0-flash".to_string()), model: Some("gemini-2.0-flash".to_string()),
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -319,7 +261,6 @@ fn test_build_gemini_provider_without_model() {
api_key: Some("test-api-key".to_string()), api_key: Some("test-api-key".to_string()),
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -366,7 +307,6 @@ fn test_deeplink_usage_script_does_not_copy_provider_credentials() {
api_key: Some("sk-main".to_string()), api_key: Some("sk-main".to_string()),
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -414,7 +354,6 @@ fn usage_script_request(code: &str, usage_enabled: Option<bool>) -> DeepLinkImpo
api_key: Some("sk-main".to_string()), api_key: Some("sk-main".to_string()),
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -498,7 +437,6 @@ fn test_deeplink_usage_script_omits_explicit_credentials_that_match_provider() {
api_key: Some("sk-main".to_string()), api_key: Some("sk-main".to_string()),
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -547,7 +485,6 @@ fn test_deeplink_usage_script_preserves_distinct_usage_credentials() {
api_key: Some("sk-main".to_string()), api_key: Some("sk-main".to_string()),
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -601,7 +538,6 @@ fn test_parse_and_merge_config_claude() {
api_key: None, api_key: None,
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -725,7 +661,6 @@ fn test_parse_and_merge_config_url_override() {
api_key: Some("sk-new".to_string()), // URL param should override api_key: Some("sk-new".to_string()), // URL param should override
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -789,7 +724,6 @@ fn test_build_claude_provider_preserves_custom_env_fields() {
icon: None, icon: None,
// URL param: must win over the same key in config (haiku-from-config) // URL param: must win over the same key in config (haiku-from-config)
model: Some("main-model".to_string()), model: Some("main-model".to_string()),
api: None,
notes: None, notes: None,
haiku_model: Some("haiku-from-url".to_string()), haiku_model: Some("haiku-from-url".to_string()),
sonnet_model: None, sonnet_model: None,
@@ -845,7 +779,6 @@ fn test_build_claude_provider_without_config_unchanged() {
api_key: Some("sk".to_string()), api_key: Some("sk".to_string()),
icon: None, icon: None,
model: None, model: None,
api: None,
notes: None, notes: None,
haiku_model: None, haiku_model: None,
sonnet_model: None, sonnet_model: None,
@@ -1019,39 +952,6 @@ fn test_parse_multiple_endpoints_comma_separated() {
assert!(endpoint.contains("https://api3.example.com")); assert!(endpoint.contains("https://api3.example.com"));
} }
#[test]
#[serial_test::serial]
fn provider_deeplink_creates_all_initial_endpoints_in_one_aggregate() {
let _test_home = TestHomeGuard::new();
let request = parse_deeplink_url(
"ccswitch://v1/import?resource=provider&app=claude&name=Endpoint%20Aggregate&endpoint=https%3A%2F%2Fprimary.example.com,https%3A%2F%2Fsecond.example.com%2F,https%3A%2F%2Fthird.example.com&apiKey=sk-test",
)
.expect("parse provider deeplink");
let state = AppState::new(Arc::new(Database::memory().expect("create memory db")));
let provider_id =
import_provider_from_deeplink(&state, request).expect("import provider aggregate");
let aggregate = state
.db
.get_provider_aggregate(AppType::Claude.as_str(), &provider_id)
.expect("read provider aggregate")
.expect("provider exists");
assert_eq!(aggregate.endpoints.len(), 2);
assert_eq!(
aggregate.endpoints["https://second.example.com"].url,
"https://second.example.com"
);
assert_eq!(
aggregate.endpoints["https://third.example.com"].url,
"https://third.example.com"
);
assert!(aggregate
.endpoints
.values()
.all(|endpoint| endpoint.added_at.is_some() && endpoint.last_used.is_none()));
}
#[test] #[test]
fn test_parse_single_endpoint_backward_compatible() { fn test_parse_single_endpoint_backward_compatible() {
// Old format with single endpoint should still work // Old format with single endpoint should still work
-7
View File
@@ -9,13 +9,6 @@ pub enum AppError {
Config(String), Config(String),
#[error("无效输入: {0}")] #[error("无效输入: {0}")]
InvalidInput(String), InvalidInput(String),
#[error("未找到: {0}")]
NotFound(String),
/// 结构化冲突:并发前置期望失败(如 reconcile 的 ExpectAbsent 撞上竞争
/// 创建、ExpectPresent 的指纹过期)。调用方据此重读重试或上浮,不得解析
/// Database(String) 文本。由前置工程 A 认证契约引入(T9)。
#[error("并发冲突: {0}")]
Conflict(String),
#[error("IO 错误: {path}: {source}")] #[error("IO 错误: {path}: {source}")]
Io { Io {
path: String, path: String,
+12 -79
View File
@@ -25,7 +25,6 @@ mod model_capabilities;
mod openclaw_config; mod openclaw_config;
mod opencode_config; mod opencode_config;
mod panic_hook; mod panic_hook;
mod pi_config;
mod prompt; mod prompt;
mod prompt_files; mod prompt_files;
mod provider; mod provider;
@@ -39,9 +38,6 @@ mod tray;
mod usage_events; mod usage_events;
mod usage_script; mod usage_script;
#[cfg(test)]
mod architecture_tests;
pub use app_config::{AppType, InstalledSkill, McpApps, McpServer, MultiAppConfig, SkillApps}; pub use app_config::{AppType, InstalledSkill, McpApps, McpServer, MultiAppConfig, SkillApps};
pub use codex_config::{ pub use codex_config::{
get_codex_auth_path, get_codex_config_path, read_codex_live_settings, write_codex_live_atomic, get_codex_auth_path, get_codex_config_path, read_codex_live_settings, write_codex_live_atomic,
@@ -49,10 +45,7 @@ pub use codex_config::{
pub use commands::open_provider_terminal; pub use commands::open_provider_terminal;
pub use commands::*; pub use commands::*;
pub use config::{get_claude_mcp_path, get_claude_settings_path, read_json_file}; pub use config::{get_claude_mcp_path, get_claude_settings_path, read_json_file};
pub use database::{ pub use database::{Database, Profile};
Database, NewEndpoint, NewProviderAggregate, Profile, ProviderKey, ProviderRowUpdate,
RenameProvider,
};
pub use deeplink::{import_provider_from_deeplink, parse_deeplink_url, DeepLinkImportRequest}; pub use deeplink::{import_provider_from_deeplink, parse_deeplink_url, DeepLinkImportRequest};
pub use error::AppError; pub use error::AppError;
pub use grok_config::get_grok_config_path; pub use grok_config::get_grok_config_path;
@@ -64,7 +57,7 @@ pub use mcp::{
sync_single_server_to_gemini, sync_single_server_to_grokbuild, sync_single_server_to_gemini, sync_single_server_to_grokbuild,
}; };
pub use prompt::Prompt; pub use prompt::Prompt;
pub use provider::{Provider, ProviderAggregate, ProviderMeta, ProviderMutationInput}; pub use provider::{Provider, ProviderMeta};
pub use services::{ pub use services::{
profile::{ProfilePayload, ProfileScope, ProfileService}, profile::{ProfilePayload, ProfileScope, ProfileService},
provider::reapply_current_codex_official_live, provider::reapply_current_codex_official_live,
@@ -953,7 +946,6 @@ pub fn run() {
crate::app_config::AppType::OpenCode, crate::app_config::AppType::OpenCode,
crate::app_config::AppType::OpenClaw, crate::app_config::AppType::OpenClaw,
crate::app_config::AppType::Hermes, crate::app_config::AppType::Hermes,
crate::app_config::AppType::Pi,
] { ] {
match crate::services::prompt::PromptService::import_from_file_on_first_launch( match crate::services::prompt::PromptService::import_from_file_on_first_launch(
&app_state, &app_state,
@@ -1332,12 +1324,6 @@ pub fn run() {
commands::remove_provider_from_live_config, commands::remove_provider_from_live_config,
commands::switch_provider, commands::switch_provider,
commands::import_default_config, 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_status,
commands::get_claude_desktop_default_routes, commands::get_claude_desktop_default_routes,
commands::import_claude_desktop_providers_from_claude, commands::import_claude_desktop_providers_from_claude,
@@ -1422,14 +1408,6 @@ pub fn run() {
commands::enable_prompt, commands::enable_prompt,
commands::import_prompt_from_file, commands::import_prompt_from_file,
commands::get_current_prompt_file_content, 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 (项目配置方案) // Profile management (项目配置方案)
commands::list_profiles, commands::list_profiles,
commands::create_profile, commands::create_profile,
@@ -1484,7 +1462,6 @@ pub fn run() {
commands::restore_env_backup, commands::restore_env_backup,
// Skill management (v3.10.0+ unified) // Skill management (v3.10.0+ unified)
commands::get_installed_skills, commands::get_installed_skills,
commands::get_pi_skill_statuses,
commands::get_skill_backups, commands::get_skill_backups,
commands::delete_skill_backup, commands::delete_skill_backup,
commands::install_skill_unified, commands::install_skill_unified,
@@ -1850,11 +1827,7 @@ pub async fn cleanup_before_exit(app_handle: &tauri::AppHandle) {
} }
}; };
let live_taken_over = proxy_service.detect_takeover_in_live_configs(); let live_taken_over = proxy_service.detect_takeover_in_live_configs();
let needs_restore = cleanup_before_exit_needed( let needs_restore = has_backups || live_taken_over;
has_backups,
live_taken_over,
crate::settings::pi_takeover_enabled(),
);
if needs_restore { if needs_restore {
log::info!("检测到接管残留,开始恢复 Live 配置(保留代理状态)..."); log::info!("检测到接管残留,开始恢复 Live 配置(保留代理状态)...");
@@ -1878,14 +1851,6 @@ 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()` /// `std::process::exit` 会绕过 Tauri 运行时,触发不了 `TrayIcon::drop()`
@@ -1916,10 +1881,7 @@ pub(crate) fn remove_tray_icon_before_exit(app_handle: &tauri::AppHandle) {
/// 则自动启动代理服务并接管对应应用的 Live 配置。 /// 则自动启动代理服务并接管对应应用的 Live 配置。
const PROXY_STARTUP_APP_TYPES: [&str; 4] = ["claude", "codex", "gemini", "grokbuild"]; const PROXY_STARTUP_APP_TYPES: [&str; 4] = ["claude", "codex", "gemini", "grokbuild"];
async fn enabled_proxy_apps_on_startup( async fn enabled_proxy_apps_on_startup(db: &database::Database) -> Vec<&'static str> {
db: &database::Database,
pi_takeover_enabled: bool,
) -> Vec<&'static str> {
let mut apps = Vec::new(); let mut apps = Vec::new();
for app_type in PROXY_STARTUP_APP_TYPES { for app_type in PROXY_STARTUP_APP_TYPES {
if db if db
@@ -1930,16 +1892,12 @@ async fn enabled_proxy_apps_on_startup(
apps.push(app_type); apps.push(app_type);
} }
} }
if pi_takeover_enabled {
apps.push("pi");
}
apps apps
} }
async fn restore_proxy_state_on_startup(state: &store::AppState) { async fn restore_proxy_state_on_startup(state: &store::AppState) {
// 收集需要恢复接管的应用列表(从 proxy_config.enabled 读取) // 收集需要恢复接管的应用列表(从 proxy_config.enabled 读取)
let apps_to_restore = let apps_to_restore = enabled_proxy_apps_on_startup(&state.db).await;
enabled_proxy_apps_on_startup(&state.db, crate::settings::pi_takeover_enabled()).await;
if apps_to_restore.is_empty() { if apps_to_restore.is_empty() {
log::debug!("启动时无需恢复代理状态"); log::debug!("启动时无需恢复代理状态");
@@ -1960,15 +1918,7 @@ async fn restore_proxy_state_on_startup(state: &store::AppState) {
} }
Err(e) => { Err(e) => {
log::error!("✗ 恢复 {app_type} 的代理接管状态失败: {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 if let Err(clear_err) = state
.proxy_service .proxy_service
.set_takeover_for_app(app_type, false) .set_takeover_for_app(app_type, false)
@@ -2038,7 +1988,6 @@ fn initialize_common_config_snippets(state: &store::AppState) {
.unwrap_or(true); .unwrap_or(true);
if should_run_legacy_migration { if should_run_legacy_migration {
let mut legacy_migration_succeeded = true;
for app_type in [ for app_type in [
crate::app_config::AppType::Claude, crate::app_config::AppType::Claude,
crate::app_config::AppType::Codex, crate::app_config::AppType::Codex,
@@ -2052,14 +2001,11 @@ fn initialize_common_config_snippets(state: &store::AppState) {
"✗ Failed to migrate legacy common-config usage for {}: {e}", "✗ Failed to migrate legacy common-config usage for {}: {e}",
app_type.as_str() app_type.as_str()
); );
legacy_migration_succeeded = false;
} }
} }
if legacy_migration_succeeded { if let Err(e) = state.db.set_legacy_common_config_migrated(true) {
if let Err(e) = state.db.set_legacy_common_config_migrated(true) { log::warn!("✗ Failed to persist legacy common-config migration flag: {e}");
log::warn!("✗ Failed to persist legacy common-config migration flag: {e}");
}
} }
} }
} }
@@ -2270,9 +2216,9 @@ pub fn restart_process(app_handle: &tauri::AppHandle) -> ! {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{ use super::{
classify_exit_request, cleanup_before_exit_needed, enabled_proxy_apps_on_startup, classify_exit_request, enabled_proxy_apps_on_startup, redact_url_for_log,
redact_url_for_log, redact_url_for_log_with_secrets, redact_url_origin_for_log, redact_url_for_log_with_secrets, redact_url_origin_for_log, runtime_log_level_allows,
runtime_log_level_allows, ExitRequestAction, ExitRequestAction,
}; };
use crate::database::Database; use crate::database::Database;
@@ -2394,21 +2340,8 @@ mod tests {
.await .await
.expect("enable Grok Build proxy config"); .expect("enable Grok Build proxy config");
let apps = enabled_proxy_apps_on_startup(&db, false).await; let apps = enabled_proxy_apps_on_startup(&db).await;
assert_eq!(apps, vec!["grokbuild"]); 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));
}
} }
-979
View File
@@ -1,979 +0,0 @@
//! 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<Value>,
pub input: Value,
pub cost: Value,
pub context_window: Value,
pub max_tokens: Value,
pub headers: BTreeMap<String, String>,
pub provider_headers: Vec<PiComposedHeader>,
pub model_headers: Vec<PiComposedHeader>,
pub compat: Option<Value>,
pub api_key: Option<String>,
pub oauth: Option<Value>,
pub auth_header: bool,
pub provider_extra: BTreeMap<String, Value>,
pub model_extra: BTreeMap<String, Value>,
pub override_extra: BTreeMap<String, Value>,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct PiNativeComposition {
pub status: PiComposerStatus,
pub provider_id: Option<String>,
pub provider_name: Option<String>,
pub provider_base_url: Option<String>,
pub models: Vec<PiComposedNativeModel>,
pub ignored_override_keys: Vec<String>,
pub reasons: Vec<PiComposerReason>,
}
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<String>) -> 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<String>) -> 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::<BTreeMap<_, _>>();
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<PiComposedNativeModel> = 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::<HashSet<_>>();
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<String, Value>) -> 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<PiComposedHeader> {
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<PiComposedHeader>,
overlay: impl IntoIterator<Item = PiComposedHeader>,
) {
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<String, Value>, recognized: &[&str]) -> BTreeMap<String, Value> {
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<ComposerOracleCase>,
fail_closed_cases: Vec<FailClosedCase>,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct ComposerOracleCase {
id: String,
provider_id: String,
input: Value,
execution: Execution,
#[serde(default)]
auth_execution: Option<Value>,
#[serde(default)]
expected: Option<Value>,
#[serde(default)]
expected_error: Option<String>,
}
#[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::<Vec<_>>(),
"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<_>>(),
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"}])
);
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
-218
View File
@@ -1,218 +0,0 @@
//! 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<String, JavaScriptSpreadValue>;
/// 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<Value>,
overlay: Option<Value>,
) -> Result<Option<Value>, 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<Map<String, Value>, 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
}
})))
);
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,834 +0,0 @@
#![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: <key>`;
//! - anthropic `sk-ant-oat...` → `authorization: Bearer <token>` +
//! `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"
);
}
-351
View File
@@ -1,351 +0,0 @@
//! 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<Mutex<()>> = 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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub default_model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub session_dir: Option<String>,
}
#[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<PiNativeDefaultsRollback, AppError> {
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<PathBuf, AppError> {
Ok(super::native::get_pi_agent_dir()?.join("settings.json"))
}
pub(crate) fn read_pi_native_defaults() -> Result<PiNativeDefaults, AppError> {
read_pi_native_defaults_at(&get_pi_settings_path()?)
}
pub(crate) fn read_pi_native_defaults_at(path: &Path) -> Result<PiNativeDefaults, AppError> {
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<PiNativeDefaultsReceipt, AppError> {
// Native model identifiers are opaque exact strings under pinned Pi's
// schema. Do not trim a schema-valid whitespace or edge-whitespace ID.
if provider_key.trim().is_empty() || model_id.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<String, Value>,
key: &str,
path: &Path,
) -> Result<Option<String>, 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<String, Value>) -> Result<(), AppError>,
) -> Result<PiNativeDefaultsReceipt, AppError> {
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<Value, AppError> {
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]
#[serial_test::serial]
fn native_default_preserves_a_pinned_exact_whitespace_model_id() {
struct EnvGuard(Option<std::ffi::OsString>);
impl Drop for EnvGuard {
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 temp = tempfile::tempdir().expect("tempdir");
let _home = EnvGuard(std::env::var_os("CC_SWITCH_TEST_HOME"));
std::env::set_var("CC_SWITCH_TEST_HOME", temp.path());
set_pi_native_default_with_receipt("provider", " ")
.expect("pinned exact model id is accepted");
let defaults = read_pi_native_defaults().expect("native defaults");
assert_eq!(defaults.default_provider.as_deref(), Some("provider"));
assert_eq!(defaults.default_model.as_deref(), Some(" "));
assert!(set_pi_native_default_with_receipt("provider", "").is_err());
}
#[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"
);
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -1,6 +1,6 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Prompt { pub struct Prompt {
pub id: String, pub id: String,
pub name: String, pub name: String,
-12
View File
@@ -26,7 +26,6 @@ pub fn prompt_file_path(app: &AppType) -> Result<PathBuf, AppError> {
AppType::OpenCode => get_opencode_dir(), AppType::OpenCode => get_opencode_dir(),
AppType::OpenClaw => get_openclaw_dir(), AppType::OpenClaw => get_openclaw_dir(),
AppType::Hermes => crate::hermes_config::get_hermes_dir(), AppType::Hermes => crate::hermes_config::get_hermes_dir(),
AppType::Pi => crate::pi_config::native::get_pi_agent_dir()?,
AppType::ClaudeDesktop => unreachable!("handled above"), AppType::ClaudeDesktop => unreachable!("handled above"),
}; };
@@ -36,7 +35,6 @@ pub fn prompt_file_path(app: &AppType) -> Result<PathBuf, AppError> {
AppType::Gemini => "GEMINI.md", AppType::Gemini => "GEMINI.md",
AppType::GrokBuild | AppType::OpenCode | AppType::OpenClaw => "AGENTS.md", AppType::GrokBuild | AppType::OpenCode | AppType::OpenClaw => "AGENTS.md",
AppType::Hermes => "SOUL.md", AppType::Hermes => "SOUL.md",
AppType::Pi => "AGENTS.md",
AppType::ClaudeDesktop => unreachable!("handled above"), AppType::ClaudeDesktop => unreachable!("handled above"),
}; };
@@ -56,16 +54,6 @@ mod tests {
Some("SOUL.md") Some("SOUL.md")
); );
} }
#[test]
fn pi_prompt_file_uses_agents_md() {
let path = prompt_file_path(&AppType::Pi).expect("Pi prompt path");
assert_eq!(
path.file_name().and_then(|name| name.to_str()),
Some("AGENTS.md")
);
}
} }
fn get_base_dir_with_fallback( fn get_base_dir_with_fallback(
-148
View File
@@ -43,84 +43,6 @@ pub struct Provider {
pub in_failover_queue: bool, pub in_failover_queue: bool,
} }
/// IPC/service input for creating or editing a provider.
///
/// This deliberately is not the hydrated [`Provider`] read projection. In
/// particular, callers cannot pass a DAO aggregate back into the provider-row
/// writer without first crossing the service boundary, where endpoint
/// ownership is checked.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderMutationInput {
pub id: String,
pub name: String,
#[serde(rename = "settingsConfig")]
pub settings_config: Value,
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "websiteUrl")]
pub website_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub category: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "createdAt")]
pub created_at: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "sortIndex")]
pub sort_index: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub notes: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub meta: Option<ProviderMeta>,
#[serde(skip_serializing_if = "Option::is_none")]
pub icon: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "iconColor")]
pub icon_color: Option<String>,
#[serde(default)]
#[serde(rename = "inFailoverQueue")]
pub in_failover_queue: bool,
}
impl From<ProviderMutationInput> for Provider {
fn from(input: ProviderMutationInput) -> Self {
Self {
id: input.id,
name: input.name,
settings_config: input.settings_config,
website_url: input.website_url,
category: input.category,
created_at: input.created_at,
sort_index: input.sort_index,
notes: input.notes,
meta: input.meta,
icon: input.icon,
icon_color: input.icon_color,
in_failover_queue: input.in_failover_queue,
}
}
}
/// A provider row and every endpoint owned by that row.
///
/// SQLite stores endpoints separately from provider metadata. This aggregate
/// is the only lossless DAO boundary; legacy `Provider` reads are projections
/// of it for API compatibility.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderAggregate {
pub provider: Provider,
#[serde(default)]
pub endpoints: IndexMap<String, crate::settings::CustomEndpoint>,
}
impl ProviderAggregate {
pub(crate) fn into_provider(mut self) -> Provider {
self.provider
.meta
.get_or_insert_with(ProviderMeta::default)
.custom_endpoints = self.endpoints.into_iter().collect();
self.provider
}
}
impl Provider { impl Provider {
/// 从现有ID创建供应商 /// 从现有ID创建供应商
pub fn with_id( pub fn with_id(
@@ -145,72 +67,6 @@ impl Provider {
} }
} }
pub(crate) fn row_content_fingerprint(&self) -> String {
use sha2::{Digest, Sha256};
fn hash_canonical(value: &serde_json::Value, hasher: &mut Sha256) {
match value {
serde_json::Value::Null => hasher.update(b"n"),
serde_json::Value::Bool(value) => {
hasher.update(b"b");
hasher.update([*value as u8]);
}
serde_json::Value::Number(value) => {
let text = value.to_string();
hasher.update(b"#");
hasher.update((text.len() as u64).to_le_bytes());
hasher.update(text.as_bytes());
}
serde_json::Value::String(value) => {
hasher.update(b"s");
hasher.update((value.len() as u64).to_le_bytes());
hasher.update(value.as_bytes());
}
serde_json::Value::Array(items) => {
hasher.update(b"[");
hasher.update((items.len() as u64).to_le_bytes());
for item in items {
hash_canonical(item, hasher);
}
hasher.update(b"]");
}
serde_json::Value::Object(map) => {
hasher.update(b"{");
hasher.update((map.len() as u64).to_le_bytes());
let mut keys: Vec<&String> = map.keys().collect();
keys.sort();
for key in keys {
hasher.update((key.len() as u64).to_le_bytes());
hasher.update(key.as_bytes());
hash_canonical(&map[key.as_str()], hasher);
}
hasher.update(b"}");
}
}
}
let mut meta = serde_json::to_value(&self.meta).unwrap_or(serde_json::Value::Null);
if let serde_json::Value::Object(map) = &mut meta {
map.remove("custom_endpoints");
map.remove("customEndpoints");
}
let mut hasher = Sha256::new();
for part in [
serde_json::Value::String(self.name.clone()),
self.settings_config.clone(),
serde_json::to_value(&self.website_url).unwrap_or(serde_json::Value::Null),
serde_json::to_value(&self.category).unwrap_or(serde_json::Value::Null),
serde_json::to_value(&self.notes).unwrap_or(serde_json::Value::Null),
serde_json::to_value(&self.icon).unwrap_or(serde_json::Value::Null),
serde_json::to_value(&self.icon_color).unwrap_or(serde_json::Value::Null),
meta,
] {
hash_canonical(&part, &mut hasher);
hasher.update([0u8]);
}
format!("{:x}", hasher.finalize())
}
pub fn is_codex_oauth(&self) -> bool { pub fn is_codex_oauth(&self) -> bool {
self.provider_type() == Some("codex_oauth") self.provider_type() == Some("codex_oauth")
} }
@@ -346,10 +202,6 @@ impl Provider {
str_at(settings.get("base_url")), str_at(settings.get("base_url")),
str_at(settings.get("api_key")), 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. // OpenClaw (openclaw.json) flattens credentials at the top level, camelCase.
AppType::OpenClaw => ( AppType::OpenClaw => (
str_at(settings.get("baseUrl")), str_at(settings.get("baseUrl")),
-7
View File
@@ -4,7 +4,6 @@
use crate::app_config::AppType; use crate::app_config::AppType;
use crate::proxy::usage::parser::TokenUsage; use crate::proxy::usage::parser::TokenUsage;
use crate::proxy::usage::InputTokenSemantics;
use serde_json::Value; use serde_json::Value;
/// 使用量解析器类型别名 /// 使用量解析器类型别名
@@ -32,8 +31,6 @@ pub struct UsageParserConfig {
pub model_extractor: StreamModelExtractor, pub model_extractor: StreamModelExtractor,
/// 流式 usage 事件预过滤器 /// 流式 usage 事件预过滤器
pub stream_event_filter: Option<StreamUsageEventFilter>, pub stream_event_filter: Option<StreamUsageEventFilter>,
/// Semantics of `TokenUsage::input_tokens` produced by these parsers.
pub input_token_semantics: InputTokenSemantics,
/// 应用类型字符串(用于日志记录) /// 应用类型字符串(用于日志记录)
pub app_type_str: &'static str, pub app_type_str: &'static str,
} }
@@ -144,7 +141,6 @@ pub const CLAUDE_PARSER_CONFIG: UsageParserConfig = UsageParserConfig {
response_parser: TokenUsage::from_claude_response, response_parser: TokenUsage::from_claude_response,
model_extractor: claude_model_extractor, model_extractor: claude_model_extractor,
stream_event_filter: Some(claude_stream_usage_event_filter), stream_event_filter: Some(claude_stream_usage_event_filter),
input_token_semantics: InputTokenSemantics::FreshExcludesCache,
app_type_str: "claude", app_type_str: "claude",
}; };
@@ -154,7 +150,6 @@ pub const OPENAI_PARSER_CONFIG: UsageParserConfig = UsageParserConfig {
response_parser: TokenUsage::from_openai_response, response_parser: TokenUsage::from_openai_response,
model_extractor: openai_model_extractor, model_extractor: openai_model_extractor,
stream_event_filter: Some(openai_stream_usage_event_filter), stream_event_filter: Some(openai_stream_usage_event_filter),
input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets,
app_type_str: "codex", app_type_str: "codex",
}; };
@@ -164,7 +159,6 @@ pub const CODEX_PARSER_CONFIG: UsageParserConfig = UsageParserConfig {
response_parser: TokenUsage::from_codex_response_auto, response_parser: TokenUsage::from_codex_response_auto,
model_extractor: codex_auto_model_extractor, model_extractor: codex_auto_model_extractor,
stream_event_filter: Some(codex_stream_usage_event_filter), stream_event_filter: Some(codex_stream_usage_event_filter),
input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets,
app_type_str: "codex", app_type_str: "codex",
}; };
@@ -174,7 +168,6 @@ pub const GEMINI_PARSER_CONFIG: UsageParserConfig = UsageParserConfig {
response_parser: TokenUsage::from_gemini_response, response_parser: TokenUsage::from_gemini_response,
model_extractor: gemini_model_extractor, model_extractor: gemini_model_extractor,
stream_event_filter: Some(gemini_stream_usage_event_filter), stream_event_filter: Some(gemini_stream_usage_event_filter),
input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets,
app_type_str: "gemini", app_type_str: "gemini",
}; };
+1 -11
View File
@@ -39,7 +39,7 @@ use super::{
server::ProxyState, server::ProxyState,
sse::{strip_sse_field, take_sse_block}, sse::{strip_sse_field, take_sse_block},
types::*, types::*,
usage::{parser::TokenUsage, InputTokenSemantics}, usage::parser::TokenUsage,
ProxyError, ProxyError,
}; };
use crate::app_config::AppType; use crate::app_config::AppType;
@@ -338,7 +338,6 @@ async fn write_claude_usage_log(state: &ProxyState, log: ClaudeUsageLog) {
&log.model, &log.model,
&log.request_model, &log.request_model,
&log.outbound_model, &log.outbound_model,
InputTokenSemantics::FreshExcludesCache,
log.usage, log.usage,
log.latency_ms, log.latency_ms,
None, None,
@@ -466,7 +465,6 @@ async fn handle_claude_transform(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
InputTokenSemantics::FreshExcludesCache,
usage, usage,
latency_ms, latency_ms,
first_token_ms, first_token_ms,
@@ -1135,7 +1133,6 @@ async fn handle_codex_responses_namespace_restore(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
InputTokenSemantics::TotalIncludesCacheBuckets,
usage, usage,
latency_ms, latency_ms,
None, None,
@@ -1248,7 +1245,6 @@ async fn handle_codex_chat_to_responses_transform(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
InputTokenSemantics::TotalIncludesCacheBuckets,
usage, usage,
latency_ms, latency_ms,
first_token_ms, first_token_ms,
@@ -1370,7 +1366,6 @@ async fn handle_codex_chat_to_responses_transform(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
InputTokenSemantics::TotalIncludesCacheBuckets,
usage, usage,
latency_ms, latency_ms,
None, None,
@@ -1536,7 +1531,6 @@ async fn handle_codex_anthropic_to_responses_transform(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
InputTokenSemantics::TotalIncludesCacheBuckets,
usage, usage,
latency_ms, latency_ms,
None, None,
@@ -1624,7 +1618,6 @@ fn build_codex_anthropic_sse_response(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
InputTokenSemantics::TotalIncludesCacheBuckets,
usage, usage,
latency_ms, latency_ms,
first_token_ms, first_token_ms,
@@ -2597,7 +2590,6 @@ fn log_forward_error(
is_streaming, is_streaming,
Some(ctx.session_id.clone()), Some(ctx.session_id.clone()),
None, None,
InputTokenSemantics::FreshExcludesCache,
) { ) {
log::warn!("记录失败请求日志失败: {e}"); log::warn!("记录失败请求日志失败: {e}");
} }
@@ -2615,7 +2607,6 @@ async fn log_usage(
model: &str, model: &str,
request_model: &str, request_model: &str,
outbound_model: &str, outbound_model: &str,
input_token_semantics: InputTokenSemantics,
usage: TokenUsage, usage: TokenUsage,
latency_ms: u64, latency_ms: u64,
first_token_ms: Option<u64>, first_token_ms: Option<u64>,
@@ -2649,7 +2640,6 @@ async fn log_usage(
model.to_string(), model.to_string(),
request_model.to_string(), request_model.to_string(),
pricing_model.to_string(), pricing_model.to_string(),
input_token_semantics,
usage, usage,
multiplier, multiplier,
latency_ms, latency_ms,
+45 -109
View File
@@ -10,15 +10,8 @@ use std::net::IpAddr;
use std::sync::RwLock; use std::sync::RwLock;
use std::time::Duration; use std::time::Duration;
#[derive(Clone)] /// 全局 HTTP 客户端实例
struct GlobalClients { static GLOBAL_CLIENT: OnceCell<RwLock<Client>> = OnceCell::new();
standard: Client,
no_redirect: Client,
}
/// 全局 HTTP 客户端实例。Pi 网关使用同一代理配置下的 no-redirect 客户端,
/// 防止 307/308 把凭证、自定义头和请求体重放到另一个 origin。
static GLOBAL_CLIENTS: OnceCell<RwLock<GlobalClients>> = OnceCell::new();
/// 当前代理 URL(用于日志和状态查询) /// 当前代理 URL(用于日志和状态查询)
static CURRENT_PROXY_URL: OnceCell<RwLock<Option<String>>> = OnceCell::new(); static CURRENT_PROXY_URL: OnceCell<RwLock<Option<String>>> = OnceCell::new();
@@ -59,10 +52,10 @@ fn get_proxy_port() -> u16 {
/// 传入 None 或空字符串表示直连 /// 传入 None 或空字符串表示直连
pub fn init(proxy_url: Option<&str>) -> Result<(), String> { pub fn init(proxy_url: Option<&str>) -> Result<(), String> {
let effective_url = proxy_url.filter(|s| !s.trim().is_empty()); let effective_url = proxy_url.filter(|s| !s.trim().is_empty());
let clients = build_clients(effective_url)?; let client = build_client(effective_url)?;
// 尝试初始化全局客户端,如果已存在则记录警告并使用 apply_proxy 更新 // 尝试初始化全局客户端,如果已存在则记录警告并使用 apply_proxy 更新
if GLOBAL_CLIENTS.set(RwLock::new(clients)).is_err() { if GLOBAL_CLIENT.set(RwLock::new(client.clone())).is_err() {
log::warn!( log::warn!(
"[GlobalProxy] [GP-003] Already initialized, updating instead: {}", "[GlobalProxy] [GP-003] Already initialized, updating instead: {}",
effective_url effective_url
@@ -98,8 +91,8 @@ pub fn init(proxy_url: Option<&str>) -> Result<(), String> {
/// 验证成功返回 Ok(()),失败返回错误信息 /// 验证成功返回 Ok(()),失败返回错误信息
pub fn validate_proxy(proxy_url: Option<&str>) -> Result<(), String> { pub fn validate_proxy(proxy_url: Option<&str>) -> Result<(), String> {
let effective_url = proxy_url.filter(|s| !s.trim().is_empty()); let effective_url = proxy_url.filter(|s| !s.trim().is_empty());
// 同时验证标准与 no-redirect 客户端,保证应用配置时不会只更新一半。 // 只调用 build_client 来验证,但不应用
build_clients(effective_url)?; build_client(effective_url)?;
Ok(()) Ok(())
} }
@@ -112,15 +105,15 @@ pub fn validate_proxy(proxy_url: Option<&str>) -> Result<(), String> {
/// * `proxy_url` - 代理 URLNone 或空字符串表示直连 /// * `proxy_url` - 代理 URLNone 或空字符串表示直连
pub fn apply_proxy(proxy_url: Option<&str>) -> Result<(), String> { pub fn apply_proxy(proxy_url: Option<&str>) -> Result<(), String> {
let effective_url = proxy_url.filter(|s| !s.trim().is_empty()); let effective_url = proxy_url.filter(|s| !s.trim().is_empty());
let new_clients = build_clients(effective_url)?; let new_client = build_client(effective_url)?;
// 更新客户端 // 更新客户端
if let Some(lock) = GLOBAL_CLIENTS.get() { if let Some(lock) = GLOBAL_CLIENT.get() {
let mut clients = lock.write().map_err(|e| { let mut client = lock.write().map_err(|e| {
log::error!("[GlobalProxy] [GP-001] Failed to acquire write lock: {e}"); log::error!("[GlobalProxy] [GP-001] Failed to acquire write lock: {e}");
"Failed to update proxy: lock poisoned".to_string() "Failed to update proxy: lock poisoned".to_string()
})?; })?;
*clients = new_clients; *client = new_client;
} else { } else {
// 如果还没初始化,则初始化 // 如果还没初始化,则初始化
return init(proxy_url); return init(proxy_url);
@@ -155,42 +148,54 @@ pub fn apply_proxy(proxy_url: Option<&str>) -> Result<(), String> {
/// * `proxy_url` - 新的代理 URLNone 或空字符串表示直连 /// * `proxy_url` - 新的代理 URLNone 或空字符串表示直连
#[allow(dead_code)] #[allow(dead_code)]
pub fn update_proxy(proxy_url: Option<&str>) -> Result<(), String> { pub fn update_proxy(proxy_url: Option<&str>) -> Result<(), String> {
apply_proxy(proxy_url) let effective_url = proxy_url.filter(|s| !s.trim().is_empty());
let new_client = build_client(effective_url)?;
// 更新客户端
if let Some(lock) = GLOBAL_CLIENT.get() {
let mut client = lock.write().map_err(|e| {
log::error!("[GlobalProxy] [GP-001] Failed to acquire write lock: {e}");
"Failed to update proxy: lock poisoned".to_string()
})?;
*client = new_client;
} else {
// 如果还没初始化,则初始化
return init(proxy_url);
}
// 更新代理 URL 记录
if let Some(lock) = CURRENT_PROXY_URL.get() {
let mut url = lock.write().map_err(|e| {
log::error!("[GlobalProxy] [GP-002] Failed to acquire URL write lock: {e}");
"Failed to update proxy URL record: lock poisoned".to_string()
})?;
*url = effective_url.map(|s| s.to_string());
}
log::info!(
"[GlobalProxy] Updated: {}",
effective_url
.map(mask_url)
.unwrap_or_else(|| "direct connection".to_string())
);
Ok(())
} }
/// 获取全局 HTTP 客户端 /// 获取全局 HTTP 客户端
/// ///
/// 返回配置了代理的客户端(如果已配置代理),否则返回跟随系统代理的客户端。 /// 返回配置了代理的客户端(如果已配置代理),否则返回跟随系统代理的客户端。
pub fn get() -> Client { pub fn get() -> Client {
GLOBAL_CLIENTS GLOBAL_CLIENT
.get() .get()
.and_then(|lock| lock.read().ok()) .and_then(|lock| lock.read().ok())
.map(|clients| clients.standard.clone()) .map(|c| c.clone())
.unwrap_or_else(|| { .unwrap_or_else(|| {
log::warn!("[GlobalProxy] [GP-004] Client not initialized, using fallback"); log::warn!("[GlobalProxy] [GP-004] Client not initialized, using fallback");
build_client(None).unwrap_or_default() build_client(None).unwrap_or_default()
}) })
} }
/// 获取禁用自动重定向的全局客户端。
///
/// 与 `get` 不同,这个安全边界在锁毒化或构建失败时 fail closed;调用方不得
/// 回退到会自动跟随跨 origin 重定向的默认客户端。
pub(crate) fn get_no_redirect() -> Result<Client, String> {
if let Some(lock) = GLOBAL_CLIENTS.get() {
return lock
.read()
.map(|clients| clients.no_redirect.clone())
.map_err(|error| {
log::error!("[GlobalProxy] [GP-005] Failed to acquire read lock: {error}");
"Failed to read no-redirect HTTP client: lock poisoned".to_string()
});
}
log::warn!("[GlobalProxy] [GP-006] Client not initialized, using no-redirect fallback");
build_no_redirect_client(None)
}
/// 获取当前代理 URL /// 获取当前代理 URL
/// ///
/// 返回当前配置的代理 URL,None 表示直连。 /// 返回当前配置的代理 URL,None 表示直连。
@@ -209,24 +214,6 @@ pub fn is_proxy_enabled() -> bool {
/// 构建 HTTP 客户端 /// 构建 HTTP 客户端
fn build_client(proxy_url: Option<&str>) -> Result<Client, String> { fn build_client(proxy_url: Option<&str>) -> Result<Client, String> {
build_client_with_redirect_policy(proxy_url, true)
}
fn build_no_redirect_client(proxy_url: Option<&str>) -> Result<Client, String> {
build_client_with_redirect_policy(proxy_url, false)
}
fn build_clients(proxy_url: Option<&str>) -> Result<GlobalClients, String> {
Ok(GlobalClients {
standard: build_client(proxy_url)?,
no_redirect: build_no_redirect_client(proxy_url)?,
})
}
fn build_client_with_redirect_policy(
proxy_url: Option<&str>,
follow_redirects: bool,
) -> Result<Client, String> {
let mut builder = Client::builder() let mut builder = Client::builder()
.timeout(Duration::from_secs(600)) .timeout(Duration::from_secs(600))
.connect_timeout(Duration::from_secs(30)) .connect_timeout(Duration::from_secs(30))
@@ -238,9 +225,6 @@ fn build_client_with_redirect_policy(
.no_brotli() .no_brotli()
.no_deflate() .no_deflate()
.no_zstd(); .no_zstd();
if !follow_redirects {
builder = builder.redirect(reqwest::redirect::Policy::none());
}
// 有代理地址则使用代理,否则跟随系统代理 // 有代理地址则使用代理,否则跟随系统代理
if let Some(url) = proxy_url { if let Some(url) = proxy_url {
@@ -353,9 +337,7 @@ pub fn mask_url(url: &str) -> String {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Mutex, OnceLock};
use std::sync::{Arc, Mutex, OnceLock};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
fn env_lock() -> &'static Mutex<()> { fn env_lock() -> &'static Mutex<()> {
static LOCK: OnceLock<Mutex<()>> = OnceLock::new(); static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
@@ -410,52 +392,6 @@ mod tests {
assert!(result.is_err(), "Should reject invalid proxy scheme"); assert!(result.is_err(), "Should reject invalid proxy scheme");
} }
#[tokio::test]
async fn no_redirect_client_exposes_redirect_without_replaying_request() {
let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0))
.await
.expect("bind redirect server");
let address = listener.local_addr().expect("server address");
let hits = Arc::new(AtomicUsize::new(0));
let server_hits = Arc::clone(&hits);
let server = tokio::spawn(async move {
loop {
let (mut stream, _) = listener.accept().await.expect("accept request");
let hit = server_hits.fetch_add(1, Ordering::SeqCst);
let mut request = [0_u8; 2048];
let _ = stream.read(&mut request).await.expect("read request");
let response = if hit == 0 {
format!(
"HTTP/1.1 307 Temporary Redirect\r\n\
Location: http://{address}/redirected\r\n\
Content-Length: 0\r\nConnection: close\r\n\r\n"
)
} else {
"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n".to_string()
};
stream
.write_all(response.as_bytes())
.await
.expect("write response");
}
});
let client = build_no_redirect_client(None).expect("build no-redirect client");
let response = client
.get(format!("http://{address}/initial"))
.send()
.await
.expect("send request");
assert_eq!(response.status(), reqwest::StatusCode::TEMPORARY_REDIRECT);
tokio::time::sleep(Duration::from_millis(25)).await;
assert_eq!(
hits.load(Ordering::SeqCst),
1,
"the redirected endpoint must not receive a replay"
);
server.abort();
}
#[test] #[test]
fn test_proxy_points_to_loopback() { fn test_proxy_points_to_loopback() {
// 设置 CC Switch 代理端口为 15721(默认值) // 设置 CC Switch 代理端口为 15721(默认值)
-2
View File
@@ -21,8 +21,6 @@ pub(crate) mod json_canonical;
pub mod log_codes; pub mod log_codes;
pub mod media_sanitizer; pub mod media_sanitizer;
pub mod model_mapper; pub mod model_mapper;
pub(crate) mod pi_handler;
pub(crate) mod pi_runtime;
pub mod provider_router; pub mod provider_router;
pub mod providers; pub mod providers;
pub mod response_processor; pub mod response_processor;
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+12 -31
View File
@@ -29,16 +29,6 @@ impl ProviderRouter {
} }
} }
async fn app_proxy_config(
&self,
app_type: &str,
) -> Result<crate::proxy::types::AppProxyConfig, AppError> {
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
}
/// 选择可用的供应商(支持故障转移) /// 选择可用的供应商(支持故障转移)
/// ///
/// 返回按优先级排序的可用供应商列表: /// 返回按优先级排序的可用供应商列表:
@@ -50,7 +40,7 @@ impl ProviderRouter {
let mut circuit_open_count = 0usize; let mut circuit_open_count = 0usize;
// 检查该应用的自动故障转移开关是否开启(从 proxy_config 表读取) // 检查该应用的自动故障转移开关是否开启(从 proxy_config 表读取)
let auto_failover_enabled = match self.app_proxy_config(app_type).await { let auto_failover_enabled = match self.db.get_proxy_config_for_app(app_type).await {
Ok(config) => config.auto_failover_enabled, Ok(config) => config.auto_failover_enabled,
Err(e) => { Err(e) => {
log::error!("[{app_type}] 读取 proxy_config 失败: {e},默认禁用故障转移"); log::error!("[{app_type}] 读取 proxy_config 失败: {e},默认禁用故障转移");
@@ -142,7 +132,7 @@ impl ProviderRouter {
error_msg: Option<String>, error_msg: Option<String>,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
// 1. 按应用独立获取熔断器配置 // 1. 按应用独立获取熔断器配置
let failure_threshold = match self.app_proxy_config(app_type).await { let failure_threshold = match self.db.get_proxy_config_for_app(app_type).await {
Ok(app_config) => app_config.circuit_failure_threshold, Ok(app_config) => app_config.circuit_failure_threshold,
Err(_) => 5, // 默认值 Err(_) => 5, // 默认值
}; };
@@ -261,7 +251,7 @@ impl ProviderRouter {
let app_type = key.split(':').next().unwrap_or("claude"); let app_type = key.split(':').next().unwrap_or("claude");
// 按应用独立读取熔断器配置 // 按应用独立读取熔断器配置
let config = match self.app_proxy_config(app_type).await { let config = match self.db.get_proxy_config_for_app(app_type).await {
Ok(app_config) => crate::proxy::circuit_breaker::CircuitBreakerConfig { Ok(app_config) => crate::proxy::circuit_breaker::CircuitBreakerConfig {
failure_threshold: app_config.circuit_failure_threshold, failure_threshold: app_config.circuit_failure_threshold,
success_threshold: app_config.circuit_success_threshold, success_threshold: app_config.circuit_success_threshold,
@@ -358,10 +348,8 @@ mod tests {
let provider_b = let provider_b =
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
db.reconcile_provider_fixture("claude", &provider_a) db.save_provider("claude", &provider_a).unwrap();
.unwrap(); db.save_provider("claude", &provider_b).unwrap();
db.reconcile_provider_fixture("claude", &provider_b)
.unwrap();
db.set_current_provider("claude", "a").unwrap(); db.set_current_provider("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap(); db.add_to_failover_queue("claude", "b").unwrap();
@@ -386,10 +374,8 @@ mod tests {
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
provider_b.sort_index = Some(1); provider_b.sort_index = Some(1);
db.reconcile_provider_fixture("claude", &provider_a) db.save_provider("claude", &provider_a).unwrap();
.unwrap(); db.save_provider("claude", &provider_b).unwrap();
db.reconcile_provider_fixture("claude", &provider_b)
.unwrap();
db.set_current_provider("claude", "a").unwrap(); db.set_current_provider("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap(); db.add_to_failover_queue("claude", "b").unwrap();
@@ -421,10 +407,8 @@ mod tests {
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
provider_b.sort_index = Some(1); provider_b.sort_index = Some(1);
db.reconcile_provider_fixture("claude", &provider_a) db.save_provider("claude", &provider_a).unwrap();
.unwrap(); db.save_provider("claude", &provider_b).unwrap();
db.reconcile_provider_fixture("claude", &provider_b)
.unwrap();
db.set_current_provider("claude", "a").unwrap(); db.set_current_provider("claude", "a").unwrap();
// 只把 b 加入故障转移队列(模拟“当前供应商不在队列里”的常见配置) // 只把 b 加入故障转移队列(模拟“当前供应商不在队列里”的常见配置)
@@ -460,10 +444,8 @@ mod tests {
let provider_b = let provider_b =
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
db.reconcile_provider_fixture("claude", &provider_a) db.save_provider("claude", &provider_a).unwrap();
.unwrap(); db.save_provider("claude", &provider_b).unwrap();
db.reconcile_provider_fixture("claude", &provider_b)
.unwrap();
db.add_to_failover_queue("claude", "a").unwrap(); db.add_to_failover_queue("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap(); db.add_to_failover_queue("claude", "b").unwrap();
@@ -503,8 +485,7 @@ mod tests {
let provider_a = let provider_a =
Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None); Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None);
db.reconcile_provider_fixture("claude", &provider_a) db.save_provider("claude", &provider_a).unwrap();
.unwrap();
db.add_to_failover_queue("claude", "a").unwrap(); db.add_to_failover_queue("claude", "a").unwrap();
// 启用自动故障转移 // 启用自动故障转移
+2 -10
View File
@@ -205,11 +205,7 @@ impl ProviderType {
ProviderType::Gemini ProviderType::Gemini
} }
AppType::GrokBuild => ProviderType::Codex, AppType::GrokBuild => ProviderType::Codex,
AppType::OpenCode | AppType::OpenClaw | AppType::Hermes | AppType::Pi => { AppType::OpenCode | AppType::OpenClaw | AppType::Hermes => ProviderType::Codex,
// Generic callers cannot infer Pi's wire family from AppType;
// the dedicated Pi runtime routes by effective model API.
ProviderType::Codex
}
} }
} }
@@ -263,11 +259,7 @@ pub fn get_adapter(app_type: &AppType) -> Box<dyn ProviderAdapter> {
AppType::Codex => Box::new(CodexAdapter::new()), AppType::Codex => Box::new(CodexAdapter::new()),
AppType::Gemini => Box::new(GeminiAdapter::new()), AppType::Gemini => Box::new(GeminiAdapter::new()),
AppType::GrokBuild => Box::new(CodexAdapter::new()), AppType::GrokBuild => Box::new(CodexAdapter::new()),
AppType::OpenCode | AppType::OpenClaw | AppType::Hermes | AppType::Pi => { AppType::OpenCode | AppType::OpenClaw | AppType::Hermes => Box::new(CodexAdapter::new()),
// Pi requests use the dedicated per-model adapter path. Keep the
// generic fallback deterministic for non-routing utilities.
Box::new(CodexAdapter::new())
}
} }
} }
+5 -16
View File
@@ -254,11 +254,11 @@ pub async fn handle_non_streaming(
spawn_log_usage( spawn_log_usage(
state, state,
ctx, ctx,
parser_config.input_token_semantics,
usage, usage,
&model, &model,
&ctx.request_model, &ctx.request_model,
status.as_u16(), status.as_u16(),
false,
); );
} else { } else {
let model = json_value let model = json_value
@@ -271,11 +271,11 @@ pub async fn handle_non_streaming(
spawn_log_usage( spawn_log_usage(
state, state,
ctx, ctx,
parser_config.input_token_semantics,
TokenUsage::default(), TokenUsage::default(),
&model, &model,
&ctx.request_model, &ctx.request_model,
status.as_u16(), status.as_u16(),
false,
); );
log::debug!( log::debug!(
"[{}] 未能解析 usage 信息,跳过记录", "[{}] 未能解析 usage 信息,跳过记录",
@@ -291,11 +291,11 @@ pub async fn handle_non_streaming(
spawn_log_usage( spawn_log_usage(
state, state,
ctx, ctx,
parser_config.input_token_semantics,
TokenUsage::default(), TokenUsage::default(),
ctx.outbound_model.as_deref().unwrap_or(&ctx.request_model), ctx.outbound_model.as_deref().unwrap_or(&ctx.request_model),
&ctx.request_model, &ctx.request_model,
status.as_u16(), status.as_u16(),
false,
); );
} }
} else { } else {
@@ -488,7 +488,6 @@ pub(crate) fn create_usage_collector(
let start_time = ctx.start_time; let start_time = ctx.start_time;
let stream_parser = parser_config.stream_parser; let stream_parser = parser_config.stream_parser;
let model_extractor = parser_config.model_extractor; let model_extractor = parser_config.model_extractor;
let input_token_semantics = parser_config.input_token_semantics;
let session_id = ctx.session_id.clone(); let session_id = ctx.session_id.clone();
Some(SseUsageCollector::new( Some(SseUsageCollector::new(
@@ -513,7 +512,6 @@ pub(crate) fn create_usage_collector(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
input_token_semantics,
usage, usage,
latency_ms, latency_ms,
first_token_ms, first_token_ms,
@@ -540,7 +538,6 @@ pub(crate) fn create_usage_collector(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
input_token_semantics,
TokenUsage::default(), TokenUsage::default(),
latency_ms, latency_ms,
first_token_ms, first_token_ms,
@@ -560,11 +557,11 @@ pub(crate) fn create_usage_collector(
fn spawn_log_usage( fn spawn_log_usage(
state: &ProxyState, state: &ProxyState,
ctx: &RequestContext, ctx: &RequestContext,
input_token_semantics: super::usage::InputTokenSemantics,
usage: TokenUsage, usage: TokenUsage,
model: &str, model: &str,
request_model: &str, request_model: &str,
status_code: u16, status_code: u16,
is_streaming: bool,
) { ) {
// Check enable_logging before spawning the log task // Check enable_logging before spawning the log task
if let Ok(config) = state.config.try_read() { if let Ok(config) = state.config.try_read() {
@@ -594,11 +591,10 @@ fn spawn_log_usage(
&model, &model,
&request_model, &request_model,
&outbound_model, &outbound_model,
input_token_semantics,
usage, usage,
latency_ms, latency_ms,
None, None,
false, is_streaming,
status_code, status_code,
Some(session_id), Some(session_id),
) )
@@ -628,7 +624,6 @@ async fn log_usage_internal(
model: &str, model: &str,
request_model: &str, request_model: &str,
outbound_model: &str, outbound_model: &str,
input_token_semantics: super::usage::InputTokenSemantics,
usage: TokenUsage, usage: TokenUsage,
latency_ms: u64, latency_ms: u64,
first_token_ms: Option<u64>, first_token_ms: Option<u64>,
@@ -666,7 +661,6 @@ async fn log_usage_internal(
model.to_string(), model.to_string(),
request_model.to_string(), request_model.to_string(),
pricing_model.to_string(), pricing_model.to_string(),
input_token_semantics,
usage, usage,
multiplier, multiplier,
latency_ms, latency_ms,
@@ -1007,8 +1001,6 @@ mod tests {
codex_chat_history: Arc::new(CodexChatHistoryStore::default()), codex_chat_history: Arc::new(CodexChatHistoryStore::default()),
app_handle: None, app_handle: None,
failover_manager: Arc::new(FailoverSwitchManager::new(db)), failover_manager: Arc::new(FailoverSwitchManager::new(db)),
pi_runtime: Arc::new(crate::proxy::pi_runtime::PiRuntimeStore::default()),
pi_server_generation: 0,
} }
} }
@@ -1080,7 +1072,6 @@ mod tests {
"resp-model", "resp-model",
"req-model", "req-model",
"req-model", "req-model",
crate::proxy::usage::InputTokenSemantics::FreshExcludesCache,
usage, usage,
10, 10,
None, None,
@@ -1151,7 +1142,6 @@ mod tests {
"resp-model", "resp-model",
"req-model", "req-model",
"outbound-model", "outbound-model",
crate::proxy::usage::InputTokenSemantics::FreshExcludesCache,
usage, usage,
10, 10,
None, None,
@@ -1232,7 +1222,6 @@ mod tests {
"resp-model", "resp-model",
"req-model", "req-model",
"req-model", "req-model",
crate::proxy::usage::InputTokenSemantics::FreshExcludesCache,
usage, usage,
10, 10,
None, None,
-22
View File
@@ -12,7 +12,6 @@ use super::{
failover_switch::FailoverSwitchManager, failover_switch::FailoverSwitchManager,
handlers, handlers,
log_codes::srv as log_srv, log_codes::srv as log_srv,
pi_runtime::PiRuntimeStore,
provider_router::ProviderRouter, provider_router::ProviderRouter,
providers::{codex_chat_history::CodexChatHistoryStore, gemini_shadow::GeminiShadowStore}, providers::{codex_chat_history::CodexChatHistoryStore, gemini_shadow::GeminiShadowStore},
types::*, types::*,
@@ -49,11 +48,6 @@ pub struct ProxyState {
pub app_handle: Option<tauri::AppHandle>, pub app_handle: Option<tauri::AppHandle>,
/// 故障转移切换管理器 /// 故障转移切换管理器
pub failover_manager: Arc<FailoverSwitchManager>, pub failover_manager: Arc<FailoverSwitchManager>,
/// Immutable Pi catalog publication point shared with `ProxyService`.
pub pi_runtime: Arc<PiRuntimeStore>,
/// Listener instance identity. A runtime built for an older listener can
/// never admit requests through this state.
pub pi_server_generation: u64,
} }
/// 代理HTTP服务器 /// 代理HTTP服务器
@@ -63,7 +57,6 @@ pub struct ProxyServer {
shutdown_tx: Arc<RwLock<Option<oneshot::Sender<()>>>>, shutdown_tx: Arc<RwLock<Option<oneshot::Sender<()>>>>,
/// 服务器任务句柄,用于等待服务器实际关闭 /// 服务器任务句柄,用于等待服务器实际关闭
server_handle: Arc<RwLock<Option<JoinHandle<()>>>>, server_handle: Arc<RwLock<Option<JoinHandle<()>>>>,
pi_server_generation: u64,
} }
impl ProxyServer { impl ProxyServer {
@@ -71,8 +64,6 @@ impl ProxyServer {
config: ProxyConfig, config: ProxyConfig,
db: Arc<Database>, db: Arc<Database>,
app_handle: Option<tauri::AppHandle>, app_handle: Option<tauri::AppHandle>,
pi_runtime: Arc<PiRuntimeStore>,
pi_server_generation: u64,
) -> Self { ) -> Self {
// 创建共享的 ProviderRouter(熔断器状态将跨所有请求保持) // 创建共享的 ProviderRouter(熔断器状态将跨所有请求保持)
let provider_router = Arc::new(ProviderRouter::new(db.clone())); let provider_router = Arc::new(ProviderRouter::new(db.clone()));
@@ -90,8 +81,6 @@ impl ProxyServer {
codex_chat_history: Arc::new(CodexChatHistoryStore::default()), codex_chat_history: Arc::new(CodexChatHistoryStore::default()),
app_handle, app_handle,
failover_manager, failover_manager,
pi_runtime,
pi_server_generation,
}; };
Self { Self {
@@ -99,14 +88,9 @@ impl ProxyServer {
state, state,
shutdown_tx: Arc::new(RwLock::new(None)), shutdown_tx: Arc::new(RwLock::new(None)),
server_handle: 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<ProxyServerInfo, ProxyError> { pub async fn start(&self) -> Result<ProxyServerInfo, ProxyError> {
// 检查是否已在运行 // 检查是否已在运行
if self.shutdown_tx.read().await.is_some() { if self.shutdown_tx.read().await.is_some() {
@@ -380,12 +364,6 @@ impl ProxyServer {
.route("/gemini/v1beta/*path", any(handlers::handle_gemini)) .route("/gemini/v1beta/*path", any(handlers::handle_gemini))
// Gemini 的 GA 版本也叫 /v1,给原 SDK 留一条出口 // Gemini 的 GA 版本也叫 /v1,给原 SDK 留一条出口
.route("/gemini/v1/*path", any(handlers::handle_gemini)) .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 // 提高默认请求体大小限制(避免 413 Payload Too Large
.layer(DefaultBodyLimit::max(200 * 1024 * 1024)) .layer(DefaultBodyLimit::max(200 * 1024 * 1024))
.with_state(self.state.clone()) .with_state(self.state.clone())
+3 -20
View File
@@ -4,32 +4,15 @@
//! 防止并发切换导致 is_current 与 Live 备份不一致。 //! 防止并发切换导致 is_current 与 Live 备份不一致。
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::{Arc, OnceLock}; use std::sync::Arc;
use tokio::sync::{Mutex, OwnedMutexGuard, RwLock}; use tokio::sync::{Mutex, OwnedMutexGuard, RwLock};
type PerAppLocks = Arc<RwLock<HashMap<String, Arc<Mutex<()>>>>>;
/// 每个应用类型一把互斥锁,保证同一应用的切换操作串行执行。 /// 每个应用类型一把互斥锁,保证同一应用的切换操作串行执行。
/// ///
/// 不同应用之间(如 Claude 和 Codex)可以并行切换。 /// 不同应用之间(如 Claude 和 Codex)可以并行切换。
#[derive(Clone)] #[derive(Clone, Default)]
pub struct SwitchLockManager { pub struct SwitchLockManager {
locks: PerAppLocks, locks: Arc<RwLock<HashMap<String, Arc<Mutex<()>>>>>,
}
impl Default for SwitchLockManager {
fn default() -> Self {
// Some commands construct a short-lived AppState around the shared
// database before running a blocking sync. A per-ProxyService map
// would give those paths a different lock and defeat serialization
// with provider rename/switch operations in the primary AppState.
static LOCKS: OnceLock<PerAppLocks> = OnceLock::new();
Self {
locks: LOCKS
.get_or_init(|| Arc::new(RwLock::new(HashMap::new())))
.clone(),
}
}
} }
impl SwitchLockManager { impl SwitchLockManager {
-11
View File
@@ -116,17 +116,6 @@ pub struct ProxyTakeoverStatus {
pub grokbuild: bool, pub grokbuild: bool,
pub opencode: bool, pub opencode: bool,
pub openclaw: 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健康状态 /// Provider健康状态
+20 -32
View File
@@ -3,7 +3,6 @@
//! 使用高精度 Decimal 类型避免浮点数精度问题 //! 使用高精度 Decimal 类型避免浮点数精度问题
use super::parser::TokenUsage; use super::parser::TokenUsage;
use super::semantics::InputTokenSemantics;
use rust_decimal::Decimal; use rust_decimal::Decimal;
use std::str::FromStr; use std::str::FromStr;
@@ -47,17 +46,13 @@ impl CostCalculator {
pricing: &ModelPricing, pricing: &ModelPricing,
cost_multiplier: Decimal, cost_multiplier: Decimal,
) -> CostBreakdown { ) -> CostBreakdown {
Self::calculate_with_input_semantics( Self::calculate_with_cache_semantics(usage, pricing, cost_multiplier, false)
InputTokenSemantics::FreshExcludesCache,
usage,
pricing,
cost_multiplier,
)
} }
/// Compatibility helper for existing callers. Live request paths use /// 按 app_type 选择输入 token 语义后计算成本。
/// [`Self::calculate_with_input_semantics`] so product app ownership never ///
/// stands in for the actual response parser/wire family. /// Codex/OpenAI Responses 与 Gemini 的输入 token 字段包含 cache read 部分;
/// Claude/Anthropic 的 input_tokens 已经是 fresh input。
pub fn calculate_for_app( pub fn calculate_for_app(
app_type: &str, app_type: &str,
usage: &TokenUsage, usage: &TokenUsage,
@@ -66,37 +61,32 @@ impl CostCalculator {
) -> CostBreakdown { ) -> CostBreakdown {
let input_includes_cache_read = let input_includes_cache_read =
crate::services::sql_helpers::is_cache_inclusive_app(app_type); crate::services::sql_helpers::is_cache_inclusive_app(app_type);
Self::calculate_with_input_semantics( Self::calculate_with_cache_semantics(
if input_includes_cache_read {
InputTokenSemantics::TotalIncludesCacheBuckets
} else {
InputTokenSemantics::FreshExcludesCache
},
usage, usage,
pricing, pricing,
cost_multiplier, cost_multiplier,
input_includes_cache_read,
) )
} }
pub fn calculate_with_input_semantics( fn calculate_with_cache_semantics(
input_semantics: InputTokenSemantics,
usage: &TokenUsage, usage: &TokenUsage,
pricing: &ModelPricing, pricing: &ModelPricing,
cost_multiplier: Decimal, cost_multiplier: Decimal,
input_includes_cache_read: bool,
) -> CostBreakdown { ) -> CostBreakdown {
let million = Decimal::from(1_000_000); let million = Decimal::from(1_000_000);
// OpenAI/Gemini 风格的 input_tokens 包含缓存读取和写入,需要扣除后再按输入价计费; // OpenAI/Gemini 风格的 input_tokens 包含缓存读取和写入,需要扣除后再按输入价计费;
// Claude/Anthropic 风格的 input_tokens 已经是 fresh input,不能再次扣减。 // Claude/Anthropic 风格的 input_tokens 已经是 fresh input,不能再次扣减。
let billable_input_tokens = let billable_input_tokens = if input_includes_cache_read {
if input_semantics == InputTokenSemantics::TotalIncludesCacheBuckets { usage
usage .input_tokens
.input_tokens .saturating_sub(usage.cache_read_tokens)
.saturating_sub(usage.cache_read_tokens) .saturating_sub(usage.cache_creation_tokens)
.saturating_sub(usage.cache_creation_tokens) } else {
} else { usage.input_tokens
usage.input_tokens };
};
// 各项基础成本(不含倍率) // 各项基础成本(不含倍率)
let input_cost = let input_cost =
@@ -122,15 +112,13 @@ impl CostCalculator {
} }
} }
pub fn try_calculate_with_input_semantics( pub fn try_calculate_for_app(
input_semantics: InputTokenSemantics, app_type: &str,
usage: &TokenUsage, usage: &TokenUsage,
pricing: Option<&ModelPricing>, pricing: Option<&ModelPricing>,
cost_multiplier: Decimal, cost_multiplier: Decimal,
) -> Option<CostBreakdown> { ) -> Option<CostBreakdown> {
pricing.map(|pricing| { pricing.map(|p| Self::calculate_for_app(app_type, usage, p, cost_multiplier))
Self::calculate_with_input_semantics(input_semantics, usage, pricing, cost_multiplier)
})
} }
} }
+14 -31
View File
@@ -2,9 +2,9 @@
use super::calculator::{CostBreakdown, CostCalculator, ModelPricing}; use super::calculator::{CostBreakdown, CostCalculator, ModelPricing};
use super::parser::TokenUsage; use super::parser::TokenUsage;
use super::semantics::InputTokenSemantics;
use crate::database::{Database, PRICING_SOURCE_REQUEST, PRICING_SOURCE_RESPONSE}; use crate::database::{Database, PRICING_SOURCE_REQUEST, PRICING_SOURCE_RESPONSE};
use crate::error::AppError; 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 crate::services::usage_stats::{find_model_pricing_row, is_placeholder_pricing_model};
use rusqlite::OptionalExtension; use rusqlite::OptionalExtension;
use rust_decimal::Decimal; use rust_decimal::Decimal;
@@ -72,9 +72,6 @@ pub struct RequestLog {
/// 用 model/request_model 猜——路由接管下三者可能各不相同。 /// 用 model/request_model 猜——路由接管下三者可能各不相同。
/// 错误行(未计价)为空字符串。 /// 错误行(未计价)为空字符串。
pub pricing_model: String, 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 usage: TokenUsage,
pub cost: Option<CostBreakdown>, pub cost: Option<CostBreakdown>,
pub latency_ms: u64, pub latency_ms: u64,
@@ -124,7 +121,12 @@ impl<'a> UsageLogger<'a> {
}; };
let created_at = chrono::Utc::now().timestamp(); let created_at = chrono::Utc::now().timestamp();
let input_token_semantics = log.input_token_semantics.stored_value(); 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 semantic = UsageSemantic::from_log(log, input_token_semantics); let semantic = UsageSemantic::from_log(log, input_token_semantics);
let existing = Self::load_existing_semantic(&conn, &log.request_id)?; let existing = Self::load_existing_semantic(&conn, &log.request_id)?;
@@ -264,7 +266,6 @@ impl<'a> UsageLogger<'a> {
status_code: u16, status_code: u16,
error_message: String, error_message: String,
latency_ms: u64, latency_ms: u64,
input_token_semantics: InputTokenSemantics,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
let request_model = model.clone(); let request_model = model.clone();
let log = RequestLog { let log = RequestLog {
@@ -275,7 +276,6 @@ impl<'a> UsageLogger<'a> {
request_model, request_model,
// 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行) // 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行)
pricing_model: String::new(), pricing_model: String::new(),
input_token_semantics,
usage: TokenUsage::default(), usage: TokenUsage::default(),
cost: None, cost: None,
latency_ms, latency_ms,
@@ -307,7 +307,6 @@ impl<'a> UsageLogger<'a> {
is_streaming: bool, is_streaming: bool,
session_id: Option<String>, session_id: Option<String>,
provider_type: Option<String>, provider_type: Option<String>,
input_token_semantics: InputTokenSemantics,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
let request_model = model.clone(); let request_model = model.clone();
let log = RequestLog { let log = RequestLog {
@@ -318,7 +317,6 @@ impl<'a> UsageLogger<'a> {
request_model, request_model,
// 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行) // 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行)
pricing_model: String::new(), pricing_model: String::new(),
input_token_semantics,
usage: TokenUsage::default(), usage: TokenUsage::default(),
cost: None, cost: None,
latency_ms, latency_ms,
@@ -362,17 +360,14 @@ impl<'a> UsageLogger<'a> {
} else { } else {
app_type app_type
}; };
let default_multiplier_raw = if default_app_type == "pi" { let default_multiplier_raw =
crate::settings::get_pi_default_cost_multiplier()
} else {
match self.db.get_default_cost_multiplier(default_app_type).await { match self.db.get_default_cost_multiplier(default_app_type).await {
Ok(value) => value, Ok(value) => value,
Err(e) => { Err(e) => {
log::warn!("[USG-003] 获取默认倍率失败 (app_type={app_type}): {e}"); log::warn!("[USG-003] 获取默认倍率失败 (app_type={app_type}): {e}");
"1".to_string() "1".to_string()
} }
} };
};
let default_multiplier = match Decimal::from_str(&default_multiplier_raw) { let default_multiplier = match Decimal::from_str(&default_multiplier_raw) {
Ok(value) => value, Ok(value) => value,
Err(e) => { Err(e) => {
@@ -383,17 +378,14 @@ impl<'a> UsageLogger<'a> {
} }
}; };
let default_pricing_source_raw = if default_app_type == "pi" { let default_pricing_source_raw =
crate::settings::get_pi_pricing_model_source()
} else {
match self.db.get_pricing_model_source(default_app_type).await { match self.db.get_pricing_model_source(default_app_type).await {
Ok(value) => value, Ok(value) => value,
Err(e) => { Err(e) => {
log::warn!("[USG-003] 获取默认计费模式失败 (app_type={app_type}): {e}"); log::warn!("[USG-003] 获取默认计费模式失败 (app_type={app_type}): {e}");
PRICING_SOURCE_RESPONSE.to_string() PRICING_SOURCE_RESPONSE.to_string()
} }
} };
};
let default_pricing_source = if default_pricing_source_raw == PRICING_SOURCE_RESPONSE let default_pricing_source = if default_pricing_source_raw == PRICING_SOURCE_RESPONSE
|| default_pricing_source_raw == PRICING_SOURCE_REQUEST || default_pricing_source_raw == PRICING_SOURCE_REQUEST
{ {
@@ -459,7 +451,6 @@ impl<'a> UsageLogger<'a> {
model: String, model: String,
request_model: String, request_model: String,
pricing_model: String, pricing_model: String,
input_token_semantics: InputTokenSemantics,
usage: TokenUsage, usage: TokenUsage,
cost_multiplier: Decimal, cost_multiplier: Decimal,
latency_ms: u64, latency_ms: u64,
@@ -480,8 +471,8 @@ impl<'a> UsageLogger<'a> {
log::warn!("[USG-002] 模型定价未找到,成本将记录为 0: {pricing_model}"); log::warn!("[USG-002] 模型定价未找到,成本将记录为 0: {pricing_model}");
} }
let cost = CostCalculator::try_calculate_with_input_semantics( let cost = CostCalculator::try_calculate_for_app(
input_token_semantics, &app_type,
&usage, &usage,
pricing.as_ref(), pricing.as_ref(),
cost_multiplier, cost_multiplier,
@@ -494,7 +485,6 @@ impl<'a> UsageLogger<'a> {
model, model,
request_model, request_model,
pricing_model, pricing_model,
input_token_semantics,
usage, usage,
cost, cost,
latency_ms, latency_ms,
@@ -523,7 +513,6 @@ mod tests {
model: "gpt-5.6".to_string(), model: "gpt-5.6".to_string(),
request_model: "gpt-5.6".to_string(), request_model: "gpt-5.6".to_string(),
pricing_model: "gpt-5.6".to_string(), pricing_model: "gpt-5.6".to_string(),
input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets,
usage: TokenUsage { usage: TokenUsage {
input_tokens, input_tokens,
output_tokens: 5, output_tokens: 5,
@@ -577,7 +566,6 @@ mod tests {
"test-model".to_string(), "test-model".to_string(),
"req-model".to_string(), "req-model".to_string(),
"test-model".to_string(), "test-model".to_string(),
InputTokenSemantics::FreshExcludesCache,
usage, usage,
Decimal::from(1), Decimal::from(1),
100, 100,
@@ -763,7 +751,6 @@ mod tests {
500, 500,
"Internal Server Error".to_string(), "Internal Server Error".to_string(),
50, 50,
InputTokenSemantics::FreshExcludesCache,
)?; )?;
// 验证错误记录已插入 // 验证错误记录已插入
@@ -791,7 +778,6 @@ mod tests {
model: "grok-4.5".to_string(), model: "grok-4.5".to_string(),
request_model: "grok-4.5".to_string(), request_model: "grok-4.5".to_string(),
pricing_model: String::new(), pricing_model: String::new(),
input_token_semantics: InputTokenSemantics::TotalIncludesCacheBuckets,
usage: TokenUsage::default(), usage: TokenUsage::default(),
cost: None, cost: None,
latency_ms: 1, latency_ms: 1,
@@ -812,10 +798,7 @@ mod tests {
[], [],
|row| row.get(0), |row| row.get(0),
)?; )?;
assert_eq!( assert_eq!(semantics, INPUT_TOKEN_SEMANTICS_TOTAL);
semantics,
InputTokenSemantics::TotalIncludesCacheBuckets.stored_value()
);
Ok(()) Ok(())
} }
} }
-3
View File
@@ -5,7 +5,6 @@
pub mod calculator; pub mod calculator;
pub mod logger; pub mod logger;
pub mod parser; pub mod parser;
pub mod semantics;
// 仅导出内部使用的类型,避免未使用警告 // 仅导出内部使用的类型,避免未使用警告
#[allow(unused_imports)] #[allow(unused_imports)]
@@ -14,5 +13,3 @@ pub use calculator::{CostBreakdown, CostCalculator, ModelPricing};
pub use logger::{RequestLog, UsageLogger}; pub use logger::{RequestLog, UsageLogger};
#[allow(unused_imports)] #[allow(unused_imports)]
pub use parser::TokenUsage; pub use parser::TokenUsage;
#[allow(unused_imports)]
pub use semantics::InputTokenSemantics;
-32
View File
@@ -1,32 +0,0 @@
//! 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,
}
}
}
-4
View File
@@ -138,10 +138,6 @@ impl ConfigService {
AppType::Hermes => { AppType::Hermes => {
// Hermes uses additive mode, no live sync needed // 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(()) Ok(())
+1 -47
View File
@@ -147,13 +147,6 @@ impl McpService {
AppType::Hermes => { AppType::Hermes => {
mcp::sync_single_server_to_hermes(&Default::default(), &server.id, &server.server)?; 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(()) Ok(())
} }
@@ -190,13 +183,6 @@ impl McpService {
AppType::Hermes => { AppType::Hermes => {
mcp::remove_server_from_hermes(id)?; 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(()) Ok(())
} }
@@ -241,10 +227,7 @@ impl McpService {
servers: &IndexMap<String, McpServer>, servers: &IndexMap<String, McpServer>,
app: &AppType, app: &AppType,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
if matches!( if matches!(app, AppType::OpenClaw | AppType::ClaudeDesktop) {
app,
AppType::OpenClaw | AppType::ClaudeDesktop | AppType::Pi
) {
return Ok(()); return Ok(());
} }
@@ -561,32 +544,3 @@ 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");
}
}
-3
View File
@@ -8,8 +8,6 @@ pub mod mcp;
pub mod model_fetch; pub mod model_fetch;
pub mod model_pricing; pub mod model_pricing;
pub mod omo; pub mod omo;
pub(crate) mod pi_catalog;
pub mod pi_prompt_files;
pub mod profile; pub mod profile;
pub mod prompt; pub mod prompt;
pub mod provider; pub mod provider;
@@ -23,7 +21,6 @@ pub mod session_usage_gemini;
pub mod session_usage_grokbuild; pub mod session_usage_grokbuild;
pub mod session_usage_opencode; pub mod session_usage_opencode;
pub mod skill; pub mod skill;
pub(crate) mod skill_deployment;
pub mod speedtest; pub mod speedtest;
pub mod sql_helpers; pub mod sql_helpers;
pub mod stream_check; pub mod stream_check;
+1 -5
View File
@@ -1,5 +1,4 @@
use crate::config::{atomic_write, write_json_file}; use crate::config::{atomic_write, write_json_file};
use crate::database::NewProviderAggregate;
use crate::error::AppError; use crate::error::AppError;
use crate::opencode_config::get_opencode_dir; use crate::opencode_config::get_opencode_dir;
use crate::provider::Provider; use crate::provider::Provider;
@@ -289,10 +288,7 @@ impl OmoService {
in_failover_queue: false, in_failover_queue: false,
}; };
state.db.create_provider(NewProviderAggregate::from_input( state.db.save_provider("opencode", &provider)?;
"opencode",
crate::services::provider::provider_to_mutation_input(provider.clone()),
)?)?;
state state
.db .db
.set_omo_provider_current("opencode", &provider.id, v.category)?; .set_omo_provider_current("opencode", &provider.id, v.category)?;
File diff suppressed because it is too large Load Diff
-483
View File
@@ -1,483 +0,0 @@
//! 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<Arc<Mutex<()>>> = LazyLock::new(|| Arc::new(Mutex::new(())));
pub(crate) type PiInstructionFileGuard = OwnedMutexGuard<()>;
pub(crate) fn lock_instruction_files() -> Result<PiInstructionFileGuard, AppError> {
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<PiPromptFileSnapshot, AppError> {
let guard = lock_instruction_files()?;
Self::read_under_guard(&guard, kind)
}
pub fn replace(
kind: PiPromptFileKind,
expected_revision: &str,
content: &str,
) -> Result<PiPromptFileSnapshot, AppError> {
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<bool, AppError> {
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<PiPromptFileSnapshot, AppError> {
Self::read_at(&get_pi_agent_dir()?, kind)
}
pub(crate) fn replace_under_guard(
_guard: &PiInstructionFileGuard,
kind: PiPromptFileKind,
expected_revision: &str,
content: &str,
) -> Result<PiPromptFileSnapshot, AppError> {
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<bool, AppError> {
Self::delete_at(&get_pi_agent_dir()?, kind, expected_revision)
}
fn read_at(root: &Path, kind: PiPromptFileKind) -> Result<PiPromptFileSnapshot, AppError> {
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<PiPromptFileSnapshot, AppError> {
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<bool, AppError> {
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<Vec<PiPromptTemplate>, AppError> {
let _guard = lock_instruction_files()?;
Self::list_at(&get_pi_agent_dir()?.join("prompts"))
}
pub fn upsert(
slug: &str,
expected_revision: &str,
content: &str,
) -> Result<PiPromptTemplate, AppError> {
let _guard = lock_instruction_files()?;
Self::upsert_at(
&get_pi_agent_dir()?.join("prompts"),
slug,
expected_revision,
content,
)
}
pub fn delete(slug: &str, expected_revision: &str) -> Result<bool, AppError> {
validate_template_slug(slug)?;
let _guard = lock_instruction_files()?;
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<Vec<PiPromptTemplate>, 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<PiPromptTemplate, AppError> {
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> {
// Pinned Pi's request-capture discovers `release notes.md`, but
// expandPromptTemplate("/release notes", ...) leaves the command
// unchanged because slash-command names are one token. Keep the managed
// namespace both callable and portable across Unix and Windows.
let windows_basename = slug
.split_once('.')
.map_or(slug, |(basename, _extension)| basename);
let windows_basename = windows_basename.to_ascii_lowercase();
let windows_reserved = matches!(
windows_basename.as_str(),
"con"
| "prn"
| "aux"
| "nul"
| "com1"
| "com2"
| "com3"
| "com4"
| "com5"
| "com6"
| "com7"
| "com8"
| "com9"
| "lpt1"
| "lpt2"
| "lpt3"
| "lpt4"
| "lpt5"
| "lpt6"
| "lpt7"
| "lpt8"
| "lpt9"
);
let valid = !slug.is_empty()
&& slug.len() <= MAX_TEMPLATE_SLUG_BYTES
&& slug != "."
&& slug != ".."
&& !slug.starts_with('.')
&& !slug.ends_with('.')
&& !windows_reserved
&& !slug.chars().any(|character| {
character.is_control()
|| character.is_whitespace()
|| matches!(
character,
'<' | '>' | ':' | '"' | '/' | '\\' | '|' | '?' | '*'
)
});
if valid {
Ok(())
} else {
Err(AppError::InvalidInput(
"Pi prompt-template slug must be one portable slash-command token (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",
"internal space",
"tab\tname",
"a/b",
r"a\b",
"bad:name",
"bad*name",
"CON",
"con.anything",
"LPT9",
"nul.json",
] {
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]);
}
}
+1 -1
View File
@@ -459,7 +459,7 @@ impl ProfileService {
.set_current_profile_id(scope.as_str(), Some(profile_id))?; .set_current_profile_id(scope.as_str(), Some(profile_id))?;
// 当前分组内所有接管已关闭;若其它应用也无接管,可停止代理服务。 // 当前分组内所有接管已关闭;若其它应用也无接管,可停止代理服务。
let should_stop_proxy = !state.proxy_service.shared_listener_is_desired_sync(); let should_stop_proxy = !state.db.is_live_takeover_active_sync();
Ok((warnings, should_stop_proxy)) Ok((warnings, should_stop_proxy))
} }
File diff suppressed because it is too large Load Diff
+15 -9
View File
@@ -5,7 +5,6 @@
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
use crate::app_config::AppType; use crate::app_config::AppType;
use crate::database::{NewEndpoint, ProviderKey};
use crate::error::AppError; use crate::error::AppError;
use crate::settings::CustomEndpoint; use crate::settings::CustomEndpoint;
use crate::store::AppState; use crate::store::AppState;
@@ -48,10 +47,9 @@ pub fn add_custom_endpoint(
)); ));
} }
let key = ProviderKey::new(app_type.as_str(), provider_id)?;
state state
.db .db
.add_provider_endpoint(&key, NewEndpoint::now(normalized)?)?; .add_custom_endpoint(app_type.as_str(), provider_id, &normalized)?;
Ok(()) Ok(())
} }
@@ -63,8 +61,9 @@ pub fn remove_custom_endpoint(
url: String, url: String,
) -> Result<(), AppError> { ) -> Result<(), AppError> {
let normalized = url.trim().trim_end_matches('/').to_string(); let normalized = url.trim().trim_end_matches('/').to_string();
let key = ProviderKey::new(app_type.as_str(), provider_id)?; state
state.db.remove_provider_endpoint(&key, &normalized)?; .db
.remove_custom_endpoint(app_type.as_str(), provider_id, &normalized)?;
Ok(()) Ok(())
} }
@@ -77,10 +76,17 @@ pub fn update_endpoint_last_used(
) -> Result<(), AppError> { ) -> Result<(), AppError> {
let normalized = url.trim().trim_end_matches('/').to_string(); let normalized = url.trim().trim_end_matches('/').to_string();
let key = ProviderKey::new(app_type.as_str(), provider_id)?; // Get provider, update last_used, save back
state let mut providers = state.db.get_all_providers(app_type.as_str())?;
.db if let Some(provider) = providers.get_mut(provider_id) {
.touch_provider_endpoint(&key, &normalized, now_millis()) if let Some(meta) = provider.meta.as_mut() {
if let Some(endpoint) = meta.custom_endpoints.get_mut(&normalized) {
endpoint.last_used = Some(now_millis());
state.db.save_provider(app_type.as_str(), provider)?;
}
}
}
Ok(())
} }
/// Get current timestamp in milliseconds /// Get current timestamp in milliseconds
+24 -131
View File
@@ -19,10 +19,7 @@ use crate::store::AppState;
use super::gemini_auth::{ use super::gemini_auth::{
detect_gemini_auth_type, ensure_google_oauth_security_flag, GeminiAuthType, detect_gemini_auth_type, ensure_google_oauth_security_flag, GeminiAuthType,
}; };
use super::{ use super::normalize_claude_models_in_value;
normalize_claude_models_in_value, provider_row_fingerprint, provider_to_mutation_input,
reconcile_provider_record_with_precondition, ReconcilePrecondition,
};
/// ChatGPT Codex catalogs gpt-5.6 at a 372K context window with a ~353K /// ChatGPT Codex catalogs gpt-5.6 at a 372K context window with a ~353K
/// effective budget (openai/codex#31860), far below the 1.05M API spec. /// effective budget (openai/codex#31860), far below the 1.05M API spec.
@@ -530,7 +527,6 @@ fn settings_contain_common_config(app_type: &AppType, settings: &Value, snippet:
| AppType::OpenCode | AppType::OpenCode
| AppType::OpenClaw | AppType::OpenClaw
| AppType::Hermes | AppType::Hermes
| AppType::Pi
| AppType::ClaudeDesktop => false, | AppType::ClaudeDesktop => false,
} }
} }
@@ -605,7 +601,6 @@ pub(crate) fn remove_common_config_from_settings(
| AppType::OpenCode | AppType::OpenCode
| AppType::OpenClaw | AppType::OpenClaw
| AppType::Hermes | AppType::Hermes
| AppType::Pi
| AppType::ClaudeDesktop => Ok(settings.clone()), | AppType::ClaudeDesktop => Ok(settings.clone()),
} }
} }
@@ -665,7 +660,6 @@ fn apply_common_config_to_settings(
| AppType::OpenCode | AppType::OpenCode
| AppType::OpenClaw | AppType::OpenClaw
| AppType::Hermes | AppType::Hermes
| AppType::Pi
| AppType::ClaudeDesktop => Ok(settings.clone()), | AppType::ClaudeDesktop => Ok(settings.clone()),
} }
} }
@@ -1168,13 +1162,6 @@ pub(crate) fn write_live_snapshot(app_type: &AppType, provider: &Provider) -> Re
crate::hermes_config::set_provider(&provider.id, provider.settings_config.clone())?; crate::hermes_config::set_provider(&provider.id, provider.settings_config.clone())?;
log::debug!("Hermes provider '{}' written to live config", provider.id); 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(()) Ok(())
} }
@@ -1289,34 +1276,10 @@ fn sync_current_provider_for_app_respecting_takeover(
/// ///
/// For additive mode apps (OpenCode), all providers are synced instead of just the current one. /// For additive mode apps (OpenCode), all providers are synced instead of just the current one.
pub fn sync_current_to_live(state: &AppState) -> Result<(), AppError> { pub fn sync_current_to_live(state: &AppState) -> Result<(), AppError> {
// Pi's portable-import boundary closes runtime admission before replacing // Sync providers based on mode
// SQLite. Recover it first so an unrelated application's broken live file for app_type in AppType::all() {
// cannot strand Pi in that closed state.
let mut pi_failures = Vec::new();
if let Err(error) =
crate::services::pi_catalog::PiCatalogCoordinator::reconcile_portable_import(state)
{
pi_failures.push(format!("provider={error}"));
}
if let Err(error) = crate::services::skill::SkillService::sync_to_app(&state.db, &AppType::Pi) {
pi_failures.push(format!("skill={error}"));
}
if !pi_failures.is_empty() {
return Err(AppError::Config(format!(
"Pi live reconciliation incomplete: {}",
pi_failures.join("; ")
)));
}
// Preserve the existing fail-fast behavior for all other provider views.
for app_type in AppType::all().filter(|app_type| !matches!(app_type, AppType::Pi)) {
if app_type.is_additive_mode() { if app_type.is_additive_mode() {
// Provider rename and every additive live mutation share this // Additive mode: sync ALL providers
// per-app lock. Acquire it before reading the catalog so a key
// cannot be renamed after this sync captured a stale provider map.
let _guard = futures::executor::block_on(
state.proxy_service.lock_switch_for_app(app_type.as_str()),
);
sync_all_providers_to_live(state, &app_type)?; sync_all_providers_to_live(state, &app_type)?;
} else { } else {
// Switch mode: sync only current provider. During proxy takeover, // Switch mode: sync only current provider. During proxy takeover,
@@ -1326,28 +1289,20 @@ pub fn sync_current_to_live(state: &AppState) -> Result<(), AppError> {
} }
} }
let mut failures = Vec::new(); // MCP syncbest-effort 逐应用投影,内部已聚合失败)。错误暂存到
if let Err(error) = McpService::sync_all_enabled(state) { // Skill 同步之后再返回:MCP 的失败不该跳过 Skill 同步,但调用方
failures.push(format!("mcp={error}")); //(配置导入 / 云同步恢复)需要知道结果不完整。
} let mcp_result = McpService::sync_all_enabled(state);
// Continue through all apps so one collision cannot hide unrelated Skills. // Skill sync
for app_type in AppType::all().filter(|app_type| !matches!(app_type, AppType::Pi)) { for app_type in AppType::all() {
if let Err(error) = crate::services::skill::SkillService::sync_to_app(&state.db, &app_type) if let Err(e) = crate::services::skill::SkillService::sync_to_app(&state.db, &app_type) {
{ log::warn!("同步 Skill 到 {app_type:?} 失败: {e}");
log::warn!("同步 Skill 到 {app_type:?} 失败: {error}"); // Continue syncing other apps, don't abort
failures.push(format!("skill:{}={error}", app_type.as_str()));
} }
} }
if failures.is_empty() { mcp_result
Ok(())
} else {
Err(AppError::Config(format!(
"live synchronization incomplete: {}",
failures.join("; ")
)))
}
} }
/// Read current live settings for an app type /// Read current live settings for an app type
@@ -1462,11 +1417,6 @@ pub fn read_live_settings(app_type: AppType) -> Result<Value, AppError> {
let config = crate::hermes_config::yaml_to_json(&yaml_config)?; let config = crate::hermes_config::yaml_to_json(&yaml_config)?;
Ok(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",
)),
} }
} }
@@ -1575,13 +1525,6 @@ pub fn import_default_config(state: &AppState, app_type: AppType) -> Result<bool
"config": config_obj "config": config_obj
}) })
} }
AppType::Pi => {
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 // OpenCode, OpenClaw and Hermes use additive mode and are handled by early return above
AppType::OpenCode | AppType::OpenClaw | AppType::Hermes => { AppType::OpenCode | AppType::OpenClaw | AppType::Hermes => {
unreachable!("additive mode apps are handled by early return") unreachable!("additive mode apps are handled by early return")
@@ -1621,12 +1564,7 @@ pub fn import_default_config(state: &AppState, app_type: AppType) -> Result<bool
.to_string(), .to_string(),
); );
reconcile_provider_record_with_precondition( state.db.save_provider(app_type.as_str(), &provider)?;
&state.db,
app_type.as_str(),
provider_to_mutation_input(provider.clone()),
ReconcilePrecondition::ExpectAbsent,
)?;
state state
.db .db
.set_current_provider(app_type.as_str(), &provider.id)?; .set_current_provider(app_type.as_str(), &provider.id)?;
@@ -1794,25 +1732,15 @@ pub fn import_opencode_providers_from_live(state: &AppState) -> Result<usize, Ap
}; };
if existing_ids.contains(&id) { if existing_ids.contains(&id) {
match state.db.get_provider_aggregate("opencode", &id) { match state.db.get_provider_by_id(&id, "opencode") {
Ok(Some(existing)) => { Ok(Some(existing)) => {
let existing = existing.provider;
let display_name = config.name.clone().unwrap_or_else(|| existing.name.clone()); let display_name = config.name.clone().unwrap_or_else(|| existing.name.clone());
if existing.settings_config != settings_config || existing.name != display_name if existing.settings_config != settings_config || existing.name != display_name
{ {
let fingerprint = provider_row_fingerprint(&existing);
let mut provider = existing; let mut provider = existing;
provider.name = display_name; provider.name = display_name;
provider.settings_config = settings_config; provider.settings_config = settings_config;
if let Some(meta) = provider.meta.as_mut() { if let Err(e) = state.db.save_provider("opencode", &provider) {
meta.custom_endpoints.clear();
}
if let Err(e) = reconcile_provider_record_with_precondition(
&state.db,
"opencode",
provider_to_mutation_input(provider),
ReconcilePrecondition::ExpectPresent { fingerprint },
) {
log::warn!( log::warn!(
"Failed to update OpenCode provider '{id}' from live config: {e}" "Failed to update OpenCode provider '{id}' from live config: {e}"
); );
@@ -1839,12 +1767,7 @@ pub fn import_opencode_providers_from_live(state: &AppState) -> Result<usize, Ap
}); });
// Save to database // Save to database
if let Err(e) = reconcile_provider_record_with_precondition( if let Err(e) = state.db.save_provider("opencode", &provider) {
&state.db,
"opencode",
provider_to_mutation_input(provider),
ReconcilePrecondition::ExpectAbsent,
) {
log::warn!("Failed to import OpenCode provider '{id}': {e}"); log::warn!("Failed to import OpenCode provider '{id}': {e}");
continue; continue;
} }
@@ -1894,22 +1817,12 @@ pub fn import_openclaw_providers_from_live(state: &AppState) -> Result<usize, Ap
}; };
if existing_ids.contains(&id) { if existing_ids.contains(&id) {
match state.db.get_provider_aggregate("openclaw", &id) { match state.db.get_provider_by_id(&id, "openclaw") {
Ok(Some(existing)) => { Ok(Some(existing)) => {
let existing = existing.provider;
if existing.settings_config != settings_config { if existing.settings_config != settings_config {
let fingerprint = provider_row_fingerprint(&existing);
let mut provider = existing; let mut provider = existing;
provider.settings_config = settings_config; provider.settings_config = settings_config;
if let Some(meta) = provider.meta.as_mut() { if let Err(e) = state.db.save_provider("openclaw", &provider) {
meta.custom_endpoints.clear();
}
if let Err(e) = reconcile_provider_record_with_precondition(
&state.db,
"openclaw",
provider_to_mutation_input(provider),
ReconcilePrecondition::ExpectPresent { fingerprint },
) {
log::warn!( log::warn!(
"Failed to update OpenClaw provider '{id}' from live config: {e}" "Failed to update OpenClaw provider '{id}' from live config: {e}"
); );
@@ -1942,12 +1855,7 @@ pub fn import_openclaw_providers_from_live(state: &AppState) -> Result<usize, Ap
}); });
// Save to database // Save to database
if let Err(e) = reconcile_provider_record_with_precondition( if let Err(e) = state.db.save_provider("openclaw", &provider) {
&state.db,
"openclaw",
provider_to_mutation_input(provider),
ReconcilePrecondition::ExpectAbsent,
) {
log::warn!("Failed to import OpenClaw provider '{id}': {e}"); log::warn!("Failed to import OpenClaw provider '{id}': {e}");
continue; continue;
} }
@@ -1984,22 +1892,12 @@ pub fn import_hermes_providers_from_live(state: &AppState) -> Result<usize, AppE
} }
if existing_ids.contains(&name) { if existing_ids.contains(&name) {
match state.db.get_provider_aggregate("hermes", &name) { match state.db.get_provider_by_id(&name, "hermes") {
Ok(Some(existing)) => { Ok(Some(existing)) => {
let existing = existing.provider;
if existing.settings_config != config { if existing.settings_config != config {
let fingerprint = provider_row_fingerprint(&existing);
let mut provider = existing; let mut provider = existing;
provider.settings_config = config; provider.settings_config = config;
if let Some(meta) = provider.meta.as_mut() { if let Err(e) = state.db.save_provider("hermes", &provider) {
meta.custom_endpoints.clear();
}
if let Err(e) = reconcile_provider_record_with_precondition(
&state.db,
"hermes",
provider_to_mutation_input(provider),
ReconcilePrecondition::ExpectPresent { fingerprint },
) {
log::warn!( log::warn!(
"Failed to update Hermes provider '{name}' from live config: {e}" "Failed to update Hermes provider '{name}' from live config: {e}"
); );
@@ -2025,12 +1923,7 @@ pub fn import_hermes_providers_from_live(state: &AppState) -> Result<usize, AppE
}); });
// Save to database // Save to database
if let Err(e) = reconcile_provider_record_with_precondition( if let Err(e) = state.db.save_provider("hermes", &provider) {
&state.db,
"hermes",
provider_to_mutation_input(provider),
ReconcilePrecondition::ExpectAbsent,
) {
log::warn!("Failed to import Hermes provider '{name}': {e}"); log::warn!("Failed to import Hermes provider '{name}': {e}");
continue; continue;
} }
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+2 -35
View File
@@ -54,7 +54,8 @@ pub fn fresh_input_sql(alias: &str) -> String {
format!( format!(
"CASE \ "CASE \
WHEN {prefix}input_token_semantics = {INPUT_TOKEN_SEMANTICS_FRESH} THEN {prefix}input_tokens \ WHEN {prefix}input_token_semantics = {INPUT_TOKEN_SEMANTICS_FRESH} THEN {prefix}input_tokens \
WHEN {prefix}input_token_semantics = {INPUT_TOKEN_SEMANTICS_TOTAL} \ WHEN {prefix}app_type IN ({app_type_list}) \
AND {prefix}input_token_semantics = {INPUT_TOKEN_SEMANTICS_TOTAL} \
AND {prefix}input_tokens >= ({prefix}cache_read_tokens + {prefix}cache_creation_tokens) \ AND {prefix}input_tokens >= ({prefix}cache_read_tokens + {prefix}cache_creation_tokens) \
THEN ({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}) \ WHEN {prefix}app_type IN ({app_type_list}) \
@@ -143,40 +144,6 @@ mod tests {
assert_eq!(total, 400 + 500 + 450 + 200); 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::<Result<_, _>>()
.unwrap();
assert_eq!(
values,
vec![
("pi-anthropic".to_string(), 1000),
("pi-openai".to_string(), 200),
]
);
}
#[test] #[test]
fn fresh_input_handles_codex_with_cache_exceeding_input() { fn fresh_input_handles_codex_with_cache_exceeding_input() {
// Defensive: if a malformed Codex row somehow has cache > input, // Defensive: if a malformed Codex row somehow has cache > input,
-56
View File
@@ -181,7 +181,6 @@ impl StreamCheckService {
} }
AppType::OpenClaw => Self::extract_openclaw_base_url(provider), AppType::OpenClaw => Self::extract_openclaw_base_url(provider),
AppType::Hermes => Self::extract_hermes_base_url(provider), AppType::Hermes => Self::extract_hermes_base_url(provider),
AppType::Pi => Self::extract_pi_base_url(provider),
AppType::ClaudeDesktop => ClaudeAdapter::new() AppType::ClaudeDesktop => ClaudeAdapter::new()
.extract_base_url(provider) .extract_base_url(provider)
.map_err(|e| AppError::Message(format!("Failed to extract base_url: {e}"))), .map_err(|e| AppError::Message(format!("Failed to extract base_url: {e}"))),
@@ -324,31 +323,6 @@ 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<String, AppError> {
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 }, ... }` /// OpenCode: `{ npm, options: { baseURL, apiKey }, ... }`
/// ///
/// 用户未显式填 `options.baseURL` 时,按 `npm`AI SDK 包)回退到包自带默认端点。 /// 用户未显式填 `options.baseURL` 时,按 `npm`AI SDK 包)回退到包自带默认端点。
@@ -527,36 +501,6 @@ 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] #[test]
fn test_resolve_base_url_uses_explicit_url_or_errors_when_missing() { fn test_resolve_base_url_uses_explicit_url_or_errors_when_missing() {
// 有显式 base_url → 直接用 // 有显式 base_url → 直接用
+2 -43
View File
@@ -106,7 +106,7 @@ pub(crate) fn build_local_snapshot(
db: &crate::database::Database, db: &crate::database::Database,
) -> Result<LocalSnapshot, AppError> { ) -> Result<LocalSnapshot, AppError> {
// Export database to SQL string // Export database to SQL string
let sql_string = db.export_portable_sql_string_for_sync()?; let sql_string = db.export_sql_string_for_sync()?;
let db_sql = sql_string.into_bytes(); let db_sql = sql_string.into_bytes();
// Pack skills into deterministic ZIP // Pack skills into deterministic ZIP
@@ -310,16 +310,6 @@ pub(crate) fn apply_snapshot(
db: &crate::database::Database, db: &crate::database::Database,
db_sql: &[u8], db_sql: &[u8],
skills_zip: &[u8], skills_zip: &[u8],
) -> Result<(), AppError> {
crate::services::skill_deployment::PiSkillDeploymentService::coordinate_portable_import(|| {
apply_snapshot_under_pi_skill_guard(db, db_sql, skills_zip)
})
}
fn apply_snapshot_under_pi_skill_guard(
db: &crate::database::Database,
db_sql: &[u8],
skills_zip: &[u8],
) -> Result<(), AppError> { ) -> Result<(), AppError> {
let sql_str = std::str::from_utf8(db_sql).map_err(|e| { let sql_str = std::str::from_utf8(db_sql).map_err(|e| {
localized( localized(
@@ -333,7 +323,7 @@ fn apply_snapshot_under_pi_skill_guard(
// Replace skills first, then import database; roll back skills on DB failure. // Replace skills first, then import database; roll back skills on DB failure.
restore_skills_zip(skills_zip)?; restore_skills_zip(skills_zip)?;
if let Err(db_err) = db.import_portable_sql_string_for_sync(sql_str) { if let Err(db_err) = db.import_sql_string_for_sync(sql_str) {
if let Err(rollback_err) = restore_skills_from_backup(&skills_backup) { if let Err(rollback_err) = restore_skills_from_backup(&skills_backup) {
return Err(localized( return Err(localized(
"sync.db_import_and_rollback_failed", "sync.db_import_and_rollback_failed",
@@ -429,37 +419,6 @@ where
mod tests { mod tests {
use super::*; use super::*;
#[test]
#[serial_test::serial]
fn snapshot_application_waits_for_the_pi_skill_ownership_boundary() {
use std::sync::mpsc;
use std::time::Duration;
let db = crate::database::Database::memory().expect("database");
let guard = crate::services::skill_deployment::PiSkillDeploymentService::operation_guard();
let (ready_tx, ready_rx) = mpsc::channel();
let (result_tx, result_rx) = mpsc::channel();
let worker = std::thread::spawn(move || {
ready_tx.send(()).expect("signal worker ready");
let result = apply_snapshot(&db, &[0xff], &[]);
result_tx.send(result).expect("signal snapshot result");
});
ready_rx
.recv_timeout(Duration::from_secs(2))
.expect("worker reaches snapshot entry");
assert!(
result_rx.recv_timeout(Duration::from_millis(100)).is_err(),
"WebDAV/S3 snapshot application must wait for a concurrent Pi Skill mutation"
);
drop(guard);
let result = result_rx
.recv_timeout(Duration::from_secs(2))
.expect("snapshot proceeds after ownership boundary release");
assert!(result.is_err(), "the intentionally invalid SQL must fail");
worker.join().expect("worker");
}
fn artifact(sha256: &str, size: u64) -> ArtifactMeta { fn artifact(sha256: &str, size: u64) -> ArtifactMeta {
ArtifactMeta { ArtifactMeta {
sha256: sha256.to_string(), sha256: sha256.to_string(),
+19 -67
View File
@@ -137,9 +137,8 @@ pub struct RequestLogDetail {
pub output_tokens: u32, pub output_tokens: u32,
pub cache_read_tokens: u32, pub cache_read_tokens: u32,
pub cache_creation_tokens: u32, pub cache_creation_tokens: u32,
/// Persisted request-level semantics used by both pricing and UI cache /// Internal storage semantics; omitted from the UI/API payload.
/// normalization. This must cross IPC; app-type inference is only a legacy #[serde(skip)]
/// fallback for rows written before the semantics column existed.
pub input_token_semantics: i64, pub input_token_semantics: i64,
pub input_cost_usd: String, pub input_cost_usd: String,
pub output_cost_usd: String, pub output_cost_usd: String,
@@ -1654,10 +1653,10 @@ impl Database {
let detail_sql = format!( let detail_sql = format!(
"SELECT l.request_id, l.provider_id, {detail_pname} as provider_name, l.app_type, l.model, "SELECT l.request_id, l.provider_id, {detail_pname} as provider_name, l.app_type, l.model,
l.request_model, l.cost_multiplier, l.request_model, l.cost_multiplier,
l.input_tokens, l.output_tokens, l.cache_read_tokens, l.cache_creation_tokens, input_tokens, output_tokens, cache_read_tokens, 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, input_cost_usd, output_cost_usd, cache_read_cost_usd, cache_creation_cost_usd, total_cost_usd,
l.is_streaming, l.latency_ms, l.first_token_ms, l.duration_ms, is_streaming, latency_ms, first_token_ms, duration_ms,
l.status_code, l.error_message, l.created_at, l.data_source, l.pricing_model, status_code, error_message, created_at, l.data_source, l.pricing_model,
l.input_token_semantics l.input_token_semantics
FROM proxy_request_logs l FROM proxy_request_logs l
LEFT JOIN providers p ON l.provider_id = p.id AND l.app_type = p.app_type LEFT JOIN providers p ON l.provider_id = p.id AND l.app_type = p.app_type
@@ -1898,18 +1897,19 @@ impl Database {
// 1. 历史 cache-inclusive 行只包含 cache read;新 total 行还包含 cache write。 // 1. 历史 cache-inclusive 行只包含 cache read;新 total 行还包含 cache write。
// 2. Claude/Anthropic 的 input_tokens 已经是 fresh input,不能再次扣减 // 2. Claude/Anthropic 的 input_tokens 已经是 fresh input,不能再次扣减
// 3. 各项成本是基础成本(不含倍率),倍率只作用于最终总价 // 3. 各项成本是基础成本(不含倍率),倍率只作用于最终总价
let billable_input_tokens = if log.input_token_semantics == INPUT_TOKEN_SEMANTICS_FRESH { let cache_inclusive_app =
log.input_tokens as u64 crate::services::sql_helpers::is_cache_inclusive_app(log.app_type.as_str());
} else if log.input_token_semantics == INPUT_TOKEN_SEMANTICS_TOTAL { let billable_input_tokens =
(log.input_tokens as u64) if !cache_inclusive_app || log.input_token_semantics == INPUT_TOKEN_SEMANTICS_FRESH {
.saturating_sub(log.cache_read_tokens as u64) log.input_tokens as u64
.saturating_sub(log.cache_creation_tokens as u64) } else if log.input_token_semantics == INPUT_TOKEN_SEMANTICS_TOTAL {
} else if crate::services::sql_helpers::is_cache_inclusive_app(log.app_type.as_str()) { (log.input_tokens as u64)
// v12 and earlier: input included cache reads but excluded cache writes. .saturating_sub(log.cache_read_tokens as u64)
(log.input_tokens as u64).saturating_sub(log.cache_read_tokens as u64) .saturating_sub(log.cache_creation_tokens as u64)
} else { } else {
log.input_tokens as u64 // 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 input_cost = let input_cost =
rust_decimal::Decimal::from(billable_input_tokens) * pricing.input / million; rust_decimal::Decimal::from(billable_input_tokens) * pricing.input / million;
let output_cost = let output_cost =
@@ -2407,54 +2407,6 @@ mod tests {
Ok(()) 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> { fn create_legacy_nullable_logs_table(conn: &Connection) -> Result<(), AppError> {
conn.execute( conn.execute(
"CREATE TABLE proxy_request_logs ( "CREATE TABLE proxy_request_logs (
+2 -8
View File
@@ -4,7 +4,7 @@ pub mod terminal;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use providers::{claude, codex, gemini, grokbuild, hermes, openclaw, opencode, pi}; use providers::{claude, codex, gemini, grokbuild, hermes, openclaw, opencode};
#[derive(Debug, Clone, Serialize)] #[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
@@ -56,7 +56,7 @@ pub struct DeleteSessionOutcome {
} }
pub fn scan_sessions() -> Vec<SessionMeta> { pub fn scan_sessions() -> Vec<SessionMeta> {
let (r1, r2, r3, r4, r5, r6, r7, r8) = std::thread::scope(|s| { let (r1, r2, r3, r4, r5, r6, r7) = std::thread::scope(|s| {
let h1 = s.spawn(codex::scan_sessions); let h1 = s.spawn(codex::scan_sessions);
let h2 = s.spawn(claude::scan_sessions); let h2 = s.spawn(claude::scan_sessions);
let h3 = s.spawn(opencode::scan_sessions); let h3 = s.spawn(opencode::scan_sessions);
@@ -64,7 +64,6 @@ pub fn scan_sessions() -> Vec<SessionMeta> {
let h5 = s.spawn(gemini::scan_sessions); let h5 = s.spawn(gemini::scan_sessions);
let h6 = s.spawn(hermes::scan_sessions); let h6 = s.spawn(hermes::scan_sessions);
let h7 = s.spawn(grokbuild::scan_sessions); let h7 = s.spawn(grokbuild::scan_sessions);
let h8 = s.spawn(pi::scan_sessions);
( (
h1.join().unwrap_or_default(), h1.join().unwrap_or_default(),
h2.join().unwrap_or_default(), h2.join().unwrap_or_default(),
@@ -73,7 +72,6 @@ pub fn scan_sessions() -> Vec<SessionMeta> {
h5.join().unwrap_or_default(), h5.join().unwrap_or_default(),
h6.join().unwrap_or_default(), h6.join().unwrap_or_default(),
h7.join().unwrap_or_default(), h7.join().unwrap_or_default(),
h8.join().unwrap_or_default(),
) )
}); });
@@ -85,7 +83,6 @@ pub fn scan_sessions() -> Vec<SessionMeta> {
sessions.extend(r5); sessions.extend(r5);
sessions.extend(r6); sessions.extend(r6);
sessions.extend(r7); sessions.extend(r7);
sessions.extend(r8);
sessions.sort_by(|a, b| { sessions.sort_by(|a, b| {
let a_ts = a.last_active_at.or(a.created_at).unwrap_or(0); let a_ts = a.last_active_at.or(a.created_at).unwrap_or(0);
@@ -114,7 +111,6 @@ pub fn load_messages(provider_id: &str, source_path: &str) -> Result<Vec<Session
"gemini" => gemini::load_messages(path), "gemini" => gemini::load_messages(path),
"grokbuild" => grokbuild::load_messages(path), "grokbuild" => grokbuild::load_messages(path),
"hermes" => hermes::load_messages(path), "hermes" => hermes::load_messages(path),
"pi" => pi::load_messages(path),
_ => Err(format!("Unsupported provider: {provider_id}")), _ => Err(format!("Unsupported provider: {provider_id}")),
} }
} }
@@ -177,7 +173,6 @@ fn delete_session_with_roots(
grokbuild::delete_session(&validated_root, &validated_source, session_id) grokbuild::delete_session(&validated_root, &validated_source, session_id)
} }
"hermes" => hermes::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}")), _ => Err(format!("Unsupported provider: {provider_id}")),
}; };
} }
@@ -208,7 +203,6 @@ fn provider_roots(provider_id: &str) -> Result<Vec<PathBuf>, String> {
"gemini" => vec![crate::gemini_config::get_gemini_dir().join("tmp")], "gemini" => vec![crate::gemini_config::get_gemini_dir().join("tmp")],
"grokbuild" => grokbuild::session_roots(), "grokbuild" => grokbuild::session_roots(),
"hermes" => vec![crate::hermes_config::get_hermes_dir().join("sessions")], "hermes" => vec![crate::hermes_config::get_hermes_dir().join("sessions")],
"pi" => pi::session_roots(),
_ => return Err(format!("Unsupported provider: {provider_id}")), _ => return Err(format!("Unsupported provider: {provider_id}")),
}; };
@@ -5,5 +5,4 @@ pub mod grokbuild;
pub mod hermes; pub mod hermes;
pub mod openclaw; pub mod openclaw;
pub mod opencode; pub mod opencode;
pub mod pi;
mod utils; mod utils;
@@ -1,731 +0,0 @@
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<i64>,
version: u64,
}
#[derive(Debug)]
struct SessionTree {
header: SessionHeader,
active_ids: HashSet<String>,
}
#[derive(Default)]
struct ActiveSessionData {
messages: Vec<SessionMessage>,
first_user_message: Option<String>,
last_message: Option<String>,
explicit_name: Option<Option<String>>,
last_active_at: Option<i64>,
}
/// 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<PathBuf> {
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<PathBuf> {
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<SessionMeta> {
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<SessionMeta> {
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<Vec<SessionMessage>, 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<Vec<SessionMessage>, 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<bool, String> {
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<SessionMeta, String> {
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<SessionTree, String> {
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::<String, Option<String>>::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::<Value>(&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<ActiveSessionData, String> {
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::<Value>(&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<SessionMessage>, content: &str, ts: Option<i64>) {
if !content.trim().is_empty() {
messages.push(SessionMessage {
role: "system".to_string(),
content: content.to_string(),
ts,
});
}
}
fn parse_header(value: &Value) -> Result<SessionHeader, String> {
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<String>)> {
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<PathBuf>) {
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<_>>(),
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", <SessionManager.getSessionFile()>] 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<_>>(),
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"));
}
}
@@ -333,7 +333,7 @@ fn build_shell_command(command: &str, cwd: Option<&str>) -> String {
/// ///
/// 单引号内不做任何展开,唯一的特例是 `'` 自身无法被表示:用「闭合-转义-重开」 /// 单引号内不做任何展开,唯一的特例是 `'` 自身无法被表示:用「闭合-转义-重开」
/// 的 `'\''` 序列绕过。 /// 的 `'\''` 序列绕过。
pub(crate) fn shell_escape(value: &str) -> String { fn shell_escape(value: &str) -> String {
format!("'{}'", value.replace('\'', r"'\''")) format!("'{}'", value.replace('\'', r"'\''"))
} }
+29 -340
View File
@@ -8,179 +8,19 @@ use crate::error::AppError;
use crate::services::skill::{SkillStorageLocation, SyncMethod}; use crate::services::skill::{SkillStorageLocation, SyncMethod};
/// 自定义端点配置(历史兼容,实际存储在 provider.meta.custom_endpoints /// 自定义端点配置(历史兼容,实际存储在 provider.meta.custom_endpoints
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
pub struct CustomEndpoint { pub struct CustomEndpoint {
pub url: String, pub url: String,
pub added_at: Option<i64>, pub added_at: i64,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub last_used: Option<i64>, pub last_used: Option<i64>,
} }
/// 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_cost_multiplier")]
pub default_cost_multiplier: String,
#[serde(default = "default_pi_pricing_model_source")]
pub pricing_model_source: String,
#[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,
}
fn default_pi_cost_multiplier() -> String {
"1".to_string()
}
fn default_pi_pricing_model_source() -> String {
crate::database::PRICING_SOURCE_RESPONSE.to_string()
}
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,
default_cost_multiplier: default_pi_cost_multiplier(),
pricing_model_source: default_pi_pricing_model_source(),
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> {
crate::database::validate_cost_multiplier(&self.default_cost_multiplier)?;
crate::database::validate_pricing_source(&self.pricing_model_source)?;
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 { fn default_true() -> bool {
true 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(<redacted>)")
}
}
/// 主页面显示的应用配置 /// 主页面显示的应用配置
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
@@ -206,8 +46,6 @@ pub struct VisibleApps {
pub openclaw: bool, pub openclaw: bool,
#[serde(default)] #[serde(default)]
pub hermes: bool, pub hermes: bool,
#[serde(default = "default_true")]
pub pi: bool,
} }
impl Default for VisibleApps { impl Default for VisibleApps {
@@ -221,7 +59,6 @@ impl Default for VisibleApps {
opencode: true, opencode: true,
openclaw: true, openclaw: true,
hermes: false, // 默认不显示,需用户手动启用 hermes: false, // 默认不显示,需用户手动启用
pi: true,
} }
} }
} }
@@ -238,7 +75,6 @@ impl VisibleApps {
AppType::OpenCode => self.opencode, AppType::OpenCode => self.opencode,
AppType::OpenClaw => self.openclaw, AppType::OpenClaw => self.openclaw,
AppType::Hermes => self.hermes, AppType::Hermes => self.hermes,
AppType::Pi => self.pi,
} }
} }
} }
@@ -586,8 +422,6 @@ pub struct AppSettings {
pub openclaw_config_dir: Option<String>, pub openclaw_config_dir: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
pub hermes_config_dir: Option<String>, pub hermes_config_dir: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub pi_config_dir: Option<String>,
// ===== 当前供应商 ID(设备级)===== // ===== 当前供应商 ID(设备级)=====
/// 当前 Claude 供应商 ID(本地存储,优先于数据库 is_current /// 当前 Claude 供应商 ID(本地存储,优先于数据库 is_current
@@ -614,29 +448,6 @@ pub struct AppSettings {
/// 当前 Hermes 供应商 ID(本地存储,保持结构一致) /// 当前 Hermes 供应商 ID(本地存储,保持结构一致)
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
pub current_provider_hermes: Option<String>, pub current_provider_hermes: Option<String>,
/// 当前 Pi 供应商 ID(本地存储)
#[serde(default, skip_serializing_if = "Option::is_none")]
pub current_provider_pi: Option<String>,
/// 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<GatewayToken>,
// ===== Skill 同步设置 ===== // ===== Skill 同步设置 =====
/// Skill 同步方式:auto(默认,优先 symlink)、symlink、copy /// Skill 同步方式:auto(默认,优先 symlink)、symlink、copy
@@ -722,7 +533,6 @@ impl Default for AppSettings {
opencode_config_dir: None, opencode_config_dir: None,
openclaw_config_dir: None, openclaw_config_dir: None,
hermes_config_dir: None, hermes_config_dir: None,
pi_config_dir: None,
current_provider_claude: None, current_provider_claude: None,
current_provider_claude_desktop: None, current_provider_claude_desktop: None,
current_provider_codex: None, current_provider_codex: None,
@@ -731,10 +541,6 @@ impl Default for AppSettings {
current_provider_opencode: None, current_provider_opencode: None,
current_provider_openclaw: None, current_provider_openclaw: None,
current_provider_hermes: 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_sync_method: SyncMethod::default(),
skill_storage_location: SkillStorageLocation::default(), skill_storage_location: SkillStorageLocation::default(),
webdav_sync: None, webdav_sync: None,
@@ -808,13 +614,6 @@ impl AppSettings {
.filter(|s| !s.is_empty()) .filter(|s| !s.is_empty())
.map(|s| s.to_string()); .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 self.language = self
.language .language
.as_ref() .as_ref()
@@ -873,9 +672,31 @@ fn save_settings_file(settings: &AppSettings) -> Result<(), AppError> {
fs::create_dir_all(parent).map_err(|e| AppError::io(parent, e))?; fs::create_dir_all(parent).map_err(|e| AppError::io(parent, e))?;
} }
let json = serde_json::to_vec_pretty(&normalized) let json = serde_json::to_string_pretty(&normalized)
.map_err(|e| AppError::JsonSerialize { source: e })?; .map_err(|e| AppError::JsonSerialize { source: e })?;
crate::config::atomic_write_durable(&path, &json, Some(0o600)) #[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(())
} }
static SETTINGS_STORE: OnceLock<RwLock<AppSettings>> = OnceLock::new(); static SETTINGS_STORE: OnceLock<RwLock<AppSettings>> = OnceLock::new();
@@ -884,7 +705,7 @@ fn settings_store() -> &'static RwLock<AppSettings> {
SETTINGS_STORE.get_or_init(|| RwLock::new(AppSettings::load_from_file())) SETTINGS_STORE.get_or_init(|| RwLock::new(AppSettings::load_from_file()))
} }
pub(crate) fn resolve_override_path(raw: &str) -> PathBuf { fn resolve_override_path(raw: &str) -> PathBuf {
if raw == "~" { if raw == "~" {
if let Some(home) = dirs::home_dir() { if let Some(home) = dirs::home_dir() {
return home; return home;
@@ -921,17 +742,17 @@ pub fn get_settings_for_frontend() -> AppSettings {
s3.secret_access_key.clear(); s3.secret_access_key.clear();
} }
settings.webdav_backup = None; settings.webdav_backup = None;
settings.pi_gateway_token = None;
settings settings
} }
pub fn update_settings(mut new_settings: AppSettings) -> Result<(), AppError> { pub fn update_settings(mut new_settings: AppSettings) -> Result<(), AppError> {
new_settings.normalize_paths(); new_settings.normalize_paths();
save_settings_file(&new_settings)?;
let mut guard = settings_store().write().unwrap_or_else(|e| { let mut guard = settings_store().write().unwrap_or_else(|e| {
log::warn!("设置锁已毒化,使用恢复值: {e}"); log::warn!("设置锁已毒化,使用恢复值: {e}");
e.into_inner() e.into_inner()
}); });
save_settings_file(&new_settings)?;
*guard = new_settings; *guard = new_settings;
Ok(()) Ok(())
} }
@@ -1112,83 +933,6 @@ pub fn get_hermes_override_dir() -> Option<PathBuf> {
.map(|p| resolve_override_path(p)) .map(|p| resolve_override_path(p))
} }
pub fn get_pi_override_dir() -> Option<PathBuf> {
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<GatewayToken, AppError> {
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<GatewayToken, AppError> {
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<GatewayToken, AppError> {
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_default_cost_multiplier() -> String {
get_pi_proxy_settings().default_cost_multiplier
}
pub(crate) fn set_pi_default_cost_multiplier(value: &str) -> Result<(), AppError> {
crate::database::validate_cost_multiplier(value)?;
let value = value.trim().to_string();
mutate_settings(move |settings| settings.pi_proxy.default_cost_multiplier = value)
}
pub(crate) fn get_pi_pricing_model_source() -> String {
get_pi_proxy_settings().pricing_model_source
}
pub(crate) fn set_pi_pricing_model_source(value: &str) -> Result<(), AppError> {
let value = crate::database::validate_pricing_source(value)?.to_string();
mutate_settings(move |settings| settings.pi_proxy.pricing_model_source = value)
}
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 { pub fn preserve_codex_official_auth_on_switch() -> bool {
settings_store() settings_store()
.read() .read()
@@ -1226,7 +970,6 @@ pub fn get_current_provider(app_type: &AppType) -> Option<String> {
AppType::OpenCode => settings.current_provider_opencode.clone(), AppType::OpenCode => settings.current_provider_opencode.clone(),
AppType::OpenClaw => settings.current_provider_openclaw.clone(), AppType::OpenClaw => settings.current_provider_openclaw.clone(),
AppType::Hermes => settings.current_provider_hermes.clone(), AppType::Hermes => settings.current_provider_hermes.clone(),
AppType::Pi => settings.current_provider_pi.clone(),
} }
} }
@@ -1245,7 +988,6 @@ pub fn set_current_provider(app_type: &AppType, id: Option<&str>) -> Result<(),
AppType::OpenCode => settings.current_provider_opencode = id_owned.clone(), AppType::OpenCode => settings.current_provider_opencode = id_owned.clone(),
AppType::OpenClaw => settings.current_provider_openclaw = id_owned.clone(), AppType::OpenClaw => settings.current_provider_openclaw = id_owned.clone(),
AppType::Hermes => settings.current_provider_hermes = id_owned.clone(), AppType::Hermes => settings.current_provider_hermes = id_owned.clone(),
AppType::Pi => settings.current_provider_pi = id_owned.clone(),
}) })
} }
@@ -1406,55 +1148,6 @@ mod tests {
use super::*; use super::*;
use crate::app_config::AppType; use crate::app_config::AppType;
#[cfg(unix)]
#[test]
#[serial_test::serial]
fn saving_a_gateway_token_tightens_legacy_settings_permissions() {
use std::os::unix::fs::PermissionsExt;
struct EnvGuard(Option<std::ffi::OsString>);
impl Drop for EnvGuard {
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 temp = tempfile::tempdir().expect("tempdir");
let _home = EnvGuard(std::env::var_os("CC_SWITCH_TEST_HOME"));
std::env::set_var("CC_SWITCH_TEST_HOME", temp.path());
let path = AppSettings::settings_path().expect("settings path");
fs::create_dir_all(path.parent().expect("settings parent")).expect("settings directory");
fs::write(&path, b"{}").expect("legacy settings");
fs::set_permissions(&path, fs::Permissions::from_mode(0o644))
.expect("make legacy settings permissive");
let token = GatewayToken::generate();
let expected = token.expose().to_string();
let settings = AppSettings {
pi_gateway_token: Some(token),
..AppSettings::default()
};
save_settings_file(&settings).expect("save sensitive settings");
assert_eq!(
fs::metadata(&path)
.expect("settings metadata")
.permissions()
.mode()
& 0o777,
0o600
);
assert!(
fs::read_to_string(&path)
.expect("settings body")
.contains(&expected),
"the permission assertion must exercise a persisted Pi gateway token"
);
}
#[test] #[test]
fn visible_apps_old_settings_default_claude_desktop_visible() { fn visible_apps_old_settings_default_claude_desktop_visible() {
let visible: VisibleApps = serde_json::from_value(serde_json::json!({ let visible: VisibleApps = serde_json::from_value(serde_json::json!({
@@ -1468,10 +1161,6 @@ mod tests {
.expect("visible apps"); .expect("visible apps");
assert!(visible.is_visible(&AppType::ClaudeDesktop)); 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] #[test]
-1
View File
@@ -3,7 +3,6 @@ use crate::services::{ProxyService, UsageCache};
use std::sync::Arc; use std::sync::Arc;
/// 全局应用状态 /// 全局应用状态
#[derive(Clone)]
pub struct AppState { pub struct AppState {
pub db: Arc<Database>, pub db: Arc<Database>,
pub proxy_service: ProxyService, pub proxy_service: ProxyService,
+59 -72
View File
@@ -7,14 +7,13 @@ use std::fs;
use serde_json::json; use serde_json::json;
use cc_switch_lib::{ use cc_switch_lib::{
AppType, InstalledSkill, McpServer, McpService, NewProviderAggregate, ProfilePayload, AppType, InstalledSkill, McpServer, McpService, ProfilePayload, ProfileScope, ProfileService,
ProfileScope, ProfileService, Prompt, PromptService, Provider, ProviderService, SkillApps, Prompt, PromptService, Provider, ProviderService, SkillApps, SkillService,
SkillService,
}; };
#[path = "support.rs"] #[path = "support.rs"]
mod support; mod support;
use support::{create_test_state, ensure_test_home, new_provider_input, reset_test_fs, test_mutex}; use support::{create_test_state, ensure_test_home, reset_test_fs, test_mutex};
fn claude_provider(id: &str, token: &str) -> Provider { fn claude_provider(id: &str, token: &str) -> Provider {
Provider::with_id( Provider::with_id(
@@ -108,41 +107,34 @@ fn profile_snapshot_apply_roundtrip_restores_configuration() {
let state = create_test_state().expect("create test state"); let state = create_test_state().expect("create test state");
// ---- 种子数据:2 个 Claude 供应商(p1 为当前)+ 2 个 MCP + 1 个 Skill + 2 个 Prompt ---- // ---- 种子数据:2 个 Claude 供应商(p1 为当前)+ 2 个 MCP + 1 个 Skill + 2 个 Prompt ----
ProviderService::add( state
&state, .db
AppType::Claude, .save_provider(AppType::Claude.as_str(), &claude_provider("p1", "key-1"))
new_provider_input(claude_provider("p1", "key-1")), .expect("save provider p1");
false, state
) .db
.expect("create provider p1"); .save_provider(AppType::Claude.as_str(), &claude_provider("p2", "key-2"))
ProviderService::add( .expect("save provider p2");
&state,
AppType::Claude,
new_provider_input(claude_provider("p2", "key-2")),
false,
)
.expect("create provider p2");
state state
.db .db
.set_current_provider(AppType::Claude.as_str(), "p1") .set_current_provider(AppType::Claude.as_str(), "p1")
.expect("set current provider p1"); .expect("set current provider p1");
// Claude Desktop 只有供应商一个活跃维度(MCP/Skills/Prompt 对它不适用) // Claude Desktop 只有供应商一个活跃维度(MCP/Skills/Prompt 对它不适用)
for provider in [ state
desktop_provider("d1", "dk-1"), .db
desktop_provider("d2", "dk-2"), .save_provider(
] { AppType::ClaudeDesktop.as_str(),
state &desktop_provider("d1", "dk-1"),
.db )
.create_provider( .expect("save desktop provider d1");
NewProviderAggregate::from_input( state
AppType::ClaudeDesktop.as_str(), .db
new_provider_input(provider), .save_provider(
) AppType::ClaudeDesktop.as_str(),
.expect("build typed desktop create"), &desktop_provider("d2", "dk-2"),
) )
.expect("create desktop provider"); .expect("save desktop provider d2");
}
state state
.db .db
.set_current_provider(AppType::ClaudeDesktop.as_str(), "d1") .set_current_provider(AppType::ClaudeDesktop.as_str(), "d1")
@@ -295,13 +287,10 @@ fn shared_profile_sides_are_isolated_and_mergeable() {
let state = create_test_state().expect("create test state"); let state = create_test_state().expect("create test state");
// 种子:Claude 侧有当前供应商 + 启用的 MCP // 种子:Claude 侧有当前供应商 + 启用的 MCP
ProviderService::add( state
&state, .db
AppType::Claude, .save_provider(AppType::Claude.as_str(), &claude_provider("p1", "key-1"))
new_provider_input(claude_provider("p1", "key-1")), .expect("save provider p1");
false,
)
.expect("create provider p1");
state state
.db .db
.set_current_provider(AppType::Claude.as_str(), "p1") .set_current_provider(AppType::Claude.as_str(), "p1")
@@ -507,20 +496,14 @@ fn switching_profile_autosaves_previous_profile_state() {
let state = create_test_state().expect("create test state"); let state = create_test_state().expect("create test state");
// ---- 种子:Claude 侧两套供应商 / MCP / Prompt ---- // ---- 种子:Claude 侧两套供应商 / MCP / Prompt ----
ProviderService::add( state
&state, .db
AppType::Claude, .save_provider(AppType::Claude.as_str(), &claude_provider("p1", "key-1"))
new_provider_input(claude_provider("p1", "key-1")), .expect("save provider p1");
false, state
) .db
.expect("create provider p1"); .save_provider(AppType::Claude.as_str(), &claude_provider("p2", "key-2"))
ProviderService::add( .expect("save provider p2");
&state,
AppType::Claude,
new_provider_input(claude_provider("p2", "key-2")),
false,
)
.expect("create provider p2");
state state
.db .db
.set_current_provider(AppType::Claude.as_str(), "p1") .set_current_provider(AppType::Claude.as_str(), "p1")
@@ -682,13 +665,17 @@ fn profile_switch_auto_disables_takeover_before_apply() {
// ---- 两个 Claude 供应商:custom1 与 custom2 ---- // ---- 两个 Claude 供应商:custom1 与 custom2 ----
let mut custom1 = claude_provider("custom1", "custom-key-1"); let mut custom1 = claude_provider("custom1", "custom-key-1");
custom1.category = Some("custom".to_string()); custom1.category = Some("custom".to_string());
ProviderService::add(&state, AppType::Claude, new_provider_input(custom1), false) state
.expect("create custom1 provider"); .db
.save_provider(AppType::Claude.as_str(), &custom1)
.expect("save custom1 provider");
let mut custom2 = claude_provider("custom2", "custom-key-2"); let mut custom2 = claude_provider("custom2", "custom-key-2");
custom2.category = Some("custom".to_string()); custom2.category = Some("custom".to_string());
ProviderService::add(&state, AppType::Claude, new_provider_input(custom2), false) state
.expect("create custom2 provider"); .db
.save_provider(AppType::Claude.as_str(), &custom2)
.expect("save custom2 provider");
// 初始状态:custom1 + 代理接管 // 初始状态:custom1 + 代理接管
ProviderService::switch(&state, AppType::Claude, "custom1").expect("switch to custom1"); ProviderService::switch(&state, AppType::Claude, "custom1").expect("switch to custom1");
@@ -770,20 +757,20 @@ fn claude_desktop_profile_scope_is_independent() {
let state = create_test_state().expect("create test state"); let state = create_test_state().expect("create test state");
ProviderService::add( state
&state, .db
AppType::ClaudeDesktop, .save_provider(
new_provider_input(desktop_provider("d1", "dk-1")), AppType::ClaudeDesktop.as_str(),
false, &desktop_provider("d1", "dk-1"),
) )
.expect("create desktop provider d1"); .expect("save desktop provider d1");
ProviderService::add( state
&state, .db
AppType::ClaudeDesktop, .save_provider(
new_provider_input(desktop_provider("d2", "dk-2")), AppType::ClaudeDesktop.as_str(),
false, &desktop_provider("d2", "dk-2"),
) )
.expect("create desktop provider d2"); .expect("save desktop provider d2");
state state
.db .db
.set_current_provider(AppType::ClaudeDesktop.as_str(), "d1") .set_current_provider(AppType::ClaudeDesktop.as_str(), "d1")
+13 -13
View File
@@ -12,7 +12,7 @@ mod support;
use std::collections::HashMap; use std::collections::HashMap;
use support::{ use support::{
create_test_state, create_test_state_with_config, enable_codex_official_auth_preservation, create_test_state, create_test_state_with_config, enable_codex_official_auth_preservation,
ensure_test_home, new_provider_input, reset_test_fs, test_mutex, ensure_test_home, reset_test_fs, test_mutex,
}; };
fn settings_path(home: &Path) -> PathBuf { fn settings_path(home: &Path) -> PathBuf {
@@ -64,18 +64,18 @@ fn grokbuild_import_and_switch_write_live_config() {
); );
let next_config = grokbuild_config("Relay", "https://new.example/v1", "new-key"); let next_config = grokbuild_config("Relay", "https://new.example/v1", "new-key");
ProviderService::add( state
&state, .db
AppType::GrokBuild, .save_provider(
new_provider_input(Provider::with_id( AppType::GrokBuild.as_str(),
"relay".to_string(), &Provider::with_id(
"Relay".to_string(), "relay".to_string(),
json!({ "config": next_config }), "Relay".to_string(),
None, json!({ "config": next_config }),
)), None,
false, ),
) )
.expect("create second Grok Build provider"); .expect("save second Grok Build provider");
switch_provider_test_hook(&state, AppType::GrokBuild, "relay") switch_provider_test_hook(&state, AppType::GrokBuild, "relay")
.expect("switch Grok Build provider"); .expect("switch Grok Build provider");
+5 -3
View File
@@ -9,7 +9,7 @@ use cc_switch_lib::{
mod support; mod support;
use support::{ use support::{
create_test_state, create_test_state_with_config, enable_codex_official_auth_preservation, create_test_state, create_test_state_with_config, enable_codex_official_auth_preservation,
ensure_test_home, new_provider_input, reset_test_fs, test_mutex, ensure_test_home, reset_test_fs, test_mutex,
}; };
fn sanitize_provider_name(name: &str) -> String { fn sanitize_provider_name(name: &str) -> String {
@@ -3084,8 +3084,10 @@ fn recover_from_crash_without_backup_cleans_placeholder_instead_of_writing_it_ba
taken_over_live.clone(), taken_over_live.clone(),
None, None,
); );
ProviderService::add(&state, AppType::Claude, new_provider_input(provider), false) state
.expect("create placeholder provider"); .db
.save_provider(AppType::Claude.as_str(), &provider)
.expect("save placeholder provider");
state state
.db .db
.set_current_provider(AppType::Claude.as_str(), "default") .set_current_provider(AppType::Claude.as_str(), "default")
+1 -25
View File
@@ -1,31 +1,7 @@
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, OnceLock}; use std::sync::{Arc, Mutex, OnceLock};
use cc_switch_lib::{ use cc_switch_lib::{update_settings, AppSettings, AppState, Database, MultiAppConfig};
update_settings, AppSettings, AppState, Database, MultiAppConfig, Provider,
ProviderMutationInput,
};
/// Build the public write DTO explicitly for integration tests. Keeping this
/// conversion test-only avoids reintroducing a production `From<Provider>`
/// path from hydrated read projections to provider mutations.
#[allow(dead_code)]
pub fn new_provider_input(provider: Provider) -> ProviderMutationInput {
ProviderMutationInput {
id: provider.id,
name: provider.name,
settings_config: provider.settings_config,
website_url: provider.website_url,
category: provider.category,
created_at: provider.created_at,
sort_index: provider.sort_index,
notes: provider.notes,
meta: provider.meta,
icon: provider.icon,
icon_color: provider.icon_color,
in_failover_queue: provider.in_failover_queue,
}
}
/// 为测试设置隔离的 HOME 目录,避免污染真实用户数据。 /// 为测试设置隔离的 HOME 目录,避免污染真实用户数据。
pub fn ensure_test_home() -> &'static Path { pub fn ensure_test_home() -> &'static Path {

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