Merge remote-tracking branch 'origin/main' into pr-385-merge

# Conflicts:
#	server/tests/interrupt.rs
This commit is contained in:
leokun
2026-08-31 16:12:11 +08:00
333 changed files with 14679 additions and 48321 deletions
+2 -1
View File
@@ -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;
},
};
+210
View File
@@ -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");
});
+16
View File
@@ -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;
},
};
+148
View File
@@ -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 } : {}),
};
},
};
+31 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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)
);
}
}
+40
View File
@@ -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),
+124
View File
@@ -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()))
}
+114 -8
View File
@@ -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();
+76
View File
@@ -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('<', "&lt;")
.replace('>', "&gt;")
}
#[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());
}
}
+7 -1
View File
@@ -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,
+8 -1
View File
@@ -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,
};
+1 -1
View File
@@ -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) => {
+356
View File
@@ -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(&timestamp(&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
}
}
}
+1
View File
@@ -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;
+219 -94
View File
@@ -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(),
+21 -1
View File
@@ -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 -1
View File
@@ -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");
}
}
+2 -1
View File
@@ -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
+49
View File
@@ -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(&notebook_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(&notebook_call("hi"), &single_cell_notebook()).unwrap();
assert!(edited.contains("print('replacement')"));
}
}
+54
View File
@@ -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
}
+2
View File
@@ -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",
+1
View File
@@ -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;
+4
View File
@@ -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"
)
+16
View File
@@ -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,
+3 -1
View File
@@ -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()),
+88
View File
@@ -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());
}
}
+235
View File
@@ -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());
}
}
+317
View File
@@ -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());
}
}
+246
View File
@@ -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());
}
}
+217
View File
@@ -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());
}
}
+230
View File
@@ -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);
}
}
+273
View File
@@ -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")
);
}
}
+213
View File
@@ -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"));
}
}
+32
View File
@@ -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) {}
+92
View File
@@ -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"),
}
}
}
+959
View File
@@ -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
))
})
}
+229
View File
@@ -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;
}
}
+5
View File
@@ -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())));
+5
View File
@@ -0,0 +1,5 @@
{
"fmt": {
"lineWidth": 200
}
}
+10
View File
@@ -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"
}
}
+29
View File
@@ -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[]>;
};
+100
View File
@@ -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");
}
}
+129
View File
@@ -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>;
};
+118
View File
@@ -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>;
};
+185
View File
@@ -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));
}
}
}
+457
View File
@@ -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");
}
}
+297
View File
@@ -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());
}
}
+647
View File
@@ -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: &params,
},
)
.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(&params, "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, &params).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, &params).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(&params, "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}'")))
}
+6 -4
View File
@@ -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)))?,
_ => {}
}
}
+130
View File
@@ -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);
}
}
+8 -4
View File
@@ -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));
}
+8 -4
View File
@@ -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)))?,
_ => {}
}
}
+1
View File
@@ -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,
-3
View File
@@ -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
View File
@@ -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)"
);
}
}
+15 -10
View File
@@ -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)
));
}
}
+1 -1
View File
@@ -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(
+1
View File
@@ -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,
+49
View File
@@ -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);
}
}
+1 -1
View File
@@ -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);
}
}
+72 -4
View File
@@ -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| {
+8 -1
View File
@@ -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?;
}
+1
View File
@@ -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,
+63
View File
@@ -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,
+213
View File
@@ -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);
}
+133
View File
@@ -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()
},
)),
}
}