mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-06 21:52:51 +08:00
Merge remote-tracking branch 'origin/main' into pr-385-merge
# Conflicts: # server/tests/interrupt.rs
This commit is contained in:
+2
-1
@@ -48,7 +48,7 @@ similar = "2"
|
||||
sqlx = { version = "0.8", features = ["runtime-tokio", "sqlite"] }
|
||||
thiserror = "2"
|
||||
time = "0.3"
|
||||
tokio = { version = "1", features = ["macros", "rt-multi-thread", "signal", "sync", "time", "net"] }
|
||||
tokio = { version = "1", features = ["fs", "io-util", "macros", "process", "rt-multi-thread", "signal", "sync", "time", "net"] }
|
||||
tokio-stream = { version = "0.1", features = ["sync"] }
|
||||
tokio-util = "0.7"
|
||||
tracing = "0.1"
|
||||
@@ -57,6 +57,7 @@ tower-http = { version = "0.6", features = ["cors", "decompression-gzip", "fs"]
|
||||
url = "2"
|
||||
uuid = { version = "1", features = ["v4"] }
|
||||
x509-parser = "0.18"
|
||||
zip = { version = "4", default-features = false, features = ["deflate"] }
|
||||
[build-dependencies]
|
||||
prost-build = "0.13"
|
||||
protoc-bin-vendored = "3"
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
-- llm_calls 是历史记录:model_hash 现在既可指向内置 model_configs,
|
||||
-- 也可携带插件稳定模型 ID(plugin:<plugin>/<provider>/<model>)。
|
||||
-- 去掉指向 model_configs 的外键;SQLite 不支持删除约束,按整表重建执行。
|
||||
PRAGMA defer_foreign_keys = ON;
|
||||
|
||||
CREATE TABLE llm_calls_new (
|
||||
call_id TEXT PRIMARY KEY,
|
||||
run_id TEXT NOT NULL,
|
||||
conversation_id TEXT NOT NULL,
|
||||
provider_call_index INTEGER NOT NULL,
|
||||
model_hash TEXT,
|
||||
provider_type TEXT NOT NULL,
|
||||
provider_url TEXT NOT NULL,
|
||||
request_type TEXT NOT NULL,
|
||||
request_url TEXT NOT NULL,
|
||||
model_id TEXT NOT NULL,
|
||||
display_name TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
finish_reason TEXT,
|
||||
created_at_ms INTEGER NOT NULL,
|
||||
request_started_at_ms INTEGER,
|
||||
response_headers_at_ms INTEGER,
|
||||
first_event_at_ms INTEGER,
|
||||
first_text_at_ms INTEGER,
|
||||
finished_at_ms INTEGER,
|
||||
queue_ms INTEGER,
|
||||
ttfb_ms INTEGER,
|
||||
ttft_ms INTEGER,
|
||||
duration_ms INTEGER,
|
||||
input_tokens INTEGER,
|
||||
output_tokens INTEGER,
|
||||
total_tokens INTEGER,
|
||||
cache_read_tokens INTEGER,
|
||||
cache_write_tokens INTEGER,
|
||||
reasoning_tokens INTEGER,
|
||||
usage_json TEXT,
|
||||
message_count INTEGER NOT NULL,
|
||||
tool_count INTEGER NOT NULL,
|
||||
request_bytes INTEGER,
|
||||
response_bytes INTEGER NOT NULL DEFAULT 0,
|
||||
stream_event_count INTEGER NOT NULL DEFAULT 0,
|
||||
http_status INTEGER,
|
||||
error_kind TEXT,
|
||||
error_message TEXT,
|
||||
detailed INTEGER NOT NULL,
|
||||
reasoning_effort TEXT,
|
||||
fast INTEGER NOT NULL DEFAULT 0 CHECK (fast IN (0, 1)),
|
||||
first_valid_response_at_ms INTEGER,
|
||||
ttfr_ms INTEGER
|
||||
);
|
||||
|
||||
INSERT INTO llm_calls_new (
|
||||
call_id, run_id, conversation_id, provider_call_index, model_hash, provider_type,
|
||||
provider_url, request_type, request_url, model_id, display_name, status, finish_reason,
|
||||
created_at_ms, request_started_at_ms, response_headers_at_ms, first_event_at_ms,
|
||||
first_text_at_ms, finished_at_ms, queue_ms, ttfb_ms, ttft_ms, duration_ms,
|
||||
input_tokens, output_tokens, total_tokens, cache_read_tokens, cache_write_tokens,
|
||||
reasoning_tokens, usage_json, message_count, tool_count, request_bytes, response_bytes,
|
||||
stream_event_count, http_status, error_kind, error_message, detailed, reasoning_effort, fast,
|
||||
first_valid_response_at_ms, ttfr_ms
|
||||
)
|
||||
SELECT
|
||||
call_id, run_id, conversation_id, provider_call_index, model_hash, provider_type,
|
||||
provider_url, request_type, request_url, model_id, display_name, status, finish_reason,
|
||||
created_at_ms, request_started_at_ms, response_headers_at_ms, first_event_at_ms,
|
||||
first_text_at_ms, finished_at_ms, queue_ms, ttfb_ms, ttft_ms, duration_ms,
|
||||
input_tokens, output_tokens, total_tokens, cache_read_tokens, cache_write_tokens,
|
||||
reasoning_tokens, usage_json, message_count, tool_count, request_bytes, response_bytes,
|
||||
stream_event_count, http_status, error_kind, error_message, detailed, reasoning_effort, fast,
|
||||
first_valid_response_at_ms, ttfr_ms
|
||||
FROM llm_calls;
|
||||
|
||||
CREATE TABLE llm_call_requests_new (
|
||||
call_id TEXT PRIMARY KEY,
|
||||
headers_json TEXT NOT NULL,
|
||||
body_json TEXT NOT NULL,
|
||||
byte_count INTEGER NOT NULL,
|
||||
FOREIGN KEY(call_id) REFERENCES llm_calls_new(call_id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
INSERT INTO llm_call_requests_new(call_id, headers_json, body_json, byte_count)
|
||||
SELECT call_id, headers_json, body_json, byte_count FROM llm_call_requests;
|
||||
|
||||
CREATE TABLE llm_call_response_chunks_new (
|
||||
call_id TEXT NOT NULL,
|
||||
seq INTEGER NOT NULL,
|
||||
received_offset_ms INTEGER NOT NULL,
|
||||
data BLOB NOT NULL,
|
||||
byte_count INTEGER NOT NULL,
|
||||
PRIMARY KEY(call_id, seq),
|
||||
FOREIGN KEY(call_id) REFERENCES llm_calls_new(call_id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
INSERT INTO llm_call_response_chunks_new(call_id, seq, received_offset_ms, data, byte_count)
|
||||
SELECT call_id, seq, received_offset_ms, data, byte_count FROM llm_call_response_chunks;
|
||||
|
||||
DROP TABLE llm_call_requests;
|
||||
DROP TABLE llm_call_response_chunks;
|
||||
DROP TABLE llm_calls;
|
||||
|
||||
ALTER TABLE llm_calls_new RENAME TO llm_calls;
|
||||
ALTER TABLE llm_call_requests_new RENAME TO llm_call_requests;
|
||||
ALTER TABLE llm_call_response_chunks_new RENAME TO llm_call_response_chunks;
|
||||
|
||||
CREATE INDEX llm_calls_created ON llm_calls(created_at_ms DESC);
|
||||
CREATE INDEX llm_calls_run ON llm_calls(run_id, provider_call_index);
|
||||
CREATE INDEX llm_calls_model ON llm_calls(model_hash, created_at_ms DESC);
|
||||
@@ -0,0 +1,4 @@
|
||||
-- Custom provider-group display name shared by models with the same upstream host.
|
||||
-- NULL means no custom name; the UI falls back to the base_url hostname and the
|
||||
-- Cursor model picker badge falls back to the model type label.
|
||||
ALTER TABLE model_configs ADD COLUMN group_name TEXT;
|
||||
File diff suppressed because one or more lines are too long
|
After Width: | Height: | Size: 43 KiB |
@@ -0,0 +1,383 @@
|
||||
import type {
|
||||
JsonValue,
|
||||
NetworkEventStream,
|
||||
NetworkResponse,
|
||||
PluginContext,
|
||||
} from "cursor-byok:plugin";
|
||||
import type { LlmRequest, ModelEvent } from "cursor-byok:provider";
|
||||
import type { ResourceSnapshot } from "cursor-byok:resource";
|
||||
import { codexDeviceOAuth } from "./oauth.ts";
|
||||
import { parseOfficialModels } from "./models.ts";
|
||||
import { codexProvider, isQuotaError } from "./provider.ts";
|
||||
import {
|
||||
accountIdentity,
|
||||
credentialDraft,
|
||||
parseCodexUsage,
|
||||
parseCredentialFiles,
|
||||
presentAccount,
|
||||
quotaState,
|
||||
RESOURCE_TYPE,
|
||||
} from "./resources.ts";
|
||||
|
||||
function assert(condition: unknown, message = "assertion failed"): asserts condition {
|
||||
if (!condition) throw new Error(message);
|
||||
}
|
||||
|
||||
function assertEquals(actual: unknown, expected: unknown): void {
|
||||
const left = JSON.stringify(actual);
|
||||
const right = JSON.stringify(expected);
|
||||
if (left !== right) throw new Error(`expected ${right}, received ${left}`);
|
||||
}
|
||||
|
||||
function jwt(payload: Record<string, unknown>): string {
|
||||
const encoded = btoa(JSON.stringify(payload)).replace(/=/g, "").replace(/\+/g, "-").replace(
|
||||
/\//g,
|
||||
"_",
|
||||
);
|
||||
return `header.${encoded}.signature`;
|
||||
}
|
||||
|
||||
type RequestInit = { body?: string; headers?: Record<string, string> };
|
||||
type FetchHandler = (url: string, init?: RequestInit) => NetworkResponse;
|
||||
type StreamHandler = (url: string, init?: RequestInit) => NetworkEventStream;
|
||||
|
||||
function context(handlers: { fetch?: FetchHandler; stream?: StreamHandler }): PluginContext {
|
||||
return {
|
||||
network: {
|
||||
fetch: (url, init) => {
|
||||
if (!handlers.fetch) throw new Error("fetch was not expected");
|
||||
return Promise.resolve(handlers.fetch(url, init));
|
||||
},
|
||||
stream: (url, init) => {
|
||||
if (!handlers.stream) throw new Error("stream was not expected");
|
||||
return Promise.resolve(handlers.stream(url, init));
|
||||
},
|
||||
},
|
||||
signal: new AbortController().signal,
|
||||
};
|
||||
}
|
||||
|
||||
function snapshot(privateData: JsonValue): ResourceSnapshot {
|
||||
return {
|
||||
id: "resource-1",
|
||||
type: RESOURCE_TYPE,
|
||||
key: "codex:acct-1",
|
||||
privateData,
|
||||
state: { status: "ready" },
|
||||
};
|
||||
}
|
||||
|
||||
async function* sse(lines: string[]): AsyncGenerator<string> {
|
||||
for (const line of lines) yield line;
|
||||
}
|
||||
|
||||
function request(): LlmRequest {
|
||||
return {
|
||||
instructions: "You are a coding assistant.",
|
||||
messages: [{ role: "user", content: [{ type: "text", text: "hi" }] }],
|
||||
tools: [],
|
||||
reasoning: { enabled: true, effort: "medium" },
|
||||
latency: "fast",
|
||||
maxOutputTokens: 128_000,
|
||||
cacheKey: "conversation-1",
|
||||
};
|
||||
}
|
||||
|
||||
Deno.test("account identity prioritizes ChatGPT account ID and drafts keep tokens private-side", async () => {
|
||||
const token = jwt({
|
||||
"https://api.openai.com/auth": { chatgpt_account_id: "acct-1" },
|
||||
sub: "subject-1",
|
||||
email: "person@example.com",
|
||||
});
|
||||
assertEquals(await accountIdentity(token), {
|
||||
key: "codex:acct-1",
|
||||
displayName: "person@example.com",
|
||||
});
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
assertEquals(draft.key, "codex:acct-1");
|
||||
const view = presentAccount(snapshot(draft.privateData));
|
||||
assert(!JSON.stringify(view).includes(token), "resource view exposed an access token");
|
||||
assertEquals(view.displayName, "person@example.com");
|
||||
});
|
||||
|
||||
Deno.test("credential import accepts Codex auth JSON files", () => {
|
||||
const { credentials, warnings } = parseCredentialFiles([
|
||||
{
|
||||
name: "auth.json",
|
||||
content: JSON.stringify({
|
||||
tokens: {
|
||||
access_token: "access-secret",
|
||||
refresh_token: "refresh-secret",
|
||||
id_token: jwt({ email: "person@example.com" }),
|
||||
},
|
||||
}),
|
||||
},
|
||||
{ name: "broken.json", content: "{not json" },
|
||||
]);
|
||||
assertEquals(credentials, [{
|
||||
accessToken: "access-secret",
|
||||
refreshToken: "refresh-secret",
|
||||
displayName: "person@example.com",
|
||||
}]);
|
||||
assertEquals(warnings, ["broken.json: not valid JSON"]);
|
||||
});
|
||||
|
||||
Deno.test("usage maps secondary to weekly and primary to five-hour quota", () => {
|
||||
const quota = parseCodexUsage({
|
||||
plan_type: "plus",
|
||||
rate_limit: {
|
||||
primary_window: { used_percent: 80, reset_at: 1_800_000_000 },
|
||||
secondary_window: { used_percent: 25, reset_at: 1_900_000_000 },
|
||||
},
|
||||
}, 1_700_000_000_000);
|
||||
assertEquals(quota.planLabel, "ChatGPT Plus");
|
||||
assertEquals(quota.weekly?.remainingPercent, 75);
|
||||
assertEquals(quota.fiveHour?.remainingPercent, 20);
|
||||
assertEquals(quota.weekly?.resetAtMs, 1_900_000_000_000);
|
||||
assertEquals(quotaState(quota, 1_700_000_000_000), { status: "ready" });
|
||||
});
|
||||
|
||||
Deno.test("exhausted quota projects a cooling state until the latest reset", () => {
|
||||
const quota = parseCodexUsage({
|
||||
rate_limit: {
|
||||
primary_window: { used_percent: 100, reset_at: 1_800_000_000 },
|
||||
secondary_window: { used_percent: 100, reset_at: 1_900_000_000 },
|
||||
},
|
||||
}, 1_700_000_000_000);
|
||||
assertEquals(quotaState(quota, 1_700_000_000_000), {
|
||||
status: "cooling",
|
||||
retryAtMs: 1_900_000_000_000,
|
||||
message: "ChatGPT quota is exhausted",
|
||||
});
|
||||
});
|
||||
|
||||
Deno.test("official model discovery excludes hidden models and puts the default first", () => {
|
||||
const models = parseOfficialModels({
|
||||
default_model: "gpt-second",
|
||||
models: [
|
||||
{
|
||||
slug: "gpt-first",
|
||||
display_name: "GPT First",
|
||||
supported_in_api: true,
|
||||
visibility: "list",
|
||||
supported_reasoning_levels: [
|
||||
{ effort: "low", description: "Fast responses" },
|
||||
{ effort: "medium", description: "Balanced" },
|
||||
],
|
||||
},
|
||||
{ slug: "gpt-second", supported_in_api: true, visibility: "list" },
|
||||
{ slug: "gpt-hidden", supported_in_api: true, visibility: "hidden" },
|
||||
{ slug: "gpt-internal", supported_in_api: false, visibility: "list" },
|
||||
],
|
||||
});
|
||||
assertEquals(models.map((model) => model.id), ["gpt-second", "gpt-first"]);
|
||||
assertEquals(models[1].capabilities, { images: true });
|
||||
assertEquals(models[1].privateData, { reasoningEfforts: ["low", "medium"] });
|
||||
});
|
||||
|
||||
Deno.test("device OAuth begins with a host-held session and completes with a resource draft", async () => {
|
||||
const accessToken = jwt({
|
||||
"https://api.openai.com/auth": { chatgpt_account_id: "acct-oauth" },
|
||||
email: "oauth@example.com",
|
||||
});
|
||||
let requestNumber = 0;
|
||||
const flowContext = context({
|
||||
fetch: (url, init) => {
|
||||
requestNumber += 1;
|
||||
if (requestNumber === 1) {
|
||||
assertEquals(url, "https://auth.openai.com/api/accounts/deviceauth/usercode");
|
||||
return {
|
||||
status: 200,
|
||||
headers: {},
|
||||
body: JSON.stringify({
|
||||
device_auth_id: "private-device-id",
|
||||
user_code: "ABCD-EFGH",
|
||||
expires_in: 900,
|
||||
interval: 5,
|
||||
}),
|
||||
};
|
||||
}
|
||||
if (requestNumber === 2) {
|
||||
assertEquals(url, "https://auth.openai.com/api/accounts/deviceauth/token");
|
||||
return {
|
||||
status: 200,
|
||||
headers: {},
|
||||
body: JSON.stringify({
|
||||
authorization_code: "authorization-code",
|
||||
code_verifier: "pkce-verifier",
|
||||
}),
|
||||
};
|
||||
}
|
||||
assertEquals(url, "https://auth.openai.com/oauth/token");
|
||||
assert(init?.body?.includes("grant_type=authorization_code"));
|
||||
assert(init?.body?.includes("code_verifier=pkce-verifier"));
|
||||
return {
|
||||
status: 200,
|
||||
headers: {},
|
||||
body: JSON.stringify({ access_token: accessToken, refresh_token: "refresh-secret" }),
|
||||
};
|
||||
},
|
||||
});
|
||||
|
||||
const begun = await codexDeviceOAuth.begin(flowContext);
|
||||
assertEquals(begun.userCode, "ABCD-EFGH");
|
||||
assertEquals(begun.pollIntervalMs, 5000);
|
||||
|
||||
const polled = await codexDeviceOAuth.poll(begun.session, flowContext);
|
||||
assert(polled.status === "completed", `expected completed, received ${polled.status}`);
|
||||
assertEquals(polled.resources[0].key, "codex:acct-oauth");
|
||||
assertEquals(requestNumber, 3);
|
||||
});
|
||||
|
||||
Deno.test("invoke streams normalized events from the Codex Responses API", async () => {
|
||||
const token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
let requestBody = "";
|
||||
let requestHeaders: Record<string, string> = {};
|
||||
const events: ModelEvent[] = [];
|
||||
const result = await codexProvider.invoke(
|
||||
{
|
||||
model: {
|
||||
id: "gpt-test",
|
||||
displayName: "GPT Test",
|
||||
privateData: { reasoningEfforts: ["medium"] },
|
||||
},
|
||||
resource: snapshot(draft.privateData),
|
||||
request: request(),
|
||||
},
|
||||
{ emit: (event) => events.push(event) },
|
||||
context({
|
||||
stream: (url, init) => {
|
||||
assertEquals(url, "https://chatgpt.com/backend-api/codex/responses");
|
||||
requestBody = init?.body ?? "";
|
||||
requestHeaders = init?.headers ?? {};
|
||||
return {
|
||||
status: 200,
|
||||
headers: {},
|
||||
lines: sse([
|
||||
'data: {"type":"response.output_text.delta","delta":"Hel"}',
|
||||
'data: {"type":"response.output_text.delta","delta":"lo"}',
|
||||
'data: {"type":"response.completed","response":{"usage":{"input_tokens":10,"output_tokens":2,"input_tokens_details":{"cached_tokens":4}}}}',
|
||||
]),
|
||||
};
|
||||
},
|
||||
}),
|
||||
);
|
||||
assertEquals(result, { status: "completed" });
|
||||
const body = JSON.parse(requestBody) as Record<string, unknown>;
|
||||
assertEquals(body.model, "gpt-test");
|
||||
assertEquals(body.store, false);
|
||||
assertEquals(body.reasoning, { summary: "auto", effort: "medium" });
|
||||
assertEquals(body.instructions, "You are a coding assistant.");
|
||||
assertEquals(body.include, ["reasoning.encrypted_content"]);
|
||||
assert(!("max_output_tokens" in body), "Codex endpoint rejects max_output_tokens");
|
||||
assertEquals(body.service_tier, "priority");
|
||||
assertEquals(body.prompt_cache_key, "conversation-1");
|
||||
// 缓存亲和头与 prompt_cache_key 同源。
|
||||
assertEquals(requestHeaders["session-id"], "conversation-1");
|
||||
assertEquals(requestHeaders["thread-id"], "conversation-1");
|
||||
assertEquals(requestHeaders["x-client-request-id"], "conversation-1");
|
||||
assertEquals(events, [
|
||||
{ type: "text-start" },
|
||||
{ type: "text-delta", text: "Hel" },
|
||||
{ type: "text-delta", text: "lo" },
|
||||
{
|
||||
type: "usage",
|
||||
usage: {
|
||||
inputTokens: 10,
|
||||
outputTokens: 2,
|
||||
totalTokens: null,
|
||||
cacheReadTokens: 4,
|
||||
cacheWriteTokens: null,
|
||||
reasoningTokens: null,
|
||||
},
|
||||
},
|
||||
{ type: "text-end" },
|
||||
{ type: "done", reason: "stop" },
|
||||
]);
|
||||
});
|
||||
|
||||
Deno.test("invoke streams incremental tool calls and replays reasoning items", async () => {
|
||||
const token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
const events: ModelEvent[] = [];
|
||||
const result = await codexProvider.invoke(
|
||||
{
|
||||
model: { id: "gpt-test", displayName: "GPT Test" },
|
||||
resource: snapshot(draft.privateData),
|
||||
request: request(),
|
||||
},
|
||||
{ emit: (event) => events.push(event) },
|
||||
context({
|
||||
stream: () => ({
|
||||
status: 200,
|
||||
headers: {},
|
||||
lines: sse([
|
||||
'data: {"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","call_id":"call-1","name":"read_file"}}',
|
||||
'data: {"type":"response.function_call_arguments.delta","output_index":0,"delta":"{\\"path\\":"}',
|
||||
'data: {"type":"response.function_call_arguments.delta","output_index":0,"delta":"\\"a.ts\\"}"}',
|
||||
'data: {"type":"response.output_item.done","output_index":0,"item":{"type":"function_call","call_id":"call-1","name":"read_file","arguments":"{\\"path\\":\\"a.ts\\"}"}}',
|
||||
'data: {"type":"response.output_item.done","output_index":1,"item":{"type":"reasoning","encrypted_content":"opaque"}}',
|
||||
'data: {"type":"response.completed","response":{}}',
|
||||
]),
|
||||
}),
|
||||
}),
|
||||
);
|
||||
assertEquals(result, { status: "completed" });
|
||||
assertEquals(events, [
|
||||
{ type: "tool-call-start", index: 0, callId: "call-1", name: "read_file" },
|
||||
{ type: "tool-call-arguments-delta", index: 0, delta: '{"path":' },
|
||||
{ type: "tool-call-arguments-delta", index: 0, delta: '"a.ts"}' },
|
||||
{ type: "tool-call-end", index: 0 },
|
||||
{
|
||||
type: "replay-state",
|
||||
providerKind: "openai_responses",
|
||||
value: { items: [{ type: "reasoning", encrypted_content: "opaque" }] },
|
||||
},
|
||||
{ type: "done", reason: "tool-use" },
|
||||
]);
|
||||
});
|
||||
|
||||
Deno.test("invoke maps quota failures to a cooling resource error", async () => {
|
||||
assert(!isQuotaError("429 rate_limit_reached"));
|
||||
assert(isQuotaError("429 usage_limit_reached: 5-hour limit"));
|
||||
const token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
const result = await codexProvider.invoke(
|
||||
{
|
||||
model: { id: "gpt-test", displayName: "GPT Test" },
|
||||
resource: snapshot(draft.privateData),
|
||||
request: request(),
|
||||
},
|
||||
{ emit: () => {} },
|
||||
context({
|
||||
stream: () => ({
|
||||
status: 429,
|
||||
headers: {},
|
||||
lines: sse(['{"detail":"usage_limit_reached","reset_after_seconds":600}']),
|
||||
}),
|
||||
}),
|
||||
);
|
||||
assert(result.status === "resource-error", `expected resource-error, received ${result.status}`);
|
||||
assert(result.patch.state?.status === "cooling", "quota failure should cool the resource");
|
||||
assert(
|
||||
result.patch.state.retryAtMs !== undefined && result.patch.state.retryAtMs > Date.now(),
|
||||
"cooling should carry the parsed reset time",
|
||||
);
|
||||
});
|
||||
@@ -0,0 +1,13 @@
|
||||
{
|
||||
"imports": {
|
||||
"cursor-byok:plugin": "../../../src/plugin/sdk/plugin.ts",
|
||||
"cursor-byok:provider": "../../../src/plugin/sdk/provider.ts",
|
||||
"cursor-byok:model": "../../../src/plugin/sdk/model.ts",
|
||||
"cursor-byok:resource": "../../../src/plugin/sdk/resource.ts",
|
||||
"cursor-byok:protocol/openai-responses": "../../../src/plugin/sdk/protocol/openai_responses.ts"
|
||||
},
|
||||
"fmt": {
|
||||
"lineWidth": 100,
|
||||
"exclude": ["assets"]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
import { defineProviderPlugin } from "cursor-byok:plugin";
|
||||
import { codexDeviceOAuth } from "./oauth.ts";
|
||||
import { codexProvider } from "./provider.ts";
|
||||
import { credentialImport, presentAccount, refreshAccount, RESOURCE_TYPE } from "./resources.ts";
|
||||
|
||||
export default defineProviderPlugin({
|
||||
providers: [codexProvider],
|
||||
resources: [{
|
||||
type: RESOURCE_TYPE,
|
||||
displayName: { "en-US": "ChatGPT accounts", "zh-CN": "ChatGPT 账号" },
|
||||
add: [codexDeviceOAuth],
|
||||
import: credentialImport,
|
||||
present: presentAccount,
|
||||
refresh: refreshAccount,
|
||||
}],
|
||||
});
|
||||
@@ -0,0 +1,128 @@
|
||||
import type { JsonValue } from "cursor-byok:plugin";
|
||||
import type { ModelDefinition, ModelSnapshot, ModelSupport } from "cursor-byok:model";
|
||||
import { accountData, accountHeaders } from "./resources.ts";
|
||||
|
||||
const MODELS_URL = "https://chatgpt.com/backend-api/codex/models?client_version=1.0.0";
|
||||
|
||||
function object(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value)
|
||||
? value as Record<string, unknown>
|
||||
: null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" && value.trim() ? value.trim() : null;
|
||||
}
|
||||
|
||||
function positiveInteger(value: unknown): number | null {
|
||||
const parsed = typeof value === "number"
|
||||
? value
|
||||
: typeof value === "string"
|
||||
? Number(value)
|
||||
: NaN;
|
||||
return Number.isFinite(parsed) && parsed > 0 ? Math.floor(parsed) : null;
|
||||
}
|
||||
|
||||
function parseReasoningEfforts(model: Record<string, unknown>): string[] {
|
||||
const source = model.supported_reasoning_levels ??
|
||||
model.supportedReasoningLevels ??
|
||||
model.reasoning_levels ??
|
||||
model.reasoningLevels ??
|
||||
model.supported_reasoning_efforts ??
|
||||
model.supportedReasoningEfforts ??
|
||||
model.reasoning_efforts ??
|
||||
model.reasoningEfforts;
|
||||
if (!Array.isArray(source)) return [];
|
||||
const values = source.flatMap((item) => {
|
||||
if (typeof item === "string") return [item.trim()];
|
||||
const entry = object(item);
|
||||
const value = text(entry?.effort ?? entry?.id ?? entry?.value ?? entry?.name);
|
||||
return value ? [value] : [];
|
||||
}).filter(Boolean);
|
||||
return [...new Set(values)];
|
||||
}
|
||||
|
||||
function modelId(value: unknown): string | null {
|
||||
if (typeof value === "string") return text(value);
|
||||
const model = object(value);
|
||||
return model ? text(model.slug ?? model.id ?? model.model ?? model.name) : null;
|
||||
}
|
||||
|
||||
export function parseOfficialModels(body: unknown): ModelDefinition[] {
|
||||
const root = object(body);
|
||||
const source = root?.models ?? root?.data ?? body;
|
||||
if (!Array.isArray(source)) {
|
||||
throw new Error("Codex model discovery response does not contain a model list");
|
||||
}
|
||||
const seen = new Set<string>();
|
||||
const models: ModelDefinition[] = [];
|
||||
for (const raw of source) {
|
||||
const model = object(raw);
|
||||
if (!model || model.supported_in_api === false || model.supportedInApi === false) continue;
|
||||
if (text(model.visibility)?.toLowerCase() === "hidden") continue;
|
||||
const id = modelId(model);
|
||||
if (!id || seen.has(id)) continue;
|
||||
seen.add(id);
|
||||
const efforts = parseReasoningEfforts(model);
|
||||
const description = text(model.description);
|
||||
const maxOutputTokens = positiveInteger(
|
||||
model.max_output_tokens ?? model.maxOutputTokens ?? model.max_completion_tokens ??
|
||||
model.maxCompletionTokens,
|
||||
);
|
||||
models.push({
|
||||
id,
|
||||
displayName: text(model.display_name ?? model.displayName ?? model.title ?? model.name) ??
|
||||
id,
|
||||
...(description ? { description } : {}),
|
||||
...(maxOutputTokens !== null ? { maxOutputTokens } : {}),
|
||||
capabilities: { images: true },
|
||||
privateData: { reasoningEfforts: efforts },
|
||||
});
|
||||
}
|
||||
const defaultModel = modelId(
|
||||
root?.default_model ??
|
||||
root?.defaultModel ??
|
||||
root?.default_model_slug ??
|
||||
root?.defaultModelSlug ??
|
||||
root?.primary_model ??
|
||||
root?.primaryModel,
|
||||
);
|
||||
// 把上游默认模型排在最前,让宿主自然选中它。
|
||||
if (defaultModel) {
|
||||
models.sort((left, right) =>
|
||||
Number(right.id === defaultModel) - Number(left.id === defaultModel)
|
||||
);
|
||||
}
|
||||
return models;
|
||||
}
|
||||
|
||||
export function reasoningEfforts(model: ModelSnapshot): string[] {
|
||||
const data = object(model.privateData);
|
||||
const efforts = data?.reasoningEfforts;
|
||||
return Array.isArray(efforts) ? efforts.filter((item) => typeof item === "string") : [];
|
||||
}
|
||||
|
||||
export const codexModels: ModelSupport = {
|
||||
list: async ({ resource }, context): Promise<ModelDefinition[]> => {
|
||||
if (!resource) throw new Error("add a ChatGPT account before syncing Codex models");
|
||||
const data = accountData(resource);
|
||||
const response = await context.network.fetch(MODELS_URL, {
|
||||
method: "GET",
|
||||
headers: accountHeaders(data),
|
||||
});
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
throw new Error(`Codex model discovery failed (HTTP ${response.status}): ${response.body}`);
|
||||
}
|
||||
let body: unknown;
|
||||
try {
|
||||
body = JSON.parse(response.body) as JsonValue;
|
||||
} catch {
|
||||
throw new Error("Codex model discovery returned invalid JSON");
|
||||
}
|
||||
const models = parseOfficialModels(body);
|
||||
if (models.length === 0) {
|
||||
throw new Error("Codex model discovery returned no supported models");
|
||||
}
|
||||
return models;
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,210 @@
|
||||
import type { JsonValue, PluginContext } from "cursor-byok:plugin";
|
||||
import type { OAuth2AddMethod, OAuth2Begin, OAuth2Poll } from "cursor-byok:resource";
|
||||
import { type CredentialCandidate, credentialDraft } from "./resources.ts";
|
||||
|
||||
const CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann";
|
||||
const DEVICE_CODE_URL = "https://auth.openai.com/api/accounts/deviceauth/usercode";
|
||||
const DEVICE_TOKEN_URL = "https://auth.openai.com/api/accounts/deviceauth/token";
|
||||
const OAUTH_TOKEN_URL = "https://auth.openai.com/oauth/token";
|
||||
const REDIRECT_URI = "https://auth.openai.com/deviceauth/callback";
|
||||
const VERIFICATION_URI = "https://auth.openai.com/codex/device";
|
||||
|
||||
type Session = {
|
||||
deviceAuthId: string;
|
||||
userCode: string;
|
||||
};
|
||||
|
||||
function object(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value)
|
||||
? value as Record<string, unknown>
|
||||
: null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" && value.trim() ? value.trim() : null;
|
||||
}
|
||||
|
||||
function number(value: unknown): number | null {
|
||||
if (typeof value === "number" && Number.isFinite(value)) return value;
|
||||
if (typeof value === "string" && value.trim()) {
|
||||
const parsed = Number(value);
|
||||
return Number.isFinite(parsed) ? parsed : null;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function parseBody(body: string): Record<string, unknown> {
|
||||
try {
|
||||
return object(JSON.parse(body)) ?? {};
|
||||
} catch {
|
||||
return {};
|
||||
}
|
||||
}
|
||||
|
||||
function parseSession(value: JsonValue): Session {
|
||||
const session = object(value);
|
||||
const deviceAuthId = text(session?.deviceAuthId);
|
||||
const userCode = text(session?.userCode);
|
||||
if (!deviceAuthId || !userCode) throw new Error("Codex OAuth session is invalid");
|
||||
return { deviceAuthId, userCode };
|
||||
}
|
||||
|
||||
function errorCode(body: Record<string, unknown>): string {
|
||||
const error = body.error;
|
||||
if (typeof error === "string") return error;
|
||||
const nested = object(error);
|
||||
return text(nested?.code ?? nested?.type ?? body.status ?? body.state) ?? "";
|
||||
}
|
||||
|
||||
function errorMessage(body: Record<string, unknown>): string | null {
|
||||
const error = object(body.error);
|
||||
return text(body.error_description ?? body.message ?? error?.message);
|
||||
}
|
||||
|
||||
function pendingMessage(message: string): boolean {
|
||||
const lower = message.toLowerCase();
|
||||
return lower.includes("authorization is pending") ||
|
||||
lower.includes("authorization_pending") ||
|
||||
lower.includes("device authorization is pending");
|
||||
}
|
||||
|
||||
async function begin(context: PluginContext): Promise<OAuth2Begin> {
|
||||
const response = await context.network.fetch(DEVICE_CODE_URL, {
|
||||
method: "POST",
|
||||
headers: { accept: "application/json", "content-type": "application/json" },
|
||||
body: JSON.stringify({ client_id: CLIENT_ID }),
|
||||
});
|
||||
const body = parseBody(response.body);
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
throw new Error(
|
||||
`Failed to request OpenAI Codex device code (HTTP ${response.status}): ${response.body}`,
|
||||
);
|
||||
}
|
||||
const deviceAuthId = text(body.device_auth_id ?? body.device_code);
|
||||
const userCode = text(body.user_code ?? body.usercode);
|
||||
if (!deviceAuthId || !userCode) {
|
||||
throw new Error("OpenAI Codex device authorization response is incomplete");
|
||||
}
|
||||
const session: Session = { deviceAuthId, userCode };
|
||||
return {
|
||||
session: session as unknown as JsonValue,
|
||||
userCode,
|
||||
verificationUrl: VERIFICATION_URI,
|
||||
verificationUrlComplete: VERIFICATION_URI,
|
||||
expiresAtMs: Date.now() + Math.max(1, number(body.expires_in) ?? 900) * 1000,
|
||||
pollIntervalMs: Math.max(1, number(body.interval) ?? 5) * 1000,
|
||||
};
|
||||
}
|
||||
|
||||
async function exchangeAuthorizationCode(
|
||||
context: PluginContext,
|
||||
authorizationCode: string,
|
||||
codeVerifier: string,
|
||||
): Promise<CredentialCandidate> {
|
||||
const response = await context.network.fetch(OAUTH_TOKEN_URL, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
accept: "application/json",
|
||||
"content-type": "application/x-www-form-urlencoded",
|
||||
},
|
||||
body: new URLSearchParams({
|
||||
grant_type: "authorization_code",
|
||||
code: authorizationCode,
|
||||
redirect_uri: REDIRECT_URI,
|
||||
client_id: CLIENT_ID,
|
||||
code_verifier: codeVerifier,
|
||||
}).toString(),
|
||||
});
|
||||
const body = parseBody(response.body);
|
||||
const accessToken = text(body.access_token);
|
||||
if (!accessToken) {
|
||||
throw new Error(
|
||||
errorMessage(body) ?? `Failed to exchange Codex authorization code (HTTP ${response.status})`,
|
||||
);
|
||||
}
|
||||
return { accessToken, refreshToken: text(body.refresh_token), displayName: null };
|
||||
}
|
||||
|
||||
async function completed(credential: CredentialCandidate): Promise<OAuth2Poll> {
|
||||
return { status: "completed", resources: [await credentialDraft(credential)] };
|
||||
}
|
||||
|
||||
async function poll(sessionValue: JsonValue, context: PluginContext): Promise<OAuth2Poll> {
|
||||
const session = parseSession(sessionValue);
|
||||
const response = await context.network.fetch(DEVICE_TOKEN_URL, {
|
||||
method: "POST",
|
||||
headers: { accept: "application/json", "content-type": "application/json" },
|
||||
body: JSON.stringify({
|
||||
device_auth_id: session.deviceAuthId,
|
||||
user_code: session.userCode,
|
||||
}),
|
||||
});
|
||||
const body = parseBody(response.body);
|
||||
// 该端点用 403/404 表示"尚未完成授权"。
|
||||
if (response.status === 403 || response.status === 404) return { status: "pending" };
|
||||
|
||||
const code = errorCode(body);
|
||||
const message = errorMessage(body);
|
||||
if (
|
||||
["authorization_pending", "pending", "waiting", "in_progress", "device_authorization_pending"]
|
||||
.includes(code) ||
|
||||
(message !== null && pendingMessage(message))
|
||||
) {
|
||||
return { status: "pending" };
|
||||
}
|
||||
if (code === "slow_down") return { status: "slow-down" };
|
||||
if (code === "expired_token" || code === "expired") {
|
||||
return { status: "failed", message: message ?? "Device authorization code expired" };
|
||||
}
|
||||
if (code === "access_denied" || code === "denied") {
|
||||
return { status: "denied", message: message ?? undefined };
|
||||
}
|
||||
|
||||
const directToken = text(body.access_token);
|
||||
if (directToken) {
|
||||
return await completed({
|
||||
accessToken: directToken,
|
||||
refreshToken: text(body.refresh_token),
|
||||
displayName: null,
|
||||
});
|
||||
}
|
||||
|
||||
const authorizationCode = text(body.authorization_code);
|
||||
const codeVerifier = text(body.code_verifier);
|
||||
if (response.status >= 200 && response.status < 300 && authorizationCode && codeVerifier) {
|
||||
try {
|
||||
return await completed(
|
||||
await exchangeAuthorizationCode(context, authorizationCode, codeVerifier),
|
||||
);
|
||||
} catch (error) {
|
||||
return {
|
||||
status: "failed",
|
||||
message: error instanceof Error ? error.message : String(error),
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if (!code && body.error === undefined && response.status >= 400) return { status: "pending" };
|
||||
return {
|
||||
status: "failed",
|
||||
message: message ??
|
||||
(code
|
||||
? `OAuth error: ${code}`
|
||||
: `Codex device authorization failed (HTTP ${response.status})`),
|
||||
};
|
||||
}
|
||||
|
||||
export const codexDeviceOAuth: OAuth2AddMethod = {
|
||||
type: "oauth2.0",
|
||||
id: "chatgpt-device",
|
||||
displayName: {
|
||||
"en-US": "Sign in with ChatGPT",
|
||||
"zh-CN": "使用 ChatGPT 登录",
|
||||
},
|
||||
description: {
|
||||
"en-US": "Authorize this device with OpenAI, then add the resulting ChatGPT account.",
|
||||
"zh-CN": "在 OpenAI 完成设备授权后,自动添加对应的 ChatGPT 账号。",
|
||||
},
|
||||
begin,
|
||||
poll,
|
||||
};
|
||||
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"apiVersion": 1,
|
||||
"id": "dev.cursorbyok.examples.codex-auth",
|
||||
"name": "Codex",
|
||||
"version": "0.1.0",
|
||||
"author": "@leookun",
|
||||
"minAppVersion": "0.1.0",
|
||||
"icon": "assets/codex.svg",
|
||||
"entry": "main.ts",
|
||||
"permissions": {
|
||||
"network": [
|
||||
"auth.openai.com",
|
||||
"chatgpt.com"
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
import type {
|
||||
ProviderInvokeInput,
|
||||
ProviderOutput,
|
||||
ProviderResult,
|
||||
ProviderSupport,
|
||||
} from "cursor-byok:provider";
|
||||
import type { PluginContext } from "cursor-byok:plugin";
|
||||
import { HttpError, streamOpenAiResponses } from "cursor-byok:protocol/openai-responses";
|
||||
import { codexModels, reasoningEfforts } from "./models.ts";
|
||||
import {
|
||||
type AccountData,
|
||||
accountData,
|
||||
chatGptAccountId,
|
||||
quotaExhaustedPatch,
|
||||
RESOURCE_TYPE,
|
||||
} from "./resources.ts";
|
||||
|
||||
const RESPONSES_URL = "https://chatgpt.com/backend-api/codex/responses";
|
||||
|
||||
/** 流内错误只有文本可用,按额度关键词分类。 */
|
||||
export function isQuotaError(error: string): boolean {
|
||||
const message = error.toLowerCase();
|
||||
return message.includes("insufficient_quota") ||
|
||||
message.includes("usage_limit_reached") ||
|
||||
message.includes("exceeded your current quota") ||
|
||||
message.includes("quota_exceeded") ||
|
||||
message.includes("5-hour") ||
|
||||
message.includes("5 hour") ||
|
||||
(message.includes("429") &&
|
||||
(message.includes("quota") || message.includes("usage_limit") ||
|
||||
message.includes("insufficient")));
|
||||
}
|
||||
|
||||
/** HTTP 失败携带结构化状态码,429 时放宽响应体的匹配条件。 */
|
||||
function isQuotaHttpError(error: HttpError): boolean {
|
||||
const body = error.body.toLowerCase();
|
||||
return body.includes("insufficient_quota") ||
|
||||
body.includes("usage_limit_reached") ||
|
||||
body.includes("exceeded your current quota") ||
|
||||
body.includes("quota_exceeded") ||
|
||||
body.includes("5-hour") ||
|
||||
body.includes("5 hour") ||
|
||||
(error.status === 429 &&
|
||||
(body.includes("quota") || body.includes("usage_limit") || body.includes("insufficient")));
|
||||
}
|
||||
|
||||
function invalidResult(message: string, stateMessage: string): ProviderResult {
|
||||
return {
|
||||
status: "resource-error",
|
||||
message,
|
||||
patch: { state: { status: "invalid", message: stateMessage } },
|
||||
};
|
||||
}
|
||||
|
||||
function headers(data: AccountData, cacheKey: string | null): Record<string, string> {
|
||||
const result: Record<string, string> = {
|
||||
authorization: `Bearer ${data.accessToken}`,
|
||||
originator: "codex_cli_rs",
|
||||
};
|
||||
const accountId = chatGptAccountId(data.accessToken);
|
||||
if (accountId) result["ChatGPT-Account-Id"] = accountId;
|
||||
// Codex 后端的缓存亲和契约:session-id / thread-id / prompt_cache_key
|
||||
// 三者同源(见 codex-rs client.rs);缺头会导致请求落在随机分片上。
|
||||
if (cacheKey !== null) {
|
||||
result["session-id"] = cacheKey;
|
||||
result["thread-id"] = cacheKey;
|
||||
result["x-client-request-id"] = cacheKey;
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
async function invoke(
|
||||
input: ProviderInvokeInput,
|
||||
output: ProviderOutput,
|
||||
context: PluginContext,
|
||||
): Promise<ProviderResult> {
|
||||
if (!input.resource) {
|
||||
return { status: "request-error", message: "add a ChatGPT account before calling Codex" };
|
||||
}
|
||||
let data: AccountData;
|
||||
try {
|
||||
data = accountData(input.resource);
|
||||
} catch (error) {
|
||||
const message = error instanceof Error ? error.message : String(error);
|
||||
return invalidResult(message, message);
|
||||
}
|
||||
const efforts = reasoningEfforts(input.model);
|
||||
const reasoning = input.request.reasoning;
|
||||
const effort = reasoning.effort !== null && efforts.includes(reasoning.effort)
|
||||
? reasoning.effort
|
||||
: null;
|
||||
try {
|
||||
await streamOpenAiResponses(
|
||||
{
|
||||
url: RESPONSES_URL,
|
||||
model: input.model.id,
|
||||
// Codex 订阅端点不接受 max_output_tokens;fast 档位经协议库映射为
|
||||
// service_tier: "priority" 后透传。
|
||||
request: {
|
||||
...input.request,
|
||||
reasoning: { enabled: reasoning.enabled, effort },
|
||||
maxOutputTokens: null,
|
||||
},
|
||||
headers: headers(data, input.request.cacheKey),
|
||||
extraBody: { store: false },
|
||||
},
|
||||
output,
|
||||
context,
|
||||
);
|
||||
return { status: "completed" };
|
||||
} catch (error) {
|
||||
if (error instanceof HttpError) {
|
||||
if (error.status === 401) {
|
||||
return invalidResult(error.message, "ChatGPT authorization expired; sign in again");
|
||||
}
|
||||
if (isQuotaHttpError(error)) {
|
||||
return {
|
||||
status: "resource-error",
|
||||
message: error.message,
|
||||
patch: quotaExhaustedPatch(data, error.body),
|
||||
};
|
||||
}
|
||||
return { status: "request-error", message: error.message };
|
||||
}
|
||||
const message = error instanceof Error ? error.message : String(error);
|
||||
if (isQuotaError(message)) {
|
||||
return { status: "resource-error", message, patch: quotaExhaustedPatch(data, message) };
|
||||
}
|
||||
return { status: "request-error", message };
|
||||
}
|
||||
}
|
||||
|
||||
export const codexProvider: ProviderSupport = {
|
||||
id: "codex",
|
||||
displayName: "OpenAI Codex",
|
||||
description: {
|
||||
"en-US": "ChatGPT subscription access through the official Codex Responses API.",
|
||||
"zh-CN": "通过官方 Codex Responses API 使用 ChatGPT 订阅。",
|
||||
},
|
||||
providerType: "openai",
|
||||
resourceType: RESOURCE_TYPE,
|
||||
models: codexModels,
|
||||
invoke,
|
||||
};
|
||||
@@ -0,0 +1,432 @@
|
||||
import type { JsonValue, PluginContext } from "cursor-byok:plugin";
|
||||
import type {
|
||||
ResourceDraft,
|
||||
ResourceImportFile,
|
||||
ResourceImportResult,
|
||||
ResourceImportSupport,
|
||||
ResourceMetric,
|
||||
ResourcePatch,
|
||||
ResourceSnapshot,
|
||||
ResourceState,
|
||||
ResourceView,
|
||||
} from "cursor-byok:resource";
|
||||
|
||||
export const RESOURCE_TYPE = "chatgpt-account";
|
||||
|
||||
const USAGE_URL = "https://chatgpt.com/backend-api/wham/usage";
|
||||
const FIVE_HOURS_MS = 5 * 60 * 60 * 1000;
|
||||
|
||||
export type QuotaWindow = {
|
||||
usedPercent: number | null;
|
||||
remainingPercent: number | null;
|
||||
resetAtMs: number | null;
|
||||
};
|
||||
|
||||
export type AccountQuota = {
|
||||
planLabel: string | null;
|
||||
weekly: QuotaWindow | null;
|
||||
fiveHour: QuotaWindow | null;
|
||||
limitReached: boolean;
|
||||
updatedAtMs: number;
|
||||
};
|
||||
|
||||
/** 单条 chatgpt-account 资源的 privateData 形状。 */
|
||||
export type AccountData = {
|
||||
accessToken: string;
|
||||
refreshToken: string | null;
|
||||
displayName: string;
|
||||
quota: AccountQuota | null;
|
||||
};
|
||||
|
||||
export type CredentialCandidate = {
|
||||
accessToken: string;
|
||||
refreshToken: string | null;
|
||||
displayName: string | null;
|
||||
};
|
||||
|
||||
function object(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value)
|
||||
? value as Record<string, unknown>
|
||||
: null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" && value.trim() ? value.trim() : null;
|
||||
}
|
||||
|
||||
function number(value: unknown): number | null {
|
||||
if (typeof value === "number" && Number.isFinite(value)) return value;
|
||||
if (typeof value === "string" && value.trim()) {
|
||||
const parsed = Number(value);
|
||||
return Number.isFinite(parsed) ? parsed : null;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function decodeJwtPayload(token: string): Record<string, unknown> | null {
|
||||
const encoded = token.split(".")[1];
|
||||
if (!encoded) return null;
|
||||
try {
|
||||
const normalized = encoded.replace(/-/g, "+").replace(/_/g, "/");
|
||||
const padded = normalized.padEnd(Math.ceil(normalized.length / 4) * 4, "=");
|
||||
const bytes = Uint8Array.from(atob(padded), (character) => character.charCodeAt(0));
|
||||
return object(JSON.parse(new TextDecoder().decode(bytes)));
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
function claim(payload: Record<string, unknown> | null, key: string): string | null {
|
||||
return payload ? text(payload[key]) : null;
|
||||
}
|
||||
|
||||
export function chatGptAccountId(accessToken: string): string | null {
|
||||
const payload = decodeJwtPayload(accessToken);
|
||||
const auth = object(payload?.["https://api.openai.com/auth"]);
|
||||
return text(auth?.chatgpt_account_id) ?? claim(payload, "chatgpt_account_id");
|
||||
}
|
||||
|
||||
async function tokenFingerprint(token: string): Promise<string> {
|
||||
const digest = await crypto.subtle.digest("SHA-256", new TextEncoder().encode(token));
|
||||
return Array.from(
|
||||
new Uint8Array(digest).slice(0, 8),
|
||||
(byte) => byte.toString(16).padStart(2, "0"),
|
||||
).join("");
|
||||
}
|
||||
|
||||
/** ChatGPT access token 的邮箱通常在 OpenAI 的 profile 声明里,而不是顶层 email。 */
|
||||
function profileEmail(payload: Record<string, unknown> | null): string | null {
|
||||
const profile = object(payload?.["https://api.openai.com/profile"]);
|
||||
return text(profile?.email);
|
||||
}
|
||||
|
||||
export async function accountIdentity(
|
||||
accessToken: string,
|
||||
): Promise<{ key: string; displayName: string }> {
|
||||
const payload = decodeJwtPayload(accessToken);
|
||||
const identity = chatGptAccountId(accessToken) ??
|
||||
claim(payload, "sub") ??
|
||||
claim(payload, "email") ??
|
||||
await tokenFingerprint(accessToken);
|
||||
const displayName = claim(payload, "email") ??
|
||||
profileEmail(payload) ??
|
||||
claim(payload, "preferred_username") ??
|
||||
claim(payload, "name") ??
|
||||
identity;
|
||||
return { key: `codex:${identity}`, displayName };
|
||||
}
|
||||
|
||||
export async function credentialDraft(credential: CredentialCandidate): Promise<ResourceDraft> {
|
||||
const identity = await accountIdentity(credential.accessToken);
|
||||
const data: AccountData = {
|
||||
accessToken: credential.accessToken,
|
||||
refreshToken: credential.refreshToken,
|
||||
displayName: credential.displayName ?? identity.displayName,
|
||||
quota: null,
|
||||
};
|
||||
return { key: identity.key, privateData: data as unknown as JsonValue };
|
||||
}
|
||||
|
||||
export function accountData(resource: ResourceSnapshot): AccountData {
|
||||
const data = object(resource.privateData);
|
||||
const accessToken = text(data?.accessToken);
|
||||
if (!accessToken) throw new Error("ChatGPT account resource is missing its access token");
|
||||
return {
|
||||
accessToken,
|
||||
refreshToken: text(data?.refreshToken),
|
||||
displayName: text(data?.displayName) ?? "ChatGPT account",
|
||||
quota: (data?.quota ?? null) as AccountQuota | null,
|
||||
};
|
||||
}
|
||||
|
||||
export function accountHeaders(data: AccountData): Record<string, string> {
|
||||
const headers: Record<string, string> = {
|
||||
accept: "application/json",
|
||||
originator: "codex_cli_rs",
|
||||
authorization: `Bearer ${data.accessToken}`,
|
||||
};
|
||||
const accountId = chatGptAccountId(data.accessToken);
|
||||
if (accountId) headers["ChatGPT-Account-Id"] = accountId;
|
||||
return headers;
|
||||
}
|
||||
|
||||
function clampPercent(value: number): number {
|
||||
return Math.max(0, Math.min(100, value));
|
||||
}
|
||||
|
||||
function resetAtMs(window: Record<string, unknown>, nowMs: number): number | null {
|
||||
const resetAt = window.reset_at ?? window.resetAt;
|
||||
const numeric = number(resetAt);
|
||||
if (numeric !== null) return numeric > 10_000_000_000 ? numeric : numeric * 1000;
|
||||
if (typeof resetAt === "string") {
|
||||
const parsed = Date.parse(resetAt);
|
||||
if (Number.isFinite(parsed)) return parsed;
|
||||
}
|
||||
const afterSeconds = number(window.reset_after_seconds ?? window.resetAfterSeconds);
|
||||
return afterSeconds === null ? null : nowMs + afterSeconds * 1000;
|
||||
}
|
||||
|
||||
function quotaWindow(value: unknown, nowMs: number): QuotaWindow | null {
|
||||
const window = object(value);
|
||||
if (!window) return null;
|
||||
const used = number(window.used_percent ?? window.usedPercent);
|
||||
const remaining = used === null
|
||||
? number(window.remaining_percent ?? window.remainingPercent)
|
||||
: clampPercent(100 - used);
|
||||
return {
|
||||
usedPercent: used === null
|
||||
? (remaining === null ? null : clampPercent(100 - remaining))
|
||||
: clampPercent(used),
|
||||
remainingPercent: remaining === null ? null : clampPercent(remaining),
|
||||
resetAtMs: resetAtMs(window, nowMs),
|
||||
};
|
||||
}
|
||||
|
||||
function planLabel(value: unknown): string | null {
|
||||
const plan = text(value);
|
||||
if (!plan) return null;
|
||||
const labels: Record<string, string> = {
|
||||
plus: "ChatGPT Plus",
|
||||
pro: "ChatGPT Pro",
|
||||
team: "ChatGPT Team",
|
||||
business: "ChatGPT Business",
|
||||
enterprise: "ChatGPT Enterprise",
|
||||
free: "ChatGPT Free",
|
||||
go: "ChatGPT Go",
|
||||
};
|
||||
return labels[plan.toLowerCase()] ?? plan;
|
||||
}
|
||||
|
||||
export function parseCodexUsage(body: unknown, nowMs = Date.now()): AccountQuota {
|
||||
const root = object(body) ?? {};
|
||||
const rateLimit = object(root.rate_limit ?? root.rateLimit) ?? root;
|
||||
const primary = rateLimit.primary_window ?? rateLimit.primaryWindow;
|
||||
const secondary = rateLimit.secondary_window ?? rateLimit.secondaryWindow;
|
||||
const weekly = quotaWindow(secondary ?? primary, nowMs);
|
||||
const fiveHour = secondary === undefined || secondary === null
|
||||
? null
|
||||
: quotaWindow(primary, nowMs);
|
||||
const explicitLimit = rateLimit.limit_reached ?? rateLimit.limitReached;
|
||||
return {
|
||||
planLabel: planLabel(root.plan_type ?? root.planType),
|
||||
weekly,
|
||||
fiveHour,
|
||||
limitReached: typeof explicitLimit === "boolean" ? explicitLimit : [weekly, fiveHour].some(
|
||||
(window) => window?.remainingPercent !== null && window?.remainingPercent === 0,
|
||||
),
|
||||
updatedAtMs: nowMs,
|
||||
};
|
||||
}
|
||||
|
||||
function windowCoolingUntil(window: QuotaWindow | null, nowMs: number): number | null {
|
||||
if (!window || window.remainingPercent === null || window.remainingPercent > 0) return null;
|
||||
if (window.resetAtMs !== null && window.resetAtMs <= nowMs) return null;
|
||||
return window.resetAtMs ?? nowMs + FIVE_HOURS_MS;
|
||||
}
|
||||
|
||||
export function quotaCoolingUntil(quota: AccountQuota, nowMs = Date.now()): number | null {
|
||||
const resets = [
|
||||
windowCoolingUntil(quota.weekly, nowMs),
|
||||
windowCoolingUntil(quota.fiveHour, nowMs),
|
||||
].filter((value): value is number => value !== null);
|
||||
if (resets.length > 0) return Math.max(...resets);
|
||||
return quota.limitReached ? nowMs + FIVE_HOURS_MS : null;
|
||||
}
|
||||
|
||||
export function quotaState(quota: AccountQuota | null, nowMs = Date.now()): ResourceState {
|
||||
if (!quota) return { status: "ready" };
|
||||
const coolingUntil = quotaCoolingUntil(quota, nowMs);
|
||||
return coolingUntil === null
|
||||
? { status: "ready" }
|
||||
: { status: "cooling", retryAtMs: coolingUntil, message: "ChatGPT quota is exhausted" };
|
||||
}
|
||||
|
||||
/** 从上游错误文本中提取重置时间;拿不到时回退 5 小时。 */
|
||||
function resetFromError(error: string, nowMs: number): number {
|
||||
const resetAt = error.match(/["']?reset_at["']?\s*[:=]\s*["']?(\d+(?:\.\d+)?)/i)?.[1];
|
||||
if (resetAt) {
|
||||
const value = Number(resetAt);
|
||||
if (Number.isFinite(value)) return value > 10_000_000_000 ? value : value * 1000;
|
||||
}
|
||||
const resetAfter = error.match(/["']?reset_after_seconds["']?\s*[:=]\s*["']?(\d+(?:\.\d+)?)/i)
|
||||
?.[1];
|
||||
if (resetAfter) {
|
||||
const value = Number(resetAfter);
|
||||
if (Number.isFinite(value)) return nowMs + value * 1000;
|
||||
}
|
||||
return nowMs + FIVE_HOURS_MS;
|
||||
}
|
||||
|
||||
/** 额度耗尽时的资源补丁:标记 5 小时窗口耗尽并按重置时间进入冷却。 */
|
||||
export function quotaExhaustedPatch(
|
||||
data: AccountData,
|
||||
error: string,
|
||||
nowMs = Date.now(),
|
||||
): ResourcePatch {
|
||||
const quota: AccountQuota = {
|
||||
planLabel: data.quota?.planLabel ?? null,
|
||||
weekly: data.quota?.weekly ?? null,
|
||||
fiveHour: {
|
||||
usedPercent: 100,
|
||||
remainingPercent: 0,
|
||||
resetAtMs: resetFromError(error, nowMs),
|
||||
},
|
||||
limitReached: true,
|
||||
updatedAtMs: nowMs,
|
||||
};
|
||||
return {
|
||||
privateData: { ...data, quota } as unknown as JsonValue,
|
||||
state: quotaState(quota, nowMs),
|
||||
};
|
||||
}
|
||||
|
||||
export function presentAccount(resource: ResourceSnapshot): ResourceView {
|
||||
const data = accountData(resource);
|
||||
const metrics: ResourceMetric[] = [];
|
||||
const weekly = data.quota?.weekly;
|
||||
if (weekly && weekly.remainingPercent !== null) {
|
||||
metrics.push({
|
||||
id: "weekly",
|
||||
label: { "en-US": "Weekly quota", "zh-CN": "周额度" },
|
||||
unit: "percent",
|
||||
value: weekly.remainingPercent,
|
||||
...(weekly.resetAtMs !== null ? { resetAtMs: weekly.resetAtMs } : {}),
|
||||
});
|
||||
}
|
||||
const fiveHour = data.quota?.fiveHour;
|
||||
if (fiveHour && fiveHour.remainingPercent !== null) {
|
||||
metrics.push({
|
||||
id: "five-hour",
|
||||
label: { "en-US": "5-hour window", "zh-CN": "5 小时窗口" },
|
||||
unit: "percent",
|
||||
value: fiveHour.remainingPercent,
|
||||
...(fiveHour.resetAtMs !== null ? { resetAtMs: fiveHour.resetAtMs } : {}),
|
||||
});
|
||||
}
|
||||
return {
|
||||
// 旧记录可能存的是账号 ID;展示时优先从 token 现算邮箱。
|
||||
displayName: jwtDisplayName(data.accessToken) ?? data.displayName,
|
||||
...(data.quota?.planLabel ? { description: data.quota.planLabel } : {}),
|
||||
...(metrics.length > 0 ? { metrics } : {}),
|
||||
};
|
||||
}
|
||||
|
||||
export async function refreshAccount(
|
||||
resource: ResourceSnapshot,
|
||||
context: PluginContext,
|
||||
): Promise<ResourcePatch> {
|
||||
const data = accountData(resource);
|
||||
const response = await context.network.fetch(USAGE_URL, {
|
||||
method: "GET",
|
||||
headers: accountHeaders(data),
|
||||
});
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
if (response.status === 401) {
|
||||
return {
|
||||
state: { status: "invalid", message: "ChatGPT authorization expired; sign in again" },
|
||||
};
|
||||
}
|
||||
throw new Error(`Codex usage lookup failed (HTTP ${response.status}): ${response.body}`);
|
||||
}
|
||||
let body: unknown;
|
||||
try {
|
||||
body = JSON.parse(response.body);
|
||||
} catch {
|
||||
throw new Error("Codex usage lookup returned invalid JSON");
|
||||
}
|
||||
const quota = parseCodexUsage(body);
|
||||
return {
|
||||
privateData: { ...data, quota } as unknown as JsonValue,
|
||||
state: quotaState(quota),
|
||||
};
|
||||
}
|
||||
|
||||
function firstText(source: Record<string, unknown>, keys: string[]): string | null {
|
||||
for (const key of keys) {
|
||||
const value = text(source[key]);
|
||||
if (value) return value;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function jwtDisplayName(token: string | null): string | null {
|
||||
if (!token) return null;
|
||||
const payload = decodeJwtPayload(token);
|
||||
return claim(payload, "email") ?? profileEmail(payload) ??
|
||||
claim(payload, "preferred_username") ?? claim(payload, "name");
|
||||
}
|
||||
|
||||
function collectCredentials(value: unknown, output: CredentialCandidate[]): void {
|
||||
if (Array.isArray(value)) {
|
||||
for (const item of value) collectCredentials(item, output);
|
||||
return;
|
||||
}
|
||||
const item = object(value);
|
||||
if (!item || item.disabled === true) return;
|
||||
for (const key of ["accounts", "credentials", "items"]) {
|
||||
if (Array.isArray(item[key])) {
|
||||
collectCredentials(item[key], output);
|
||||
return;
|
||||
}
|
||||
}
|
||||
const tokens = object(item.tokens) ?? item;
|
||||
const accessToken = firstText(tokens, ["access_token", "accessToken", "token", "key"]) ??
|
||||
firstText(item, ["access_token", "accessToken", "token", "key", "OPENAI_API_KEY"]);
|
||||
if (!accessToken) return;
|
||||
const refreshToken = firstText(tokens, ["refresh_token", "refreshToken"]) ??
|
||||
firstText(item, ["refresh_token", "refreshToken"]);
|
||||
const idToken = firstText(tokens, ["id_token", "idToken"]) ??
|
||||
firstText(item, ["id_token", "idToken"]);
|
||||
const displayName = firstText(item, ["email", "display_name", "displayName", "name"]) ??
|
||||
firstText(tokens, ["email", "display_name", "displayName", "name"]) ??
|
||||
jwtDisplayName(idToken);
|
||||
output.push({ accessToken, refreshToken, displayName });
|
||||
}
|
||||
|
||||
export function parseCredentialFiles(files: ResourceImportFile[]): {
|
||||
credentials: CredentialCandidate[];
|
||||
warnings: string[];
|
||||
} {
|
||||
const credentials: CredentialCandidate[] = [];
|
||||
const warnings: string[] = [];
|
||||
for (const file of files) {
|
||||
let content: unknown;
|
||||
try {
|
||||
content = JSON.parse(file.content);
|
||||
} catch {
|
||||
warnings.push(`${file.name}: not valid JSON`);
|
||||
continue;
|
||||
}
|
||||
const found: CredentialCandidate[] = [];
|
||||
collectCredentials(content, found);
|
||||
if (found.length === 0) {
|
||||
warnings.push(`${file.name}: no ChatGPT access token found`);
|
||||
continue;
|
||||
}
|
||||
credentials.push(...found);
|
||||
}
|
||||
return { credentials, warnings };
|
||||
}
|
||||
|
||||
export const credentialImport: ResourceImportSupport = {
|
||||
displayName: {
|
||||
"en-US": "Import Codex credentials",
|
||||
"zh-CN": "导入 Codex 凭证",
|
||||
},
|
||||
description: {
|
||||
"en-US": "Import one or more Codex JSON credential files.",
|
||||
"zh-CN": "导入一个或多个 Codex JSON 凭证文件。",
|
||||
},
|
||||
accept: [".json"],
|
||||
multiple: true,
|
||||
parse: async (files: ResourceImportFile[]): Promise<ResourceImportResult> => {
|
||||
const { credentials, warnings } = parseCredentialFiles(files);
|
||||
if (credentials.length === 0) {
|
||||
throw new Error(warnings.join("; ") || "credential JSON does not contain an access token");
|
||||
}
|
||||
return {
|
||||
resources: await Promise.all(credentials.map(credentialDraft)),
|
||||
...(warnings.length > 0 ? { warnings } : {}),
|
||||
};
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,10 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<svg width="230" height="230" viewBox="0 0 230 230" xmlns="http://www.w3.org/2000/svg">
|
||||
<title>grok</title>
|
||||
<rect x="22" y="21" width="187" height="187" rx="42" fill="#000000"/>
|
||||
<g fill="#FFFFFF">
|
||||
<path d="M96.5 137.5 L152 82 L166 96 L110.5 151.5 Z"/>
|
||||
<path d="M64 82 L106 124 L92 138 L64 110 Z"/>
|
||||
<path d="M166 148 L166 96 L152 110 L152 148 Z"/>
|
||||
</g>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 418 B |
@@ -0,0 +1,13 @@
|
||||
{
|
||||
"imports": {
|
||||
"cursor-byok:plugin": "../../../src/plugin/sdk/plugin.ts",
|
||||
"cursor-byok:provider": "../../../src/plugin/sdk/provider.ts",
|
||||
"cursor-byok:model": "../../../src/plugin/sdk/model.ts",
|
||||
"cursor-byok:resource": "../../../src/plugin/sdk/resource.ts",
|
||||
"cursor-byok:protocol/openai-chat": "../../../src/plugin/sdk/protocol/openai_chat.ts"
|
||||
},
|
||||
"fmt": {
|
||||
"lineWidth": 100,
|
||||
"exclude": ["assets"]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,379 @@
|
||||
import type {
|
||||
JsonValue,
|
||||
NetworkEventStream,
|
||||
NetworkResponse,
|
||||
PluginContext,
|
||||
} from "cursor-byok:plugin";
|
||||
import type { LlmRequest, ModelEvent } from "cursor-byok:provider";
|
||||
import type { ResourceSnapshot } from "cursor-byok:resource";
|
||||
import { grokDeviceOAuth } from "./oauth.ts";
|
||||
import { FALLBACK_MODELS, grokModels, parseGrokModels } from "./models.ts";
|
||||
import { grokProvider, isQuotaError } from "./provider.ts";
|
||||
import {
|
||||
accountIdentity,
|
||||
credentialDraft,
|
||||
parseCredentialFiles,
|
||||
parseGrokUsage,
|
||||
presentAccount,
|
||||
quotaState,
|
||||
RESOURCE_TYPE,
|
||||
} from "./resources.ts";
|
||||
|
||||
function assert(condition: unknown, message = "assertion failed"): asserts condition {
|
||||
if (!condition) throw new Error(message);
|
||||
}
|
||||
|
||||
function assertEquals(actual: unknown, expected: unknown): void {
|
||||
const left = JSON.stringify(actual);
|
||||
const right = JSON.stringify(expected);
|
||||
if (left !== right) throw new Error(`expected ${right}, received ${left}`);
|
||||
}
|
||||
|
||||
function jwt(payload: Record<string, unknown>): string {
|
||||
const encoded = btoa(JSON.stringify(payload)).replace(/=/g, "").replace(/\+/g, "-").replace(
|
||||
/\//g,
|
||||
"_",
|
||||
);
|
||||
return `header.${encoded}.signature`;
|
||||
}
|
||||
|
||||
type RequestInit = { body?: string; headers?: Record<string, string> };
|
||||
type FetchHandler = (url: string, init?: RequestInit) => NetworkResponse;
|
||||
type StreamHandler = (url: string, init?: RequestInit) => NetworkEventStream;
|
||||
|
||||
function context(handlers: { fetch?: FetchHandler; stream?: StreamHandler }): PluginContext {
|
||||
return {
|
||||
network: {
|
||||
fetch: (url, init) => {
|
||||
if (!handlers.fetch) throw new Error("fetch was not expected");
|
||||
return Promise.resolve(handlers.fetch(url, init));
|
||||
},
|
||||
stream: (url, init) => {
|
||||
if (!handlers.stream) throw new Error("stream was not expected");
|
||||
return Promise.resolve(handlers.stream(url, init));
|
||||
},
|
||||
},
|
||||
signal: new AbortController().signal,
|
||||
};
|
||||
}
|
||||
|
||||
function snapshot(privateData: JsonValue): ResourceSnapshot {
|
||||
return {
|
||||
id: "resource-1",
|
||||
type: RESOURCE_TYPE,
|
||||
key: "grok:user-1",
|
||||
privateData,
|
||||
state: { status: "ready" },
|
||||
};
|
||||
}
|
||||
|
||||
async function* sse(lines: string[]): AsyncGenerator<string> {
|
||||
for (const line of lines) yield line;
|
||||
}
|
||||
|
||||
function request(): LlmRequest {
|
||||
return {
|
||||
instructions: "You are a coding assistant.",
|
||||
messages: [{ role: "user", content: [{ type: "text", text: "hi" }] }],
|
||||
tools: [],
|
||||
reasoning: { enabled: true, effort: "medium" },
|
||||
latency: "fast",
|
||||
maxOutputTokens: 32_000,
|
||||
cacheKey: "conversation-1",
|
||||
};
|
||||
}
|
||||
|
||||
Deno.test("account identity uses the JWT subject and drafts keep tokens private-side", async () => {
|
||||
const token = jwt({ sub: "user-1", email: "person@x.ai" });
|
||||
assertEquals(await accountIdentity(token), {
|
||||
key: "grok:user-1",
|
||||
displayName: "person@x.ai",
|
||||
});
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
assertEquals(draft.key, "grok:user-1");
|
||||
const view = presentAccount(snapshot(draft.privateData));
|
||||
assert(!JSON.stringify(view).includes(token), "resource view exposed an access token");
|
||||
assertEquals(view.displayName, "person@x.ai");
|
||||
});
|
||||
|
||||
Deno.test("credential import accepts Grok credential JSON files", () => {
|
||||
const { credentials, warnings } = parseCredentialFiles([
|
||||
{
|
||||
name: "accounts.json",
|
||||
content: JSON.stringify({
|
||||
accounts: [
|
||||
{ access_token: "token-1", refresh_token: "refresh-1", email: "a@x.ai" },
|
||||
{ access_token: "token-2", disabled: true },
|
||||
],
|
||||
}),
|
||||
},
|
||||
{ name: "broken.json", content: "{not json" },
|
||||
]);
|
||||
assertEquals(credentials, [{
|
||||
accessToken: "token-1",
|
||||
refreshToken: "refresh-1",
|
||||
displayName: "a@x.ai",
|
||||
}]);
|
||||
assertEquals(warnings, ["broken.json: not valid JSON"]);
|
||||
});
|
||||
|
||||
Deno.test("credit usage percent is inverted to remaining and drives cooling", () => {
|
||||
const quota = parseGrokUsage({
|
||||
config: {
|
||||
creditUsagePercent: 34,
|
||||
subscriptionTierDisplay: "SuperGrok",
|
||||
currentPeriod: { end: "2026-09-01T00:00:00Z" },
|
||||
},
|
||||
}, 1_700_000_000_000);
|
||||
assertEquals(quota.planLabel, "SuperGrok");
|
||||
assertEquals(quota.remainingPercent, 66);
|
||||
assertEquals(quota.resetAtMs, Date.parse("2026-09-01T00:00:00Z"));
|
||||
assertEquals(quotaState(quota, 1_700_000_000_000), { status: "ready" });
|
||||
|
||||
const exhausted = parseGrokUsage({
|
||||
config: { creditUsagePercent: 100, currentPeriod: { end: 1_900_000_000 } },
|
||||
}, 1_700_000_000_000);
|
||||
assertEquals(quotaState(exhausted, 1_700_000_000_000), {
|
||||
status: "cooling",
|
||||
retryAtMs: 1_900_000_000_000,
|
||||
message: "Grok credits are exhausted",
|
||||
});
|
||||
});
|
||||
|
||||
Deno.test("missing usage with a billing period counts as unused", () => {
|
||||
const quota = parseGrokUsage({ config: { currentPeriod: { end: 1_900_000_000 } } });
|
||||
assertEquals(quota.remainingPercent, 100);
|
||||
assertEquals(quota.limitReached, false);
|
||||
});
|
||||
|
||||
Deno.test("model discovery parses both language-models and standard list shapes", () => {
|
||||
const richModels = parseGrokModels({
|
||||
models: [
|
||||
{ id: "grok-4", input_modalities: ["text", "image"], context_window: 256_000 },
|
||||
{ id: "grok-3-mini", input_modalities: ["text"] },
|
||||
{ id: "grok-4" },
|
||||
],
|
||||
});
|
||||
assertEquals(richModels.map((model) => model.id), ["grok-4", "grok-3-mini"]);
|
||||
assertEquals(richModels[0].displayName, "Grok 4");
|
||||
assertEquals(richModels[0].capabilities, { images: true });
|
||||
assertEquals(richModels[1].capabilities, { images: false });
|
||||
|
||||
const plainModels = parseGrokModels({ data: [{ id: "grok-4-fast" }] });
|
||||
assertEquals(plainModels.map((model) => model.id), ["grok-4-fast"]);
|
||||
assertEquals(plainModels[0].displayName, "Grok 4 Fast");
|
||||
});
|
||||
|
||||
Deno.test("model discovery falls back to known models when the account cannot list", async () => {
|
||||
const token = jwt({ sub: "user-1" });
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
const models = await grokModels.list(
|
||||
{ resource: snapshot(draft.privateData) },
|
||||
context({
|
||||
fetch: () => ({
|
||||
status: 403,
|
||||
headers: {},
|
||||
body: JSON.stringify({ code: "personal-team-blocked:spending-limit" }),
|
||||
}),
|
||||
}),
|
||||
);
|
||||
assertEquals(models, FALLBACK_MODELS);
|
||||
});
|
||||
|
||||
Deno.test("device OAuth begins with a host-held session and completes with a resource draft", async () => {
|
||||
const accessToken = jwt({ sub: "user-oauth", email: "oauth@x.ai" });
|
||||
let requestNumber = 0;
|
||||
const flowContext = context({
|
||||
fetch: (url, init) => {
|
||||
requestNumber += 1;
|
||||
if (requestNumber === 1) {
|
||||
assertEquals(url, "https://auth.x.ai/oauth2/device/code");
|
||||
assert(init?.body?.includes("scope="), "device code request must carry the scope");
|
||||
return {
|
||||
status: 200,
|
||||
headers: {},
|
||||
body: JSON.stringify({
|
||||
device_code: "private-device-code",
|
||||
user_code: "ABCD-EFGH",
|
||||
verification_uri: "https://accounts.x.ai/activate",
|
||||
verification_uri_complete: "https://accounts.x.ai/activate?code=ABCD-EFGH",
|
||||
expires_in: 900,
|
||||
interval: 5,
|
||||
}),
|
||||
};
|
||||
}
|
||||
assertEquals(url, "https://auth.x.ai/oauth2/token");
|
||||
assert(init?.body?.includes("device_code=private-device-code"));
|
||||
if (requestNumber === 2) {
|
||||
return {
|
||||
status: 400,
|
||||
headers: {},
|
||||
body: JSON.stringify({ error: "authorization_pending" }),
|
||||
};
|
||||
}
|
||||
return {
|
||||
status: 200,
|
||||
headers: {},
|
||||
body: JSON.stringify({ access_token: accessToken, refresh_token: "refresh-secret" }),
|
||||
};
|
||||
},
|
||||
});
|
||||
|
||||
const begun = await grokDeviceOAuth.begin(flowContext);
|
||||
assertEquals(begun.userCode, "ABCD-EFGH");
|
||||
assertEquals(begun.pollIntervalMs, 5000);
|
||||
|
||||
const pending = await grokDeviceOAuth.poll(begun.session, flowContext);
|
||||
assertEquals(pending.status, "pending");
|
||||
|
||||
const polled = await grokDeviceOAuth.poll(begun.session, flowContext);
|
||||
assert(polled.status === "completed", `expected completed, received ${polled.status}`);
|
||||
assertEquals(polled.resources[0].key, "grok:user-oauth");
|
||||
assertEquals(requestNumber, 3);
|
||||
});
|
||||
|
||||
Deno.test("invoke streams normalized events from the xAI Chat Completions API", async () => {
|
||||
const token = jwt({ sub: "user-1" });
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
let requestBody = "";
|
||||
let requestHeaders: Record<string, string> = {};
|
||||
const events: ModelEvent[] = [];
|
||||
const result = await grokProvider.invoke(
|
||||
{
|
||||
model: { id: "grok-4", displayName: "Grok 4" },
|
||||
resource: snapshot(draft.privateData),
|
||||
request: request(),
|
||||
},
|
||||
{ emit: (event) => events.push(event) },
|
||||
context({
|
||||
stream: (url, init) => {
|
||||
assertEquals(url, "https://api.x.ai/v1/chat/completions");
|
||||
requestBody = init?.body ?? "";
|
||||
requestHeaders = init?.headers ?? {};
|
||||
return {
|
||||
status: 200,
|
||||
headers: {},
|
||||
lines: sse([
|
||||
'data: {"choices":[{"delta":{"content":"Hel"}}]}',
|
||||
'data: {"choices":[{"delta":{"content":"lo"}}]}',
|
||||
'data: {"choices":[{"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":2,"prompt_tokens_details":{"cached_tokens":4}}}',
|
||||
"data: [DONE]",
|
||||
]),
|
||||
};
|
||||
},
|
||||
}),
|
||||
);
|
||||
assertEquals(result, { status: "completed" });
|
||||
const body = JSON.parse(requestBody) as Record<string, unknown>;
|
||||
assertEquals(body.model, "grok-4");
|
||||
assertEquals(body.stream, true);
|
||||
assertEquals(body.prompt_cache_key, "conversation-1");
|
||||
assert(!("reasoning_effort" in body), "xAI endpoint rejects reasoning_effort");
|
||||
assert(!("service_tier" in body), "xAI endpoint rejects service_tier");
|
||||
assertEquals(requestHeaders["authorization"], `Bearer ${token}`);
|
||||
assertEquals(events, [
|
||||
{ type: "text-start" },
|
||||
{ type: "text-delta", text: "Hel" },
|
||||
{ type: "text-delta", text: "lo" },
|
||||
{ type: "text-end" },
|
||||
{
|
||||
type: "usage",
|
||||
usage: {
|
||||
inputTokens: 10,
|
||||
outputTokens: 2,
|
||||
totalTokens: null,
|
||||
cacheReadTokens: 4,
|
||||
cacheWriteTokens: null,
|
||||
reasoningTokens: null,
|
||||
},
|
||||
},
|
||||
{ type: "done", reason: "stop" },
|
||||
]);
|
||||
});
|
||||
|
||||
Deno.test("invoke streams incremental tool calls and reasoning replay state", async () => {
|
||||
const token = jwt({ sub: "user-1" });
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
const events: ModelEvent[] = [];
|
||||
const result = await grokProvider.invoke(
|
||||
{
|
||||
model: { id: "grok-4", displayName: "Grok 4" },
|
||||
resource: snapshot(draft.privateData),
|
||||
request: request(),
|
||||
},
|
||||
{ emit: (event) => events.push(event) },
|
||||
context({
|
||||
stream: () => ({
|
||||
status: 200,
|
||||
headers: {},
|
||||
lines: sse([
|
||||
'data: {"choices":[{"delta":{"reasoning_content":"thinking"}}]}',
|
||||
'data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call-1","function":{"name":"read_file","arguments":"{\\"path\\":"}}]}}]}',
|
||||
'data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\\"a.ts\\"}"}}]}}]}',
|
||||
'data: {"choices":[{"delta":{},"finish_reason":"tool_calls"}]}',
|
||||
"data: [DONE]",
|
||||
]),
|
||||
}),
|
||||
}),
|
||||
);
|
||||
assertEquals(result, { status: "completed" });
|
||||
assertEquals(events, [
|
||||
{ type: "thinking-start" },
|
||||
{ type: "thinking-delta", text: "thinking" },
|
||||
{ type: "tool-call-start", index: 0, callId: "call-1", name: "read_file" },
|
||||
{ type: "tool-call-arguments-delta", index: 0, delta: '{"path":' },
|
||||
{ type: "tool-call-arguments-delta", index: 0, delta: '"a.ts"}' },
|
||||
{ type: "thinking-end" },
|
||||
{ type: "tool-call-end", index: 0 },
|
||||
{
|
||||
type: "replay-state",
|
||||
providerKind: "openai_chat",
|
||||
value: { reasoning_content: "thinking" },
|
||||
},
|
||||
{ type: "done", reason: "tool-use" },
|
||||
]);
|
||||
});
|
||||
|
||||
Deno.test("invoke maps quota failures to a cooling resource error", async () => {
|
||||
assert(!isQuotaError("400 invalid request"));
|
||||
assert(isQuotaError("429 credits exhausted"));
|
||||
const token = jwt({ sub: "user-1" });
|
||||
const draft = await credentialDraft({
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
displayName: null,
|
||||
});
|
||||
const result = await grokProvider.invoke(
|
||||
{
|
||||
model: { id: "grok-4", displayName: "Grok 4" },
|
||||
resource: snapshot(draft.privateData),
|
||||
request: request(),
|
||||
},
|
||||
{ emit: () => {} },
|
||||
context({
|
||||
stream: () => ({
|
||||
status: 429,
|
||||
headers: {},
|
||||
lines: sse(['{"error":"credits exhausted"}']),
|
||||
}),
|
||||
}),
|
||||
);
|
||||
assert(result.status === "resource-error", `expected resource-error, received ${result.status}`);
|
||||
assert(result.patch.state?.status === "cooling", "quota failure should cool the resource");
|
||||
});
|
||||
@@ -0,0 +1,16 @@
|
||||
import { defineProviderPlugin } from "cursor-byok:plugin";
|
||||
import { grokDeviceOAuth } from "./oauth.ts";
|
||||
import { grokProvider } from "./provider.ts";
|
||||
import { credentialImport, presentAccount, refreshAccount, RESOURCE_TYPE } from "./resources.ts";
|
||||
|
||||
export default defineProviderPlugin({
|
||||
providers: [grokProvider],
|
||||
resources: [{
|
||||
type: RESOURCE_TYPE,
|
||||
displayName: { "en-US": "Grok accounts", "zh-CN": "Grok 账号" },
|
||||
add: [grokDeviceOAuth],
|
||||
import: credentialImport,
|
||||
present: presentAccount,
|
||||
refresh: refreshAccount,
|
||||
}],
|
||||
});
|
||||
@@ -0,0 +1,96 @@
|
||||
import type { ModelDefinition, ModelSupport } from "cursor-byok:model";
|
||||
import { accountData } from "./resources.ts";
|
||||
|
||||
const LANGUAGE_MODELS_URL = "https://api.x.ai/v1/language-models";
|
||||
const MODELS_URL = "https://api.x.ai/v1/models";
|
||||
|
||||
/** 免费账号无权调用模型列表接口(403 spending-limit);退回已知模型。 */
|
||||
export const FALLBACK_MODELS: ModelDefinition[] = [
|
||||
{
|
||||
id: "grok-4.6",
|
||||
displayName: "Grok 4.6",
|
||||
capabilities: { images: true },
|
||||
},
|
||||
{
|
||||
id: "grok-4.5",
|
||||
displayName: "Grok 4.5",
|
||||
capabilities: { images: true },
|
||||
},
|
||||
];
|
||||
|
||||
function object(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value)
|
||||
? value as Record<string, unknown>
|
||||
: null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" && value.trim() ? value.trim() : null;
|
||||
}
|
||||
|
||||
function modalities(value: unknown): string[] {
|
||||
return Array.isArray(value)
|
||||
? value.flatMap((item) => (typeof item === "string" ? [item.toLowerCase()] : []))
|
||||
: [];
|
||||
}
|
||||
|
||||
/** 把模型 ID 变成可读名称,如 grok-4-fast → Grok 4 Fast。 */
|
||||
function displayName(id: string): string {
|
||||
return id
|
||||
.split("-")
|
||||
.map((part) => (/^\d/.test(part) ? part : part.charAt(0).toUpperCase() + part.slice(1)))
|
||||
.join(" ");
|
||||
}
|
||||
|
||||
/** 兼容 /v1/language-models 的 models 数组与 /v1/models 的 data 数组。 */
|
||||
export function parseGrokModels(body: unknown): ModelDefinition[] {
|
||||
const root = object(body);
|
||||
const source = root?.models ?? root?.data ?? body;
|
||||
if (!Array.isArray(source)) {
|
||||
throw new Error("Grok model discovery response does not contain a model list");
|
||||
}
|
||||
const seen = new Set<string>();
|
||||
const models: ModelDefinition[] = [];
|
||||
for (const raw of source) {
|
||||
const model = object(raw);
|
||||
const id = model ? text(model.id ?? model.name) : null;
|
||||
if (!id || seen.has(id)) continue;
|
||||
seen.add(id);
|
||||
const inputs = modalities(model?.input_modalities ?? model?.inputModalities);
|
||||
models.push({
|
||||
id,
|
||||
displayName: displayName(id),
|
||||
capabilities: {
|
||||
images: inputs.length === 0 || inputs.includes("image"),
|
||||
},
|
||||
});
|
||||
}
|
||||
return models;
|
||||
}
|
||||
|
||||
export const grokModels: ModelSupport = {
|
||||
list: async ({ resource }, context): Promise<ModelDefinition[]> => {
|
||||
if (!resource) throw new Error("add a Grok account before syncing models");
|
||||
const data = accountData(resource);
|
||||
const headers = {
|
||||
accept: "application/json",
|
||||
authorization: `Bearer ${data.accessToken}`,
|
||||
};
|
||||
// language-models 带模态与上下文元数据;不可用时回退到标准列表。
|
||||
let response = await context.network.fetch(LANGUAGE_MODELS_URL, { method: "GET", headers });
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
response = await context.network.fetch(MODELS_URL, { method: "GET", headers });
|
||||
}
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
return FALLBACK_MODELS;
|
||||
}
|
||||
let body: unknown;
|
||||
try {
|
||||
body = JSON.parse(response.body);
|
||||
} catch {
|
||||
throw new Error("Grok model discovery returned invalid JSON");
|
||||
}
|
||||
const models = parseGrokModels(body);
|
||||
return models.length > 0 ? models : FALLBACK_MODELS;
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,148 @@
|
||||
import type { JsonValue, PluginContext } from "cursor-byok:plugin";
|
||||
import type { OAuth2AddMethod, OAuth2Begin, OAuth2Poll } from "cursor-byok:resource";
|
||||
import { credentialDraft } from "./resources.ts";
|
||||
|
||||
const CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828";
|
||||
const DEVICE_CODE_URL = "https://auth.x.ai/oauth2/device/code";
|
||||
const TOKEN_URL = "https://auth.x.ai/oauth2/token";
|
||||
const SCOPE = "openid profile email offline_access grok-cli:access api:access";
|
||||
|
||||
type Session = {
|
||||
deviceCode: string;
|
||||
};
|
||||
|
||||
function object(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value)
|
||||
? value as Record<string, unknown>
|
||||
: null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" && value.trim() ? value.trim() : null;
|
||||
}
|
||||
|
||||
function number(value: unknown): number | null {
|
||||
if (typeof value === "number" && Number.isFinite(value)) return value;
|
||||
if (typeof value === "string" && value.trim()) {
|
||||
const parsed = Number(value);
|
||||
return Number.isFinite(parsed) ? parsed : null;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function parseBody(body: string): Record<string, unknown> {
|
||||
try {
|
||||
return object(JSON.parse(body)) ?? {};
|
||||
} catch {
|
||||
return {};
|
||||
}
|
||||
}
|
||||
|
||||
function parseSession(value: JsonValue): Session {
|
||||
const session = object(value);
|
||||
const deviceCode = text(session?.deviceCode);
|
||||
if (!deviceCode) throw new Error("Grok OAuth session is invalid");
|
||||
return { deviceCode };
|
||||
}
|
||||
|
||||
async function begin(context: PluginContext): Promise<OAuth2Begin> {
|
||||
const response = await context.network.fetch(DEVICE_CODE_URL, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
accept: "application/json",
|
||||
"content-type": "application/x-www-form-urlencoded",
|
||||
},
|
||||
body: new URLSearchParams({ client_id: CLIENT_ID, scope: SCOPE }).toString(),
|
||||
});
|
||||
const body = parseBody(response.body);
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
throw new Error(
|
||||
`Failed to request xAI device code (HTTP ${response.status}): ${response.body}`,
|
||||
);
|
||||
}
|
||||
const deviceCode = text(body.device_code);
|
||||
const userCode = text(body.user_code);
|
||||
const verificationUrl = text(body.verification_uri);
|
||||
if (!deviceCode || !userCode || !verificationUrl) {
|
||||
throw new Error("xAI device authorization response is incomplete");
|
||||
}
|
||||
const session: Session = { deviceCode };
|
||||
return {
|
||||
session: session as unknown as JsonValue,
|
||||
userCode,
|
||||
verificationUrl,
|
||||
...(text(body.verification_uri_complete)
|
||||
? { verificationUrlComplete: text(body.verification_uri_complete)! }
|
||||
: {}),
|
||||
expiresAtMs: Date.now() + Math.max(1, number(body.expires_in) ?? 900) * 1000,
|
||||
pollIntervalMs: Math.max(1, number(body.interval) ?? 5) * 1000,
|
||||
};
|
||||
}
|
||||
|
||||
async function poll(sessionValue: JsonValue, context: PluginContext): Promise<OAuth2Poll> {
|
||||
const session = parseSession(sessionValue);
|
||||
const response = await context.network.fetch(TOKEN_URL, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
accept: "application/json",
|
||||
"content-type": "application/x-www-form-urlencoded",
|
||||
},
|
||||
body: new URLSearchParams({
|
||||
grant_type: "urn:ietf:params:oauth:grant-type:device_code",
|
||||
client_id: CLIENT_ID,
|
||||
device_code: session.deviceCode,
|
||||
}).toString(),
|
||||
});
|
||||
const body = parseBody(response.body);
|
||||
if (response.status >= 200 && response.status < 300) {
|
||||
const accessToken = text(body.access_token);
|
||||
if (!accessToken) {
|
||||
return { status: "failed", message: "xAI token response is missing access_token" };
|
||||
}
|
||||
return {
|
||||
status: "completed",
|
||||
resources: [
|
||||
await credentialDraft({
|
||||
accessToken,
|
||||
refreshToken: text(body.refresh_token),
|
||||
displayName: null,
|
||||
}),
|
||||
],
|
||||
};
|
||||
}
|
||||
const code = text(body.error) ?? "";
|
||||
const message = text(body.error_description);
|
||||
switch (code) {
|
||||
case "authorization_pending":
|
||||
return { status: "pending" };
|
||||
case "slow_down":
|
||||
return { status: "slow-down" };
|
||||
case "expired_token":
|
||||
return { status: "failed", message: message ?? "Device authorization code expired" };
|
||||
case "access_denied":
|
||||
return { status: "denied", ...(message ? { message } : {}) };
|
||||
default:
|
||||
return {
|
||||
status: "failed",
|
||||
message: message ??
|
||||
(code
|
||||
? `OAuth error: ${code}`
|
||||
: `xAI device authorization failed (HTTP ${response.status})`),
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
export const grokDeviceOAuth: OAuth2AddMethod = {
|
||||
type: "oauth2.0",
|
||||
id: "xai-device",
|
||||
displayName: {
|
||||
"en-US": "Sign in with xAI",
|
||||
"zh-CN": "使用 xAI 登录",
|
||||
},
|
||||
description: {
|
||||
"en-US": "Authorize this device with xAI, then add the resulting Grok account.",
|
||||
"zh-CN": "在 xAI 完成设备授权后,自动添加对应的 Grok 账号。",
|
||||
},
|
||||
begin,
|
||||
poll,
|
||||
};
|
||||
@@ -0,0 +1,17 @@
|
||||
{
|
||||
"apiVersion": 1,
|
||||
"id": "dev.cursorbyok.examples.grok-auth",
|
||||
"name": "Grok",
|
||||
"version": "0.1.0",
|
||||
"author": "@leookun",
|
||||
"minAppVersion": "0.1.0",
|
||||
"icon": "assets/grok.svg",
|
||||
"entry": "main.ts",
|
||||
"permissions": {
|
||||
"network": [
|
||||
"auth.x.ai",
|
||||
"api.x.ai",
|
||||
"cli-chat-proxy.grok.com"
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
import type {
|
||||
ProviderInvokeInput,
|
||||
ProviderOutput,
|
||||
ProviderResult,
|
||||
ProviderSupport,
|
||||
} from "cursor-byok:provider";
|
||||
import type { PluginContext } from "cursor-byok:plugin";
|
||||
import { HttpError, streamOpenAiChat } from "cursor-byok:protocol/openai-chat";
|
||||
import { grokModels } from "./models.ts";
|
||||
import { type AccountData, accountData, quotaExhaustedPatch, RESOURCE_TYPE } from "./resources.ts";
|
||||
|
||||
const CHAT_URL = "https://api.x.ai/v1/chat/completions";
|
||||
|
||||
/** 流内错误只有文本可用,按积分/额度关键词分类。 */
|
||||
export function isQuotaError(error: string): boolean {
|
||||
const message = error.toLowerCase();
|
||||
return message.includes("insufficient_quota") ||
|
||||
message.includes("credits exhausted") ||
|
||||
message.includes("out of credits") ||
|
||||
message.includes("quota_exceeded") ||
|
||||
(message.includes("429") &&
|
||||
(message.includes("quota") || message.includes("credit") ||
|
||||
message.includes("insufficient")));
|
||||
}
|
||||
|
||||
/** HTTP 失败携带结构化状态码,429 一律按额度耗尽处理并冷却账号。 */
|
||||
function isQuotaHttpError(error: HttpError): boolean {
|
||||
if (error.status === 429) return true;
|
||||
const body = error.body.toLowerCase();
|
||||
return body.includes("insufficient_quota") ||
|
||||
body.includes("credits exhausted") ||
|
||||
body.includes("out of credits") ||
|
||||
// 免费账号触达消费上限时返回 403 spending-limit,属于额度而非授权问题。
|
||||
body.includes("spending-limit") ||
|
||||
body.includes("run out of credits") ||
|
||||
body.includes("quota_exceeded");
|
||||
}
|
||||
|
||||
function invalidResult(message: string, stateMessage: string): ProviderResult {
|
||||
return {
|
||||
status: "resource-error",
|
||||
message,
|
||||
patch: { state: { status: "invalid", message: stateMessage } },
|
||||
};
|
||||
}
|
||||
|
||||
async function invoke(
|
||||
input: ProviderInvokeInput,
|
||||
output: ProviderOutput,
|
||||
context: PluginContext,
|
||||
): Promise<ProviderResult> {
|
||||
if (!input.resource) {
|
||||
return { status: "request-error", message: "add a Grok account before calling Grok" };
|
||||
}
|
||||
let data: AccountData;
|
||||
try {
|
||||
data = accountData(input.resource);
|
||||
} catch (error) {
|
||||
const message = error instanceof Error ? error.message : String(error);
|
||||
return invalidResult(message, message);
|
||||
}
|
||||
try {
|
||||
await streamOpenAiChat(
|
||||
{
|
||||
url: CHAT_URL,
|
||||
model: input.model.id,
|
||||
// xAI 不接受 reasoning_effort 与 service_tier;思考由模型自身决定。
|
||||
request: {
|
||||
...input.request,
|
||||
reasoning: { enabled: false, effort: null },
|
||||
latency: "standard",
|
||||
},
|
||||
headers: { authorization: `Bearer ${data.accessToken}` },
|
||||
},
|
||||
output,
|
||||
context,
|
||||
);
|
||||
return { status: "completed" };
|
||||
} catch (error) {
|
||||
if (error instanceof HttpError) {
|
||||
if ((error.status === 401 || error.status === 403) && !isQuotaHttpError(error)) {
|
||||
return invalidResult(error.message, "Grok authorization expired; sign in again");
|
||||
}
|
||||
if (isQuotaHttpError(error)) {
|
||||
return {
|
||||
status: "resource-error",
|
||||
message: error.message,
|
||||
patch: quotaExhaustedPatch(data),
|
||||
};
|
||||
}
|
||||
return { status: "request-error", message: error.message };
|
||||
}
|
||||
const message = error instanceof Error ? error.message : String(error);
|
||||
if (isQuotaError(message)) {
|
||||
return { status: "resource-error", message, patch: quotaExhaustedPatch(data) };
|
||||
}
|
||||
return { status: "request-error", message };
|
||||
}
|
||||
}
|
||||
|
||||
export const grokProvider: ProviderSupport = {
|
||||
id: "grok",
|
||||
displayName: "xAI Grok",
|
||||
description: {
|
||||
"en-US": "SuperGrok subscription access through the official Grok CLI endpoint.",
|
||||
"zh-CN": "通过官方 Grok CLI 接口使用 SuperGrok 订阅。",
|
||||
},
|
||||
providerType: "xai",
|
||||
resourceType: RESOURCE_TYPE,
|
||||
models: grokModels,
|
||||
invoke,
|
||||
};
|
||||
@@ -0,0 +1,342 @@
|
||||
import type { JsonValue, PluginContext } from "cursor-byok:plugin";
|
||||
import type {
|
||||
ResourceDraft,
|
||||
ResourceImportFile,
|
||||
ResourceImportResult,
|
||||
ResourceImportSupport,
|
||||
ResourceMetric,
|
||||
ResourcePatch,
|
||||
ResourceSnapshot,
|
||||
ResourceState,
|
||||
ResourceView,
|
||||
} from "cursor-byok:resource";
|
||||
|
||||
export const RESOURCE_TYPE = "grok-account";
|
||||
|
||||
const CREDITS_URL = "https://cli-chat-proxy.grok.com/v1/billing?format=credits";
|
||||
const ONE_HOUR_MS = 60 * 60 * 1000;
|
||||
|
||||
export type AccountQuota = {
|
||||
planLabel: string | null;
|
||||
usedPercent: number | null;
|
||||
remainingPercent: number | null;
|
||||
resetAtMs: number | null;
|
||||
limitReached: boolean;
|
||||
updatedAtMs: number;
|
||||
};
|
||||
|
||||
/** 单条 grok-account 资源的 privateData 形状。 */
|
||||
export type AccountData = {
|
||||
accessToken: string;
|
||||
refreshToken: string | null;
|
||||
displayName: string;
|
||||
quota: AccountQuota | null;
|
||||
};
|
||||
|
||||
export type CredentialCandidate = {
|
||||
accessToken: string;
|
||||
refreshToken: string | null;
|
||||
displayName: string | null;
|
||||
};
|
||||
|
||||
function object(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value)
|
||||
? value as Record<string, unknown>
|
||||
: null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" && value.trim() ? value.trim() : null;
|
||||
}
|
||||
|
||||
function number(value: unknown): number | null {
|
||||
if (typeof value === "number" && Number.isFinite(value)) return value;
|
||||
if (typeof value === "string" && value.trim()) {
|
||||
const parsed = Number(value);
|
||||
return Number.isFinite(parsed) ? parsed : null;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function decodeJwtPayload(token: string): Record<string, unknown> | null {
|
||||
const encoded = token.split(".")[1];
|
||||
if (!encoded) return null;
|
||||
try {
|
||||
const normalized = encoded.replace(/-/g, "+").replace(/_/g, "/");
|
||||
const padded = normalized.padEnd(Math.ceil(normalized.length / 4) * 4, "=");
|
||||
const bytes = Uint8Array.from(atob(padded), (character) => character.charCodeAt(0));
|
||||
return object(JSON.parse(new TextDecoder().decode(bytes)));
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
function claim(payload: Record<string, unknown> | null, key: string): string | null {
|
||||
return payload ? text(payload[key]) : null;
|
||||
}
|
||||
|
||||
async function tokenFingerprint(token: string): Promise<string> {
|
||||
const digest = await crypto.subtle.digest("SHA-256", new TextEncoder().encode(token));
|
||||
return Array.from(
|
||||
new Uint8Array(digest).slice(0, 8),
|
||||
(byte) => byte.toString(16).padStart(2, "0"),
|
||||
).join("");
|
||||
}
|
||||
|
||||
export async function accountIdentity(
|
||||
accessToken: string,
|
||||
): Promise<{ key: string; displayName: string }> {
|
||||
const payload = decodeJwtPayload(accessToken);
|
||||
const identity = claim(payload, "sub") ??
|
||||
claim(payload, "email") ??
|
||||
await tokenFingerprint(accessToken);
|
||||
const displayName = claim(payload, "email") ??
|
||||
claim(payload, "preferred_username") ??
|
||||
claim(payload, "name") ??
|
||||
identity;
|
||||
return { key: `grok:${identity}`, displayName };
|
||||
}
|
||||
|
||||
export async function credentialDraft(credential: CredentialCandidate): Promise<ResourceDraft> {
|
||||
const identity = await accountIdentity(credential.accessToken);
|
||||
const data: AccountData = {
|
||||
accessToken: credential.accessToken,
|
||||
refreshToken: credential.refreshToken,
|
||||
displayName: credential.displayName ?? identity.displayName,
|
||||
quota: null,
|
||||
};
|
||||
return { key: identity.key, privateData: data as unknown as JsonValue };
|
||||
}
|
||||
|
||||
export function accountData(resource: ResourceSnapshot): AccountData {
|
||||
const data = object(resource.privateData);
|
||||
const accessToken = text(data?.accessToken);
|
||||
if (!accessToken) throw new Error("Grok account resource is missing its access token");
|
||||
return {
|
||||
accessToken,
|
||||
refreshToken: text(data?.refreshToken),
|
||||
displayName: text(data?.displayName) ?? "Grok account",
|
||||
quota: (data?.quota ?? null) as AccountQuota | null,
|
||||
};
|
||||
}
|
||||
|
||||
function clampPercent(value: number): number {
|
||||
return Math.max(0, Math.min(100, value));
|
||||
}
|
||||
|
||||
function resetAtMs(value: unknown): number | null {
|
||||
const numeric = number(value);
|
||||
if (numeric !== null) return numeric > 10_000_000_000 ? numeric : numeric * 1000;
|
||||
if (typeof value === "string") {
|
||||
const parsed = Date.parse(value);
|
||||
if (Number.isFinite(parsed)) return parsed;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
/** 解析 Grok CLI 计费接口的积分响应;creditUsagePercent 表示已用占比。 */
|
||||
export function parseGrokUsage(body: unknown, nowMs = Date.now()): AccountQuota {
|
||||
const root = object(body) ?? {};
|
||||
const config = object(root.config) ?? root;
|
||||
let used = number(config.creditUsagePercent ?? config.credit_usage_percent);
|
||||
if (used === null) {
|
||||
const onDemandUsed = number(config.onDemandUsed ?? config.on_demand_used);
|
||||
const onDemandCap = number(config.onDemandCap ?? config.on_demand_cap);
|
||||
if (onDemandUsed !== null && onDemandCap !== null && onDemandCap > 0) {
|
||||
used = (onDemandUsed / onDemandCap) * 100;
|
||||
}
|
||||
}
|
||||
// 存在计费周期但没有用量字段时视为未使用。
|
||||
if (used === null && (config.currentPeriod ?? config.current_period) !== undefined) {
|
||||
used = 0;
|
||||
}
|
||||
const remaining = used === null ? null : clampPercent(100 - used);
|
||||
const period = object(config.currentPeriod ?? config.current_period);
|
||||
return {
|
||||
planLabel: text(
|
||||
config.subscriptionTierDisplay ?? config.subscription_tier_display ??
|
||||
config.subscriptionTier ?? config.product,
|
||||
),
|
||||
usedPercent: used === null ? null : clampPercent(used),
|
||||
remainingPercent: remaining,
|
||||
resetAtMs: resetAtMs(period?.end ?? config.billingPeriodEnd ?? config.billing_period_end),
|
||||
limitReached: remaining !== null && remaining <= 0,
|
||||
updatedAtMs: nowMs,
|
||||
};
|
||||
}
|
||||
|
||||
export function quotaState(quota: AccountQuota | null, nowMs = Date.now()): ResourceState {
|
||||
if (!quota || !quota.limitReached) return { status: "ready" };
|
||||
if (quota.resetAtMs !== null && quota.resetAtMs <= nowMs) return { status: "ready" };
|
||||
return {
|
||||
status: "cooling",
|
||||
retryAtMs: quota.resetAtMs ?? nowMs + ONE_HOUR_MS,
|
||||
message: "Grok credits are exhausted",
|
||||
};
|
||||
}
|
||||
|
||||
/** 额度耗尽时的资源补丁:标记积分耗尽并进入冷却,重置时间未知时回退 1 小时。 */
|
||||
export function quotaExhaustedPatch(data: AccountData, nowMs = Date.now()): ResourcePatch {
|
||||
const quota: AccountQuota = {
|
||||
planLabel: data.quota?.planLabel ?? null,
|
||||
usedPercent: 100,
|
||||
remainingPercent: 0,
|
||||
resetAtMs: data.quota?.resetAtMs !== undefined && data.quota?.resetAtMs !== null &&
|
||||
data.quota.resetAtMs > nowMs
|
||||
? data.quota.resetAtMs
|
||||
: null,
|
||||
limitReached: true,
|
||||
updatedAtMs: nowMs,
|
||||
};
|
||||
return {
|
||||
privateData: { ...data, quota } as unknown as JsonValue,
|
||||
state: quotaState(quota, nowMs),
|
||||
};
|
||||
}
|
||||
|
||||
export function accountHeaders(data: AccountData): Record<string, string> {
|
||||
return {
|
||||
accept: "application/json",
|
||||
authorization: `Bearer ${data.accessToken}`,
|
||||
// Grok CLI 计费接口要求该头标识客户端来源。
|
||||
"x-xai-token-auth": "xai-grok-cli",
|
||||
};
|
||||
}
|
||||
|
||||
function jwtDisplayName(token: string | null): string | null {
|
||||
if (!token) return null;
|
||||
const payload = decodeJwtPayload(token);
|
||||
return claim(payload, "email") ?? claim(payload, "preferred_username") ??
|
||||
claim(payload, "name");
|
||||
}
|
||||
|
||||
export function presentAccount(resource: ResourceSnapshot): ResourceView {
|
||||
const data = accountData(resource);
|
||||
const metrics: ResourceMetric[] = [];
|
||||
const quota = data.quota;
|
||||
if (quota && quota.remainingPercent !== null) {
|
||||
metrics.push({
|
||||
id: "credits",
|
||||
label: { "en-US": "Credits", "zh-CN": "积分额度" },
|
||||
unit: "percent",
|
||||
value: quota.remainingPercent,
|
||||
...(quota.resetAtMs !== null ? { resetAtMs: quota.resetAtMs } : {}),
|
||||
});
|
||||
}
|
||||
return {
|
||||
// 旧记录可能存的是账号 ID;展示时优先从 token 现算邮箱。
|
||||
displayName: jwtDisplayName(data.accessToken) ?? data.displayName,
|
||||
...(quota?.planLabel ? { description: quota.planLabel } : {}),
|
||||
...(metrics.length > 0 ? { metrics } : {}),
|
||||
};
|
||||
}
|
||||
|
||||
export async function refreshAccount(
|
||||
resource: ResourceSnapshot,
|
||||
context: PluginContext,
|
||||
): Promise<ResourcePatch> {
|
||||
const data = accountData(resource);
|
||||
const response = await context.network.fetch(CREDITS_URL, {
|
||||
method: "GET",
|
||||
headers: accountHeaders(data),
|
||||
});
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
if (response.status === 401 || response.status === 403) {
|
||||
return {
|
||||
state: { status: "invalid", message: "Grok authorization expired; sign in again" },
|
||||
};
|
||||
}
|
||||
throw new Error(`Grok usage lookup failed (HTTP ${response.status}): ${response.body}`);
|
||||
}
|
||||
let body: unknown;
|
||||
try {
|
||||
body = JSON.parse(response.body);
|
||||
} catch {
|
||||
throw new Error("Grok usage lookup returned invalid JSON");
|
||||
}
|
||||
const quota = parseGrokUsage(body);
|
||||
return {
|
||||
privateData: { ...data, quota } as unknown as JsonValue,
|
||||
state: quotaState(quota),
|
||||
};
|
||||
}
|
||||
|
||||
function firstText(source: Record<string, unknown>, keys: string[]): string | null {
|
||||
for (const key of keys) {
|
||||
const value = text(source[key]);
|
||||
if (value) return value;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function collectCredentials(value: unknown, output: CredentialCandidate[]): void {
|
||||
if (Array.isArray(value)) {
|
||||
for (const item of value) collectCredentials(item, output);
|
||||
return;
|
||||
}
|
||||
const item = object(value);
|
||||
if (!item || item.disabled === true) return;
|
||||
for (const key of ["accounts", "credentials", "items"]) {
|
||||
if (Array.isArray(item[key])) {
|
||||
collectCredentials(item[key], output);
|
||||
return;
|
||||
}
|
||||
}
|
||||
const tokens = object(item.tokens) ?? item;
|
||||
const accessToken = firstText(tokens, ["access_token", "accessToken", "token", "key"]) ??
|
||||
firstText(item, ["access_token", "accessToken", "token", "key", "XAI_API_KEY"]);
|
||||
if (!accessToken) return;
|
||||
const refreshToken = firstText(tokens, ["refresh_token", "refreshToken"]) ??
|
||||
firstText(item, ["refresh_token", "refreshToken"]);
|
||||
const displayName = firstText(item, ["email", "display_name", "displayName", "name"]) ??
|
||||
firstText(tokens, ["email", "display_name", "displayName", "name"]);
|
||||
output.push({ accessToken, refreshToken, displayName });
|
||||
}
|
||||
|
||||
export function parseCredentialFiles(files: ResourceImportFile[]): {
|
||||
credentials: CredentialCandidate[];
|
||||
warnings: string[];
|
||||
} {
|
||||
const credentials: CredentialCandidate[] = [];
|
||||
const warnings: string[] = [];
|
||||
for (const file of files) {
|
||||
let content: unknown;
|
||||
try {
|
||||
content = JSON.parse(file.content);
|
||||
} catch {
|
||||
warnings.push(`${file.name}: not valid JSON`);
|
||||
continue;
|
||||
}
|
||||
const found: CredentialCandidate[] = [];
|
||||
collectCredentials(content, found);
|
||||
if (found.length === 0) {
|
||||
warnings.push(`${file.name}: no Grok access token found`);
|
||||
continue;
|
||||
}
|
||||
credentials.push(...found);
|
||||
}
|
||||
return { credentials, warnings };
|
||||
}
|
||||
|
||||
export const credentialImport: ResourceImportSupport = {
|
||||
displayName: {
|
||||
"en-US": "Import Grok credentials",
|
||||
"zh-CN": "导入 Grok 凭证",
|
||||
},
|
||||
description: {
|
||||
"en-US": "Import one or more Grok JSON credential files.",
|
||||
"zh-CN": "导入一个或多个 Grok JSON 凭证文件。",
|
||||
},
|
||||
accept: [".json"],
|
||||
multiple: true,
|
||||
parse: async (files: ResourceImportFile[]): Promise<ResourceImportResult> => {
|
||||
const { credentials, warnings } = parseCredentialFiles(files);
|
||||
if (credentials.length === 0) {
|
||||
throw new Error(warnings.join("; ") || "credential JSON does not contain an access token");
|
||||
}
|
||||
return {
|
||||
resources: await Promise.all(credentials.map(credentialDraft)),
|
||||
...(warnings.length > 0 ? { warnings } : {}),
|
||||
};
|
||||
},
|
||||
};
|
||||
@@ -19,7 +19,9 @@ use crate::{
|
||||
connect,
|
||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
},
|
||||
services::{account, analytics, model_catalog, observability::CursorTraceRecorder, tab},
|
||||
services::{
|
||||
account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab,
|
||||
},
|
||||
transport::{TransportParent, TransportRegistry},
|
||||
},
|
||||
Result,
|
||||
@@ -27,10 +29,15 @@ use crate::{
|
||||
|
||||
pub fn router(registry: TransportRegistry) -> Result<Router> {
|
||||
let proxy = CursorProxy::cursor(registry.store().clone())?;
|
||||
Ok(router_with_proxy(registry, proxy))
|
||||
let knowledge = knowledge::KnowledgeService::managed()?;
|
||||
Ok(router_with_proxy(registry, proxy, knowledge))
|
||||
}
|
||||
|
||||
fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router {
|
||||
fn router_with_proxy(
|
||||
registry: TransportRegistry,
|
||||
proxy: CursorProxy,
|
||||
knowledge_service: knowledge::KnowledgeService,
|
||||
) -> Router {
|
||||
let web_cache = registry.web_cache().router();
|
||||
Router::new()
|
||||
.route("/__byok-api__/healthz", get(health))
|
||||
@@ -69,6 +76,22 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
|
||||
"/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
|
||||
post(account::usage_limit_status),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseAdd",
|
||||
post(knowledge::add),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseList",
|
||||
post(knowledge::list),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseUpdate",
|
||||
post(knowledge::update),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseRemove",
|
||||
post(knowledge::remove),
|
||||
)
|
||||
.route(
|
||||
analytics::BOOTSTRAP_STATSIG_PATH,
|
||||
post(analytics::bootstrap_statsig),
|
||||
@@ -80,6 +103,7 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
|
||||
.fallback(proxy::forward)
|
||||
.method_not_allowed_fallback(proxy::forward)
|
||||
.layer(Extension(proxy))
|
||||
.layer(Extension(knowledge_service))
|
||||
.with_state(registry)
|
||||
.merge(web_cache)
|
||||
}
|
||||
@@ -133,7 +157,10 @@ async fn bidi_handler(
|
||||
let conversation_id = decoded.conversation_id().map(str::to_owned);
|
||||
let trace_metadata = decoded.trace_metadata();
|
||||
let local = if let Some(model_id) = decoded.model_id() {
|
||||
if registry.store().model(model_id).await?.is_some() {
|
||||
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
|
||||
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX)
|
||||
|| registry.store().model(model_id).await?.is_some()
|
||||
{
|
||||
tracing::info!(
|
||||
request_id = decoded.request_id,
|
||||
model_id,
|
||||
|
||||
+14
-2
@@ -13,6 +13,7 @@ use crate::{
|
||||
transport::TransportRegistry,
|
||||
},
|
||||
local_app::CursorHarness,
|
||||
plugin::{PluginRegistry, PluginRuntime},
|
||||
provider::ProviderRouter,
|
||||
search::WebCache,
|
||||
store::Store,
|
||||
@@ -37,17 +38,28 @@ impl App {
|
||||
}
|
||||
let assets = PromptAssets::embedded()?;
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let plugin_runtime = PluginRuntime::managed()?;
|
||||
let plugins = PluginRegistry::managed(
|
||||
store.clone(),
|
||||
plugin_runtime.clone(),
|
||||
config.app_version.clone(),
|
||||
)?;
|
||||
let provider = std::sync::Arc::new(ProviderRouter::new(
|
||||
store.clone(),
|
||||
plugins.clone(),
|
||||
config.provider_request_timeout,
|
||||
config.provider_stream_idle_timeout,
|
||||
));
|
||||
let registry = TransportRegistry::with_web_cache(
|
||||
let registry = TransportRegistry::with_plugins(
|
||||
store.clone(),
|
||||
provider.clone(),
|
||||
compiler,
|
||||
WebCache::managed()?,
|
||||
plugins.clone(),
|
||||
crate::config::managed_data_dir()?.join("rules"),
|
||||
);
|
||||
let control = control::ControlService::new(store.clone(), provider)?;
|
||||
let control =
|
||||
control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?;
|
||||
let harness = control.cursor_harness().clone();
|
||||
let mut router = api::router(registry.clone())?;
|
||||
router = match &config.console {
|
||||
|
||||
+28
-1
@@ -10,7 +10,8 @@ const DATA_DIR_NAME: &str = ".cursor-byok-v3";
|
||||
const DATABASE_FILE_NAME: &str = "cursor-byok.db";
|
||||
const V0049_DATA_DIR_NAME: &str = ".cursor-local-assistant-v2";
|
||||
const V0049_CONFIG_FILE_NAME: &str = "config.yaml";
|
||||
const DEFAULT_PROVIDER_REQUEST_TIMEOUT: Duration = Duration::from_secs(3000);
|
||||
const DEFAULT_PROVIDER_REQUEST_TIMEOUT: Duration = Duration::from_secs(60 * 60);
|
||||
const DEFAULT_PROVIDER_STREAM_IDLE_TIMEOUT: Duration = Duration::from_secs(30 * 60);
|
||||
|
||||
pub fn managed_data_dir() -> Result<PathBuf> {
|
||||
let home_dir = dirs::home_dir()
|
||||
@@ -45,6 +46,8 @@ pub struct ProviderConfig {
|
||||
pub custom_headers: reqwest::header::HeaderMap,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
pub request_timeout: Duration,
|
||||
pub retry_count: u32,
|
||||
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -52,8 +55,11 @@ pub struct Config {
|
||||
pub listen_addr: SocketAddr,
|
||||
pub database_url: String,
|
||||
pub provider_request_timeout: Duration,
|
||||
pub provider_stream_idle_timeout: Duration,
|
||||
pub console: Option<ConsoleSource>,
|
||||
pub use_persisted_ports: bool,
|
||||
/// 面向用户的应用版本;桌面壳会覆盖为自身版本,用于插件 minAppVersion 门控。
|
||||
pub app_version: String,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -102,8 +108,10 @@ impl Config {
|
||||
listen_addr,
|
||||
database_url: database_url_from_env()?,
|
||||
provider_request_timeout: request_timeout,
|
||||
provider_stream_idle_timeout: DEFAULT_PROVIDER_STREAM_IDLE_TIMEOUT,
|
||||
console,
|
||||
use_persisted_ports: false,
|
||||
app_version: env!("CARGO_PKG_VERSION").into(),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -114,8 +122,10 @@ impl Config {
|
||||
.expect("desktop listen address is static"),
|
||||
database_url: default_database_url()?,
|
||||
provider_request_timeout: DEFAULT_PROVIDER_REQUEST_TIMEOUT,
|
||||
provider_stream_idle_timeout: DEFAULT_PROVIDER_STREAM_IDLE_TIMEOUT,
|
||||
console: None,
|
||||
use_persisted_ports: true,
|
||||
app_version: env!("CARGO_PKG_VERSION").into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -142,3 +152,20 @@ fn database_url_for_dir(data_dir: &std::path::Path) -> Result<String> {
|
||||
.ok_or_else(|| Error::Config("database path is not valid UTF-8".into()))?;
|
||||
Ok(format!("sqlite://{database_path}"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn provider_timeout_defaults_match_runtime_boundaries() {
|
||||
assert_eq!(
|
||||
DEFAULT_PROVIDER_STREAM_IDLE_TIMEOUT,
|
||||
Duration::from_secs(30 * 60)
|
||||
);
|
||||
assert_eq!(
|
||||
DEFAULT_PROVIDER_REQUEST_TIMEOUT,
|
||||
Duration::from_secs(60 * 60)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ mod calls;
|
||||
mod harness;
|
||||
mod models;
|
||||
mod overview;
|
||||
mod plugins;
|
||||
mod service;
|
||||
mod settings;
|
||||
|
||||
@@ -136,6 +137,45 @@ pub fn api_router(service: ControlService) -> Router {
|
||||
)
|
||||
.route("/__byok-api__/api/llm-calls", get(calls::list))
|
||||
.route("/__byok-api__/api/llm-calls/{call_id}", get(calls::detail))
|
||||
.route("/__byok-api__/api/plugins", get(plugins::list))
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/runtime",
|
||||
get(plugins::runtime_status)
|
||||
.post(plugins::initialize_runtime)
|
||||
.delete(plugins::cancel_runtime_initialization),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/oauth/{session_id}/poll",
|
||||
post(plugins::oauth_poll),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/{plugin_id}",
|
||||
axum::routing::delete(plugins::remove),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/add/{method_id}/begin",
|
||||
post(plugins::oauth_begin),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/import",
|
||||
post(plugins::import),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/export",
|
||||
get(plugins::export_resources),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/{resource_id}",
|
||||
axum::routing::delete(plugins::delete_resource),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/{resource_id}/refresh",
|
||||
post(plugins::refresh_resource),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/plugins/{plugin_id}/providers/{provider_id}/models/sync",
|
||||
post(plugins::sync_models),
|
||||
)
|
||||
.route(
|
||||
"/__byok-api__/api/settings/observability",
|
||||
get(settings::get).put(settings::update),
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
//! Exposes plugin discovery, resource lifecycle, model sync, and runtime endpoints.
|
||||
use axum::{
|
||||
extract::{Path, State},
|
||||
http::StatusCode,
|
||||
Json,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
plugin::{
|
||||
ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginDescriptor,
|
||||
PluginRuntimeStatus,
|
||||
},
|
||||
Result,
|
||||
};
|
||||
|
||||
use super::ControlService;
|
||||
|
||||
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<PluginDescriptor>>> {
|
||||
Ok(Json(service.plugins().await))
|
||||
}
|
||||
|
||||
pub async fn remove(
|
||||
State(service): State<ControlService>,
|
||||
Path(plugin_id): Path<String>,
|
||||
) -> Result<StatusCode> {
|
||||
service.remove_plugin_configuration(&plugin_id).await?;
|
||||
Ok(StatusCode::NO_CONTENT)
|
||||
}
|
||||
|
||||
pub async fn oauth_begin(
|
||||
State(service): State<ControlService>,
|
||||
Path((plugin_id, resource_type, method_id)): Path<(String, String, String)>,
|
||||
) -> Result<Json<OAuthBeginResponse>> {
|
||||
Ok(Json(
|
||||
service
|
||||
.plugin_oauth_begin(&plugin_id, &resource_type, &method_id)
|
||||
.await?,
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn oauth_poll(
|
||||
State(service): State<ControlService>,
|
||||
Path(session_id): Path<String>,
|
||||
) -> Result<Json<OAuthPollResponse>> {
|
||||
Ok(Json(service.plugin_oauth_poll(&session_id).await?))
|
||||
}
|
||||
|
||||
pub async fn import(
|
||||
State(service): State<ControlService>,
|
||||
Path((plugin_id, resource_type)): Path<(String, String)>,
|
||||
Json(files): Json<serde_json::Value>,
|
||||
) -> Result<Json<ImportResponse>> {
|
||||
Ok(Json(
|
||||
service
|
||||
.plugin_import(&plugin_id, &resource_type, files)
|
||||
.await?,
|
||||
))
|
||||
}
|
||||
|
||||
/// 以附件形式返回账号资源导出文件,便于浏览器直接下载。
|
||||
pub async fn export_resources(
|
||||
State(service): State<ControlService>,
|
||||
Path((plugin_id, resource_type)): Path<(String, String)>,
|
||||
) -> Result<axum::response::Response> {
|
||||
let value = service
|
||||
.plugin_export_resources(&plugin_id, &resource_type)
|
||||
.await?;
|
||||
let body = serde_json::to_vec_pretty(&value)?;
|
||||
let response = axum::response::Response::builder()
|
||||
.header(axum::http::header::CONTENT_TYPE, "application/json")
|
||||
.header(
|
||||
axum::http::header::CONTENT_DISPOSITION,
|
||||
format!("attachment; filename=\"{plugin_id}-{resource_type}.json\""),
|
||||
)
|
||||
.body(axum::body::Body::from(body))
|
||||
.expect("static export response");
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
pub async fn refresh_resource(
|
||||
State(service): State<ControlService>,
|
||||
Path((plugin_id, resource_type, resource_id)): Path<(String, String, String)>,
|
||||
) -> Result<StatusCode> {
|
||||
service
|
||||
.plugin_refresh_resource(&plugin_id, &resource_type, &resource_id)
|
||||
.await?;
|
||||
Ok(StatusCode::NO_CONTENT)
|
||||
}
|
||||
|
||||
pub async fn delete_resource(
|
||||
State(service): State<ControlService>,
|
||||
Path((plugin_id, resource_type, resource_id)): Path<(String, String, String)>,
|
||||
) -> Result<StatusCode> {
|
||||
service
|
||||
.plugin_delete_resource(&plugin_id, &resource_type, &resource_id)
|
||||
.await?;
|
||||
Ok(StatusCode::NO_CONTENT)
|
||||
}
|
||||
|
||||
pub async fn sync_models(
|
||||
State(service): State<ControlService>,
|
||||
Path((plugin_id, provider_id)): Path<(String, String)>,
|
||||
) -> Result<Json<serde_json::Value>> {
|
||||
let count = service.plugin_sync_models(&plugin_id, &provider_id).await?;
|
||||
Ok(Json(serde_json::json!({ "models": count })))
|
||||
}
|
||||
|
||||
pub async fn runtime_status(
|
||||
State(service): State<ControlService>,
|
||||
) -> Result<Json<PluginRuntimeStatus>> {
|
||||
Ok(Json(service.plugin_runtime_status()))
|
||||
}
|
||||
|
||||
pub async fn initialize_runtime(
|
||||
State(service): State<ControlService>,
|
||||
) -> Result<Json<PluginRuntimeStatus>> {
|
||||
Ok(Json(service.initialize_plugin_runtime()))
|
||||
}
|
||||
|
||||
pub async fn cancel_runtime_initialization(
|
||||
State(service): State<ControlService>,
|
||||
) -> Result<Json<PluginRuntimeStatus>> {
|
||||
Ok(Json(service.cancel_plugin_runtime_initialization()))
|
||||
}
|
||||
@@ -25,6 +25,7 @@ use crate::{
|
||||
ModelRequest, ModelSpec, ModelType, Overview, ProjectedContent, ProjectedMessage,
|
||||
PromptSpec, ProviderType, Role,
|
||||
},
|
||||
plugin::{PluginDescriptor, PluginRegistry, PluginRuntime, PluginRuntimeStatus},
|
||||
provider::{is_valid_response_event, ModelEvent, Provider},
|
||||
store::{
|
||||
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store,
|
||||
@@ -38,6 +39,8 @@ pub struct ControlService {
|
||||
store: Store,
|
||||
cursor_harness: CursorHarness,
|
||||
provider: Arc<dyn Provider>,
|
||||
plugin_runtime: PluginRuntime,
|
||||
plugins: PluginRegistry,
|
||||
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
|
||||
}
|
||||
|
||||
@@ -143,11 +146,18 @@ pub struct ObservabilitySettings {
|
||||
}
|
||||
|
||||
impl ControlService {
|
||||
pub fn new(store: Store, provider: Arc<dyn Provider>) -> Result<Self> {
|
||||
pub fn new(
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
plugin_runtime: PluginRuntime,
|
||||
plugins: PluginRegistry,
|
||||
) -> Result<Self> {
|
||||
Ok(Self {
|
||||
cursor_harness: CursorHarness::new(store.clone())?,
|
||||
store,
|
||||
provider,
|
||||
plugin_runtime,
|
||||
plugins,
|
||||
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
|
||||
})
|
||||
}
|
||||
@@ -156,6 +166,91 @@ impl ControlService {
|
||||
&self.cursor_harness
|
||||
}
|
||||
|
||||
pub async fn plugins(&self) -> Vec<PluginDescriptor> {
|
||||
self.plugins.plugins().await
|
||||
}
|
||||
|
||||
pub async fn plugin_oauth_begin(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
method_id: &str,
|
||||
) -> Result<crate::plugin::OAuthBeginResponse> {
|
||||
self.plugins
|
||||
.oauth_begin(plugin_id, resource_type, method_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn plugin_oauth_poll(
|
||||
&self,
|
||||
session_id: &str,
|
||||
) -> Result<crate::plugin::OAuthPollResponse> {
|
||||
self.plugins.oauth_poll(session_id).await
|
||||
}
|
||||
|
||||
pub async fn plugin_import(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
files: serde_json::Value,
|
||||
) -> Result<crate::plugin::ImportResponse> {
|
||||
self.plugins
|
||||
.import_resources(plugin_id, resource_type, files)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn plugin_export_resources(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
) -> Result<serde_json::Value> {
|
||||
self.plugins
|
||||
.export_resources(plugin_id, resource_type)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn plugin_refresh_resource(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<()> {
|
||||
self.plugins
|
||||
.refresh_resource(plugin_id, resource_type, resource_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn plugin_delete_resource(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<()> {
|
||||
self.plugins
|
||||
.delete_resource(plugin_id, resource_type, resource_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn plugin_sync_models(&self, plugin_id: &str, provider_id: &str) -> Result<usize> {
|
||||
self.plugins.sync_models(plugin_id, provider_id).await
|
||||
}
|
||||
|
||||
pub async fn remove_plugin_configuration(&self, plugin_id: &str) -> Result<()> {
|
||||
self.plugins.remove(plugin_id).await
|
||||
}
|
||||
|
||||
pub fn plugin_runtime_status(&self) -> PluginRuntimeStatus {
|
||||
self.plugin_runtime.status()
|
||||
}
|
||||
|
||||
pub fn initialize_plugin_runtime(&self) -> PluginRuntimeStatus {
|
||||
self.plugin_runtime.initialize(self.store.clone())
|
||||
}
|
||||
|
||||
pub fn cancel_plugin_runtime_initialization(&self) -> PluginRuntimeStatus {
|
||||
self.plugin_runtime.cancel_initialization()
|
||||
}
|
||||
|
||||
pub(super) async fn ads(
|
||||
&self,
|
||||
disabled_ad_ids: Option<&str>,
|
||||
@@ -293,14 +388,20 @@ impl ControlService {
|
||||
const TEST_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(45);
|
||||
const TEST_PROMPT: &str = "Output the numbers 1 through 120 separated by a single space. No commas, no newlines, no explanation.";
|
||||
|
||||
let configured = self
|
||||
.store
|
||||
.model(model_hash)
|
||||
.await?
|
||||
.ok_or_else(|| Error::RunNotFound(format!("model {model_hash}")))?;
|
||||
let mut model = ModelSpec::new(model_hash);
|
||||
configured.configure(&mut model);
|
||||
model.max_output_tokens = Some(configured.max_output_tokens().unwrap_or(65_536));
|
||||
if model_hash.starts_with(crate::plugin::ADAPTER_ID_PREFIX) {
|
||||
let descriptor = self.plugins.model_descriptor(model_hash).await?;
|
||||
model.display_name = Some(descriptor.display_name);
|
||||
model.max_output_tokens = Some(descriptor.max_output_tokens.unwrap_or(65_536));
|
||||
} else {
|
||||
let configured = self
|
||||
.store
|
||||
.model(model_hash)
|
||||
.await?
|
||||
.ok_or_else(|| Error::RunNotFound(format!("model {model_hash}")))?;
|
||||
configured.configure(&mut model);
|
||||
model.max_output_tokens = Some(configured.max_output_tokens().unwrap_or(65_536));
|
||||
}
|
||||
let call_id = format!("model-test-{}", uuid::Uuid::new_v4());
|
||||
let invocation = ModelInvocation {
|
||||
call_id: call_id.clone(),
|
||||
@@ -699,6 +800,11 @@ async fn discover_models_from_endpoint(
|
||||
ProviderType::Anthropic => {
|
||||
anthropic_models(client, base_url, api_key, custom_headers).await?
|
||||
}
|
||||
ProviderType::Plugin => {
|
||||
return Err(Error::Config(
|
||||
"plugin providers discover models through their plugin".into(),
|
||||
))
|
||||
}
|
||||
};
|
||||
models.sort();
|
||||
models.dedup();
|
||||
|
||||
@@ -146,6 +146,37 @@ async fn decode_part<T: Message + Default>(
|
||||
.map_err(|error| Error::Protocol(format!("invalid {name} context Blob: {error}")))
|
||||
}
|
||||
|
||||
/// 把本地 md 规则目录(rules 服务的存储)合并进请求上下文,
|
||||
/// 使 BYOK 运行在 IDE 未携带这些规则时也能消费它们。
|
||||
/// 与 IDE 已发规则按内容去重;读取失败只告警,不影响运行。
|
||||
pub fn merge_local_rules(context: &mut pb::RequestContext, rules_dir: &Path) {
|
||||
let records = match crate::cursor::services::knowledge::RuleStore::open(rules_dir.into())
|
||||
.and_then(|store| store.list())
|
||||
{
|
||||
Ok(records) => records,
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "cannot read local rules; continuing without them");
|
||||
return;
|
||||
}
|
||||
};
|
||||
let existing = context
|
||||
.rules
|
||||
.iter()
|
||||
.chain(context.non_file_rules.iter())
|
||||
.map(|rule| rule.content.trim().to_owned())
|
||||
.chain(context.cloud_rule.iter().map(|rule| rule.trim().to_owned()))
|
||||
.collect::<HashSet<_>>();
|
||||
for record in records {
|
||||
if record.knowledge.trim().is_empty() || existing.contains(record.knowledge.trim()) {
|
||||
continue;
|
||||
}
|
||||
context.non_file_rules.push(pb::CursorRule {
|
||||
content: record.knowledge,
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub fn request_context(request: &pb::AgentRunRequest) -> Option<&pb::RequestContext> {
|
||||
let action = request.action.as_ref()?;
|
||||
action
|
||||
@@ -563,3 +594,48 @@ fn xml(value: &str) -> String {
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn rule(content: &str) -> pb::CursorRule {
|
||||
pb::CursorRule {
|
||||
content: content.into(),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_local_rules_appends_and_dedupes_by_content() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
std::fs::write(directory.path().join("a.md"), "shared rule").unwrap();
|
||||
std::fs::write(directory.path().join("b.md"), "local only rule").unwrap();
|
||||
std::fs::write(directory.path().join("c.md"), " \n").unwrap();
|
||||
|
||||
let mut context = pb::RequestContext {
|
||||
non_file_rules: vec![rule(" shared rule ")],
|
||||
..Default::default()
|
||||
};
|
||||
merge_local_rules(&mut context, directory.path());
|
||||
|
||||
let contents = context
|
||||
.non_file_rules
|
||||
.iter()
|
||||
.map(|rule| rule.content.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
contents,
|
||||
[" shared rule ", "local only rule"],
|
||||
"IDE-sent duplicate is kept once and blank local rules are skipped"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_local_rules_survives_a_missing_directory() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let mut context = pb::RequestContext::default();
|
||||
merge_local_rules(&mut context, &directory.path().join("nested/rules"));
|
||||
assert!(context.non_file_rules.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,6 +51,7 @@ pub(crate) struct PrepareDependencies<'a> {
|
||||
pub checkpoint: &'a CheckpointBuilder,
|
||||
pub blob_sync: &'a BlobSynchronizer,
|
||||
pub context_sync: &'a RequestContextSynchronizer,
|
||||
pub local_rules_dir: Option<&'a std::path::Path>,
|
||||
}
|
||||
|
||||
pub(crate) async fn prepare(
|
||||
@@ -64,6 +65,7 @@ pub(crate) async fn prepare(
|
||||
checkpoint,
|
||||
blob_sync,
|
||||
context_sync,
|
||||
local_rules_dir,
|
||||
} = dependencies;
|
||||
checkpoint
|
||||
.import_prefetched(&request.pre_fetched_blobs)
|
||||
@@ -119,7 +121,11 @@ pub(crate) async fn prepare(
|
||||
.artifact("history_projection", "byok_server", &encoded, summary)
|
||||
.await;
|
||||
}
|
||||
let request_context = context::hydrate(request, context_sync).await?;
|
||||
let mut request_context = context::hydrate(request, context_sync).await?;
|
||||
if let Some(rules_dir) = local_rules_dir {
|
||||
context::merge_local_rules(&mut request_context, rules_dir);
|
||||
}
|
||||
let request_context = request_context;
|
||||
let ActionProjection {
|
||||
mode: mode_number,
|
||||
mut turn_user,
|
||||
|
||||
@@ -343,7 +343,14 @@ impl ConversationOutput {
|
||||
let call = calls.get_mut(&index).ok_or_else(|| {
|
||||
Error::Protocol(format!("unknown completed tool index: {index}"))
|
||||
})?;
|
||||
call.arguments = serde_json::from_str(&call.arguments_text)?;
|
||||
// A tool call with no arguments streams no argument text.
|
||||
// Treat empty text as an empty object, matching the model
|
||||
// cycle, instead of failing the run on `from_str("")`.
|
||||
call.arguments = if call.arguments_text.trim().is_empty() {
|
||||
serde_json::json!({})
|
||||
} else {
|
||||
serde_json::from_str(&call.arguments_text)?
|
||||
};
|
||||
}
|
||||
RunEvent::Usage(usage) => {
|
||||
if !self.context.compacting {
|
||||
|
||||
@@ -26,6 +26,8 @@ pub(crate) struct ConversationDependencies {
|
||||
pub provider: Arc<dyn Provider>,
|
||||
pub compiler: PromptCompiler,
|
||||
pub web_cache: WebCache,
|
||||
/// 本地 rules 服务的 md 存储目录;编译请求上下文时合并其中的规则。
|
||||
pub local_rules_dir: Option<std::path::PathBuf>,
|
||||
}
|
||||
|
||||
struct RegistryInner {
|
||||
@@ -47,6 +49,7 @@ impl ConversationRegistry {
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
local_rules_dir: Option<std::path::PathBuf>,
|
||||
) -> Self {
|
||||
Self {
|
||||
inner: Arc::new(RegistryInner {
|
||||
@@ -58,6 +61,7 @@ impl ConversationRegistry {
|
||||
provider,
|
||||
compiler,
|
||||
web_cache,
|
||||
local_rules_dir,
|
||||
},
|
||||
}),
|
||||
}
|
||||
|
||||
@@ -471,6 +471,7 @@ fn spawn_run_request(
|
||||
checkpoint: &checkpoint,
|
||||
blob_sync: &blob_sync,
|
||||
context_sync: &context_sync,
|
||||
local_rules_dir: dependencies.local_rules_dir.as_deref(),
|
||||
},
|
||||
) => prepared,
|
||||
};
|
||||
|
||||
@@ -254,7 +254,7 @@ async fn forward_or(
|
||||
match proxy::forward_buffered(&upstream, request).await {
|
||||
Ok(response) if response.status.is_success() => Ok(response.into_response()),
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "Cursor identity upstream rejected request; using local identity");
|
||||
tracing::debug!(status = %response.status, "Cursor identity upstream rejected request; using local identity");
|
||||
fallback()
|
||||
}
|
||||
Err(error) => {
|
||||
|
||||
@@ -0,0 +1,356 @@
|
||||
//! Serves Cursor user rules: upstream-first with an offline markdown cache.
|
||||
//!
|
||||
//! 每个请求先回放离线日志再尝试上游;上游成功时把结果写穿到本地镜像,
|
||||
//! 上游不可达时降级为本地 md 存储并记录日志等待回放。
|
||||
mod store;
|
||||
mod sync;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::{
|
||||
body::{to_bytes, Body, Bytes},
|
||||
extract::Extension,
|
||||
http::{header, Request, Response},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
use crate::{api::cursor::proxy, config, cursor::protocol::connect, Result};
|
||||
|
||||
pub(crate) use store::{RuleRecord, RuleStore};
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseAddRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "2")]
|
||||
title: String,
|
||||
#[prost(string, tag = "3")]
|
||||
git_origin: String,
|
||||
#[prost(string, optional, tag = "4")]
|
||||
composer_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseAddResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
#[prost(string, tag = "2")]
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseListRequest {
|
||||
#[prost(int32, optional, tag = "1")]
|
||||
limit: Option<i32>,
|
||||
#[prost(string, optional, tag = "2")]
|
||||
git_origin: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseListResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
#[prost(message, repeated, tag = "2")]
|
||||
all_results: Vec<KnowledgeBaseListItem>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseListItem {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
#[prost(string, tag = "2")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "3")]
|
||||
title: String,
|
||||
#[prost(string, tag = "4")]
|
||||
created_at: String,
|
||||
#[prost(bool, tag = "5")]
|
||||
is_generated: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseUpdateRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
#[prost(string, tag = "2")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "3")]
|
||||
title: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseUpdateResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseRemoveRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseRemoveResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
}
|
||||
|
||||
/// 规则存储与并发锁;经 axum Extension 注入四个 handler。
|
||||
#[derive(Clone)]
|
||||
pub struct KnowledgeService {
|
||||
inner: Arc<Inner>,
|
||||
}
|
||||
|
||||
struct Inner {
|
||||
store: RuleStore,
|
||||
lock: tokio::sync::Mutex<()>,
|
||||
}
|
||||
|
||||
impl KnowledgeService {
|
||||
pub fn managed() -> Result<Self> {
|
||||
Self::with_root(config::managed_data_dir()?.join("rules"))
|
||||
}
|
||||
|
||||
/// 指定存储根目录构造;managed() 与集成测试共用。
|
||||
pub fn with_root(root: std::path::PathBuf) -> Result<Self> {
|
||||
Ok(Self {
|
||||
inner: Arc::new(Inner {
|
||||
store: RuleStore::open(root)?,
|
||||
lock: tokio::sync::Mutex::new(()),
|
||||
}),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn add(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseAddRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&response.body)
|
||||
{
|
||||
if reply.success && !reply.id.is_empty() {
|
||||
store.upsert(&RuleRecord {
|
||||
id: reply.id,
|
||||
knowledge: message.knowledge,
|
||||
title: message.title,
|
||||
created_at: now(),
|
||||
is_generated: false,
|
||||
git_origin: message.git_origin,
|
||||
})?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected add; storing locally");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for add; storing locally");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let id = format!("{}{}", store::LOCAL_ID_PREFIX, uuid::Uuid::new_v4());
|
||||
store.upsert(&RuleRecord {
|
||||
id: id.clone(),
|
||||
knowledge: message.knowledge,
|
||||
title: message.title,
|
||||
created_at: now(),
|
||||
is_generated: false,
|
||||
git_origin: message.git_origin,
|
||||
})?;
|
||||
store.record_add(&id)?;
|
||||
proto(KnowledgeBaseAddResponse { success: true, id })
|
||||
}
|
||||
|
||||
pub async fn list(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseListRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
let git_origin = message.git_origin.unwrap_or_default();
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) =
|
||||
connect::decode_unary::<KnowledgeBaseListResponse>(&response.body)
|
||||
{
|
||||
// 带 git_origin 过滤的列表只是子集,整体覆盖会误删其他规则。
|
||||
if reply.success && git_origin.is_empty() {
|
||||
sync::mirror(store, reply.all_results)?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected list; serving local cache");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for list; serving local cache");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut records = store.list()?;
|
||||
if !git_origin.is_empty() {
|
||||
records.retain(|record| record.git_origin == git_origin);
|
||||
}
|
||||
if let Some(limit) = message.limit {
|
||||
if limit >= 0 {
|
||||
records.truncate(limit as usize);
|
||||
}
|
||||
}
|
||||
proto(KnowledgeBaseListResponse {
|
||||
success: true,
|
||||
all_results: records
|
||||
.into_iter()
|
||||
.map(|record| KnowledgeBaseListItem {
|
||||
id: record.id,
|
||||
knowledge: record.knowledge,
|
||||
title: record.title,
|
||||
created_at: record.created_at,
|
||||
is_generated: record.is_generated,
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn update(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseUpdateRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) =
|
||||
connect::decode_unary::<KnowledgeBaseUpdateResponse>(&response.body)
|
||||
{
|
||||
if reply.success {
|
||||
let existing = store.get(&message.id)?;
|
||||
store.upsert(&RuleRecord {
|
||||
id: message.id,
|
||||
knowledge: message.knowledge,
|
||||
title: message.title,
|
||||
created_at: existing
|
||||
.as_ref()
|
||||
.map_or_else(now, |record| record.created_at.clone()),
|
||||
is_generated: existing
|
||||
.as_ref()
|
||||
.is_some_and(|record| record.is_generated),
|
||||
git_origin: existing
|
||||
.map(|record| record.git_origin)
|
||||
.unwrap_or_default(),
|
||||
})?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected update; storing locally");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for update; storing locally");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let Some(mut record) = store.get(&message.id)? else {
|
||||
return proto(KnowledgeBaseUpdateResponse { success: false });
|
||||
};
|
||||
record.knowledge = message.knowledge;
|
||||
record.title = message.title;
|
||||
store.upsert(&record)?;
|
||||
store.record_update(&message.id)?;
|
||||
proto(KnowledgeBaseUpdateResponse { success: true })
|
||||
}
|
||||
|
||||
pub async fn remove(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseRemoveRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) =
|
||||
connect::decode_unary::<KnowledgeBaseRemoveResponse>(&response.body)
|
||||
{
|
||||
if reply.success {
|
||||
store.remove(&message.id)?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected remove; removing locally");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for remove; removing locally");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
store.remove(&message.id)?;
|
||||
store.record_remove(&message.id)?;
|
||||
proto(KnowledgeBaseRemoveResponse { success: true })
|
||||
}
|
||||
|
||||
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
|
||||
let (parts, body) = request.into_parts();
|
||||
let body = to_bytes(body, usize::MAX)
|
||||
.await
|
||||
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
|
||||
Ok((parts, body))
|
||||
}
|
||||
|
||||
fn now() -> String {
|
||||
chrono::Utc::now().to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
|
||||
}
|
||||
|
||||
fn proto(message: impl Message) -> Result<Response<Body>> {
|
||||
let body = message.encode_to_vec();
|
||||
let length = body.len();
|
||||
let mut response = Response::new(Body::from(body));
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_TYPE,
|
||||
axum::http::HeaderValue::from_static("application/proto"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_LENGTH,
|
||||
length
|
||||
.to_string()
|
||||
.parse()
|
||||
.expect("body length is always a valid header value"),
|
||||
);
|
||||
Ok(response)
|
||||
}
|
||||
@@ -0,0 +1,487 @@
|
||||
//! Persists rules as markdown files with a JSON metadata sidecar.
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
path::{Path, PathBuf},
|
||||
};
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
const META_FILE: &str = "meta.json";
|
||||
const RULE_EXTENSION: &str = "md";
|
||||
pub const LOCAL_ID_PREFIX: &str = "local-";
|
||||
|
||||
/// 一条规则的完整视图:knowledge 来自 md 文件,其余字段来自 meta.json。
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct RuleRecord {
|
||||
pub id: String,
|
||||
pub knowledge: String,
|
||||
pub title: String,
|
||||
pub created_at: String,
|
||||
pub is_generated: bool,
|
||||
pub git_origin: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum JournalOp {
|
||||
Add,
|
||||
Update,
|
||||
Remove,
|
||||
}
|
||||
|
||||
/// 离线期间未同步到上游的一次变更。
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct JournalEntry {
|
||||
pub op: JournalOp,
|
||||
pub id: String,
|
||||
}
|
||||
|
||||
#[derive(Default, Serialize, Deserialize)]
|
||||
struct Meta {
|
||||
#[serde(default)]
|
||||
rules: BTreeMap<String, RuleMeta>,
|
||||
#[serde(default)]
|
||||
journal: Vec<JournalEntry>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Serialize, Deserialize)]
|
||||
struct RuleMeta {
|
||||
#[serde(default)]
|
||||
title: String,
|
||||
#[serde(default)]
|
||||
created_at: String,
|
||||
#[serde(default)]
|
||||
is_generated: bool,
|
||||
#[serde(default)]
|
||||
git_origin: String,
|
||||
}
|
||||
|
||||
/// md 文件为核心的规则存储;调用方需自行串行化并发访问。
|
||||
pub struct RuleStore {
|
||||
root: PathBuf,
|
||||
}
|
||||
|
||||
impl RuleStore {
|
||||
pub fn open(root: PathBuf) -> Result<Self> {
|
||||
std::fs::create_dir_all(&root)?;
|
||||
Ok(Self { root })
|
||||
}
|
||||
|
||||
pub fn list(&self) -> Result<Vec<RuleRecord>> {
|
||||
let meta = self.read_meta();
|
||||
let mut records = Vec::new();
|
||||
for entry in std::fs::read_dir(&self.root)? {
|
||||
let path = entry?.path();
|
||||
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
|
||||
continue;
|
||||
}
|
||||
let Some(id) = path.file_stem().and_then(|value| value.to_str()) else {
|
||||
continue;
|
||||
};
|
||||
if validate_id(id).is_err() {
|
||||
continue;
|
||||
}
|
||||
let knowledge = std::fs::read_to_string(&path)?;
|
||||
records.push(assemble(id, knowledge, meta.rules.get(id), &path));
|
||||
}
|
||||
records.sort_by(|left, right| {
|
||||
timestamp(&right.created_at)
|
||||
.cmp(×tamp(&left.created_at))
|
||||
.then_with(|| left.id.cmp(&right.id))
|
||||
});
|
||||
Ok(records)
|
||||
}
|
||||
|
||||
pub fn get(&self, id: &str) -> Result<Option<RuleRecord>> {
|
||||
validate_id(id)?;
|
||||
let path = self.rule_path(id);
|
||||
let knowledge = match std::fs::read_to_string(&path) {
|
||||
Ok(knowledge) => knowledge,
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
|
||||
Err(error) => return Err(error.into()),
|
||||
};
|
||||
let meta = self.read_meta();
|
||||
Ok(Some(assemble(id, knowledge, meta.rules.get(id), &path)))
|
||||
}
|
||||
|
||||
pub fn upsert(&self, record: &RuleRecord) -> Result<()> {
|
||||
validate_id(&record.id)?;
|
||||
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
|
||||
let mut meta = self.read_meta();
|
||||
meta.rules.insert(record.id.clone(), rule_meta(record));
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
pub fn remove(&self, id: &str) -> Result<()> {
|
||||
validate_id(id)?;
|
||||
remove_file_if_exists(&self.rule_path(id))?;
|
||||
let mut meta = self.read_meta();
|
||||
if meta.rules.remove(id).is_some() {
|
||||
self.write_meta(&meta)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 离线新增的规则在上游落地后,把本地临时 id 换成上游分配的真实 id。
|
||||
pub fn promote(&self, old_id: &str, new_id: &str) -> Result<()> {
|
||||
validate_id(old_id)?;
|
||||
validate_id(new_id)?;
|
||||
let source = self.rule_path(old_id);
|
||||
let target = self.rule_path(new_id);
|
||||
#[cfg(windows)]
|
||||
remove_file_if_exists(&target)?;
|
||||
std::fs::rename(&source, &target)?;
|
||||
let mut meta = self.read_meta();
|
||||
if let Some(rule) = meta.rules.remove(old_id) {
|
||||
meta.rules.insert(new_id.into(), rule);
|
||||
}
|
||||
for entry in &mut meta.journal {
|
||||
if entry.id == old_id {
|
||||
entry.id = new_id.into();
|
||||
}
|
||||
}
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
/// 用上游的完整列表覆盖本地镜像;仅应在日志为空(已全部回放)时调用。
|
||||
pub fn replace_all(&self, records: &[RuleRecord]) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
meta.rules.clear();
|
||||
for record in records {
|
||||
validate_id(&record.id)?;
|
||||
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
|
||||
meta.rules.insert(record.id.clone(), rule_meta(record));
|
||||
}
|
||||
for entry in std::fs::read_dir(&self.root)? {
|
||||
let path = entry?.path();
|
||||
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
|
||||
continue;
|
||||
}
|
||||
let keep = path
|
||||
.file_stem()
|
||||
.and_then(|value| value.to_str())
|
||||
.is_some_and(|id| meta.rules.contains_key(id));
|
||||
if !keep {
|
||||
remove_file_if_exists(&path)?;
|
||||
}
|
||||
}
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
pub fn journal_front(&self) -> Result<Option<JournalEntry>> {
|
||||
Ok(self.read_meta().journal.first().cloned())
|
||||
}
|
||||
|
||||
pub fn pop_journal(&self) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
if !meta.journal.is_empty() {
|
||||
meta.journal.remove(0);
|
||||
self.write_meta(&meta)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn record_add(&self, id: &str) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
meta.journal.push(JournalEntry {
|
||||
op: JournalOp::Add,
|
||||
id: id.into(),
|
||||
});
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
pub fn record_update(&self, id: &str) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
if journal_contains(&meta.journal, id, JournalOp::Add) {
|
||||
// 回放 add 时会读取最新内容,无需单独的 update 日志。
|
||||
return Ok(());
|
||||
}
|
||||
let op = if id.starts_with(LOCAL_ID_PREFIX) {
|
||||
// 本地临时 id 没有对应的 add 日志(如镜像覆盖后的残留),按新增回放。
|
||||
JournalOp::Add
|
||||
} else {
|
||||
JournalOp::Update
|
||||
};
|
||||
if !journal_contains(&meta.journal, id, op) {
|
||||
meta.journal.push(JournalEntry { op, id: id.into() });
|
||||
self.write_meta(&meta)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn record_remove(&self, id: &str) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
let never_synced = journal_contains(&meta.journal, id, JournalOp::Add);
|
||||
meta.journal.retain(|entry| entry.id != id);
|
||||
if !never_synced && !id.starts_with(LOCAL_ID_PREFIX) {
|
||||
meta.journal.push(JournalEntry {
|
||||
op: JournalOp::Remove,
|
||||
id: id.into(),
|
||||
});
|
||||
}
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
fn rule_path(&self, id: &str) -> PathBuf {
|
||||
self.root.join(format!("{id}.{RULE_EXTENSION}"))
|
||||
}
|
||||
|
||||
fn meta_path(&self) -> PathBuf {
|
||||
self.root.join(META_FILE)
|
||||
}
|
||||
|
||||
fn read_meta(&self) -> Meta {
|
||||
match std::fs::read(self.meta_path()) {
|
||||
Ok(bytes) => serde_json::from_slice(&bytes).unwrap_or_else(|error| {
|
||||
tracing::warn!(%error, "rules meta.json is corrupt; starting from empty metadata");
|
||||
Meta::default()
|
||||
}),
|
||||
Err(_) => Meta::default(),
|
||||
}
|
||||
}
|
||||
|
||||
fn write_meta(&self, meta: &Meta) -> Result<()> {
|
||||
write_atomic(&self.meta_path(), &serde_json::to_vec_pretty(meta)?)
|
||||
}
|
||||
}
|
||||
|
||||
fn assemble(id: &str, knowledge: String, meta: Option<&RuleMeta>, path: &Path) -> RuleRecord {
|
||||
match meta {
|
||||
Some(meta) => RuleRecord {
|
||||
id: id.into(),
|
||||
knowledge,
|
||||
title: meta.title.clone(),
|
||||
created_at: meta.created_at.clone(),
|
||||
is_generated: meta.is_generated,
|
||||
git_origin: meta.git_origin.clone(),
|
||||
},
|
||||
// 用户手放的 md 文件没有元数据,用文件名当标题、修改时间当创建时间。
|
||||
None => RuleRecord {
|
||||
id: id.into(),
|
||||
knowledge,
|
||||
title: id.into(),
|
||||
created_at: file_modified_at(path),
|
||||
is_generated: false,
|
||||
git_origin: String::new(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn rule_meta(record: &RuleRecord) -> RuleMeta {
|
||||
RuleMeta {
|
||||
title: record.title.clone(),
|
||||
created_at: record.created_at.clone(),
|
||||
is_generated: record.is_generated,
|
||||
git_origin: record.git_origin.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn journal_contains(journal: &[JournalEntry], id: &str, op: JournalOp) -> bool {
|
||||
journal.iter().any(|entry| entry.id == id && entry.op == op)
|
||||
}
|
||||
|
||||
fn timestamp(created_at: &str) -> i64 {
|
||||
chrono::DateTime::parse_from_rfc3339(created_at)
|
||||
.map(|time| time.timestamp_millis())
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
fn file_modified_at(path: &Path) -> String {
|
||||
let modified = std::fs::metadata(path)
|
||||
.and_then(|meta| meta.modified())
|
||||
.unwrap_or_else(|_| std::time::SystemTime::now());
|
||||
chrono::DateTime::<chrono::Utc>::from(modified)
|
||||
.to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
|
||||
}
|
||||
|
||||
fn validate_id(id: &str) -> Result<()> {
|
||||
if id.is_empty()
|
||||
|| id.len() > 128
|
||||
|| !id
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'))
|
||||
{
|
||||
return Err(Error::Protocol(format!("invalid rule id: {id:?}")));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn remove_file_if_exists(path: &Path) -> Result<()> {
|
||||
match std::fs::remove_file(path) {
|
||||
Ok(()) => Ok(()),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
fn write_atomic(path: &Path, bytes: &[u8]) -> Result<()> {
|
||||
use std::io::Write;
|
||||
let directory = path.parent().expect("rule path has a parent");
|
||||
let temporary = directory.join(format!(".{}.tmp", uuid::Uuid::new_v4()));
|
||||
let mut file = std::fs::File::create(&temporary)?;
|
||||
file.write_all(bytes)?;
|
||||
file.sync_all()?;
|
||||
drop(file);
|
||||
#[cfg(windows)]
|
||||
remove_file_if_exists(path)?;
|
||||
std::fs::rename(&temporary, path).inspect_err(|_| {
|
||||
let _ = std::fs::remove_file(&temporary);
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn record(id: &str, knowledge: &str, created_at: &str) -> RuleRecord {
|
||||
RuleRecord {
|
||||
id: id.into(),
|
||||
knowledge: knowledge.into(),
|
||||
title: format!("title-{id}"),
|
||||
created_at: created_at.into(),
|
||||
is_generated: false,
|
||||
git_origin: String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn journal(store: &RuleStore) -> Vec<JournalEntry> {
|
||||
store.read_meta().journal
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn upserts_lists_and_removes_rules() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
store
|
||||
.upsert(&record("100", "older", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
store
|
||||
.upsert(&record("200", "newer", "2026-02-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
|
||||
let listed = store.list().unwrap();
|
||||
assert_eq!(
|
||||
listed
|
||||
.iter()
|
||||
.map(|rule| rule.id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["200", "100"],
|
||||
"list is sorted by created_at descending"
|
||||
);
|
||||
assert_eq!(listed[0].knowledge, "newer");
|
||||
assert_eq!(listed[0].title, "title-200");
|
||||
|
||||
store.remove("200").unwrap();
|
||||
assert!(store.get("200").unwrap().is_none());
|
||||
assert_eq!(store.list().unwrap().len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_path_traversal_ids() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
assert!(store.get("../escape").is_err());
|
||||
assert!(store.get("a/b").is_err());
|
||||
assert!(store.get("").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compacts_offline_journal() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
|
||||
// 离线新增后再更新:回放 add 即可携带最新内容,不产生 update 日志。
|
||||
store
|
||||
.upsert(&record("local-a", "v1", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
store.record_add("local-a").unwrap();
|
||||
store.record_update("local-a").unwrap();
|
||||
assert_eq!(
|
||||
journal(&store),
|
||||
vec![JournalEntry {
|
||||
op: JournalOp::Add,
|
||||
id: "local-a".into()
|
||||
}]
|
||||
);
|
||||
|
||||
// 离线新增后又删除:上游从未见过它,日志清空。
|
||||
store.record_remove("local-a").unwrap();
|
||||
assert!(journal(&store).is_empty());
|
||||
|
||||
// 更新上游已有规则:多次更新合并为一条;删除后 update 日志被顶替。
|
||||
store.record_update("42").unwrap();
|
||||
store.record_update("42").unwrap();
|
||||
assert_eq!(
|
||||
journal(&store),
|
||||
vec![JournalEntry {
|
||||
op: JournalOp::Update,
|
||||
id: "42".into()
|
||||
}]
|
||||
);
|
||||
store.record_remove("42").unwrap();
|
||||
assert_eq!(
|
||||
journal(&store),
|
||||
vec![JournalEntry {
|
||||
op: JournalOp::Remove,
|
||||
id: "42".into()
|
||||
}]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn promote_renames_rule_and_journal_ids() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
store
|
||||
.upsert(&record("local-a", "content", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
store.record_add("local-a").unwrap();
|
||||
|
||||
store.promote("local-a", "17353272").unwrap();
|
||||
|
||||
assert!(store.get("local-a").unwrap().is_none());
|
||||
let promoted = store.get("17353272").unwrap().unwrap();
|
||||
assert_eq!(promoted.knowledge, "content");
|
||||
assert_eq!(promoted.title, "title-local-a");
|
||||
assert_eq!(journal(&store)[0].id, "17353272");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replace_all_mirrors_upstream_state() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
store
|
||||
.upsert(&record("stale", "gone soon", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
|
||||
store
|
||||
.replace_all(&[record(
|
||||
"17353272",
|
||||
"from upstream",
|
||||
"2026-02-01T00:00:00.000Z",
|
||||
)])
|
||||
.unwrap();
|
||||
|
||||
let listed = store.list().unwrap();
|
||||
assert_eq!(listed.len(), 1);
|
||||
assert_eq!(listed[0].id, "17353272");
|
||||
assert_eq!(listed[0].knowledge, "from upstream");
|
||||
assert!(store.get("stale").unwrap().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lists_hand_written_markdown_without_metadata() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
std::fs::write(root.path().join("rules/manual_rule.md"), "hand written").unwrap();
|
||||
|
||||
let listed = store.list().unwrap();
|
||||
assert_eq!(listed.len(), 1);
|
||||
assert_eq!(listed[0].id, "manual_rule");
|
||||
assert_eq!(listed[0].title, "manual_rule");
|
||||
assert_eq!(listed[0].knowledge, "hand written");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,184 @@
|
||||
//! Replays the offline journal to upstream and mirrors upstream list state.
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http::{header, HeaderMap, HeaderValue, Method, Request},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
use crate::{api::cursor::proxy, cursor::protocol::connect, Result};
|
||||
|
||||
use super::{
|
||||
store::{JournalOp, RuleRecord, RuleStore},
|
||||
KnowledgeBaseAddRequest, KnowledgeBaseAddResponse, KnowledgeBaseListItem,
|
||||
KnowledgeBaseRemoveRequest, KnowledgeBaseRemoveResponse, KnowledgeBaseUpdateRequest,
|
||||
KnowledgeBaseUpdateResponse,
|
||||
};
|
||||
|
||||
const ADD_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseAdd";
|
||||
const UPDATE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseUpdate";
|
||||
const REMOVE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseRemove";
|
||||
|
||||
/// 逐条把离线日志推送到上游。返回 true 表示日志已清空(上游可用),
|
||||
/// false 表示上游不可达,剩余日志保留、调用方应降级到本地。
|
||||
pub async fn replay(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
) -> Result<bool> {
|
||||
while let Some(entry) = store.journal_front()? {
|
||||
let advanced = match entry.op {
|
||||
JournalOp::Add => replay_add(upstream, headers, store, &entry.id).await?,
|
||||
JournalOp::Update => replay_update(upstream, headers, store, &entry.id).await?,
|
||||
JournalOp::Remove => replay_remove(upstream, headers, store, &entry.id).await?,
|
||||
};
|
||||
if !advanced {
|
||||
return Ok(false);
|
||||
}
|
||||
}
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
/// 用上游返回的完整列表覆盖本地镜像。仅应在日志已清空时调用。
|
||||
pub fn mirror(store: &RuleStore, items: Vec<KnowledgeBaseListItem>) -> Result<()> {
|
||||
let records = items
|
||||
.into_iter()
|
||||
.map(|item| RuleRecord {
|
||||
id: item.id,
|
||||
knowledge: item.knowledge,
|
||||
title: item.title,
|
||||
created_at: item.created_at,
|
||||
is_generated: item.is_generated,
|
||||
git_origin: String::new(),
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
store.replace_all(&records)
|
||||
}
|
||||
|
||||
async fn replay_add(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
id: &str,
|
||||
) -> Result<bool> {
|
||||
let Some(record) = store.get(id)? else {
|
||||
// 规则文件已不在(被手动删除等),日志作废。
|
||||
store.pop_journal()?;
|
||||
return Ok(true);
|
||||
};
|
||||
let message = KnowledgeBaseAddRequest {
|
||||
knowledge: record.knowledge,
|
||||
title: record.title,
|
||||
git_origin: record.git_origin,
|
||||
composer_id: None,
|
||||
};
|
||||
let Some(body) = send(upstream, headers, ADD_PATH, &message).await else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&body) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !reply.success || reply.id.is_empty() {
|
||||
tracing::warn!(
|
||||
id,
|
||||
"rules upstream declined replayed add; dropping journal entry"
|
||||
);
|
||||
store.pop_journal()?;
|
||||
return Ok(true);
|
||||
}
|
||||
store.promote(id, &reply.id)?;
|
||||
store.pop_journal()?;
|
||||
tracing::info!(
|
||||
local_id = id,
|
||||
upstream_id = reply.id,
|
||||
"replayed offline rule add to upstream"
|
||||
);
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn replay_update(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
id: &str,
|
||||
) -> Result<bool> {
|
||||
let Some(record) = store.get(id)? else {
|
||||
store.pop_journal()?;
|
||||
return Ok(true);
|
||||
};
|
||||
let message = KnowledgeBaseUpdateRequest {
|
||||
id: id.into(),
|
||||
knowledge: record.knowledge,
|
||||
title: record.title,
|
||||
};
|
||||
let Some(body) = send(upstream, headers, UPDATE_PATH, &message).await else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Ok(reply) = connect::decode_unary::<KnowledgeBaseUpdateResponse>(&body) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !reply.success {
|
||||
tracing::warn!(
|
||||
id,
|
||||
"rules upstream declined replayed update; dropping journal entry"
|
||||
);
|
||||
}
|
||||
store.pop_journal()?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn replay_remove(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
id: &str,
|
||||
) -> Result<bool> {
|
||||
let message = KnowledgeBaseRemoveRequest { id: id.into() };
|
||||
let Some(body) = send(upstream, headers, REMOVE_PATH, &message).await else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Ok(reply) = connect::decode_unary::<KnowledgeBaseRemoveResponse>(&body) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !reply.success {
|
||||
tracing::warn!(
|
||||
id,
|
||||
"rules upstream declined replayed remove; dropping journal entry"
|
||||
);
|
||||
}
|
||||
store.pop_journal()?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
/// 以当前请求的头为模板向上游发起一次 unary RPC。
|
||||
/// 成功(2xx)返回响应体;不可达或被拒绝返回 None,由调用方保留日志。
|
||||
async fn send(
|
||||
upstream: &proxy::CursorProxy,
|
||||
template: &HeaderMap,
|
||||
path: &str,
|
||||
message: &impl Message,
|
||||
) -> Option<Bytes> {
|
||||
let mut headers = template.clone();
|
||||
// 模板里的上游 URL 头指向原始 RPC 路径,必须移除才能命中回放路径。
|
||||
headers.remove(proxy::UPSTREAM_URL_HEADER);
|
||||
headers.remove(header::CONTENT_LENGTH);
|
||||
headers.insert(
|
||||
header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("application/proto"),
|
||||
);
|
||||
let mut request = Request::new(Body::from(message.encode_to_vec()));
|
||||
*request.method_mut() = Method::POST;
|
||||
*request.uri_mut() = path.parse().expect("replay path is a valid URI");
|
||||
*request.headers_mut() = headers;
|
||||
|
||||
match proxy::forward_buffered(upstream, request).await {
|
||||
Ok(response) if response.status.is_success() => Some(response.body),
|
||||
Ok(response) => {
|
||||
tracing::warn!(path, status = %response.status, "rules journal replay rejected by upstream");
|
||||
None
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(path, %error, "rules journal replay cannot reach upstream");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ pub mod account;
|
||||
pub mod analytics;
|
||||
pub mod blob_sync;
|
||||
pub mod context_sync;
|
||||
pub mod knowledge;
|
||||
pub mod model_catalog;
|
||||
pub mod observability;
|
||||
pub mod tab;
|
||||
|
||||
@@ -10,7 +10,8 @@ use prost::Message;
|
||||
use crate::{
|
||||
api::cursor::proxy::{self, CursorProxy},
|
||||
cursor::{protocol::proto::agent::v1 as agent, transport::TransportRegistry},
|
||||
model::{format_token_count, parse_token_count, ModelConfig, ModelType},
|
||||
model::{format_token_count, parse_token_count, ModelConfig},
|
||||
plugin::PluginModelDescriptor,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
@@ -186,9 +187,10 @@ struct UsableModelsAddition {
|
||||
models: Vec<agent::ModelDetails>,
|
||||
}
|
||||
|
||||
const CONTEXTS: [(&str, &str); 4] = [
|
||||
const CONTEXTS: [(&str, &str); 5] = [
|
||||
("200k", "200K"),
|
||||
("356k", "356K"),
|
||||
("500k", "500K"),
|
||||
("800k", "800K"),
|
||||
("1m", "1M"),
|
||||
];
|
||||
@@ -201,12 +203,12 @@ const EFFORTS: [(&str, &str); 5] = [
|
||||
];
|
||||
const DEFAULT_CONTEXT: &str = "200k";
|
||||
|
||||
fn context_options(model: &ModelConfig) -> Vec<(String, String)> {
|
||||
fn context_options(context_window_tokens: Option<u64>) -> Vec<(String, String)> {
|
||||
let mut contexts = CONTEXTS
|
||||
.into_iter()
|
||||
.map(|(value, display_name)| (value.to_owned(), display_name.to_owned()))
|
||||
.collect::<Vec<_>>();
|
||||
if let Some(tokens) = model.context_window_tokens {
|
||||
if let Some(tokens) = context_window_tokens {
|
||||
let value = tokens.to_string();
|
||||
let duplicate = contexts
|
||||
.iter()
|
||||
@@ -224,15 +226,22 @@ pub async fn available_models(
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let models = registry.store().models().await?;
|
||||
let plugin_models = match registry.plugins() {
|
||||
Some(plugins) => plugins.configured_models().await,
|
||||
None => Vec::new(),
|
||||
};
|
||||
tracing::info!(
|
||||
model_count = models.len(),
|
||||
plugin_model_count = plugin_models.len(),
|
||||
"appending BYOK models to Cursor AvailableModels"
|
||||
);
|
||||
let available_models = models.iter().map(available_model).collect::<Vec<_>>();
|
||||
let mut available_models = models.iter().map(available_model).collect::<Vec<_>>();
|
||||
available_models.extend(plugin_models.iter().map(available_plugin_model));
|
||||
let local = AvailableModelsAddition {
|
||||
model_names: models
|
||||
.iter()
|
||||
.map(|model| model.model_hash.clone())
|
||||
.chain(plugin_models.iter().map(|model| model.id.clone()))
|
||||
.collect(),
|
||||
models: available_models,
|
||||
}
|
||||
@@ -252,12 +261,21 @@ pub async fn usable_models(
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let models = registry.store().models().await?;
|
||||
let plugin_models = match registry.plugins() {
|
||||
Some(plugins) => plugins.configured_models().await,
|
||||
None => Vec::new(),
|
||||
};
|
||||
tracing::info!(
|
||||
model_count = models.len(),
|
||||
plugin_model_count = plugin_models.len(),
|
||||
"appending BYOK models to Cursor GetUsableModels"
|
||||
);
|
||||
let local = UsableModelsAddition {
|
||||
models: models.iter().map(usable_model).collect(),
|
||||
models: models
|
||||
.iter()
|
||||
.map(usable_model)
|
||||
.chain(plugin_models.iter().map(usable_plugin_model))
|
||||
.collect(),
|
||||
}
|
||||
.encode_to_vec();
|
||||
match proxy::forward_buffered(&proxy, request).await {
|
||||
@@ -319,13 +337,19 @@ fn unary_payload(body: &Bytes) -> Result<(bool, &[u8])> {
|
||||
}
|
||||
|
||||
fn available_model(model: &ModelConfig) -> AvailableModel {
|
||||
let contexts = context_options(model);
|
||||
let variants = model_variants(model, &contexts);
|
||||
let contexts = context_options(model.context_window_tokens);
|
||||
let tooltip = model_tooltip(model);
|
||||
let variants = model_variants(
|
||||
&model.model_hash,
|
||||
&model.display_name,
|
||||
&tooltip,
|
||||
&contexts,
|
||||
true,
|
||||
);
|
||||
let legacy_slugs = variants
|
||||
.iter()
|
||||
.filter_map(|variant| variant.legacy_slug.clone())
|
||||
.collect();
|
||||
let tooltip = model_tooltip(model);
|
||||
AvailableModel {
|
||||
name: model.model_hash.clone(),
|
||||
default_on: true,
|
||||
@@ -344,7 +368,7 @@ fn available_model(model: &ModelConfig) -> AvailableModel {
|
||||
inputbox_short_model_name: Some(model.display_name.clone()),
|
||||
supports_sandboxing: Some(true),
|
||||
supports_cmd_k: Some(false),
|
||||
parameter_definitions: model_parameters(&contexts),
|
||||
parameter_definitions: model_parameters(&contexts, true),
|
||||
variants,
|
||||
legacy_slugs,
|
||||
named_model_section_index: Some(1),
|
||||
@@ -354,37 +378,49 @@ fn available_model(model: &ModelConfig) -> AvailableModel {
|
||||
display_name: "Cursor".into(),
|
||||
}),
|
||||
model_picker_badges: vec![ModelPickerBadge {
|
||||
label: match model.model_type {
|
||||
ModelType::OpenAi => "OpenAI".into(),
|
||||
ModelType::Anthropic => "Anthropic".into(),
|
||||
},
|
||||
label: model
|
||||
.group_name
|
||||
.clone()
|
||||
.unwrap_or_else(|| provider_host(&model.base_url)),
|
||||
variant: 1,
|
||||
dismiss_on_selection: false,
|
||||
}],
|
||||
}
|
||||
}
|
||||
|
||||
fn model_parameters(contexts: &[(String, String)]) -> Vec<ModelParameterDefinition> {
|
||||
vec![
|
||||
ModelParameterDefinition {
|
||||
id: "context".into(),
|
||||
name: "Context".into(),
|
||||
markdown_tooltip: Some("Context size used to trigger conversation compaction.".into()),
|
||||
parameter_type: Some(ModelParameterType {
|
||||
boolean_parameter: None,
|
||||
enum_parameter: Some(EnumParameter {
|
||||
values: contexts
|
||||
.iter()
|
||||
.map(|(value, display_name)| EnumParameterValue {
|
||||
value: value.clone(),
|
||||
display_name: Some(display_name.clone()),
|
||||
})
|
||||
.collect(),
|
||||
}),
|
||||
/// 徽章回退标签:base_url 的主机名。入库时已校验为带主机的 HTTP(S) URL,
|
||||
/// 解析失败仅是理论分支,此时原样返回 base_url。
|
||||
fn provider_host(base_url: &str) -> String {
|
||||
reqwest::Url::parse(base_url.trim())
|
||||
.ok()
|
||||
.and_then(|url| url.host_str().map(str::to_lowercase))
|
||||
.unwrap_or_else(|| base_url.trim().into())
|
||||
}
|
||||
|
||||
fn model_parameters(
|
||||
contexts: &[(String, String)],
|
||||
thinking: bool,
|
||||
) -> Vec<ModelParameterDefinition> {
|
||||
let mut parameters = vec![ModelParameterDefinition {
|
||||
id: "context".into(),
|
||||
name: "Context".into(),
|
||||
markdown_tooltip: Some("Context size used to trigger conversation compaction.".into()),
|
||||
parameter_type: Some(ModelParameterType {
|
||||
boolean_parameter: None,
|
||||
enum_parameter: Some(EnumParameter {
|
||||
values: contexts
|
||||
.iter()
|
||||
.map(|(value, display_name)| EnumParameterValue {
|
||||
value: value.clone(),
|
||||
display_name: Some(display_name.clone()),
|
||||
})
|
||||
.collect(),
|
||||
}),
|
||||
is_cycleable_by_hotkey: Some(false),
|
||||
},
|
||||
ModelParameterDefinition {
|
||||
}),
|
||||
is_cycleable_by_hotkey: Some(false),
|
||||
}];
|
||||
if thinking {
|
||||
parameters.push(ModelParameterDefinition {
|
||||
id: "reasoning".into(),
|
||||
name: "Effort".into(),
|
||||
markdown_tooltip: Some("Effort the model uses to generate its response.".into()),
|
||||
@@ -401,44 +437,64 @@ fn model_parameters(contexts: &[(String, String)]) -> Vec<ModelParameterDefiniti
|
||||
}),
|
||||
}),
|
||||
is_cycleable_by_hotkey: Some(true),
|
||||
},
|
||||
ModelParameterDefinition {
|
||||
id: "fast".into(),
|
||||
name: "Fast".into(),
|
||||
markdown_tooltip: Some("Significantly faster but consumes more usage".into()),
|
||||
parameter_type: Some(ModelParameterType {
|
||||
boolean_parameter: Some(BooleanParameter {
|
||||
values: vec![
|
||||
BooleanParameterValue {
|
||||
value: "false".into(),
|
||||
display_name: None,
|
||||
increases_model_cost: None,
|
||||
},
|
||||
BooleanParameterValue {
|
||||
value: "true".into(),
|
||||
display_name: Some("Fast".into()),
|
||||
increases_model_cost: Some(true),
|
||||
},
|
||||
],
|
||||
}),
|
||||
enum_parameter: None,
|
||||
});
|
||||
}
|
||||
parameters.push(ModelParameterDefinition {
|
||||
id: "fast".into(),
|
||||
name: "Fast".into(),
|
||||
markdown_tooltip: Some("Significantly faster but consumes more usage".into()),
|
||||
parameter_type: Some(ModelParameterType {
|
||||
boolean_parameter: Some(BooleanParameter {
|
||||
values: vec![
|
||||
BooleanParameterValue {
|
||||
value: "false".into(),
|
||||
display_name: None,
|
||||
increases_model_cost: None,
|
||||
},
|
||||
BooleanParameterValue {
|
||||
value: "true".into(),
|
||||
display_name: Some("Fast".into()),
|
||||
increases_model_cost: Some(true),
|
||||
},
|
||||
],
|
||||
}),
|
||||
is_cycleable_by_hotkey: Some(false),
|
||||
},
|
||||
]
|
||||
enum_parameter: None,
|
||||
}),
|
||||
is_cycleable_by_hotkey: Some(false),
|
||||
});
|
||||
parameters
|
||||
}
|
||||
|
||||
fn model_variants(model: &ModelConfig, contexts: &[(String, String)]) -> Vec<ModelVariant> {
|
||||
let mut variants = Vec::with_capacity(contexts.len() * EFFORTS.len() * 2);
|
||||
fn model_variants(
|
||||
name: &str,
|
||||
display_name: &str,
|
||||
tooltip: &TooltipData,
|
||||
contexts: &[(String, String)],
|
||||
thinking: bool,
|
||||
) -> Vec<ModelVariant> {
|
||||
// 非思考模型没有 Effort 轴,变体网格只剩 Context × Fast。
|
||||
let efforts: &[Option<(&str, &str)>] = if thinking {
|
||||
&[
|
||||
Some(EFFORTS[0]),
|
||||
Some(EFFORTS[1]),
|
||||
Some(EFFORTS[2]),
|
||||
Some(EFFORTS[3]),
|
||||
Some(EFFORTS[4]),
|
||||
]
|
||||
} else {
|
||||
&[None]
|
||||
};
|
||||
let mut variants = Vec::with_capacity(contexts.len() * efforts.len() * 2);
|
||||
for (context, context_name) in contexts {
|
||||
for (effort, effort_name) in EFFORTS {
|
||||
for effort in efforts {
|
||||
for fast in [false, true] {
|
||||
variants.push(model_variant(
|
||||
model,
|
||||
name,
|
||||
display_name,
|
||||
tooltip,
|
||||
context,
|
||||
context_name,
|
||||
effort,
|
||||
effort_name,
|
||||
*effort,
|
||||
fast,
|
||||
));
|
||||
}
|
||||
@@ -448,55 +504,67 @@ fn model_variants(model: &ModelConfig, contexts: &[(String, String)]) -> Vec<Mod
|
||||
}
|
||||
|
||||
fn model_variant(
|
||||
model: &ModelConfig,
|
||||
name: &str,
|
||||
display_name: &str,
|
||||
tooltip: &TooltipData,
|
||||
context: &str,
|
||||
context_name: &str,
|
||||
effort: &str,
|
||||
effort_name: &str,
|
||||
effort: Option<(&str, &str)>,
|
||||
fast: bool,
|
||||
) -> ModelVariant {
|
||||
let mut suffix = Vec::with_capacity(3);
|
||||
if context != DEFAULT_CONTEXT {
|
||||
suffix.push(context_name);
|
||||
}
|
||||
suffix.push(effort_name);
|
||||
if let Some((_, effort_name)) = effort {
|
||||
suffix.push(effort_name);
|
||||
}
|
||||
if fast {
|
||||
suffix.push("Fast");
|
||||
}
|
||||
let suffix = suffix.join(" ");
|
||||
let display_name = format!(
|
||||
"{} <span style=\"color: var(--cursor-text-tertiary);\">{suffix}</span>",
|
||||
model.display_name
|
||||
);
|
||||
let is_default = context == DEFAULT_CONTEXT && effort == "high" && !fast;
|
||||
let display_name = if suffix.is_empty() {
|
||||
display_name.to_owned()
|
||||
} else {
|
||||
format!(
|
||||
"{display_name} <span style=\"color: var(--cursor-text-tertiary);\">{suffix}</span>"
|
||||
)
|
||||
};
|
||||
let is_default =
|
||||
context == DEFAULT_CONTEXT && !fast && effort.is_none_or(|(effort, _)| effort == "high");
|
||||
let mut parameter_values = vec![ModelParameterValue {
|
||||
id: "context".into(),
|
||||
value: context.into(),
|
||||
}];
|
||||
if let Some((effort, _)) = effort {
|
||||
parameter_values.push(ModelParameterValue {
|
||||
id: "reasoning".into(),
|
||||
value: effort.into(),
|
||||
});
|
||||
}
|
||||
parameter_values.push(ModelParameterValue {
|
||||
id: "fast".into(),
|
||||
value: fast.to_string(),
|
||||
});
|
||||
ModelVariant {
|
||||
parameter_values: vec![
|
||||
ModelParameterValue {
|
||||
id: "context".into(),
|
||||
value: context.into(),
|
||||
},
|
||||
ModelParameterValue {
|
||||
id: "reasoning".into(),
|
||||
value: effort.into(),
|
||||
},
|
||||
ModelParameterValue {
|
||||
id: "fast".into(),
|
||||
value: fast.to_string(),
|
||||
},
|
||||
],
|
||||
parameter_values,
|
||||
display_name: display_name.clone(),
|
||||
is_max_mode: false,
|
||||
is_default_max_config: is_default.then_some(true),
|
||||
is_default_non_max_config: is_default.then_some(true),
|
||||
tooltip_data: Some(model_tooltip(model)),
|
||||
tooltip_data: Some(tooltip.clone()),
|
||||
display_name_outside_picker: Some(display_name),
|
||||
variant_string_representation: Some(format!(
|
||||
"{}[context={context},reasoning={effort},fast={fast}]",
|
||||
model.model_hash
|
||||
)),
|
||||
variant_string_representation: Some(match effort {
|
||||
Some((effort, _)) => {
|
||||
format!("{name}[context={context},reasoning={effort},fast={fast}]")
|
||||
}
|
||||
None => format!("{name}[context={context},fast={fast}]"),
|
||||
}),
|
||||
legacy_slug: Some(format!(
|
||||
"{}-{context}-{effort}{}",
|
||||
model.model_hash,
|
||||
"{name}-{context}{}{}",
|
||||
effort
|
||||
.map(|(effort, _)| format!("-{effort}"))
|
||||
.unwrap_or_default(),
|
||||
if fast { "-fast" } else { "" }
|
||||
)),
|
||||
}
|
||||
@@ -508,6 +576,63 @@ fn model_tooltip(model: &ModelConfig) -> TooltipData {
|
||||
}
|
||||
}
|
||||
|
||||
fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
|
||||
let tooltip = TooltipData {
|
||||
markdown_content: model.description.clone(),
|
||||
};
|
||||
// Effort 与上下文档位由宿主统一提供,与内置模型一致;插件不再声明这两项。
|
||||
let contexts = context_options(None);
|
||||
let variants = model_variants(&model.id, &model.display_name, &tooltip, &contexts, true);
|
||||
let legacy_slugs = variants
|
||||
.iter()
|
||||
.filter_map(|variant| variant.legacy_slug.clone())
|
||||
.collect();
|
||||
AvailableModel {
|
||||
name: model.id.clone(),
|
||||
default_on: true,
|
||||
supports_agent: Some(true),
|
||||
degradation_status: Some(0),
|
||||
tooltip_data: Some(tooltip.clone()),
|
||||
supports_thinking: Some(true),
|
||||
supports_images: Some(model.images),
|
||||
supports_max_mode: Some(false),
|
||||
client_display_name: Some(model.display_name.clone()),
|
||||
server_model_name: Some(model.id.clone()),
|
||||
supports_non_max_mode: Some(true),
|
||||
tooltip_data_for_max_mode: Some(tooltip.clone()),
|
||||
is_recommended_for_background_composer: Some(false),
|
||||
supports_plan_mode: Some(true),
|
||||
inputbox_short_model_name: Some(model.display_name.clone()),
|
||||
supports_sandboxing: Some(true),
|
||||
supports_cmd_k: Some(false),
|
||||
parameter_definitions: model_parameters(&contexts, true),
|
||||
variants,
|
||||
legacy_slugs,
|
||||
named_model_section_index: Some(1),
|
||||
vendor_name: Some(model.provider_type.clone()),
|
||||
vendor: Some(AvailableModelVendor {
|
||||
id: 6,
|
||||
display_name: model.provider_type.clone(),
|
||||
}),
|
||||
model_picker_badges: vec![ModelPickerBadge {
|
||||
label: model.plugin_name.clone(),
|
||||
variant: 1,
|
||||
dismiss_on_selection: false,
|
||||
}],
|
||||
}
|
||||
}
|
||||
|
||||
fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails {
|
||||
agent::ModelDetails {
|
||||
model_id: model.id.clone(),
|
||||
display_model_id: model.id.clone(),
|
||||
display_name: model.display_name.clone(),
|
||||
display_name_short: model.display_name.clone(),
|
||||
thinking_details: Some(agent::ThinkingDetails::default()),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
fn usable_model(model: &ModelConfig) -> agent::ModelDetails {
|
||||
agent::ModelDetails {
|
||||
model_id: model.model_hash.clone(),
|
||||
|
||||
@@ -167,7 +167,7 @@ pub fn tool_completed(call: &ToolCall, completion: &ToolCompletion) -> pb::Agent
|
||||
pub fn tool_placeholder(name: &str, call_id: &str) -> Result<pb::ToolCall> {
|
||||
use pb::tool_call::Tool;
|
||||
let tool = match normalized(name).as_str() {
|
||||
"shell" => Tool::ShellToolCall(pb::ShellToolCall::default()),
|
||||
"shell" | "bash" => Tool::ShellToolCall(pb::ShellToolCall::default()),
|
||||
"delete" => Tool::DeleteToolCall(pb::DeleteToolCall::default()),
|
||||
"glob" => Tool::GlobToolCall(pb::GlobToolCall::default()),
|
||||
"grep" => Tool::GrepToolCall(pb::GrepToolCall::default()),
|
||||
@@ -527,3 +527,23 @@ fn now_ms() -> u64 {
|
||||
.unwrap_or_default()
|
||||
.as_millis() as u64
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::tool_placeholder;
|
||||
use crate::cursor::protocol::proto::agent::v1 as pb;
|
||||
|
||||
#[test]
|
||||
fn bash_renders_as_a_shell_placeholder() {
|
||||
// The dispatcher treats `bash`/`Bash` as a Shell alias, so the streaming
|
||||
// placeholder must too; otherwise a `Bash` tool call aborts the turn with
|
||||
// `unsupported tool: bash` before it ever runs.
|
||||
for name in ["shell", "Shell", "bash", "Bash"] {
|
||||
let tool = tool_placeholder(name, "call-1").unwrap().tool;
|
||||
assert!(
|
||||
matches!(tool, Some(pb::tool_call::Tool::ShellToolCall(_))),
|
||||
"{name} should render as a Shell tool"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -35,7 +35,7 @@ pub fn request(id: u32, call: &ToolCall, context: &ExecContext) -> Result<pb::Ag
|
||||
.map(|v| v as i32)
|
||||
};
|
||||
let message = match normalize(&call.name).as_str() {
|
||||
"shell" => {
|
||||
"shell" | "bash" => {
|
||||
let command = string("command")?;
|
||||
let (simple_commands, parsing_result) = shell_command_metadata(&command);
|
||||
Message::ShellStreamArgs(pb::ShellArgs {
|
||||
@@ -520,3 +520,37 @@ fn prost_value(value: &Value) -> prost_types::Value {
|
||||
};
|
||||
ProstValue { kind: Some(kind) }
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::request;
|
||||
use crate::cursor::protocol::proto::agent::v1 as pb;
|
||||
use crate::cursor::tools::runtime::ExecContext;
|
||||
use crate::model::ToolCall;
|
||||
|
||||
#[test]
|
||||
fn bash_is_encoded_as_a_shell_exec_request() {
|
||||
// The dispatcher routes `bash`/`Bash` to the shell executor, so the
|
||||
// request codec must encode it as a Shell stream instead of erroring
|
||||
// with `tool bash is not executed through ExecServerMessage`.
|
||||
let call = ToolCall {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
model_call_id: "model-1".into(),
|
||||
name: "Bash".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({ "command": "ls -la" }),
|
||||
};
|
||||
let message = request(1, &call, &ExecContext::default()).unwrap();
|
||||
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message
|
||||
else {
|
||||
panic!("expected an ExecServerMessage");
|
||||
};
|
||||
let Some(pb::exec_server_message::Message::ShellStreamArgs(args)) = exec.message else {
|
||||
panic!("expected ShellStreamArgs");
|
||||
};
|
||||
assert_eq!(args.command, "ls -la");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -144,7 +144,8 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
|
||||
return Ok(None);
|
||||
};
|
||||
let error = "Cursor Exec stream closed before returning a terminal result";
|
||||
if entry.call.name.eq_ignore_ascii_case("Shell") {
|
||||
if entry.call.name.eq_ignore_ascii_case("Shell") || entry.call.name.eq_ignore_ascii_case("Bash")
|
||||
{
|
||||
let command = entry
|
||||
.call
|
||||
.arguments
|
||||
|
||||
@@ -183,6 +183,9 @@ fn edit_notebook(call: &ToolCall, before: &str) -> std::result::Result<String, S
|
||||
.unwrap_or_default();
|
||||
let old =
|
||||
normalize_newlines(&string(call, "old_string").map_err(|error| error.to_string())?);
|
||||
if old.is_empty() {
|
||||
return Err("old_string must not be empty".into());
|
||||
}
|
||||
let occurrences = source.match_indices(&old).count();
|
||||
let edited = match occurrences {
|
||||
0 => return Err("old_string was not found in the notebook cell".into()),
|
||||
@@ -241,3 +244,49 @@ fn normalized(value: &str) -> String {
|
||||
.flat_map(char::to_lowercase)
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::edit_notebook;
|
||||
use crate::model::ToolCall;
|
||||
|
||||
fn notebook_call(old_string: &str) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: "call".into(),
|
||||
model_call_id: "model".into(),
|
||||
name: "EditNotebook".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({
|
||||
"target_notebook": "/notebook.ipynb",
|
||||
"cell_idx": 0,
|
||||
"old_string": old_string,
|
||||
"new_string": "replacement",
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn single_cell_notebook() -> String {
|
||||
json!({
|
||||
"cells": [{"cell_type": "code", "source": ["print('hi')\n"]}],
|
||||
})
|
||||
.to_string()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn edit_notebook_rejects_empty_old_string() {
|
||||
// StrReplace rejects an empty old_string; EditNotebook must do the same
|
||||
// instead of prepending new_string (empty cell) or reporting a
|
||||
// misleading "not unique" error (non-empty cell).
|
||||
let error = edit_notebook(¬ebook_call(""), &single_cell_notebook()).unwrap_err();
|
||||
assert_eq!(error, "old_string must not be empty");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn edit_notebook_replaces_a_unique_old_string() {
|
||||
let edited = edit_notebook(¬ebook_call("hi"), &single_cell_notebook()).unwrap();
|
||||
assert!(edited.contains("print('replacement')"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ use crate::{
|
||||
conversation::ConversationRegistry, prompting::PromptCompiler,
|
||||
services::observability::CursorTraceRecorder,
|
||||
},
|
||||
plugin::PluginRegistry,
|
||||
provider::Provider,
|
||||
search::WebCache,
|
||||
store::Store,
|
||||
@@ -28,6 +29,7 @@ struct RegistryInner {
|
||||
route_changed: Notify,
|
||||
store: Store,
|
||||
web_cache: WebCache,
|
||||
plugins: Option<PluginRegistry>,
|
||||
conversations: ConversationRegistry,
|
||||
}
|
||||
|
||||
@@ -47,6 +49,52 @@ impl TransportRegistry {
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
) -> Self {
|
||||
Self::build(store, provider, compiler, web_cache, None, None)
|
||||
}
|
||||
|
||||
/// 附带本地 rules 目录的构造;编译请求上下文时会合并该目录下的 md 规则。
|
||||
pub fn with_local_rules(
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
local_rules_dir: std::path::PathBuf,
|
||||
) -> Self {
|
||||
Self::build(
|
||||
store,
|
||||
provider,
|
||||
compiler,
|
||||
WebCache::default(),
|
||||
None,
|
||||
Some(local_rules_dir),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn with_plugins(
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
plugins: PluginRegistry,
|
||||
local_rules_dir: std::path::PathBuf,
|
||||
) -> Self {
|
||||
Self::build(
|
||||
store,
|
||||
provider,
|
||||
compiler,
|
||||
web_cache,
|
||||
Some(plugins),
|
||||
Some(local_rules_dir),
|
||||
)
|
||||
}
|
||||
|
||||
fn build(
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
plugins: Option<PluginRegistry>,
|
||||
local_rules_dir: Option<std::path::PathBuf>,
|
||||
) -> Self {
|
||||
Self {
|
||||
inner: Arc::new(RegistryInner {
|
||||
@@ -58,9 +106,11 @@ impl TransportRegistry {
|
||||
provider,
|
||||
compiler,
|
||||
web_cache.clone(),
|
||||
local_rules_dir,
|
||||
),
|
||||
store,
|
||||
web_cache,
|
||||
plugins,
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -73,6 +123,10 @@ impl TransportRegistry {
|
||||
&self.inner.web_cache
|
||||
}
|
||||
|
||||
pub fn plugins(&self) -> Option<&PluginRegistry> {
|
||||
self.inner.plugins.as_ref()
|
||||
}
|
||||
|
||||
pub fn conversations(&self) -> &ConversationRegistry {
|
||||
&self.inner.conversations
|
||||
}
|
||||
|
||||
@@ -57,6 +57,8 @@ impl IntoResponse for Error {
|
||||
| Self::Encode(_)
|
||||
| Self::Io(_) => StatusCode::INTERNAL_SERVER_ERROR,
|
||||
};
|
||||
// 所有回给 UI 的错误统一落日志,否则失败原因只出现在前端提示里。
|
||||
tracing::warn!(%status, error = %self, "request failed");
|
||||
let code = match status {
|
||||
StatusCode::BAD_REQUEST => "invalid_argument",
|
||||
StatusCode::NOT_FOUND => "not_found",
|
||||
|
||||
@@ -8,6 +8,7 @@ pub mod error;
|
||||
pub mod local_app;
|
||||
pub mod model;
|
||||
pub mod network;
|
||||
pub mod plugin;
|
||||
pub mod provider;
|
||||
pub mod run;
|
||||
pub mod search;
|
||||
|
||||
@@ -161,6 +161,10 @@ fn is_local_path(path: &str) -> bool {
|
||||
| "/aiserver.v1.DashboardService/GetUserProfile"
|
||||
| "/aiserver.v1.DashboardService/GetCurrentPeriodUsage"
|
||||
| "/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseAdd"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseList"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseUpdate"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseRemove"
|
||||
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
|
||||
| "/auth/full_stripe_profile"
|
||||
)
|
||||
|
||||
@@ -18,6 +18,9 @@ pub enum ProviderType {
|
||||
OpenAiResponses,
|
||||
#[serde(rename = "anthropic")]
|
||||
Anthropic,
|
||||
/// 插件执行的调用;协议细节在插件内部,核心只按统一事件流记录。
|
||||
#[serde(rename = "plugin")]
|
||||
Plugin,
|
||||
}
|
||||
|
||||
impl ProviderType {
|
||||
@@ -26,6 +29,7 @@ impl ProviderType {
|
||||
Self::OpenAiChat => "openai-chat",
|
||||
Self::OpenAiResponses => "openai-responses",
|
||||
Self::Anthropic => "anthropic",
|
||||
Self::Plugin => "plugin",
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -44,6 +48,7 @@ impl FromStr for ProviderType {
|
||||
"openai-chat" => Ok(Self::OpenAiChat),
|
||||
"openai-responses" => Ok(Self::OpenAiResponses),
|
||||
"anthropic" => Ok(Self::Anthropic),
|
||||
"plugin" => Ok(Self::Plugin),
|
||||
_ => Err(Error::Config(format!("unsupported provider type: {value}"))),
|
||||
}
|
||||
}
|
||||
@@ -82,6 +87,9 @@ pub struct ModelConfigInput {
|
||||
#[serde(default)]
|
||||
pub sort_order: i64,
|
||||
pub display_name: String,
|
||||
/// 供应商分组的自定义显示名;同一 base_url 主机下的模型共享。
|
||||
#[serde(default)]
|
||||
pub group_name: Option<String>,
|
||||
#[serde(rename = "type")]
|
||||
pub model_type: ModelType,
|
||||
pub base_url: String,
|
||||
@@ -119,6 +127,7 @@ pub struct ModelConfig {
|
||||
pub model_hash: String,
|
||||
pub sort_order: i64,
|
||||
pub display_name: String,
|
||||
pub group_name: Option<String>,
|
||||
#[serde(rename = "type")]
|
||||
pub model_type: ModelType,
|
||||
pub base_url: String,
|
||||
@@ -199,6 +208,12 @@ impl ModelConfig {
|
||||
|
||||
pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInput> {
|
||||
let display_name = required(&input.display_name, "model display name")?;
|
||||
let group_name = input
|
||||
.group_name
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(String::from);
|
||||
let base_url = normalize_request_url(&input.base_url)?;
|
||||
let api_key = required(&input.api_key, "model API key")?;
|
||||
let tooltip_data = required(&input.tooltip_data, "model tooltip")?;
|
||||
@@ -225,6 +240,7 @@ pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInpu
|
||||
let normalized = ModelConfigInput {
|
||||
sort_order: input.sort_order.max(0),
|
||||
display_name,
|
||||
group_name,
|
||||
model_type: input.model_type,
|
||||
base_url,
|
||||
use_full_url: input.use_full_url,
|
||||
|
||||
@@ -23,7 +23,9 @@ mod usage {
|
||||
pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option<u64> {
|
||||
let input = self.input_tokens?;
|
||||
match provider {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => Some(input),
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses | ProviderType::Plugin => {
|
||||
Some(input)
|
||||
}
|
||||
ProviderType::Anthropic => input
|
||||
.checked_add(self.cache_read_tokens.unwrap_or_default())?
|
||||
.checked_add(self.cache_write_tokens.unwrap_or_default()),
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
//! Maps supported desktop platforms to pinned Deno release assets.
|
||||
pub(super) const DENO_VERSION: &str = "2.9.6";
|
||||
|
||||
#[derive(Clone, Copy, Debug)]
|
||||
pub(super) struct RuntimeAsset {
|
||||
pub target: &'static str,
|
||||
pub sha256: &'static str,
|
||||
}
|
||||
|
||||
impl RuntimeAsset {
|
||||
pub fn current() -> Option<Self> {
|
||||
Self::for_platform(std::env::consts::OS, std::env::consts::ARCH)
|
||||
}
|
||||
|
||||
pub(super) fn for_platform(os: &str, arch: &str) -> Option<Self> {
|
||||
let (target, sha256) = match (os, arch) {
|
||||
("macos", "aarch64") => (
|
||||
"aarch64-apple-darwin",
|
||||
"213a2f304f04d3c9cb5220669afad138f60a5aab1fe80962abdeb8f35807a472",
|
||||
),
|
||||
("macos", "x86_64") => (
|
||||
"x86_64-apple-darwin",
|
||||
"7d4524b82bcc557fe020a1a5b56956ed42b992ae5b28026e8ad5d17329533f5f",
|
||||
),
|
||||
("windows", "aarch64") => (
|
||||
"aarch64-pc-windows-msvc",
|
||||
"acb014afe2299847764e232b4993e162e3946cdeec36603e3f1a0b548cd1ea55",
|
||||
),
|
||||
("windows", "x86_64") => (
|
||||
"x86_64-pc-windows-msvc",
|
||||
"15e5300b0ba3c3695a7621d90160a746ec9e710228cee639afa9d580f6e3cd11",
|
||||
),
|
||||
("linux", "aarch64") => (
|
||||
"aarch64-unknown-linux-gnu",
|
||||
"9a46afc6c392c7cd2ff71a31558935545b46408d0e87f7a86908c712721c046e",
|
||||
),
|
||||
("linux", "x86_64") => (
|
||||
"x86_64-unknown-linux-gnu",
|
||||
"394f07f4da2bebe6ce6f1e7ce0fa16429b29b08c35e3fac3fe25972676dff4b2",
|
||||
),
|
||||
_ => return None,
|
||||
};
|
||||
Some(Self { target, sha256 })
|
||||
}
|
||||
|
||||
pub fn archive_name(self) -> String {
|
||||
format!("deno-{}.zip", self.target)
|
||||
}
|
||||
|
||||
pub fn download_url(self) -> String {
|
||||
format!(
|
||||
"https://github.com/denoland/deno/releases/download/v{DENO_VERSION}/{}",
|
||||
self.archive_name()
|
||||
)
|
||||
}
|
||||
|
||||
pub fn executable_name(self) -> &'static str {
|
||||
if self.target.contains("windows") {
|
||||
"deno.exe"
|
||||
} else {
|
||||
"deno"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn maps_every_supported_desktop_target() {
|
||||
let cases = [
|
||||
("macos", "aarch64", "aarch64-apple-darwin"),
|
||||
("macos", "x86_64", "x86_64-apple-darwin"),
|
||||
("windows", "aarch64", "aarch64-pc-windows-msvc"),
|
||||
("windows", "x86_64", "x86_64-pc-windows-msvc"),
|
||||
("linux", "aarch64", "aarch64-unknown-linux-gnu"),
|
||||
("linux", "x86_64", "x86_64-unknown-linux-gnu"),
|
||||
];
|
||||
for (os, arch, expected) in cases {
|
||||
assert_eq!(
|
||||
RuntimeAsset::for_platform(os, arch).unwrap().target,
|
||||
expected
|
||||
);
|
||||
}
|
||||
assert!(RuntimeAsset::for_platform("linux", "x86").is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
//! Pre-installs bundled built-in plugins into the user's installed directory.
|
||||
use std::path::Path;
|
||||
|
||||
use super::definition::write_if_changed;
|
||||
use crate::Result;
|
||||
|
||||
/// 随二进制打包的内置插件文件;发布构建没有源码目录,靠这里预装。
|
||||
const CODEX_AUTH: &[(&str, &str)] = &[
|
||||
(
|
||||
"plugin.json",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/codex-auth/plugin.json"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"main.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/codex-auth/main.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"provider.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/codex-auth/provider.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"models.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/codex-auth/models.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"oauth.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/codex-auth/oauth.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"resources.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/codex-auth/resources.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"assets/codex.svg",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/codex-auth/assets/codex.svg"
|
||||
)),
|
||||
),
|
||||
];
|
||||
|
||||
const GROK_AUTH: &[(&str, &str)] = &[
|
||||
(
|
||||
"plugin.json",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/grok-auth/plugin.json"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"main.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/grok-auth/main.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"provider.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/grok-auth/provider.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"models.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/grok-auth/models.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"oauth.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/grok-auth/oauth.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"resources.ts",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/grok-auth/resources.ts"
|
||||
)),
|
||||
),
|
||||
(
|
||||
"assets/grok.svg",
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/plugins/build-in/grok-auth/assets/grok.svg"
|
||||
)),
|
||||
),
|
||||
];
|
||||
|
||||
const PLUGINS: &[(&str, &[(&str, &str)])] = &[("codex-auth", CODEX_AUTH), ("grok-auth", GROK_AUTH)];
|
||||
|
||||
/// 把内置插件预装到 installed 目录。manifest 的 version 是缓存键:
|
||||
/// 版本一致时零写盘;版本变化时整目录同步并清理旧版本残留文件。
|
||||
pub(super) fn install(installed: &Path) -> Result<()> {
|
||||
for (name, files) in PLUGINS {
|
||||
let directory = installed.join(name);
|
||||
if disk_version(&directory) == Some(embedded_version(files)?) {
|
||||
continue;
|
||||
}
|
||||
write_plugin(&directory, files)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn embedded_version(files: &[(&str, &str)]) -> Result<String> {
|
||||
let manifest = files
|
||||
.iter()
|
||||
.find(|(name, _)| *name == "plugin.json")
|
||||
.map(|(_, content)| *content)
|
||||
.expect("built-in plugin bundles plugin.json");
|
||||
let value: serde_json::Value = serde_json::from_str(manifest)?;
|
||||
value
|
||||
.get("version")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_owned)
|
||||
.ok_or_else(|| crate::Error::Config("built-in plugin manifest requires version".into()))
|
||||
}
|
||||
|
||||
fn disk_version(directory: &Path) -> Option<String> {
|
||||
let manifest = std::fs::read_to_string(directory.join("plugin.json")).ok()?;
|
||||
let value: serde_json::Value = serde_json::from_str(&manifest).ok()?;
|
||||
Some(value.get("version")?.as_str()?.to_owned())
|
||||
}
|
||||
|
||||
fn write_plugin(directory: &Path, files: &[(&str, &str)]) -> Result<()> {
|
||||
for (relative, content) in files {
|
||||
let path = directory.join(relative);
|
||||
let parent = path.parent().expect("plugin file path has a parent");
|
||||
std::fs::create_dir_all(parent)?;
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(parent, std::fs::Permissions::from_mode(0o700))?;
|
||||
}
|
||||
write_if_changed(&path, content)?;
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600))?;
|
||||
}
|
||||
}
|
||||
prune_unknown_files(directory, directory, files)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 删除插件目录中不在嵌入清单里的文件与空目录(旧版本残留)。
|
||||
fn prune_unknown_files(root: &Path, directory: &Path, files: &[(&str, &str)]) -> Result<()> {
|
||||
for entry in std::fs::read_dir(directory)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
if entry.file_type()?.is_dir() {
|
||||
prune_unknown_files(root, &path, files)?;
|
||||
if std::fs::read_dir(&path)?.next().is_none() {
|
||||
std::fs::remove_dir(&path)?;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
let known = files
|
||||
.iter()
|
||||
.any(|(relative, _)| root.join(relative) == path);
|
||||
if !known {
|
||||
std::fs::remove_file(&path)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn embedded_main() -> &'static str {
|
||||
CODEX_AUTH
|
||||
.iter()
|
||||
.find(|(name, _)| *name == "main.ts")
|
||||
.unwrap()
|
||||
.1
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn install_is_version_gated_and_syncs_on_version_change() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let plugin = root.path().join("codex-auth");
|
||||
|
||||
install(root.path()).unwrap();
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(plugin.join("main.ts")).unwrap(),
|
||||
embedded_main()
|
||||
);
|
||||
|
||||
// 版本一致:本地改动与额外文件保持原样,不发生任何写盘。
|
||||
std::fs::write(plugin.join("main.ts"), "edited").unwrap();
|
||||
std::fs::write(plugin.join("stale.ts"), "extra").unwrap();
|
||||
install(root.path()).unwrap();
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(plugin.join("main.ts")).unwrap(),
|
||||
"edited"
|
||||
);
|
||||
assert!(plugin.join("stale.ts").exists());
|
||||
|
||||
// 版本变化:整目录同步回嵌入内容并清理残留。
|
||||
let manifest = std::fs::read_to_string(plugin.join("plugin.json")).unwrap();
|
||||
let mut value: serde_json::Value = serde_json::from_str(&manifest).unwrap();
|
||||
value["version"] = serde_json::Value::String("0.0.1".into());
|
||||
std::fs::write(plugin.join("plugin.json"), value.to_string()).unwrap();
|
||||
install(root.path()).unwrap();
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(plugin.join("main.ts")).unwrap(),
|
||||
embedded_main()
|
||||
);
|
||||
assert!(!plugin.join("stale.ts").exists());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
//! Discovers plugin manifests and evaluates serializable TypeScript definitions.
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
fs,
|
||||
path::{Path, PathBuf},
|
||||
};
|
||||
|
||||
use base64::{engine::general_purpose::STANDARD, Engine};
|
||||
|
||||
use super::{
|
||||
definition::PluginDefinitionLoader,
|
||||
descriptor::PluginModuleDefinition,
|
||||
manifest::{validate_id, PluginManifest},
|
||||
};
|
||||
use crate::{config, Error, Result};
|
||||
|
||||
const MANIFEST_FILE_NAME: &str = "plugin.json";
|
||||
const MAX_ICON_BYTES: u64 = 1024 * 1024;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PluginCatalog {
|
||||
roots: Vec<PathBuf>,
|
||||
definition_loader: PluginDefinitionLoader,
|
||||
app_version: String,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct PluginEntry {
|
||||
pub directory: PathBuf,
|
||||
pub entry: PathBuf,
|
||||
pub manifest: PluginManifest,
|
||||
pub definition: PluginModuleDefinition,
|
||||
pub icon: String,
|
||||
}
|
||||
|
||||
impl PluginCatalog {
|
||||
pub fn managed(app_version: String) -> Result<Self> {
|
||||
let installed = config::managed_data_dir()?.join("plugins/installed");
|
||||
fs::create_dir_all(&installed)?;
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
fs::set_permissions(&installed, fs::Permissions::from_mode(0o700))?;
|
||||
}
|
||||
// 内置插件按版本预装进 installed;版本一致时不写盘。
|
||||
super::builtin::install(&installed)?;
|
||||
// 扫描顺序即优先级:debug 下源码目录优先,保证内置插件热改生效;
|
||||
// 发布构建只有 installed 一个根。
|
||||
#[cfg(debug_assertions)]
|
||||
let roots = vec![
|
||||
PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("plugins/build-in"),
|
||||
installed,
|
||||
];
|
||||
#[cfg(not(debug_assertions))]
|
||||
let roots = vec![installed];
|
||||
Ok(Self {
|
||||
roots,
|
||||
definition_loader: PluginDefinitionLoader::managed()?,
|
||||
app_version,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn loader(&self) -> &PluginDefinitionLoader {
|
||||
&self.definition_loader
|
||||
}
|
||||
|
||||
pub(crate) async fn entries(&self, executable: &Path) -> Vec<PluginEntry> {
|
||||
let mut plugins = BTreeMap::new();
|
||||
for root in &self.roots {
|
||||
let mut directories = match child_directories(root) {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
tracing::warn!(path = %root.display(), %error, "failed to scan plugin directory");
|
||||
continue;
|
||||
}
|
||||
};
|
||||
directories.sort();
|
||||
for directory in directories {
|
||||
match load_plugin(
|
||||
&directory,
|
||||
&self.definition_loader,
|
||||
executable,
|
||||
&self.app_version,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(entry) => {
|
||||
if plugins.contains_key(&entry.manifest.id) {
|
||||
tracing::warn!(plugin = %entry.manifest.id, path = %directory.display(), "ignoring duplicate plugin");
|
||||
} else {
|
||||
plugins.insert(entry.manifest.id.clone(), entry);
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(path = %directory.display(), %error, "ignoring invalid plugin")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
plugins.into_values().collect()
|
||||
}
|
||||
|
||||
pub(crate) fn manifests(&self) -> Vec<(PluginManifest, String)> {
|
||||
let mut plugins = BTreeMap::new();
|
||||
for root in &self.roots {
|
||||
let Ok(mut directories) = child_directories(root) else {
|
||||
continue;
|
||||
};
|
||||
directories.sort();
|
||||
for directory in directories {
|
||||
let loaded = (|| -> Result<_> {
|
||||
let manifest: PluginManifest =
|
||||
serde_json::from_slice(&fs::read(directory.join(MANIFEST_FILE_NAME))?)?;
|
||||
manifest.validate(&directory)?;
|
||||
require_app_version(&manifest, &self.app_version)?;
|
||||
let icon = icon_data_url(&directory, &manifest.icon)?;
|
||||
Ok((manifest, icon))
|
||||
})();
|
||||
if let Ok((manifest, icon)) = loaded {
|
||||
plugins
|
||||
.entry(manifest.id.clone())
|
||||
.or_insert((manifest, icon));
|
||||
}
|
||||
}
|
||||
}
|
||||
plugins.into_values().collect()
|
||||
}
|
||||
}
|
||||
|
||||
fn child_directories(root: &Path) -> Result<Vec<PathBuf>> {
|
||||
if !root.exists() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut directories = Vec::new();
|
||||
for entry in fs::read_dir(root)? {
|
||||
let entry = entry?;
|
||||
if entry.file_type()?.is_dir() && !entry.file_name().to_string_lossy().starts_with('.') {
|
||||
directories.push(entry.path());
|
||||
}
|
||||
}
|
||||
Ok(directories)
|
||||
}
|
||||
|
||||
/// 应用过旧时拒绝加载,让插件的 minAppVersion 声明生效。
|
||||
fn require_app_version(manifest: &PluginManifest, app_version: &str) -> Result<()> {
|
||||
let Some(minimum) = &manifest.min_app_version else {
|
||||
return Ok(());
|
||||
};
|
||||
if super::manifest::version_at_least(app_version, minimum) {
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::Config(format!(
|
||||
"plugin '{}' requires app version {minimum} or newer (current {app_version})",
|
||||
manifest.id
|
||||
)))
|
||||
}
|
||||
|
||||
async fn load_plugin(
|
||||
directory: &Path,
|
||||
loader: &PluginDefinitionLoader,
|
||||
executable: &Path,
|
||||
app_version: &str,
|
||||
) -> Result<PluginEntry> {
|
||||
let manifest: PluginManifest =
|
||||
serde_json::from_slice(&fs::read(directory.join(MANIFEST_FILE_NAME))?)?;
|
||||
manifest.validate(directory)?;
|
||||
require_app_version(&manifest, app_version)?;
|
||||
let icon = icon_data_url(directory, &manifest.icon)?;
|
||||
let entry = directory.join(&manifest.entry).canonicalize()?;
|
||||
let definition = loader.load(executable, directory, &entry).await?;
|
||||
validate_definition(&manifest.id, &definition)?;
|
||||
Ok(PluginEntry {
|
||||
directory: directory.to_path_buf(),
|
||||
entry,
|
||||
manifest,
|
||||
definition,
|
||||
icon,
|
||||
})
|
||||
}
|
||||
|
||||
/// 显示文本必须是非空字符串,或全为非空字符串的 locale 映射。
|
||||
fn validate_localized_text(value: &serde_json::Value, label: &str) -> Result<()> {
|
||||
match value {
|
||||
serde_json::Value::String(text) if !text.trim().is_empty() => Ok(()),
|
||||
serde_json::Value::Object(map)
|
||||
if !map.is_empty()
|
||||
&& map
|
||||
.values()
|
||||
.all(|entry| entry.as_str().is_some_and(|text| !text.trim().is_empty())) =>
|
||||
{
|
||||
Ok(())
|
||||
}
|
||||
_ => Err(Error::Config(format!(
|
||||
"{label} must be a non-empty string or a locale map of non-empty strings"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_definition(plugin_id: &str, definition: &PluginModuleDefinition) -> Result<()> {
|
||||
if definition.providers.is_empty() {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' must define at least one provider"
|
||||
)));
|
||||
}
|
||||
let mut provider_ids = std::collections::HashSet::new();
|
||||
for provider in &definition.providers {
|
||||
validate_id(&provider.id, "plugin provider id")?;
|
||||
validate_localized_text(
|
||||
&provider.display_name,
|
||||
&format!(
|
||||
"plugin '{plugin_id}' provider '{}' displayName",
|
||||
provider.id
|
||||
),
|
||||
)?;
|
||||
if provider.provider_type.trim().is_empty() {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' provider '{}' requires providerType",
|
||||
provider.id
|
||||
)));
|
||||
}
|
||||
if !provider_ids.insert(provider.id.clone()) {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' contains duplicate provider '{}'",
|
||||
provider.id
|
||||
)));
|
||||
}
|
||||
if let Some(resource_type) = &provider.resource_type {
|
||||
if !definition
|
||||
.resources
|
||||
.iter()
|
||||
.any(|resource| &resource.resource_type == resource_type)
|
||||
{
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' provider '{}' consumes undeclared resource '{resource_type}'",
|
||||
provider.id
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
let mut resource_types = std::collections::HashSet::new();
|
||||
for resource in &definition.resources {
|
||||
validate_id(&resource.resource_type, "plugin resource type")?;
|
||||
validate_localized_text(
|
||||
&resource.display_name,
|
||||
&format!(
|
||||
"plugin '{plugin_id}' resource '{}' displayName",
|
||||
resource.resource_type
|
||||
),
|
||||
)?;
|
||||
if !resource_types.insert(resource.resource_type.clone()) {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' contains duplicate resource type '{}'",
|
||||
resource.resource_type
|
||||
)));
|
||||
}
|
||||
for method in &resource.add {
|
||||
validate_id(&method.id, "plugin add method id")?;
|
||||
if method.method_type != super::descriptor::OAUTH2_ADD_METHOD {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' add method '{}' uses unsupported type '{}'",
|
||||
method.id, method.method_type
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn icon_data_url(directory: &Path, relative: &str) -> Result<String> {
|
||||
let root = directory.canonicalize()?;
|
||||
let path = directory.join(relative).canonicalize()?;
|
||||
if !path.starts_with(&root) {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin icon escapes its directory: {relative}"
|
||||
)));
|
||||
}
|
||||
if fs::metadata(&path)?.len() > MAX_ICON_BYTES {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin icon exceeds {MAX_ICON_BYTES} bytes: {relative}"
|
||||
)));
|
||||
}
|
||||
let extension = path
|
||||
.extension()
|
||||
.and_then(|value| value.to_str())
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase();
|
||||
let mime = match extension.as_str() {
|
||||
"svg" => "image/svg+xml",
|
||||
"png" => "image/png",
|
||||
"webp" => "image/webp",
|
||||
_ => {
|
||||
return Err(Error::Config(format!(
|
||||
"unsupported plugin icon: {relative}"
|
||||
)))
|
||||
}
|
||||
};
|
||||
Ok(format!(
|
||||
"data:{mime};base64,{}",
|
||||
STANDARD.encode(fs::read(path)?)
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[test]
|
||||
fn repository_examples_have_valid_static_manifests() {
|
||||
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("plugins/build-in");
|
||||
let sdk = tempfile::tempdir().unwrap();
|
||||
let catalog = PluginCatalog {
|
||||
roots: vec![root],
|
||||
definition_loader: PluginDefinitionLoader::for_test(sdk.path()).unwrap(),
|
||||
app_version: env!("CARGO_PKG_VERSION").into(),
|
||||
};
|
||||
assert!(!catalog.manifests().is_empty());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,246 @@
|
||||
//! Stores plugin-owned JSON with private permissions and atomic replacement.
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
path::{Path, PathBuf},
|
||||
sync::Arc,
|
||||
};
|
||||
|
||||
use parking_lot::Mutex;
|
||||
use tokio::sync::Mutex as AsyncMutex;
|
||||
|
||||
use crate::{config, Error, Result};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PluginDataStore {
|
||||
root: PathBuf,
|
||||
locks: Arc<Mutex<HashMap<String, Arc<AsyncMutex<()>>>>>,
|
||||
}
|
||||
|
||||
impl PluginDataStore {
|
||||
pub fn managed() -> Result<Self> {
|
||||
Self::new(config::managed_data_dir()?.join("plugins/data"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(super) fn for_test(root: PathBuf) -> Result<Self> {
|
||||
Self::new(root)
|
||||
}
|
||||
|
||||
fn new(root: PathBuf) -> Result<Self> {
|
||||
std::fs::create_dir_all(&root)?;
|
||||
set_directory_permissions(&root)?;
|
||||
Ok(Self {
|
||||
root,
|
||||
locks: Arc::new(Mutex::new(HashMap::new())),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn read(&self, plugin_id: &str, key: &str) -> Result<serde_json::Value> {
|
||||
let path = self.path(plugin_id, key)?;
|
||||
let lock = self.lock(plugin_id);
|
||||
let _guard = lock.lock().await;
|
||||
match tokio::fs::read(&path).await {
|
||||
Ok(bytes) => Ok(serde_json::from_slice(&bytes)?),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
|
||||
Ok(serde_json::Value::Null)
|
||||
}
|
||||
Err(error) => Err(Error::Config(format!(
|
||||
"plugin data read failed at {}: {error}",
|
||||
path.display()
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn update(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
key: &str,
|
||||
value: &serde_json::Value,
|
||||
) -> Result<()> {
|
||||
let path = self.path(plugin_id, key)?;
|
||||
let lock = self.lock(plugin_id);
|
||||
let _guard = lock.lock().await;
|
||||
self.write_locked(&path, key, value)
|
||||
.await
|
||||
// 带上具体路径,Windows 上的拒绝访问才能定位到是哪一步。
|
||||
.map_err(|error| {
|
||||
Error::Config(format!(
|
||||
"plugin data write failed at {}: {error}",
|
||||
path.display()
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
/// 全程使用同步 IO 在阻塞线程完成:tokio 异步文件的关闭是延迟的,
|
||||
/// 替换前句柄可能仍被本进程持有;同步写入保证替换时句柄已确定关闭。
|
||||
async fn write_locked(&self, path: &Path, key: &str, value: &serde_json::Value) -> Result<()> {
|
||||
let directory = path
|
||||
.parent()
|
||||
.expect("plugin data path has a parent")
|
||||
.to_owned();
|
||||
let temporary = directory.join(format!(".{key}.{}.tmp", uuid::Uuid::new_v4()));
|
||||
let target = path.to_owned();
|
||||
let bytes = serde_json::to_vec_pretty(value)?;
|
||||
tokio::task::spawn_blocking(move || {
|
||||
// Windows 上杀软或索引器会短暂锁住新建文件,任何一步都可能
|
||||
// 拒绝访问,因此把整个序列作为一个整体重试。
|
||||
let mut attempts = 0;
|
||||
loop {
|
||||
match write_once(&directory, &temporary, &target, &bytes) {
|
||||
Ok(()) => return Ok(()),
|
||||
Err((step, error)) if attempts < 20 && transient(&error) => {
|
||||
attempts += 1;
|
||||
tracing::debug!(step, attempts, %error, "retrying plugin data write");
|
||||
std::thread::sleep(std::time::Duration::from_millis(100));
|
||||
}
|
||||
Err((step, error)) => {
|
||||
let _ = std::fs::remove_file(&temporary);
|
||||
tracing::warn!(
|
||||
path = %target.display(),
|
||||
step,
|
||||
attempts,
|
||||
%error,
|
||||
"plugin data write failed"
|
||||
);
|
||||
return Err(Error::Config(format!("{step}: {error}")));
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("plugin data write task panicked")
|
||||
}
|
||||
|
||||
pub async fn clear(&self, plugin_id: &str) -> Result<()> {
|
||||
validate_component(plugin_id, "plugin id")?;
|
||||
let lock = self.lock(plugin_id);
|
||||
let _guard = lock.lock().await;
|
||||
let path = self.root.join(plugin_id);
|
||||
match tokio::fs::remove_dir_all(&path).await {
|
||||
Ok(()) => Ok(()),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
|
||||
Err(error) => Err(Error::Config(format!(
|
||||
"plugin data cleanup failed at {}: {error}",
|
||||
path.display()
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn path(&self, plugin_id: &str, key: &str) -> Result<PathBuf> {
|
||||
validate_component(plugin_id, "plugin id")?;
|
||||
validate_component(key, "plugin data key")?;
|
||||
Ok(self.root.join(plugin_id).join(format!("{key}.json")))
|
||||
}
|
||||
|
||||
fn lock(&self, plugin_id: &str) -> Arc<AsyncMutex<()>> {
|
||||
self.locks
|
||||
.lock()
|
||||
.entry(plugin_id.to_owned())
|
||||
.or_insert_with(|| Arc::new(AsyncMutex::new(())))
|
||||
.clone()
|
||||
}
|
||||
}
|
||||
|
||||
/// 单次完整写入:建目录、写临时文件、落盘、原子替换。
|
||||
/// 失败时返回失败步骤的标签,供上层区分重试与报错。
|
||||
fn write_once(
|
||||
directory: &Path,
|
||||
temporary: &Path,
|
||||
target: &Path,
|
||||
bytes: &[u8],
|
||||
) -> std::result::Result<(), (&'static str, std::io::Error)> {
|
||||
use std::io::Write;
|
||||
std::fs::create_dir_all(directory).map_err(|error| ("create data directory", error))?;
|
||||
let _ = set_directory_permissions(directory);
|
||||
let mut file =
|
||||
std::fs::File::create(temporary).map_err(|error| ("create temporary file", error))?;
|
||||
file.write_all(bytes)
|
||||
.map_err(|error| ("write temporary file", error))?;
|
||||
file.sync_all()
|
||||
.map_err(|error| ("sync temporary file", error))?;
|
||||
drop(file);
|
||||
let _ = set_file_permissions(temporary);
|
||||
// Windows 的 rename 不覆盖已存在文件,先删除旧文件。
|
||||
#[cfg(windows)]
|
||||
match std::fs::remove_file(target) {
|
||||
Ok(()) => {}
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
|
||||
Err(error) => return Err(("remove previous file", error)),
|
||||
}
|
||||
std::fs::rename(temporary, target).map_err(|error| ("replace target file", error))?;
|
||||
let _ = set_file_permissions(target);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Windows 下拒绝访问(5)与共享冲突(32)通常是杀软或索引器的
|
||||
/// 瞬时锁定,值得重试;其余错误与其他平台一律直接失败。
|
||||
fn transient(error: &std::io::Error) -> bool {
|
||||
#[cfg(windows)]
|
||||
{
|
||||
const ACCESS_DENIED: i32 = 5;
|
||||
const SHARING_VIOLATION: i32 = 32;
|
||||
matches!(
|
||||
error.raw_os_error(),
|
||||
Some(ACCESS_DENIED | SHARING_VIOLATION)
|
||||
)
|
||||
}
|
||||
#[cfg(not(windows))]
|
||||
{
|
||||
let _ = error;
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_component(value: &str, label: &str) -> Result<()> {
|
||||
if value.is_empty()
|
||||
|| value.len() > 128
|
||||
|| !value
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
|
||||
{
|
||||
return Err(Error::Config(format!("invalid {label}: {value}")));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn set_directory_permissions(path: &Path) -> Result<()> {
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o700))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn set_file_permissions(path: &Path) -> Result<()> {
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[tokio::test]
|
||||
async fn writes_reads_and_removes_json() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = PluginDataStore::new(root.path().join("data")).unwrap();
|
||||
store
|
||||
.update(
|
||||
"com.example",
|
||||
"state",
|
||||
&serde_json::json!({"token":"secret"}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
store.read("com.example", "state").await.unwrap()["token"],
|
||||
"secret"
|
||||
);
|
||||
store.clear("com.example").await.unwrap();
|
||||
assert!(store.read("com.example", "state").await.unwrap().is_null());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,217 @@
|
||||
//! Evaluates TypeScript plugin definitions through the host-owned virtual module.
|
||||
use std::{
|
||||
path::{Path, PathBuf},
|
||||
process::Stdio,
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use tokio::io::AsyncReadExt;
|
||||
|
||||
use super::descriptor::PluginModuleDefinition;
|
||||
use crate::{config, Error, Result};
|
||||
|
||||
const DEFINITION_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
const MAX_OUTPUT_BYTES: u64 = 2 * 1024 * 1024;
|
||||
const OUTPUT_PREFIX: &str = "CURSOR_BYOK_PLUGIN_DEFINITION:";
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PluginDefinitionLoader {
|
||||
sdk_dir: PathBuf,
|
||||
import_map: PathBuf,
|
||||
collector: PathBuf,
|
||||
worker: PathBuf,
|
||||
deno_dir: PathBuf,
|
||||
}
|
||||
|
||||
impl PluginDefinitionLoader {
|
||||
pub fn managed() -> Result<Self> {
|
||||
Self::in_directory(config::managed_data_dir()?.join("plugins/runtime/sdk/v1"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(super) fn for_test(root: &Path) -> Result<Self> {
|
||||
Self::in_directory(root.join(".plugin-sdk"))
|
||||
}
|
||||
|
||||
fn in_directory(sdk_dir: PathBuf) -> Result<Self> {
|
||||
std::fs::create_dir_all(&sdk_dir)?;
|
||||
std::fs::create_dir_all(sdk_dir.join("protocol"))?;
|
||||
let import_map = sdk_dir.join("import-map.json");
|
||||
let collector = sdk_dir.join("collect.ts");
|
||||
let worker = sdk_dir.join("worker.ts");
|
||||
let deno_dir = sdk_dir.join("cache");
|
||||
std::fs::create_dir_all(&deno_dir)?;
|
||||
let modules = [
|
||||
(&import_map, include_str!("sdk/import-map.json")),
|
||||
(&collector, include_str!("sdk/collect.ts")),
|
||||
(&worker, include_str!("sdk/worker.ts")),
|
||||
(&sdk_dir.join("plugin.ts"), include_str!("sdk/plugin.ts")),
|
||||
(
|
||||
&sdk_dir.join("provider.ts"),
|
||||
include_str!("sdk/provider.ts"),
|
||||
),
|
||||
(&sdk_dir.join("model.ts"), include_str!("sdk/model.ts")),
|
||||
(
|
||||
&sdk_dir.join("resource.ts"),
|
||||
include_str!("sdk/resource.ts"),
|
||||
),
|
||||
(
|
||||
&sdk_dir.join("protocol/openai_responses.ts"),
|
||||
include_str!("sdk/protocol/openai_responses.ts"),
|
||||
),
|
||||
(
|
||||
&sdk_dir.join("protocol/openai_chat.ts"),
|
||||
include_str!("sdk/protocol/openai_chat.ts"),
|
||||
),
|
||||
];
|
||||
for (path, content) in &modules {
|
||||
write_if_changed(path, content)?;
|
||||
}
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(&sdk_dir, std::fs::Permissions::from_mode(0o700))?;
|
||||
std::fs::set_permissions(
|
||||
sdk_dir.join("protocol"),
|
||||
std::fs::Permissions::from_mode(0o700),
|
||||
)?;
|
||||
std::fs::set_permissions(&deno_dir, std::fs::Permissions::from_mode(0o700))?;
|
||||
for (path, _) in &modules {
|
||||
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
|
||||
}
|
||||
}
|
||||
Ok(Self {
|
||||
sdk_dir,
|
||||
import_map,
|
||||
collector,
|
||||
worker,
|
||||
deno_dir,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn worker_path(&self) -> &Path {
|
||||
&self.worker
|
||||
}
|
||||
pub fn import_map(&self) -> &Path {
|
||||
&self.import_map
|
||||
}
|
||||
pub fn sdk_dir(&self) -> &Path {
|
||||
&self.sdk_dir
|
||||
}
|
||||
pub fn deno_dir(&self) -> &Path {
|
||||
&self.deno_dir
|
||||
}
|
||||
|
||||
pub async fn load(
|
||||
&self,
|
||||
executable: &Path,
|
||||
plugin_directory: &Path,
|
||||
entry: &Path,
|
||||
) -> Result<PluginModuleDefinition> {
|
||||
let entry_url = file_url(entry)?;
|
||||
let mut command = tokio::process::Command::new(executable);
|
||||
super::detach_console(&mut command);
|
||||
command
|
||||
.arg("run")
|
||||
.arg("--quiet")
|
||||
.arg("--no-config")
|
||||
.arg("--no-lock")
|
||||
.arg("--no-npm")
|
||||
.arg("--no-remote")
|
||||
.arg("--no-prompt")
|
||||
.arg(format!("--allow-read={}", plugin_directory.display()))
|
||||
.arg(format!("--allow-read={}", self.sdk_dir.display()))
|
||||
.arg(format!("--import-map={}", self.import_map.display()))
|
||||
.arg(&self.collector)
|
||||
.arg(entry_url.as_str())
|
||||
.env("DENO_DIR", &self.deno_dir)
|
||||
.env("DENO_NO_UPDATE_CHECK", "1")
|
||||
.stdin(Stdio::null())
|
||||
.stdout(Stdio::piped())
|
||||
.stderr(Stdio::piped())
|
||||
.kill_on_drop(true);
|
||||
let mut child = command.spawn()?;
|
||||
let stdout = child
|
||||
.stdout
|
||||
.take()
|
||||
.ok_or_else(|| Error::Config("cannot capture plugin definition output".into()))?;
|
||||
let stderr = child
|
||||
.stderr
|
||||
.take()
|
||||
.ok_or_else(|| Error::Config("cannot capture plugin definition error output".into()))?;
|
||||
let (stdout, stderr, status) = tokio::time::timeout(DEFINITION_TIMEOUT, async move {
|
||||
let (stdout, stderr, status) =
|
||||
tokio::join!(read_limited(stdout), read_limited(stderr), child.wait());
|
||||
Ok::<_, Error>((stdout?, stderr?, status?))
|
||||
})
|
||||
.await
|
||||
.map_err(|_| Error::Config("plugin definition evaluation timed out".into()))??;
|
||||
if !status.success() {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin definition evaluation failed: {}",
|
||||
String::from_utf8_lossy(&stderr).trim()
|
||||
)));
|
||||
}
|
||||
parse_definition_output(&stdout)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn file_url(path: &Path) -> Result<url::Url> {
|
||||
url::Url::from_file_path(path).map_err(|_| {
|
||||
Error::Config(format!(
|
||||
"plugin entry path is not a valid file URL: {}",
|
||||
path.display()
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn read_limited(reader: impl tokio::io::AsyncRead + Unpin) -> Result<Vec<u8>> {
|
||||
let mut bytes = Vec::new();
|
||||
reader
|
||||
.take(MAX_OUTPUT_BYTES + 1)
|
||||
.read_to_end(&mut bytes)
|
||||
.await?;
|
||||
if bytes.len() as u64 > MAX_OUTPUT_BYTES {
|
||||
return Err(Error::Config(
|
||||
"plugin definition output is larger than allowed".into(),
|
||||
));
|
||||
}
|
||||
Ok(bytes)
|
||||
}
|
||||
|
||||
fn parse_definition_output(output: &[u8]) -> Result<PluginModuleDefinition> {
|
||||
let output = String::from_utf8(output.to_vec()).map_err(|error| {
|
||||
Error::Config(format!("plugin definition output is not UTF-8: {error}"))
|
||||
})?;
|
||||
let json = output
|
||||
.lines()
|
||||
.rev()
|
||||
.find_map(|line| line.strip_prefix(OUTPUT_PREFIX))
|
||||
.ok_or_else(|| Error::Config("plugin definition did not produce a descriptor".into()))?;
|
||||
Ok(serde_json::from_str(json)?)
|
||||
}
|
||||
|
||||
pub(super) fn write_if_changed(path: &Path, content: &str) -> Result<()> {
|
||||
if std::fs::read(path).is_ok_and(|current| current == content.as_bytes()) {
|
||||
return Ok(());
|
||||
}
|
||||
std::fs::write(path, content)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[test]
|
||||
fn parses_descriptor_marker() {
|
||||
let output = br#"CURSOR_BYOK_PLUGIN_DEFINITION:{"providers":[{"id":"codex","displayName":"OpenAI Codex","description":null,"providerType":"openai","resourceType":"chatgpt-account","hasModels":true}],"resources":[{"type":"chatgpt-account","displayName":"ChatGPT accounts","add":[{"type":"oauth2.0","id":"chatgpt-device","displayName":"Sign in","description":null}],"import":{"displayName":"Import","description":null,"accept":[".json"],"multiple":true},"canRefresh":true,"canRemove":false}]}"#;
|
||||
let descriptor = parse_definition_output(output).unwrap();
|
||||
assert_eq!(descriptor.providers[0].id, "codex");
|
||||
assert_eq!(
|
||||
descriptor.providers[0].resource_type.as_deref(),
|
||||
Some("chatgpt-account")
|
||||
);
|
||||
assert_eq!(descriptor.resources[0].add[0].method_type, "oauth2.0");
|
||||
assert!(descriptor.resources[0].import.is_some());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
//! Defines serializable plugin capability definitions and desktop descriptors.
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::state::{ResourceRecord, ResourceState, StoredModel};
|
||||
|
||||
/// 由 collect.ts 输出的能力摘要;不含任何可执行内容。
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct PluginModuleDefinition {
|
||||
pub providers: Vec<ProviderDefinition>,
|
||||
#[serde(default)]
|
||||
pub resources: Vec<ResourceDefinition>,
|
||||
}
|
||||
|
||||
/// 插件提供的显示文本:纯字符串或 locale → 文本映射;核心原样透传,由前端解析。
|
||||
pub type LocalizedText = serde_json::Value;
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct ProviderDefinition {
|
||||
pub id: String,
|
||||
pub display_name: LocalizedText,
|
||||
#[serde(default)]
|
||||
pub description: LocalizedText,
|
||||
pub provider_type: String,
|
||||
#[serde(default)]
|
||||
pub resource_type: Option<String>,
|
||||
pub has_models: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct ResourceDefinition {
|
||||
#[serde(rename = "type")]
|
||||
pub resource_type: String,
|
||||
pub display_name: LocalizedText,
|
||||
#[serde(default)]
|
||||
pub add: Vec<AddMethodDefinition>,
|
||||
#[serde(default)]
|
||||
pub import: Option<ImportDefinition>,
|
||||
pub can_refresh: bool,
|
||||
pub can_remove: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct AddMethodDefinition {
|
||||
#[serde(rename = "type")]
|
||||
pub method_type: String,
|
||||
pub id: String,
|
||||
pub display_name: LocalizedText,
|
||||
#[serde(default)]
|
||||
pub description: LocalizedText,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct ImportDefinition {
|
||||
pub display_name: LocalizedText,
|
||||
#[serde(default)]
|
||||
pub description: LocalizedText,
|
||||
pub accept: Vec<String>,
|
||||
pub multiple: bool,
|
||||
}
|
||||
|
||||
pub const OAUTH2_ADD_METHOD: &str = "oauth2.0";
|
||||
|
||||
/// 桌面端看到的插件全貌。
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginDescriptor {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
pub version: String,
|
||||
pub author: Option<String>,
|
||||
pub icon: String,
|
||||
pub providers: Vec<PluginProviderDescriptor>,
|
||||
pub resources: Vec<PluginResourceDescriptor>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginProviderDescriptor {
|
||||
pub id: String,
|
||||
pub plugin_id: String,
|
||||
pub display_name: LocalizedText,
|
||||
pub description: LocalizedText,
|
||||
pub provider_type: String,
|
||||
pub resource_type: Option<String>,
|
||||
pub has_models: bool,
|
||||
/// 已满足调用条件:模型目录非空,且需要资源时至少有一条资源。
|
||||
pub configured: bool,
|
||||
pub models: Vec<PluginModelDescriptor>,
|
||||
}
|
||||
|
||||
/// 一个可直接被 Cursor 调用的插件模型。
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginModelDescriptor {
|
||||
/// 稳定模型 ID:`plugin:<plugin>/<provider>/<model>`。
|
||||
pub id: String,
|
||||
pub plugin_id: String,
|
||||
pub plugin_name: String,
|
||||
pub provider_id: String,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub description: Option<String>,
|
||||
pub icon: String,
|
||||
pub provider_type: String,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
pub images: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginResourceDescriptor {
|
||||
#[serde(rename = "type")]
|
||||
pub resource_type: String,
|
||||
pub display_name: LocalizedText,
|
||||
pub add: Vec<AddMethodDefinition>,
|
||||
pub import: Option<ImportDefinition>,
|
||||
pub can_refresh: bool,
|
||||
pub can_remove: bool,
|
||||
pub resources: Vec<PluginResourceView>,
|
||||
}
|
||||
|
||||
/// 单条资源的对外投影;凭证保留在核心存储,不进入该结构。
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct PluginResourceView {
|
||||
pub id: String,
|
||||
pub state: ResourceState,
|
||||
pub display_name: String,
|
||||
pub description: LocalizedText,
|
||||
pub metrics: Vec<ResourceMetric>,
|
||||
pub created_at_ms: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct ResourceMetric {
|
||||
pub id: String,
|
||||
pub label: LocalizedText,
|
||||
pub unit: String,
|
||||
pub value: f64,
|
||||
#[serde(default)]
|
||||
pub reset_at_ms: Option<i64>,
|
||||
}
|
||||
|
||||
/// 插件对一条资源的展示投影(resource.present 的返回值)。
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct ResourcePresentation {
|
||||
pub display_name: String,
|
||||
#[serde(default)]
|
||||
pub description: LocalizedText,
|
||||
#[serde(default)]
|
||||
pub metrics: Vec<ResourceMetric>,
|
||||
}
|
||||
|
||||
impl PluginResourceView {
|
||||
pub fn from_record(record: &ResourceRecord, presentation: ResourcePresentation) -> Self {
|
||||
Self {
|
||||
id: record.id.clone(),
|
||||
state: record.state.clone(),
|
||||
display_name: presentation.display_name,
|
||||
description: presentation.description,
|
||||
metrics: presentation.metrics,
|
||||
created_at_ms: record.created_at_ms,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub const ADAPTER_ID_PREFIX: &str = "plugin:";
|
||||
|
||||
pub fn model_id(plugin_id: &str, provider_id: &str, model_id: &str) -> String {
|
||||
format!("{ADAPTER_ID_PREFIX}{plugin_id}/{provider_id}/{model_id}")
|
||||
}
|
||||
|
||||
/// 解析稳定模型 ID;上游模型段允许包含 `/`。
|
||||
pub fn parse_model_id(value: &str) -> Option<(&str, &str, &str)> {
|
||||
let rest = value.strip_prefix(ADAPTER_ID_PREFIX)?;
|
||||
let (plugin_id, rest) = rest.split_once('/')?;
|
||||
let (provider_id, model_id) = rest.split_once('/')?;
|
||||
(!plugin_id.is_empty() && !provider_id.is_empty() && !model_id.is_empty()).then_some((
|
||||
plugin_id,
|
||||
provider_id,
|
||||
model_id,
|
||||
))
|
||||
}
|
||||
|
||||
impl PluginModelDescriptor {
|
||||
pub fn new(
|
||||
plugin_id: &str,
|
||||
plugin_name: &str,
|
||||
icon: &str,
|
||||
provider: &ProviderDefinition,
|
||||
model: &StoredModel,
|
||||
) -> Self {
|
||||
Self {
|
||||
id: model_id(plugin_id, &provider.id, &model.id),
|
||||
plugin_id: plugin_id.to_owned(),
|
||||
plugin_name: plugin_name.to_owned(),
|
||||
provider_id: provider.id.clone(),
|
||||
model_id: model.id.clone(),
|
||||
display_name: model.display_name.clone(),
|
||||
description: model.description.clone(),
|
||||
icon: icon.to_owned(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
max_output_tokens: model.max_output_tokens,
|
||||
images: model.images,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn parses_stable_model_ids_with_slashes() {
|
||||
let id = model_id("dev.example", "codex", "org/gpt-5");
|
||||
assert_eq!(
|
||||
parse_model_id(&id),
|
||||
Some(("dev.example", "codex", "org/gpt-5"))
|
||||
);
|
||||
assert_eq!(parse_model_id("plugin:only/one"), None);
|
||||
assert_eq!(parse_model_id("model-hash"), None);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,273 @@
|
||||
//! Downloads, verifies, extracts, and validates a pinned Deno runtime.
|
||||
use std::{
|
||||
io,
|
||||
path::{Path, PathBuf},
|
||||
process::Stdio,
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use futures_util::StreamExt;
|
||||
use sha2::{Digest, Sha256};
|
||||
use tokio::io::AsyncWriteExt;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use super::{
|
||||
asset::{RuntimeAsset, DENO_VERSION},
|
||||
runtime::PluginRuntimePhase,
|
||||
};
|
||||
use crate::{network, store::Store, Error, Result};
|
||||
|
||||
const DOWNLOAD_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
||||
const VALIDATION_TIMEOUT: Duration = Duration::from_secs(15);
|
||||
const MAX_ARCHIVE_BYTES: u64 = 128 * 1024 * 1024;
|
||||
|
||||
pub(super) async fn install(
|
||||
root: &Path,
|
||||
store: &Store,
|
||||
asset: RuntimeAsset,
|
||||
cancellation: CancellationToken,
|
||||
on_progress: impl Fn(PluginRuntimePhase, u64, Option<u64>),
|
||||
) -> Result<()> {
|
||||
let paths = RuntimePaths::new(root, asset);
|
||||
tokio::fs::create_dir_all(&paths.download_dir).await?;
|
||||
tokio::fs::create_dir_all(&paths.install_dir).await?;
|
||||
remove_if_exists(&paths.archive).await?;
|
||||
remove_if_exists(&paths.executable_staging).await?;
|
||||
remove_if_exists(&paths.ready_marker).await?;
|
||||
|
||||
let result = download_and_install(store, asset, &paths, &cancellation, &on_progress).await;
|
||||
if result.is_err() {
|
||||
let _ = remove_if_exists(&paths.archive).await;
|
||||
let _ = remove_if_exists(&paths.executable_staging).await;
|
||||
let _ = remove_if_exists(&paths.executable).await;
|
||||
let _ = remove_if_exists(&paths.ready_marker).await;
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
pub(super) fn runtime_complete(root: &Path, asset: RuntimeAsset) -> bool {
|
||||
let paths = RuntimePaths::new(root, asset);
|
||||
paths.executable.is_file() && paths.ready_marker.is_file()
|
||||
}
|
||||
|
||||
pub(super) fn runtime_executable(root: &Path, asset: RuntimeAsset) -> PathBuf {
|
||||
RuntimePaths::new(root, asset).executable
|
||||
}
|
||||
|
||||
async fn download_and_install(
|
||||
store: &Store,
|
||||
asset: RuntimeAsset,
|
||||
paths: &RuntimePaths,
|
||||
cancellation: &CancellationToken,
|
||||
on_progress: &impl Fn(PluginRuntimePhase, u64, Option<u64>),
|
||||
) -> Result<()> {
|
||||
ensure_not_cancelled(cancellation)?;
|
||||
let client = network::client(store).await?;
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = client
|
||||
.get(asset.download_url())
|
||||
.timeout(DOWNLOAD_TIMEOUT)
|
||||
.send() => response?,
|
||||
}
|
||||
.error_for_status()?;
|
||||
let total_bytes = response.content_length();
|
||||
if total_bytes.is_some_and(|size| size > MAX_ARCHIVE_BYTES) {
|
||||
return Err(Error::Config(
|
||||
"Deno runtime archive is larger than allowed".into(),
|
||||
));
|
||||
}
|
||||
|
||||
on_progress(PluginRuntimePhase::Downloading, 0, total_bytes);
|
||||
let mut archive = tokio::fs::File::create(&paths.archive).await?;
|
||||
let mut hasher = Sha256::new();
|
||||
let mut downloaded_bytes = 0_u64;
|
||||
let mut stream = response.bytes_stream();
|
||||
loop {
|
||||
let next = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
next = stream.next() => next,
|
||||
};
|
||||
let Some(chunk) = next else { break };
|
||||
let chunk = chunk?;
|
||||
downloaded_bytes = downloaded_bytes.saturating_add(chunk.len() as u64);
|
||||
if downloaded_bytes > MAX_ARCHIVE_BYTES {
|
||||
return Err(Error::Config(
|
||||
"Deno runtime archive is larger than allowed".into(),
|
||||
));
|
||||
}
|
||||
archive.write_all(&chunk).await?;
|
||||
hasher.update(&chunk);
|
||||
on_progress(
|
||||
PluginRuntimePhase::Downloading,
|
||||
downloaded_bytes,
|
||||
total_bytes,
|
||||
);
|
||||
}
|
||||
archive.flush().await?;
|
||||
archive.sync_all().await?;
|
||||
drop(archive);
|
||||
ensure_not_cancelled(cancellation)?;
|
||||
|
||||
on_progress(PluginRuntimePhase::Verifying, downloaded_bytes, total_bytes);
|
||||
let actual_hash = hex::encode(hasher.finalize());
|
||||
if actual_hash != asset.sha256 {
|
||||
return Err(Error::Config(format!(
|
||||
"Deno runtime checksum mismatch: expected {}, received {actual_hash}",
|
||||
asset.sha256
|
||||
)));
|
||||
}
|
||||
|
||||
on_progress(
|
||||
PluginRuntimePhase::Installing,
|
||||
downloaded_bytes,
|
||||
total_bytes,
|
||||
);
|
||||
let archive_path = paths.archive.clone();
|
||||
let staging_path = paths.executable_staging.clone();
|
||||
let executable_name = asset.executable_name();
|
||||
tokio::task::spawn_blocking(move || {
|
||||
extract_runtime_archive(&archive_path, &staging_path, executable_name)
|
||||
})
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("Deno extraction task failed: {error}")))??;
|
||||
ensure_not_cancelled(cancellation)?;
|
||||
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
tokio::fs::set_permissions(
|
||||
&paths.executable_staging,
|
||||
std::fs::Permissions::from_mode(0o700),
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
remove_if_exists(&paths.executable).await?;
|
||||
tokio::fs::rename(&paths.executable_staging, &paths.executable).await?;
|
||||
ensure_not_cancelled(cancellation)?;
|
||||
|
||||
on_progress(
|
||||
PluginRuntimePhase::Validating,
|
||||
downloaded_bytes,
|
||||
total_bytes,
|
||||
);
|
||||
validate_runtime(&paths.executable, cancellation).await?;
|
||||
ensure_not_cancelled(cancellation)?;
|
||||
tokio::fs::write(&paths.ready_marker, format!("deno {DENO_VERSION}\n")).await?;
|
||||
remove_if_exists(&paths.archive).await?;
|
||||
tracing::info!(
|
||||
version = DENO_VERSION,
|
||||
target = asset.target,
|
||||
path = %paths.executable.display(),
|
||||
"plugin runtime initialized"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
struct RuntimePaths {
|
||||
download_dir: PathBuf,
|
||||
install_dir: PathBuf,
|
||||
archive: PathBuf,
|
||||
executable: PathBuf,
|
||||
executable_staging: PathBuf,
|
||||
ready_marker: PathBuf,
|
||||
}
|
||||
|
||||
impl RuntimePaths {
|
||||
fn new(root: &Path, asset: RuntimeAsset) -> Self {
|
||||
let download_dir = root.join(".downloads");
|
||||
let install_dir = root
|
||||
.join("deno")
|
||||
.join(format!("v{DENO_VERSION}"))
|
||||
.join(asset.target);
|
||||
let executable = install_dir.join(asset.executable_name());
|
||||
Self {
|
||||
archive: download_dir.join(format!("{}.part", asset.archive_name())),
|
||||
executable_staging: install_dir.join(format!("{}.part", asset.executable_name())),
|
||||
ready_marker: install_dir.join(".ready"),
|
||||
download_dir,
|
||||
install_dir,
|
||||
executable,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn extract_runtime_archive(archive: &Path, output: &Path, executable_name: &str) -> Result<()> {
|
||||
let file = std::fs::File::open(archive)?;
|
||||
let mut archive = zip::ZipArchive::new(file)
|
||||
.map_err(|error| Error::Config(format!("invalid Deno runtime archive: {error}")))?;
|
||||
let mut executable = archive
|
||||
.by_name(executable_name)
|
||||
.map_err(|error| Error::Config(format!("Deno executable missing from archive: {error}")))?;
|
||||
let mut destination = std::fs::File::create(output)?;
|
||||
io::copy(&mut executable, &mut destination)?;
|
||||
destination.sync_all()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn validate_runtime(executable: &Path, cancellation: &CancellationToken) -> Result<()> {
|
||||
let mut command = tokio::process::Command::new(executable);
|
||||
super::detach_console(&mut command);
|
||||
command
|
||||
.arg("--version")
|
||||
.stdin(Stdio::null())
|
||||
.stderr(Stdio::piped())
|
||||
.stdout(Stdio::piped())
|
||||
.kill_on_drop(true);
|
||||
let output = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
result = tokio::time::timeout(VALIDATION_TIMEOUT, command.output()) => {
|
||||
result.map_err(|_| Error::Config("Deno runtime validation timed out".into()))??
|
||||
}
|
||||
};
|
||||
if !output.status.success() {
|
||||
return Err(Error::Config(format!(
|
||||
"Deno runtime validation failed: {}",
|
||||
String::from_utf8_lossy(&output.stderr).trim()
|
||||
)));
|
||||
}
|
||||
let expected = format!("deno {DENO_VERSION}");
|
||||
let stdout = String::from_utf8_lossy(&output.stdout);
|
||||
let version_line = stdout.lines().next().unwrap_or_default().trim();
|
||||
if version_line != expected && !version_line.starts_with(&format!("{expected} ")) {
|
||||
return Err(Error::Config(format!(
|
||||
"unexpected Deno runtime version: {version_line}"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn ensure_not_cancelled(cancellation: &CancellationToken) -> Result<()> {
|
||||
if cancellation.is_cancelled() {
|
||||
Err(Error::Cancelled)
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
async fn remove_if_exists(path: &Path) -> Result<()> {
|
||||
match tokio::fs::remove_file(path).await {
|
||||
Ok(()) => Ok(()),
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn uses_versioned_runtime_directory() {
|
||||
let root = PathBuf::from("/tmp/plugin-runtime");
|
||||
let asset = super::super::asset::RuntimeAsset::for_platform("macos", "aarch64").unwrap();
|
||||
let paths = RuntimePaths::new(&root, asset);
|
||||
assert_eq!(
|
||||
paths.executable,
|
||||
root.join("deno")
|
||||
.join(format!("v{DENO_VERSION}"))
|
||||
.join(asset.target)
|
||||
.join("deno")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
//! Defines and validates the static filesystem plugin manifest.
|
||||
use std::{collections::HashSet, path::Path};
|
||||
|
||||
use regex::Regex;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
pub const PLUGIN_API_VERSION: u32 = 1;
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct PluginManifest {
|
||||
pub api_version: u32,
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
/// 插件自身版本;内置插件预装时以它为缓存键决定是否重新落盘。
|
||||
pub version: String,
|
||||
#[serde(default)]
|
||||
pub author: Option<String>,
|
||||
/// 插件要求的最低应用版本;应用过旧时插件被忽略。
|
||||
#[serde(default)]
|
||||
pub min_app_version: Option<String>,
|
||||
pub icon: String,
|
||||
pub entry: String,
|
||||
#[serde(default)]
|
||||
pub permissions: PluginPermissions,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct PluginPermissions {
|
||||
#[serde(default)]
|
||||
pub network: Vec<String>,
|
||||
}
|
||||
|
||||
impl PluginManifest {
|
||||
pub fn validate(&self, directory: &Path) -> Result<()> {
|
||||
if self.api_version != PLUGIN_API_VERSION {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{}' uses unsupported API version {}",
|
||||
self.id, self.api_version
|
||||
)));
|
||||
}
|
||||
validate_id(&self.id, "plugin id")?;
|
||||
required(&self.name, "plugin name")?;
|
||||
parse_version(&self.version)
|
||||
.ok_or_else(|| Error::Config(format!("invalid plugin version: {}", self.version)))?;
|
||||
if let Some(minimum) = &self.min_app_version {
|
||||
parse_version(minimum)
|
||||
.ok_or_else(|| Error::Config(format!("invalid plugin minAppVersion: {minimum}")))?;
|
||||
}
|
||||
validate_entry_path(directory, &self.entry)?;
|
||||
validate_asset_path(directory, &self.icon)?;
|
||||
let mut hosts = HashSet::new();
|
||||
for host in &self.permissions.network {
|
||||
validate_network_host(host)?;
|
||||
if !hosts.insert(host.to_ascii_lowercase()) {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{}' contains duplicate network host '{host}'",
|
||||
self.id
|
||||
)));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// 解析 semver 的核心三段(忽略预发布/构建后缀),格式非法返回 None。
|
||||
pub(super) fn parse_version(value: &str) -> Option<(u64, u64, u64)> {
|
||||
let core = value.split(['-', '+']).next()?;
|
||||
let mut parts = core.split('.');
|
||||
let major = parts.next()?.parse().ok()?;
|
||||
let minor = parts.next()?.parse().ok()?;
|
||||
let patch = parts.next()?.parse().ok()?;
|
||||
parts.next().is_none().then_some((major, minor, patch))
|
||||
}
|
||||
|
||||
pub(super) fn version_at_least(actual: &str, minimum: &str) -> bool {
|
||||
match (parse_version(actual), parse_version(minimum)) {
|
||||
(Some(actual), Some(minimum)) => actual >= minimum,
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn validate_id(value: &str, label: &str) -> Result<()> {
|
||||
static ID: std::sync::OnceLock<Regex> = std::sync::OnceLock::new();
|
||||
let expression = ID.get_or_init(|| Regex::new(r"^[a-z0-9]+(?:[._-][a-z0-9]+)*$").unwrap());
|
||||
if expression.is_match(value) {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(Error::Config(format!("invalid {label}: {value}")))
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_network_host(value: &str) -> Result<()> {
|
||||
if value.is_empty()
|
||||
|| value.contains('/')
|
||||
|| value.contains(':')
|
||||
|| value.starts_with('.')
|
||||
|| value.ends_with('.')
|
||||
{
|
||||
return Err(Error::Config(format!(
|
||||
"invalid plugin network host: {value}"
|
||||
)));
|
||||
}
|
||||
let parsed = url::Url::parse(&format!("https://{value}")).map_err(|error| {
|
||||
Error::Config(format!("invalid plugin network host '{value}': {error}"))
|
||||
})?;
|
||||
if parsed.host_str() != Some(value) {
|
||||
return Err(Error::Config(format!(
|
||||
"invalid plugin network host: {value}"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn required<'a>(value: &'a str, label: &str) -> Result<&'a str> {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
Err(Error::Config(format!("{label} is required")))
|
||||
} else {
|
||||
Ok(value)
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_entry_path(directory: &Path, value: &str) -> Result<()> {
|
||||
let path = Path::new(value);
|
||||
if !is_safe_relative_path(path) {
|
||||
return Err(Error::Config(format!("invalid plugin entry path: {value}")));
|
||||
}
|
||||
let extension = path
|
||||
.extension()
|
||||
.and_then(|value| value.to_str())
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase();
|
||||
if !matches!(extension.as_str(), "js" | "mjs" | "ts" | "mts") {
|
||||
return Err(Error::Config(format!(
|
||||
"unsupported plugin entry format: {value}"
|
||||
)));
|
||||
}
|
||||
let entry = directory.join(path);
|
||||
if !entry.is_file() {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin entry does not exist: {value}"
|
||||
)));
|
||||
}
|
||||
let root = directory.canonicalize()?;
|
||||
let entry = entry.canonicalize()?;
|
||||
if !entry.starts_with(root) {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin entry escapes its directory: {value}"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn is_safe_relative_path(path: &Path) -> bool {
|
||||
!path.is_absolute()
|
||||
&& !path.components().any(|component| {
|
||||
matches!(
|
||||
component,
|
||||
std::path::Component::ParentDir
|
||||
| std::path::Component::RootDir
|
||||
| std::path::Component::Prefix(_)
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn validate_asset_path(directory: &Path, value: &str) -> Result<()> {
|
||||
let path = Path::new(value);
|
||||
if !is_safe_relative_path(path) {
|
||||
return Err(Error::Config(format!("invalid plugin asset path: {value}")));
|
||||
}
|
||||
let extension = path
|
||||
.extension()
|
||||
.and_then(|value| value.to_str())
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase();
|
||||
if !matches!(extension.as_str(), "svg" | "png" | "webp") {
|
||||
return Err(Error::Config(format!(
|
||||
"unsupported plugin icon format: {value}"
|
||||
)));
|
||||
}
|
||||
if !directory.join(path).is_file() {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin icon does not exist: {value}"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn rejects_urls_in_network_host_allowlist() {
|
||||
assert!(validate_network_host("https://example.com").is_err());
|
||||
assert!(validate_network_host("example.com:443").is_err());
|
||||
assert!(validate_network_host("example.com").is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compares_semver_cores_and_ignores_prerelease_suffixes() {
|
||||
assert_eq!(parse_version("0.1.5-beta.1"), Some((0, 1, 5)));
|
||||
assert_eq!(parse_version("1.2"), None);
|
||||
assert!(version_at_least("0.1.5-beta.1", "0.1.5"));
|
||||
assert!(version_at_least("0.2.0", "0.1.9"));
|
||||
assert!(!version_at_least("0.1.4", "0.1.5"));
|
||||
assert!(!version_at_least("bogus", "0.1.0"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
//! Owns filesystem plugin discovery, sandboxed workers, and plugin providers.
|
||||
mod asset;
|
||||
mod builtin;
|
||||
mod catalog;
|
||||
mod data;
|
||||
mod definition;
|
||||
mod descriptor;
|
||||
mod installation;
|
||||
mod manifest;
|
||||
mod protocol;
|
||||
mod registry;
|
||||
mod runtime;
|
||||
mod state;
|
||||
mod wire;
|
||||
mod worker;
|
||||
|
||||
pub use descriptor::{
|
||||
parse_model_id, PluginDescriptor, PluginModelDescriptor, PluginProviderDescriptor,
|
||||
PluginResourceDescriptor, PluginResourceView, ADAPTER_ID_PREFIX,
|
||||
};
|
||||
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
|
||||
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
|
||||
pub(crate) use wire::llm_request as plugin_llm_request;
|
||||
|
||||
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
|
||||
#[cfg(windows)]
|
||||
fn detach_console(command: &mut tokio::process::Command) {
|
||||
command.creation_flags(0x0800_0000);
|
||||
}
|
||||
|
||||
#[cfg(not(windows))]
|
||||
fn detach_console(_command: &mut tokio::process::Command) {}
|
||||
@@ -0,0 +1,92 @@
|
||||
//! Defines newline-delimited messages exchanged with a plugin worker.
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum HostMessage<'a> {
|
||||
Request {
|
||||
id: &'a str,
|
||||
method: &'a str,
|
||||
params: &'a serde_json::Value,
|
||||
},
|
||||
Cancel {
|
||||
id: &'a str,
|
||||
},
|
||||
HostResult {
|
||||
id: &'a str,
|
||||
result: &'a serde_json::Value,
|
||||
},
|
||||
HostError {
|
||||
id: &'a str,
|
||||
error: &'a str,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum WorkerMessage {
|
||||
Result {
|
||||
id: String,
|
||||
#[serde(default)]
|
||||
result: serde_json::Value,
|
||||
#[serde(default)]
|
||||
error: Option<String>,
|
||||
},
|
||||
/// 流式请求(provider.invoke)在最终 Result 之前发出的模型事件。
|
||||
Event {
|
||||
id: String,
|
||||
event: serde_json::Value,
|
||||
},
|
||||
HostCall {
|
||||
id: String,
|
||||
#[serde(rename = "requestId")]
|
||||
request_id: String,
|
||||
method: String,
|
||||
#[serde(default)]
|
||||
params: serde_json::Value,
|
||||
},
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn parses_multiplexed_host_call_and_events() {
|
||||
let message: WorkerMessage = serde_json::from_value(serde_json::json!({
|
||||
"type": "host_call",
|
||||
"id": "host-2",
|
||||
"requestId": "request-1",
|
||||
"method": "network.fetch",
|
||||
"params": { "url": "https://example.com" }
|
||||
}))
|
||||
.unwrap();
|
||||
match message {
|
||||
WorkerMessage::HostCall {
|
||||
id,
|
||||
request_id,
|
||||
method,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(id, "host-2");
|
||||
assert_eq!(request_id, "request-1");
|
||||
assert_eq!(method, "network.fetch");
|
||||
}
|
||||
_ => panic!("expected host call"),
|
||||
}
|
||||
|
||||
let message: WorkerMessage = serde_json::from_value(serde_json::json!({
|
||||
"type": "event",
|
||||
"id": "request-1",
|
||||
"event": { "type": "text-delta", "text": "hi" }
|
||||
}))
|
||||
.unwrap();
|
||||
match message {
|
||||
WorkerMessage::Event { id, event } => {
|
||||
assert_eq!(id, "request-1");
|
||||
assert_eq!(event["type"], "text-delta");
|
||||
}
|
||||
_ => panic!("expected event"),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,959 @@
|
||||
//! Orchestrates plugin capabilities: resources, model catalogs, and invocation.
|
||||
use std::{collections::HashMap, path::Path, sync::Arc};
|
||||
|
||||
use async_stream::try_stream;
|
||||
use serde::Serialize;
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use super::{
|
||||
catalog::{PluginCatalog, PluginEntry},
|
||||
data::PluginDataStore,
|
||||
descriptor::{
|
||||
parse_model_id, PluginDescriptor, PluginModelDescriptor, PluginProviderDescriptor,
|
||||
PluginResourceDescriptor, PluginResourceView, ProviderDefinition, ResourceDefinition,
|
||||
ResourcePresentation, OAUTH2_ADD_METHOD,
|
||||
},
|
||||
runtime::PluginRuntime,
|
||||
state::{now_ms, PluginStateStore, ResourceDraft, ResourcePatch, ResourceRecord, StoredModel},
|
||||
wire,
|
||||
worker::{PluginWorker, WorkerStreamItem},
|
||||
};
|
||||
use crate::{
|
||||
model::ModelInvocation, provider::ModelEvent, provider::ProviderStream, store::Store, Error,
|
||||
Result,
|
||||
};
|
||||
|
||||
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
|
||||
const MAX_IMPORT_DRAFTS: usize = 256;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PluginRegistry {
|
||||
inner: Arc<RegistryInner>,
|
||||
}
|
||||
|
||||
struct RegistryInner {
|
||||
store: Store,
|
||||
runtime: PluginRuntime,
|
||||
catalog: PluginCatalog,
|
||||
state: PluginStateStore,
|
||||
entries: RwLock<Option<Vec<PluginEntry>>>,
|
||||
workers: Mutex<HashMap<String, Arc<PluginWorker>>>,
|
||||
oauth_sessions: Mutex<HashMap<String, OAuthSession>>,
|
||||
}
|
||||
|
||||
struct OAuthSession {
|
||||
plugin_id: String,
|
||||
resource_type: String,
|
||||
method_id: String,
|
||||
session: serde_json::Value,
|
||||
expires_at_ms: i64,
|
||||
poll_interval_ms: i64,
|
||||
next_poll_at_ms: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct OAuthBeginResponse {
|
||||
pub session_id: String,
|
||||
pub user_code: String,
|
||||
pub verification_url: String,
|
||||
pub verification_url_complete: Option<String>,
|
||||
pub expires_at_ms: i64,
|
||||
pub poll_interval_ms: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase", tag = "status")]
|
||||
pub enum OAuthPollResponse {
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Pending { poll_interval_ms: i64 },
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Completed {
|
||||
added: usize,
|
||||
updated: usize,
|
||||
model_sync_error: Option<String>,
|
||||
},
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Denied { message: Option<String> },
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Failed { message: String },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ImportResponse {
|
||||
pub added: usize,
|
||||
pub updated: usize,
|
||||
pub warnings: Vec<String>,
|
||||
pub model_sync_error: Option<String>,
|
||||
}
|
||||
|
||||
/// 路由分支在建立 Recorder 时需要的插件模型元数据。
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct PluginInvocationPlan {
|
||||
pub model: PluginModelDescriptor,
|
||||
pub request_url: String,
|
||||
}
|
||||
|
||||
impl PluginRegistry {
|
||||
pub fn managed(store: Store, runtime: PluginRuntime, app_version: String) -> Result<Self> {
|
||||
let data = PluginDataStore::managed()?;
|
||||
Ok(Self {
|
||||
inner: Arc::new(RegistryInner {
|
||||
store,
|
||||
runtime,
|
||||
catalog: PluginCatalog::managed(app_version)?,
|
||||
state: PluginStateStore::new(data),
|
||||
entries: RwLock::new(None),
|
||||
workers: Mutex::new(HashMap::new()),
|
||||
oauth_sessions: Mutex::new(HashMap::new()),
|
||||
}),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn plugins(&self) -> Vec<PluginDescriptor> {
|
||||
let Some(executable) = self.inner.runtime.executable() else {
|
||||
return self
|
||||
.inner
|
||||
.catalog
|
||||
.manifests()
|
||||
.into_iter()
|
||||
.map(|(manifest, icon)| PluginDescriptor {
|
||||
id: manifest.id,
|
||||
name: manifest.name,
|
||||
version: manifest.version,
|
||||
author: manifest.author,
|
||||
icon,
|
||||
providers: Vec::new(),
|
||||
resources: Vec::new(),
|
||||
})
|
||||
.collect();
|
||||
};
|
||||
let mut plugins = Vec::new();
|
||||
for entry in self.entries(&executable).await {
|
||||
plugins.push(self.descriptor(&entry, &executable).await);
|
||||
}
|
||||
plugins
|
||||
}
|
||||
|
||||
/// 已满足调用条件的全部插件模型;每个模型独立进入 Cursor 目录。
|
||||
pub async fn configured_models(&self) -> Vec<PluginModelDescriptor> {
|
||||
let Some(executable) = self.inner.runtime.executable() else {
|
||||
return Vec::new();
|
||||
};
|
||||
let mut models = Vec::new();
|
||||
for entry in self.entries(&executable).await {
|
||||
for provider in &entry.definition.providers {
|
||||
if !self.provider_configured(&entry, provider).await {
|
||||
continue;
|
||||
}
|
||||
let stored = self
|
||||
.inner
|
||||
.state
|
||||
.models(&entry.manifest.id, &provider.id)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
models.extend(stored.iter().map(|model| {
|
||||
PluginModelDescriptor::new(
|
||||
&entry.manifest.id,
|
||||
&entry.manifest.name,
|
||||
&entry.icon,
|
||||
provider,
|
||||
model,
|
||||
)
|
||||
}));
|
||||
}
|
||||
}
|
||||
models
|
||||
}
|
||||
|
||||
pub async fn model_descriptor(&self, model_id: &str) -> Result<PluginModelDescriptor> {
|
||||
let (plugin_id, provider_id, upstream_id) = parse_model_id(model_id)
|
||||
.ok_or_else(|| Error::Provider(format!("invalid plugin model ID: {model_id}")))?;
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, plugin_id).await?;
|
||||
let provider = find_provider(&entry, provider_id)?;
|
||||
let stored = self
|
||||
.inner
|
||||
.state
|
||||
.models(plugin_id, provider_id)
|
||||
.await?
|
||||
.into_iter()
|
||||
.find(|model| model.id == upstream_id)
|
||||
.ok_or_else(|| Error::RunNotFound(format!("plugin model {model_id}")))?;
|
||||
Ok(PluginModelDescriptor::new(
|
||||
plugin_id,
|
||||
&entry.manifest.name,
|
||||
&entry.icon,
|
||||
provider,
|
||||
&stored,
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn plan_model(&self, model_id: &str) -> Result<PluginInvocationPlan> {
|
||||
let model = self.model_descriptor(model_id).await?;
|
||||
let request_url = format!("plugin://{}/{}", model.plugin_id, model.provider_id);
|
||||
Ok(PluginInvocationPlan { model, request_url })
|
||||
}
|
||||
|
||||
/// 插件模型的统一 Provider 流:选首个可用资源,经 Worker 执行,
|
||||
/// 事件与内置 Provider 走同一管道。未来的负载均衡在这里换资源重试。
|
||||
pub fn stream_model(
|
||||
&self,
|
||||
invocation: ModelInvocation,
|
||||
cancellation: CancellationToken,
|
||||
) -> ProviderStream {
|
||||
let registry = self.clone();
|
||||
Box::pin(try_stream! {
|
||||
let model_id = invocation.request.model.model_id.clone();
|
||||
let (plugin_id, provider_id, upstream_id) = parse_model_id(&model_id)
|
||||
.map(|(plugin, provider, model)| (plugin.to_owned(), provider.to_owned(), model.to_owned()))
|
||||
.ok_or_else(|| Error::Provider(format!("invalid plugin model ID: {model_id}")))?;
|
||||
let executable = registry.executable()?;
|
||||
let entry = registry.find_entry(&executable, &plugin_id).await?;
|
||||
let provider = find_provider(&entry, &provider_id)?.clone();
|
||||
let stored = registry.inner.state.models(&plugin_id, &provider_id).await?
|
||||
.into_iter()
|
||||
.find(|model| model.id == upstream_id)
|
||||
.ok_or_else(|| Error::RunNotFound(format!("plugin model {model_id}")))?;
|
||||
let resource = match &provider.resource_type {
|
||||
Some(resource_type) => Some((
|
||||
resource_type.clone(),
|
||||
registry.select_resource(&plugin_id, resource_type).await?,
|
||||
)),
|
||||
None => None,
|
||||
};
|
||||
let request = wire::llm_request(&invocation)?;
|
||||
let params = serde_json::json!({
|
||||
"providerId": provider_id,
|
||||
"model": stored.snapshot(),
|
||||
"resource": resource.as_ref().map(|(resource_type, record)| record.snapshot(resource_type)),
|
||||
"request": request,
|
||||
});
|
||||
let worker = registry.worker(&entry, &executable).await;
|
||||
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone()).await?;
|
||||
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
|
||||
while let Some(item) = items.recv().await {
|
||||
match item {
|
||||
WorkerStreamItem::Event(event) => {
|
||||
yield wire::model_event(&event)?;
|
||||
}
|
||||
WorkerStreamItem::Result(result) => {
|
||||
let value = result?;
|
||||
let status = value.get("status").and_then(serde_json::Value::as_str).unwrap_or_default();
|
||||
let patch = value.get("patch")
|
||||
.filter(|patch| !patch.is_null())
|
||||
.map(|patch| serde_json::from_value::<ResourcePatch>(patch.clone()))
|
||||
.transpose()?;
|
||||
if let (Some(patch), Some((resource_type, record))) = (patch, resource.as_ref()) {
|
||||
if let Err(error) = registry.inner.state
|
||||
.apply_patch(&plugin_id, resource_type, &record.id, patch).await
|
||||
{
|
||||
tracing::warn!(plugin = %plugin_id, %error, "failed to apply plugin resource patch");
|
||||
}
|
||||
}
|
||||
match status {
|
||||
"completed" => return,
|
||||
"resource-error" | "request-error" => {
|
||||
let message = value.get("message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("plugin provider call failed");
|
||||
Err(Error::Provider(message.to_owned()))?;
|
||||
}
|
||||
status => {
|
||||
Err(Error::Protocol(format!("unknown plugin provider result: {status}")))?;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(Error::Provider(format!("plugin '{plugin_id}' worker stopped mid-stream")))?;
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn oauth_begin(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
method_id: &str,
|
||||
) -> Result<OAuthBeginResponse> {
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, plugin_id).await?;
|
||||
let resource = find_resource(&entry, resource_type)?;
|
||||
let method = resource
|
||||
.add
|
||||
.iter()
|
||||
.find(|method| method.id == method_id && method.method_type == OAUTH2_ADD_METHOD)
|
||||
.ok_or_else(|| {
|
||||
Error::Config(format!(
|
||||
"plugin '{plugin_id}' does not define OAuth method '{method_id}'"
|
||||
))
|
||||
})?;
|
||||
let value = self
|
||||
.worker(&entry, &executable)
|
||||
.await
|
||||
.invoke(
|
||||
"oauth.begin",
|
||||
serde_json::json!({ "resourceType": resource_type, "methodId": method.id }),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await?;
|
||||
let begin: OAuth2Begin = serde_json::from_value(value)?;
|
||||
let session_id = uuid::Uuid::new_v4().to_string();
|
||||
self.inner.oauth_sessions.lock().await.insert(
|
||||
session_id.clone(),
|
||||
OAuthSession {
|
||||
plugin_id: plugin_id.to_owned(),
|
||||
resource_type: resource_type.to_owned(),
|
||||
method_id: method_id.to_owned(),
|
||||
session: begin.session,
|
||||
expires_at_ms: begin.expires_at_ms,
|
||||
poll_interval_ms: begin.poll_interval_ms.max(1_000),
|
||||
next_poll_at_ms: now_ms() + begin.poll_interval_ms.max(1_000),
|
||||
},
|
||||
);
|
||||
Ok(OAuthBeginResponse {
|
||||
session_id,
|
||||
user_code: begin.user_code,
|
||||
verification_url: begin.verification_url,
|
||||
verification_url_complete: begin.verification_url_complete,
|
||||
expires_at_ms: begin.expires_at_ms,
|
||||
poll_interval_ms: begin.poll_interval_ms.max(1_000),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn oauth_poll(&self, session_id: &str) -> Result<OAuthPollResponse> {
|
||||
let now = now_ms();
|
||||
let (plugin_id, resource_type, method_id, session, poll_interval_ms) = {
|
||||
let mut sessions = self.inner.oauth_sessions.lock().await;
|
||||
let Some(state) = sessions.get_mut(session_id) else {
|
||||
return Ok(OAuthPollResponse::Failed {
|
||||
message: "authorization session no longer exists".into(),
|
||||
});
|
||||
};
|
||||
if now >= state.expires_at_ms {
|
||||
sessions.remove(session_id);
|
||||
return Ok(OAuthPollResponse::Failed {
|
||||
message: "device authorization expired".into(),
|
||||
});
|
||||
}
|
||||
if now < state.next_poll_at_ms {
|
||||
return Ok(OAuthPollResponse::Pending {
|
||||
poll_interval_ms: state.poll_interval_ms,
|
||||
});
|
||||
}
|
||||
state.next_poll_at_ms = now + state.poll_interval_ms;
|
||||
(
|
||||
state.plugin_id.clone(),
|
||||
state.resource_type.clone(),
|
||||
state.method_id.clone(),
|
||||
state.session.clone(),
|
||||
state.poll_interval_ms,
|
||||
)
|
||||
};
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, &plugin_id).await?;
|
||||
let value = self
|
||||
.worker(&entry, &executable)
|
||||
.await
|
||||
.invoke(
|
||||
"oauth.poll",
|
||||
serde_json::json!({
|
||||
"resourceType": resource_type,
|
||||
"methodId": method_id,
|
||||
"session": session,
|
||||
}),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await?;
|
||||
let poll: OAuth2Poll = serde_json::from_value(value)?;
|
||||
match poll {
|
||||
OAuth2Poll::Pending { session } => {
|
||||
self.update_session(session_id, session, None).await;
|
||||
Ok(OAuthPollResponse::Pending { poll_interval_ms })
|
||||
}
|
||||
OAuth2Poll::SlowDown { session } => {
|
||||
let interval = poll_interval_ms + OAUTH_SLOW_DOWN_STEP_MS;
|
||||
self.update_session(session_id, session, Some(interval))
|
||||
.await;
|
||||
Ok(OAuthPollResponse::Pending {
|
||||
poll_interval_ms: interval,
|
||||
})
|
||||
}
|
||||
OAuth2Poll::Completed { resources } => {
|
||||
// 持久化成功后才销毁会话:写盘瞬时失败时下次轮询还能重试。
|
||||
let outcome = self
|
||||
.inner
|
||||
.state
|
||||
.upsert_resources(&plugin_id, &resource_type, resources)
|
||||
.await?;
|
||||
self.inner.oauth_sessions.lock().await.remove(session_id);
|
||||
let model_sync_error = self
|
||||
.sync_provider_models_for_resource(&entry, &executable, &resource_type)
|
||||
.await;
|
||||
Ok(OAuthPollResponse::Completed {
|
||||
added: outcome.added,
|
||||
updated: outcome.updated,
|
||||
model_sync_error,
|
||||
})
|
||||
}
|
||||
OAuth2Poll::Denied { message } => {
|
||||
self.inner.oauth_sessions.lock().await.remove(session_id);
|
||||
Ok(OAuthPollResponse::Denied { message })
|
||||
}
|
||||
OAuth2Poll::Failed { message } => {
|
||||
self.inner.oauth_sessions.lock().await.remove(session_id);
|
||||
Ok(OAuthPollResponse::Failed { message })
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn import_resources(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
files: serde_json::Value,
|
||||
) -> Result<ImportResponse> {
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, plugin_id).await?;
|
||||
let resource = find_resource(&entry, resource_type)?;
|
||||
if resource.import.is_none() {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' resource '{resource_type}' does not support import"
|
||||
)));
|
||||
}
|
||||
let value = self
|
||||
.worker(&entry, &executable)
|
||||
.await
|
||||
.invoke(
|
||||
"import.parse",
|
||||
serde_json::json!({ "resourceType": resource_type, "files": files }),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await?;
|
||||
let parsed: ImportParseResult = serde_json::from_value(value)?;
|
||||
if parsed.resources.is_empty() {
|
||||
return Err(Error::Config(
|
||||
parsed
|
||||
.warnings
|
||||
.first()
|
||||
.cloned()
|
||||
.unwrap_or_else(|| "import produced no resources".into()),
|
||||
));
|
||||
}
|
||||
if parsed.resources.len() > MAX_IMPORT_DRAFTS {
|
||||
return Err(Error::Config(format!(
|
||||
"import produced more than {MAX_IMPORT_DRAFTS} resources"
|
||||
)));
|
||||
}
|
||||
let outcome = self
|
||||
.inner
|
||||
.state
|
||||
.upsert_resources(plugin_id, resource_type, parsed.resources)
|
||||
.await?;
|
||||
let model_sync_error = self
|
||||
.sync_provider_models_for_resource(&entry, &executable, resource_type)
|
||||
.await;
|
||||
Ok(ImportResponse {
|
||||
added: outcome.added,
|
||||
updated: outcome.updated,
|
||||
warnings: parsed.warnings,
|
||||
model_sync_error,
|
||||
})
|
||||
}
|
||||
|
||||
/// 导出某资源类型的全部私有数据,供备份或迁移;格式与批量导入兼容。
|
||||
pub async fn export_resources(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
) -> Result<serde_json::Value> {
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, plugin_id).await?;
|
||||
find_resource(&entry, resource_type)?;
|
||||
let records = self.inner.state.resources(plugin_id, resource_type).await?;
|
||||
Ok(serde_json::json!({
|
||||
"accounts": records
|
||||
.iter()
|
||||
.map(|record| record.private_data.clone())
|
||||
.collect::<Vec<_>>(),
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn refresh_resource(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<()> {
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, plugin_id).await?;
|
||||
let resource = find_resource(&entry, resource_type)?;
|
||||
if !resource.can_refresh {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{plugin_id}' resource '{resource_type}' does not support refresh"
|
||||
)));
|
||||
}
|
||||
let record = self
|
||||
.find_record(plugin_id, resource_type, resource_id)
|
||||
.await?;
|
||||
let value = self
|
||||
.worker(&entry, &executable)
|
||||
.await
|
||||
.invoke(
|
||||
"resource.refresh",
|
||||
serde_json::json!({
|
||||
"resourceType": resource_type,
|
||||
"resource": record.snapshot(resource_type),
|
||||
}),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await?;
|
||||
let patch: ResourcePatch = serde_json::from_value(value)?;
|
||||
self.inner
|
||||
.state
|
||||
.apply_patch(plugin_id, resource_type, resource_id, patch)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn delete_resource(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<()> {
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, plugin_id).await?;
|
||||
let resource = find_resource(&entry, resource_type)?;
|
||||
let record = self
|
||||
.find_record(plugin_id, resource_type, resource_id)
|
||||
.await?;
|
||||
if resource.can_remove {
|
||||
// 上游撤销失败不阻塞本地删除:用户必须能移除已失效的资源。
|
||||
if let Err(error) = self
|
||||
.worker(&entry, &executable)
|
||||
.await
|
||||
.invoke(
|
||||
"resource.remove",
|
||||
serde_json::json!({
|
||||
"resourceType": resource_type,
|
||||
"resource": record.snapshot(resource_type),
|
||||
}),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(plugin = %plugin_id, %error, "plugin resource remove hook failed");
|
||||
}
|
||||
}
|
||||
self.inner
|
||||
.state
|
||||
.remove_resource(plugin_id, resource_type, resource_id)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn sync_models(&self, plugin_id: &str, provider_id: &str) -> Result<usize> {
|
||||
let executable = self.executable()?;
|
||||
let entry = self.find_entry(&executable, plugin_id).await?;
|
||||
let provider = find_provider(&entry, provider_id)?.clone();
|
||||
self.sync_provider_models(&entry, &executable, &provider)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn remove(&self, plugin_id: &str) -> Result<()> {
|
||||
if let Some(worker) = self.inner.workers.lock().await.remove(plugin_id) {
|
||||
worker.stop().await;
|
||||
}
|
||||
self.inner.state.clear(plugin_id).await
|
||||
}
|
||||
|
||||
async fn descriptor(&self, entry: &PluginEntry, executable: &Path) -> PluginDescriptor {
|
||||
let plugin_id = &entry.manifest.id;
|
||||
let mut providers = Vec::new();
|
||||
for provider in &entry.definition.providers {
|
||||
let stored = self
|
||||
.inner
|
||||
.state
|
||||
.models(plugin_id, &provider.id)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let configured = self.provider_configured(entry, provider).await;
|
||||
providers.push(PluginProviderDescriptor {
|
||||
id: provider.id.clone(),
|
||||
plugin_id: plugin_id.clone(),
|
||||
display_name: provider.display_name.clone(),
|
||||
description: provider.description.clone(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
resource_type: provider.resource_type.clone(),
|
||||
has_models: provider.has_models,
|
||||
configured,
|
||||
models: stored
|
||||
.iter()
|
||||
.map(|model| {
|
||||
PluginModelDescriptor::new(
|
||||
plugin_id,
|
||||
&entry.manifest.name,
|
||||
&entry.icon,
|
||||
provider,
|
||||
model,
|
||||
)
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
}
|
||||
let mut resources = Vec::new();
|
||||
for definition in &entry.definition.resources {
|
||||
let records = self
|
||||
.inner
|
||||
.state
|
||||
.resources(plugin_id, &definition.resource_type)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let views = self
|
||||
.present_resources(entry, executable, definition, &records)
|
||||
.await;
|
||||
resources.push(PluginResourceDescriptor {
|
||||
resource_type: definition.resource_type.clone(),
|
||||
display_name: definition.display_name.clone(),
|
||||
add: definition.add.clone(),
|
||||
import: definition.import.clone(),
|
||||
can_refresh: definition.can_refresh,
|
||||
can_remove: definition.can_remove,
|
||||
resources: views,
|
||||
});
|
||||
}
|
||||
PluginDescriptor {
|
||||
id: plugin_id.clone(),
|
||||
name: entry.manifest.name.clone(),
|
||||
version: entry.manifest.version.clone(),
|
||||
author: entry.manifest.author.clone(),
|
||||
icon: entry.icon.clone(),
|
||||
providers,
|
||||
resources,
|
||||
}
|
||||
}
|
||||
|
||||
async fn present_resources(
|
||||
&self,
|
||||
entry: &PluginEntry,
|
||||
executable: &Path,
|
||||
definition: &ResourceDefinition,
|
||||
records: &[ResourceRecord],
|
||||
) -> Vec<PluginResourceView> {
|
||||
if records.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
let snapshots = records
|
||||
.iter()
|
||||
.map(|record| record.snapshot(&definition.resource_type))
|
||||
.collect::<Vec<_>>();
|
||||
let presented = self
|
||||
.worker(entry, executable)
|
||||
.await
|
||||
.invoke(
|
||||
"resource.present",
|
||||
serde_json::json!({
|
||||
"resourceType": definition.resource_type,
|
||||
"resources": snapshots,
|
||||
}),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.and_then(|value| {
|
||||
serde_json::from_value::<Vec<ResourcePresentation>>(value).map_err(Error::from)
|
||||
});
|
||||
match presented {
|
||||
Ok(views) if views.len() == records.len() => records
|
||||
.iter()
|
||||
.zip(views)
|
||||
.map(|(record, view)| PluginResourceView::from_record(record, view))
|
||||
.collect(),
|
||||
Ok(_) | Err(_) => records
|
||||
.iter()
|
||||
.map(|record| {
|
||||
PluginResourceView::from_record(
|
||||
record,
|
||||
ResourcePresentation {
|
||||
display_name: record.key.clone(),
|
||||
description: serde_json::Value::Null,
|
||||
metrics: Vec::new(),
|
||||
},
|
||||
)
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn provider_configured(
|
||||
&self,
|
||||
entry: &PluginEntry,
|
||||
provider: &ProviderDefinition,
|
||||
) -> bool {
|
||||
let plugin_id = &entry.manifest.id;
|
||||
if provider.has_models {
|
||||
let models = self
|
||||
.inner
|
||||
.state
|
||||
.models(plugin_id, &provider.id)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
if models.is_empty() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
match &provider.resource_type {
|
||||
Some(resource_type) => !self
|
||||
.inner
|
||||
.state
|
||||
.resources(plugin_id, resource_type)
|
||||
.await
|
||||
.unwrap_or_default()
|
||||
.is_empty(),
|
||||
None => true,
|
||||
}
|
||||
}
|
||||
|
||||
/// 资源到位后刷新使用该资源类型的 Provider 模型目录;失败只报告不中断。
|
||||
async fn sync_provider_models_for_resource(
|
||||
&self,
|
||||
entry: &PluginEntry,
|
||||
executable: &Path,
|
||||
resource_type: &str,
|
||||
) -> Option<String> {
|
||||
let mut errors = Vec::new();
|
||||
for provider in entry.definition.providers.clone() {
|
||||
if provider.resource_type.as_deref() != Some(resource_type) || !provider.has_models {
|
||||
continue;
|
||||
}
|
||||
if let Err(error) = self
|
||||
.sync_provider_models(entry, executable, &provider)
|
||||
.await
|
||||
{
|
||||
errors.push(format!("{}: {error}", provider.id));
|
||||
}
|
||||
}
|
||||
(!errors.is_empty()).then(|| errors.join("; "))
|
||||
}
|
||||
|
||||
async fn sync_provider_models(
|
||||
&self,
|
||||
entry: &PluginEntry,
|
||||
executable: &Path,
|
||||
provider: &ProviderDefinition,
|
||||
) -> Result<usize> {
|
||||
if !provider.has_models {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin provider '{}' does not enumerate models",
|
||||
provider.id
|
||||
)));
|
||||
}
|
||||
let plugin_id = &entry.manifest.id;
|
||||
let resource = match &provider.resource_type {
|
||||
Some(resource_type) => {
|
||||
let record = self.select_resource(plugin_id, resource_type).await?;
|
||||
Some(record.snapshot(resource_type))
|
||||
}
|
||||
None => None,
|
||||
};
|
||||
let value = self
|
||||
.worker(entry, executable)
|
||||
.await
|
||||
.invoke(
|
||||
"models.list",
|
||||
serde_json::json!({ "providerId": provider.id, "resource": resource }),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await?;
|
||||
let definitions = value
|
||||
.as_array()
|
||||
.ok_or_else(|| Error::Protocol("plugin models.list must return an array".into()))?;
|
||||
let mut models = Vec::with_capacity(definitions.len());
|
||||
let mut seen = std::collections::HashSet::new();
|
||||
for definition in definitions {
|
||||
let model = StoredModel::from_definition(definition)?;
|
||||
if seen.insert(model.id.clone()) {
|
||||
models.push(model);
|
||||
}
|
||||
}
|
||||
if models.is_empty() {
|
||||
return Err(Error::Provider(format!(
|
||||
"plugin provider '{}' returned no models",
|
||||
provider.id
|
||||
)));
|
||||
}
|
||||
self.inner
|
||||
.state
|
||||
.replace_models(plugin_id, &provider.id, &models)
|
||||
.await?;
|
||||
Ok(models.len())
|
||||
}
|
||||
|
||||
/// 第一版选择策略:按创建顺序取首个可用资源;冷却到期视为可用。
|
||||
async fn select_resource(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
) -> Result<ResourceRecord> {
|
||||
let records = self.inner.state.resources(plugin_id, resource_type).await?;
|
||||
if records.is_empty() {
|
||||
return Err(Error::Provider(format!(
|
||||
"plugin '{plugin_id}' has no '{resource_type}' resource; add one first"
|
||||
)));
|
||||
}
|
||||
let now = now_ms();
|
||||
records
|
||||
.iter()
|
||||
.find(|record| record.state.is_ready(now))
|
||||
.or_else(|| records.first())
|
||||
.cloned()
|
||||
.ok_or_else(|| Error::Provider("no plugin resource is available".into()))
|
||||
}
|
||||
|
||||
async fn find_record(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<ResourceRecord> {
|
||||
self.inner
|
||||
.state
|
||||
.resources(plugin_id, resource_type)
|
||||
.await?
|
||||
.into_iter()
|
||||
.find(|record| record.id == resource_id)
|
||||
.ok_or_else(|| Error::RunNotFound(format!("plugin resource {resource_id}")))
|
||||
}
|
||||
|
||||
async fn update_session(
|
||||
&self,
|
||||
session_id: &str,
|
||||
session: Option<serde_json::Value>,
|
||||
poll_interval_ms: Option<i64>,
|
||||
) {
|
||||
let mut sessions = self.inner.oauth_sessions.lock().await;
|
||||
if let Some(state) = sessions.get_mut(session_id) {
|
||||
if let Some(session) = session {
|
||||
state.session = session;
|
||||
}
|
||||
if let Some(interval) = poll_interval_ms {
|
||||
state.poll_interval_ms = interval;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn executable(&self) -> Result<std::path::PathBuf> {
|
||||
self.inner
|
||||
.runtime
|
||||
.executable()
|
||||
.ok_or_else(|| Error::Config("plugin runtime is not ready".into()))
|
||||
}
|
||||
|
||||
async fn entries(&self, executable: &Path) -> Vec<PluginEntry> {
|
||||
if let Some(entries) = self.inner.entries.read().await.as_ref() {
|
||||
return entries.clone();
|
||||
}
|
||||
let loaded = self.inner.catalog.entries(executable).await;
|
||||
*self.inner.entries.write().await = Some(loaded.clone());
|
||||
loaded
|
||||
}
|
||||
|
||||
async fn find_entry(&self, executable: &Path, plugin_id: &str) -> Result<PluginEntry> {
|
||||
self.entries(executable)
|
||||
.await
|
||||
.into_iter()
|
||||
.find(|entry| entry.manifest.id == plugin_id)
|
||||
.ok_or_else(|| Error::RunNotFound(format!("plugin {plugin_id}")))
|
||||
}
|
||||
|
||||
async fn worker(&self, entry: &PluginEntry, executable: &Path) -> Arc<PluginWorker> {
|
||||
let mut workers = self.inner.workers.lock().await;
|
||||
workers
|
||||
.entry(entry.manifest.id.clone())
|
||||
.or_insert_with(|| {
|
||||
Arc::new(PluginWorker::new(
|
||||
entry,
|
||||
executable.to_path_buf(),
|
||||
self.inner.catalog.loader().clone(),
|
||||
self.inner.store.clone(),
|
||||
))
|
||||
})
|
||||
.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, serde::Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
struct OAuth2Begin {
|
||||
session: serde_json::Value,
|
||||
user_code: String,
|
||||
verification_url: String,
|
||||
#[serde(default)]
|
||||
verification_url_complete: Option<String>,
|
||||
expires_at_ms: i64,
|
||||
poll_interval_ms: i64,
|
||||
}
|
||||
|
||||
#[derive(Debug, serde::Deserialize)]
|
||||
#[serde(rename_all = "kebab-case", tag = "status")]
|
||||
enum OAuth2Poll {
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Pending {
|
||||
#[serde(default)]
|
||||
session: Option<serde_json::Value>,
|
||||
},
|
||||
#[serde(rename_all = "camelCase")]
|
||||
SlowDown {
|
||||
#[serde(default)]
|
||||
session: Option<serde_json::Value>,
|
||||
},
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Completed { resources: Vec<ResourceDraft> },
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Denied {
|
||||
#[serde(default)]
|
||||
message: Option<String>,
|
||||
},
|
||||
#[serde(rename_all = "camelCase")]
|
||||
Failed { message: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, serde::Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
struct ImportParseResult {
|
||||
resources: Vec<ResourceDraft>,
|
||||
#[serde(default)]
|
||||
warnings: Vec<String>,
|
||||
}
|
||||
|
||||
fn find_provider<'a>(entry: &'a PluginEntry, provider_id: &str) -> Result<&'a ProviderDefinition> {
|
||||
entry
|
||||
.definition
|
||||
.providers
|
||||
.iter()
|
||||
.find(|provider| provider.id == provider_id)
|
||||
.ok_or_else(|| {
|
||||
Error::RunNotFound(format!(
|
||||
"plugin '{}' provider {provider_id}",
|
||||
entry.manifest.id
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn find_resource<'a>(
|
||||
entry: &'a PluginEntry,
|
||||
resource_type: &str,
|
||||
) -> Result<&'a ResourceDefinition> {
|
||||
entry
|
||||
.definition
|
||||
.resources
|
||||
.iter()
|
||||
.find(|resource| resource.resource_type == resource_type)
|
||||
.ok_or_else(|| {
|
||||
Error::RunNotFound(format!(
|
||||
"plugin '{}' resource type {resource_type}",
|
||||
entry.manifest.id
|
||||
))
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
//! Tracks Deno runtime readiness and coordinates one initialization task.
|
||||
use std::{
|
||||
path::PathBuf,
|
||||
sync::{
|
||||
atomic::{AtomicBool, Ordering},
|
||||
Arc,
|
||||
},
|
||||
};
|
||||
|
||||
use parking_lot::{Mutex, RwLock};
|
||||
use serde::Serialize;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use super::{
|
||||
asset::{RuntimeAsset, DENO_VERSION},
|
||||
installation,
|
||||
};
|
||||
use crate::{config, store::Store, Error, Result};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PluginRuntime {
|
||||
inner: Arc<PluginRuntimeInner>,
|
||||
}
|
||||
|
||||
struct PluginRuntimeInner {
|
||||
root: PathBuf,
|
||||
asset: Option<RuntimeAsset>,
|
||||
status: RwLock<PluginRuntimeStatus>,
|
||||
initializing: AtomicBool,
|
||||
cancellation: Mutex<Option<CancellationToken>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum PluginRuntimeState {
|
||||
Uninitialized,
|
||||
Initializing,
|
||||
Ready,
|
||||
Failed,
|
||||
Unsupported,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum PluginRuntimePhase {
|
||||
Checking,
|
||||
Downloading,
|
||||
Verifying,
|
||||
Installing,
|
||||
Validating,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, PartialEq, Eq)]
|
||||
pub struct PluginRuntimeStatus {
|
||||
pub state: PluginRuntimeState,
|
||||
pub version: String,
|
||||
pub target: Option<String>,
|
||||
pub phase: Option<PluginRuntimePhase>,
|
||||
pub downloaded_bytes: u64,
|
||||
pub total_bytes: Option<u64>,
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
impl PluginRuntimeStatus {
|
||||
fn uninitialized(asset: RuntimeAsset) -> Self {
|
||||
Self::new(PluginRuntimeState::Uninitialized, Some(asset))
|
||||
}
|
||||
|
||||
fn ready(asset: RuntimeAsset) -> Self {
|
||||
Self::new(PluginRuntimeState::Ready, Some(asset))
|
||||
}
|
||||
|
||||
fn unsupported() -> Self {
|
||||
Self {
|
||||
state: PluginRuntimeState::Unsupported,
|
||||
version: DENO_VERSION.into(),
|
||||
target: None,
|
||||
phase: None,
|
||||
downloaded_bytes: 0,
|
||||
total_bytes: None,
|
||||
error: Some(format!(
|
||||
"unsupported platform: {}/{}",
|
||||
std::env::consts::OS,
|
||||
std::env::consts::ARCH
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn new(state: PluginRuntimeState, asset: Option<RuntimeAsset>) -> Self {
|
||||
Self {
|
||||
state,
|
||||
version: DENO_VERSION.into(),
|
||||
target: asset.map(|value| value.target.into()),
|
||||
phase: None,
|
||||
downloaded_bytes: 0,
|
||||
total_bytes: None,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PluginRuntime {
|
||||
pub fn managed() -> Result<Self> {
|
||||
Self::new(config::managed_data_dir()?.join("plugins").join("runtime"))
|
||||
}
|
||||
|
||||
fn new(root: PathBuf) -> Result<Self> {
|
||||
std::fs::create_dir_all(&root)?;
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o700))?;
|
||||
}
|
||||
let asset = RuntimeAsset::current();
|
||||
let status = match asset {
|
||||
Some(asset) if installation::runtime_complete(&root, asset) => {
|
||||
PluginRuntimeStatus::ready(asset)
|
||||
}
|
||||
Some(asset) => PluginRuntimeStatus::uninitialized(asset),
|
||||
None => PluginRuntimeStatus::unsupported(),
|
||||
};
|
||||
Ok(Self {
|
||||
inner: Arc::new(PluginRuntimeInner {
|
||||
root,
|
||||
asset,
|
||||
status: RwLock::new(status),
|
||||
initializing: AtomicBool::new(false),
|
||||
cancellation: Mutex::new(None),
|
||||
}),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn status(&self) -> PluginRuntimeStatus {
|
||||
let mut status = self.inner.status.write();
|
||||
if status.state == PluginRuntimeState::Ready {
|
||||
if let Some(asset) = self.inner.asset {
|
||||
if !installation::runtime_complete(&self.inner.root, asset) {
|
||||
*status = PluginRuntimeStatus::uninitialized(asset);
|
||||
}
|
||||
}
|
||||
}
|
||||
status.clone()
|
||||
}
|
||||
|
||||
pub fn executable(&self) -> Option<PathBuf> {
|
||||
let asset = self.inner.asset?;
|
||||
if self.status().state != PluginRuntimeState::Ready {
|
||||
return None;
|
||||
}
|
||||
Some(installation::runtime_executable(&self.inner.root, asset))
|
||||
}
|
||||
|
||||
pub fn initialize(&self, store: Store) -> PluginRuntimeStatus {
|
||||
let Some(asset) = self.inner.asset else {
|
||||
return self.status();
|
||||
};
|
||||
if installation::runtime_complete(&self.inner.root, asset) {
|
||||
let ready = PluginRuntimeStatus::ready(asset);
|
||||
*self.inner.status.write() = ready.clone();
|
||||
return ready;
|
||||
}
|
||||
if self
|
||||
.inner
|
||||
.initializing
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||
.is_err()
|
||||
{
|
||||
return self.status();
|
||||
}
|
||||
|
||||
let mut initializing =
|
||||
PluginRuntimeStatus::new(PluginRuntimeState::Initializing, Some(asset));
|
||||
initializing.phase = Some(PluginRuntimePhase::Checking);
|
||||
*self.inner.status.write() = initializing.clone();
|
||||
|
||||
let cancellation = CancellationToken::new();
|
||||
*self.inner.cancellation.lock() = Some(cancellation.clone());
|
||||
let runtime = self.clone();
|
||||
tokio::spawn(async move {
|
||||
let result = installation::install(
|
||||
&runtime.inner.root,
|
||||
&store,
|
||||
asset,
|
||||
cancellation,
|
||||
|phase, downloaded, total| {
|
||||
runtime.update_progress(phase, downloaded, total);
|
||||
},
|
||||
)
|
||||
.await;
|
||||
let status = match result {
|
||||
Ok(()) => PluginRuntimeStatus::ready(asset),
|
||||
Err(Error::Cancelled) => PluginRuntimeStatus::uninitialized(asset),
|
||||
Err(error) => {
|
||||
tracing::error!(%error, target = asset.target, "plugin runtime initialization failed");
|
||||
let mut failed =
|
||||
PluginRuntimeStatus::new(PluginRuntimeState::Failed, Some(asset));
|
||||
failed.error = Some("plugin runtime initialization failed".into());
|
||||
failed
|
||||
}
|
||||
};
|
||||
*runtime.inner.status.write() = status;
|
||||
runtime.inner.cancellation.lock().take();
|
||||
runtime.inner.initializing.store(false, Ordering::Release);
|
||||
});
|
||||
|
||||
initializing
|
||||
}
|
||||
|
||||
pub fn cancel_initialization(&self) -> PluginRuntimeStatus {
|
||||
if let Some(cancellation) = self.inner.cancellation.lock().as_ref() {
|
||||
cancellation.cancel();
|
||||
}
|
||||
self.status()
|
||||
}
|
||||
|
||||
fn update_progress(
|
||||
&self,
|
||||
phase: PluginRuntimePhase,
|
||||
downloaded_bytes: u64,
|
||||
total_bytes: Option<u64>,
|
||||
) {
|
||||
let mut status = self.inner.status.write();
|
||||
status.state = PluginRuntimeState::Initializing;
|
||||
status.phase = Some(phase);
|
||||
status.downloaded_bytes = downloaded_bytes;
|
||||
status.total_bytes = total_bytes;
|
||||
status.error = None;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
import { __descriptor, __getRegisteredPlugin } from "cursor-byok:plugin";
|
||||
|
||||
if (Deno.args.length !== 1) throw new Error("plugin entry URL is required");
|
||||
await import(Deno.args[0]);
|
||||
console.log("CURSOR_BYOK_PLUGIN_DEFINITION:" + JSON.stringify(__descriptor(__getRegisteredPlugin())));
|
||||
@@ -0,0 +1,5 @@
|
||||
{
|
||||
"fmt": {
|
||||
"lineWidth": 200
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
{
|
||||
"imports": {
|
||||
"cursor-byok:plugin": "./plugin.ts",
|
||||
"cursor-byok:provider": "./provider.ts",
|
||||
"cursor-byok:model": "./model.ts",
|
||||
"cursor-byok:resource": "./resource.ts",
|
||||
"cursor-byok:protocol/openai-responses": "./protocol/openai_responses.ts",
|
||||
"cursor-byok:protocol/openai-chat": "./protocol/openai_chat.ts"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
import type { JsonValue, PluginContext } from "./plugin.ts";
|
||||
import type { ResourceSnapshot } from "./resource.ts";
|
||||
|
||||
export type ModelCapabilities = {
|
||||
images?: boolean;
|
||||
};
|
||||
|
||||
export type ModelDefinition = {
|
||||
id: string;
|
||||
displayName: string;
|
||||
description?: string;
|
||||
maxOutputTokens?: number;
|
||||
capabilities?: ModelCapabilities;
|
||||
/** 之后的调用原样传回;永远不会展示给用户。 */
|
||||
privateData?: JsonValue;
|
||||
};
|
||||
|
||||
/** 宿主目录中持久化的一条模型。 */
|
||||
export type ModelSnapshot = ModelDefinition;
|
||||
|
||||
export type ModelListInput = {
|
||||
/** 模型发现需要认证时为首个可用资源,否则为 null。 */
|
||||
resource: ResourceSnapshot | null;
|
||||
};
|
||||
|
||||
export type ModelSupport = {
|
||||
/** 列举成功后,宿主用返回值整体替换该 Provider 的模型目录。 */
|
||||
list(input: ModelListInput, context: PluginContext): Promise<ModelDefinition[]>;
|
||||
};
|
||||
@@ -0,0 +1,100 @@
|
||||
import type { ProviderSupport } from "./provider.ts";
|
||||
import type { ResourceSupport } from "./resource.ts";
|
||||
|
||||
export type JsonPrimitive = string | number | boolean | null;
|
||||
export type JsonValue = JsonPrimitive | JsonValue[] | { [key: string]: JsonValue };
|
||||
|
||||
/**
|
||||
* 可本地化文本:纯字符串,或 locale → 文本 的映射
|
||||
* (如 { "zh-CN": "账号", "en-US": "Accounts" })。
|
||||
* 宿主原样透传,由界面按当前语言解析;模型名等来自上游的数据保持纯字符串。
|
||||
*/
|
||||
export type LocalizedText = string | { [locale: string]: string };
|
||||
|
||||
export type NetworkRequestInit = {
|
||||
method?: string;
|
||||
headers?: Record<string, string>;
|
||||
body?: string;
|
||||
};
|
||||
|
||||
export type NetworkResponse = {
|
||||
status: number;
|
||||
headers: Record<string, string>;
|
||||
body: string;
|
||||
};
|
||||
|
||||
/** 流式响应体,按行随到随交付(用于 SSE)。 */
|
||||
export type NetworkEventStream = {
|
||||
status: number;
|
||||
headers: Record<string, string>;
|
||||
lines: AsyncIterable<string>;
|
||||
};
|
||||
|
||||
/**
|
||||
* 每次能力调用收到的宿主服务。网络请求仅限 plugin.json 声明的 HTTPS 主机;
|
||||
* 宿主取消本次调用时通过 `signal` 中止。
|
||||
*/
|
||||
export type PluginContext = {
|
||||
network: {
|
||||
fetch(url: string, init?: NetworkRequestInit): Promise<NetworkResponse>;
|
||||
stream(url: string, init?: NetworkRequestInit): Promise<NetworkEventStream>;
|
||||
};
|
||||
signal: AbortSignal;
|
||||
};
|
||||
|
||||
/**
|
||||
* Provider 插件定义:一组能力实现的集合。插件不持有任何持久状态——
|
||||
* 资源与模型目录由宿主存储,每次调用所需的数据都通过参数传入。
|
||||
*/
|
||||
export type ProviderPluginDefinition = {
|
||||
providers: ProviderSupport[];
|
||||
resources?: ResourceSupport[];
|
||||
};
|
||||
|
||||
let registered: ProviderPluginDefinition | undefined;
|
||||
|
||||
/** 注册 Provider 插件;每个插件入口只能调用一次。 */
|
||||
export function defineProviderPlugin(definition: ProviderPluginDefinition): ProviderPluginDefinition {
|
||||
if (registered) throw new Error("defineProviderPlugin can only be called once");
|
||||
registered = definition;
|
||||
return definition;
|
||||
}
|
||||
|
||||
export function __getRegisteredPlugin(): ProviderPluginDefinition {
|
||||
if (!registered) throw new Error("plugin entry must call defineProviderPlugin");
|
||||
return registered;
|
||||
}
|
||||
|
||||
/** 可序列化的能力摘要,宿主收集它时不调用任何能力方法。 */
|
||||
export function __descriptor(definition: ProviderPluginDefinition) {
|
||||
return {
|
||||
providers: definition.providers.map((provider) => ({
|
||||
id: provider.id,
|
||||
displayName: provider.displayName,
|
||||
description: provider.description ?? null,
|
||||
providerType: provider.providerType,
|
||||
resourceType: provider.resourceType ?? null,
|
||||
hasModels: provider.models !== undefined,
|
||||
})),
|
||||
resources: (definition.resources ?? []).map((resource) => ({
|
||||
type: resource.type,
|
||||
displayName: resource.displayName,
|
||||
add: (resource.add ?? []).map((method) => ({
|
||||
type: method.type,
|
||||
id: method.id,
|
||||
displayName: method.displayName,
|
||||
description: method.description ?? null,
|
||||
})),
|
||||
import: resource.import
|
||||
? {
|
||||
displayName: resource.import.displayName,
|
||||
description: resource.import.description ?? null,
|
||||
accept: resource.import.accept,
|
||||
multiple: resource.import.multiple ?? false,
|
||||
}
|
||||
: null,
|
||||
canRefresh: resource.refresh !== undefined,
|
||||
canRemove: resource.remove !== undefined,
|
||||
})),
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
import type { JsonValue, PluginContext } from "../plugin.ts";
|
||||
import type { LlmContentPart, LlmRequest, ModelEvent, ProviderOutput } from "../provider.ts";
|
||||
|
||||
/** 本协议产生的回放状态种类;与宿主内置 Chat Provider 一致,可互相回放。 */
|
||||
export const REPLAY_KIND = "openai_chat";
|
||||
|
||||
/** 上游返回非 2xx 时抛出,携带完整响应体供调用方分类。 */
|
||||
export class HttpError extends Error {
|
||||
constructor(readonly status: number, readonly body: string) {
|
||||
super(`HTTP ${status}: ${body}`);
|
||||
}
|
||||
}
|
||||
|
||||
export type OpenAiChatCall = {
|
||||
url: string;
|
||||
model: string;
|
||||
request: LlmRequest;
|
||||
headers?: Record<string, string>;
|
||||
/** 最后合并进请求体。 */
|
||||
extraBody?: Record<string, JsonValue>;
|
||||
};
|
||||
|
||||
function record(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value)
|
||||
? value as Record<string, unknown>
|
||||
: null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" ? value : null;
|
||||
}
|
||||
|
||||
function count(value: unknown): number | null {
|
||||
return typeof value === "number" && Number.isFinite(value) ? value : null;
|
||||
}
|
||||
|
||||
/** 纯文本消息保持字符串形式;混合图片时展开为分块数组。 */
|
||||
function chatContent(parts: LlmContentPart[]): JsonValue {
|
||||
if (parts.every((part) => part.type === "text")) {
|
||||
return parts.map((part) => part.type === "text" ? part.text : "").join("");
|
||||
}
|
||||
return parts.map((part): JsonValue =>
|
||||
part.type === "text" ? { type: "text", text: part.text } : {
|
||||
type: "image_url",
|
||||
image_url: { url: `data:${part.mediaType};base64,${part.dataBase64}` },
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
export function buildChatBody(call: OpenAiChatCall): Record<string, JsonValue> {
|
||||
const messages: JsonValue[] = [];
|
||||
if (call.request.instructions) {
|
||||
messages.push({ role: "system", content: call.request.instructions });
|
||||
}
|
||||
for (const message of call.request.messages) {
|
||||
if (message.role === "assistant") {
|
||||
const reasoning = message.replayState?.providerKind === REPLAY_KIND
|
||||
? text(record(message.replayState.value)?.reasoning_content)
|
||||
: null;
|
||||
// Chat Completions 拒绝空字符串的 assistant content;完全无可见
|
||||
// 内容的 assistant 消息不需要发送。
|
||||
if (!message.text && message.toolCalls.length === 0 && !reasoning) continue;
|
||||
const value: Record<string, JsonValue> = {
|
||||
role: "assistant",
|
||||
content: message.text ? message.text : null,
|
||||
};
|
||||
if (reasoning) value.reasoning_content = reasoning;
|
||||
if (message.toolCalls.length > 0) {
|
||||
value.tool_calls = message.toolCalls.map((toolCall) => ({
|
||||
id: toolCall.callId,
|
||||
type: "function",
|
||||
function: { name: toolCall.name, arguments: JSON.stringify(toolCall.arguments) },
|
||||
}));
|
||||
}
|
||||
messages.push(value);
|
||||
} else if (message.role === "tool") {
|
||||
messages.push({
|
||||
role: "tool",
|
||||
content: message.parts.length === 0 ? message.content : chatContent(message.parts),
|
||||
tool_call_id: message.callId,
|
||||
});
|
||||
} else {
|
||||
messages.push({ role: message.role, content: chatContent(message.content) });
|
||||
}
|
||||
}
|
||||
const body: Record<string, JsonValue> = {
|
||||
model: call.model,
|
||||
messages,
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
};
|
||||
if (call.request.tools.length > 0) {
|
||||
body.tools = call.request.tools.map((tool) => ({
|
||||
type: "function",
|
||||
function: { name: tool.name, description: tool.description, parameters: tool.parameters },
|
||||
}));
|
||||
}
|
||||
if (call.request.maxOutputTokens !== null) {
|
||||
body.max_completion_tokens = call.request.maxOutputTokens;
|
||||
}
|
||||
if (call.request.reasoning.effort !== null) {
|
||||
body.reasoning_effort = call.request.reasoning.effort;
|
||||
}
|
||||
if (call.request.latency === "fast") body.service_tier = "fast";
|
||||
if (call.request.cacheKey !== null) body.prompt_cache_key = call.request.cacheKey;
|
||||
return { ...body, ...call.extraBody };
|
||||
}
|
||||
|
||||
type ToolState = {
|
||||
callId: string;
|
||||
name: string;
|
||||
arguments: string;
|
||||
emitted: number;
|
||||
started: boolean;
|
||||
};
|
||||
|
||||
/** 部分上游会重发完整片段而不是增量;去重后再拼接。 */
|
||||
function mergeFragment(target: string, fragment: string): string {
|
||||
if (target === fragment || target.endsWith(fragment)) return target;
|
||||
if (fragment.startsWith(target)) return fragment;
|
||||
return target + fragment;
|
||||
}
|
||||
|
||||
function updateTool(
|
||||
index: number,
|
||||
callId: string | null,
|
||||
name: string | null,
|
||||
argumentsDelta: string | null,
|
||||
tools: Map<number, ToolState>,
|
||||
): ModelEvent[] {
|
||||
let tool = tools.get(index);
|
||||
if (!tool) {
|
||||
tool = { callId: "", name: "", arguments: "", emitted: 0, started: false };
|
||||
tools.set(index, tool);
|
||||
}
|
||||
if (callId !== null) tool.callId = mergeFragment(tool.callId, callId);
|
||||
if (name !== null) tool.name = mergeFragment(tool.name, name);
|
||||
if (argumentsDelta !== null) tool.arguments += argumentsDelta;
|
||||
|
||||
const events: ModelEvent[] = [];
|
||||
if (!tool.started && tool.callId && tool.name) {
|
||||
tool.started = true;
|
||||
events.push({ type: "tool-call-start", index, callId: tool.callId, name: tool.name });
|
||||
}
|
||||
if (tool.started && tool.emitted < tool.arguments.length) {
|
||||
events.push({
|
||||
type: "tool-call-arguments-delta",
|
||||
index,
|
||||
delta: tool.arguments.slice(tool.emitted),
|
||||
});
|
||||
tool.emitted = tool.arguments.length;
|
||||
}
|
||||
return events;
|
||||
}
|
||||
|
||||
function eventError(value: Record<string, unknown>): string | null {
|
||||
const error = value.error;
|
||||
if (error === undefined || error === null) return null;
|
||||
if (typeof error === "string") return error;
|
||||
return text(record(error)?.message) ?? JSON.stringify(error);
|
||||
}
|
||||
|
||||
function usageEvent(value: unknown): ModelEvent {
|
||||
const usage = record(value) ?? {};
|
||||
return {
|
||||
type: "usage",
|
||||
usage: {
|
||||
inputTokens: count(usage.prompt_tokens),
|
||||
outputTokens: count(usage.completion_tokens),
|
||||
totalTokens: count(usage.total_tokens),
|
||||
cacheReadTokens: count(record(usage.prompt_tokens_details)?.cached_tokens),
|
||||
cacheWriteTokens: null,
|
||||
reasoningTokens: count(record(usage.completion_tokens_details)?.reasoning_tokens),
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
function mapFinish(reason: string, hasTools: boolean): "stop" | "length" | "tool-use" {
|
||||
if (reason === "tool_calls" || reason === "function_call") return "tool-use";
|
||||
if (reason === "length") return "length";
|
||||
if (reason === "stop" || reason === "content_filter") return "stop";
|
||||
return hasTools ? "tool-use" : "stop";
|
||||
}
|
||||
|
||||
async function readBody(lines: AsyncIterable<string>): Promise<string> {
|
||||
const collected: string[] = [];
|
||||
for await (const line of lines) collected.push(line);
|
||||
return collected.join("\n");
|
||||
}
|
||||
|
||||
/**
|
||||
* 执行一次 Chat Completions 流式调用,发出与宿主统一事件集一致的标准化事件,
|
||||
* 包括文本/思考边界与工具参数增量。非 2xx 响应抛出 `HttpError`,
|
||||
* 流内失败抛出 `Error`,由调用方分类额度与授权问题。
|
||||
*/
|
||||
export async function streamOpenAiChat(
|
||||
call: OpenAiChatCall,
|
||||
output: ProviderOutput,
|
||||
context: PluginContext,
|
||||
): Promise<void> {
|
||||
const response = await context.network.stream(call.url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
accept: "text/event-stream",
|
||||
"content-type": "application/json",
|
||||
...call.headers,
|
||||
},
|
||||
body: JSON.stringify(buildChatBody(call)),
|
||||
});
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
throw new HttpError(response.status, await readBody(response.lines));
|
||||
}
|
||||
|
||||
let textOpen = false;
|
||||
let thinkingOpen = false;
|
||||
let reasoning = "";
|
||||
const tools = new Map<number, ToolState>();
|
||||
let finalUsage: ModelEvent | null = null;
|
||||
let finish: "stop" | "length" | "tool-use" | null = null;
|
||||
let sawDoneMarker = false;
|
||||
|
||||
for await (const line of response.lines) {
|
||||
if (!line.startsWith("data:")) continue;
|
||||
const payload = line.slice(5).trim();
|
||||
if (!payload) continue;
|
||||
if (payload === "[DONE]") {
|
||||
sawDoneMarker = true;
|
||||
break;
|
||||
}
|
||||
let value: Record<string, unknown>;
|
||||
try {
|
||||
value = record(JSON.parse(payload)) ?? {};
|
||||
} catch {
|
||||
throw new Error("OpenAI Chat SSE returned invalid JSON");
|
||||
}
|
||||
const error = eventError(value);
|
||||
if (error !== null) throw new Error(`OpenAI Chat error: ${error}`);
|
||||
if (value.usage !== undefined && value.usage !== null) {
|
||||
finalUsage = usageEvent(value.usage);
|
||||
}
|
||||
const choice = Array.isArray(value.choices) ? record(value.choices[0]) : null;
|
||||
if (!choice) continue;
|
||||
const delta = record(choice.delta) ?? {};
|
||||
const reasoningDelta = text(delta.reasoning_content) ?? text(delta.reasoning);
|
||||
if (reasoningDelta) {
|
||||
if (!thinkingOpen) {
|
||||
thinkingOpen = true;
|
||||
output.emit({ type: "thinking-start" });
|
||||
}
|
||||
reasoning += reasoningDelta;
|
||||
output.emit({ type: "thinking-delta", text: reasoningDelta });
|
||||
}
|
||||
const content = text(delta.content);
|
||||
if (content) {
|
||||
if (thinkingOpen) {
|
||||
thinkingOpen = false;
|
||||
output.emit({ type: "thinking-end" });
|
||||
}
|
||||
if (!textOpen) {
|
||||
textOpen = true;
|
||||
output.emit({ type: "text-start" });
|
||||
}
|
||||
output.emit({ type: "text-delta", text: content });
|
||||
}
|
||||
if (Array.isArray(delta.tool_calls)) {
|
||||
for (const [position, rawTool] of delta.tool_calls.entries()) {
|
||||
const toolDelta = record(rawTool);
|
||||
if (!toolDelta) continue;
|
||||
const index = count(toolDelta.index) ?? position;
|
||||
const fn = record(toolDelta.function) ?? {};
|
||||
for (
|
||||
const event of updateTool(
|
||||
index,
|
||||
text(toolDelta.id),
|
||||
text(fn.name),
|
||||
text(fn.arguments),
|
||||
tools,
|
||||
)
|
||||
) {
|
||||
output.emit(event);
|
||||
}
|
||||
}
|
||||
}
|
||||
const finishReason = text(choice.finish_reason);
|
||||
if (finishReason !== null) finish = mapFinish(finishReason, tools.size > 0);
|
||||
}
|
||||
|
||||
if (thinkingOpen) output.emit({ type: "thinking-end" });
|
||||
if (textOpen) output.emit({ type: "text-end" });
|
||||
for (const [index, tool] of tools) {
|
||||
if (!tool.started) {
|
||||
if (!tool.name) throw new Error("OpenAI Chat tool call is missing name");
|
||||
if (!tool.callId) tool.callId = `call-${index}`;
|
||||
tool.started = true;
|
||||
output.emit({ type: "tool-call-start", index, callId: tool.callId, name: tool.name });
|
||||
if (tool.arguments) {
|
||||
tool.emitted = tool.arguments.length;
|
||||
output.emit({ type: "tool-call-arguments-delta", index, delta: tool.arguments });
|
||||
}
|
||||
}
|
||||
output.emit({ type: "tool-call-end", index });
|
||||
}
|
||||
if (finalUsage !== null) output.emit(finalUsage);
|
||||
if (reasoning) {
|
||||
output.emit({
|
||||
type: "replay-state",
|
||||
providerKind: REPLAY_KIND,
|
||||
value: { reasoning_content: reasoning },
|
||||
});
|
||||
}
|
||||
const reason = finish ??
|
||||
(sawDoneMarker ? (tools.size > 0 ? "tool-use" : "stop") : null);
|
||||
if (reason === null) {
|
||||
throw new Error("OpenAI Chat stream ended without finish_reason");
|
||||
}
|
||||
output.emit({ type: "done", reason });
|
||||
}
|
||||
@@ -0,0 +1,425 @@
|
||||
import type { JsonValue, PluginContext } from "../plugin.ts";
|
||||
import type { LlmContentPart, LlmRequest, ModelEvent, ProviderOutput } from "../provider.ts";
|
||||
|
||||
/** 本协议产生的回放状态种类;与宿主内置 Responses Provider 一致,可互相回放。 */
|
||||
export const REPLAY_KIND = "openai_responses";
|
||||
|
||||
/** 上游返回非 2xx 时抛出,携带完整响应体供调用方分类。 */
|
||||
export class HttpError extends Error {
|
||||
constructor(readonly status: number, readonly body: string) {
|
||||
super(`HTTP ${status}: ${body}`);
|
||||
}
|
||||
}
|
||||
|
||||
export type OpenAiResponsesCall = {
|
||||
url: string;
|
||||
model: string;
|
||||
request: LlmRequest;
|
||||
headers?: Record<string, string>;
|
||||
/** 最后合并进请求体,如 { store: false }。 */
|
||||
extraBody?: Record<string, JsonValue>;
|
||||
};
|
||||
|
||||
function record(value: unknown): Record<string, unknown> | null {
|
||||
return value !== null && typeof value === "object" && !Array.isArray(value) ? value as Record<string, unknown> : null;
|
||||
}
|
||||
|
||||
function text(value: unknown): string | null {
|
||||
return typeof value === "string" ? value : null;
|
||||
}
|
||||
|
||||
function count(value: unknown): number | null {
|
||||
return typeof value === "number" && Number.isFinite(value) ? value : null;
|
||||
}
|
||||
|
||||
function contentParts(parts: LlmContentPart[], textType: "input_text" | "output_text"): JsonValue[] {
|
||||
const content: JsonValue[] = [];
|
||||
for (const part of parts) {
|
||||
if (part.type === "text") {
|
||||
if (part.text) content.push({ type: textType, text: part.text });
|
||||
} else {
|
||||
content.push({
|
||||
type: "input_image",
|
||||
detail: "auto",
|
||||
image_url: `data:${part.mediaType};base64,${part.dataBase64}`,
|
||||
});
|
||||
}
|
||||
}
|
||||
return content;
|
||||
}
|
||||
|
||||
function replayItems(value: JsonValue): JsonValue[] {
|
||||
const items = record(value)?.items;
|
||||
if (!Array.isArray(items)) {
|
||||
throw new Error("OpenAI Responses replay state is missing items");
|
||||
}
|
||||
return items;
|
||||
}
|
||||
|
||||
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
|
||||
const input: JsonValue[] = [];
|
||||
for (const message of call.request.messages) {
|
||||
if (message.role === "assistant") {
|
||||
if (message.replayState?.providerKind === REPLAY_KIND) {
|
||||
input.push(...replayItems(message.replayState.value));
|
||||
}
|
||||
if (message.text) {
|
||||
input.push({
|
||||
type: "message",
|
||||
role: "assistant",
|
||||
content: [{ type: "output_text", text: message.text }],
|
||||
});
|
||||
}
|
||||
for (const toolCall of message.toolCalls) {
|
||||
input.push({
|
||||
type: "function_call",
|
||||
call_id: toolCall.callId,
|
||||
name: toolCall.name,
|
||||
arguments: JSON.stringify(toolCall.arguments),
|
||||
});
|
||||
}
|
||||
} else if (message.role === "tool") {
|
||||
input.push({
|
||||
type: "function_call_output",
|
||||
call_id: message.callId,
|
||||
output: message.parts.length === 0 ? message.content : contentParts(message.parts, "input_text"),
|
||||
});
|
||||
} else {
|
||||
const content = contentParts(message.content, "input_text");
|
||||
if (content.length > 0) input.push({ type: "message", role: message.role, content });
|
||||
}
|
||||
}
|
||||
const body: Record<string, JsonValue> = {
|
||||
model: call.model,
|
||||
input,
|
||||
stream: true,
|
||||
instructions: call.request.instructions,
|
||||
include: ["reasoning.encrypted_content"],
|
||||
};
|
||||
if (call.request.tools.length > 0) {
|
||||
body.tools = call.request.tools.map((tool) => ({
|
||||
type: "function",
|
||||
name: tool.name,
|
||||
description: tool.description,
|
||||
parameters: tool.parameters,
|
||||
strict: false,
|
||||
}));
|
||||
}
|
||||
if (call.request.maxOutputTokens !== null) body.max_output_tokens = call.request.maxOutputTokens;
|
||||
const reasoning = call.request.reasoning;
|
||||
if (reasoning.enabled || reasoning.effort !== null) {
|
||||
body.reasoning = {
|
||||
summary: "auto",
|
||||
...(reasoning.effort !== null ? { effort: reasoning.effort } : {}),
|
||||
};
|
||||
}
|
||||
// OpenAI 的规范 tier 值是 priority;"fast" 只是客户端别名,上游不接受。
|
||||
if (call.request.latency === "fast") body.service_tier = "priority";
|
||||
// 会话级缓存键把请求钉到同一缓存分片,前缀缓存才能稳定命中。
|
||||
if (call.request.cacheKey !== null) body.prompt_cache_key = call.request.cacheKey;
|
||||
return { ...body, ...call.extraBody };
|
||||
}
|
||||
|
||||
type ToolState = {
|
||||
callId: string | null;
|
||||
name: string | null;
|
||||
arguments: string;
|
||||
emitted: number;
|
||||
started: boolean;
|
||||
ended: boolean;
|
||||
};
|
||||
|
||||
type ToolArguments =
|
||||
| { kind: "none" }
|
||||
| { kind: "delta"; delta: string }
|
||||
| { kind: "snapshot"; snapshot: string };
|
||||
|
||||
function updateTool(
|
||||
index: number,
|
||||
item: Record<string, unknown> | null,
|
||||
args: ToolArguments,
|
||||
done: boolean,
|
||||
tools: Map<number, ToolState>,
|
||||
): ModelEvent[] {
|
||||
let tool = tools.get(index);
|
||||
if (!tool) {
|
||||
tool = { callId: null, name: null, arguments: "", emitted: 0, started: false, ended: false };
|
||||
tools.set(index, tool);
|
||||
}
|
||||
tool.callId ??= text(item?.call_id);
|
||||
tool.name ??= text(item?.name);
|
||||
if (args.kind === "delta") {
|
||||
tool.arguments += args.delta;
|
||||
} else if (args.kind === "snapshot" && args.snapshot !== tool.arguments) {
|
||||
if (!args.snapshot.startsWith(tool.arguments)) {
|
||||
throw new Error("OpenAI Responses final tool arguments do not match streamed arguments");
|
||||
}
|
||||
tool.arguments += args.snapshot.slice(tool.arguments.length);
|
||||
}
|
||||
|
||||
const events: ModelEvent[] = [];
|
||||
if (!tool.started && tool.callId !== null && tool.name !== null) {
|
||||
tool.started = true;
|
||||
events.push({ type: "tool-call-start", index, callId: tool.callId, name: tool.name });
|
||||
}
|
||||
if (tool.started && tool.emitted < tool.arguments.length) {
|
||||
events.push({ type: "tool-call-arguments-delta", index, delta: tool.arguments.slice(tool.emitted) });
|
||||
tool.emitted = tool.arguments.length;
|
||||
}
|
||||
if (done && !tool.ended) {
|
||||
if (!tool.started) {
|
||||
throw new Error("OpenAI Responses function call is missing call_id or name");
|
||||
}
|
||||
tool.ended = true;
|
||||
events.push({ type: "tool-call-end", index });
|
||||
}
|
||||
return events;
|
||||
}
|
||||
|
||||
function itemText(item: Record<string, unknown>): string | null {
|
||||
const content = item.content;
|
||||
if (!Array.isArray(content)) return null;
|
||||
return content
|
||||
.map((part) => record(part))
|
||||
.filter((part) => part?.type === "output_text")
|
||||
.map((part) => text(part?.text) ?? "")
|
||||
.join("");
|
||||
}
|
||||
|
||||
function requiredIndex(value: Record<string, unknown>): number {
|
||||
const index = count(value.output_index);
|
||||
if (index === null) throw new Error("OpenAI Responses event is missing output_index");
|
||||
return index;
|
||||
}
|
||||
|
||||
function usageEvent(value: unknown): ModelEvent {
|
||||
const usage = record(value) ?? {};
|
||||
return {
|
||||
type: "usage",
|
||||
usage: {
|
||||
inputTokens: count(usage.input_tokens),
|
||||
outputTokens: count(usage.output_tokens),
|
||||
totalTokens: count(usage.total_tokens),
|
||||
cacheReadTokens: count(record(usage.input_tokens_details)?.cached_tokens),
|
||||
cacheWriteTokens: null,
|
||||
reasoningTokens: count(record(usage.output_tokens_details)?.reasoning_tokens),
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
async function readBody(lines: AsyncIterable<string>): Promise<string> {
|
||||
const collected: string[] = [];
|
||||
for await (const line of lines) collected.push(line);
|
||||
return collected.join("\n");
|
||||
}
|
||||
|
||||
/**
|
||||
* 执行一次 Responses API 流式调用,发出与宿主统一事件集一致的标准化事件,
|
||||
* 包括文本/思考边界、工具参数增量与加密推理回放状态。非 2xx 响应抛出
|
||||
* `HttpError`,流内失败抛出 `Error`,由调用方分类额度与授权问题。
|
||||
*/
|
||||
export async function streamOpenAiResponses(
|
||||
call: OpenAiResponsesCall,
|
||||
output: ProviderOutput,
|
||||
context: PluginContext,
|
||||
): Promise<void> {
|
||||
const response = await context.network.stream(call.url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
accept: "text/event-stream",
|
||||
"content-type": "application/json",
|
||||
...call.headers,
|
||||
},
|
||||
body: JSON.stringify(buildResponsesBody(call)),
|
||||
});
|
||||
if (response.status < 200 || response.status >= 300) {
|
||||
throw new HttpError(response.status, await readBody(response.lines));
|
||||
}
|
||||
|
||||
let textOpen = false;
|
||||
let streamedText = "";
|
||||
let thinkingOpen = false;
|
||||
const tools = new Map<number, ToolState>();
|
||||
const reasoningItems: JsonValue[] = [];
|
||||
let sawTool = false;
|
||||
let sawCompletedItem = false;
|
||||
let terminal = false;
|
||||
|
||||
const closeThinking = () => {
|
||||
if (thinkingOpen) {
|
||||
thinkingOpen = false;
|
||||
output.emit({ type: "thinking-end" });
|
||||
}
|
||||
};
|
||||
const closeText = () => {
|
||||
if (textOpen) {
|
||||
textOpen = false;
|
||||
output.emit({ type: "text-end" });
|
||||
}
|
||||
};
|
||||
// 流式增量可能落后于最终文本;补发缺失的后缀。
|
||||
const reconcileText = (finalText: string) => {
|
||||
if (finalText.startsWith(streamedText) && finalText.length > streamedText.length) {
|
||||
if (!textOpen) {
|
||||
textOpen = true;
|
||||
output.emit({ type: "text-start" });
|
||||
}
|
||||
output.emit({ type: "text-delta", text: finalText.slice(streamedText.length) });
|
||||
streamedText = finalText;
|
||||
}
|
||||
};
|
||||
const endStartedTools = () => {
|
||||
for (const [index, tool] of tools) {
|
||||
if (tool.started && !tool.ended) {
|
||||
tool.ended = true;
|
||||
output.emit({ type: "tool-call-end", index });
|
||||
}
|
||||
}
|
||||
};
|
||||
const emitReplayState = () => {
|
||||
if (reasoningItems.length > 0) {
|
||||
output.emit({ type: "replay-state", providerKind: REPLAY_KIND, value: { items: reasoningItems.slice() } });
|
||||
reasoningItems.length = 0;
|
||||
}
|
||||
};
|
||||
|
||||
for await (const line of response.lines) {
|
||||
if (!line.startsWith("data:")) continue;
|
||||
const payload = line.slice(5).trim();
|
||||
if (!payload) continue;
|
||||
if (payload === "[DONE]") break;
|
||||
let value: Record<string, unknown>;
|
||||
try {
|
||||
value = record(JSON.parse(payload)) ?? {};
|
||||
} catch {
|
||||
throw new Error("OpenAI Responses SSE returned invalid JSON");
|
||||
}
|
||||
switch (value.type) {
|
||||
case "response.output_text.delta": {
|
||||
closeThinking();
|
||||
if (!textOpen) {
|
||||
textOpen = true;
|
||||
output.emit({ type: "text-start" });
|
||||
}
|
||||
const delta = text(value.delta);
|
||||
if (delta !== null) {
|
||||
streamedText += delta;
|
||||
output.emit({ type: "text-delta", text: delta });
|
||||
}
|
||||
break;
|
||||
}
|
||||
case "response.output_text.done": {
|
||||
const finalText = text(value.text);
|
||||
if (finalText !== null) reconcileText(finalText);
|
||||
closeText();
|
||||
break;
|
||||
}
|
||||
case "response.reasoning_summary_text.delta":
|
||||
case "response.reasoning_text.delta": {
|
||||
if (!thinkingOpen) {
|
||||
thinkingOpen = true;
|
||||
output.emit({ type: "thinking-start" });
|
||||
}
|
||||
const delta = text(value.delta);
|
||||
if (delta !== null) output.emit({ type: "thinking-delta", text: delta });
|
||||
break;
|
||||
}
|
||||
case "response.reasoning_summary_text.done":
|
||||
case "response.reasoning_text.done":
|
||||
closeThinking();
|
||||
break;
|
||||
case "response.output_item.added": {
|
||||
const item = record(value.item);
|
||||
if (item?.type !== "function_call") break;
|
||||
sawTool = true;
|
||||
for (const event of updateTool(requiredIndex(value), item, { kind: "none" }, false, tools)) {
|
||||
output.emit(event);
|
||||
}
|
||||
break;
|
||||
}
|
||||
case "response.output_item.done": {
|
||||
const item = record(value.item);
|
||||
if (item?.type === "reasoning") {
|
||||
closeThinking();
|
||||
reasoningItems.push(item as JsonValue);
|
||||
} else if (item?.type === "message") {
|
||||
sawCompletedItem = true;
|
||||
const finalText = itemText(item);
|
||||
if (finalText !== null) reconcileText(finalText);
|
||||
closeText();
|
||||
} else if (item?.type === "function_call") {
|
||||
sawCompletedItem = true;
|
||||
sawTool = true;
|
||||
const snapshot = text(item.arguments);
|
||||
const args: ToolArguments = snapshot === null ? { kind: "none" } : { kind: "snapshot", snapshot };
|
||||
for (const event of updateTool(requiredIndex(value), item, args, true, tools)) {
|
||||
output.emit(event);
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
case "response.function_call_arguments.delta": {
|
||||
const delta = text(value.delta);
|
||||
if (delta === null) break;
|
||||
sawTool = true;
|
||||
for (const event of updateTool(requiredIndex(value), null, { kind: "delta", delta }, false, tools)) {
|
||||
output.emit(event);
|
||||
}
|
||||
break;
|
||||
}
|
||||
case "response.function_call_arguments.done": {
|
||||
const snapshot = text(value.arguments);
|
||||
// 空快照不代表结束;等 output_item.done 收尾。
|
||||
const args: ToolArguments = snapshot === null || snapshot === "" ? { kind: "none" } : { kind: "snapshot", snapshot };
|
||||
const done = snapshot !== null && snapshot !== "";
|
||||
for (const event of updateTool(requiredIndex(value), null, args, done, tools)) {
|
||||
output.emit(event);
|
||||
}
|
||||
break;
|
||||
}
|
||||
case "response.completed": {
|
||||
const usage = record(value.response)?.usage;
|
||||
if (usage !== undefined) output.emit(usageEvent(usage));
|
||||
closeThinking();
|
||||
closeText();
|
||||
endStartedTools();
|
||||
for (const tool of tools.values()) {
|
||||
if (!tool.started) {
|
||||
throw new Error("OpenAI Responses completed with incomplete tool metadata");
|
||||
}
|
||||
}
|
||||
terminal = true;
|
||||
emitReplayState();
|
||||
output.emit({ type: "done", reason: sawTool ? "tool-use" : "stop" });
|
||||
break;
|
||||
}
|
||||
case "response.incomplete": {
|
||||
closeThinking();
|
||||
closeText();
|
||||
endStartedTools();
|
||||
terminal = true;
|
||||
output.emit({ type: "done", reason: "length" });
|
||||
break;
|
||||
}
|
||||
case "response.failed":
|
||||
throw new Error(`OpenAI Responses failed: ${payload}`);
|
||||
}
|
||||
if (terminal) break;
|
||||
}
|
||||
|
||||
if (!terminal && sawCompletedItem) {
|
||||
closeThinking();
|
||||
closeText();
|
||||
for (const tool of tools.values()) {
|
||||
if (!tool.ended) {
|
||||
throw new Error("OpenAI Responses stream ended with an incomplete tool call");
|
||||
}
|
||||
}
|
||||
terminal = true;
|
||||
emitReplayState();
|
||||
output.emit({ type: "done", reason: sawTool ? "tool-use" : "stop" });
|
||||
}
|
||||
if (!terminal) {
|
||||
throw new Error("OpenAI Responses stream ended without response.completed or response.incomplete");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
import type { JsonValue, LocalizedText, PluginContext } from "./plugin.ts";
|
||||
import type { ModelSnapshot, ModelSupport } from "./model.ts";
|
||||
import type { ResourcePatch, ResourceSnapshot } from "./resource.ts";
|
||||
|
||||
/**
|
||||
* LLM 请求契约。宿主把它的规范会话(ProjectedMessage)投影成这个形状;
|
||||
* 插件负责把它适配成上游 Provider 的协议。
|
||||
*/
|
||||
export type LlmContentPart =
|
||||
| { type: "text"; text: string }
|
||||
| { type: "image"; mediaType: string; dataBase64: string };
|
||||
|
||||
/** 不透明的 Provider 回放状态(如加密推理项);回放时按 providerKind 过滤。 */
|
||||
export type LlmReplayState = {
|
||||
providerKind: string;
|
||||
value: JsonValue;
|
||||
};
|
||||
|
||||
export type LlmToolCall = {
|
||||
/** 同一轮内的稳定序号。 */
|
||||
index: number;
|
||||
callId: string;
|
||||
name: string;
|
||||
/** 已解析的 JSON 参数。 */
|
||||
arguments: JsonValue;
|
||||
};
|
||||
|
||||
export type LlmMessage =
|
||||
| { role: "system" | "user"; content: LlmContentPart[] }
|
||||
| {
|
||||
role: "assistant";
|
||||
text: string;
|
||||
thinking: string;
|
||||
replayState: LlmReplayState | null;
|
||||
toolCalls: LlmToolCall[];
|
||||
}
|
||||
| {
|
||||
role: "tool";
|
||||
callId: string;
|
||||
name: string;
|
||||
content: string;
|
||||
isError: boolean;
|
||||
/** 非空时优先于 content,承载图片等富工具结果。 */
|
||||
parts: LlmContentPart[];
|
||||
};
|
||||
|
||||
export type LlmTool = {
|
||||
name: string;
|
||||
description: string;
|
||||
/** 工具参数的 JSON Schema。 */
|
||||
parameters: JsonValue;
|
||||
};
|
||||
|
||||
export type LlmRequest = {
|
||||
/** 系统指令;空字符串表示没有。 */
|
||||
instructions: string;
|
||||
messages: LlmMessage[];
|
||||
tools: LlmTool[];
|
||||
reasoning: { enabled: boolean; effort: string | null };
|
||||
latency: "fast" | "standard";
|
||||
maxOutputTokens: number | null;
|
||||
/** 会话级稳定缓存键,用于上游前缀缓存的路由亲和(如 prompt_cache_key)。 */
|
||||
cacheKey: string | null;
|
||||
};
|
||||
|
||||
export type ModelUsage = {
|
||||
inputTokens: number | null;
|
||||
outputTokens: number | null;
|
||||
totalTokens: number | null;
|
||||
cacheReadTokens: number | null;
|
||||
cacheWriteTokens: number | null;
|
||||
reasoningTokens: number | null;
|
||||
};
|
||||
|
||||
/**
|
||||
* 标准化输出契约,与宿主统一流事件一一对应。插件边接收上游数据边发出事件;
|
||||
* 文本、思考和每个工具调用都有显式的开始/结束边界,工具参数以增量交付。
|
||||
* 回放状态在流结束前发出一次,宿主存入 assistant 消息供下一轮回放。
|
||||
*/
|
||||
export type ModelEvent =
|
||||
| { type: "text-start" }
|
||||
| { type: "text-delta"; text: string }
|
||||
| { type: "text-end" }
|
||||
| { type: "thinking-start" }
|
||||
| { type: "thinking-delta"; text: string }
|
||||
| { type: "thinking-end" }
|
||||
| { type: "tool-call-start"; index: number; callId: string; name: string }
|
||||
| { type: "tool-call-arguments-delta"; index: number; delta: string }
|
||||
| { type: "tool-call-end"; index: number }
|
||||
| { type: "replay-state"; providerKind: string; value: JsonValue }
|
||||
| { type: "usage"; usage: ModelUsage }
|
||||
| { type: "done"; reason: "stop" | "length" | "tool-use" };
|
||||
|
||||
export type ProviderOutput = {
|
||||
emit(event: ModelEvent): void;
|
||||
};
|
||||
|
||||
export type ProviderInvokeInput = {
|
||||
model: ModelSnapshot;
|
||||
/** 宿主为本次调用选中的资源;无资源 Provider 为 null。 */
|
||||
resource: ResourceSnapshot | null;
|
||||
request: LlmRequest;
|
||||
};
|
||||
|
||||
/**
|
||||
* `resource-error` 把失败归因到选中的资源,宿主据此更新资源状态,
|
||||
* 并可在尚未发出任何事件时(未来)换一个资源重试。`patch` 同时用于
|
||||
* 持久化成功调用的副作用,例如刷新后的 access token。
|
||||
*/
|
||||
export type ProviderResult =
|
||||
| { status: "completed"; patch?: ResourcePatch }
|
||||
| { status: "resource-error"; message: string; patch: ResourcePatch }
|
||||
| { status: "request-error"; message: string; patch?: ResourcePatch };
|
||||
|
||||
export type ProviderSupport = {
|
||||
id: string;
|
||||
displayName: LocalizedText;
|
||||
description?: LocalizedText;
|
||||
/** 产品身份,用于归类与图标,如 "openai"。 */
|
||||
providerType: string;
|
||||
/** 每次调用消费的资源类型;无资源 Provider 可省略。 */
|
||||
resourceType?: string;
|
||||
models?: ModelSupport;
|
||||
invoke(
|
||||
input: ProviderInvokeInput,
|
||||
output: ProviderOutput,
|
||||
context: PluginContext,
|
||||
): Promise<ProviderResult>;
|
||||
};
|
||||
@@ -0,0 +1,118 @@
|
||||
import type { JsonValue, LocalizedText, PluginContext } from "./plugin.ts";
|
||||
|
||||
/**
|
||||
* 资源是插件定义的私有记录(通常是上游账号),由 Provider 消费。
|
||||
* 宿主负责持久化、列表和每次调用的资源选择;插件只负责创建、投影和解释资源。
|
||||
*/
|
||||
export type ResourceState =
|
||||
| { status: "ready" }
|
||||
| { status: "cooling"; retryAtMs?: number; message?: string }
|
||||
| { status: "invalid"; message?: string };
|
||||
|
||||
/** 由添加流程或导入产生的新资源。 */
|
||||
export type ResourceDraft = {
|
||||
/** 去重键:宿主按 (资源类型, key) 执行 upsert。 */
|
||||
key: string;
|
||||
/** 凭证与插件私有字段;永远不会展示给用户。 */
|
||||
privateData: JsonValue;
|
||||
/** 缺省为 ready。 */
|
||||
state?: ResourceState;
|
||||
};
|
||||
|
||||
/** 宿主已持久化的一条资源。 */
|
||||
export type ResourceSnapshot = {
|
||||
/** 宿主分配的标识,区别于插件的去重键。 */
|
||||
id: string;
|
||||
type: string;
|
||||
key: string;
|
||||
privateData: JsonValue;
|
||||
state: ResourceState;
|
||||
};
|
||||
|
||||
/** 宿主原子应用到单条资源上的部分更新。 */
|
||||
export type ResourcePatch = {
|
||||
privateData?: JsonValue;
|
||||
state?: ResourceState;
|
||||
};
|
||||
|
||||
export type ResourceMetric = {
|
||||
id: string;
|
||||
label: LocalizedText;
|
||||
unit: "percent" | "count";
|
||||
/** percent 指标表示剩余占比,0..100。 */
|
||||
value: number;
|
||||
resetAtMs?: number;
|
||||
};
|
||||
|
||||
/** 单条资源的用户可见投影;不得泄露凭证。displayName 是数据(如邮箱),保持纯字符串。 */
|
||||
export type ResourceView = {
|
||||
displayName: string;
|
||||
description?: LocalizedText;
|
||||
metrics?: ResourceMetric[];
|
||||
};
|
||||
|
||||
/**
|
||||
* OAuth 2.0 设备码式添加流程。宿主负责绘制 UI、驱动轮询循环
|
||||
* (间隔、slow-down 退避、超时判定),并在流程存续期内在内存中持有
|
||||
* `session`;插件只实现两次 HTTP 状态转移。
|
||||
*/
|
||||
export type OAuth2AddMethod = {
|
||||
type: "oauth2.0";
|
||||
id: string;
|
||||
displayName: LocalizedText;
|
||||
description?: LocalizedText;
|
||||
begin(context: PluginContext): Promise<OAuth2Begin>;
|
||||
poll(session: JsonValue, context: PluginContext): Promise<OAuth2Poll>;
|
||||
};
|
||||
|
||||
export type OAuth2Begin = {
|
||||
/** 不透明流程状态(设备码、PKCE verifier 等);永远不会持久化。 */
|
||||
session: JsonValue;
|
||||
userCode: string;
|
||||
verificationUrl: string;
|
||||
verificationUrlComplete?: string;
|
||||
expiresAtMs: number;
|
||||
pollIntervalMs: number;
|
||||
};
|
||||
|
||||
export type OAuth2Poll =
|
||||
| { status: "pending"; session?: JsonValue }
|
||||
| { status: "slow-down"; session?: JsonValue }
|
||||
| { status: "completed"; resources: ResourceDraft[] }
|
||||
| { status: "denied"; message?: string }
|
||||
| { status: "failed"; message: string };
|
||||
|
||||
export type ResourceAddMethod = OAuth2AddMethod;
|
||||
|
||||
export type ResourceImportFile = {
|
||||
name: string;
|
||||
/** 文件原文;解析和校验由插件负责。 */
|
||||
content: string;
|
||||
};
|
||||
|
||||
export type ResourceImportSupport = {
|
||||
displayName: LocalizedText;
|
||||
description?: LocalizedText;
|
||||
/** 宿主文件选择器接受的扩展名,如 [".json"]。 */
|
||||
accept: string[];
|
||||
multiple?: boolean;
|
||||
parse(files: ResourceImportFile[], context: PluginContext): Promise<ResourceImportResult>;
|
||||
};
|
||||
|
||||
export type ResourceImportResult = {
|
||||
resources: ResourceDraft[];
|
||||
/** 单个文件的问题,值得提示但不必使整次导入失败。 */
|
||||
warnings?: string[];
|
||||
};
|
||||
|
||||
export type ResourceSupport = {
|
||||
type: string;
|
||||
displayName: LocalizedText;
|
||||
add?: ResourceAddMethod[];
|
||||
import?: ResourceImportSupport;
|
||||
present(resource: ResourceSnapshot): ResourceView;
|
||||
/** 用户主动触发时重新读取上游状态(额度、凭证有效性)。 */
|
||||
refresh?(resource: ResourceSnapshot, context: PluginContext): Promise<ResourcePatch>;
|
||||
/** 可选的上游撤销;宿主随后删除本地记录。 */
|
||||
remove?(resource: ResourceSnapshot, context: PluginContext): Promise<void>;
|
||||
};
|
||||
@@ -0,0 +1,185 @@
|
||||
import { __getRegisteredPlugin, type JsonValue, type NetworkEventStream, type PluginContext } from "cursor-byok:plugin";
|
||||
import type { ModelEvent, ProviderSupport } from "cursor-byok:provider";
|
||||
import type { ResourceAddMethod, ResourceSupport } from "cursor-byok:resource";
|
||||
|
||||
if (Deno.args.length !== 1) throw new Error("plugin entry URL is required");
|
||||
await import(Deno.args[0]);
|
||||
const plugin = __getRegisteredPlugin();
|
||||
const encoder = new TextEncoder();
|
||||
const writer = Deno.stdout.writable.getWriter();
|
||||
const pendingHost = new Map<string, { resolve(value: unknown): void; reject(error: Error): void }>();
|
||||
const controllers = new Map<string, AbortController>();
|
||||
let hostSequence = 0;
|
||||
// 事件与最终结果共用一条串行写队列,保证顺序。
|
||||
let writeQueue = Promise.resolve();
|
||||
|
||||
function send(value: unknown): Promise<void> {
|
||||
const operation = writeQueue.then(() => writer.write(encoder.encode(JSON.stringify(value) + "\n")));
|
||||
writeQueue = operation.catch(() => undefined);
|
||||
return operation;
|
||||
}
|
||||
|
||||
function hostCall(requestId: string, method: string, params: unknown): Promise<unknown> {
|
||||
const id = `${requestId}:host:${++hostSequence}`;
|
||||
return new Promise((resolve, reject) => {
|
||||
pendingHost.set(id, { resolve, reject });
|
||||
void send({ type: "host_call", id, requestId, method, params });
|
||||
});
|
||||
}
|
||||
|
||||
async function* streamLines(requestId: string, streamId: string): AsyncGenerator<string> {
|
||||
try {
|
||||
for (;;) {
|
||||
const chunk = await hostCall(requestId, "network.stream.read", { streamId }) as {
|
||||
lines: string[];
|
||||
done: boolean;
|
||||
};
|
||||
for (const line of chunk.lines) yield line;
|
||||
if (chunk.done) return;
|
||||
}
|
||||
} finally {
|
||||
void hostCall(requestId, "network.stream.close", { streamId }).catch(() => undefined);
|
||||
}
|
||||
}
|
||||
|
||||
function contextFor(requestId: string, signal: AbortSignal): PluginContext {
|
||||
return {
|
||||
network: {
|
||||
fetch: (url, init = {}) => hostCall(requestId, "network.fetch", { url, ...init }) as ReturnType<PluginContext["network"]["fetch"]>,
|
||||
stream: async (url, init = {}): Promise<NetworkEventStream> => {
|
||||
const opened = await hostCall(requestId, "network.stream.open", { url, ...init }) as {
|
||||
streamId: string;
|
||||
status: number;
|
||||
headers: Record<string, string>;
|
||||
};
|
||||
return {
|
||||
status: opened.status,
|
||||
headers: opened.headers,
|
||||
lines: streamLines(requestId, opened.streamId),
|
||||
};
|
||||
},
|
||||
},
|
||||
signal,
|
||||
};
|
||||
}
|
||||
|
||||
function provider(id: unknown): ProviderSupport {
|
||||
const found = plugin.providers.find((provider) => provider.id === id);
|
||||
if (!found) throw new Error(`unknown plugin provider: ${id}`);
|
||||
return found;
|
||||
}
|
||||
|
||||
function resourceSupport(type: unknown): ResourceSupport {
|
||||
const found = (plugin.resources ?? []).find((resource) => resource.type === type);
|
||||
if (!found) throw new Error(`unknown plugin resource type: ${type}`);
|
||||
return found;
|
||||
}
|
||||
|
||||
function addMethod(support: ResourceSupport, methodId: unknown): ResourceAddMethod {
|
||||
const found = (support.add ?? []).find((method) => method.id === methodId);
|
||||
if (!found) throw new Error(`unknown plugin add method: ${methodId}`);
|
||||
return found;
|
||||
}
|
||||
|
||||
async function dispatch(message: { id: string; method: string; params?: JsonValue }) {
|
||||
const controller = new AbortController();
|
||||
controllers.set(message.id, controller);
|
||||
const context = contextFor(message.id, controller.signal);
|
||||
const params = (message.params ?? {}) as Record<string, JsonValue>;
|
||||
try {
|
||||
let result: unknown;
|
||||
switch (message.method) {
|
||||
case "provider.invoke": {
|
||||
const output = {
|
||||
emit: (event: ModelEvent) => void send({ type: "event", id: message.id, event }),
|
||||
};
|
||||
result = await provider(params.providerId).invoke(
|
||||
{
|
||||
model: params.model as never,
|
||||
resource: (params.resource ?? null) as never,
|
||||
request: params.request as never,
|
||||
},
|
||||
output,
|
||||
context,
|
||||
);
|
||||
break;
|
||||
}
|
||||
case "models.list": {
|
||||
const models = provider(params.providerId).models;
|
||||
if (!models) throw new Error(`plugin provider ${params.providerId} has no models`);
|
||||
result = await models.list({ resource: (params.resource ?? null) as never }, context);
|
||||
break;
|
||||
}
|
||||
case "resource.present": {
|
||||
const support = resourceSupport(params.resourceType);
|
||||
const resources = Array.isArray(params.resources) ? params.resources : [];
|
||||
result = resources.map((resource) => support.present(resource as never));
|
||||
break;
|
||||
}
|
||||
case "resource.refresh": {
|
||||
const support = resourceSupport(params.resourceType);
|
||||
if (!support.refresh) throw new Error(`resource ${params.resourceType} has no refresh`);
|
||||
result = await support.refresh(params.resource as never, context);
|
||||
break;
|
||||
}
|
||||
case "resource.remove": {
|
||||
const support = resourceSupport(params.resourceType);
|
||||
await support.remove?.(params.resource as never, context);
|
||||
result = null;
|
||||
break;
|
||||
}
|
||||
case "oauth.begin": {
|
||||
const support = resourceSupport(params.resourceType);
|
||||
result = await addMethod(support, params.methodId).begin(context);
|
||||
break;
|
||||
}
|
||||
case "oauth.poll": {
|
||||
const support = resourceSupport(params.resourceType);
|
||||
result = await addMethod(support, params.methodId).poll(params.session ?? null, context);
|
||||
break;
|
||||
}
|
||||
case "import.parse": {
|
||||
const support = resourceSupport(params.resourceType);
|
||||
if (!support.import) throw new Error(`resource ${params.resourceType} has no import`);
|
||||
const files = Array.isArray(params.files) ? params.files : [];
|
||||
result = await support.import.parse(files as never, context);
|
||||
break;
|
||||
}
|
||||
default:
|
||||
throw new Error(`unknown plugin method: ${message.method}`);
|
||||
}
|
||||
await send({ type: "result", id: message.id, result: result ?? null });
|
||||
} catch (error) {
|
||||
await send({ type: "result", id: message.id, error: error instanceof Error ? error.message : String(error) });
|
||||
} finally {
|
||||
controllers.delete(message.id);
|
||||
}
|
||||
}
|
||||
|
||||
let buffered = "";
|
||||
for await (const chunk of Deno.stdin.readable.pipeThrough(new TextDecoderStream())) {
|
||||
buffered += chunk;
|
||||
for (;;) {
|
||||
const newline = buffered.indexOf("\n");
|
||||
if (newline < 0) break;
|
||||
const line = buffered.slice(0, newline);
|
||||
buffered = buffered.slice(newline + 1);
|
||||
if (!line.trim()) continue;
|
||||
const message = JSON.parse(line);
|
||||
if (message.type === "request") {
|
||||
void dispatch(message);
|
||||
} else if (message.type === "cancel") {
|
||||
controllers.get(message.id)?.abort();
|
||||
} else if (message.type === "host_result") {
|
||||
const pending = pendingHost.get(message.id);
|
||||
if (!pending) continue;
|
||||
pendingHost.delete(message.id);
|
||||
pending.resolve(message.result);
|
||||
} else if (message.type === "host_error") {
|
||||
const pending = pendingHost.get(message.id);
|
||||
if (!pending) continue;
|
||||
pendingHost.delete(message.id);
|
||||
pending.reject(new Error(message.error));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,457 @@
|
||||
//! Owns core-side persistence of plugin resources and model catalogs.
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::data::PluginDataStore;
|
||||
use crate::{Error, Result};
|
||||
|
||||
/// 核心理解的资源运行状态;插件只能通过 draft/patch/report 改变它。
|
||||
#[derive(Clone, Debug, Deserialize, Serialize, PartialEq)]
|
||||
#[serde(tag = "status", rename_all = "snake_case")]
|
||||
pub enum ResourceState {
|
||||
Ready,
|
||||
Cooling {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
retry_at_ms: Option<i64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
message: Option<String>,
|
||||
},
|
||||
Invalid {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
message: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl ResourceState {
|
||||
/// 冷却到期后自动恢复可用。
|
||||
pub fn is_ready(&self, now_ms: i64) -> bool {
|
||||
match self {
|
||||
Self::Ready => true,
|
||||
Self::Cooling { retry_at_ms, .. } => retry_at_ms.is_some_and(|at| at <= now_ms),
|
||||
Self::Invalid { .. } => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 核心持久化的一条插件资源。`private_data` 只回传给插件。
|
||||
#[derive(Clone, Debug, Deserialize, Serialize)]
|
||||
pub struct ResourceRecord {
|
||||
pub id: String,
|
||||
pub key: String,
|
||||
pub private_data: serde_json::Value,
|
||||
pub state: ResourceState,
|
||||
pub created_at_ms: i64,
|
||||
pub updated_at_ms: i64,
|
||||
}
|
||||
|
||||
impl ResourceRecord {
|
||||
/// 传给插件的快照形状(SDK 的 ResourceSnapshot)。
|
||||
pub fn snapshot(&self, resource_type: &str) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"id": self.id,
|
||||
"type": resource_type,
|
||||
"key": self.key,
|
||||
"privateData": self.private_data,
|
||||
"state": state_json(&self.state),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn state_json(state: &ResourceState) -> serde_json::Value {
|
||||
match state {
|
||||
ResourceState::Ready => serde_json::json!({ "status": "ready" }),
|
||||
ResourceState::Cooling {
|
||||
retry_at_ms,
|
||||
message,
|
||||
} => serde_json::json!({
|
||||
"status": "cooling",
|
||||
"retryAtMs": retry_at_ms,
|
||||
"message": message,
|
||||
}),
|
||||
ResourceState::Invalid { message } => serde_json::json!({
|
||||
"status": "invalid",
|
||||
"message": message,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
/// 插件返回的新资源(SDK 的 ResourceDraft)。
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct ResourceDraft {
|
||||
pub key: String,
|
||||
pub private_data: serde_json::Value,
|
||||
#[serde(default)]
|
||||
pub state: Option<ResourceStateInput>,
|
||||
}
|
||||
|
||||
/// 插件对单条资源的部分更新(SDK 的 ResourcePatch)。
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct ResourcePatch {
|
||||
#[serde(default)]
|
||||
pub private_data: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub state: Option<ResourceStateInput>,
|
||||
}
|
||||
|
||||
/// SDK 侧 camelCase 状态输入,转换成核心存储形状。
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(tag = "status", rename_all = "kebab-case", deny_unknown_fields)]
|
||||
pub enum ResourceStateInput {
|
||||
Ready,
|
||||
Cooling {
|
||||
#[serde(default, rename = "retryAtMs")]
|
||||
retry_at_ms: Option<i64>,
|
||||
#[serde(default)]
|
||||
message: Option<String>,
|
||||
},
|
||||
Invalid {
|
||||
#[serde(default)]
|
||||
message: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl From<ResourceStateInput> for ResourceState {
|
||||
fn from(input: ResourceStateInput) -> Self {
|
||||
match input {
|
||||
ResourceStateInput::Ready => Self::Ready,
|
||||
ResourceStateInput::Cooling {
|
||||
retry_at_ms,
|
||||
message,
|
||||
} => Self::Cooling {
|
||||
retry_at_ms,
|
||||
message,
|
||||
},
|
||||
ResourceStateInput::Invalid { message } => Self::Invalid { message },
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 插件发现的一个模型(SDK 的 ModelDefinition),由核心整体替换目录。
|
||||
#[derive(Clone, Debug, Deserialize, Serialize)]
|
||||
pub struct StoredModel {
|
||||
pub id: String,
|
||||
pub display_name: String,
|
||||
#[serde(default)]
|
||||
pub description: Option<String>,
|
||||
#[serde(default)]
|
||||
pub max_output_tokens: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub images: bool,
|
||||
#[serde(default)]
|
||||
pub private_data: serde_json::Value,
|
||||
}
|
||||
|
||||
impl StoredModel {
|
||||
pub fn from_definition(value: &serde_json::Value) -> Result<Self> {
|
||||
let object = value
|
||||
.as_object()
|
||||
.ok_or_else(|| Error::Protocol("plugin model definition must be an object".into()))?;
|
||||
let id = object
|
||||
.get("id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.filter(|id| !id.trim().is_empty())
|
||||
.ok_or_else(|| Error::Protocol("plugin model definition requires id".into()))?;
|
||||
let display_name = object
|
||||
.get("displayName")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.filter(|name| !name.trim().is_empty())
|
||||
.ok_or_else(|| {
|
||||
Error::Protocol("plugin model definition requires displayName".into())
|
||||
})?;
|
||||
let capabilities = object
|
||||
.get("capabilities")
|
||||
.and_then(|value| value.as_object());
|
||||
let capability = |name: &str| {
|
||||
capabilities
|
||||
.and_then(|value| value.get(name))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
};
|
||||
Ok(Self {
|
||||
id: id.to_owned(),
|
||||
display_name: display_name.to_owned(),
|
||||
description: object
|
||||
.get("description")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_owned),
|
||||
max_output_tokens: object
|
||||
.get("maxOutputTokens")
|
||||
.and_then(serde_json::Value::as_u64),
|
||||
images: capability("images"),
|
||||
private_data: object
|
||||
.get("privateData")
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null),
|
||||
})
|
||||
}
|
||||
|
||||
/// 传给插件的模型快照(SDK 的 ModelSnapshot)。
|
||||
pub fn snapshot(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"id": self.id,
|
||||
"displayName": self.display_name,
|
||||
"description": self.description,
|
||||
"maxOutputTokens": self.max_output_tokens,
|
||||
"capabilities": { "images": self.images },
|
||||
"privateData": self.private_data,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 资源与模型目录的核心存储,构建在插件私有 JSON 文件之上。
|
||||
#[derive(Clone)]
|
||||
pub struct PluginStateStore {
|
||||
data: PluginDataStore,
|
||||
}
|
||||
|
||||
pub struct UpsertOutcome {
|
||||
pub added: usize,
|
||||
pub updated: usize,
|
||||
}
|
||||
|
||||
impl PluginStateStore {
|
||||
pub fn new(data: PluginDataStore) -> Self {
|
||||
Self { data }
|
||||
}
|
||||
|
||||
pub async fn resources(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
) -> Result<Vec<ResourceRecord>> {
|
||||
let value = self
|
||||
.data
|
||||
.read(plugin_id, &resource_key(resource_type))
|
||||
.await?;
|
||||
if value.is_null() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
Ok(serde_json::from_value(value)?)
|
||||
}
|
||||
|
||||
pub async fn upsert_resources(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
drafts: Vec<ResourceDraft>,
|
||||
) -> Result<UpsertOutcome> {
|
||||
let mut records = self.resources(plugin_id, resource_type).await?;
|
||||
let now = now_ms();
|
||||
let mut outcome = UpsertOutcome {
|
||||
added: 0,
|
||||
updated: 0,
|
||||
};
|
||||
for draft in drafts {
|
||||
if draft.key.trim().is_empty() {
|
||||
return Err(Error::Protocol("plugin resource draft requires key".into()));
|
||||
}
|
||||
let state = draft
|
||||
.state
|
||||
.map_or(ResourceState::Ready, ResourceState::from);
|
||||
match records.iter_mut().find(|record| record.key == draft.key) {
|
||||
Some(existing) => {
|
||||
existing.private_data = draft.private_data;
|
||||
existing.state = state;
|
||||
existing.updated_at_ms = now;
|
||||
outcome.updated += 1;
|
||||
}
|
||||
None => {
|
||||
records.push(ResourceRecord {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
key: draft.key,
|
||||
private_data: draft.private_data,
|
||||
state,
|
||||
created_at_ms: now,
|
||||
updated_at_ms: now,
|
||||
});
|
||||
outcome.added += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
self.save_resources(plugin_id, resource_type, &records)
|
||||
.await?;
|
||||
Ok(outcome)
|
||||
}
|
||||
|
||||
pub async fn apply_patch(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
patch: ResourcePatch,
|
||||
) -> Result<()> {
|
||||
let mut records = self.resources(plugin_id, resource_type).await?;
|
||||
let record = records
|
||||
.iter_mut()
|
||||
.find(|record| record.id == resource_id)
|
||||
.ok_or_else(|| Error::RunNotFound(format!("plugin resource {resource_id}")))?;
|
||||
if let Some(private_data) = patch.private_data {
|
||||
record.private_data = private_data;
|
||||
}
|
||||
if let Some(state) = patch.state {
|
||||
record.state = state.into();
|
||||
}
|
||||
record.updated_at_ms = now_ms();
|
||||
self.save_resources(plugin_id, resource_type, &records)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn remove_resource(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<ResourceRecord> {
|
||||
let mut records = self.resources(plugin_id, resource_type).await?;
|
||||
let index = records
|
||||
.iter()
|
||||
.position(|record| record.id == resource_id)
|
||||
.ok_or_else(|| Error::RunNotFound(format!("plugin resource {resource_id}")))?;
|
||||
let removed = records.remove(index);
|
||||
self.save_resources(plugin_id, resource_type, &records)
|
||||
.await?;
|
||||
Ok(removed)
|
||||
}
|
||||
|
||||
pub async fn models(&self, plugin_id: &str, provider_id: &str) -> Result<Vec<StoredModel>> {
|
||||
let value = self.data.read(plugin_id, &model_key(provider_id)).await?;
|
||||
if value.is_null() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
Ok(serde_json::from_value(value)?)
|
||||
}
|
||||
|
||||
pub async fn replace_models(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
provider_id: &str,
|
||||
models: &[StoredModel],
|
||||
) -> Result<()> {
|
||||
self.data
|
||||
.update(
|
||||
plugin_id,
|
||||
&model_key(provider_id),
|
||||
&serde_json::to_value(models)?,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn clear(&self, plugin_id: &str) -> Result<()> {
|
||||
self.data.clear(plugin_id).await
|
||||
}
|
||||
|
||||
async fn save_resources(
|
||||
&self,
|
||||
plugin_id: &str,
|
||||
resource_type: &str,
|
||||
records: &[ResourceRecord],
|
||||
) -> Result<()> {
|
||||
self.data
|
||||
.update(
|
||||
plugin_id,
|
||||
&resource_key(resource_type),
|
||||
&serde_json::to_value(records)?,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
pub fn now_ms() -> i64 {
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.map(|duration| duration.as_millis().min(i64::MAX as u128) as i64)
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn resource_key(resource_type: &str) -> String {
|
||||
format!("resources-{resource_type}")
|
||||
}
|
||||
|
||||
fn model_key(provider_id: &str) -> String {
|
||||
format!("models-{provider_id}")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn store() -> (tempfile::TempDir, PluginStateStore) {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let data = PluginDataStore::for_test(root.path().join("data")).unwrap();
|
||||
(root, PluginStateStore::new(data))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upserts_resources_by_key_and_applies_patches() {
|
||||
let (_root, store) = store();
|
||||
let outcome = store
|
||||
.upsert_resources(
|
||||
"dev.example",
|
||||
"account",
|
||||
vec![ResourceDraft {
|
||||
key: "acct-1".into(),
|
||||
private_data: serde_json::json!({"token":"one"}),
|
||||
state: None,
|
||||
}],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(outcome.added, 1);
|
||||
let outcome = store
|
||||
.upsert_resources(
|
||||
"dev.example",
|
||||
"account",
|
||||
vec![ResourceDraft {
|
||||
key: "acct-1".into(),
|
||||
private_data: serde_json::json!({"token":"two"}),
|
||||
state: None,
|
||||
}],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(outcome.updated, 1);
|
||||
let records = store.resources("dev.example", "account").await.unwrap();
|
||||
assert_eq!(records.len(), 1);
|
||||
assert_eq!(records[0].private_data["token"], "two");
|
||||
|
||||
store
|
||||
.apply_patch(
|
||||
"dev.example",
|
||||
"account",
|
||||
&records[0].id,
|
||||
ResourcePatch {
|
||||
private_data: None,
|
||||
state: Some(ResourceStateInput::Cooling {
|
||||
retry_at_ms: Some(200),
|
||||
message: None,
|
||||
}),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let records = store.resources("dev.example", "account").await.unwrap();
|
||||
assert!(!records[0].state.is_ready(100));
|
||||
assert!(records[0].state.is_ready(300), "cooling expires over time");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn replaces_model_catalogs() {
|
||||
let (_root, store) = store();
|
||||
let model = StoredModel::from_definition(&serde_json::json!({
|
||||
"id": "gpt-test",
|
||||
"displayName": "GPT Test",
|
||||
"capabilities": {"images": true},
|
||||
"privateData": {"reasoningEfforts": ["low"]},
|
||||
}))
|
||||
.unwrap();
|
||||
store
|
||||
.replace_models("dev.example", "codex", &[model])
|
||||
.await
|
||||
.unwrap();
|
||||
let models = store.models("dev.example", "codex").await.unwrap();
|
||||
assert_eq!(models.len(), 1);
|
||||
assert!(models[0].images);
|
||||
assert_eq!(models[0].private_data["reasoningEfforts"][0], "low");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,297 @@
|
||||
//! Translates between core model types and the plugin SDK wire contract.
|
||||
use base64::{engine::general_purpose::STANDARD, Engine};
|
||||
|
||||
use crate::{
|
||||
model::{
|
||||
ContentPart, ModelInvocation, ModelLatency, ProjectedContent, ProjectedMessage,
|
||||
ProviderReplayState, Role, Usage,
|
||||
},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
/// 把一次核心模型调用投影成 SDK 的 LlmRequest。
|
||||
pub fn llm_request(invocation: &ModelInvocation) -> Result<serde_json::Value> {
|
||||
let request = &invocation.request;
|
||||
let messages = request
|
||||
.history
|
||||
.iter()
|
||||
.map(wire_message)
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
Ok(serde_json::json!({
|
||||
"instructions": request.prompt.instructions,
|
||||
"messages": messages,
|
||||
"tools": request.prompt.tools.iter().map(|tool| serde_json::json!({
|
||||
"name": tool.name,
|
||||
"description": tool.description,
|
||||
"parameters": tool.parameters,
|
||||
})).collect::<Vec<_>>(),
|
||||
"reasoning": {
|
||||
"enabled": request.model.reasoning.enabled,
|
||||
"effort": request.model.reasoning.effort,
|
||||
},
|
||||
"latency": match request.model.latency {
|
||||
ModelLatency::Fast => "fast",
|
||||
_ => "standard",
|
||||
},
|
||||
"maxOutputTokens": request.model.max_output_tokens,
|
||||
"cacheKey": invocation.conversation_id,
|
||||
}))
|
||||
}
|
||||
|
||||
fn wire_message(message: &ProjectedMessage) -> Result<serde_json::Value> {
|
||||
match &message.content {
|
||||
ProjectedContent::Parts(parts) => match message.role {
|
||||
Role::System | Role::User => Ok(serde_json::json!({
|
||||
"role": if message.role == Role::System { "system" } else { "user" },
|
||||
"content": wire_parts(parts),
|
||||
})),
|
||||
// 纯文本 assistant 历史消息投影成无工具调用的 assistant。
|
||||
Role::Assistant => Ok(serde_json::json!({
|
||||
"role": "assistant",
|
||||
"text": joined_text(parts),
|
||||
"thinking": "",
|
||||
"replayState": serde_json::Value::Null,
|
||||
"toolCalls": [],
|
||||
})),
|
||||
Role::Tool => Err(Error::Protocol(
|
||||
"tool messages must carry a tool result".into(),
|
||||
)),
|
||||
},
|
||||
ProjectedContent::Assistant {
|
||||
text,
|
||||
thinking,
|
||||
replay_state,
|
||||
calls,
|
||||
} => Ok(serde_json::json!({
|
||||
"role": "assistant",
|
||||
"text": text,
|
||||
"thinking": thinking,
|
||||
"replayState": replay_state.as_ref().map(|state| serde_json::json!({
|
||||
"providerKind": state.provider_kind,
|
||||
"value": state.value,
|
||||
})),
|
||||
"toolCalls": calls.iter().map(|call| serde_json::json!({
|
||||
"index": call.index,
|
||||
"callId": call.call_id,
|
||||
"name": call.name,
|
||||
"arguments": call.arguments,
|
||||
})).collect::<Vec<_>>(),
|
||||
})),
|
||||
ProjectedContent::ToolResult(result) => Ok(serde_json::json!({
|
||||
"role": "tool",
|
||||
"callId": result.call_id,
|
||||
"name": result.name,
|
||||
"content": result.content,
|
||||
"isError": result.is_error,
|
||||
"parts": wire_parts(&result.provider_parts),
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
fn wire_parts(parts: &[ContentPart]) -> Vec<serde_json::Value> {
|
||||
parts
|
||||
.iter()
|
||||
.map(|part| match part {
|
||||
ContentPart::Text { text } => serde_json::json!({ "type": "text", "text": text }),
|
||||
ContentPart::Image { mime_type, data } => serde_json::json!({
|
||||
"type": "image",
|
||||
"mediaType": mime_type,
|
||||
"dataBase64": STANDARD.encode(data),
|
||||
}),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn joined_text(parts: &[ContentPart]) -> String {
|
||||
parts
|
||||
.iter()
|
||||
.filter_map(|part| match part {
|
||||
ContentPart::Text { text } => Some(text.as_str()),
|
||||
ContentPart::Image { .. } => None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// 把插件发出的标准化事件解析为核心 ModelEvent。
|
||||
pub fn model_event(value: &serde_json::Value) -> Result<ModelEvent> {
|
||||
let kind = value
|
||||
.get("type")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or_else(|| Error::Protocol("plugin model event requires type".into()))?;
|
||||
let event = match kind {
|
||||
"text-start" => ModelEvent::TextStart,
|
||||
"text-delta" => ModelEvent::TextDelta(required_str(value, "text")?.to_owned()),
|
||||
"text-end" => ModelEvent::TextEnd,
|
||||
"thinking-start" => ModelEvent::ThinkingStart,
|
||||
"thinking-delta" => ModelEvent::ThinkingDelta(required_str(value, "text")?.to_owned()),
|
||||
"thinking-end" => ModelEvent::ThinkingEnd,
|
||||
"tool-call-start" => ModelEvent::ToolCallStart {
|
||||
index: required_index(value)?,
|
||||
call_id: required_str(value, "callId")?.to_owned(),
|
||||
name: required_str(value, "name")?.to_owned(),
|
||||
},
|
||||
"tool-call-arguments-delta" => ModelEvent::ToolCallArgumentsDelta {
|
||||
index: required_index(value)?,
|
||||
delta: required_str(value, "delta")?.to_owned(),
|
||||
},
|
||||
"tool-call-end" => ModelEvent::ToolCallEnd {
|
||||
index: required_index(value)?,
|
||||
},
|
||||
"replay-state" => ModelEvent::ProviderReplayState(ProviderReplayState {
|
||||
provider_kind: required_str(value, "providerKind")?.to_owned(),
|
||||
value: value.get("value").cloned().unwrap_or_default(),
|
||||
}),
|
||||
"usage" => {
|
||||
let usage = value
|
||||
.get("usage")
|
||||
.ok_or_else(|| Error::Protocol("plugin usage event requires usage".into()))?;
|
||||
let tokens = |name: &str| usage.get(name).and_then(serde_json::Value::as_u64);
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: tokens("inputTokens"),
|
||||
output_tokens: tokens("outputTokens"),
|
||||
total_tokens: tokens("totalTokens"),
|
||||
cache_read_tokens: tokens("cacheReadTokens"),
|
||||
cache_write_tokens: tokens("cacheWriteTokens"),
|
||||
reasoning_tokens: tokens("reasoningTokens"),
|
||||
})
|
||||
}
|
||||
"done" => ModelEvent::Done(match required_str(value, "reason")? {
|
||||
"stop" => FinishReason::Stop,
|
||||
"length" => FinishReason::Length,
|
||||
"tool-use" => FinishReason::ToolUse,
|
||||
reason => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"unknown plugin finish reason: {reason}"
|
||||
)))
|
||||
}
|
||||
}),
|
||||
kind => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"unknown plugin model event: {kind}"
|
||||
)))
|
||||
}
|
||||
};
|
||||
Ok(event)
|
||||
}
|
||||
|
||||
fn required_str<'a>(value: &'a serde_json::Value, key: &str) -> Result<&'a str> {
|
||||
value
|
||||
.get(key)
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or_else(|| Error::Protocol(format!("plugin model event requires string '{key}'")))
|
||||
}
|
||||
|
||||
fn required_index(value: &serde_json::Value) -> Result<usize> {
|
||||
value
|
||||
.get("index")
|
||||
.and_then(serde_json::Value::as_u64)
|
||||
.map(|index| index as usize)
|
||||
.ok_or_else(|| Error::Protocol("plugin model event requires index".into()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::model::{
|
||||
ModelRequest, ModelSpec, ProjectedContent, PromptSpec, ToolCallContent, ToolResultContent,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn projects_history_into_wire_messages() {
|
||||
let invocation = ModelInvocation {
|
||||
call_id: "call".into(),
|
||||
run_id: "run".into(),
|
||||
conversation_id: "conversation".into(),
|
||||
provider_call_index: 0,
|
||||
request: ModelRequest {
|
||||
prompt: PromptSpec {
|
||||
instructions: "be brief".into(),
|
||||
tools: Vec::new(),
|
||||
},
|
||||
model: ModelSpec::new("plugin:p/c/m"),
|
||||
history: vec![
|
||||
ProjectedMessage {
|
||||
message_id: "m1".into(),
|
||||
role: Role::User,
|
||||
content: ProjectedContent::Parts(vec![ContentPart::Text {
|
||||
text: "hi".into(),
|
||||
}]),
|
||||
},
|
||||
ProjectedMessage {
|
||||
message_id: "m2".into(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text: "".into(),
|
||||
thinking: "t".into(),
|
||||
replay_state: Some(ProviderReplayState {
|
||||
provider_kind: "openai_responses".into(),
|
||||
value: serde_json::json!({"items": []}),
|
||||
}),
|
||||
calls: vec![ToolCallContent {
|
||||
index: 0,
|
||||
call_id: "c1".into(),
|
||||
name: "read".into(),
|
||||
arguments: serde_json::json!({"path":"a"}),
|
||||
}],
|
||||
},
|
||||
},
|
||||
ProjectedMessage {
|
||||
message_id: "m3".into(),
|
||||
role: Role::Tool,
|
||||
content: ProjectedContent::ToolResult(ToolResultContent {
|
||||
call_id: "c1".into(),
|
||||
name: "read".into(),
|
||||
content: "data".into(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: Vec::new(),
|
||||
}),
|
||||
},
|
||||
],
|
||||
},
|
||||
};
|
||||
let request = llm_request(&invocation).unwrap();
|
||||
assert_eq!(request["instructions"], "be brief");
|
||||
assert_eq!(request["latency"], "standard");
|
||||
assert_eq!(request["cacheKey"], "conversation");
|
||||
let messages = request["messages"].as_array().unwrap();
|
||||
assert_eq!(messages[0]["role"], "user");
|
||||
assert_eq!(
|
||||
messages[1]["replayState"]["providerKind"],
|
||||
"openai_responses"
|
||||
);
|
||||
assert_eq!(messages[1]["toolCalls"][0]["callId"], "c1");
|
||||
assert_eq!(messages[2]["role"], "tool");
|
||||
assert_eq!(messages[2]["isError"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_plugin_events_into_model_events() {
|
||||
assert_eq!(
|
||||
model_event(&serde_json::json!({"type":"text-delta","text":"hi"})).unwrap(),
|
||||
ModelEvent::TextDelta("hi".into())
|
||||
);
|
||||
assert_eq!(
|
||||
model_event(&serde_json::json!({"type":"done","reason":"tool-use"})).unwrap(),
|
||||
ModelEvent::Done(FinishReason::ToolUse)
|
||||
);
|
||||
let usage = model_event(&serde_json::json!({
|
||||
"type":"usage",
|
||||
"usage":{"inputTokens":10,"outputTokens":2,"cacheReadTokens":4}
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
usage,
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(10),
|
||||
output_tokens: Some(2),
|
||||
total_tokens: None,
|
||||
cache_read_tokens: Some(4),
|
||||
cache_write_tokens: None,
|
||||
reasoning_tokens: None,
|
||||
})
|
||||
);
|
||||
assert!(model_event(&serde_json::json!({"type":"mystery"})).is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,647 @@
|
||||
//! Runs one long-lived, sandboxed Deno process per active plugin.
|
||||
use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
path::PathBuf,
|
||||
process::Stdio,
|
||||
sync::Arc,
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use tokio::{
|
||||
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
|
||||
process::{Child, ChildStdin},
|
||||
sync::{mpsc, Mutex},
|
||||
};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use super::{
|
||||
catalog::PluginEntry,
|
||||
definition::{file_url, PluginDefinitionLoader},
|
||||
protocol::{HostMessage, WorkerMessage},
|
||||
};
|
||||
use crate::{store::Store, Error, Result};
|
||||
|
||||
const INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
||||
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
|
||||
const MAX_STREAM_BYTES: u64 = 256 * 1024 * 1024;
|
||||
|
||||
/// 一次流式调用的输出:零或多个事件,然后恰好一个最终结果。
|
||||
#[derive(Debug)]
|
||||
pub enum WorkerStreamItem {
|
||||
Event(serde_json::Value),
|
||||
Result(Result<serde_json::Value>),
|
||||
}
|
||||
|
||||
type Pending = Arc<Mutex<HashMap<String, mpsc::UnboundedSender<WorkerStreamItem>>>>;
|
||||
type StreamLines = Arc<Mutex<mpsc::Receiver<Result<String>>>>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PluginWorker {
|
||||
inner: Arc<PluginWorkerInner>,
|
||||
}
|
||||
|
||||
struct PluginWorkerInner {
|
||||
plugin_id: String,
|
||||
executable: PathBuf,
|
||||
directory: PathBuf,
|
||||
entry: PathBuf,
|
||||
loader: PluginDefinitionLoader,
|
||||
host: HostContext,
|
||||
process: Mutex<Option<WorkerProcess>>,
|
||||
pending: Pending,
|
||||
}
|
||||
|
||||
struct WorkerProcess {
|
||||
child: Child,
|
||||
stdin: Arc<Mutex<ChildStdin>>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct HostContext {
|
||||
plugin_id: String,
|
||||
network_hosts: Arc<HashSet<String>>,
|
||||
store: Store,
|
||||
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
|
||||
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
|
||||
}
|
||||
|
||||
impl PluginWorker {
|
||||
pub fn new(
|
||||
plugin: &PluginEntry,
|
||||
executable: PathBuf,
|
||||
loader: PluginDefinitionLoader,
|
||||
store: Store,
|
||||
) -> Self {
|
||||
let plugin_id = plugin.manifest.id.clone();
|
||||
Self {
|
||||
inner: Arc::new(PluginWorkerInner {
|
||||
host: HostContext {
|
||||
plugin_id: plugin_id.clone(),
|
||||
network_hosts: Arc::new(
|
||||
plugin
|
||||
.manifest
|
||||
.permissions
|
||||
.network
|
||||
.iter()
|
||||
.map(|host| host.to_ascii_lowercase())
|
||||
.collect(),
|
||||
),
|
||||
store,
|
||||
cancellations: Arc::new(Mutex::new(HashMap::new())),
|
||||
streams: Arc::new(Mutex::new(HashMap::new())),
|
||||
},
|
||||
plugin_id,
|
||||
executable,
|
||||
directory: plugin.directory.clone(),
|
||||
entry: plugin.entry.clone(),
|
||||
loader,
|
||||
process: Mutex::new(None),
|
||||
pending: Arc::new(Mutex::new(HashMap::new())),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
/// 一元调用:忽略事件,等待最终结果,受统一超时约束。
|
||||
pub async fn invoke(
|
||||
&self,
|
||||
method: &str,
|
||||
params: serde_json::Value,
|
||||
cancellation: CancellationToken,
|
||||
) -> Result<serde_json::Value> {
|
||||
let mut items = self.invoke_streaming(method, params, cancellation).await?;
|
||||
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
|
||||
while let Some(item) = items.recv().await {
|
||||
if let WorkerStreamItem::Result(result) = item {
|
||||
return result;
|
||||
}
|
||||
}
|
||||
Err(Error::Provider(format!(
|
||||
"plugin '{}' worker stopped",
|
||||
self.inner.plugin_id
|
||||
)))
|
||||
})
|
||||
.await;
|
||||
match result {
|
||||
Ok(result) => result,
|
||||
Err(_) => Err(Error::Provider(format!(
|
||||
"plugin '{}' invocation timed out",
|
||||
self.inner.plugin_id
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
/// 流式调用:事件按序转发,最终以恰好一个 Result 收尾。
|
||||
/// 取消通过传入的令牌传播到 Worker 与其挂起的宿主网络请求。
|
||||
pub async fn invoke_streaming(
|
||||
&self,
|
||||
method: &str,
|
||||
params: serde_json::Value,
|
||||
cancellation: CancellationToken,
|
||||
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
|
||||
let id = uuid::Uuid::new_v4().to_string();
|
||||
let request_cancellation = CancellationToken::new();
|
||||
self.inner
|
||||
.host
|
||||
.cancellations
|
||||
.lock()
|
||||
.await
|
||||
.insert(id.clone(), request_cancellation.clone());
|
||||
let (sender, receiver) = mpsc::unbounded_channel();
|
||||
self.inner
|
||||
.pending
|
||||
.lock()
|
||||
.await
|
||||
.insert(id.clone(), sender.clone());
|
||||
let send_result = async {
|
||||
let stdin = self.stdin().await?;
|
||||
write_message(
|
||||
&stdin,
|
||||
&HostMessage::Request {
|
||||
id: &id,
|
||||
method,
|
||||
params: ¶ms,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
.await;
|
||||
if let Err(error) = send_result {
|
||||
self.cleanup(&id).await;
|
||||
return Err(error);
|
||||
}
|
||||
|
||||
// 取消监视:通知 Worker,同时中止该请求挂起的宿主网络调用。
|
||||
let inner = self.inner.clone();
|
||||
let request_id = id.clone();
|
||||
tokio::spawn(async move {
|
||||
tokio::select! {
|
||||
_ = cancellation.cancelled() => {
|
||||
request_cancellation.cancel();
|
||||
if let Some(process) = inner.process.lock().await.as_ref() {
|
||||
let _ = write_message(&process.stdin, &HostMessage::Cancel { id: &request_id }).await;
|
||||
}
|
||||
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
|
||||
inner.pending.lock().await.remove(&request_id);
|
||||
inner.host.cancellations.lock().await.remove(&request_id);
|
||||
}
|
||||
_ = sender.closed() => {
|
||||
inner.host.cancellations.lock().await.remove(&request_id);
|
||||
}
|
||||
}
|
||||
});
|
||||
Ok(receiver)
|
||||
}
|
||||
|
||||
pub async fn stop(&self) {
|
||||
if let Some(mut process) = self.inner.process.lock().await.take() {
|
||||
let _ = process.child.kill().await;
|
||||
}
|
||||
fail_pending(&self.inner.pending, "plugin worker stopped").await;
|
||||
}
|
||||
|
||||
async fn cleanup(&self, id: &str) {
|
||||
self.inner.pending.lock().await.remove(id);
|
||||
self.inner.host.cancellations.lock().await.remove(id);
|
||||
}
|
||||
|
||||
async fn stdin(&self) -> Result<Arc<Mutex<ChildStdin>>> {
|
||||
let mut process = self.inner.process.lock().await;
|
||||
let dead = match process.as_mut() {
|
||||
Some(current) => current
|
||||
.child
|
||||
.try_wait()
|
||||
.map_err(|error| {
|
||||
Error::Config(format!("cannot check plugin worker status: {error}"))
|
||||
})?
|
||||
.is_some(),
|
||||
None => true,
|
||||
};
|
||||
if dead {
|
||||
*process = Some(self.spawn().await?);
|
||||
}
|
||||
Ok(process
|
||||
.as_ref()
|
||||
.expect("plugin worker was started")
|
||||
.stdin
|
||||
.clone())
|
||||
}
|
||||
|
||||
async fn spawn(&self) -> Result<WorkerProcess> {
|
||||
let entry_url = file_url(&self.inner.entry)?;
|
||||
let mut command = tokio::process::Command::new(&self.inner.executable);
|
||||
super::detach_console(&mut command);
|
||||
command
|
||||
.arg("run")
|
||||
.arg("--quiet")
|
||||
.arg("--no-config")
|
||||
.arg("--no-lock")
|
||||
.arg("--no-npm")
|
||||
.arg("--no-remote")
|
||||
.arg("--no-prompt")
|
||||
.arg(format!("--allow-read={}", self.inner.directory.display()))
|
||||
.arg(format!(
|
||||
"--allow-read={}",
|
||||
self.inner.loader.sdk_dir().display()
|
||||
))
|
||||
.arg(format!(
|
||||
"--import-map={}",
|
||||
self.inner.loader.import_map().display()
|
||||
))
|
||||
.arg(self.inner.loader.worker_path())
|
||||
.arg(entry_url.as_str())
|
||||
.env("DENO_DIR", self.inner.loader.deno_dir())
|
||||
.env("DENO_NO_UPDATE_CHECK", "1")
|
||||
.current_dir(&self.inner.directory)
|
||||
.stdin(Stdio::piped())
|
||||
.stdout(Stdio::piped())
|
||||
.stderr(Stdio::piped())
|
||||
.kill_on_drop(true);
|
||||
let mut child = command.spawn().map_err(|error| {
|
||||
Error::Config(format!(
|
||||
"cannot start plugin worker {}: {error}",
|
||||
self.inner.executable.display()
|
||||
))
|
||||
})?;
|
||||
let stdin =
|
||||
Arc::new(Mutex::new(child.stdin.take().ok_or_else(|| {
|
||||
Error::Config("cannot open plugin worker stdin".into())
|
||||
})?));
|
||||
let stdout = child
|
||||
.stdout
|
||||
.take()
|
||||
.ok_or_else(|| Error::Config("cannot open plugin worker stdout".into()))?;
|
||||
let stderr = child
|
||||
.stderr
|
||||
.take()
|
||||
.ok_or_else(|| Error::Config("cannot open plugin worker stderr".into()))?;
|
||||
spawn_stdout_reader(
|
||||
self.inner.plugin_id.clone(),
|
||||
stdout,
|
||||
stdin.clone(),
|
||||
self.inner.pending.clone(),
|
||||
self.inner.host.clone(),
|
||||
);
|
||||
spawn_stderr_reader(self.inner.plugin_id.clone(), stderr);
|
||||
Ok(WorkerProcess { child, stdin })
|
||||
}
|
||||
}
|
||||
|
||||
fn spawn_stdout_reader(
|
||||
plugin_id: String,
|
||||
stdout: tokio::process::ChildStdout,
|
||||
stdin: Arc<Mutex<ChildStdin>>,
|
||||
pending: Pending,
|
||||
host: HostContext,
|
||||
) {
|
||||
tokio::spawn(async move {
|
||||
let mut lines = BufReader::new(stdout).lines();
|
||||
while let Ok(Some(line)) = lines.next_line().await {
|
||||
let message = match serde_json::from_str::<WorkerMessage>(&line) {
|
||||
Ok(message) => message,
|
||||
Err(error) => {
|
||||
tracing::warn!(plugin = %plugin_id, %error, "plugin worker wrote an invalid message");
|
||||
continue;
|
||||
}
|
||||
};
|
||||
match message {
|
||||
WorkerMessage::Result { id, result, error } => {
|
||||
if let Some(sender) = pending.lock().await.remove(&id) {
|
||||
let value = match error {
|
||||
Some(error) => {
|
||||
Err(Error::Provider(format!("plugin '{plugin_id}': {error}")))
|
||||
}
|
||||
None => Ok(result),
|
||||
};
|
||||
let _ = sender.send(WorkerStreamItem::Result(value));
|
||||
}
|
||||
}
|
||||
WorkerMessage::Event { id, event } => {
|
||||
if let Some(sender) = pending.lock().await.get(&id) {
|
||||
let _ = sender.send(WorkerStreamItem::Event(event));
|
||||
}
|
||||
}
|
||||
WorkerMessage::HostCall {
|
||||
id,
|
||||
request_id,
|
||||
method,
|
||||
params,
|
||||
} => {
|
||||
let host = host.clone();
|
||||
let stdin = stdin.clone();
|
||||
tokio::spawn(async move {
|
||||
let result = host.call(&request_id, &method, params).await;
|
||||
match result {
|
||||
Ok(result) => {
|
||||
let _ = write_message(
|
||||
&stdin,
|
||||
&HostMessage::HostResult {
|
||||
id: &id,
|
||||
result: &result,
|
||||
},
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Err(error) => {
|
||||
let text = error.to_string();
|
||||
let _ = write_message(
|
||||
&stdin,
|
||||
&HostMessage::HostError {
|
||||
id: &id,
|
||||
error: &text,
|
||||
},
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
fail_pending(&pending, &format!("plugin '{plugin_id}' worker exited")).await;
|
||||
});
|
||||
}
|
||||
|
||||
fn spawn_stderr_reader(plugin_id: String, stderr: tokio::process::ChildStderr) {
|
||||
tokio::spawn(async move {
|
||||
let mut lines = BufReader::new(stderr).lines();
|
||||
while let Ok(Some(line)) = lines.next_line().await {
|
||||
tracing::warn!(plugin = %plugin_id, message = %line, "plugin worker stderr");
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
async fn write_message(stdin: &Arc<Mutex<ChildStdin>>, message: &HostMessage<'_>) -> Result<()> {
|
||||
let mut bytes = serde_json::to_vec(message)?;
|
||||
bytes.push(b'\n');
|
||||
let mut stdin = stdin.lock().await;
|
||||
stdin
|
||||
.write_all(&bytes)
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("cannot write to plugin worker: {error}")))?;
|
||||
stdin
|
||||
.flush()
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("cannot flush plugin worker stdin: {error}")))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn fail_pending(pending: &Pending, message: &str) {
|
||||
for (_, sender) in std::mem::take(&mut *pending.lock().await) {
|
||||
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Provider(
|
||||
message.into(),
|
||||
))));
|
||||
}
|
||||
}
|
||||
|
||||
impl HostContext {
|
||||
async fn call(
|
||||
&self,
|
||||
request_id: &str,
|
||||
method: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value> {
|
||||
match method {
|
||||
"network.fetch" => self.fetch(request_id, params).await,
|
||||
"network.stream.open" => self.stream_open(request_id, params).await,
|
||||
"network.stream.read" => self.stream_read(params).await,
|
||||
"network.stream.close" => {
|
||||
self.streams
|
||||
.lock()
|
||||
.await
|
||||
.remove(required_string(¶ms, "streamId")?);
|
||||
Ok(serde_json::Value::Null)
|
||||
}
|
||||
_ => Err(Error::Protocol(format!(
|
||||
"unsupported plugin host method: {method}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
async fn request(
|
||||
&self,
|
||||
request_id: &str,
|
||||
params: &serde_json::Value,
|
||||
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
|
||||
let raw_url = required_string(params, "url")?;
|
||||
let url = url::Url::parse(raw_url)
|
||||
.map_err(|error| Error::Config(format!("invalid plugin network URL: {error}")))?;
|
||||
if url.scheme() != "https" || !url.username().is_empty() || url.password().is_some() {
|
||||
return Err(Error::Config(
|
||||
"plugin network URL must be HTTPS without credentials".into(),
|
||||
));
|
||||
}
|
||||
let host = url
|
||||
.host_str()
|
||||
.ok_or_else(|| Error::Config("plugin network URL has no host".into()))?
|
||||
.to_ascii_lowercase();
|
||||
if !self.network_hosts.contains(&host) {
|
||||
return Err(Error::Config(format!(
|
||||
"plugin '{}' cannot access host '{host}'",
|
||||
self.plugin_id
|
||||
)));
|
||||
}
|
||||
let method = params
|
||||
.get("method")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("GET")
|
||||
.parse::<reqwest::Method>()
|
||||
.map_err(|error| Error::Config(format!("invalid plugin HTTP method: {error}")))?;
|
||||
let client = crate::network::client_builder(&self.store)
|
||||
.await?
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.connect_timeout(Duration::from_secs(30))
|
||||
.build()?;
|
||||
let mut request = client.request(method, url);
|
||||
if let Some(headers) = params.get("headers").and_then(serde_json::Value::as_object) {
|
||||
for (name, value) in headers {
|
||||
let value = value.as_str().ok_or_else(|| {
|
||||
Error::Config(format!("plugin HTTP header '{name}' must be a string"))
|
||||
})?;
|
||||
request = request.header(name, value);
|
||||
}
|
||||
}
|
||||
if let Some(body) = params.get("body").and_then(serde_json::Value::as_str) {
|
||||
request = request.body(body.to_owned());
|
||||
}
|
||||
let cancellation = self
|
||||
.cancellations
|
||||
.lock()
|
||||
.await
|
||||
.get(request_id)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
Ok((request, cancellation))
|
||||
}
|
||||
|
||||
async fn fetch(
|
||||
&self,
|
||||
request_id: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value> {
|
||||
let (request, cancellation) = self.request(request_id, ¶ms).await?;
|
||||
let request = request.timeout(Duration::from_secs(60));
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
if response
|
||||
.content_length()
|
||||
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
|
||||
{
|
||||
return Err(Error::Provider(
|
||||
"plugin network response is larger than allowed".into(),
|
||||
));
|
||||
}
|
||||
let headers = header_map(&response);
|
||||
let body = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
body = response.bytes() => body?,
|
||||
};
|
||||
if body.len() as u64 > MAX_NETWORK_RESPONSE_BYTES {
|
||||
return Err(Error::Provider(
|
||||
"plugin network response is larger than allowed".into(),
|
||||
));
|
||||
}
|
||||
Ok(
|
||||
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
|
||||
)
|
||||
}
|
||||
|
||||
/// 打开流式响应:立即返回状态与响应头,响应体按行经 stream.read 拉取。
|
||||
async fn stream_open(
|
||||
&self,
|
||||
request_id: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value> {
|
||||
let (request, cancellation) = self.request(request_id, ¶ms).await?;
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
let headers = header_map(&response);
|
||||
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
|
||||
tokio::spawn(async move {
|
||||
use futures_util::StreamExt;
|
||||
let mut body = response.bytes_stream();
|
||||
let mut buffered = Vec::<u8>::new();
|
||||
let mut total = 0_u64;
|
||||
loop {
|
||||
let chunk = tokio::select! {
|
||||
_ = cancellation.cancelled() => {
|
||||
let _ = sender.send(Err(Error::Cancelled)).await;
|
||||
return;
|
||||
}
|
||||
chunk = body.next() => chunk,
|
||||
};
|
||||
let Some(chunk) = chunk else { break };
|
||||
let chunk = match chunk {
|
||||
Ok(chunk) => chunk,
|
||||
Err(error) => {
|
||||
let _ = sender.send(Err(Error::from(error))).await;
|
||||
return;
|
||||
}
|
||||
};
|
||||
total += chunk.len() as u64;
|
||||
if total > MAX_STREAM_BYTES {
|
||||
let _ = sender
|
||||
.send(Err(Error::Provider(
|
||||
"plugin network stream is larger than allowed".into(),
|
||||
)))
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
buffered.extend_from_slice(&chunk);
|
||||
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
|
||||
let mut line = buffered.drain(..=position).collect::<Vec<u8>>();
|
||||
line.pop();
|
||||
if line.last() == Some(&b'\r') {
|
||||
line.pop();
|
||||
}
|
||||
if sender
|
||||
.send(Ok(String::from_utf8_lossy(&line).into_owned()))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
if !buffered.is_empty() {
|
||||
let _ = sender
|
||||
.send(Ok(String::from_utf8_lossy(&buffered).into_owned()))
|
||||
.await;
|
||||
}
|
||||
});
|
||||
let stream_id = uuid::Uuid::new_v4().to_string();
|
||||
self.streams
|
||||
.lock()
|
||||
.await
|
||||
.insert(stream_id.clone(), Arc::new(Mutex::new(receiver)));
|
||||
Ok(serde_json::json!({
|
||||
"streamId": stream_id,
|
||||
"status": status,
|
||||
"headers": headers,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn stream_read(&self, params: serde_json::Value) -> Result<serde_json::Value> {
|
||||
let stream_id = required_string(¶ms, "streamId")?;
|
||||
let lines_handle = self
|
||||
.streams
|
||||
.lock()
|
||||
.await
|
||||
.get(stream_id)
|
||||
.cloned()
|
||||
.ok_or_else(|| Error::Protocol(format!("unknown plugin stream: {stream_id}")))?;
|
||||
let mut receiver = lines_handle.lock().await;
|
||||
let mut lines = Vec::new();
|
||||
match receiver.recv().await {
|
||||
Some(Ok(line)) => lines.push(line),
|
||||
Some(Err(error)) => {
|
||||
drop(receiver);
|
||||
self.streams.lock().await.remove(stream_id);
|
||||
return Err(error);
|
||||
}
|
||||
None => {
|
||||
drop(receiver);
|
||||
self.streams.lock().await.remove(stream_id);
|
||||
return Ok(serde_json::json!({ "lines": [], "done": true }));
|
||||
}
|
||||
}
|
||||
// 把已就绪的行一并带走,减少往返。
|
||||
while lines.len() < 256 {
|
||||
match receiver.try_recv() {
|
||||
Ok(Ok(line)) => lines.push(line),
|
||||
Ok(Err(error)) => {
|
||||
drop(receiver);
|
||||
self.streams.lock().await.remove(stream_id);
|
||||
return Err(error);
|
||||
}
|
||||
Err(_) => break,
|
||||
}
|
||||
}
|
||||
Ok(serde_json::json!({ "lines": lines, "done": false }))
|
||||
}
|
||||
}
|
||||
|
||||
fn header_map(response: &reqwest::Response) -> std::collections::BTreeMap<String, String> {
|
||||
response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| {
|
||||
value
|
||||
.to_str()
|
||||
.ok()
|
||||
.map(|value| (name.to_string(), value.to_string()))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn required_string<'a>(params: &'a serde_json::Value, key: &str) -> Result<&'a str> {
|
||||
params
|
||||
.get(key)
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
|
||||
}
|
||||
@@ -12,7 +12,7 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
merge_extra_params,
|
||||
map_sse_error, merge_extra_params, provider_event_error,
|
||||
recorder::recorded_headers,
|
||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||
@@ -96,7 +96,7 @@ impl Provider for AnthropicProvider {
|
||||
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
|
||||
.headers(config.custom_headers.clone())
|
||||
.json(&body),
|
||||
RetryPolicy::default(),
|
||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
||||
&cancellation,
|
||||
recorder.as_ref(),
|
||||
request_headers,
|
||||
@@ -129,8 +129,11 @@ impl Provider for AnthropicProvider {
|
||||
_ = cancellation.cancelled() => { return; }
|
||||
event = source.next() => event,
|
||||
} {
|
||||
let event = event.map_err(|error| Error::Provider(format!("Anthropic SSE: {error}")))?;
|
||||
let event = event.map_err(|error| map_sse_error("Anthropic", error))?;
|
||||
let value: Value = serde_json::from_str(&event.data)?;
|
||||
if let Some(error) = provider_event_error("Anthropic", &value) {
|
||||
Err(error)?;
|
||||
}
|
||||
let data_kind = value.get("type").and_then(Value::as_str);
|
||||
let kind = match event.event.as_str() {
|
||||
"" | "message" => data_kind.unwrap_or(event.event.as_str()),
|
||||
@@ -242,7 +245,6 @@ impl Provider for AnthropicProvider {
|
||||
};
|
||||
yield ModelEvent::Done(finish);
|
||||
}
|
||||
"error" => Err(Error::Provider(format!("Anthropic stream error: {}", event.data)))?,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -32,6 +32,52 @@ pub trait Provider: Send + Sync {
|
||||
) -> ProviderStream;
|
||||
}
|
||||
|
||||
fn map_sse_error(
|
||||
label: &str,
|
||||
error: eventsource_stream::EventStreamError<crate::Error>,
|
||||
) -> crate::Error {
|
||||
match error {
|
||||
eventsource_stream::EventStreamError::Transport(error) => error,
|
||||
eventsource_stream::EventStreamError::Utf8(error) => {
|
||||
crate::Error::Provider(format!("{label} SSE UTF-8 error: {error}"))
|
||||
}
|
||||
eventsource_stream::EventStreamError::Parser(error) => {
|
||||
crate::Error::Provider(format!("{label} SSE parse error: {error}"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_event_error(label: &str, value: &serde_json::Value) -> Option<crate::Error> {
|
||||
let kind = value.get("type").and_then(serde_json::Value::as_str);
|
||||
let direct_error = value.get("error").filter(|error| !error.is_null());
|
||||
if !matches!(kind, Some("error" | "response.failed")) && direct_error.is_none() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let message = value
|
||||
.get("message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.or_else(|| {
|
||||
value
|
||||
.pointer("/error/message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
})
|
||||
.or_else(|| {
|
||||
value
|
||||
.pointer("/response/error/message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
})
|
||||
.or_else(|| direct_error.and_then(serde_json::Value::as_str))
|
||||
.or_else(|| {
|
||||
value
|
||||
.pointer("/response/error")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
})
|
||||
.unwrap_or("provider returned an error event without a message");
|
||||
|
||||
Some(crate::Error::Provider(format!("{label} error: {message}")))
|
||||
}
|
||||
|
||||
fn merge_extra_params(body: &mut serde_json::Value, extra: &serde_json::Value) -> Result<()> {
|
||||
let extra = extra
|
||||
.as_object()
|
||||
@@ -60,6 +106,19 @@ fn merge_extra_params(body: &mut serde_json::Value, extra: &serde_json::Value) -
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn apply_body_allowlist(
|
||||
body: &mut serde_json::Value,
|
||||
allowed: Option<&std::collections::HashSet<String>>,
|
||||
) -> Result<()> {
|
||||
let Some(allowed) = allowed else {
|
||||
return Ok(());
|
||||
};
|
||||
body.as_object_mut()
|
||||
.ok_or_else(|| crate::Error::Provider("provider request body must be an object".into()))?
|
||||
.retain(|name, _| allowed.contains(name));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn apply_openai_prompt_cache_key(body: &mut serde_json::Value, model_id: &str) -> Result<()> {
|
||||
if !model_id.to_ascii_lowercase().contains("gpt") {
|
||||
return Ok(());
|
||||
@@ -72,3 +131,74 @@ fn apply_openai_prompt_cache_key(body: &mut serde_json::Value, model_id: &str) -
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn sse_transport_errors_are_not_relabelled_as_parse_errors() {
|
||||
let error = map_sse_error(
|
||||
"test provider",
|
||||
eventsource_stream::EventStreamError::Transport(crate::Error::Provider(
|
||||
"connection closed".into(),
|
||||
)),
|
||||
);
|
||||
|
||||
let crate::Error::Provider(message) = error else {
|
||||
panic!("transport error category must be preserved");
|
||||
};
|
||||
assert_eq!(message, "connection closed");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_error_events_extract_flat_and_nested_messages() {
|
||||
assert_provider_error(
|
||||
"OpenAI Responses",
|
||||
serde_json::json!({
|
||||
"type": "error",
|
||||
"message": "Internal error during token generation"
|
||||
}),
|
||||
"OpenAI Responses error: Internal error during token generation",
|
||||
);
|
||||
assert_provider_error(
|
||||
"OpenAI Chat",
|
||||
serde_json::json!({
|
||||
"error": {"message": "quota exceeded", "type": "server_error"}
|
||||
}),
|
||||
"OpenAI Chat error: quota exceeded",
|
||||
);
|
||||
assert_provider_error(
|
||||
"Anthropic",
|
||||
serde_json::json!({
|
||||
"type": "error",
|
||||
"error": {"type": "overloaded_error", "message": "Overloaded"}
|
||||
}),
|
||||
"Anthropic error: Overloaded",
|
||||
);
|
||||
assert_provider_error(
|
||||
"OpenAI Responses",
|
||||
serde_json::json!({
|
||||
"type": "response.failed",
|
||||
"response": {"error": {"message": "generation failed"}}
|
||||
}),
|
||||
"OpenAI Responses error: generation failed",
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn successful_provider_events_are_not_errors() {
|
||||
assert!(provider_event_error(
|
||||
"OpenAI Responses",
|
||||
&serde_json::json!({"type": "response.completed", "error": null})
|
||||
)
|
||||
.is_none());
|
||||
}
|
||||
|
||||
fn assert_provider_error(label: &str, value: serde_json::Value, expected: &str) {
|
||||
let Some(crate::Error::Provider(message)) = provider_event_error(label, &value) else {
|
||||
panic!("expected provider error");
|
||||
};
|
||||
assert_eq!(message, expected);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,7 +17,8 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
apply_openai_prompt_cache_key, merge_extra_params,
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
||||
provider_event_error,
|
||||
recorder::recorded_headers,
|
||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||
@@ -86,6 +87,7 @@ impl Provider for OpenAiChatProvider {
|
||||
apply_model(&mut body, &request.model, config.max_output_tokens)?;
|
||||
merge_extra_params(&mut body, &request.model.extra_params)?;
|
||||
apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?;
|
||||
apply_body_allowlist(&mut body, config.allowed_body_fields.as_ref())?;
|
||||
let request_headers = recorded_headers(&config, &[("content-type", "application/json")]);
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.request(request_headers.clone(), &body).await?;
|
||||
@@ -94,7 +96,7 @@ impl Provider for OpenAiChatProvider {
|
||||
"OpenAI Chat",
|
||||
|| client.post(&config.request_url)
|
||||
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
||||
RetryPolicy::default(),
|
||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
||||
&cancellation,
|
||||
recorder.as_ref(),
|
||||
request_headers,
|
||||
@@ -146,12 +148,14 @@ impl Provider for OpenAiChatProvider {
|
||||
break;
|
||||
};
|
||||
let event = event.map_err(|error| {
|
||||
let err_msg = error.to_string();
|
||||
tracing::debug!(iteration = loop_iteration, error = %error, "OpenAI Chat SSE event failed");
|
||||
Error::Provider(format!("OpenAI Chat SSE: {err_msg}"))
|
||||
map_sse_error("OpenAI Chat", error)
|
||||
})?;
|
||||
if event.data == "[DONE]" { saw_done_marker = true; break; }
|
||||
let value: Value = serde_json::from_str(&event.data)?;
|
||||
if let Some(error) = provider_event_error("OpenAI Chat", &value) {
|
||||
Err(error)?;
|
||||
}
|
||||
if let Some(usage) = value.get("usage").filter(|value| !value.is_null()) {
|
||||
final_usage = Some(openai_usage(usage));
|
||||
}
|
||||
|
||||
@@ -14,7 +14,8 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
apply_openai_prompt_cache_key, merge_extra_params,
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
||||
provider_event_error,
|
||||
recorder::recorded_headers,
|
||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||
@@ -83,6 +84,7 @@ impl Provider for OpenAiResponsesProvider {
|
||||
apply_model(&mut body, &request.model, config.max_output_tokens)?;
|
||||
merge_extra_params(&mut body, &request.model.extra_params)?;
|
||||
apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?;
|
||||
apply_body_allowlist(&mut body, config.allowed_body_fields.as_ref())?;
|
||||
let request_headers = recorded_headers(&config, &[("content-type", "application/json")]);
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.request(request_headers.clone(), &body).await?;
|
||||
@@ -91,7 +93,7 @@ impl Provider for OpenAiResponsesProvider {
|
||||
"OpenAI Responses",
|
||||
|| client.post(&config.request_url)
|
||||
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
||||
RetryPolicy::default(),
|
||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
||||
&cancellation,
|
||||
recorder.as_ref(),
|
||||
request_headers,
|
||||
@@ -126,9 +128,12 @@ impl Provider for OpenAiResponsesProvider {
|
||||
event = source.next() => event,
|
||||
};
|
||||
let Some(event) = event else { break };
|
||||
let event = event.map_err(|error| Error::Provider(format!("OpenAI Responses SSE: {error}")))?;
|
||||
let event = event.map_err(|error| map_sse_error("OpenAI Responses", error))?;
|
||||
if event.data == "[DONE]" { break; }
|
||||
let value: Value = serde_json::from_str(&event.data)?;
|
||||
if let Some(error) = provider_event_error("OpenAI Responses", &value) {
|
||||
Err(error)?;
|
||||
}
|
||||
let kind = value.get("type").and_then(Value::as_str).unwrap_or(&event.event);
|
||||
match kind {
|
||||
"response.output_text.delta" => {
|
||||
@@ -247,7 +252,6 @@ impl Provider for OpenAiResponsesProvider {
|
||||
terminal = true;
|
||||
yield ModelEvent::Done(FinishReason::Length);
|
||||
}
|
||||
"response.failed" => Err(Error::Provider(format!("OpenAI Responses failed: {}", event.data)))?,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -423,6 +423,7 @@ mod tests {
|
||||
.create_model(&ModelConfigInput {
|
||||
sort_order: 0,
|
||||
display_name: "Test Model".into(),
|
||||
group_name: None,
|
||||
model_type: ModelType::OpenAi,
|
||||
base_url: "https://example.com/v1/chat/completions".into(),
|
||||
use_full_url: true,
|
||||
|
||||
@@ -60,9 +60,6 @@ where
|
||||
String::from_utf8_lossy(&bytes)
|
||||
));
|
||||
if attempt == policy.retries {
|
||||
if let Some(recorder) = recorder {
|
||||
recorder.failed(&error).await?;
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
tracing::warn!(
|
||||
|
||||
+248
-90
@@ -1,4 +1,4 @@
|
||||
//! Routes model requests to the configured provider.
|
||||
//! Routes model requests to built-in configurations or stable plugin model IDs.
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use async_stream::try_stream;
|
||||
@@ -8,25 +8,37 @@ use tokio_util::sync::CancellationToken;
|
||||
use crate::{
|
||||
config::{ProviderConfig, ProviderKind},
|
||||
model::{ModelInvocation, ModelLatency, NewLlmCall, ProviderType},
|
||||
plugin::{PluginRegistry, ADAPTER_ID_PREFIX},
|
||||
store::Store,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{
|
||||
normalize::NormalizedProvider, AnthropicProvider, CallRecorder, OpenAiChatProvider,
|
||||
OpenAiResponsesProvider, Provider, ProviderStream,
|
||||
normalize::NormalizedProvider, recorder::CancelOnDrop, AnthropicProvider, CallRecorder,
|
||||
OpenAiChatProvider, OpenAiResponsesProvider, Provider, ProviderStream,
|
||||
};
|
||||
|
||||
const BUILTIN_PROVIDER_RETRIES: u32 = 5;
|
||||
|
||||
pub struct ProviderRouter {
|
||||
store: Store,
|
||||
plugins: PluginRegistry,
|
||||
request_timeout: Duration,
|
||||
stream_idle_timeout: Duration,
|
||||
}
|
||||
|
||||
impl ProviderRouter {
|
||||
pub fn new(store: Store, request_timeout: Duration) -> Self {
|
||||
pub fn new(
|
||||
store: Store,
|
||||
plugins: PluginRegistry,
|
||||
request_timeout: Duration,
|
||||
stream_idle_timeout: Duration,
|
||||
) -> Self {
|
||||
Self {
|
||||
store,
|
||||
plugins,
|
||||
request_timeout,
|
||||
stream_idle_timeout,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -34,100 +46,96 @@ impl ProviderRouter {
|
||||
impl Provider for ProviderRouter {
|
||||
fn stream(
|
||||
&self,
|
||||
mut invocation: ModelInvocation,
|
||||
invocation: ModelInvocation,
|
||||
cancellation: CancellationToken,
|
||||
) -> ProviderStream {
|
||||
let store = self.store.clone();
|
||||
let plugins = self.plugins.clone();
|
||||
let request_timeout = self.request_timeout;
|
||||
let stream_idle_timeout = self.stream_idle_timeout;
|
||||
Box::pin(try_stream! {
|
||||
let selected = invocation.request.model.model_id.clone();
|
||||
let model = store
|
||||
.model(&selected)
|
||||
.await?
|
||||
.ok_or_else(|| Error::Provider(format!("unknown model: {selected}")))?;
|
||||
let provider_type = model.provider_type();
|
||||
let request_url = model.request_url()?;
|
||||
model.configure(&mut invocation.request.model);
|
||||
invocation.request.model.extra_params = model.extra_params().clone();
|
||||
invocation.request.model.model_id = model.model_id.clone();
|
||||
let recorder = CallRecorder::start(store.clone(), NewLlmCall {
|
||||
call_id: invocation.call_id.clone(),
|
||||
run_id: invocation.run_id.clone(),
|
||||
conversation_id: invocation.conversation_id.clone(),
|
||||
provider_call_index: invocation.provider_call_index.min(i64::MAX as u64) as i64,
|
||||
model_hash: model.model_hash.clone(),
|
||||
provider_type,
|
||||
provider_url: model.base_url.clone(),
|
||||
request_type: provider_type,
|
||||
request_url: request_url.clone(),
|
||||
model_id: model.model_id.clone(),
|
||||
display_name: model.display_name.clone(),
|
||||
reasoning_effort: invocation.request.model.reasoning.effort.clone(),
|
||||
fast: invocation.request.model.latency == ModelLatency::Fast,
|
||||
message_count: invocation.request.history.len(),
|
||||
tool_count: invocation.request.prompt.tools.len(),
|
||||
detailed: false,
|
||||
}).await?;
|
||||
let _cancel_on_drop = recorder.cancel_on_drop();
|
||||
let config = ProviderConfig {
|
||||
kind: match provider_type {
|
||||
ProviderType::OpenAiChat => ProviderKind::OpenAiChat,
|
||||
ProviderType::OpenAiResponses => ProviderKind::OpenAiResponses,
|
||||
ProviderType::Anthropic => ProviderKind::Anthropic,
|
||||
},
|
||||
request_url,
|
||||
api_key: model.api_key.clone(),
|
||||
custom_headers: if model.custom_headers_enabled {
|
||||
custom_headers(&model.custom_headers)?
|
||||
// 两条分支只负责装配 Recorder 与 Provider 流;
|
||||
// 事件消费(空闲超时看门狗、记录、错误规范化)对两者完全一致。
|
||||
let (recorder, _cancel_on_drop, mut stream): (CallRecorder, CancelOnDrop, ProviderStream) =
|
||||
if selected.starts_with(ADAPTER_ID_PREFIX) {
|
||||
// 插件模型与内置模型走完全相同的流程:资源选择与将来的
|
||||
// 负载均衡都在插件 Provider 内部。
|
||||
let plan = plugins.plan_model(&selected).await?;
|
||||
let recorder = start_recorder(&store, &invocation, &selected, &plan.model.display_name, ProviderType::Plugin, &plan.request_url, &plan.model.model_id).await?;
|
||||
let guard = recorder.cancel_on_drop();
|
||||
recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?;
|
||||
let mut routed = invocation.clone();
|
||||
routed.request.model.display_name = Some(plan.model.display_name.clone());
|
||||
if let Some(tokens) = plan.model.max_output_tokens {
|
||||
routed.request.model.max_output_tokens.get_or_insert(tokens);
|
||||
}
|
||||
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
|
||||
registry: plugins.clone(),
|
||||
})));
|
||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||
} else {
|
||||
reqwest::header::HeaderMap::new()
|
||||
},
|
||||
max_output_tokens: model.max_output_tokens(),
|
||||
request_timeout,
|
||||
};
|
||||
let client = crate::network::client_builder(&store)
|
||||
.await?
|
||||
.timeout(config.request_timeout)
|
||||
.build()?;
|
||||
let provider = build_observed(&config, recorder.clone(), client)?;
|
||||
let stream_cancellation = cancellation.clone();
|
||||
let mut stream = provider.stream(invocation, cancellation);
|
||||
let mut routed = invocation.clone();
|
||||
let model = store.model(&selected).await?.ok_or_else(|| Error::Provider(format!("unknown model: {selected}")))?;
|
||||
let provider_type = model.provider_type();
|
||||
let request_url = model.request_url()?;
|
||||
model.configure(&mut routed.request.model);
|
||||
routed.request.model.extra_params = model.extra_params().clone();
|
||||
routed.request.model.model_id = model.model_id.clone();
|
||||
let recorder = start_recorder(&store, &invocation, &model.model_hash, &model.display_name, provider_type, &request_url, &model.model_id).await?;
|
||||
let guard = recorder.cancel_on_drop();
|
||||
let config = ProviderConfig {
|
||||
kind: provider_kind(provider_type),
|
||||
request_url,
|
||||
api_key: model.api_key.clone(),
|
||||
custom_headers: if model.custom_headers_enabled { custom_headers(&model.custom_headers)? } else { reqwest::header::HeaderMap::new() },
|
||||
max_output_tokens: model.max_output_tokens(),
|
||||
request_timeout,
|
||||
retry_count: BUILTIN_PROVIDER_RETRIES,
|
||||
allowed_body_fields: None,
|
||||
};
|
||||
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
|
||||
let provider = build_observed(&config, recorder.clone(), client)?;
|
||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||
};
|
||||
|
||||
let stream_started = std::time::Instant::now();
|
||||
tracing::debug!(
|
||||
model = %selected,
|
||||
provider_type = ?provider_type,
|
||||
timeout_ms = config.request_timeout.as_millis() as u64,
|
||||
request_timeout_ms = request_timeout.as_millis() as u64,
|
||||
stream_idle_timeout_ms = stream_idle_timeout.as_millis() as u64,
|
||||
"provider stream created"
|
||||
);
|
||||
let mut last_event_time = std::time::Instant::now();
|
||||
let mut event_count: u64 = 0;
|
||||
while let Some(event) = stream.next().await {
|
||||
loop {
|
||||
let event = match next_provider_event(&mut stream, stream_idle_timeout).await {
|
||||
Ok(Some(event)) => event,
|
||||
Ok(None) => break,
|
||||
Err(_) => {
|
||||
let elapsed_ms = stream_started.elapsed().as_millis() as u64;
|
||||
let error = stream_idle_timeout_error(stream_idle_timeout);
|
||||
tracing::warn!(
|
||||
error = %error,
|
||||
elapsed_ms,
|
||||
event_count,
|
||||
idle_timeout_ms = stream_idle_timeout.as_millis() as u64,
|
||||
"provider stream idle timeout"
|
||||
);
|
||||
Err(error)
|
||||
}
|
||||
};
|
||||
let now = std::time::Instant::now();
|
||||
let gap_ms = now.duration_since(last_event_time).as_millis() as u64;
|
||||
let elapsed_ms = now.duration_since(stream_started).as_millis() as u64;
|
||||
event_count += 1;
|
||||
match event {
|
||||
Ok(event) => {
|
||||
let event_name = match &event {
|
||||
super::ModelEvent::Start { .. } => "Start",
|
||||
super::ModelEvent::TextStart => "TextStart",
|
||||
super::ModelEvent::TextDelta(_) => "TextDelta",
|
||||
super::ModelEvent::TextEnd => "TextEnd",
|
||||
super::ModelEvent::ThinkingStart => "ThinkingStart",
|
||||
super::ModelEvent::ThinkingDelta(_) => "ThinkingDelta",
|
||||
super::ModelEvent::ThinkingEnd => "ThinkingEnd",
|
||||
super::ModelEvent::ToolCallStart { .. } => "ToolCallStart",
|
||||
super::ModelEvent::ToolCallArgumentsDelta { .. } => "ToolCallArgsDelta",
|
||||
super::ModelEvent::ToolCallEnd { .. } => "ToolCallEnd",
|
||||
super::ModelEvent::ProviderReplayState(_) => "ReplayState",
|
||||
super::ModelEvent::Usage(_) => "Usage",
|
||||
super::ModelEvent::Done(_) => "Done",
|
||||
};
|
||||
if gap_ms > 5000 {
|
||||
tracing::debug!(
|
||||
gap_ms,
|
||||
elapsed_ms,
|
||||
event = event_name,
|
||||
event = event_name(&event),
|
||||
event_count,
|
||||
"slow gap detected between provider events"
|
||||
);
|
||||
@@ -137,6 +145,7 @@ impl Provider for ProviderRouter {
|
||||
yield event;
|
||||
}
|
||||
Err(error) => {
|
||||
let error = normalize_provider_stream_error(error, request_timeout);
|
||||
tracing::debug!(
|
||||
error = %error,
|
||||
elapsed_ms,
|
||||
@@ -149,26 +158,142 @@ impl Provider for ProviderRouter {
|
||||
}
|
||||
}
|
||||
}
|
||||
if !recorder.is_finished() {
|
||||
let elapsed_ms = stream_started.elapsed().as_millis() as u64;
|
||||
if stream_cancellation.is_cancelled() {
|
||||
tracing::debug!(elapsed_ms, event_count, "provider stream ended after cancellation");
|
||||
recorder.cancelled().await?;
|
||||
} else {
|
||||
let error = Error::Provider("provider stream ended without Done".into());
|
||||
tracing::warn!(
|
||||
elapsed_ms,
|
||||
event_count,
|
||||
"provider stream ended without Done"
|
||||
);
|
||||
recorder.failed(&error).await?;
|
||||
Err(error)?;
|
||||
}
|
||||
}
|
||||
finish_stream(&recorder, &cancellation).await?;
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn event_name(event: &super::ModelEvent) -> &'static str {
|
||||
match event {
|
||||
super::ModelEvent::Start { .. } => "Start",
|
||||
super::ModelEvent::TextStart => "TextStart",
|
||||
super::ModelEvent::TextDelta(_) => "TextDelta",
|
||||
super::ModelEvent::TextEnd => "TextEnd",
|
||||
super::ModelEvent::ThinkingStart => "ThinkingStart",
|
||||
super::ModelEvent::ThinkingDelta(_) => "ThinkingDelta",
|
||||
super::ModelEvent::ThinkingEnd => "ThinkingEnd",
|
||||
super::ModelEvent::ToolCallStart { .. } => "ToolCallStart",
|
||||
super::ModelEvent::ToolCallArgumentsDelta { .. } => "ToolCallArgsDelta",
|
||||
super::ModelEvent::ToolCallEnd { .. } => "ToolCallEnd",
|
||||
super::ModelEvent::ProviderReplayState(_) => "ReplayState",
|
||||
super::ModelEvent::Usage(_) => "Usage",
|
||||
super::ModelEvent::Done(_) => "Done",
|
||||
}
|
||||
}
|
||||
|
||||
async fn start_recorder(
|
||||
store: &Store,
|
||||
invocation: &ModelInvocation,
|
||||
model_hash: &str,
|
||||
display_name: &str,
|
||||
provider_type: ProviderType,
|
||||
request_url: &str,
|
||||
model_id: &str,
|
||||
) -> Result<CallRecorder> {
|
||||
CallRecorder::start(
|
||||
store.clone(),
|
||||
NewLlmCall {
|
||||
call_id: invocation.call_id.clone(),
|
||||
run_id: invocation.run_id.clone(),
|
||||
conversation_id: invocation.conversation_id.clone(),
|
||||
provider_call_index: invocation.provider_call_index.min(i64::MAX as u64) as i64,
|
||||
model_hash: model_hash.into(),
|
||||
provider_type,
|
||||
provider_url: request_url.into(),
|
||||
request_type: provider_type,
|
||||
request_url: request_url.into(),
|
||||
model_id: model_id.into(),
|
||||
display_name: display_name.into(),
|
||||
reasoning_effort: invocation.request.model.reasoning.effort.clone(),
|
||||
fast: invocation.request.model.latency == ModelLatency::Fast,
|
||||
message_count: invocation.request.history.len(),
|
||||
tool_count: invocation.request.prompt.tools.len(),
|
||||
detailed: false,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken) -> Result<()> {
|
||||
if recorder.is_finished() {
|
||||
return Ok(());
|
||||
}
|
||||
if cancellation.is_cancelled() {
|
||||
recorder.cancelled().await
|
||||
} else {
|
||||
let error = Error::Provider("provider stream ended without Done".into());
|
||||
recorder.failed(&error).await?;
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
|
||||
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
|
||||
struct PluginModelProvider {
|
||||
registry: PluginRegistry,
|
||||
}
|
||||
|
||||
impl Provider for PluginModelProvider {
|
||||
fn stream(
|
||||
&self,
|
||||
invocation: ModelInvocation,
|
||||
cancellation: CancellationToken,
|
||||
) -> ProviderStream {
|
||||
self.registry.stream_model(invocation, cancellation)
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_kind(provider_type: ProviderType) -> ProviderKind {
|
||||
match provider_type {
|
||||
ProviderType::OpenAiChat => ProviderKind::OpenAiChat,
|
||||
ProviderType::OpenAiResponses => ProviderKind::OpenAiResponses,
|
||||
ProviderType::Anthropic => ProviderKind::Anthropic,
|
||||
// 内置模型的 provider_type 只来自 ModelType,不可能是插件。
|
||||
ProviderType::Plugin => unreachable!("plugin models never use built-in provider configs"),
|
||||
}
|
||||
}
|
||||
|
||||
async fn next_provider_event(
|
||||
stream: &mut ProviderStream,
|
||||
idle_timeout: Duration,
|
||||
) -> std::result::Result<Option<Result<super::ModelEvent>>, tokio::time::error::Elapsed> {
|
||||
tokio::time::timeout(idle_timeout, stream.next()).await
|
||||
}
|
||||
|
||||
fn stream_idle_timeout_error(idle_timeout: Duration) -> Error {
|
||||
Error::Provider(format!(
|
||||
"provider stream idle timeout: no events received for {} seconds ({} minutes)",
|
||||
idle_timeout.as_secs(),
|
||||
idle_timeout.as_secs() / 60
|
||||
))
|
||||
}
|
||||
|
||||
fn request_timeout_error(request_timeout: Duration) -> Error {
|
||||
Error::Provider(format!(
|
||||
"provider request timed out after {} seconds ({} minutes)",
|
||||
request_timeout.as_secs(),
|
||||
request_timeout.as_secs() / 60
|
||||
))
|
||||
}
|
||||
|
||||
fn normalize_provider_stream_error(error: Error, request_timeout: Duration) -> Error {
|
||||
match error {
|
||||
Error::Http(source) if source.is_timeout() => request_timeout_error(request_timeout),
|
||||
Error::Http(source) if source.is_body() => Error::Provider(format!(
|
||||
"provider stream transport failed while reading the response body: {}",
|
||||
root_error_message(&source)
|
||||
)),
|
||||
error => error,
|
||||
}
|
||||
}
|
||||
|
||||
fn root_error_message(error: &(dyn std::error::Error + 'static)) -> String {
|
||||
let mut current = error;
|
||||
while let Some(source) = current.source() {
|
||||
current = source;
|
||||
}
|
||||
current.to_string()
|
||||
}
|
||||
|
||||
fn custom_headers(value: &serde_json::Value) -> Result<reqwest::header::HeaderMap> {
|
||||
let object = value
|
||||
.as_object()
|
||||
@@ -223,3 +348,36 @@ fn build_inner(
|
||||
};
|
||||
Ok(Arc::new(NormalizedProvider::new(provider)))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn pending_provider_event_hits_the_idle_timeout() {
|
||||
let mut stream: ProviderStream = Box::pin(futures_util::stream::pending());
|
||||
|
||||
let result = next_provider_event(&mut stream, Duration::from_millis(1)).await;
|
||||
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn timeout_errors_state_the_boundary_and_duration() {
|
||||
let Error::Provider(idle) = stream_idle_timeout_error(Duration::from_secs(30 * 60)) else {
|
||||
panic!("idle timeout must be a provider error");
|
||||
};
|
||||
assert_eq!(
|
||||
idle,
|
||||
"provider stream idle timeout: no events received for 1800 seconds (30 minutes)"
|
||||
);
|
||||
|
||||
let Error::Provider(request) = request_timeout_error(Duration::from_secs(60 * 60)) else {
|
||||
panic!("request timeout must be a provider error");
|
||||
};
|
||||
assert_eq!(
|
||||
request,
|
||||
"provider request timed out after 3600 seconds (60 minutes)"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,9 +2,7 @@
|
||||
|
||||
use std::collections::HashSet;
|
||||
|
||||
use crate::model::{
|
||||
CanonicalMessage, LlmCallUsageAnchor, PreparedRun, ProjectedMessage, RunAction,
|
||||
};
|
||||
use crate::model::{CanonicalMessage, LlmCallUsageAnchor, PreparedRun, ProjectedMessage};
|
||||
|
||||
const FALLBACK_CHARS: usize = 12_000;
|
||||
|
||||
@@ -34,9 +32,6 @@ pub(super) fn should_compact(
|
||||
projected_messages: &[ProjectedMessage],
|
||||
anchor: Option<ContextUsageAnchor>,
|
||||
) -> bool {
|
||||
if prepared.action != RunAction::Start {
|
||||
return false;
|
||||
}
|
||||
let Some(context_window) = prepared.model.context_window_tokens else {
|
||||
return false;
|
||||
};
|
||||
@@ -113,15 +108,15 @@ fn estimate_serialized_tokens(serialized: &str) -> u64 {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::model::{
|
||||
project_messages, CheckpointId, ConversationId, ModelSpec, Origin, PromptSpec, Role, RunId,
|
||||
RunKind,
|
||||
project_messages, CheckpointId, ConversationId, ModelSpec, Origin, PromptSpec, Role,
|
||||
RunAction, RunId, RunKind,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn automatic_compaction_starts_only_after_the_context_window_is_exceeded() {
|
||||
fn automatic_compaction_runs_for_start_and_resume_actions_after_the_limit() {
|
||||
let mut model = ModelSpec::new("model");
|
||||
model.context_window_tokens = Some(200_000);
|
||||
let prepared = PreparedRun {
|
||||
let mut prepared = PreparedRun {
|
||||
run_id: RunId::new("run"),
|
||||
cursor_request_id: None,
|
||||
conversation_id: ConversationId::new("conversation"),
|
||||
@@ -169,5 +164,15 @@ mod tests {
|
||||
&projected,
|
||||
anchor(200_001)
|
||||
));
|
||||
|
||||
prepared.action = RunAction::Resume {
|
||||
pending_tool_round: None,
|
||||
};
|
||||
assert!(should_compact(
|
||||
&prepared,
|
||||
&messages,
|
||||
&projected,
|
||||
anchor(200_001)
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -170,7 +170,7 @@ impl RunEngine {
|
||||
Ok(messages) => messages,
|
||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
||||
};
|
||||
let context_anchor = if !auto_compacted && prepared.action == RunAction::Start {
|
||||
let context_anchor = if !auto_compacted {
|
||||
match self
|
||||
.store
|
||||
.latest_llm_call_usage_anchor(
|
||||
|
||||
@@ -158,6 +158,7 @@ fn model_input(model: LegacyModel) -> Result<ModelConfigInput> {
|
||||
Ok(ModelConfigInput {
|
||||
sort_order: model.sort,
|
||||
display_name: model.display_name.clone(),
|
||||
group_name: None,
|
||||
model_type,
|
||||
base_url,
|
||||
use_full_url,
|
||||
|
||||
@@ -412,3 +412,52 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result<LlmCallSummary> {
|
||||
detailed: row.try_get("detailed")?,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
|
||||
#[tokio::test]
|
||||
async fn plugin_calls_record_without_a_model_config_row() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("test.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let plugin_model = "plugin:dev.example/codex/gpt-test";
|
||||
store
|
||||
.start_llm_call(&NewLlmCall {
|
||||
call_id: "plugin-call".into(),
|
||||
run_id: "run".into(),
|
||||
conversation_id: "conversation".into(),
|
||||
provider_call_index: 0,
|
||||
model_hash: plugin_model.into(),
|
||||
provider_type: ProviderType::Plugin,
|
||||
provider_url: "plugin://dev.example/codex".into(),
|
||||
request_type: ProviderType::Plugin,
|
||||
request_url: "plugin://dev.example/codex".into(),
|
||||
model_id: "gpt-test".into(),
|
||||
display_name: "GPT Test".into(),
|
||||
reasoning_effort: None,
|
||||
fast: false,
|
||||
message_count: 1,
|
||||
tool_count: 0,
|
||||
detailed: false,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.finish_llm_call("plugin-call", "completed", None, 10, None, None)
|
||||
.await
|
||||
.unwrap();
|
||||
let overview = store
|
||||
.overview(None, None, Some(&format!("[\"{plugin_model}\"]")))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(overview.metrics.llm_calls, 1);
|
||||
assert_eq!(overview.metrics.successful_calls, 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -446,7 +446,7 @@ mod tests {
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(checksum_after, checksum_before);
|
||||
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6]);
|
||||
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7, 8]);
|
||||
assert_eq!(checkpoint_table_exists, 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,7 +11,7 @@ use crate::{
|
||||
use super::{now_ms, Store};
|
||||
|
||||
const MODEL_COLUMNS: &str = r#"
|
||||
model_hash, sort_order, display_name, model_type, base_url, use_full_url, api_key, tooltip_data,
|
||||
model_hash, sort_order, display_name, group_name, model_type, base_url, use_full_url, api_key, tooltip_data,
|
||||
model_id, reasoning_effort, openai_endpoint, openai_extra_params_enabled,
|
||||
openai_extra_params_json, custom_headers_enabled, custom_headers_json,
|
||||
anthropic_extra_params_enabled, anthropic_extra_params_json, context_window_tokens,
|
||||
@@ -123,7 +123,7 @@ impl Store {
|
||||
}
|
||||
let result = sqlx::query(
|
||||
r#"UPDATE model_configs SET
|
||||
model_hash = ?, sort_order = ?, display_name = ?, model_type = ?, base_url = ?,
|
||||
model_hash = ?, sort_order = ?, display_name = ?, group_name = ?, model_type = ?, base_url = ?,
|
||||
use_full_url = ?, api_key = ?, tooltip_data = ?, model_id = ?, reasoning_effort = ?,
|
||||
openai_endpoint = ?, openai_extra_params_enabled = ?, openai_extra_params_json = ?,
|
||||
custom_headers_enabled = ?, custom_headers_json = ?,
|
||||
@@ -135,6 +135,7 @@ impl Store {
|
||||
.bind(&next_hash)
|
||||
.bind(input.sort_order)
|
||||
.bind(&input.display_name)
|
||||
.bind(&input.group_name)
|
||||
.bind(input.model_type.as_str())
|
||||
.bind(&input.base_url)
|
||||
.bind(input.use_full_url)
|
||||
@@ -242,13 +243,13 @@ async fn insert_model_with_conflict(
|
||||
) -> Result<bool> {
|
||||
let mut statement = String::from(
|
||||
r#"INSERT INTO model_configs(
|
||||
model_hash, sort_order, display_name, model_type, base_url, use_full_url, api_key, tooltip_data,
|
||||
model_hash, sort_order, display_name, group_name, model_type, base_url, use_full_url, api_key, tooltip_data,
|
||||
model_id, reasoning_effort, openai_endpoint, openai_extra_params_enabled,
|
||||
openai_extra_params_json, custom_headers_enabled, custom_headers_json,
|
||||
anthropic_extra_params_enabled, anthropic_extra_params_json, context_window_tokens,
|
||||
max_completion_tokens, anthropic_max_tokens, anthropic_thinking_effort,
|
||||
thinking_budget_tokens, created_at_ms, updated_at_ms
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"#,
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"#,
|
||||
);
|
||||
if ignore_existing {
|
||||
statement.push_str(" ON CONFLICT(model_hash) DO NOTHING");
|
||||
@@ -257,6 +258,7 @@ async fn insert_model_with_conflict(
|
||||
.bind(hash)
|
||||
.bind(input.sort_order)
|
||||
.bind(&input.display_name)
|
||||
.bind(&input.group_name)
|
||||
.bind(input.model_type.as_str())
|
||||
.bind(&input.base_url)
|
||||
.bind(input.use_full_url)
|
||||
@@ -288,6 +290,7 @@ fn model_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ModelConfig> {
|
||||
model_hash: row.try_get("model_hash")?,
|
||||
sort_order: row.try_get("sort_order")?,
|
||||
display_name: row.try_get("display_name")?,
|
||||
group_name: row.try_get("group_name")?,
|
||||
model_type: ModelType::from_str(row.try_get("model_type")?)?,
|
||||
base_url: row.try_get("base_url")?,
|
||||
use_full_url: row.try_get("use_full_url")?,
|
||||
@@ -320,6 +323,71 @@ fn model_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ModelConfig> {
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn model_input(group_name: Option<&str>) -> ModelConfigInput {
|
||||
ModelConfigInput {
|
||||
sort_order: 0,
|
||||
display_name: "Test Model".into(),
|
||||
group_name: group_name.map(String::from),
|
||||
model_type: ModelType::OpenAi,
|
||||
base_url: "https://example.com/v1/chat/completions".into(),
|
||||
use_full_url: true,
|
||||
api_key: "test-key".into(),
|
||||
tooltip_data: "Test Model".into(),
|
||||
model_id: "test-model".into(),
|
||||
reasoning_effort: None,
|
||||
openai_endpoint: crate::model::OPENAI_CHAT_ENDPOINT.into(),
|
||||
openai_extra_params_enabled: false,
|
||||
openai_extra_params: serde_json::json!({}),
|
||||
custom_headers_enabled: false,
|
||||
custom_headers: serde_json::json!({}),
|
||||
anthropic_extra_params_enabled: false,
|
||||
anthropic_extra_params: serde_json::json!({}),
|
||||
context_window_tokens: None,
|
||||
max_completion_tokens: None,
|
||||
anthropic_max_tokens: None,
|
||||
anthropic_thinking_effort: None,
|
||||
thinking_budget_tokens: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 分组名是纯展示字段:入库时去除首尾空白、空串归一为 NULL,
|
||||
/// 更新分组名不得改变模型身份哈希。
|
||||
#[tokio::test]
|
||||
async fn group_name_round_trips_without_changing_model_identity() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("test.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let created = store
|
||||
.create_model(&model_input(Some(" My Group ")))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(created.group_name.as_deref(), Some("My Group"));
|
||||
|
||||
let renamed = store
|
||||
.update_model(&created.model_hash, &model_input(Some("Renamed")))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(renamed.model_hash, created.model_hash);
|
||||
assert_eq!(renamed.group_name.as_deref(), Some("Renamed"));
|
||||
|
||||
let cleared = store
|
||||
.update_model(&created.model_hash, &model_input(Some(" ")))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(cleared.model_hash, created.model_hash);
|
||||
assert_eq!(cleared.group_name, None);
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_u64(row: &sqlx::sqlite::SqliteRow, column: &str) -> Result<Option<u64>> {
|
||||
row.try_get::<Option<i64>, _>(column)?
|
||||
.map(|value| {
|
||||
|
||||
@@ -101,7 +101,14 @@ impl Store {
|
||||
.bind(&call.call_id)
|
||||
.bind(&call.model_call_id)
|
||||
.bind(&call.name)
|
||||
.bind(&call.arguments_text)
|
||||
// A no-argument tool call streams no argument text; persist it as an
|
||||
// empty object so the `arguments_json` column always holds valid JSON
|
||||
// and can be re-parsed on load.
|
||||
.bind(if call.arguments_text.trim().is_empty() {
|
||||
"{}"
|
||||
} else {
|
||||
call.arguments_text.as_str()
|
||||
})
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
}
|
||||
|
||||
@@ -27,6 +27,7 @@ async fn summarize_replaces_model_history_and_preserves_cursor_history() {
|
||||
.create_model(&ModelConfigInput {
|
||||
sort_order: 0,
|
||||
display_name: "Test Model".into(),
|
||||
group_name: None,
|
||||
model_type: ModelType::OpenAi,
|
||||
base_url: "https://example.com/v1/chat/completions".into(),
|
||||
use_full_url: true,
|
||||
|
||||
@@ -622,6 +622,68 @@ async fn runtime_user_message_reports_delivered_and_appended() {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tool_call_with_empty_arguments_does_not_fail_the_run() {
|
||||
// A tool call that carries no arguments streams no argument text. Parsing it
|
||||
// as JSON must yield an empty object (as the model cycle already does), not
|
||||
// fail the run with `EOF while parsing a value`.
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(tool_response("call-1", "UpdateCurrentStep", ""));
|
||||
provider.push(text_response("done after empty-argument tool"));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("empty-args-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run_for(
|
||||
"empty-args-request",
|
||||
"empty-args-conversation",
|
||||
)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut saw_done = false;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("RunSSE closed before successful EndStream");
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
assert_eq!(
|
||||
payload.as_ref(),
|
||||
b"{}",
|
||||
"run failed: {}",
|
||||
String::from_utf8_lossy(&payload)
|
||||
);
|
||||
break;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
if let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = server.message {
|
||||
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
|
||||
saw_done |= delta.text.contains("done after empty-argument tool");
|
||||
}
|
||||
}
|
||||
acknowledge_kv(&handle, &mut append_seqno, &frame).await;
|
||||
}
|
||||
assert!(saw_done);
|
||||
assert_eq!(provider.requests().len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn injected_user_context_restarts_only_the_active_model_cycle() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
@@ -944,6 +1006,7 @@ async fn injected_user_context_interrupts_automatic_compaction() {
|
||||
.create_model(&ModelConfigInput {
|
||||
sort_order: 0,
|
||||
display_name: "Test Model".into(),
|
||||
group_name: None,
|
||||
model_type: ModelType::OpenAi,
|
||||
base_url: "https://example.com/v1/chat/completions".into(),
|
||||
use_full_url: true,
|
||||
|
||||
@@ -0,0 +1,213 @@
|
||||
//! Verifies KnowledgeBase rules CRUD falls back to local markdown storage
|
||||
//! when the Cursor upstream is unreachable or rejects the request.
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
extract::Extension,
|
||||
http::{header, Request, Response},
|
||||
};
|
||||
use cursor_server::{
|
||||
api::cursor::proxy::CursorProxy,
|
||||
cursor::services::knowledge::{self, KnowledgeService},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
// 测试侧的镜像消息定义,同时充当 wire 兼容性检查。
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct AddRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "2")]
|
||||
title: String,
|
||||
#[prost(string, tag = "3")]
|
||||
git_origin: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct AddResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
#[prost(string, tag = "2")]
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct ListRequest {
|
||||
#[prost(int32, optional, tag = "1")]
|
||||
limit: Option<i32>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct ListResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
#[prost(message, repeated, tag = "2")]
|
||||
all_results: Vec<ListItem>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct ListItem {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
#[prost(string, tag = "2")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "3")]
|
||||
title: String,
|
||||
#[prost(string, tag = "4")]
|
||||
created_at: String,
|
||||
#[prost(bool, tag = "5")]
|
||||
is_generated: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct UpdateRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
#[prost(string, tag = "2")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "3")]
|
||||
title: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct UpdateResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct RemoveRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
struct RemoveResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
}
|
||||
|
||||
fn proto_request(message: &impl Message) -> Request<Body> {
|
||||
Request::post("/test")
|
||||
.header(header::CONTENT_TYPE, "application/proto")
|
||||
.body(Body::from(message.encode_to_vec()))
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
async fn decode<M: Message + Default>(response: Response<Body>) -> M {
|
||||
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||
M::decode(body.as_ref()).unwrap()
|
||||
}
|
||||
|
||||
/// 无凭据请求上游必然失败(网络错误或 401),四个接口全部走本地降级,
|
||||
/// 覆盖 md 持久化、离线日志压缩与增删改查闭环。
|
||||
#[tokio::test]
|
||||
async fn offline_crud_round_trip_persists_markdown() {
|
||||
let (_store_dir, store) = fixtures::temp_store().await;
|
||||
let upstream = CursorProxy::cursor(store).unwrap();
|
||||
let rules_dir = tempfile::tempdir().unwrap();
|
||||
let rules_root = rules_dir.path().join("rules");
|
||||
let service = KnowledgeService::with_root(rules_root.clone()).unwrap();
|
||||
|
||||
// Add:得到本地临时 id,md 文件落盘。
|
||||
let response = knowledge::add(
|
||||
Extension(upstream.clone()),
|
||||
Extension(service.clone()),
|
||||
proto_request(&AddRequest {
|
||||
knowledge: "always answer in haiku".into(),
|
||||
title: "haiku rule".into(),
|
||||
git_origin: String::new(),
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let added: AddResponse = decode(response).await;
|
||||
assert!(added.success);
|
||||
assert!(added.id.starts_with("local-"), "offline add uses a local id");
|
||||
let markdown = rules_root.join(format!("{}.md", added.id));
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(&markdown).unwrap(),
|
||||
"always answer in haiku"
|
||||
);
|
||||
|
||||
// List:本地缓存返回刚写入的规则。
|
||||
let response = knowledge::list(
|
||||
Extension(upstream.clone()),
|
||||
Extension(service.clone()),
|
||||
proto_request(&ListRequest { limit: Some(100) }),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let listed: ListResponse = decode(response).await;
|
||||
assert!(listed.success);
|
||||
assert_eq!(listed.all_results.len(), 1);
|
||||
assert_eq!(listed.all_results[0].id, added.id);
|
||||
assert_eq!(listed.all_results[0].title, "haiku rule");
|
||||
|
||||
// Update:内容与标题都更新到 md 与元数据。
|
||||
let response = knowledge::update(
|
||||
Extension(upstream.clone()),
|
||||
Extension(service.clone()),
|
||||
proto_request(&UpdateRequest {
|
||||
id: added.id.clone(),
|
||||
knowledge: "always answer in sonnets".into(),
|
||||
title: "sonnet rule".into(),
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let updated: UpdateResponse = decode(response).await;
|
||||
assert!(updated.success);
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(&markdown).unwrap(),
|
||||
"always answer in sonnets"
|
||||
);
|
||||
|
||||
// Remove:文件删除,列表为空。
|
||||
let response = knowledge::remove(
|
||||
Extension(upstream.clone()),
|
||||
Extension(service.clone()),
|
||||
proto_request(&RemoveRequest {
|
||||
id: added.id.clone(),
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let removed: RemoveResponse = decode(response).await;
|
||||
assert!(removed.success);
|
||||
assert!(!markdown.exists());
|
||||
|
||||
let response = knowledge::list(
|
||||
Extension(upstream),
|
||||
Extension(service),
|
||||
proto_request(&ListRequest { limit: Some(100) }),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let listed: ListResponse = decode(response).await;
|
||||
assert!(listed.all_results.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn updating_missing_rule_reports_failure() {
|
||||
let (_store_dir, store) = fixtures::temp_store().await;
|
||||
let upstream = CursorProxy::cursor(store).unwrap();
|
||||
let rules_dir = tempfile::tempdir().unwrap();
|
||||
let service = KnowledgeService::with_root(rules_dir.path().join("rules")).unwrap();
|
||||
|
||||
let response = knowledge::update(
|
||||
Extension(upstream),
|
||||
Extension(service),
|
||||
proto_request(&UpdateRequest {
|
||||
id: "17353272".into(),
|
||||
knowledge: "anything".into(),
|
||||
title: "anything".into(),
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let updated: UpdateResponse = decode(response).await;
|
||||
assert!(!updated.success);
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
//! Verifies local markdown rules are merged into the request-context message.
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
protocol::connect,
|
||||
protocol::proto::agent::v1 as pb,
|
||||
TransportCommand, TransportRegistry,
|
||||
},
|
||||
model::{ContentPart, ProjectedContent},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn local_markdown_rules_land_in_the_request_context_message() {
|
||||
let (_store_dir, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "call-1".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("ok".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let rules_dir = tempfile::tempdir().unwrap();
|
||||
let rules_root = rules_dir.path().join("rules");
|
||||
std::fs::create_dir_all(&rules_root).unwrap();
|
||||
std::fs::write(rules_root.join("17353272.md"), "Always answer in haiku.").unwrap();
|
||||
|
||||
let registry = TransportRegistry::with_local_rules(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
rules_root,
|
||||
);
|
||||
let handle = registry.get_or_create("rules-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(user_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||
.await
|
||||
.expect("run finishes within timeout")
|
||||
.expect("output stays open until EndStream");
|
||||
let ended = connect::decode_frames(&frame)
|
||||
.unwrap()
|
||||
.iter()
|
||||
.any(|(flags, _)| flags & connect::END_STREAM_FLAG != 0);
|
||||
if ended {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 1);
|
||||
let context_texts = requests[0]
|
||||
.history
|
||||
.iter()
|
||||
.filter(|message| message.message_id.starts_with("request-context:"))
|
||||
.map(|message| {
|
||||
let ProjectedContent::Parts(parts) = &message.content else {
|
||||
panic!("request context message must be parts")
|
||||
};
|
||||
let [ContentPart::Text { text }] = parts.as_slice() else {
|
||||
panic!("request context message must be one text part")
|
||||
};
|
||||
text.clone()
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
context_texts.len(),
|
||||
1,
|
||||
"exactly one request-context message is projected"
|
||||
);
|
||||
assert!(
|
||||
context_texts[0].contains("<user_rule>\nAlways answer in haiku.\n</user_rule>"),
|
||||
"local markdown rule must appear as a user rule: {}",
|
||||
context_texts[0]
|
||||
);
|
||||
|
||||
registry.shutdown().await;
|
||||
}
|
||||
|
||||
fn user_run() -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
pb::UserMessageAction {
|
||||
user_message: Some(pb::UserMessage {
|
||||
text: "hello".into(),
|
||||
message_id: "rules-user".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("rules-conversation".into()),
|
||||
run_id: Some("rules-request".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user